{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":47317,"databundleVersionId":5799376},{"sourceType":"datasetVersion","sourceId":5926189,"datasetId":3291966,"databundleVersionId":6003684},{"sourceType":"datasetVersion","sourceId":3492503,"datasetId":2102274,"databundleVersionId":3544941},{"sourceType":"datasetVersion","sourceId":12469743,"datasetId":7866985,"databundleVersionId":13044596},{"sourceType":"datasetVersion","sourceId":5914240,"datasetId":3397020,"databundleVersionId":5991623},{"sourceType":"datasetVersion","sourceId":12469757,"datasetId":7866997,"databundleVersionId":13044612},{"sourceType":"datasetVersion","sourceId":5013702,"datasetId":2909244,"databundleVersionId":5083617},{"sourceType":"datasetVersion","sourceId":699609,"datasetId":255887,"databundleVersionId":719625},{"sourceType":"datasetVersion","sourceId":3492463,"datasetId":2102244,"databundleVersionId":3544900},{"sourceType":"datasetVersion","sourceId":3951115,"datasetId":1027206,"databundleVersionId":4006592},{"sourceType":"datasetVersion","sourceId":5912994,"datasetId":3291965,"databundleVersionId":5990374},{"sourceType":"datasetVersion","sourceId":5621228,"datasetId":3204855,"databundleVersionId":5696423},{"sourceType":"datasetVersion","sourceId":5309119,"datasetId":3085847,"databundleVersionId":5382305}],"dockerImageVersionId":30408,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"HI\n#/kaggle/input/vesuvius-challenge-ink-detection/train\n!pip install segmentation-models-pytorch==0.2.0\n!pip install monai\n!pip install einops\n!pip install segmentation-models-pytorch==0.2.0","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS INK DETECTION - V5.3 \"SINGLE-MODEL, DIAGNOSTICS-DRIVEN\"\n# ============================================================\n# Built on V5.2. ONE change from V5.2: the ensemble is dropped down to its\n# single stronger member -- tu-convnext_tiny + window (12,38,1). The second\n# member (efficientnet-b4 + a wide window) validated worse AND roughly\n# doubled total run time, so it was pure cost with no benefit. The ensemble\n# loop/calibration code is left in place (harmless at N=1) so a second member\n# can be re-added later with a one-line CFG change (N_ENSEMBLE_MODELS,\n# ensemble_seeds/encoders/windows).\n#\n# Also added: a dedicated comparison figure per Fragment 1 --\n# input / ground truth / probability / prediction / TP-FP-FN overlay -- saved\n# as its own function (see PART 18) with clearly distinct overlay colors\n# (green=TP, red=FP, yellow=FN), independent of the run's console output.\n#\n# Everything else is V5.2 as diagnosed from your actual run. Every change\n# below is tied to a number in that diagnostics run:\n#\n#  PART 1 (fragment 1): AP=0.50 vs prevalence 0.18 -> real signal, but RANKING is\n#  the ceiling (oracle F0.5 0.513 == best threshold you used). However the\n#  ensemble's OTSU threshold (0.424) scored F0.5 0.408, i.e. ~0.10 WORSE than the\n#  validation thresholds, and best-Dice threshold (0.51) != best-F0.5 threshold\n#  (0.67). => (1) Otsu is no longer used to decide anything. (2) All thresholds\n#  are chosen by F0.5 (the competition metric), ON THE ENSEMBLED, TTA'd\n#  VALIDATION PREDICTIONS, then frozen. (3) Checkpoints are selected by\n#  best-threshold F0.5 instead of Dice@0.5.\n#\n#  PART 2 (depth): fragments 1 and 3 share almost the same mean depth profile\n#  (r=0.996) while fragment 2 differs (r=0.83; broader/lower peak, trough ~6\n#  slices deeper). The fixed window (12..37) leaves 43-55% of the ink-vs-\n#  background contrast outside it, and fragment 3's strongest contrast sits on\n#  the window's last slice. => (4) member windows now cover ~8..58 (stride 2 keeps\n#  the channel count and I/O identical); (5) per-fragment PER-DEPTH mean-profile\n#  normalization removes the fragment-specific depth profile (label-free, same\n#  status as the old per-fragment z-score); (6) depth-warp augmentation\n#  (random z-scale/shift) targets the fragment-2-style profile shift;\n#  (7) domain-balanced sampling stops fragment 2 (78% of patches) from\n#  dominating fragment 3, the fragment whose depth profile matches the test.\n#\n#  ALSO: (8) validation-selected Gaussian smoothing of the probability map\n#  (ink strokes are spatially coherent; sigma chosen by val F0.5, 0 allowed),\n#  (9) predictions are masked by the tissue mask before scoring (labels are),\n#  (10) cache warm-up per volume (fixes cold random-read I/O on Kaggle).\n#\n# NOT included on purpose: the auxiliary \"ink differential\" head. Part 2 shows\n# the ink depth signature differs between fragments 2 and 3 (peak slice 17 vs\n# 24, different sign lobes), so supervising a fixed differential signature would\n# mostly learn fragment-specific patterns.\n#\n# HONEST EXPECTATION: (1)-(3) and (9) are corrections and should recover the\n# ~0.10 lost by Otsu plus a small ensemble gain. (4)-(8) are evidence-motivated\n# HYPOTHESES about the ranking ceiling; each has its own CFG flag so you can\n# ablate them. Fragment 1's ground truth is only used for reporting at the end.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.amp import autocast as _autocast_new, GradScaler as _GradScaler_new\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\n\ndef autocast(enabled=True):\n    \"\"\"Thin wrapper over torch.amp (torch.cuda.amp is deprecated).\"\"\"\n    return _autocast_new(\"cuda\" if torch.cuda.is_available() else \"cpu\", enabled=enabled)\n\n\nclass GradScaler(_GradScaler_new):\n    def __init__(self, enabled=True):\n        super().__init__(\"cuda\" if torch.cuda.is_available() else \"cpu\", enabled=enabled)\n\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    # --- depth windows (start, stop_exclusive, step); every window must give n_slices\n    #     channels. Diagnostics: 43-55% of ink contrast lies OUTSIDE (12..37), so the\n    #     wide windows use stride 2 over ~6..58 (same channels, same I/O cost). -----\n    n_slices = 26\n    in_channels = n_slices\n    max_slice_index = 64\n    # V5.3: single-model run -- efficientnet-b4 (former member 2) validated\n    # worse AND cost roughly 2x total time for no benefit, so it's dropped.\n    # Kept only the tu-convnext_tiny + (12,38,1) config that was actually the\n    # stronger of the two. The ensemble loop below still runs (harmless at\n    # N=1) so re-enabling a second member later is a one-line CFG change.\n    ensemble_windows = [(12, 38, 1)]\n    mask_reference_slice = 24          # used only for the fallback tissue mask + overview image\n\n    patch_size = 320\n    train_stride = 128\n    test_stride = 64\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.45\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 30\n    early_stop_patience = 5\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n\n    # --- backbone (verified: smp.Unet + tu-convnext_tiny; not UnetPlusPlus) ---------\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    USE_DEPTH_FUSION_STEM = True\n    depth_stem_out_channels = 26\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- ensemble (member i uses window i) ------------------------------------------\n    N_ENSEMBLE_MODELS = 1                                  # raise + extend the lists above to re-enable\n    ensemble_seeds = [42]\n    ensemble_encoders = [\"tu-convnext_tiny\"]\n\n    # --- LOSS -------------------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.30\n    focal_tversky_weight = 0.25\n    topology_weight = 0.15\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n    cldice_iters = 8\n\n    # --- augmentation -----------------------------------------------------------------\n    USE_PHYSICAL_SYNTHETIC_INK = True\n    physical_ink_p = 0.15\n    physical_ink_max_amplitude = 60\n\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n    USE_SHADOW = True\n    shadow_p = 0.20\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # Depth-warp: random z-scale / z-shift of the channel axis (linear interpolation\n    # between slices). Targets the fragment-2-style profile shift seen in Part 2.\n    USE_DEPTH_WARP_AUG = True\n    depth_warp_p = 0.40\n    depth_warp_scale_range = (0.75, 1.33)\n    depth_warp_shift_range = (-2.0, 2.0)      # in channel units\n\n    # Histogram-match aug: pool from TRAIN fragments only (strict protocol).\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # Domain-balanced sampling: share_f = (1-s)*natural_share + s*(1/k). Total sample\n    # count is unchanged (bigger fragment subsampled, smaller one oversampled).\n    DOMAIN_BALANCE = True\n    domain_balance_strength = 0.5\n\n    # --- normalization ----------------------------------------------------------------\n    # \"per_fragment_per_slice\": subtract the fragment's own mean tissue intensity PER\n    #     DEPTH SLICE, divide by one per-fragment scale (label-free, computed from\n    #     tissue pixels of each fragment separately -- same status as the old z-score).\n    # \"per_fragment_zscore\": legacy scalar mean/std from the middle slice.\n    NORMALIZATION_MODE = \"per_fragment_per_slice\"\n    PER_SLICE_STD_SCALING = False\n    profile_row_stride = 8\n    frag_stats_sample_patches = 60\n\n    # --- validation calibration (threshold + smoothing chosen on ENSEMBLED, TTA'd val) -\n    val_calibration_max_patches = 600\n    smooth_sigma_candidates = [0.0, 1.0, 2.0, 3.0, 4.0, 6.0]\n    threshold_smooth_bins = 2\n\n    # --- TTA: full 8-view D4 group ------------------------------------------------------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"hvflip\", \"rot90\", \"rot180\", \"rot270\", \"transpose\"]\n\n    use_adabn = False\n    adabn_max_patches = 2000\n\n    USE_EMA = True\n    ema_decay = 0.999\n\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    CLAHE_MODE = \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # Forces one fast sequential read of every opened slice (cold random reads on\n    # Kaggle's input filesystem can cost seconds per patch).\n    WARM_UP_VOLUME_CACHE = True\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    viz_dir = os.path.join(out_dir, \"v5_3_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v5_3_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        ref_path = os.path.join(frag_dir, \"surface_volume\", f\"{CFG.mask_reference_slice:02d}.tif\")\n        ref = tifffile.imread(ref_path)\n        mask = (ref > ref.mean() * 0.15).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion emptied the mask for {frag_dir}; keeping un-eroded mask.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / BALANCING\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            if mask[y:y + patch_size, x:x + patch_size].mean() > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3, rng=random):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = float(labels_full[fid][y:y + patch_size, x:x + patch_size].mean())\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = max(int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio), len(positive))\n    negative_selected = rng.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    balanced = positive * repeat + negative_selected\n    rng.shuffle(balanced)\n    print(f\"  Positive patches: {len(positive)} | Negative patches: {len(negative)} | \"\n          f\"Balanced: {len(balanced)} (positive ratio={len(positive)*repeat/max(len(balanced),1):.3f})\")\n    return balanced\n\n\ndef domain_balance(samples, rng):\n    \"\"\"Blend natural fragment shares toward equal shares, keeping the TOTAL constant\n    (subsample the larger fragment, oversample the smaller one).\"\"\"\n    if not CFG.DOMAIN_BALANCE or len(CFG.train_frags) < 2 or not samples:\n        return samples\n    by = {}\n    for s in samples:\n        by.setdefault(s[0], []).append(s)\n    fids = sorted(by)\n    n_total = len(samples)\n    k = len(fids)\n    natural = np.array([len(by[f]) for f in fids], dtype=np.float64)\n    natural /= natural.sum()\n    lam = CFG.domain_balance_strength\n    share = (1.0 - lam) * natural + lam * (1.0 / k)\n    out = []\n    for f, sh in zip(fids, share):\n        n = int(round(sh * n_total))\n        pool = by[f]\n        if n <= len(pool):\n            out.extend(rng.sample(pool, n))\n        else:\n            out.extend(pool + [rng.choice(pool) for _ in range(n - len(pool))])\n    rng.shuffle(out)\n    print(\"  Domain balance: \" + \" | \".join(\n        f\"frag {f}: {len(by[f])} -> {int(round(sh * n_total))}\" for f, sh in zip(fids, share)))\n    return out\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME (per-depth normalization, cache warm-up)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None\n        self.frag_mean = None\n        self.frag_std = None\n        self.slice_mean = None      # per-depth tissue mean (0..255 units, after contrast LUT)\n        self.slice_std = None\n        self.norm_scale = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _to_uint8(block):\n        block = np.asarray(block)\n        if block.dtype != np.uint8:\n            block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        return np.ascontiguousarray(block)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid = len(self._slices) // 2\n        ref_slice = self._to_uint8(self._slices[mid])\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def _apply_contrast(self, block):\n        if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n            return self._clahe.apply(block)\n        if CFG.CLAHE_MODE == \"global_shared\":\n            return cv2.LUT(block, self._shared_clahe_lut)\n        return block\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        out = np.zeros((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = self._to_uint8(s[y:y + size, x:x + size])\n            block = self._apply_contrast(block)\n            hh, ww = block.shape\n            out[i, :hh, :ww] = block[:size, :size]\n        return out\n\n    def warm_up_cache(self, tag=\"\"):\n        \"\"\"One fast sequential read of every opened slice so later random patch reads\n        are page-cache hits (result of .sum() is discarded; nothing is retained).\"\"\"\n        self._ensure_open()\n        t0 = time.time()\n        total = 0\n        for arr in self._slices:\n            _ = np.asarray(arr).sum(dtype=np.int64)\n            total += arr.nbytes\n        dt = time.time() - t0\n        print(f\"  [cache warm-up{(' ' + tag) if tag else ''}] {len(self._slices)} slices, \"\n              f\"{total/1e9:.2f} GB in {dt:.1f}s ({total/1e6/max(dt,1e-6):.0f} MB/s)\")\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        \"\"\"Legacy scalar stats (middle slice) + per-depth tissue profile.\"\"\"\n        self._ensure_open()\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n        else:\n            sample_coords = random.sample(coords, min(n_samples, len(coords)))\n            mid = len(self._slices) // 2\n            vals = []\n            for (y, x) in sample_coords:\n                block = self._to_uint8(self._slices[mid][y:y + patch_size, x:x + patch_size])\n                block = self._apply_contrast(block)\n                vals.append(block.astype(np.float32).ravel())\n            vals = np.concatenate(vals)\n            self.frag_mean = float(vals.mean())\n            self.frag_std = float(vals.std() + 1e-6)\n\n        if CFG.NORMALIZATION_MODE == \"per_fragment_per_slice\":\n            self._compute_depth_profile(tissue_mask)\n\n    def _compute_depth_profile(self, tissue_mask):\n        \"\"\"Per-depth mean/std over tissue pixels, from every `profile_row_stride`-th row\n        (reads ~1/stride of the data). Label-free.\"\"\"\n        stride = max(int(CFG.profile_row_stride), 1)\n        H, W = self._h, self._w\n        sel = tissue_mask[:H:stride, :W] > 0\n        n = len(self._slices)\n        means, stds = np.zeros(n), np.zeros(n)\n        for i, s in enumerate(self._slices):\n            block = self._to_uint8(s[:H:stride, :W])\n            block = self._apply_contrast(block)\n            vals = block[sel] if sel.shape == block.shape and sel.sum() > 1000 else block.ravel()\n            means[i] = float(vals.mean())\n            stds[i] = float(vals.std() + 1e-6)\n        self.slice_mean = means\n        self.slice_std = stds\n        self.norm_scale = float(np.sqrt(np.mean(stds ** 2)))\n        print(f\"    depth profile: mean range [{means.min():.1f}, {means.max():.1f}] \"\n              f\"(peak slice #{int(np.argmax(means))}, trough slice #{int(np.argmin(means))} of window) \"\n              f\"| scale={self.norm_scale:.2f}\")\n\n    def normalize(self, img_dhw):\n        \"\"\"img_dhw: float32 (D,H,W) in [0,1].\"\"\"\n        mode = CFG.NORMALIZATION_MODE\n        if mode == \"per_fragment_per_slice\" and self.slice_mean is not None:\n            m = (self.slice_mean / 255.0).astype(np.float32)[:, None, None]\n            if CFG.PER_SLICE_STD_SCALING:\n                s = (np.maximum(self.slice_std, 1.0) / 255.0).astype(np.float32)[:, None, None]\n            else:\n                s = np.float32(self.norm_scale / 255.0 + 1e-6)\n            return ((img_dhw - m) / s).astype(np.float32)\n        if mode == \"per_fragment_zscore\" and self.frag_mean is not None:\n            return ((img_dhw - self.frag_mean / 255.0) / (self.frag_std / 255.0 + 1e-6)).astype(np.float32)\n        return ((img_dhw - img_dhw.mean()) / (img_dhw.std() + 1e-6)).astype(np.float32)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\n# ============================================================\n# 6. AUGMENTATION\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Label-FREE distractor (never marked positive).\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)], axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        n_affected = max(1, int(d * random.uniform(0.3, 0.7)))\n        affected = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_physical_synthetic_ink(img_hwd, label, n_strokes=None, max_amplitude=CFG.physical_ink_max_amplitude):\n    \"\"\"Label-POSITIVE synthetic ink with a Gaussian depth profile\n    I(z)=A*exp(-(z-z0)^2/(2 sigma^2)); the label mask is updated.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    label_out = label.copy()\n    n_strokes = n_strokes or random.randint(1, 2)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)], axis=1).astype(np.int32)\n        thickness = random.randint(2, 5)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n\n        z0 = random.uniform(d * 0.3, d * 0.7)\n        sigma = random.uniform(1.5, 4.0)\n        amplitude = random.uniform(20, max_amplitude) * random.choice([-1, 1])\n        profile = amplitude * np.exp(-((np.arange(d) - z0) ** 2) / (2 * sigma ** 2))\n        for zi in range(d):\n            delta = profile[zi]\n            if abs(delta) < 1.0:\n                continue\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n        label_out[stroke_mask > 0] = 1\n    return out, label_out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist_norm = np.clip(np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2) / radius, 0, 1)\n    factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\ndef warp_depth(img_hwd, scale, shift):\n    \"\"\"Resample the channel (depth) axis: output channel j reads source coordinate\n    (j-c)*scale + c + shift (clamped), with linear interpolation between slices.\"\"\"\n    d = img_hwd.shape[2]\n    c = (d - 1) / 2.0\n    src = np.clip((np.arange(d) - c) * scale + c + shift, 0, d - 1)\n    lo = np.floor(src).astype(np.int64)\n    hi = np.minimum(lo + 1, d - 1)\n    frac = (src - lo).astype(np.float32)\n    out = img_hwd[:, :, lo].astype(np.float32) * (1.0 - frac) + img_hwd[:, :, hi].astype(np.float32) * frac\n    return np.clip(out + 0.5, 0, 255).astype(np.uint8)\n\n\ndef _safe_transform(builder_new, builder_old, name):\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable ({e1} / {e2}); skipping.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma, border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03, p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 7. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            y = max(0, min(y + random.randint(-self.jitter, self.jitter), H - size))\n            x = max(0, min(x + random.randint(-self.jitter, self.jitter), W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size]\n        img = np.transpose(patch, (1, 2, 0))          # HWD uint8\n\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref_hwd = np.transpose(random.choice(self.hist_match_pool), (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n\n        if self.train_mode and CFG.USE_PHYSICAL_SYNTHETIC_INK and random.random() < CFG.physical_ink_p:\n            img, label = inject_physical_synthetic_ink(img, label)\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.train_mode and CFG.USE_DEPTH_WARP_AUG and random.random() < CFG.depth_warp_p:\n            lo, hi = CFG.depth_warp_scale_range\n            scale = math.exp(random.uniform(math.log(lo), math.log(hi)))\n            shift = random.uniform(*CFG.depth_warp_shift_range)\n            img = warp_depth(img, scale, shift)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = np.ascontiguousarray(np.transpose(img.astype(np.float32) / 255.0, (2, 0, 1)))   # DHW\n        img = vol.normalize(img)\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ============================================================\n# 8. MODEL\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"Raw + first/second finite differences along the ordered depth axis, mixed by a\n    1x1 conv (V4's stem).\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = (in_depth if use_raw else 0) + ((in_depth - 1) if use_grad else 0) + \\\n                   (max(in_depth - 2, 0) if use_curv else 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: enable at least one of raw/grad/curv\")\n        self.mix = nn.Sequential(nn.Conv2d(total_in, out_channels, kernel_size=1),\n                                 nn.BatchNorm2d(out_channels), nn.GELU())\n        self.out_channels = out_channels\n\n    def forward(self, x):\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:] - x[:, :-1])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:] - 2 * x[:, 1:-1] + x[:, :-2])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass DepthAwareSegModel(nn.Module):\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):\n        return self.seg_model(self.depth_stem(x) if self.depth_stem is not None else x)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    return arch_cls(encoder_name=encoder_name, encoder_weights=CFG.encoder_weights, in_channels=in_channels,\n                     classes=1, decoder_channels=CFG.decoder_channels,\n                     decoder_attention_type=CFG.decoder_attention_type)\n\n\ndef build_model(encoder_name):\n    stem, seg_in = None, CFG.in_channels\n    if CFG.USE_DEPTH_FUSION_STEM:\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in = CFG.depth_stem_out_channels\n    try:\n        seg_model = _build_seg_backbone(encoder_name, seg_in)\n        print(f\"[backbone] using {encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {encoder_name} failed ({e}); falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in)\n    return DepthAwareSegModel(stem, seg_model)\n\n\n@torch.no_grad()\ndef run_architecture_report(model, tag, full=True):\n    print(f\"\\n{'='*70}\\nARCHITECTURE INSPECTION [{tag}]\\n{'='*70}\")\n    print(f\"Parameters: {sum(p.numel() for p in model.parameters())/1e6:.2f}M\")\n    for name, m in model.named_modules():\n        if isinstance(m, (nn.Conv1d, nn.Conv2d, nn.Conv3d)) and (m.in_channels <= 0 or m.out_channels <= 0):\n            raise RuntimeError(f\"Zero-channel layer found: {name}\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n    if full:\n        model.eval()\n        s = min(CFG.patch_size, 256)\n        out = model(torch.zeros(1, CFG.in_channels, s, s, device=CFG.device))\n        bad = torch.isnan(out).any().item() or torch.isinf(out).any().item()\n        print(f\"Dry-run {s}x{s}: {tuple(out.shape)} | NaN/Inf: {'FAIL' if bad else 'PASS'}\")\n        if bad:\n            raise RuntimeError(\"Model produced NaN/Inf on a dry run.\")\n        model.train()\n        if torch.cuda.is_available():\n            print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")\n    print(\"=\" * 70)\n\n\n# ============================================================\n# 9. LOSSES\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (probs.sum(dim=1) + t.sum(dim=1) + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    return torch.pow(1.0 - (tp + eps) / (tp + alpha * fp + beta * fn + eps), gamma).mean()\n\n\ndef soft_erode(I):\n    p1 = -F.max_pool2d(-I, (3, 1), (1, 1), (1, 0))\n    p2 = -F.max_pool2d(-I, (1, 3), (1, 1), (0, 1))\n    return torch.min(p1, p2)\n\n\ndef soft_dilate(I):\n    return F.max_pool2d(I, (3, 3), (1, 1), (1, 1))\n\n\ndef soft_open(I):\n    return soft_dilate(soft_erode(I))\n\n\ndef soft_skeletonize(I, iters=CFG.cldice_iters):\n    I1 = soft_open(I)\n    skel = F.relu(I - I1)\n    for _ in range(iters):\n        I = soft_erode(I)\n        I1 = soft_open(I)\n        delta = F.relu(I - I1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef soft_cldice(pred_probs, target, iters=CFG.cldice_iters, smooth=1.0):\n    pred_sk = soft_skeletonize(pred_probs, iters)\n    targ_sk = soft_skeletonize(target, iters)\n    tprec = (torch.sum(pred_sk * target) + smooth) / (torch.sum(pred_sk) + smooth)\n    tsens = (torch.sum(targ_sk * pred_probs) + smooth) / (torch.sum(targ_sk) + smooth)\n    return 1.0 - 2.0 * (tprec * tsens) / (tprec + tsens + 1e-8)\n\n\nclass V5ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        topo = soft_cldice(torch.sigmoid(logits), targets)\n        return (CFG.bce_weight * bce + CFG.dice_weight * dice +\n                CFG.focal_tversky_weight * tv + CFG.topology_weight * topo)\n\n\n# ============================================================\n# 10. METRICS (fixed-threshold + threshold-free histogram based)\n# ============================================================\n\ndef calculate_metrics(tp, fp, fn):\n    eps = 1e-6\n    dice = (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return calculate_metrics(self.tp, self.fp, self.fn)\n\n\ndef summarize_hist(pos, neg, beta2=0.25, smooth_bins=None):\n    \"\"\"pos/neg: 256-bin histograms of quantized probabilities (q=round(p*255)) for\n    positive / negative pixels. Rule: predict positive iff q >= k, k=1..255.\n    Returns the best-F0.5 operating point (argmax on a lightly smoothed curve, to avoid\n    razor-thin optima) plus average precision.\"\"\"\n    smooth_bins = CFG.threshold_smooth_bins if smooth_bins is None else smooth_bins\n    pos = np.asarray(pos, dtype=np.float64)\n    neg = np.asarray(neg, dtype=np.float64)\n    tp = np.cumsum(pos[::-1])[::-1][1:]\n    fp = np.cumsum(neg[::-1])[::-1][1:]\n    p_total = max(pos.sum(), 1.0)\n    prec = tp / np.maximum(tp + fp, 1.0)\n    rec = tp / p_total\n    f = (1.0 + beta2) * prec * rec / np.maximum(beta2 * prec + rec, 1e-12)\n    if smooth_bins > 0:\n        w = 2 * smooth_bins + 1\n        fs = np.convolve(np.pad(f, smooth_bins, mode=\"edge\"), np.ones(w) / w, mode=\"valid\")\n    else:\n        fs = f\n    k = int(np.argmax(fs))\n    rec_next = np.append(rec[1:], 0.0)\n    ap = float(np.sum((rec - rec_next) * prec))\n    prev = p_total / max(pos.sum() + neg.sum(), 1.0)\n    ap += float((1.0 - rec[0]) * prev)          # segment below the lowest threshold\n    return {\"f05\": float(f[k]), \"thr\": float((k + 0.5) / 255.0),\n            \"precision\": float(prec[k]), \"recall\": float(rec[k]), \"ap\": ap}\n\n\nclass ScoreHistogram:\n    def __init__(self):\n        self.pos = np.zeros(256, dtype=np.int64)\n        self.neg = np.zeros(256, dtype=np.int64)\n\n    @torch.no_grad()\n    def update(self, probs, targets):\n        q = (probs.float().clamp(0, 1) * 255.0 + 0.5).long().flatten()\n        t = targets.flatten() > 0.5\n        self.pos += torch.bincount(q[t], minlength=256).cpu().numpy()\n        self.neg += torch.bincount(q[~t], minlength=256).cpu().numpy()\n\n    def summary(self):\n        return summarize_hist(self.pos, self.neg)\n\n\ndef hist_from_stack(prob_u8_stack, labels, masks, sigma=0.0):\n    \"\"\"Histograms over TISSUE pixels of a stack of uint8 probability patches,\n    optionally Gaussian-smoothed first.\"\"\"\n    pos = np.zeros(256, dtype=np.int64)\n    neg = np.zeros(256, dtype=np.int64)\n    for j in range(len(prob_u8_stack)):\n        q = prob_u8_stack[j]\n        if sigma > 0:\n            pf = cv2.GaussianBlur(q.astype(np.float32) / 255.0, (0, 0), sigmaX=float(sigma))\n            q = np.clip(pf * 255.0 + 0.5, 0, 255).astype(np.uint8)\n        m = masks[j] > 0\n        lab = labels[j] > 0\n        pos += np.bincount(q[m & lab], minlength=256)\n        neg += np.bincount(q[m & ~lab], minlength=256)\n    return pos, neg\n\n\ndef evaluate_prob_map_hist(prob_map, gt, mask):\n    q = np.clip(prob_map * 255.0 + 0.5, 0, 255).astype(np.uint8)\n    m = mask > 0\n    lab = gt > 0\n    return summarize_hist(np.bincount(q[m & lab], minlength=256), np.bincount(q[m & ~lab], minlength=256))\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    return calculate_metrics((preds * gt).sum(), (preds * (1.0 - gt)).sum(), ((1.0 - preds) * gt).sum())\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 11. SHARED DATA (built once, reused by every member)\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING SHARED TRAIN/VAL PATCH GRID\")\nprint(\"=\" * 70)\n\n_shared_masks, _shared_labels = {}, {}\nall_samples = []\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {mask.shape} {len(coords)} candidate patches\")\n    _shared_masks[fid] = mask\n    _shared_labels[fid] = labels\n    all_samples.extend([(fid, y, x) for y, x in coords])\n\n_groups = {}\nfor fid, y, x in all_samples:\n    _groups.setdefault((fid, y // 768, x // 768), []).append((fid, y, x))\n_keys = list(_groups.keys())\nrandom.Random(CFG.seed).shuffle(_keys)\n_val_keys = set(_keys[:max(1, int(len(_keys) * CFG.val_fraction))])\ntrain_grid, val_samples = [], []\nfor key, items in _groups.items():\n    (val_samples if key in _val_keys else train_grid).extend(items)\nprint(f\"Spatial train: {len(train_grid)} | Spatial val: {len(val_samples)}\")\n\n\ndef build_train_samples(seed):\n    r = random.Random(seed)\n    s = balance_positive_patches(train_grid, _shared_labels, CFG.patch_size,\n                                  positive_threshold=CFG.positive_patch_fraction,\n                                  target_positive_ratio=CFG.target_positive_patch_ratio,\n                                  max_positive_repeat=CFG.max_positive_repeat, rng=r)\n    return domain_balance(s, r)\n\n\n# --- validation-calibration subset (evenly spaced), with tissue-masked labels ---------\nVAL_CAL_IDXS = np.unique(np.linspace(0, len(val_samples) - 1,\n                                     min(CFG.val_calibration_max_patches, len(val_samples))).astype(int))\n_S = CFG.patch_size\nVAL_LABELS = np.stack([_shared_labels[val_samples[i][0]][val_samples[i][1]:val_samples[i][1] + _S,\n                                                          val_samples[i][2]:val_samples[i][2] + _S]\n                       for i in VAL_CAL_IDXS]).astype(np.uint8)\nVAL_MASKS = np.stack([_shared_masks[val_samples[i][0]][val_samples[i][1]:val_samples[i][1] + _S,\n                                                        val_samples[i][2]:val_samples[i][2] + _S]\n                      for i in VAL_CAL_IDXS]).astype(np.uint8)\nprint(f\"Validation calibration subset: {len(VAL_CAL_IDXS)} patches\")\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\nprint(f\"Test fragment {CFG.test_frag}: {test_mask.shape}\")\n\npool_source = []\nif CFG.USE_HIST_MATCH_AUG:\n    for fid in CFG.train_frags:\n        c = generate_grid_coords(_shared_masks[fid], CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n        pool_source.extend([(fid, y, x) for y, x in c])\n    print(f\"[strict protocol] histogram-match pool from TRAIN fragments only ({len(pool_source)} candidates).\")\n\n\n# ============================================================\n# 12. INFERENCE HELPERS\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    g1d = np.exp(-(ax ** 2) / (2.0 * (size * sigma_frac) ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\": return x\n    if mode == \"hflip\": return torch.flip(x, dims=[3])\n    if mode == \"vflip\": return torch.flip(x, dims=[2])\n    if mode == \"hvflip\": return torch.flip(x, dims=[2, 3])\n    if mode == \"rot90\": return torch.rot90(x, k=1, dims=[2, 3])\n    if mode == \"rot180\": return torch.rot90(x, k=2, dims=[2, 3])\n    if mode == \"rot270\": return torch.rot90(x, k=3, dims=[2, 3])\n    if mode == \"transpose\": return x.transpose(2, 3)\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\": return pred\n    if mode == \"hflip\": return torch.flip(pred, dims=[2])\n    if mode == \"vflip\": return torch.flip(pred, dims=[1])\n    if mode == \"hvflip\": return torch.flip(pred, dims=[1, 2])\n    if mode == \"rot90\": return torch.rot90(pred, k=-1, dims=[1, 2])\n    if mode == \"rot180\": return torch.rot90(pred, k=-2, dims=[1, 2])\n    if mode == \"rot270\": return torch.rot90(pred, k=-3, dims=[1, 2])\n    if mode == \"transpose\": return pred.transpose(1, 2)\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    acc = None\n    for mode in modes:\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            probs = torch.sigmoid(model(apply_tta_tensor(x, mode)))\n        probs = invert_tta_prediction(probs[:, 0], mode).float() / len(modes)\n        acc = probs if acc is None else acc + probs\n    return acc.cpu().numpy().astype(np.float32)\n\n\n@torch.no_grad()\ndef predict_val_stack(model, val_ds, idxs):\n    \"\"\"TTA'd probabilities (uint8-quantized) for the calibration subset of validation.\"\"\"\n    loader = DataLoader(Subset(val_ds, [int(i) for i in idxs]), batch_size=CFG.infer_batch, shuffle=False,\n                        num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"))\n    model.eval()\n    out = []\n    for imgs, _ in loader:\n        p = predict_batch_tta(model, imgs.numpy())\n        out.append(np.clip(p * 255.0 + 0.5, 0, 255).astype(np.uint8))\n    return np.concatenate(out, axis=0)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"  Inference patches: {len(coords)} (TTA views: {len(CFG.tta_modes) if CFG.use_tta else 1})\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        probs = predict_batch_tta(model, np.stack(batch_imgs).astype(np.float32))\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear()\n        batch_coords.clear()\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        batch_imgs.append(vol.normalize(raw))\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm(model, vol, mask, patch_size, stride, max_patches, batch_size):\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        imgs = [vol.normalize(vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0)\n                for y, x in coords[start:start + batch_size]]\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            model(torch.from_numpy(np.stack(imgs)).to(CFG.device))\n    model.eval()\n    cleanup_memory()\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    n, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if n <= 1:\n        return binary\n    out = np.zeros_like(binary, dtype=np.uint8)\n    for i in range(1, n):\n        if stats[i, cv2.CC_STAT_AREA] >= min_size:\n            out[labels == i] = 1\n    return out\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    return remove_small_components(binary, CFG.min_component_size).astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\n# ============================================================\n# 13. TRAIN + INFER ONE ENSEMBLE MEMBER\n# ============================================================\n\nclass EMAModel:\n    def __init__(self, model, decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n\ndef train_and_infer_one_member(member_idx, seed, encoder_name, window):\n    start, stop, step = window\n    depth_indices = list(range(start, stop, step))\n    if len(depth_indices) != CFG.n_slices or max(depth_indices) > CFG.max_slice_index:\n        raise ValueError(f\"window {window} must give {CFG.n_slices} slices <= {CFG.max_slice_index}\")\n    print(f\"\\n{'#'*70}\\n# MEMBER {member_idx+1}/{CFG.N_ENSEMBLE_MODELS} seed={seed} encoder={encoder_name} \"\n          f\"window={window} (slices {depth_indices[0]}..{depth_indices[-1]}, step {step})\\n{'#'*70}\")\n    set_seed(seed)\n\n    train_volumes = {}\n    for fid in CFG.train_frags:\n        v = FragmentVolume(os.path.join(CFG.base_dir, fid), depth_indices)\n        if CFG.WARM_UP_VOLUME_CACHE:\n            v.warm_up_cache(f\"fragment {fid}\")\n        v.compute_fragment_stats(_shared_masks[fid], CFG.patch_size)\n        print(f\"  fragment {fid}: legacy mean={v.frag_mean:.2f} std={v.frag_std:.2f}\")\n        train_volumes[fid] = v\n\n    train_samples = build_train_samples(seed)\n\n    hist_pool = None\n    if CFG.USE_HIST_MATCH_AUG and pool_source:\n        picked = random.sample(pool_source, min(CFG.hist_match_pool_size, len(pool_source)))\n        hist_pool = [train_volumes[fid].read_patch(y, x, CFG.patch_size) for fid, y, x in picked]\n\n    train_ds = InkPatchDataset(train_volumes, _shared_labels, train_samples, CFG.patch_size,\n                                transform=build_train_transform(), jitter=CFG.train_jitter,\n                                hist_match_pool=hist_pool, train_mode=True)\n    val_ds = InkPatchDataset(train_volumes, _shared_labels, val_samples, CFG.patch_size,\n                              transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n    train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True, num_workers=CFG.num_workers,\n                               pin_memory=(CFG.device == \"cuda\"), drop_last=CFG.drop_last,\n                               persistent_workers=CFG.num_workers > 0)\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers,\n                             pin_memory=(CFG.device == \"cuda\"), persistent_workers=CFG.num_workers > 0)\n\n    pos_frac = estimate_positive_fraction(_shared_labels, train_samples, CFG.patch_size)\n    initial_bias = math.log(pos_frac / max(1 - pos_frac, 1e-6))\n    pos_weight = float(np.clip(np.sqrt((1 - pos_frac) / pos_frac), 1.0, 8.0))\n    print(f\"  positive fraction={pos_frac:.5f} bias={initial_bias:.3f} pos_weight={pos_weight:.2f}\")\n\n    model = build_model(encoder_name).to(CFG.device)\n    run_architecture_report(model, f\"member{member_idx}\", full=(member_idx == 0))\n    with torch.no_grad():\n        try:\n            model.seg_model.segmentation_head[0].bias.fill_(initial_bias)\n        except Exception:\n            pass\n\n    criterion = V5ComboLoss(pos_weight=torch.tensor([pos_weight], dtype=torch.float32, device=CFG.device))\n\n    enc_p, dec_p, stem_p = [], [], []\n    for name, p in model.named_parameters():\n        if p.requires_grad:\n            (stem_p if name.startswith(\"depth_stem.\") else\n             enc_p if name.startswith(\"seg_model.encoder.\") else dec_p).append(p)\n    groups = [{\"params\": enc_p, \"lr\": CFG.encoder_lr}, {\"params\": dec_p, \"lr\": CFG.decoder_lr}]\n    if stem_p:\n        groups.append({\"params\": stem_p, \"lr\": CFG.depth_stem_lr})\n    optimizer = torch.optim.AdamW(groups, weight_decay=CFG.weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n    ema = EMAModel(model, CFG.ema_decay) if CFG.USE_EMA else None\n\n    def run_epoch(loader, train_mode, want_hist=False):\n        model.train(train_mode)\n        total_loss = 0.0\n        acc = GlobalConfusionAccumulator()\n        hist = ScoreHistogram() if want_hist else None\n        for imgs, masks in loader:\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            masks = masks.to(CFG.device, non_blocking=True)\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n            with torch.set_grad_enabled(train_mode):\n                with autocast(enabled=(CFG.device == \"cuda\")):\n                    logits = model(imgs)\n                    loss = criterion(logits, masks)\n                if train_mode:\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    if ema is not None:\n                        ema.update(model)\n            probs = torch.sigmoid(logits.detach())\n            acc.update(probs, masks, 0.5)\n            if hist is not None:\n                hist.update(probs, masks)\n            total_loss += float(loss.item())\n            del imgs, masks, logits, probs\n        return total_loss / max(len(loader), 1), acc.compute(), hist\n\n    print(f\"\\n  Training up to {CFG.epochs} epochs (patch={CFG.patch_size}, batch={CFG.batch_size}); \"\n          f\"checkpoint metric = val best-threshold F0.5 ...\")\n    best_score, no_improve = -1.0, 0\n    best_state, best_used = None, None\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n        train_loss, train_m, _ = run_epoch(train_loader, True)\n\n        # Evaluate raw AND EMA weights; checkpoint the better one; ALWAYS continue from raw.\n        raw_state = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        _, _, h_raw = run_epoch(val_loader, False, want_hist=True)\n        s_raw = h_raw.summary()\n        sel, cand, used = s_raw, raw_state, \"raw\"\n        if ema is not None:\n            model.load_state_dict(ema.shadow)\n            _, _, h_ema = run_epoch(val_loader, False, want_hist=True)\n            model.load_state_dict(raw_state)\n            s_ema = h_ema.summary()\n            print(f\"    [checkpoint choice] EMA F0.5*={s_ema['f05']:.4f} raw F0.5*={s_raw['f05']:.4f}\")\n            if s_ema[\"f05\"] >= s_raw[\"f05\"]:\n                sel, cand, used = s_ema, {k: v.clone() for k, v in ema.shadow.items()}, \"ema\"\n        scheduler.step()\n\n        print(f\"  [member {member_idx+1}][{epoch:02d}/{CFG.epochs}] {time.time()-t0:.0f}s \"\n              f\"train_loss={train_loss:.4f} train_dice={train_m['dice']:.4f} | \"\n              f\"val F0.5*={sel['f05']:.4f} (thr {sel['thr']:.2f}, P={sel['precision']:.3f}, \"\n              f\"R={sel['recall']:.3f}) AP={sel['ap']:.4f} [{used}]\")\n\n        if sel[\"f05\"] > best_score:\n            best_score, no_improve, best_state, best_used = sel[\"f05\"], 0, cand, used\n            print(f\"    *** new best val F0.5*={best_score:.4f} ({used}) ***\")\n        else:\n            no_improve += 1\n            if no_improve >= CFG.early_stop_patience:\n                print(\"    Early stopping.\")\n                break\n        cleanup_memory()\n\n    model.load_state_dict(best_state)\n    model.eval()\n    print(f\"  Member {member_idx+1} best val F0.5*: {best_score:.4f} ({best_used})\")\n\n    # --- TTA'd validation predictions for ensemble-level calibration ---\n    val_stack = predict_val_stack(model, val_ds, VAL_CAL_IDXS)\n    solo = summarize_hist(*hist_from_stack(val_stack, VAL_LABELS, VAL_MASKS, 0.0))\n    print(f\"  Member {member_idx+1} TTA'd val (tissue only): F0.5*={solo['f05']:.4f} thr={solo['thr']:.3f} AP={solo['ap']:.4f}\")\n\n    for v in train_volumes.values():\n        v.close()\n    del train_loader, val_loader, train_ds, val_ds\n    cleanup_memory()\n\n    # --- fragment-1 inference with this member's own depth window ---\n    test_vol = FragmentVolume(test_dir, depth_indices)\n    if CFG.WARM_UP_VOLUME_CACHE:\n        test_vol.warm_up_cache(f\"fragment {CFG.test_frag} (held-out)\")\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    print(f\"  [member {member_idx+1}] {len(CFG.tta_modes) if CFG.use_tta else 1}-view TTA inference ...\")\n    prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.infer_batch)\n    if CFG.use_adabn:\n        model = recalibrate_batchnorm(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                       CFG.adabn_max_patches, CFG.infer_batch)\n        prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.infer_batch)\n    test_vol.close()\n    cleanup_memory()\n\n    return {\"member_idx\": member_idx, \"seed\": seed, \"encoder\": encoder_name, \"window\": list(window),\n            \"best_val_f05\": float(best_score), \"solo_val\": solo, \"val_stack\": val_stack, \"probability\": prob}\n\n\n# ============================================================\n# 14. RUN THE ENSEMBLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(f\"STARTING V5.3 ({'single model' if CFG.N_ENSEMBLE_MODELS == 1 else f'{CFG.N_ENSEMBLE_MODELS}-model ensemble'}, single split)\")\nprint(\"=\" * 70)\n\nmember_results = []\nfor i in range(CFG.N_ENSEMBLE_MODELS):\n    member_results.append(train_and_infer_one_member(\n        i, CFG.ensemble_seeds[i], CFG.ensemble_encoders[i], CFG.ensemble_windows[i]))\n    cleanup_memory()\n\n\n# ============================================================\n# 15. ENSEMBLE CALIBRATION ON VALIDATION (labels of fragment 1 not used here)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"ENSEMBLE CALIBRATION (validation only): smoothing sigma + F0.5 threshold\")\nprint(\"=\" * 70)\n\nens_val = np.mean(np.stack([r[\"val_stack\"] for r in member_results]).astype(np.float32), axis=0)\nens_val = np.clip(ens_val + 0.5, 0, 255).astype(np.uint8)\n\nbest_sigma, best_sum = 0.0, None\nfor sigma in CFG.smooth_sigma_candidates:\n    s = summarize_hist(*hist_from_stack(ens_val, VAL_LABELS, VAL_MASKS, sigma))\n    print(f\"  sigma={sigma:>4}: val F0.5*={s['f05']:.4f} thr={s['thr']:.3f} \"\n          f\"P={s['precision']:.3f} R={s['recall']:.3f} AP={s['ap']:.4f}\")\n    if best_sum is None or s[\"f05\"] > best_sum[\"f05\"] + 1e-4:\n        best_sigma, best_sum = sigma, s\nfinal_threshold = best_sum[\"thr\"]\nprint(f\"\\n  -> chosen sigma={best_sigma}, threshold={final_threshold:.3f} \"\n      f\"(val F0.5*={best_sum['f05']:.4f}); frozen before touching fragment-1 labels.\")\n\n\n# ============================================================\n# 16. FINAL TEST MAP + DIAGNOSTIC REPORT (labels loaded only now)\n# ============================================================\n\ndef finalize_map(prob, sigma):\n    p = cv2.GaussianBlur(prob, (0, 0), sigmaX=float(sigma)) if sigma > 0 else prob\n    return (p * test_mask).astype(np.float32)        # labels are tissue-masked, so predictions are too\n\n\nensembled_probability = finalize_map(np.mean([r[\"probability\"] for r in member_results], axis=0).astype(np.float32),\n                                     best_sigma)\nfinal_prediction = postprocess(ensembled_probability, final_threshold)\notsu_thr = compute_otsu_threshold(ensembled_probability, test_mask, fallback=final_threshold)\n\ntest_labels = load_ink_labels(test_dir)\ngt_test = (test_labels * test_mask).astype(np.float32) if test_labels is not None else None\n\nraw_metrics = post_metrics = otsu_metrics = oracle = None\nmember_solo = []\nif gt_test is not None:\n    print(\"\\n\" + \"=\" * 70)\n    print(\"FRAGMENT 1 RESULTS (labels used for REPORTING only)\")\n    print(\"=\" * 70)\n    for r in member_results:\n        thr = r[\"solo_val\"][\"thr\"]\n        m = evaluate_probability_map(r[\"probability\"] * test_mask, gt_test, thr)\n        member_solo.append(m)\n        print(f\"  member {r['member_idx']+1} ({r['encoder']}, window {r['window']}) at its own val thr={thr:.3f}: \"\n              f\"F0.5={m['fbeta0.5']:.4f} P={m['precision']:.3f} R={m['recall']:.3f} Dice={m['dice']:.4f}\")\n\n    raw_metrics = evaluate_probability_map(ensembled_probability, gt_test, final_threshold)\n    pf = final_prediction.astype(np.float32)\n    post_metrics = calculate_metrics((pf * gt_test).sum(), (pf * (1 - gt_test)).sum(), ((1 - pf) * gt_test).sum())\n    otsu_metrics = evaluate_probability_map(ensembled_probability, gt_test, otsu_thr)\n    oracle = evaluate_prob_map_hist(ensembled_probability, gt_test, test_mask)\n\n    print(f\"\\n  ENSEMBLE @ val-chosen thr={final_threshold:.3f}, sigma={best_sigma}\")\n    print(f\"    raw:           {raw_metrics}\")\n    print(f\"    postprocessed: {post_metrics}\")\n    print(f\"  [diagnostic] Otsu thr={otsu_thr:.3f} would give F0.5={otsu_metrics['fbeta0.5']:.4f} (NOT used)\")\n    print(f\"  [analysis only] ranking quality: AP={oracle['ap']:.4f} | oracle F0.5={oracle['f05']:.4f} \"\n          f\"at thr={oracle['thr']:.3f} (upper bound, NOT a reportable result)\")\n\n\n# ============================================================\n# 17. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_v5_3.npy\")\nnp.save(prob_path, ensembled_probability)\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_v5_3.png\")\ncv2.imwrite(pred_path, (final_prediction * 255).astype(np.uint8))\nfor r in member_results:\n    np.save(os.path.join(CFG.out_dir, f\"fragment1_probability_member{r['member_idx']}_v5_3.npy\"), r[\"probability\"])\n\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump({\n        \"chosen_sigma\": best_sigma, \"chosen_threshold\": final_threshold, \"val_calibration\": best_sum,\n        \"ensembled_raw\": raw_metrics, \"ensembled_postprocessed\": post_metrics,\n        \"otsu_diagnostic\": otsu_metrics, \"ranking_analysis_only\": oracle,\n        \"members\": [{\"idx\": r[\"member_idx\"], \"seed\": r[\"seed\"], \"encoder\": r[\"encoder\"], \"window\": r[\"window\"],\n                     \"best_val_f05\": r[\"best_val_f05\"], \"solo_val\": r[\"solo_val\"]} for r in member_results],\n        \"config\": cfg_to_dict(CFG),\n    }, f, indent=2, default=float)\nprint(f\"\\nSaved metrics: {CFG.metrics_path}\")\n\n\n# ============================================================\n# 18. VISUALIZATION\n# ============================================================\n\ndef make_confusion_overlay(gray_u8, gt_bin, pred_bin, alpha=0.55):\n    base = np.stack([gray_u8] * 3, axis=-1).astype(np.float32)\n    ov = base.copy()\n    g, p = gt_bin > 0, pred_bin > 0\n    for m, col in ((p & g, [0, 255, 0]), (p & ~g, [255, 0, 0]), (~p & g, [255, 255, 0])):\n        ov[m] = (1 - alpha) * base[m] + alpha * np.array(col, dtype=np.float32)\n    return np.clip(ov, 0, 255).astype(np.uint8)\n\n\ndef save_overview():\n    \"\"\"Saves ONE figure with, left to right: input slice, ground truth,\n    probability map, thresholded prediction, and a TP/FP/FN overlay with an\n    explicit color-swatch legend (not just a title) -- green=TP, red=FP,\n    yellow=FN, unlabeled tissue stays grayscale.\"\"\"\n    ref = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{CFG.mask_reference_slice:02d}.tif\"))\n    sc = 2000 / max(ref.shape)\n    small = cv2.resize(ref, None, fx=sc, fy=sc, interpolation=cv2.INTER_AREA)\n    small_u8 = cv2.normalize(small, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n    pred_s = cv2.resize((final_prediction * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    prob_s = cv2.resize((ensembled_probability * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_s = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n        ov = make_confusion_overlay(small_u8, (gt_s > 127).astype(np.uint8), (pred_s > 127).astype(np.uint8))\n        fig, ax = plt.subplots(1, 5, figsize=(28, 6))\n        ax[0].imshow(small, cmap=\"gray\"); ax[0].set_title(f\"Input (slice {CFG.mask_reference_slice})\", fontsize=12)\n        ax[1].imshow(gt_s, cmap=\"gray\"); ax[1].set_title(\"Ground truth\", fontsize=12)\n        ax[2].imshow(prob_s, cmap=\"gray\"); ax[2].set_title(f\"Probability (sigma={best_sigma})\", fontsize=12)\n        ax[3].imshow(pred_s, cmap=\"gray\"); ax[3].set_title(f\"Prediction (thr={final_threshold:.2f})\", fontsize=12)\n        ax[4].imshow(ov); ax[4].set_title(\"Prediction vs. ground truth\", fontsize=12)\n        from matplotlib.patches import Patch\n        legend_handles = [\n            Patch(facecolor=np.array([0, 255, 0]) / 255, label=\"True Positive (correct ink)\"),\n            Patch(facecolor=np.array([255, 0, 0]) / 255, label=\"False Positive (over-predicted)\"),\n            Patch(facecolor=np.array([255, 255, 0]) / 255, label=\"False Negative (missed ink)\"),\n        ]\n        ax[4].legend(handles=legend_handles, loc=\"upper center\", bbox_to_anchor=(0.5, -0.05),\n                     ncol=1, fontsize=10, frameon=True, framealpha=0.9)\n    else:\n        fig, ax = plt.subplots(1, 3, figsize=(18, 6))\n        ax[0].imshow(small, cmap=\"gray\"); ax[0].set_title(f\"Input (slice {CFG.mask_reference_slice})\", fontsize=12)\n        ax[1].imshow(prob_s, cmap=\"gray\"); ax[1].set_title(f\"Probability (sigma={best_sigma})\", fontsize=12)\n        ax[2].imshow(pred_s, cmap=\"gray\"); ax[2].set_title(f\"Prediction (thr={final_threshold:.2f})\", fontsize=12)\n    for a in ax:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_v5_3_overview.png\")\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(\"Saved comparison figure (input | ground truth | probability | prediction | TP/FP/FN overlay):\", path)\n\n\nsave_overview()\ncleanup_memory()\n\n\n# ============================================================\n# 19. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V5.3 COMPLETE\")\nprint(\"=\" * 70)\nfor r in member_results:\n    print(f\"Member {r['member_idx']+1}: {r['encoder']} window={r['window']} \"\n          f\"best val F0.5*={r['best_val_f05']:.4f} | TTA'd val F0.5*={r['solo_val']['f05']:.4f}\")\nprint(f\"\\nEnsemble: sigma={best_sigma} threshold={final_threshold:.3f} (chosen on validation by F0.5)\")\nif post_metrics is not None:\n    print(f\"Fragment 1 F0.5 (raw): {raw_metrics['fbeta0.5']:.4f} | (postprocessed): {post_metrics['fbeta0.5']:.4f} \"\n          f\"| Dice: {post_metrics['dice']:.4f} | AP: {oracle['ap']:.4f}\")\nprint(f\"\\nProbability map: {prob_path}\\nPrediction: {pred_path}\\nMetrics: {CFG.metrics_path}\")\nprint(\"\\n=== V5.3 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T05:29:11.895854Z","iopub.execute_input":"2026-09-30T05:29:11.896516Z"}},"outputs":[{"name":"stdout","text":"   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 4.0 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.7/2.7 MB 37.6 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.30\n======================================================================\nBUILDING SHARED TRAIN/VAL PATCH GRID\n======================================================================\nFragment 2: (14830, 9506) 6146 candidate patches\nFragment 3: (7606, 5249) 1607 candidate patches\nSpatial train: 6151 | Spatial val: 1602\nValidation calibration subset: 600 patches\nTest fragment 1: (8181, 6330)\n[strict protocol] histogram-match pool from TRAIN fragments only (1247 candidates).\n\n======================================================================\nSTARTING V5.3 (single model, single split)\n======================================================================\n\n######################################################################\n# MEMBER 1/1 seed=42 encoder=tu-convnext_tiny window=(12, 38, 1) (slices 12..37, step 1)\n######################################################################\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS INK DETECTION - V5.3 \"SINGLE-MODEL, DIAGNOSTICS-DRIVEN\"\n# ============================================================\n# Built on V5.2. ONE change from V5.2: the ensemble is dropped down to its\n# single stronger member -- tu-convnext_tiny + window (12,38,1). The second\n# member (efficientnet-b4 + a wide window) validated worse AND roughly\n# doubled total run time, so it was pure cost with no benefit. The ensemble\n# loop/calibration code is left in place (harmless at N=1) so a second member\n# can be re-added later with a one-line CFG change (N_ENSEMBLE_MODELS,\n# ensemble_seeds/encoders/windows).\n#\n# Also added: a dedicated comparison figure per Fragment 1 --\n# input / ground truth / probability / prediction / TP-FP-FN overlay -- saved\n# as its own function (see PART 18) with clearly distinct overlay colors\n# (green=TP, red=FP, yellow=FN), independent of the run's console output.\n#\n# Everything else is V5.2 as diagnosed from your actual run. Every change\n# below is tied to a number in that diagnostics run:\n#\n#  PART 1 (fragment 1): AP=0.50 vs prevalence 0.18 -> real signal, but RANKING is\n#  the ceiling (oracle F0.5 0.513 == best threshold you used). However the\n#  ensemble's OTSU threshold (0.424) scored F0.5 0.408, i.e. ~0.10 WORSE than the\n#  validation thresholds, and best-Dice threshold (0.51) != best-F0.5 threshold\n#  (0.67). => (1) Otsu is no longer used to decide anything. (2) All thresholds\n#  are chosen by F0.5 (the competition metric), ON THE ENSEMBLED, TTA'd\n#  VALIDATION PREDICTIONS, then frozen. (3) Checkpoints are selected by\n#  best-threshold F0.5 instead of Dice@0.5.\n#\n#  PART 2 (depth): fragments 1 and 3 share almost the same mean depth profile\n#  (r=0.996) while fragment 2 differs (r=0.83; broader/lower peak, trough ~6\n#  slices deeper). The fixed window (12..37) leaves 43-55% of the ink-vs-\n#  background contrast outside it, and fragment 3's strongest contrast sits on\n#  the window's last slice. => (4) member windows now cover ~8..58 (stride 2 keeps\n#  the channel count and I/O identical); (5) per-fragment PER-DEPTH mean-profile\n#  normalization removes the fragment-specific depth profile (label-free, same\n#  status as the old per-fragment z-score); (6) depth-warp augmentation\n#  (random z-scale/shift) targets the fragment-2-style profile shift;\n#  (7) domain-balanced sampling stops fragment 2 (78% of patches) from\n#  dominating fragment 3, the fragment whose depth profile matches the test.\n#\n#  ALSO: (8) validation-selected Gaussian smoothing of the probability map\n#  (ink strokes are spatially coherent; sigma chosen by val F0.5, 0 allowed),\n#  (9) predictions are masked by the tissue mask before scoring (labels are),\n#  (10) cache warm-up per volume (fixes cold random-read I/O on Kaggle).\n#\n# NOT included on purpose: the auxiliary \"ink differential\" head. Part 2 shows\n# the ink depth signature differs between fragments 2 and 3 (peak slice 17 vs\n# 24, different sign lobes), so supervising a fixed differential signature would\n# mostly learn fragment-specific patterns.\n#\n# HONEST EXPECTATION: (1)-(3) and (9) are corrections and should recover the\n# ~0.10 lost by Otsu plus a small ensemble gain. (4)-(8) are evidence-motivated\n# HYPOTHESES about the ranking ceiling; each has its own CFG flag so you can\n# ablate them. Fragment 1's ground truth is only used for reporting at the end.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.amp import autocast as _autocast_new, GradScaler as _GradScaler_new\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\n\ndef autocast(enabled=True):\n    \"\"\"Thin wrapper over torch.amp (torch.cuda.amp is deprecated).\"\"\"\n    return _autocast_new(\"cuda\" if torch.cuda.is_available() else \"cpu\", enabled=enabled)\n\n\nclass GradScaler(_GradScaler_new):\n    def __init__(self, enabled=True):\n        super().__init__(\"cuda\" if torch.cuda.is_available() else \"cpu\", enabled=enabled)\n\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    # --- depth windows (start, stop_exclusive, step); every window must give n_slices\n    #     channels. Diagnostics: 43-55% of ink contrast lies OUTSIDE (12..37), so the\n    #     wide windows use stride 2 over ~6..58 (same channels, same I/O cost). -----\n    n_slices = 26\n    in_channels = n_slices\n    max_slice_index = 64\n    # V5.3: single-model run -- efficientnet-b4 (former member 2) validated\n    # worse AND cost roughly 2x total time for no benefit, so it's dropped.\n    # Kept only the tu-convnext_tiny + (12,38,1) config that was actually the\n    # stronger of the two. The ensemble loop below still runs (harmless at\n    # N=1) so re-enabling a second member later is a one-line CFG change.\n    ensemble_windows = [(12, 38, 1)]\n    mask_reference_slice = 24          # used only for the fallback tissue mask + overview image\n\n    patch_size = 320\n    train_stride = 128\n    test_stride = 64\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.45\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 10\n    early_stop_patience = 3\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n\n    # --- backbone (verified: smp.Unet + tu-convnext_tiny; not UnetPlusPlus) ---------\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    USE_DEPTH_FUSION_STEM = True\n    depth_stem_out_channels = 26\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- ensemble (member i uses window i) ------------------------------------------\n    N_ENSEMBLE_MODELS = 1                                  # raise + extend the lists above to re-enable\n    ensemble_seeds = [42]\n    ensemble_encoders = [\"tu-convnext_tiny\"]\n\n    # --- LOSS -------------------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.30\n    focal_tversky_weight = 0.25\n    topology_weight = 0.15\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n    cldice_iters = 8\n\n    # --- augmentation -----------------------------------------------------------------\n    USE_PHYSICAL_SYNTHETIC_INK = True\n    physical_ink_p = 0.15\n    physical_ink_max_amplitude = 60\n\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n    USE_SHADOW = True\n    shadow_p = 0.20\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # Depth-warp: random z-scale / z-shift of the channel axis (linear interpolation\n    # between slices). Targets the fragment-2-style profile shift seen in Part 2.\n    USE_DEPTH_WARP_AUG = True\n    depth_warp_p = 0.40\n    depth_warp_scale_range = (0.75, 1.33)\n    depth_warp_shift_range = (-2.0, 2.0)      # in channel units\n\n    # Histogram-match aug: pool from TRAIN fragments only (strict protocol).\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # Domain-balanced sampling: share_f = (1-s)*natural_share + s*(1/k). Total sample\n    # count is unchanged (bigger fragment subsampled, smaller one oversampled).\n    DOMAIN_BALANCE = True\n    domain_balance_strength = 0.5\n\n    # --- normalization ----------------------------------------------------------------\n    # \"per_fragment_per_slice\": subtract the fragment's own mean tissue intensity PER\n    #     DEPTH SLICE, divide by one per-fragment scale (label-free, computed from\n    #     tissue pixels of each fragment separately -- same status as the old z-score).\n    # \"per_fragment_zscore\": legacy scalar mean/std from the middle slice.\n    NORMALIZATION_MODE = \"per_fragment_per_slice\"\n    PER_SLICE_STD_SCALING = False\n    profile_row_stride = 8\n    frag_stats_sample_patches = 60\n\n    # --- validation calibration (threshold + smoothing chosen on ENSEMBLED, TTA'd val) -\n    val_calibration_max_patches = 600\n    smooth_sigma_candidates = [0.0, 1.0, 2.0, 3.0, 4.0, 6.0]\n    threshold_smooth_bins = 2\n\n    # --- TTA: full 8-view D4 group ------------------------------------------------------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"hvflip\", \"rot90\", \"rot180\", \"rot270\", \"transpose\"]\n\n    use_adabn = False\n    adabn_max_patches = 2000\n\n    USE_EMA = True\n    ema_decay = 0.999\n\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    CLAHE_MODE = \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # Forces one fast sequential read of every opened slice (cold random reads on\n    # Kaggle's input filesystem can cost seconds per patch).\n    WARM_UP_VOLUME_CACHE = True\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    viz_dir = os.path.join(out_dir, \"v5_3_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v5_3_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        ref_path = os.path.join(frag_dir, \"surface_volume\", f\"{CFG.mask_reference_slice:02d}.tif\")\n        ref = tifffile.imread(ref_path)\n        mask = (ref > ref.mean() * 0.15).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion emptied the mask for {frag_dir}; keeping un-eroded mask.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / BALANCING\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            if mask[y:y + patch_size, x:x + patch_size].mean() > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3, rng=random):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = float(labels_full[fid][y:y + patch_size, x:x + patch_size].mean())\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = max(int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio), len(positive))\n    negative_selected = rng.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    balanced = positive * repeat + negative_selected\n    rng.shuffle(balanced)\n    print(f\"  Positive patches: {len(positive)} | Negative patches: {len(negative)} | \"\n          f\"Balanced: {len(balanced)} (positive ratio={len(positive)*repeat/max(len(balanced),1):.3f})\")\n    return balanced\n\n\ndef domain_balance(samples, rng):\n    \"\"\"Blend natural fragment shares toward equal shares, keeping the TOTAL constant\n    (subsample the larger fragment, oversample the smaller one).\"\"\"\n    if not CFG.DOMAIN_BALANCE or len(CFG.train_frags) < 2 or not samples:\n        return samples\n    by = {}\n    for s in samples:\n        by.setdefault(s[0], []).append(s)\n    fids = sorted(by)\n    n_total = len(samples)\n    k = len(fids)\n    natural = np.array([len(by[f]) for f in fids], dtype=np.float64)\n    natural /= natural.sum()\n    lam = CFG.domain_balance_strength\n    share = (1.0 - lam) * natural + lam * (1.0 / k)\n    out = []\n    for f, sh in zip(fids, share):\n        n = int(round(sh * n_total))\n        pool = by[f]\n        if n <= len(pool):\n            out.extend(rng.sample(pool, n))\n        else:\n            out.extend(pool + [rng.choice(pool) for _ in range(n - len(pool))])\n    rng.shuffle(out)\n    print(\"  Domain balance: \" + \" | \".join(\n        f\"frag {f}: {len(by[f])} -> {int(round(sh * n_total))}\" for f, sh in zip(fids, share)))\n    return out\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME (per-depth normalization, cache warm-up)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None\n        self.frag_mean = None\n        self.frag_std = None\n        self.slice_mean = None      # per-depth tissue mean (0..255 units, after contrast LUT)\n        self.slice_std = None\n        self.norm_scale = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _to_uint8(block):\n        block = np.asarray(block)\n        if block.dtype != np.uint8:\n            block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        return np.ascontiguousarray(block)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid = len(self._slices) // 2\n        ref_slice = self._to_uint8(self._slices[mid])\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def _apply_contrast(self, block):\n        if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n            return self._clahe.apply(block)\n        if CFG.CLAHE_MODE == \"global_shared\":\n            return cv2.LUT(block, self._shared_clahe_lut)\n        return block\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        out = np.zeros((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = self._to_uint8(s[y:y + size, x:x + size])\n            block = self._apply_contrast(block)\n            hh, ww = block.shape\n            out[i, :hh, :ww] = block[:size, :size]\n        return out\n\n    def warm_up_cache(self, tag=\"\"):\n        \"\"\"One fast sequential read of every opened slice so later random patch reads\n        are page-cache hits (result of .sum() is discarded; nothing is retained).\"\"\"\n        self._ensure_open()\n        t0 = time.time()\n        total = 0\n        for arr in self._slices:\n            _ = np.asarray(arr).sum(dtype=np.int64)\n            total += arr.nbytes\n        dt = time.time() - t0\n        print(f\"  [cache warm-up{(' ' + tag) if tag else ''}] {len(self._slices)} slices, \"\n              f\"{total/1e9:.2f} GB in {dt:.1f}s ({total/1e6/max(dt,1e-6):.0f} MB/s)\")\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        \"\"\"Legacy scalar stats (middle slice) + per-depth tissue profile.\"\"\"\n        self._ensure_open()\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n        else:\n            sample_coords = random.sample(coords, min(n_samples, len(coords)))\n            mid = len(self._slices) // 2\n            vals = []\n            for (y, x) in sample_coords:\n                block = self._to_uint8(self._slices[mid][y:y + patch_size, x:x + patch_size])\n                block = self._apply_contrast(block)\n                vals.append(block.astype(np.float32).ravel())\n            vals = np.concatenate(vals)\n            self.frag_mean = float(vals.mean())\n            self.frag_std = float(vals.std() + 1e-6)\n\n        if CFG.NORMALIZATION_MODE == \"per_fragment_per_slice\":\n            self._compute_depth_profile(tissue_mask)\n\n    def _compute_depth_profile(self, tissue_mask):\n        \"\"\"Per-depth mean/std over tissue pixels, from every `profile_row_stride`-th row\n        (reads ~1/stride of the data). Label-free.\"\"\"\n        stride = max(int(CFG.profile_row_stride), 1)\n        H, W = self._h, self._w\n        sel = tissue_mask[:H:stride, :W] > 0\n        n = len(self._slices)\n        means, stds = np.zeros(n), np.zeros(n)\n        for i, s in enumerate(self._slices):\n            block = self._to_uint8(s[:H:stride, :W])\n            block = self._apply_contrast(block)\n            vals = block[sel] if sel.shape == block.shape and sel.sum() > 1000 else block.ravel()\n            means[i] = float(vals.mean())\n            stds[i] = float(vals.std() + 1e-6)\n        self.slice_mean = means\n        self.slice_std = stds\n        self.norm_scale = float(np.sqrt(np.mean(stds ** 2)))\n        print(f\"    depth profile: mean range [{means.min():.1f}, {means.max():.1f}] \"\n              f\"(peak slice #{int(np.argmax(means))}, trough slice #{int(np.argmin(means))} of window) \"\n              f\"| scale={self.norm_scale:.2f}\")\n\n    def normalize(self, img_dhw):\n        \"\"\"img_dhw: float32 (D,H,W) in [0,1].\"\"\"\n        mode = CFG.NORMALIZATION_MODE\n        if mode == \"per_fragment_per_slice\" and self.slice_mean is not None:\n            m = (self.slice_mean / 255.0).astype(np.float32)[:, None, None]\n            if CFG.PER_SLICE_STD_SCALING:\n                s = (np.maximum(self.slice_std, 1.0) / 255.0).astype(np.float32)[:, None, None]\n            else:\n                s = np.float32(self.norm_scale / 255.0 + 1e-6)\n            return ((img_dhw - m) / s).astype(np.float32)\n        if mode == \"per_fragment_zscore\" and self.frag_mean is not None:\n            return ((img_dhw - self.frag_mean / 255.0) / (self.frag_std / 255.0 + 1e-6)).astype(np.float32)\n        return ((img_dhw - img_dhw.mean()) / (img_dhw.std() + 1e-6)).astype(np.float32)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\n# ============================================================\n# 6. AUGMENTATION\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Label-FREE distractor (never marked positive).\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)], axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        n_affected = max(1, int(d * random.uniform(0.3, 0.7)))\n        affected = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_physical_synthetic_ink(img_hwd, label, n_strokes=None, max_amplitude=CFG.physical_ink_max_amplitude):\n    \"\"\"Label-POSITIVE synthetic ink with a Gaussian depth profile\n    I(z)=A*exp(-(z-z0)^2/(2 sigma^2)); the label mask is updated.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    label_out = label.copy()\n    n_strokes = n_strokes or random.randint(1, 2)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)], axis=1).astype(np.int32)\n        thickness = random.randint(2, 5)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n\n        z0 = random.uniform(d * 0.3, d * 0.7)\n        sigma = random.uniform(1.5, 4.0)\n        amplitude = random.uniform(20, max_amplitude) * random.choice([-1, 1])\n        profile = amplitude * np.exp(-((np.arange(d) - z0) ** 2) / (2 * sigma ** 2))\n        for zi in range(d):\n            delta = profile[zi]\n            if abs(delta) < 1.0:\n                continue\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n        label_out[stroke_mask > 0] = 1\n    return out, label_out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist_norm = np.clip(np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2) / radius, 0, 1)\n    factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\ndef warp_depth(img_hwd, scale, shift):\n    \"\"\"Resample the channel (depth) axis: output channel j reads source coordinate\n    (j-c)*scale + c + shift (clamped), with linear interpolation between slices.\"\"\"\n    d = img_hwd.shape[2]\n    c = (d - 1) / 2.0\n    src = np.clip((np.arange(d) - c) * scale + c + shift, 0, d - 1)\n    lo = np.floor(src).astype(np.int64)\n    hi = np.minimum(lo + 1, d - 1)\n    frac = (src - lo).astype(np.float32)\n    out = img_hwd[:, :, lo].astype(np.float32) * (1.0 - frac) + img_hwd[:, :, hi].astype(np.float32) * frac\n    return np.clip(out + 0.5, 0, 255).astype(np.uint8)\n\n\ndef _safe_transform(builder_new, builder_old, name):\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable ({e1} / {e2}); skipping.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma, border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03, p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 7. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            y = max(0, min(y + random.randint(-self.jitter, self.jitter), H - size))\n            x = max(0, min(x + random.randint(-self.jitter, self.jitter), W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size]\n        img = np.transpose(patch, (1, 2, 0))          # HWD uint8\n\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref_hwd = np.transpose(random.choice(self.hist_match_pool), (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n\n        if self.train_mode and CFG.USE_PHYSICAL_SYNTHETIC_INK and random.random() < CFG.physical_ink_p:\n            img, label = inject_physical_synthetic_ink(img, label)\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.train_mode and CFG.USE_DEPTH_WARP_AUG and random.random() < CFG.depth_warp_p:\n            lo, hi = CFG.depth_warp_scale_range\n            scale = math.exp(random.uniform(math.log(lo), math.log(hi)))\n            shift = random.uniform(*CFG.depth_warp_shift_range)\n            img = warp_depth(img, scale, shift)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = np.ascontiguousarray(np.transpose(img.astype(np.float32) / 255.0, (2, 0, 1)))   # DHW\n        img = vol.normalize(img)\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ============================================================\n# 8. MODEL\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"Raw + first/second finite differences along the ordered depth axis, mixed by a\n    1x1 conv (V4's stem).\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = (in_depth if use_raw else 0) + ((in_depth - 1) if use_grad else 0) + \\\n                   (max(in_depth - 2, 0) if use_curv else 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: enable at least one of raw/grad/curv\")\n        self.mix = nn.Sequential(nn.Conv2d(total_in, out_channels, kernel_size=1),\n                                 nn.BatchNorm2d(out_channels), nn.GELU())\n        self.out_channels = out_channels\n\n    def forward(self, x):\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:] - x[:, :-1])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:] - 2 * x[:, 1:-1] + x[:, :-2])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass DepthAwareSegModel(nn.Module):\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):\n        return self.seg_model(self.depth_stem(x) if self.depth_stem is not None else x)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    return arch_cls(encoder_name=encoder_name, encoder_weights=CFG.encoder_weights, in_channels=in_channels,\n                     classes=1, decoder_channels=CFG.decoder_channels,\n                     decoder_attention_type=CFG.decoder_attention_type)\n\n\ndef build_model(encoder_name):\n    stem, seg_in = None, CFG.in_channels\n    if CFG.USE_DEPTH_FUSION_STEM:\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in = CFG.depth_stem_out_channels\n    try:\n        seg_model = _build_seg_backbone(encoder_name, seg_in)\n        print(f\"[backbone] using {encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {encoder_name} failed ({e}); falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in)\n    return DepthAwareSegModel(stem, seg_model)\n\n\n@torch.no_grad()\ndef run_architecture_report(model, tag, full=True):\n    print(f\"\\n{'='*70}\\nARCHITECTURE INSPECTION [{tag}]\\n{'='*70}\")\n    print(f\"Parameters: {sum(p.numel() for p in model.parameters())/1e6:.2f}M\")\n    for name, m in model.named_modules():\n        if isinstance(m, (nn.Conv1d, nn.Conv2d, nn.Conv3d)) and (m.in_channels <= 0 or m.out_channels <= 0):\n            raise RuntimeError(f\"Zero-channel layer found: {name}\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n    if full:\n        model.eval()\n        s = min(CFG.patch_size, 256)\n        out = model(torch.zeros(1, CFG.in_channels, s, s, device=CFG.device))\n        bad = torch.isnan(out).any().item() or torch.isinf(out).any().item()\n        print(f\"Dry-run {s}x{s}: {tuple(out.shape)} | NaN/Inf: {'FAIL' if bad else 'PASS'}\")\n        if bad:\n            raise RuntimeError(\"Model produced NaN/Inf on a dry run.\")\n        model.train()\n        if torch.cuda.is_available():\n            print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")\n    print(\"=\" * 70)\n\n\n# ============================================================\n# 9. LOSSES\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (probs.sum(dim=1) + t.sum(dim=1) + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    return torch.pow(1.0 - (tp + eps) / (tp + alpha * fp + beta * fn + eps), gamma).mean()\n\n\ndef soft_erode(I):\n    p1 = -F.max_pool2d(-I, (3, 1), (1, 1), (1, 0))\n    p2 = -F.max_pool2d(-I, (1, 3), (1, 1), (0, 1))\n    return torch.min(p1, p2)\n\n\ndef soft_dilate(I):\n    return F.max_pool2d(I, (3, 3), (1, 1), (1, 1))\n\n\ndef soft_open(I):\n    return soft_dilate(soft_erode(I))\n\n\ndef soft_skeletonize(I, iters=CFG.cldice_iters):\n    I1 = soft_open(I)\n    skel = F.relu(I - I1)\n    for _ in range(iters):\n        I = soft_erode(I)\n        I1 = soft_open(I)\n        delta = F.relu(I - I1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef soft_cldice(pred_probs, target, iters=CFG.cldice_iters, smooth=1.0):\n    pred_sk = soft_skeletonize(pred_probs, iters)\n    targ_sk = soft_skeletonize(target, iters)\n    tprec = (torch.sum(pred_sk * target) + smooth) / (torch.sum(pred_sk) + smooth)\n    tsens = (torch.sum(targ_sk * pred_probs) + smooth) / (torch.sum(targ_sk) + smooth)\n    return 1.0 - 2.0 * (tprec * tsens) / (tprec + tsens + 1e-8)\n\n\nclass V5ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        topo = soft_cldice(torch.sigmoid(logits), targets)\n        return (CFG.bce_weight * bce + CFG.dice_weight * dice +\n                CFG.focal_tversky_weight * tv + CFG.topology_weight * topo)\n\n\n# ============================================================\n# 10. METRICS (fixed-threshold + threshold-free histogram based)\n# ============================================================\n\ndef calculate_metrics(tp, fp, fn):\n    eps = 1e-6\n    dice = (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return calculate_metrics(self.tp, self.fp, self.fn)\n\n\ndef summarize_hist(pos, neg, beta2=0.25, smooth_bins=None):\n    \"\"\"pos/neg: 256-bin histograms of quantized probabilities (q=round(p*255)) for\n    positive / negative pixels. Rule: predict positive iff q >= k, k=1..255.\n    Returns the best-F0.5 operating point (argmax on a lightly smoothed curve, to avoid\n    razor-thin optima) plus average precision.\"\"\"\n    smooth_bins = CFG.threshold_smooth_bins if smooth_bins is None else smooth_bins\n    pos = np.asarray(pos, dtype=np.float64)\n    neg = np.asarray(neg, dtype=np.float64)\n    tp = np.cumsum(pos[::-1])[::-1][1:]\n    fp = np.cumsum(neg[::-1])[::-1][1:]\n    p_total = max(pos.sum(), 1.0)\n    prec = tp / np.maximum(tp + fp, 1.0)\n    rec = tp / p_total\n    f = (1.0 + beta2) * prec * rec / np.maximum(beta2 * prec + rec, 1e-12)\n    if smooth_bins > 0:\n        w = 2 * smooth_bins + 1\n        fs = np.convolve(np.pad(f, smooth_bins, mode=\"edge\"), np.ones(w) / w, mode=\"valid\")\n    else:\n        fs = f\n    k = int(np.argmax(fs))\n    rec_next = np.append(rec[1:], 0.0)\n    ap = float(np.sum((rec - rec_next) * prec))\n    prev = p_total / max(pos.sum() + neg.sum(), 1.0)\n    ap += float((1.0 - rec[0]) * prev)          # segment below the lowest threshold\n    return {\"f05\": float(f[k]), \"thr\": float((k + 0.5) / 255.0),\n            \"precision\": float(prec[k]), \"recall\": float(rec[k]), \"ap\": ap}\n\n\nclass ScoreHistogram:\n    def __init__(self):\n        self.pos = np.zeros(256, dtype=np.int64)\n        self.neg = np.zeros(256, dtype=np.int64)\n\n    @torch.no_grad()\n    def update(self, probs, targets):\n        q = (probs.float().clamp(0, 1) * 255.0 + 0.5).long().flatten()\n        t = targets.flatten() > 0.5\n        self.pos += torch.bincount(q[t], minlength=256).cpu().numpy()\n        self.neg += torch.bincount(q[~t], minlength=256).cpu().numpy()\n\n    def summary(self):\n        return summarize_hist(self.pos, self.neg)\n\n\ndef hist_from_stack(prob_u8_stack, labels, masks, sigma=0.0):\n    \"\"\"Histograms over TISSUE pixels of a stack of uint8 probability patches,\n    optionally Gaussian-smoothed first.\"\"\"\n    pos = np.zeros(256, dtype=np.int64)\n    neg = np.zeros(256, dtype=np.int64)\n    for j in range(len(prob_u8_stack)):\n        q = prob_u8_stack[j]\n        if sigma > 0:\n            pf = cv2.GaussianBlur(q.astype(np.float32) / 255.0, (0, 0), sigmaX=float(sigma))\n            q = np.clip(pf * 255.0 + 0.5, 0, 255).astype(np.uint8)\n        m = masks[j] > 0\n        lab = labels[j] > 0\n        pos += np.bincount(q[m & lab], minlength=256)\n        neg += np.bincount(q[m & ~lab], minlength=256)\n    return pos, neg\n\n\ndef evaluate_prob_map_hist(prob_map, gt, mask):\n    q = np.clip(prob_map * 255.0 + 0.5, 0, 255).astype(np.uint8)\n    m = mask > 0\n    lab = gt > 0\n    return summarize_hist(np.bincount(q[m & lab], minlength=256), np.bincount(q[m & ~lab], minlength=256))\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    return calculate_metrics((preds * gt).sum(), (preds * (1.0 - gt)).sum(), ((1.0 - preds) * gt).sum())\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 11. SHARED DATA (built once, reused by every member)\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING SHARED TRAIN/VAL PATCH GRID\")\nprint(\"=\" * 70)\n\n_shared_masks, _shared_labels = {}, {}\nall_samples = []\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {mask.shape} {len(coords)} candidate patches\")\n    _shared_masks[fid] = mask\n    _shared_labels[fid] = labels\n    all_samples.extend([(fid, y, x) for y, x in coords])\n\n_groups = {}\nfor fid, y, x in all_samples:\n    _groups.setdefault((fid, y // 768, x // 768), []).append((fid, y, x))\n_keys = list(_groups.keys())\nrandom.Random(CFG.seed).shuffle(_keys)\n_val_keys = set(_keys[:max(1, int(len(_keys) * CFG.val_fraction))])\ntrain_grid, val_samples = [], []\nfor key, items in _groups.items():\n    (val_samples if key in _val_keys else train_grid).extend(items)\nprint(f\"Spatial train: {len(train_grid)} | Spatial val: {len(val_samples)}\")\n\n\ndef build_train_samples(seed):\n    r = random.Random(seed)\n    s = balance_positive_patches(train_grid, _shared_labels, CFG.patch_size,\n                                  positive_threshold=CFG.positive_patch_fraction,\n                                  target_positive_ratio=CFG.target_positive_patch_ratio,\n                                  max_positive_repeat=CFG.max_positive_repeat, rng=r)\n    return domain_balance(s, r)\n\n\n# --- validation-calibration subset (evenly spaced), with tissue-masked labels ---------\nVAL_CAL_IDXS = np.unique(np.linspace(0, len(val_samples) - 1,\n                                     min(CFG.val_calibration_max_patches, len(val_samples))).astype(int))\n_S = CFG.patch_size\nVAL_LABELS = np.stack([_shared_labels[val_samples[i][0]][val_samples[i][1]:val_samples[i][1] + _S,\n                                                          val_samples[i][2]:val_samples[i][2] + _S]\n                       for i in VAL_CAL_IDXS]).astype(np.uint8)\nVAL_MASKS = np.stack([_shared_masks[val_samples[i][0]][val_samples[i][1]:val_samples[i][1] + _S,\n                                                        val_samples[i][2]:val_samples[i][2] + _S]\n                      for i in VAL_CAL_IDXS]).astype(np.uint8)\nprint(f\"Validation calibration subset: {len(VAL_CAL_IDXS)} patches\")\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\nprint(f\"Test fragment {CFG.test_frag}: {test_mask.shape}\")\n\npool_source = []\nif CFG.USE_HIST_MATCH_AUG:\n    for fid in CFG.train_frags:\n        c = generate_grid_coords(_shared_masks[fid], CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n        pool_source.extend([(fid, y, x) for y, x in c])\n    print(f\"[strict protocol] histogram-match pool from TRAIN fragments only ({len(pool_source)} candidates).\")\n\n\n# ============================================================\n# 12. INFERENCE HELPERS\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    g1d = np.exp(-(ax ** 2) / (2.0 * (size * sigma_frac) ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\": return x\n    if mode == \"hflip\": return torch.flip(x, dims=[3])\n    if mode == \"vflip\": return torch.flip(x, dims=[2])\n    if mode == \"hvflip\": return torch.flip(x, dims=[2, 3])\n    if mode == \"rot90\": return torch.rot90(x, k=1, dims=[2, 3])\n    if mode == \"rot180\": return torch.rot90(x, k=2, dims=[2, 3])\n    if mode == \"rot270\": return torch.rot90(x, k=3, dims=[2, 3])\n    if mode == \"transpose\": return x.transpose(2, 3)\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\": return pred\n    if mode == \"hflip\": return torch.flip(pred, dims=[2])\n    if mode == \"vflip\": return torch.flip(pred, dims=[1])\n    if mode == \"hvflip\": return torch.flip(pred, dims=[1, 2])\n    if mode == \"rot90\": return torch.rot90(pred, k=-1, dims=[1, 2])\n    if mode == \"rot180\": return torch.rot90(pred, k=-2, dims=[1, 2])\n    if mode == \"rot270\": return torch.rot90(pred, k=-3, dims=[1, 2])\n    if mode == \"transpose\": return pred.transpose(1, 2)\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    acc = None\n    for mode in modes:\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            probs = torch.sigmoid(model(apply_tta_tensor(x, mode)))\n        probs = invert_tta_prediction(probs[:, 0], mode).float() / len(modes)\n        acc = probs if acc is None else acc + probs\n    return acc.cpu().numpy().astype(np.float32)\n\n\n@torch.no_grad()\ndef predict_val_stack(model, val_ds, idxs):\n    \"\"\"TTA'd probabilities (uint8-quantized) for the calibration subset of validation.\"\"\"\n    loader = DataLoader(Subset(val_ds, [int(i) for i in idxs]), batch_size=CFG.infer_batch, shuffle=False,\n                        num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"))\n    model.eval()\n    out = []\n    for imgs, _ in loader:\n        p = predict_batch_tta(model, imgs.numpy())\n        out.append(np.clip(p * 255.0 + 0.5, 0, 255).astype(np.uint8))\n    return np.concatenate(out, axis=0)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"  Inference patches: {len(coords)} (TTA views: {len(CFG.tta_modes) if CFG.use_tta else 1})\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        probs = predict_batch_tta(model, np.stack(batch_imgs).astype(np.float32))\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear()\n        batch_coords.clear()\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        batch_imgs.append(vol.normalize(raw))\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm(model, vol, mask, patch_size, stride, max_patches, batch_size):\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        imgs = [vol.normalize(vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0)\n                for y, x in coords[start:start + batch_size]]\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            model(torch.from_numpy(np.stack(imgs)).to(CFG.device))\n    model.eval()\n    cleanup_memory()\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    n, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if n <= 1:\n        return binary\n    out = np.zeros_like(binary, dtype=np.uint8)\n    for i in range(1, n):\n        if stats[i, cv2.CC_STAT_AREA] >= min_size:\n            out[labels == i] = 1\n    return out\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    return remove_small_components(binary, CFG.min_component_size).astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\n# ============================================================\n# 13. TRAIN + INFER ONE ENSEMBLE MEMBER\n# ============================================================\n\nclass EMAModel:\n    def __init__(self, model, decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n\ndef train_and_infer_one_member(member_idx, seed, encoder_name, window):\n    start, stop, step = window\n    depth_indices = list(range(start, stop, step))\n    if len(depth_indices) != CFG.n_slices or max(depth_indices) > CFG.max_slice_index:\n        raise ValueError(f\"window {window} must give {CFG.n_slices} slices <= {CFG.max_slice_index}\")\n    print(f\"\\n{'#'*70}\\n# MEMBER {member_idx+1}/{CFG.N_ENSEMBLE_MODELS} seed={seed} encoder={encoder_name} \"\n          f\"window={window} (slices {depth_indices[0]}..{depth_indices[-1]}, step {step})\\n{'#'*70}\")\n    set_seed(seed)\n\n    train_volumes = {}\n    for fid in CFG.train_frags:\n        v = FragmentVolume(os.path.join(CFG.base_dir, fid), depth_indices)\n        if CFG.WARM_UP_VOLUME_CACHE:\n            v.warm_up_cache(f\"fragment {fid}\")\n        v.compute_fragment_stats(_shared_masks[fid], CFG.patch_size)\n        print(f\"  fragment {fid}: legacy mean={v.frag_mean:.2f} std={v.frag_std:.2f}\")\n        train_volumes[fid] = v\n\n    train_samples = build_train_samples(seed)\n\n    hist_pool = None\n    if CFG.USE_HIST_MATCH_AUG and pool_source:\n        picked = random.sample(pool_source, min(CFG.hist_match_pool_size, len(pool_source)))\n        hist_pool = [train_volumes[fid].read_patch(y, x, CFG.patch_size) for fid, y, x in picked]\n\n    train_ds = InkPatchDataset(train_volumes, _shared_labels, train_samples, CFG.patch_size,\n                                transform=build_train_transform(), jitter=CFG.train_jitter,\n                                hist_match_pool=hist_pool, train_mode=True)\n    val_ds = InkPatchDataset(train_volumes, _shared_labels, val_samples, CFG.patch_size,\n                              transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n    train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True, num_workers=CFG.num_workers,\n                               pin_memory=(CFG.device == \"cuda\"), drop_last=CFG.drop_last,\n                               persistent_workers=CFG.num_workers > 0)\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers,\n                             pin_memory=(CFG.device == \"cuda\"), persistent_workers=CFG.num_workers > 0)\n\n    pos_frac = estimate_positive_fraction(_shared_labels, train_samples, CFG.patch_size)\n    initial_bias = math.log(pos_frac / max(1 - pos_frac, 1e-6))\n    pos_weight = float(np.clip(np.sqrt((1 - pos_frac) / pos_frac), 1.0, 8.0))\n    print(f\"  positive fraction={pos_frac:.5f} bias={initial_bias:.3f} pos_weight={pos_weight:.2f}\")\n\n    model = build_model(encoder_name).to(CFG.device)\n    run_architecture_report(model, f\"member{member_idx}\", full=(member_idx == 0))\n    with torch.no_grad():\n        try:\n            model.seg_model.segmentation_head[0].bias.fill_(initial_bias)\n        except Exception:\n            pass\n\n    criterion = V5ComboLoss(pos_weight=torch.tensor([pos_weight], dtype=torch.float32, device=CFG.device))\n\n    enc_p, dec_p, stem_p = [], [], []\n    for name, p in model.named_parameters():\n        if p.requires_grad:\n            (stem_p if name.startswith(\"depth_stem.\") else\n             enc_p if name.startswith(\"seg_model.encoder.\") else dec_p).append(p)\n    groups = [{\"params\": enc_p, \"lr\": CFG.encoder_lr}, {\"params\": dec_p, \"lr\": CFG.decoder_lr}]\n    if stem_p:\n        groups.append({\"params\": stem_p, \"lr\": CFG.depth_stem_lr})\n    optimizer = torch.optim.AdamW(groups, weight_decay=CFG.weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n    ema = EMAModel(model, CFG.ema_decay) if CFG.USE_EMA else None\n\n    def run_epoch(loader, train_mode, want_hist=False):\n        model.train(train_mode)\n        total_loss = 0.0\n        acc = GlobalConfusionAccumulator()\n        hist = ScoreHistogram() if want_hist else None\n        for imgs, masks in loader:\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            masks = masks.to(CFG.device, non_blocking=True)\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n            with torch.set_grad_enabled(train_mode):\n                with autocast(enabled=(CFG.device == \"cuda\")):\n                    logits = model(imgs)\n                    loss = criterion(logits, masks)\n                if train_mode:\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    if ema is not None:\n                        ema.update(model)\n            probs = torch.sigmoid(logits.detach())\n            acc.update(probs, masks, 0.5)\n            if hist is not None:\n                hist.update(probs, masks)\n            total_loss += float(loss.item())\n            del imgs, masks, logits, probs\n        return total_loss / max(len(loader), 1), acc.compute(), hist\n\n    print(f\"\\n  Training up to {CFG.epochs} epochs (patch={CFG.patch_size}, batch={CFG.batch_size}); \"\n          f\"checkpoint metric = val best-threshold F0.5 ...\")\n    best_score, no_improve = -1.0, 0\n    best_state, best_used = None, None\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n        train_loss, train_m, _ = run_epoch(train_loader, True)\n\n        # Evaluate raw AND EMA weights; checkpoint the better one; ALWAYS continue from raw.\n        raw_state = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        _, _, h_raw = run_epoch(val_loader, False, want_hist=True)\n        s_raw = h_raw.summary()\n        sel, cand, used = s_raw, raw_state, \"raw\"\n        if ema is not None:\n            model.load_state_dict(ema.shadow)\n            _, _, h_ema = run_epoch(val_loader, False, want_hist=True)\n            model.load_state_dict(raw_state)\n            s_ema = h_ema.summary()\n            print(f\"    [checkpoint choice] EMA F0.5*={s_ema['f05']:.4f} raw F0.5*={s_raw['f05']:.4f}\")\n            if s_ema[\"f05\"] >= s_raw[\"f05\"]:\n                sel, cand, used = s_ema, {k: v.clone() for k, v in ema.shadow.items()}, \"ema\"\n        scheduler.step()\n\n        print(f\"  [member {member_idx+1}][{epoch:02d}/{CFG.epochs}] {time.time()-t0:.0f}s \"\n              f\"train_loss={train_loss:.4f} train_dice={train_m['dice']:.4f} | \"\n              f\"val F0.5*={sel['f05']:.4f} (thr {sel['thr']:.2f}, P={sel['precision']:.3f}, \"\n              f\"R={sel['recall']:.3f}) AP={sel['ap']:.4f} [{used}]\")\n\n        if sel[\"f05\"] > best_score:\n            best_score, no_improve, best_state, best_used = sel[\"f05\"], 0, cand, used\n            print(f\"    *** new best val F0.5*={best_score:.4f} ({used}) ***\")\n        else:\n            no_improve += 1\n            if no_improve >= CFG.early_stop_patience:\n                print(\"    Early stopping.\")\n                break\n        cleanup_memory()\n\n    model.load_state_dict(best_state)\n    model.eval()\n    print(f\"  Member {member_idx+1} best val F0.5*: {best_score:.4f} ({best_used})\")\n\n    # --- TTA'd validation predictions for ensemble-level calibration ---\n    val_stack = predict_val_stack(model, val_ds, VAL_CAL_IDXS)\n    solo = summarize_hist(*hist_from_stack(val_stack, VAL_LABELS, VAL_MASKS, 0.0))\n    print(f\"  Member {member_idx+1} TTA'd val (tissue only): F0.5*={solo['f05']:.4f} thr={solo['thr']:.3f} AP={solo['ap']:.4f}\")\n\n    for v in train_volumes.values():\n        v.close()\n    del train_loader, val_loader, train_ds, val_ds\n    cleanup_memory()\n\n    # --- fragment-1 inference with this member's own depth window ---\n    test_vol = FragmentVolume(test_dir, depth_indices)\n    if CFG.WARM_UP_VOLUME_CACHE:\n        test_vol.warm_up_cache(f\"fragment {CFG.test_frag} (held-out)\")\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    print(f\"  [member {member_idx+1}] {len(CFG.tta_modes) if CFG.use_tta else 1}-view TTA inference ...\")\n    prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.infer_batch)\n    if CFG.use_adabn:\n        model = recalibrate_batchnorm(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                       CFG.adabn_max_patches, CFG.infer_batch)\n        prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.infer_batch)\n    test_vol.close()\n    cleanup_memory()\n\n    return {\"member_idx\": member_idx, \"seed\": seed, \"encoder\": encoder_name, \"window\": list(window),\n            \"best_val_f05\": float(best_score), \"solo_val\": solo, \"val_stack\": val_stack, \"probability\": prob}\n\n\n# ============================================================\n# 14. RUN THE ENSEMBLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(f\"STARTING V5.3 ({'single model' if CFG.N_ENSEMBLE_MODELS == 1 else f'{CFG.N_ENSEMBLE_MODELS}-model ensemble'}, single split)\")\nprint(\"=\" * 70)\n\nmember_results = []\nfor i in range(CFG.N_ENSEMBLE_MODELS):\n    member_results.append(train_and_infer_one_member(\n        i, CFG.ensemble_seeds[i], CFG.ensemble_encoders[i], CFG.ensemble_windows[i]))\n    cleanup_memory()\n\n\n# ============================================================\n# 15. ENSEMBLE CALIBRATION ON VALIDATION (labels of fragment 1 not used here)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"ENSEMBLE CALIBRATION (validation only): smoothing sigma + F0.5 threshold\")\nprint(\"=\" * 70)\n\nens_val = np.mean(np.stack([r[\"val_stack\"] for r in member_results]).astype(np.float32), axis=0)\nens_val = np.clip(ens_val + 0.5, 0, 255).astype(np.uint8)\n\nbest_sigma, best_sum = 0.0, None\nfor sigma in CFG.smooth_sigma_candidates:\n    s = summarize_hist(*hist_from_stack(ens_val, VAL_LABELS, VAL_MASKS, sigma))\n    print(f\"  sigma={sigma:>4}: val F0.5*={s['f05']:.4f} thr={s['thr']:.3f} \"\n          f\"P={s['precision']:.3f} R={s['recall']:.3f} AP={s['ap']:.4f}\")\n    if best_sum is None or s[\"f05\"] > best_sum[\"f05\"] + 1e-4:\n        best_sigma, best_sum = sigma, s\nfinal_threshold = best_sum[\"thr\"]\nprint(f\"\\n  -> chosen sigma={best_sigma}, threshold={final_threshold:.3f} \"\n      f\"(val F0.5*={best_sum['f05']:.4f}); frozen before touching fragment-1 labels.\")\n\n\n# ============================================================\n# 16. FINAL TEST MAP + DIAGNOSTIC REPORT (labels loaded only now)\n# ============================================================\n\ndef finalize_map(prob, sigma):\n    p = cv2.GaussianBlur(prob, (0, 0), sigmaX=float(sigma)) if sigma > 0 else prob\n    return (p * test_mask).astype(np.float32)        # labels are tissue-masked, so predictions are too\n\n\nensembled_probability = finalize_map(np.mean([r[\"probability\"] for r in member_results], axis=0).astype(np.float32),\n                                     best_sigma)\nfinal_prediction = postprocess(ensembled_probability, final_threshold)\notsu_thr = compute_otsu_threshold(ensembled_probability, test_mask, fallback=final_threshold)\n\ntest_labels = load_ink_labels(test_dir)\ngt_test = (test_labels * test_mask).astype(np.float32) if test_labels is not None else None\n\nraw_metrics = post_metrics = otsu_metrics = oracle = None\nmember_solo = []\nif gt_test is not None:\n    print(\"\\n\" + \"=\" * 70)\n    print(\"FRAGMENT 1 RESULTS (labels used for REPORTING only)\")\n    print(\"=\" * 70)\n    for r in member_results:\n        thr = r[\"solo_val\"][\"thr\"]\n        m = evaluate_probability_map(r[\"probability\"] * test_mask, gt_test, thr)\n        member_solo.append(m)\n        print(f\"  member {r['member_idx']+1} ({r['encoder']}, window {r['window']}) at its own val thr={thr:.3f}: \"\n              f\"F0.5={m['fbeta0.5']:.4f} P={m['precision']:.3f} R={m['recall']:.3f} Dice={m['dice']:.4f}\")\n\n    raw_metrics = evaluate_probability_map(ensembled_probability, gt_test, final_threshold)\n    pf = final_prediction.astype(np.float32)\n    post_metrics = calculate_metrics((pf * gt_test).sum(), (pf * (1 - gt_test)).sum(), ((1 - pf) * gt_test).sum())\n    otsu_metrics = evaluate_probability_map(ensembled_probability, gt_test, otsu_thr)\n    oracle = evaluate_prob_map_hist(ensembled_probability, gt_test, test_mask)\n\n    print(f\"\\n  ENSEMBLE @ val-chosen thr={final_threshold:.3f}, sigma={best_sigma}\")\n    print(f\"    raw:           {raw_metrics}\")\n    print(f\"    postprocessed: {post_metrics}\")\n    print(f\"  [diagnostic] Otsu thr={otsu_thr:.3f} would give F0.5={otsu_metrics['fbeta0.5']:.4f} (NOT used)\")\n    print(f\"  [analysis only] ranking quality: AP={oracle['ap']:.4f} | oracle F0.5={oracle['f05']:.4f} \"\n          f\"at thr={oracle['thr']:.3f} (upper bound, NOT a reportable result)\")\n\n\n# ============================================================\n# 17. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_v5_3.npy\")\nnp.save(prob_path, ensembled_probability)\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_v5_3.png\")\ncv2.imwrite(pred_path, (final_prediction * 255).astype(np.uint8))\nfor r in member_results:\n    np.save(os.path.join(CFG.out_dir, f\"fragment1_probability_member{r['member_idx']}_v5_3.npy\"), r[\"probability\"])\n\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump({\n        \"chosen_sigma\": best_sigma, \"chosen_threshold\": final_threshold, \"val_calibration\": best_sum,\n        \"ensembled_raw\": raw_metrics, \"ensembled_postprocessed\": post_metrics,\n        \"otsu_diagnostic\": otsu_metrics, \"ranking_analysis_only\": oracle,\n        \"members\": [{\"idx\": r[\"member_idx\"], \"seed\": r[\"seed\"], \"encoder\": r[\"encoder\"], \"window\": r[\"window\"],\n                     \"best_val_f05\": r[\"best_val_f05\"], \"solo_val\": r[\"solo_val\"]} for r in member_results],\n        \"config\": cfg_to_dict(CFG),\n    }, f, indent=2, default=float)\nprint(f\"\\nSaved metrics: {CFG.metrics_path}\")\n\n\n# ============================================================\n# 18. VISUALIZATION\n# ============================================================\n\ndef make_confusion_overlay(gray_u8, gt_bin, pred_bin, alpha=0.55):\n    base = np.stack([gray_u8] * 3, axis=-1).astype(np.float32)\n    ov = base.copy()\n    g, p = gt_bin > 0, pred_bin > 0\n    for m, col in ((p & g, [0, 255, 0]), (p & ~g, [255, 0, 0]), (~p & g, [255, 255, 0])):\n        ov[m] = (1 - alpha) * base[m] + alpha * np.array(col, dtype=np.float32)\n    return np.clip(ov, 0, 255).astype(np.uint8)\n\n\ndef save_overview():\n    \"\"\"Saves ONE figure with, left to right: input slice, ground truth,\n    probability map, thresholded prediction, and a TP/FP/FN overlay with an\n    explicit color-swatch legend (not just a title) -- green=TP, red=FP,\n    yellow=FN, unlabeled tissue stays grayscale.\"\"\"\n    ref = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{CFG.mask_reference_slice:02d}.tif\"))\n    sc = 2000 / max(ref.shape)\n    small = cv2.resize(ref, None, fx=sc, fy=sc, interpolation=cv2.INTER_AREA)\n    small_u8 = cv2.normalize(small, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n    pred_s = cv2.resize((final_prediction * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    prob_s = cv2.resize((ensembled_probability * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_s = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n        ov = make_confusion_overlay(small_u8, (gt_s > 127).astype(np.uint8), (pred_s > 127).astype(np.uint8))\n        fig, ax = plt.subplots(1, 5, figsize=(28, 6))\n        ax[0].imshow(small, cmap=\"gray\"); ax[0].set_title(f\"Input (slice {CFG.mask_reference_slice})\", fontsize=12)\n        ax[1].imshow(gt_s, cmap=\"gray\"); ax[1].set_title(\"Ground truth\", fontsize=12)\n        ax[2].imshow(prob_s, cmap=\"gray\"); ax[2].set_title(f\"Probability (sigma={best_sigma})\", fontsize=12)\n        ax[3].imshow(pred_s, cmap=\"gray\"); ax[3].set_title(f\"Prediction (thr={final_threshold:.2f})\", fontsize=12)\n        ax[4].imshow(ov); ax[4].set_title(\"Prediction vs. ground truth\", fontsize=12)\n        from matplotlib.patches import Patch\n        legend_handles = [\n            Patch(facecolor=np.array([0, 255, 0]) / 255, label=\"True Positive (correct ink)\"),\n            Patch(facecolor=np.array([255, 0, 0]) / 255, label=\"False Positive (over-predicted)\"),\n            Patch(facecolor=np.array([255, 255, 0]) / 255, label=\"False Negative (missed ink)\"),\n        ]\n        ax[4].legend(handles=legend_handles, loc=\"upper center\", bbox_to_anchor=(0.5, -0.05),\n                     ncol=1, fontsize=10, frameon=True, framealpha=0.9)\n    else:\n        fig, ax = plt.subplots(1, 3, figsize=(18, 6))\n        ax[0].imshow(small, cmap=\"gray\"); ax[0].set_title(f\"Input (slice {CFG.mask_reference_slice})\", fontsize=12)\n        ax[1].imshow(prob_s, cmap=\"gray\"); ax[1].set_title(f\"Probability (sigma={best_sigma})\", fontsize=12)\n        ax[2].imshow(pred_s, cmap=\"gray\"); ax[2].set_title(f\"Prediction (thr={final_threshold:.2f})\", fontsize=12)\n    for a in ax:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_v5_3_overview.png\")\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(\"Saved comparison figure (input | ground truth | probability | prediction | TP/FP/FN overlay):\", path)\n\n\nsave_overview()\ncleanup_memory()\n\n\n# ============================================================\n# 19. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V5.3 COMPLETE\")\nprint(\"=\" * 70)\nfor r in member_results:\n    print(f\"Member {r['member_idx']+1}: {r['encoder']} window={r['window']} \"\n          f\"best val F0.5*={r['best_val_f05']:.4f} | TTA'd val F0.5*={r['solo_val']['f05']:.4f}\")\nprint(f\"\\nEnsemble: sigma={best_sigma} threshold={final_threshold:.3f} (chosen on validation by F0.5)\")\nif post_metrics is not None:\n    print(f\"Fragment 1 F0.5 (raw): {raw_metrics['fbeta0.5']:.4f} | (postprocessed): {post_metrics['fbeta0.5']:.4f} \"\n          f\"| Dice: {post_metrics['dice']:.4f} | AP: {oracle['ap']:.4f}\")\nprint(f\"\\nProbability map: {prob_path}\\nPrediction: {pred_path}\\nMetrics: {CFG.metrics_path}\")\nprint(\"\\n=== V5.3 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T07:25:33.179726Z","iopub.execute_input":"2026-09-29T07:25:33.180533Z","iopub.status.idle":"2026-09-29T08:48:36.785882Z","shell.execute_reply.started":"2026-09-29T07:25:33.180502Z","shell.execute_reply":"2026-09-29T08:48:36.785092Z"}},"outputs":[{"name":"stdout","text":"   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 5.0 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.7/2.7 MB 45.6 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.30\n======================================================================\nBUILDING SHARED TRAIN/VAL PATCH GRID\n======================================================================\nFragment 2: (14830, 9506) 6146 candidate patches\nFragment 3: (7606, 5249) 1607 candidate patches\nSpatial train: 6151 | Spatial val: 1602\nValidation calibration subset: 600 patches\nTest fragment 1: (8181, 6330)\n[strict protocol] histogram-match pool from TRAIN fragments only (1247 candidates).\n\n======================================================================\nSTARTING V5.3 (single model, single split)\n======================================================================\n\n######################################################################\n# MEMBER 1/1 seed=42 encoder=tu-convnext_tiny window=(12, 38, 1) (slices 12..37, step 1)\n######################################################################\n  [cache warm-up fragment 2] 26 slices, 7.33 GB in 247.4s (30 MB/s)\n    depth profile: mean range [94.4, 121.0] (peak slice #14, trough slice #25 of window) | scale=54.39\n  fragment 2: legacy mean=118.25 std=59.13\n  [cache warm-up fragment 3] 26 slices, 2.08 GB in 69.1s (30 MB/s)\n    depth profile: mean range [57.9, 135.1] (peak slice #15, trough slice #25 of window) | scale=52.32\n  fragment 3: legacy mean=119.91 std=61.59\n  Positive patches: 4015 | Negative patches: 2136 | Balanced: 6151 (positive ratio=0.653)\n  Domain balance: frag 2: 4919 -> 3997 | frag 3: 1232 -> 2154\n  positive fraction=0.12190 bias=-1.975 pos_weight=2.68\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"23107e05a69543d9816f12c6f4a9c8bf"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nARCHITECTURE INSPECTION [member0]\n======================================================================\nParameters: 32.18M\nZero-channel layers: 0 (PASS)\nDry-run 256x256: (1, 1, 256, 256) | NaN/Inf: PASS\nCUDA device: Tesla T4\n======================================================================\n\n  Training up to 10 epochs (patch=320, batch=8); checkpoint metric = val best-threshold F0.5 ...\n    [checkpoint choice] EMA F0.5*=0.2964 raw F0.5*=0.3478\n  [member 1][01/10] 420s train_loss=0.7950 train_dice=0.3015 | val F0.5*=0.3478 (thr 0.31, P=0.454, R=0.180) AP=0.3368 [raw]\n    *** new best val F0.5*=0.3478 (raw) ***\n    [checkpoint choice] EMA F0.5*=0.5231 raw F0.5*=0.5291\n  [member 1][02/10] 365s train_loss=0.7358 train_dice=0.4209 | val F0.5*=0.5291 (thr 0.58, P=0.607, R=0.349) AP=0.5111 [raw]\n    *** new best val F0.5*=0.5291 (raw) ***\n    [checkpoint choice] EMA F0.5*=0.5813 raw F0.5*=0.5929\n  [member 1][03/10] 364s train_loss=0.7157 train_dice=0.4568 | val F0.5*=0.5929 (thr 0.62, P=0.692, R=0.377) AP=0.5746 [raw]\n    *** new best val F0.5*=0.5929 (raw) ***\n    [checkpoint choice] EMA F0.5*=0.6103 raw F0.5*=0.6122\n  [member 1][04/10] 374s train_loss=0.6946 train_dice=0.4911 | val F0.5*=0.6122 (thr 0.82, P=0.701, R=0.406) AP=0.5893 [raw]\n    *** new best val F0.5*=0.6122 (raw) ***\n    [checkpoint choice] EMA F0.5*=0.6283 raw F0.5*=0.6075\n  [member 1][05/10] 366s train_loss=0.6782 train_dice=0.5149 | val F0.5*=0.6283 (thr 0.77, P=0.729, R=0.404) AP=0.6273 [ema]\n    *** new best val F0.5*=0.6283 (ema) ***\n    [checkpoint choice] EMA F0.5*=0.6405 raw F0.5*=0.6347\n  [member 1][06/10] 365s train_loss=0.6611 train_dice=0.5389 | val F0.5*=0.6405 (thr 0.81, P=0.734, R=0.424) AP=0.6435 [ema]\n    *** new best val F0.5*=0.6405 (ema) ***\n    [checkpoint choice] EMA F0.5*=0.6482 raw F0.5*=0.6593\n  [member 1][07/10] 364s train_loss=0.6426 train_dice=0.5703 | val F0.5*=0.6593 (thr 0.75, P=0.767, R=0.422) AP=0.6604 [raw]\n    *** new best val F0.5*=0.6593 (raw) ***\n    [checkpoint choice] EMA F0.5*=0.6559 raw F0.5*=0.6665\n  [member 1][08/10] 368s train_loss=0.6269 train_dice=0.5957 | val F0.5*=0.6665 (thr 0.88, P=0.763, R=0.443) AP=0.6640 [raw]\n    *** new best val F0.5*=0.6665 (raw) ***\n    [checkpoint choice] EMA F0.5*=0.6630 raw F0.5*=0.6665\n  [member 1][09/10] 367s train_loss=0.6198 train_dice=0.6182 | val F0.5*=0.6665 (thr 0.84, P=0.766, R=0.438) AP=0.6632 [raw]\n    [checkpoint choice] EMA F0.5*=0.6680 raw F0.5*=0.6711\n  [member 1][10/10] 364s train_loss=0.6092 train_dice=0.6307 | val F0.5*=0.6711 (thr 0.89, P=0.773, R=0.440) AP=0.6722 [raw]\n    *** new best val F0.5*=0.6711 (raw) ***\n  Member 1 best val F0.5*: 0.6711 (raw)\n  Member 1 TTA'd val (tissue only): F0.5*=0.7037 thr=0.841 AP=0.7045\n  [cache warm-up fragment 1 (held-out)] 26 slices, 2.69 GB in 102.3s (26 MB/s)\n    depth profile: mean range [60.2, 131.7] (peak slice #16, trough slice #25 of window) | scale=58.06\n  [member 1] 8-view TTA inference ...\n  Inference patches: 7705 (TTA views: 8)\n\n======================================================================\nENSEMBLE CALIBRATION (validation only): smoothing sigma + F0.5 threshold\n======================================================================\n  sigma= 0.0: val F0.5*=0.7037 thr=0.841 P=0.789 R=0.491 AP=0.7045\n  sigma= 1.0: val F0.5*=0.7096 thr=0.861 P=0.810 R=0.474 AP=0.7026\n  sigma= 2.0: val F0.5*=0.7101 thr=0.865 P=0.815 R=0.468 AP=0.7021\n  sigma= 3.0: val F0.5*=0.7099 thr=0.861 P=0.812 R=0.473 AP=0.7021\n  sigma= 4.0: val F0.5*=0.7099 thr=0.861 P=0.813 R=0.471 AP=0.7022\n  sigma= 6.0: val F0.5*=0.7097 thr=0.861 P=0.815 R=0.467 AP=0.7025\n\n  -> chosen sigma=2.0, threshold=0.865 (val F0.5*=0.7101); frozen before touching fragment-1 labels.\n\n======================================================================\nFRAGMENT 1 RESULTS (labels used for REPORTING only)\n======================================================================\n  member 1 (tu-convnext_tiny, window [12, 38, 1]) at its own val thr=0.841: F0.5=0.5148 P=0.659 R=0.274 Dice=0.3874\n\n  ENSEMBLE @ val-chosen thr=0.865, sigma=2.0\n    raw:           {'dice': 0.3611622750759125, 'iou': 0.2203770875930786, 'precision': 0.6964507102966309, 'recall': 0.24379391968250275, 'fbeta0.5': 0.5078611969947815}\n    postprocessed: {'dice': 0.36117199063301086, 'iou': 0.22038431465625763, 'precision': 0.6964387893676758, 'recall': 0.2438042163848877, 'fbeta0.5': 0.5078650712966919}\n  [diagnostic] Otsu thr=0.436 would give F0.5=0.4074 (NOT used)\n  [analysis only] ranking quality: AP=0.5047 | oracle F0.5=0.5205 at thr=0.806 (upper bound, NOT a reportable result)\n\nSaved metrics: /kaggle/working/v5_3_metrics_summary.json\nSaved comparison figure (input | ground truth | probability | prediction | TP/FP/FN overlay): /kaggle/working/v5_3_visualizations/fragment1_v5_3_overview.png\n\n======================================================================\nV5.3 COMPLETE\n======================================================================\nMember 1: tu-convnext_tiny window=[12, 38, 1] best val F0.5*=0.6711 | TTA'd val F0.5*=0.7037\n\nEnsemble: sigma=2.0 threshold=0.865 (chosen on validation by F0.5)\nFragment 1 F0.5 (raw): 0.5079 | (postprocessed): 0.5079 | Dice: 0.3612 | AP: 0.5047\n\nProbability map: /kaggle/working/fragment1_probability_v5_3.npy\nPrediction: /kaggle/working/fragment1_prediction_v5_3.png\nMetrics: /kaggle/working/v5_3_metrics_summary.json\n\n=== V5.3 DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# هيحفظ الفرق التفاضلي بين الslices لنفس الفراجمنت للتدريب\n# وبيعمل ensemble of 2 models\n# ============================================================\n# VESUVIUS INK DETECTION - V5 \"LEAN ENSEMBLE\"\n# ============================================================\n# Built on V4 (test F0.5=0.523), NOT on the DSN/LOFO/curriculum experiment\n# (test F0.5=0.48-0.49 after 8h/fold -- see rationale below).\n#\n# ------------------------------------------------------------------\n# WHY THIS DROPS THE DEPTH-SIGNATURE-MODULE / LOFO / CURRICULUM MACHINERY\n# ------------------------------------------------------------------\n# Your own experiment is the evidence here, not a guess:\n#   V4 (ConvNeXt + simple depth-fusion stem, single split): test F0.5 = 0.523\n#   V6 (+ LDDC + depth attention + PE + fiber-consistency + curriculum + LOFO):\n#       test F0.5 = 0.48-0.49, 8 HOURS for a single fold\n# Three concrete things in that run point at *why*:\n#   1. Curriculum sampling: epoch-1 F0.5 was 0.085 -- \"easy examples first\"\n#      delayed learning the hard cases the model is actually judged on.\n#   2. EMA-vs-raw checkpointing was picking EMA even in epochs where raw\n#      validated better (e.g. epoch 9: EMA=0.503 vs raw=0.536) -- a real bug,\n#      fixed below (see \"ADAPTIVE EMA\" in the training loop).\n#   3. The fiber-consistency loss's extra forward pass roughly doubled\n#      per-step cost for a component with no demonstrated benefit.\n# LOFO cross-validation tripled wall-clock cost on top of all that, which is\n# what actually made reaching a usable result within a Kaggle session\n# infeasible. This is a legitimate finding worth stating directly in the\n# paper: added depth-attention modeling did not improve cross-fragment\n# generalization over a simpler depth-fusion baseline here, suggesting the\n# bottleneck is domain shift, not model capacity -- a real negative ablation\n# result, not a failure to report.\n#\n# ------------------------------------------------------------------\n# WHAT THIS VERSION ADDS INSTEAD (single train[2,3]/test[1] split, no LOFO)\n# ------------------------------------------------------------------\n#  1. N-model ENSEMBLE (varied seed + encoder + small Z-depth offset) -- the\n#     single most reliably effective lever in this whole project's history.\n#  2. ADAPTIVE EMA: evaluates EMA and raw weights EVERY epoch and checkpoints\n#     whichever actually validates better that epoch (fixes the bug above).\n#     Training always continues from the raw/optimizer trajectory regardless\n#     of which gets checkpointed -- EMA is a shadow copy for the saved\n#     checkpoint only, never something training jumps onto mid-run.\n#  3. Full 8-view (D4 dihedral group) TTA at inference: original, h-flip,\n#     v-flip, both flips, and the three non-trivial rotations, plus transpose\n#     -- verified exactly self-consistent (apply-then-invert round-trips to\n#     the identity for all 8 views).\n#  4. Topology-aware (clDice) loss -- cheap (no extra forward pass), penalizes\n#     broken/disconnected thin strokes even when pixel Dice looks fine.\n#  5. Physics-informed synthetic ink -- Gaussian depth profile\n#     I(z)=A*exp(-(z-z0)^2/(2*sigma^2)) instead of an arbitrary intensity\n#     delta, AND (unlike the old fake-ink distractor) updates the label mask,\n#     so it's extra positive training signal, not just a distractor.\n#  6. Histogram-match augmentation reference pool built from TRAIN fragments\n#     only by default (Protocol A / strict) -- fragment 1's pixels never\n#     touch training. This is a correctness fix independent of whether DSN\n#     helped; worth keeping regardless.\n#  7. Everything proven-useful from V4 kept as-is: ConvNeXt+Unet (verified\n#     working, not UnetPlusPlus -- see prior fix), simple DepthFusionStem,\n#     architecture inspection report + zero-channel detector, AdaBN with a\n#     BatchNorm-location report, Otsu diagnostic, overlay visualizations.\n#\n# HONEST EXPECTATION: I cannot promise F0.5 >= 0.65 -- that remains a real,\n# open question about how much a 2-fragment training set can generalize to\n# a third, genuinely different fragment. These are the highest-confidence\n# levers given everything tested so far in this project; if you still fall\n# short after this, that's itself informative for the paper's discussion\n# section (a concrete, quantified generalization ceiling under this protocol).\n#\n# RUNTIME: 2 ensemble members, single split, no curriculum, no fiber-\n# consistency's extra forward pass -- should be substantially faster per\n# epoch than the V6 run in your log. Increase CFG.epochs now that the\n# per-epoch cost is back down; reduce N_ENSEMBLE_MODELS to 1 if you need to\n# fit within a tighter time budget.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    depth_indices_base = list(range(12, 38))     # 26 slices\n    in_channels = len(depth_indices_base)\n\n    patch_size = 320          # between V4's 480 and V6's 320 -- a runtime/context compromise\n    train_stride = 112\n    test_stride = 64\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 8                    # raised now that per-epoch cost is back down (no fiber-consistency 2nd forward pass, no curriculum)\n    early_stop_patience = 3\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- backbone / architecture (verified working: smp.Unet + tu-convnext_tiny;\n    #     do NOT switch to UnetPlusPlus with tu-* encoders -- see prior fix) -------\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- depth-aware fusion stem (V4's simpler, empirically-better-performing\n    #     stem -- the attention/LDDC/PE \"DepthSignatureModule\" variant was\n    #     tested and did not outperform this, at much higher cost; not included\n    #     here, see module docstring) -------------------------------------------\n    USE_DEPTH_FUSION_STEM = True\n    depth_stem_out_channels = 36\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- V5.1: auxiliary differential-supervision head ------------------------\n    # Directly implements \"teach the model what an ink-like depth differential\n    # looks like vs. a fiber/noise-like one, using the ground truth\". Branches\n    # a small head off the RAW grad/curvature features already computed inside\n    # DepthFusionStem (no extra forward pass -- those tensors exist either\n    # way), supervised with its own BCE against the same label mask. This is a\n    # weaker, much cheaper version of what the DSN experiment already tried\n    # (which added attention/LDDC on top of these same differentials and did\n    # not improve cross-fragment generalization) -- so treat a null result\n    # here as a real, informative outcome too, not a sign to add more capacity.\n    # Kept small and cheap deliberately: if it doesn't help, it cost almost\n    # nothing to find out.\n    USE_AUX_DIFFERENTIAL_LOSS = True\n    aux_differential_weight = 0.10\n    aux_head_hidden_channels = 16\n\n    # --- ensemble --------------------------------------------------------------\n    N_ENSEMBLE_MODELS = 1                                    # raise to 3 only if you have session time to spare\n    ensemble_seeds         = [42, 123, 2024][:3]\n    ensemble_encoders      = [\"tu-convnext_tiny\", \"efficientnet-b4\", \"tu-convnext_tiny\"][:3]\n    ensemble_depth_offsets = [0, 0, 3][:3]                    # small Z-shift for the 3rd member, if used\n\n    # --- LOSS --------------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.30\n    focal_tversky_weight = 0.25\n    topology_weight = 0.15            # clDice\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n    cldice_iters = 8\n\n    # --- physics-informed synthetic ink (adds label-positive training signal,\n    #     unlike the old label-free fake-ink distractor) --------------------------\n    USE_PHYSICAL_SYNTHETIC_INK = True\n    physical_ink_p = 0.15\n    physical_ink_max_amplitude = 60\n\n    # --- domain-randomization augmentation (kept from V3/V4, all cheap, no\n    #     extra forward pass) -----------------------------------------------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n    USE_SHADOW = True\n    shadow_p = 0.20\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # --- histogram-matching style augmentation -- Protocol A (strict): reference\n    #     pool built from TRAIN fragments only, fragment 1 never touches training\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- THRESHOLD --------------------------------------------------------------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n\n    # --- TTA: full 8-view D4 dihedral group (verified exactly self-consistent) --\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"hvflip\", \"rot90\", \"rot180\", \"rot270\", \"transpose\"]\n\n    # --- AdaBN (see the \"where are the BN layers\" report at model-build time) --\n    use_adabn = False\n    adabn_max_patches = 2000\n\n    # --- EMA (see ADAPTIVE EMA note in the training loop for the bugfix) -------\n    USE_EMA = True\n    ema_decay = 0.999\n\n    # --- POSTPROCESS ----------------------------------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    # --- edge-artifact cropping ---------------------------------------------------\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- CLAHE mode (see V4 rationale: \"global_shared\" preserves inter-slice\n    #     depth relationships far better than independent per-slice CLAHE) -------\n    CLAHE_MODE = \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization ------------------------------------------------------------\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    viz_dir = os.path.join(out_dir, \"v5_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v5_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS (+ mask erosion)\n# ============================================================\n\ndef load_tissue_mask(frag_dir, depth_indices):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = depth_indices[len(depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion emptied the mask for {frag_dir}; keeping un-eroded mask.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            if mask[y:y + patch_size, x:x + patch_size].mean() > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef clip_depth_indices(base_indices, offset, max_slice_idx=64):\n    shifted = [i + offset for i in base_indices]\n    lo, hi = min(shifted), max(shifted)\n    if lo < 0:\n        shifted = [i - lo for i in shifted]\n    elif hi > max_slice_idx:\n        shifted = [i - (hi - max_slice_idx) for i in shifted]\n    return shifted\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3, rng=random):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = float(labels_full[fid][y:y + patch_size, x:x + patch_size].mean())\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = max(int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio), len(positive))\n    negative_selected = rng.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    samples_balanced = positive * repeat + negative_selected\n    rng.shuffle(samples_balanced)\n    print(f\"  Positive patches: {len(positive)} | Negative patches: {len(negative)} | \"\n          f\"Balanced: {len(samples_balanced)} (positive ratio={len(positive)*repeat/max(len(samples_balanced),1):.3f})\")\n    return samples_balanced\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME (CLAHE modes, per-fragment stats)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid = len(self._slices) // 2\n        ref_slice = np.asarray(self._slices[mid])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def read_patch(self, y, x, size, apply_clahe=None):\n        self._ensure_open()\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0]); ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[i] = block\n        return out\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        mid = len(self._slices) // 2\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            block = self._slices[mid][y:y + patch_size, x:x + patch_size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                block = self._clahe.apply(block)\n            elif CFG.CLAHE_MODE == \"global_shared\":\n                block = cv2.LUT(block, self._shared_clahe_lut)\n            vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 6. AUGMENTATION: domain-randomization + physics-informed synthetic ink\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Label-FREE distractor: looks locally ink-like but is never marked\n    positive, so the model can't shortcut on raw brightness alone.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)], axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        n_affected = max(1, int(d * random.uniform(0.3, 0.7)))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_physical_synthetic_ink(img_hwd, label, n_strokes=None, max_amplitude=CFG.physical_ink_max_amplitude):\n    \"\"\"Label-POSITIVE synthetic ink: models a stroke's depth profile as a Gaussian\n    I(z) = A*exp(-(z-z0)^2/(2*sigma^2)) instead of an arbitrary flat intensity\n    delta -- physically closer to how a real ink deposit attenuates through\n    nearby depth slices -- AND updates the label mask, giving extra positive\n    training signal with a controllable depth profile (unlike inject_fake_ink,\n    which is deliberately label-free).\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    label_out = label.copy()\n    n_strokes = n_strokes or random.randint(1, 2)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)], axis=1).astype(np.int32)\n        thickness = random.randint(2, 5)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n\n        z0 = random.uniform(d * 0.3, d * 0.7)\n        sigma = random.uniform(1.5, 4.0)\n        amplitude = random.uniform(20, max_amplitude) * random.choice([-1, 1])\n        z_idx = np.arange(d)\n        depth_profile = amplitude * np.exp(-((z_idx - z0) ** 2) / (2 * sigma ** 2))\n\n        for zi in range(d):\n            delta = depth_profile[zi]\n            if abs(delta) < 1.0:\n                continue\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n        label_out[stroke_mask > 0] = 1\n    return out, label_out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    \"\"\"Channel-count-agnostic shadow augmentation. albumentations' RandomShadow\n    hard-requires 3-channel RGB in every version and crashes on multi-channel\n    depth-stack data -- this replaces it entirely.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist_norm = np.clip(np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2) / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\ndef _safe_transform(builder_new, builder_old, name):\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable ({e1} / {e2}); skipping.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma, border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None: transforms.append(t)\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03, p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None: transforms.append(t)\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None: transforms.append(t)\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 7. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size]\n        img = np.transpose(patch, (1, 2, 0))\n\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n\n        if self.train_mode and CFG.USE_PHYSICAL_SYNTHETIC_INK and random.random() < CFG.physical_ink_p:\n            img, label = inject_physical_synthetic_ink(img, label)\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img, vol.frag_mean, vol.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ============================================================\n# 8. MODEL: DepthFusionStem + verified-working ConvNeXt+Unet\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"Raw + first/second finite differences along the ordered depth axis,\n    mixed via a 1x1 conv into a compact learned representation. This is V4's\n    stem (empirically the better performer vs. the attention-based variant\n    tested separately -- see module docstring).\n\n    V5.1: also stores the RAW (pre-mix) grad+curvature tensor as\n    `self.last_diff_features`, purely as a hook for AuxDifferentialHead to\n    read -- this is the exact \"differentiation of every pixel across depth\"\n    signal, computed once, at zero extra cost (it's already computed for the\n    main mix() below; storing a reference costs nothing).\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = (in_depth if use_raw else 0) + ((in_depth - 1) if use_grad else 0) + \\\n                   (max(in_depth - 2, 0) if use_curv else 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),   # kept for AdaBN to have something real to recalibrate\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.diff_channels = ((in_depth - 1) if use_grad else 0) + (max(in_depth - 2, 0) if use_curv else 0)\n        self.last_diff_features = None\n\n    def forward(self, x):\n        parts = []\n        diff_parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            grad = x[:, 1:, :, :] - x[:, :-1, :, :]\n            parts.append(grad)\n            diff_parts.append(grad)\n        if self.use_curv and x.shape[1] > 2:\n            curv = x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :]\n            parts.append(curv)\n            diff_parts.append(curv)\n        self.last_diff_features = torch.cat(diff_parts, dim=1) if diff_parts else None\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass AuxDifferentialHead(nn.Module):\n    \"\"\"V5.1: a small head operating DIRECTLY on the raw depth-differential\n    features (grad + curvature across the 26 slices, before the main stem's\n    1x1 mix), supervised with its own loss against the ground-truth mask.\n    This is the concrete implementation of \"does this differential pattern\n    look like ink or like fiber/noise\" -- a 3x3 conv gives it a little local\n    spatial context (a single pixel's differential alone is noisy; its\n    neighborhood's differential pattern is more informative), then a 1x1\n    conv to a single logit map at the SAME spatial resolution as the input\n    patch, so it can be supervised directly with no resizing.\"\"\"\n    def __init__(self, in_channels, hidden=16):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_channels, hidden, kernel_size=3, padding=1),\n            nn.BatchNorm2d(hidden),\n            nn.GELU(),\n            nn.Conv2d(hidden, 1, kernel_size=1),\n        )\n\n    def forward(self, diff_features):\n        return self.net(diff_features)\n\n\nclass DepthAwareSegModel(nn.Module):\n    def __init__(self, depth_stem, seg_model, aux_head=None):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n        self.aux_head = aux_head\n        self.last_aux_logits = None   # populated on forward() when aux_head is set\n\n    def forward(self, x):\n        feat = self.depth_stem(x) if self.depth_stem is not None else x\n        if (self.aux_head is not None and self.depth_stem is not None\n                and self.depth_stem.last_diff_features is not None):\n            self.last_aux_logits = self.aux_head(self.depth_stem.last_diff_features)\n        else:\n            self.last_aux_logits = None\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' -- only 'unet' is verified \"\n              f\"to handle tu-* timm encoders' 0-channel placeholder stage correctly.\")\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    return arch_cls(encoder_name=encoder_name, encoder_weights=CFG.encoder_weights, in_channels=in_channels,\n                     classes=1, decoder_channels=CFG.decoder_channels, decoder_attention_type=CFG.decoder_attention_type)\n\n\ndef build_model(encoder_name):\n    stem = None\n    aux_head = None\n    seg_in_channels = CFG.in_channels\n    if CFG.USE_DEPTH_FUSION_STEM:\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = CFG.depth_stem_out_channels\n        if CFG.USE_AUX_DIFFERENTIAL_LOSS and stem.diff_channels > 0:\n            aux_head = AuxDifferentialHead(stem.diff_channels, hidden=CFG.aux_head_hidden_channels)\n            print(f\"[backbone] AuxDifferentialHead attached: {stem.diff_channels} raw differential \"\n                  f\"channels -> {CFG.aux_head_hidden_channels} hidden -> 1 logit (weight=\"\n                  f\"{CFG.aux_differential_weight})\")\n\n    try:\n        seg_model = _build_seg_backbone(encoder_name, seg_in_channels)\n        print(f\"[backbone] using {encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {encoder_name} unavailable/failed ({e}); falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model, aux_head=aux_head)\n\n\ndef inspect_model_channels(model):\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\ndef report_batchnorm_locations(model):\n    counts = {\"depth_stem\": 0, \"encoder\": 0, \"decoder\": 0, \"other\": 0}\n    for name, module in model.named_modules():\n        if isinstance(module, nn.BatchNorm2d):\n            if name.startswith(\"depth_stem\"):\n                counts[\"depth_stem\"] += 1\n            elif \".encoder.\" in f\".{name}.\" or name.endswith(\".encoder\"):\n                counts[\"encoder\"] += 1\n            elif \".decoder.\" in f\".{name}.\" or name.endswith(\".decoder\"):\n                counts[\"decoder\"] += 1\n            else:\n                counts[\"other\"] += 1\n    total = sum(counts.values())\n    print(f\"  [AdaBN report] BatchNorm2d layers: depth_stem={counts['depth_stem']} encoder={counts['encoder']} \"\n          f\"decoder={counts['decoder']} other={counts['other']} (total={total})\")\n    return counts\n\n\n@torch.no_grad()\ndef run_architecture_report(model, member_tag, full=True):\n    print(f\"\\n{'='*70}\\nARCHITECTURE INSPECTION [{member_tag}]\\n{'='*70}\")\n    n_params = sum(p.numel() for p in model.parameters())\n    print(f\"Parameters: {n_params/1e6:.2f}M total\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"ZERO/INVALID CHANNEL LAYERS FOUND:\", bad_layers)\n        raise RuntimeError(\"Model contains invalid zero-channel layers -- fix before training.\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    if full:\n        model.eval()\n        small_size = min(CFG.patch_size, 256)\n        dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n        out = model(dummy)\n        has_nan = torch.isnan(out).any().item()\n        has_inf = torch.isinf(out).any().item()\n        print(f\"Dry-run ({small_size}x{small_size}) output: {tuple(out.shape)} | \"\n              f\"NaN: {'FAIL' if has_nan else 'PASS'} | Inf: {'FAIL' if has_inf else 'PASS'}\")\n        if has_nan or has_inf:\n            raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass.\")\n        model.train()\n        if torch.cuda.is_available():\n            print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")\n    report_batchnorm_locations(model)\n    print(\"=\" * 70)\n\n\n# ============================================================\n# 9. LOSSES: BCE + Dice + Focal Tversky + clDice (topology)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\ndef soft_erode(I):\n    p1 = -F.max_pool2d(-I, (3, 1), (1, 1), (1, 0))\n    p2 = -F.max_pool2d(-I, (1, 3), (1, 1), (0, 1))\n    return torch.min(p1, p2)\n\n\ndef soft_dilate(I):\n    return F.max_pool2d(I, (3, 3), (1, 1), (1, 1))\n\n\ndef soft_open(I):\n    return soft_dilate(soft_erode(I))\n\n\ndef soft_skeletonize(I, iters=CFG.cldice_iters):\n    \"\"\"Differentiable soft skeletonization (Shit et al., 2021, clDice).\"\"\"\n    I1 = soft_open(I)\n    skel = F.relu(I - I1)\n    for _ in range(iters):\n        I = soft_erode(I)\n        I1 = soft_open(I)\n        delta = F.relu(I - I1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef soft_cldice(pred_probs, target, iters=CFG.cldice_iters, smooth=1.0):\n    \"\"\"Penalizes broken/disconnected thin strokes even when pixel Dice looks\n    fine -- no extra forward pass, just extra pooling ops on the mask.\"\"\"\n    pred_sk = soft_skeletonize(pred_probs, iters)\n    targ_sk = soft_skeletonize(target, iters)\n    tprec = (torch.sum(pred_sk * target) + smooth) / (torch.sum(pred_sk) + smooth)\n    tsens = (torch.sum(targ_sk * pred_probs) + smooth) / (torch.sum(targ_sk) + smooth)\n    return 1.0 - 2.0 * (tprec * tsens) / (tprec + tsens + 1e-8)\n\n\nclass V5ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets, aux_logits=None):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        probs = torch.sigmoid(logits)\n        topo = soft_cldice(probs, targets)\n        loss = (CFG.bce_weight * bce + CFG.dice_weight * dice +\n                CFG.focal_tversky_weight * tv + CFG.topology_weight * topo)\n\n        aux_loss_val = None\n        if aux_logits is not None and CFG.USE_AUX_DIFFERENTIAL_LOSS:\n            # V5.1: plain BCE, no dice/tversky here -- this head only needs to\n            # learn \"does this raw differential pattern look ink-like\", a much\n            # simpler target than the main segmentation output, so a heavier\n            # loss would be overkill and could destabilize the small head.\n            aux_bce = F.binary_cross_entropy_with_logits(aux_logits, targets, pos_weight=self.pos_weight)\n            loss = loss + CFG.aux_differential_weight * aux_bce\n            aux_loss_val = aux_bce.item()\n        return loss, aux_loss_val\n\n\n# ============================================================\n# 10. METRICS\n# ============================================================\n\ndef calculate_metrics(tp, fp, fn):\n    eps = 1e-6\n    dice = (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return calculate_metrics(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 11. SHARED DATA (computed ONCE, reused across ensemble members)\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING SHARED TRAIN/VAL PATCH GRID (fragments 2 & 3)\")\nprint(\"=\" * 70)\n\n_shared_masks, _shared_labels = {}, {}\nall_samples = []\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir, CFG.depth_indices_base)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {mask.shape} {len(coords)} candidate patches\")\n    _shared_masks[fid] = mask\n    _shared_labels[fid] = labels\n    all_samples.extend([(fid, y, x) for y, x in coords])\n\ngroups = {}\nfor fid, y, x in all_samples:\n    key = (fid, y // 768, x // 768)\n    groups.setdefault(key, []).append((fid, y, x))\nkeys = list(groups.keys())\nrng0 = random.Random(CFG.seed)\nrng0.shuffle(keys)\nn_val_groups = max(1, int(len(keys) * CFG.val_fraction))\nval_keys = set(keys[:n_val_groups])\ntrain_grid, val_samples = [], []\nfor key, items in groups.items():\n    (val_samples if key in val_keys else train_grid).extend(items)\nprint(f\"Spatial train: {len(train_grid)} | Spatial val: {len(val_samples)}\")\n\n\ndef build_train_samples(seed):\n    r = random.Random(seed)\n    return balance_positive_patches(train_grid, _shared_labels, CFG.patch_size,\n                                     positive_threshold=CFG.positive_patch_fraction,\n                                     target_positive_ratio=CFG.target_positive_patch_ratio,\n                                     max_positive_repeat=CFG.max_positive_repeat, rng=r)\n\n\n# --- test fragment (unlabeled use only during training: never touches training\n#     except through per-fragment normalization stats, which are self-derived,\n#     not learned-weight-affecting -- see PROTOCOL note in the module docstring) --\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir, CFG.depth_indices_base)\nprint(f\"\\nTest fragment {CFG.test_frag}: {test_mask.shape}\")\n\n# Protocol A (strict): histogram-match augmentation reference pool built from\n# TRAIN fragments only -- fragment 1's pixels never touch training.\npool_source = []\nif CFG.USE_HIST_MATCH_AUG:\n    for fid in CFG.train_frags:\n        c = generate_grid_coords(_shared_masks[fid], CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n        pool_source.extend([(fid, y, x) for y, x in c])\n    print(f\"[Protocol A - strict] histogram-match reference pool built from TRAIN fragments only \"\n          f\"({len(pool_source)} candidates) -- fragment 1 untouched during training.\")\n\n\n# ============================================================\n# 12. INFERENCE HELPERS (shared across members)\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\": return x\n    if mode == \"hflip\": return torch.flip(x, dims=[3])\n    if mode == \"vflip\": return torch.flip(x, dims=[2])\n    if mode == \"hvflip\": return torch.flip(x, dims=[2, 3])\n    if mode == \"rot90\": return torch.rot90(x, k=1, dims=[2, 3])\n    if mode == \"rot180\": return torch.rot90(x, k=2, dims=[2, 3])\n    if mode == \"rot270\": return torch.rot90(x, k=3, dims=[2, 3])\n    if mode == \"transpose\": return x.transpose(2, 3)\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\": return pred\n    if mode == \"hflip\": return torch.flip(pred, dims=[2])\n    if mode == \"vflip\": return torch.flip(pred, dims=[1])\n    if mode == \"hvflip\": return torch.flip(pred, dims=[1, 2])\n    if mode == \"rot90\": return torch.rot90(pred, k=-1, dims=[1, 2])\n    if mode == \"rot180\": return torch.rot90(pred, k=-2, dims=[1, 2])\n    if mode == \"rot270\": return torch.rot90(pred, k=-3, dims=[1, 2])\n    if mode == \"transpose\": return pred.transpose(1, 2)\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, batch_size, pre_transform=None):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"  Inference patches: {len(coords)} (TTA views: {len(CFG.tta_modes) if CFG.use_tta else 1})\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        probs = predict_batch_tta(model, np.stack(batch_imgs).astype(np.float32))\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        cleanup_memory()\n\n    for y, x in coords:\n        raw_u8 = vol.read_patch(y, x, patch_size)\n        if pre_transform is not None:\n            raw_u8 = pre_transform(raw_u8)\n        raw = raw_u8.astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm(model, vol, mask, patch_size, stride, max_patches, batch_size):\n    n_reset = 0\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n            n_reset += 1\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"  AdaBN: resetting {n_reset} BatchNorm2d layers, recalibrating on {len(coords)} patches\")\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        batch = coords[start:start + batch_size]\n        imgs = [normalize_patch(vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0, vol.frag_mean, vol.frag_std)\n                for y, x in batch]\n        inp = torch.from_numpy(np.stack(imgs)).to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            model(inp)\n        del inp, imgs\n    model.eval()\n    cleanup_memory()\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return calculate_metrics(tp, fp, fn)\n\n\n# ============================================================\n# 13. TRAIN + INFER ONE ENSEMBLE MEMBER\n# ============================================================\n\ndef train_and_infer_one_member(member_idx, seed, encoder_name, depth_offset):\n    print(f\"\\n{'#'*70}\\n# ENSEMBLE MEMBER {member_idx+1}/{CFG.N_ENSEMBLE_MODELS}  \"\n          f\"seed={seed} encoder={encoder_name} depth_offset={depth_offset}\\n{'#'*70}\")\n    set_seed(seed)\n    depth_indices = clip_depth_indices(CFG.depth_indices_base, depth_offset)\n\n    train_volumes = {fid: FragmentVolume(os.path.join(CFG.base_dir, fid), depth_indices) for fid in CFG.train_frags}\n    for fid in CFG.train_frags:\n        train_volumes[fid].compute_fragment_stats(_shared_masks[fid], CFG.patch_size)\n        print(f\"  fragment {fid} norm stats: mean={train_volumes[fid].frag_mean:.2f} std={train_volumes[fid].frag_std:.2f}\")\n\n    train_samples = build_train_samples(seed)\n\n    member_hist_pool = None\n    if CFG.USE_HIST_MATCH_AUG and pool_source:\n        sampled = random.sample(pool_source, min(CFG.hist_match_pool_size, len(pool_source)))\n        member_hist_pool = [train_volumes[fid].read_patch(y, x, CFG.patch_size) for fid, y, x in sampled]\n\n    train_ds = InkPatchDataset(train_volumes, _shared_labels, train_samples, CFG.patch_size,\n                                transform=build_train_transform(), jitter=CFG.train_jitter,\n                                hist_match_pool=member_hist_pool, train_mode=True)\n    val_ds = InkPatchDataset(train_volumes, _shared_labels, val_samples, CFG.patch_size,\n                              transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\n    train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True, num_workers=CFG.num_workers,\n                               pin_memory=(CFG.device == \"cuda\"), drop_last=CFG.drop_last,\n                               persistent_workers=CFG.num_workers > 0)\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers,\n                             pin_memory=(CFG.device == \"cuda\"), persistent_workers=CFG.num_workers > 0)\n\n    pos_frac = estimate_positive_fraction(_shared_labels, train_samples, CFG.patch_size, n_samples=100)\n    initial_bias = math.log(pos_frac / max(1 - pos_frac, 1e-6))\n    bce_pos_weight = float(np.clip(np.sqrt((1 - pos_frac) / pos_frac), 1.0, 8.0))\n    print(f\"  positive fraction={pos_frac:.5f} bias={initial_bias:.3f} pos_weight={bce_pos_weight:.2f}\")\n\n    model = build_model(encoder_name).to(CFG.device)\n    run_architecture_report(model, member_tag=f\"member{member_idx}\", full=(member_idx == 0))\n\n    with torch.no_grad():\n        try:\n            model.seg_model.segmentation_head[0].bias.fill_(initial_bias)\n        except Exception:\n            pass\n\n    pos_weight_t = torch.tensor([bce_pos_weight], dtype=torch.float32, device=CFG.device)\n    criterion = V5ComboLoss(pos_weight=pos_weight_t)\n\n    encoder_params, decoder_params, stem_params = [], [], []\n    for name, p in model.named_parameters():\n        if not p.requires_grad:\n            continue\n        (stem_params if name.startswith(\"depth_stem.\") else\n         encoder_params if name.startswith(\"seg_model.encoder.\") else decoder_params).append(p)\n    param_groups = [{\"params\": encoder_params, \"lr\": CFG.encoder_lr}, {\"params\": decoder_params, \"lr\": CFG.decoder_lr}]\n    if stem_params:\n        param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\n    optimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\n    class EMAModel:\n        def __init__(self, model, decay):\n            self.decay = decay\n            self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n        @torch.no_grad()\n        def update(self, model):\n            for k, v in model.state_dict().items():\n                if v.dtype.is_floating_point:\n                    self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n                else:\n                    self.shadow[k] = v.detach().clone()\n\n    ema = EMAModel(model, CFG.ema_decay) if CFG.USE_EMA else None\n\n    def run_epoch(loader, train_mode, threshold=0.5):\n        model.train(train_mode)\n        total_loss = 0.0\n        total_aux_loss = 0.0\n        n_aux = 0\n        acc = GlobalConfusionAccumulator()\n        for imgs, masks in loader:\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            masks = masks.to(CFG.device, non_blocking=True)\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n            with torch.set_grad_enabled(train_mode):\n                with autocast(enabled=(CFG.device == \"cuda\")):\n                    logits = model(imgs)\n                    # V5.1: model.last_aux_logits is set by DepthAwareSegModel's\n                    # forward() as a side effect of this SAME call above -- no\n                    # extra forward pass, just reading an attribute it already\n                    # computed.\n                    aux_logits = getattr(model, \"last_aux_logits\", None)\n                    loss, aux_loss_val = criterion(logits, masks, aux_logits=aux_logits)\n                if train_mode:\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    if ema is not None:\n                        ema.update(model)\n            probs = torch.sigmoid(logits.detach())\n            acc.update(probs, masks, threshold)\n            total_loss += float(loss.item())\n            if aux_loss_val is not None:\n                total_aux_loss += aux_loss_val\n                n_aux += 1\n            del imgs, masks, logits, probs\n        avg_aux = total_aux_loss / max(n_aux, 1)\n        return total_loss / max(len(loader), 1), acc.compute(), avg_aux\n\n    print(f\"\\n  Training member {member_idx+1} for up to {CFG.epochs} epochs \"\n          f\"(patch={CFG.patch_size}, batch={CFG.batch_size}) ...\")\n    best_val_dice, epochs_no_improve = -1.0, 0\n    best_state, best_used = None, None\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n        train_loss, train_metrics, train_aux_loss = run_epoch(train_loader, train_mode=True)\n\n        # --- ADAPTIVE EMA: evaluate BOTH raw and EMA weights every epoch, checkpoint\n        # whichever actually validates better THIS epoch. Training always continues\n        # from raw regardless of which gets checkpointed (fixes the bug in the V6\n        # log where a worse-validating EMA snapshot was checkpointed anyway). ---\n        raw_state = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        val_loss_raw, val_metrics_raw, _ = run_epoch(val_loader, train_mode=False)\n\n        if ema is not None:\n            model.load_state_dict(ema.shadow)\n            val_loss_ema, val_metrics_ema, _ = run_epoch(val_loader, train_mode=False)\n            model.load_state_dict(raw_state)   # ALWAYS restore raw before next training epoch\n\n            if val_metrics_ema[\"dice\"] >= val_metrics_raw[\"dice\"]:\n                val_metrics, candidate_state, used = val_metrics_ema, {k: v.clone() for k, v in ema.shadow.items()}, \"ema\"\n            else:\n                val_metrics, candidate_state, used = val_metrics_raw, raw_state, \"raw\"\n            print(f\"    [checkpoint choice] EMA dice={val_metrics_ema['dice']:.4f} \"\n                  f\"raw dice={val_metrics_raw['dice']:.4f} -> using {used}\")\n        else:\n            val_metrics, candidate_state, used = val_metrics_raw, raw_state, \"raw\"\n\n        scheduler.step()\n        print(f\"  [member {member_idx+1}][{epoch:02d}/{CFG.epochs}] time={time.time()-t0:.1f}s \"\n              f\"train_loss={train_loss:.5f} train_dice={train_metrics['dice']:.5f} \"\n              f\"aux_loss={train_aux_loss:.5f} | \"\n              f\"val_dice={val_metrics['dice']:.5f} val_P={val_metrics['precision']:.4f} \"\n              f\"val_R={val_metrics['recall']:.4f} val_f0.5={val_metrics['fbeta0.5']:.4f} [{used}]\")\n\n        if val_metrics[\"dice\"] > best_val_dice:\n            best_val_dice, epochs_no_improve = val_metrics[\"dice\"], 0\n            best_state, best_used = candidate_state, used\n            print(f\"    *** new best (val_dice={best_val_dice:.5f}, {used} weights) ***\")\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= CFG.early_stop_patience:\n                print(\"    Early stopping.\")\n                break\n        cleanup_memory()\n\n    model.load_state_dict(best_state)\n    model.eval()\n    print(f\"  Member {member_idx+1} best validation Dice: {best_val_dice:.5f} ({best_used} weights)\")\n\n    # --- validation threshold search ---\n    @torch.no_grad()\n    def find_best_dice_threshold():\n        thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n        tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n        for imgs, masks in val_loader:\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            masks_np = masks.numpy()\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n            probs = torch.sigmoid(logits).float().cpu().numpy()\n            for i, t in enumerate(thresholds):\n                preds = (probs > t).astype(np.float32)\n                tp[i] += (preds * masks_np).sum()\n                fp[i] += (preds * (1.0 - masks_np)).sum()\n                fn[i] += ((1.0 - preds) * masks_np).sum()\n            del imgs, masks, logits, probs\n        dice_scores = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n        best_idx = int(np.argmax(dice_scores))\n        return float(thresholds[best_idx]), float(dice_scores[best_idx])\n\n    val_threshold, val_threshold_dice = find_best_dice_threshold()\n    print(f\"  Member {member_idx+1} val threshold={val_threshold:.2f} (val dice={val_threshold_dice:.4f})\")\n\n    for v in train_volumes.values():\n        v.close()\n    del train_loader, val_loader, train_ds, val_ds\n    cleanup_memory()\n\n    # --- fragment-1 inference: this member's own (possibly Z-shifted) test volume ---\n    test_vol = FragmentVolume(test_dir, depth_indices)\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n\n    print(f\"\\n  [member {member_idx+1}] baseline + {len(CFG.tta_modes) if CFG.use_tta else 1}-view TTA inference ...\")\n    baseline_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.infer_batch)\n\n    adabn_prob = None\n    if CFG.use_adabn:\n        model = recalibrate_batchnorm(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                       CFG.adabn_max_patches, CFG.infer_batch)\n        adabn_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.infer_batch)\n\n    final_member_prob = adabn_prob if adabn_prob is not None else baseline_prob\n    otsu_thr = compute_otsu_threshold(final_member_prob, test_mask, fallback=val_threshold)\n\n    test_vol.close()\n    cleanup_memory()\n\n    return {\n        \"member_idx\": member_idx, \"seed\": seed, \"encoder\": encoder_name, \"depth_offset\": depth_offset,\n        \"best_val_dice\": best_val_dice, \"val_threshold\": val_threshold, \"otsu_threshold\": otsu_thr,\n        \"probability\": final_member_prob,\n    }\n\n\n# ============================================================\n# 14. RUN THE ENSEMBLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(f\"STARTING V5 ENSEMBLE ({CFG.N_ENSEMBLE_MODELS} members, single split, no LOFO)\")\nprint(\"=\" * 70)\n\nmember_results = []\nfor i in range(CFG.N_ENSEMBLE_MODELS):\n    member_results.append(train_and_infer_one_member(\n        member_idx=i, seed=CFG.ensemble_seeds[i],\n        encoder_name=CFG.ensemble_encoders[i], depth_offset=CFG.ensemble_depth_offsets[i]))\n    cleanup_memory()\n\n\n# ============================================================\n# 15. LOAD TEST LABELS (diagnostics only, loaded AFTER all training) + ENSEMBLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST LABELS (diagnostics only -- never used during training/model selection)\")\nprint(\"=\" * 70)\ntest_labels = load_ink_labels(test_dir)\ngt_test = (test_labels * test_mask).astype(np.float32) if test_labels is not None else None\nprint(\"Test GT available:\", test_labels is not None)\n\nensembled_probability = np.mean([r[\"probability\"] for r in member_results], axis=0).astype(np.float32)\nensembled_threshold = compute_otsu_threshold(\n    ensembled_probability, test_mask, fallback=float(np.mean([r[\"val_threshold\"] for r in member_results])))\n\nfinal_prediction = postprocess(ensembled_probability, ensembled_threshold)\n\nprint(f\"\\nEnsembled ({CFG.N_ENSEMBLE_MODELS} members) + Otsu threshold={ensembled_threshold:.3f}\")\nfor r in member_results:\n    single_metrics = evaluate_probability_map(r[\"probability\"], gt_test, r[\"otsu_threshold\"]) if gt_test is not None else None\n    print(f\"  member {r['member_idx']+1} ({r['encoder']}, seed={r['seed']}): val_dice={r['best_val_dice']:.4f} \"\n          f\"-> solo test dice={single_metrics['dice'] if single_metrics else 'n/a'}\")\n\nensembled_metrics_raw = evaluate_probability_map(ensembled_probability, gt_test, ensembled_threshold) if gt_test is not None else None\nif ensembled_metrics_raw is not None:\n    print(\"Ensembled RAW metrics:\", ensembled_metrics_raw)\n\npost_metrics = None\nif gt_test is not None:\n    pred_f = final_prediction.astype(np.float32)\n    tp = (pred_f * gt_test).sum(); fp = (pred_f * (1 - gt_test)).sum(); fn = ((1 - pred_f) * gt_test).sum()\n    post_metrics = calculate_metrics(tp, fp, fn)\n    print(\"Ensembled POSTPROCESSED metrics:\", post_metrics)\n\n\n# ============================================================\n# 16. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_ensembled_v5.npy\")\nnp.save(prob_path, ensembled_probability)\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_ensembled_v5.png\")\ncv2.imwrite(pred_path, (final_prediction * 255).astype(np.uint8))\n\nfor r in member_results:\n    np.save(os.path.join(CFG.out_dir, f\"fragment1_probability_member{r['member_idx']}_v5.npy\"), r[\"probability\"])\n\nmetrics_summary = {\n    \"n_ensemble_models\": CFG.N_ENSEMBLE_MODELS,\n    \"ensembled_threshold\": float(ensembled_threshold),\n    \"ensembled_raw\": ensembled_metrics_raw,\n    \"ensembled_postprocessed\": post_metrics,\n    \"members\": [{\"idx\": r[\"member_idx\"], \"seed\": r[\"seed\"], \"encoder\": r[\"encoder\"],\n                 \"depth_offset\": r[\"depth_offset\"], \"best_val_dice\": r[\"best_val_dice\"],\n                 \"val_threshold\": r[\"val_threshold\"], \"otsu_threshold\": r[\"otsu_threshold\"]}\n                for r in member_results],\n    \"config\": cfg_to_dict(CFG),\n}\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump(metrics_summary, f, indent=2)\nprint(f\"\\nSaved metrics: {CFG.metrics_path}\")\n\n\n# ============================================================\n# 17. VISUALIZATION (with prediction-vs-ground-truth overlay)\n# ============================================================\n\ndef make_confusion_overlay(input_gray_u8, gt_binary, pred_binary, alpha=0.55):\n    base = np.stack([input_gray_u8] * 3, axis=-1).astype(np.float32)\n    overlay = base.copy()\n    gt_b, pred_b = gt_binary > 0, pred_binary > 0\n    tp, fp, fn = pred_b & gt_b, pred_b & ~gt_b, ~pred_b & gt_b\n    overlay[tp] = (1 - alpha) * base[tp] + alpha * np.array([0, 255, 0], dtype=np.float32)\n    overlay[fp] = (1 - alpha) * base[fp] + alpha * np.array([255, 0, 0], dtype=np.float32)\n    overlay[fn] = (1 - alpha) * base[fn] + alpha * np.array([255, 255, 0], dtype=np.float32)\n    return np.clip(overlay, 0, 255).astype(np.uint8)\n\n\ndef save_overview():\n    mid_idx = CFG.depth_indices_base[len(CFG.depth_indices_base) // 2]\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    small_u8 = cv2.normalize(small, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n    pred_small = cv2.resize((final_prediction * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    prob_small = cv2.resize((ensembled_probability * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_small = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n        overlay = make_confusion_overlay(small_u8, (gt_small > 127).astype(np.uint8), (pred_small > 127).astype(np.uint8))\n        fig, axes = plt.subplots(1, 5, figsize=(27, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n        axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability thr={ensembled_threshold:.2f}\")\n        axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"Ensembled Prediction\")\n        axes[4].imshow(overlay); axes[4].set_title(\"Overlay: green=TP red=FP yellow=FN\")\n    else:\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(prob_small, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n        axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"Ensembled Prediction\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_v5_overview.png\")\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(\"Saved overview:\", path)\n\n\nsave_overview()\ncleanup_memory()\n\n\n# ============================================================\n# 18. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V5 LEAN ENSEMBLE COMPLETE\")\nprint(\"=\" * 70)\nfor r in member_results:\n    print(f\"Member {r['member_idx']+1}: {r['encoder']} seed={r['seed']} depth_offset={r['depth_offset']} \"\n          f\"val_dice={r['best_val_dice']:.5f}\")\nprint(f\"\\nEnsembled threshold (Otsu): {ensembled_threshold:.3f}\")\nif ensembled_metrics_raw is not None:\n    print(f\"Ensembled test F0.5 (raw): {ensembled_metrics_raw['fbeta0.5']:.5f} | Dice: {ensembled_metrics_raw['dice']:.5f}\")\nif post_metrics is not None:\n    print(f\"Ensembled test F0.5 (postprocessed): {post_metrics['fbeta0.5']:.5f} | Dice: {post_metrics['dice']:.5f}\")\nprint(f\"\\nProbability map: {prob_path}\")\nprint(f\"Prediction: {pred_path}\")\nprint(f\"Metrics: {CFG.metrics_path}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")\nprint(\"\\n=== V5 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T09:42:02.213316Z","iopub.execute_input":"2026-09-29T09:42:02.214159Z"}},"outputs":[{"name":"stdout","text":"   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 4.6 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.7/2.7 MB 54.4 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.30\n======================================================================\nBUILDING SHARED TRAIN/VAL PATCH GRID (fragments 2 & 3)\n======================================================================\nFragment 2: (14830, 9506) 8049 candidate patches\nFragment 3: (7606, 5249) 2123 candidate patches\nSpatial train: 8083 | Spatial val: 2089\n\nTest fragment 1: (8181, 6330)\n[Protocol A - strict] histogram-match reference pool built from TRAIN fragments only (1247 candidates) -- fragment 1 untouched during training.\n\n======================================================================\nSTARTING V5 ENSEMBLE (1 members, single split, no LOFO)\n======================================================================\n\n######################################################################\n# ENSEMBLE MEMBER 1/1  seed=42 encoder=tu-convnext_tiny depth_offset=0\n######################################################################\n  fragment 2 norm stats: mean=118.25 std=59.13\n  fragment 3 norm stats: mean=119.91 std=61.59\n  Positive patches: 5375 | Negative patches: 2708 | Balanced: 8083 (positive ratio=0.665)\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n","output_type":"stream"},{"name":"stdout","text":"  positive fraction=0.15838 bias=-1.670 pos_weight=2.31\n[backbone] AuxDifferentialHead attached: 49 raw differential channels -> 16 hidden -> 1 logit (weight=0.1)\n","output_type":"stream"},{"name":"stderr","text":"Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"d10ac3c461f440e5beda9dcdd53c10ba"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nARCHITECTURE INSPECTION [member0]\n======================================================================\nParameters: 32.20M total\nZero-channel layers: 0 (PASS)\nDry-run (256x256) output: (1, 1, 256, 256) | NaN: PASS | Inf: PASS\nCUDA device: Tesla T4\n  [AdaBN report] BatchNorm2d layers: depth_stem=1 encoder=0 decoder=10 other=1 (total=12)\n======================================================================\n\n  Training member 1 for up to 8 epochs (patch=320, batch=8) ...\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/896222920.py:1340: FutureWarning: `torch.cuda.amp.GradScaler(args...)` is deprecated. Please use `torch.amp.GradScaler('cuda', args...)` instead.\n  scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n/tmp/ipykernel_58/896222920.py:1369: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"    [checkpoint choice] EMA dice=0.0000 raw dice=0.4340 -> using raw\n  [member 1][01/8] time=619.7s train_loss=0.85844 train_dice=0.30051 aux_loss=0.84959 | val_dice=0.43402 val_P=0.4741 val_R=0.4002 val_f0.5=0.4572 [raw]\n    *** new best (val_dice=0.43402, raw weights) ***\n    [checkpoint choice] EMA dice=0.4230 raw dice=0.4795 -> using raw\n  [member 1][02/8] time=453.0s train_loss=0.78868 train_dice=0.45244 aux_loss=0.78652 | val_dice=0.47953 val_P=0.3683 val_R=0.6869 val_f0.5=0.4060 [raw]\n    *** new best (val_dice=0.47953, raw weights) ***\n    [checkpoint choice] EMA dice=0.5352 raw dice=0.4886 -> using ema\n  [member 1][03/8] time=453.0s train_loss=0.76023 train_dice=0.49932 aux_loss=0.77342 | val_dice=0.53515 val_P=0.4723 val_R=0.6174 val_f0.5=0.4956 [ema]\n    *** new best (val_dice=0.53515, ema weights) ***\n    [checkpoint choice] EMA dice=0.5463 raw dice=0.5088 -> using ema\n  [member 1][04/8] time=456.7s train_loss=0.73667 train_dice=0.52894 aux_loss=0.76584 | val_dice=0.54630 val_P=0.4484 val_R=0.6989 val_f0.5=0.4830 [ema]\n    *** new best (val_dice=0.54630, ema weights) ***\n    [checkpoint choice] EMA dice=0.5691 raw dice=0.5993 -> using raw\n  [member 1][05/8] time=458.9s train_loss=0.71218 train_dice=0.57376 aux_loss=0.76319 | val_dice=0.59927 val_P=0.5507 val_R=0.6572 val_f0.5=0.5692 [raw]\n    *** new best (val_dice=0.59927, raw weights) ***\n    [checkpoint choice] EMA dice=0.5893 raw dice=0.6099 -> using raw\n  [member 1][06/8] time=456.4s train_loss=0.68608 train_dice=0.61717 aux_loss=0.76263 | val_dice=0.60990 val_P=0.5863 val_R=0.6355 val_f0.5=0.5955 [raw]\n    *** new best (val_dice=0.60990, raw weights) ***\n    [checkpoint choice] EMA dice=0.6060 raw dice=0.6181 -> using raw\n  [member 1][07/8] time=453.6s train_loss=0.66963 train_dice=0.64885 aux_loss=0.76036 | val_dice=0.61811 val_P=0.5622 val_R=0.6864 val_f0.5=0.5833 [raw]\n    *** new best (val_dice=0.61811, raw weights) ***\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS INK DETECTION - V4 \"DEPTH-AWARE CONVNEXT\"\n# ============================================================\n# Built on V3 (Dice=0.499). Fixes the crash and addresses the deeper\n# architectural critique: the model was treating 26 ordered depth\n# slices as unordered multispectral channels.\n#\n# ------------------------------------------------------------------\n# WHAT ACTUALLY CAUSED THE CRASH (verified empirically, not guessed)\n# ------------------------------------------------------------------\n# The proposed fix of adding `decoder_channels=(256,128,64,32,16)` to\n# smp.UnetPlusPlus was tested directly against this exact setup and\n# it does NOT fix the crash -- I reproduced the identical\n# `weight of size [0, 96, 3, 3]` error with it. The real cause:\n# ConvNeXt's stem is a stride-4 patchify with no separate stride-2\n# stage (unlike ResNet), so smp's generic \"tu-\" timm-encoder wrapper\n# fabricates an EMPTY 0-channel placeholder feature map to keep a\n# uniform 5-stage pyramid API. `smp.Unet`'s decoder already handles\n# that 0-channel stage gracefully (confirmed working); `smp.UnetPlusPlus`'s\n# dense skip-connections do not (confirmed failing, with or without\n# explicit decoder_channels). So V4 uses `smp.Unet`, not UnetPlusPlus,\n# with tu-convnext_tiny. This is a real library limitation, not a\n# parameter you can configure around.\n#\n# ------------------------------------------------------------------\n# DEPTH-AWARE CHANGES (addressing \"26 slices as unordered channels\")\n# ------------------------------------------------------------------\n#  1. DepthFusionStem: computes explicit first/second finite\n#     differences along the physically-ordered depth axis (how\n#     intensity changes through the papyrus), concatenates them with\n#     the raw stack, and learns a compact (16-32 channel) mixed\n#     representation via a 1x1 conv BEFORE the 2D ConvNeXt encoder\n#     ever sees the data -- giving the network an explicit inductive\n#     bias toward depth structure instead of hoping a from-scratch\n#     first-conv discovers it.\n#  2. CLAHE_MODE: \"per_slice\" (V3's old behavior, independent CLAHE\n#     per slice -- can distort inter-slice relationships) vs.\n#     \"global_shared\" (ONE contrast-remapping LUT derived from a\n#     representative slice, applied identically to every slice --\n#     preserves relative depth relationships). Both available for\n#     the ablation you suggested; default is now \"global_shared\".\n#  3. Histogram-matching augmentation now uses ONE shared mapping\n#     across the whole depth stack (`channel_axis=None`) instead of\n#     26 independent per-slice mappings (`channel_axis=2`). I verified\n#     this empirically: independent per-channel matching compressed\n#     slice-to-slice differences unevenly (2.3-3.8 range in a test),\n#     while the shared mapping preserved them far more consistently\n#     (4.3-7.3 range, proportional to the original 7.5-8.6 spacing).\n#  4. Architecture inspection report + zero-channel detector run\n#     BEFORE the optimizer/training loop are created -- this is\n#     exactly the check that would have caught the original crash\n#     immediately instead of after a dummy forward pass deep in setup.\n#  5. AdaBN now prints exactly where its BatchNorm2d layers are found\n#     (depth stem / decoder / encoder) since ConvNeXt itself is\n#     LayerNorm-only and has none -- so you know what it's actually\n#     recalibrating before deciding whether to enable it.\n#  6. V4 baseline: USE_DANN / USE_MIXSTYLE / use_adabn default to\n#     False so the depth-aware architecture change can be evaluated\n#     cleanly on its own first. Turn them back on one at a time\n#     afterward -- CFG has a comment marking each one.\n#  7. One explicit environment setup (current smp + timm, no more\n#     pinning smp==0.2.0 and upgrading from inside build_model()).\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS  (single explicit environment, no version-pin dance)\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom scipy import ndimage\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    patch_size = 320\n    train_stride = 112\n    test_stride = 32\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 10\n    early_stop_patience = 2\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- V4: backbone / architecture -----------------------------------------\n    # VERIFIED WORKING with tu-convnext_tiny: smp.Unet (NOT UnetPlusPlus -- see\n    # module docstring for the empirical reason). Falls back to efficientnet-b4\n    # (a \"real\" smp encoder with no 0-channel-stage quirk) if ConvNeXt/timm are\n    # unavailable in this session for any reason.\n    encoder_name = \"tu-convnext_tiny\"\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"                          # do NOT set to \"unetplusplus\" with tu-* encoders\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- V4: depth-aware fusion stem -----------------------------------------\n    USE_DEPTH_FUSION_STEM = True\n    depth_stem_out_channels = 46\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- LOSS -----------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.35\n    focal_tversky_weight = 0.35\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n\n    # --- THRESHOLD --------------------------------------------------------------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n    threshold = 0.50\n\n    # --- TTA --------------------------------------------------------------------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN (see the \"where are the BN layers\" report at model-build time) --\n    # ConvNeXt is LayerNorm-only; V4 defaults this OFF until the printed report\n    # shows there's something meaningful (decoder / depth-stem BN) for it to do.\n    use_adabn = False\n    adabn_max_patches = 2000\n\n    # --- POSTPROCESS ----------------------------------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    # --- edge-artifact cropping -------------------------------------------------\n    #USE_MASK_EROSION = True\n    USE_MASK_EROSION = False\n    mask_erode_px = 24\n\n    # --- V4: CLAHE mode ---------------------------------------------------------\n    # \"off\"           : no contrast enhancement\n    # \"per_slice\"     : V3's old behavior -- independent CLAHE per slice, can\n    #                   distort inter-slice depth relationships\n    # \"global_shared\" : ONE remapping LUT derived from a representative slice,\n    #                   applied identically to every slice -- preserves relative\n    #                   depth relationships. DEFAULT for V4.\n    CLAHE_MODE = \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization ------------------------------------------------------------\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n\n    # --- domain-randomization augmentation --------------------------------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n\n    USE_SHADOW = True\n    shadow_p = 0.20\n\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # histogram-matching style augmentation -- NOW uses one shared mapping across\n    # depth (channel_axis=None) instead of 26 independent per-slice mappings.\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- EMA --------------------------------------------------------------------\n    USE_EMA = True\n    ema_decay = 0.999\n\n    # --- DANN (OFF for the V4 baseline -- turn on only after the depth-aware\n    #     architecture change is validated on its own; see module docstring) ----\n    USE_DANN = False\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n\n    # --- MixStyle (OFF for the V4 baseline, same reasoning as DANN) -------------\n    USE_MIXSTYLE = True\n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    ckpt_path = os.path.join(out_dir, \"vesuviusnet_v4_best.pth\")\n    viz_dir = os.path.join(out_dir, \"v4_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v4_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\nset_seed(CFG.seed)\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS  (+ mask erosion)\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING (unchanged)\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME  (+ V4: CLAHE modes, per-fragment stats)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None   # built lazily for CLAHE_MODE == \"global_shared\"\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        \"\"\"Derives ONE 256-entry intensity-remapping lookup table from CLAHE applied\n        to a single representative slice, then this exact LUT is applied identically\n        to every depth slice via cv2.LUT. Unlike calling .apply() independently per\n        slice, this guarantees the same monotonic mapping everywhere, so relative\n        inter-slice intensity relationships (the actual depth signal) are preserved.\"\"\"\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid = len(self._slices) // 2\n        ref_slice = np.asarray(self._slices[mid])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(\n            ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def read_patch(self, y, x, size, apply_clahe=None):\n        self._ensure_open()\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[i] = block\n        return out\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        mid = len(self._slices) // 2\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            block = self._slices[mid][y:y + patch_size, x:x + patch_size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                block = self._clahe.apply(block)\n            elif CFG.CLAHE_MODE == \"global_shared\":\n                block = cv2.LUT(block, self._shared_clahe_lut)\n            vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION (fake ink / fiber / shadow)\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\n# ============================================================\n# 8. AUGMENTATION\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    \"\"\"Try constructing a transform with the current albumentations API; if that\n    fails (parameter names changed across versions), fall back to the older API.\n    If both fail, skip the transform rather than crashing the whole pipeline.\"\"\"\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    # NOTE: albumentations' RandomShadow hard-requires 3-channel RGB images in\n    # every version and raises ValueError on multi-channel depth-stack data.\n    # inject_fake_shadow() (applied directly in the dataset, image-only) is used\n    # instead -- see USE_SHADOW below.\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size]\n        img = np.transpose(patch, (1, 2, 0))   # HWD\n\n        # --- histogram-matching style augmentation: ONE shared mapping across\n        # depth (channel_axis=None), not 26 independent per-slice mappings. ---\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img, vol.frag_mean, vol.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN: yields normalized (D,H,W) tensors from fragment 1, NO labels.\n    Uses the SAME basic preprocessing path (CLAHE mode, per-fragment normalization)\n    as the source dataset for consistency -- it intentionally skips the AUGMENTATION\n    pipeline (elastic/dropout/fake-ink/histogram-match), since DANN needs to see the\n    target domain's natural distribution, not an augmented/hallucinated version of it.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL: DEPTH FUSION STEM + VERIFIED-WORKING ConvNeXt+Unet\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"Computes explicit depth-derivative features (first and second finite\n    differences along the physically-ordered depth axis) alongside the raw stack,\n    then learns a compact mixed representation via a 1x1 conv, producing\n    `out_channels` channels to feed into the 2D encoder. This is the \"ink isn't\n    just absolute intensity, it's how intensity changes through the papyrus\"\n    inductive bias, made explicit instead of hoping a from-scratch first-conv\n    layer discovers it purely from data.\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = 0\n        if use_raw:\n            total_in += in_depth\n        if use_grad:\n            total_in += (in_depth - 1)\n        if use_curv:\n            total_in += max(in_depth - 2, 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        # BatchNorm2d kept HERE deliberately (even though the ConvNeXt backbone\n        # itself is all LayerNorm) so AdaBN has a real, meaningful place to\n        # recalibrate target-domain statistics -- see the model-build report.\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.in_depth = in_depth\n\n    def forward(self, x):   # x: (B, D, H, W)\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:, :, :] - x[:, :-1, :, :])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass DepthAwareSegModel(nn.Module):\n    \"\"\"Wraps an smp segmentation model with a DepthFusionStem in front of it.\"\"\"\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):   # x: (B, D, H, W)\n        feat = self.depth_stem(x) if self.depth_stem is not None else x.unsqueeze(1) if x.dim() == 3 else x\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' requested, but only \"\n              f\"'unet' has been verified to handle tu-* timm encoders' 0-channel placeholder \"\n              f\"stage correctly (see module docstring) -- using UnetPlusPlus may crash with \"\n              f\"ConvNeXt/other tu-* encoders.\")\n    return arch_cls(\n        encoder_name=encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=in_channels,\n        classes=1,\n        decoder_channels=CFG.decoder_channels,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\ndef build_model():\n    stem = None\n    seg_in_channels = CFG.in_channels\n    if CFG.USE_DEPTH_FUSION_STEM:\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = CFG.depth_stem_out_channels\n        print(f\"[backbone] DepthFusionStem: {CFG.in_channels} depth slices \"\n              f\"(raw={CFG.depth_stem_use_raw}, grad={CFG.depth_stem_use_grad}, curv={CFG.depth_stem_use_curv}) \"\n              f\"-> {seg_in_channels} learned channels\")\n\n    try:\n        seg_model = _build_seg_backbone(CFG.encoder_name, seg_in_channels)\n        print(f\"[backbone] using {CFG.encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {CFG.encoder_name} unavailable/failed ({e}); \"\n              f\"falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model)\n\n\nclass GradReverse(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x, lambd):\n        ctx.lambd = lambd\n        return x.view_as(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return -ctx.lambd * grad_output, None\n\n\ndef grad_reverse(x, lambd=1.0):\n    return GradReverse.apply(x, lambd)\n\n\nclass DomainClassifier(nn.Module):\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(in_ch, 128), nn.ReLU(inplace=True), nn.Dropout(0.3),\n            nn.Linear(128, 1),\n        )\n\n    def forward(self, feat, lambd):\n        return self.net(grad_reverse(feat, lambd))\n\n\nclass EncoderFeatureCapture:\n    \"\"\"Architecture-agnostic hook capturing the segmentation backbone's encoder\n    output feature list on every forward pass. Attaches to model.seg_model.encoder\n    (the DepthAwareSegModel wrapper's inner smp model), not model.encoder directly,\n    since V4 wraps the smp model with a depth-fusion stem in front of it.\"\"\"\n    def __init__(self, model):\n        self.features = None\n        target = model.seg_model.encoder if hasattr(model, \"seg_model\") else model.encoder\n        self.handle = target.register_forward_hook(self._hook)\n\n    def _hook(self, module, inp, out):\n        self.features = out\n\n    def remove(self):\n        self.handle.remove()\n\n\ndef mixstyle_batch(imgs, p=CFG.mixstyle_p, alpha=CFG.mixstyle_alpha):\n    if torch.rand(1).item() > p:\n        return imgs\n    B = imgs.size(0)\n    if B < 2:\n        return imgs\n    mu = imgs.mean(dim=[2, 3], keepdim=True)\n    var = imgs.var(dim=[2, 3], keepdim=True)\n    sig = (var + 1e-6).sqrt()\n    x_norm = (imgs - mu) / sig\n    perm = torch.randperm(B, device=imgs.device)\n    mu2, sig2 = mu[perm], sig[perm]\n    lam = torch.distributions.Beta(alpha, alpha).sample((B, 1, 1, 1)).to(imgs.device)\n    mu_mix = lam * mu + (1 - lam) * mu2\n    sig_mix = lam * sig + (1 - lam) * sig2\n    return x_norm * sig_mix + mu_mix\n\n\nclass EMAModel:\n    def __init__(self, model, decay=CFG.ema_decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. V4: PRE-TRAINING ARCHITECTURE INSPECTION\n#     (this is exactly what would have caught the original crash\n#     immediately, before the optimizer/training loop existed)\n# ============================================================\n\ndef inspect_model_channels(model):\n    \"\"\"Scans every Conv/Linear layer for zero or negative in/out channels -- the\n    exact failure mode behind the original UnetPlusPlus/ConvNeXt crash.\"\"\"\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\ndef report_batchnorm_locations(model):\n    \"\"\"Prints where BatchNorm2d layers actually live in the model -- ConvNeXt's\n    own encoder is LayerNorm-only, so AdaBN (which recalibrates BatchNorm running\n    stats) has nothing to do there. It CAN still do something meaningful in the\n    depth-fusion stem and/or the smp decoder, if those have BatchNorm2d layers.\"\"\"\n    counts = {\"depth_stem\": 0, \"encoder\": 0, \"decoder\": 0, \"other\": 0}\n    for name, module in model.named_modules():\n        if isinstance(module, nn.BatchNorm2d):\n            if name.startswith(\"depth_stem\"):\n                counts[\"depth_stem\"] += 1\n            elif \".encoder.\" in f\".{name}.\" or name.endswith(\".encoder\"):\n                counts[\"encoder\"] += 1\n            elif \".decoder.\" in f\".{name}.\" or name.endswith(\".decoder\"):\n                counts[\"decoder\"] += 1\n            else:\n                counts[\"other\"] += 1\n    total = sum(counts.values())\n    print(f\"[AdaBN report] BatchNorm2d layers found: depth_stem={counts['depth_stem']} \"\n          f\"encoder={counts['encoder']} decoder={counts['decoder']} other={counts['other']} \"\n          f\"(total={total})\")\n    if counts[\"encoder\"] == 0 and total > 0:\n        print(\"  -> the ConvNeXt encoder itself has none (it's LayerNorm-only, as expected). \"\n              \"AdaBN would only recalibrate the depth-stem/decoder BN layers listed above.\")\n    if total == 0:\n        print(\"  -> NO BatchNorm2d layers anywhere in this model. Enabling use_adabn would \"\n              \"currently be a no-op. (Kept OFF by default in V4 CFG for exactly this reason.)\")\n    return counts\n\n\n@torch.no_grad()\ndef run_architecture_report(model):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"V4 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\")\n    print(\"=\" * 70)\n\n    n_params = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Input:  depth slices={CFG.in_channels} | patch={CFG.patch_size}x{CFG.patch_size} | \"\n          f\"input tensor=[1, {CFG.in_channels}, {CFG.patch_size}, {CFG.patch_size}]\")\n    print(f\"Encoder: backbone={CFG.encoder_name} | pretrained={CFG.encoder_weights} | \"\n          f\"architecture={CFG.architecture}\")\n    if CFG.USE_DEPTH_FUSION_STEM:\n        print(f\"Depth stem: {CFG.in_channels} -> {model.depth_stem.out_channels} channels \"\n              f\"(raw={model.depth_stem.use_raw}, grad={model.depth_stem.use_grad}, \"\n              f\"curv={model.depth_stem.use_curv})\")\n    print(f\"Decoder channels: {CFG.decoder_channels}\")\n    print(f\"Parameters: {n_params/1e6:.2f}M total | {n_trainable/1e6:.2f}M trainable\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"\\nZERO/INVALID CHANNEL LAYERS FOUND:\")\n        for item in bad_layers:\n            print(\"  \", item)\n        raise RuntimeError(\n            \"Model contains invalid zero-channel layers -- fix architecture before training. \"\n            \"(If this happened with architecture='unetplusplus' and a tu-* encoder, switch to \"\n            \"architecture='unet' -- see module docstring for why.)\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    model.eval()\n    small_size = min(CFG.patch_size, 256)   # keep the dry-run cheap regardless of real patch_size\n    dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n    capture = EncoderFeatureCapture(model)\n    try:\n        out = model(dummy)\n        print(f\"\\nEncoder feature stages (dry run at {small_size}x{small_size} for speed):\")\n        for i, f in enumerate(capture.features):\n            print(f\"  stage {i}: {tuple(f.shape)}\")\n        print(f\"\\nOutput logits shape: {tuple(out.shape)}\")\n        has_nan = torch.isnan(out).any().item()\n        has_inf = torch.isinf(out).any().item()\n        print(f\"NaN check: {'FAIL' if has_nan else 'PASS'} | Inf check: {'FAIL' if has_inf else 'PASS'}\")\n        if has_nan or has_inf:\n            raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass -- fix before training.\")\n    finally:\n        capture.remove()\n    model.train()\n\n    if torch.cuda.is_available():\n        print(f\"\\nCUDA device: {torch.cuda.get_device_name(0)} | \"\n              f\"capability: {torch.cuda.get_device_capability(0)}\")\n    print(\"Forward pass: PASS\")\n\n    report_batchnorm_locations(model)\n    print(\"=\" * 70 + \"\\n\")\n\n\n# ============================================================\n# 12. LOSSES (unchanged)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\nclass V2ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        return CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n\n\n# ============================================================\n# 13. METRICS\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 14. BUILD TRAIN DATA\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING TRAIN DATA\")\nprint(\"=\" * 70)\n\ntrain_volumes = {}\ntrain_labels_full = {}\ntrain_masks_full = {}\ntrain_samples_raw = []\nval_samples = []\n\nfor fid in CFG.train_frags:\n    print(f\"\\nProcessing fragment {fid} ...\")\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        vol.compute_fragment_stats(mask, CFG.patch_size)\n        print(f\"  fragment {fid} normalization stats: mean={vol.frag_mean:.2f} std={vol.frag_std:.2f}\")\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {len(coords)} candidate patches (post mask-erosion)\")\n\n    tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n    print(f\"  spatial train={len(tr_coords)} validation={len(va_coords)}\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    train_masks_full[fid] = mask\n    train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n    val_samples.extend([(fid, y, x) for y, x in va_coords])\n    del mask\n    cleanup_memory()\n\nprint(\"\\nRaw train samples:\", len(train_samples_raw))\nprint(\"Validation samples:\", len(val_samples))\n\ntrain_samples = balance_positive_patches(\n    train_samples_raw, train_labels_full, CFG.patch_size,\n    positive_threshold=CFG.positive_patch_fraction,\n    target_positive_ratio=CFG.target_positive_patch_ratio,\n    max_positive_repeat=CFG.max_positive_repeat,\n)\n\n\n# ============================================================\n# 15. TEST-FRAGMENT MASK/VOLUME LOADED EARLY (unlabeled use only)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\")\nprint(\"=\" * 70)\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\nif CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    print(f\"  fragment {CFG.test_frag} normalization stats: \"\n          f\"mean={test_vol.frag_mean:.2f} std={test_vol.frag_std:.2f}\")\n\ntest_coords_for_unlabeled_use = generate_grid_coords(\n    test_mask, CFG.patch_size, CFG.test_stride, CFG.min_tissue_frac_train)\nprint(f\"Fragment {CFG.test_frag}: {len(test_coords_for_unlabeled_use)} unlabeled candidate patches\")\n\nhist_match_pool = None\nif CFG.USE_HIST_MATCH_AUG:\n    pool_coords = random.sample(test_coords_for_unlabeled_use,\n                                 min(CFG.hist_match_pool_size, len(test_coords_for_unlabeled_use)))\n    hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n    print(f\"Built histogram-matching reference pool: {len(hist_match_pool)} patches \"\n          f\"({sum(p.nbytes for p in hist_match_pool)/1e6:.1f} MB)\")\n\ndann_loader = None\nif CFG.USE_DANN:\n    dann_ds = UnlabeledPatchDataset(test_vol, test_coords_for_unlabeled_use, CFG.patch_size)\n    dann_loader = DataLoader(dann_ds, batch_size=CFG.dann_target_batch_size, shuffle=True,\n                              num_workers=max(1, CFG.num_workers - 1), pin_memory=(CFG.device == \"cuda\"),\n                              drop_last=True, persistent_workers=True)\n\n    def infinite_dann_loader():\n        while True:\n            for batch in dann_loader:\n                yield batch\n    dann_iter = infinite_dann_loader()\n\n\n# ============================================================\n# 16. DATASETS / LOADERS\n# ============================================================\n\ntrain_transform = build_train_transform()\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                            transform=train_transform, jitter=CFG.train_jitter,\n                            hist_match_pool=hist_match_pool, train_mode=True)\n\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                          transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0)\n\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         drop_last=False, persistent_workers=CFG.num_workers > 0)\n\n\n# ============================================================\n# 17. MODEL / ARCHITECTURE REPORT / LOSS / OPTIMIZER\n# ============================================================\n\nprint(\"\\nBuilding V4 depth-aware model ...\")\nmodel = build_model().to(CFG.device)\n\nrun_architecture_report(model)   # <-- catches zero-channel / NaN / shape problems HERE\n\nprint(\"\\nEstimating positive-pixel fraction ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\nprint(f\"Estimated positive fraction: {pos_frac:.6f}\")\n\nwith torch.no_grad():\n    bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n    try:\n        model.seg_model.segmentation_head[0].bias.fill_(bias_val)\n    except Exception as e:\n        print(f\"  (could not set output bias directly: {e})\")\nprint(f\"Output bias initialized to {bias_val:.4f}\")\n\nraw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\npos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\npos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\nprint(f\"BCE positive weight: {pos_weight_val:.3f}\")\n\ncriterion = V2ComboLoss(pos_weight=pos_weight)\n\nencoder_params, decoder_params, stem_params = [], [], []\nfor name, param in model.named_parameters():\n    if not param.requires_grad:\n        continue\n    if name.startswith(\"depth_stem.\"):\n        stem_params.append(param)\n    elif name.startswith(\"seg_model.encoder.\"):\n        encoder_params.append(param)\n    else:\n        decoder_params.append(param)\nprint(f\"Depth-stem parameters: {len(stem_params)} | Encoder parameters: {len(encoder_params)} | \"\n      f\"Decoder/head parameters: {len(decoder_params)} tensors\")\n\nparam_groups = [\n    {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n    {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n]\nif stem_params:\n    param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\ndomain_classifier = None\nfeature_capture = None\nif CFG.USE_DANN:\n    feature_capture = EncoderFeatureCapture(model)\n    with torch.no_grad():\n        dummy = torch.zeros(1, CFG.in_channels, CFG.patch_size, CFG.patch_size, device=CFG.device)\n        model.eval()\n        _ = model(dummy)\n        deepest_ch = feature_capture.features[-1].shape[1]\n        model.train()\n    domain_classifier = DomainClassifier(deepest_ch).to(CFG.device)\n    param_groups.append({\"params\": domain_classifier.parameters(), \"lr\": CFG.dann_lr})\n    print(f\"[DANN] domain classifier attached on {deepest_ch}-channel bottleneck features\")\n    del dummy\n    cleanup_memory()\n\noptimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nema = EMAModel(model, decay=CFG.ema_decay) if CFG.USE_EMA else None\n\n\n# ============================================================\n# 18. TRAIN / VALIDATION EPOCH\n# ============================================================\n\n_global_step = 0\n_total_steps = CFG.epochs * max(1, len(train_loader) // CFG.accumulation_steps)\n\n\ndef dann_lambda_schedule():\n    progress = min(_global_step / max(_total_steps, 1), 1.0)\n    return CFG.dann_lambda_max * (2.0 / (1.0 + math.exp(-10.0 * progress)) - 1.0)\n\n\ndef run_epoch(loader, train_mode=True, threshold=0.5):\n    global _global_step\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_dann_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    for batch_idx, (imgs, masks) in enumerate(loader):\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        if train_mode and CFG.USE_MIXSTYLE:\n            imgs = mixstyle_batch(imgs)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n                dann_loss_val = 0.0\n\n                if train_mode and CFG.USE_DANN:\n                    src_feat = feature_capture.features[-1]\n                    lambd = dann_lambda_schedule()\n                    src_domain_logits = domain_classifier(src_feat, lambd)\n                    src_domain_target = torch.zeros_like(src_domain_logits)\n\n                    tgt_imgs = next(dann_iter).to(CFG.device, non_blocking=True)\n                    _ = model(tgt_imgs)\n                    tgt_feat = feature_capture.features[-1]\n                    tgt_domain_logits = domain_classifier(tgt_feat, lambd)\n                    tgt_domain_target = torch.ones_like(tgt_domain_logits)\n\n                    dann_loss = 0.5 * (\n                        F.binary_cross_entropy_with_logits(src_domain_logits, src_domain_target) +\n                        F.binary_cross_entropy_with_logits(tgt_domain_logits, tgt_domain_target)\n                    )\n                    dann_loss_val = dann_loss.item()\n                    loss = loss + dann_loss\n\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    params_to_clip = list(model.parameters())\n                    if domain_classifier is not None:\n                        params_to_clip += list(domain_classifier.parameters())\n                    torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    _global_step += 1\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_dann_loss += dann_loss_val\n\n        del imgs, masks, logits, probs\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        params_to_clip = list(model.parameters())\n        if domain_classifier is not None:\n            params_to_clip += list(domain_classifier.parameters())\n        torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_dann_loss = total_dann_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_dann_loss\n\n\n# ============================================================\n# 19. TRAINING LOOP\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING V4 TRAINING\")\nprint(\"=\" * 70)\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"train_dice\": [], \"val_dice\": [],\n           \"val_iou\": [], \"val_precision\": [], \"val_recall\": [], \"dann_loss\": []}\n\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_metrics, dann_loss_avg = run_epoch(train_loader, train_mode=True, threshold=0.50)\n\n    if ema is not None:\n        backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        model.load_state_dict(ema.state_dict())\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n        model.load_state_dict(backup)\n        del backup\n    else:\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"train_dice\"].append(train_metrics[\"dice\"])\n    history[\"val_dice\"].append(val_metrics[\"dice\"])\n    history[\"val_iou\"].append(val_metrics[\"iou\"])\n    history[\"val_precision\"].append(val_metrics[\"precision\"])\n    history[\"val_recall\"].append(val_metrics[\"recall\"])\n    history[\"dann_loss\"].append(dann_loss_avg)\n\n    print(f\"\\n[{epoch:02d}/{CFG.epochs}] time={time.time()-t0:.1f}s\")\n    print(f\"train_loss={train_loss:.5f} train_dice={train_metrics['dice']:.5f} \"\n          f\"dann_loss={dann_loss_avg:.5f}\")\n    print(f\"val_loss={val_loss:.5f} val_dice={val_metrics['dice']:.5f} val_iou={val_metrics['iou']:.5f}\")\n    print(f\"precision={val_metrics['precision']:.5f} recall={val_metrics['recall']:.5f}\")\n    print(f\"encoder_lr={optimizer.param_groups[0]['lr']:.7f} decoder_lr={optimizer.param_groups[1]['lr']:.7f}\")\n\n    if val_metrics[\"dice\"] > best_val_dice:\n        best_val_dice = val_metrics[\"dice\"]\n        epochs_no_improve = 0\n        save_state = ema.state_dict() if ema is not None else model.state_dict()\n        checkpoint = {\"model\": save_state, \"cfg\": cfg_to_dict(CFG), \"best_val_dice\": best_val_dice,\n                      \"history\": history, \"pos_frac\": pos_frac, \"pos_weight\": pos_weight_val,\n                      \"used_ema\": CFG.USE_EMA}\n        torch.save(checkpoint, CFG.ckpt_path)\n        print(f\"*** NEW BEST CHECKPOINT val_dice={best_val_dice:.5f} \"\n              f\"({'EMA' if CFG.USE_EMA else 'raw'} weights) ***\")\n    else:\n        epochs_no_improve += 1\n        print(f\"No improvement: {epochs_no_improve}/{CFG.early_stop_patience}\")\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(\"Early stopping.\")\n            break\n\n    cleanup_memory()\n\nprint(\"\\nBest validation Dice:\", best_val_dice)\n\nif feature_capture is not None:\n    feature_capture.remove()\n\n\n# ============================================================\n# 20. LOAD BEST MODEL\n# ============================================================\n\ncheckpoint = torch.load(CFG.ckpt_path, map_location=CFG.device)\nmodel.load_state_dict(checkpoint[\"model\"])\nprint(f\"\\nBest checkpoint loaded ({'EMA' if checkpoint.get('used_ema') else 'raw'} weights).\")\n\n\n# ============================================================\n# 21. FINE DICE THRESHOLD SEARCH\n# ============================================================\n\n@torch.no_grad()\ndef find_best_dice_threshold(model, loader):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice_scores = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    best_idx = int(np.argmax(dice_scores))\n    return float(thresholds[best_idx]), float(dice_scores[best_idx]), thresholds, dice_scores\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINE DICE THRESHOLD SEARCH\")\nprint(\"=\" * 70)\n\nbest_threshold, threshold_dice, threshold_grid, threshold_scores = find_best_dice_threshold(model, val_loader)\nprint(f\"BEST VALIDATION THRESHOLD = {best_threshold:.2f}\")\nprint(f\"DICE AT BEST THRESHOLD = {threshold_dice:.5f}\")\n\n_, final_val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal validation metrics at optimized Dice threshold:\", final_val_metrics)\n\n\n# ============================================================\n# 22. INFERENCE HELPERS\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference_v2(model, vol, mask, patch_size, stride, device, batch_size,\n                                 pre_transform=None):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw_u8 = vol.read_patch(y, x, patch_size)\n        if pre_transform is not None:\n            raw_u8 = pre_transform(raw_u8)\n        raw = raw_u8.astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm_v2(model, vol, mask, patch_size, stride, device, max_patches, batch_size):\n    print(\"\\nStarting AdaBN...\")\n    n_reset = 0\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n            n_reset += 1\n    print(f\"Reset {n_reset} BatchNorm2d layers (see the architecture report above for where they live).\")\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"AdaBN patches: {len(coords)}\")\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        batch = coords[start:start + batch_size]\n        imgs = []\n        for y, x in batch:\n            raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n            raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n            imgs.append(raw)\n        inp = torch.from_numpy(np.stack(imgs)).to(device, non_blocking=True)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp)\n        del inp, imgs\n    model.eval()\n    cleanup_memory()\n    print(\"AdaBN finished.\")\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess_v2(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 23. LOAD TEST LABELS (local diagnostics only, loaded late)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST LABELS (diagnostics only)\")\nprint(\"=\" * 70)\n\ntest_labels = load_ink_labels(test_dir)\nif test_labels is not None:\n    print(\"Test GT found: local diagnostic evaluation enabled.\")\n    gt_test = (test_labels * test_mask).astype(np.float32)\nelse:\n    print(\"No test GT found: running competition-style inference.\")\n    gt_test = None\n\n\n# ============================================================\n# 24. INFERENCE A: BASELINE + TTA\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE A: BASELINE + TTA\")\nprint(\"=\" * 70)\n\ntest_prob_baseline = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n\ntest_metrics_baseline = None\nif test_labels is not None:\n    test_metrics_baseline = evaluate_probability_map(test_prob_baseline, gt_test, best_threshold)\n    print(\"\\nBASELINE TEST METRICS:\", test_metrics_baseline)\n\n\n# ============================================================\n# 25. INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\")\nprint(\"=\" * 70)\n\ntrain_hist_pool = []\nfor fid in CFG.train_frags:\n    v = train_volumes[fid]\n    m = train_masks_full[fid]\n    coords = generate_grid_coords(m, CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n    if coords:\n        for (y, x) in random.sample(coords, min(10, len(coords))):\n            train_hist_pool.append(v.read_patch(y, x, CFG.patch_size))\nprint(f\"Built train-domain reference pool for diagnostic: {len(train_hist_pool)} patches\")\n\n\ndef histogram_match_to_train_domain(raw_u8_dhw):\n    if not train_hist_pool:\n        return raw_u8_dhw\n    ref = random.choice(train_hist_pool)\n    img_hwd = np.transpose(raw_u8_dhw, (1, 2, 0))\n    ref_hwd = np.transpose(ref, (1, 2, 0))\n    try:\n        matched = match_histograms(img_hwd, ref_hwd, channel_axis=None).astype(np.uint8)\n        return np.transpose(matched, (2, 0, 1))\n    except Exception:\n        return raw_u8_dhw\n\n\ntest_prob_histmatch = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch,\n    pre_transform=histogram_match_to_train_domain)\n\ntest_metrics_histmatch = None\nif test_labels is not None:\n    test_metrics_histmatch = evaluate_probability_map(test_prob_histmatch, gt_test, best_threshold)\n    print(\"\\nHISTOGRAM-MATCHING DIAGNOSTIC TEST METRICS:\", test_metrics_histmatch)\n    if test_metrics_baseline is not None:\n        delta = test_metrics_histmatch[\"dice\"] - test_metrics_baseline[\"dice\"]\n        print(f\"\\n>>> Histogram matching alone changed local test Dice by {delta:+.4f} \"\n              f\"(baseline {test_metrics_baseline['dice']:.4f} -> {test_metrics_histmatch['dice']:.4f})\")\n\n# --- save the histogram-matched probability map + thresholded prediction, same\n# treatment as the other prediction variants ---\ntest_pred_histmatch_bin = postprocess_v2(test_prob_histmatch, best_threshold)\nhistmatch_prob_path = os.path.join(CFG.out_dir, \"fragment1_probability_histmatch_v4.npy\")\nnp.save(histmatch_prob_path, test_prob_histmatch)\nprint(\"Saved histogram-matched probability map:\", histmatch_prob_path)\nhistmatch_pred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_histmatch_v4.png\")\ncv2.imwrite(histmatch_pred_path, (test_pred_histmatch_bin * 255).astype(np.uint8))\nprint(\"Saved histogram-matched prediction:\", histmatch_pred_path)\n\n\n# ============================================================\n# 26. INFERENCE C: ADABN + TTA (only if enabled -- see architecture report)\n# ============================================================\n\ntest_prob_adabn = None\ntest_metrics_adabn = None\n\nif CFG.use_adabn:\n    print(\"\\n\" + \"=\" * 70)\n    print(\"INFERENCE C: ADABN + TTA\")\n    print(\"=\" * 70)\n    model = recalibrate_batchnorm_v2(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                      CFG.device, CFG.adabn_max_patches, CFG.infer_batch)\n    test_prob_adabn = sliding_window_inference_v2(\n        model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n    if test_labels is not None:\n        test_metrics_adabn = evaluate_probability_map(test_prob_adabn, gt_test, best_threshold)\n        print(\"\\nADABN TEST METRICS:\", test_metrics_adabn)\nelse:\n    print(\"\\n(Skipping AdaBN inference -- CFG.use_adabn=False. See the architecture report's \"\n          \"BatchNorm2d location summary above for why/whether it would help.)\")\n\n\n# ============================================================\n# 27. CHOOSE FINAL PROBABILITY MAP\n# ============================================================\n\ntest_prob = test_prob_adabn if (CFG.use_adabn and test_prob_adabn is not None) else test_prob_baseline\n\notsu_threshold = compute_otsu_threshold(test_prob, test_mask, fallback=best_threshold)\nprint(\"\\nValidation-tuned threshold:\", best_threshold)\nprint(\"Unsupervised Otsu threshold:\", otsu_threshold)\n\nfinal_threshold = best_threshold\ntest_pred_bin = postprocess_v2(test_prob, final_threshold)\n\ntest_metrics_raw = test_metrics_otsu = test_metrics_post = None\nif test_labels is not None:\n    test_metrics_raw = evaluate_probability_map(test_prob, gt_test, final_threshold)\n    test_metrics_otsu = evaluate_probability_map(test_prob, gt_test, otsu_threshold)\n    post_preds = test_pred_bin.astype(np.float32)\n    tp = (post_preds * gt_test).sum()\n    fp = (post_preds * (1.0 - gt_test)).sum()\n    fn = ((1.0 - post_preds) * gt_test).sum()\n    test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LOCAL TEST DIAGNOSTICS SUMMARY\")\n    print(\"=\" * 70)\n    print(\"A) Baseline + val threshold:            \", test_metrics_baseline)\n    print(\"B) Histogram-matched (no retrain) + val threshold:\", test_metrics_histmatch)\n    print(\"C) AdaBN + val threshold:                \", test_metrics_adabn)\n    print(\"Final (chosen) raw + val threshold:      \", test_metrics_raw)\n    print(\"Final (chosen) raw + Otsu threshold:      \", test_metrics_otsu)\n    print(\"Final (chosen) postprocessed + val threshold:\", test_metrics_post)\n\n\n# ============================================================\n# 28. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_v4.npy\")\nnp.save(prob_path, test_prob)\nprint(\"\\nSaved probability map:\", prob_path)\n\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_v4.png\")\ncv2.imwrite(pred_path, (test_pred_bin * 255).astype(np.uint8))\nprint(\"Saved prediction:\", pred_path)\n\nmetrics_summary = {\n    \"best_validation_dice_at_0.50\": best_val_dice,\n    \"best_validation_threshold\": best_threshold,\n    \"validation_dice_at_best_threshold\": threshold_dice,\n    \"otsu_threshold\": otsu_threshold,\n    \"final_threshold\": final_threshold,\n    \"final_validation_metrics\": final_val_metrics,\n    \"test_baseline\": test_metrics_baseline,\n    \"test_histogram_matched_diagnostic_no_retrain\": test_metrics_histmatch,\n    \"test_adabn\": test_metrics_adabn,\n    \"test_raw\": test_metrics_raw,\n    \"test_otsu\": test_metrics_otsu,\n    \"test_postprocessed\": test_metrics_post,\n    \"config\": cfg_to_dict(CFG),\n    \"history\": history,\n}\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump(metrics_summary, f, indent=2)\nprint(\"Saved metrics:\", CFG.metrics_path)\n\n\n# ============================================================\n# 29. VISUALIZATION  (+ prediction-vs-ground-truth overlay on every comparison)\n# ============================================================\n\ndef make_confusion_overlay(input_gray_u8, gt_binary, pred_binary, alpha=0.55):\n    \"\"\"RGB error-map overlay on top of the grayscale input:\n       green  = true positive  (prediction AND ground truth)\n       red    = false positive (prediction only)\n       yellow = false negative (ground truth only)\n    Makes prediction-vs-ground-truth mismatches immediately visible, instead of\n    having to mentally compare two separate side-by-side panels.\"\"\"\n    base = np.stack([input_gray_u8] * 3, axis=-1).astype(np.float32)\n    overlay = base.copy()\n    gt_b = gt_binary > 0\n    pred_b = pred_binary > 0\n    tp = pred_b & gt_b\n    fp = pred_b & ~gt_b\n    fn = ~pred_b & gt_b\n    color_tp = np.array([0, 255, 0], dtype=np.float32)\n    color_fp = np.array([255, 0, 0], dtype=np.float32)\n    color_fn = np.array([255, 255, 0], dtype=np.float32)\n    overlay[tp] = (1 - alpha) * base[tp] + alpha * color_tp\n    overlay[fp] = (1 - alpha) * base[fp] + alpha * color_fp\n    overlay[fn] = (1 - alpha) * base[fn] + alpha * color_fn\n    return np.clip(overlay, 0, 255).astype(np.uint8)\n\n\ndef make_prediction_only_overlay(input_gray_u8, pred_binary, alpha=0.5, color=(255, 0, 0)):\n    \"\"\"Single-color prediction overlay for when no ground truth is available (real\n    competition-style inference on an unlabeled fragment).\"\"\"\n    base = np.stack([input_gray_u8] * 3, axis=-1).astype(np.float32)\n    overlay = base.copy()\n    mask = pred_binary > 0\n    overlay[mask] = (1 - alpha) * base[mask] + alpha * np.array(color, dtype=np.float32)\n    return np.clip(overlay, 0, 255).astype(np.uint8)\n\n\ndef save_overview(prob_map, pred_bin, tag, threshold_label):\n    \"\"\"Saves a full-fragment overview for ONE prediction variant (baseline /\n    histogram-matched / AdaBN / final), including a confusion-map overlay panel\n    against ground truth when available.\"\"\"\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    small_u8 = cv2.normalize(small, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n\n    pred_small = cv2.resize((pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    pred_small_bin = (pred_small > 127).astype(np.uint8)\n    prob_small = cv2.resize((prob_map * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_small = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1],\n                               interpolation=cv2.INTER_NEAREST)\n        gt_small_bin = (gt_small > 127).astype(np.uint8)\n        overlay = make_confusion_overlay(small_u8, gt_small_bin, pred_small_bin)\n\n        fig, axes = plt.subplots(1, 5, figsize=(27, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n        axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability ({threshold_label})\")\n        axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(f\"Prediction [{tag}]\")\n        axes[4].imshow(overlay); axes[4].set_title(\"Overlay: green=TP red=FP yellow=FN\")\n    else:\n        overlay = make_prediction_only_overlay(small_u8, pred_small_bin)\n        fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(prob_small, cmap=\"gray\"); axes[1].set_title(f\"Probability ({threshold_label})\")\n        axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(f\"Prediction [{tag}]\")\n        axes[3].imshow(overlay); axes[3].set_title(\"Prediction overlay (no GT available)\")\n\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    overview_path = os.path.join(CFG.viz_dir, f\"fragment1_v4_overview_{tag}.png\")\n    plt.savefig(overview_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved {tag} overview:\", overview_path)\n\n\n# --- one overview (with overlay) per prediction variant --------------------------\ntest_pred_baseline_bin = postprocess_v2(test_prob_baseline, best_threshold)\nsave_overview(test_prob_baseline, test_pred_baseline_bin, tag=\"baseline\",\n              threshold_label=f\"val_thr={best_threshold:.2f}\")\n\nsave_overview(test_prob_histmatch, test_pred_histmatch_bin, tag=\"histmatch\",\n              threshold_label=f\"val_thr={best_threshold:.2f}\")\n\nif test_prob_adabn is not None:\n    test_pred_adabn_bin = postprocess_v2(test_prob_adabn, best_threshold)\n    save_overview(test_prob_adabn, test_pred_adabn_bin, tag=\"adabn\",\n                  threshold_label=f\"val_thr={best_threshold:.2f}\")\n\nsave_overview(test_prob, test_pred_bin, tag=\"final\",\n              threshold_label=f\"final_thr={final_threshold:.2f}\")\n\n\ndef save_patch_comparisons(n=6):\n    coords = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(coords)\n    coords = coords[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(coords):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        input_slice_u8 = cv2.normalize(input_slice, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n        prob_patch = test_prob[y:y + size, x:x + size]\n\n        if test_labels is not None:\n            gt_patch = test_labels[y:y + size, x:x + size]\n            overlay_patch = make_confusion_overlay(input_slice_u8, gt_patch, pred_patch)\n            fig, axes = plt.subplots(1, 5, figsize=(20, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n            axes[2].imshow(prob_patch, cmap=\"gray\"); axes[2].set_title(\"Probability\")\n            axes[3].imshow(pred_patch, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n            axes[4].imshow(overlay_patch); axes[4].set_title(\"Overlay: green=TP red=FP yellow=FN\")\n        else:\n            overlay_patch = make_prediction_only_overlay(input_slice_u8, pred_patch)\n            fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(prob_patch, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n            axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n            axes[3].imshow(overlay_patch); axes[3].set_title(\"Prediction overlay (no GT)\")\n\n        for ax in axes:\n            ax.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n        plt.close(fig)\n    print(f\"Saved {len(coords)} patch comparisons.\")\n\n\nsave_patch_comparisons(n=6)\n\n\ndef save_training_curves():\n    epochs_axis = np.arange(1, len(history[\"train_loss\"]) + 1)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_loss\"], label=\"Train Loss\")\n    plt.plot(epochs_axis, history[\"val_loss\"], label=\"Val Loss\")\n    if CFG.USE_DANN:\n        plt.plot(epochs_axis, history[\"dann_loss\"], label=\"DANN domain loss\", linestyle=\"--\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Loss\"); plt.title(\"V4 Training Loss\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_loss.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_dice\"], label=\"Train Dice\")\n    plt.plot(epochs_axis, history[\"val_dice\"], label=\"Val Dice (EMA)\" if CFG.USE_EMA else \"Val Dice\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Dice\"); plt.title(\"V4 Dice\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_dice.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(threshold_grid, threshold_scores)\n    plt.axvline(best_threshold, linestyle=\"--\", label=f\"best={best_threshold:.2f}\")\n    plt.xlabel(\"Threshold\"); plt.ylabel(\"Validation Dice\"); plt.title(\"Dice Threshold Search\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"threshold_search.png\"), dpi=150)\n    plt.close(fig)\n\n    print(\"Saved training curves.\")\n\n\nsave_training_curves()\n\n\n# ============================================================\n# 30. CLEANUP\n# ============================================================\n\ntest_vol.close()\nfor v in train_volumes.values():\n    v.close()\ncleanup_memory()\n\n\n# ============================================================\n# 31. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V4 COMPLETE\")\nprint(\"=\" * 70)\nprint(f\"Best validation Dice @ 0.50: {best_val_dice:.5f}\")\nprint(f\"Best Dice threshold: {best_threshold:.2f}\")\nprint(f\"Validation Dice @ optimized threshold: {threshold_dice:.5f}\")\nprint(f\"Backbone: {CFG.encoder_name} / architecture={CFG.architecture} / \"\n      f\"depth_stem={CFG.USE_DEPTH_FUSION_STEM}\")\nprint(f\"EMA: {CFG.USE_EMA} | DANN: {CFG.USE_DANN} | MixStyle: {CFG.USE_MIXSTYLE} | AdaBN: {CFG.use_adabn}\")\nprint(f\"CLAHE mode: {CFG.CLAHE_MODE} | Normalization: {CFG.NORMALIZATION_MODE} | \"\n      f\"Mask erosion: {CFG.mask_erode_px}px\")\nif test_metrics_baseline is not None:\n    print(f\"\\nLocal test Dice -- baseline: {test_metrics_baseline['dice']:.5f}\")\nif test_metrics_histmatch is not None:\n    print(f\"Local test Dice -- histogram-matched (diagnostic, no retrain): \"\n          f\"{test_metrics_histmatch['dice']:.5f}\")\nif test_metrics_adabn is not None:\n    print(f\"Local test Dice -- AdaBN: {test_metrics_adabn['dice']:.5f}\")\nif test_metrics_post is not None:\n    print(f\"Local test Dice -- final postprocessed: {test_metrics_post['dice']:.5f}\")\nprint(f\"\\nBest checkpoint: {CFG.ckpt_path}\")\nprint(f\"Probability map: {prob_path}\")\nprint(f\"Prediction: {pred_path}\")\nprint(f\"Metrics: {CFG.metrics_path}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")\nprint(\"\\n=== V4 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-27T06:23:44.198565Z","iopub.execute_input":"2026-09-27T06:23:44.198988Z","iopub.status.idle":"2026-09-27T09:01:36.096594Z","shell.execute_reply.started":"2026-09-27T06:23:44.198957Z","shell.execute_reply":"2026-09-27T09:01:36.095693Z"}},"outputs":[{"name":"stdout","text":"   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 3.1 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.7/2.7 MB 41.3 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.30\n======================================================================\nBUILDING TRAIN DATA\n======================================================================\n\nProcessing fragment 2 ...\n  fragment 2 normalization stats: mean=114.66 std=62.76\nFragment 2: 8131 candidate patches (post mask-erosion)\n  spatial train=6081 validation=1673\n\nProcessing fragment 3 ...\n  fragment 3 normalization stats: mean=115.07 std=64.02\nFragment 3: 2157 candidate patches (post mask-erosion)\n  spatial train=1693 validation=221\n\nRaw train samples: 7774\nValidation samples: 1894\nPositive patches: 5252 | Negative patches: 2522\nBalanced dataset: 7774 | positive ratio=0.676\n\n======================================================================\nLOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\n======================================================================\n  fragment 1 normalization stats: mean=110.75 std=69.02\nFragment 1: 30674 unlabeled candidate patches\nBuilt histogram-matching reference pool: 20 patches (53.2 MB)\n\nBuilding V4 depth-aware model ...\n[backbone] DepthFusionStem: 26 depth slices (raw=True, grad=True, curv=True) -> 46 learned channels\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\nWarning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"cd9bda88f3794f8cbbca06f72bfdf127"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV4 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320 | input tensor=[1, 26, 320, 320]\nEncoder: backbone=tu-convnext_tiny | pretrained=imagenet | architecture=unet\nDepth stem: 26 -> 46 channels (raw=True, grad=True, curv=True)\nDecoder channels: (256, 128, 64, 32, 16)\nParameters: 32.21M total | 32.21M trainable\nZero-channel layers: 0 (PASS)\n\nEncoder feature stages (dry run at 256x256 for speed):\n  stage 0: (1, 46, 256, 256)\n  stage 1: (1, 0, 128, 128)\n  stage 2: (1, 96, 64, 64)\n  stage 3: (1, 192, 32, 32)\n  stage 4: (1, 384, 16, 16)\n  stage 5: (1, 768, 8, 8)\n\nOutput logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\n\nCUDA device: Tesla T4 | capability: (7, 5)\nForward pass: PASS\n[AdaBN report] BatchNorm2d layers found: depth_stem=1 encoder=0 decoder=10 other=0 (total=11)\n  -> the ConvNeXt encoder itself has none (it's LayerNorm-only, as expected). AdaBN would only recalibrate the depth-stem/decoder BN layers listed above.\n======================================================================\n\n\nEstimating positive-pixel fraction ...\nEstimated positive fraction: 0.183978\nOutput bias initialized to -1.4896\nBCE positive weight: 2.106\nDepth-stem parameters: 4 | Encoder parameters: 178 | Decoder/head parameters: 92 tensors\n\n======================================================================\nSTARTING V4 TRAINING\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/359890098.py:1303: FutureWarning: `torch.cuda.amp.GradScaler(args...)` is deprecated. Please use `torch.amp.GradScaler('cuda', args...)` instead.\n  scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n/tmp/ipykernel_58/359890098.py:1342: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"\n[01/10] time=658.9s\ntrain_loss=0.76679 train_dice=0.36091 dann_loss=0.00000\nval_loss=0.83073 val_dice=0.00000 val_iou=0.00000\nprecision=0.00000 recall=0.00000\nencoder_lr=0.0000488 decoder_lr=0.0000976\n*** NEW BEST CHECKPOINT val_dice=0.00000 (EMA weights) ***\n\n[02/10] time=381.0s\ntrain_loss=0.71339 train_dice=0.47697 dann_loss=0.00000\nval_loss=0.76816 val_dice=0.17703 val_iou=0.09711\nprecision=0.85941 recall=0.09868\nencoder_lr=0.0000453 decoder_lr=0.0000905\n*** NEW BEST CHECKPOINT val_dice=0.17703 (EMA weights) ***\n\n[03/10] time=379.8s\ntrain_loss=0.68286 train_dice=0.52766 dann_loss=0.00000\nval_loss=0.69725 val_dice=0.52923 val_iou=0.35983\nprecision=0.64335 recall=0.44949\nencoder_lr=0.0000399 decoder_lr=0.0000796\n*** NEW BEST CHECKPOINT val_dice=0.52923 (EMA weights) ***\n\n[04/10] time=403.2s\ntrain_loss=0.64990 train_dice=0.57949 dann_loss=0.00000\nval_loss=0.66912 val_dice=0.56682 val_iou=0.39550\nprecision=0.54940 recall=0.58537\nencoder_lr=0.0000331 decoder_lr=0.0000658\n*** NEW BEST CHECKPOINT val_dice=0.56682 (EMA weights) ***\n\n[05/10] time=390.4s\ntrain_loss=0.61737 train_dice=0.62845 dann_loss=0.00000\nval_loss=0.65738 val_dice=0.56788 val_iou=0.39653\nprecision=0.51785 recall=0.62860\nencoder_lr=0.0000255 decoder_lr=0.0000505\n*** NEW BEST CHECKPOINT val_dice=0.56788 (EMA weights) ***\n\n[06/10] time=391.4s\ntrain_loss=0.58126 train_dice=0.67672 dann_loss=0.00000\nval_loss=0.65912 val_dice=0.56949 val_iou=0.39810\nprecision=0.51645 recall=0.63467\nencoder_lr=0.0000179 decoder_lr=0.0000352\n*** NEW BEST CHECKPOINT val_dice=0.56949 (EMA weights) ***\n\n[07/10] time=390.2s\ntrain_loss=0.54883 train_dice=0.72186 dann_loss=0.00000\nval_loss=0.67117 val_dice=0.56639 val_iou=0.39508\nprecision=0.52537 recall=0.61436\nencoder_lr=0.0000111 decoder_lr=0.0000214\nNo improvement: 1/2\n\n[08/10] time=390.0s\ntrain_loss=0.52731 train_dice=0.74905 dann_loss=0.00000\nval_loss=0.68263 val_dice=0.56043 val_iou=0.38930\nprecision=0.51568 recall=0.61368\nencoder_lr=0.0000057 decoder_lr=0.0000105\nNo improvement: 2/2\nEarly stopping.\n\nBest validation Dice: 0.5694896144873579\n\nBest checkpoint loaded (EMA weights).\n\n======================================================================\nFINE DICE THRESHOLD SEARCH\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/359890098.py:1498: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"BEST VALIDATION THRESHOLD = 0.63\nDICE AT BEST THRESHOLD = 0.57950\n\nFinal validation metrics at optimized Dice threshold: {'dice': 0.5794964910315127, 'iou': 0.40795146747108674, 'precision': 0.5832399208102367, 'recall': 0.5758008079784775, 'fbeta0.5': 0.5817373398485572}\n\n======================================================================\nLOADING TEST LABELS (diagnostics only)\n======================================================================\nTest GT found: local diagnostic evaluation enabled.\n\n======================================================================\nINFERENCE A: BASELINE + TTA\n======================================================================\nInference patches: 31438\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/359890098.py:1567: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"\nBASELINE TEST METRICS: {'dice': 0.4506455659866333, 'iou': 0.29086020588874817, 'precision': 0.521607518196106, 'recall': 0.3966794013977051, 'fbeta0.5': 0.4907008111476898}\n\n======================================================================\nINFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\n======================================================================\nBuilt train-domain reference pool for diagnostic: 20 patches\nInference patches: 31438\n\nHISTOGRAM-MATCHING DIAGNOSTIC TEST METRICS: {'dice': 0.4487503170967102, 'iou': 0.2892830967903137, 'precision': 0.5447622537612915, 'recall': 0.38151076436042786, 'fbeta0.5': 0.501816987991333}\n\n>>> Histogram matching alone changed local test Dice by -0.0019 (baseline 0.4506 -> 0.4488)\nSaved histogram-matched probability map: /kaggle/working/fragment1_probability_histmatch_v4.npy\nSaved histogram-matched prediction: /kaggle/working/fragment1_prediction_histmatch_v4.png\n\n(Skipping AdaBN inference -- CFG.use_adabn=False. See the architecture report's BatchNorm2d location summary above for why/whether it would help.)\n\nValidation-tuned threshold: 0.6299999999999997\nUnsupervised Otsu threshold: 0.39355987310409546\n\n======================================================================\nLOCAL TEST DIAGNOSTICS SUMMARY\n======================================================================\nA) Baseline + val threshold:             {'dice': 0.4506455659866333, 'iou': 0.29086020588874817, 'precision': 0.521607518196106, 'recall': 0.3966794013977051, 'fbeta0.5': 0.4907008111476898}\nB) Histogram-matched (no retrain) + val threshold: {'dice': 0.4487503170967102, 'iou': 0.2892830967903137, 'precision': 0.5447622537612915, 'recall': 0.38151076436042786, 'fbeta0.5': 0.501816987991333}\nC) AdaBN + val threshold:                 None\nFinal (chosen) raw + val threshold:       {'dice': 0.4506455659866333, 'iou': 0.29086020588874817, 'precision': 0.521607518196106, 'recall': 0.3966794013977051, 'fbeta0.5': 0.4907008111476898}\nFinal (chosen) raw + Otsu threshold:       {'dice': 0.4820377230644226, 'iou': 0.3175557851791382, 'precision': 0.4195447862148285, 'recall': 0.56640625, 'fbeta0.5': 0.4424920678138733}\nFinal (chosen) postprocessed + val threshold: {'dice': 0.4507356882095337, 'iou': 0.29093530774116516, 'precision': 0.5212509632110596, 'recall': 0.3970257043838501, 'fbeta0.5': 0.4905541241168976}\n\nSaved probability map: /kaggle/working/fragment1_probability_v4.npy\nSaved prediction: /kaggle/working/fragment1_prediction_v4.png\nSaved metrics: /kaggle/working/v4_metrics_summary.json\nSaved baseline overview: /kaggle/working/v4_visualizations/fragment1_v4_overview_baseline.png\nSaved histmatch overview: /kaggle/working/v4_visualizations/fragment1_v4_overview_histmatch.png\nSaved final overview: /kaggle/working/v4_visualizations/fragment1_v4_overview_final.png\nSaved 6 patch comparisons.\nSaved training curves.\n\n======================================================================\nV4 COMPLETE\n======================================================================\nBest validation Dice @ 0.50: 0.56949\nBest Dice threshold: 0.63\nValidation Dice @ optimized threshold: 0.57950\nBackbone: tu-convnext_tiny / architecture=unet / depth_stem=True\nEMA: True | DANN: False | MixStyle: True | AdaBN: False\nCLAHE mode: global_shared | Normalization: per_fragment_zscore | Mask erosion: 24px\n\nLocal test Dice -- baseline: 0.45065\nLocal test Dice -- histogram-matched (diagnostic, no retrain): 0.44875\nLocal test Dice -- final postprocessed: 0.45074\n\nBest checkpoint: /kaggle/working/vesuviusnet_v4_best.pth\nProbability map: /kaggle/working/fragment1_probability_v4.npy\nPrediction: /kaggle/working/fragment1_prediction_v4.png\nMetrics: /kaggle/working/v4_metrics_summary.json\nVisualizations: /kaggle/working/v4_visualizations\n\n=== V4 DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# Test F1 F0.5=0.48 , F0.5 postprocessed=0.49\n# ============================================================\n# VESUVIUS INK DETECTION - V6 \"DEPTH-SIGNATURE RESEARCH MODEL\"\n# ============================================================\n# Built on V4/V5. V5 declared a rich set of research-grade modules in CFG\n# (DepthSignatureModule / LDDC / depth positional encoding / multi-head depth\n# attention / depth-statistics channels, physics-informed synthetic ink,\n# fiber-consistency loss, clDice topology loss, MC-dropout uncertainty) but the\n# actual forward/training graph only ever ran V4's simpler DepthFusionStem and\n# V2ComboLoss (BCE+Dice+FocalTversky). That mismatch is exactly what a Q1\n# reviewer (or a careful re-implementation) would catch immediately, so V6's\n# only job is to close that gap and add the methodological rigor (leave-one-\n# fragment-out CV, ablation matrix, held-out threshold discipline, MC-dropout\n# uncertainty as a genuine diagnostic) needed to defend this as a paper rather\n# than a Kaggle-optimization script.\n#\n# ------------------------------------------------------------------\n# WHAT IS NOW ACTUALLY WIRED (vs. V5's CFG-only declarations)\n# ------------------------------------------------------------------\n#  1. DepthSignatureModule (depth_module_type=\"signature\") is a real nn.Module\n#     that participates in the forward pass:\n#       - LDDC: learnable 1D convolutions along the depth axis whose kernels\n#         are re-parameterized to sum to zero every forward pass (a hard\n#         constraint, not a hope), giving genuinely learnable derivative-like\n#         operators instead of V4's fixed finite differences.\n#       - Depth positional encoding: sinusoidal encoding of each of the 26\n#         physical depth positions, broadcast spatially and concatenated in.\n#       - Multi-head depth attention: an HONEST engineering compromise is\n#         documented in the class docstring -- full per-pixel (H,W,D,D)\n#         attention is computed at a pooled resolution (not full patch\n#         resolution) because full-resolution per-pixel attention at\n#         patch_size=480 would need ~2TB of activation memory on a single\n#         Tesla T4. The pooled attention map is bilinearly upsampled back to\n#         full resolution. This trade-off is exactly the kind of thing a\n#         methods section must state explicitly rather than let the CFG\n#         silently imply otherwise.\n#       - Depth-statistics channels (mean/std/max/depth-centroid/gradient\n#         energy/curvature energy) computed analytically, not learned.\n#       - MC-dropout channel + optional gradient checkpointing, both real.\n#  2. Physics-informed synthetic ink (USE_PHYSICAL_SYNTHETIC_INK) now ALSO\n#     updates the label at the injected stroke, unlike the old distractor-only\n#     inject_fake_ink (which deliberately never touched the label). The\n#     Gaussian depth profile I(z) = A*exp(-(z-z0)^2/(2*sigma^2)) determines a\n#     continuous per-slice intensity weight, not a hard \"affected slices\"\n#     subset.\n#  3. Fiber-consistency loss (USE_FIBER_CONSISTENCY_LOSS) does a genuine\n#     second forward pass per training step on a fiber-perturbed copy of the\n#     batch and penalizes prediction drift -- this costs real compute, exactly\n#     as documented, and is now actually added into the backward graph.\n#  4. clDice topology loss (USE_TOPOLOGY_LOSS) is a real differentiable soft-\n#     skeletonization loss (Shit et al. 2021) added into the combo loss.\n#  5. Depth-shift invariance (USE_DEPTH_SHIFT_AUG, new in V6, implements\n#     critique section 10): FragmentVolume can read an alternate depth window\n#     shifted by +/- depth_shift_max physical slices; training does a second\n#     forward pass on the shifted window and penalizes prediction drift, so\n#     the model is pushed toward learning \"ink signature\" rather than\n#     \"ink lives at absolute index 18\".\n#  6. MC-dropout uncertainty (USE_MC_DROPOUT_UNCERTAINTY) is a genuine\n#     multi-pass inference routine (only Dropout stays in train mode) that\n#     produces a real per-pixel variance map, plus a precision-vs-uncertainty\n#     bucket analysis (does the model's self-disagreement predict its errors?).\n#  7. Leave-one-fragment-out cross-validation: with 3 fragments there are 3\n#     folds (train on 2, test on the held-out one). V6 wraps the whole\n#     data/model/train/eval pipeline into `run_one_fold(...)` and drives it\n#     three times, reporting mean +/- std and a paired significance test\n#     across folds instead of a single \"best validation Dice\" number.\n#  8. Threshold discipline: the Dice/F0.5 threshold is selected ONLY on that\n#     fold's validation split and then frozen before touching the held-out\n#     fragment. It is never re-tuned on the held-out fragment.\n#  9. F0.5 (not Dice) is treated as the PRIMARY reported metric, matching the\n#     competition's own precision-weighted metric; Dice/IoU/precision/recall/\n#     clDice are reported alongside it.\n# 10. Ablation harness (`run_ablation_matrix`): a small set of named CFG\n#     overrides (baseline -> +DepthFusion -> +LDDC -> +PE -> +Attention ->\n#     +Physics -> +Topology -> +Consistency) run back-to-back on a single\n#     fold's data (cached, not re-downloaded/re-normalized per arm) so the\n#     component-by-component contribution can actually be reported in a\n#     table, which is what a reviewer will ask for first.\n#\n# ------------------------------------------------------------------\n# WHAT V6 DELIBERATELY DOES NOT DO (kept out on purpose, per the review)\n# ------------------------------------------------------------------\n#  - No grid search over hyperparameters, no ad hoc encoder swapping, no\n#    10-model ensemble. TTA/EMA/postprocessing morphology remain labeled as\n#    inference-engineering details, not scientific contributions.\n#  - DANN / MixStyle / AdaBN remain OFF by default and are kept in a clearly\n#    separate \"Protocol B (transductive)\" code path -- they are never silently\n#    mixed into the strict-protocol numbers used for the main CV table.\n# ============================================================\n\n\n# ============================================================\n# V7 CHANGE LOG (applied on top of V6, not a reversion to V4)\n# ============================================================\n# A second review came in on the plain V4 code and proposed a mix of generic\n# and specific fixes. Applied on top of V6 rather than V4, so nothing already\n# fixed (honest depth-signature wiring, LOFO-CV, ablation matrix, threshold\n# discipline) gets thrown away. Adopted vs. pushed-back-on, explicitly:\n#\n# ADOPTED:\n#  - Longer training with a real schedule: CFG.LR_SCHEDULE supports\n#    \"onecycle\" (default) alongside the existing \"cosine\"; CFG.epochs raised\n#    from 8 -- which was genuinely too short for a pretrained ConvNeXt\n#    encoder -- to a configurable default of 30, governed by the SAME\n#    early-stopping patience so it doesn't just run needlessly long.\n#  - A genuinely higher-capacity depth stem as an ADDITIONAL ablation arm:\n#    Conv3DDepthStem (two 3D-conv branches over raw + first-difference\n#    volumes, depth-mean-pooled) -- this is a real alternative to\n#    DepthSignatureModule's attention-based pooling, not a redundant restate\n#    of it, so it's wired in as depth_module_type=\"conv3d\" and added to the\n#    ablation matrix rather than replacing the signature module.\n#  - CutMix (image+mask cut-and-paste) as an optional additional batch-level\n#    augmentation (USE_CUTMIX), cheap and orthogonal to what's already there.\n#  - Curriculum sampling (USE_CURRICULUM): trains on easier (higher ink\n#    fraction) patches proportionally more early on, shifting toward the\n#    full distribution over curriculum_warmup_epochs -- implemented as a\n#    real per-epoch WeightedRandomSampler rebuild, not just a comment.\n#  - A cheap frequency-domain feature (USE_FREQUENCY_FEATURES): radial\n#    high-frequency energy from the depth-averaged image's FFT magnitude,\n#    added as one extra analytic channel in DepthStatsBranch -- thin ink\n#    strokes contribute disproportionately to high spatial frequencies, so\n#    this is a legitimate cheap signal, not the heavier per-slice FFT-fusion\n#    module the critique sketched (which would multiply memory cost for\n#    unclear extra benefit over a single depth-averaged FFT).\n#\n# PUSHED BACK ON (implemented as clearly-labeled OPTIONAL/transductive-only,\n# NOT folded into the main strict-protocol LOFO-CV numbers, with the reason\n# stated here rather than silently ignored):\n#  - \"Turn DANN + MixStyle on, no clear reason they're off\": there IS a clear\n#    reason, stated in V5/V6 already -- they mix held-out-fragment statistics\n#    into the trained weights, which is exactly the leakage the strict/\n#    transductive protocol split exists to prevent. They stay OFF for the\n#    strict-protocol CV table. CFG.PROTOCOL=\"transductive\" remains the\n#    explicit, separate path for anyone who wants to report transductive\n#    numbers ALONGSIDE (never instead of) the strict ones.\n#  - \"Test-Time Training via entropy minimization on each test patch\": this\n#    is *also* a transductive technique (it fits model weights to the target\n#    fragment's unlabeled statistics at test time) -- implemented as\n#    `test_time_training_adapt`, callable only when CFG.PROTOCOL==\"transductive\"\n#    AND CFG.USE_TEST_TIME_TRAINING, producing a separately-labeled\n#    `test_metrics_ttt` that is never averaged into the strict LOFO-CV\n#    summary. Per-single-patch fine-tuning (as literally proposed) would\n#    re-run an optimizer step per inference patch, which is both\n#    prohibitively slow on a T4 for a ~2000-patch fragment and statistically\n#    dubious (no early-stopping signal without labels) -- so this adapts\n#    once, briefly, over unlabeled target-fragment batches instead of per\n#    patch, then runs one normal inference pass.\n#  - \"patch_size=480 is too small for context\" and \"60% negative patches is\n#    imbalance/domain fooling\": both are tunable CFG values already\n#    (patch_size, target_positive_patch_ratio) rather than bugs -- 480px at\n#    this resolution is a substantial spatial context for a segmentation\n#    patch, and the negative/positive ratio is a deliberate class-balance\n#    choice, not an artifact. Left as configurable rather than \"fixed\",\n#    since there's no evidence in the critique that the current values are\n#    actually wrong for this data.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image scipy\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\nfrom collections import defaultdict\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast as _autocast_new, GradScaler as _GradScaler_new\n\n\ndef autocast(enabled=True):\n    \"\"\"V7.3: thin wrapper so every existing `autocast(enabled=...)` call site\n    keeps working unchanged while using torch>=2.x's non-deprecated\n    torch.amp API instead of the deprecated torch.cuda.amp one.\"\"\"\n    device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    return _autocast_new(device_type, enabled=enabled)\n\n\nclass GradScaler(_GradScaler_new):\n    def __init__(self, enabled=True):\n        device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        super().__init__(device_type, enabled=enabled)\nfrom torch.utils.checkpoint import checkpoint as grad_checkpoint\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\nfrom scipy import stats as sstats\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    all_fragments = [\"1\", \"2\", \"3\"]     # used to build the 3-fold LOFO schedule\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    # V6: depth-shift invariance needs slices OUTSIDE depth_indices to shift into.\n    depth_shift_max = 4\n    depth_pool_indices = list(range(min(depth_indices) - depth_shift_max,\n                                     max(depth_indices) + depth_shift_max + 1))\n\n    patch_size = 320\n    train_stride = 128\n    test_stride = 64\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    # V7.1 (perf): 2 workers was almost certainly starving the GPU given this\n    # augmentation pipeline (elastic transform, histogram matching, physics-ink\n    # injection looping over 26 slices -- all CPU-bound). Raise this to your\n    # actual CPU core count minus 1-2 (Kaggle T4 sessions typically give 4\n    # cores -> try 4; a local box with more cores can go higher).\n    num_workers = 2\n    # V7.4 (perf): 4 queued batches per worker was likely contributing to the\n    # system-RAM OOM that forced num_workers down to 2 -- each queued batch\n    # holds (batch_size, 26, H, W) float32 for BOTH img and shifted_img, so at\n    # num_workers=4 this alone was ~2-3GB just sitting in the prefetch queue.\n    # Halving it frees headroom to raise num_workers back up if you want to;\n    # the tradeoff is a smaller read-ahead buffer, which matters only if your\n    # per-sample CPU cost is spiky (use PROFILE_DATASET below to check).\n    prefetch_factor = 2       # only used when num_workers > 0\n    drop_last = True\n\n    # V7: 8 epochs was genuinely too short for a pretrained ConvNeXt encoder.\n    # Raised to 30, still governed by early_stop_patience so it doesn't run\n    # needlessly long once validation F0.5 plateaus.\n    epochs = 15\n    early_stop_patience = 3\n\n    # V7: LR schedule. \"cosine\" = V6's CosineAnnealingLR. \"onecycle\" = warmup\n    # + cosine anneal in one cycle (Smith 2018), generally a better fit for\n    # a short-ish fine-tuning run than plain cosine-from-the-start.\n    LR_SCHEDULE = \"onecycle\"\n    onecycle_pct_start = 0.10\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- backbone / architecture -------------------------------------------\n    encoder_name = \"tu-convnext_tiny\"\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"          # do NOT use \"unetplusplus\" with tu-* encoders (see V4 notes)\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- V4 depth-fusion stem (kept only as an ablation arm / fallback) -----\n    depth_stem_out_channels = 16\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- V7: Conv3DDepthStem output width (see class docstring) -------------\n    conv3d_stem_out_channels = 48\n\n    # --- V6/V7: which depth-aware front end actually builds into the model --\n    # \"signature\" = DepthSignatureModule (LDDC + depth PE + depth attention + stats)\n    # \"conv3d\"    = V7's Conv3DDepthStem (two 3D-conv branches, depth-mean-pooled)\n    # \"fusion\"    = V4's DepthFusionStem (raw + finite-difference grad/curv)\n    # \"none\"      = raw depth stack straight into the encoder\n    depth_module_type = \"signature\"\n\n    lddc_num_filters = 4\n    lddc_kernel_size = 3\n    depth_pe_dim = 8\n    depth_attention_heads = 4\n    depth_attention_pool = 8        # pooled resolution for depth attention (see docstring)\n    USE_DEPTH_STATS = True\n    depth_signature_dropout_p = 0.2\n    # V7.4 (perf): gradient checkpointing trades GPU compute for GPU memory --\n    # it was needed as a safety margin at patch_size=480, but at 320 (your\n    # current setting) the DepthSignatureModule's activations are ~2.25x\n    # smaller, so the memory pressure it was guarding against is much less\n    # likely. Turning it off means one fewer recompute pass through the\n    # depth-signature module per forward, which is a straightforward speed\n    # win at no cost to what the model learns (checkpointing only changes\n    # memory/compute tradeoff, never numerical results). Re-enable it if you\n    # hit a CUDA out-of-memory error.\n    USE_DEPTH_SIGNATURE_CHECKPOINT = False\n    depth_signature_out_channels = 24\n\n    # --- protocol (strict vs transductive; see V5 CFG notes) ----------------\n    PROTOCOL = \"strict\"\n\n    # --- V6: physics-informed synthetic ink (now WITH label update) --------\n    USE_PHYSICAL_SYNTHETIC_INK = True\n    physical_ink_p = 0.15\n    physical_ink_min_amplitude = 10\n    physical_ink_max_amplitude = 60\n    physical_ink_sigma_range = (2.0, 6.0)   # in depth-slice units\n\n    # --- V6: fiber-invariant consistency loss (real second forward pass) ---\n    USE_FIBER_CONSISTENCY_LOSS = True\n    fiber_consistency_weight = 0.10\n    fiber_consistency_amplitude = 0.3       # in normalized (z-scored) units\n\n    # --- V6: depth-shift invariance consistency (new, critique section 10) -\n    USE_DEPTH_SHIFT_CONSISTENCY = True\n    depth_shift_consistency_weight = 0.10\n    # V7.4 (perf): the depth-shift consistency loss is only USED on 1-in-\n    # `consistency_every_n_steps` training steps, but the Dataset was\n    # generating the shifted patch (a full extra disk read across 26 slices,\n    # in a separate worker process) for ~50% of SAMPLES regardless -- wasted\n    # I/O on the ~3-in-4 steps where it's computed and immediately discarded.\n    # Scaling this down to roughly match how often it's actually consumed\n    # keeps the same effective training signal at a fraction of the CPU cost.\n    # This is a coarse per-sample approximation of \"1 in N batches\" (a\n    # dataset __getitem__ has no visibility into which batch/step it's part\n    # of), not an exact match -- raise it back toward 0.5 only if you disable\n    # per-step throttling (consistency_every_n_steps=1).\n    depth_shift_p = 0.5 / max(4, 1)          # ~0.125 by default; ties to the\n                                               # default consistency_every_n_steps below\n\n    # V7.1 (perf): fiber-consistency and depth-shift-consistency each cost a\n    # full extra forward pass through the WHOLE model (ConvNeXt+U-Net), not\n    # just the depth stem -- doing that every single step is why an epoch\n    # went from \"slow\" to \"3159s\". Computing them every Nth step instead keeps\n    # the same training signal (it's a regularizer, not the primary loss) at\n    # a fraction of the cost. Set to 1 to restore V6/V7's original\n    # every-step behavior once you've confirmed this isn't your bottleneck.\n    consistency_every_n_steps = 4\n\n    # V7.1 (perf): prints a one-time timing breakdown (dataloader wait vs.\n    # main forward vs. consistency forward(s) vs. backward) for the first\n    # PROFILE_TIMING_STEPS steps of fold training, then stops. Use this to see\n    # whether YOUR bottleneck is actually the GPU compute added above, or the\n    # CPU-side augmentation pipeline / dataloader instead of guessing.\n    PROFILE_TIMING = False   # V7.5: diagnosis complete (I/O, fixed by warm_up_cache) -- turn back on if needed\n    PROFILE_TIMING_STEPS = 8\n\n    # V7.4: per-stage CPU timing inside InkPatchDataset.__getitem__, printed\n    # for the first PROFILE_DATASET_CALLS calls PER WORKER PROCESS (so with\n    # num_workers=2 you'll see ~2x that many interleaved lines -- expected).\n    # Tells you which augmentation stage actually dominates CPU cost instead\n    # of guessing. Turn off once you've identified the bottleneck.\n    PROFILE_DATASET = False  # V7.5: diagnosis complete -- turn back on if warm_up_cache doesn't fully fix it\n    PROFILE_DATASET_CALLS = 6\n\n    # V7.5: the diagnosis is in -- main_patch_read varied 9ms-4880ms for\n    # identical-sized reads, which is the signature of cold random-access\n    # reads against Kaggle's backing filesystem, not CPU cost. This forces\n    # one fast sequential read per slice up front so subsequent random-access\n    # patch reads hit a warm OS cache instead. See FragmentVolume.warm_up_cache.\n    WARM_UP_VOLUME_CACHE = True\n\n    # --- V6: topology-aware (clDice) loss -----------------------------------\n    USE_TOPOLOGY_LOSS = True\n    topology_weight = 0.15\n    cldice_iters = 8\n\n    # --- V6: MC-Dropout uncertainty (genuine multi-pass diagnostic) --------\n    USE_MC_DROPOUT_UNCERTAINTY = True\n    mc_dropout_passes = 8\n\n    # --- LOSS weights (BCE / Dice / FocalTversky always on; clDice optional) -\n    bce_weight = 0.25\n    dice_weight = 0.30\n    focal_tversky_weight = 0.30\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n    # NOTE: bce_weight+dice_weight+focal_tversky_weight+topology_weight should\n    # sum to ~1.0; topology_weight is added on top and the others renormalized\n    # implicitly by training dynamics -- kept explicit rather than hidden.\n\n    # --- THRESHOLD (selected on validation only, frozen before test) --------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n\n    # --- TTA (inference engineering, not a scientific contribution) --------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN / DANN / MixStyle: OFF for the main strict-protocol CV table -\n    use_adabn = False\n    adabn_max_patches = 2000\n    USE_DANN = False\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n    USE_MIXSTYLE = False\n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    # --- POSTPROCESS (inference engineering) --------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- CLAHE ---------------------------------------------------------------\n    CLAHE_MODE = \"global_shared\"     # \"off\" | \"per_slice\" | \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization: multi-slice sampled stats, not just the mid slice ---\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n    frag_stats_sample_slices = 9     # V6: sample across depth, not just mid slice\n\n    # --- domain-randomization augmentation (unchanged from V4) -------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n    USE_SHADOW = True\n    shadow_p = 0.20\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15      # distractor-only fake ink (no label update) -- kept\n                            # for robustness training, separate from physical ink\n    fake_fiber_p = 0.15\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- V7: CutMix (batch-level, orthogonal to the existing per-patch augs) -\n    USE_CUTMIX = True\n    cutmix_p = 0.20\n    cutmix_alpha = 1.0\n\n    # --- V7: curriculum sampling (easy -> full distribution over N epochs) --\n    USE_CURRICULUM = True\n    curriculum_warmup_epochs = 8\n    curriculum_easy_ink_frac = 0.10     # >= this ink fraction counts \"easy\"\n\n    # --- V7: cheap frequency-domain feature (radial high-freq FFT energy) ---\n    USE_FREQUENCY_FEATURES = True\n\n    # --- V7: Test-Time Training -- PROTOCOL=\"transductive\" ONLY. Never mixed\n    # into the strict-protocol LOFO-CV numbers (see V7 change-log docstring).\n    USE_TEST_TIME_TRAINING = False\n    ttt_steps = 20\n    ttt_lr = 1e-5\n    ttt_batch_size = 4\n\n    USE_EMA = True\n    # V7.6 (bugfix): a FIXED ema_decay is silently wrong whenever CFG.epochs\n    # changes. EMAModel initializes its shadow weights as a copy of the\n    # UNTRAINED model -- with decay=0.999 (a ~1000-step effective averaging\n    # window) and a short run (e.g. 6 epochs x ~340-750 steps = ~2000-4500\n    # total steps), the EMA weights used for validation stay heavily anchored\n    # to the early, near-random model for most of training. This is the\n    # direct explanation for fold 2's val_f0.5 spiking to 0.92 at epoch 3 then\n    # collapsing to ~0 on the actual held-out fragment: that spike was a\n    # lucky momentary EMA snapshot, not a converged state. Fixed: ema_decay is\n    # now DERIVED from the actual number of optimizer steps this fold will\n    # run, targeting an effective averaging window of\n    # `ema_target_window_fraction` of total steps, clamped to\n    # [ema_decay_min, ema_decay_max]. This scales sensibly whether you run 6\n    # epochs or 60 without hand-tuning. Set ema_decay_min == ema_decay_max to\n    # pin a fixed value again if you ever want the old behavior.\n    ema_target_window_fraction = 0.03\n    ema_decay_min = 0.90\n    ema_decay_max = 0.9995\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    viz_dir = os.path.join(out_dir, \"v6_visualizations\")\n\n    # --- V6: what to run -----------------------------------------------------\n    # \"lofo_cv\"   : 3-fold leave-one-fragment-out CV (main scientific result)\n    # \"ablation\"  : component ablation matrix on ONE fold\n    # \"single\"    : one train_frags/test_frag run (fast debugging)\n    RUN_MODE = \"lofo_cv\"\n    single_train_frags = [\"2\", \"3\"]\n    single_test_frag = \"1\"\n    ablation_test_frag = \"1\"\n    ablation_train_frags = [\"2\", \"3\"]\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\nclass CfgOverride:\n    \"\"\"Context manager: temporarily overrides CFG attributes, restores them on\n    exit. Used by the ablation harness so each arm is a clean, reproducible\n    CFG state rather than hand-editing globals between runs.\"\"\"\n    def __init__(self, **overrides):\n        self.overrides = overrides\n        self.previous = {}\n\n    def __enter__(self):\n        for k, v in self.overrides.items():\n            self.previous[k] = getattr(CFG, k)\n            setattr(CFG, k, v)\n        return CFG\n\n    def __exit__(self, *exc):\n        for k, v in self.previous.items():\n            setattr(CFG, k, v)\n        return False\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\ndef compute_sample_difficulty(samples, labels_full, patch_size, easy_ink_frac):\n    \"\"\"V7: per-sample difficulty tier for curriculum sampling -- 0.0 (easy:\n    wide ink coverage), 0.5 (medium), 1.0 (hard: thin/sparse ink or none).\n    Index-aligned with `samples` (the same list used to build the training\n    Dataset), so it can be used directly as WeightedRandomSampler weights.\"\"\"\n    difficulties = np.zeros(len(samples), dtype=np.float32)\n    for i, (fid, y, x) in enumerate(samples):\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        if frac >= easy_ink_frac:\n            difficulties[i] = 0.0\n        elif frac > 0:\n            difficulties[i] = 0.5\n        else:\n            difficulties[i] = 1.0\n    return difficulties\n\n\ndef build_curriculum_sampler(difficulties, epoch, total_epochs, warmup_epochs):\n    \"\"\"Returns per-sample WeightedRandomSampler weights. Early epochs\n    strongly favor easy/medium samples; by `warmup_epochs` the weighting has\n    linearly relaxed to uniform (i.e. the full, already-class-balanced\n    distribution `balance_positive_patches` built) -- curriculum learning is\n    meant to warm the model up, not to permanently exclude hard examples.\"\"\"\n    progress = min(epoch / max(warmup_epochs, 1), 1.0)   # 0 -> 1 over warmup\n    # weight(difficulty=1.0) goes from a small floor up to 1.0 (uniform) as\n    # progress -> 1; weight(difficulty=0.0) stays at 1.0 throughout.\n    hard_floor = 0.15\n    weights = 1.0 - (1.0 - hard_floor) * (1.0 - progress) * difficulties\n    weights = np.clip(weights, hard_floor, 1.0)\n    return torch.DoubleTensor(weights)\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME (extended: depth pool for shift-invariance training)\n# ============================================================\n\nclass FragmentVolume:\n    \"\"\"V6 change: opens every slice in CFG.depth_pool_indices (the default\n    26-slice window PLUS +/- depth_shift_max on each side), not just\n    CFG.depth_indices. read_patch() accepts an explicit z_indices list so the\n    depth-shift-consistency training step can request a shifted window of the\n    SAME physical stack without re-opening files.\"\"\"\n\n    def __init__(self, frag_dir, pool_indices, default_indices):\n        self.pool_indices = pool_indices\n        self.default_indices = default_indices\n        self.index_to_pos = {z: i for i, z in enumerate(pool_indices)}\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in pool_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid_pos = self.index_to_pos.get(\n            self.default_indices[len(self.default_indices) // 2],\n            len(self._slices) // 2)\n        ref_slice = np.asarray(self._slices[mid_pos])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(\n            ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def _slice_positions(self, z_indices):\n        positions = []\n        for z in z_indices:\n            z_clamped = min(max(z, self.pool_indices[0]), self.pool_indices[-1])\n            positions.append(self.index_to_pos[z_clamped])\n        return positions\n\n    def read_patch(self, y, x, size, z_indices=None, apply_clahe=None):\n        self._ensure_open()\n        z_indices = self.default_indices if z_indices is None else z_indices\n        positions = self._slice_positions(z_indices)\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(positions), size, size), dtype=np.uint8)\n        for out_i, pos in enumerate(positions):\n            s = self._slices[pos]\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[out_i] = block\n        return out\n\n    def sample_shifted_indices(self, max_shift):\n        \"\"\"Returns a physically-shifted (but still contiguous, still ordered)\n        window of depth indices, clamped to stay inside the opened pool.\"\"\"\n        delta = random.randint(-max_shift, max_shift)\n        shifted = [z + delta for z in self.default_indices]\n        lo, hi = self.pool_indices[0], self.pool_indices[-1]\n        shifted = [min(max(z, lo), hi) for z in shifted]\n        return shifted\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        \"\"\"V6: samples statistics across several depth slices (not just the mid\n        slice) for a more robust per-fragment normalization constant.\"\"\"\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        n_slices = min(CFG.frag_stats_sample_slices, len(self.default_indices))\n        sample_z = sorted(random.sample(self.default_indices, n_slices))\n        sample_positions = self._slice_positions(sample_z)\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            for pos in sample_positions:\n                block = self._slices[pos][y:y + patch_size, x:x + patch_size]\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n                vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def warm_up_cache(self):\n        \"\"\"V7.5: forces every opened slice into the OS page cache with ONE\n        fast SEQUENTIAL read, up front. This is the actual fix for the\n        dataset-profile evidence (`main_patch_read` ranging 9ms-4880ms for\n        the same kind of read): patches are small, scattered, random-access\n        windows into 26 separate TIFF files, and on Kaggle's backing\n        filesystem each cold access to a new region can cost hundreds of ms\n        to multiple seconds. A single sequential pass per slice is a\n        completely different (much faster) I/O pattern, and afterwards every\n        patch read becomes a cache hit -- with train_stride=128 and\n        patch_size=320, patches overlap heavily anyway, so this isn't wasted\n        work. `.sum()` is used purely to force every byte to be read; the\n        result is discarded. This touches the OS-level page cache (not a\n        Python-owned buffer), so it's reclaimed automatically under memory\n        pressure rather than risking an OOM the way permanently loading\n        everything into a retained array would.\"\"\"\n        self._ensure_open()\n        t0 = time.time()\n        total_bytes = 0\n        for arr in self._slices:\n            _ = np.asarray(arr).sum(dtype=np.int64)\n            total_bytes += arr.nbytes\n        elapsed = time.time() - t0\n        rate = total_bytes / 1e6 / max(elapsed, 1e-6)\n        print(f\"  [cache warm-up] {len(self._slices)} slices, {total_bytes/1e9:.2f} GB read \"\n              f\"sequentially in {elapsed:.1f}s ({rate:.1f} MB/s) -- subsequent random-access \"\n              f\"patch reads should now mostly be cache hits.\")\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\ndef make_fragment_volume(frag_dir):\n    return FragmentVolume(frag_dir, CFG.depth_pool_indices, CFG.depth_indices)\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Distractor-only fake ink: perturbs intensity but NEVER touches the\n    label. Used purely as a robustness/negative-hallucination stress test --\n    kept separate from the physics-informed version below, which DOES update\n    the label because it is meant to represent an actual (simulated) ink\n    deposit, not a distractor.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_physical_ink_with_label(img_hwd, label_hw, depth_indices, n_strokes=None):\n    \"\"\"V6: the physics-informed synthetic ink model actually promised in V5's\n    CFG. Models a stroke's cross-depth intensity profile as a Gaussian\n    I(z) = A * exp(-(z - z0)^2 / (2*sigma^2)) with A, sigma, z0 sampled per\n    stroke, applies it as a CONTINUOUS per-slice weight (not a hard \"these N\n    slices are affected\" cutoff), and -- unlike inject_fake_ink -- writes the\n    stroke into the label as well, since this is meant to represent a\n    simulated real ink deposit rather than a distractor.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    label_out = label_hw.copy()\n    n_strokes = n_strokes or random.randint(1, 2)\n    z_arr = np.asarray(depth_indices, dtype=np.float32)\n\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n\n        A_amp = random.uniform(CFG.physical_ink_min_amplitude, CFG.physical_ink_max_amplitude)\n        sign = random.choice([-1, 1])\n        sigma = random.uniform(*CFG.physical_ink_sigma_range)\n        z0 = random.uniform(z_arr.min(), z_arr.max())\n\n        profile = sign * A_amp * np.exp(-((z_arr - z0) ** 2) / (2.0 * sigma ** 2))  # (d,)\n\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        stroke_bool = stroke_mask > 0\n        if not stroke_bool.any():\n            continue\n\n        for zi in range(d):\n            delta = profile[zi]\n            if abs(delta) < 0.5:\n                continue\n            sl = out[:, :, zi].astype(np.float32)\n            sl[stroke_bool] = np.clip(sl[stroke_bool] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n\n        # Only mark label where the Gaussian profile has real amplitude near\n        # the profile's peak (>= 40% of |A|) -- a near-zero-weight slice at\n        # the tail of the Gaussian isn't meaningfully \"ink\" at that slice, but\n        # the 2D label is depth-collapsed anyway, so we mark the stroke\n        # footprint whenever the profile is non-trivial anywhere in depth.\n        if np.abs(profile).max() >= 0.4 * A_amp:\n            label_out[stroke_bool] = 1\n\n    return out, label_out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\ndef cutmix_batch(imgs, masks):\n    \"\"\"V7: batch-level CutMix on the already-loaded tensors (image + label\n    mask cut together, so this stays label-consistent unlike a naive\n    image-only cutmix). Orthogonal to the existing per-patch augmentations\n    (which perturb ONE patch); this mixes TWO patches within a batch.\"\"\"\n    B = imgs.size(0)\n    if B < 2:\n        return imgs, masks\n    lam = float(np.random.beta(CFG.cutmix_alpha, CFG.cutmix_alpha))\n    rand_index = torch.randperm(B, device=imgs.device)\n\n    H, W = imgs.shape[-2:]\n    cut_ratio = math.sqrt(max(1.0 - lam, 1e-6))\n    cut_h, cut_w = int(H * cut_ratio), int(W * cut_ratio)\n    cy, cx = np.random.randint(H), np.random.randint(W)\n    y1, y2 = max(0, cy - cut_h // 2), min(H, cy + cut_h // 2)\n    x1, x2 = max(0, cx - cut_w // 2), min(W, cx + cut_w // 2)\n\n    imgs = imgs.clone()\n    masks = masks.clone()\n    imgs[:, :, y1:y2, x1:x2] = imgs[rand_index][:, :, y1:y2, x1:x2]\n    masks[:, :, y1:y2, x1:x2] = masks[rand_index][:, :, y1:y2, x1:x2]\n    return imgs, masks\n\n\ndef add_fiber_pattern_tensor(img_tensor, amplitude):\n    \"\"\"Tensor-level fiber perturbation (in normalized z-scored units) used by\n    the fiber-consistency loss so the perturbation stays inside the autograd\n    graph without a second numpy round-trip. img_tensor: (B, D, H, W).\"\"\"\n    B, D, H, W = img_tensor.shape\n    device = img_tensor.device\n    theta = torch.rand(B, device=device) * math.pi\n    freq = 0.02 + torch.rand(B, device=device) * 0.04\n    amp = amplitude * (0.5 + torch.rand(B, device=device))\n\n    yy, xx = torch.meshgrid(torch.arange(H, device=device, dtype=torch.float32),\n                             torch.arange(W, device=device, dtype=torch.float32), indexing=\"ij\")\n    yy = yy.unsqueeze(0)   # (1,H,W)\n    xx = xx.unsqueeze(0)\n\n    phase = (xx * torch.cos(theta).view(B, 1, 1) + yy * torch.sin(theta).view(B, 1, 1)) \\\n        * freq.view(B, 1, 1)\n    pattern = amp.view(B, 1, 1) * torch.sin(2 * math.pi * phase)   # (B,H,W)\n    pattern = pattern.unsqueeze(1).expand(-1, D, -1, -1)            # (B,D,H,W)\n    return img_tensor + pattern\n\n\n# ============================================================\n# 8. AUGMENTATION\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    \"\"\"V6 changes: (1) physics-informed ink now updates the label and is\n    applied via inject_physical_ink_with_label, kept independent from the\n    label-blind inject_fake_ink distractor; (2) optionally returns a\n    depth-shifted alternate view of the same patch for the shift-consistency\n    loss (train_mode only).\"\"\"\n\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n        self._profile_calls_left = CFG.PROFILE_DATASET_CALLS if CFG.PROFILE_DATASET else 0\n\n    def __len__(self):\n        return len(self.samples)\n\n    def _prep(self, img_u8_hwd, label_hw, normalize_stats):\n        img = img_u8_hwd.astype(np.float32) / 255.0\n        img = normalize_patch(img, *normalize_stats)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label_hw > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n    def __getitem__(self, idx):\n        # V7.4: per-stage CPU timing, printed for the first\n        # CFG.PROFILE_DATASET_CALLS calls in EACH worker process (so with\n        # num_workers=2 you'll see ~2x that many lines total, interleaved --\n        # that's expected, not a bug). This exists to answer \"which\n        # augmentation is actually slow\" empirically instead of guessing;\n        # set CFG.PROFILE_DATASET=False once you've identified the culprit.\n        do_prof = self.train_mode and self._profile_calls_left > 0\n        if do_prof:\n            self._profile_calls_left -= 1\n            _t = time.time()\n            def _lap(label, timings=[]):\n                nonlocal _t\n                now = time.time()\n                timings.append((label, now - _t))\n                _t = now\n                return timings\n            _timings = []\n\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size].copy()\n        img = np.transpose(patch, (1, 2, 0))   # HWD\n        if do_prof:\n            _timings = _lap(\"main_patch_read\")\n\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n        if do_prof:\n            _timings = _lap(\"hist_match\")\n\n        if self.train_mode and CFG.USE_PHYSICAL_SYNTHETIC_INK and random.random() < CFG.physical_ink_p:\n            img, label = inject_physical_ink_with_label(img, label, CFG.depth_indices)\n        if do_prof:\n            _timings = _lap(\"physical_ink\")\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n        if do_prof:\n            _timings = _lap(\"fake_ink_fiber_shadow\")\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n        if do_prof:\n            _timings = _lap(\"albumentations_transform\")\n\n        img_t, label_t = self._prep(img, label, (vol.frag_mean, vol.frag_std))\n        if do_prof:\n            _timings = _lap(\"prep_to_tensor\")\n\n        shifted_t = None\n        if self.train_mode and CFG.USE_DEPTH_SHIFT_CONSISTENCY and random.random() < CFG.depth_shift_p:\n            shifted_z = vol.sample_shifted_indices(CFG.depth_shift_max)\n            shifted_patch = vol.read_patch(y, x, size, z_indices=shifted_z)\n            shifted_hwd = np.transpose(shifted_patch, (1, 2, 0)).astype(np.float32) / 255.0\n            shifted_hwd = normalize_patch(shifted_hwd, vol.frag_mean, vol.frag_std)\n            shifted_t = torch.from_numpy(np.ascontiguousarray(np.transpose(shifted_hwd, (2, 0, 1))))\n        if do_prof:\n            _timings = _lap(\"depth_shift_patch_read\")\n\n        if shifted_t is None:\n            has_shift = torch.tensor(False)\n            shifted_t = torch.zeros_like(img_t)\n        else:\n            has_shift = torch.tensor(True)\n\n        if do_prof:\n            total = sum(t for _, t in _timings)\n            breakdown = \" | \".join(f\"{name}={t*1000:.1f}ms\" for name, t in _timings)\n            print(f\"    [dataset profile pid={os.getpid()}] TOTAL={total*1000:.1f}ms :: {breakdown}\")\n\n        return img_t, label_t, shifted_t, has_shift\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN (Protocol B only): yields normalized (D,H,W) tensors from the\n    held-out fragment, NO labels.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"V4's simpler depth-aware stem, kept as an ablation arm / fallback.\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = 0\n        if use_raw:\n            total_in += in_depth\n        if use_grad:\n            total_in += (in_depth - 1)\n        if use_curv:\n            total_in += max(in_depth - 2, 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.in_depth = in_depth\n\n    def forward(self, x):\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:, :, :] - x[:, :-1, :, :])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass Conv3DDepthStem(nn.Module):\n    \"\"\"V7 addition: a genuinely higher-capacity alternative depth stem,\n    addressing the '26->16 is a severe information bottleneck' critique with\n    real extra convolutional capacity rather than just a wider 1x1 mix. Two\n    branches (raw depth stack, first-difference depth stack) each go through\n    two 3D conv layers BEFORE any depth-collapsing, then are mean-pooled over\n    depth and concatenated -- unlike a single 1x1 conv straight from 70\n    finite-difference channels down to 16, the 3D convs get real learnable\n    interaction across depth and space before anything is collapsed. This is\n    wired in as its own `depth_module_type=\"conv3d\"` ablation arm alongside\n    DepthSignatureModule, not a replacement for it -- they represent two\n    different hypotheses (attention-based depth pooling vs. convolutional\n    depth pooling) worth comparing empirically, which is exactly what the\n    ablation matrix is for.\"\"\"\n    def __init__(self, in_depth, out_channels=48):\n        super().__init__()\n        half = out_channels // 2\n        self.raw_branch = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=3, padding=1), nn.BatchNorm3d(16), nn.GELU(),\n            nn.Conv3d(16, half, kernel_size=3, padding=1), nn.BatchNorm3d(half), nn.GELU(),\n        )\n        self.grad_branch = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=3, padding=1), nn.BatchNorm3d(16), nn.GELU(),\n            nn.Conv3d(16, out_channels - half, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels - half), nn.GELU(),\n        )\n        self.out_channels = out_channels\n\n    def forward(self, x):   # x: (B, D, H, W)\n        raw = x.unsqueeze(1)                                  # (B,1,D,H,W)\n        grad = (x[:, 1:, :, :] - x[:, :-1, :, :]).unsqueeze(1)  # (B,1,D-1,H,W)\n\n        raw_feat = self.raw_branch(raw).mean(dim=2)            # (B,half,H,W)\n        grad_feat = self.grad_branch(grad).mean(dim=2)         # (B,out-half,H,W)\n        return torch.cat([raw_feat, grad_feat], dim=1)\n\n\nclass LDDC(nn.Module):\n    \"\"\"Learnable Depth Differential Convolution. Learns `num_filters` 1D\n    kernels along the physically-ordered depth axis. Each kernel is\n    re-parameterized every forward pass to sum to zero:\n        w' = w - mean(w)\n    which is a hard constraint (not a training-time regularizer that might\n    only be approximately satisfied) -- guaranteeing every filter behaves like\n    a (learned) derivative operator regardless of what the raw weights drift\n    to during optimization. Implemented as a Conv3d with kernel (k,1,1) over\n    x.unsqueeze(1): (B,1,D,H,W) -> (B,num_filters,D,H,W).\n    \"\"\"\n    def __init__(self, num_filters=4, kernel_size=3):\n        super().__init__()\n        self.num_filters = num_filters\n        self.kernel_size = kernel_size\n        self.weight = nn.Parameter(torch.randn(num_filters, 1, kernel_size, 1, 1) * 0.1)\n        self.bias = nn.Parameter(torch.zeros(num_filters))\n\n    def forward(self, x):   # x: (B, D, H, W)\n        w = self.weight - self.weight.mean(dim=2, keepdim=True)   # zero-sum constraint\n        xin = x.unsqueeze(1)   # (B,1,D,H,W)\n        out = F.conv3d(xin, w, bias=self.bias, padding=(self.kernel_size // 2, 0, 0))\n        return out   # (B, num_filters, D, H, W)\n\n\nclass DepthPositionalEncoding(nn.Module):\n    \"\"\"Sinusoidal encoding of each physical depth position z, broadcast\n    spatially. Returns (B, pe_dim, D, H, W) so it can be concatenated\n    alongside LDDC's depth-preserving output before depth attention pools it\n    down to a 2D feature map.\"\"\"\n    def __init__(self, pe_dim=8):\n        super().__init__()\n        assert pe_dim % 2 == 0\n        self.pe_dim = pe_dim\n        div_term = torch.exp(torch.arange(0, pe_dim, 2).float() * (-math.log(10000.0) / pe_dim))\n        self.register_buffer(\"div_term\", div_term, persistent=False)\n\n    def forward(self, depth_positions, B, H, W, device):\n        # depth_positions: 1D float tensor of physical z indices, length D\n        z = depth_positions.to(device).view(-1, 1)                     # (D,1)\n        angles = z * self.div_term.view(1, -1).to(device)               # (D, pe_dim/2)\n        pe = torch.zeros(z.shape[0], self.pe_dim, device=device)\n        pe[:, 0::2] = torch.sin(angles)\n        pe[:, 1::2] = torch.cos(angles)\n        # (D, pe_dim) -> (1, pe_dim, D, 1, 1) -> broadcast to (B, pe_dim, D, H, W)\n        pe = pe.transpose(0, 1).view(1, self.pe_dim, -1, 1, 1)\n        return pe.expand(B, -1, -1, H, W)\n\n\nclass PooledDepthAttention(nn.Module):\n    \"\"\"Multi-head attention ACROSS the depth axis, at every spatial location.\n\n    Honest engineering note (this is exactly the kind of thing the review\n    flagged as needing to be stated explicitly): true per-pixel attention\n    across D=26 depth positions needs an (H, W, heads, D, D) tensor. At\n    patch_size=480 with heads=4, D=26, batch=8 that is B*H*W*heads*D*D*4 bytes\n    ~= 2 TB of activation memory -- not something a single Tesla T4 (14.5GB)\n    can hold, full stop. So this module computes genuine per-pixel softmax\n    attention across depth at a POOLED spatial resolution (default 480/8=60),\n    where the memory cost (8*60*60*4*26*26*4 bytes ~= 93MB) is trivial, and\n    then bilinearly upsamples the resulting per-position depth-attention\n    output back to full resolution. This still gives spatially-varying,\n    per-location depth attention (unlike a pure squeeze-and-excite global\n    version), just not literally per-pixel -- state this trade-off in the\n    methods section rather than letting the module name imply otherwise.\n    \"\"\"\n    def __init__(self, in_channels, heads=4, pool_size=8):\n        super().__init__()\n        self.heads = heads\n        self.pool_size = pool_size\n        self.head_dim = max(in_channels // heads, 4)\n        inner = self.heads * self.head_dim\n        self.to_q = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.to_k = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.to_v = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.out_proj = nn.Conv1d(inner, in_channels, kernel_size=1)\n\n    def forward(self, x):   # x: (B, C, D, H, W)\n        B, C, D, H, W = x.shape\n        ph = max(1, round(H / self.pool_size))\n        pw = max(1, round(W / self.pool_size))\n        x_pooled = F.adaptive_avg_pool3d(x, output_size=(D, max(1, H // ph), max(1, W // pw)))\n        _, _, _, ph_, pw_ = x_pooled.shape\n\n        # reshape depth axis into the \"sequence\" dimension for a 1D attention\n        # per spatial location: (B*ph_*pw_, C, D)\n        xp = x_pooled.permute(0, 3, 4, 1, 2).reshape(B * ph_ * pw_, C, D)\n        q = self.to_q(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n        k = self.to_k(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n        v = self.to_v(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n\n        attn = torch.einsum(\"nhdi,nhdj->nhij\", q, k) / math.sqrt(self.head_dim)\n        attn = attn.softmax(dim=-1)                                    # (N, heads, D, D)\n        out = torch.einsum(\"nhij,nhdj->nhdi\", attn, v)                 # (N, heads, head_dim, D)\n        out = out.reshape(B * ph_ * pw_, self.heads * self.head_dim, D)\n        out = self.out_proj(out)                                       # (N, C, D)\n        out = out.view(B, ph_, pw_, C, D).permute(0, 3, 4, 1, 2)        # (B, C, D, ph_, pw_)\n\n        out = out.reshape(B, C * D, ph_, pw_)\n        out = F.interpolate(out, size=(H, W), mode=\"bilinear\", align_corners=False)\n        out = out.view(B, C, D, H, W)\n        return out\n\n\nclass DepthStatsBranch(nn.Module):\n    \"\"\"Analytic (non-learned) depth-statistics channels: mean, std, max,\n    depth-centroid (intensity-weighted mean z), gradient energy, curvature\n    energy, and (V7, optional) radial high-frequency FFT energy. Computed\n    directly from the raw depth stack so they carry signal even before\n    LDDC/attention have learned anything useful early in training.\n\n    V7's frequency channel: thin ink strokes contribute disproportionately to\n    high spatial frequencies compared to broad fiber/background texture, so a\n    single cheap FFT magnitude computed on the depth-AVERAGED image (not a\n    separate FFT per slice -- that would multiply memory cost for unclear\n    extra benefit) gives one extra, genuinely informative analytic channel.\n    \"\"\"\n    def __init__(self, use_frequency=False):\n        super().__init__()\n        self.use_frequency = use_frequency\n        self.out_channels = 6 + (1 if use_frequency else 0)\n\n    @staticmethod\n    def _radial_high_freq_energy(mean_img, high_freq_frac=0.5):\n        # mean_img: (B, 1, H, W). Returns (B, 1, H, W) -- the same scalar\n        # (per-sample high-frequency energy fraction) broadcast spatially, so\n        # it can be concatenated as a \"channel\" alongside genuinely spatial\n        # stats without pretending to carry spatial variation it doesn't have.\n        B, _, H, W = mean_img.shape\n        fft = torch.fft.rfft2(mean_img.squeeze(1).float(), norm=\"ortho\")\n        mag = torch.abs(fft)   # (B, H, W//2+1)\n        fy = torch.fft.fftfreq(H, device=mean_img.device).view(H, 1)\n        fx = torch.fft.rfftfreq(W, device=mean_img.device).view(1, -1)\n        radius = torch.sqrt(fy ** 2 + fx ** 2)\n        radius = radius / radius.max().clamp_min(1e-6)\n        high_mask = (radius >= high_freq_frac).float()\n        high_energy = (mag * high_mask).sum(dim=(1, 2))\n        total_energy = mag.sum(dim=(1, 2)).clamp_min(1e-6)\n        frac = (high_energy / total_energy).view(B, 1, 1, 1).expand(-1, 1, H, W)\n        return frac.to(mean_img.dtype)\n\n    def forward(self, x, depth_positions):   # x: (B, D, H, W)\n        B, D, H, W = x.shape\n        mean = x.mean(dim=1, keepdim=True)\n        std = x.std(dim=1, keepdim=True)\n        maxv = x.max(dim=1, keepdim=True).values\n\n        z = depth_positions.to(x.device).view(1, D, 1, 1)\n        weights = F.softmax(x, dim=1)\n        centroid = (weights * z).sum(dim=1, keepdim=True)\n\n        grad = x[:, 1:, :, :] - x[:, :-1, :, :]\n        grad_energy = (grad ** 2).mean(dim=1, keepdim=True)\n\n        if D > 2:\n            curv = x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :]\n            curv_energy = (curv ** 2).mean(dim=1, keepdim=True)\n        else:\n            curv_energy = torch.zeros_like(mean)\n\n        parts = [mean, std, maxv, centroid, grad_energy, curv_energy]\n        if self.use_frequency:\n            parts.append(self._radial_high_freq_energy(mean))\n        return torch.cat(parts, dim=1)   # (B, 6 or 7, H, W)\n\n\nclass DepthSignatureModule(nn.Module):\n    \"\"\"V5's promised (and now actually implemented) depth-aware front end:\n    LDDC + depth positional encoding + pooled multi-head depth attention +\n    analytic depth-statistics channels, mixed down to `out_channels` for the\n    2D encoder. Includes the MC-dropout channel used both for regularization\n    during training and for genuine uncertainty estimation at inference (see\n    run_mc_dropout_uncertainty). Gradient checkpointing is applied to the\n    (D-preserving, memory-heavy) LDDC+PE+attention stack when\n    USE_DEPTH_SIGNATURE_CHECKPOINT is True.\"\"\"\n    def __init__(self, in_depth, depth_positions, out_channels=24):\n        super().__init__()\n        self.in_depth = in_depth\n        self.register_buffer(\"depth_positions\", torch.tensor(depth_positions, dtype=torch.float32),\n                              persistent=False)\n\n        self.lddc = LDDC(CFG.lddc_num_filters, CFG.lddc_kernel_size)\n        self.pe = DepthPositionalEncoding(CFG.depth_pe_dim)\n        lddc_pe_channels = CFG.lddc_num_filters + CFG.depth_pe_dim\n        self.attn = PooledDepthAttention(lddc_pe_channels, heads=CFG.depth_attention_heads,\n                                          pool_size=CFG.depth_attention_pool)\n        self.stats = DepthStatsBranch(use_frequency=CFG.USE_FREQUENCY_FEATURES) if CFG.USE_DEPTH_STATS else None\n\n        depth_collapsed_channels = lddc_pe_channels * in_depth       # attention output, D collapsed via mean\n        stats_channels = self.stats.out_channels if self.stats is not None else 0\n        mix_in = depth_collapsed_channels + stats_channels + in_depth  # + raw stack\n\n        self.dropout = nn.Dropout2d(CFG.depth_signature_dropout_p)\n        self.mix = nn.Sequential(\n            nn.Conv2d(mix_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n\n    def _lddc_pe_attn(self, x):\n        B, D, H, W = x.shape\n        lddc_out = self.lddc(x)                                          # (B, F, D, H, W)\n        pe_out = self.pe(self.depth_positions, B, H, W, x.device)          # (B, pe, D, H, W)\n        combined = torch.cat([lddc_out, pe_out], dim=1)                    # (B, F+pe, D, H, W)\n        attended = self.attn(combined)                                     # (B, F+pe, D, H, W)\n        # collapse depth by taking mean over the attended depth axis for the\n        # final 2D feature map, but keep the FULL (F+pe)*D as concatenated\n        # channels too -- mean loses information, so we use both.\n        collapsed = attended.reshape(B, -1, H, W)                          # (B, (F+pe)*D, H, W)\n        return collapsed\n\n    def forward(self, x):   # x: (B, D, H, W)\n        if CFG.USE_DEPTH_SIGNATURE_CHECKPOINT and self.training:\n            collapsed = grad_checkpoint(self._lddc_pe_attn, x, use_reentrant=False)\n        else:\n            collapsed = self._lddc_pe_attn(x)\n\n        parts = [x, collapsed]\n        if self.stats is not None:\n            parts.append(self.stats(x, self.depth_positions))\n\n        feat = torch.cat(parts, dim=1)\n        feat = self.dropout(feat)\n        return self.mix(feat)\n\n\nclass DepthAwareSegModel(nn.Module):\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):   # x: (B, D, H, W)\n        feat = self.depth_stem(x) if self.depth_stem is not None else x\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' requested, but only \"\n              f\"'unet' has been verified to handle tu-* timm encoders' 0-channel placeholder \"\n              f\"stage correctly -- using UnetPlusPlus may crash with ConvNeXt/other tu-* encoders.\")\n    return arch_cls(\n        encoder_name=encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=in_channels,\n        classes=1,\n        decoder_channels=CFG.decoder_channels,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\ndef build_model():\n    stem = None\n    seg_in_channels = CFG.in_channels\n\n    if CFG.depth_module_type == \"signature\":\n        stem = DepthSignatureModule(CFG.in_channels, CFG.depth_indices, CFG.depth_signature_out_channels)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] DepthSignatureModule: {CFG.in_channels} depth slices -> \"\n              f\"{seg_in_channels} learned channels (LDDC={CFG.lddc_num_filters} filters, \"\n              f\"PE dim={CFG.depth_pe_dim}, attention heads={CFG.depth_attention_heads}, \"\n              f\"pooled to {CFG.depth_attention_pool}x{CFG.depth_attention_pool}, \"\n              f\"stats={CFG.USE_DEPTH_STATS})\")\n    elif CFG.depth_module_type == \"fusion\":\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] DepthFusionStem: {CFG.in_channels} -> {seg_in_channels} channels\")\n    elif CFG.depth_module_type == \"conv3d\":\n        stem = Conv3DDepthStem(CFG.in_channels, CFG.conv3d_stem_out_channels)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] Conv3DDepthStem: {CFG.in_channels} -> {seg_in_channels} channels \"\n              f\"(two 3D-conv branches, depth-mean-pooled)\")\n    elif CFG.depth_module_type == \"none\":\n        stem = None\n        seg_in_channels = CFG.in_channels\n        print(f\"[backbone] no depth-aware stem: raw {CFG.in_channels}-channel stack into encoder\")\n    else:\n        raise ValueError(f\"Unknown depth_module_type={CFG.depth_module_type}\")\n\n    try:\n        seg_model = _build_seg_backbone(CFG.encoder_name, seg_in_channels)\n        print(f\"[backbone] using {CFG.encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {CFG.encoder_name} unavailable/failed ({e}); \"\n              f\"falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model)\n\n\nclass EMAModel:\n    def __init__(self, model, decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. ARCHITECTURE INSPECTION\n# ============================================================\n\ndef inspect_model_channels(model):\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\n@torch.no_grad()\ndef run_architecture_report(model):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"V6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\")\n    print(\"=\" * 70)\n\n    n_params = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Input:  depth slices={CFG.in_channels} | patch={CFG.patch_size}x{CFG.patch_size}\")\n    print(f\"depth_module_type={CFG.depth_module_type} | encoder={CFG.encoder_name} | \"\n          f\"architecture={CFG.architecture}\")\n    print(f\"Parameters: {n_params/1e6:.2f}M total | {n_trainable/1e6:.2f}M trainable\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"\\nZERO/INVALID CHANNEL LAYERS FOUND:\")\n        for item in bad_layers:\n            print(\"  \", item)\n        raise RuntimeError(\"Model contains invalid zero-channel layers -- fix before training.\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    model.eval()\n    small_size = min(CFG.patch_size, 256)\n    dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n    out = model(dummy)\n    print(f\"Dry-run ({small_size}x{small_size}) output logits shape: {tuple(out.shape)}\")\n    has_nan = torch.isnan(out).any().item()\n    has_inf = torch.isinf(out).any().item()\n    print(f\"NaN check: {'FAIL' if has_nan else 'PASS'} | Inf check: {'FAIL' if has_inf else 'PASS'}\")\n    if has_nan or has_inf:\n        raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass -- fix before training.\")\n    model.train()\n\n    if torch.cuda.is_available():\n        print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")\n    print(\"Forward pass: PASS\")\n    print(\"=\" * 70 + \"\\n\")\n\n\n# ============================================================\n# 12. LOSSES (BCE + Dice + FocalTversky + clDice)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\ndef _soft_erode(img):\n    return -F.max_pool2d(-img, kernel_size=3, stride=1, padding=1)\n\n\ndef _soft_dilate(img):\n    return F.max_pool2d(img, kernel_size=3, stride=1, padding=1)\n\n\ndef _soft_open(img):\n    return _soft_dilate(_soft_erode(img))\n\n\ndef soft_skeletonize(img, iters):\n    \"\"\"Differentiable soft-skeletonization (Shit et al. 2021, 'clDice -- A\n    Novel Topology-Preserving Loss Function for Tubular Structure\n    Segmentation'), used for the clDice topology-aware loss below.\"\"\"\n    img1 = _soft_open(img)\n    skel = F.relu(img - img1)\n    for _ in range(iters):\n        img = _soft_erode(img)\n        img1 = _soft_open(img)\n        delta = F.relu(img - img1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef cldice_loss(logits, targets, iters=8, eps=1e-6):\n    probs = torch.sigmoid(logits)\n    skel_pred = soft_skeletonize(probs, iters)\n    skel_true = soft_skeletonize(targets, iters)\n    t_prec = (skel_pred * targets).sum() / (skel_pred.sum() + eps)\n    t_sens = (skel_true * probs).sum() / (skel_true.sum() + eps)\n    cldice = 1.0 - (2.0 * t_prec * t_sens) / (t_prec + t_sens + eps)\n    return cldice\n\n\nclass VXComboLoss(nn.Module):\n    \"\"\"BCE + Dice + FocalTversky, with clDice added when USE_TOPOLOGY_LOSS is\n    on -- unlike V5's CFG, topology_weight now actually reaches the backward\n    graph.\"\"\"\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        loss = CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n        cldice_val = None\n        if CFG.USE_TOPOLOGY_LOSS:\n            cldice_val = cldice_loss(logits, targets, CFG.cldice_iters)\n            loss = loss + CFG.topology_weight * cldice_val\n        return loss, {\"bce\": bce.item(), \"dice\": dice.item(), \"focal_tversky\": tv.item(),\n                       \"cldice\": (cldice_val.item() if cldice_val is not None else None)}\n\n\n# ============================================================\n# 13. METRICS  (F0.5 is the PRIMARY reported metric)\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25   # F0.5\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"f0.5\": float(fbeta), \"dice\": float(dice), \"iou\": float(iou),\n            \"precision\": float(precision), \"recall\": float(recall)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 14. TRAIN / VALIDATION EPOCH (fiber + depth-shift consistency wired in)\n# ============================================================\n\ndef run_epoch(model, criterion, loader, optimizer, scaler, ema, domain_classifier,\n              feature_capture, dann_iter, global_step_holder, total_steps,\n              train_mode=True, threshold=0.5, scheduler=None, scheduler_steps_per_batch=False,\n              profile_steps=0):\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_consistency_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    # V7.4 (essential, not cosmetic): without this, a slow epoch prints\n    # NOTHING between the profiled first `profile_steps` batches and the\n    # epoch-summary line -- exactly the \"goes dark for 90 minutes with no way\n    # to tell if it's working or stuck\" problem. This is separate from\n    # `profile_steps`: it runs every epoch (not just epoch 1), has no\n    # cuda.synchronize() overhead, and reports wall-clock throughput/ETA so a\n    # slow-but-alive run is distinguishable from a genuinely hung one.\n    epoch_t_start = time.time()\n    n_loader_batches = len(loader)\n    progress_every = max(1, n_loader_batches // 20)   # ~20 prints per epoch\n\n    t_data_end = time.time()\n    for batch_idx, (imgs, masks, shifted_imgs, has_shift) in enumerate(loader):\n        if train_mode and n_loader_batches > 0 and (batch_idx + 1) % progress_every == 0:\n            elapsed = time.time() - epoch_t_start\n            rate = (batch_idx + 1) / max(elapsed, 1e-6)\n            eta = (n_loader_batches - (batch_idx + 1)) / max(rate, 1e-6)\n            print(f\"    step {batch_idx + 1}/{n_loader_batches} | \"\n                  f\"{elapsed:.0f}s elapsed | {rate:.2f} steps/s | ETA {eta:.0f}s\")\n\n        do_profile = train_mode and profile_steps > 0 and batch_idx < profile_steps\n        if do_profile:\n            t0 = time.time()\n            data_wait = t0 - t_data_end\n            if CFG.device == \"cuda\":\n                torch.cuda.synchronize()\n\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        shifted_imgs = shifted_imgs.to(CFG.device, non_blocking=True)\n        has_shift = has_shift.to(CFG.device, non_blocking=True)\n\n        # V7.1 (perf): consistency losses only computed every Nth step -- see\n        # CFG.consistency_every_n_steps docstring for why.\n        do_consistency = train_mode and (batch_idx % max(CFG.consistency_every_n_steps, 1) == 0)\n\n        # V7.2 (bugfix, flagged by review): CutMix is applied to `imgs` here,\n        # but `shifted_imgs` comes straight from the dataset and is NEVER\n        # cutmixed (the Dataset builds it independently of this loop). If\n        # depth-shift consistency then compared model(shifted_imgs) against\n        # model(imgs) on a step where imgs got cutmixed, it would be\n        # comparing predictions on a DIFFERENT spatial composition -- not the\n        # same sample at a shifted depth window, which defeats the point of\n        # the consistency loss (and would train it against noise). Fiber\n        # consistency is unaffected: fiber_imgs is built FROM `imgs` after\n        # cutmix, so it stays the same sample as `logits`. Fix: track\n        # whether cutmix fired this step and skip ONLY depth-shift\n        # consistency when it did.\n        did_cutmix = False\n        if train_mode and CFG.USE_CUTMIX and random.random() < CFG.cutmix_p:\n            imgs, masks = cutmix_batch(imgs, masks)\n            did_cutmix = True\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss, loss_parts = criterion(logits, masks)\n                consistency_total = torch.zeros((), device=CFG.device)\n\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_main_fwd = time.time()\n\n                # --- fiber-consistency loss (real 2nd forward pass) ---------\n                # safe under CutMix: fiber_imgs is derived FROM imgs (the\n                # possibly-cutmixed tensor), so both sides of the comparison\n                # are the same sample.\n                if do_consistency and CFG.USE_FIBER_CONSISTENCY_LOSS:\n                    fiber_imgs = add_fiber_pattern_tensor(imgs, CFG.fiber_consistency_amplitude)\n                    fiber_logits = model(fiber_imgs)\n                    fiber_consistency = F.mse_loss(torch.sigmoid(fiber_logits), torch.sigmoid(logits.detach()))\n                    consistency_total = consistency_total + CFG.fiber_consistency_weight * fiber_consistency\n\n                # --- depth-shift consistency loss (real 2nd forward pass) ---\n                # NOT safe under CutMix (see comment above) -- skipped this step.\n                #\n                # V7.3 (critical perf bugfix): this USED to forward only the\n                # has_shift subset (`shifted_imgs[idx]`), whose size is random\n                # every time (0-8, since depth_shift_p=0.5 per sample). With\n                # cudnn.benchmark=True, EVERY new batch size the model sees\n                # triggers a fresh convolution-algorithm benchmark search --\n                # for this model (ConvNeXt+U-Net+3D depth stem) that can cost\n                # seconds to tens of seconds PER NEW SIZE, plus growing\n                # per-shape workspace memory. Over hundreds of steps hitting\n                # sizes 1..8 repeatedly in no particular order, this compounds\n                # into exactly the kind of multi-hour stall that doesn't show\n                # up in a short profiling window (the profiled steps 0 and 4\n                # already show this: 94s and 4s of one-off cost). Fixed: ALWAYS\n                # forward the full, fixed-size batch (same shape as the main/\n                # fiber paths, which is why THOSE stayed fast at ~0.27s/1.2s\n                # steady-state) and mask out the non-shifted samples in the\n                # loss instead of indexing them out of the tensor.\n                if do_consistency and (not did_cutmix) and CFG.USE_DEPTH_SHIFT_CONSISTENCY:\n                    shift_logits = model(shifted_imgs)   # fixed shape: (CFG.batch_size, D, H, W)\n                    shift_probs = torch.sigmoid(shift_logits)\n                    main_probs_detached = torch.sigmoid(logits.detach())\n                    per_sample_mse = F.mse_loss(shift_probs, main_probs_detached,\n                                                 reduction=\"none\").mean(dim=[1, 2, 3])\n                    weight = has_shift.float()\n                    denom = weight.sum().clamp_min(1.0)\n                    shift_consistency = (per_sample_mse * weight).sum() / denom\n                    consistency_total = consistency_total + CFG.depth_shift_consistency_weight * shift_consistency\n                    del shift_logits, shift_probs, main_probs_detached, per_sample_mse\n\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_consist_fwd = time.time()\n\n                loss = loss + consistency_total\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_backward = time.time()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    # V7.3: GradScaler can SKIP the actual optimizer.step() on\n                    # steps where it detects inf/nan gradients (routine during\n                    # AMP's initial loss-scale calibration, typically just the\n                    # first few iterations) -- if we call scheduler.step()\n                    # unconditionally after that, OneCycleLR advances one step\n                    # further than the optimizer actually did, which is what\n                    # the \"lr_scheduler.step() before optimizer.step()\"\n                    # warning is reporting. Comparing the scaler's scale\n                    # before/after detects a skipped step so we skip the\n                    # scheduler step too, keeping the two in sync.\n                    prev_scale = scaler.get_scale()\n                    scaler.step(optimizer)\n                    scaler.update()\n                    step_was_skipped = scaler.get_scale() < prev_scale\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    if scheduler_steps_per_batch and scheduler is not None and not step_was_skipped:\n                        scheduler.step()\n                    global_step_holder[0] += 1\n\n        if do_profile:\n            print(f\"  [profile step {batch_idx}] data_wait={data_wait:.3f}s \"\n                  f\"main_fwd={t_main_fwd - t0:.3f}s \"\n                  f\"consistency_fwd={t_consist_fwd - t_main_fwd:.3f}s \"\n                  f\"(consistency_computed={do_consistency}) \"\n                  f\"backward+step={t_backward - t_consist_fwd:.3f}s \"\n                  f\"TOTAL={t_backward - t0:.3f}s\")\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_consistency_loss += float(consistency_total.detach().item()) if train_mode else 0.0\n\n        del imgs, masks, logits, probs, shifted_imgs, has_shift\n        t_data_end = time.time()\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n        prev_scale = scaler.get_scale()\n        scaler.step(optimizer)\n        scaler.update()\n        step_was_skipped = scaler.get_scale() < prev_scale\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n        # V7.2: this trailing partial-accumulation step is also a real\n        # optimizer step -- OneCycleLR needs scheduler.step() called here\n        # too, or the last step of every epoch silently goes unaccounted for.\n        if scheduler_steps_per_batch and scheduler is not None and not step_was_skipped:\n            scheduler.step()\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_consistency = total_consistency_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_consistency\n\n\n# ============================================================\n# 15. THRESHOLD SEARCH (validation-only, frozen before test)\n# ============================================================\n\n@torch.no_grad()\ndef find_best_threshold(model, loader, metric=\"f0.5\"):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks, _, _ in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    prec = (tp + 1e-6) / (tp + fp + 1e-6)\n    rec = (tp + 1e-6) / (tp + fn + 1e-6)\n    f05 = (1.25 * prec * rec + 1e-6) / (0.25 * prec + rec + 1e-6)\n    scores = f05 if metric == \"f0.5\" else dice\n    best_idx = int(np.argmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), thresholds, scores\n\n\n# ============================================================\n# 16. MC-DROPOUT UNCERTAINTY (genuine multi-pass diagnostic)\n# ============================================================\n\ndef _set_dropout_train(model):\n    for m in model.modules():\n        if isinstance(m, (nn.Dropout, nn.Dropout2d, nn.Dropout3d)):\n            m.train()\n\n\n@torch.no_grad()\ndef run_mc_dropout_uncertainty(model, vol, mask, patch_size, stride, passes):\n    \"\"\"Runs `passes` stochastic forward passes (only Dropout layers stay in\n    train mode; BatchNorm/LayerNorm stay in eval mode) over a sliding window\n    and returns (mean_prob_map, uncertainty_map) where uncertainty is the\n    per-pixel variance across passes.\"\"\"\n    model.eval()\n    _set_dropout_train(model)\n\n    H, W = mask.shape\n    sum_map = np.zeros((H, W), dtype=np.float32)\n    sumsq_map = np.zeros((H, W), dtype=np.float32)\n    weight_map = np.zeros((H, W), dtype=np.float32)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        x_t = torch.from_numpy(raw).unsqueeze(0).to(CFG.device)\n        pass_probs = []\n        for _ in range(passes):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(x_t)\n            pass_probs.append(torch.sigmoid(logits)[0, 0].float().cpu().numpy())\n        pass_probs = np.stack(pass_probs, axis=0)   # (passes, size, size)\n        mean_p = pass_probs.mean(axis=0)\n        sq_p = (pass_probs ** 2).mean(axis=0)\n        sum_map[y:y + patch_size, x:x + patch_size] += mean_p\n        sumsq_map[y:y + patch_size, x:x + patch_size] += sq_p\n        weight_map[y:y + patch_size, x:x + patch_size] += 1.0\n        del x_t, pass_probs\n\n    weight_map[weight_map <= 1e-8] = 1.0\n    mean_prob = sum_map / weight_map\n    mean_sq = sumsq_map / weight_map\n    uncertainty = np.clip(mean_sq - mean_prob ** 2, 0, None)\n\n    model.eval()   # restore full eval mode (dropout off) for any subsequent calls\n    return mean_prob, uncertainty\n\n\ndef analyze_uncertainty_vs_error(prob_map, uncertainty_map, gt, mask, threshold, n_buckets=3):\n    \"\"\"Buckets pixels by uncertainty and reports precision in each bucket --\n    directly answers 'does uncertainty predict where the model is wrong?'\"\"\"\n    valid = mask > 0\n    unc = uncertainty_map[valid]\n    prob = prob_map[valid]\n    gtv = gt[valid]\n    preds = (prob > threshold).astype(np.float32)\n\n    quantiles = np.quantile(unc, np.linspace(0, 1, n_buckets + 1))\n    report = []\n    for i in range(n_buckets):\n        lo, hi = quantiles[i], quantiles[i + 1]\n        bucket = (unc >= lo) & (unc <= hi) if i == n_buckets - 1 else (unc >= lo) & (unc < hi)\n        if bucket.sum() == 0:\n            continue\n        p_bucket = preds[bucket]\n        g_bucket = gtv[bucket]\n        tp = (p_bucket * g_bucket).sum()\n        fp = (p_bucket * (1 - g_bucket)).sum()\n        precision = (tp + 1e-6) / (tp + fp + 1e-6)\n        report.append({\"bucket\": i, \"uncertainty_range\": (float(lo), float(hi)),\n                        \"n_pixels\": int(bucket.sum()), \"precision\": float(precision)})\n    return report\n\n\n# ============================================================\n# 16b. TEST-TIME TRAINING (V7, Protocol=\"transductive\" ONLY -- see change log)\n# ============================================================\n\ndef test_time_training_adapt(model, test_vol, test_coords_unlabeled):\n    \"\"\"Adapts a COPY of the trained model to the held-out fragment's\n    unlabeled statistics via entropy minimization, for CFG.ttt_steps batches.\n    Returns the adapted copy; the caller's original `model` (and therefore\n    the strict-protocol evaluation) is never touched.\n\n    This is a transductive technique -- it fits weights to target-domain\n    unlabeled data -- so it is ONLY ever invoked when\n    CFG.PROTOCOL == \"transductive\" and CFG.USE_TEST_TIME_TRAINING, and its\n    output is stored as a separately-labeled `test_metrics_ttt`, never\n    averaged into the strict LOFO-CV summary (see `summarize_folds`, which\n    only reads `test_metrics_postprocessed`).\"\"\"\n    assert CFG.PROTOCOL == \"transductive\", (\n        \"test_time_training_adapt called outside Protocol B -- refusing, since \"\n        \"this would silently leak target-fragment statistics into a \"\n        \"strict-protocol result.\")\n\n    adapted = copy.deepcopy(model)\n    adapted.train()\n    optimizer = torch.optim.SGD(adapted.parameters(), lr=CFG.ttt_lr)\n\n    ds = UnlabeledPatchDataset(test_vol, test_coords_unlabeled, CFG.patch_size)\n    loader = DataLoader(ds, batch_size=CFG.ttt_batch_size, shuffle=True, num_workers=1, drop_last=True)\n    loader_iter = iter(loader)\n\n    print(f\"[TTT] adapting on {CFG.ttt_steps} unlabeled target-fragment batches \"\n          f\"(lr={CFG.ttt_lr}) ...\")\n    for step in range(CFG.ttt_steps):\n        try:\n            batch = next(loader_iter)\n        except StopIteration:\n            loader_iter = iter(loader)\n            batch = next(loader_iter)\n        batch = batch.to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = adapted(batch)\n            probs = torch.sigmoid(logits).clamp(1e-6, 1 - 1e-6)\n            entropy = -(probs * torch.log(probs) + (1 - probs) * torch.log(1 - probs)).mean()\n        entropy.backward()\n        torch.nn.utils.clip_grad_norm_(adapted.parameters(), CFG.grad_clip)\n        optimizer.step()\n        optimizer.zero_grad(set_to_none=True)\n        del batch, logits, probs\n\n    adapted.eval()\n    return adapted\n\n\n# ============================================================\n# 17. INFERENCE HELPERS (sliding window, TTA, postprocessing)\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats_, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats_[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 18. ONE-FOLD TRAIN/EVAL (the core unit LOFO-CV and ablation both call)\n# ============================================================\n\ndef run_one_fold(train_frags, test_frag, fold_tag, save_viz=False):\n    \"\"\"Builds data, trains, selects threshold on validation only, evaluates\n    on the held-out fragment, and returns a metrics dict. This is the single\n    unit both `run_lofo_cv` and `run_ablation_matrix` call, so every arm goes\n    through IDENTICAL code -- only CFG differs between calls.\"\"\"\n    set_seed(CFG.seed)\n    print(\"\\n\" + \"#\" * 70)\n    print(f\"# FOLD [{fold_tag}]  train={train_frags}  test={test_frag}\")\n    print(\"#\" * 70)\n\n    # ---- build train data ----\n    train_volumes, train_labels_full, train_masks_full = {}, {}, {}\n    train_samples_raw, val_samples = [], []\n\n    for fid in train_frags:\n        frag_dir = os.path.join(CFG.base_dir, fid)\n        mask = load_tissue_mask(frag_dir)\n        labels = load_ink_labels(frag_dir)\n        if labels is None:\n            raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n        vol = make_fragment_volume(frag_dir)\n        if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n            vol.compute_fragment_stats(mask, CFG.patch_size)\n        if CFG.WARM_UP_VOLUME_CACHE:\n            print(f\"  fragment {fid}:\", end=\" \")\n            vol.warm_up_cache()\n\n        coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n        tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n\n        train_volumes[fid] = vol\n        train_labels_full[fid] = labels\n        train_masks_full[fid] = mask\n        train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n        val_samples.extend([(fid, y, x) for y, x in va_coords])\n        print(f\"  fragment {fid}: train={len(tr_coords)} val={len(va_coords)} \"\n              f\"mean={vol.frag_mean:.1f} std={vol.frag_std:.1f}\")\n        del mask\n        cleanup_memory()\n\n    train_samples = balance_positive_patches(\n        train_samples_raw, train_labels_full, CFG.patch_size,\n        positive_threshold=CFG.positive_patch_fraction,\n        target_positive_ratio=CFG.target_positive_patch_ratio,\n        max_positive_repeat=CFG.max_positive_repeat)\n\n    # ---- held-out fragment: unlabeled use only, labels loaded LATE ----\n    test_dir = os.path.join(CFG.base_dir, test_frag)\n    test_mask = load_tissue_mask(test_dir)\n    test_vol = make_fragment_volume(test_dir)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    if CFG.WARM_UP_VOLUME_CACHE:\n        print(f\"  fragment {test_frag} (held-out):\", end=\" \")\n        test_vol.warm_up_cache()\n\n    test_coords_unlabeled = generate_grid_coords(test_mask, CFG.patch_size, CFG.test_stride,\n                                                  CFG.min_tissue_frac_train)\n\n    # V7.2 (bugfix, flagged by review): the histogram-match reference pool was\n    # being built from the HELD-OUT fragment's own image patches and fed\n    # straight into the TRAINING dataset as an augmentation reference -- even\n    # without labels, that lets the model's training-time inputs be reshaped\n    # to look like the test fragment's intensity distribution, which is\n    # exactly the leakage the strict/transductive protocol split exists to\n    # prevent (and which the CFG.PROTOCOL docstring already claimed doesn't\n    # happen). Fixed: in \"strict\" protocol the pool is built from the\n    # TRAINING fragments' own patches; the held-out fragment's patches are\n    # only used for this purpose under PROTOCOL == \"transductive\", where\n    # that's the explicit, clearly-labeled point of the experiment.\n    hist_match_pool = None\n    if CFG.USE_HIST_MATCH_AUG:\n        if CFG.PROTOCOL == \"strict\":\n            pool_source_coords = []\n            for fid in train_frags:\n                coords = generate_grid_coords(train_masks_full[fid], CFG.patch_size, CFG.patch_size,\n                                               CFG.min_tissue_frac_train)\n                pool_source_coords.extend([(fid, y, x) for y, x in coords])\n            n_pool = min(CFG.hist_match_pool_size, len(pool_source_coords))\n            pool_samples = random.sample(pool_source_coords, n_pool) if pool_source_coords else []\n            hist_match_pool = [train_volumes[fid].read_patch(y, x, CFG.patch_size)\n                                for fid, y, x in pool_samples]\n            print(f\"  hist-match pool (strict protocol): {len(hist_match_pool)} patches from \"\n                  f\"TRAINING fragments {train_frags} only -- held-out fragment {test_frag} untouched.\")\n        elif CFG.PROTOCOL == \"transductive\":\n            pool_coords = random.sample(test_coords_unlabeled,\n                                         min(CFG.hist_match_pool_size, len(test_coords_unlabeled)))\n            hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n            print(f\"  hist-match pool (TRANSDUCTIVE protocol, by design): {len(hist_match_pool)} \"\n                  f\"patches from held-out fragment {test_frag}.\")\n        else:\n            raise ValueError(f\"Unknown CFG.PROTOCOL={CFG.PROTOCOL}\")\n\n    # ---- datasets / loaders ----\n    train_transform = build_train_transform()\n    train_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                                transform=train_transform, jitter=CFG.train_jitter,\n                                hist_match_pool=hist_match_pool, train_mode=True)\n    val_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                              transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\n    sample_difficulties = None\n    curriculum_sampler = None\n    if CFG.USE_CURRICULUM:\n        sample_difficulties = compute_sample_difficulty(\n            train_samples, train_labels_full, CFG.patch_size, CFG.curriculum_easy_ink_frac)\n        print(f\"  curriculum sampling ON: warmup_epochs={CFG.curriculum_warmup_epochs} \"\n              f\"easy/medium/hard counts = \"\n              f\"{int((sample_difficulties==0).sum())}/{int((sample_difficulties==0.5).sum())}/\"\n              f\"{int((sample_difficulties==1).sum())}\")\n        # V7.1 (perf): build ONE sampler object and mutate its `.weights`\n        # in-place each epoch instead of recreating the DataLoader (which\n        # respawns worker processes from scratch every epoch -- expensive\n        # with tifffile-backed volumes and a heavy CPU augmentation pipeline).\n        curriculum_sampler = torch.utils.data.WeightedRandomSampler(\n            build_curriculum_sampler(sample_difficulties, 0, CFG.epochs, CFG.curriculum_warmup_epochs),\n            num_samples=len(train_ds), replacement=True)\n        train_loader = DataLoader(\n            train_ds, batch_size=CFG.batch_size, sampler=curriculum_sampler,\n            num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n            drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0,\n            prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n    else:\n        train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                                   num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                                   drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0,\n                                   prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                             num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                             drop_last=False, persistent_workers=CFG.num_workers > 0,\n                             prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n\n    # ---- model / loss / optimizer ----\n    model = build_model().to(CFG.device)\n    run_architecture_report(model)\n\n    pos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\n    with torch.no_grad():\n        bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n        try:\n            model.seg_model.segmentation_head[0].bias.fill_(bias_val)\n        except Exception as e:\n            print(f\"  (could not set output bias directly: {e})\")\n\n    raw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\n    pos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\n    pos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\n    criterion = VXComboLoss(pos_weight=pos_weight)\n\n    encoder_params, decoder_params, stem_params = [], [], []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if name.startswith(\"depth_stem.\"):\n            stem_params.append(param)\n        elif name.startswith(\"seg_model.encoder.\"):\n            encoder_params.append(param)\n        else:\n            decoder_params.append(param)\n\n    param_groups = [\n        {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n        {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n    ]\n    if stem_params:\n        param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\n    optimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\n    # V7.2/V7.6: computed once, up front, so both the LR schedule and the EMA\n    # decay derive from the SAME actual optimizer-step count for this fold\n    # (they used to disagree: OneCycleLR's own fix used ceil() while the old\n    # `total_steps` below used floor() -- fixed to be consistent everywhere).\n    optimizer_steps_per_epoch = math.ceil(len(train_loader) / CFG.accumulation_steps)\n    total_optimizer_steps = CFG.epochs * optimizer_steps_per_epoch\n\n    if CFG.LR_SCHEDULE == \"onecycle\":\n        max_lrs = [g[\"lr\"] for g in param_groups]\n        # V7.2 (bugfix, flagged by review): scheduler.step() only fires once\n        # per OPTIMIZER update (i.e. once every accumulation_steps batches),\n        # not once per batch -- so telling OneCycleLR steps_per_epoch=\n        # len(train_loader) overstates the cycle length whenever\n        # accumulation_steps > 1, desynchronizing the LR curve from the\n        # actual number of optimizer steps taken. (With the default\n        # accumulation_steps=1 this was a no-op, but it's a real bug for\n        # anyone who raises accumulation_steps, which is a normal thing to\n        # do on a memory-constrained T4.)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=max_lrs, epochs=CFG.epochs, steps_per_epoch=optimizer_steps_per_epoch,\n            pct_start=CFG.onecycle_pct_start, anneal_strategy=\"cos\")\n        scheduler_steps_per_batch = True\n    else:\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n        scheduler_steps_per_batch = False\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\n    # V7.6: ema_decay derived from actual step count for THIS fold -- see the\n    # CFG field's docstring for why a fixed decay silently breaks on short runs.\n    if CFG.USE_EMA:\n        window_steps = max(CFG.ema_target_window_fraction * total_optimizer_steps, 1.0)\n        ema_decay = float(np.clip(1.0 - 1.0 / window_steps, CFG.ema_decay_min, CFG.ema_decay_max))\n        print(f\"  EMA: {total_optimizer_steps} total optimizer steps this fold -> \"\n              f\"target window {window_steps:.0f} steps -> decay={ema_decay:.5f}\")\n        ema = EMAModel(model, decay=ema_decay)\n    else:\n        ema = None\n    global_step_holder = [0]\n    total_steps = total_optimizer_steps\n\n    # ---- training loop ----\n    best_val_f05 = -1.0\n    epochs_no_improve = 0\n    best_state = None\n    history = defaultdict(list)\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n\n        if CFG.USE_CURRICULUM and sample_difficulties is not None:\n            # V7.1 (perf): mutate the existing sampler's weights in place --\n            # does NOT recreate the DataLoader, so persistent worker\n            # processes are kept warm across epochs instead of respawned.\n            curriculum_sampler.weights = build_curriculum_sampler(\n                sample_difficulties, epoch - 1, CFG.epochs, CFG.curriculum_warmup_epochs)\n\n        train_loss, train_metrics, train_consist = run_epoch(\n            model, criterion, train_loader, optimizer, scaler, ema,\n            None, None, None, global_step_holder, total_steps, train_mode=True, threshold=0.50,\n            scheduler=scheduler, scheduler_steps_per_batch=scheduler_steps_per_batch,\n            profile_steps=(CFG.PROFILE_TIMING_STEPS if (CFG.PROFILE_TIMING and epoch == 1) else 0))\n\n        if ema is not None:\n            backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n            model.load_state_dict(ema.state_dict())\n            val_loss, val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, ema,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n            model.load_state_dict(backup)\n            del backup\n            # V7.6 (diagnostic): also evaluate the RAW (non-EMA) weights on\n            # the same validation set. If these two disagree a lot, that's\n            # EMA-driven instability (like fold 2's 0.92-then-collapse\n            # pattern) rather than genuine model instability -- this print\n            # makes that distinction visible instead of having to guess at it\n            # from val_f0.5 alone. Small added cost (one extra val pass).\n            _, raw_val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, None,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n            print(f\"    [EMA vs raw] EMA val_f0.5={val_metrics['f0.5']:.4f} | \"\n                  f\"raw val_f0.5={raw_val_metrics['f0.5']:.4f} | \"\n                  f\"gap={abs(val_metrics['f0.5'] - raw_val_metrics['f0.5']):.4f}\")\n        else:\n            val_loss, val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, ema,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n\n        if not scheduler_steps_per_batch:\n            scheduler.step()\n\n        for k, v in val_metrics.items():\n            history[f\"val_{k}\"].append(v)\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"train_dice\"].append(train_metrics[\"dice\"])\n        history[\"train_consistency_loss\"].append(train_consist)\n\n        print(f\"[{epoch:02d}/{CFG.epochs}] {time.time()-t0:.1f}s | \"\n              f\"train_loss={train_loss:.4f} (consistency={train_consist:.4f}) | \"\n              f\"val_f0.5={val_metrics['f0.5']:.4f} val_dice={val_metrics['dice']:.4f} \"\n              f\"val_prec={val_metrics['precision']:.4f} val_rec={val_metrics['recall']:.4f}\")\n\n        if val_metrics[\"f0.5\"] > best_val_f05:\n            best_val_f05 = val_metrics[\"f0.5\"]\n            epochs_no_improve = 0\n            best_state = copy.deepcopy(ema.state_dict() if ema is not None else model.state_dict())\n            print(f\"  *** new best (val F0.5={best_val_f05:.4f}) ***\")\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= CFG.early_stop_patience:\n                print(\"  early stopping.\")\n                break\n        cleanup_memory()\n\n    model.load_state_dict(best_state)\n\n    # ---- threshold selection: VALIDATION ONLY, then frozen ----\n    best_threshold, val_score_at_best, thr_grid, thr_scores = find_best_threshold(\n        model, val_loader, metric=\"f0.5\")\n    print(f\"\\nSelected threshold (validation-only) = {best_threshold:.2f} \"\n          f\"(val F0.5={val_score_at_best:.4f})\")\n\n    # ---- held-out fragment: load labels now, evaluate at the FROZEN threshold ----\n    test_labels = load_ink_labels(test_dir)\n    fold_result = {\"fold_tag\": fold_tag, \"train_frags\": list(train_frags), \"test_frag\": test_frag,\n                   \"best_val_f05\": best_val_f05, \"selected_threshold\": best_threshold,\n                   \"history\": dict(history)}\n\n    if test_labels is not None:\n        gt_test = (test_labels * test_mask).astype(np.float32)\n        test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                              CFG.test_stride, CFG.infer_batch)\n        test_metrics_raw = evaluate_probability_map(test_prob, gt_test, best_threshold)\n        test_pred_bin = postprocess(test_prob, best_threshold)\n        post_preds = test_pred_bin.astype(np.float32)\n        tp = (post_preds * gt_test).sum(); fp = (post_preds * (1 - gt_test)).sum()\n        fn = ((1 - post_preds) * gt_test).sum()\n        test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n        print(f\"\\nHELD-OUT fragment {test_frag} metrics (frozen threshold={best_threshold:.2f}):\")\n        print(f\"  raw:          {test_metrics_raw}\")\n        print(f\"  postprocessed:{test_metrics_post}\")\n\n        fold_result[\"test_metrics_raw\"] = test_metrics_raw\n        fold_result[\"test_metrics_postprocessed\"] = test_metrics_post\n\n        if CFG.USE_MC_DROPOUT_UNCERTAINTY:\n            mean_prob, uncertainty = run_mc_dropout_uncertainty(\n                model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.mc_dropout_passes)\n            unc_report = analyze_uncertainty_vs_error(mean_prob, uncertainty, gt_test, test_mask,\n                                                       best_threshold, n_buckets=3)\n            print(\"  MC-dropout uncertainty vs. precision (low/med/high uncertainty buckets):\")\n            for r in unc_report:\n                print(f\"    bucket {r['bucket']}: n={r['n_pixels']} precision={r['precision']:.3f}\")\n            fold_result[\"mc_dropout_uncertainty_report\"] = unc_report\n\n        if CFG.PROTOCOL == \"transductive\" and CFG.USE_TEST_TIME_TRAINING:\n            ttt_model = test_time_training_adapt(model, test_vol, test_coords_unlabeled)\n            test_prob_ttt = sliding_window_inference(ttt_model, test_vol, test_mask, CFG.patch_size,\n                                                       CFG.test_stride, CFG.infer_batch)\n            test_metrics_ttt = evaluate_probability_map(test_prob_ttt, gt_test, best_threshold)\n            print(f\"  [TTT, Protocol B, NOT part of strict CV] test metrics: {test_metrics_ttt}\")\n            fold_result[\"test_metrics_ttt_protocol_b_only\"] = test_metrics_ttt\n            del ttt_model\n            cleanup_memory()\n\n        if save_viz:\n            _save_fold_overview(test_dir, test_prob, test_pred_bin, test_labels, fold_tag, best_threshold)\n    else:\n        print(\"No ground-truth labels for held-out fragment -- competition-style inference only.\")\n        test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                              CFG.test_stride, CFG.infer_batch)\n        fold_result[\"test_metrics_raw\"] = None\n\n    prob_path = os.path.join(CFG.out_dir, f\"fragment{test_frag}_probability_{fold_tag}.npy\")\n    np.save(prob_path, test_prob)\n    fold_result[\"probability_map_path\"] = prob_path\n\n    # ---- cleanup ----\n    test_vol.close()\n    for v in train_volumes.values():\n        v.close()\n    del model, optimizer, scheduler\n    cleanup_memory()\n\n    return fold_result\n\n\ndef _save_fold_overview(test_dir, prob_map, pred_bin, gt_labels, tag, threshold):\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 1600 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize((gt_labels * 255).astype(np.uint8), small.shape[::-1],\n                           interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    prob_small = cv2.resize((prob_map * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n    axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability (thr={threshold:.2f})\")\n    axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, f\"overview_{tag}.png\")\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"  saved overview: {path}\")\n\n\n# ============================================================\n# 19. LEAVE-ONE-FRAGMENT-OUT CV  (main scientific result)\n# ============================================================\n\ndef run_lofo_cv(fragments):\n    \"\"\"3 fragments -> 3 folds. Reports mean +/- std across folds instead of a\n    single 'best validation Dice' number, plus a paired t-test / Wilcoxon\n    signed-rank comparison IS available via compare_fold_results below if you\n    run two configurations (e.g. baseline vs. +DepthSignature) through this\n    same function and diff their fold-level F0.5 lists.\"\"\"\n    fold_results = []\n    for held_out in fragments:\n        train_frags = [f for f in fragments if f != held_out]\n        result = run_one_fold(train_frags, held_out, fold_tag=f\"lofo_test{held_out}\",\n                               save_viz=True)\n        fold_results.append(result)\n\n    summary = summarize_folds(fold_results)\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LEAVE-ONE-FRAGMENT-OUT CV SUMMARY\")\n    print(\"=\" * 70)\n    for metric, (mean, std, vals) in summary.items():\n        print(f\"  {metric:12s} = {mean:.4f} +/- {std:.4f}   (per-fold: {['%.4f' % v for v in vals]})\")\n\n    out_path = os.path.join(CFG.out_dir, \"lofo_cv_results.json\")\n    with open(out_path, \"w\") as f:\n        json.dump({\"folds\": fold_results, \"summary\": {k: (v[0], v[1]) for k, v in summary.items()},\n                    \"config\": cfg_to_dict(CFG)}, f, indent=2, default=str)\n    print(f\"\\nSaved: {out_path}\")\n    return fold_results, summary\n\n\ndef summarize_folds(fold_results, metric_keys=(\"f0.5\", \"dice\", \"iou\", \"precision\", \"recall\")):\n    summary = {}\n    for key in metric_keys:\n        vals = [fr[\"test_metrics_postprocessed\"][key] for fr in fold_results\n                 if fr.get(\"test_metrics_postprocessed\") is not None]\n        if not vals:\n            continue\n        summary[key] = (float(np.mean(vals)), float(np.std(vals)), vals)\n    return summary\n\n\ndef compare_fold_results(fold_results_a, fold_results_b, metric=\"f0.5\"):\n    \"\"\"Paired comparison across folds (same held-out fragments in the same\n    order for both configurations) -- Wilcoxon signed-rank test, with a\n    paired t-test reported alongside since n=3 folds is too small for the\n    Wilcoxon test's own asymptotics to be trustworthy; report both and let\n    the reader see they agree in direction.\"\"\"\n    a = [fr[\"test_metrics_postprocessed\"][metric] for fr in fold_results_a]\n    b = [fr[\"test_metrics_postprocessed\"][metric] for fr in fold_results_b]\n    t_stat, t_p = sstats.ttest_rel(a, b)\n    try:\n        w_stat, w_p = sstats.wilcoxon(a, b)\n    except Exception:\n        w_stat, w_p = float(\"nan\"), float(\"nan\")\n    print(f\"Paired comparison on {metric}: A={np.mean(a):.4f} B={np.mean(b):.4f} \"\n          f\"| paired t-test p={t_p:.4f} | Wilcoxon p={w_p:.4f}\")\n    return {\"metric\": metric, \"mean_a\": float(np.mean(a)), \"mean_b\": float(np.mean(b)),\n            \"t_stat\": float(t_stat), \"t_p\": float(t_p), \"w_stat\": float(w_stat), \"w_p\": float(w_p)}\n\n\n# ============================================================\n# 20. ABLATION MATRIX (component-by-component, single fold)\n# ============================================================\n\nABLATION_ARMS = [\n    (\"A_baseline_no_depth_stem\",       dict(depth_module_type=\"none\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"B_depth_fusion_stem\",            dict(depth_module_type=\"fusion\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"C_depth_signature_no_extras\",    dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"D_signature_plus_physics_ink\",   dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"E_signature_plus_topology\",      dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=True)),\n    (\"F_signature_plus_consistency\",   dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=False, USE_CURRICULUM=False)),   # = full V6 model\n    (\"G_conv3d_stem_instead\",          dict(depth_module_type=\"conv3d\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=False, USE_CURRICULUM=False)),\n    (\"H_full_v7_plus_cutmix_curriculum\", dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=True, USE_CURRICULUM=True)),\n]\n\n\ndef run_ablation_matrix(train_frags, test_frag, arms=ABLATION_ARMS):\n    \"\"\"Runs each named arm on the SAME fold (same train/test fragment split)\n    so the differences are attributable to the listed components, not to a\n    different data split. Each arm is a fresh CfgOverride, so arms don't leak\n    settings into each other.\"\"\"\n    results = {}\n    for name, overrides in arms:\n        with CfgOverride(**overrides):\n            print(f\"\\n>>> ABLATION ARM: {name}  overrides={overrides}\")\n            fold_result = run_one_fold(train_frags, test_frag, fold_tag=f\"ablation_{name}\")\n            results[name] = fold_result\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"ABLATION MATRIX SUMMARY (test fragment = %s)\" % test_frag)\n    print(\"=\" * 70)\n    print(f\"{'arm':32s} {'F0.5':>8s} {'Dice':>8s} {'IoU':>8s} {'Prec':>8s} {'Rec':>8s}\")\n    for name, fr in results.items():\n        m = fr.get(\"test_metrics_postprocessed\")\n        if m is None:\n            print(f\"{name:32s}  (no GT available)\")\n            continue\n        print(f\"{name:32s} {m['f0.5']:8.4f} {m['dice']:8.4f} {m['iou']:8.4f} \"\n              f\"{m['precision']:8.4f} {m['recall']:8.4f}\")\n\n    out_path = os.path.join(CFG.out_dir, \"ablation_matrix_results.json\")\n    with open(out_path, \"w\") as f:\n        json.dump(results, f, indent=2, default=str)\n    print(f\"\\nSaved: {out_path}\")\n    return results\n\n\n# ============================================================\n# 21. MAIN\n# ============================================================\n\nif __name__ == \"__main__\":\n    set_seed(CFG.seed)\n\n    if CFG.RUN_MODE == \"lofo_cv\":\n        fold_results, summary = run_lofo_cv(CFG.all_fragments)\n\n    elif CFG.RUN_MODE == \"ablation\":\n        ablation_results = run_ablation_matrix(CFG.ablation_train_frags, CFG.ablation_test_frag)\n\n    elif CFG.RUN_MODE == \"single\":\n        result = run_one_fold(CFG.single_train_frags, CFG.single_test_frag,\n                               fold_tag=\"single_run\", save_viz=True)\n        print(\"\\nSingle-run result:\", json.dumps(\n            {k: v for k, v in result.items() if k != \"history\"}, indent=2, default=str))\n\n    else:\n        raise ValueError(f\"Unknown CFG.RUN_MODE={CFG.RUN_MODE}\")\n\n    print(\"\\n=== V6 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-23T06:12:10.01403Z","iopub.execute_input":"2026-09-23T06:12:10.01475Z"}},"outputs":[{"name":"stdout","text":"   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 11.6 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.7/2.7 MB 86.7 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.30\n\n######################################################################\n# FOLD [lofo_test1]  train=['2', '3']  test=1\n######################################################################\n  fragment 2:   [cache warm-up] 34 slices, 9.59 GB read sequentially in 123.9s (77.4 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 2: train=4559 val=1260 mean=110.0 std=57.6\n  fragment 3:   [cache warm-up] 34 slices, 2.71 GB read sequentially in 29.9s (90.7 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 3: train=1272 val=160 mean=100.3 std=62.4\nPositive patches: 3953 | Negative patches: 1878\nBalanced dataset: 5831 | positive ratio=0.678\n  fragment 1 (held-out):   [cache warm-up] 34 slices, 3.52 GB read sequentially in 49.3s (71.4 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  hist-match pool (strict protocol): 20 patches from TRAINING fragments ['2', '3'] only -- held-out fragment 1 untouched.\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n","output_type":"stream"},{"name":"stdout","text":"  curriculum sampling ON: warmup_epochs=8 easy/medium/hard counts = 2736/1287/1808\n[backbone] DepthSignatureModule: 26 depth slices -> 24 learned channels (LDDC=4 filters, PE dim=8, attention heads=4, pooled to 8x8, stats=True)\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"ffb14c0cbeea4a538f68e89a1f1cbb94"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320\ndepth_module_type=signature | encoder=tu-convnext_tiny | architecture=unet\nParameters: 32.18M total | 32.18M trainable\nZero-channel layers: 0 (PASS)\nDry-run (256x256) output logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\nCUDA device: Tesla T4\nForward pass: PASS\n======================================================================\n\n  EMA: 10920 total optimizer steps this fold -> target window 328 steps -> decay=0.99695\n    step 36/728 | 130s elapsed | 0.28 steps/s | ETA 2496s\n    step 72/728 | 210s elapsed | 0.34 steps/s | ETA 1912s\n    step 108/728 | 288s elapsed | 0.37 steps/s | ETA 1655s\n    step 144/728 | 365s elapsed | 0.39 steps/s | ETA 1482s\n    step 180/728 | 443s elapsed | 0.41 steps/s | ETA 1347s\n    step 216/728 | 517s elapsed | 0.42 steps/s | ETA 1225s\n    step 252/728 | 592s elapsed | 0.43 steps/s | ETA 1119s\n    step 288/728 | 669s elapsed | 0.43 steps/s | ETA 1023s\n    step 324/728 | 747s elapsed | 0.43 steps/s | ETA 931s\n    step 360/728 | 822s elapsed | 0.44 steps/s | ETA 840s\n    step 396/728 | 899s elapsed | 0.44 steps/s | ETA 754s\n    step 432/728 | 976s elapsed | 0.44 steps/s | ETA 669s\n    step 468/728 | 1055s elapsed | 0.44 steps/s | ETA 586s\n    step 504/728 | 1133s elapsed | 0.44 steps/s | ETA 504s\n    step 540/728 | 1213s elapsed | 0.45 steps/s | ETA 422s\n    step 576/728 | 1292s elapsed | 0.45 steps/s | ETA 341s\n    step 612/728 | 1370s elapsed | 0.45 steps/s | ETA 260s\n    step 648/728 | 1450s elapsed | 0.45 steps/s | ETA 179s\n    step 684/728 | 1529s elapsed | 0.45 steps/s | ETA 98s\n    step 720/728 | 1604s elapsed | 0.45 steps/s | ETA 18s\n    [EMA vs raw] EMA val_f0.5=0.0850 | raw val_f0.5=0.2540 | gap=0.1690\n[01/15] 1724.9s | train_loss=0.8346 (consistency=0.0012) | val_f0.5=0.0850 val_dice=0.0394 val_prec=0.3713 val_rec=0.0208\n  *** new best (val F0.5=0.0850) ***\n    step 36/728 | 78s elapsed | 0.46 steps/s | ETA 1497s\n    step 72/728 | 156s elapsed | 0.46 steps/s | ETA 1425s\n    step 108/728 | 235s elapsed | 0.46 steps/s | ETA 1348s\n    step 144/728 | 315s elapsed | 0.46 steps/s | ETA 1277s\n    step 180/728 | 390s elapsed | 0.46 steps/s | ETA 1189s\n    step 216/728 | 468s elapsed | 0.46 steps/s | ETA 1108s\n    step 252/728 | 546s elapsed | 0.46 steps/s | ETA 1031s\n    step 288/728 | 617s elapsed | 0.47 steps/s | ETA 943s\n    step 324/728 | 693s elapsed | 0.47 steps/s | ETA 864s\n    step 360/728 | 770s elapsed | 0.47 steps/s | ETA 787s\n    step 396/728 | 846s elapsed | 0.47 steps/s | ETA 709s\n    step 432/728 | 920s elapsed | 0.47 steps/s | ETA 630s\n    step 468/728 | 1000s elapsed | 0.47 steps/s | ETA 555s\n    step 504/728 | 1077s elapsed | 0.47 steps/s | ETA 479s\n    step 540/728 | 1152s elapsed | 0.47 steps/s | ETA 401s\n    step 576/728 | 1231s elapsed | 0.47 steps/s | ETA 325s\n    step 612/728 | 1304s elapsed | 0.47 steps/s | ETA 247s\n    step 648/728 | 1382s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1462s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1539s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.2886 | raw val_f0.5=0.3568 | gap=0.0682\n[02/15] 1654.7s | train_loss=0.7161 (consistency=0.0011) | val_f0.5=0.2886 val_dice=0.3829 val_prec=0.2479 val_rec=0.8410\n  *** new best (val F0.5=0.2886) ***\n    step 36/728 | 75s elapsed | 0.48 steps/s | ETA 1444s\n    step 72/728 | 152s elapsed | 0.47 steps/s | ETA 1387s\n    step 108/728 | 226s elapsed | 0.48 steps/s | ETA 1300s\n    step 144/728 | 301s elapsed | 0.48 steps/s | ETA 1219s\n    step 180/728 | 378s elapsed | 0.48 steps/s | ETA 1150s\n    step 216/728 | 455s elapsed | 0.48 steps/s | ETA 1078s\n    step 252/728 | 532s elapsed | 0.47 steps/s | ETA 1004s\n    step 288/728 | 609s elapsed | 0.47 steps/s | ETA 930s\n    step 324/728 | 687s elapsed | 0.47 steps/s | ETA 857s\n    step 360/728 | 766s elapsed | 0.47 steps/s | ETA 783s\n    step 396/728 | 840s elapsed | 0.47 steps/s | ETA 704s\n    step 432/728 | 918s elapsed | 0.47 steps/s | ETA 629s\n    step 468/728 | 998s elapsed | 0.47 steps/s | ETA 555s\n    step 504/728 | 1073s elapsed | 0.47 steps/s | ETA 477s\n    step 540/728 | 1151s elapsed | 0.47 steps/s | ETA 401s\n    step 576/728 | 1228s elapsed | 0.47 steps/s | ETA 324s\n    step 612/728 | 1304s elapsed | 0.47 steps/s | ETA 247s\n    step 648/728 | 1381s elapsed | 0.47 steps/s | ETA 170s\n    step 684/728 | 1458s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1535s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.3158 | raw val_f0.5=0.3924 | gap=0.0766\n[03/15] 1650.6s | train_loss=0.6675 (consistency=0.0014) | val_f0.5=0.3158 val_dice=0.4136 val_prec=0.2728 val_rec=0.8541\n  *** new best (val F0.5=0.3158) ***\n    step 36/728 | 78s elapsed | 0.46 steps/s | ETA 1496s\n    step 72/728 | 155s elapsed | 0.46 steps/s | ETA 1412s\n    step 108/728 | 233s elapsed | 0.46 steps/s | ETA 1340s\n    step 144/728 | 313s elapsed | 0.46 steps/s | ETA 1271s\n    step 180/728 | 392s elapsed | 0.46 steps/s | ETA 1193s\n    step 216/728 | 469s elapsed | 0.46 steps/s | ETA 1111s\n    step 252/728 | 545s elapsed | 0.46 steps/s | ETA 1029s\n    step 288/728 | 620s elapsed | 0.46 steps/s | ETA 947s\n    step 324/728 | 699s elapsed | 0.46 steps/s | ETA 871s\n    step 360/728 | 777s elapsed | 0.46 steps/s | ETA 794s\n    step 396/728 | 856s elapsed | 0.46 steps/s | ETA 717s\n    step 432/728 | 933s elapsed | 0.46 steps/s | ETA 639s\n    step 468/728 | 1011s elapsed | 0.46 steps/s | ETA 562s\n    step 504/728 | 1090s elapsed | 0.46 steps/s | ETA 484s\n    step 540/728 | 1167s elapsed | 0.46 steps/s | ETA 406s\n    step 576/728 | 1244s elapsed | 0.46 steps/s | ETA 328s\n    step 612/728 | 1318s elapsed | 0.46 steps/s | ETA 250s\n    step 648/728 | 1395s elapsed | 0.46 steps/s | ETA 172s\n    step 684/728 | 1471s elapsed | 0.47 steps/s | ETA 95s\n    step 720/728 | 1549s elapsed | 0.46 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.3669 | raw val_f0.5=0.2999 | gap=0.0670\n[04/15] 1665.0s | train_loss=0.6435 (consistency=0.0017) | val_f0.5=0.3669 val_dice=0.4631 val_prec=0.3223 val_rec=0.8223\n  *** new best (val F0.5=0.3669) ***\n    step 36/728 | 74s elapsed | 0.49 steps/s | ETA 1414s\n    step 72/728 | 154s elapsed | 0.47 steps/s | ETA 1399s\n    step 108/728 | 231s elapsed | 0.47 steps/s | ETA 1324s\n    step 144/728 | 311s elapsed | 0.46 steps/s | ETA 1259s\n    step 180/728 | 388s elapsed | 0.46 steps/s | ETA 1180s\n    step 216/728 | 462s elapsed | 0.47 steps/s | ETA 1094s\n    step 252/728 | 540s elapsed | 0.47 steps/s | ETA 1020s\n    step 288/728 | 614s elapsed | 0.47 steps/s | ETA 939s\n    step 324/728 | 693s elapsed | 0.47 steps/s | ETA 864s\n    step 360/728 | 770s elapsed | 0.47 steps/s | ETA 787s\n    step 396/728 | 848s elapsed | 0.47 steps/s | ETA 711s\n    step 432/728 | 927s elapsed | 0.47 steps/s | ETA 635s\n    step 468/728 | 1005s elapsed | 0.47 steps/s | ETA 559s\n    step 504/728 | 1084s elapsed | 0.47 steps/s | ETA 482s\n    step 540/728 | 1164s elapsed | 0.46 steps/s | ETA 405s\n    step 576/728 | 1238s elapsed | 0.47 steps/s | ETA 327s\n    step 612/728 | 1316s elapsed | 0.46 steps/s | ETA 250s\n    step 648/728 | 1393s elapsed | 0.47 steps/s | ETA 172s\n    step 684/728 | 1469s elapsed | 0.47 steps/s | ETA 95s\n    step 720/728 | 1548s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.4020 | raw val_f0.5=0.2676 | gap=0.1345\n[05/15] 1664.7s | train_loss=0.6330 (consistency=0.0018) | val_f0.5=0.4020 val_dice=0.4922 val_prec=0.3583 val_rec=0.7860\n  *** new best (val F0.5=0.4020) ***\n    step 36/728 | 79s elapsed | 0.46 steps/s | ETA 1521s\n    step 72/728 | 159s elapsed | 0.45 steps/s | ETA 1449s\n    step 108/728 | 238s elapsed | 0.45 steps/s | ETA 1364s\n    step 144/728 | 313s elapsed | 0.46 steps/s | ETA 1270s\n    step 180/728 | 390s elapsed | 0.46 steps/s | ETA 1188s\n    step 216/728 | 466s elapsed | 0.46 steps/s | ETA 1104s\n    step 252/728 | 544s elapsed | 0.46 steps/s | ETA 1028s\n    step 288/728 | 622s elapsed | 0.46 steps/s | ETA 950s\n    step 324/728 | 701s elapsed | 0.46 steps/s | ETA 875s\n    step 360/728 | 777s elapsed | 0.46 steps/s | ETA 794s\n    step 396/728 | 856s elapsed | 0.46 steps/s | ETA 717s\n    step 432/728 | 934s elapsed | 0.46 steps/s | ETA 640s\n    step 468/728 | 1011s elapsed | 0.46 steps/s | ETA 562s\n    step 504/728 | 1085s elapsed | 0.46 steps/s | ETA 482s\n    step 540/728 | 1161s elapsed | 0.47 steps/s | ETA 404s\n    step 576/728 | 1238s elapsed | 0.47 steps/s | ETA 327s\n    step 612/728 | 1317s elapsed | 0.46 steps/s | ETA 250s\n    step 648/728 | 1394s elapsed | 0.46 steps/s | ETA 172s\n    step 684/728 | 1469s elapsed | 0.47 steps/s | ETA 95s\n    step 720/728 | 1549s elapsed | 0.46 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.4313 | raw val_f0.5=0.4536 | gap=0.0223\n[06/15] 1666.3s | train_loss=0.6203 (consistency=0.0020) | val_f0.5=0.4313 val_dice=0.5091 val_prec=0.3914 val_rec=0.7281\n  *** new best (val F0.5=0.4313) ***\n    step 36/728 | 75s elapsed | 0.48 steps/s | ETA 1444s\n    step 72/728 | 154s elapsed | 0.47 steps/s | ETA 1400s\n    step 108/728 | 232s elapsed | 0.47 steps/s | ETA 1333s\n    step 144/728 | 308s elapsed | 0.47 steps/s | ETA 1248s\n    step 180/728 | 383s elapsed | 0.47 steps/s | ETA 1168s\n    step 216/728 | 461s elapsed | 0.47 steps/s | ETA 1092s\n    step 252/728 | 539s elapsed | 0.47 steps/s | ETA 1018s\n    step 288/728 | 618s elapsed | 0.47 steps/s | ETA 944s\n    step 324/728 | 698s elapsed | 0.46 steps/s | ETA 870s\n    step 360/728 | 773s elapsed | 0.47 steps/s | ETA 790s\n    step 396/728 | 852s elapsed | 0.46 steps/s | ETA 714s\n    step 432/728 | 929s elapsed | 0.47 steps/s | ETA 636s\n    step 468/728 | 1004s elapsed | 0.47 steps/s | ETA 558s\n    step 504/728 | 1083s elapsed | 0.47 steps/s | ETA 481s\n    step 540/728 | 1163s elapsed | 0.46 steps/s | ETA 405s\n    step 576/728 | 1241s elapsed | 0.46 steps/s | ETA 328s\n    step 612/728 | 1316s elapsed | 0.47 steps/s | ETA 249s\n    step 648/728 | 1396s elapsed | 0.46 steps/s | ETA 172s\n    step 684/728 | 1474s elapsed | 0.46 steps/s | ETA 95s\n    step 720/728 | 1548s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.4633 | raw val_f0.5=0.4722 | gap=0.0089\n[07/15] 1665.5s | train_loss=0.6006 (consistency=0.0023) | val_f0.5=0.4633 val_dice=0.5265 val_prec=0.4290 val_rec=0.6813\n  *** new best (val F0.5=0.4633) ***\n    step 36/728 | 78s elapsed | 0.46 steps/s | ETA 1500s\n    step 72/728 | 157s elapsed | 0.46 steps/s | ETA 1426s\n    step 108/728 | 234s elapsed | 0.46 steps/s | ETA 1341s\n    step 144/728 | 312s elapsed | 0.46 steps/s | ETA 1266s\n    step 180/728 | 385s elapsed | 0.47 steps/s | ETA 1172s\n    step 216/728 | 462s elapsed | 0.47 steps/s | ETA 1095s\n    step 252/728 | 540s elapsed | 0.47 steps/s | ETA 1021s\n    step 288/728 | 617s elapsed | 0.47 steps/s | ETA 943s\n    step 324/728 | 696s elapsed | 0.47 steps/s | ETA 868s\n    step 360/728 | 776s elapsed | 0.46 steps/s | ETA 793s\n    step 396/728 | 854s elapsed | 0.46 steps/s | ETA 716s\n    step 432/728 | 929s elapsed | 0.47 steps/s | ETA 636s\n    step 468/728 | 1007s elapsed | 0.46 steps/s | ETA 559s\n    step 504/728 | 1083s elapsed | 0.47 steps/s | ETA 481s\n    step 540/728 | 1160s elapsed | 0.47 steps/s | ETA 404s\n    step 576/728 | 1234s elapsed | 0.47 steps/s | ETA 326s\n    step 612/728 | 1312s elapsed | 0.47 steps/s | ETA 249s\n    step 648/728 | 1391s elapsed | 0.47 steps/s | ETA 172s\n    step 684/728 | 1469s elapsed | 0.47 steps/s | ETA 95s\n    step 720/728 | 1548s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.4769 | raw val_f0.5=0.4222 | gap=0.0547\n[08/15] 1665.1s | train_loss=0.5835 (consistency=0.0021) | val_f0.5=0.4769 val_dice=0.5339 val_prec=0.4451 val_rec=0.6670\n  *** new best (val F0.5=0.4769) ***\n    step 36/728 | 78s elapsed | 0.46 steps/s | ETA 1495s\n    step 72/728 | 156s elapsed | 0.46 steps/s | ETA 1424s\n    step 108/728 | 232s elapsed | 0.47 steps/s | ETA 1332s\n    step 144/728 | 310s elapsed | 0.46 steps/s | ETA 1259s\n    step 180/728 | 389s elapsed | 0.46 steps/s | ETA 1184s\n    step 216/728 | 463s elapsed | 0.47 steps/s | ETA 1098s\n    step 252/728 | 543s elapsed | 0.46 steps/s | ETA 1026s\n    step 288/728 | 619s elapsed | 0.47 steps/s | ETA 945s\n    step 324/728 | 696s elapsed | 0.47 steps/s | ETA 868s\n    step 360/728 | 773s elapsed | 0.47 steps/s | ETA 790s\n    step 396/728 | 851s elapsed | 0.47 steps/s | ETA 714s\n    step 432/728 | 930s elapsed | 0.46 steps/s | ETA 637s\n    step 468/728 | 1006s elapsed | 0.47 steps/s | ETA 559s\n    step 504/728 | 1083s elapsed | 0.47 steps/s | ETA 481s\n    step 540/728 | 1160s elapsed | 0.47 steps/s | ETA 404s\n    step 576/728 | 1232s elapsed | 0.47 steps/s | ETA 325s\n    step 612/728 | 1309s elapsed | 0.47 steps/s | ETA 248s\n    step 648/728 | 1387s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1464s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1544s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.5029 | raw val_f0.5=0.5361 | gap=0.0332\n[09/15] 1660.7s | train_loss=0.5748 (consistency=0.0022) | val_f0.5=0.5029 val_dice=0.5443 val_prec=0.4786 val_rec=0.6311\n  *** new best (val F0.5=0.5029) ***\n    step 36/728 | 77s elapsed | 0.47 steps/s | ETA 1473s\n    step 72/728 | 155s elapsed | 0.46 steps/s | ETA 1413s\n    step 108/728 | 231s elapsed | 0.47 steps/s | ETA 1325s\n    step 144/728 | 309s elapsed | 0.47 steps/s | ETA 1254s\n    step 180/728 | 386s elapsed | 0.47 steps/s | ETA 1176s\n    step 216/728 | 465s elapsed | 0.46 steps/s | ETA 1102s\n    step 252/728 | 540s elapsed | 0.47 steps/s | ETA 1021s\n    step 288/728 | 618s elapsed | 0.47 steps/s | ETA 943s\n    step 324/728 | 696s elapsed | 0.47 steps/s | ETA 868s\n    step 360/728 | 773s elapsed | 0.47 steps/s | ETA 790s\n    step 396/728 | 852s elapsed | 0.46 steps/s | ETA 714s\n    step 432/728 | 930s elapsed | 0.46 steps/s | ETA 637s\n    step 468/728 | 1007s elapsed | 0.46 steps/s | ETA 560s\n    step 504/728 | 1086s elapsed | 0.46 steps/s | ETA 483s\n    step 540/728 | 1160s elapsed | 0.47 steps/s | ETA 404s\n    step 576/728 | 1237s elapsed | 0.47 steps/s | ETA 326s\n    step 612/728 | 1316s elapsed | 0.47 steps/s | ETA 249s\n    step 648/728 | 1391s elapsed | 0.47 steps/s | ETA 172s\n    step 684/728 | 1467s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1544s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.5079 | raw val_f0.5=0.5243 | gap=0.0164\n[10/15] 1659.6s | train_loss=0.5511 (consistency=0.0020) | val_f0.5=0.5079 val_dice=0.5414 val_prec=0.4877 val_rec=0.6084\n  *** new best (val F0.5=0.5079) ***\n    step 36/728 | 78s elapsed | 0.46 steps/s | ETA 1496s\n    step 72/728 | 156s elapsed | 0.46 steps/s | ETA 1425s\n    step 108/728 | 235s elapsed | 0.46 steps/s | ETA 1349s\n    step 144/728 | 315s elapsed | 0.46 steps/s | ETA 1277s\n    step 180/728 | 393s elapsed | 0.46 steps/s | ETA 1198s\n    step 216/728 | 470s elapsed | 0.46 steps/s | ETA 1115s\n    step 252/728 | 549s elapsed | 0.46 steps/s | ETA 1037s\n    step 288/728 | 627s elapsed | 0.46 steps/s | ETA 959s\n    step 324/728 | 707s elapsed | 0.46 steps/s | ETA 882s\n    step 360/728 | 787s elapsed | 0.46 steps/s | ETA 805s\n    step 396/728 | 867s elapsed | 0.46 steps/s | ETA 727s\n    step 432/728 | 943s elapsed | 0.46 steps/s | ETA 646s\n    step 468/728 | 1017s elapsed | 0.46 steps/s | ETA 565s\n    step 504/728 | 1096s elapsed | 0.46 steps/s | ETA 487s\n    step 540/728 | 1170s elapsed | 0.46 steps/s | ETA 407s\n    step 576/728 | 1250s elapsed | 0.46 steps/s | ETA 330s\n    step 612/728 | 1327s elapsed | 0.46 steps/s | ETA 252s\n    step 648/728 | 1406s elapsed | 0.46 steps/s | ETA 174s\n    step 684/728 | 1481s elapsed | 0.46 steps/s | ETA 95s\n    step 720/728 | 1560s elapsed | 0.46 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.5012 | raw val_f0.5=0.5200 | gap=0.0188\n[11/15] 1677.1s | train_loss=0.5385 (consistency=0.0024) | val_f0.5=0.5012 val_dice=0.5378 val_prec=0.4794 val_rec=0.6123\n    step 36/728 | 77s elapsed | 0.47 steps/s | ETA 1473s\n    step 72/728 | 152s elapsed | 0.47 steps/s | ETA 1388s\n    step 108/728 | 225s elapsed | 0.48 steps/s | ETA 1292s\n    step 144/728 | 305s elapsed | 0.47 steps/s | ETA 1237s\n    step 180/728 | 379s elapsed | 0.47 steps/s | ETA 1155s\n    step 216/728 | 456s elapsed | 0.47 steps/s | ETA 1082s\n    step 252/728 | 533s elapsed | 0.47 steps/s | ETA 1008s\n    step 288/728 | 609s elapsed | 0.47 steps/s | ETA 931s\n    step 324/728 | 688s elapsed | 0.47 steps/s | ETA 857s\n    step 360/728 | 766s elapsed | 0.47 steps/s | ETA 783s\n    step 396/728 | 840s elapsed | 0.47 steps/s | ETA 705s\n    step 432/728 | 916s elapsed | 0.47 steps/s | ETA 628s\n    step 468/728 | 990s elapsed | 0.47 steps/s | ETA 550s\n    step 504/728 | 1069s elapsed | 0.47 steps/s | ETA 475s\n    step 540/728 | 1146s elapsed | 0.47 steps/s | ETA 399s\n    step 576/728 | 1223s elapsed | 0.47 steps/s | ETA 323s\n    step 612/728 | 1299s elapsed | 0.47 steps/s | ETA 246s\n    step 648/728 | 1379s elapsed | 0.47 steps/s | ETA 170s\n    step 684/728 | 1456s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1531s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.5081 | raw val_f0.5=0.5199 | gap=0.0118\n[12/15] 1647.0s | train_loss=0.5292 (consistency=0.0020) | val_f0.5=0.5081 val_dice=0.5404 val_prec=0.4886 val_rec=0.6046\n  *** new best (val F0.5=0.5081) ***\n    step 36/728 | 78s elapsed | 0.46 steps/s | ETA 1498s\n    step 72/728 | 155s elapsed | 0.46 steps/s | ETA 1412s\n    step 108/728 | 229s elapsed | 0.47 steps/s | ETA 1316s\n    step 144/728 | 308s elapsed | 0.47 steps/s | ETA 1248s\n    step 180/728 | 386s elapsed | 0.47 steps/s | ETA 1176s\n    step 216/728 | 465s elapsed | 0.46 steps/s | ETA 1102s\n    step 252/728 | 545s elapsed | 0.46 steps/s | ETA 1029s\n    step 288/728 | 622s elapsed | 0.46 steps/s | ETA 950s\n    step 324/728 | 699s elapsed | 0.46 steps/s | ETA 871s\n    step 360/728 | 776s elapsed | 0.46 steps/s | ETA 793s\n    step 396/728 | 853s elapsed | 0.46 steps/s | ETA 715s\n    step 432/728 | 926s elapsed | 0.47 steps/s | ETA 634s\n    step 468/728 | 1003s elapsed | 0.47 steps/s | ETA 557s\n    step 504/728 | 1080s elapsed | 0.47 steps/s | ETA 480s\n    step 540/728 | 1151s elapsed | 0.47 steps/s | ETA 401s\n    step 576/728 | 1228s elapsed | 0.47 steps/s | ETA 324s\n    step 612/728 | 1304s elapsed | 0.47 steps/s | ETA 247s\n    step 648/728 | 1381s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1458s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1532s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.5126 | raw val_f0.5=0.5217 | gap=0.0091\n[13/15] 1649.5s | train_loss=0.5239 (consistency=0.0020) | val_f0.5=0.5126 val_dice=0.5425 val_prec=0.4944 val_rec=0.6010\n  *** new best (val F0.5=0.5126) ***\n    step 36/728 | 79s elapsed | 0.45 steps/s | ETA 1525s\n    step 72/728 | 156s elapsed | 0.46 steps/s | ETA 1425s\n    step 108/728 | 232s elapsed | 0.47 steps/s | ETA 1332s\n    step 144/728 | 309s elapsed | 0.47 steps/s | ETA 1254s\n    step 180/728 | 383s elapsed | 0.47 steps/s | ETA 1167s\n    step 216/728 | 461s elapsed | 0.47 steps/s | ETA 1092s\n    step 252/728 | 539s elapsed | 0.47 steps/s | ETA 1018s\n    step 288/728 | 616s elapsed | 0.47 steps/s | ETA 941s\n    step 324/728 | 693s elapsed | 0.47 steps/s | ETA 864s\n    step 360/728 | 772s elapsed | 0.47 steps/s | ETA 789s\n    step 396/728 | 849s elapsed | 0.47 steps/s | ETA 712s\n    step 432/728 | 929s elapsed | 0.47 steps/s | ETA 636s\n    step 468/728 | 1006s elapsed | 0.47 steps/s | ETA 559s\n    step 504/728 | 1082s elapsed | 0.47 steps/s | ETA 481s\n    step 540/728 | 1160s elapsed | 0.47 steps/s | ETA 404s\n    step 576/728 | 1237s elapsed | 0.47 steps/s | ETA 326s\n    step 612/728 | 1313s elapsed | 0.47 steps/s | ETA 249s\n    step 648/728 | 1389s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1466s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1543s elapsed | 0.47 steps/s | ETA 17s\n    [EMA vs raw] EMA val_f0.5=0.5068 | raw val_f0.5=0.5004 | gap=0.0064\n[14/15] 1659.8s | train_loss=0.5223 (consistency=0.0019) | val_f0.5=0.5068 val_dice=0.5428 val_prec=0.4854 val_rec=0.6157\n    step 36/728 | 79s elapsed | 0.45 steps/s | ETA 1524s\n    step 72/728 | 156s elapsed | 0.46 steps/s | ETA 1425s\n    step 108/728 | 235s elapsed | 0.46 steps/s | ETA 1349s\n    step 144/728 | 311s elapsed | 0.46 steps/s | ETA 1260s\n    step 180/728 | 386s elapsed | 0.47 steps/s | ETA 1176s\n    step 216/728 | 465s elapsed | 0.46 steps/s | ETA 1102s\n    step 252/728 | 542s elapsed | 0.47 steps/s | ETA 1023s\n    step 288/728 | 619s elapsed | 0.47 steps/s | ETA 946s\n    step 324/728 | 699s elapsed | 0.46 steps/s | ETA 872s\n    step 360/728 | 775s elapsed | 0.46 steps/s | ETA 792s\n    step 396/728 | 853s elapsed | 0.46 steps/s | ETA 715s\n    step 432/728 | 930s elapsed | 0.46 steps/s | ETA 637s\n    step 468/728 | 1004s elapsed | 0.47 steps/s | ETA 558s\n    step 504/728 | 1081s elapsed | 0.47 steps/s | ETA 481s\n    step 540/728 | 1159s elapsed | 0.47 steps/s | ETA 403s\n    step 576/728 | 1236s elapsed | 0.47 steps/s | ETA 326s\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS INK DETECTION - V7 \"DEPTH-SIGNATURE RESEARCH MODEL\"\n# ============================================================\n# Built on V4/V5. V5 declared a rich set of research-grade modules in CFG\n# (DepthSignatureModule / LDDC / depth positional encoding / multi-head depth\n# attention / depth-statistics channels, physics-informed synthetic ink,\n# fiber-consistency loss, clDice topology loss, MC-dropout uncertainty) but the\n# actual forward/training graph only ever ran V4's simpler DepthFusionStem and\n# V2ComboLoss (BCE+Dice+FocalTversky). That mismatch is exactly what a Q1\n# reviewer (or a careful re-implementation) would catch immediately, so V6's\n# only job is to close that gap and add the methodological rigor (leave-one-\n# fragment-out CV, ablation matrix, held-out threshold discipline, MC-dropout\n# uncertainty as a genuine diagnostic) needed to defend this as a paper rather\n# than a Kaggle-optimization script.\n#\n# ------------------------------------------------------------------\n# WHAT IS NOW ACTUALLY WIRED (vs. V5's CFG-only declarations)\n# ------------------------------------------------------------------\n#  1. DepthSignatureModule (depth_module_type=\"signature\") is a real nn.Module\n#     that participates in the forward pass:\n#       - LDDC: learnable 1D convolutions along the depth axis whose kernels\n#         are re-parameterized to sum to zero every forward pass (a hard\n#         constraint, not a hope), giving genuinely learnable derivative-like\n#         operators instead of V4's fixed finite differences.\n#       - Depth positional encoding: sinusoidal encoding of each of the 26\n#         physical depth positions, broadcast spatially and concatenated in.\n#       - Multi-head depth attention: an HONEST engineering compromise is\n#         documented in the class docstring -- full per-pixel (H,W,D,D)\n#         attention is computed at a pooled resolution (not full patch\n#         resolution) because full-resolution per-pixel attention at\n#         patch_size=480 would need ~2TB of activation memory on a single\n#         Tesla T4. The pooled attention map is bilinearly upsampled back to\n#         full resolution. This trade-off is exactly the kind of thing a\n#         methods section must state explicitly rather than let the CFG\n#         silently imply otherwise.\n#       - Depth-statistics channels (mean/std/max/depth-centroid/gradient\n#         energy/curvature energy) computed analytically, not learned.\n#       - MC-dropout channel + optional gradient checkpointing, both real.\n#  2. Physics-informed synthetic ink (USE_PHYSICAL_SYNTHETIC_INK) now ALSO\n#     updates the label at the injected stroke, unlike the old distractor-only\n#     inject_fake_ink (which deliberately never touched the label). The\n#     Gaussian depth profile I(z) = A*exp(-(z-z0)^2/(2*sigma^2)) determines a\n#     continuous per-slice intensity weight, not a hard \"affected slices\"\n#     subset.\n#  3. Fiber-consistency loss (USE_FIBER_CONSISTENCY_LOSS) does a genuine\n#     second forward pass per training step on a fiber-perturbed copy of the\n#     batch and penalizes prediction drift -- this costs real compute, exactly\n#     as documented, and is now actually added into the backward graph.\n#  4. clDice topology loss (USE_TOPOLOGY_LOSS) is a real differentiable soft-\n#     skeletonization loss (Shit et al. 2021) added into the combo loss.\n#  5. Depth-shift invariance (USE_DEPTH_SHIFT_AUG, new in V6, implements\n#     critique section 10): FragmentVolume can read an alternate depth window\n#     shifted by +/- depth_shift_max physical slices; training does a second\n#     forward pass on the shifted window and penalizes prediction drift, so\n#     the model is pushed toward learning \"ink signature\" rather than\n#     \"ink lives at absolute index 18\".\n#  6. MC-dropout uncertainty (USE_MC_DROPOUT_UNCERTAINTY) is a genuine\n#     multi-pass inference routine (only Dropout stays in train mode) that\n#     produces a real per-pixel variance map, plus a precision-vs-uncertainty\n#     bucket analysis (does the model's self-disagreement predict its errors?).\n#  7. Leave-one-fragment-out cross-validation: with 3 fragments there are 3\n#     folds (train on 2, test on the held-out one). V6 wraps the whole\n#     data/model/train/eval pipeline into `run_one_fold(...)` and drives it\n#     three times, reporting mean +/- std and a paired significance test\n#     across folds instead of a single \"best validation Dice\" number.\n#  8. Threshold discipline: the Dice/F0.5 threshold is selected ONLY on that\n#     fold's validation split and then frozen before touching the held-out\n#     fragment. It is never re-tuned on the held-out fragment.\n#  9. F0.5 (not Dice) is treated as the PRIMARY reported metric, matching the\n#     competition's own precision-weighted metric; Dice/IoU/precision/recall/\n#     clDice are reported alongside it.\n# 10. Ablation harness (`run_ablation_matrix`): a small set of named CFG\n#     overrides (baseline -> +DepthFusion -> +LDDC -> +PE -> +Attention ->\n#     +Physics -> +Topology -> +Consistency) run back-to-back on a single\n#     fold's data (cached, not re-downloaded/re-normalized per arm) so the\n#     component-by-component contribution can actually be reported in a\n#     table, which is what a reviewer will ask for first.\n#\n# ------------------------------------------------------------------\n# WHAT V6 DELIBERATELY DOES NOT DO (kept out on purpose, per the review)\n# ------------------------------------------------------------------\n#  - No grid search over hyperparameters, no ad hoc encoder swapping, no\n#    10-model ensemble. TTA/EMA/postprocessing morphology remain labeled as\n#    inference-engineering details, not scientific contributions.\n#  - DANN / MixStyle / AdaBN remain OFF by default and are kept in a clearly\n#    separate \"Protocol B (transductive)\" code path -- they are never silently\n#    mixed into the strict-protocol numbers used for the main CV table.\n# ============================================================\n\n\n# ============================================================\n# V7 CHANGE LOG (applied on top of V6, not a reversion to V4)\n# ============================================================\n# A second review came in on the plain V4 code and proposed a mix of generic\n# and specific fixes. Applied on top of V6 rather than V4, so nothing already\n# fixed (honest depth-signature wiring, LOFO-CV, ablation matrix, threshold\n# discipline) gets thrown away. Adopted vs. pushed-back-on, explicitly:\n#\n# ADOPTED:\n#  - Longer training with a real schedule: CFG.LR_SCHEDULE supports\n#    \"onecycle\" (default) alongside the existing \"cosine\"; CFG.epochs raised\n#    from 8 -- which was genuinely too short for a pretrained ConvNeXt\n#    encoder -- to a configurable default of 30, governed by the SAME\n#    early-stopping patience so it doesn't just run needlessly long.\n#  - A genuinely higher-capacity depth stem as an ADDITIONAL ablation arm:\n#    Conv3DDepthStem (two 3D-conv branches over raw + first-difference\n#    volumes, depth-mean-pooled) -- this is a real alternative to\n#    DepthSignatureModule's attention-based pooling, not a redundant restate\n#    of it, so it's wired in as depth_module_type=\"conv3d\" and added to the\n#    ablation matrix rather than replacing the signature module.\n#  - CutMix (image+mask cut-and-paste) as an optional additional batch-level\n#    augmentation (USE_CUTMIX), cheap and orthogonal to what's already there.\n#  - Curriculum sampling (USE_CURRICULUM): trains on easier (higher ink\n#    fraction) patches proportionally more early on, shifting toward the\n#    full distribution over curriculum_warmup_epochs -- implemented as a\n#    real per-epoch WeightedRandomSampler rebuild, not just a comment.\n#  - A cheap frequency-domain feature (USE_FREQUENCY_FEATURES): radial\n#    high-frequency energy from the depth-averaged image's FFT magnitude,\n#    added as one extra analytic channel in DepthStatsBranch -- thin ink\n#    strokes contribute disproportionately to high spatial frequencies, so\n#    this is a legitimate cheap signal, not the heavier per-slice FFT-fusion\n#    module the critique sketched (which would multiply memory cost for\n#    unclear extra benefit over a single depth-averaged FFT).\n#\n# PUSHED BACK ON (implemented as clearly-labeled OPTIONAL/transductive-only,\n# NOT folded into the main strict-protocol LOFO-CV numbers, with the reason\n# stated here rather than silently ignored):\n#  - \"Turn DANN + MixStyle on, no clear reason they're off\": there IS a clear\n#    reason, stated in V5/V6 already -- they mix held-out-fragment statistics\n#    into the trained weights, which is exactly the leakage the strict/\n#    transductive protocol split exists to prevent. They stay OFF for the\n#    strict-protocol CV table. CFG.PROTOCOL=\"transductive\" remains the\n#    explicit, separate path for anyone who wants to report transductive\n#    numbers ALONGSIDE (never instead of) the strict ones.\n#  - \"Test-Time Training via entropy minimization on each test patch\": this\n#    is *also* a transductive technique (it fits model weights to the target\n#    fragment's unlabeled statistics at test time) -- implemented as\n#    `test_time_training_adapt`, callable only when CFG.PROTOCOL==\"transductive\"\n#    AND CFG.USE_TEST_TIME_TRAINING, producing a separately-labeled\n#    `test_metrics_ttt` that is never averaged into the strict LOFO-CV\n#    summary. Per-single-patch fine-tuning (as literally proposed) would\n#    re-run an optimizer step per inference patch, which is both\n#    prohibitively slow on a T4 for a ~2000-patch fragment and statistically\n#    dubious (no early-stopping signal without labels) -- so this adapts\n#    once, briefly, over unlabeled target-fragment batches instead of per\n#    patch, then runs one normal inference pass.\n#  - \"patch_size=480 is too small for context\" and \"60% negative patches is\n#    imbalance/domain fooling\": both are tunable CFG values already\n#    (patch_size, target_positive_patch_ratio) rather than bugs -- 480px at\n#    this resolution is a substantial spatial context for a segmentation\n#    patch, and the negative/positive ratio is a deliberate class-balance\n#    choice, not an artifact. Left as configurable rather than \"fixed\",\n#    since there's no evidence in the critique that the current values are\n#    actually wrong for this data.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image scipy\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\nfrom collections import defaultdict\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast as _autocast_new, GradScaler as _GradScaler_new\n\n\ndef autocast(enabled=True):\n    \"\"\"V7.3: thin wrapper so every existing `autocast(enabled=...)` call site\n    keeps working unchanged while using torch>=2.x's non-deprecated\n    torch.amp API instead of the deprecated torch.cuda.amp one.\"\"\"\n    device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    return _autocast_new(device_type, enabled=enabled)\n\n\nclass GradScaler(_GradScaler_new):\n    def __init__(self, enabled=True):\n        device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        super().__init__(device_type, enabled=enabled)\nfrom torch.utils.checkpoint import checkpoint as grad_checkpoint\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\nfrom scipy import stats as sstats\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    all_fragments = [\"1\", \"2\", \"3\"]     # used to build the 3-fold LOFO schedule\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    # V6: depth-shift invariance needs slices OUTSIDE depth_indices to shift into.\n    depth_shift_max = 4\n    depth_pool_indices = list(range(min(depth_indices) - depth_shift_max,\n                                     max(depth_indices) + depth_shift_max + 1))\n\n    patch_size = 320\n    train_stride = 128\n    test_stride = 128\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    # V7.1 (perf): 2 workers was almost certainly starving the GPU given this\n    # augmentation pipeline (elastic transform, histogram matching, physics-ink\n    # injection looping over 26 slices -- all CPU-bound). Raise this to your\n    # actual CPU core count minus 1-2 (Kaggle T4 sessions typically give 4\n    # cores -> try 4; a local box with more cores can go higher).\n    num_workers = 2\n    # V7.4 (perf): 4 queued batches per worker was likely contributing to the\n    # system-RAM OOM that forced num_workers down to 2 -- each queued batch\n    # holds (batch_size, 26, H, W) float32 for BOTH img and shifted_img, so at\n    # num_workers=4 this alone was ~2-3GB just sitting in the prefetch queue.\n    # Halving it frees headroom to raise num_workers back up if you want to;\n    # the tradeoff is a smaller read-ahead buffer, which matters only if your\n    # per-sample CPU cost is spiky (use PROFILE_DATASET below to check).\n    prefetch_factor = 2       # only used when num_workers > 0\n    drop_last = True\n\n    # V7: 8 epochs was genuinely too short for a pretrained ConvNeXt encoder.\n    # Raised to 30, still governed by early_stop_patience so it doesn't run\n    # needlessly long once validation F0.5 plateaus.\n    epochs = 6\n    early_stop_patience = 3\n\n    # V7: LR schedule. \"cosine\" = V6's CosineAnnealingLR. \"onecycle\" = warmup\n    # + cosine anneal in one cycle (Smith 2018), generally a better fit for\n    # a short-ish fine-tuning run than plain cosine-from-the-start.\n    LR_SCHEDULE = \"onecycle\"\n    onecycle_pct_start = 0.10\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- backbone / architecture -------------------------------------------\n    encoder_name = \"tu-convnext_tiny\"\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"          # do NOT use \"unetplusplus\" with tu-* encoders (see V4 notes)\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- V4 depth-fusion stem (kept only as an ablation arm / fallback) -----\n    depth_stem_out_channels = 16\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- V7: Conv3DDepthStem output width (see class docstring) -------------\n    conv3d_stem_out_channels = 48\n\n    # --- V6/V7: which depth-aware front end actually builds into the model --\n    # \"signature\" = DepthSignatureModule (LDDC + depth PE + depth attention + stats)\n    # \"conv3d\"    = V7's Conv3DDepthStem (two 3D-conv branches, depth-mean-pooled)\n    # \"fusion\"    = V4's DepthFusionStem (raw + finite-difference grad/curv)\n    # \"none\"      = raw depth stack straight into the encoder\n    depth_module_type = \"signature\"\n\n    lddc_num_filters = 4\n    lddc_kernel_size = 3\n    depth_pe_dim = 8\n    depth_attention_heads = 4\n    depth_attention_pool = 8        # pooled resolution for depth attention (see docstring)\n    USE_DEPTH_STATS = True\n    depth_signature_dropout_p = 0.2\n    # V7.4 (perf): gradient checkpointing trades GPU compute for GPU memory --\n    # it was needed as a safety margin at patch_size=480, but at 320 (your\n    # current setting) the DepthSignatureModule's activations are ~2.25x\n    # smaller, so the memory pressure it was guarding against is much less\n    # likely. Turning it off means one fewer recompute pass through the\n    # depth-signature module per forward, which is a straightforward speed\n    # win at no cost to what the model learns (checkpointing only changes\n    # memory/compute tradeoff, never numerical results). Re-enable it if you\n    # hit a CUDA out-of-memory error.\n    USE_DEPTH_SIGNATURE_CHECKPOINT = False\n    depth_signature_out_channels = 24\n\n    # --- protocol (strict vs transductive; see V5 CFG notes) ----------------\n    PROTOCOL = \"strict\"\n\n    # --- V6: physics-informed synthetic ink (now WITH label update) --------\n    USE_PHYSICAL_SYNTHETIC_INK = True\n    physical_ink_p = 0.15\n    physical_ink_min_amplitude = 10\n    physical_ink_max_amplitude = 60\n    physical_ink_sigma_range = (2.0, 6.0)   # in depth-slice units\n\n    # --- V6: fiber-invariant consistency loss (real second forward pass) ---\n    USE_FIBER_CONSISTENCY_LOSS = True\n    fiber_consistency_weight = 0.10\n    fiber_consistency_amplitude = 0.3       # in normalized (z-scored) units\n\n    # --- V6: depth-shift invariance consistency (new, critique section 10) -\n    USE_DEPTH_SHIFT_CONSISTENCY = True\n    depth_shift_consistency_weight = 0.10\n    # V7.4 (perf): the depth-shift consistency loss is only USED on 1-in-\n    # `consistency_every_n_steps` training steps, but the Dataset was\n    # generating the shifted patch (a full extra disk read across 26 slices,\n    # in a separate worker process) for ~50% of SAMPLES regardless -- wasted\n    # I/O on the ~3-in-4 steps where it's computed and immediately discarded.\n    # Scaling this down to roughly match how often it's actually consumed\n    # keeps the same effective training signal at a fraction of the CPU cost.\n    # This is a coarse per-sample approximation of \"1 in N batches\" (a\n    # dataset __getitem__ has no visibility into which batch/step it's part\n    # of), not an exact match -- raise it back toward 0.5 only if you disable\n    # per-step throttling (consistency_every_n_steps=1).\n    depth_shift_p = 0.5 / max(4, 1)          # ~0.125 by default; ties to the\n                                               # default consistency_every_n_steps below\n\n    # V7.1 (perf): fiber-consistency and depth-shift-consistency each cost a\n    # full extra forward pass through the WHOLE model (ConvNeXt+U-Net), not\n    # just the depth stem -- doing that every single step is why an epoch\n    # went from \"slow\" to \"3159s\". Computing them every Nth step instead keeps\n    # the same training signal (it's a regularizer, not the primary loss) at\n    # a fraction of the cost. Set to 1 to restore V6/V7's original\n    # every-step behavior once you've confirmed this isn't your bottleneck.\n    consistency_every_n_steps = 4\n\n    # V7.1 (perf): prints a one-time timing breakdown (dataloader wait vs.\n    # main forward vs. consistency forward(s) vs. backward) for the first\n    # PROFILE_TIMING_STEPS steps of fold training, then stops. Use this to see\n    # whether YOUR bottleneck is actually the GPU compute added above, or the\n    # CPU-side augmentation pipeline / dataloader instead of guessing.\n    PROFILE_TIMING = False   # V7.5: diagnosis complete (I/O, fixed by warm_up_cache) -- turn back on if needed\n    PROFILE_TIMING_STEPS = 8\n\n    # V7.4: per-stage CPU timing inside InkPatchDataset.__getitem__, printed\n    # for the first PROFILE_DATASET_CALLS calls PER WORKER PROCESS (so with\n    # num_workers=2 you'll see ~2x that many interleaved lines -- expected).\n    # Tells you which augmentation stage actually dominates CPU cost instead\n    # of guessing. Turn off once you've identified the bottleneck.\n    PROFILE_DATASET = False  # V7.5: diagnosis complete -- turn back on if warm_up_cache doesn't fully fix it\n    PROFILE_DATASET_CALLS = 6\n\n    # V7.5: the diagnosis is in -- main_patch_read varied 9ms-4880ms for\n    # identical-sized reads, which is the signature of cold random-access\n    # reads against Kaggle's backing filesystem, not CPU cost. This forces\n    # one fast sequential read per slice up front so subsequent random-access\n    # patch reads hit a warm OS cache instead. See FragmentVolume.warm_up_cache.\n    WARM_UP_VOLUME_CACHE = True\n\n    # --- V6: topology-aware (clDice) loss -----------------------------------\n    USE_TOPOLOGY_LOSS = True\n    topology_weight = 0.15\n    cldice_iters = 8\n\n    # --- V6: MC-Dropout uncertainty (genuine multi-pass diagnostic) --------\n    USE_MC_DROPOUT_UNCERTAINTY = True\n    mc_dropout_passes = 8\n\n    # --- LOSS weights (BCE / Dice / FocalTversky always on; clDice optional) -\n    bce_weight = 0.25\n    dice_weight = 0.30\n    focal_tversky_weight = 0.30\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n    # NOTE: bce_weight+dice_weight+focal_tversky_weight+topology_weight should\n    # sum to ~1.0; topology_weight is added on top and the others renormalized\n    # implicitly by training dynamics -- kept explicit rather than hidden.\n\n    # --- THRESHOLD (selected on validation only, frozen before test) --------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n\n    # --- TTA (inference engineering, not a scientific contribution) --------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN / DANN / MixStyle: OFF for the main strict-protocol CV table -\n    use_adabn = False\n    adabn_max_patches = 2000\n    USE_DANN = False\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n    USE_MIXSTYLE = True\n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    # --- POSTPROCESS (inference engineering) --------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- CLAHE ---------------------------------------------------------------\n    CLAHE_MODE = \"global_shared\"     # \"off\" | \"per_slice\" | \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization: multi-slice sampled stats, not just the mid slice ---\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n    frag_stats_sample_slices = 9     # V6: sample across depth, not just mid slice\n\n    # --- domain-randomization augmentation (unchanged from V4) -------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n    USE_SHADOW = True\n    shadow_p = 0.20\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15      # distractor-only fake ink (no label update) -- kept\n                            # for robustness training, separate from physical ink\n    fake_fiber_p = 0.15\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- V7: CutMix (batch-level, orthogonal to the existing per-patch augs) -\n    USE_CUTMIX = True\n    cutmix_p = 0.20\n    cutmix_alpha = 1.0\n\n    # --- V7: curriculum sampling (easy -> full distribution over N epochs) --\n    USE_CURRICULUM = True\n    curriculum_warmup_epochs = 8\n    curriculum_easy_ink_frac = 0.10     # >= this ink fraction counts \"easy\"\n\n    # --- V7: cheap frequency-domain feature (radial high-freq FFT energy) ---\n    USE_FREQUENCY_FEATURES = True\n\n    # --- V7: Test-Time Training -- PROTOCOL=\"transductive\" ONLY. Never mixed\n    # into the strict-protocol LOFO-CV numbers (see V7 change-log docstring).\n    USE_TEST_TIME_TRAINING = False\n    ttt_steps = 20\n    ttt_lr = 1e-5\n    ttt_batch_size = 4\n\n    USE_EMA = True\n    ema_decay = 0.999\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    viz_dir = os.path.join(out_dir, \"v6_visualizations\")\n\n    # --- V6: what to run -----------------------------------------------------\n    # \"lofo_cv\"   : 3-fold leave-one-fragment-out CV (main scientific result)\n    # \"ablation\"  : component ablation matrix on ONE fold\n    # \"single\"    : one train_frags/test_frag run (fast debugging)\n    RUN_MODE = \"lofo_cv\"\n    single_train_frags = [\"2\", \"3\"]\n    single_test_frag = \"1\"\n    ablation_test_frag = \"1\"\n    ablation_train_frags = [\"2\", \"3\"]\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\nclass CfgOverride:\n    \"\"\"Context manager: temporarily overrides CFG attributes, restores them on\n    exit. Used by the ablation harness so each arm is a clean, reproducible\n    CFG state rather than hand-editing globals between runs.\"\"\"\n    def __init__(self, **overrides):\n        self.overrides = overrides\n        self.previous = {}\n\n    def __enter__(self):\n        for k, v in self.overrides.items():\n            self.previous[k] = getattr(CFG, k)\n            setattr(CFG, k, v)\n        return CFG\n\n    def __exit__(self, *exc):\n        for k, v in self.previous.items():\n            setattr(CFG, k, v)\n        return False\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\ndef compute_sample_difficulty(samples, labels_full, patch_size, easy_ink_frac):\n    \"\"\"V7: per-sample difficulty tier for curriculum sampling -- 0.0 (easy:\n    wide ink coverage), 0.5 (medium), 1.0 (hard: thin/sparse ink or none).\n    Index-aligned with `samples` (the same list used to build the training\n    Dataset), so it can be used directly as WeightedRandomSampler weights.\"\"\"\n    difficulties = np.zeros(len(samples), dtype=np.float32)\n    for i, (fid, y, x) in enumerate(samples):\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        if frac >= easy_ink_frac:\n            difficulties[i] = 0.0\n        elif frac > 0:\n            difficulties[i] = 0.5\n        else:\n            difficulties[i] = 1.0\n    return difficulties\n\n\ndef build_curriculum_sampler(difficulties, epoch, total_epochs, warmup_epochs):\n    \"\"\"Returns per-sample WeightedRandomSampler weights. Early epochs\n    strongly favor easy/medium samples; by `warmup_epochs` the weighting has\n    linearly relaxed to uniform (i.e. the full, already-class-balanced\n    distribution `balance_positive_patches` built) -- curriculum learning is\n    meant to warm the model up, not to permanently exclude hard examples.\"\"\"\n    progress = min(epoch / max(warmup_epochs, 1), 1.0)   # 0 -> 1 over warmup\n    # weight(difficulty=1.0) goes from a small floor up to 1.0 (uniform) as\n    # progress -> 1; weight(difficulty=0.0) stays at 1.0 throughout.\n    hard_floor = 0.15\n    weights = 1.0 - (1.0 - hard_floor) * (1.0 - progress) * difficulties\n    weights = np.clip(weights, hard_floor, 1.0)\n    return torch.DoubleTensor(weights)\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME (extended: depth pool for shift-invariance training)\n# ============================================================\n\nclass FragmentVolume:\n    \"\"\"V6 change: opens every slice in CFG.depth_pool_indices (the default\n    26-slice window PLUS +/- depth_shift_max on each side), not just\n    CFG.depth_indices. read_patch() accepts an explicit z_indices list so the\n    depth-shift-consistency training step can request a shifted window of the\n    SAME physical stack without re-opening files.\"\"\"\n\n    def __init__(self, frag_dir, pool_indices, default_indices):\n        self.pool_indices = pool_indices\n        self.default_indices = default_indices\n        self.index_to_pos = {z: i for i, z in enumerate(pool_indices)}\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in pool_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid_pos = self.index_to_pos.get(\n            self.default_indices[len(self.default_indices) // 2],\n            len(self._slices) // 2)\n        ref_slice = np.asarray(self._slices[mid_pos])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(\n            ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def _slice_positions(self, z_indices):\n        positions = []\n        for z in z_indices:\n            z_clamped = min(max(z, self.pool_indices[0]), self.pool_indices[-1])\n            positions.append(self.index_to_pos[z_clamped])\n        return positions\n\n    def read_patch(self, y, x, size, z_indices=None, apply_clahe=None):\n        self._ensure_open()\n        z_indices = self.default_indices if z_indices is None else z_indices\n        positions = self._slice_positions(z_indices)\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(positions), size, size), dtype=np.uint8)\n        for out_i, pos in enumerate(positions):\n            s = self._slices[pos]\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[out_i] = block\n        return out\n\n    def sample_shifted_indices(self, max_shift):\n        \"\"\"Returns a physically-shifted (but still contiguous, still ordered)\n        window of depth indices, clamped to stay inside the opened pool.\"\"\"\n        delta = random.randint(-max_shift, max_shift)\n        shifted = [z + delta for z in self.default_indices]\n        lo, hi = self.pool_indices[0], self.pool_indices[-1]\n        shifted = [min(max(z, lo), hi) for z in shifted]\n        return shifted\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        \"\"\"V6: samples statistics across several depth slices (not just the mid\n        slice) for a more robust per-fragment normalization constant.\"\"\"\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        n_slices = min(CFG.frag_stats_sample_slices, len(self.default_indices))\n        sample_z = sorted(random.sample(self.default_indices, n_slices))\n        sample_positions = self._slice_positions(sample_z)\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            for pos in sample_positions:\n                block = self._slices[pos][y:y + patch_size, x:x + patch_size]\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n                vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def warm_up_cache(self):\n        \"\"\"V7.5: forces every opened slice into the OS page cache with ONE\n        fast SEQUENTIAL read, up front. This is the actual fix for the\n        dataset-profile evidence (`main_patch_read` ranging 9ms-4880ms for\n        the same kind of read): patches are small, scattered, random-access\n        windows into 26 separate TIFF files, and on Kaggle's backing\n        filesystem each cold access to a new region can cost hundreds of ms\n        to multiple seconds. A single sequential pass per slice is a\n        completely different (much faster) I/O pattern, and afterwards every\n        patch read becomes a cache hit -- with train_stride=128 and\n        patch_size=320, patches overlap heavily anyway, so this isn't wasted\n        work. `.sum()` is used purely to force every byte to be read; the\n        result is discarded. This touches the OS-level page cache (not a\n        Python-owned buffer), so it's reclaimed automatically under memory\n        pressure rather than risking an OOM the way permanently loading\n        everything into a retained array would.\"\"\"\n        self._ensure_open()\n        t0 = time.time()\n        total_bytes = 0\n        for arr in self._slices:\n            _ = np.asarray(arr).sum(dtype=np.int64)\n            total_bytes += arr.nbytes\n        elapsed = time.time() - t0\n        rate = total_bytes / 1e6 / max(elapsed, 1e-6)\n        print(f\"  [cache warm-up] {len(self._slices)} slices, {total_bytes/1e9:.2f} GB read \"\n              f\"sequentially in {elapsed:.1f}s ({rate:.1f} MB/s) -- subsequent random-access \"\n              f\"patch reads should now mostly be cache hits.\")\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\ndef make_fragment_volume(frag_dir):\n    return FragmentVolume(frag_dir, CFG.depth_pool_indices, CFG.depth_indices)\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Distractor-only fake ink: perturbs intensity but NEVER touches the\n    label. Used purely as a robustness/negative-hallucination stress test --\n    kept separate from the physics-informed version below, which DOES update\n    the label because it is meant to represent an actual (simulated) ink\n    deposit, not a distractor.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_physical_ink_with_label(img_hwd, label_hw, depth_indices, n_strokes=None):\n    \"\"\"V6: the physics-informed synthetic ink model actually promised in V5's\n    CFG. Models a stroke's cross-depth intensity profile as a Gaussian\n    I(z) = A * exp(-(z - z0)^2 / (2*sigma^2)) with A, sigma, z0 sampled per\n    stroke, applies it as a CONTINUOUS per-slice weight (not a hard \"these N\n    slices are affected\" cutoff), and -- unlike inject_fake_ink -- writes the\n    stroke into the label as well, since this is meant to represent a\n    simulated real ink deposit rather than a distractor.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    label_out = label_hw.copy()\n    n_strokes = n_strokes or random.randint(1, 2)\n    z_arr = np.asarray(depth_indices, dtype=np.float32)\n\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n\n        A_amp = random.uniform(CFG.physical_ink_min_amplitude, CFG.physical_ink_max_amplitude)\n        sign = random.choice([-1, 1])\n        sigma = random.uniform(*CFG.physical_ink_sigma_range)\n        z0 = random.uniform(z_arr.min(), z_arr.max())\n\n        profile = sign * A_amp * np.exp(-((z_arr - z0) ** 2) / (2.0 * sigma ** 2))  # (d,)\n\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        stroke_bool = stroke_mask > 0\n        if not stroke_bool.any():\n            continue\n\n        for zi in range(d):\n            delta = profile[zi]\n            if abs(delta) < 0.5:\n                continue\n            sl = out[:, :, zi].astype(np.float32)\n            sl[stroke_bool] = np.clip(sl[stroke_bool] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n\n        # Only mark label where the Gaussian profile has real amplitude near\n        # the profile's peak (>= 40% of |A|) -- a near-zero-weight slice at\n        # the tail of the Gaussian isn't meaningfully \"ink\" at that slice, but\n        # the 2D label is depth-collapsed anyway, so we mark the stroke\n        # footprint whenever the profile is non-trivial anywhere in depth.\n        if np.abs(profile).max() >= 0.4 * A_amp:\n            label_out[stroke_bool] = 1\n\n    return out, label_out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\ndef cutmix_batch(imgs, masks):\n    \"\"\"V7: batch-level CutMix on the already-loaded tensors (image + label\n    mask cut together, so this stays label-consistent unlike a naive\n    image-only cutmix). Orthogonal to the existing per-patch augmentations\n    (which perturb ONE patch); this mixes TWO patches within a batch.\"\"\"\n    B = imgs.size(0)\n    if B < 2:\n        return imgs, masks\n    lam = float(np.random.beta(CFG.cutmix_alpha, CFG.cutmix_alpha))\n    rand_index = torch.randperm(B, device=imgs.device)\n\n    H, W = imgs.shape[-2:]\n    cut_ratio = math.sqrt(max(1.0 - lam, 1e-6))\n    cut_h, cut_w = int(H * cut_ratio), int(W * cut_ratio)\n    cy, cx = np.random.randint(H), np.random.randint(W)\n    y1, y2 = max(0, cy - cut_h // 2), min(H, cy + cut_h // 2)\n    x1, x2 = max(0, cx - cut_w // 2), min(W, cx + cut_w // 2)\n\n    imgs = imgs.clone()\n    masks = masks.clone()\n    imgs[:, :, y1:y2, x1:x2] = imgs[rand_index][:, :, y1:y2, x1:x2]\n    masks[:, :, y1:y2, x1:x2] = masks[rand_index][:, :, y1:y2, x1:x2]\n    return imgs, masks\n\n\ndef add_fiber_pattern_tensor(img_tensor, amplitude):\n    \"\"\"Tensor-level fiber perturbation (in normalized z-scored units) used by\n    the fiber-consistency loss so the perturbation stays inside the autograd\n    graph without a second numpy round-trip. img_tensor: (B, D, H, W).\"\"\"\n    B, D, H, W = img_tensor.shape\n    device = img_tensor.device\n    theta = torch.rand(B, device=device) * math.pi\n    freq = 0.02 + torch.rand(B, device=device) * 0.04\n    amp = amplitude * (0.5 + torch.rand(B, device=device))\n\n    yy, xx = torch.meshgrid(torch.arange(H, device=device, dtype=torch.float32),\n                             torch.arange(W, device=device, dtype=torch.float32), indexing=\"ij\")\n    yy = yy.unsqueeze(0)   # (1,H,W)\n    xx = xx.unsqueeze(0)\n\n    phase = (xx * torch.cos(theta).view(B, 1, 1) + yy * torch.sin(theta).view(B, 1, 1)) \\\n        * freq.view(B, 1, 1)\n    pattern = amp.view(B, 1, 1) * torch.sin(2 * math.pi * phase)   # (B,H,W)\n    pattern = pattern.unsqueeze(1).expand(-1, D, -1, -1)            # (B,D,H,W)\n    return img_tensor + pattern\n\n\n# ============================================================\n# 8. AUGMENTATION\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    \"\"\"V6 changes: (1) physics-informed ink now updates the label and is\n    applied via inject_physical_ink_with_label, kept independent from the\n    label-blind inject_fake_ink distractor; (2) optionally returns a\n    depth-shifted alternate view of the same patch for the shift-consistency\n    loss (train_mode only).\"\"\"\n\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n        self._profile_calls_left = CFG.PROFILE_DATASET_CALLS if CFG.PROFILE_DATASET else 0\n\n    def __len__(self):\n        return len(self.samples)\n\n    def _prep(self, img_u8_hwd, label_hw, normalize_stats):\n        img = img_u8_hwd.astype(np.float32) / 255.0\n        img = normalize_patch(img, *normalize_stats)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label_hw > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n    def __getitem__(self, idx):\n        # V7.4: per-stage CPU timing, printed for the first\n        # CFG.PROFILE_DATASET_CALLS calls in EACH worker process (so with\n        # num_workers=2 you'll see ~2x that many lines total, interleaved --\n        # that's expected, not a bug). This exists to answer \"which\n        # augmentation is actually slow\" empirically instead of guessing;\n        # set CFG.PROFILE_DATASET=False once you've identified the culprit.\n        do_prof = self.train_mode and self._profile_calls_left > 0\n        if do_prof:\n            self._profile_calls_left -= 1\n            _t = time.time()\n            def _lap(label, timings=[]):\n                nonlocal _t\n                now = time.time()\n                timings.append((label, now - _t))\n                _t = now\n                return timings\n            _timings = []\n\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size].copy()\n        img = np.transpose(patch, (1, 2, 0))   # HWD\n        if do_prof:\n            _timings = _lap(\"main_patch_read\")\n\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n        if do_prof:\n            _timings = _lap(\"hist_match\")\n\n        if self.train_mode and CFG.USE_PHYSICAL_SYNTHETIC_INK and random.random() < CFG.physical_ink_p:\n            img, label = inject_physical_ink_with_label(img, label, CFG.depth_indices)\n        if do_prof:\n            _timings = _lap(\"physical_ink\")\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n        if do_prof:\n            _timings = _lap(\"fake_ink_fiber_shadow\")\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n        if do_prof:\n            _timings = _lap(\"albumentations_transform\")\n\n        img_t, label_t = self._prep(img, label, (vol.frag_mean, vol.frag_std))\n        if do_prof:\n            _timings = _lap(\"prep_to_tensor\")\n\n        shifted_t = None\n        if self.train_mode and CFG.USE_DEPTH_SHIFT_CONSISTENCY and random.random() < CFG.depth_shift_p:\n            shifted_z = vol.sample_shifted_indices(CFG.depth_shift_max)\n            shifted_patch = vol.read_patch(y, x, size, z_indices=shifted_z)\n            shifted_hwd = np.transpose(shifted_patch, (1, 2, 0)).astype(np.float32) / 255.0\n            shifted_hwd = normalize_patch(shifted_hwd, vol.frag_mean, vol.frag_std)\n            shifted_t = torch.from_numpy(np.ascontiguousarray(np.transpose(shifted_hwd, (2, 0, 1))))\n        if do_prof:\n            _timings = _lap(\"depth_shift_patch_read\")\n\n        if shifted_t is None:\n            has_shift = torch.tensor(False)\n            shifted_t = torch.zeros_like(img_t)\n        else:\n            has_shift = torch.tensor(True)\n\n        if do_prof:\n            total = sum(t for _, t in _timings)\n            breakdown = \" | \".join(f\"{name}={t*1000:.1f}ms\" for name, t in _timings)\n            print(f\"    [dataset profile pid={os.getpid()}] TOTAL={total*1000:.1f}ms :: {breakdown}\")\n\n        return img_t, label_t, shifted_t, has_shift\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN (Protocol B only): yields normalized (D,H,W) tensors from the\n    held-out fragment, NO labels.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"V4's simpler depth-aware stem, kept as an ablation arm / fallback.\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = 0\n        if use_raw:\n            total_in += in_depth\n        if use_grad:\n            total_in += (in_depth - 1)\n        if use_curv:\n            total_in += max(in_depth - 2, 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.in_depth = in_depth\n\n    def forward(self, x):\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:, :, :] - x[:, :-1, :, :])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass Conv3DDepthStem(nn.Module):\n    \"\"\"V7 addition: a genuinely higher-capacity alternative depth stem,\n    addressing the '26->16 is a severe information bottleneck' critique with\n    real extra convolutional capacity rather than just a wider 1x1 mix. Two\n    branches (raw depth stack, first-difference depth stack) each go through\n    two 3D conv layers BEFORE any depth-collapsing, then are mean-pooled over\n    depth and concatenated -- unlike a single 1x1 conv straight from 70\n    finite-difference channels down to 16, the 3D convs get real learnable\n    interaction across depth and space before anything is collapsed. This is\n    wired in as its own `depth_module_type=\"conv3d\"` ablation arm alongside\n    DepthSignatureModule, not a replacement for it -- they represent two\n    different hypotheses (attention-based depth pooling vs. convolutional\n    depth pooling) worth comparing empirically, which is exactly what the\n    ablation matrix is for.\"\"\"\n    def __init__(self, in_depth, out_channels=48):\n        super().__init__()\n        half = out_channels // 2\n        self.raw_branch = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=3, padding=1), nn.BatchNorm3d(16), nn.GELU(),\n            nn.Conv3d(16, half, kernel_size=3, padding=1), nn.BatchNorm3d(half), nn.GELU(),\n        )\n        self.grad_branch = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=3, padding=1), nn.BatchNorm3d(16), nn.GELU(),\n            nn.Conv3d(16, out_channels - half, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels - half), nn.GELU(),\n        )\n        self.out_channels = out_channels\n\n    def forward(self, x):   # x: (B, D, H, W)\n        raw = x.unsqueeze(1)                                  # (B,1,D,H,W)\n        grad = (x[:, 1:, :, :] - x[:, :-1, :, :]).unsqueeze(1)  # (B,1,D-1,H,W)\n\n        raw_feat = self.raw_branch(raw).mean(dim=2)            # (B,half,H,W)\n        grad_feat = self.grad_branch(grad).mean(dim=2)         # (B,out-half,H,W)\n        return torch.cat([raw_feat, grad_feat], dim=1)\n\n\nclass LDDC(nn.Module):\n    \"\"\"Learnable Depth Differential Convolution. Learns `num_filters` 1D\n    kernels along the physically-ordered depth axis. Each kernel is\n    re-parameterized every forward pass to sum to zero:\n        w' = w - mean(w)\n    which is a hard constraint (not a training-time regularizer that might\n    only be approximately satisfied) -- guaranteeing every filter behaves like\n    a (learned) derivative operator regardless of what the raw weights drift\n    to during optimization. Implemented as a Conv3d with kernel (k,1,1) over\n    x.unsqueeze(1): (B,1,D,H,W) -> (B,num_filters,D,H,W).\n    \"\"\"\n    def __init__(self, num_filters=4, kernel_size=3):\n        super().__init__()\n        self.num_filters = num_filters\n        self.kernel_size = kernel_size\n        self.weight = nn.Parameter(torch.randn(num_filters, 1, kernel_size, 1, 1) * 0.1)\n        self.bias = nn.Parameter(torch.zeros(num_filters))\n\n    def forward(self, x):   # x: (B, D, H, W)\n        w = self.weight - self.weight.mean(dim=2, keepdim=True)   # zero-sum constraint\n        xin = x.unsqueeze(1)   # (B,1,D,H,W)\n        out = F.conv3d(xin, w, bias=self.bias, padding=(self.kernel_size // 2, 0, 0))\n        return out   # (B, num_filters, D, H, W)\n\n\nclass DepthPositionalEncoding(nn.Module):\n    \"\"\"Sinusoidal encoding of each physical depth position z, broadcast\n    spatially. Returns (B, pe_dim, D, H, W) so it can be concatenated\n    alongside LDDC's depth-preserving output before depth attention pools it\n    down to a 2D feature map.\"\"\"\n    def __init__(self, pe_dim=8):\n        super().__init__()\n        assert pe_dim % 2 == 0\n        self.pe_dim = pe_dim\n        div_term = torch.exp(torch.arange(0, pe_dim, 2).float() * (-math.log(10000.0) / pe_dim))\n        self.register_buffer(\"div_term\", div_term, persistent=False)\n\n    def forward(self, depth_positions, B, H, W, device):\n        # depth_positions: 1D float tensor of physical z indices, length D\n        z = depth_positions.to(device).view(-1, 1)                     # (D,1)\n        angles = z * self.div_term.view(1, -1).to(device)               # (D, pe_dim/2)\n        pe = torch.zeros(z.shape[0], self.pe_dim, device=device)\n        pe[:, 0::2] = torch.sin(angles)\n        pe[:, 1::2] = torch.cos(angles)\n        # (D, pe_dim) -> (1, pe_dim, D, 1, 1) -> broadcast to (B, pe_dim, D, H, W)\n        pe = pe.transpose(0, 1).view(1, self.pe_dim, -1, 1, 1)\n        return pe.expand(B, -1, -1, H, W)\n\n\nclass PooledDepthAttention(nn.Module):\n    \"\"\"Multi-head attention ACROSS the depth axis, at every spatial location.\n\n    Honest engineering note (this is exactly the kind of thing the review\n    flagged as needing to be stated explicitly): true per-pixel attention\n    across D=26 depth positions needs an (H, W, heads, D, D) tensor. At\n    patch_size=480 with heads=4, D=26, batch=8 that is B*H*W*heads*D*D*4 bytes\n    ~= 2 TB of activation memory -- not something a single Tesla T4 (14.5GB)\n    can hold, full stop. So this module computes genuine per-pixel softmax\n    attention across depth at a POOLED spatial resolution (default 480/8=60),\n    where the memory cost (8*60*60*4*26*26*4 bytes ~= 93MB) is trivial, and\n    then bilinearly upsamples the resulting per-position depth-attention\n    output back to full resolution. This still gives spatially-varying,\n    per-location depth attention (unlike a pure squeeze-and-excite global\n    version), just not literally per-pixel -- state this trade-off in the\n    methods section rather than letting the module name imply otherwise.\n    \"\"\"\n    def __init__(self, in_channels, heads=4, pool_size=8):\n        super().__init__()\n        self.heads = heads\n        self.pool_size = pool_size\n        self.head_dim = max(in_channels // heads, 4)\n        inner = self.heads * self.head_dim\n        self.to_q = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.to_k = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.to_v = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.out_proj = nn.Conv1d(inner, in_channels, kernel_size=1)\n\n    def forward(self, x):   # x: (B, C, D, H, W)\n        B, C, D, H, W = x.shape\n        ph = max(1, round(H / self.pool_size))\n        pw = max(1, round(W / self.pool_size))\n        x_pooled = F.adaptive_avg_pool3d(x, output_size=(D, max(1, H // ph), max(1, W // pw)))\n        _, _, _, ph_, pw_ = x_pooled.shape\n\n        # reshape depth axis into the \"sequence\" dimension for a 1D attention\n        # per spatial location: (B*ph_*pw_, C, D)\n        xp = x_pooled.permute(0, 3, 4, 1, 2).reshape(B * ph_ * pw_, C, D)\n        q = self.to_q(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n        k = self.to_k(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n        v = self.to_v(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n\n        attn = torch.einsum(\"nhdi,nhdj->nhij\", q, k) / math.sqrt(self.head_dim)\n        attn = attn.softmax(dim=-1)                                    # (N, heads, D, D)\n        out = torch.einsum(\"nhij,nhdj->nhdi\", attn, v)                 # (N, heads, head_dim, D)\n        out = out.reshape(B * ph_ * pw_, self.heads * self.head_dim, D)\n        out = self.out_proj(out)                                       # (N, C, D)\n        out = out.view(B, ph_, pw_, C, D).permute(0, 3, 4, 1, 2)        # (B, C, D, ph_, pw_)\n\n        out = out.reshape(B, C * D, ph_, pw_)\n        out = F.interpolate(out, size=(H, W), mode=\"bilinear\", align_corners=False)\n        out = out.view(B, C, D, H, W)\n        return out\n\n\nclass DepthStatsBranch(nn.Module):\n    \"\"\"Analytic (non-learned) depth-statistics channels: mean, std, max,\n    depth-centroid (intensity-weighted mean z), gradient energy, curvature\n    energy, and (V7, optional) radial high-frequency FFT energy. Computed\n    directly from the raw depth stack so they carry signal even before\n    LDDC/attention have learned anything useful early in training.\n\n    V7's frequency channel: thin ink strokes contribute disproportionately to\n    high spatial frequencies compared to broad fiber/background texture, so a\n    single cheap FFT magnitude computed on the depth-AVERAGED image (not a\n    separate FFT per slice -- that would multiply memory cost for unclear\n    extra benefit) gives one extra, genuinely informative analytic channel.\n    \"\"\"\n    def __init__(self, use_frequency=False):\n        super().__init__()\n        self.use_frequency = use_frequency\n        self.out_channels = 6 + (1 if use_frequency else 0)\n\n    @staticmethod\n    def _radial_high_freq_energy(mean_img, high_freq_frac=0.5):\n        # mean_img: (B, 1, H, W). Returns (B, 1, H, W) -- the same scalar\n        # (per-sample high-frequency energy fraction) broadcast spatially, so\n        # it can be concatenated as a \"channel\" alongside genuinely spatial\n        # stats without pretending to carry spatial variation it doesn't have.\n        B, _, H, W = mean_img.shape\n        fft = torch.fft.rfft2(mean_img.squeeze(1).float(), norm=\"ortho\")\n        mag = torch.abs(fft)   # (B, H, W//2+1)\n        fy = torch.fft.fftfreq(H, device=mean_img.device).view(H, 1)\n        fx = torch.fft.rfftfreq(W, device=mean_img.device).view(1, -1)\n        radius = torch.sqrt(fy ** 2 + fx ** 2)\n        radius = radius / radius.max().clamp_min(1e-6)\n        high_mask = (radius >= high_freq_frac).float()\n        high_energy = (mag * high_mask).sum(dim=(1, 2))\n        total_energy = mag.sum(dim=(1, 2)).clamp_min(1e-6)\n        frac = (high_energy / total_energy).view(B, 1, 1, 1).expand(-1, 1, H, W)\n        return frac.to(mean_img.dtype)\n\n    def forward(self, x, depth_positions):   # x: (B, D, H, W)\n        B, D, H, W = x.shape\n        mean = x.mean(dim=1, keepdim=True)\n        std = x.std(dim=1, keepdim=True)\n        maxv = x.max(dim=1, keepdim=True).values\n\n        z = depth_positions.to(x.device).view(1, D, 1, 1)\n        weights = F.softmax(x, dim=1)\n        centroid = (weights * z).sum(dim=1, keepdim=True)\n\n        grad = x[:, 1:, :, :] - x[:, :-1, :, :]\n        grad_energy = (grad ** 2).mean(dim=1, keepdim=True)\n\n        if D > 2:\n            curv = x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :]\n            curv_energy = (curv ** 2).mean(dim=1, keepdim=True)\n        else:\n            curv_energy = torch.zeros_like(mean)\n\n        parts = [mean, std, maxv, centroid, grad_energy, curv_energy]\n        if self.use_frequency:\n            parts.append(self._radial_high_freq_energy(mean))\n        return torch.cat(parts, dim=1)   # (B, 6 or 7, H, W)\n\n\nclass DepthSignatureModule(nn.Module):\n    \"\"\"V5's promised (and now actually implemented) depth-aware front end:\n    LDDC + depth positional encoding + pooled multi-head depth attention +\n    analytic depth-statistics channels, mixed down to `out_channels` for the\n    2D encoder. Includes the MC-dropout channel used both for regularization\n    during training and for genuine uncertainty estimation at inference (see\n    run_mc_dropout_uncertainty). Gradient checkpointing is applied to the\n    (D-preserving, memory-heavy) LDDC+PE+attention stack when\n    USE_DEPTH_SIGNATURE_CHECKPOINT is True.\"\"\"\n    def __init__(self, in_depth, depth_positions, out_channels=24):\n        super().__init__()\n        self.in_depth = in_depth\n        self.register_buffer(\"depth_positions\", torch.tensor(depth_positions, dtype=torch.float32),\n                              persistent=False)\n\n        self.lddc = LDDC(CFG.lddc_num_filters, CFG.lddc_kernel_size)\n        self.pe = DepthPositionalEncoding(CFG.depth_pe_dim)\n        lddc_pe_channels = CFG.lddc_num_filters + CFG.depth_pe_dim\n        self.attn = PooledDepthAttention(lddc_pe_channels, heads=CFG.depth_attention_heads,\n                                          pool_size=CFG.depth_attention_pool)\n        self.stats = DepthStatsBranch(use_frequency=CFG.USE_FREQUENCY_FEATURES) if CFG.USE_DEPTH_STATS else None\n\n        depth_collapsed_channels = lddc_pe_channels * in_depth       # attention output, D collapsed via mean\n        stats_channels = self.stats.out_channels if self.stats is not None else 0\n        mix_in = depth_collapsed_channels + stats_channels + in_depth  # + raw stack\n\n        self.dropout = nn.Dropout2d(CFG.depth_signature_dropout_p)\n        self.mix = nn.Sequential(\n            nn.Conv2d(mix_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n\n    def _lddc_pe_attn(self, x):\n        B, D, H, W = x.shape\n        lddc_out = self.lddc(x)                                          # (B, F, D, H, W)\n        pe_out = self.pe(self.depth_positions, B, H, W, x.device)          # (B, pe, D, H, W)\n        combined = torch.cat([lddc_out, pe_out], dim=1)                    # (B, F+pe, D, H, W)\n        attended = self.attn(combined)                                     # (B, F+pe, D, H, W)\n        # collapse depth by taking mean over the attended depth axis for the\n        # final 2D feature map, but keep the FULL (F+pe)*D as concatenated\n        # channels too -- mean loses information, so we use both.\n        collapsed = attended.reshape(B, -1, H, W)                          # (B, (F+pe)*D, H, W)\n        return collapsed\n\n    def forward(self, x):   # x: (B, D, H, W)\n        if CFG.USE_DEPTH_SIGNATURE_CHECKPOINT and self.training:\n            collapsed = grad_checkpoint(self._lddc_pe_attn, x, use_reentrant=False)\n        else:\n            collapsed = self._lddc_pe_attn(x)\n\n        parts = [x, collapsed]\n        if self.stats is not None:\n            parts.append(self.stats(x, self.depth_positions))\n\n        feat = torch.cat(parts, dim=1)\n        feat = self.dropout(feat)\n        return self.mix(feat)\n\n\nclass DepthAwareSegModel(nn.Module):\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):   # x: (B, D, H, W)\n        feat = self.depth_stem(x) if self.depth_stem is not None else x\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' requested, but only \"\n              f\"'unet' has been verified to handle tu-* timm encoders' 0-channel placeholder \"\n              f\"stage correctly -- using UnetPlusPlus may crash with ConvNeXt/other tu-* encoders.\")\n    return arch_cls(\n        encoder_name=encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=in_channels,\n        classes=1,\n        decoder_channels=CFG.decoder_channels,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\ndef build_model():\n    stem = None\n    seg_in_channels = CFG.in_channels\n\n    if CFG.depth_module_type == \"signature\":\n        stem = DepthSignatureModule(CFG.in_channels, CFG.depth_indices, CFG.depth_signature_out_channels)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] DepthSignatureModule: {CFG.in_channels} depth slices -> \"\n              f\"{seg_in_channels} learned channels (LDDC={CFG.lddc_num_filters} filters, \"\n              f\"PE dim={CFG.depth_pe_dim}, attention heads={CFG.depth_attention_heads}, \"\n              f\"pooled to {CFG.depth_attention_pool}x{CFG.depth_attention_pool}, \"\n              f\"stats={CFG.USE_DEPTH_STATS})\")\n    elif CFG.depth_module_type == \"fusion\":\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] DepthFusionStem: {CFG.in_channels} -> {seg_in_channels} channels\")\n    elif CFG.depth_module_type == \"conv3d\":\n        stem = Conv3DDepthStem(CFG.in_channels, CFG.conv3d_stem_out_channels)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] Conv3DDepthStem: {CFG.in_channels} -> {seg_in_channels} channels \"\n              f\"(two 3D-conv branches, depth-mean-pooled)\")\n    elif CFG.depth_module_type == \"none\":\n        stem = None\n        seg_in_channels = CFG.in_channels\n        print(f\"[backbone] no depth-aware stem: raw {CFG.in_channels}-channel stack into encoder\")\n    else:\n        raise ValueError(f\"Unknown depth_module_type={CFG.depth_module_type}\")\n\n    try:\n        seg_model = _build_seg_backbone(CFG.encoder_name, seg_in_channels)\n        print(f\"[backbone] using {CFG.encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {CFG.encoder_name} unavailable/failed ({e}); \"\n              f\"falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model)\n\n\nclass EMAModel:\n    def __init__(self, model, decay=CFG.ema_decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. ARCHITECTURE INSPECTION\n# ============================================================\n\ndef inspect_model_channels(model):\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\n@torch.no_grad()\ndef run_architecture_report(model):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"V6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\")\n    print(\"=\" * 70)\n\n    n_params = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Input:  depth slices={CFG.in_channels} | patch={CFG.patch_size}x{CFG.patch_size}\")\n    print(f\"depth_module_type={CFG.depth_module_type} | encoder={CFG.encoder_name} | \"\n          f\"architecture={CFG.architecture}\")\n    print(f\"Parameters: {n_params/1e6:.2f}M total | {n_trainable/1e6:.2f}M trainable\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"\\nZERO/INVALID CHANNEL LAYERS FOUND:\")\n        for item in bad_layers:\n            print(\"  \", item)\n        raise RuntimeError(\"Model contains invalid zero-channel layers -- fix before training.\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    model.eval()\n    small_size = min(CFG.patch_size, 256)\n    dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n    out = model(dummy)\n    print(f\"Dry-run ({small_size}x{small_size}) output logits shape: {tuple(out.shape)}\")\n    has_nan = torch.isnan(out).any().item()\n    has_inf = torch.isinf(out).any().item()\n    print(f\"NaN check: {'FAIL' if has_nan else 'PASS'} | Inf check: {'FAIL' if has_inf else 'PASS'}\")\n    if has_nan or has_inf:\n        raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass -- fix before training.\")\n    model.train()\n\n    if torch.cuda.is_available():\n        print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")\n    print(\"Forward pass: PASS\")\n    print(\"=\" * 70 + \"\\n\")\n\n\n# ============================================================\n# 12. LOSSES (BCE + Dice + FocalTversky + clDice)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\ndef _soft_erode(img):\n    return -F.max_pool2d(-img, kernel_size=3, stride=1, padding=1)\n\n\ndef _soft_dilate(img):\n    return F.max_pool2d(img, kernel_size=3, stride=1, padding=1)\n\n\ndef _soft_open(img):\n    return _soft_dilate(_soft_erode(img))\n\n\ndef soft_skeletonize(img, iters):\n    \"\"\"Differentiable soft-skeletonization (Shit et al. 2021, 'clDice -- A\n    Novel Topology-Preserving Loss Function for Tubular Structure\n    Segmentation'), used for the clDice topology-aware loss below.\"\"\"\n    img1 = _soft_open(img)\n    skel = F.relu(img - img1)\n    for _ in range(iters):\n        img = _soft_erode(img)\n        img1 = _soft_open(img)\n        delta = F.relu(img - img1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef cldice_loss(logits, targets, iters=8, eps=1e-6):\n    probs = torch.sigmoid(logits)\n    skel_pred = soft_skeletonize(probs, iters)\n    skel_true = soft_skeletonize(targets, iters)\n    t_prec = (skel_pred * targets).sum() / (skel_pred.sum() + eps)\n    t_sens = (skel_true * probs).sum() / (skel_true.sum() + eps)\n    cldice = 1.0 - (2.0 * t_prec * t_sens) / (t_prec + t_sens + eps)\n    return cldice\n\n\nclass VXComboLoss(nn.Module):\n    \"\"\"BCE + Dice + FocalTversky, with clDice added when USE_TOPOLOGY_LOSS is\n    on -- unlike V5's CFG, topology_weight now actually reaches the backward\n    graph.\"\"\"\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        loss = CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n        cldice_val = None\n        if CFG.USE_TOPOLOGY_LOSS:\n            cldice_val = cldice_loss(logits, targets, CFG.cldice_iters)\n            loss = loss + CFG.topology_weight * cldice_val\n        return loss, {\"bce\": bce.item(), \"dice\": dice.item(), \"focal_tversky\": tv.item(),\n                       \"cldice\": (cldice_val.item() if cldice_val is not None else None)}\n\n\n# ============================================================\n# 13. METRICS  (F0.5 is the PRIMARY reported metric)\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25   # F0.5\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"f0.5\": float(fbeta), \"dice\": float(dice), \"iou\": float(iou),\n            \"precision\": float(precision), \"recall\": float(recall)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 14. TRAIN / VALIDATION EPOCH (fiber + depth-shift consistency wired in)\n# ============================================================\n\ndef run_epoch(model, criterion, loader, optimizer, scaler, ema, domain_classifier,\n              feature_capture, dann_iter, global_step_holder, total_steps,\n              train_mode=True, threshold=0.5, scheduler=None, scheduler_steps_per_batch=False,\n              profile_steps=0):\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_consistency_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    # V7.4 (essential, not cosmetic): without this, a slow epoch prints\n    # NOTHING between the profiled first `profile_steps` batches and the\n    # epoch-summary line -- exactly the \"goes dark for 90 minutes with no way\n    # to tell if it's working or stuck\" problem. This is separate from\n    # `profile_steps`: it runs every epoch (not just epoch 1), has no\n    # cuda.synchronize() overhead, and reports wall-clock throughput/ETA so a\n    # slow-but-alive run is distinguishable from a genuinely hung one.\n    epoch_t_start = time.time()\n    n_loader_batches = len(loader)\n    progress_every = max(1, n_loader_batches // 20)   # ~20 prints per epoch\n\n    t_data_end = time.time()\n    for batch_idx, (imgs, masks, shifted_imgs, has_shift) in enumerate(loader):\n        if train_mode and n_loader_batches > 0 and (batch_idx + 1) % progress_every == 0:\n            elapsed = time.time() - epoch_t_start\n            rate = (batch_idx + 1) / max(elapsed, 1e-6)\n            eta = (n_loader_batches - (batch_idx + 1)) / max(rate, 1e-6)\n            print(f\"    step {batch_idx + 1}/{n_loader_batches} | \"\n                  f\"{elapsed:.0f}s elapsed | {rate:.2f} steps/s | ETA {eta:.0f}s\")\n\n        do_profile = train_mode and profile_steps > 0 and batch_idx < profile_steps\n        if do_profile:\n            t0 = time.time()\n            data_wait = t0 - t_data_end\n            if CFG.device == \"cuda\":\n                torch.cuda.synchronize()\n\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        shifted_imgs = shifted_imgs.to(CFG.device, non_blocking=True)\n        has_shift = has_shift.to(CFG.device, non_blocking=True)\n\n        # V7.1 (perf): consistency losses only computed every Nth step -- see\n        # CFG.consistency_every_n_steps docstring for why.\n        do_consistency = train_mode and (batch_idx % max(CFG.consistency_every_n_steps, 1) == 0)\n\n        # V7.2 (bugfix, flagged by review): CutMix is applied to `imgs` here,\n        # but `shifted_imgs` comes straight from the dataset and is NEVER\n        # cutmixed (the Dataset builds it independently of this loop). If\n        # depth-shift consistency then compared model(shifted_imgs) against\n        # model(imgs) on a step where imgs got cutmixed, it would be\n        # comparing predictions on a DIFFERENT spatial composition -- not the\n        # same sample at a shifted depth window, which defeats the point of\n        # the consistency loss (and would train it against noise). Fiber\n        # consistency is unaffected: fiber_imgs is built FROM `imgs` after\n        # cutmix, so it stays the same sample as `logits`. Fix: track\n        # whether cutmix fired this step and skip ONLY depth-shift\n        # consistency when it did.\n        did_cutmix = False\n        if train_mode and CFG.USE_CUTMIX and random.random() < CFG.cutmix_p:\n            imgs, masks = cutmix_batch(imgs, masks)\n            did_cutmix = True\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss, loss_parts = criterion(logits, masks)\n                consistency_total = torch.zeros((), device=CFG.device)\n\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_main_fwd = time.time()\n\n                # --- fiber-consistency loss (real 2nd forward pass) ---------\n                # safe under CutMix: fiber_imgs is derived FROM imgs (the\n                # possibly-cutmixed tensor), so both sides of the comparison\n                # are the same sample.\n                if do_consistency and CFG.USE_FIBER_CONSISTENCY_LOSS:\n                    fiber_imgs = add_fiber_pattern_tensor(imgs, CFG.fiber_consistency_amplitude)\n                    fiber_logits = model(fiber_imgs)\n                    fiber_consistency = F.mse_loss(torch.sigmoid(fiber_logits), torch.sigmoid(logits.detach()))\n                    consistency_total = consistency_total + CFG.fiber_consistency_weight * fiber_consistency\n\n                # --- depth-shift consistency loss (real 2nd forward pass) ---\n                # NOT safe under CutMix (see comment above) -- skipped this step.\n                #\n                # V7.3 (critical perf bugfix): this USED to forward only the\n                # has_shift subset (`shifted_imgs[idx]`), whose size is random\n                # every time (0-8, since depth_shift_p=0.5 per sample). With\n                # cudnn.benchmark=True, EVERY new batch size the model sees\n                # triggers a fresh convolution-algorithm benchmark search --\n                # for this model (ConvNeXt+U-Net+3D depth stem) that can cost\n                # seconds to tens of seconds PER NEW SIZE, plus growing\n                # per-shape workspace memory. Over hundreds of steps hitting\n                # sizes 1..8 repeatedly in no particular order, this compounds\n                # into exactly the kind of multi-hour stall that doesn't show\n                # up in a short profiling window (the profiled steps 0 and 4\n                # already show this: 94s and 4s of one-off cost). Fixed: ALWAYS\n                # forward the full, fixed-size batch (same shape as the main/\n                # fiber paths, which is why THOSE stayed fast at ~0.27s/1.2s\n                # steady-state) and mask out the non-shifted samples in the\n                # loss instead of indexing them out of the tensor.\n                if do_consistency and (not did_cutmix) and CFG.USE_DEPTH_SHIFT_CONSISTENCY:\n                    shift_logits = model(shifted_imgs)   # fixed shape: (CFG.batch_size, D, H, W)\n                    shift_probs = torch.sigmoid(shift_logits)\n                    main_probs_detached = torch.sigmoid(logits.detach())\n                    per_sample_mse = F.mse_loss(shift_probs, main_probs_detached,\n                                                 reduction=\"none\").mean(dim=[1, 2, 3])\n                    weight = has_shift.float()\n                    denom = weight.sum().clamp_min(1.0)\n                    shift_consistency = (per_sample_mse * weight).sum() / denom\n                    consistency_total = consistency_total + CFG.depth_shift_consistency_weight * shift_consistency\n                    del shift_logits, shift_probs, main_probs_detached, per_sample_mse\n\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_consist_fwd = time.time()\n\n                loss = loss + consistency_total\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_backward = time.time()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    # V7.3: GradScaler can SKIP the actual optimizer.step() on\n                    # steps where it detects inf/nan gradients (routine during\n                    # AMP's initial loss-scale calibration, typically just the\n                    # first few iterations) -- if we call scheduler.step()\n                    # unconditionally after that, OneCycleLR advances one step\n                    # further than the optimizer actually did, which is what\n                    # the \"lr_scheduler.step() before optimizer.step()\"\n                    # warning is reporting. Comparing the scaler's scale\n                    # before/after detects a skipped step so we skip the\n                    # scheduler step too, keeping the two in sync.\n                    prev_scale = scaler.get_scale()\n                    scaler.step(optimizer)\n                    scaler.update()\n                    step_was_skipped = scaler.get_scale() < prev_scale\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    if scheduler_steps_per_batch and scheduler is not None and not step_was_skipped:\n                        scheduler.step()\n                    global_step_holder[0] += 1\n\n        if do_profile:\n            print(f\"  [profile step {batch_idx}] data_wait={data_wait:.3f}s \"\n                  f\"main_fwd={t_main_fwd - t0:.3f}s \"\n                  f\"consistency_fwd={t_consist_fwd - t_main_fwd:.3f}s \"\n                  f\"(consistency_computed={do_consistency}) \"\n                  f\"backward+step={t_backward - t_consist_fwd:.3f}s \"\n                  f\"TOTAL={t_backward - t0:.3f}s\")\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_consistency_loss += float(consistency_total.detach().item()) if train_mode else 0.0\n\n        del imgs, masks, logits, probs, shifted_imgs, has_shift\n        t_data_end = time.time()\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n        prev_scale = scaler.get_scale()\n        scaler.step(optimizer)\n        scaler.update()\n        step_was_skipped = scaler.get_scale() < prev_scale\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n        # V7.2: this trailing partial-accumulation step is also a real\n        # optimizer step -- OneCycleLR needs scheduler.step() called here\n        # too, or the last step of every epoch silently goes unaccounted for.\n        if scheduler_steps_per_batch and scheduler is not None and not step_was_skipped:\n            scheduler.step()\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_consistency = total_consistency_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_consistency\n\n\n# ============================================================\n# 15. THRESHOLD SEARCH (validation-only, frozen before test)\n# ============================================================\n\n@torch.no_grad()\ndef find_best_threshold(model, loader, metric=\"f0.5\"):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks, _, _ in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    prec = (tp + 1e-6) / (tp + fp + 1e-6)\n    rec = (tp + 1e-6) / (tp + fn + 1e-6)\n    f05 = (1.25 * prec * rec + 1e-6) / (0.25 * prec + rec + 1e-6)\n    scores = f05 if metric == \"f0.5\" else dice\n    best_idx = int(np.argmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), thresholds, scores\n\n\n# ============================================================\n# 16. MC-DROPOUT UNCERTAINTY (genuine multi-pass diagnostic)\n# ============================================================\n\ndef _set_dropout_train(model):\n    for m in model.modules():\n        if isinstance(m, (nn.Dropout, nn.Dropout2d, nn.Dropout3d)):\n            m.train()\n\n\n@torch.no_grad()\ndef run_mc_dropout_uncertainty(model, vol, mask, patch_size, stride, passes):\n    \"\"\"Runs `passes` stochastic forward passes (only Dropout layers stay in\n    train mode; BatchNorm/LayerNorm stay in eval mode) over a sliding window\n    and returns (mean_prob_map, uncertainty_map) where uncertainty is the\n    per-pixel variance across passes.\"\"\"\n    model.eval()\n    _set_dropout_train(model)\n\n    H, W = mask.shape\n    sum_map = np.zeros((H, W), dtype=np.float32)\n    sumsq_map = np.zeros((H, W), dtype=np.float32)\n    weight_map = np.zeros((H, W), dtype=np.float32)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        x_t = torch.from_numpy(raw).unsqueeze(0).to(CFG.device)\n        pass_probs = []\n        for _ in range(passes):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(x_t)\n            pass_probs.append(torch.sigmoid(logits)[0, 0].float().cpu().numpy())\n        pass_probs = np.stack(pass_probs, axis=0)   # (passes, size, size)\n        mean_p = pass_probs.mean(axis=0)\n        sq_p = (pass_probs ** 2).mean(axis=0)\n        sum_map[y:y + patch_size, x:x + patch_size] += mean_p\n        sumsq_map[y:y + patch_size, x:x + patch_size] += sq_p\n        weight_map[y:y + patch_size, x:x + patch_size] += 1.0\n        del x_t, pass_probs\n\n    weight_map[weight_map <= 1e-8] = 1.0\n    mean_prob = sum_map / weight_map\n    mean_sq = sumsq_map / weight_map\n    uncertainty = np.clip(mean_sq - mean_prob ** 2, 0, None)\n\n    model.eval()   # restore full eval mode (dropout off) for any subsequent calls\n    return mean_prob, uncertainty\n\n\ndef analyze_uncertainty_vs_error(prob_map, uncertainty_map, gt, mask, threshold, n_buckets=3):\n    \"\"\"Buckets pixels by uncertainty and reports precision in each bucket --\n    directly answers 'does uncertainty predict where the model is wrong?'\"\"\"\n    valid = mask > 0\n    unc = uncertainty_map[valid]\n    prob = prob_map[valid]\n    gtv = gt[valid]\n    preds = (prob > threshold).astype(np.float32)\n\n    quantiles = np.quantile(unc, np.linspace(0, 1, n_buckets + 1))\n    report = []\n    for i in range(n_buckets):\n        lo, hi = quantiles[i], quantiles[i + 1]\n        bucket = (unc >= lo) & (unc <= hi) if i == n_buckets - 1 else (unc >= lo) & (unc < hi)\n        if bucket.sum() == 0:\n            continue\n        p_bucket = preds[bucket]\n        g_bucket = gtv[bucket]\n        tp = (p_bucket * g_bucket).sum()\n        fp = (p_bucket * (1 - g_bucket)).sum()\n        precision = (tp + 1e-6) / (tp + fp + 1e-6)\n        report.append({\"bucket\": i, \"uncertainty_range\": (float(lo), float(hi)),\n                        \"n_pixels\": int(bucket.sum()), \"precision\": float(precision)})\n    return report\n\n\n# ============================================================\n# 16b. TEST-TIME TRAINING (V7, Protocol=\"transductive\" ONLY -- see change log)\n# ============================================================\n\ndef test_time_training_adapt(model, test_vol, test_coords_unlabeled):\n    \"\"\"Adapts a COPY of the trained model to the held-out fragment's\n    unlabeled statistics via entropy minimization, for CFG.ttt_steps batches.\n    Returns the adapted copy; the caller's original `model` (and therefore\n    the strict-protocol evaluation) is never touched.\n\n    This is a transductive technique -- it fits weights to target-domain\n    unlabeled data -- so it is ONLY ever invoked when\n    CFG.PROTOCOL == \"transductive\" and CFG.USE_TEST_TIME_TRAINING, and its\n    output is stored as a separately-labeled `test_metrics_ttt`, never\n    averaged into the strict LOFO-CV summary (see `summarize_folds`, which\n    only reads `test_metrics_postprocessed`).\"\"\"\n    assert CFG.PROTOCOL == \"transductive\", (\n        \"test_time_training_adapt called outside Protocol B -- refusing, since \"\n        \"this would silently leak target-fragment statistics into a \"\n        \"strict-protocol result.\")\n\n    adapted = copy.deepcopy(model)\n    adapted.train()\n    optimizer = torch.optim.SGD(adapted.parameters(), lr=CFG.ttt_lr)\n\n    ds = UnlabeledPatchDataset(test_vol, test_coords_unlabeled, CFG.patch_size)\n    loader = DataLoader(ds, batch_size=CFG.ttt_batch_size, shuffle=True, num_workers=1, drop_last=True)\n    loader_iter = iter(loader)\n\n    print(f\"[TTT] adapting on {CFG.ttt_steps} unlabeled target-fragment batches \"\n          f\"(lr={CFG.ttt_lr}) ...\")\n    for step in range(CFG.ttt_steps):\n        try:\n            batch = next(loader_iter)\n        except StopIteration:\n            loader_iter = iter(loader)\n            batch = next(loader_iter)\n        batch = batch.to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = adapted(batch)\n            probs = torch.sigmoid(logits).clamp(1e-6, 1 - 1e-6)\n            entropy = -(probs * torch.log(probs) + (1 - probs) * torch.log(1 - probs)).mean()\n        entropy.backward()\n        torch.nn.utils.clip_grad_norm_(adapted.parameters(), CFG.grad_clip)\n        optimizer.step()\n        optimizer.zero_grad(set_to_none=True)\n        del batch, logits, probs\n\n    adapted.eval()\n    return adapted\n\n\n# ============================================================\n# 17. INFERENCE HELPERS (sliding window, TTA, postprocessing)\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats_, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats_[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 18. ONE-FOLD TRAIN/EVAL (the core unit LOFO-CV and ablation both call)\n# ============================================================\n\ndef run_one_fold(train_frags, test_frag, fold_tag, save_viz=False):\n    \"\"\"Builds data, trains, selects threshold on validation only, evaluates\n    on the held-out fragment, and returns a metrics dict. This is the single\n    unit both `run_lofo_cv` and `run_ablation_matrix` call, so every arm goes\n    through IDENTICAL code -- only CFG differs between calls.\"\"\"\n    set_seed(CFG.seed)\n    print(\"\\n\" + \"#\" * 70)\n    print(f\"# FOLD [{fold_tag}]  train={train_frags}  test={test_frag}\")\n    print(\"#\" * 70)\n\n    # ---- build train data ----\n    train_volumes, train_labels_full, train_masks_full = {}, {}, {}\n    train_samples_raw, val_samples = [], []\n\n    for fid in train_frags:\n        frag_dir = os.path.join(CFG.base_dir, fid)\n        mask = load_tissue_mask(frag_dir)\n        labels = load_ink_labels(frag_dir)\n        if labels is None:\n            raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n        vol = make_fragment_volume(frag_dir)\n        if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n            vol.compute_fragment_stats(mask, CFG.patch_size)\n        if CFG.WARM_UP_VOLUME_CACHE:\n            print(f\"  fragment {fid}:\", end=\" \")\n            vol.warm_up_cache()\n\n        coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n        tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n\n        train_volumes[fid] = vol\n        train_labels_full[fid] = labels\n        train_masks_full[fid] = mask\n        train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n        val_samples.extend([(fid, y, x) for y, x in va_coords])\n        print(f\"  fragment {fid}: train={len(tr_coords)} val={len(va_coords)} \"\n              f\"mean={vol.frag_mean:.1f} std={vol.frag_std:.1f}\")\n        del mask\n        cleanup_memory()\n\n    train_samples = balance_positive_patches(\n        train_samples_raw, train_labels_full, CFG.patch_size,\n        positive_threshold=CFG.positive_patch_fraction,\n        target_positive_ratio=CFG.target_positive_patch_ratio,\n        max_positive_repeat=CFG.max_positive_repeat)\n\n    # ---- held-out fragment: unlabeled use only, labels loaded LATE ----\n    test_dir = os.path.join(CFG.base_dir, test_frag)\n    test_mask = load_tissue_mask(test_dir)\n    test_vol = make_fragment_volume(test_dir)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    if CFG.WARM_UP_VOLUME_CACHE:\n        print(f\"  fragment {test_frag} (held-out):\", end=\" \")\n        test_vol.warm_up_cache()\n\n    test_coords_unlabeled = generate_grid_coords(test_mask, CFG.patch_size, CFG.test_stride,\n                                                  CFG.min_tissue_frac_train)\n\n    # V7.2 (bugfix, flagged by review): the histogram-match reference pool was\n    # being built from the HELD-OUT fragment's own image patches and fed\n    # straight into the TRAINING dataset as an augmentation reference -- even\n    # without labels, that lets the model's training-time inputs be reshaped\n    # to look like the test fragment's intensity distribution, which is\n    # exactly the leakage the strict/transductive protocol split exists to\n    # prevent (and which the CFG.PROTOCOL docstring already claimed doesn't\n    # happen). Fixed: in \"strict\" protocol the pool is built from the\n    # TRAINING fragments' own patches; the held-out fragment's patches are\n    # only used for this purpose under PROTOCOL == \"transductive\", where\n    # that's the explicit, clearly-labeled point of the experiment.\n    hist_match_pool = None\n    if CFG.USE_HIST_MATCH_AUG:\n        if CFG.PROTOCOL == \"strict\":\n            pool_source_coords = []\n            for fid in train_frags:\n                coords = generate_grid_coords(train_masks_full[fid], CFG.patch_size, CFG.patch_size,\n                                               CFG.min_tissue_frac_train)\n                pool_source_coords.extend([(fid, y, x) for y, x in coords])\n            n_pool = min(CFG.hist_match_pool_size, len(pool_source_coords))\n            pool_samples = random.sample(pool_source_coords, n_pool) if pool_source_coords else []\n            hist_match_pool = [train_volumes[fid].read_patch(y, x, CFG.patch_size)\n                                for fid, y, x in pool_samples]\n            print(f\"  hist-match pool (strict protocol): {len(hist_match_pool)} patches from \"\n                  f\"TRAINING fragments {train_frags} only -- held-out fragment {test_frag} untouched.\")\n        elif CFG.PROTOCOL == \"transductive\":\n            pool_coords = random.sample(test_coords_unlabeled,\n                                         min(CFG.hist_match_pool_size, len(test_coords_unlabeled)))\n            hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n            print(f\"  hist-match pool (TRANSDUCTIVE protocol, by design): {len(hist_match_pool)} \"\n                  f\"patches from held-out fragment {test_frag}.\")\n        else:\n            raise ValueError(f\"Unknown CFG.PROTOCOL={CFG.PROTOCOL}\")\n\n    # ---- datasets / loaders ----\n    train_transform = build_train_transform()\n    train_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                                transform=train_transform, jitter=CFG.train_jitter,\n                                hist_match_pool=hist_match_pool, train_mode=True)\n    val_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                              transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\n    sample_difficulties = None\n    curriculum_sampler = None\n    if CFG.USE_CURRICULUM:\n        sample_difficulties = compute_sample_difficulty(\n            train_samples, train_labels_full, CFG.patch_size, CFG.curriculum_easy_ink_frac)\n        print(f\"  curriculum sampling ON: warmup_epochs={CFG.curriculum_warmup_epochs} \"\n              f\"easy/medium/hard counts = \"\n              f\"{int((sample_difficulties==0).sum())}/{int((sample_difficulties==0.5).sum())}/\"\n              f\"{int((sample_difficulties==1).sum())}\")\n        # V7.1 (perf): build ONE sampler object and mutate its `.weights`\n        # in-place each epoch instead of recreating the DataLoader (which\n        # respawns worker processes from scratch every epoch -- expensive\n        # with tifffile-backed volumes and a heavy CPU augmentation pipeline).\n        curriculum_sampler = torch.utils.data.WeightedRandomSampler(\n            build_curriculum_sampler(sample_difficulties, 0, CFG.epochs, CFG.curriculum_warmup_epochs),\n            num_samples=len(train_ds), replacement=True)\n        train_loader = DataLoader(\n            train_ds, batch_size=CFG.batch_size, sampler=curriculum_sampler,\n            num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n            drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0,\n            prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n    else:\n        train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                                   num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                                   drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0,\n                                   prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                             num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                             drop_last=False, persistent_workers=CFG.num_workers > 0,\n                             prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n\n    # ---- model / loss / optimizer ----\n    model = build_model().to(CFG.device)\n    run_architecture_report(model)\n\n    pos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\n    with torch.no_grad():\n        bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n        try:\n            model.seg_model.segmentation_head[0].bias.fill_(bias_val)\n        except Exception as e:\n            print(f\"  (could not set output bias directly: {e})\")\n\n    raw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\n    pos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\n    pos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\n    criterion = VXComboLoss(pos_weight=pos_weight)\n\n    encoder_params, decoder_params, stem_params = [], [], []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if name.startswith(\"depth_stem.\"):\n            stem_params.append(param)\n        elif name.startswith(\"seg_model.encoder.\"):\n            encoder_params.append(param)\n        else:\n            decoder_params.append(param)\n\n    param_groups = [\n        {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n        {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n    ]\n    if stem_params:\n        param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\n    optimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\n    if CFG.LR_SCHEDULE == \"onecycle\":\n        max_lrs = [g[\"lr\"] for g in param_groups]\n        # V7.2 (bugfix, flagged by review): scheduler.step() only fires once\n        # per OPTIMIZER update (i.e. once every accumulation_steps batches),\n        # not once per batch -- so telling OneCycleLR steps_per_epoch=\n        # len(train_loader) overstates the cycle length whenever\n        # accumulation_steps > 1, desynchronizing the LR curve from the\n        # actual number of optimizer steps taken. (With the default\n        # accumulation_steps=1 this was a no-op, but it's a real bug for\n        # anyone who raises accumulation_steps, which is a normal thing to\n        # do on a memory-constrained T4.)\n        optimizer_steps_per_epoch = math.ceil(len(train_loader) / CFG.accumulation_steps)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=max_lrs, epochs=CFG.epochs, steps_per_epoch=optimizer_steps_per_epoch,\n            pct_start=CFG.onecycle_pct_start, anneal_strategy=\"cos\")\n        scheduler_steps_per_batch = True\n    else:\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n        scheduler_steps_per_batch = False\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n    ema = EMAModel(model, decay=CFG.ema_decay) if CFG.USE_EMA else None\n    global_step_holder = [0]\n    total_steps = CFG.epochs * max(1, len(train_loader) // CFG.accumulation_steps)\n\n    # ---- training loop ----\n    best_val_f05 = -1.0\n    epochs_no_improve = 0\n    best_state = None\n    history = defaultdict(list)\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n\n        if CFG.USE_CURRICULUM and sample_difficulties is not None:\n            # V7.1 (perf): mutate the existing sampler's weights in place --\n            # does NOT recreate the DataLoader, so persistent worker\n            # processes are kept warm across epochs instead of respawned.\n            curriculum_sampler.weights = build_curriculum_sampler(\n                sample_difficulties, epoch - 1, CFG.epochs, CFG.curriculum_warmup_epochs)\n\n        train_loss, train_metrics, train_consist = run_epoch(\n            model, criterion, train_loader, optimizer, scaler, ema,\n            None, None, None, global_step_holder, total_steps, train_mode=True, threshold=0.50,\n            scheduler=scheduler, scheduler_steps_per_batch=scheduler_steps_per_batch,\n            profile_steps=(CFG.PROFILE_TIMING_STEPS if (CFG.PROFILE_TIMING and epoch == 1) else 0))\n\n        if ema is not None:\n            backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n            model.load_state_dict(ema.state_dict())\n            val_loss, val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, ema,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n            model.load_state_dict(backup)\n            del backup\n        else:\n            val_loss, val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, ema,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n\n        if not scheduler_steps_per_batch:\n            scheduler.step()\n\n        for k, v in val_metrics.items():\n            history[f\"val_{k}\"].append(v)\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"train_dice\"].append(train_metrics[\"dice\"])\n        history[\"train_consistency_loss\"].append(train_consist)\n\n        print(f\"[{epoch:02d}/{CFG.epochs}] {time.time()-t0:.1f}s | \"\n              f\"train_loss={train_loss:.4f} (consistency={train_consist:.4f}) | \"\n              f\"val_f0.5={val_metrics['f0.5']:.4f} val_dice={val_metrics['dice']:.4f} \"\n              f\"val_prec={val_metrics['precision']:.4f} val_rec={val_metrics['recall']:.4f}\")\n\n        if val_metrics[\"f0.5\"] > best_val_f05:\n            best_val_f05 = val_metrics[\"f0.5\"]\n            epochs_no_improve = 0\n            best_state = copy.deepcopy(ema.state_dict() if ema is not None else model.state_dict())\n            print(f\"  *** new best (val F0.5={best_val_f05:.4f}) ***\")\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= CFG.early_stop_patience:\n                print(\"  early stopping.\")\n                break\n        cleanup_memory()\n\n    model.load_state_dict(best_state)\n\n    # ---- threshold selection: VALIDATION ONLY, then frozen ----\n    best_threshold, val_score_at_best, thr_grid, thr_scores = find_best_threshold(\n        model, val_loader, metric=\"f0.5\")\n    print(f\"\\nSelected threshold (validation-only) = {best_threshold:.2f} \"\n          f\"(val F0.5={val_score_at_best:.4f})\")\n\n    # ---- held-out fragment: load labels now, evaluate at the FROZEN threshold ----\n    test_labels = load_ink_labels(test_dir)\n    fold_result = {\"fold_tag\": fold_tag, \"train_frags\": list(train_frags), \"test_frag\": test_frag,\n                   \"best_val_f05\": best_val_f05, \"selected_threshold\": best_threshold,\n                   \"history\": dict(history)}\n\n    if test_labels is not None:\n        gt_test = (test_labels * test_mask).astype(np.float32)\n        test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                              CFG.test_stride, CFG.infer_batch)\n        test_metrics_raw = evaluate_probability_map(test_prob, gt_test, best_threshold)\n        test_pred_bin = postprocess(test_prob, best_threshold)\n        post_preds = test_pred_bin.astype(np.float32)\n        tp = (post_preds * gt_test).sum(); fp = (post_preds * (1 - gt_test)).sum()\n        fn = ((1 - post_preds) * gt_test).sum()\n        test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n        print(f\"\\nHELD-OUT fragment {test_frag} metrics (frozen threshold={best_threshold:.2f}):\")\n        print(f\"  raw:          {test_metrics_raw}\")\n        print(f\"  postprocessed:{test_metrics_post}\")\n\n        fold_result[\"test_metrics_raw\"] = test_metrics_raw\n        fold_result[\"test_metrics_postprocessed\"] = test_metrics_post\n\n        if CFG.USE_MC_DROPOUT_UNCERTAINTY:\n            mean_prob, uncertainty = run_mc_dropout_uncertainty(\n                model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.mc_dropout_passes)\n            unc_report = analyze_uncertainty_vs_error(mean_prob, uncertainty, gt_test, test_mask,\n                                                       best_threshold, n_buckets=3)\n            print(\"  MC-dropout uncertainty vs. precision (low/med/high uncertainty buckets):\")\n            for r in unc_report:\n                print(f\"    bucket {r['bucket']}: n={r['n_pixels']} precision={r['precision']:.3f}\")\n            fold_result[\"mc_dropout_uncertainty_report\"] = unc_report\n\n        if CFG.PROTOCOL == \"transductive\" and CFG.USE_TEST_TIME_TRAINING:\n            ttt_model = test_time_training_adapt(model, test_vol, test_coords_unlabeled)\n            test_prob_ttt = sliding_window_inference(ttt_model, test_vol, test_mask, CFG.patch_size,\n                                                       CFG.test_stride, CFG.infer_batch)\n            test_metrics_ttt = evaluate_probability_map(test_prob_ttt, gt_test, best_threshold)\n            print(f\"  [TTT, Protocol B, NOT part of strict CV] test metrics: {test_metrics_ttt}\")\n            fold_result[\"test_metrics_ttt_protocol_b_only\"] = test_metrics_ttt\n            del ttt_model\n            cleanup_memory()\n\n        if save_viz:\n            _save_fold_overview(test_dir, test_prob, test_pred_bin, test_labels, fold_tag, best_threshold)\n    else:\n        print(\"No ground-truth labels for held-out fragment -- competition-style inference only.\")\n        test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                              CFG.test_stride, CFG.infer_batch)\n        fold_result[\"test_metrics_raw\"] = None\n\n    prob_path = os.path.join(CFG.out_dir, f\"fragment{test_frag}_probability_{fold_tag}.npy\")\n    np.save(prob_path, test_prob)\n    fold_result[\"probability_map_path\"] = prob_path\n\n    # ---- cleanup ----\n    test_vol.close()\n    for v in train_volumes.values():\n        v.close()\n    del model, optimizer, scheduler\n    cleanup_memory()\n\n    return fold_result\n\n\ndef _save_fold_overview(test_dir, prob_map, pred_bin, gt_labels, tag, threshold):\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 1600 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize((gt_labels * 255).astype(np.uint8), small.shape[::-1],\n                           interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    prob_small = cv2.resize((prob_map * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n    axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability (thr={threshold:.2f})\")\n    axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, f\"overview_{tag}.png\")\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"  saved overview: {path}\")\n\n\n# ============================================================\n# 19. LEAVE-ONE-FRAGMENT-OUT CV  (main scientific result)\n# ============================================================\n\ndef run_lofo_cv(fragments):\n    \"\"\"3 fragments -> 3 folds. Reports mean +/- std across folds instead of a\n    single 'best validation Dice' number, plus a paired t-test / Wilcoxon\n    signed-rank comparison IS available via compare_fold_results below if you\n    run two configurations (e.g. baseline vs. +DepthSignature) through this\n    same function and diff their fold-level F0.5 lists.\"\"\"\n    fold_results = []\n    for held_out in fragments:\n        train_frags = [f for f in fragments if f != held_out]\n        result = run_one_fold(train_frags, held_out, fold_tag=f\"lofo_test{held_out}\",\n                               save_viz=True)\n        fold_results.append(result)\n\n    summary = summarize_folds(fold_results)\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LEAVE-ONE-FRAGMENT-OUT CV SUMMARY\")\n    print(\"=\" * 70)\n    for metric, (mean, std, vals) in summary.items():\n        print(f\"  {metric:12s} = {mean:.4f} +/- {std:.4f}   (per-fold: {['%.4f' % v for v in vals]})\")\n\n    out_path = os.path.join(CFG.out_dir, \"lofo_cv_results.json\")\n    with open(out_path, \"w\") as f:\n        json.dump({\"folds\": fold_results, \"summary\": {k: (v[0], v[1]) for k, v in summary.items()},\n                    \"config\": cfg_to_dict(CFG)}, f, indent=2, default=str)\n    print(f\"\\nSaved: {out_path}\")\n    return fold_results, summary\n\n\ndef summarize_folds(fold_results, metric_keys=(\"f0.5\", \"dice\", \"iou\", \"precision\", \"recall\")):\n    summary = {}\n    for key in metric_keys:\n        vals = [fr[\"test_metrics_postprocessed\"][key] for fr in fold_results\n                 if fr.get(\"test_metrics_postprocessed\") is not None]\n        if not vals:\n            continue\n        summary[key] = (float(np.mean(vals)), float(np.std(vals)), vals)\n    return summary\n\n\ndef compare_fold_results(fold_results_a, fold_results_b, metric=\"f0.5\"):\n    \"\"\"Paired comparison across folds (same held-out fragments in the same\n    order for both configurations) -- Wilcoxon signed-rank test, with a\n    paired t-test reported alongside since n=3 folds is too small for the\n    Wilcoxon test's own asymptotics to be trustworthy; report both and let\n    the reader see they agree in direction.\"\"\"\n    a = [fr[\"test_metrics_postprocessed\"][metric] for fr in fold_results_a]\n    b = [fr[\"test_metrics_postprocessed\"][metric] for fr in fold_results_b]\n    t_stat, t_p = sstats.ttest_rel(a, b)\n    try:\n        w_stat, w_p = sstats.wilcoxon(a, b)\n    except Exception:\n        w_stat, w_p = float(\"nan\"), float(\"nan\")\n    print(f\"Paired comparison on {metric}: A={np.mean(a):.4f} B={np.mean(b):.4f} \"\n          f\"| paired t-test p={t_p:.4f} | Wilcoxon p={w_p:.4f}\")\n    return {\"metric\": metric, \"mean_a\": float(np.mean(a)), \"mean_b\": float(np.mean(b)),\n            \"t_stat\": float(t_stat), \"t_p\": float(t_p), \"w_stat\": float(w_stat), \"w_p\": float(w_p)}\n\n\n# ============================================================\n# 20. ABLATION MATRIX (component-by-component, single fold)\n# ============================================================\n\nABLATION_ARMS = [\n    (\"A_baseline_no_depth_stem\",       dict(depth_module_type=\"none\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"B_depth_fusion_stem\",            dict(depth_module_type=\"fusion\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"C_depth_signature_no_extras\",    dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"D_signature_plus_physics_ink\",   dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"E_signature_plus_topology\",      dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=True)),\n    (\"F_signature_plus_consistency\",   dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=False, USE_CURRICULUM=False)),   # = full V6 model\n    (\"G_conv3d_stem_instead\",          dict(depth_module_type=\"conv3d\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=False, USE_CURRICULUM=False)),\n    (\"H_full_v7_plus_cutmix_curriculum\", dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=True, USE_CURRICULUM=True)),\n]\n\n\ndef run_ablation_matrix(train_frags, test_frag, arms=ABLATION_ARMS):\n    \"\"\"Runs each named arm on the SAME fold (same train/test fragment split)\n    so the differences are attributable to the listed components, not to a\n    different data split. Each arm is a fresh CfgOverride, so arms don't leak\n    settings into each other.\"\"\"\n    results = {}\n    for name, overrides in arms:\n        with CfgOverride(**overrides):\n            print(f\"\\n>>> ABLATION ARM: {name}  overrides={overrides}\")\n            fold_result = run_one_fold(train_frags, test_frag, fold_tag=f\"ablation_{name}\")\n            results[name] = fold_result\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"ABLATION MATRIX SUMMARY (test fragment = %s)\" % test_frag)\n    print(\"=\" * 70)\n    print(f\"{'arm':32s} {'F0.5':>8s} {'Dice':>8s} {'IoU':>8s} {'Prec':>8s} {'Rec':>8s}\")\n    for name, fr in results.items():\n        m = fr.get(\"test_metrics_postprocessed\")\n        if m is None:\n            print(f\"{name:32s}  (no GT available)\")\n            continue\n        print(f\"{name:32s} {m['f0.5']:8.4f} {m['dice']:8.4f} {m['iou']:8.4f} \"\n              f\"{m['precision']:8.4f} {m['recall']:8.4f}\")\n\n    out_path = os.path.join(CFG.out_dir, \"ablation_matrix_results.json\")\n    with open(out_path, \"w\") as f:\n        json.dump(results, f, indent=2, default=str)\n    print(f\"\\nSaved: {out_path}\")\n    return results\n\n\n# ============================================================\n# 21. MAIN\n# ============================================================\n\nif __name__ == \"__main__\":\n    set_seed(CFG.seed)\n\n    if CFG.RUN_MODE == \"lofo_cv\":\n        fold_results, summary = run_lofo_cv(CFG.all_fragments)\n\n    elif CFG.RUN_MODE == \"ablation\":\n        ablation_results = run_ablation_matrix(CFG.ablation_train_frags, CFG.ablation_test_frag)\n\n    elif CFG.RUN_MODE == \"single\":\n        result = run_one_fold(CFG.single_train_frags, CFG.single_test_frag,\n                               fold_tag=\"single_run\", save_viz=True)\n        print(\"\\nSingle-run result:\", json.dumps(\n            {k: v for k, v in result.items() if k != \"history\"}, indent=2, default=str))\n\n    else:\n        raise ValueError(f\"Unknown CFG.RUN_MODE={CFG.RUN_MODE}\")\n\n    print(\"\\n=== V6 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-18T12:09:27.716542Z","iopub.execute_input":"2026-09-18T12:09:27.716988Z","iopub.status.idle":"2026-09-18T19:57:07.007359Z","shell.execute_reply.started":"2026-09-18T12:09:27.716956Z","shell.execute_reply":"2026-09-18T19:57:07.006507Z"}},"outputs":[{"name":"stdout","text":"     ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 43.7/43.7 kB 3.0 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 11.0 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.6/2.6 MB 77.9 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.29\n\n######################################################################\n# FOLD [lofo_test1]  train=['2', '3']  test=1\n######################################################################\n  fragment 2:   [cache warm-up] 34 slices, 9.59 GB read sequentially in 103.8s (92.3 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 2: train=4559 val=1260 mean=110.0 std=57.6\n  fragment 3:   [cache warm-up] 34 slices, 2.71 GB read sequentially in 29.0s (93.6 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 3: train=1272 val=160 mean=100.3 std=62.4\nPositive patches: 3953 | Negative patches: 1878\nBalanced dataset: 5831 | positive ratio=0.678\n  fragment 1 (held-out):   [cache warm-up] 34 slices, 3.52 GB read sequentially in 37.4s (94.3 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  hist-match pool (strict protocol): 20 patches from TRAINING fragments ['2', '3'] only -- held-out fragment 1 untouched.\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n","output_type":"stream"},{"name":"stdout","text":"  curriculum sampling ON: warmup_epochs=8 easy/medium/hard counts = 2736/1287/1808\n[backbone] DepthSignatureModule: 26 depth slices -> 24 learned channels (LDDC=4 filters, PE dim=8, attention heads=4, pooled to 8x8, stats=True)\n","output_type":"stream"},{"name":"stderr","text":"Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"f97ec7d0fef74ec890c5c3d6af752c1d"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320\ndepth_module_type=signature | encoder=tu-convnext_tiny | architecture=unet\nParameters: 32.18M total | 32.18M trainable\nZero-channel layers: 0 (PASS)\nDry-run (256x256) output logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\nCUDA device: Tesla T4\nForward pass: PASS\n======================================================================\n\n    step 36/728 | 132s elapsed | 0.27 steps/s | ETA 2543s\n    step 72/728 | 211s elapsed | 0.34 steps/s | ETA 1925s\n    step 108/728 | 289s elapsed | 0.37 steps/s | ETA 1659s\n    step 144/728 | 365s elapsed | 0.39 steps/s | ETA 1482s\n    step 180/728 | 442s elapsed | 0.41 steps/s | ETA 1345s\n    step 216/728 | 515s elapsed | 0.42 steps/s | ETA 1222s\n    step 252/728 | 591s elapsed | 0.43 steps/s | ETA 1115s\n    step 288/728 | 667s elapsed | 0.43 steps/s | ETA 1019s\n    step 324/728 | 744s elapsed | 0.44 steps/s | ETA 927s\n    step 360/728 | 819s elapsed | 0.44 steps/s | ETA 837s\n    step 396/728 | 895s elapsed | 0.44 steps/s | ETA 750s\n    step 432/728 | 972s elapsed | 0.44 steps/s | ETA 666s\n    step 468/728 | 1050s elapsed | 0.45 steps/s | ETA 583s\n    step 504/728 | 1127s elapsed | 0.45 steps/s | ETA 501s\n    step 540/728 | 1207s elapsed | 0.45 steps/s | ETA 420s\n    step 576/728 | 1285s elapsed | 0.45 steps/s | ETA 339s\n    step 612/728 | 1363s elapsed | 0.45 steps/s | ETA 258s\n    step 648/728 | 1442s elapsed | 0.45 steps/s | ETA 178s\n    step 684/728 | 1520s elapsed | 0.45 steps/s | ETA 98s\n    step 720/728 | 1595s elapsed | 0.45 steps/s | ETA 18s\n[01/6] 1665.7s | train_loss=0.8012 (consistency=0.0011) | val_f0.5=0.0027 val_dice=0.0000 val_prec=0.0015 val_rec=0.0000\n  *** new best (val F0.5=0.0027) ***\n    step 36/728 | 77s elapsed | 0.46 steps/s | ETA 1489s\n    step 72/728 | 155s elapsed | 0.46 steps/s | ETA 1417s\n    step 108/728 | 233s elapsed | 0.46 steps/s | ETA 1340s\n    step 144/728 | 313s elapsed | 0.46 steps/s | ETA 1269s\n    step 180/728 | 388s elapsed | 0.46 steps/s | ETA 1181s\n    step 216/728 | 464s elapsed | 0.47 steps/s | ETA 1101s\n    step 252/728 | 542s elapsed | 0.46 steps/s | ETA 1024s\n    step 288/728 | 613s elapsed | 0.47 steps/s | ETA 937s\n    step 324/728 | 688s elapsed | 0.47 steps/s | ETA 858s\n    step 360/728 | 765s elapsed | 0.47 steps/s | ETA 782s\n    step 396/728 | 840s elapsed | 0.47 steps/s | ETA 704s\n    step 432/728 | 913s elapsed | 0.47 steps/s | ETA 626s\n    step 468/728 | 993s elapsed | 0.47 steps/s | ETA 552s\n    step 504/728 | 1069s elapsed | 0.47 steps/s | ETA 475s\n    step 540/728 | 1144s elapsed | 0.47 steps/s | ETA 398s\n    step 576/728 | 1222s elapsed | 0.47 steps/s | ETA 323s\n    step 612/728 | 1295s elapsed | 0.47 steps/s | ETA 245s\n    step 648/728 | 1372s elapsed | 0.47 steps/s | ETA 169s\n    step 684/728 | 1452s elapsed | 0.47 steps/s | ETA 93s\n    step 720/728 | 1528s elapsed | 0.47 steps/s | ETA 17s\n[02/6] 1593.9s | train_loss=0.7088 (consistency=0.0011) | val_f0.5=0.0045 val_dice=0.0018 val_prec=0.8999 val_rec=0.0009\n  *** new best (val F0.5=0.0045) ***\n    step 36/728 | 75s elapsed | 0.48 steps/s | ETA 1434s\n    step 72/728 | 151s elapsed | 0.48 steps/s | ETA 1377s\n    step 108/728 | 225s elapsed | 0.48 steps/s | ETA 1291s\n    step 144/728 | 299s elapsed | 0.48 steps/s | ETA 1211s\n    step 180/728 | 375s elapsed | 0.48 steps/s | ETA 1142s\n    step 216/728 | 452s elapsed | 0.48 steps/s | ETA 1070s\n    step 252/728 | 528s elapsed | 0.48 steps/s | ETA 997s\n    step 288/728 | 604s elapsed | 0.48 steps/s | ETA 923s\n    step 324/728 | 682s elapsed | 0.47 steps/s | ETA 851s\n    step 360/728 | 760s elapsed | 0.47 steps/s | ETA 777s\n    step 396/728 | 834s elapsed | 0.47 steps/s | ETA 699s\n    step 432/728 | 912s elapsed | 0.47 steps/s | ETA 625s\n    step 468/728 | 991s elapsed | 0.47 steps/s | ETA 551s\n    step 504/728 | 1065s elapsed | 0.47 steps/s | ETA 473s\n    step 540/728 | 1143s elapsed | 0.47 steps/s | ETA 398s\n    step 576/728 | 1219s elapsed | 0.47 steps/s | ETA 322s\n    step 612/728 | 1294s elapsed | 0.47 steps/s | ETA 245s\n    step 648/728 | 1371s elapsed | 0.47 steps/s | ETA 169s\n    step 684/728 | 1447s elapsed | 0.47 steps/s | ETA 93s\n    step 720/728 | 1524s elapsed | 0.47 steps/s | ETA 17s\n[03/6] 1589.5s | train_loss=0.6752 (consistency=0.0014) | val_f0.5=0.4653 val_dice=0.4759 val_prec=0.4585 val_rec=0.4947\n  *** new best (val F0.5=0.4653) ***\n    step 36/728 | 77s elapsed | 0.46 steps/s | ETA 1489s\n    step 72/728 | 154s elapsed | 0.47 steps/s | ETA 1403s\n    step 108/728 | 232s elapsed | 0.47 steps/s | ETA 1332s\n    step 144/728 | 311s elapsed | 0.46 steps/s | ETA 1263s\n    step 180/728 | 389s elapsed | 0.46 steps/s | ETA 1185s\n    step 216/728 | 466s elapsed | 0.46 steps/s | ETA 1104s\n    step 252/728 | 541s elapsed | 0.47 steps/s | ETA 1022s\n    step 288/728 | 616s elapsed | 0.47 steps/s | ETA 941s\n    step 324/728 | 694s elapsed | 0.47 steps/s | ETA 865s\n    step 360/728 | 772s elapsed | 0.47 steps/s | ETA 789s\n    step 396/728 | 850s elapsed | 0.47 steps/s | ETA 712s\n    step 432/728 | 926s elapsed | 0.47 steps/s | ETA 635s\n    step 468/728 | 1004s elapsed | 0.47 steps/s | ETA 558s\n    step 504/728 | 1082s elapsed | 0.47 steps/s | ETA 481s\n    step 540/728 | 1159s elapsed | 0.47 steps/s | ETA 403s\n    step 576/728 | 1235s elapsed | 0.47 steps/s | ETA 326s\n    step 612/728 | 1309s elapsed | 0.47 steps/s | ETA 248s\n    step 648/728 | 1385s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1460s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1538s elapsed | 0.47 steps/s | ETA 17s\n[04/6] 1604.1s | train_loss=0.6534 (consistency=0.0016) | val_f0.5=0.3777 val_dice=0.4630 val_prec=0.3364 val_rec=0.7425\n    step 36/728 | 73s elapsed | 0.49 steps/s | ETA 1406s\n    step 72/728 | 153s elapsed | 0.47 steps/s | ETA 1390s\n    step 108/728 | 229s elapsed | 0.47 steps/s | ETA 1315s\n    step 144/728 | 308s elapsed | 0.47 steps/s | ETA 1251s\n    step 180/728 | 385s elapsed | 0.47 steps/s | ETA 1172s\n    step 216/728 | 459s elapsed | 0.47 steps/s | ETA 1087s\n    step 252/728 | 537s elapsed | 0.47 steps/s | ETA 1014s\n    step 288/728 | 610s elapsed | 0.47 steps/s | ETA 932s\n    step 324/728 | 688s elapsed | 0.47 steps/s | ETA 858s\n    step 360/728 | 765s elapsed | 0.47 steps/s | ETA 782s\n    step 396/728 | 843s elapsed | 0.47 steps/s | ETA 706s\n    step 432/728 | 921s elapsed | 0.47 steps/s | ETA 631s\n    step 468/728 | 999s elapsed | 0.47 steps/s | ETA 555s\n    step 504/728 | 1076s elapsed | 0.47 steps/s | ETA 478s\n    step 540/728 | 1156s elapsed | 0.47 steps/s | ETA 402s\n    step 576/728 | 1229s elapsed | 0.47 steps/s | ETA 324s\n    step 612/728 | 1307s elapsed | 0.47 steps/s | ETA 248s\n    step 648/728 | 1384s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1459s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1537s elapsed | 0.47 steps/s | ETA 17s\n[05/6] 1604.1s | train_loss=0.6428 (consistency=0.0018) | val_f0.5=0.3482 val_dice=0.4439 val_prec=0.3044 val_rec=0.8190\n    step 36/728 | 79s elapsed | 0.46 steps/s | ETA 1516s\n    step 72/728 | 158s elapsed | 0.45 steps/s | ETA 1442s\n    step 108/728 | 236s elapsed | 0.46 steps/s | ETA 1356s\n    step 144/728 | 311s elapsed | 0.46 steps/s | ETA 1263s\n    step 180/728 | 388s elapsed | 0.46 steps/s | ETA 1181s\n    step 216/728 | 463s elapsed | 0.47 steps/s | ETA 1098s\n    step 252/728 | 541s elapsed | 0.47 steps/s | ETA 1022s\n    step 288/728 | 617s elapsed | 0.47 steps/s | ETA 943s\n    step 324/728 | 697s elapsed | 0.46 steps/s | ETA 869s\n    step 360/728 | 772s elapsed | 0.47 steps/s | ETA 789s\n    step 396/728 | 850s elapsed | 0.47 steps/s | ETA 713s\n    step 432/728 | 928s elapsed | 0.47 steps/s | ETA 636s\n    step 468/728 | 1004s elapsed | 0.47 steps/s | ETA 558s\n    step 504/728 | 1078s elapsed | 0.47 steps/s | ETA 479s\n    step 540/728 | 1153s elapsed | 0.47 steps/s | ETA 401s\n    step 576/728 | 1230s elapsed | 0.47 steps/s | ETA 324s\n    step 612/728 | 1307s elapsed | 0.47 steps/s | ETA 248s\n    step 648/728 | 1384s elapsed | 0.47 steps/s | ETA 171s\n    step 684/728 | 1459s elapsed | 0.47 steps/s | ETA 94s\n    step 720/728 | 1538s elapsed | 0.47 steps/s | ETA 17s\n[06/6] 1605.6s | train_loss=0.6425 (consistency=0.0020) | val_f0.5=0.3598 val_dice=0.4551 val_prec=0.3158 val_rec=0.8143\n  early stopping.\n\nSelected threshold (validation-only) = 0.55 (val F0.5=0.4952)\nInference patches: 1932\n\nHELD-OUT fragment 1 metrics (frozen threshold=0.55):\n  raw:          {'f0.5': 0.46326783299446106, 'dice': 0.38643747568130493, 'iou': 0.23949334025382996, 'precision': 0.5340511798858643, 'recall': 0.3027549088001251}\n  postprocessed:{'f0.5': 0.4637351930141449, 'dice': 0.3879503011703491, 'iou': 0.2406565248966217, 'precision': 0.5331680774688721, 'recall': 0.3049042224884033}\n  MC-dropout uncertainty vs. precision (low/med/high uncertainty buckets):\n    bucket 0: n=9433101 precision=0.563\n    bucket 1: n=9433321 precision=0.547\n    bucket 2: n=9433291 precision=0.593\n  saved overview: /kaggle/working/v6_visualizations/overview_lofo_test1.png\n\n######################################################################\n# FOLD [lofo_test2]  train=['1', '3']  test=2\n######################################################################\n  fragment 1:   [cache warm-up] 34 slices, 3.52 GB read sequentially in 1.3s (2685.8 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 1: train=1458 val=277 mean=95.1 std=64.9\n  fragment 3:   [cache warm-up] 34 slices, 2.71 GB read sequentially in 0.8s (3325.4 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 3: train=1272 val=160 mean=85.8 std=56.1\nPositive patches: 1722 | Negative patches: 1008\nBalanced dataset: 2730 | positive ratio=0.631\n  fragment 2 (held-out):   [cache warm-up] 34 slices, 9.59 GB read sequentially in 3.1s (3059.8 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  hist-match pool (strict protocol): 20 patches from TRAINING fragments ['1', '3'] only -- held-out fragment 2 untouched.\n  curriculum sampling ON: warmup_epochs=8 easy/medium/hard counts = 1224/521/985\n[backbone] DepthSignatureModule: 26 depth slices -> 24 learned channels (LDDC=4 filters, PE dim=8, attention heads=4, pooled to 8x8, stats=True)\n[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320\ndepth_module_type=signature | encoder=tu-convnext_tiny | architecture=unet\nParameters: 32.18M total | 32.18M trainable\nZero-channel layers: 0 (PASS)\nDry-run (256x256) output logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\nCUDA device: Tesla T4\nForward pass: PASS\n======================================================================\n\n    step 17/341 | 33s elapsed | 0.51 steps/s | ETA 638s\n    step 34/341 | 73s elapsed | 0.47 steps/s | ETA 660s\n    step 51/341 | 107s elapsed | 0.48 steps/s | ETA 608s\n    step 68/341 | 142s elapsed | 0.48 steps/s | ETA 571s\n    step 85/341 | 179s elapsed | 0.47 steps/s | ETA 539s\n    step 102/341 | 219s elapsed | 0.47 steps/s | ETA 512s\n    step 119/341 | 254s elapsed | 0.47 steps/s | ETA 474s\n    step 136/341 | 289s elapsed | 0.47 steps/s | ETA 436s\n    step 153/341 | 326s elapsed | 0.47 steps/s | ETA 401s\n    step 170/341 | 363s elapsed | 0.47 steps/s | ETA 365s\n    step 187/341 | 397s elapsed | 0.47 steps/s | ETA 327s\n    step 204/341 | 430s elapsed | 0.47 steps/s | ETA 289s\n    step 221/341 | 467s elapsed | 0.47 steps/s | ETA 254s\n    step 238/341 | 507s elapsed | 0.47 steps/s | ETA 219s\n    step 255/341 | 544s elapsed | 0.47 steps/s | ETA 183s\n    step 272/341 | 580s elapsed | 0.47 steps/s | ETA 147s\n    step 289/341 | 616s elapsed | 0.47 steps/s | ETA 111s\n    step 306/341 | 652s elapsed | 0.47 steps/s | ETA 75s\n    step 323/341 | 689s elapsed | 0.47 steps/s | ETA 38s\n    step 340/341 | 724s elapsed | 0.47 steps/s | ETA 2s\n[01/6] 749.2s | train_loss=0.8456 (consistency=0.0013) | val_f0.5=0.0045 val_dice=0.0018 val_prec=0.1345 val_rec=0.0009\n  *** new best (val F0.5=0.0045) ***\n    step 17/341 | 35s elapsed | 0.49 steps/s | ETA 663s\n    step 34/341 | 74s elapsed | 0.46 steps/s | ETA 672s\n    step 51/341 | 108s elapsed | 0.47 steps/s | ETA 616s\n    step 68/341 | 144s elapsed | 0.47 steps/s | ETA 577s\n    step 85/341 | 180s elapsed | 0.47 steps/s | ETA 543s\n    step 102/341 | 219s elapsed | 0.47 steps/s | ETA 512s\n    step 119/341 | 254s elapsed | 0.47 steps/s | ETA 474s\n    step 136/341 | 289s elapsed | 0.47 steps/s | ETA 436s\n    step 153/341 | 325s elapsed | 0.47 steps/s | ETA 399s\n    step 170/341 | 364s elapsed | 0.47 steps/s | ETA 366s\n    step 187/341 | 400s elapsed | 0.47 steps/s | ETA 329s\n    step 204/341 | 435s elapsed | 0.47 steps/s | ETA 292s\n    step 221/341 | 469s elapsed | 0.47 steps/s | ETA 255s\n    step 238/341 | 506s elapsed | 0.47 steps/s | ETA 219s\n    step 255/341 | 538s elapsed | 0.47 steps/s | ETA 181s\n    step 272/341 | 573s elapsed | 0.47 steps/s | ETA 145s\n    step 289/341 | 610s elapsed | 0.47 steps/s | ETA 110s\n    step 306/341 | 648s elapsed | 0.47 steps/s | ETA 74s\n    step 323/341 | 685s elapsed | 0.47 steps/s | ETA 38s\n    step 340/341 | 719s elapsed | 0.47 steps/s | ETA 2s\n[02/6] 738.6s | train_loss=0.7355 (consistency=0.0014) | val_f0.5=0.0001 val_dice=0.0000 val_prec=0.0868 val_rec=0.0000\n    step 17/341 | 36s elapsed | 0.47 steps/s | ETA 690s\n    step 34/341 | 74s elapsed | 0.46 steps/s | ETA 672s\n    step 51/341 | 111s elapsed | 0.46 steps/s | ETA 632s\n    step 68/341 | 147s elapsed | 0.46 steps/s | ETA 588s\n    step 85/341 | 182s elapsed | 0.47 steps/s | ETA 548s\n    step 102/341 | 222s elapsed | 0.46 steps/s | ETA 519s\n    step 119/341 | 257s elapsed | 0.46 steps/s | ETA 479s\n    step 136/341 | 292s elapsed | 0.47 steps/s | ETA 440s\n    step 153/341 | 329s elapsed | 0.47 steps/s | ETA 404s\n    step 170/341 | 366s elapsed | 0.46 steps/s | ETA 368s\n    step 187/341 | 400s elapsed | 0.47 steps/s | ETA 329s\n    step 204/341 | 435s elapsed | 0.47 steps/s | ETA 292s\n    step 221/341 | 470s elapsed | 0.47 steps/s | ETA 255s\n    step 238/341 | 510s elapsed | 0.47 steps/s | ETA 221s\n    step 255/341 | 544s elapsed | 0.47 steps/s | ETA 183s\n    step 272/341 | 579s elapsed | 0.47 steps/s | ETA 147s\n    step 289/341 | 615s elapsed | 0.47 steps/s | ETA 111s\n    step 306/341 | 653s elapsed | 0.47 steps/s | ETA 75s\n    step 323/341 | 687s elapsed | 0.47 steps/s | ETA 38s\n    step 340/341 | 723s elapsed | 0.47 steps/s | ETA 2s\n[03/6] 744.3s | train_loss=0.6960 (consistency=0.0014) | val_f0.5=0.9231 val_dice=0.0000 val_prec=0.0000 val_rec=0.0000\n  *** new best (val F0.5=0.9231) ***\n    step 17/341 | 36s elapsed | 0.47 steps/s | ETA 688s\n    step 34/341 | 74s elapsed | 0.46 steps/s | ETA 671s\n    step 51/341 | 111s elapsed | 0.46 steps/s | ETA 632s\n    step 68/341 | 148s elapsed | 0.46 steps/s | ETA 594s\n    step 85/341 | 182s elapsed | 0.47 steps/s | ETA 548s\n    step 102/341 | 219s elapsed | 0.47 steps/s | ETA 512s\n    step 119/341 | 254s elapsed | 0.47 steps/s | ETA 474s\n    step 136/341 | 291s elapsed | 0.47 steps/s | ETA 438s\n    step 153/341 | 326s elapsed | 0.47 steps/s | ETA 401s\n    step 170/341 | 364s elapsed | 0.47 steps/s | ETA 366s\n    step 187/341 | 401s elapsed | 0.47 steps/s | ETA 330s\n    step 204/341 | 438s elapsed | 0.47 steps/s | ETA 294s\n    step 221/341 | 474s elapsed | 0.47 steps/s | ETA 258s\n    step 238/341 | 514s elapsed | 0.46 steps/s | ETA 222s\n    step 255/341 | 551s elapsed | 0.46 steps/s | ETA 186s\n    step 272/341 | 588s elapsed | 0.46 steps/s | ETA 149s\n    step 289/341 | 623s elapsed | 0.46 steps/s | ETA 112s\n    step 306/341 | 661s elapsed | 0.46 steps/s | ETA 76s\n    step 323/341 | 697s elapsed | 0.46 steps/s | ETA 39s\n    step 340/341 | 733s elapsed | 0.46 steps/s | ETA 2s\n[04/6] 754.4s | train_loss=0.6845 (consistency=0.0015) | val_f0.5=0.0273 val_dice=0.0111 val_prec=0.9572 val_rec=0.0056\n    step 17/341 | 35s elapsed | 0.49 steps/s | ETA 663s\n    step 34/341 | 73s elapsed | 0.47 steps/s | ETA 659s\n    step 51/341 | 110s elapsed | 0.46 steps/s | ETA 624s\n    step 68/341 | 144s elapsed | 0.47 steps/s | ETA 577s\n    step 85/341 | 179s elapsed | 0.47 steps/s | ETA 539s\n    step 102/341 | 219s elapsed | 0.47 steps/s | ETA 512s\n    step 119/341 | 255s elapsed | 0.47 steps/s | ETA 477s\n    step 136/341 | 291s elapsed | 0.47 steps/s | ETA 438s\n    step 153/341 | 323s elapsed | 0.47 steps/s | ETA 397s\n    step 170/341 | 361s elapsed | 0.47 steps/s | ETA 364s\n    step 187/341 | 397s elapsed | 0.47 steps/s | ETA 327s\n    step 204/341 | 431s elapsed | 0.47 steps/s | ETA 289s\n    step 221/341 | 465s elapsed | 0.48 steps/s | ETA 252s\n    step 238/341 | 501s elapsed | 0.47 steps/s | ETA 217s\n    step 255/341 | 538s elapsed | 0.47 steps/s | ETA 181s\n    step 272/341 | 575s elapsed | 0.47 steps/s | ETA 146s\n    step 289/341 | 609s elapsed | 0.47 steps/s | ETA 110s\n    step 306/341 | 648s elapsed | 0.47 steps/s | ETA 74s\n    step 323/341 | 685s elapsed | 0.47 steps/s | ETA 38s\n    step 340/341 | 720s elapsed | 0.47 steps/s | ETA 2s\n[05/6] 741.4s | train_loss=0.6774 (consistency=0.0016) | val_f0.5=0.4422 val_dice=0.3728 val_prec=0.5048 val_rec=0.2955\n    step 17/341 | 33s elapsed | 0.51 steps/s | ETA 634s\n    step 34/341 | 71s elapsed | 0.48 steps/s | ETA 645s\n    step 51/341 | 105s elapsed | 0.48 steps/s | ETA 599s\n    step 68/341 | 141s elapsed | 0.48 steps/s | ETA 565s\n    step 85/341 | 175s elapsed | 0.49 steps/s | ETA 526s\n    step 102/341 | 214s elapsed | 0.48 steps/s | ETA 502s\n    step 119/341 | 248s elapsed | 0.48 steps/s | ETA 463s\n    step 136/341 | 285s elapsed | 0.48 steps/s | ETA 430s\n    step 153/341 | 322s elapsed | 0.48 steps/s | ETA 395s\n    step 170/341 | 360s elapsed | 0.47 steps/s | ETA 362s\n    step 187/341 | 397s elapsed | 0.47 steps/s | ETA 327s\n    step 204/341 | 433s elapsed | 0.47 steps/s | ETA 291s\n    step 221/341 | 470s elapsed | 0.47 steps/s | ETA 255s\n    step 238/341 | 507s elapsed | 0.47 steps/s | ETA 219s\n    step 255/341 | 542s elapsed | 0.47 steps/s | ETA 183s\n    step 272/341 | 579s elapsed | 0.47 steps/s | ETA 147s\n    step 289/341 | 616s elapsed | 0.47 steps/s | ETA 111s\n    step 306/341 | 654s elapsed | 0.47 steps/s | ETA 75s\n    step 323/341 | 691s elapsed | 0.47 steps/s | ETA 38s\n    step 340/341 | 726s elapsed | 0.47 steps/s | ETA 2s\n\nSelected threshold (validation-only) = 0.43 (val F0.5=0.9985)\nInference patches: 6227\n\nHELD-OUT fragment 2 metrics (frozen threshold=0.43):\n  raw:          {'f0.5': 3.99998407374369e-06, 'dice': 5.930621797692673e-14, 'iou': 5.930621797692673e-14, 'precision': 1.0, 'recall': 5.930621797692673e-14}\n  postprocessed:{'f0.5': 3.99998407374369e-06, 'dice': 5.930621797692673e-14, 'iou': 5.930621797692673e-14, 'precision': 1.0, 'recall': 5.930621797692673e-14}\n  MC-dropout uncertainty vs. precision (low/med/high uncertainty buckets):\n    bucket 0: n=32212795 precision=1.000\n    bucket 1: n=32212792 precision=1.000\n    bucket 2: n=32212822 precision=1.000\n  saved overview: /kaggle/working/v6_visualizations/overview_lofo_test2.png\n\n######################################################################\n# FOLD [lofo_test3]  train=['1', '2']  test=3\n######################################################################\n  fragment 1:   [cache warm-up] 34 slices, 3.52 GB read sequentially in 1.3s (2709.3 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 1: train=1458 val=277 mean=95.1 std=64.9\n  fragment 2:   [cache warm-up] 34 slices, 9.59 GB read sequentially in 2.9s (3331.7 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  fragment 2: train=4559 val=1260 mean=110.9 std=57.8\nPositive patches: 4143 | Negative patches: 1874\nBalanced dataset: 6017 | positive ratio=0.689\n  fragment 3 (held-out):   [cache warm-up] 34 slices, 2.71 GB read sequentially in 0.9s (3110.6 MB/s) -- subsequent random-access patch reads should now mostly be cache hits.\n  hist-match pool (strict protocol): 20 patches from TRAINING fragments ['1', '2'] only -- held-out fragment 3 untouched.\n  curriculum sampling ON: warmup_epochs=8 easy/medium/hard counts = 2898/1320/1799\n[backbone] DepthSignatureModule: 26 depth slices -> 24 learned channels (LDDC=4 filters, PE dim=8, attention heads=4, pooled to 8x8, stats=True)\n[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320\ndepth_module_type=signature | encoder=tu-convnext_tiny | architecture=unet\nParameters: 32.18M total | 32.18M trainable\nZero-channel layers: 0 (PASS)\nDry-run (256x256) output logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\nCUDA device: Tesla T4\nForward pass: PASS\n======================================================================\n\n    step 37/752 | 81s elapsed | 0.46 steps/s | ETA 1560s\n    step 74/752 | 159s elapsed | 0.47 steps/s | ETA 1454s\n    step 111/752 | 234s elapsed | 0.47 steps/s | ETA 1351s\n    step 148/752 | 312s elapsed | 0.47 steps/s | ETA 1273s\n    step 185/752 | 391s elapsed | 0.47 steps/s | ETA 1199s\n    step 222/752 | 471s elapsed | 0.47 steps/s | ETA 1124s\n    step 259/752 | 547s elapsed | 0.47 steps/s | ETA 1042s\n    step 296/752 | 625s elapsed | 0.47 steps/s | ETA 963s\n    step 333/752 | 705s elapsed | 0.47 steps/s | ETA 887s\n    step 370/752 | 786s elapsed | 0.47 steps/s | ETA 811s\n    step 407/752 | 867s elapsed | 0.47 steps/s | ETA 735s\n    step 444/752 | 947s elapsed | 0.47 steps/s | ETA 657s\n    step 481/752 | 1025s elapsed | 0.47 steps/s | ETA 578s\n    step 518/752 | 1109s elapsed | 0.47 steps/s | ETA 501s\n    step 555/752 | 1190s elapsed | 0.47 steps/s | ETA 422s\n    step 592/752 | 1265s elapsed | 0.47 steps/s | ETA 342s\n    step 629/752 | 1345s elapsed | 0.47 steps/s | ETA 263s\n    step 666/752 | 1427s elapsed | 0.47 steps/s | ETA 184s\n    step 703/752 | 1506s elapsed | 0.47 steps/s | ETA 105s\n    step 740/752 | 1587s elapsed | 0.47 steps/s | ETA 26s\n[01/6] 1667.2s | train_loss=0.8063 (consistency=0.0011) | val_f0.5=0.0006 val_dice=0.0000 val_prec=0.0070 val_rec=0.0000\n  *** new best (val F0.5=0.0006) ***\n    step 37/752 | 76s elapsed | 0.49 steps/s | ETA 1472s\n    step 74/752 | 157s elapsed | 0.47 steps/s | ETA 1439s\n    step 111/752 | 237s elapsed | 0.47 steps/s | ETA 1366s\n    step 148/752 | 309s elapsed | 0.48 steps/s | ETA 1261s\n    step 185/752 | 384s elapsed | 0.48 steps/s | ETA 1177s\n    step 222/752 | 465s elapsed | 0.48 steps/s | ETA 1110s\n    step 259/752 | 543s elapsed | 0.48 steps/s | ETA 1033s\n    step 296/752 | 618s elapsed | 0.48 steps/s | ETA 952s\n    step 333/752 | 699s elapsed | 0.48 steps/s | ETA 879s\n    step 370/752 | 780s elapsed | 0.47 steps/s | ETA 805s\n    step 407/752 | 856s elapsed | 0.48 steps/s | ETA 726s\n    step 444/752 | 934s elapsed | 0.48 steps/s | ETA 648s\n    step 481/752 | 1010s elapsed | 0.48 steps/s | ETA 569s\n    step 518/752 | 1092s elapsed | 0.47 steps/s | ETA 493s\n    step 555/752 | 1171s elapsed | 0.47 steps/s | ETA 416s\n    step 592/752 | 1249s elapsed | 0.47 steps/s | ETA 338s\n    step 629/752 | 1326s elapsed | 0.47 steps/s | ETA 259s\n    step 666/752 | 1404s elapsed | 0.47 steps/s | ETA 181s\n    step 703/752 | 1481s elapsed | 0.47 steps/s | ETA 103s\n    step 740/752 | 1554s elapsed | 0.48 steps/s | ETA 25s\n[02/6] 1634.2s | train_loss=0.7215 (consistency=0.0009) | val_f0.5=0.0295 val_dice=0.0121 val_prec=0.7406 val_rec=0.0061\n  *** new best (val F0.5=0.0295) ***\n    step 37/752 | 78s elapsed | 0.48 steps/s | ETA 1501s\n    step 74/752 | 160s elapsed | 0.46 steps/s | ETA 1466s\n    step 111/752 | 237s elapsed | 0.47 steps/s | ETA 1367s\n    step 148/752 | 316s elapsed | 0.47 steps/s | ETA 1290s\n    step 185/752 | 396s elapsed | 0.47 steps/s | ETA 1212s\n    step 222/752 | 474s elapsed | 0.47 steps/s | ETA 1131s\n    step 259/752 | 553s elapsed | 0.47 steps/s | ETA 1053s\n    step 296/752 | 634s elapsed | 0.47 steps/s | ETA 976s\n    step 333/752 | 709s elapsed | 0.47 steps/s | ETA 892s\n    step 370/752 | 791s elapsed | 0.47 steps/s | ETA 817s\n    step 407/752 | 869s elapsed | 0.47 steps/s | ETA 737s\n    step 444/752 | 946s elapsed | 0.47 steps/s | ETA 656s\n    step 481/752 | 1024s elapsed | 0.47 steps/s | ETA 577s\n    step 518/752 | 1105s elapsed | 0.47 steps/s | ETA 499s\n    step 555/752 | 1181s elapsed | 0.47 steps/s | ETA 419s\n    step 592/752 | 1261s elapsed | 0.47 steps/s | ETA 341s\n    step 629/752 | 1339s elapsed | 0.47 steps/s | ETA 262s\n    step 666/752 | 1421s elapsed | 0.47 steps/s | ETA 184s\n    step 703/752 | 1502s elapsed | 0.47 steps/s | ETA 105s\n    step 740/752 | 1582s elapsed | 0.47 steps/s | ETA 26s\n[03/6] 1661.5s | train_loss=0.6960 (consistency=0.0013) | val_f0.5=0.4240 val_dice=0.4429 val_prec=0.4122 val_rec=0.4784\n  *** new best (val F0.5=0.4240) ***\n    step 37/752 | 76s elapsed | 0.49 steps/s | ETA 1472s\n    step 74/752 | 153s elapsed | 0.48 steps/s | ETA 1400s\n    step 111/752 | 234s elapsed | 0.47 steps/s | ETA 1350s\n    step 148/752 | 312s elapsed | 0.47 steps/s | ETA 1272s\n    step 185/752 | 393s elapsed | 0.47 steps/s | ETA 1203s\n    step 222/752 | 472s elapsed | 0.47 steps/s | ETA 1127s\n    step 259/752 | 553s elapsed | 0.47 steps/s | ETA 1053s\n    step 296/752 | 631s elapsed | 0.47 steps/s | ETA 972s\n    step 333/752 | 711s elapsed | 0.47 steps/s | ETA 894s\n    step 370/752 | 790s elapsed | 0.47 steps/s | ETA 816s\n    step 407/752 | 867s elapsed | 0.47 steps/s | ETA 735s\n    step 444/752 | 945s elapsed | 0.47 steps/s | ETA 655s\n    step 481/752 | 1021s elapsed | 0.47 steps/s | ETA 575s\n    step 518/752 | 1102s elapsed | 0.47 steps/s | ETA 498s\n    step 555/752 | 1177s elapsed | 0.47 steps/s | ETA 418s\n    step 592/752 | 1257s elapsed | 0.47 steps/s | ETA 340s\n    step 629/752 | 1336s elapsed | 0.47 steps/s | ETA 261s\n    step 666/752 | 1418s elapsed | 0.47 steps/s | ETA 183s\n    step 703/752 | 1498s elapsed | 0.47 steps/s | ETA 104s\n    step 740/752 | 1572s elapsed | 0.47 steps/s | ETA 25s\n[04/6] 1651.4s | train_loss=0.6829 (consistency=0.0014) | val_f0.5=0.3373 val_dice=0.4265 val_prec=0.2960 val_rec=0.7626\n    step 37/752 | 79s elapsed | 0.47 steps/s | ETA 1527s\n    step 74/752 | 157s elapsed | 0.47 steps/s | ETA 1439s\n    step 111/752 | 235s elapsed | 0.47 steps/s | ETA 1358s\n    step 148/752 | 315s elapsed | 0.47 steps/s | ETA 1284s\n    step 185/752 | 394s elapsed | 0.47 steps/s | ETA 1208s\n    step 222/752 | 475s elapsed | 0.47 steps/s | ETA 1134s\n    step 259/752 | 556s elapsed | 0.47 steps/s | ETA 1058s\n    step 296/752 | 635s elapsed | 0.47 steps/s | ETA 979s\n    step 333/752 | 715s elapsed | 0.47 steps/s | ETA 899s\n    step 370/752 | 793s elapsed | 0.47 steps/s | ETA 819s\n    step 407/752 | 872s elapsed | 0.47 steps/s | ETA 739s\n    step 444/752 | 949s elapsed | 0.47 steps/s | ETA 658s\n    step 481/752 | 1027s elapsed | 0.47 steps/s | ETA 579s\n    step 518/752 | 1111s elapsed | 0.47 steps/s | ETA 502s\n    step 555/752 | 1191s elapsed | 0.47 steps/s | ETA 423s\n    step 592/752 | 1272s elapsed | 0.47 steps/s | ETA 344s\n    step 629/752 | 1352s elapsed | 0.47 steps/s | ETA 264s\n    step 666/752 | 1430s elapsed | 0.47 steps/s | ETA 185s\n    step 703/752 | 1508s elapsed | 0.47 steps/s | ETA 105s\n    step 740/752 | 1586s elapsed | 0.47 steps/s | ETA 26s\n[05/6] 1664.4s | train_loss=0.6627 (consistency=0.0015) | val_f0.5=0.3175 val_dice=0.4129 val_prec=0.2751 val_rec=0.8274\n    step 37/752 | 78s elapsed | 0.48 steps/s | ETA 1500s\n    step 74/752 | 161s elapsed | 0.46 steps/s | ETA 1479s\n    step 111/752 | 239s elapsed | 0.46 steps/s | ETA 1383s\n    step 148/752 | 318s elapsed | 0.47 steps/s | ETA 1296s\n    step 185/752 | 397s elapsed | 0.47 steps/s | ETA 1217s\n    step 222/752 | 478s elapsed | 0.46 steps/s | ETA 1141s\n    step 259/752 | 553s elapsed | 0.47 steps/s | ETA 1053s\n    step 296/752 | 630s elapsed | 0.47 steps/s | ETA 970s\n    step 333/752 | 708s elapsed | 0.47 steps/s | ETA 890s\n    step 370/752 | 790s elapsed | 0.47 steps/s | ETA 816s\n    step 407/752 | 868s elapsed | 0.47 steps/s | ETA 736s\n    step 444/752 | 945s elapsed | 0.47 steps/s | ETA 655s\n    step 481/752 | 1025s elapsed | 0.47 steps/s | ETA 578s\n    step 518/752 | 1106s elapsed | 0.47 steps/s | ETA 500s\n    step 555/752 | 1184s elapsed | 0.47 steps/s | ETA 420s\n    step 592/752 | 1264s elapsed | 0.47 steps/s | ETA 342s\n    step 629/752 | 1340s elapsed | 0.47 steps/s | ETA 262s\n    step 666/752 | 1420s elapsed | 0.47 steps/s | ETA 183s\n    step 703/752 | 1498s elapsed | 0.47 steps/s | ETA 104s\n    step 740/752 | 1577s elapsed | 0.47 steps/s | ETA 26s\n[06/6] 1657.1s | train_loss=0.6632 (consistency=0.0016) | val_f0.5=0.3263 val_dice=0.4223 val_prec=0.2834 val_rec=0.8289\n  early stopping.\n\nSelected threshold (validation-only) = 0.54 (val F0.5=0.4707)\nInference patches: 1645\n\nHELD-OUT fragment 3 metrics (frozen threshold=0.54):\n  raw:          {'f0.5': 0.44898954033851624, 'dice': 0.42522695660591125, 'iou': 0.27002429962158203, 'precision': 0.4663618505001068, 'recall': 0.3907604217529297}\n  postprocessed:{'f0.5': 0.4484195113182068, 'dice': 0.42611682415008545, 'iou': 0.2707423269748688, 'precision': 0.46462997794151306, 'recall': 0.3934996426105499}\n  MC-dropout uncertainty vs. precision (low/med/high uncertainty buckets):\n    bucket 0: n=8138216 precision=0.470\n    bucket 1: n=8138407 precision=0.363\n    bucket 2: n=8138336 precision=0.499\n  saved overview: /kaggle/working/v6_visualizations/overview_lofo_test3.png\n\n======================================================================\nLEAVE-ONE-FRAGMENT-OUT CV SUMMARY\n======================================================================\n  f0.5         = 0.3041 +/- 0.2151   (per-fold: ['0.4637', '0.0000', '0.4484'])\n  dice         = 0.2714 +/- 0.1925   (per-fold: ['0.3880', '0.0000', '0.4261'])\n  iou          = 0.1705 +/- 0.1212   (per-fold: ['0.2407', '0.0000', '0.2707'])\n  precision    = 0.6659 +/- 0.2379   (per-fold: ['0.5332', '1.0000', '0.4646'])\n  recall       = 0.2328 +/- 0.1685   (per-fold: ['0.3049', '0.0000', '0.3935'])\n\nSaved: /kaggle/working/lofo_cv_results.json\n\n=== V6 DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"!pip install zarr==2.12.0","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!pip install segmentation-models-pytorch==0.4.0\nimport segmentation_models_pytorch as smp\nauto_strong_convnext\nmodel = smp.Unet(\n    #encoder_name=\"tu-convnext_tiny\",\n    encoder_name=\"auto_strong_convnext\",\n    encoder_weights=True,\n    in_channels=3,\n    classes=1,\n)\n\nprint(\"SUCCESS!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-06T07:24:00.227918Z","iopub.execute_input":"2026-09-06T07:24:00.228668Z","iopub.status.idle":"2026-09-06T07:24:00.235151Z","shell.execute_reply.started":"2026-09-06T07:24:00.228637Z","shell.execute_reply":"2026-09-06T07:24:00.234155Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_58/3234761623.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m      1\u001b[0m \u001b[0;31m#!pip install segmentation-models-pytorch==0.4.0\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      2\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0msegmentation_models_pytorch\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0msmp\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 3\u001b[0;31m \u001b[0mauto_strong_convnext\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m      4\u001b[0m model = smp.Unet(\n\u001b[1;32m      5\u001b[0m     \u001b[0;31m#encoder_name=\"tu-convnext_tiny\",\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mNameError\u001b[0m: name 'auto_strong_convnext' is not defined"],"ename":"NameError","evalue":"name 'auto_strong_convnext' is not defined","output_type":"error"}],"execution_count":6},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.4.0\n\nimport torch\nimport segmentation_models_pytorch as smp\nimport timm\nimport sys\n\nprint(\"=\" * 70)\nprint(\"ENVIRONMENT\")\nprint(\"=\" * 70)\n\nprint(\"Python       :\", sys.version)\nprint(\"PyTorch      :\", torch.__version__)\nprint(\"Torch CUDA   :\", torch.version.cuda)\nprint(\"SMP          :\", smp.__version__)\nprint(\"timm         :\", timm.__version__)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"GPU\")\nprint(\"=\" * 70)\n\nprint(\"CUDA available:\", torch.cuda.is_available())\n\nif torch.cuda.is_available():\n    print(\"GPU           :\", torch.cuda.get_device_name(0))\n    print(\"Capability    :\", torch.cuda.get_device_capability(0))\n    \n    print(\"\\nCompiled CUDA architectures:\")\n    print(torch.cuda.get_arch_list())\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CONVNEXT\")\nprint(\"=\" * 70)\n\ntry:\n    model = smp.Unet(\n        encoder_name=\"tu-convnext_tiny\",\n        encoder_weights=True,\n        in_channels=3,\n        classes=1,\n    )\n    print(\"tu-convnext_tiny: OK\")\nexcept Exception as e:\n    print(\"tu-convnext_tiny: FAILED\")\n    print(repr(e))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-06T07:36:19.503558Z","iopub.execute_input":"2026-09-06T07:36:19.504395Z","iopub.status.idle":"2026-09-06T07:36:45.278395Z","shell.execute_reply.started":"2026-09-06T07:36:19.504361Z","shell.execute_reply":"2026-09-06T07:36:45.277695Z"}},"outputs":[{"name":"stdout","text":"Collecting segmentation-models-pytorch==0.4.0\n  Downloading segmentation_models_pytorch-0.4.0-py3-none-any.whl.metadata (32 kB)\nCollecting efficientnet-pytorch>=0.6.1 (from segmentation-models-pytorch==0.4.0)\n  Downloading efficientnet_pytorch-0.7.1.tar.gz (21 kB)\n  Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\nRequirement already satisfied: huggingface-hub>=0.24 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.11.0)\nRequirement already satisfied: numpy>=1.19.3 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (2.0.2)\nRequirement already satisfied: pillow>=8 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (11.3.0)\nCollecting pretrainedmodels>=0.7.1 (from segmentation-models-pytorch==0.4.0)\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m5.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\nRequirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.17.0)\nRequirement already satisfied: timm>=0.9 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.0.26)\nRequirement already satisfied: torch>=1.8 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (2.10.0+cu128)\nRequirement already satisfied: torchvision>=0.9 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (0.25.0+cu128)\nRequirement already satisfied: tqdm>=4.42.1 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (4.67.3)\nRequirement already satisfied: filelock>=3.10.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (3.29.0)\nRequirement already satisfied: fsspec>=2023.5.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2025.3.0)\nRequirement already satisfied: hf-xet<2.0.0,>=1.4.3 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.4.3)\nRequirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.28.1)\nRequirement already satisfied: packaging>=20.9 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (26.1)\nRequirement already satisfied: pyyaml>=5.1 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (6.0.3)\nRequirement already satisfied: typer in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.24.2)\nRequirement already satisfied: typing-extensions>=4.1.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.15.0)\nCollecting munch (from pretrainedmodels>=0.7.1->segmentation-models-pytorch==0.4.0)\n  Downloading munch-4.0.0-py2.py3-none-any.whl.metadata (5.9 kB)\nRequirement already satisfied: safetensors in /usr/local/lib/python3.12/dist-packages (from timm>=0.9->segmentation-models-pytorch==0.4.0) (0.7.0)\nRequirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (81.0.0)\nRequirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (1.14.0)\nRequirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.6.1)\nRequirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.1.6)\nRequirement already satisfied: cuda-bindings==12.9.4 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.9.4)\nRequirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.93)\nRequirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (9.10.2.21)\nRequirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.4.1)\nRequirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (11.3.3.83)\nRequirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (10.3.9.90)\nRequirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (11.7.3.90)\nRequirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.5.8.93)\nRequirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (0.7.1)\nRequirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (2.27.5)\nRequirement already satisfied: nvidia-nvshmem-cu12==3.4.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.4.5)\nRequirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.93)\nRequirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (1.13.1.3)\nRequirement already satisfied: triton==3.6.0 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.6.0)\nRequirement already satisfied: cuda-pathfinder~=1.1 in /usr/local/lib/python3.12/dist-packages (from cuda-bindings==12.9.4->torch>=1.8->segmentation-models-pytorch==0.4.0) (1.5.3)\nRequirement already satisfied: anyio in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.13.0)\nRequirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2026.4.22)\nRequirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.0.9)\nRequirement already satisfied: idna in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (3.13)\nRequirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.16.0)\nRequirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch>=1.8->segmentation-models-pytorch==0.4.0) (1.3.0)\nRequirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch>=1.8->segmentation-models-pytorch==0.4.0) (3.0.3)\nRequirement already satisfied: click>=8.2.1 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (8.3.3)\nRequirement already satisfied: shellingham>=1.3.0 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.5.4)\nRequirement already satisfied: rich>=12.3.0 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (13.9.4)\nRequirement already satisfied: annotated-doc>=0.0.2 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.0.4)\nRequirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.0.0)\nRequirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2.20.0)\nRequirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.1.2)\nDownloading segmentation_models_pytorch-0.4.0-py3-none-any.whl (121 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m121.3/121.3 kB\u001b[0m \u001b[31m11.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading munch-4.0.0-py2.py3-none-any.whl (9.9 kB)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25l\u001b[?25hdone\n  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.7.1-py3-none-any.whl size=16477 sha256=def061acabdd5fa96cd29bb2970167b634397d2c981bcccbcb07cb2d02bb7cb9\n  Stored in directory: /root/.cache/pip/wheels/9c/3f/43/e6271c7026fe08c185da2be23c98c8e87477d3db63f41f32ad\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25l\u001b[?25hdone\n  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60990 sha256=a3954cc8956a43b86cd1ae1609d55550dc50661764720ee1fc6acd9975a9e446\n  Stored in directory: /root/.cache/pip/wheels/4c/01/56/40a48f75dbdfe167a0cb70d3b48913369a00ec5c4e9fed5f2b\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: munch, efficientnet-pytorch, pretrainedmodels, segmentation-models-pytorch\nSuccessfully installed efficientnet-pytorch-0.7.1 munch-4.0.0 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.4.0\n======================================================================\nENVIRONMENT\n======================================================================\nPython       : 3.12.13 (main, Mar  4 2026, 09:23:07) [GCC 11.4.0]\nPyTorch      : 2.10.0+cu128\nTorch CUDA   : 12.8\nSMP          : 0.4.0\ntimm         : 1.0.26\n\n======================================================================\nGPU\n======================================================================\nCUDA available: True\nGPU           : Tesla T4\nCapability    : (7, 5)\n\nCompiled CUDA architectures:\n['sm_70', 'sm_75', 'sm_80', 'sm_86', 'sm_90', 'sm_100', 'sm_120']\n\n======================================================================\nCONVNEXT\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"4cc3630776f94f138b42b81f7c57a8e3"}},"metadata":{}},{"name":"stdout","text":"tu-convnext_tiny: OK\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS INK DETECTION - V6 \"DEPTH-SIGNATURE RESEARCH MODEL\"\n# ============================================================\n# Built on V4/V5. V5 declared a rich set of research-grade modules in CFG\n# (DepthSignatureModule / LDDC / depth positional encoding / multi-head depth\n# attention / depth-statistics channels, physics-informed synthetic ink,\n# fiber-consistency loss, clDice topology loss, MC-dropout uncertainty) but the\n# actual forward/training graph only ever ran V4's simpler DepthFusionStem and\n# V2ComboLoss (BCE+Dice+FocalTversky). That mismatch is exactly what a Q1\n# reviewer (or a careful re-implementation) would catch immediately, so V6's\n# only job is to close that gap and add the methodological rigor (leave-one-\n# fragment-out CV, ablation matrix, held-out threshold discipline, MC-dropout\n# uncertainty as a genuine diagnostic) needed to defend this as a paper rather\n# than a Kaggle-optimization script.\n#\n# ------------------------------------------------------------------\n# WHAT IS NOW ACTUALLY WIRED (vs. V5's CFG-only declarations)\n# ------------------------------------------------------------------\n#  1. DepthSignatureModule (depth_module_type=\"signature\") is a real nn.Module\n#     that participates in the forward pass:\n#       - LDDC: learnable 1D convolutions along the depth axis whose kernels\n#         are re-parameterized to sum to zero every forward pass (a hard\n#         constraint, not a hope), giving genuinely learnable derivative-like\n#         operators instead of V4's fixed finite differences.\n#       - Depth positional encoding: sinusoidal encoding of each of the 26\n#         physical depth positions, broadcast spatially and concatenated in.\n#       - Multi-head depth attention: an HONEST engineering compromise is\n#         documented in the class docstring -- full per-pixel (H,W,D,D)\n#         attention is computed at a pooled resolution (not full patch\n#         resolution) because full-resolution per-pixel attention at\n#         patch_size=480 would need ~2TB of activation memory on a single\n#         Tesla T4. The pooled attention map is bilinearly upsampled back to\n#         full resolution. This trade-off is exactly the kind of thing a\n#         methods section must state explicitly rather than let the CFG\n#         silently imply otherwise.\n#       - Depth-statistics channels (mean/std/max/depth-centroid/gradient\n#         energy/curvature energy) computed analytically, not learned.\n#       - MC-dropout channel + optional gradient checkpointing, both real.\n#  2. Physics-informed synthetic ink (USE_PHYSICAL_SYNTHETIC_INK) now ALSO\n#     updates the label at the injected stroke, unlike the old distractor-only\n#     inject_fake_ink (which deliberately never touched the label). The\n#     Gaussian depth profile I(z) = A*exp(-(z-z0)^2/(2*sigma^2)) determines a\n#     continuous per-slice intensity weight, not a hard \"affected slices\"\n#     subset.\n#  3. Fiber-consistency loss (USE_FIBER_CONSISTENCY_LOSS) does a genuine\n#     second forward pass per training step on a fiber-perturbed copy of the\n#     batch and penalizes prediction drift -- this costs real compute, exactly\n#     as documented, and is now actually added into the backward graph.\n#  4. clDice topology loss (USE_TOPOLOGY_LOSS) is a real differentiable soft-\n#     skeletonization loss (Shit et al. 2021) added into the combo loss.\n#  5. Depth-shift invariance (USE_DEPTH_SHIFT_AUG, new in V6, implements\n#     critique section 10): FragmentVolume can read an alternate depth window\n#     shifted by +/- depth_shift_max physical slices; training does a second\n#     forward pass on the shifted window and penalizes prediction drift, so\n#     the model is pushed toward learning \"ink signature\" rather than\n#     \"ink lives at absolute index 18\".\n#  6. MC-dropout uncertainty (USE_MC_DROPOUT_UNCERTAINTY) is a genuine\n#     multi-pass inference routine (only Dropout stays in train mode) that\n#     produces a real per-pixel variance map, plus a precision-vs-uncertainty\n#     bucket analysis (does the model's self-disagreement predict its errors?).\n#  7. Leave-one-fragment-out cross-validation: with 3 fragments there are 3\n#     folds (train on 2, test on the held-out one). V6 wraps the whole\n#     data/model/train/eval pipeline into `run_one_fold(...)` and drives it\n#     three times, reporting mean +/- std and a paired significance test\n#     across folds instead of a single \"best validation Dice\" number.\n#  8. Threshold discipline: the Dice/F0.5 threshold is selected ONLY on that\n#     fold's validation split and then frozen before touching the held-out\n#     fragment. It is never re-tuned on the held-out fragment.\n#  9. F0.5 (not Dice) is treated as the PRIMARY reported metric, matching the\n#     competition's own precision-weighted metric; Dice/IoU/precision/recall/\n#     clDice are reported alongside it.\n# 10. Ablation harness (`run_ablation_matrix`): a small set of named CFG\n#     overrides (baseline -> +DepthFusion -> +LDDC -> +PE -> +Attention ->\n#     +Physics -> +Topology -> +Consistency) run back-to-back on a single\n#     fold's data (cached, not re-downloaded/re-normalized per arm) so the\n#     component-by-component contribution can actually be reported in a\n#     table, which is what a reviewer will ask for first.\n#\n# ------------------------------------------------------------------\n# WHAT V6 DELIBERATELY DOES NOT DO (kept out on purpose, per the review)\n# ------------------------------------------------------------------\n#  - No grid search over hyperparameters, no ad hoc encoder swapping, no\n#    10-model ensemble. TTA/EMA/postprocessing morphology remain labeled as\n#    inference-engineering details, not scientific contributions.\n#  - DANN / MixStyle / AdaBN remain OFF by default and are kept in a clearly\n#    separate \"Protocol B (transductive)\" code path -- they are never silently\n#    mixed into the strict-protocol numbers used for the main CV table.\n# ============================================================\n\n\n# ============================================================\n# V7 CHANGE LOG (applied on top of V6, not a reversion to V4)\n# ============================================================\n# A second review came in on the plain V4 code and proposed a mix of generic\n# and specific fixes. Applied on top of V6 rather than V4, so nothing already\n# fixed (honest depth-signature wiring, LOFO-CV, ablation matrix, threshold\n# discipline) gets thrown away. Adopted vs. pushed-back-on, explicitly:\n#\n# ADOPTED:\n#  - Longer training with a real schedule: CFG.LR_SCHEDULE supports\n#    \"onecycle\" (default) alongside the existing \"cosine\"; CFG.epochs raised\n#    from 8 -- which was genuinely too short for a pretrained ConvNeXt\n#    encoder -- to a configurable default of 30, governed by the SAME\n#    early-stopping patience so it doesn't just run needlessly long.\n#  - A genuinely higher-capacity depth stem as an ADDITIONAL ablation arm:\n#    Conv3DDepthStem (two 3D-conv branches over raw + first-difference\n#    volumes, depth-mean-pooled) -- this is a real alternative to\n#    DepthSignatureModule's attention-based pooling, not a redundant restate\n#    of it, so it's wired in as depth_module_type=\"conv3d\" and added to the\n#    ablation matrix rather than replacing the signature module.\n#  - CutMix (image+mask cut-and-paste) as an optional additional batch-level\n#    augmentation (USE_CUTMIX), cheap and orthogonal to what's already there.\n#  - Curriculum sampling (USE_CURRICULUM): trains on easier (higher ink\n#    fraction) patches proportionally more early on, shifting toward the\n#    full distribution over curriculum_warmup_epochs -- implemented as a\n#    real per-epoch WeightedRandomSampler rebuild, not just a comment.\n#  - A cheap frequency-domain feature (USE_FREQUENCY_FEATURES): radial\n#    high-frequency energy from the depth-averaged image's FFT magnitude,\n#    added as one extra analytic channel in DepthStatsBranch -- thin ink\n#    strokes contribute disproportionately to high spatial frequencies, so\n#    this is a legitimate cheap signal, not the heavier per-slice FFT-fusion\n#    module the critique sketched (which would multiply memory cost for\n#    unclear extra benefit over a single depth-averaged FFT).\n#\n# PUSHED BACK ON (implemented as clearly-labeled OPTIONAL/transductive-only,\n# NOT folded into the main strict-protocol LOFO-CV numbers, with the reason\n# stated here rather than silently ignored):\n#  - \"Turn DANN + MixStyle on, no clear reason they're off\": there IS a clear\n#    reason, stated in V5/V6 already -- they mix held-out-fragment statistics\n#    into the trained weights, which is exactly the leakage the strict/\n#    transductive protocol split exists to prevent. They stay OFF for the\n#    strict-protocol CV table. CFG.PROTOCOL=\"transductive\" remains the\n#    explicit, separate path for anyone who wants to report transductive\n#    numbers ALONGSIDE (never instead of) the strict ones.\n#  - \"Test-Time Training via entropy minimization on each test patch\": this\n#    is *also* a transductive technique (it fits model weights to the target\n#    fragment's unlabeled statistics at test time) -- implemented as\n#    `test_time_training_adapt`, callable only when CFG.PROTOCOL==\"transductive\"\n#    AND CFG.USE_TEST_TIME_TRAINING, producing a separately-labeled\n#    `test_metrics_ttt` that is never averaged into the strict LOFO-CV\n#    summary. Per-single-patch fine-tuning (as literally proposed) would\n#    re-run an optimizer step per inference patch, which is both\n#    prohibitively slow on a T4 for a ~2000-patch fragment and statistically\n#    dubious (no early-stopping signal without labels) -- so this adapts\n#    once, briefly, over unlabeled target-fragment batches instead of per\n#    patch, then runs one normal inference pass.\n#  - \"patch_size=480 is too small for context\" and \"60% negative patches is\n#    imbalance/domain fooling\": both are tunable CFG values already\n#    (patch_size, target_positive_patch_ratio) rather than bugs -- 480px at\n#    this resolution is a substantial spatial context for a segmentation\n#    patch, and the negative/positive ratio is a deliberate class-balance\n#    choice, not an artifact. Left as configurable rather than \"fixed\",\n#    since there's no evidence in the critique that the current values are\n#    actually wrong for this data.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image scipy\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\nfrom collections import defaultdict\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast as _autocast_new, GradScaler as _GradScaler_new\n\n\ndef autocast(enabled=True):\n    \"\"\"V7.3: thin wrapper so every existing `autocast(enabled=...)` call site\n    keeps working unchanged while using torch>=2.x's non-deprecated\n    torch.amp API instead of the deprecated torch.cuda.amp one.\"\"\"\n    device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    return _autocast_new(device_type, enabled=enabled)\n\n\nclass GradScaler(_GradScaler_new):\n    def __init__(self, enabled=True):\n        device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        super().__init__(device_type, enabled=enabled)\nfrom torch.utils.checkpoint import checkpoint as grad_checkpoint\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\nfrom scipy import stats as sstats\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    all_fragments = [\"1\", \"2\", \"3\"]     # used to build the 3-fold LOFO schedule\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    # V6: depth-shift invariance needs slices OUTSIDE depth_indices to shift into.\n    depth_shift_max = 4\n    depth_pool_indices = list(range(min(depth_indices) - depth_shift_max,\n                                     max(depth_indices) + depth_shift_max + 1))\n\n    patch_size = 320\n    train_stride = 128\n    test_stride = 128\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    # V7.1 (perf): 2 workers was almost certainly starving the GPU given this\n    # augmentation pipeline (elastic transform, histogram matching, physics-ink\n    # injection looping over 26 slices -- all CPU-bound). Raise this to your\n    # actual CPU core count minus 1-2 (Kaggle T4 sessions typically give 4\n    # cores -> try 4; a local box with more cores can go higher).\n    num_workers = 2\n    # V7.4 (perf): 4 queued batches per worker was likely contributing to the\n    # system-RAM OOM that forced num_workers down to 2 -- each queued batch\n    # holds (batch_size, 26, H, W) float32 for BOTH img and shifted_img, so at\n    # num_workers=4 this alone was ~2-3GB just sitting in the prefetch queue.\n    # Halving it frees headroom to raise num_workers back up if you want to;\n    # the tradeoff is a smaller read-ahead buffer, which matters only if your\n    # per-sample CPU cost is spiky (use PROFILE_DATASET below to check).\n    prefetch_factor = 2       # only used when num_workers > 0\n    drop_last = True\n\n    # V7: 8 epochs was genuinely too short for a pretrained ConvNeXt encoder.\n    # Raised to 30, still governed by early_stop_patience so it doesn't run\n    # needlessly long once validation F0.5 plateaus.\n    epochs = 6\n    early_stop_patience = 3\n\n    # V7: LR schedule. \"cosine\" = V6's CosineAnnealingLR. \"onecycle\" = warmup\n    # + cosine anneal in one cycle (Smith 2018), generally a better fit for\n    # a short-ish fine-tuning run than plain cosine-from-the-start.\n    LR_SCHEDULE = \"onecycle\"\n    onecycle_pct_start = 0.10\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- backbone / architecture -------------------------------------------\n    encoder_name = \"tu-convnext_tiny\"\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"          # do NOT use \"unetplusplus\" with tu-* encoders (see V4 notes)\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- V4 depth-fusion stem (kept only as an ablation arm / fallback) -----\n    depth_stem_out_channels = 26\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- V7: Conv3DDepthStem output width (see class docstring) -------------\n    conv3d_stem_out_channels = 48\n\n    # --- V6/V7: which depth-aware front end actually builds into the model --\n    # \"signature\" = DepthSignatureModule (LDDC + depth PE + depth attention + stats)\n    # \"conv3d\"    = V7's Conv3DDepthStem (two 3D-conv branches, depth-mean-pooled)\n    # \"fusion\"    = V4's DepthFusionStem (raw + finite-difference grad/curv)\n    # \"none\"      = raw depth stack straight into the encoder\n    depth_module_type = \"signature\"\n\n    lddc_num_filters = 4\n    lddc_kernel_size = 3\n    depth_pe_dim = 8\n    depth_attention_heads = 4\n    depth_attention_pool = 8        # pooled resolution for depth attention (see docstring)\n    USE_DEPTH_STATS = True\n    depth_signature_dropout_p = 0.2\n    # V7.4 (perf): gradient checkpointing trades GPU compute for GPU memory --\n    # it was needed as a safety margin at patch_size=480, but at 320 (your\n    # current setting) the DepthSignatureModule's activations are ~2.25x\n    # smaller, so the memory pressure it was guarding against is much less\n    # likely. Turning it off means one fewer recompute pass through the\n    # depth-signature module per forward, which is a straightforward speed\n    # win at no cost to what the model learns (checkpointing only changes\n    # memory/compute tradeoff, never numerical results). Re-enable it if you\n    # hit a CUDA out-of-memory error.\n    USE_DEPTH_SIGNATURE_CHECKPOINT = False\n    depth_signature_out_channels = 24\n\n    # --- protocol (strict vs transductive; see V5 CFG notes) ----------------\n    PROTOCOL = \"strict\"\n\n    # --- V6: physics-informed synthetic ink (now WITH label update) --------\n    USE_PHYSICAL_SYNTHETIC_INK = True\n    physical_ink_p = 0.15\n    physical_ink_min_amplitude = 10\n    physical_ink_max_amplitude = 60\n    physical_ink_sigma_range = (2.0, 6.0)   # in depth-slice units\n\n    # --- V6: fiber-invariant consistency loss (real second forward pass) ---\n    USE_FIBER_CONSISTENCY_LOSS = True\n    fiber_consistency_weight = 0.10\n    fiber_consistency_amplitude = 0.3       # in normalized (z-scored) units\n\n    # --- V6: depth-shift invariance consistency (new, critique section 10) -\n    USE_DEPTH_SHIFT_CONSISTENCY = True\n    depth_shift_consistency_weight = 0.10\n    # V7.4 (perf): the depth-shift consistency loss is only USED on 1-in-\n    # `consistency_every_n_steps` training steps, but the Dataset was\n    # generating the shifted patch (a full extra disk read across 26 slices,\n    # in a separate worker process) for ~50% of SAMPLES regardless -- wasted\n    # I/O on the ~3-in-4 steps where it's computed and immediately discarded.\n    # Scaling this down to roughly match how often it's actually consumed\n    # keeps the same effective training signal at a fraction of the CPU cost.\n    # This is a coarse per-sample approximation of \"1 in N batches\" (a\n    # dataset __getitem__ has no visibility into which batch/step it's part\n    # of), not an exact match -- raise it back toward 0.5 only if you disable\n    # per-step throttling (consistency_every_n_steps=1).\n    depth_shift_p = 0.5 / max(4, 1)          # ~0.125 by default; ties to the\n                                               # default consistency_every_n_steps below\n\n    # V7.1 (perf): fiber-consistency and depth-shift-consistency each cost a\n    # full extra forward pass through the WHOLE model (ConvNeXt+U-Net), not\n    # just the depth stem -- doing that every single step is why an epoch\n    # went from \"slow\" to \"3159s\". Computing them every Nth step instead keeps\n    # the same training signal (it's a regularizer, not the primary loss) at\n    # a fraction of the cost. Set to 1 to restore V6/V7's original\n    # every-step behavior once you've confirmed this isn't your bottleneck.\n    consistency_every_n_steps = 4\n\n    # V7.1 (perf): prints a one-time timing breakdown (dataloader wait vs.\n    # main forward vs. consistency forward(s) vs. backward) for the first\n    # PROFILE_TIMING_STEPS steps of fold training, then stops. Use this to see\n    # whether YOUR bottleneck is actually the GPU compute added above, or the\n    # CPU-side augmentation pipeline / dataloader instead of guessing.\n    PROFILE_TIMING = True\n    PROFILE_TIMING_STEPS = 8\n\n    # V7.4: per-stage CPU timing inside InkPatchDataset.__getitem__, printed\n    # for the first PROFILE_DATASET_CALLS calls PER WORKER PROCESS (so with\n    # num_workers=2 you'll see ~2x that many interleaved lines -- expected).\n    # Tells you which augmentation stage actually dominates CPU cost instead\n    # of guessing. Turn off once you've identified the bottleneck.\n    PROFILE_DATASET = True\n    PROFILE_DATASET_CALLS = 6\n\n    # --- V6: topology-aware (clDice) loss -----------------------------------\n    USE_TOPOLOGY_LOSS = True\n    topology_weight = 0.15\n    cldice_iters = 8\n\n    # --- V6: MC-Dropout uncertainty (genuine multi-pass diagnostic) --------\n    USE_MC_DROPOUT_UNCERTAINTY = True\n    mc_dropout_passes = 8\n\n    # --- LOSS weights (BCE / Dice / FocalTversky always on; clDice optional) -\n    bce_weight = 0.25\n    dice_weight = 0.30\n    focal_tversky_weight = 0.30\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n    # NOTE: bce_weight+dice_weight+focal_tversky_weight+topology_weight should\n    # sum to ~1.0; topology_weight is added on top and the others renormalized\n    # implicitly by training dynamics -- kept explicit rather than hidden.\n\n    # --- THRESHOLD (selected on validation only, frozen before test) --------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n\n    # --- TTA (inference engineering, not a scientific contribution) --------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN / DANN / MixStyle: OFF for the main strict-protocol CV table -\n    use_adabn = False\n    adabn_max_patches = 2000\n    USE_DANN = False\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n    #USE_MIXSTYLE = False\n    USE_MIXSTYLE = True \n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    # --- POSTPROCESS (inference engineering) --------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- CLAHE ---------------------------------------------------------------\n    CLAHE_MODE = \"global_shared\"     # \"off\" | \"per_slice\" | \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization: multi-slice sampled stats, not just the mid slice ---\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n    frag_stats_sample_slices = 9     # V6: sample across depth, not just mid slice\n\n    # --- domain-randomization augmentation (unchanged from V4) -------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n    USE_SHADOW = True\n    shadow_p = 0.20\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15      # distractor-only fake ink (no label update) -- kept\n                            # for robustness training, separate from physical ink\n    fake_fiber_p = 0.15\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- V7: CutMix (batch-level, orthogonal to the existing per-patch augs) -\n    USE_CUTMIX = True\n    cutmix_p = 0.20\n    cutmix_alpha = 1.0\n\n    # --- V7: curriculum sampling (easy -> full distribution over N epochs) --\n    USE_CURRICULUM = True\n    curriculum_warmup_epochs = 8\n    curriculum_easy_ink_frac = 0.10     # >= this ink fraction counts \"easy\"\n\n    # --- V7: cheap frequency-domain feature (radial high-freq FFT energy) ---\n    USE_FREQUENCY_FEATURES = True\n\n    # --- V7: Test-Time Training -- PROTOCOL=\"transductive\" ONLY. Never mixed\n    # into the strict-protocol LOFO-CV numbers (see V7 change-log docstring).\n    USE_TEST_TIME_TRAINING = False\n    ttt_steps = 20\n    ttt_lr = 1e-5\n    ttt_batch_size = 4\n\n    USE_EMA = True\n    ema_decay = 0.999\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    viz_dir = os.path.join(out_dir, \"v6_visualizations\")\n\n    # --- V6: what to run -----------------------------------------------------\n    # \"lofo_cv\"   : 3-fold leave-one-fragment-out CV (main scientific result)\n    # \"ablation\"  : component ablation matrix on ONE fold\n    # \"single\"    : one train_frags/test_frag run (fast debugging)\n    RUN_MODE = \"lofo_cv\"\n    single_train_frags = [\"2\", \"3\"]\n    single_test_frag = \"1\"\n    ablation_test_frag = \"1\"\n    ablation_train_frags = [\"2\", \"3\"]\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\nclass CfgOverride:\n    \"\"\"Context manager: temporarily overrides CFG attributes, restores them on\n    exit. Used by the ablation harness so each arm is a clean, reproducible\n    CFG state rather than hand-editing globals between runs.\"\"\"\n    def __init__(self, **overrides):\n        self.overrides = overrides\n        self.previous = {}\n\n    def __enter__(self):\n        for k, v in self.overrides.items():\n            self.previous[k] = getattr(CFG, k)\n            setattr(CFG, k, v)\n        return CFG\n\n    def __exit__(self, *exc):\n        for k, v in self.previous.items():\n            setattr(CFG, k, v)\n        return False\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\ndef compute_sample_difficulty(samples, labels_full, patch_size, easy_ink_frac):\n    \"\"\"V7: per-sample difficulty tier for curriculum sampling -- 0.0 (easy:\n    wide ink coverage), 0.5 (medium), 1.0 (hard: thin/sparse ink or none).\n    Index-aligned with `samples` (the same list used to build the training\n    Dataset), so it can be used directly as WeightedRandomSampler weights.\"\"\"\n    difficulties = np.zeros(len(samples), dtype=np.float32)\n    for i, (fid, y, x) in enumerate(samples):\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        if frac >= easy_ink_frac:\n            difficulties[i] = 0.0\n        elif frac > 0:\n            difficulties[i] = 0.5\n        else:\n            difficulties[i] = 1.0\n    return difficulties\n\n\ndef build_curriculum_sampler(difficulties, epoch, total_epochs, warmup_epochs):\n    \"\"\"Returns per-sample WeightedRandomSampler weights. Early epochs\n    strongly favor easy/medium samples; by `warmup_epochs` the weighting has\n    linearly relaxed to uniform (i.e. the full, already-class-balanced\n    distribution `balance_positive_patches` built) -- curriculum learning is\n    meant to warm the model up, not to permanently exclude hard examples.\"\"\"\n    progress = min(epoch / max(warmup_epochs, 1), 1.0)   # 0 -> 1 over warmup\n    # weight(difficulty=1.0) goes from a small floor up to 1.0 (uniform) as\n    # progress -> 1; weight(difficulty=0.0) stays at 1.0 throughout.\n    hard_floor = 0.15\n    weights = 1.0 - (1.0 - hard_floor) * (1.0 - progress) * difficulties\n    weights = np.clip(weights, hard_floor, 1.0)\n    return torch.DoubleTensor(weights)\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME (extended: depth pool for shift-invariance training)\n# ============================================================\n\nclass FragmentVolume:\n    \"\"\"V6 change: opens every slice in CFG.depth_pool_indices (the default\n    26-slice window PLUS +/- depth_shift_max on each side), not just\n    CFG.depth_indices. read_patch() accepts an explicit z_indices list so the\n    depth-shift-consistency training step can request a shifted window of the\n    SAME physical stack without re-opening files.\"\"\"\n\n    def __init__(self, frag_dir, pool_indices, default_indices):\n        self.pool_indices = pool_indices\n        self.default_indices = default_indices\n        self.index_to_pos = {z: i for i, z in enumerate(pool_indices)}\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in pool_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid_pos = self.index_to_pos.get(\n            self.default_indices[len(self.default_indices) // 2],\n            len(self._slices) // 2)\n        ref_slice = np.asarray(self._slices[mid_pos])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(\n            ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def _slice_positions(self, z_indices):\n        positions = []\n        for z in z_indices:\n            z_clamped = min(max(z, self.pool_indices[0]), self.pool_indices[-1])\n            positions.append(self.index_to_pos[z_clamped])\n        return positions\n\n    def read_patch(self, y, x, size, z_indices=None, apply_clahe=None):\n        self._ensure_open()\n        z_indices = self.default_indices if z_indices is None else z_indices\n        positions = self._slice_positions(z_indices)\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(positions), size, size), dtype=np.uint8)\n        for out_i, pos in enumerate(positions):\n            s = self._slices[pos]\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[out_i] = block\n        return out\n\n    def sample_shifted_indices(self, max_shift):\n        \"\"\"Returns a physically-shifted (but still contiguous, still ordered)\n        window of depth indices, clamped to stay inside the opened pool.\"\"\"\n        delta = random.randint(-max_shift, max_shift)\n        shifted = [z + delta for z in self.default_indices]\n        lo, hi = self.pool_indices[0], self.pool_indices[-1]\n        shifted = [min(max(z, lo), hi) for z in shifted]\n        return shifted\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        \"\"\"V6: samples statistics across several depth slices (not just the mid\n        slice) for a more robust per-fragment normalization constant.\"\"\"\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        n_slices = min(CFG.frag_stats_sample_slices, len(self.default_indices))\n        sample_z = sorted(random.sample(self.default_indices, n_slices))\n        sample_positions = self._slice_positions(sample_z)\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            for pos in sample_positions:\n                block = self._slices[pos][y:y + patch_size, x:x + patch_size]\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n                vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\ndef make_fragment_volume(frag_dir):\n    return FragmentVolume(frag_dir, CFG.depth_pool_indices, CFG.depth_indices)\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Distractor-only fake ink: perturbs intensity but NEVER touches the\n    label. Used purely as a robustness/negative-hallucination stress test --\n    kept separate from the physics-informed version below, which DOES update\n    the label because it is meant to represent an actual (simulated) ink\n    deposit, not a distractor.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_physical_ink_with_label(img_hwd, label_hw, depth_indices, n_strokes=None):\n    \"\"\"V6: the physics-informed synthetic ink model actually promised in V5's\n    CFG. Models a stroke's cross-depth intensity profile as a Gaussian\n    I(z) = A * exp(-(z - z0)^2 / (2*sigma^2)) with A, sigma, z0 sampled per\n    stroke, applies it as a CONTINUOUS per-slice weight (not a hard \"these N\n    slices are affected\" cutoff), and -- unlike inject_fake_ink -- writes the\n    stroke into the label as well, since this is meant to represent a\n    simulated real ink deposit rather than a distractor.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    label_out = label_hw.copy()\n    n_strokes = n_strokes or random.randint(1, 2)\n    z_arr = np.asarray(depth_indices, dtype=np.float32)\n\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n\n        A_amp = random.uniform(CFG.physical_ink_min_amplitude, CFG.physical_ink_max_amplitude)\n        sign = random.choice([-1, 1])\n        sigma = random.uniform(*CFG.physical_ink_sigma_range)\n        z0 = random.uniform(z_arr.min(), z_arr.max())\n\n        profile = sign * A_amp * np.exp(-((z_arr - z0) ** 2) / (2.0 * sigma ** 2))  # (d,)\n\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        stroke_bool = stroke_mask > 0\n        if not stroke_bool.any():\n            continue\n\n        for zi in range(d):\n            delta = profile[zi]\n            if abs(delta) < 0.5:\n                continue\n            sl = out[:, :, zi].astype(np.float32)\n            sl[stroke_bool] = np.clip(sl[stroke_bool] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n\n        # Only mark label where the Gaussian profile has real amplitude near\n        # the profile's peak (>= 40% of |A|) -- a near-zero-weight slice at\n        # the tail of the Gaussian isn't meaningfully \"ink\" at that slice, but\n        # the 2D label is depth-collapsed anyway, so we mark the stroke\n        # footprint whenever the profile is non-trivial anywhere in depth.\n        if np.abs(profile).max() >= 0.4 * A_amp:\n            label_out[stroke_bool] = 1\n\n    return out, label_out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\ndef cutmix_batch(imgs, masks):\n    \"\"\"V7: batch-level CutMix on the already-loaded tensors (image + label\n    mask cut together, so this stays label-consistent unlike a naive\n    image-only cutmix). Orthogonal to the existing per-patch augmentations\n    (which perturb ONE patch); this mixes TWO patches within a batch.\"\"\"\n    B = imgs.size(0)\n    if B < 2:\n        return imgs, masks\n    lam = float(np.random.beta(CFG.cutmix_alpha, CFG.cutmix_alpha))\n    rand_index = torch.randperm(B, device=imgs.device)\n\n    H, W = imgs.shape[-2:]\n    cut_ratio = math.sqrt(max(1.0 - lam, 1e-6))\n    cut_h, cut_w = int(H * cut_ratio), int(W * cut_ratio)\n    cy, cx = np.random.randint(H), np.random.randint(W)\n    y1, y2 = max(0, cy - cut_h // 2), min(H, cy + cut_h // 2)\n    x1, x2 = max(0, cx - cut_w // 2), min(W, cx + cut_w // 2)\n\n    imgs = imgs.clone()\n    masks = masks.clone()\n    imgs[:, :, y1:y2, x1:x2] = imgs[rand_index][:, :, y1:y2, x1:x2]\n    masks[:, :, y1:y2, x1:x2] = masks[rand_index][:, :, y1:y2, x1:x2]\n    return imgs, masks\n\n\ndef add_fiber_pattern_tensor(img_tensor, amplitude):\n    \"\"\"Tensor-level fiber perturbation (in normalized z-scored units) used by\n    the fiber-consistency loss so the perturbation stays inside the autograd\n    graph without a second numpy round-trip. img_tensor: (B, D, H, W).\"\"\"\n    B, D, H, W = img_tensor.shape\n    device = img_tensor.device\n    theta = torch.rand(B, device=device) * math.pi\n    freq = 0.02 + torch.rand(B, device=device) * 0.04\n    amp = amplitude * (0.5 + torch.rand(B, device=device))\n\n    yy, xx = torch.meshgrid(torch.arange(H, device=device, dtype=torch.float32),\n                             torch.arange(W, device=device, dtype=torch.float32), indexing=\"ij\")\n    yy = yy.unsqueeze(0)   # (1,H,W)\n    xx = xx.unsqueeze(0)\n\n    phase = (xx * torch.cos(theta).view(B, 1, 1) + yy * torch.sin(theta).view(B, 1, 1)) \\\n        * freq.view(B, 1, 1)\n    pattern = amp.view(B, 1, 1) * torch.sin(2 * math.pi * phase)   # (B,H,W)\n    pattern = pattern.unsqueeze(1).expand(-1, D, -1, -1)            # (B,D,H,W)\n    return img_tensor + pattern\n\n\n# ============================================================\n# 8. AUGMENTATION\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    \"\"\"V6 changes: (1) physics-informed ink now updates the label and is\n    applied via inject_physical_ink_with_label, kept independent from the\n    label-blind inject_fake_ink distractor; (2) optionally returns a\n    depth-shifted alternate view of the same patch for the shift-consistency\n    loss (train_mode only).\"\"\"\n\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n        self._profile_calls_left = CFG.PROFILE_DATASET_CALLS if CFG.PROFILE_DATASET else 0\n\n    def __len__(self):\n        return len(self.samples)\n\n    def _prep(self, img_u8_hwd, label_hw, normalize_stats):\n        img = img_u8_hwd.astype(np.float32) / 255.0\n        img = normalize_patch(img, *normalize_stats)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label_hw > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n    def __getitem__(self, idx):\n        # V7.4: per-stage CPU timing, printed for the first\n        # CFG.PROFILE_DATASET_CALLS calls in EACH worker process (so with\n        # num_workers=2 you'll see ~2x that many lines total, interleaved --\n        # that's expected, not a bug). This exists to answer \"which\n        # augmentation is actually slow\" empirically instead of guessing;\n        # set CFG.PROFILE_DATASET=False once you've identified the culprit.\n        do_prof = self.train_mode and self._profile_calls_left > 0\n        if do_prof:\n            self._profile_calls_left -= 1\n            _t = time.time()\n            def _lap(label, timings=[]):\n                nonlocal _t\n                now = time.time()\n                timings.append((label, now - _t))\n                _t = now\n                return timings\n            _timings = []\n\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size].copy()\n        img = np.transpose(patch, (1, 2, 0))   # HWD\n        if do_prof:\n            _timings = _lap(\"main_patch_read\")\n\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n        if do_prof:\n            _timings = _lap(\"hist_match\")\n\n        if self.train_mode and CFG.USE_PHYSICAL_SYNTHETIC_INK and random.random() < CFG.physical_ink_p:\n            img, label = inject_physical_ink_with_label(img, label, CFG.depth_indices)\n        if do_prof:\n            _timings = _lap(\"physical_ink\")\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n        if do_prof:\n            _timings = _lap(\"fake_ink_fiber_shadow\")\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n        if do_prof:\n            _timings = _lap(\"albumentations_transform\")\n\n        img_t, label_t = self._prep(img, label, (vol.frag_mean, vol.frag_std))\n        if do_prof:\n            _timings = _lap(\"prep_to_tensor\")\n\n        shifted_t = None\n        if self.train_mode and CFG.USE_DEPTH_SHIFT_CONSISTENCY and random.random() < CFG.depth_shift_p:\n            shifted_z = vol.sample_shifted_indices(CFG.depth_shift_max)\n            shifted_patch = vol.read_patch(y, x, size, z_indices=shifted_z)\n            shifted_hwd = np.transpose(shifted_patch, (1, 2, 0)).astype(np.float32) / 255.0\n            shifted_hwd = normalize_patch(shifted_hwd, vol.frag_mean, vol.frag_std)\n            shifted_t = torch.from_numpy(np.ascontiguousarray(np.transpose(shifted_hwd, (2, 0, 1))))\n        if do_prof:\n            _timings = _lap(\"depth_shift_patch_read\")\n\n        if shifted_t is None:\n            has_shift = torch.tensor(False)\n            shifted_t = torch.zeros_like(img_t)\n        else:\n            has_shift = torch.tensor(True)\n\n        if do_prof:\n            total = sum(t for _, t in _timings)\n            breakdown = \" | \".join(f\"{name}={t*1000:.1f}ms\" for name, t in _timings)\n            print(f\"    [dataset profile pid={os.getpid()}] TOTAL={total*1000:.1f}ms :: {breakdown}\")\n\n        return img_t, label_t, shifted_t, has_shift\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN (Protocol B only): yields normalized (D,H,W) tensors from the\n    held-out fragment, NO labels.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"V4's simpler depth-aware stem, kept as an ablation arm / fallback.\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = 0\n        if use_raw:\n            total_in += in_depth\n        if use_grad:\n            total_in += (in_depth - 1)\n        if use_curv:\n            total_in += max(in_depth - 2, 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.in_depth = in_depth\n\n    def forward(self, x):\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:, :, :] - x[:, :-1, :, :])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass Conv3DDepthStem(nn.Module):\n    \"\"\"V7 addition: a genuinely higher-capacity alternative depth stem,\n    addressing the '26->16 is a severe information bottleneck' critique with\n    real extra convolutional capacity rather than just a wider 1x1 mix. Two\n    branches (raw depth stack, first-difference depth stack) each go through\n    two 3D conv layers BEFORE any depth-collapsing, then are mean-pooled over\n    depth and concatenated -- unlike a single 1x1 conv straight from 70\n    finite-difference channels down to 16, the 3D convs get real learnable\n    interaction across depth and space before anything is collapsed. This is\n    wired in as its own `depth_module_type=\"conv3d\"` ablation arm alongside\n    DepthSignatureModule, not a replacement for it -- they represent two\n    different hypotheses (attention-based depth pooling vs. convolutional\n    depth pooling) worth comparing empirically, which is exactly what the\n    ablation matrix is for.\"\"\"\n    def __init__(self, in_depth, out_channels=48):\n        super().__init__()\n        half = out_channels // 2\n        self.raw_branch = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=3, padding=1), nn.BatchNorm3d(16), nn.GELU(),\n            nn.Conv3d(16, half, kernel_size=3, padding=1), nn.BatchNorm3d(half), nn.GELU(),\n        )\n        self.grad_branch = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=3, padding=1), nn.BatchNorm3d(16), nn.GELU(),\n            nn.Conv3d(16, out_channels - half, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels - half), nn.GELU(),\n        )\n        self.out_channels = out_channels\n\n    def forward(self, x):   # x: (B, D, H, W)\n        raw = x.unsqueeze(1)                                  # (B,1,D,H,W)\n        grad = (x[:, 1:, :, :] - x[:, :-1, :, :]).unsqueeze(1)  # (B,1,D-1,H,W)\n\n        raw_feat = self.raw_branch(raw).mean(dim=2)            # (B,half,H,W)\n        grad_feat = self.grad_branch(grad).mean(dim=2)         # (B,out-half,H,W)\n        return torch.cat([raw_feat, grad_feat], dim=1)\n\n\nclass LDDC(nn.Module):\n    \"\"\"Learnable Depth Differential Convolution. Learns `num_filters` 1D\n    kernels along the physically-ordered depth axis. Each kernel is\n    re-parameterized every forward pass to sum to zero:\n        w' = w - mean(w)\n    which is a hard constraint (not a training-time regularizer that might\n    only be approximately satisfied) -- guaranteeing every filter behaves like\n    a (learned) derivative operator regardless of what the raw weights drift\n    to during optimization. Implemented as a Conv3d with kernel (k,1,1) over\n    x.unsqueeze(1): (B,1,D,H,W) -> (B,num_filters,D,H,W).\n    \"\"\"\n    def __init__(self, num_filters=4, kernel_size=3):\n        super().__init__()\n        self.num_filters = num_filters\n        self.kernel_size = kernel_size\n        self.weight = nn.Parameter(torch.randn(num_filters, 1, kernel_size, 1, 1) * 0.1)\n        self.bias = nn.Parameter(torch.zeros(num_filters))\n\n    def forward(self, x):   # x: (B, D, H, W)\n        w = self.weight - self.weight.mean(dim=2, keepdim=True)   # zero-sum constraint\n        xin = x.unsqueeze(1)   # (B,1,D,H,W)\n        out = F.conv3d(xin, w, bias=self.bias, padding=(self.kernel_size // 2, 0, 0))\n        return out   # (B, num_filters, D, H, W)\n\n\nclass DepthPositionalEncoding(nn.Module):\n    \"\"\"Sinusoidal encoding of each physical depth position z, broadcast\n    spatially. Returns (B, pe_dim, D, H, W) so it can be concatenated\n    alongside LDDC's depth-preserving output before depth attention pools it\n    down to a 2D feature map.\"\"\"\n    def __init__(self, pe_dim=8):\n        super().__init__()\n        assert pe_dim % 2 == 0\n        self.pe_dim = pe_dim\n        div_term = torch.exp(torch.arange(0, pe_dim, 2).float() * (-math.log(10000.0) / pe_dim))\n        self.register_buffer(\"div_term\", div_term, persistent=False)\n\n    def forward(self, depth_positions, B, H, W, device):\n        # depth_positions: 1D float tensor of physical z indices, length D\n        z = depth_positions.to(device).view(-1, 1)                     # (D,1)\n        angles = z * self.div_term.view(1, -1).to(device)               # (D, pe_dim/2)\n        pe = torch.zeros(z.shape[0], self.pe_dim, device=device)\n        pe[:, 0::2] = torch.sin(angles)\n        pe[:, 1::2] = torch.cos(angles)\n        # (D, pe_dim) -> (1, pe_dim, D, 1, 1) -> broadcast to (B, pe_dim, D, H, W)\n        pe = pe.transpose(0, 1).view(1, self.pe_dim, -1, 1, 1)\n        return pe.expand(B, -1, -1, H, W)\n\n\nclass PooledDepthAttention(nn.Module):\n    \"\"\"Multi-head attention ACROSS the depth axis, at every spatial location.\n\n    Honest engineering note (this is exactly the kind of thing the review\n    flagged as needing to be stated explicitly): true per-pixel attention\n    across D=26 depth positions needs an (H, W, heads, D, D) tensor. At\n    patch_size=480 with heads=4, D=26, batch=8 that is B*H*W*heads*D*D*4 bytes\n    ~= 2 TB of activation memory -- not something a single Tesla T4 (14.5GB)\n    can hold, full stop. So this module computes genuine per-pixel softmax\n    attention across depth at a POOLED spatial resolution (default 480/8=60),\n    where the memory cost (8*60*60*4*26*26*4 bytes ~= 93MB) is trivial, and\n    then bilinearly upsamples the resulting per-position depth-attention\n    output back to full resolution. This still gives spatially-varying,\n    per-location depth attention (unlike a pure squeeze-and-excite global\n    version), just not literally per-pixel -- state this trade-off in the\n    methods section rather than letting the module name imply otherwise.\n    \"\"\"\n    def __init__(self, in_channels, heads=4, pool_size=8):\n        super().__init__()\n        self.heads = heads\n        self.pool_size = pool_size\n        self.head_dim = max(in_channels // heads, 4)\n        inner = self.heads * self.head_dim\n        self.to_q = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.to_k = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.to_v = nn.Conv1d(in_channels, inner, kernel_size=1)\n        self.out_proj = nn.Conv1d(inner, in_channels, kernel_size=1)\n\n    def forward(self, x):   # x: (B, C, D, H, W)\n        B, C, D, H, W = x.shape\n        ph = max(1, round(H / self.pool_size))\n        pw = max(1, round(W / self.pool_size))\n        x_pooled = F.adaptive_avg_pool3d(x, output_size=(D, max(1, H // ph), max(1, W // pw)))\n        _, _, _, ph_, pw_ = x_pooled.shape\n\n        # reshape depth axis into the \"sequence\" dimension for a 1D attention\n        # per spatial location: (B*ph_*pw_, C, D)\n        xp = x_pooled.permute(0, 3, 4, 1, 2).reshape(B * ph_ * pw_, C, D)\n        q = self.to_q(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n        k = self.to_k(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n        v = self.to_v(xp).view(B * ph_ * pw_, self.heads, self.head_dim, D)\n\n        attn = torch.einsum(\"nhdi,nhdj->nhij\", q, k) / math.sqrt(self.head_dim)\n        attn = attn.softmax(dim=-1)                                    # (N, heads, D, D)\n        out = torch.einsum(\"nhij,nhdj->nhdi\", attn, v)                 # (N, heads, head_dim, D)\n        out = out.reshape(B * ph_ * pw_, self.heads * self.head_dim, D)\n        out = self.out_proj(out)                                       # (N, C, D)\n        out = out.view(B, ph_, pw_, C, D).permute(0, 3, 4, 1, 2)        # (B, C, D, ph_, pw_)\n\n        out = out.reshape(B, C * D, ph_, pw_)\n        out = F.interpolate(out, size=(H, W), mode=\"bilinear\", align_corners=False)\n        out = out.view(B, C, D, H, W)\n        return out\n\n\nclass DepthStatsBranch(nn.Module):\n    \"\"\"Analytic (non-learned) depth-statistics channels: mean, std, max,\n    depth-centroid (intensity-weighted mean z), gradient energy, curvature\n    energy, and (V7, optional) radial high-frequency FFT energy. Computed\n    directly from the raw depth stack so they carry signal even before\n    LDDC/attention have learned anything useful early in training.\n\n    V7's frequency channel: thin ink strokes contribute disproportionately to\n    high spatial frequencies compared to broad fiber/background texture, so a\n    single cheap FFT magnitude computed on the depth-AVERAGED image (not a\n    separate FFT per slice -- that would multiply memory cost for unclear\n    extra benefit) gives one extra, genuinely informative analytic channel.\n    \"\"\"\n    def __init__(self, use_frequency=False):\n        super().__init__()\n        self.use_frequency = use_frequency\n        self.out_channels = 6 + (1 if use_frequency else 0)\n\n    @staticmethod\n    def _radial_high_freq_energy(mean_img, high_freq_frac=0.5):\n        # mean_img: (B, 1, H, W). Returns (B, 1, H, W) -- the same scalar\n        # (per-sample high-frequency energy fraction) broadcast spatially, so\n        # it can be concatenated as a \"channel\" alongside genuinely spatial\n        # stats without pretending to carry spatial variation it doesn't have.\n        B, _, H, W = mean_img.shape\n        fft = torch.fft.rfft2(mean_img.squeeze(1).float(), norm=\"ortho\")\n        mag = torch.abs(fft)   # (B, H, W//2+1)\n        fy = torch.fft.fftfreq(H, device=mean_img.device).view(H, 1)\n        fx = torch.fft.rfftfreq(W, device=mean_img.device).view(1, -1)\n        radius = torch.sqrt(fy ** 2 + fx ** 2)\n        radius = radius / radius.max().clamp_min(1e-6)\n        high_mask = (radius >= high_freq_frac).float()\n        high_energy = (mag * high_mask).sum(dim=(1, 2))\n        total_energy = mag.sum(dim=(1, 2)).clamp_min(1e-6)\n        frac = (high_energy / total_energy).view(B, 1, 1, 1).expand(-1, 1, H, W)\n        return frac.to(mean_img.dtype)\n\n    def forward(self, x, depth_positions):   # x: (B, D, H, W)\n        B, D, H, W = x.shape\n        mean = x.mean(dim=1, keepdim=True)\n        std = x.std(dim=1, keepdim=True)\n        maxv = x.max(dim=1, keepdim=True).values\n\n        z = depth_positions.to(x.device).view(1, D, 1, 1)\n        weights = F.softmax(x, dim=1)\n        centroid = (weights * z).sum(dim=1, keepdim=True)\n\n        grad = x[:, 1:, :, :] - x[:, :-1, :, :]\n        grad_energy = (grad ** 2).mean(dim=1, keepdim=True)\n\n        if D > 2:\n            curv = x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :]\n            curv_energy = (curv ** 2).mean(dim=1, keepdim=True)\n        else:\n            curv_energy = torch.zeros_like(mean)\n\n        parts = [mean, std, maxv, centroid, grad_energy, curv_energy]\n        if self.use_frequency:\n            parts.append(self._radial_high_freq_energy(mean))\n        return torch.cat(parts, dim=1)   # (B, 6 or 7, H, W)\n\n\nclass DepthSignatureModule(nn.Module):\n    \"\"\"V5's promised (and now actually implemented) depth-aware front end:\n    LDDC + depth positional encoding + pooled multi-head depth attention +\n    analytic depth-statistics channels, mixed down to `out_channels` for the\n    2D encoder. Includes the MC-dropout channel used both for regularization\n    during training and for genuine uncertainty estimation at inference (see\n    run_mc_dropout_uncertainty). Gradient checkpointing is applied to the\n    (D-preserving, memory-heavy) LDDC+PE+attention stack when\n    USE_DEPTH_SIGNATURE_CHECKPOINT is True.\"\"\"\n    def __init__(self, in_depth, depth_positions, out_channels=24):\n        super().__init__()\n        self.in_depth = in_depth\n        self.register_buffer(\"depth_positions\", torch.tensor(depth_positions, dtype=torch.float32),\n                              persistent=False)\n\n        self.lddc = LDDC(CFG.lddc_num_filters, CFG.lddc_kernel_size)\n        self.pe = DepthPositionalEncoding(CFG.depth_pe_dim)\n        lddc_pe_channels = CFG.lddc_num_filters + CFG.depth_pe_dim\n        self.attn = PooledDepthAttention(lddc_pe_channels, heads=CFG.depth_attention_heads,\n                                          pool_size=CFG.depth_attention_pool)\n        self.stats = DepthStatsBranch(use_frequency=CFG.USE_FREQUENCY_FEATURES) if CFG.USE_DEPTH_STATS else None\n\n        depth_collapsed_channels = lddc_pe_channels * in_depth       # attention output, D collapsed via mean\n        stats_channels = self.stats.out_channels if self.stats is not None else 0\n        mix_in = depth_collapsed_channels + stats_channels + in_depth  # + raw stack\n\n        self.dropout = nn.Dropout2d(CFG.depth_signature_dropout_p)\n        self.mix = nn.Sequential(\n            nn.Conv2d(mix_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n\n    def _lddc_pe_attn(self, x):\n        B, D, H, W = x.shape\n        lddc_out = self.lddc(x)                                          # (B, F, D, H, W)\n        pe_out = self.pe(self.depth_positions, B, H, W, x.device)          # (B, pe, D, H, W)\n        combined = torch.cat([lddc_out, pe_out], dim=1)                    # (B, F+pe, D, H, W)\n        attended = self.attn(combined)                                     # (B, F+pe, D, H, W)\n        # collapse depth by taking mean over the attended depth axis for the\n        # final 2D feature map, but keep the FULL (F+pe)*D as concatenated\n        # channels too -- mean loses information, so we use both.\n        collapsed = attended.reshape(B, -1, H, W)                          # (B, (F+pe)*D, H, W)\n        return collapsed\n\n    def forward(self, x):   # x: (B, D, H, W)\n        if CFG.USE_DEPTH_SIGNATURE_CHECKPOINT and self.training:\n            collapsed = grad_checkpoint(self._lddc_pe_attn, x, use_reentrant=False)\n        else:\n            collapsed = self._lddc_pe_attn(x)\n\n        parts = [x, collapsed]\n        if self.stats is not None:\n            parts.append(self.stats(x, self.depth_positions))\n\n        feat = torch.cat(parts, dim=1)\n        feat = self.dropout(feat)\n        return self.mix(feat)\n\n\nclass DepthAwareSegModel(nn.Module):\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):   # x: (B, D, H, W)\n        feat = self.depth_stem(x) if self.depth_stem is not None else x\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' requested, but only \"\n              f\"'unet' has been verified to handle tu-* timm encoders' 0-channel placeholder \"\n              f\"stage correctly -- using UnetPlusPlus may crash with ConvNeXt/other tu-* encoders.\")\n    return arch_cls(\n        encoder_name=encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=in_channels,\n        classes=1,\n        decoder_channels=CFG.decoder_channels,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\ndef build_model():\n    stem = None\n    seg_in_channels = CFG.in_channels\n\n    if CFG.depth_module_type == \"signature\":\n        stem = DepthSignatureModule(CFG.in_channels, CFG.depth_indices, CFG.depth_signature_out_channels)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] DepthSignatureModule: {CFG.in_channels} depth slices -> \"\n              f\"{seg_in_channels} learned channels (LDDC={CFG.lddc_num_filters} filters, \"\n              f\"PE dim={CFG.depth_pe_dim}, attention heads={CFG.depth_attention_heads}, \"\n              f\"pooled to {CFG.depth_attention_pool}x{CFG.depth_attention_pool}, \"\n              f\"stats={CFG.USE_DEPTH_STATS})\")\n    elif CFG.depth_module_type == \"fusion\":\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] DepthFusionStem: {CFG.in_channels} -> {seg_in_channels} channels\")\n    elif CFG.depth_module_type == \"conv3d\":\n        stem = Conv3DDepthStem(CFG.in_channels, CFG.conv3d_stem_out_channels)\n        seg_in_channels = stem.out_channels\n        print(f\"[backbone] Conv3DDepthStem: {CFG.in_channels} -> {seg_in_channels} channels \"\n              f\"(two 3D-conv branches, depth-mean-pooled)\")\n    elif CFG.depth_module_type == \"none\":\n        stem = None\n        seg_in_channels = CFG.in_channels\n        print(f\"[backbone] no depth-aware stem: raw {CFG.in_channels}-channel stack into encoder\")\n    else:\n        raise ValueError(f\"Unknown depth_module_type={CFG.depth_module_type}\")\n\n    try:\n        seg_model = _build_seg_backbone(CFG.encoder_name, seg_in_channels)\n        print(f\"[backbone] using {CFG.encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {CFG.encoder_name} unavailable/failed ({e}); \"\n              f\"falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model)\n\n\nclass EMAModel:\n    def __init__(self, model, decay=CFG.ema_decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. ARCHITECTURE INSPECTION\n# ============================================================\n\ndef inspect_model_channels(model):\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\n@torch.no_grad()\ndef run_architecture_report(model):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"V6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\")\n    print(\"=\" * 70)\n\n    n_params = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Input:  depth slices={CFG.in_channels} | patch={CFG.patch_size}x{CFG.patch_size}\")\n    print(f\"depth_module_type={CFG.depth_module_type} | encoder={CFG.encoder_name} | \"\n          f\"architecture={CFG.architecture}\")\n    print(f\"Parameters: {n_params/1e6:.2f}M total | {n_trainable/1e6:.2f}M trainable\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"\\nZERO/INVALID CHANNEL LAYERS FOUND:\")\n        for item in bad_layers:\n            print(\"  \", item)\n        raise RuntimeError(\"Model contains invalid zero-channel layers -- fix before training.\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    model.eval()\n    small_size = min(CFG.patch_size, 256)\n    dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n    out = model(dummy)\n    print(f\"Dry-run ({small_size}x{small_size}) output logits shape: {tuple(out.shape)}\")\n    has_nan = torch.isnan(out).any().item()\n    has_inf = torch.isinf(out).any().item()\n    print(f\"NaN check: {'FAIL' if has_nan else 'PASS'} | Inf check: {'FAIL' if has_inf else 'PASS'}\")\n    if has_nan or has_inf:\n        raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass -- fix before training.\")\n    model.train()\n\n    if torch.cuda.is_available():\n        print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")\n    print(\"Forward pass: PASS\")\n    print(\"=\" * 70 + \"\\n\")\n\n\n# ============================================================\n# 12. LOSSES (BCE + Dice + FocalTversky + clDice)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\ndef _soft_erode(img):\n    return -F.max_pool2d(-img, kernel_size=3, stride=1, padding=1)\n\n\ndef _soft_dilate(img):\n    return F.max_pool2d(img, kernel_size=3, stride=1, padding=1)\n\n\ndef _soft_open(img):\n    return _soft_dilate(_soft_erode(img))\n\n\ndef soft_skeletonize(img, iters):\n    \"\"\"Differentiable soft-skeletonization (Shit et al. 2021, 'clDice -- A\n    Novel Topology-Preserving Loss Function for Tubular Structure\n    Segmentation'), used for the clDice topology-aware loss below.\"\"\"\n    img1 = _soft_open(img)\n    skel = F.relu(img - img1)\n    for _ in range(iters):\n        img = _soft_erode(img)\n        img1 = _soft_open(img)\n        delta = F.relu(img - img1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef cldice_loss(logits, targets, iters=8, eps=1e-6):\n    probs = torch.sigmoid(logits)\n    skel_pred = soft_skeletonize(probs, iters)\n    skel_true = soft_skeletonize(targets, iters)\n    t_prec = (skel_pred * targets).sum() / (skel_pred.sum() + eps)\n    t_sens = (skel_true * probs).sum() / (skel_true.sum() + eps)\n    cldice = 1.0 - (2.0 * t_prec * t_sens) / (t_prec + t_sens + eps)\n    return cldice\n\n\nclass VXComboLoss(nn.Module):\n    \"\"\"BCE + Dice + FocalTversky, with clDice added when USE_TOPOLOGY_LOSS is\n    on -- unlike V5's CFG, topology_weight now actually reaches the backward\n    graph.\"\"\"\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        loss = CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n        cldice_val = None\n        if CFG.USE_TOPOLOGY_LOSS:\n            cldice_val = cldice_loss(logits, targets, CFG.cldice_iters)\n            loss = loss + CFG.topology_weight * cldice_val\n        return loss, {\"bce\": bce.item(), \"dice\": dice.item(), \"focal_tversky\": tv.item(),\n                       \"cldice\": (cldice_val.item() if cldice_val is not None else None)}\n\n\n# ============================================================\n# 13. METRICS  (F0.5 is the PRIMARY reported metric)\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25   # F0.5\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"f0.5\": float(fbeta), \"dice\": float(dice), \"iou\": float(iou),\n            \"precision\": float(precision), \"recall\": float(recall)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 14. TRAIN / VALIDATION EPOCH (fiber + depth-shift consistency wired in)\n# ============================================================\n\ndef run_epoch(model, criterion, loader, optimizer, scaler, ema, domain_classifier,\n              feature_capture, dann_iter, global_step_holder, total_steps,\n              train_mode=True, threshold=0.5, scheduler=None, scheduler_steps_per_batch=False,\n              profile_steps=0):\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_consistency_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    # V7.4 (essential, not cosmetic): without this, a slow epoch prints\n    # NOTHING between the profiled first `profile_steps` batches and the\n    # epoch-summary line -- exactly the \"goes dark for 90 minutes with no way\n    # to tell if it's working or stuck\" problem. This is separate from\n    # `profile_steps`: it runs every epoch (not just epoch 1), has no\n    # cuda.synchronize() overhead, and reports wall-clock throughput/ETA so a\n    # slow-but-alive run is distinguishable from a genuinely hung one.\n    epoch_t_start = time.time()\n    n_loader_batches = len(loader)\n    progress_every = max(1, n_loader_batches // 20)   # ~20 prints per epoch\n\n    t_data_end = time.time()\n    for batch_idx, (imgs, masks, shifted_imgs, has_shift) in enumerate(loader):\n        if train_mode and n_loader_batches > 0 and (batch_idx + 1) % progress_every == 0:\n            elapsed = time.time() - epoch_t_start\n            rate = (batch_idx + 1) / max(elapsed, 1e-6)\n            eta = (n_loader_batches - (batch_idx + 1)) / max(rate, 1e-6)\n            print(f\"    step {batch_idx + 1}/{n_loader_batches} | \"\n                  f\"{elapsed:.0f}s elapsed | {rate:.2f} steps/s | ETA {eta:.0f}s\")\n\n        do_profile = train_mode and profile_steps > 0 and batch_idx < profile_steps\n        if do_profile:\n            t0 = time.time()\n            data_wait = t0 - t_data_end\n            if CFG.device == \"cuda\":\n                torch.cuda.synchronize()\n\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        shifted_imgs = shifted_imgs.to(CFG.device, non_blocking=True)\n        has_shift = has_shift.to(CFG.device, non_blocking=True)\n\n        # V7.1 (perf): consistency losses only computed every Nth step -- see\n        # CFG.consistency_every_n_steps docstring for why.\n        do_consistency = train_mode and (batch_idx % max(CFG.consistency_every_n_steps, 1) == 0)\n\n        # V7.2 (bugfix, flagged by review): CutMix is applied to `imgs` here,\n        # but `shifted_imgs` comes straight from the dataset and is NEVER\n        # cutmixed (the Dataset builds it independently of this loop). If\n        # depth-shift consistency then compared model(shifted_imgs) against\n        # model(imgs) on a step where imgs got cutmixed, it would be\n        # comparing predictions on a DIFFERENT spatial composition -- not the\n        # same sample at a shifted depth window, which defeats the point of\n        # the consistency loss (and would train it against noise). Fiber\n        # consistency is unaffected: fiber_imgs is built FROM `imgs` after\n        # cutmix, so it stays the same sample as `logits`. Fix: track\n        # whether cutmix fired this step and skip ONLY depth-shift\n        # consistency when it did.\n        did_cutmix = False\n        if train_mode and CFG.USE_CUTMIX and random.random() < CFG.cutmix_p:\n            imgs, masks = cutmix_batch(imgs, masks)\n            did_cutmix = True\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss, loss_parts = criterion(logits, masks)\n                consistency_total = torch.zeros((), device=CFG.device)\n\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_main_fwd = time.time()\n\n                # --- fiber-consistency loss (real 2nd forward pass) ---------\n                # safe under CutMix: fiber_imgs is derived FROM imgs (the\n                # possibly-cutmixed tensor), so both sides of the comparison\n                # are the same sample.\n                if do_consistency and CFG.USE_FIBER_CONSISTENCY_LOSS:\n                    fiber_imgs = add_fiber_pattern_tensor(imgs, CFG.fiber_consistency_amplitude)\n                    fiber_logits = model(fiber_imgs)\n                    fiber_consistency = F.mse_loss(torch.sigmoid(fiber_logits), torch.sigmoid(logits.detach()))\n                    consistency_total = consistency_total + CFG.fiber_consistency_weight * fiber_consistency\n\n                # --- depth-shift consistency loss (real 2nd forward pass) ---\n                # NOT safe under CutMix (see comment above) -- skipped this step.\n                #\n                # V7.3 (critical perf bugfix): this USED to forward only the\n                # has_shift subset (`shifted_imgs[idx]`), whose size is random\n                # every time (0-8, since depth_shift_p=0.5 per sample). With\n                # cudnn.benchmark=True, EVERY new batch size the model sees\n                # triggers a fresh convolution-algorithm benchmark search --\n                # for this model (ConvNeXt+U-Net+3D depth stem) that can cost\n                # seconds to tens of seconds PER NEW SIZE, plus growing\n                # per-shape workspace memory. Over hundreds of steps hitting\n                # sizes 1..8 repeatedly in no particular order, this compounds\n                # into exactly the kind of multi-hour stall that doesn't show\n                # up in a short profiling window (the profiled steps 0 and 4\n                # already show this: 94s and 4s of one-off cost). Fixed: ALWAYS\n                # forward the full, fixed-size batch (same shape as the main/\n                # fiber paths, which is why THOSE stayed fast at ~0.27s/1.2s\n                # steady-state) and mask out the non-shifted samples in the\n                # loss instead of indexing them out of the tensor.\n                if do_consistency and (not did_cutmix) and CFG.USE_DEPTH_SHIFT_CONSISTENCY:\n                    shift_logits = model(shifted_imgs)   # fixed shape: (CFG.batch_size, D, H, W)\n                    shift_probs = torch.sigmoid(shift_logits)\n                    main_probs_detached = torch.sigmoid(logits.detach())\n                    per_sample_mse = F.mse_loss(shift_probs, main_probs_detached,\n                                                 reduction=\"none\").mean(dim=[1, 2, 3])\n                    weight = has_shift.float()\n                    denom = weight.sum().clamp_min(1.0)\n                    shift_consistency = (per_sample_mse * weight).sum() / denom\n                    consistency_total = consistency_total + CFG.depth_shift_consistency_weight * shift_consistency\n                    del shift_logits, shift_probs, main_probs_detached, per_sample_mse\n\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_consist_fwd = time.time()\n\n                loss = loss + consistency_total\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if do_profile:\n                    if CFG.device == \"cuda\":\n                        torch.cuda.synchronize()\n                    t_backward = time.time()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    # V7.3: GradScaler can SKIP the actual optimizer.step() on\n                    # steps where it detects inf/nan gradients (routine during\n                    # AMP's initial loss-scale calibration, typically just the\n                    # first few iterations) -- if we call scheduler.step()\n                    # unconditionally after that, OneCycleLR advances one step\n                    # further than the optimizer actually did, which is what\n                    # the \"lr_scheduler.step() before optimizer.step()\"\n                    # warning is reporting. Comparing the scaler's scale\n                    # before/after detects a skipped step so we skip the\n                    # scheduler step too, keeping the two in sync.\n                    prev_scale = scaler.get_scale()\n                    scaler.step(optimizer)\n                    scaler.update()\n                    step_was_skipped = scaler.get_scale() < prev_scale\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    if scheduler_steps_per_batch and scheduler is not None and not step_was_skipped:\n                        scheduler.step()\n                    global_step_holder[0] += 1\n\n        if do_profile:\n            print(f\"  [profile step {batch_idx}] data_wait={data_wait:.3f}s \"\n                  f\"main_fwd={t_main_fwd - t0:.3f}s \"\n                  f\"consistency_fwd={t_consist_fwd - t_main_fwd:.3f}s \"\n                  f\"(consistency_computed={do_consistency}) \"\n                  f\"backward+step={t_backward - t_consist_fwd:.3f}s \"\n                  f\"TOTAL={t_backward - t0:.3f}s\")\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_consistency_loss += float(consistency_total.detach().item()) if train_mode else 0.0\n\n        del imgs, masks, logits, probs, shifted_imgs, has_shift\n        t_data_end = time.time()\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n        prev_scale = scaler.get_scale()\n        scaler.step(optimizer)\n        scaler.update()\n        step_was_skipped = scaler.get_scale() < prev_scale\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n        # V7.2: this trailing partial-accumulation step is also a real\n        # optimizer step -- OneCycleLR needs scheduler.step() called here\n        # too, or the last step of every epoch silently goes unaccounted for.\n        if scheduler_steps_per_batch and scheduler is not None and not step_was_skipped:\n            scheduler.step()\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_consistency = total_consistency_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_consistency\n\n\n# ============================================================\n# 15. THRESHOLD SEARCH (validation-only, frozen before test)\n# ============================================================\n\n@torch.no_grad()\ndef find_best_threshold(model, loader, metric=\"f0.5\"):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks, _, _ in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    prec = (tp + 1e-6) / (tp + fp + 1e-6)\n    rec = (tp + 1e-6) / (tp + fn + 1e-6)\n    f05 = (1.25 * prec * rec + 1e-6) / (0.25 * prec + rec + 1e-6)\n    scores = f05 if metric == \"f0.5\" else dice\n    best_idx = int(np.argmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), thresholds, scores\n\n\n# ============================================================\n# 16. MC-DROPOUT UNCERTAINTY (genuine multi-pass diagnostic)\n# ============================================================\n\ndef _set_dropout_train(model):\n    for m in model.modules():\n        if isinstance(m, (nn.Dropout, nn.Dropout2d, nn.Dropout3d)):\n            m.train()\n\n\n@torch.no_grad()\ndef run_mc_dropout_uncertainty(model, vol, mask, patch_size, stride, passes):\n    \"\"\"Runs `passes` stochastic forward passes (only Dropout layers stay in\n    train mode; BatchNorm/LayerNorm stay in eval mode) over a sliding window\n    and returns (mean_prob_map, uncertainty_map) where uncertainty is the\n    per-pixel variance across passes.\"\"\"\n    model.eval()\n    _set_dropout_train(model)\n\n    H, W = mask.shape\n    sum_map = np.zeros((H, W), dtype=np.float32)\n    sumsq_map = np.zeros((H, W), dtype=np.float32)\n    weight_map = np.zeros((H, W), dtype=np.float32)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        x_t = torch.from_numpy(raw).unsqueeze(0).to(CFG.device)\n        pass_probs = []\n        for _ in range(passes):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(x_t)\n            pass_probs.append(torch.sigmoid(logits)[0, 0].float().cpu().numpy())\n        pass_probs = np.stack(pass_probs, axis=0)   # (passes, size, size)\n        mean_p = pass_probs.mean(axis=0)\n        sq_p = (pass_probs ** 2).mean(axis=0)\n        sum_map[y:y + patch_size, x:x + patch_size] += mean_p\n        sumsq_map[y:y + patch_size, x:x + patch_size] += sq_p\n        weight_map[y:y + patch_size, x:x + patch_size] += 1.0\n        del x_t, pass_probs\n\n    weight_map[weight_map <= 1e-8] = 1.0\n    mean_prob = sum_map / weight_map\n    mean_sq = sumsq_map / weight_map\n    uncertainty = np.clip(mean_sq - mean_prob ** 2, 0, None)\n\n    model.eval()   # restore full eval mode (dropout off) for any subsequent calls\n    return mean_prob, uncertainty\n\n\ndef analyze_uncertainty_vs_error(prob_map, uncertainty_map, gt, mask, threshold, n_buckets=3):\n    \"\"\"Buckets pixels by uncertainty and reports precision in each bucket --\n    directly answers 'does uncertainty predict where the model is wrong?'\"\"\"\n    valid = mask > 0\n    unc = uncertainty_map[valid]\n    prob = prob_map[valid]\n    gtv = gt[valid]\n    preds = (prob > threshold).astype(np.float32)\n\n    quantiles = np.quantile(unc, np.linspace(0, 1, n_buckets + 1))\n    report = []\n    for i in range(n_buckets):\n        lo, hi = quantiles[i], quantiles[i + 1]\n        bucket = (unc >= lo) & (unc <= hi) if i == n_buckets - 1 else (unc >= lo) & (unc < hi)\n        if bucket.sum() == 0:\n            continue\n        p_bucket = preds[bucket]\n        g_bucket = gtv[bucket]\n        tp = (p_bucket * g_bucket).sum()\n        fp = (p_bucket * (1 - g_bucket)).sum()\n        precision = (tp + 1e-6) / (tp + fp + 1e-6)\n        report.append({\"bucket\": i, \"uncertainty_range\": (float(lo), float(hi)),\n                        \"n_pixels\": int(bucket.sum()), \"precision\": float(precision)})\n    return report\n\n\n# ============================================================\n# 16b. TEST-TIME TRAINING (V7, Protocol=\"transductive\" ONLY -- see change log)\n# ============================================================\n\ndef test_time_training_adapt(model, test_vol, test_coords_unlabeled):\n    \"\"\"Adapts a COPY of the trained model to the held-out fragment's\n    unlabeled statistics via entropy minimization, for CFG.ttt_steps batches.\n    Returns the adapted copy; the caller's original `model` (and therefore\n    the strict-protocol evaluation) is never touched.\n\n    This is a transductive technique -- it fits weights to target-domain\n    unlabeled data -- so it is ONLY ever invoked when\n    CFG.PROTOCOL == \"transductive\" and CFG.USE_TEST_TIME_TRAINING, and its\n    output is stored as a separately-labeled `test_metrics_ttt`, never\n    averaged into the strict LOFO-CV summary (see `summarize_folds`, which\n    only reads `test_metrics_postprocessed`).\"\"\"\n    assert CFG.PROTOCOL == \"transductive\", (\n        \"test_time_training_adapt called outside Protocol B -- refusing, since \"\n        \"this would silently leak target-fragment statistics into a \"\n        \"strict-protocol result.\")\n\n    adapted = copy.deepcopy(model)\n    adapted.train()\n    optimizer = torch.optim.SGD(adapted.parameters(), lr=CFG.ttt_lr)\n\n    ds = UnlabeledPatchDataset(test_vol, test_coords_unlabeled, CFG.patch_size)\n    loader = DataLoader(ds, batch_size=CFG.ttt_batch_size, shuffle=True, num_workers=1, drop_last=True)\n    loader_iter = iter(loader)\n\n    print(f\"[TTT] adapting on {CFG.ttt_steps} unlabeled target-fragment batches \"\n          f\"(lr={CFG.ttt_lr}) ...\")\n    for step in range(CFG.ttt_steps):\n        try:\n            batch = next(loader_iter)\n        except StopIteration:\n            loader_iter = iter(loader)\n            batch = next(loader_iter)\n        batch = batch.to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = adapted(batch)\n            probs = torch.sigmoid(logits).clamp(1e-6, 1 - 1e-6)\n            entropy = -(probs * torch.log(probs) + (1 - probs) * torch.log(1 - probs)).mean()\n        entropy.backward()\n        torch.nn.utils.clip_grad_norm_(adapted.parameters(), CFG.grad_clip)\n        optimizer.step()\n        optimizer.zero_grad(set_to_none=True)\n        del batch, logits, probs\n\n    adapted.eval()\n    return adapted\n\n\n# ============================================================\n# 17. INFERENCE HELPERS (sliding window, TTA, postprocessing)\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats_, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats_[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 18. ONE-FOLD TRAIN/EVAL (the core unit LOFO-CV and ablation both call)\n# ============================================================\n\ndef run_one_fold(train_frags, test_frag, fold_tag, save_viz=False):\n    \"\"\"Builds data, trains, selects threshold on validation only, evaluates\n    on the held-out fragment, and returns a metrics dict. This is the single\n    unit both `run_lofo_cv` and `run_ablation_matrix` call, so every arm goes\n    through IDENTICAL code -- only CFG differs between calls.\"\"\"\n    set_seed(CFG.seed)\n    print(\"\\n\" + \"#\" * 70)\n    print(f\"# FOLD [{fold_tag}]  train={train_frags}  test={test_frag}\")\n    print(\"#\" * 70)\n\n    # ---- build train data ----\n    train_volumes, train_labels_full, train_masks_full = {}, {}, {}\n    train_samples_raw, val_samples = [], []\n\n    for fid in train_frags:\n        frag_dir = os.path.join(CFG.base_dir, fid)\n        mask = load_tissue_mask(frag_dir)\n        labels = load_ink_labels(frag_dir)\n        if labels is None:\n            raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n        vol = make_fragment_volume(frag_dir)\n        if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n            vol.compute_fragment_stats(mask, CFG.patch_size)\n\n        coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n        tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n\n        train_volumes[fid] = vol\n        train_labels_full[fid] = labels\n        train_masks_full[fid] = mask\n        train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n        val_samples.extend([(fid, y, x) for y, x in va_coords])\n        print(f\"  fragment {fid}: train={len(tr_coords)} val={len(va_coords)} \"\n              f\"mean={vol.frag_mean:.1f} std={vol.frag_std:.1f}\")\n        del mask\n        cleanup_memory()\n\n    train_samples = balance_positive_patches(\n        train_samples_raw, train_labels_full, CFG.patch_size,\n        positive_threshold=CFG.positive_patch_fraction,\n        target_positive_ratio=CFG.target_positive_patch_ratio,\n        max_positive_repeat=CFG.max_positive_repeat)\n\n    # ---- held-out fragment: unlabeled use only, labels loaded LATE ----\n    test_dir = os.path.join(CFG.base_dir, test_frag)\n    test_mask = load_tissue_mask(test_dir)\n    test_vol = make_fragment_volume(test_dir)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n\n    test_coords_unlabeled = generate_grid_coords(test_mask, CFG.patch_size, CFG.test_stride,\n                                                  CFG.min_tissue_frac_train)\n\n    # V7.2 (bugfix, flagged by review): the histogram-match reference pool was\n    # being built from the HELD-OUT fragment's own image patches and fed\n    # straight into the TRAINING dataset as an augmentation reference -- even\n    # without labels, that lets the model's training-time inputs be reshaped\n    # to look like the test fragment's intensity distribution, which is\n    # exactly the leakage the strict/transductive protocol split exists to\n    # prevent (and which the CFG.PROTOCOL docstring already claimed doesn't\n    # happen). Fixed: in \"strict\" protocol the pool is built from the\n    # TRAINING fragments' own patches; the held-out fragment's patches are\n    # only used for this purpose under PROTOCOL == \"transductive\", where\n    # that's the explicit, clearly-labeled point of the experiment.\n    hist_match_pool = None\n    if CFG.USE_HIST_MATCH_AUG:\n        if CFG.PROTOCOL == \"strict\":\n            pool_source_coords = []\n            for fid in train_frags:\n                coords = generate_grid_coords(train_masks_full[fid], CFG.patch_size, CFG.patch_size,\n                                               CFG.min_tissue_frac_train)\n                pool_source_coords.extend([(fid, y, x) for y, x in coords])\n            n_pool = min(CFG.hist_match_pool_size, len(pool_source_coords))\n            pool_samples = random.sample(pool_source_coords, n_pool) if pool_source_coords else []\n            hist_match_pool = [train_volumes[fid].read_patch(y, x, CFG.patch_size)\n                                for fid, y, x in pool_samples]\n            print(f\"  hist-match pool (strict protocol): {len(hist_match_pool)} patches from \"\n                  f\"TRAINING fragments {train_frags} only -- held-out fragment {test_frag} untouched.\")\n        elif CFG.PROTOCOL == \"transductive\":\n            pool_coords = random.sample(test_coords_unlabeled,\n                                         min(CFG.hist_match_pool_size, len(test_coords_unlabeled)))\n            hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n            print(f\"  hist-match pool (TRANSDUCTIVE protocol, by design): {len(hist_match_pool)} \"\n                  f\"patches from held-out fragment {test_frag}.\")\n        else:\n            raise ValueError(f\"Unknown CFG.PROTOCOL={CFG.PROTOCOL}\")\n\n    # ---- datasets / loaders ----\n    train_transform = build_train_transform()\n    train_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                                transform=train_transform, jitter=CFG.train_jitter,\n                                hist_match_pool=hist_match_pool, train_mode=True)\n    val_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                              transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\n    sample_difficulties = None\n    curriculum_sampler = None\n    if CFG.USE_CURRICULUM:\n        sample_difficulties = compute_sample_difficulty(\n            train_samples, train_labels_full, CFG.patch_size, CFG.curriculum_easy_ink_frac)\n        print(f\"  curriculum sampling ON: warmup_epochs={CFG.curriculum_warmup_epochs} \"\n              f\"easy/medium/hard counts = \"\n              f\"{int((sample_difficulties==0).sum())}/{int((sample_difficulties==0.5).sum())}/\"\n              f\"{int((sample_difficulties==1).sum())}\")\n        # V7.1 (perf): build ONE sampler object and mutate its `.weights`\n        # in-place each epoch instead of recreating the DataLoader (which\n        # respawns worker processes from scratch every epoch -- expensive\n        # with tifffile-backed volumes and a heavy CPU augmentation pipeline).\n        curriculum_sampler = torch.utils.data.WeightedRandomSampler(\n            build_curriculum_sampler(sample_difficulties, 0, CFG.epochs, CFG.curriculum_warmup_epochs),\n            num_samples=len(train_ds), replacement=True)\n        train_loader = DataLoader(\n            train_ds, batch_size=CFG.batch_size, sampler=curriculum_sampler,\n            num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n            drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0,\n            prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n    else:\n        train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                                   num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                                   drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0,\n                                   prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                             num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                             drop_last=False, persistent_workers=CFG.num_workers > 0,\n                             prefetch_factor=(CFG.prefetch_factor if CFG.num_workers > 0 else None))\n\n    # ---- model / loss / optimizer ----\n    model = build_model().to(CFG.device)\n    run_architecture_report(model)\n\n    pos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\n    with torch.no_grad():\n        bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n        try:\n            model.seg_model.segmentation_head[0].bias.fill_(bias_val)\n        except Exception as e:\n            print(f\"  (could not set output bias directly: {e})\")\n\n    raw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\n    pos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\n    pos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\n    criterion = VXComboLoss(pos_weight=pos_weight)\n\n    encoder_params, decoder_params, stem_params = [], [], []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if name.startswith(\"depth_stem.\"):\n            stem_params.append(param)\n        elif name.startswith(\"seg_model.encoder.\"):\n            encoder_params.append(param)\n        else:\n            decoder_params.append(param)\n\n    param_groups = [\n        {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n        {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n    ]\n    if stem_params:\n        param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\n    optimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\n    if CFG.LR_SCHEDULE == \"onecycle\":\n        max_lrs = [g[\"lr\"] for g in param_groups]\n        # V7.2 (bugfix, flagged by review): scheduler.step() only fires once\n        # per OPTIMIZER update (i.e. once every accumulation_steps batches),\n        # not once per batch -- so telling OneCycleLR steps_per_epoch=\n        # len(train_loader) overstates the cycle length whenever\n        # accumulation_steps > 1, desynchronizing the LR curve from the\n        # actual number of optimizer steps taken. (With the default\n        # accumulation_steps=1 this was a no-op, but it's a real bug for\n        # anyone who raises accumulation_steps, which is a normal thing to\n        # do on a memory-constrained T4.)\n        optimizer_steps_per_epoch = math.ceil(len(train_loader) / CFG.accumulation_steps)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=max_lrs, epochs=CFG.epochs, steps_per_epoch=optimizer_steps_per_epoch,\n            pct_start=CFG.onecycle_pct_start, anneal_strategy=\"cos\")\n        scheduler_steps_per_batch = True\n    else:\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n        scheduler_steps_per_batch = False\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n    ema = EMAModel(model, decay=CFG.ema_decay) if CFG.USE_EMA else None\n    global_step_holder = [0]\n    total_steps = CFG.epochs * max(1, len(train_loader) // CFG.accumulation_steps)\n\n    # ---- training loop ----\n    best_val_f05 = -1.0\n    epochs_no_improve = 0\n    best_state = None\n    history = defaultdict(list)\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n\n        if CFG.USE_CURRICULUM and sample_difficulties is not None:\n            # V7.1 (perf): mutate the existing sampler's weights in place --\n            # does NOT recreate the DataLoader, so persistent worker\n            # processes are kept warm across epochs instead of respawned.\n            curriculum_sampler.weights = build_curriculum_sampler(\n                sample_difficulties, epoch - 1, CFG.epochs, CFG.curriculum_warmup_epochs)\n\n        train_loss, train_metrics, train_consist = run_epoch(\n            model, criterion, train_loader, optimizer, scaler, ema,\n            None, None, None, global_step_holder, total_steps, train_mode=True, threshold=0.50,\n            scheduler=scheduler, scheduler_steps_per_batch=scheduler_steps_per_batch,\n            profile_steps=(CFG.PROFILE_TIMING_STEPS if (CFG.PROFILE_TIMING and epoch == 1) else 0))\n\n        if ema is not None:\n            backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n            model.load_state_dict(ema.state_dict())\n            val_loss, val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, ema,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n            model.load_state_dict(backup)\n            del backup\n        else:\n            val_loss, val_metrics, _ = run_epoch(\n                model, criterion, val_loader, optimizer, scaler, ema,\n                None, None, None, global_step_holder, total_steps, train_mode=False, threshold=0.50)\n\n        if not scheduler_steps_per_batch:\n            scheduler.step()\n\n        for k, v in val_metrics.items():\n            history[f\"val_{k}\"].append(v)\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"train_dice\"].append(train_metrics[\"dice\"])\n        history[\"train_consistency_loss\"].append(train_consist)\n\n        print(f\"[{epoch:02d}/{CFG.epochs}] {time.time()-t0:.1f}s | \"\n              f\"train_loss={train_loss:.4f} (consistency={train_consist:.4f}) | \"\n              f\"val_f0.5={val_metrics['f0.5']:.4f} val_dice={val_metrics['dice']:.4f} \"\n              f\"val_prec={val_metrics['precision']:.4f} val_rec={val_metrics['recall']:.4f}\")\n\n        if val_metrics[\"f0.5\"] > best_val_f05:\n            best_val_f05 = val_metrics[\"f0.5\"]\n            epochs_no_improve = 0\n            best_state = copy.deepcopy(ema.state_dict() if ema is not None else model.state_dict())\n            print(f\"  *** new best (val F0.5={best_val_f05:.4f}) ***\")\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= CFG.early_stop_patience:\n                print(\"  early stopping.\")\n                break\n        cleanup_memory()\n\n    model.load_state_dict(best_state)\n\n    # ---- threshold selection: VALIDATION ONLY, then frozen ----\n    best_threshold, val_score_at_best, thr_grid, thr_scores = find_best_threshold(\n        model, val_loader, metric=\"f0.5\")\n    print(f\"\\nSelected threshold (validation-only) = {best_threshold:.2f} \"\n          f\"(val F0.5={val_score_at_best:.4f})\")\n\n    # ---- held-out fragment: load labels now, evaluate at the FROZEN threshold ----\n    test_labels = load_ink_labels(test_dir)\n    fold_result = {\"fold_tag\": fold_tag, \"train_frags\": list(train_frags), \"test_frag\": test_frag,\n                   \"best_val_f05\": best_val_f05, \"selected_threshold\": best_threshold,\n                   \"history\": dict(history)}\n\n    if test_labels is not None:\n        gt_test = (test_labels * test_mask).astype(np.float32)\n        test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                              CFG.test_stride, CFG.infer_batch)\n        test_metrics_raw = evaluate_probability_map(test_prob, gt_test, best_threshold)\n        test_pred_bin = postprocess(test_prob, best_threshold)\n        post_preds = test_pred_bin.astype(np.float32)\n        tp = (post_preds * gt_test).sum(); fp = (post_preds * (1 - gt_test)).sum()\n        fn = ((1 - post_preds) * gt_test).sum()\n        test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n        print(f\"\\nHELD-OUT fragment {test_frag} metrics (frozen threshold={best_threshold:.2f}):\")\n        print(f\"  raw:          {test_metrics_raw}\")\n        print(f\"  postprocessed:{test_metrics_post}\")\n\n        fold_result[\"test_metrics_raw\"] = test_metrics_raw\n        fold_result[\"test_metrics_postprocessed\"] = test_metrics_post\n\n        if CFG.USE_MC_DROPOUT_UNCERTAINTY:\n            mean_prob, uncertainty = run_mc_dropout_uncertainty(\n                model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.mc_dropout_passes)\n            unc_report = analyze_uncertainty_vs_error(mean_prob, uncertainty, gt_test, test_mask,\n                                                       best_threshold, n_buckets=3)\n            print(\"  MC-dropout uncertainty vs. precision (low/med/high uncertainty buckets):\")\n            for r in unc_report:\n                print(f\"    bucket {r['bucket']}: n={r['n_pixels']} precision={r['precision']:.3f}\")\n            fold_result[\"mc_dropout_uncertainty_report\"] = unc_report\n\n        if CFG.PROTOCOL == \"transductive\" and CFG.USE_TEST_TIME_TRAINING:\n            ttt_model = test_time_training_adapt(model, test_vol, test_coords_unlabeled)\n            test_prob_ttt = sliding_window_inference(ttt_model, test_vol, test_mask, CFG.patch_size,\n                                                       CFG.test_stride, CFG.infer_batch)\n            test_metrics_ttt = evaluate_probability_map(test_prob_ttt, gt_test, best_threshold)\n            print(f\"  [TTT, Protocol B, NOT part of strict CV] test metrics: {test_metrics_ttt}\")\n            fold_result[\"test_metrics_ttt_protocol_b_only\"] = test_metrics_ttt\n            del ttt_model\n            cleanup_memory()\n\n        if save_viz:\n            _save_fold_overview(test_dir, test_prob, test_pred_bin, test_labels, fold_tag, best_threshold)\n    else:\n        print(\"No ground-truth labels for held-out fragment -- competition-style inference only.\")\n        test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                              CFG.test_stride, CFG.infer_batch)\n        fold_result[\"test_metrics_raw\"] = None\n\n    prob_path = os.path.join(CFG.out_dir, f\"fragment{test_frag}_probability_{fold_tag}.npy\")\n    np.save(prob_path, test_prob)\n    fold_result[\"probability_map_path\"] = prob_path\n\n    # ---- cleanup ----\n    test_vol.close()\n    for v in train_volumes.values():\n        v.close()\n    del model, optimizer, scheduler\n    cleanup_memory()\n\n    return fold_result\n\n\ndef _save_fold_overview(test_dir, prob_map, pred_bin, gt_labels, tag, threshold):\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 1600 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize((gt_labels * 255).astype(np.uint8), small.shape[::-1],\n                           interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    prob_small = cv2.resize((prob_map * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n    axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability (thr={threshold:.2f})\")\n    axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, f\"overview_{tag}.png\")\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"  saved overview: {path}\")\n\n\n# ============================================================\n# 19. LEAVE-ONE-FRAGMENT-OUT CV  (main scientific result)\n# ============================================================\n\ndef run_lofo_cv(fragments):\n    \"\"\"3 fragments -> 3 folds. Reports mean +/- std across folds instead of a\n    single 'best validation Dice' number, plus a paired t-test / Wilcoxon\n    signed-rank comparison IS available via compare_fold_results below if you\n    run two configurations (e.g. baseline vs. +DepthSignature) through this\n    same function and diff their fold-level F0.5 lists.\"\"\"\n    fold_results = []\n    for held_out in fragments:\n        train_frags = [f for f in fragments if f != held_out]\n        result = run_one_fold(train_frags, held_out, fold_tag=f\"lofo_test{held_out}\",\n                               save_viz=True)\n        fold_results.append(result)\n\n    summary = summarize_folds(fold_results)\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LEAVE-ONE-FRAGMENT-OUT CV SUMMARY\")\n    print(\"=\" * 70)\n    for metric, (mean, std, vals) in summary.items():\n        print(f\"  {metric:12s} = {mean:.4f} +/- {std:.4f}   (per-fold: {['%.4f' % v for v in vals]})\")\n\n    out_path = os.path.join(CFG.out_dir, \"lofo_cv_results.json\")\n    with open(out_path, \"w\") as f:\n        json.dump({\"folds\": fold_results, \"summary\": {k: (v[0], v[1]) for k, v in summary.items()},\n                    \"config\": cfg_to_dict(CFG)}, f, indent=2, default=str)\n    print(f\"\\nSaved: {out_path}\")\n    return fold_results, summary\n\n\ndef summarize_folds(fold_results, metric_keys=(\"f0.5\", \"dice\", \"iou\", \"precision\", \"recall\")):\n    summary = {}\n    for key in metric_keys:\n        vals = [fr[\"test_metrics_postprocessed\"][key] for fr in fold_results\n                 if fr.get(\"test_metrics_postprocessed\") is not None]\n        if not vals:\n            continue\n        summary[key] = (float(np.mean(vals)), float(np.std(vals)), vals)\n    return summary\n\n\ndef compare_fold_results(fold_results_a, fold_results_b, metric=\"f0.5\"):\n    \"\"\"Paired comparison across folds (same held-out fragments in the same\n    order for both configurations) -- Wilcoxon signed-rank test, with a\n    paired t-test reported alongside since n=3 folds is too small for the\n    Wilcoxon test's own asymptotics to be trustworthy; report both and let\n    the reader see they agree in direction.\"\"\"\n    a = [fr[\"test_metrics_postprocessed\"][metric] for fr in fold_results_a]\n    b = [fr[\"test_metrics_postprocessed\"][metric] for fr in fold_results_b]\n    t_stat, t_p = sstats.ttest_rel(a, b)\n    try:\n        w_stat, w_p = sstats.wilcoxon(a, b)\n    except Exception:\n        w_stat, w_p = float(\"nan\"), float(\"nan\")\n    print(f\"Paired comparison on {metric}: A={np.mean(a):.4f} B={np.mean(b):.4f} \"\n          f\"| paired t-test p={t_p:.4f} | Wilcoxon p={w_p:.4f}\")\n    return {\"metric\": metric, \"mean_a\": float(np.mean(a)), \"mean_b\": float(np.mean(b)),\n            \"t_stat\": float(t_stat), \"t_p\": float(t_p), \"w_stat\": float(w_stat), \"w_p\": float(w_p)}\n\n\n# ============================================================\n# 20. ABLATION MATRIX (component-by-component, single fold)\n# ============================================================\n\nABLATION_ARMS = [\n    (\"A_baseline_no_depth_stem\",       dict(depth_module_type=\"none\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"B_depth_fusion_stem\",            dict(depth_module_type=\"fusion\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"C_depth_signature_no_extras\",    dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=False,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"D_signature_plus_physics_ink\",   dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=False)),\n    (\"E_signature_plus_topology\",      dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=False,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=False,\n                                             USE_TOPOLOGY_LOSS=True)),\n    (\"F_signature_plus_consistency\",   dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=False, USE_CURRICULUM=False)),   # = full V6 model\n    (\"G_conv3d_stem_instead\",          dict(depth_module_type=\"conv3d\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=False, USE_CURRICULUM=False)),\n    (\"H_full_v7_plus_cutmix_curriculum\", dict(depth_module_type=\"signature\",\n                                             USE_PHYSICAL_SYNTHETIC_INK=True,\n                                             USE_FIBER_CONSISTENCY_LOSS=True,\n                                             USE_DEPTH_SHIFT_CONSISTENCY=True,\n                                             USE_TOPOLOGY_LOSS=True,\n                                             USE_CUTMIX=True, USE_CURRICULUM=True)),\n]\n\n\ndef run_ablation_matrix(train_frags, test_frag, arms=ABLATION_ARMS):\n    \"\"\"Runs each named arm on the SAME fold (same train/test fragment split)\n    so the differences are attributable to the listed components, not to a\n    different data split. Each arm is a fresh CfgOverride, so arms don't leak\n    settings into each other.\"\"\"\n    results = {}\n    for name, overrides in arms:\n        with CfgOverride(**overrides):\n            print(f\"\\n>>> ABLATION ARM: {name}  overrides={overrides}\")\n            fold_result = run_one_fold(train_frags, test_frag, fold_tag=f\"ablation_{name}\")\n            results[name] = fold_result\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"ABLATION MATRIX SUMMARY (test fragment = %s)\" % test_frag)\n    print(\"=\" * 70)\n    print(f\"{'arm':32s} {'F0.5':>8s} {'Dice':>8s} {'IoU':>8s} {'Prec':>8s} {'Rec':>8s}\")\n    for name, fr in results.items():\n        m = fr.get(\"test_metrics_postprocessed\")\n        if m is None:\n            print(f\"{name:32s}  (no GT available)\")\n            continue\n        print(f\"{name:32s} {m['f0.5']:8.4f} {m['dice']:8.4f} {m['iou']:8.4f} \"\n              f\"{m['precision']:8.4f} {m['recall']:8.4f}\")\n\n    out_path = os.path.join(CFG.out_dir, \"ablation_matrix_results.json\")\n    with open(out_path, \"w\") as f:\n        json.dump(results, f, indent=2, default=str)\n    print(f\"\\nSaved: {out_path}\")\n    return results\n\n\n# ============================================================\n# 21. MAIN\n# ============================================================\n\nif __name__ == \"__main__\":\n    set_seed(CFG.seed)\n\n    if CFG.RUN_MODE == \"lofo_cv\":\n        fold_results, summary = run_lofo_cv(CFG.all_fragments)\n\n    elif CFG.RUN_MODE == \"ablation\":\n        ablation_results = run_ablation_matrix(CFG.ablation_train_frags, CFG.ablation_test_frag)\n\n    elif CFG.RUN_MODE == \"single\":\n        result = run_one_fold(CFG.single_train_frags, CFG.single_test_frag,\n                               fold_tag=\"single_run\", save_viz=True)\n        print(\"\\nSingle-run result:\", json.dumps(\n            {k: v for k, v in result.items() if k != \"history\"}, indent=2, default=str))\n\n    else:\n        raise ValueError(f\"Unknown CFG.RUN_MODE={CFG.RUN_MODE}\")\n\n    print(\"\\n=== V6 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-17T07:42:32.394504Z","iopub.execute_input":"2026-09-17T07:42:32.395387Z"}},"outputs":[{"name":"stdout","text":"[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.29\n\n######################################################################\n# FOLD [lofo_test1]  train=['2', '3']  test=1\n######################################################################\n  fragment 2: train=4559 val=1260 mean=110.0 std=57.6\n  fragment 3: train=1272 val=160 mean=100.3 std=62.4\nPositive patches: 3953 | Negative patches: 1878\nBalanced dataset: 5831 | positive ratio=0.678\n  hist-match pool (strict protocol): 20 patches from TRAINING fragments ['2', '3'] only -- held-out fragment 1 untouched.\n  curriculum sampling ON: warmup_epochs=8 easy/medium/hard counts = 2736/1287/1808\n[backbone] DepthSignatureModule: 26 depth slices -> 24 learned channels (LDDC=4 filters, PE dim=8, attention heads=4, pooled to 8x8, stats=True)\n","output_type":"stream"},{"name":"stderr","text":"Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"508b6af361ef488486350aa5fb95f0a0"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV6 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320\ndepth_module_type=signature | encoder=tu-convnext_tiny | architecture=unet\nParameters: 32.18M total | 32.18M trainable\nZero-channel layers: 0 (PASS)\nDry-run (256x256) output logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\nCUDA device: Tesla T4\nForward pass: PASS\n======================================================================\n\n    [dataset profile pid=167] TOTAL=945.9ms :: main_patch_read=904.1ms | hist_match=0.0ms | physical_ink=0.0ms | fake_ink_fiber_shadow=13.1ms | albumentations_transform=7.4ms | prep_to_tensor=21.4ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=166] TOTAL=3999.3ms :: main_patch_read=3952.4ms | hist_match=0.0ms | physical_ink=0.0ms | fake_ink_fiber_shadow=23.1ms | albumentations_transform=7.8ms | prep_to_tensor=15.9ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=167] TOTAL=4956.9ms :: main_patch_read=4879.9ms | hist_match=35.1ms | physical_ink=0.0ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=26.1ms | prep_to_tensor=15.8ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=166] TOTAL=2582.2ms :: main_patch_read=2539.0ms | hist_match=0.0ms | physical_ink=0.0ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=23.2ms | prep_to_tensor=20.0ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=166] TOTAL=74.9ms :: main_patch_read=18.5ms | hist_match=34.1ms | physical_ink=0.0ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=6.7ms | prep_to_tensor=15.6ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=167] TOTAL=2836.6ms :: main_patch_read=2806.3ms | hist_match=0.0ms | physical_ink=0.0ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=10.2ms | prep_to_tensor=20.0ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=167] TOTAL=535.2ms :: main_patch_read=426.1ms | hist_match=0.0ms | physical_ink=0.0ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=95.7ms | prep_to_tensor=13.4ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=166] TOTAL=3494.9ms :: main_patch_read=3369.7ms | hist_match=0.0ms | physical_ink=19.7ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=92.2ms | prep_to_tensor=13.4ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=167] TOTAL=2700.1ms :: main_patch_read=2587.5ms | hist_match=31.3ms | physical_ink=15.2ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=51.5ms | prep_to_tensor=14.5ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=166] TOTAL=1941.8ms :: main_patch_read=1836.8ms | hist_match=32.4ms | physical_ink=6.4ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=51.7ms | prep_to_tensor=14.5ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=166] TOTAL=135.6ms :: main_patch_read=9.6ms | hist_match=31.1ms | physical_ink=0.0ms | fake_ink_fiber_shadow=27.9ms | albumentations_transform=53.6ms | prep_to_tensor=13.4ms | depth_shift_patch_read=0.0ms\n    [dataset profile pid=167] TOTAL=3274.4ms :: main_patch_read=3188.5ms | hist_match=0.0ms | physical_ink=0.0ms | fake_ink_fiber_shadow=0.0ms | albumentations_transform=71.4ms | prep_to_tensor=14.6ms | depth_shift_patch_read=0.0ms\n  [profile step 0] data_wait=15.053s main_fwd=9.303s consistency_fwd=0.506s (consistency_computed=True) backward+step=49.052s TOTAL=58.860s\n  [profile step 1] data_wait=0.000s main_fwd=0.269s consistency_fwd=0.000s (consistency_computed=False) backward+step=1.173s TOTAL=1.442s\n  [profile step 2] data_wait=0.000s main_fwd=0.316s consistency_fwd=0.000s (consistency_computed=False) backward+step=1.174s TOTAL=1.490s\n  [profile step 3] data_wait=0.000s main_fwd=0.272s consistency_fwd=0.000s (consistency_computed=False) backward+step=1.173s TOTAL=1.445s\n  [profile step 4] data_wait=0.000s main_fwd=0.271s consistency_fwd=0.502s (consistency_computed=True) backward+step=3.502s TOTAL=4.275s\n  [profile step 5] data_wait=1.511s main_fwd=0.271s consistency_fwd=0.000s (consistency_computed=False) backward+step=1.175s TOTAL=1.446s\n  [profile step 6] data_wait=0.000s main_fwd=0.270s consistency_fwd=0.000s (consistency_computed=False) backward+step=1.175s TOTAL=1.445s\n  [profile step 7] data_wait=4.182s main_fwd=0.309s consistency_fwd=0.000s (consistency_computed=False) backward+step=1.175s TOTAL=1.484s\n    step 36/728 | 152s elapsed | 0.24 steps/s | ETA 2914s\n    step 72/728 | 232s elapsed | 0.31 steps/s | ETA 2112s\n    step 108/728 | 311s elapsed | 0.35 steps/s | ETA 1784s\n    step 144/728 | 388s elapsed | 0.37 steps/s | ETA 1574s\n    step 180/728 | 466s elapsed | 0.39 steps/s | ETA 1417s\n    step 216/728 | 540s elapsed | 0.40 steps/s | ETA 1280s\n    step 252/728 | 616s elapsed | 0.41 steps/s | ETA 1164s\n    step 288/728 | 694s elapsed | 0.42 steps/s | ETA 1060s\n    step 324/728 | 771s elapsed | 0.42 steps/s | ETA 961s\n    step 360/728 | 847s elapsed | 0.43 steps/s | ETA 866s\n    step 396/728 | 924s elapsed | 0.43 steps/s | ETA 775s\n    step 432/728 | 1002s elapsed | 0.43 steps/s | ETA 686s\n    step 468/728 | 1080s elapsed | 0.43 steps/s | ETA 600s\n    step 504/728 | 1159s elapsed | 0.43 steps/s | ETA 515s\n    step 540/728 | 1239s elapsed | 0.44 steps/s | ETA 431s\n    step 576/728 | 1318s elapsed | 0.44 steps/s | ETA 348s\n    step 612/728 | 1397s elapsed | 0.44 steps/s | ETA 265s\n    step 684/728 | 1556s elapsed | 0.44 steps/s | ETA 100s\n    step 720/728 | 1632s elapsed | 0.44 steps/s | ETA 18s\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# SAVE OVERLAY and & HIstogram matched prediction\n#'fbeta0.5': 0.4907547533512 with img=480, stride=128\n# ============================================================\n# VESUVIUS INK DETECTION - V4 \"DEPTH-AWARE CONVNEXT\"\n# ============================================================\n# Built on V3 (Dice=0.499). Fixes the crash and addresses the deeper\n# architectural critique: the model was treating 26 ordered depth\n# slices as unordered multispectral channels.\n#\n# ------------------------------------------------------------------\n# WHAT ACTUALLY CAUSED THE CRASH (verified empirically, not guessed)\n# ------------------------------------------------------------------\n# The proposed fix of adding `decoder_channels=(256,128,64,32,16)` to\n# smp.UnetPlusPlus was tested directly against this exact setup and\n# it does NOT fix the crash -- I reproduced the identical\n# `weight of size [0, 96, 3, 3]` error with it. The real cause:\n# ConvNeXt's stem is a stride-4 patchify with no separate stride-2\n# stage (unlike ResNet), so smp's generic \"tu-\" timm-encoder wrapper\n# fabricates an EMPTY 0-channel placeholder feature map to keep a\n# uniform 5-stage pyramid API. `smp.Unet`'s decoder already handles\n# that 0-channel stage gracefully (confirmed working); `smp.UnetPlusPlus`'s\n# dense skip-connections do not (confirmed failing, with or without\n# explicit decoder_channels). So V4 uses `smp.Unet`, not UnetPlusPlus,\n# with tu-convnext_tiny. This is a real library limitation, not a\n# parameter you can configure around.\n#\n# ------------------------------------------------------------------\n# DEPTH-AWARE CHANGES (addressing \"26 slices as unordered channels\")\n# ------------------------------------------------------------------\n#  1. DepthFusionStem: computes explicit first/second finite\n#     differences along the physically-ordered depth axis (how\n#     intensity changes through the papyrus), concatenates them with\n#     the raw stack, and learns a compact (16-32 channel) mixed\n#     representation via a 1x1 conv BEFORE the 2D ConvNeXt encoder\n#     ever sees the data -- giving the network an explicit inductive\n#     bias toward depth structure instead of hoping a from-scratch\n#     first-conv discovers it.\n#  2. CLAHE_MODE: \"per_slice\" (V3's old behavior, independent CLAHE\n#     per slice -- can distort inter-slice relationships) vs.\n#     \"global_shared\" (ONE contrast-remapping LUT derived from a\n#     representative slice, applied identically to every slice --\n#     preserves relative depth relationships). Both available for\n#     the ablation you suggested; default is now \"global_shared\".\n#  3. Histogram-matching augmentation now uses ONE shared mapping\n#     across the whole depth stack (`channel_axis=None`) instead of\n#     26 independent per-slice mappings (`channel_axis=2`). I verified\n#     this empirically: independent per-channel matching compressed\n#     slice-to-slice differences unevenly (2.3-3.8 range in a test),\n#     while the shared mapping preserved them far more consistently\n#     (4.3-7.3 range, proportional to the original 7.5-8.6 spacing).\n#  4. Architecture inspection report + zero-channel detector run\n#     BEFORE the optimizer/training loop are created -- this is\n#     exactly the check that would have caught the original crash\n#     immediately instead of after a dummy forward pass deep in setup.\n#  5. AdaBN now prints exactly where its BatchNorm2d layers are found\n#     (depth stem / decoder / encoder) since ConvNeXt itself is\n#     LayerNorm-only and has none -- so you know what it's actually\n#     recalibrating before deciding whether to enable it.\n#  6. V4 baseline: USE_DANN / USE_MIXSTYLE / use_adabn default to\n#     False so the depth-aware architecture change can be evaluated\n#     cleanly on its own first. Turn them back on one at a time\n#     afterward -- CFG has a comment marking each one.\n#  7. One explicit environment setup (current smp + timm, no more\n#     pinning smp==0.2.0 and upgrading from inside build_model()).\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS  (single explicit environment, no version-pin dance)\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n!pip install segmentation-models-pytorch==0.4.0\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom scipy import ndimage\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    patch_size = 320\n    train_stride = 60\n    test_stride = 60\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 8\n    early_stop_patience = 3\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- V4: backbone / architecture -----------------------------------------\n    # VERIFIED WORKING with tu-convnext_tiny: smp.Unet (NOT UnetPlusPlus -- see\n    # module docstring for the empirical reason). Falls back to efficientnet-b4\n    # (a \"real\" smp encoder with no 0-channel-stage quirk) if ConvNeXt/timm are\n    # unavailable in this session for any reason.\n    encoder_name = \"resnet50\"\n    #encoder_name = \"tu-convnext_tiny\"\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"                          # do NOT set to \"unetplusplus\" with tu-* encoders\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- V4: depth-aware fusion stem -----------------------------------------\n    USE_DEPTH_FUSION_STEM = True\n    depth_stem_out_channels = 16\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- LOSS -----------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.35\n    focal_tversky_weight = 0.35\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n\n    # --- THRESHOLD --------------------------------------------------------------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n    threshold = 0.50\n\n    # --- TTA --------------------------------------------------------------------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN (see the \"where are the BN layers\" report at model-build time) --\n    # ConvNeXt is LayerNorm-only; V4 defaults this OFF until the printed report\n    # shows there's something meaningful (decoder / depth-stem BN) for it to do.\n    use_adabn = False\n    adabn_max_patches = 2000\n\n    # --- POSTPROCESS ----------------------------------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    # --- edge-artifact cropping -------------------------------------------------\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- V4: CLAHE mode ---------------------------------------------------------\n    # \"off\"           : no contrast enhancement\n    # \"per_slice\"     : V3's old behavior -- independent CLAHE per slice, can\n    #                   distort inter-slice depth relationships\n    # \"global_shared\" : ONE remapping LUT derived from a representative slice,\n    #                   applied identically to every slice -- preserves relative\n    #                   depth relationships. DEFAULT for V4.\n    CLAHE_MODE = \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization ------------------------------------------------------------\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n\n    # --- domain-randomization augmentation --------------------------------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n\n    USE_SHADOW = True\n    shadow_p = 0.20\n\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # histogram-matching style augmentation -- NOW uses one shared mapping across\n    # depth (channel_axis=None) instead of 26 independent per-slice mappings.\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- EMA --------------------------------------------------------------------\n    USE_EMA = True\n    ema_decay = 0.999\n\n    # --- DANN (OFF for the V4 baseline -- turn on only after the depth-aware\n    #     architecture change is validated on its own; see module docstring) ----\n    USE_DANN = False\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n\n    # --- MixStyle (OFF for the V4 baseline, same reasoning as DANN) -------------\n    #USE_MIXSTYLE = False\n    USE_MIXSTYLE = True\n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    ckpt_path = os.path.join(out_dir, \"vesuviusnet_v4_best.pth\")\n    viz_dir = os.path.join(out_dir, \"v4_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v4_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\nset_seed(CFG.seed)\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS  (+ mask erosion)\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING (unchanged)\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME  (+ V4: CLAHE modes, per-fragment stats)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None   # built lazily for CLAHE_MODE == \"global_shared\"\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        \"\"\"Derives ONE 256-entry intensity-remapping lookup table from CLAHE applied\n        to a single representative slice, then this exact LUT is applied identically\n        to every depth slice via cv2.LUT. Unlike calling .apply() independently per\n        slice, this guarantees the same monotonic mapping everywhere, so relative\n        inter-slice intensity relationships (the actual depth signal) are preserved.\"\"\"\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid = len(self._slices) // 2\n        ref_slice = np.asarray(self._slices[mid])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(\n            ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def read_patch(self, y, x, size, apply_clahe=None):\n        self._ensure_open()\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[i] = block\n        return out\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        mid = len(self._slices) // 2\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            block = self._slices[mid][y:y + patch_size, x:x + patch_size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                block = self._clahe.apply(block)\n            elif CFG.CLAHE_MODE == \"global_shared\":\n                block = cv2.LUT(block, self._shared_clahe_lut)\n            vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION (fake ink / fiber / shadow)\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\n# ============================================================\n# 8. AUGMENTATION\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    \"\"\"Try constructing a transform with the current albumentations API; if that\n    fails (parameter names changed across versions), fall back to the older API.\n    If both fail, skip the transform rather than crashing the whole pipeline.\"\"\"\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    # NOTE: albumentations' RandomShadow hard-requires 3-channel RGB images in\n    # every version and raises ValueError on multi-channel depth-stack data.\n    # inject_fake_shadow() (applied directly in the dataset, image-only) is used\n    # instead -- see USE_SHADOW below.\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size]\n        img = np.transpose(patch, (1, 2, 0))   # HWD\n\n        # --- histogram-matching style augmentation: ONE shared mapping across\n        # depth (channel_axis=None), not 26 independent per-slice mappings. ---\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img, vol.frag_mean, vol.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN: yields normalized (D,H,W) tensors from fragment 1, NO labels.\n    Uses the SAME basic preprocessing path (CLAHE mode, per-fragment normalization)\n    as the source dataset for consistency -- it intentionally skips the AUGMENTATION\n    pipeline (elastic/dropout/fake-ink/histogram-match), since DANN needs to see the\n    target domain's natural distribution, not an augmented/hallucinated version of it.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL: DEPTH FUSION STEM + VERIFIED-WORKING ConvNeXt+Unet\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"Computes explicit depth-derivative features (first and second finite\n    differences along the physically-ordered depth axis) alongside the raw stack,\n    then learns a compact mixed representation via a 1x1 conv, producing\n    `out_channels` channels to feed into the 2D encoder. This is the \"ink isn't\n    just absolute intensity, it's how intensity changes through the papyrus\"\n    inductive bias, made explicit instead of hoping a from-scratch first-conv\n    layer discovers it purely from data.\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = 0\n        if use_raw:\n            total_in += in_depth\n        if use_grad:\n            total_in += (in_depth - 1)\n        if use_curv:\n            total_in += max(in_depth - 2, 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        # BatchNorm2d kept HERE deliberately (even though the ConvNeXt backbone\n        # itself is all LayerNorm) so AdaBN has a real, meaningful place to\n        # recalibrate target-domain statistics -- see the model-build report.\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.in_depth = in_depth\n\n    def forward(self, x):   # x: (B, D, H, W)\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:, :, :] - x[:, :-1, :, :])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass DepthAwareSegModel(nn.Module):\n    \"\"\"Wraps an smp segmentation model with a DepthFusionStem in front of it.\"\"\"\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):   # x: (B, D, H, W)\n        feat = self.depth_stem(x) if self.depth_stem is not None else x.unsqueeze(1) if x.dim() == 3 else x\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' requested, but only \"\n              f\"'unet' has been verified to handle tu-* timm encoders' 0-channel placeholder \"\n              f\"stage correctly (see module docstring) -- using UnetPlusPlus may crash with \"\n              f\"ConvNeXt/other tu-* encoders.\")\n    return arch_cls(\n        encoder_name=encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=in_channels,\n        classes=1,\n        decoder_channels=CFG.decoder_channels,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\ndef build_model():\n    stem = None\n    seg_in_channels = CFG.in_channels\n    if CFG.USE_DEPTH_FUSION_STEM:\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = CFG.depth_stem_out_channels\n        print(f\"[backbone] DepthFusionStem: {CFG.in_channels} depth slices \"\n              f\"(raw={CFG.depth_stem_use_raw}, grad={CFG.depth_stem_use_grad}, curv={CFG.depth_stem_use_curv}) \"\n              f\"-> {seg_in_channels} learned channels\")\n\n    try:\n        seg_model = _build_seg_backbone(CFG.encoder_name, seg_in_channels)\n        print(f\"[backbone] using {CFG.encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {CFG.encoder_name} unavailable/failed ({e}); \"\n              f\"falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model)\n\n\nclass GradReverse(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x, lambd):\n        ctx.lambd = lambd\n        return x.view_as(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return -ctx.lambd * grad_output, None\n\n\ndef grad_reverse(x, lambd=1.0):\n    return GradReverse.apply(x, lambd)\n\n\nclass DomainClassifier(nn.Module):\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(in_ch, 128), nn.ReLU(inplace=True), nn.Dropout(0.3),\n            nn.Linear(128, 1),\n        )\n\n    def forward(self, feat, lambd):\n        return self.net(grad_reverse(feat, lambd))\n\n\nclass EncoderFeatureCapture:\n    \"\"\"Architecture-agnostic hook capturing the segmentation backbone's encoder\n    output feature list on every forward pass. Attaches to model.seg_model.encoder\n    (the DepthAwareSegModel wrapper's inner smp model), not model.encoder directly,\n    since V4 wraps the smp model with a depth-fusion stem in front of it.\"\"\"\n    def __init__(self, model):\n        self.features = None\n        target = model.seg_model.encoder if hasattr(model, \"seg_model\") else model.encoder\n        self.handle = target.register_forward_hook(self._hook)\n\n    def _hook(self, module, inp, out):\n        self.features = out\n\n    def remove(self):\n        self.handle.remove()\n\n\ndef mixstyle_batch(imgs, p=CFG.mixstyle_p, alpha=CFG.mixstyle_alpha):\n    if torch.rand(1).item() > p:\n        return imgs\n    B = imgs.size(0)\n    if B < 2:\n        return imgs\n    mu = imgs.mean(dim=[2, 3], keepdim=True)\n    var = imgs.var(dim=[2, 3], keepdim=True)\n    sig = (var + 1e-6).sqrt()\n    x_norm = (imgs - mu) / sig\n    perm = torch.randperm(B, device=imgs.device)\n    mu2, sig2 = mu[perm], sig[perm]\n    lam = torch.distributions.Beta(alpha, alpha).sample((B, 1, 1, 1)).to(imgs.device)\n    mu_mix = lam * mu + (1 - lam) * mu2\n    sig_mix = lam * sig + (1 - lam) * sig2\n    return x_norm * sig_mix + mu_mix\n\n\nclass EMAModel:\n    def __init__(self, model, decay=CFG.ema_decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. V4: PRE-TRAINING ARCHITECTURE INSPECTION\n#     (this is exactly what would have caught the original crash\n#     immediately, before the optimizer/training loop existed)\n# ============================================================\n\ndef inspect_model_channels(model):\n    \"\"\"Scans every Conv/Linear layer for zero or negative in/out channels -- the\n    exact failure mode behind the original UnetPlusPlus/ConvNeXt crash.\"\"\"\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\ndef report_batchnorm_locations(model):\n    \"\"\"Prints where BatchNorm2d layers actually live in the model -- ConvNeXt's\n    own encoder is LayerNorm-only, so AdaBN (which recalibrates BatchNorm running\n    stats) has nothing to do there. It CAN still do something meaningful in the\n    depth-fusion stem and/or the smp decoder, if those have BatchNorm2d layers.\"\"\"\n    counts = {\"depth_stem\": 0, \"encoder\": 0, \"decoder\": 0, \"other\": 0}\n    for name, module in model.named_modules():\n        if isinstance(module, nn.BatchNorm2d):\n            if name.startswith(\"depth_stem\"):\n                counts[\"depth_stem\"] += 1\n            elif \".encoder.\" in f\".{name}.\" or name.endswith(\".encoder\"):\n                counts[\"encoder\"] += 1\n            elif \".decoder.\" in f\".{name}.\" or name.endswith(\".decoder\"):\n                counts[\"decoder\"] += 1\n            else:\n                counts[\"other\"] += 1\n    total = sum(counts.values())\n    print(f\"[AdaBN report] BatchNorm2d layers found: depth_stem={counts['depth_stem']} \"\n          f\"encoder={counts['encoder']} decoder={counts['decoder']} other={counts['other']} \"\n          f\"(total={total})\")\n    if counts[\"encoder\"] == 0 and total > 0:\n        print(\"  -> the ConvNeXt encoder itself has none (it's LayerNorm-only, as expected). \"\n              \"AdaBN would only recalibrate the depth-stem/decoder BN layers listed above.\")\n    if total == 0:\n        print(\"  -> NO BatchNorm2d layers anywhere in this model. Enabling use_adabn would \"\n              \"currently be a no-op. (Kept OFF by default in V4 CFG for exactly this reason.)\")\n    return counts\n\n\n@torch.no_grad()\ndef run_architecture_report(model):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"V4 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\")\n    print(\"=\" * 70)\n\n    n_params = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Input:  depth slices={CFG.in_channels} | patch={CFG.patch_size}x{CFG.patch_size} | \"\n          f\"input tensor=[1, {CFG.in_channels}, {CFG.patch_size}, {CFG.patch_size}]\")\n    print(f\"Encoder: backbone={CFG.encoder_name} | pretrained={CFG.encoder_weights} | \"\n          f\"architecture={CFG.architecture}\")\n    if CFG.USE_DEPTH_FUSION_STEM:\n        print(f\"Depth stem: {CFG.in_channels} -> {model.depth_stem.out_channels} channels \"\n              f\"(raw={model.depth_stem.use_raw}, grad={model.depth_stem.use_grad}, \"\n              f\"curv={model.depth_stem.use_curv})\")\n    print(f\"Decoder channels: {CFG.decoder_channels}\")\n    print(f\"Parameters: {n_params/1e6:.2f}M total | {n_trainable/1e6:.2f}M trainable\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"\\nZERO/INVALID CHANNEL LAYERS FOUND:\")\n        for item in bad_layers:\n            print(\"  \", item)\n        raise RuntimeError(\n            \"Model contains invalid zero-channel layers -- fix architecture before training. \"\n            \"(If this happened with architecture='unetplusplus' and a tu-* encoder, switch to \"\n            \"architecture='unet' -- see module docstring for why.)\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    model.eval()\n    small_size = min(CFG.patch_size, 256)   # keep the dry-run cheap regardless of real patch_size\n    dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n    capture = EncoderFeatureCapture(model)\n    try:\n        out = model(dummy)\n        print(f\"\\nEncoder feature stages (dry run at {small_size}x{small_size} for speed):\")\n        for i, f in enumerate(capture.features):\n            print(f\"  stage {i}: {tuple(f.shape)}\")\n        print(f\"\\nOutput logits shape: {tuple(out.shape)}\")\n        has_nan = torch.isnan(out).any().item()\n        has_inf = torch.isinf(out).any().item()\n        print(f\"NaN check: {'FAIL' if has_nan else 'PASS'} | Inf check: {'FAIL' if has_inf else 'PASS'}\")\n        if has_nan or has_inf:\n            raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass -- fix before training.\")\n    finally:\n        capture.remove()\n    model.train()\n\n    if torch.cuda.is_available():\n        print(f\"\\nCUDA device: {torch.cuda.get_device_name(0)} | \"\n              f\"capability: {torch.cuda.get_device_capability(0)}\")\n    print(\"Forward pass: PASS\")\n\n    report_batchnorm_locations(model)\n    print(\"=\" * 70 + \"\\n\")\n\n\n# ============================================================\n# 12. LOSSES (unchanged)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\nclass V2ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        return CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n\n\n# ============================================================\n# 13. METRICS\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 14. BUILD TRAIN DATA\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING TRAIN DATA\")\nprint(\"=\" * 70)\n\ntrain_volumes = {}\ntrain_labels_full = {}\ntrain_masks_full = {}\ntrain_samples_raw = []\nval_samples = []\n\nfor fid in CFG.train_frags:\n    print(f\"\\nProcessing fragment {fid} ...\")\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        vol.compute_fragment_stats(mask, CFG.patch_size)\n        print(f\"  fragment {fid} normalization stats: mean={vol.frag_mean:.2f} std={vol.frag_std:.2f}\")\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {len(coords)} candidate patches (post mask-erosion)\")\n\n    tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n    print(f\"  spatial train={len(tr_coords)} validation={len(va_coords)}\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    train_masks_full[fid] = mask\n    train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n    val_samples.extend([(fid, y, x) for y, x in va_coords])\n    del mask\n    cleanup_memory()\n\nprint(\"\\nRaw train samples:\", len(train_samples_raw))\nprint(\"Validation samples:\", len(val_samples))\n\ntrain_samples = balance_positive_patches(\n    train_samples_raw, train_labels_full, CFG.patch_size,\n    positive_threshold=CFG.positive_patch_fraction,\n    target_positive_ratio=CFG.target_positive_patch_ratio,\n    max_positive_repeat=CFG.max_positive_repeat,\n)\n\n\n# ============================================================\n# 15. TEST-FRAGMENT MASK/VOLUME LOADED EARLY (unlabeled use only)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\")\nprint(\"=\" * 70)\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\nif CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    print(f\"  fragment {CFG.test_frag} normalization stats: \"\n          f\"mean={test_vol.frag_mean:.2f} std={test_vol.frag_std:.2f}\")\n\ntest_coords_for_unlabeled_use = generate_grid_coords(\n    test_mask, CFG.patch_size, CFG.test_stride, CFG.min_tissue_frac_train)\nprint(f\"Fragment {CFG.test_frag}: {len(test_coords_for_unlabeled_use)} unlabeled candidate patches\")\n\nhist_match_pool = None\nif CFG.USE_HIST_MATCH_AUG:\n    pool_coords = random.sample(test_coords_for_unlabeled_use,\n                                 min(CFG.hist_match_pool_size, len(test_coords_for_unlabeled_use)))\n    hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n    print(f\"Built histogram-matching reference pool: {len(hist_match_pool)} patches \"\n          f\"({sum(p.nbytes for p in hist_match_pool)/1e6:.1f} MB)\")\n\ndann_loader = None\nif CFG.USE_DANN:\n    dann_ds = UnlabeledPatchDataset(test_vol, test_coords_for_unlabeled_use, CFG.patch_size)\n    dann_loader = DataLoader(dann_ds, batch_size=CFG.dann_target_batch_size, shuffle=True,\n                              num_workers=max(1, CFG.num_workers - 1), pin_memory=(CFG.device == \"cuda\"),\n                              drop_last=True, persistent_workers=True)\n\n    def infinite_dann_loader():\n        while True:\n            for batch in dann_loader:\n                yield batch\n    dann_iter = infinite_dann_loader()\n\n\n# ============================================================\n# 16. DATASETS / LOADERS\n# ============================================================\n\ntrain_transform = build_train_transform()\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                            transform=train_transform, jitter=CFG.train_jitter,\n                            hist_match_pool=hist_match_pool, train_mode=True)\n\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                          transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0)\n\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         drop_last=False, persistent_workers=CFG.num_workers > 0)\n\n\n# ============================================================\n# 17. MODEL / ARCHITECTURE REPORT / LOSS / OPTIMIZER\n# ============================================================\n\nprint(\"\\nBuilding V4 depth-aware model ...\")\nmodel = build_model().to(CFG.device)\n\nrun_architecture_report(model)   # <-- catches zero-channel / NaN / shape problems HERE\n\nprint(\"\\nEstimating positive-pixel fraction ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\nprint(f\"Estimated positive fraction: {pos_frac:.6f}\")\n\nwith torch.no_grad():\n    bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n    try:\n        model.seg_model.segmentation_head[0].bias.fill_(bias_val)\n    except Exception as e:\n        print(f\"  (could not set output bias directly: {e})\")\nprint(f\"Output bias initialized to {bias_val:.4f}\")\n\nraw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\npos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\npos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\nprint(f\"BCE positive weight: {pos_weight_val:.3f}\")\n\ncriterion = V2ComboLoss(pos_weight=pos_weight)\n\nencoder_params, decoder_params, stem_params = [], [], []\nfor name, param in model.named_parameters():\n    if not param.requires_grad:\n        continue\n    if name.startswith(\"depth_stem.\"):\n        stem_params.append(param)\n    elif name.startswith(\"seg_model.encoder.\"):\n        encoder_params.append(param)\n    else:\n        decoder_params.append(param)\nprint(f\"Depth-stem parameters: {len(stem_params)} | Encoder parameters: {len(encoder_params)} | \"\n      f\"Decoder/head parameters: {len(decoder_params)} tensors\")\n\nparam_groups = [\n    {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n    {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n]\nif stem_params:\n    param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\ndomain_classifier = None\nfeature_capture = None\nif CFG.USE_DANN:\n    feature_capture = EncoderFeatureCapture(model)\n    with torch.no_grad():\n        dummy = torch.zeros(1, CFG.in_channels, CFG.patch_size, CFG.patch_size, device=CFG.device)\n        model.eval()\n        _ = model(dummy)\n        deepest_ch = feature_capture.features[-1].shape[1]\n        model.train()\n    domain_classifier = DomainClassifier(deepest_ch).to(CFG.device)\n    param_groups.append({\"params\": domain_classifier.parameters(), \"lr\": CFG.dann_lr})\n    print(f\"[DANN] domain classifier attached on {deepest_ch}-channel bottleneck features\")\n    del dummy\n    cleanup_memory()\n\noptimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nema = EMAModel(model, decay=CFG.ema_decay) if CFG.USE_EMA else None\n\n\n# ============================================================\n# 18. TRAIN / VALIDATION EPOCH\n# ============================================================\n\n_global_step = 0\n_total_steps = CFG.epochs * max(1, len(train_loader) // CFG.accumulation_steps)\n\n\ndef dann_lambda_schedule():\n    progress = min(_global_step / max(_total_steps, 1), 1.0)\n    return CFG.dann_lambda_max * (2.0 / (1.0 + math.exp(-10.0 * progress)) - 1.0)\n\n\ndef run_epoch(loader, train_mode=True, threshold=0.5):\n    global _global_step\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_dann_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    for batch_idx, (imgs, masks) in enumerate(loader):\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        if train_mode and CFG.USE_MIXSTYLE:\n            imgs = mixstyle_batch(imgs)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n                dann_loss_val = 0.0\n\n                if train_mode and CFG.USE_DANN:\n                    src_feat = feature_capture.features[-1]\n                    lambd = dann_lambda_schedule()\n                    src_domain_logits = domain_classifier(src_feat, lambd)\n                    src_domain_target = torch.zeros_like(src_domain_logits)\n\n                    tgt_imgs = next(dann_iter).to(CFG.device, non_blocking=True)\n                    _ = model(tgt_imgs)\n                    tgt_feat = feature_capture.features[-1]\n                    tgt_domain_logits = domain_classifier(tgt_feat, lambd)\n                    tgt_domain_target = torch.ones_like(tgt_domain_logits)\n\n                    dann_loss = 0.5 * (\n                        F.binary_cross_entropy_with_logits(src_domain_logits, src_domain_target) +\n                        F.binary_cross_entropy_with_logits(tgt_domain_logits, tgt_domain_target)\n                    )\n                    dann_loss_val = dann_loss.item()\n                    loss = loss + dann_loss\n\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    params_to_clip = list(model.parameters())\n                    if domain_classifier is not None:\n                        params_to_clip += list(domain_classifier.parameters())\n                    torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    _global_step += 1\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_dann_loss += dann_loss_val\n\n        del imgs, masks, logits, probs\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        params_to_clip = list(model.parameters())\n        if domain_classifier is not None:\n            params_to_clip += list(domain_classifier.parameters())\n        torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_dann_loss = total_dann_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_dann_loss\n\n\n# ============================================================\n# 19. TRAINING LOOP\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING V4 TRAINING\")\nprint(\"=\" * 70)\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"train_dice\": [], \"val_dice\": [],\n           \"val_iou\": [], \"val_precision\": [], \"val_recall\": [], \"dann_loss\": []}\n\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_metrics, dann_loss_avg = run_epoch(train_loader, train_mode=True, threshold=0.50)\n\n    if ema is not None:\n        backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        model.load_state_dict(ema.state_dict())\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n        model.load_state_dict(backup)\n        del backup\n    else:\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"train_dice\"].append(train_metrics[\"dice\"])\n    history[\"val_dice\"].append(val_metrics[\"dice\"])\n    history[\"val_iou\"].append(val_metrics[\"iou\"])\n    history[\"val_precision\"].append(val_metrics[\"precision\"])\n    history[\"val_recall\"].append(val_metrics[\"recall\"])\n    history[\"dann_loss\"].append(dann_loss_avg)\n\n    print(f\"\\n[{epoch:02d}/{CFG.epochs}] time={time.time()-t0:.1f}s\")\n    print(f\"train_loss={train_loss:.5f} train_dice={train_metrics['dice']:.5f} \"\n          f\"dann_loss={dann_loss_avg:.5f}\")\n    print(f\"val_loss={val_loss:.5f} val_dice={val_metrics['dice']:.5f} val_iou={val_metrics['iou']:.5f}\")\n    print(f\"precision={val_metrics['precision']:.5f} recall={val_metrics['recall']:.5f}\")\n    print(f\"encoder_lr={optimizer.param_groups[0]['lr']:.7f} decoder_lr={optimizer.param_groups[1]['lr']:.7f}\")\n\n    if val_metrics[\"dice\"] > best_val_dice:\n        best_val_dice = val_metrics[\"dice\"]\n        epochs_no_improve = 0\n        save_state = ema.state_dict() if ema is not None else model.state_dict()\n        checkpoint = {\"model\": save_state, \"cfg\": cfg_to_dict(CFG), \"best_val_dice\": best_val_dice,\n                      \"history\": history, \"pos_frac\": pos_frac, \"pos_weight\": pos_weight_val,\n                      \"used_ema\": CFG.USE_EMA}\n        torch.save(checkpoint, CFG.ckpt_path)\n        print(f\"*** NEW BEST CHECKPOINT val_dice={best_val_dice:.5f} \"\n              f\"({'EMA' if CFG.USE_EMA else 'raw'} weights) ***\")\n    else:\n        epochs_no_improve += 1\n        print(f\"No improvement: {epochs_no_improve}/{CFG.early_stop_patience}\")\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(\"Early stopping.\")\n            break\n\n    cleanup_memory()\n\nprint(\"\\nBest validation Dice:\", best_val_dice)\n\nif feature_capture is not None:\n    feature_capture.remove()\n\n\n# ============================================================\n# 20. LOAD BEST MODEL\n# ============================================================\n\ncheckpoint = torch.load(CFG.ckpt_path, map_location=CFG.device)\nmodel.load_state_dict(checkpoint[\"model\"])\nprint(f\"\\nBest checkpoint loaded ({'EMA' if checkpoint.get('used_ema') else 'raw'} weights).\")\n\n\n# ============================================================\n# 21. FINE DICE THRESHOLD SEARCH\n# ============================================================\n\n@torch.no_grad()\ndef find_best_dice_threshold(model, loader):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice_scores = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    best_idx = int(np.argmax(dice_scores))\n    return float(thresholds[best_idx]), float(dice_scores[best_idx]), thresholds, dice_scores\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINE DICE THRESHOLD SEARCH\")\nprint(\"=\" * 70)\n\nbest_threshold, threshold_dice, threshold_grid, threshold_scores = find_best_dice_threshold(model, val_loader)\nprint(f\"BEST VALIDATION THRESHOLD = {best_threshold:.2f}\")\nprint(f\"DICE AT BEST THRESHOLD = {threshold_dice:.5f}\")\n\n_, final_val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal validation metrics at optimized Dice threshold:\", final_val_metrics)\n\n\n# ============================================================\n# 22. INFERENCE HELPERS\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference_v2(model, vol, mask, patch_size, stride, device, batch_size,\n                                 pre_transform=None):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw_u8 = vol.read_patch(y, x, patch_size)\n        if pre_transform is not None:\n            raw_u8 = pre_transform(raw_u8)\n        raw = raw_u8.astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm_v2(model, vol, mask, patch_size, stride, device, max_patches, batch_size):\n    print(\"\\nStarting AdaBN...\")\n    n_reset = 0\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n            n_reset += 1\n    print(f\"Reset {n_reset} BatchNorm2d layers (see the architecture report above for where they live).\")\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"AdaBN patches: {len(coords)}\")\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        batch = coords[start:start + batch_size]\n        imgs = []\n        for y, x in batch:\n            raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n            raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n            imgs.append(raw)\n        inp = torch.from_numpy(np.stack(imgs)).to(device, non_blocking=True)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp)\n        del inp, imgs\n    model.eval()\n    cleanup_memory()\n    print(\"AdaBN finished.\")\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess_v2(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 23. LOAD TEST LABELS (local diagnostics only, loaded late)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST LABELS (diagnostics only)\")\nprint(\"=\" * 70)\n\ntest_labels = load_ink_labels(test_dir)\nif test_labels is not None:\n    print(\"Test GT found: local diagnostic evaluation enabled.\")\n    gt_test = (test_labels * test_mask).astype(np.float32)\nelse:\n    print(\"No test GT found: running competition-style inference.\")\n    gt_test = None\n\n\n# ============================================================\n# 24. INFERENCE A: BASELINE + TTA\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE A: BASELINE + TTA\")\nprint(\"=\" * 70)\n\ntest_prob_baseline = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n\ntest_metrics_baseline = None\nif test_labels is not None:\n    test_metrics_baseline = evaluate_probability_map(test_prob_baseline, gt_test, best_threshold)\n    print(\"\\nBASELINE TEST METRICS:\", test_metrics_baseline)\n\n\n# ============================================================\n# 25. INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\")\nprint(\"=\" * 70)\n\ntrain_hist_pool = []\nfor fid in CFG.train_frags:\n    v = train_volumes[fid]\n    m = train_masks_full[fid]\n    coords = generate_grid_coords(m, CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n    if coords:\n        for (y, x) in random.sample(coords, min(10, len(coords))):\n            train_hist_pool.append(v.read_patch(y, x, CFG.patch_size))\nprint(f\"Built train-domain reference pool for diagnostic: {len(train_hist_pool)} patches\")\n\n\ndef histogram_match_to_train_domain(raw_u8_dhw):\n    if not train_hist_pool:\n        return raw_u8_dhw\n    ref = random.choice(train_hist_pool)\n    img_hwd = np.transpose(raw_u8_dhw, (1, 2, 0))\n    ref_hwd = np.transpose(ref, (1, 2, 0))\n    try:\n        matched = match_histograms(img_hwd, ref_hwd, channel_axis=None).astype(np.uint8)\n        return np.transpose(matched, (2, 0, 1))\n    except Exception:\n        return raw_u8_dhw\n\n\ntest_prob_histmatch = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch,\n    pre_transform=histogram_match_to_train_domain)\n\ntest_metrics_histmatch = None\nif test_labels is not None:\n    test_metrics_histmatch = evaluate_probability_map(test_prob_histmatch, gt_test, best_threshold)\n    print(\"\\nHISTOGRAM-MATCHING DIAGNOSTIC TEST METRICS:\", test_metrics_histmatch)\n    if test_metrics_baseline is not None:\n        delta = test_metrics_histmatch[\"dice\"] - test_metrics_baseline[\"dice\"]\n        print(f\"\\n>>> Histogram matching alone changed local test Dice by {delta:+.4f} \"\n              f\"(baseline {test_metrics_baseline['dice']:.4f} -> {test_metrics_histmatch['dice']:.4f})\")\n\n# --- save the histogram-matched probability map + thresholded prediction, same\n# treatment as the other prediction variants ---\ntest_pred_histmatch_bin = postprocess_v2(test_prob_histmatch, best_threshold)\nhistmatch_prob_path = os.path.join(CFG.out_dir, \"fragment1_probability_histmatch_v4.npy\")\nnp.save(histmatch_prob_path, test_prob_histmatch)\nprint(\"Saved histogram-matched probability map:\", histmatch_prob_path)\nhistmatch_pred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_histmatch_v4.png\")\ncv2.imwrite(histmatch_pred_path, (test_pred_histmatch_bin * 255).astype(np.uint8))\nprint(\"Saved histogram-matched prediction:\", histmatch_pred_path)\n\n\n# ============================================================\n# 26. INFERENCE C: ADABN + TTA (only if enabled -- see architecture report)\n# ============================================================\n\ntest_prob_adabn = None\ntest_metrics_adabn = None\n\nif CFG.use_adabn:\n    print(\"\\n\" + \"=\" * 70)\n    print(\"INFERENCE C: ADABN + TTA\")\n    print(\"=\" * 70)\n    model = recalibrate_batchnorm_v2(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                      CFG.device, CFG.adabn_max_patches, CFG.infer_batch)\n    test_prob_adabn = sliding_window_inference_v2(\n        model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n    if test_labels is not None:\n        test_metrics_adabn = evaluate_probability_map(test_prob_adabn, gt_test, best_threshold)\n        print(\"\\nADABN TEST METRICS:\", test_metrics_adabn)\nelse:\n    print(\"\\n(Skipping AdaBN inference -- CFG.use_adabn=False. See the architecture report's \"\n          \"BatchNorm2d location summary above for why/whether it would help.)\")\n\n\n# ============================================================\n# 27. CHOOSE FINAL PROBABILITY MAP\n# ============================================================\n\ntest_prob = test_prob_adabn if (CFG.use_adabn and test_prob_adabn is not None) else test_prob_baseline\n\notsu_threshold = compute_otsu_threshold(test_prob, test_mask, fallback=best_threshold)\nprint(\"\\nValidation-tuned threshold:\", best_threshold)\nprint(\"Unsupervised Otsu threshold:\", otsu_threshold)\n\nfinal_threshold = best_threshold\ntest_pred_bin = postprocess_v2(test_prob, final_threshold)\n\ntest_metrics_raw = test_metrics_otsu = test_metrics_post = None\nif test_labels is not None:\n    test_metrics_raw = evaluate_probability_map(test_prob, gt_test, final_threshold)\n    test_metrics_otsu = evaluate_probability_map(test_prob, gt_test, otsu_threshold)\n    post_preds = test_pred_bin.astype(np.float32)\n    tp = (post_preds * gt_test).sum()\n    fp = (post_preds * (1.0 - gt_test)).sum()\n    fn = ((1.0 - post_preds) * gt_test).sum()\n    test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LOCAL TEST DIAGNOSTICS SUMMARY\")\n    print(\"=\" * 70)\n    print(\"A) Baseline + val threshold:            \", test_metrics_baseline)\n    print(\"B) Histogram-matched (no retrain) + val threshold:\", test_metrics_histmatch)\n    print(\"C) AdaBN + val threshold:                \", test_metrics_adabn)\n    print(\"Final (chosen) raw + val threshold:      \", test_metrics_raw)\n    print(\"Final (chosen) raw + Otsu threshold:      \", test_metrics_otsu)\n    print(\"Final (chosen) postprocessed + val threshold:\", test_metrics_post)\n\n\n# ============================================================\n# 28. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_v4.npy\")\nnp.save(prob_path, test_prob)\nprint(\"\\nSaved probability map:\", prob_path)\n\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_v4.png\")\ncv2.imwrite(pred_path, (test_pred_bin * 255).astype(np.uint8))\nprint(\"Saved prediction:\", pred_path)\n\nmetrics_summary = {\n    \"best_validation_dice_at_0.50\": best_val_dice,\n    \"best_validation_threshold\": best_threshold,\n    \"validation_dice_at_best_threshold\": threshold_dice,\n    \"otsu_threshold\": otsu_threshold,\n    \"final_threshold\": final_threshold,\n    \"final_validation_metrics\": final_val_metrics,\n    \"test_baseline\": test_metrics_baseline,\n    \"test_histogram_matched_diagnostic_no_retrain\": test_metrics_histmatch,\n    \"test_adabn\": test_metrics_adabn,\n    \"test_raw\": test_metrics_raw,\n    \"test_otsu\": test_metrics_otsu,\n    \"test_postprocessed\": test_metrics_post,\n    \"config\": cfg_to_dict(CFG),\n    \"history\": history,\n}\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump(metrics_summary, f, indent=2)\nprint(\"Saved metrics:\", CFG.metrics_path)\n\n\n# ============================================================\n# 29. VISUALIZATION  (+ prediction-vs-ground-truth overlay on every comparison)\n# ============================================================\n\ndef make_confusion_overlay(input_gray_u8, gt_binary, pred_binary, alpha=0.55):\n    \"\"\"RGB error-map overlay on top of the grayscale input:\n       green  = true positive  (prediction AND ground truth)\n       red    = false positive (prediction only)\n       yellow = false negative (ground truth only)\n    Makes prediction-vs-ground-truth mismatches immediately visible, instead of\n    having to mentally compare two separate side-by-side panels.\"\"\"\n    base = np.stack([input_gray_u8] * 3, axis=-1).astype(np.float32)\n    overlay = base.copy()\n    gt_b = gt_binary > 0\n    pred_b = pred_binary > 0\n    tp = pred_b & gt_b\n    fp = pred_b & ~gt_b\n    fn = ~pred_b & gt_b\n    color_tp = np.array([255, 0, 0], dtype=np.float32)\n    #color_tp = np.array([0, 255, 0], dtype=np.float32)\n    #color_fp = np.array([255, 0, 0], dtype=np.float32)\n    color_fp = np.array([0, 255, 0], dtype=np.float32)\n    color_fn = np.array([255, 255, 0], dtype=np.float32)\n    overlay[tp] = (1 - alpha) * base[tp] + alpha * color_tp\n    overlay[fp] = (1 - alpha) * base[fp] + alpha * color_fp\n    overlay[fn] = (1 - alpha) * base[fn] + alpha * color_fn\n    return np.clip(overlay, 0, 255).astype(np.uint8)\n\n\ndef make_prediction_only_overlay(input_gray_u8, pred_binary, alpha=0.5, color=(255, 0, 0)):\n    \"\"\"Single-color prediction overlay for when no ground truth is available (real\n    competition-style inference on an unlabeled fragment).\"\"\"\n    base = np.stack([input_gray_u8] * 3, axis=-1).astype(np.float32)\n    overlay = base.copy()\n    mask = pred_binary > 0\n    overlay[mask] = (1 - alpha) * base[mask] + alpha * np.array(color, dtype=np.float32)\n    return np.clip(overlay, 0, 255).astype(np.uint8)\n\n\ndef save_overview(prob_map, pred_bin, tag, threshold_label):\n    \"\"\"Saves a full-fragment overview for ONE prediction variant (baseline /\n    histogram-matched / AdaBN / final), including a confusion-map overlay panel\n    against ground truth when available.\"\"\"\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    small_u8 = cv2.normalize(small, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n\n    pred_small = cv2.resize((pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    pred_small_bin = (pred_small > 127).astype(np.uint8)\n    prob_small = cv2.resize((prob_map * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_small = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1],\n                               interpolation=cv2.INTER_NEAREST)\n        gt_small_bin = (gt_small > 127).astype(np.uint8)\n        overlay = make_confusion_overlay(small_u8, gt_small_bin, pred_small_bin)\n\n        fig, axes = plt.subplots(1, 5, figsize=(27, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n        axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability ({threshold_label})\")\n        axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(f\"Prediction [{tag}]\")\n        axes[4].imshow(overlay); axes[4].set_title(\"Overlay: green=FP red=TP yellow=FN\")\n    else:\n        overlay = make_prediction_only_overlay(small_u8, pred_small_bin)\n        fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(prob_small, cmap=\"gray\"); axes[1].set_title(f\"Probability ({threshold_label})\")\n        axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(f\"Prediction [{tag}]\")\n        axes[3].imshow(overlay); axes[3].set_title(\"Prediction overlay (no GT available)\")\n\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    overview_path = os.path.join(CFG.viz_dir, f\"fragment1_v4_overview_{tag}.png\")\n    plt.savefig(overview_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved {tag} overview:\", overview_path)\n\n\n# --- one overview (with overlay) per prediction variant --------------------------\ntest_pred_baseline_bin = postprocess_v2(test_prob_baseline, best_threshold)\nsave_overview(test_prob_baseline, test_pred_baseline_bin, tag=\"baseline\",\n              threshold_label=f\"val_thr={best_threshold:.2f}\")\n\nsave_overview(test_prob_histmatch, test_pred_histmatch_bin, tag=\"histmatch\",\n              threshold_label=f\"val_thr={best_threshold:.2f}\")\n\nif test_prob_adabn is not None:\n    test_pred_adabn_bin = postprocess_v2(test_prob_adabn, best_threshold)\n    save_overview(test_prob_adabn, test_pred_adabn_bin, tag=\"adabn\",\n                  threshold_label=f\"val_thr={best_threshold:.2f}\")\n\nsave_overview(test_prob, test_pred_bin, tag=\"final\",\n              threshold_label=f\"final_thr={final_threshold:.2f}\")\n\n\ndef save_patch_comparisons(n=6):\n    coords = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(coords)\n    coords = coords[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(coords):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        input_slice_u8 = cv2.normalize(input_slice, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n        prob_patch = test_prob[y:y + size, x:x + size]\n\n        if test_labels is not None:\n            gt_patch = test_labels[y:y + size, x:x + size]\n            overlay_patch = make_confusion_overlay(input_slice_u8, gt_patch, pred_patch)\n            fig, axes = plt.subplots(1, 5, figsize=(20, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n            axes[2].imshow(prob_patch, cmap=\"gray\"); axes[2].set_title(\"Probability\")\n            axes[3].imshow(pred_patch, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n            axes[4].imshow(overlay_patch); axes[4].set_title(\"Overlay: green=TP red=FP yellow=FN\")\n        else:\n            overlay_patch = make_prediction_only_overlay(input_slice_u8, pred_patch)\n            fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(prob_patch, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n            axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n            axes[3].imshow(overlay_patch); axes[3].set_title(\"Prediction overlay (no GT)\")\n\n        for ax in axes:\n            ax.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n        plt.close(fig)\n    print(f\"Saved {len(coords)} patch comparisons.\")\n\n\nsave_patch_comparisons(n=6)\n\n\ndef save_training_curves():\n    epochs_axis = np.arange(1, len(history[\"train_loss\"]) + 1)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_loss\"], label=\"Train Loss\")\n    plt.plot(epochs_axis, history[\"val_loss\"], label=\"Val Loss\")\n    if CFG.USE_DANN:\n        plt.plot(epochs_axis, history[\"dann_loss\"], label=\"DANN domain loss\", linestyle=\"--\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Loss\"); plt.title(\"V4 Training Loss\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_loss.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_dice\"], label=\"Train Dice\")\n    plt.plot(epochs_axis, history[\"val_dice\"], label=\"Val Dice (EMA)\" if CFG.USE_EMA else \"Val Dice\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Dice\"); plt.title(\"V4 Dice\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_dice.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(threshold_grid, threshold_scores)\n    plt.axvline(best_threshold, linestyle=\"--\", label=f\"best={best_threshold:.2f}\")\n    plt.xlabel(\"Threshold\"); plt.ylabel(\"Validation Dice\"); plt.title(\"Dice Threshold Search\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"threshold_search.png\"), dpi=150)\n    plt.close(fig)\n\n    print(\"Saved training curves.\")\n\n\nsave_training_curves()\n\n\n# ============================================================\n# 30. CLEANUP\n# ============================================================\n\ntest_vol.close()\nfor v in train_volumes.values():\n    v.close()\ncleanup_memory()\n\n\n# ============================================================\n# 31. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V4 COMPLETE\")\nprint(\"=\" * 70)\nprint(f\"Best validation Dice @ 0.50: {best_val_dice:.5f}\")\nprint(f\"Best Dice threshold: {best_threshold:.2f}\")\nprint(f\"Validation Dice @ optimized threshold: {threshold_dice:.5f}\")\nprint(f\"Backbone: {CFG.encoder_name} / architecture={CFG.architecture} / \"\n      f\"depth_stem={CFG.USE_DEPTH_FUSION_STEM}\")\nprint(f\"EMA: {CFG.USE_EMA} | DANN: {CFG.USE_DANN} | MixStyle: {CFG.USE_MIXSTYLE} | AdaBN: {CFG.use_adabn}\")\nprint(f\"CLAHE mode: {CFG.CLAHE_MODE} | Normalization: {CFG.NORMALIZATION_MODE} | \"\n      f\"Mask erosion: {CFG.mask_erode_px}px\")\nif test_metrics_baseline is not None:\n    print(f\"\\nLocal test Dice -- baseline: {test_metrics_baseline['dice']:.5f}\")\nif test_metrics_histmatch is not None:\n    print(f\"Local test Dice -- histogram-matched (diagnostic, no retrain): \"\n          f\"{test_metrics_histmatch['dice']:.5f}\")\nif test_metrics_adabn is not None:\n    print(f\"Local test Dice -- AdaBN: {test_metrics_adabn['dice']:.5f}\")\nif test_metrics_post is not None:\n    print(f\"Local test Dice -- final postprocessed: {test_metrics_post['dice']:.5f}\")\nprint(f\"\\nBest checkpoint: {CFG.ckpt_path}\")\nprint(f\"Probability map: {prob_path}\")\nprint(f\"Prediction: {pred_path}\")\nprint(f\"Metrics: {CFG.metrics_path}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")\nprint(\"\\n=== V4 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-15T06:35:34.197578Z","iopub.execute_input":"2026-09-15T06:35:34.197969Z","execution_failed":"2026-09-15T08:46:28.857Z"}},"outputs":[{"name":"stdout","text":"     ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 43.7/43.7 kB 693.7 kB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 2.2 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.6/2.6 MB 14.7 MB/s eta 0:00:00\nCollecting segmentation-models-pytorch==0.4.0\n  Downloading segmentation_models_pytorch-0.4.0-py3-none-any.whl.metadata (32 kB)\nCollecting efficientnet-pytorch>=0.6.1 (from segmentation-models-pytorch==0.4.0)\n  Downloading efficientnet_pytorch-0.7.1.tar.gz (21 kB)\n  Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\nRequirement already satisfied: huggingface-hub>=0.24 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.11.0)\nRequirement already satisfied: numpy>=1.19.3 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (2.0.2)\nRequirement already satisfied: pillow>=8 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (11.3.0)\nCollecting pretrainedmodels>=0.7.1 (from segmentation-models-pytorch==0.4.0)\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m2.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\nRequirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.17.0)\nRequirement already satisfied: timm>=0.9 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.0.29)\nRequirement already satisfied: torch>=1.8 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (2.10.0+cu128)\nRequirement already satisfied: torchvision>=0.9 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (0.25.0+cu128)\nRequirement already satisfied: tqdm>=4.42.1 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (4.67.3)\nRequirement already satisfied: filelock>=3.10.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (3.29.0)\nRequirement already satisfied: fsspec>=2023.5.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2025.3.0)\nRequirement already satisfied: hf-xet<2.0.0,>=1.4.3 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.4.3)\nRequirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.28.1)\nRequirement already satisfied: packaging>=20.9 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (26.1)\nRequirement already satisfied: pyyaml>=5.1 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (6.0.3)\nRequirement already satisfied: typer in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.24.2)\nRequirement already satisfied: typing-extensions>=4.1.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.15.0)\nCollecting munch (from pretrainedmodels>=0.7.1->segmentation-models-pytorch==0.4.0)\n  Downloading munch-4.0.0-py2.py3-none-any.whl.metadata (5.9 kB)\nRequirement already satisfied: safetensors in /usr/local/lib/python3.12/dist-packages (from timm>=0.9->segmentation-models-pytorch==0.4.0) (0.7.0)\nRequirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (81.0.0)\nRequirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (1.14.0)\nRequirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.6.1)\nRequirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.1.6)\nRequirement already satisfied: cuda-bindings==12.9.4 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.9.4)\nRequirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.93)\nRequirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (9.10.2.21)\nRequirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.4.1)\nRequirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (11.3.3.83)\nRequirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (10.3.9.90)\nRequirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (11.7.3.90)\nRequirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.5.8.93)\nRequirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (0.7.1)\nRequirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (2.27.5)\nRequirement already satisfied: nvidia-nvshmem-cu12==3.4.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.4.5)\nRequirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.93)\nRequirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (1.13.1.3)\nRequirement already satisfied: triton==3.6.0 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.6.0)\nRequirement already satisfied: cuda-pathfinder~=1.1 in /usr/local/lib/python3.12/dist-packages (from cuda-bindings==12.9.4->torch>=1.8->segmentation-models-pytorch==0.4.0) (1.5.3)\nRequirement already satisfied: anyio in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.13.0)\nRequirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2026.4.22)\nRequirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.0.9)\nRequirement already satisfied: idna in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (3.13)\nRequirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.16.0)\nRequirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch>=1.8->segmentation-models-pytorch==0.4.0) (1.3.0)\nRequirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch>=1.8->segmentation-models-pytorch==0.4.0) (3.0.3)\nRequirement already satisfied: click>=8.2.1 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (8.3.3)\nRequirement already satisfied: shellingham>=1.3.0 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.5.4)\nRequirement already satisfied: rich>=12.3.0 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (13.9.4)\nRequirement already satisfied: annotated-doc>=0.0.2 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.0.4)\nRequirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.0.0)\nRequirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2.20.0)\nRequirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.1.2)\nDownloading segmentation_models_pytorch-0.4.0-py3-none-any.whl (121 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m121.3/121.3 kB\u001b[0m \u001b[31m3.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading munch-4.0.0-py2.py3-none-any.whl (9.9 kB)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25l\u001b[?25hdone\n  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.7.1-py3-none-any.whl size=16477 sha256=c5a1f26b15c570fb025d5af5869d70c705094e7db14dfe7ac2b5639a09627909\n  Stored in directory: /root/.cache/pip/wheels/9c/3f/43/e6271c7026fe08c185da2be23c98c8e87477d3db63f41f32ad\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25l\u001b[?25hdone\n  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60990 sha256=91eb3cca4c4db74bfd59ec0459f952b1b566192016cc40e106d468f5f668075e\n  Stored in directory: /root/.cache/pip/wheels/4c/01/56/40a48f75dbdfe167a0cb70d3b48913369a00ec5c4e9fed5f2b\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: munch, efficientnet-pytorch, pretrainedmodels, segmentation-models-pytorch\n  Attempting uninstall: segmentation-models-pytorch\n    Found existing installation: segmentation_models_pytorch 0.5.0\n    Uninstalling segmentation_models_pytorch-0.5.0:\n      Successfully uninstalled segmentation_models_pytorch-0.5.0\nSuccessfully installed efficientnet-pytorch-0.7.1 munch-4.0.0 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.4.0\n[env] torch=2.10.0+cu128 | smp=0.4.0 | timm=1.0.29\n======================================================================\nBUILDING TRAIN DATA\n======================================================================\n\nProcessing fragment 2 ...\n  fragment 2 normalization stats: mean=118.25 std=59.13\nFragment 2: 27971 candidate patches (post mask-erosion)\n  spatial train=20919 validation=5676\n\nProcessing fragment 3 ...\n  fragment 3 normalization stats: mean=119.91 std=61.59\nFragment 3: 7323 candidate patches (post mask-erosion)\n  spatial train=5773 validation=737\n\nRaw train samples: 26692\nValidation samples: 6413\nPositive patches: 18196 | Negative patches: 8496\nBalanced dataset: 26692 | positive ratio=0.682\n\n======================================================================\nLOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\n======================================================================\n  fragment 1 normalization stats: mean=109.51 std=69.24\nFragment 1: 8592 unlabeled candidate patches\nBuilt histogram-matching reference pool: 20 patches (53.2 MB)\n\nBuilding V4 depth-aware model ...\n[backbone] DepthFusionStem: 26 depth slices (raw=True, grad=True, curv=True) -> 16 learned channels\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n","output_type":"stream"},{"name":"stdout","text":"Downloading: \"https://download.pytorch.org/models/resnet50-19c8e357.pth\" to /root/.cache/torch/hub/checkpoints/resnet50-19c8e357.pth\n","output_type":"stream"},{"name":"stderr","text":"100%|██████████| 97.8M/97.8M [00:00<00:00, 415MB/s]\n","output_type":"stream"},{"name":"stdout","text":"[backbone] using resnet50 (ImageNet-pretrained) + unet\n\n======================================================================\nV4 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320 | input tensor=[1, 26, 320, 320]\nEncoder: backbone=resnet50 | pretrained=imagenet | architecture=unet\nDepth stem: 26 -> 16 channels (raw=True, grad=True, curv=True)\nDecoder channels: (256, 128, 64, 32, 16)\nParameters: 33.86M total | 33.86M trainable\nZero-channel layers: 0 (PASS)\n\nEncoder feature stages (dry run at 256x256 for speed):\n  stage 0: (1, 16, 256, 256)\n  stage 1: (1, 64, 128, 128)\n  stage 2: (1, 256, 64, 64)\n  stage 3: (1, 512, 32, 32)\n  stage 4: (1, 1024, 16, 16)\n  stage 5: (1, 2048, 8, 8)\n\nOutput logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\n\nCUDA device: Tesla T4 | capability: (7, 5)\nForward pass: PASS\n[AdaBN report] BatchNorm2d layers found: depth_stem=1 encoder=53 decoder=10 other=0 (total=64)\n======================================================================\n\n\nEstimating positive-pixel fraction ...\nEstimated positive fraction: 0.162960\nOutput bias initialized to -1.6364\nBCE positive weight: 2.266\nDepth-stem parameters: 4 | Encoder parameters: 159 | Decoder/head parameters: 92 tensors\n\n======================================================================\nSTARTING V4 TRAINING\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/2996949287.py:1307: FutureWarning: `torch.cuda.amp.GradScaler(args...)` is deprecated. Please use `torch.amp.GradScaler('cuda', args...)` instead.\n  scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n/tmp/ipykernel_58/2996949287.py:1346: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"\n[01/8] time=1453.1s\ntrain_loss=0.75331 train_dice=0.40722 dann_loss=0.00000\nval_loss=0.73110 val_dice=0.46344 val_iou=0.30161\nprecision=0.54653 recall=0.40228\nencoder_lr=0.0000481 decoder_lr=0.0000962\n*** NEW BEST CHECKPOINT val_dice=0.46344 (EMA weights) ***\n\n[02/8] time=1276.1s\ntrain_loss=0.67262 train_dice=0.54684 dann_loss=0.00000\nval_loss=0.69992 val_dice=0.50170 val_iou=0.33485\nprecision=0.40315 recall=0.66402\nencoder_lr=0.0000428 decoder_lr=0.0000855\n*** NEW BEST CHECKPOINT val_dice=0.50170 (EMA weights) ***\n\n[03/8] time=1273.3s\ntrain_loss=0.61783 train_dice=0.61847 dann_loss=0.00000\nval_loss=0.71926 val_dice=0.50149 val_iou=0.33466\nprecision=0.41917 recall=0.62403\nencoder_lr=0.0000349 decoder_lr=0.0000694\nNo improvement: 1/3\n\n[04/8] time=1288.5s\ntrain_loss=0.56965 train_dice=0.68343 dann_loss=0.00000\nval_loss=0.75914 val_dice=0.49421 val_iou=0.32821\nprecision=0.44003 recall=0.56360\nencoder_lr=0.0000255 decoder_lr=0.0000505\nNo improvement: 2/3\n\n[05/8] time=1289.4s\ntrain_loss=0.52999 train_dice=0.73385 dann_loss=0.00000\nval_loss=0.79206 val_dice=0.49220 val_iou=0.32644\nprecision=0.46926 recall=0.51750\nencoder_lr=0.0000161 decoder_lr=0.0000316\nNo improvement: 3/3\nEarly stopping.\n\nBest validation Dice: 0.5017048586918491\n\nBest checkpoint loaded (EMA weights).\n\n======================================================================\nFINE DICE THRESHOLD SEARCH\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/2996949287.py:1502: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"BEST VALIDATION THRESHOLD = 0.69\nDICE AT BEST THRESHOLD = 0.52624\n\nFinal validation metrics at optimized Dice threshold: {'dice': 0.52623653767537, 'iou': 0.3570698766308928, 'precision': 0.516331398225119, 'recall': 0.5365291443464854, 'fbeta0.5': 0.520249089611028}\n\n======================================================================\nLOADING TEST LABELS (diagnostics only)\n======================================================================\nTest GT found: local diagnostic evaluation enabled.\n\n======================================================================\nINFERENCE A: BASELINE + TTA\n======================================================================\nInference patches: 8822\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/2996949287.py:1571: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"\nBASELINE TEST METRICS: {'dice': 0.40366464853286743, 'iou': 0.252869576215744, 'precision': 0.4525715410709381, 'recall': 0.364297091960907, 'fbeta0.5': 0.43165358901023865}\n\n======================================================================\nINFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\n======================================================================\nBuilt train-domain reference pool for diagnostic: 20 patches\nInference patches: 8822\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"#Local test [F0.5] -- histogram-matched: 0.4974, dice=0.44 with img=128, stride=32\n#Local test [F0.5] -- histogram-matched: 0.4994, dice=0.47 with img=224, stride=64\n\n# ============================================================\n# VESUVIUS INK DETECTION - V4 \"DEPTH-AWARE CONVNEXT\"\n# ============================================================\n# Built on V3 (Dice=0.499). Fixes the crash and addresses the deeper\n# architectural critique: the model was treating 26 ordered depth\n# slices as unordered multispectral channels.\n#\n# ------------------------------------------------------------------\n# WHAT ACTUALLY CAUSED THE CRASH (verified empirically, not guessed)\n# ------------------------------------------------------------------\n# The proposed fix of adding `decoder_channels=(256,128,64,32,16)` to\n# smp.UnetPlusPlus was tested directly against this exact setup and\n# it does NOT fix the crash -- I reproduced the identical\n# `weight of size [0, 96, 3, 3]` error with it. The real cause:\n# ConvNeXt's stem is a stride-4 patchify with no separate stride-2\n# stage (unlike ResNet), so smp's generic \"tu-\" timm-encoder wrapper\n# fabricates an EMPTY 0-channel placeholder feature map to keep a\n# uniform 5-stage pyramid API. `smp.Unet`'s decoder already handles\n# that 0-channel stage gracefully (confirmed working); `smp.UnetPlusPlus`'s\n# dense skip-connections do not (confirmed failing, with or without\n# explicit decoder_channels). So V4 uses `smp.Unet`, not UnetPlusPlus,\n# with tu-convnext_tiny. This is a real library limitation, not a\n# parameter you can configure around.\n#\n# ------------------------------------------------------------------\n# DEPTH-AWARE CHANGES (addressing \"26 slices as unordered channels\")\n# ------------------------------------------------------------------\n#  1. DepthFusionStem: computes explicit first/second finite\n#     differences along the physically-ordered depth axis (how\n#     intensity changes through the papyrus), concatenates them with\n#     the raw stack, and learns a compact (16-32 channel) mixed\n#     representation via a 1x1 conv BEFORE the 2D ConvNeXt encoder\n#     ever sees the data -- giving the network an explicit inductive\n#     bias toward depth structure instead of hoping a from-scratch\n#     first-conv discovers it.\n#  2. CLAHE_MODE: \"per_slice\" (V3's old behavior, independent CLAHE\n#     per slice -- can distort inter-slice relationships) vs.\n#     \"global_shared\" (ONE contrast-remapping LUT derived from a\n#     representative slice, applied identically to every slice --\n#     preserves relative depth relationships). Both available for\n#     the ablation you suggested; default is now \"global_shared\".\n#  3. Histogram-matching augmentation now uses ONE shared mapping\n#     across the whole depth stack (`channel_axis=None`) instead of\n#     26 independent per-slice mappings (`channel_axis=2`). I verified\n#     this empirically: independent per-channel matching compressed\n#     slice-to-slice differences unevenly (2.3-3.8 range in a test),\n#     while the shared mapping preserved them far more consistently\n#     (4.3-7.3 range, proportional to the original 7.5-8.6 spacing).\n#  4. Architecture inspection report + zero-channel detector run\n#     BEFORE the optimizer/training loop are created -- this is\n#     exactly the check that would have caught the original crash\n#     immediately instead of after a dummy forward pass deep in setup.\n#  5. AdaBN now prints exactly where its BatchNorm2d layers are found\n#     (depth stem / decoder / encoder) since ConvNeXt itself is\n#     LayerNorm-only and has none -- so you know what it's actually\n#     recalibrating before deciding whether to enable it.\n#  6. V4 baseline: USE_DANN / USE_MIXSTYLE / use_adabn default to\n#     False so the depth-aware architecture change can be evaluated\n#     cleanly on its own first. Turn them back on one at a time\n#     afterward -- CFG has a comment marking each one.\n#  7. One explicit environment setup (current smp + timm, no more\n#     pinning smp==0.2.0 and upgrading from inside build_model()).\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS  (single explicit environment, no version-pin dance)\n# ============================================================\n!pip install segmentation-models-pytorch==0.4.0\n\nimport os as _os\n_os.system(\"pip install -q -U segmentation-models-pytorch timm\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom scipy import ndimage\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport timm\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\nprint(f\"[env] torch={torch.__version__} | smp={smp.__version__} | timm={timm.__version__}\")\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    patch_size = 320\n    train_stride = 64\n    test_stride = 64\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.40\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 7\n    early_stop_patience = 3\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    depth_stem_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- V4: backbone / architecture -----------------------------------------\n    # VERIFIED WORKING with tu-convnext_tiny: smp.Unet (NOT UnetPlusPlus -- see\n    # module docstring for the empirical reason). Falls back to efficientnet-b4\n    # (a \"real\" smp encoder with no 0-channel-stage quirk) if ConvNeXt/timm are\n    # unavailable in this session for any reason.\n    encoder_name = \"tu-convnext_tiny\"\n    encoder_fallback = \"efficientnet-b4\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    architecture = \"unet\"                          # do NOT set to \"unetplusplus\" with tu-* encoders\n    decoder_channels = (256, 128, 64, 32, 16)\n\n    # --- V4: depth-aware fusion stem -----------------------------------------\n    USE_DEPTH_FUSION_STEM = True\n    depth_stem_out_channels = 16\n    depth_stem_use_raw = True\n    depth_stem_use_grad = True\n    depth_stem_use_curv = True\n\n    # --- LOSS -----------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.35\n    focal_tversky_weight = 0.35\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n\n    # --- THRESHOLD --------------------------------------------------------------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n    threshold = 0.50\n\n    # --- TTA --------------------------------------------------------------------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN (see the \"where are the BN layers\" report at model-build time) --\n    # ConvNeXt is LayerNorm-only; V4 defaults this OFF until the printed report\n    # shows there's something meaningful (decoder / depth-stem BN) for it to do.\n    use_adabn = False\n    adabn_max_patches = 2000\n\n    # --- POSTPROCESS ----------------------------------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    # --- edge-artifact cropping -------------------------------------------------\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- V4: CLAHE mode ---------------------------------------------------------\n    # \"off\"           : no contrast enhancement\n    # \"per_slice\"     : V3's old behavior -- independent CLAHE per slice, can\n    #                   distort inter-slice depth relationships\n    # \"global_shared\" : ONE remapping LUT derived from a representative slice,\n    #                   applied identically to every slice -- preserves relative\n    #                   depth relationships. DEFAULT for V4.\n    CLAHE_MODE = \"global_shared\"\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization ------------------------------------------------------------\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n\n    # --- domain-randomization augmentation --------------------------------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n\n    USE_SHADOW = True\n    shadow_p = 0.20\n\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # histogram-matching style augmentation -- NOW uses one shared mapping across\n    # depth (channel_axis=None) instead of 26 independent per-slice mappings.\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- EMA --------------------------------------------------------------------\n    USE_EMA = True\n    ema_decay = 0.999\n\n    # --- DANN (OFF for the V4 baseline -- turn on only after the depth-aware\n    #     architecture change is validated on its own; see module docstring) ----\n    USE_DANN = False\n    #USE_DANN = True\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n\n    # --- MixStyle (OFF for the V4 baseline, same reasoning as DANN) -------------\n    #USE_MIXSTYLE = False\n    USE_MIXSTYLE = True\n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    ckpt_path = os.path.join(out_dir, \"vesuviusnet_v4_best.pth\")\n    viz_dir = os.path.join(out_dir, \"v4_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v4_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\nset_seed(CFG.seed)\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS  (+ mask erosion)\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING (unchanged)\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME  (+ V4: CLAHE modes, per-fragment stats)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.CLAHE_MODE == \"per_slice\" else None)\n        self._shared_clahe_lut = None   # built lazily for CLAHE_MODE == \"global_shared\"\n        self.frag_mean = None\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    @staticmethod\n    def _compute_shared_clahe_lut(reference_slice_uint8, clip_limit, tile_grid):\n        \"\"\"Derives ONE 256-entry intensity-remapping lookup table from CLAHE applied\n        to a single representative slice, then this exact LUT is applied identically\n        to every depth slice via cv2.LUT. Unlike calling .apply() independently per\n        slice, this guarantees the same monotonic mapping everywhere, so relative\n        inter-slice intensity relationships (the actual depth signal) are preserved.\"\"\"\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n        remapped = clahe.apply(reference_slice_uint8)\n        lut = np.zeros(256, dtype=np.uint8)\n        orig = reference_slice_uint8.ravel().astype(np.int64)\n        remap = remapped.ravel().astype(np.int64)\n        sums = np.zeros(256, dtype=np.float64)\n        counts = np.zeros(256, dtype=np.int64)\n        np.add.at(sums, orig, remap)\n        np.add.at(counts, orig, 1)\n        valid = counts > 0\n        lut[valid] = np.clip(sums[valid] / counts[valid], 0, 255).astype(np.uint8)\n        if not valid.all() and valid.any():\n            idx = np.arange(256)\n            lut = np.interp(idx, idx[valid], lut[valid].astype(np.float64)).astype(np.uint8)\n        return lut\n\n    def _ensure_shared_clahe_lut(self):\n        if self._shared_clahe_lut is not None:\n            return\n        self._ensure_open()\n        mid = len(self._slices) // 2\n        ref_slice = np.asarray(self._slices[mid])\n        if ref_slice.dtype != np.uint8:\n            ref_slice = (ref_slice.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n        self._shared_clahe_lut = self._compute_shared_clahe_lut(\n            ref_slice, CFG.clahe_clip_limit, CFG.clahe_tile_grid)\n\n    def read_patch(self, y, x, size, apply_clahe=None):\n        self._ensure_open()\n        apply_clahe = (CFG.CLAHE_MODE != \"off\") if apply_clahe is None else apply_clahe\n        if apply_clahe and CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n\n            if apply_clahe:\n                if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                    block = self._clahe.apply(block)\n                elif CFG.CLAHE_MODE == \"global_shared\":\n                    block = cv2.LUT(block, self._shared_clahe_lut)\n            out[i] = block\n        return out\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        mid = len(self._slices) // 2\n        if CFG.CLAHE_MODE == \"global_shared\":\n            self._ensure_shared_clahe_lut()\n        vals = []\n        for (y, x) in sample_coords:\n            block = self._slices[mid][y:y + patch_size, x:x + patch_size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            if CFG.CLAHE_MODE == \"per_slice\" and self._clahe is not None:\n                block = self._clahe.apply(block)\n            elif CFG.CLAHE_MODE == \"global_shared\":\n                block = cv2.LUT(block, self._shared_clahe_lut)\n            vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION (fake ink / fiber / shadow)\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([np.random.randint(0, w, n_pts), np.random.randint(0, h, n_pts)],\n                        axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\n# ============================================================\n# 8. AUGMENTATION\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    \"\"\"Try constructing a transform with the current albumentations API; if that\n    fails (parameter names changed across versions), fall back to the older API.\n    If both fail, skip the transform rather than crashing the whole pipeline.\"\"\"\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    # NOTE: albumentations' RandomShadow hard-requires 3-channel RGB images in\n    # every version and raises ValueError on multi-channel depth-stack data.\n    # inject_fake_shadow() (applied directly in the dataset, image-only) is used\n    # instead -- see USE_SHADOW below.\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)\n        label = self.labels[fid][y:y + size, x:x + size]\n        img = np.transpose(patch, (1, 2, 0))   # HWD\n\n        # --- histogram-matching style augmentation: ONE shared mapping across\n        # depth (channel_axis=None), not 26 independent per-slice mappings. ---\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=None).astype(np.uint8)\n            except Exception:\n                pass\n\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img, vol.frag_mean, vol.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN: yields normalized (D,H,W) tensors from fragment 1, NO labels.\n    Uses the SAME basic preprocessing path (CLAHE mode, per-fragment normalization)\n    as the source dataset for consistency -- it intentionally skips the AUGMENTATION\n    pipeline (elastic/dropout/fake-ink/histogram-match), since DANN needs to see the\n    target domain's natural distribution, not an augmented/hallucinated version of it.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL: DEPTH FUSION STEM + VERIFIED-WORKING ConvNeXt+Unet\n# ============================================================\n\nclass DepthFusionStem(nn.Module):\n    \"\"\"Computes explicit depth-derivative features (first and second finite\n    differences along the physically-ordered depth axis) alongside the raw stack,\n    then learns a compact mixed representation via a 1x1 conv, producing\n    `out_channels` channels to feed into the 2D encoder. This is the \"ink isn't\n    just absolute intensity, it's how intensity changes through the papyrus\"\n    inductive bias, made explicit instead of hoping a from-scratch first-conv\n    layer discovers it purely from data.\"\"\"\n    def __init__(self, in_depth, out_channels=16, use_raw=True, use_grad=True, use_curv=True):\n        super().__init__()\n        self.use_raw, self.use_grad, self.use_curv = use_raw, use_grad, use_curv\n        total_in = 0\n        if use_raw:\n            total_in += in_depth\n        if use_grad:\n            total_in += (in_depth - 1)\n        if use_curv:\n            total_in += max(in_depth - 2, 0)\n        if total_in == 0:\n            raise ValueError(\"DepthFusionStem: at least one of use_raw/use_grad/use_curv must be True\")\n        # BatchNorm2d kept HERE deliberately (even though the ConvNeXt backbone\n        # itself is all LayerNorm) so AdaBN has a real, meaningful place to\n        # recalibrate target-domain statistics -- see the model-build report.\n        self.mix = nn.Sequential(\n            nn.Conv2d(total_in, out_channels, kernel_size=1),\n            nn.BatchNorm2d(out_channels),\n            nn.GELU(),\n        )\n        self.out_channels = out_channels\n        self.in_depth = in_depth\n\n    def forward(self, x):   # x: (B, D, H, W)\n        parts = []\n        if self.use_raw:\n            parts.append(x)\n        if self.use_grad:\n            parts.append(x[:, 1:, :, :] - x[:, :-1, :, :])\n        if self.use_curv and x.shape[1] > 2:\n            parts.append(x[:, 2:, :, :] - 2 * x[:, 1:-1, :, :] + x[:, :-2, :, :])\n        return self.mix(torch.cat(parts, dim=1))\n\n\nclass DepthAwareSegModel(nn.Module):\n    \"\"\"Wraps an smp segmentation model with a DepthFusionStem in front of it.\"\"\"\n    def __init__(self, depth_stem, seg_model):\n        super().__init__()\n        self.depth_stem = depth_stem\n        self.seg_model = seg_model\n\n    def forward(self, x):   # x: (B, D, H, W)\n        feat = self.depth_stem(x) if self.depth_stem is not None else x.unsqueeze(1) if x.dim() == 3 else x\n        return self.seg_model(feat)\n\n\ndef _build_seg_backbone(encoder_name, in_channels):\n    arch_cls = smp.Unet if CFG.architecture == \"unet\" else smp.UnetPlusPlus\n    if CFG.architecture != \"unet\":\n        print(f\"[backbone] WARNING: architecture='{CFG.architecture}' requested, but only \"\n              f\"'unet' has been verified to handle tu-* timm encoders' 0-channel placeholder \"\n              f\"stage correctly (see module docstring) -- using UnetPlusPlus may crash with \"\n              f\"ConvNeXt/other tu-* encoders.\")\n    return arch_cls(\n        encoder_name=encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=in_channels,\n        classes=1,\n        decoder_channels=CFG.decoder_channels,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\ndef build_model():\n    stem = None\n    seg_in_channels = CFG.in_channels\n    if CFG.USE_DEPTH_FUSION_STEM:\n        stem = DepthFusionStem(CFG.in_channels, CFG.depth_stem_out_channels,\n                                CFG.depth_stem_use_raw, CFG.depth_stem_use_grad, CFG.depth_stem_use_curv)\n        seg_in_channels = CFG.depth_stem_out_channels\n        print(f\"[backbone] DepthFusionStem: {CFG.in_channels} depth slices \"\n              f\"(raw={CFG.depth_stem_use_raw}, grad={CFG.depth_stem_use_grad}, curv={CFG.depth_stem_use_curv}) \"\n              f\"-> {seg_in_channels} learned channels\")\n\n    try:\n        seg_model = _build_seg_backbone(CFG.encoder_name, seg_in_channels)\n        print(f\"[backbone] using {CFG.encoder_name} (ImageNet-pretrained) + {CFG.architecture}\")\n    except Exception as e:\n        print(f\"[backbone] {CFG.encoder_name} unavailable/failed ({e}); \"\n              f\"falling back to {CFG.encoder_fallback}.\")\n        seg_model = _build_seg_backbone(CFG.encoder_fallback, seg_in_channels)\n\n    return DepthAwareSegModel(stem, seg_model)\n\n\nclass GradReverse(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x, lambd):\n        ctx.lambd = lambd\n        return x.view_as(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return -ctx.lambd * grad_output, None\n\n\ndef grad_reverse(x, lambd=1.0):\n    return GradReverse.apply(x, lambd)\n\n\nclass DomainClassifier(nn.Module):\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(in_ch, 128), nn.ReLU(inplace=True), nn.Dropout(0.3),\n            nn.Linear(128, 1),\n        )\n\n    def forward(self, feat, lambd):\n        return self.net(grad_reverse(feat, lambd))\n\n\nclass EncoderFeatureCapture:\n    \"\"\"Architecture-agnostic hook capturing the segmentation backbone's encoder\n    output feature list on every forward pass. Attaches to model.seg_model.encoder\n    (the DepthAwareSegModel wrapper's inner smp model), not model.encoder directly,\n    since V4 wraps the smp model with a depth-fusion stem in front of it.\"\"\"\n    def __init__(self, model):\n        self.features = None\n        target = model.seg_model.encoder if hasattr(model, \"seg_model\") else model.encoder\n        self.handle = target.register_forward_hook(self._hook)\n\n    def _hook(self, module, inp, out):\n        self.features = out\n\n    def remove(self):\n        self.handle.remove()\n\n\ndef mixstyle_batch(imgs, p=CFG.mixstyle_p, alpha=CFG.mixstyle_alpha):\n    if torch.rand(1).item() > p:\n        return imgs\n    B = imgs.size(0)\n    if B < 2:\n        return imgs\n    mu = imgs.mean(dim=[2, 3], keepdim=True)\n    var = imgs.var(dim=[2, 3], keepdim=True)\n    sig = (var + 1e-6).sqrt()\n    x_norm = (imgs - mu) / sig\n    perm = torch.randperm(B, device=imgs.device)\n    mu2, sig2 = mu[perm], sig[perm]\n    lam = torch.distributions.Beta(alpha, alpha).sample((B, 1, 1, 1)).to(imgs.device)\n    mu_mix = lam * mu + (1 - lam) * mu2\n    sig_mix = lam * sig + (1 - lam) * sig2\n    return x_norm * sig_mix + mu_mix\n\n\nclass EMAModel:\n    def __init__(self, model, decay=CFG.ema_decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. V4: PRE-TRAINING ARCHITECTURE INSPECTION\n#     (this is exactly what would have caught the original crash\n#     immediately, before the optimizer/training loop existed)\n# ============================================================\n\ndef inspect_model_channels(model):\n    \"\"\"Scans every Conv/Linear layer for zero or negative in/out channels -- the\n    exact failure mode behind the original UnetPlusPlus/ConvNeXt crash.\"\"\"\n    bad = []\n    for name, module in model.named_modules():\n        if isinstance(module, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):\n            if module.in_channels <= 0 or module.out_channels <= 0:\n                bad.append((name, \"conv\", module.in_channels, module.out_channels))\n        elif isinstance(module, nn.Linear):\n            if module.in_features <= 0 or module.out_features <= 0:\n                bad.append((name, \"linear\", module.in_features, module.out_features))\n    return bad\n\n\ndef report_batchnorm_locations(model):\n    \"\"\"Prints where BatchNorm2d layers actually live in the model -- ConvNeXt's\n    own encoder is LayerNorm-only, so AdaBN (which recalibrates BatchNorm running\n    stats) has nothing to do there. It CAN still do something meaningful in the\n    depth-fusion stem and/or the smp decoder, if those have BatchNorm2d layers.\"\"\"\n    counts = {\"depth_stem\": 0, \"encoder\": 0, \"decoder\": 0, \"other\": 0}\n    for name, module in model.named_modules():\n        if isinstance(module, nn.BatchNorm2d):\n            if name.startswith(\"depth_stem\"):\n                counts[\"depth_stem\"] += 1\n            elif \".encoder.\" in f\".{name}.\" or name.endswith(\".encoder\"):\n                counts[\"encoder\"] += 1\n            elif \".decoder.\" in f\".{name}.\" or name.endswith(\".decoder\"):\n                counts[\"decoder\"] += 1\n            else:\n                counts[\"other\"] += 1\n    total = sum(counts.values())\n    print(f\"[AdaBN report] BatchNorm2d layers found: depth_stem={counts['depth_stem']} \"\n          f\"encoder={counts['encoder']} decoder={counts['decoder']} other={counts['other']} \"\n          f\"(total={total})\")\n    if counts[\"encoder\"] == 0 and total > 0:\n        print(\"  -> the ConvNeXt encoder itself has none (it's LayerNorm-only, as expected). \"\n              \"AdaBN would only recalibrate the depth-stem/decoder BN layers listed above.\")\n    if total == 0:\n        print(\"  -> NO BatchNorm2d layers anywhere in this model. Enabling use_adabn would \"\n              \"currently be a no-op. (Kept OFF by default in V4 CFG for exactly this reason.)\")\n    return counts\n\n\n@torch.no_grad()\ndef run_architecture_report(model):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"V4 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\")\n    print(\"=\" * 70)\n\n    n_params = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Input:  depth slices={CFG.in_channels} | patch={CFG.patch_size}x{CFG.patch_size} | \"\n          f\"input tensor=[1, {CFG.in_channels}, {CFG.patch_size}, {CFG.patch_size}]\")\n    print(f\"Encoder: backbone={CFG.encoder_name} | pretrained={CFG.encoder_weights} | \"\n          f\"architecture={CFG.architecture}\")\n    if CFG.USE_DEPTH_FUSION_STEM:\n        print(f\"Depth stem: {CFG.in_channels} -> {model.depth_stem.out_channels} channels \"\n              f\"(raw={model.depth_stem.use_raw}, grad={model.depth_stem.use_grad}, \"\n              f\"curv={model.depth_stem.use_curv})\")\n    print(f\"Decoder channels: {CFG.decoder_channels}\")\n    print(f\"Parameters: {n_params/1e6:.2f}M total | {n_trainable/1e6:.2f}M trainable\")\n\n    bad_layers = inspect_model_channels(model)\n    if bad_layers:\n        print(\"\\nZERO/INVALID CHANNEL LAYERS FOUND:\")\n        for item in bad_layers:\n            print(\"  \", item)\n        raise RuntimeError(\n            \"Model contains invalid zero-channel layers -- fix architecture before training. \"\n            \"(If this happened with architecture='unetplusplus' and a tu-* encoder, switch to \"\n            \"architecture='unet' -- see module docstring for why.)\")\n    print(\"Zero-channel layers: 0 (PASS)\")\n\n    model.eval()\n    small_size = min(CFG.patch_size, 256)   # keep the dry-run cheap regardless of real patch_size\n    dummy = torch.zeros(1, CFG.in_channels, small_size, small_size, device=CFG.device)\n    capture = EncoderFeatureCapture(model)\n    try:\n        out = model(dummy)\n        print(f\"\\nEncoder feature stages (dry run at {small_size}x{small_size} for speed):\")\n        for i, f in enumerate(capture.features):\n            print(f\"  stage {i}: {tuple(f.shape)}\")\n        print(f\"\\nOutput logits shape: {tuple(out.shape)}\")\n        has_nan = torch.isnan(out).any().item()\n        has_inf = torch.isinf(out).any().item()\n        print(f\"NaN check: {'FAIL' if has_nan else 'PASS'} | Inf check: {'FAIL' if has_inf else 'PASS'}\")\n        if has_nan or has_inf:\n            raise RuntimeError(\"Model produced NaN/Inf on a dry-run forward pass -- fix before training.\")\n    finally:\n        capture.remove()\n    model.train()\n\n    if torch.cuda.is_available():\n        print(f\"\\nCUDA device: {torch.cuda.get_device_name(0)} | \"\n              f\"capability: {torch.cuda.get_device_capability(0)}\")\n    print(\"Forward pass: PASS\")\n\n    report_batchnorm_locations(model)\n    print(\"=\" * 70 + \"\\n\")\n\n\n# ============================================================\n# 12. LOSSES (unchanged)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\nclass V2ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        return CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n\n\n# ============================================================\n# 13. METRICS\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 14. BUILD TRAIN DATA\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING TRAIN DATA\")\nprint(\"=\" * 70)\n\ntrain_volumes = {}\ntrain_labels_full = {}\ntrain_masks_full = {}\ntrain_samples_raw = []\nval_samples = []\n\nfor fid in CFG.train_frags:\n    print(f\"\\nProcessing fragment {fid} ...\")\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        vol.compute_fragment_stats(mask, CFG.patch_size)\n        print(f\"  fragment {fid} normalization stats: mean={vol.frag_mean:.2f} std={vol.frag_std:.2f}\")\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {len(coords)} candidate patches (post mask-erosion)\")\n\n    tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n    print(f\"  spatial train={len(tr_coords)} validation={len(va_coords)}\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    train_masks_full[fid] = mask\n    train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n    val_samples.extend([(fid, y, x) for y, x in va_coords])\n    del mask\n    cleanup_memory()\n\nprint(\"\\nRaw train samples:\", len(train_samples_raw))\nprint(\"Validation samples:\", len(val_samples))\n\ntrain_samples = balance_positive_patches(\n    train_samples_raw, train_labels_full, CFG.patch_size,\n    positive_threshold=CFG.positive_patch_fraction,\n    target_positive_ratio=CFG.target_positive_patch_ratio,\n    max_positive_repeat=CFG.max_positive_repeat,\n)\n\n\n# ============================================================\n# 15. TEST-FRAGMENT MASK/VOLUME LOADED EARLY (unlabeled use only)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\")\nprint(\"=\" * 70)\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\nif CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    print(f\"  fragment {CFG.test_frag} normalization stats: \"\n          f\"mean={test_vol.frag_mean:.2f} std={test_vol.frag_std:.2f}\")\n\ntest_coords_for_unlabeled_use = generate_grid_coords(\n    test_mask, CFG.patch_size, CFG.test_stride, CFG.min_tissue_frac_train)\nprint(f\"Fragment {CFG.test_frag}: {len(test_coords_for_unlabeled_use)} unlabeled candidate patches\")\n\nhist_match_pool = None\nif CFG.USE_HIST_MATCH_AUG:\n    pool_coords = random.sample(test_coords_for_unlabeled_use,\n                                 min(CFG.hist_match_pool_size, len(test_coords_for_unlabeled_use)))\n    hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n    print(f\"Built histogram-matching reference pool: {len(hist_match_pool)} patches \"\n          f\"({sum(p.nbytes for p in hist_match_pool)/1e6:.1f} MB)\")\n\ndann_loader = None\nif CFG.USE_DANN:\n    dann_ds = UnlabeledPatchDataset(test_vol, test_coords_for_unlabeled_use, CFG.patch_size)\n    dann_loader = DataLoader(dann_ds, batch_size=CFG.dann_target_batch_size, shuffle=True,\n                              num_workers=max(1, CFG.num_workers - 1), pin_memory=(CFG.device == \"cuda\"),\n                              drop_last=True, persistent_workers=True)\n\n    def infinite_dann_loader():\n        while True:\n            for batch in dann_loader:\n                yield batch\n    dann_iter = infinite_dann_loader()\n\n\n# ============================================================\n# 16. DATASETS / LOADERS\n# ============================================================\n\ntrain_transform = build_train_transform()\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                            transform=train_transform, jitter=CFG.train_jitter,\n                            hist_match_pool=hist_match_pool, train_mode=True)\n\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                          transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0)\n\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         drop_last=False, persistent_workers=CFG.num_workers > 0)\n\n\n# ============================================================\n# 17. MODEL / ARCHITECTURE REPORT / LOSS / OPTIMIZER\n# ============================================================\n\nprint(\"\\nBuilding V4 depth-aware model ...\")\nmodel = build_model().to(CFG.device)\n\nrun_architecture_report(model)   # <-- catches zero-channel / NaN / shape problems HERE\n\nprint(\"\\nEstimating positive-pixel fraction ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\nprint(f\"Estimated positive fraction: {pos_frac:.6f}\")\n\nwith torch.no_grad():\n    bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n    try:\n        model.seg_model.segmentation_head[0].bias.fill_(bias_val)\n    except Exception as e:\n        print(f\"  (could not set output bias directly: {e})\")\nprint(f\"Output bias initialized to {bias_val:.4f}\")\n\nraw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\npos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\npos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\nprint(f\"BCE positive weight: {pos_weight_val:.3f}\")\n\ncriterion = V2ComboLoss(pos_weight=pos_weight)\n\nencoder_params, decoder_params, stem_params = [], [], []\nfor name, param in model.named_parameters():\n    if not param.requires_grad:\n        continue\n    if name.startswith(\"depth_stem.\"):\n        stem_params.append(param)\n    elif name.startswith(\"seg_model.encoder.\"):\n        encoder_params.append(param)\n    else:\n        decoder_params.append(param)\nprint(f\"Depth-stem parameters: {len(stem_params)} | Encoder parameters: {len(encoder_params)} | \"\n      f\"Decoder/head parameters: {len(decoder_params)} tensors\")\n\nparam_groups = [\n    {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n    {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n]\nif stem_params:\n    param_groups.append({\"params\": stem_params, \"lr\": CFG.depth_stem_lr})\n\ndomain_classifier = None\nfeature_capture = None\nif CFG.USE_DANN:\n    feature_capture = EncoderFeatureCapture(model)\n    with torch.no_grad():\n        dummy = torch.zeros(1, CFG.in_channels, CFG.patch_size, CFG.patch_size, device=CFG.device)\n        model.eval()\n        _ = model(dummy)\n        deepest_ch = feature_capture.features[-1].shape[1]\n        model.train()\n    domain_classifier = DomainClassifier(deepest_ch).to(CFG.device)\n    param_groups.append({\"params\": domain_classifier.parameters(), \"lr\": CFG.dann_lr})\n    print(f\"[DANN] domain classifier attached on {deepest_ch}-channel bottleneck features\")\n    del dummy\n    cleanup_memory()\n\noptimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nema = EMAModel(model, decay=CFG.ema_decay) if CFG.USE_EMA else None\n\n\n# ============================================================\n# 18. TRAIN / VALIDATION EPOCH\n# ============================================================\n\n_global_step = 0\n_total_steps = CFG.epochs * max(1, len(train_loader) // CFG.accumulation_steps)\n\n\ndef dann_lambda_schedule():\n    progress = min(_global_step / max(_total_steps, 1), 1.0)\n    return CFG.dann_lambda_max * (2.0 / (1.0 + math.exp(-10.0 * progress)) - 1.0)\n\n\ndef run_epoch(loader, train_mode=True, threshold=0.5):\n    global _global_step\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_dann_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    for batch_idx, (imgs, masks) in enumerate(loader):\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        if train_mode and CFG.USE_MIXSTYLE:\n            imgs = mixstyle_batch(imgs)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n                dann_loss_val = 0.0\n\n                if train_mode and CFG.USE_DANN:\n                    src_feat = feature_capture.features[-1]\n                    lambd = dann_lambda_schedule()\n                    src_domain_logits = domain_classifier(src_feat, lambd)\n                    src_domain_target = torch.zeros_like(src_domain_logits)\n\n                    tgt_imgs = next(dann_iter).to(CFG.device, non_blocking=True)\n                    _ = model(tgt_imgs)\n                    tgt_feat = feature_capture.features[-1]\n                    tgt_domain_logits = domain_classifier(tgt_feat, lambd)\n                    tgt_domain_target = torch.ones_like(tgt_domain_logits)\n\n                    dann_loss = 0.5 * (\n                        F.binary_cross_entropy_with_logits(src_domain_logits, src_domain_target) +\n                        F.binary_cross_entropy_with_logits(tgt_domain_logits, tgt_domain_target)\n                    )\n                    dann_loss_val = dann_loss.item()\n                    loss = loss + dann_loss\n\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    params_to_clip = list(model.parameters())\n                    if domain_classifier is not None:\n                        params_to_clip += list(domain_classifier.parameters())\n                    torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    _global_step += 1\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_dann_loss += dann_loss_val\n\n        del imgs, masks, logits, probs\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        params_to_clip = list(model.parameters())\n        if domain_classifier is not None:\n            params_to_clip += list(domain_classifier.parameters())\n        torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_dann_loss = total_dann_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_dann_loss\n\n\n# ============================================================\n# 19. TRAINING LOOP\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING V4 TRAINING\")\nprint(\"=\" * 70)\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"train_dice\": [], \"val_dice\": [],\n           \"val_iou\": [], \"val_precision\": [], \"val_recall\": [], \"dann_loss\": []}\n\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_metrics, dann_loss_avg = run_epoch(train_loader, train_mode=True, threshold=0.50)\n\n    if ema is not None:\n        backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        model.load_state_dict(ema.state_dict())\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n        model.load_state_dict(backup)\n        del backup\n    else:\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"train_dice\"].append(train_metrics[\"dice\"])\n    history[\"val_dice\"].append(val_metrics[\"dice\"])\n    history[\"val_iou\"].append(val_metrics[\"iou\"])\n    history[\"val_precision\"].append(val_metrics[\"precision\"])\n    history[\"val_recall\"].append(val_metrics[\"recall\"])\n    history[\"dann_loss\"].append(dann_loss_avg)\n\n    print(f\"\\n[{epoch:02d}/{CFG.epochs}] time={time.time()-t0:.1f}s\")\n    print(f\"train_loss={train_loss:.5f} train_dice={train_metrics['dice']:.5f} \"\n          f\"dann_loss={dann_loss_avg:.5f}\")\n    print(f\"val_loss={val_loss:.5f} val_dice={val_metrics['dice']:.5f} val_iou={val_metrics['iou']:.5f}\")\n    print(f\"precision={val_metrics['precision']:.5f} recall={val_metrics['recall']:.5f}\")\n    print(f\"encoder_lr={optimizer.param_groups[0]['lr']:.7f} decoder_lr={optimizer.param_groups[1]['lr']:.7f}\")\n\n    if val_metrics[\"dice\"] > best_val_dice:\n        best_val_dice = val_metrics[\"dice\"]\n        epochs_no_improve = 0\n        save_state = ema.state_dict() if ema is not None else model.state_dict()\n        checkpoint = {\"model\": save_state, \"cfg\": cfg_to_dict(CFG), \"best_val_dice\": best_val_dice,\n                      \"history\": history, \"pos_frac\": pos_frac, \"pos_weight\": pos_weight_val,\n                      \"used_ema\": CFG.USE_EMA}\n        torch.save(checkpoint, CFG.ckpt_path)\n        print(f\"*** NEW BEST CHECKPOINT val_dice={best_val_dice:.5f} \"\n              f\"({'EMA' if CFG.USE_EMA else 'raw'} weights) ***\")\n    else:\n        epochs_no_improve += 1\n        print(f\"No improvement: {epochs_no_improve}/{CFG.early_stop_patience}\")\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(\"Early stopping.\")\n            break\n\n    cleanup_memory()\n\nprint(\"\\nBest validation Dice:\", best_val_dice)\n\nif feature_capture is not None:\n    feature_capture.remove()\n\n\n# ============================================================\n# 20. LOAD BEST MODEL\n# ============================================================\n\ncheckpoint = torch.load(CFG.ckpt_path, map_location=CFG.device)\nmodel.load_state_dict(checkpoint[\"model\"])\nprint(f\"\\nBest checkpoint loaded ({'EMA' if checkpoint.get('used_ema') else 'raw'} weights).\")\n\n\n# ============================================================\n# 21. FINE DICE THRESHOLD SEARCH\n# ============================================================\n\n@torch.no_grad()\ndef find_best_dice_threshold(model, loader):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice_scores = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    best_idx = int(np.argmax(dice_scores))\n    return float(thresholds[best_idx]), float(dice_scores[best_idx]), thresholds, dice_scores\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINE DICE THRESHOLD SEARCH\")\nprint(\"=\" * 70)\n\nbest_threshold, threshold_dice, threshold_grid, threshold_scores = find_best_dice_threshold(model, val_loader)\nprint(f\"BEST VALIDATION THRESHOLD = {best_threshold:.2f}\")\nprint(f\"DICE AT BEST THRESHOLD = {threshold_dice:.5f}\")\n\n_, final_val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal validation metrics at optimized Dice threshold:\", final_val_metrics)\n\n\n# ============================================================\n# 22. INFERENCE HELPERS\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference_v2(model, vol, mask, patch_size, stride, device, batch_size,\n                                 pre_transform=None):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw_u8 = vol.read_patch(y, x, patch_size)\n        if pre_transform is not None:\n            raw_u8 = pre_transform(raw_u8)\n        raw = raw_u8.astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm_v2(model, vol, mask, patch_size, stride, device, max_patches, batch_size):\n    print(\"\\nStarting AdaBN...\")\n    n_reset = 0\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n            n_reset += 1\n    print(f\"Reset {n_reset} BatchNorm2d layers (see the architecture report above for where they live).\")\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"AdaBN patches: {len(coords)}\")\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        batch = coords[start:start + batch_size]\n        imgs = []\n        for y, x in batch:\n            raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n            raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n            imgs.append(raw)\n        inp = torch.from_numpy(np.stack(imgs)).to(device, non_blocking=True)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp)\n        del inp, imgs\n    model.eval()\n    cleanup_memory()\n    print(\"AdaBN finished.\")\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess_v2(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 23. LOAD TEST LABELS (local diagnostics only, loaded late)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST LABELS (diagnostics only)\")\nprint(\"=\" * 70)\n\ntest_labels = load_ink_labels(test_dir)\nif test_labels is not None:\n    print(\"Test GT found: local diagnostic evaluation enabled.\")\n    gt_test = (test_labels * test_mask).astype(np.float32)\nelse:\n    print(\"No test GT found: running competition-style inference.\")\n    gt_test = None\n\n\n# ============================================================\n# 24. INFERENCE A: BASELINE + TTA\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE A: BASELINE + TTA\")\nprint(\"=\" * 70)\n\ntest_prob_baseline = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n\ntest_metrics_baseline = None\nif test_labels is not None:\n    test_metrics_baseline = evaluate_probability_map(test_prob_baseline, gt_test, best_threshold)\n    print(\"\\nBASELINE TEST METRICS:\", test_metrics_baseline)\n\n\n# ============================================================\n# 25. INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\")\nprint(\"=\" * 70)\n\ntrain_hist_pool = []\nfor fid in CFG.train_frags:\n    v = train_volumes[fid]\n    m = train_masks_full[fid]\n    coords = generate_grid_coords(m, CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n    if coords:\n        for (y, x) in random.sample(coords, min(10, len(coords))):\n            train_hist_pool.append(v.read_patch(y, x, CFG.patch_size))\nprint(f\"Built train-domain reference pool for diagnostic: {len(train_hist_pool)} patches\")\n\n\ndef histogram_match_to_train_domain(raw_u8_dhw):\n    if not train_hist_pool:\n        return raw_u8_dhw\n    ref = random.choice(train_hist_pool)\n    img_hwd = np.transpose(raw_u8_dhw, (1, 2, 0))\n    ref_hwd = np.transpose(ref, (1, 2, 0))\n    try:\n        matched = match_histograms(img_hwd, ref_hwd, channel_axis=None).astype(np.uint8)\n        return np.transpose(matched, (2, 0, 1))\n    except Exception:\n        return raw_u8_dhw\n\n\ntest_prob_histmatch = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch,\n    pre_transform=histogram_match_to_train_domain)\n\ntest_metrics_histmatch = None\nif test_labels is not None:\n    test_metrics_histmatch = evaluate_probability_map(test_prob_histmatch, gt_test, best_threshold)\n    print(\"\\nHISTOGRAM-MATCHING DIAGNOSTIC TEST METRICS:\", test_metrics_histmatch)\n    if test_metrics_baseline is not None:\n        delta = test_metrics_histmatch[\"dice\"] - test_metrics_baseline[\"dice\"]\n        print(f\"\\n>>> Histogram matching alone changed local test Dice by {delta:+.4f} \"\n              f\"(baseline {test_metrics_baseline['dice']:.4f} -> {test_metrics_histmatch['dice']:.4f})\")\n\n\n# ============================================================\n# 26. INFERENCE C: ADABN + TTA (only if enabled -- see architecture report)\n# ============================================================\n\ntest_prob_adabn = None\ntest_metrics_adabn = None\n\nif CFG.use_adabn:\n    print(\"\\n\" + \"=\" * 70)\n    print(\"INFERENCE C: ADABN + TTA\")\n    print(\"=\" * 70)\n    model = recalibrate_batchnorm_v2(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                      CFG.device, CFG.adabn_max_patches, CFG.infer_batch)\n    test_prob_adabn = sliding_window_inference_v2(\n        model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n    if test_labels is not None:\n        test_metrics_adabn = evaluate_probability_map(test_prob_adabn, gt_test, best_threshold)\n        print(\"\\nADABN TEST METRICS:\", test_metrics_adabn)\nelse:\n    print(\"\\n(Skipping AdaBN inference -- CFG.use_adabn=False. See the architecture report's \"\n          \"BatchNorm2d location summary above for why/whether it would help.)\")\n\n\n# ============================================================\n# 27. CHOOSE FINAL PROBABILITY MAP\n# ============================================================\n\ntest_prob = test_prob_adabn if (CFG.use_adabn and test_prob_adabn is not None) else test_prob_baseline\n\notsu_threshold = compute_otsu_threshold(test_prob, test_mask, fallback=best_threshold)\nprint(\"\\nValidation-tuned threshold:\", best_threshold)\nprint(\"Unsupervised Otsu threshold:\", otsu_threshold)\n\nfinal_threshold = best_threshold\ntest_pred_bin = postprocess_v2(test_prob, final_threshold)\n\ntest_metrics_raw = test_metrics_otsu = test_metrics_post = None\nif test_labels is not None:\n    test_metrics_raw = evaluate_probability_map(test_prob, gt_test, final_threshold)\n    test_metrics_otsu = evaluate_probability_map(test_prob, gt_test, otsu_threshold)\n    post_preds = test_pred_bin.astype(np.float32)\n    tp = (post_preds * gt_test).sum()\n    fp = (post_preds * (1.0 - gt_test)).sum()\n    fn = ((1.0 - post_preds) * gt_test).sum()\n    test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LOCAL TEST DIAGNOSTICS SUMMARY\")\n    print(\"=\" * 70)\n    print(\"A) Baseline + val threshold:            \", test_metrics_baseline)\n    print(\"B) Histogram-matched (no retrain) + val threshold:\", test_metrics_histmatch)\n    print(\"C) AdaBN + val threshold:                \", test_metrics_adabn)\n    print(\"Final (chosen) raw + val threshold:      \", test_metrics_raw)\n    print(\"Final (chosen) raw + Otsu threshold:      \", test_metrics_otsu)\n    print(\"Final (chosen) postprocessed + val threshold:\", test_metrics_post)\n\n\n# ============================================================\n# 28. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_v4.npy\")\nnp.save(prob_path, test_prob)\nprint(\"\\nSaved probability map:\", prob_path)\n\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_v4.png\")\ncv2.imwrite(pred_path, (test_pred_bin * 255).astype(np.uint8))\nprint(\"Saved prediction:\", pred_path)\n\nmetrics_summary = {\n    \"best_validation_dice_at_0.50\": best_val_dice,\n    \"best_validation_threshold\": best_threshold,\n    \"validation_dice_at_best_threshold\": threshold_dice,\n    \"otsu_threshold\": otsu_threshold,\n    \"final_threshold\": final_threshold,\n    \"final_validation_metrics\": final_val_metrics,\n    \"test_baseline\": test_metrics_baseline,\n    \"test_histogram_matched_diagnostic_no_retrain\": test_metrics_histmatch,\n    \"test_adabn\": test_metrics_adabn,\n    \"test_raw\": test_metrics_raw,\n    \"test_otsu\": test_metrics_otsu,\n    \"test_postprocessed\": test_metrics_post,\n    \"config\": cfg_to_dict(CFG),\n    \"history\": history,\n}\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump(metrics_summary, f, indent=2)\nprint(\"Saved metrics:\", CFG.metrics_path)\n\n\n# ============================================================\n# 29. VISUALIZATION\n# ============================================================\n\ndef save_full_overview():\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    prob_small = cv2.resize((test_prob * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_small = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1],\n                               interpolation=cv2.INTER_NEAREST)\n        fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n        axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability thr={final_threshold:.2f}\")\n        axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"Final Prediction\")\n    else:\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(prob_small, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n        axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"Final Prediction\")\n\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    overview_path = os.path.join(CFG.viz_dir, \"fragment1_v4_overview.png\")\n    plt.savefig(overview_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(\"Saved overview:\", overview_path)\n\n\nsave_full_overview()\n\n\ndef save_patch_comparisons(n=6):\n    coords = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(coords)\n    coords = coords[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(coords):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n        prob_patch = test_prob[y:y + size, x:x + size]\n\n        if test_labels is not None:\n            gt_patch = test_labels[y:y + size, x:x + size]\n            fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n            axes[2].imshow(prob_patch, cmap=\"gray\"); axes[2].set_title(\"Probability\")\n            axes[3].imshow(pred_patch, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n        else:\n            fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(prob_patch, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n            axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n\n        for ax in axes:\n            ax.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n        plt.close(fig)\n    print(f\"Saved {len(coords)} patch comparisons.\")\n\n\nsave_patch_comparisons(n=6)\n\n\ndef save_training_curves():\n    epochs_axis = np.arange(1, len(history[\"train_loss\"]) + 1)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_loss\"], label=\"Train Loss\")\n    plt.plot(epochs_axis, history[\"val_loss\"], label=\"Val Loss\")\n    if CFG.USE_DANN:\n        plt.plot(epochs_axis, history[\"dann_loss\"], label=\"DANN domain loss\", linestyle=\"--\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Loss\"); plt.title(\"V4 Training Loss\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_loss.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_dice\"], label=\"Train Dice\")\n    plt.plot(epochs_axis, history[\"val_dice\"], label=\"Val Dice (EMA)\" if CFG.USE_EMA else \"Val Dice\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Dice\"); plt.title(\"V4 Dice\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_dice.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(threshold_grid, threshold_scores)\n    plt.axvline(best_threshold, linestyle=\"--\", label=f\"best={best_threshold:.2f}\")\n    plt.xlabel(\"Threshold\"); plt.ylabel(\"Validation Dice\"); plt.title(\"Dice Threshold Search\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"threshold_search.png\"), dpi=150)\n    plt.close(fig)\n\n    print(\"Saved training curves.\")\n\n\nsave_training_curves()\n\n\n# ============================================================\n# 30. CLEANUP\n# ============================================================\n\ntest_vol.close()\nfor v in train_volumes.values():\n    v.close()\ncleanup_memory()\n\n\n# ============================================================\n# 31. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V4 COMPLETE\")\nprint(\"=\" * 70)\nprint(f\"Best validation Dice @ 0.50: {best_val_dice:.5f}\")\nprint(f\"Best Dice threshold: {best_threshold:.2f}\")\nprint(f\"Validation Dice @ optimized threshold: {threshold_dice:.5f}\")\nprint(f\"Backbone: {CFG.encoder_name} / architecture={CFG.architecture} / \"\n      f\"depth_stem={CFG.USE_DEPTH_FUSION_STEM}\")\nprint(f\"EMA: {CFG.USE_EMA} | DANN: {CFG.USE_DANN} | MixStyle: {CFG.USE_MIXSTYLE} | AdaBN: {CFG.use_adabn}\")\nprint(f\"CLAHE mode: {CFG.CLAHE_MODE} | Normalization: {CFG.NORMALIZATION_MODE} | \"\n      f\"Mask erosion: {CFG.mask_erode_px}px\")\nif test_metrics_baseline is not None:\n    print(f\"\\nLocal test Dice -- baseline: {test_metrics_baseline['dice']:.5f}\")\nif test_metrics_histmatch is not None:\n    print(f\"Local test Dice -- histogram-matched (diagnostic, no retrain): \"\n          f\"{test_metrics_histmatch['dice']:.5f}\")\nif test_metrics_adabn is not None:\n    print(f\"Local test Dice -- AdaBN: {test_metrics_adabn['dice']:.5f}\")\nif test_metrics_post is not None:\n    print(f\"Local test Dice -- final postprocessed: {test_metrics_post['dice']:.5f}\")\nprint(f\"\\nBest checkpoint: {CFG.ckpt_path}\")\nprint(f\"Probability map: {prob_path}\")\nprint(f\"Prediction: {pred_path}\")\nprint(f\"Metrics: {CFG.metrics_path}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")\nprint(\"\\n=== V4 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T07:21:53.932641Z","iopub.execute_input":"2026-09-08T07:21:53.933451Z","iopub.status.idle":"2026-09-08T09:32:55.661987Z","shell.execute_reply.started":"2026-09-08T07:21:53.933416Z","shell.execute_reply":"2026-09-08T09:32:55.661126Z"}},"outputs":[{"name":"stdout","text":"Collecting segmentation-models-pytorch==0.4.0\n  Downloading segmentation_models_pytorch-0.4.0-py3-none-any.whl.metadata (32 kB)\nCollecting efficientnet-pytorch>=0.6.1 (from segmentation-models-pytorch==0.4.0)\n  Downloading efficientnet_pytorch-0.7.1.tar.gz (21 kB)\n  Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\nRequirement already satisfied: huggingface-hub>=0.24 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.11.0)\nRequirement already satisfied: numpy>=1.19.3 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (2.0.2)\nRequirement already satisfied: pillow>=8 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (11.3.0)\nCollecting pretrainedmodels>=0.7.1 (from segmentation-models-pytorch==0.4.0)\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m3.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\nRequirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.17.0)\nRequirement already satisfied: timm>=0.9 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (1.0.26)\nRequirement already satisfied: torch>=1.8 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (2.10.0+cu128)\nRequirement already satisfied: torchvision>=0.9 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (0.25.0+cu128)\nRequirement already satisfied: tqdm>=4.42.1 in /usr/local/lib/python3.12/dist-packages (from segmentation-models-pytorch==0.4.0) (4.67.3)\nRequirement already satisfied: filelock>=3.10.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (3.29.0)\nRequirement already satisfied: fsspec>=2023.5.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2025.3.0)\nRequirement already satisfied: hf-xet<2.0.0,>=1.4.3 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.4.3)\nRequirement already satisfied: httpx<1,>=0.23.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.28.1)\nRequirement already satisfied: packaging>=20.9 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (26.1)\nRequirement already satisfied: pyyaml>=5.1 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (6.0.3)\nRequirement already satisfied: typer in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.24.2)\nRequirement already satisfied: typing-extensions>=4.1.0 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.15.0)\nCollecting munch (from pretrainedmodels>=0.7.1->segmentation-models-pytorch==0.4.0)\n  Downloading munch-4.0.0-py2.py3-none-any.whl.metadata (5.9 kB)\nRequirement already satisfied: safetensors in /usr/local/lib/python3.12/dist-packages (from timm>=0.9->segmentation-models-pytorch==0.4.0) (0.7.0)\nRequirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (81.0.0)\nRequirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (1.14.0)\nRequirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.6.1)\nRequirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.1.6)\nRequirement already satisfied: cuda-bindings==12.9.4 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.9.4)\nRequirement already satisfied: nvidia-cuda-nvrtc-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.93)\nRequirement already satisfied: nvidia-cuda-runtime-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-cuda-cupti-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (9.10.2.21)\nRequirement already satisfied: nvidia-cublas-cu12==12.8.4.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.4.1)\nRequirement already satisfied: nvidia-cufft-cu12==11.3.3.83 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (11.3.3.83)\nRequirement already satisfied: nvidia-curand-cu12==10.3.9.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (10.3.9.90)\nRequirement already satisfied: nvidia-cusolver-cu12==11.7.3.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (11.7.3.90)\nRequirement already satisfied: nvidia-cusparse-cu12==12.5.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.5.8.93)\nRequirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (0.7.1)\nRequirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (2.27.5)\nRequirement already satisfied: nvidia-nvshmem-cu12==3.4.5 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.4.5)\nRequirement already satisfied: nvidia-nvtx-cu12==12.8.90 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.90)\nRequirement already satisfied: nvidia-nvjitlink-cu12==12.8.93 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (12.8.93)\nRequirement already satisfied: nvidia-cufile-cu12==1.13.1.3 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (1.13.1.3)\nRequirement already satisfied: triton==3.6.0 in /usr/local/lib/python3.12/dist-packages (from torch>=1.8->segmentation-models-pytorch==0.4.0) (3.6.0)\nRequirement already satisfied: cuda-pathfinder~=1.1 in /usr/local/lib/python3.12/dist-packages (from cuda-bindings==12.9.4->torch>=1.8->segmentation-models-pytorch==0.4.0) (1.5.3)\nRequirement already satisfied: anyio in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.13.0)\nRequirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2026.4.22)\nRequirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.0.9)\nRequirement already satisfied: idna in /usr/local/lib/python3.12/dist-packages (from httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (3.13)\nRequirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1,>=0.23.0->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.16.0)\nRequirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch>=1.8->segmentation-models-pytorch==0.4.0) (1.3.0)\nRequirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch>=1.8->segmentation-models-pytorch==0.4.0) (3.0.3)\nRequirement already satisfied: click>=8.2.1 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (8.3.3)\nRequirement already satisfied: shellingham>=1.3.0 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (1.5.4)\nRequirement already satisfied: rich>=12.3.0 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (13.9.4)\nRequirement already satisfied: annotated-doc>=0.0.2 in /usr/local/lib/python3.12/dist-packages (from typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.0.4)\nRequirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (4.0.0)\nRequirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (2.20.0)\nRequirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich>=12.3.0->typer->huggingface-hub>=0.24->segmentation-models-pytorch==0.4.0) (0.1.2)\nDownloading segmentation_models_pytorch-0.4.0-py3-none-any.whl (121 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m121.3/121.3 kB\u001b[0m \u001b[31m3.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading munch-4.0.0-py2.py3-none-any.whl (9.9 kB)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25l\u001b[?25hdone\n  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.7.1-py3-none-any.whl size=16477 sha256=de07c83ef5f3403120932add3e34afd19c84c2ea428e5402d3e5297c1e2e6eed\n  Stored in directory: /root/.cache/pip/wheels/9c/3f/43/e6271c7026fe08c185da2be23c98c8e87477d3db63f41f32ad\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25l\u001b[?25hdone\n  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60990 sha256=af4361146e7ef69048187e9296fe23dd8c70c02674b8a17f18d7a71167ec4075\n  Stored in directory: /root/.cache/pip/wheels/4c/01/56/40a48f75dbdfe167a0cb70d3b48913369a00ec5c4e9fed5f2b\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: munch, efficientnet-pytorch, pretrainedmodels, segmentation-models-pytorch\nSuccessfully installed efficientnet-pytorch-0.7.1 munch-4.0.0 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.4.0\n     ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 43.7/43.7 kB 1.0 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 154.8/154.8 kB 3.4 MB/s eta 0:00:00\n   ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 2.6/2.6 MB 26.5 MB/s eta 0:00:00\n[env] torch=2.10.0+cu128 | smp=0.5.0 | timm=1.0.29\n======================================================================\nBUILDING TRAIN DATA\n======================================================================\n\nProcessing fragment 2 ...\n  fragment 2 normalization stats: mean=118.25 std=59.13\nFragment 2: 24581 candidate patches (post mask-erosion)\n  spatial train=18331 validation=4954\n\nProcessing fragment 3 ...\n  fragment 3 normalization stats: mean=119.91 std=61.59\nFragment 3: 6434 candidate patches (post mask-erosion)\n  spatial train=5126 validation=623\n\nRaw train samples: 23457\nValidation samples: 5577\nPositive patches: 15957 | Negative patches: 7500\nBalanced dataset: 23457 | positive ratio=0.680\n\n======================================================================\nLOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\n======================================================================\n  fragment 1 normalization stats: mean=112.33 std=69.54\nFragment 1: 7514 unlabeled candidate patches\nBuilt histogram-matching reference pool: 20 patches (53.2 MB)\n\nBuilding V4 depth-aware model ...\n[backbone] DepthFusionStem: 26 depth slices (raw=True, grad=True, curv=True) -> 16 learned channels\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/114M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"5d85657b52fc472da41f249da5cf9afc"}},"metadata":{}},{"name":"stdout","text":"[backbone] using tu-convnext_tiny (ImageNet-pretrained) + unet\n\n======================================================================\nV4 ARCHITECTURE INSPECTION (runs BEFORE the optimizer/training loop)\n======================================================================\nInput:  depth slices=26 | patch=320x320 | input tensor=[1, 26, 320, 320]\nEncoder: backbone=tu-convnext_tiny | pretrained=imagenet | architecture=unet\nDepth stem: 26 -> 16 channels (raw=True, grad=True, curv=True)\nDecoder channels: (256, 128, 64, 32, 16)\nParameters: 32.16M total | 32.16M trainable\nZero-channel layers: 0 (PASS)\n\nEncoder feature stages (dry run at 256x256 for speed):\n  stage 0: (1, 16, 256, 256)\n  stage 1: (1, 0, 128, 128)\n  stage 2: (1, 96, 64, 64)\n  stage 3: (1, 192, 32, 32)\n  stage 4: (1, 384, 16, 16)\n  stage 5: (1, 768, 8, 8)\n\nOutput logits shape: (1, 1, 256, 256)\nNaN check: PASS | Inf check: PASS\n\nCUDA device: Tesla T4 | capability: (7, 5)\nForward pass: PASS\n[AdaBN report] BatchNorm2d layers found: depth_stem=1 encoder=0 decoder=10 other=0 (total=11)\n  -> the ConvNeXt encoder itself has none (it's LayerNorm-only, as expected). AdaBN would only recalibrate the depth-stem/decoder BN layers listed above.\n======================================================================\n\n\nEstimating positive-pixel fraction ...\nEstimated positive fraction: 0.137213\nOutput bias initialized to -1.8386\nBCE positive weight: 2.508\nDepth-stem parameters: 4 | Encoder parameters: 178 | Decoder/head parameters: 92 tensors\n\n======================================================================\nSTARTING V4 TRAINING\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/4177651837.py:1308: FutureWarning: `torch.cuda.amp.GradScaler(args...)` is deprecated. Please use `torch.amp.GradScaler('cuda', args...)` instead.\n  scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n/tmp/ipykernel_58/4177651837.py:1347: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"\n[01/7] time=1344.7s\ntrain_loss=0.75773 train_dice=0.42221 dann_loss=0.00000\nval_loss=0.75349 val_dice=0.40344 val_iou=0.25269\nprecision=0.63539 recall=0.29555\nencoder_lr=0.0000476 decoder_lr=0.0000951\n*** NEW BEST CHECKPOINT val_dice=0.40344 (EMA weights) ***\n\n[02/7] time=1139.5s\ntrain_loss=0.69015 train_dice=0.53639 dann_loss=0.00000\nval_loss=0.68151 val_dice=0.55774 val_iou=0.38671\nprecision=0.49880 recall=0.63247\nencoder_lr=0.0000408 decoder_lr=0.0000814\n*** NEW BEST CHECKPOINT val_dice=0.55774 (EMA weights) ***\n\n[03/7] time=1142.0s\ntrain_loss=0.63139 train_dice=0.62008 dann_loss=0.00000\nval_loss=0.67943 val_dice=0.55522 val_iou=0.38430\nprecision=0.49011 recall=0.64029\nencoder_lr=0.0000310 decoder_lr=0.0000615\nNo improvement: 1/3\n\n[04/7] time=1144.1s\ntrain_loss=0.57665 train_dice=0.69153 dann_loss=0.00000\nval_loss=0.71359 val_dice=0.55444 val_iou=0.38355\nprecision=0.51626 recall=0.59872\nencoder_lr=0.0000200 decoder_lr=0.0000395\nNo improvement: 2/3\n\n[05/7] time=1151.0s\ntrain_loss=0.52720 train_dice=0.75040 dann_loss=0.00000\nval_loss=0.73964 val_dice=0.54610 val_iou=0.37561\nprecision=0.50665 recall=0.59222\nencoder_lr=0.0000102 decoder_lr=0.0000196\nNo improvement: 3/3\nEarly stopping.\n\nBest validation Dice: 0.5577376092620798\n\nBest checkpoint loaded (EMA weights).\n\n======================================================================\nFINE DICE THRESHOLD SEARCH\n======================================================================\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/4177651837.py:1503: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"BEST VALIDATION THRESHOLD = 0.57\nDICE AT BEST THRESHOLD = 0.56178\n\nFinal validation metrics at optimized Dice threshold: {'dice': 0.5617841237164952, 'iou': 0.3906118218971448, 'precision': 0.5491922600818122, 'recall': 0.57496694863576, 'fbeta0.5': 0.5541612823359191}\n\n======================================================================\nLOADING TEST LABELS (diagnostics only)\n======================================================================\nTest GT found: local diagnostic evaluation enabled.\n\n======================================================================\nINFERENCE A: BASELINE + TTA\n======================================================================\nInference patches: 7705\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_58/4177651837.py:1572: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with autocast(enabled=(CFG.device == \"cuda\")):\n","output_type":"stream"},{"name":"stdout","text":"\nBASELINE TEST METRICS: {'dice': 0.49072492122650146, 'iou': 0.32513949275016785, 'precision': 0.513267457485199, 'recall': 0.4700791835784912, 'fbeta0.5': 0.5040072202682495}\n\n======================================================================\nINFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\n======================================================================\nBuilt train-domain reference pool for diagnostic: 20 patches\nInference patches: 7705\n\nHISTOGRAM-MATCHING DIAGNOSTIC TEST METRICS: {'dice': 0.48070546984672546, 'iou': 0.3164004385471344, 'precision': 0.5600818991661072, 'recall': 0.42103514075279236, 'fbeta0.5': 0.5253814458847046}\n\n>>> Histogram matching alone changed local test Dice by -0.0100 (baseline 0.4907 -> 0.4807)\n\n(Skipping AdaBN inference -- CFG.use_adabn=False. See the architecture report's BatchNorm2d location summary above for why/whether it would help.)\n\nValidation-tuned threshold: 0.5699999999999997\nUnsupervised Otsu threshold: 0.39700546860694885\n\n======================================================================\nLOCAL TEST DIAGNOSTICS SUMMARY\n======================================================================\nA) Baseline + val threshold:             {'dice': 0.49072492122650146, 'iou': 0.32513949275016785, 'precision': 0.513267457485199, 'recall': 0.4700791835784912, 'fbeta0.5': 0.5040072202682495}\nB) Histogram-matched (no retrain) + val threshold: {'dice': 0.48070546984672546, 'iou': 0.3164004385471344, 'precision': 0.5600818991661072, 'recall': 0.42103514075279236, 'fbeta0.5': 0.5253814458847046}\nC) AdaBN + val threshold:                 None\nFinal (chosen) raw + val threshold:       {'dice': 0.49072492122650146, 'iou': 0.32513949275016785, 'precision': 0.513267457485199, 'recall': 0.4700791835784912, 'fbeta0.5': 0.5040072202682495}\nFinal (chosen) raw + Otsu threshold:       {'dice': 0.4996477961540222, 'iou': 0.3330203592777252, 'precision': 0.4211834669113159, 'recall': 0.6140403747558594, 'fbeta0.5': 0.4494144916534424}\nFinal (chosen) postprocessed + val threshold: {'dice': 0.49086225032806396, 'iou': 0.32526007294654846, 'precision': 0.512933075428009, 'recall': 0.47061240673065186, 'fbeta0.5': 0.5038716197013855}\n\nSaved probability map: /kaggle/working/fragment1_probability_v4.npy\nSaved prediction: /kaggle/working/fragment1_prediction_v4.png\nSaved metrics: /kaggle/working/v4_metrics_summary.json\nSaved overview: /kaggle/working/v4_visualizations/fragment1_v4_overview.png\nSaved 6 patch comparisons.\nSaved training curves.\n\n======================================================================\nV4 COMPLETE\n======================================================================\nBest validation Dice @ 0.50: 0.55774\nBest Dice threshold: 0.57\nValidation Dice @ optimized threshold: 0.56178\nBackbone: tu-convnext_tiny / architecture=unet / depth_stem=True\nEMA: True | DANN: False | MixStyle: True | AdaBN: False\nCLAHE mode: global_shared | Normalization: per_fragment_zscore | Mask erosion: 24px\n\nLocal test Dice -- baseline: 0.49072\nLocal test Dice -- histogram-matched (diagnostic, no retrain): 0.48071\nLocal test Dice -- final postprocessed: 0.49086\n\nBest checkpoint: /kaggle/working/vesuviusnet_v4_best.pth\nProbability map: /kaggle/working/fragment1_probability_v4.npy\nPrediction: /kaggle/working/fragment1_prediction_v4.png\nMetrics: /kaggle/working/v4_metrics_summary.json\nVisualizations: /kaggle/working/v4_visualizations\n\n=== V4 DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"#Local test Dice -- histogram-matched: 0.48377\n#patch_size = 224, train_stride = 96, test_stride = 96, batch_size = 8,epochs = 10\n#Local test Dice -- histogram-matched: 0.49977\n#patch_size = 224, train_stride = 64, test_stride = 64, batch_size = 8,epochs = 10\n\n# ============================================================\n# VESUVIUS INK DETECTION - V3 \"DOMAIN-ROBUST\"\n# ============================================================\n# Built on your V2 pipeline. Adds two things you asked for:\n#\n#   PART A: a cheap, no-retraining histogram-matching DIAGNOSTIC\n#           (runs right after baseline inference, using the\n#           already-trained model -- tells you how much of the\n#           val->test gap is pure appearance/calibration shift\n#           vs. something more structural, BEFORE you invest in\n#           heavier fixes)\n#\n#   PART B: a full domain-randomization training pipeline, so the\n#           model is pushed to stop relying on any one fragment's\n#           absolute style during training:\n#\n#     Geometric:      Rotate90 / flip (already in V2) + ElasticTransform\n#                      (the papyrus is rolled/warped -- this matters)\n#     Photometric:     brightness/contrast (already in V2) + GaussNoise\n#                      + RandomShadow-style + CoarseDropout\n#     Domain-specific: synthetic \"fake ink\" strokes and \"fake fiber\"\n#                      texture injected into the IMAGE ONLY (never the\n#                      mask) so the model can't shortcut on raw\n#                      brightness alone\n#     Appearance gap:  CLAHE per fragment before training, per-FRAGMENT\n#                      z-score (not per-patch, not global-dataset),\n#                      and a histogram-matching AUGMENTATION that\n#                      restyles some training patches toward fragment\n#                      1's own intensity distribution during training\n#     Edge artifacts:  mask erosion crops out the outer border of each\n#                      fragment (train AND test) before patches are\n#                      even generated\n#     Stronger/pretrained backbone: tries ConvNeXt-Tiny (ImageNet),\n#                      falls back to EfficientNet-B4 if unavailable;\n#                      UNet++ decoder available as an option\n#     EMA:             exponential moving average of weights, used for\n#                      validation/threshold-search/inference instead\n#                      of the raw end-of-training weights\n#     DANN:            a small domain classifier + gradient-reversal\n#                      layer, trained adversarially against the\n#                      encoder using UNLABELED fragment-1 patches\n#                      (labels are never touched -- only raw pixels)\n#                      so the encoder is pushed toward domain-invariant\n#                      features\n#     MixStyle:        mixes per-sample intensity statistics across a\n#                      training batch so the model can't memorize one\n#                      fragment's absolute \"style\"\n#\n# HONEST NOTE: this is a lot of machinery stacked at once. Each piece\n# is individually well-motivated and gated behind its own CFG flag\n# (all default ON except where noted) so you can turn things off if\n# runtime/stability becomes a problem. I'd still recommend watching\n# the PART A diagnostic output first -- if histogram matching alone\n# barely moves local test Dice, that's a signal the appearance-focused\n# techniques here (CLAHE, per-fragment norm, histogram-match aug) will\n# matter less than the structural ones (DANN, MixStyle, stronger\n# backbone, more augmentation diversity).\n#\n# RUNTIME WARNING: DANN adds an extra unlabeled forward pass per\n# training step, and there's a lot more augmentation compute per\n# sample than V2. Expect meaningfully slower epochs than V2. Turn off\n# CFG.USE_DANN first if you need to cut runtime -- it's the most\n# expensive single piece here.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nimport os as _os\n_os.system(\"pip install -q segmentation-models-pytorch==0.4.0\")\n_os.system(\"pip install -q albumentations\")\n_os.system(\"pip install -q scikit-image\")\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom scipy import ndimage\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\n\nfrom skimage.filters import threshold_otsu\nfrom skimage.exposure import match_histograms\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    depth_indices = list(range(12, 38))\n    in_channels = len(depth_indices)\n\n    patch_size = 224\n    train_stride = 64\n    test_stride = 64\n    train_jitter = 32\n\n    val_fraction = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n    positive_patch_fraction = 0.001\n\n    target_positive_patch_ratio = 0.45\n    max_positive_repeat = 2\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n    drop_last = True\n\n    epochs = 8\n    early_stop_patience = 3\n\n    encoder_lr = 5e-5\n    decoder_lr = 1e-4\n    dann_lr = 1e-4\n    weight_decay = 1e-4\n    grad_clip = 5.0\n    accumulation_steps = 1\n\n    # --- V3: backbone choice ------------------------------------------------\n    # \"auto_strong_convnext\" tries ConvNeXt-Tiny (ImageNet) via an smp/timm\n    # upgrade, falling back to EfficientNet-B4 (supported even in the pinned\n    # smp==0.2.0) if that upgrade isn't available/possible in this session.\n    # Set to a literal smp encoder name (e.g. \"resnet50\", \"se_resnext101_32x4d\")\n    # to skip the probing and just use that encoder directly.\n   \n    #encoder_name = \"auto_strong_convnext\"\n    \n    encoder_name = \"tu-convnext_tiny\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n    # \"unet\" or \"unetplusplus\"\n    architecture = \"unetplusplus\"\n\n    # --- LOSS ---------------------------------------------------------------\n    bce_weight = 0.30\n    dice_weight = 0.35\n    focal_tversky_weight = 0.35\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n\n    # --- THRESHOLD ------------------------------------------------------------\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n    threshold = 0.50\n\n    # --- TTA ------------------------------------------------------------------\n    use_tta = True\n    tta_modes = [\"original\", \"hflip\", \"vflip\", \"rot90\"]\n\n    # --- AdaBN ------------------------------------------------------------------\n    use_adabn = True\n    adabn_max_patches = 2000\n\n    # --- POSTPROCESS --------------------------------------------------------\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_size = 8\n\n    # ==========================================================================\n    # V3 ADDITIONS\n    # ==========================================================================\n\n    # --- edge-artifact cropping: erode each fragment's tissue mask before any\n    #     patch coordinates are generated (train AND test) ---\n    USE_MASK_EROSION = True\n    mask_erode_px = 24\n\n    # --- CLAHE (per-slice, applied at patch-read time) ---\n    USE_CLAHE = True\n    clahe_clip_limit = 2.0\n    clahe_tile_grid = (8, 8)\n\n    # --- normalization mode ---\n    # \"per_fragment_zscore\": one mean/std per fragment (sampled once from the\n    #     tissue region), applied to every patch from that fragment. Removes\n    #     fragment-level brightness/contrast differences without introducing\n    #     patch-to-patch normalization noise.\n    # \"per_patch_zscore\": V2's original behavior (recompute mean/std per patch).\n    NORMALIZATION_MODE = \"per_fragment_zscore\"\n    frag_stats_sample_patches = 60\n\n    # --- domain-randomization augmentation -----------------------------------\n    USE_ELASTIC = True\n    elastic_alpha = 40\n    elastic_sigma = 6\n    elastic_p = 0.30\n\n    USE_COARSE_DROPOUT = True\n    coarse_dropout_p = 0.25\n\n    USE_GAUSS_NOISE = True\n    gauss_noise_p = 0.25\n\n    USE_SHADOW = True\n    shadow_p = 0.20\n\n    # synthetic \"fake ink\" / \"fake fiber\" distractors -- image only, mask\n    # untouched, so the model can't shortcut on raw brightness/texture alone\n    USE_FAKE_INK_FIBER = True\n    fake_ink_p = 0.15\n    fake_fiber_p = 0.15\n\n    # histogram-matching STYLE AUGMENTATION during training: restyle some\n    # training patches toward a cached pool of (unlabeled) fragment-1 patches\n    USE_HIST_MATCH_AUG = True\n    hist_match_aug_p = 0.30\n    hist_match_pool_size = 20\n\n    # --- EMA ------------------------------------------------------------------\n    USE_EMA = True\n    ema_decay = 0.999\n\n    # --- DANN (unsupervised domain-adversarial training) ---------------------\n    # Uses UNLABELED fragment-1 patches during training (never its labels).\n    USE_DANN = True\n    dann_lambda_max = 1.0\n    dann_target_batch_size = 8\n\n    # --- MixStyle (input-level batch statistics mixing, training only) -------\n    USE_MIXSTYLE = True\n    mixstyle_p = 0.5\n    mixstyle_alpha = 0.1\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working\"\n    ckpt_path = os.path.join(out_dir, \"vesuviusnet_v3_best.pth\")\n    viz_dir = os.path.join(out_dir, \"v3_visualizations\")\n    metrics_path = os.path.join(out_dir, \"v3_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\n\nset_seed(CFG.seed)\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ============================================================\n# 3. TISSUE / LABEL HELPERS  (+ V3 mask erosion)\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        thresh = mid.mean() * 0.15\n        mask = (mid > thresh).astype(np.uint8) * 255\n    mask = (mask > 0).astype(np.uint8)\n\n    if CFG.USE_MASK_EROSION and CFG.mask_erode_px > 0:\n        k = 2 * CFG.mask_erode_px + 1\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n        eroded = cv2.erode(mask, kernel, iterations=1)\n        # safety: never erode a mask down to nothing (e.g. a very small/thin\n        # fragment) -- fall back to the un-eroded mask if erosion kills it\n        if eroded.sum() > 0:\n            mask = eroded\n        else:\n            print(f\"  [mask erosion] erosion by {CFG.mask_erode_px}px emptied the \"\n                  f\"mask for {frag_dir}; keeping the un-eroded mask instead.\")\n    return mask\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID / SPATIAL SPLIT / BALANCING (unchanged from V2)\n# ============================================================\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    if H < patch_size or W < patch_size:\n        return coords\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            tissue_frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\ndef spatial_split_coords(coords, image_shape, patch_size, val_fraction=0.20):\n    H, W = image_shape\n    boundary = int(H * (1.0 - val_fraction))\n    gap = patch_size\n    train_coords, val_coords = [], []\n    for y, x in coords:\n        if y + patch_size <= boundary - gap:\n            train_coords.append((y, x))\n        elif y >= boundary:\n            val_coords.append((y, x))\n    return train_coords, val_coords\n\n\ndef get_patch_positive_fraction(labels, y, x, patch_size):\n    return float(labels[y:y + patch_size, x:x + patch_size].mean())\n\n\ndef balance_positive_patches(samples, labels_full, patch_size, positive_threshold=0.001,\n                              target_positive_ratio=0.55, max_positive_repeat=3):\n    positive, negative = [], []\n    for fid, y, x in samples:\n        frac = get_patch_positive_fraction(labels_full[fid], y, x, patch_size)\n        (positive if frac >= positive_threshold else negative).append((fid, y, x))\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n    if len(negative) == 0:\n        return samples\n\n    desired_negative = int(len(positive) * (1.0 - target_positive_ratio) / target_positive_ratio)\n    desired_negative = max(desired_negative, len(positive))\n    negative_selected = random.sample(negative, desired_negative) if desired_negative < len(negative) else negative.copy()\n\n    repeat = int(math.ceil((target_positive_ratio * len(negative_selected)) /\n                            ((1.0 - target_positive_ratio) * max(len(positive), 1))))\n    repeat = max(1, min(repeat, max_positive_repeat))\n\n    positive_selected = positive * repeat\n    samples_balanced = positive_selected + negative_selected\n    random.shuffle(samples_balanced)\n\n    print(f\"Positive patches: {len(positive)} | Negative patches: {len(negative)}\")\n    print(f\"Balanced dataset: {len(samples_balanced)} | \"\n          f\"positive ratio={len(positive_selected)/max(len(samples_balanced),1):.3f}\")\n    return samples_balanced\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME  (+ V3: CLAHE, per-fragment stats)\n# ============================================================\n\nclass FragmentVolume:\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h = None\n        self._w = None\n        self._clahe = (cv2.createCLAHE(clipLimit=CFG.clahe_clip_limit, tileGridSize=CFG.clahe_tile_grid)\n                       if CFG.USE_CLAHE else None)\n        self.frag_mean = None   # lazily computed per-fragment normalization stats\n        self.frag_std = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for p in self.paths:\n            try:\n                arr = tifffile.memmap(p, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(p)\n            slices.append(arr)\n        self._slices = slices\n        self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size, apply_clahe=None):\n        self._ensure_open()\n        apply_clahe = CFG.USE_CLAHE if apply_clahe is None else apply_clahe\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.shape != (size, size):\n                padded = np.zeros((size, size), dtype=np.uint8)\n                hh = min(size, block.shape[0])\n                ww = min(size, block.shape[1])\n                if block.dtype != np.uint8:\n                    block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n                padded[:hh, :ww] = block[:hh, :ww]\n                block = padded\n            elif block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            if apply_clahe and self._clahe is not None:\n                block = self._clahe.apply(block)\n            out[i] = block\n        return out\n\n    def compute_fragment_stats(self, tissue_mask, patch_size, n_samples=CFG.frag_stats_sample_patches):\n        \"\"\"Sample a handful of tissue-covered patches once and compute a single\n        mean/std for the whole fragment (used by NORMALIZATION_MODE='per_fragment_zscore').\n        CLAHE is applied first (if enabled) so the stats match what training/inference\n        will actually see.\"\"\"\n        self._ensure_open()\n        coords = generate_grid_coords(tissue_mask, patch_size, patch_size, CFG.min_tissue_frac_train)\n        if not coords:\n            self.frag_mean, self.frag_std = 128.0, 50.0\n            return\n        sample_coords = random.sample(coords, min(n_samples, len(coords)))\n        mid = len(self._slices) // 2\n        vals = []\n        for (y, x) in sample_coords:\n            block = self._slices[mid][y:y + patch_size, x:x + patch_size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            if CFG.USE_CLAHE and self._clahe is not None:\n                block = self._clahe.apply(block)\n            vals.append(block.astype(np.float32).ravel())\n        vals = np.concatenate(vals)\n        self.frag_mean = float(vals.mean())\n        self.frag_std = float(vals.std() + 1e-6)\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\n# ============================================================\n# 6. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float, frag_mean=None, frag_std=None):\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\" and frag_mean is not None:\n        return (img_float - frag_mean / 255.0) / (frag_std / 255.0 + 1e-6)\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ============================================================\n# 7. DOMAIN-SPECIFIC SYNTHETIC AUGMENTATION (fake ink / fake fiber)\n# ============================================================\n\ndef inject_fake_ink(img_hwd, n_strokes=None, intensity_range=(20, 60)):\n    \"\"\"Draws 1-3 random thin curved strokes into the image ONLY (mask untouched),\n    at a random subset of depth slices with a random brightness delta (can be\n    lighter or darker). Purpose: these look locally like 'ink-ish' structure but\n    are NOT labeled as ink, so the model can't shortcut on raw brightness alone --\n    it has to learn structure that's actually predictive of the real label.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    n_strokes = n_strokes or random.randint(1, 3)\n    for _ in range(n_strokes):\n        n_pts = random.randint(3, 6)\n        pts = np.stack([\n            np.random.randint(0, w, n_pts),\n            np.random.randint(0, h, n_pts),\n        ], axis=1).astype(np.int32)\n        thickness = random.randint(1, 3)\n        delta = random.choice([-1, 1]) * random.randint(*intensity_range)\n        depth_frac = random.uniform(0.3, 0.7)\n        n_affected = max(1, int(d * depth_frac))\n        affected_slices = random.sample(range(d), n_affected)\n        stroke_mask = np.zeros((h, w), dtype=np.uint8)\n        cv2.polylines(stroke_mask, [pts], isClosed=False, color=1, thickness=thickness)\n        for zi in affected_slices:\n            sl = out[:, :, zi].astype(np.int16)\n            sl[stroke_mask > 0] = np.clip(sl[stroke_mask > 0] + delta, 0, 255)\n            out[:, :, zi] = sl.astype(np.uint8)\n    return out\n\n\ndef inject_fake_shadow(img_hwd, darken_range=(0.55, 0.85)):\n    \"\"\"Custom shadow-like augmentation. albumentations' RandomShadow hard-requires\n    3-channel RGB images and raises ValueError on anything else -- it is NOT usable on\n    this 22-channel depth-stack data regardless of API version, so this replaces it\n    entirely with a plain numpy implementation. Darkens a smooth radial region of the\n    patch, applied consistently across depth slices (a plausible stand-in for\n    scan-illumination gradients across a fragment's surface). Image only -- mask\n    untouched.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.astype(np.float32)\n    cx = random.uniform(0.2, 0.8) * w\n    cy = random.uniform(0.2, 0.8) * h\n    radius = random.uniform(0.4, 0.9) * max(h, w)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)\n    dist_norm = np.clip(dist / radius, 0, 1)\n    darken_factor = random.uniform(*darken_range)\n    gradient = 1.0 - (1.0 - darken_factor) * (1.0 - dist_norm) ** 2\n    for zi in range(d):\n        out[:, :, zi] = out[:, :, zi] * gradient\n    return np.clip(out, 0, 255).astype(np.uint8)\n\n\ndef inject_fake_fiber(img_hwd, amplitude_range=(6, 18)):\n    \"\"\"Overlays a low-frequency directional sinusoidal texture (mimicking papyrus\n    fiber striations) with a consistent orientation across depth (since real fibers\n    are a physical structure, not a per-slice-random pattern), image only.\"\"\"\n    h, w, d = img_hwd.shape\n    out = img_hwd.copy()\n    theta = random.uniform(0, math.pi)\n    freq = random.uniform(0.02, 0.06)\n    amplitude = random.uniform(*amplitude_range)\n    yy, xx = np.meshgrid(np.arange(h), np.arange(w), indexing=\"ij\")\n    phase = (xx * math.cos(theta) + yy * math.sin(theta)) * freq\n    pattern = (amplitude * np.sin(2 * math.pi * phase)).astype(np.float32)\n    for zi in range(d):\n        sl = out[:, :, zi].astype(np.float32) + pattern\n        out[:, :, zi] = np.clip(sl, 0, 255).astype(np.uint8)\n    return out\n\n\n# ============================================================\n# 8. AUGMENTATION  (V2's transforms + V3 domain-randomization,\n#    version-safe construction for API drift across albumentations versions)\n# ============================================================\n\ndef _safe_transform(builder_new, builder_old, name):\n    \"\"\"Try constructing a transform with the current albumentations API; if that\n    fails (parameter names changed across versions), fall back to the older API.\n    If both fail, skip the transform rather than crashing the whole pipeline.\"\"\"\n    try:\n        return builder_new()\n    except Exception as e1:\n        try:\n            t = builder_old()\n            print(f\"  [augmentation] {name}: using legacy-API construction ({e1})\")\n            return t\n        except Exception as e2:\n            print(f\"  [augmentation] {name}: unavailable in this albumentations \"\n                  f\"version ({e1} / {e2}); skipping it.\")\n            return None\n\n\ndef build_train_transform():\n    transforms = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.10, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.40),\n        A.GridDistortion(num_steps=5, distort_limit=0.15, p=0.15),\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.25),\n    ]\n\n    if CFG.USE_ELASTIC:\n        t = _safe_transform(\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            lambda: A.ElasticTransform(alpha=CFG.elastic_alpha, sigma=CFG.elastic_sigma,\n                                        alpha_affine=CFG.elastic_sigma,\n                                        border_mode=cv2.BORDER_REFLECT, p=CFG.elastic_p),\n            \"ElasticTransform\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_COARSE_DROPOUT:\n        t = _safe_transform(\n            lambda: A.CoarseDropout(num_holes_range=(1, 4), hole_height_range=(0.03, 0.10),\n                                     hole_width_range=(0.03, 0.10), p=CFG.coarse_dropout_p),\n            lambda: A.CoarseDropout(max_holes=4, max_height=0.10, max_width=0.10,\n                                     min_holes=1, min_height=0.03, min_width=0.03,\n                                     p=CFG.coarse_dropout_p),\n            \"CoarseDropout\")\n        if t is not None:\n            transforms.append(t)\n\n    if CFG.USE_GAUSS_NOISE:\n        t = _safe_transform(\n            lambda: A.GaussNoise(std_range=(0.02, 0.08), p=CFG.gauss_noise_p),\n            lambda: A.GaussNoise(var_limit=(5.0, 40.0), p=CFG.gauss_noise_p),\n            \"GaussNoise\")\n        if t is not None:\n            transforms.append(t)\n\n    # NOTE: albumentations' RandomShadow is intentionally NOT used here -- it hard-\n    # requires 3-channel RGB images and raises ValueError on our 22-channel depth-stack\n    # data regardless of API version. See inject_fake_shadow() for the replacement,\n    # applied directly in the dataset (image-only, works on any channel count).\n\n    return A.Compose(transforms)\n\n\n# ============================================================\n# 9. DATASET  (+ V3: fragment stats, fake ink/fiber, histogram-match aug)\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None, jitter=0,\n                 hist_match_pool=None, train_mode=True):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.jitter = jitter\n        self.hist_match_pool = hist_match_pool   # list of (D,H,W) uint8 reference patches\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[fid]\n        H, W = vol.shape\n\n        if self.jitter > 0:\n            dy = random.randint(-self.jitter, self.jitter)\n            dx = random.randint(-self.jitter, self.jitter)\n            y = max(0, min(y + dy, H - size))\n            x = max(0, min(x + dx, W - size))\n\n        patch = vol.read_patch(y, x, size)              # (D,H,W) uint8, CLAHE already applied\n        label = self.labels[fid][y:y + size, x:x + size]\n\n        img = np.transpose(patch, (1, 2, 0))              # HWD\n\n        # --- V3: histogram-matching style augmentation (train only) ---\n        if (self.train_mode and CFG.USE_HIST_MATCH_AUG and self.hist_match_pool\n                and random.random() < CFG.hist_match_aug_p):\n            ref = random.choice(self.hist_match_pool)      # (D,H,W) uint8\n            ref_hwd = np.transpose(ref, (1, 2, 0))\n            try:\n                img = match_histograms(img, ref_hwd, channel_axis=2).astype(np.uint8)\n            except Exception:\n                pass  # if shapes/versions mismatch, just skip this augmentation for this sample\n\n        # --- V3: fake ink / fake fiber / fake shadow distractors (train only, image only) ---\n        if self.train_mode and CFG.USE_FAKE_INK_FIBER:\n            if random.random() < CFG.fake_ink_p:\n                img = inject_fake_ink(img)\n            if random.random() < CFG.fake_fiber_p:\n                img = inject_fake_fiber(img)\n        if self.train_mode and CFG.USE_SHADOW and random.random() < CFG.shadow_p:\n            img = inject_fake_shadow(img)\n\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img, vol.frag_mean, vol.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\nclass UnlabeledPatchDataset(Dataset):\n    \"\"\"For DANN: yields normalized (D,H,W) tensors from fragment 1 with NO labels.\"\"\"\n    def __init__(self, volume, samples, patch_size):\n        self.volume = volume\n        self.samples = samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        y, x = self.samples[idx]\n        patch = self.volume.read_patch(y, x, self.patch_size)\n        img = np.transpose(patch, (1, 2, 0)).astype(np.float32) / 255.0\n        img = normalize_patch(img, self.volume.frag_mean, self.volume.frag_std)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        return torch.from_numpy(img)\n\n\n# ============================================================\n# 10. MODEL  (+ V3: ConvNeXt/UNet++ probing, DANN domain classifier, MixStyle)\n# ============================================================\n\ndef build_model():\n    arch_cls = smp.UnetPlusPlus if CFG.architecture == \"unetplusplus\" else smp.Unet\n\n    if CFG.encoder_name == \"auto_strong_convnext\":\n        import sys\n        if sys.version_info < (3, 9):\n            print(f\"[backbone] Python {sys.version_info.major}.{sys.version_info.minor} detected \"\n                  f\"(this Kaggle image) -- current segmentation-models-pytorch/timm releases that \"\n                  f\"support ConvNeXt via 'tu-' encoders require Python >=3.9, so the upgrade attempt \"\n                  f\"would just fail. Skipping it and using efficientnet-b4 directly instead.\")\n            return arch_cls(encoder_name=\"efficientnet-b4\", encoder_weights=CFG.encoder_weights,\n                             in_channels=CFG.in_channels, classes=1,\n                             decoder_attention_type=CFG.decoder_attention_type)\n        try:\n            import subprocess, sys\n            result = subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \"-U\",\n                                      \"segmentation-models-pytorch\", \"timm\"],\n                                     capture_output=True, text=True)\n            if result.returncode != 0:\n                raise RuntimeError(f\"pip upgrade failed (exit {result.returncode}): \"\n                                    f\"{result.stderr.strip().splitlines()[-1] if result.stderr else 'unknown error'}\")\n            import importlib\n            importlib.reload(smp)\n            arch_cls = smp.UnetPlusPlus if CFG.architecture == \"unetplusplus\" else smp.Unet\n            model = arch_cls(encoder_name=\"tu-convnext_tiny\", encoder_weights=CFG.encoder_weights,\n                              in_channels=CFG.in_channels, classes=1,\n                              decoder_attention_type=CFG.decoder_attention_type)\n            print(\"[backbone] using tu-convnext_tiny (ImageNet-pretrained, upgraded smp)\")\n            return model\n        except Exception as e:\n            print(f\"[backbone] ConvNeXt unavailable in this session ({e}); \"\n                  f\"falling back to efficientnet-b4.\")\n            model = arch_cls(encoder_name=\"efficientnet-b4\", encoder_weights=CFG.encoder_weights,\n                              in_channels=CFG.in_channels, classes=1,\n                              decoder_attention_type=CFG.decoder_attention_type)\n            return model\n    else:\n        return arch_cls(encoder_name=CFG.encoder_name, encoder_weights=CFG.encoder_weights,\n                         in_channels=CFG.in_channels, classes=1,\n                         decoder_attention_type=CFG.decoder_attention_type)\n\n\nclass GradReverse(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x, lambd):\n        ctx.lambd = lambd\n        return x.view_as(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return -ctx.lambd * grad_output, None\n\n\ndef grad_reverse(x, lambd=1.0):\n    return GradReverse.apply(x, lambd)\n\n\nclass DomainClassifier(nn.Module):\n    \"\"\"Tries to tell source (train fragments) from target (fragment 1, unlabeled)\n    using the encoder's deepest feature map. Trained via a gradient-reversal layer,\n    so the ENCODER is simultaneously pushed to make that classification harder --\n    i.e. toward domain-invariant features.\"\"\"\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(in_ch, 128), nn.ReLU(inplace=True), nn.Dropout(0.3),\n            nn.Linear(128, 1),\n        )\n\n    def forward(self, feat, lambd):\n        return self.net(grad_reverse(feat, lambd))\n\n\nclass EncoderFeatureCapture:\n    \"\"\"Architecture-agnostic hook that captures model.encoder's output feature list\n    on every forward pass, without needing to know the decoder's call signature\n    (which differs across segmentation_models_pytorch versions).\"\"\"\n    def __init__(self, model):\n        self.features = None\n        self.handle = model.encoder.register_forward_hook(self._hook)\n\n    def _hook(self, module, inp, out):\n        self.features = out\n\n    def remove(self):\n        self.handle.remove()\n\n\ndef mixstyle_batch(imgs, p=CFG.mixstyle_p, alpha=CFG.mixstyle_alpha):\n    \"\"\"Input-level MixStyle: mixes each sample's per-channel(depth) mean/std with a\n    randomly paired other sample's in the same batch. Implemented at the input\n    (rather than hooked into an arbitrary internal encoder layer) so it works\n    identically regardless of which backbone is selected.\"\"\"\n    if torch.rand(1).item() > p:\n        return imgs\n    B = imgs.size(0)\n    if B < 2:\n        return imgs\n    mu = imgs.mean(dim=[2, 3], keepdim=True)\n    var = imgs.var(dim=[2, 3], keepdim=True)\n    sig = (var + 1e-6).sqrt()\n    x_norm = (imgs - mu) / sig\n    perm = torch.randperm(B, device=imgs.device)\n    mu2, sig2 = mu[perm], sig[perm]\n    lam = torch.distributions.Beta(alpha, alpha).sample((B, 1, 1, 1)).to(imgs.device)\n    mu_mix = lam * mu + (1 - lam) * mu2\n    sig_mix = lam * sig + (1 - lam) * sig2\n    return x_norm * sig_mix + mu_mix\n\n\nclass EMAModel:\n    \"\"\"Exponential moving average of model weights. Update after every optimizer\n    step; use ema.state_dict() for validation/threshold-search/final inference\n    instead of the raw (noisier, more overfit-to-late-training-batches) weights.\"\"\"\n    def __init__(self, model, decay=CFG.ema_decay):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def state_dict(self):\n        return self.shadow\n\n\n# ============================================================\n# 11. LOSSES (unchanged from V2)\n# ============================================================\n\ndef soft_dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    denom = probs.sum(dim=1) + t.sum(dim=1)\n    return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.30, beta=0.70, gamma=0.75, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1.0 - t)).sum(dim=1)\n    fn = ((1.0 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return torch.pow(1.0 - tversky, gamma).mean()\n\n\nclass V2ComboLoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n\n    def forward(self, logits, targets):\n        bce = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=self.pos_weight)\n        dice = soft_dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.focal_tversky_gamma)\n        return CFG.bce_weight * bce + CFG.dice_weight * dice + CFG.focal_tversky_weight * tv\n\n\n# ============================================================\n# 12. METRICS (unchanged from V2)\n# ============================================================\n\ndef dice_from_counts(tp, fp, fn, eps=1e-6):\n    return (2.0 * tp + eps) / (2.0 * tp + fp + fn + eps)\n\n\ndef metrics_from_counts(tp, fp, fn, eps=1e-6):\n    dice = dice_from_counts(tp, fp, fn, eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1.0 + beta2) * precision * recall + eps) / (beta2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision),\n            \"recall\": float(recall), \"fbeta0.5\": float(fbeta)}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1.0 - targets)).sum().item()\n        self.fn += ((1.0 - preds) * targets).sum().item()\n\n    def compute(self):\n        return metrics_from_counts(self.tp, self.fp, self.fn)\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=100):\n    if len(samples) == 0:\n        return 1e-4\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(float(pos / max(total, 1)), 1e-4)\n\n\n# ============================================================\n# 13. BUILD TRAIN DATA\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING TRAIN DATA\")\nprint(\"=\" * 70)\n\ntrain_volumes = {}\ntrain_labels_full = {}\ntrain_masks_full = {}\ntrain_samples_raw = []\nval_samples = []\n\nfor fid in CFG.train_frags:\n    print(f\"\\nProcessing fragment {fid} ...\")\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise RuntimeError(f\"inklabels.png missing for fragment {fid}\")\n\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n    if CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n        vol.compute_fragment_stats(mask, CFG.patch_size)\n        print(f\"  fragment {fid} normalization stats: mean={vol.frag_mean:.2f} std={vol.frag_std:.2f}\")\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {len(coords)} candidate patches (post mask-erosion)\")\n\n    tr_coords, va_coords = spatial_split_coords(coords, mask.shape, CFG.patch_size, CFG.val_fraction)\n    print(f\"  spatial train={len(tr_coords)} validation={len(va_coords)}\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    train_masks_full[fid] = mask\n    train_samples_raw.extend([(fid, y, x) for y, x in tr_coords])\n    val_samples.extend([(fid, y, x) for y, x in va_coords])\n    del mask\n    cleanup_memory()\n\nprint(\"\\nRaw train samples:\", len(train_samples_raw))\nprint(\"Validation samples:\", len(val_samples))\n\ntrain_samples = balance_positive_patches(\n    train_samples_raw, train_labels_full, CFG.patch_size,\n    positive_threshold=CFG.positive_patch_fraction,\n    target_positive_ratio=CFG.target_positive_patch_ratio,\n    max_positive_repeat=CFG.max_positive_repeat,\n)\n\n\n# ============================================================\n# 14. TEST-FRAGMENT MASK/VOLUME LOADED EARLY (unlabeled use only:\n#     histogram-match reference pool + DANN target patches).\n#     Its LABELS are loaded much later, only for local diagnostics.\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\")\nprint(\"=\" * 70)\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\nif CFG.NORMALIZATION_MODE == \"per_fragment_zscore\":\n    test_vol.compute_fragment_stats(test_mask, CFG.patch_size)\n    print(f\"  fragment {CFG.test_frag} normalization stats: \"\n          f\"mean={test_vol.frag_mean:.2f} std={test_vol.frag_std:.2f}\")\n\ntest_coords_for_unlabeled_use = generate_grid_coords(\n    test_mask, CFG.patch_size, CFG.test_stride, CFG.min_tissue_frac_train)\nprint(f\"Fragment {CFG.test_frag}: {len(test_coords_for_unlabeled_use)} unlabeled candidate patches\")\n\nhist_match_pool = None\nif CFG.USE_HIST_MATCH_AUG:\n    pool_coords = random.sample(test_coords_for_unlabeled_use,\n                                 min(CFG.hist_match_pool_size, len(test_coords_for_unlabeled_use)))\n    hist_match_pool = [test_vol.read_patch(y, x, CFG.patch_size) for y, x in pool_coords]\n    print(f\"Built histogram-matching reference pool: {len(hist_match_pool)} patches \"\n          f\"({sum(p.nbytes for p in hist_match_pool)/1e6:.1f} MB)\")\n\n\n# ============================================================\n# 15. DATASETS / LOADERS\n# ============================================================\n\ntrain_transform = build_train_transform()\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples, CFG.patch_size,\n                            transform=train_transform, jitter=CFG.train_jitter,\n                            hist_match_pool=hist_match_pool, train_mode=True)\n\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples, CFG.patch_size,\n                          transform=None, jitter=0, hist_match_pool=None, train_mode=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0)\n\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         drop_last=False, persistent_workers=CFG.num_workers > 0)\n\ndann_loader = None\nif CFG.USE_DANN:\n    dann_ds = UnlabeledPatchDataset(test_vol, test_coords_for_unlabeled_use, CFG.patch_size)\n    dann_loader = DataLoader(dann_ds, batch_size=CFG.dann_target_batch_size, shuffle=True,\n                              num_workers=max(1, CFG.num_workers - 1), pin_memory=(CFG.device == \"cuda\"),\n                              drop_last=True, persistent_workers=True)\n\n    def infinite_dann_loader():\n        while True:\n            for batch in dann_loader:\n                yield batch\n    dann_iter = infinite_dann_loader()\n\n\n# ============================================================\n# 16. MODEL / LOSS / OPTIMIZER\n# ============================================================\n\nprint(\"\\nBuilding V3 model ...\")\nmodel = build_model().to(CFG.device)\n\nprint(\"\\nEstimating positive-pixel fraction ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size, n_samples=100)\nprint(f\"Estimated positive fraction: {pos_frac:.6f}\")\n\nwith torch.no_grad():\n    bias_val = float(np.log(pos_frac / max(1.0 - pos_frac, 1e-6)))\n    try:\n        model.segmentation_head[0].bias.fill_(bias_val)\n    except Exception:\n        pass\nprint(f\"Output bias initialized to {bias_val:.4f}\")\n\nraw_ratio = (1.0 - pos_frac) / max(pos_frac, 1e-6)\npos_weight_val = float(np.clip(np.sqrt(raw_ratio), 1.0, 8.0))\npos_weight = torch.tensor([pos_weight_val], dtype=torch.float32, device=CFG.device)\nprint(f\"BCE positive weight: {pos_weight_val:.3f}\")\n\ncriterion = V2ComboLoss(pos_weight=pos_weight)\n\nencoder_params, decoder_params = [], []\nfor name, param in model.named_parameters():\n    if not param.requires_grad:\n        continue\n    (encoder_params if name.startswith(\"encoder.\") else decoder_params).append(param)\nprint(f\"Encoder parameters: {len(encoder_params)} tensors | Decoder/head parameters: {len(decoder_params)} tensors\")\n\nparam_groups = [\n    {\"params\": encoder_params, \"lr\": CFG.encoder_lr},\n    {\"params\": decoder_params, \"lr\": CFG.decoder_lr},\n]\n\n# --- V3: DANN domain classifier, sized dynamically from a dry forward pass ---\ndomain_classifier = None\nfeature_capture = None\nif CFG.USE_DANN:\n    feature_capture = EncoderFeatureCapture(model)\n    with torch.no_grad():\n        dummy = torch.zeros(1, CFG.in_channels, CFG.patch_size, CFG.patch_size, device=CFG.device)\n        model.eval()\n        _ = model(dummy)\n        deepest_ch = feature_capture.features[-1].shape[1]\n        model.train()\n    domain_classifier = DomainClassifier(deepest_ch).to(CFG.device)\n    param_groups.append({\"params\": domain_classifier.parameters(), \"lr\": CFG.dann_lr})\n    print(f\"[DANN] domain classifier attached on {deepest_ch}-channel bottleneck features\")\n    del dummy\n    cleanup_memory()\n\noptimizer = torch.optim.AdamW(param_groups, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nema = EMAModel(model, decay=CFG.ema_decay) if CFG.USE_EMA else None\n\n\n# ============================================================\n# 17. TRAIN / VALIDATION EPOCH  (+ V3: MixStyle, DANN, EMA)\n# ============================================================\n\n_global_step = 0\n_total_steps = CFG.epochs * max(1, len(train_loader) // CFG.accumulation_steps)\n\n\ndef dann_lambda_schedule():\n    progress = min(_global_step / max(_total_steps, 1), 1.0)\n    return CFG.dann_lambda_max * (2.0 / (1.0 + math.exp(-10.0 * progress)) - 1.0)\n\n\ndef run_epoch(loader, train_mode=True, threshold=0.5):\n    global _global_step\n    model.train(train_mode)\n    if domain_classifier is not None:\n        domain_classifier.train(train_mode)\n\n    total_loss = 0.0\n    total_dann_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    if train_mode:\n        optimizer.zero_grad(set_to_none=True)\n\n    for batch_idx, (imgs, masks) in enumerate(loader):\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        if train_mode and CFG.USE_MIXSTYLE:\n            imgs = mixstyle_batch(imgs)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n                dann_loss_val = 0.0\n\n                if train_mode and CFG.USE_DANN:\n                    src_feat = feature_capture.features[-1]\n                    lambd = dann_lambda_schedule()\n                    src_domain_logits = domain_classifier(src_feat, lambd)\n                    src_domain_target = torch.zeros_like(src_domain_logits)  # source = 0\n\n                    tgt_imgs = next(dann_iter).to(CFG.device, non_blocking=True)\n                    _ = model(tgt_imgs)   # populates feature_capture.features for target\n                    tgt_feat = feature_capture.features[-1]\n                    tgt_domain_logits = domain_classifier(tgt_feat, lambd)\n                    tgt_domain_target = torch.ones_like(tgt_domain_logits)   # target = 1\n\n                    dann_loss = 0.5 * (\n                        F.binary_cross_entropy_with_logits(src_domain_logits, src_domain_target) +\n                        F.binary_cross_entropy_with_logits(tgt_domain_logits, tgt_domain_target)\n                    )\n                    dann_loss_val = dann_loss.item()\n                    loss = loss + dann_loss\n\n                loss_for_backward = loss / CFG.accumulation_steps if train_mode else loss\n\n            if train_mode:\n                scaler.scale(loss_for_backward).backward()\n                if (batch_idx + 1) % CFG.accumulation_steps == 0:\n                    scaler.unscale_(optimizer)\n                    params_to_clip = list(model.parameters())\n                    if domain_classifier is not None:\n                        params_to_clip += list(domain_classifier.parameters())\n                    torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad(set_to_none=True)\n                    if ema is not None:\n                        ema.update(model)\n                    _global_step += 1\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.detach().item()\n        total_dann_loss += dann_loss_val\n\n        del imgs, masks, logits, probs\n\n    if train_mode and len(loader) % CFG.accumulation_steps != 0:\n        scaler.unscale_(optimizer)\n        params_to_clip = list(model.parameters())\n        if domain_classifier is not None:\n            params_to_clip += list(domain_classifier.parameters())\n        torch.nn.utils.clip_grad_norm_(params_to_clip, CFG.grad_clip)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        if ema is not None:\n            ema.update(model)\n\n    metrics = global_acc.compute()\n    avg_loss = total_loss / max(len(loader), 1)\n    avg_dann_loss = total_dann_loss / max(len(loader), 1)\n    return avg_loss, metrics, avg_dann_loss\n\n\n# ============================================================\n# 18. TRAINING LOOP\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING V3 TRAINING\")\nprint(\"=\" * 70)\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"train_dice\": [], \"val_dice\": [],\n           \"val_iou\": [], \"val_precision\": [], \"val_recall\": [], \"dann_loss\": []}\n\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_metrics, dann_loss_avg = run_epoch(train_loader, train_mode=True, threshold=0.50)\n\n    # --- validate using EMA weights if enabled (swap in, evaluate, swap back) ---\n    if ema is not None:\n        backup = {k: v.detach().clone() for k, v in model.state_dict().items()}\n        model.load_state_dict(ema.state_dict())\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n        model.load_state_dict(backup)\n        del backup\n    else:\n        val_loss, val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=0.50)\n\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"train_dice\"].append(train_metrics[\"dice\"])\n    history[\"val_dice\"].append(val_metrics[\"dice\"])\n    history[\"val_iou\"].append(val_metrics[\"iou\"])\n    history[\"val_precision\"].append(val_metrics[\"precision\"])\n    history[\"val_recall\"].append(val_metrics[\"recall\"])\n    history[\"dann_loss\"].append(dann_loss_avg)\n\n    print(f\"\\n[{epoch:02d}/{CFG.epochs}] time={time.time()-t0:.1f}s\")\n    print(f\"train_loss={train_loss:.5f} train_dice={train_metrics['dice']:.5f} \"\n          f\"dann_loss={dann_loss_avg:.5f}\")\n    print(f\"val_loss={val_loss:.5f} val_dice={val_metrics['dice']:.5f} val_iou={val_metrics['iou']:.5f}\")\n    print(f\"precision={val_metrics['precision']:.5f} recall={val_metrics['recall']:.5f}\")\n    print(f\"encoder_lr={optimizer.param_groups[0]['lr']:.7f} decoder_lr={optimizer.param_groups[1]['lr']:.7f}\")\n\n    if val_metrics[\"dice\"] > best_val_dice:\n        best_val_dice = val_metrics[\"dice\"]\n        epochs_no_improve = 0\n        save_state = ema.state_dict() if ema is not None else model.state_dict()\n        checkpoint = {\"model\": save_state, \"cfg\": cfg_to_dict(CFG), \"best_val_dice\": best_val_dice,\n                      \"history\": history, \"pos_frac\": pos_frac, \"pos_weight\": pos_weight_val,\n                      \"used_ema\": CFG.USE_EMA}\n        torch.save(checkpoint, CFG.ckpt_path)\n        print(f\"*** NEW BEST CHECKPOINT val_dice={best_val_dice:.5f} \"\n              f\"({'EMA' if CFG.USE_EMA else 'raw'} weights) ***\")\n    else:\n        epochs_no_improve += 1\n        print(f\"No improvement: {epochs_no_improve}/{CFG.early_stop_patience}\")\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(\"Early stopping.\")\n            break\n\n    cleanup_memory()\n\nprint(\"\\nBest validation Dice:\", best_val_dice)\n\nif feature_capture is not None:\n    feature_capture.remove()\n\n\n# ============================================================\n# 19. LOAD BEST MODEL (EMA weights, if used)\n# ============================================================\n\ncheckpoint = torch.load(CFG.ckpt_path, map_location=CFG.device)\nmodel.load_state_dict(checkpoint[\"model\"])\nprint(f\"\\nBest checkpoint loaded ({'EMA' if checkpoint.get('used_ema') else 'raw'} weights).\")\n\n\n# ============================================================\n# 20. FINE DICE THRESHOLD SEARCH (unchanged from V2)\n# ============================================================\n\n@torch.no_grad()\ndef find_best_dice_threshold(model, loader):\n    model.eval()\n    thresholds = np.arange(CFG.threshold_min, CFG.threshold_max + 0.0001, CFG.threshold_step)\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks_np = masks.numpy()\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * masks_np).sum()\n            fp[i] += (preds * (1.0 - masks_np)).sum()\n            fn[i] += ((1.0 - preds) * masks_np).sum()\n        del imgs, masks, logits, probs\n    dice_scores = (2.0 * tp + 1e-6) / (2.0 * tp + fp + fn + 1e-6)\n    best_idx = int(np.argmax(dice_scores))\n    return float(thresholds[best_idx]), float(dice_scores[best_idx]), thresholds, dice_scores\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINE DICE THRESHOLD SEARCH\")\nprint(\"=\" * 70)\n\nbest_threshold, threshold_dice, threshold_grid, threshold_scores = find_best_dice_threshold(model, val_loader)\nprint(f\"BEST VALIDATION THRESHOLD = {best_threshold:.2f}\")\nprint(f\"DICE AT BEST THRESHOLD = {threshold_dice:.5f}\")\n\n_, final_val_metrics, _ = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal validation metrics at optimized Dice threshold:\", final_val_metrics)\n\n\n# ============================================================\n# 21. INFERENCE HELPERS (unchanged from V2)\n# ============================================================\n\ndef gaussian_window(size, sigma_frac=0.45):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2.0 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / max(win.max(), 1e-8)).astype(np.float32)\n\n\ndef apply_tta_tensor(x, mode):\n    if mode == \"original\":\n        return x\n    if mode == \"hflip\":\n        return torch.flip(x, dims=[3])\n    if mode == \"vflip\":\n        return torch.flip(x, dims=[2])\n    if mode == \"rot90\":\n        return torch.rot90(x, k=1, dims=[2, 3])\n    raise ValueError(mode)\n\n\ndef invert_tta_prediction(pred, mode):\n    if mode == \"original\":\n        return pred\n    if mode == \"hflip\":\n        return torch.flip(pred, dims=[2])\n    if mode == \"vflip\":\n        return torch.flip(pred, dims=[1])\n    if mode == \"rot90\":\n        return torch.rot90(pred, k=-1, dims=[1, 2])\n    raise ValueError(mode)\n\n\n@torch.no_grad()\ndef predict_batch_tta(model, batch_np):\n    x = torch.from_numpy(batch_np).to(CFG.device, non_blocking=True)\n    modes = CFG.tta_modes if CFG.use_tta else [\"original\"]\n    prediction_sum = None\n    for mode in modes:\n        tx = apply_tta_tensor(x, mode)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(tx)\n            probs = torch.sigmoid(logits)\n        probs = invert_tta_prediction(probs[:, 0], mode).float()\n        prediction_sum = probs / len(modes) if prediction_sum is None else prediction_sum + probs / len(modes)\n        del tx, logits, probs\n    result = prediction_sum.cpu().numpy().astype(np.float32)\n    del x, prediction_sum\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference_v2(model, vol, mask, patch_size, stride, device, batch_size,\n                                 pre_transform=None):\n    \"\"\"pre_transform(raw_patch_uint8_DHW) -> raw_patch_uint8_DHW, optional, applied\n    BEFORE normalization (used by the histogram-matching diagnostic below).\"\"\"\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Inference patches: {len(coords)}\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if len(batch_imgs) == 0:\n            return\n        inp = np.stack(batch_imgs).astype(np.float32)\n        probs = predict_batch_tta(model, inp)\n        for p, (cy, cx) in zip(probs, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n        del inp, probs\n        cleanup_memory()\n\n    for y, x in coords:\n        raw_u8 = vol.read_patch(y, x, patch_size)\n        if pre_transform is not None:\n            raw_u8 = pre_transform(raw_u8)\n        raw = raw_u8.astype(np.float32) / 255.0\n        raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) >= batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum <= 1e-8] = 1.0\n    return (pred_sum / weight_sum).astype(np.float32)\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm_v2(model, vol, mask, patch_size, stride, device, max_patches, batch_size):\n    print(\"\\nStarting AdaBN...\")\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"AdaBN patches: {len(coords)}\")\n    model.train()\n    for start in range(0, len(coords), batch_size):\n        batch = coords[start:start + batch_size]\n        imgs = []\n        for y, x in batch:\n            raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n            raw = normalize_patch(raw, vol.frag_mean, vol.frag_std)\n            imgs.append(raw)\n        inp = torch.from_numpy(np.stack(imgs)).to(device, non_blocking=True)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp)\n        del inp, imgs\n    model.eval()\n    cleanup_memory()\n    print(\"AdaBN finished.\")\n    return model\n\n\ndef remove_small_components(binary, min_size):\n    if min_size <= 0:\n        return binary\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if stats[label_id, cv2.CC_STAT_AREA] >= min_size:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess_v2(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_size)\n    return binary.astype(np.uint8)\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback):\n    vals = prob_map[mask > 0]\n    vals = vals[np.isfinite(vals)]\n    if len(vals) < 100:\n        return fallback\n    try:\n        return float(np.clip(threshold_otsu(vals), 0.10, 0.90))\n    except Exception:\n        return fallback\n\n\ndef evaluate_probability_map(prob_map, gt, threshold):\n    preds = (prob_map > threshold).astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (preds * gt).sum()\n    fp = (preds * (1.0 - gt)).sum()\n    fn = ((1.0 - preds) * gt).sum()\n    return metrics_from_counts(tp, fp, fn)\n\n\n# ============================================================\n# 22. LOAD TEST LABELS (local diagnostics ONLY, loaded late and\n#     never touched during training/model-selection/threshold search)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING TEST LABELS (diagnostics only)\")\nprint(\"=\" * 70)\n\ntest_labels = load_ink_labels(test_dir)\nif test_labels is not None:\n    print(\"Test GT found: local diagnostic evaluation enabled.\")\n    gt_test = (test_labels * test_mask).astype(np.float32)\nelse:\n    print(\"No test GT found: running competition-style inference.\")\n    gt_test = None\n\n\n# ============================================================\n# 23. INFERENCE A: BASELINE + TTA\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE A: BASELINE + TTA\")\nprint(\"=\" * 70)\n\ntest_prob_baseline = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n\ntest_metrics_baseline = None\nif test_labels is not None:\n    test_metrics_baseline = evaluate_probability_map(test_prob_baseline, gt_test, best_threshold)\n    print(\"\\nBASELINE TEST METRICS:\", test_metrics_baseline)\n\n\n# ============================================================\n# 24. INFERENCE B (PART A OF YOUR ASK): CHEAP HISTOGRAM-MATCHING\n#     DIAGNOSTIC -- NO RETRAINING, uses the SAME baseline weights,\n#     just restyles fragment 1's patches toward the train fragments'\n#     intensity distribution before feeding them to the model.\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INFERENCE B: HISTOGRAM-MATCHING DIAGNOSTIC (no retraining)\")\nprint(\"=\" * 70)\n\n# small reference pool built from the TRAIN fragments (source domain)\ntrain_hist_pool = []\nfor fid in CFG.train_frags:\n    v = train_volumes[fid]\n    m = train_masks_full[fid]\n    coords = generate_grid_coords(m, CFG.patch_size, CFG.patch_size, CFG.min_tissue_frac_train)\n    if coords:\n        for (y, x) in random.sample(coords, min(10, len(coords))):\n            train_hist_pool.append(v.read_patch(y, x, CFG.patch_size))\nprint(f\"Built train-domain reference pool for diagnostic: {len(train_hist_pool)} patches\")\n\n\ndef histogram_match_to_train_domain(raw_u8_dhw):\n    if not train_hist_pool:\n        return raw_u8_dhw\n    ref = random.choice(train_hist_pool)\n    img_hwd = np.transpose(raw_u8_dhw, (1, 2, 0))\n    ref_hwd = np.transpose(ref, (1, 2, 0))\n    try:\n        matched = match_histograms(img_hwd, ref_hwd, channel_axis=2).astype(np.uint8)\n        return np.transpose(matched, (2, 0, 1))\n    except Exception:\n        return raw_u8_dhw\n\n\ntest_prob_histmatch = sliding_window_inference_v2(\n    model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch,\n    pre_transform=histogram_match_to_train_domain)\n\ntest_metrics_histmatch = None\nif test_labels is not None:\n    test_metrics_histmatch = evaluate_probability_map(test_prob_histmatch, gt_test, best_threshold)\n    print(\"\\nHISTOGRAM-MATCHING DIAGNOSTIC TEST METRICS:\", test_metrics_histmatch)\n    if test_metrics_baseline is not None:\n        delta = test_metrics_histmatch[\"dice\"] - test_metrics_baseline[\"dice\"]\n        print(f\"\\n>>> Histogram matching alone changed local test Dice by {delta:+.4f} \"\n              f\"(baseline {test_metrics_baseline['dice']:.4f} -> \"\n              f\"{test_metrics_histmatch['dice']:.4f})\")\n        if abs(delta) < 0.02:\n            print(\">>> Small effect: the gap looks more structural than pure appearance/\"\n                  \"calibration shift -- the training-time techniques below (DANN, MixStyle, \"\n                  \"stronger backbone, augmentation diversity) matter more than appearance \"\n                  \"alignment alone.\")\n        else:\n            print(\">>> Meaningful effect: appearance/calibration shift is a real contributor -- \"\n                  \"the CLAHE / per-fragment-normalization / histogram-match-augmentation training \"\n                  \"changes in this script are well-targeted at exactly this.\")\n\n\n# ============================================================\n# 25. INFERENCE C: ADABN + TTA\n# ============================================================\n\ntest_prob_adabn = None\ntest_metrics_adabn = None\n\nif CFG.use_adabn:\n    print(\"\\n\" + \"=\" * 70)\n    print(\"INFERENCE C: ADABN + TTA\")\n    print(\"=\" * 70)\n    model = recalibrate_batchnorm_v2(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                                      CFG.device, CFG.adabn_max_patches, CFG.infer_batch)\n    test_prob_adabn = sliding_window_inference_v2(\n        model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device, CFG.infer_batch)\n    if test_labels is not None:\n        test_metrics_adabn = evaluate_probability_map(test_prob_adabn, gt_test, best_threshold)\n        print(\"\\nADABN TEST METRICS:\", test_metrics_adabn)\n\n\n# ============================================================\n# 26. CHOOSE FINAL PROBABILITY MAP\n# ============================================================\n\ntest_prob = test_prob_adabn if (CFG.use_adabn and test_prob_adabn is not None) else test_prob_baseline\n\notsu_threshold = compute_otsu_threshold(test_prob, test_mask, fallback=best_threshold)\nprint(\"\\nValidation-tuned threshold:\", best_threshold)\nprint(\"Unsupervised Otsu threshold:\", otsu_threshold)\n\nfinal_threshold = best_threshold\ntest_pred_bin = postprocess_v2(test_prob, final_threshold)\n\ntest_metrics_raw = test_metrics_otsu = test_metrics_post = None\nif test_labels is not None:\n    test_metrics_raw = evaluate_probability_map(test_prob, gt_test, final_threshold)\n    test_metrics_otsu = evaluate_probability_map(test_prob, gt_test, otsu_threshold)\n    post_preds = test_pred_bin.astype(np.float32)\n    tp = (post_preds * gt_test).sum()\n    fp = (post_preds * (1.0 - gt_test)).sum()\n    fn = ((1.0 - post_preds) * gt_test).sum()\n    test_metrics_post = metrics_from_counts(tp, fp, fn)\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LOCAL TEST DIAGNOSTICS SUMMARY\")\n    print(\"=\" * 70)\n    print(\"A) Baseline + val threshold:            \", test_metrics_baseline)\n    print(\"B) Histogram-matched (no retrain) + val threshold:\", test_metrics_histmatch)\n    print(\"C) AdaBN + val threshold:                \", test_metrics_adabn)\n    print(\"Final (chosen) raw + val threshold:      \", test_metrics_raw)\n    print(\"Final (chosen) raw + Otsu threshold:      \", test_metrics_otsu)\n    print(\"Final (chosen) postprocessed + val threshold:\", test_metrics_post)\n\n\n# ============================================================\n# 27. SAVE OUTPUTS\n# ============================================================\n\nprob_path = os.path.join(CFG.out_dir, \"fragment1_probability_v3.npy\")\nnp.save(prob_path, test_prob)\nprint(\"\\nSaved probability map:\", prob_path)\n\npred_path = os.path.join(CFG.out_dir, \"fragment1_prediction_v3.png\")\ncv2.imwrite(pred_path, (test_pred_bin * 255).astype(np.uint8))\nprint(\"Saved prediction:\", pred_path)\n\nmetrics_summary = {\n    \"best_validation_dice_at_0.50\": best_val_dice,\n    \"best_validation_threshold\": best_threshold,\n    \"validation_dice_at_best_threshold\": threshold_dice,\n    \"otsu_threshold\": otsu_threshold,\n    \"final_threshold\": final_threshold,\n    \"final_validation_metrics\": final_val_metrics,\n    \"test_baseline\": test_metrics_baseline,\n    \"test_histogram_matched_diagnostic_no_retrain\": test_metrics_histmatch,\n    \"test_adabn\": test_metrics_adabn,\n    \"test_raw\": test_metrics_raw,\n    \"test_otsu\": test_metrics_otsu,\n    \"test_postprocessed\": test_metrics_post,\n    \"config\": cfg_to_dict(CFG),\n    \"history\": history,\n}\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump(metrics_summary, f, indent=2)\nprint(\"Saved metrics:\", CFG.metrics_path)\n\n\n# ============================================================\n# 28. VISUALIZATION\n# ============================================================\n\ndef save_full_overview():\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_path = os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n    mid_slice = tifffile.imread(mid_path)\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    prob_small = cv2.resize((test_prob * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_AREA)\n\n    if test_labels is not None:\n        gt_small = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1],\n                               interpolation=cv2.INTER_NEAREST)\n        fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n        axes[2].imshow(prob_small, cmap=\"gray\"); axes[2].set_title(f\"Probability thr={final_threshold:.2f}\")\n        axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"Final Prediction\")\n    else:\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n        axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n        axes[1].imshow(prob_small, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n        axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"Final Prediction\")\n\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    overview_path = os.path.join(CFG.viz_dir, \"fragment1_v3_overview.png\")\n    plt.savefig(overview_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(\"Saved overview:\", overview_path)\n\n\nsave_full_overview()\n\n\ndef save_patch_comparisons(n=6):\n    coords = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(coords)\n    coords = coords[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(coords):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n        prob_patch = test_prob[y:y + size, x:x + size]\n\n        if test_labels is not None:\n            gt_patch = test_labels[y:y + size, x:x + size]\n            fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground Truth\")\n            axes[2].imshow(prob_patch, cmap=\"gray\"); axes[2].set_title(\"Probability\")\n            axes[3].imshow(pred_patch, cmap=\"gray\"); axes[3].set_title(\"Prediction\")\n        else:\n            fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n            axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n            axes[1].imshow(prob_patch, cmap=\"gray\"); axes[1].set_title(\"Probability\")\n            axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n\n        for ax in axes:\n            ax.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n        plt.close(fig)\n    print(f\"Saved {len(coords)} patch comparisons.\")\n\n\nsave_patch_comparisons(n=6)\n\n\ndef save_training_curves():\n    epochs_axis = np.arange(1, len(history[\"train_loss\"]) + 1)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_loss\"], label=\"Train Loss\")\n    plt.plot(epochs_axis, history[\"val_loss\"], label=\"Val Loss\")\n    if CFG.USE_DANN:\n        plt.plot(epochs_axis, history[\"dann_loss\"], label=\"DANN domain loss\", linestyle=\"--\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Loss\"); plt.title(\"V3 Training Loss\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_loss.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(epochs_axis, history[\"train_dice\"], label=\"Train Dice\")\n    plt.plot(epochs_axis, history[\"val_dice\"], label=\"Val Dice (EMA)\" if CFG.USE_EMA else \"Val Dice\")\n    plt.xlabel(\"Epoch\"); plt.ylabel(\"Dice\"); plt.title(\"V3 Dice\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"training_dice.png\"), dpi=150)\n    plt.close(fig)\n\n    fig = plt.figure(figsize=(8, 5))\n    plt.plot(threshold_grid, threshold_scores)\n    plt.axvline(best_threshold, linestyle=\"--\", label=f\"best={best_threshold:.2f}\")\n    plt.xlabel(\"Threshold\"); plt.ylabel(\"Validation Dice\"); plt.title(\"Dice Threshold Search\")\n    plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, \"threshold_search.png\"), dpi=150)\n    plt.close(fig)\n\n    print(\"Saved training curves.\")\n\n\nsave_training_curves()\n\n\n# ============================================================\n# 29. CLEANUP\n# ============================================================\n\ntest_vol.close()\nfor v in train_volumes.values():\n    v.close()\ncleanup_memory()\n\n\n# ============================================================\n# 30. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V3 COMPLETE\")\nprint(\"=\" * 70)\nprint(f\"Best validation Dice @ 0.50: {best_val_dice:.5f}\")\nprint(f\"Best Dice threshold: {best_threshold:.2f}\")\nprint(f\"Validation Dice @ optimized threshold: {threshold_dice:.5f}\")\nprint(f\"Backbone: {CFG.encoder_name} / architecture={CFG.architecture}\")\nprint(f\"EMA: {CFG.USE_EMA} | DANN: {CFG.USE_DANN} | MixStyle: {CFG.USE_MIXSTYLE}\")\nprint(f\"CLAHE: {CFG.USE_CLAHE} | Normalization: {CFG.NORMALIZATION_MODE} | \"\n      f\"Mask erosion: {CFG.mask_erode_px}px\")\nif test_metrics_baseline is not None:\n    print(f\"\\nLocal test Dice -- baseline: {test_metrics_baseline['dice']:.5f}\")\nif test_metrics_histmatch is not None:\n    print(f\"Local test Dice -- histogram-matched (diagnostic, no retrain): \"\n          f\"{test_metrics_histmatch['dice']:.5f}\")\nif test_metrics_adabn is not None:\n    print(f\"Local test Dice -- AdaBN: {test_metrics_adabn['dice']:.5f}\")\nif test_metrics_post is not None:\n    print(f\"Local test Dice -- final postprocessed: {test_metrics_post['dice']:.5f}\")\nprint(f\"\\nBest checkpoint: {CFG.ckpt_path}\")\nprint(f\"Probability map: {prob_path}\")\nprint(f\"Prediction: {pred_path}\")\nprint(f\"Metrics: {CFG.metrics_path}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")\nprint(\"\\n=== V3 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-06T07:39:17.52634Z","iopub.execute_input":"2026-09-06T07:39:17.527254Z","iopub.status.idle":"2026-09-06T07:40:31.130805Z","shell.execute_reply.started":"2026-09-06T07:39:17.527216Z","shell.execute_reply":"2026-09-06T07:40:31.129657Z"}},"outputs":[{"name":"stdout","text":"======================================================================\nBUILDING TRAIN DATA\n======================================================================\n\nProcessing fragment 2 ...\n  fragment 2 normalization stats: mean=122.54 std=51.26\nFragment 2: 24450 candidate patches (post mask-erosion)\n  spatial train=18466 validation=5099\n\nProcessing fragment 3 ...\n  fragment 3 normalization stats: mean=123.01 std=57.22\nFragment 3: 6348 candidate patches (post mask-erosion)\n  spatial train=5240 validation=649\n\nRaw train samples: 23706\nValidation samples: 5748\nPositive patches: 13144 | Negative patches: 10562\nBalanced dataset: 23706 | positive ratio=0.554\n\n======================================================================\nLOADING TEST FRAGMENT (unlabeled: mask + volume only, for hist-match pool / DANN)\n======================================================================\n  fragment 1 normalization stats: mean=122.80 std=59.67\nFragment 1: 7425 unlabeled candidate patches\nBuilt histogram-matching reference pool: 20 patches (26.1 MB)\n\nBuilding V3 model ...\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/albumentations/core/validation.py:114: UserWarning: ShiftScaleRotate is a special case of Affine transform. Please use Affine transform instead.\n  original_init(self, **validated_kwargs)\n/usr/local/lib/python3.12/dist-packages/torch/nn/modules/conv.py:186: UserWarning: Initializing zero-element tensors is a no-op\n  init.kaiming_uniform_(self.weight, a=math.sqrt(5))\n/usr/local/lib/python3.12/dist-packages/segmentation_models_pytorch/base/initialization.py:7: UserWarning: Initializing zero-element tensors is a no-op\n  nn.init.kaiming_uniform_(m.weight, mode=\"fan_in\", nonlinearity=\"relu\")\n","output_type":"stream"},{"name":"stdout","text":"\nEstimating positive-pixel fraction ...\nEstimated positive fraction: 0.171812\nOutput bias initialized to -1.5728\nBCE positive weight: 2.196\nEncoder parameters: 178 tensors | Decoder/head parameters: 200 tensors\n","output_type":"stream"},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mRuntimeError\u001b[0m                              Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_58/2271209446.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m   1111\u001b[0m         \u001b[0mdummy\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mzeros\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0min_channels\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpatch_size\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpatch_size\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdevice\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdevice\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1112\u001b[0m         \u001b[0mmodel\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0meval\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1113\u001b[0;31m         \u001b[0m_\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmodel\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdummy\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1114\u001b[0m         \u001b[0mdeepest_ch\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfeature_capture\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mfeatures\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m-\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mshape\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1115\u001b[0m         \u001b[0mmodel\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtrain\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/segmentation_models_pytorch/base/model.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, x)\u001b[0m\n\u001b[1;32m     47\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     48\u001b[0m         \u001b[0mfeatures\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mencoder\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 49\u001b[0;31m         \u001b[0mdecoder_output\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdecoder\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0mfeatures\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     50\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     51\u001b[0m         \u001b[0mmasks\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msegmentation_head\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdecoder_output\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/segmentation_models_pytorch/decoders/unetplusplus/decoder.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, *features)\u001b[0m\n\u001b[1;32m    134\u001b[0m             \u001b[0;32mfor\u001b[0m \u001b[0mdepth_idx\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mrange\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdepth\u001b[0m \u001b[0;34m-\u001b[0m \u001b[0mlayer_idx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    135\u001b[0m                 \u001b[0;32mif\u001b[0m \u001b[0mlayer_idx\u001b[0m \u001b[0;34m==\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 136\u001b[0;31m                     output = self.blocks[f\"x_{depth_idx}_{depth_idx}\"](\n\u001b[0m\u001b[1;32m    137\u001b[0m                         \u001b[0mfeatures\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mdepth_idx\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mfeatures\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mdepth_idx\u001b[0m \u001b[0;34m+\u001b[0m \u001b[0;36m1\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    138\u001b[0m                     )\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/segmentation_models_pytorch/decoders/unetplusplus/decoder.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, x, skip)\u001b[0m\n\u001b[1;32m     40\u001b[0m             \u001b[0mx\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcat\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mskip\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdim\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     41\u001b[0m             \u001b[0mx\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mattention1\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 42\u001b[0;31m         \u001b[0mx\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mconv1\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     43\u001b[0m         \u001b[0mx\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mconv2\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     44\u001b[0m         \u001b[0mx\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mattention2\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/container.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, input)\u001b[0m\n\u001b[1;32m    251\u001b[0m         \"\"\"\n\u001b[1;32m    252\u001b[0m         \u001b[0;32mfor\u001b[0m \u001b[0mmodule\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 253\u001b[0;31m             \u001b[0minput\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmodule\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0minput\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    254\u001b[0m         \u001b[0;32mreturn\u001b[0m \u001b[0minput\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    255\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/conv.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, input)\u001b[0m\n\u001b[1;32m    551\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    552\u001b[0m     \u001b[0;32mdef\u001b[0m \u001b[0mforward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0minput\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mTensor\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m->\u001b[0m \u001b[0mTensor\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 553\u001b[0;31m         \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_conv_forward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0minput\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mweight\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mbias\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    554\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    555\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/conv.py\u001b[0m in \u001b[0;36m_conv_forward\u001b[0;34m(self, input, weight, bias)\u001b[0m\n\u001b[1;32m    546\u001b[0m             )\n\u001b[1;32m    547\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 548\u001b[0;31m         return F.conv2d(\n\u001b[0m\u001b[1;32m    549\u001b[0m             \u001b[0minput\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mweight\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mbias\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstride\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpadding\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdilation\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mgroups\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    550\u001b[0m         )\n","\u001b[0;31mRuntimeError\u001b[0m: Given groups=1, expected weight to be at least 1 at dimension 0, but got weight of size [0, 96, 3, 3] instead"],"ename":"RuntimeError","evalue":"Given groups=1, expected weight to be at least 1 at dimension 0, but got weight of size [0, 96, 3, 3] instead","output_type":"error"}],"execution_count":2},{"cell_type":"code","source":"print(\"CFG decoder_channels:\", CFG.decoder_channels)\nprint(\"CFG in_channels:\", CFG.in_channels)\nprint(\"CFG patch_size:\", CFG.patch_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-06T07:44:12.449189Z","iopub.execute_input":"2026-09-06T07:44:12.449448Z","iopub.status.idle":"2026-09-06T07:44:12.456291Z","shell.execute_reply.started":"2026-09-06T07:44:12.449425Z","shell.execute_reply":"2026-09-06T07:44:12.455229Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mAttributeError\u001b[0m                            Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_58/365694222.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[0;32m----> 1\u001b[0;31m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"CFG decoder_channels:\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdecoder_channels\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m      2\u001b[0m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"CFG in_channels:\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0min_channels\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      3\u001b[0m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"CFG patch_size:\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mCFG\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpatch_size\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mAttributeError\u001b[0m: type object 'CFG' has no attribute 'decoder_channels'"],"ename":"AttributeError","evalue":"type object 'CFG' has no attribute 'decoder_channels'","output_type":"error"}],"execution_count":3},{"cell_type":"code","source":"import torch\n\nprint(\"PyTorch:\", torch.__version__)\nprint(\"CUDA:\", torch.version.cuda)\nprint(\"CUDA available:\", torch.cuda.is_available())\n\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n    print(\"Capability:\", torch.cuda.get_device_capability(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-06T07:31:11.861943Z","iopub.execute_input":"2026-09-06T07:31:11.862526Z","iopub.status.idle":"2026-09-06T07:31:11.868162Z","shell.execute_reply.started":"2026-09-06T07:31:11.862497Z","shell.execute_reply":"2026-09-06T07:31:11.867401Z"}},"outputs":[{"name":"stdout","text":"PyTorch: 2.10.0+cu128\nCUDA: 12.8\nCUDA available: True\nGPU: Tesla P100-PCIE-16GB\nCapability: (6, 0)\n","output_type":"stream"}],"execution_count":8},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS CHALLENGE - INK DETECTION\n# V5 -- targeting the V4 val->test generalization gap\n#\n# V4 DIAGNOSIS (from your own run):\n#   - Validation Dice 0.703, but every independent test-threshold\n#     candidate (Otsu/positive-rate/P98/val-threshold) landed in\n#     the same ~0.37-0.44 dice band on fragment 1. That means the\n#     gap is NOT a thresholding problem -- V4 already proved that\n#     with its own independent calibration. It's a genuine\n#     train/val -> test generalization gap.\n#   - Best epoch was 10/10 (the LAST epoch) -> training hadn't\n#     converged, it just hit the epoch cap.\n#   - Your own Rot90 ablation made things much WORSE (dice 0.11\n#     vs 0.43 without it) -> blind TTA is not the fix here.\n#\n# V5 CHANGES (each one targets a specific piece of that evidence):\n#\n#   1. TRAIN TO CONVERGENCE\n#        epochs 10 -> 22, early_stop_patience 3 -> 6, min_epochs 4 -> 10.\n#        You were leaving real validation Dice on the table.\n#\n#   2. MULTI-MODEL ENSEMBLE (the single most reliable generalization\n#      lever, and literally what top Vesuvius solutions did: average\n#      3-9 varied models). N_ENSEMBLE_MODELS members, varied by seed,\n#      encoder (resnet50 vs se_resnext50_32x4d), and a small Z-slice\n#      offset (depth diversity) -- each trained independently on the\n#      SAME fragment2/3 train/val split, each doing its own AdaBN\n#      recalibration + baseline/AdaBN inference on fragment 1. Final\n#      probability map = mean across members.\n#\n#   3. WEIGHT AVERAGING (\"model soup\" / SWA-style) -- each member\n#      keeps its top-K checkpoints by validation Dice in CPU memory\n#      and averages their weights before AdaBN/inference, instead of\n#      using only the single best epoch. Cheap (no extra training),\n#      well-established to reduce overfitting to late-training noise.\n#\n#   4. ROBUST (PERCENTILE-CLIPPED) NORMALIZATION instead of raw\n#      mean/std z-scoring. Z-score is sensitive to a handful of\n#      extreme bright/dark voxels (scan artifacts, resin boundaries),\n#      whose frequency/severity plausibly differs between fragments --\n#      one concrete, testable source of the intensity-distribution\n#      shift between fragments 2/3 and fragment 1.\n#\n#   5. TTA STAYS OFF BY DEFAULT. Your own Rot90 ablation showed it\n#      hurts on this data; I'm not reintroducing it as a silent\n#      default. The Rot90 ablation machinery is kept (for one\n#      ensemble member, as a diagnostic) in case you want to keep\n#      probing it, but it never feeds into the final prediction.\n#\n# WHAT THIS DOES NOT DO: fragment 1 is never touched during training,\n# and its labels are only ever used for the printed *local diagnostics*\n# -- exactly the same discipline your V4 script already had. I'm not\n# training on fragment 1 to close the gap artificially.\n#\n# HONEST EXPECTATION: I can't promise >=0.60 test Dice -- that's a\n# real, hard, open research outcome on genuinely out-of-distribution\n# data, and even well-resourced Vesuvius Challenge teams worked for\n# months to get strong fragment-transfer results. These are the\n# highest-confidence, evidence-backed changes given what your own\n# diagnostics showed; if the gap persists after this, that itself is\n# useful information (it would point toward the fragments being more\n# fundamentally different than augmentation/ensembling alone can fix).\n#\n# RUNTIME WARNING: this trains N_ENSEMBLE_MODELS full models\n# sequentially, each for up to `epochs` epochs at patch_size=480.\n# That is a LOT more compute than V4's single 10-epoch run. Defaults\n# below (2 members, 22 epochs) are chosen to have a realistic chance\n# of finishing in one Kaggle GPU session -- raise N_ENSEMBLE_MODELS to\n# 3 (add a third seed to `ensemble_seeds`/`ensemble_encoders`/\n# `ensemble_depth_offsets`) only if you have session time to spare, or\n# plan to run members across multiple sessions and average the saved\n# probability maps yourself.\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\nget_ipython().system('pip install -q segmentation-models-pytorch==0.2.0') if 'get_ipython' in dir() else __import__('os').system('pip install -q segmentation-models-pytorch==0.2.0')\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport copy\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\n\nimport segmentation_models_pytorch as smp\n\nfrom skimage.filters import threshold_otsu\n\n\n# ============================================================\n# 2. CONFIGURATION\n# ============================================================\n\nclass CFG:\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    # base depth window; each ensemble member may shift this by its\n    # own small offset (see ensemble_depth_offsets) for depth diversity\n    depth_indices_base = list(range(16, 38))   # 22 slices\n    in_channels = len(depth_indices_base)\n\n    patch_size = 480\n    train_stride = 128\n    test_stride = 128\n\n    batch_size = 8\n    infer_batch = 8\n    num_workers = 2\n    drop_last = True\n\n    # --- V5: train to convergence (V4's best epoch was its LAST epoch) ---\n    epochs = 8\n    early_stop_patience = 3\n    min_epochs = 10\n\n    encoder_lr = 5e-5\n    decoder_lr = 2e-4\n    weight_decay = 1e-4\n    grad_clip = 1.0\n    use_amp = True\n\n    val_fraction = 0.20\n    validation_block_size = 768\n\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n\n    positive_ratio = 0.35\n    hard_negative_ratio = 0.15\n    normal_ratio = 0.50\n    positive_patch_fraction = 0.001\n    hard_negative_tissue_fraction = 0.45\n    max_train_samples = 18000\n\n    decoder_attention_type = \"scse\"\n\n    bce = 0.45\n    dice = 0.40\n    tversky = 0.15\n    tversky_alpha = 0.35\n    tversky_beta = 0.65\n    tversky_gamma = 1.0\n\n    rotate90_probability = 0.75\n\n    use_adabn = True\n    adabn_max_patches = 1200\n\n    val_thresholds = np.arange(0.20, 0.801, 0.01)\n    test_threshold_min = 0.15\n    test_threshold_max = 0.75\n    final_test_threshold_method = \"otsu\"\n\n    # TTA stays OFF for the final ensembled prediction -- V4's own Rot90\n    # ablation made things worse. Kept available as a per-run diagnostic only.\n    use_tta_for_final = False\n    run_rot90_ablation_on_first_member = True\n\n    use_postprocess = True\n    closing_kernel = 3\n    min_component_area = 8\n\n    # --- V5: robust (percentile-clipped) normalization ---\n    USE_ROBUST_NORM = True\n    norm_lo_pct = 1.0\n    norm_hi_pct = 99.0\n\n    # --- V5: ensemble ---\n    N_ENSEMBLE_MODELS = 2                                   # see runtime warning above\n    ensemble_seeds          = [42, 123, 2024][:3]\n    ensemble_encoders       = [\"resnet50\", \"se_resnext50_32x4d\", \"resnet50\"][:3]\n    ensemble_depth_offsets  = [0, 0, 3][:3]                  # small Z-shift for diversity\n\n    # --- V5: weight averaging (\"model soup\") ---\n    USE_WEIGHT_AVERAGING = True\n    weight_avg_top_k = 3   # average the top-K checkpoints by val Dice, per member\n\n    seed = 42\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    out_dir = \"/kaggle/working/vesuvius_v5\"\n    viz_dir = os.path.join(out_dir, \"visualizations\")\n    metrics_path = os.path.join(out_dir, \"v5_metrics_summary.json\")\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\ntorch.backends.cudnn.benchmark = True\n\n\ndef cleanup_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\n# ============================================================\n# 3. TISSUE MASK / LABELS / PATCH GRID  (unchanged from V4)\n# ============================================================\n\ndef load_tissue_mask(frag_dir, depth_indices):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = depth_indices[len(depth_indices) // 2]\n        mid_path = os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\")\n        mid = tifffile.imread(mid_path)\n        threshold = mid.mean() * 0.15\n        mask = (mid > threshold).astype(np.uint8) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    if not os.path.exists(path):\n        return None\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    if lbl is None:\n        return None\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    ymax = max(H - patch_size, 0)\n    xmax = max(W - patch_size, 0)\n    ys = list(range(0, ymax + 1, stride))\n    xs = list(range(0, xmax + 1, stride))\n    if not ys or ys[-1] != ymax:\n        ys.append(ymax)\n    if not xs or xs[-1] != xmax:\n        xs.append(xmax)\n    coords = []\n    for y in ys:\n        for x in xs:\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac >= min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access to a fragment's surface volume, restricted to a chosen\n    subset of depth slices. Different ensemble members can pass different (Z-shifted)\n    depth_indices to read a slightly different depth window of the SAME underlying TIFFs.\"\"\"\n\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in depth_indices]\n        self._slices = None\n        self._h = None\n        self._w = None\n\n    def _ensure_open(self):\n        if self._slices is not None:\n            return\n        slices = []\n        for path in self.paths:\n            try:\n                arr = tifffile.memmap(path, mode=\"r\")\n            except Exception:\n                arr = tifffile.imread(path)\n            slices.append(arr)\n        self._slices = slices\n        self._h = slices[0].shape[0]\n        self._w = slices[0].shape[1]\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.dtype != np.uint8:\n                if np.issubdtype(block.dtype, np.integer):\n                    max_value = np.iinfo(block.dtype).max\n                else:\n                    max_value = float(np.nanmax(block)) + 1e-6\n                block = (block.astype(np.float32) / max_value * 255.0)\n                block = np.clip(block, 0, 255).astype(np.uint8)\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        cleanup_memory()\n\n\ndef clip_depth_indices(base_indices, offset, max_slice_idx=64):\n    \"\"\"Shift the base depth window by `offset` slices, clamped to a valid TIFF index range,\n    used to give ensemble members slightly different Z-depth diversity.\"\"\"\n    shifted = [i + offset for i in base_indices]\n    lo, hi = min(shifted), max(shifted)\n    if lo < 0:\n        shifted = [i - lo for i in shifted]\n    elif hi > max_slice_idx:\n        shifted = [i - (hi - max_slice_idx) for i in shifted]\n    return shifted\n\n\n# ============================================================\n# 4. NORMALIZATION -- V5: robust percentile-clipped option\n# ============================================================\n\ndef normalize_patch_robust(img_float, lo_pct=CFG.norm_lo_pct, hi_pct=CFG.norm_hi_pct):\n    \"\"\"Percentile-clip then rescale to [-1, 1], instead of raw mean/std z-scoring. Z-score\n    is sensitive to a handful of extreme bright/dark voxels (scan artifacts, resin\n    boundaries); percentile clipping is far less sensitive to those outliers and puts\n    different fragments' intensity distributions on a more comparable footing.\"\"\"\n    lo = np.percentile(img_float, lo_pct)\n    hi = np.percentile(img_float, hi_pct)\n    if hi - lo < 1e-6:\n        hi = lo + 1e-6\n    x = np.clip(img_float, lo, hi)\n    x = (x - lo) / (hi - lo)\n    x = x * 2.0 - 1.0\n    return x.astype(np.float32)\n\n\ndef normalize_patch_zscore(img_float):\n    x = img_float.astype(np.float32)\n    mean = x.mean()\n    std = x.std() + 1e-6\n    return (x - mean) / std\n\n\ndef normalize_patch(img_float):\n    if CFG.USE_ROBUST_NORM:\n        return normalize_patch_robust(img_float)\n    return normalize_patch_zscore(img_float)\n\n\n# ============================================================\n# 5. AUGMENTATION (unchanged from V4)\n# ============================================================\n\ndef build_train_transform():\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=CFG.rotate90_probability),\n        A.Transpose(p=0.25),\n        A.ShiftScaleRotate(shift_limit=0.025, scale_limit=0.06, rotate_limit=10,\n                            border_mode=cv2.BORDER_REFLECT_101, p=0.30),\n        A.RandomBrightnessContrast(brightness_limit=0.06, contrast_limit=0.08, p=0.20),\n    ])\n\n\n# ============================================================\n# 6. DATASET (unchanged from V4)\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        patch = self.volumes[fid].read_patch(y, x, self.patch_size)\n        label = self.labels[fid][y:y + self.patch_size, x:x + self.patch_size]\n\n        img = np.transpose(patch, (1, 2, 0))\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label > 0).astype(np.float32)[None, ...]\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ============================================================\n# 7. MODEL -- V5: encoder is now parameterized per ensemble member\n# ============================================================\n\ndef build_model(encoder_name):\n    return smp.Unet(\n        encoder_name=encoder_name,\n        encoder_weights=\"imagenet\",\n        in_channels=CFG.in_channels,\n        classes=1,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n\n\n# ============================================================\n# 8. LOSSES (unchanged from V4)\n# ============================================================\n\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(dim=1)\n    union = probs.sum(dim=1) + t.sum(dim=1)\n    return 1 - ((2 * inter + eps) / (union + eps)).mean()\n\n\ndef focal_tversky_loss(logits, targets, alpha, beta, gamma, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    tp = (probs * t).sum(dim=1)\n    fp = (probs * (1 - t)).sum(dim=1)\n    fn = ((1 - probs) * t).sum(dim=1)\n    tversky = (tp + eps) / (tp + alpha * fp + beta * fn + eps)\n    return (1 - tversky).clamp_min(0).pow(gamma).mean()\n\n\nclass V4Loss(nn.Module):\n    def __init__(self, pos_weight):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n\n    def forward(self, logits, targets):\n        bce = self.bce(logits, targets)\n        dice = dice_loss(logits, targets)\n        tv = focal_tversky_loss(logits, targets, CFG.tversky_alpha, CFG.tversky_beta, CFG.tversky_gamma)\n        return CFG.bce * bce + CFG.dice * dice + CFG.tversky * tv\n\n\n# ============================================================\n# 9. METRICS (unchanged from V4)\n# ============================================================\n\ndef calculate_metrics(tp, fp, fn):\n    dice = (2 * tp + 1e-6) / (2 * tp + fp + fn + 1e-6)\n    iou = (tp + 1e-6) / (tp + fp + fn + 1e-6)\n    precision = (tp + 1e-6) / (tp + fp + 1e-6)\n    recall = (tp + 1e-6) / (tp + fn + 1e-6)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"precision\": float(precision), \"recall\": float(recall)}\n\n\ndef run_epoch(model, loader, criterion, optimizer=None, scaler=None):\n    training = optimizer is not None\n    model.train(training)\n    total_loss = 0.0\n    tp = fp = fn = 0\n\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        if training:\n            optimizer.zero_grad(set_to_none=True)\n\n        with torch.set_grad_enabled(training):\n            with autocast(enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n\n            if training:\n                if scaler is not None and scaler.is_enabled():\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                    optimizer.step()\n\n        probs = torch.sigmoid(logits.detach())\n        preds = probs > 0.5\n        gt = masks > 0.5\n        tp += int((preds & gt).sum().item())\n        fp += int((preds & (~gt)).sum().item())\n        fn += int(((~preds) & gt).sum().item())\n        total_loss += float(loss.item())\n\n        del imgs, masks, logits, probs, preds\n\n    metrics = calculate_metrics(tp, fp, fn)\n    metrics[\"loss\"] = total_loss / max(len(loader), 1)\n    cleanup_memory()\n    return metrics\n\n\n# ============================================================\n# 10. SHARED DATA: masks/labels/split/sampling pools computed ONCE\n#     (these depend only on (y, x) coordinates and labels, not on\n#     which depth slices a given ensemble member reads)\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"BUILDING SHARED TRAIN/VAL PATCH GRID (fragments 2 & 3)\")\nprint(\"=\" * 70)\n\n_shared_masks, _shared_labels = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir, CFG.depth_indices_base)\n    labels = load_ink_labels(frag_dir)\n    if labels is None:\n        raise FileNotFoundError(f\"inklabels.png missing for fragment {fid}\")\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"Fragment {fid}: {mask.shape} {len(coords)} patches\")\n\n    _shared_masks[fid] = mask\n    _shared_labels[fid] = labels\n    all_samples.extend([(fid, y, x) for y, x in coords])\n\nprint(\"Total candidate patches:\", len(all_samples))\n\n# spatial train/val split\ngroups = {}\nfor fid, y, x in all_samples:\n    key = (fid, y // CFG.validation_block_size, x // CFG.validation_block_size)\n    groups.setdefault(key, []).append((fid, y, x))\n\nkeys = list(groups.keys())\nrng = random.Random(CFG.seed)\nrng.shuffle(keys)\nn_val_groups = max(1, int(len(keys) * CFG.val_fraction))\nval_keys = set(keys[:n_val_groups])\n\ntrain_grid, val_samples = [], []\nfor key, items in groups.items():\n    (val_samples if key in val_keys else train_grid).extend(items)\n\nprint(\"Spatial train:\", len(train_grid))\nprint(\"Spatial val:\", len(val_samples))\n\n# sampling pools\npositive_pool, hard_negative_pool, normal_pool = [], [], []\nfor fid, y, x in train_grid:\n    lbl = _shared_labels[fid][y:y + CFG.patch_size, x:x + CFG.patch_size]\n    tissue = _shared_masks[fid][y:y + CFG.patch_size, x:x + CFG.patch_size]\n    ink_fraction = float(lbl.mean())\n    tissue_fraction = float(tissue.mean())\n    if ink_fraction >= CFG.positive_patch_fraction:\n        positive_pool.append((fid, y, x))\n    elif tissue_fraction >= CFG.hard_negative_tissue_fraction:\n        hard_negative_pool.append((fid, y, x))\n    else:\n        normal_pool.append((fid, y, x))\n\nprint(\"\\nSampling pools -- positive:\", len(positive_pool),\n      \"hard_negative:\", len(hard_negative_pool), \"normal:\", len(normal_pool))\n\n\ndef sample_pool(pool, n):\n    if len(pool) == 0:\n        return []\n    if len(pool) <= n:\n        return pool.copy()\n    return random.sample(pool, n)\n\n\ndef build_train_samples(seed):\n    \"\"\"Rebuilt per-member (with that member's own seed) so members see different sampled\n    subsets of the shared pools -- another small source of ensemble diversity.\"\"\"\n    r = random.Random(seed)\n    N = min(CFG.max_train_samples, len(train_grid))\n    n_positive = int(N * CFG.positive_ratio)\n    n_hard = int(N * CFG.hard_negative_ratio)\n    n_normal = N - n_positive - n_hard\n\n    samples = []\n    samples.extend(r.sample(positive_pool, min(n_positive, len(positive_pool))) if positive_pool else [])\n    samples.extend(r.sample(hard_negative_pool, min(n_hard, len(hard_negative_pool))) if hard_negative_pool else [])\n    samples.extend(r.sample(normal_pool, min(n_normal, len(normal_pool))) if normal_pool else [])\n\n    if len(samples) < N:\n        selected = set(samples)\n        remaining = [s for s in train_grid if s not in selected]\n        need = min(N - len(samples), len(remaining))\n        if need > 0:\n            samples.extend(r.sample(remaining, need))\n    r.shuffle(samples)\n    return samples\n\n\n# ============================================================\n# 11. TEST FRAGMENT METADATA (mask/labels shared; volume reopened\n#     per member with that member's own depth window)\n# ============================================================\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir, CFG.depth_indices_base)\ntest_labels = load_ink_labels(test_dir)   # local diagnostics ONLY -- never used to pick a threshold\nprint(\"\\nTest shape:\", test_mask.shape, \"| local test GT available:\", test_labels is not None)\n\n\n# ============================================================\n# 12. INFERENCE HELPERS (shared across members)\n# ============================================================\n\ndef gaussian_window(size, sigma_fraction=0.42):\n    axis = np.arange(size, dtype=np.float32) - (size - 1) / 2\n    sigma = max(size * sigma_fraction, 1.0)\n    g = np.exp(-(axis ** 2) / (2 * sigma ** 2))\n    window = np.outer(g, g).astype(np.float32)\n    window /= max(float(window.max()), 1e-6)\n    window = np.maximum(window, 0.05)\n    return window.astype(np.float32)\n\n\ndef prepare_input(raw):\n    raw = raw.astype(np.float32) / 255.0\n    return np.ascontiguousarray(normalize_patch(raw))\n\n\n@torch.no_grad()\ndef predict_batch(model, batch):\n    inp = torch.from_numpy(np.stack(batch)).to(CFG.device, non_blocking=True)\n    with autocast(enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n        logits = model(inp)\n        probs = torch.sigmoid(logits)\n    result = probs.float().cpu().numpy()[:, 0]\n    del inp, logits, probs\n    return result\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, volume, tissue_mask, stride=None, transform=None):\n    if stride is None:\n        stride = CFG.test_stride\n    H, W = tissue_mask.shape\n    prediction_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    window = gaussian_window(CFG.patch_size)\n    coords = generate_grid_coords(tissue_mask, CFG.patch_size, stride, CFG.min_tissue_frac_test)\n    print(\"Inference patches:\", len(coords))\n\n    model.eval()\n    batch_images, batch_coords = [], []\n\n    def flush():\n        if not batch_images:\n            return\n        probabilities = predict_batch(model, batch_images)\n        for p, (y, x) in zip(probabilities, batch_coords):\n            if transform is not None:\n                p = transform(p)\n            prediction_sum[y:y + CFG.patch_size, x:x + CFG.patch_size] += p * window\n            weight_sum[y:y + CFG.patch_size, x:x + CFG.patch_size] += window\n        batch_images.clear()\n        batch_coords.clear()\n        cleanup_memory()\n\n    for y, x in coords:\n        raw = volume.read_patch(y, x, CFG.patch_size)\n        image = prepare_input(raw)\n        batch_images.append(image)\n        batch_coords.append((y, x))\n        if len(batch_images) >= CFG.infer_batch:\n            flush()\n    flush()\n\n    weight_sum = np.maximum(weight_sum, 1e-6)\n    probability = prediction_sum / weight_sum\n    probability[tissue_mask == 0] = 0.0\n    return probability.astype(np.float32)\n\n\n@torch.no_grad()\ndef run_adabn(model, test_vol):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STARTING AdaBN\")\n    print(\"=\" * 70)\n\n    bn_layers = 0\n    for module in model.modules():\n        if isinstance(module, nn.BatchNorm2d):\n            module.reset_running_stats()\n            module.momentum = None\n            bn_layers += 1\n    print(\"BN layers:\", bn_layers)\n\n    coords = generate_grid_coords(test_mask, CFG.patch_size, CFG.test_stride, CFG.min_tissue_frac_test)\n    r = random.Random(CFG.seed)\n    r.shuffle(coords)\n    coords = coords[:CFG.adabn_max_patches]\n    print(\"AdaBN patches:\", len(coords))\n\n    model.train()\n    for start in range(0, len(coords), CFG.infer_batch):\n        batch_coords = coords[start:start + CFG.infer_batch]\n        images = [prepare_input(test_vol.read_patch(y, x, CFG.patch_size)) for y, x in batch_coords]\n        inp = torch.from_numpy(np.stack(images)).to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n            model(inp)\n        del inp, images\n    model.eval()\n    cleanup_memory()\n    print(\"AdaBN finished.\")\n\n\ndef local_metrics(probability, threshold):\n    if test_labels is None:\n        return None\n    gt = (test_labels > 0) & (test_mask > 0)\n    pred = (probability >= threshold) & (test_mask > 0)\n    tp = np.logical_and(pred, gt).sum()\n    fp = np.logical_and(pred, ~gt).sum()\n    fn = np.logical_and(~pred, gt).sum()\n    return calculate_metrics(tp, fp, fn)\n\n\ndef compute_otsu_threshold(probability, mask, fallback=0.5):\n    values = probability[mask > 0]\n    values = values[np.isfinite(values) & (values > 1e-6)]\n    try:\n        thr = float(threshold_otsu(values))\n    except Exception:\n        thr = fallback\n    return float(np.clip(thr, CFG.test_threshold_min, CFG.test_threshold_max))\n\n\ndef remove_small_components(binary, min_area):\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary.astype(np.uint8), connectivity=8)\n    if num_labels <= 1:\n        return binary.astype(np.uint8)\n    output = np.zeros_like(binary, dtype=np.uint8)\n    for label_id in range(1, num_labels):\n        if int(stats[label_id, cv2.CC_STAT_AREA]) >= min_area:\n            output[labels == label_id] = 1\n    return output\n\n\ndef postprocess(probability, threshold, tissue_mask):\n    binary = (probability >= threshold).astype(np.uint8)\n    binary[tissue_mask == 0] = 0\n    if not CFG.use_postprocess:\n        return binary\n    kernel = np.ones((CFG.closing_kernel, CFG.closing_kernel), dtype=np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = remove_small_components(binary, CFG.min_component_area)\n    binary[tissue_mask == 0] = 0\n    return binary.astype(np.uint8)\n\n\n# ============================================================\n# 13. TRAIN + INFER ONE ENSEMBLE MEMBER\n# ============================================================\n\ndef average_state_dicts(state_dicts):\n    \"\"\"Element-wise mean of a list of CPU state_dicts ('model soup'). All members share the\n    same architecture per model instance, so keys always match.\"\"\"\n    avg = copy.deepcopy(state_dicts[0])\n    for key in avg:\n        if avg[key].dtype.is_floating_point:\n            stacked = torch.stack([sd[key].float() for sd in state_dicts], dim=0)\n            avg[key] = stacked.mean(dim=0).to(state_dicts[0][key].dtype)\n        # non-float buffers (e.g. BatchNorm num_batches_tracked) -> just keep the best one's value\n    return avg\n\n\ndef train_and_infer_one_member(member_idx, seed, encoder_name, depth_offset):\n    print(\"\\n\" + \"#\" * 70)\n    print(f\"# ENSEMBLE MEMBER {member_idx+1}/{CFG.N_ENSEMBLE_MODELS}  \"\n          f\"seed={seed} encoder={encoder_name} depth_offset={depth_offset}\")\n    print(\"#\" * 70)\n\n    set_seed(seed)\n    depth_indices = clip_depth_indices(CFG.depth_indices_base, depth_offset)\n\n    # --- open this member's own (possibly Z-shifted) volumes ---\n    train_volumes = {}\n    for fid in CFG.train_frags:\n        frag_dir = os.path.join(CFG.base_dir, fid)\n        train_volumes[fid] = FragmentVolume(frag_dir, depth_indices)\n\n    train_samples = build_train_samples(seed)\n    print(\"Member train patches:\", len(train_samples), \"| val patches:\", len(val_samples))\n\n    train_dataset = InkPatchDataset(train_volumes, _shared_labels, train_samples,\n                                     CFG.patch_size, transform=build_train_transform())\n    val_dataset = InkPatchDataset(train_volumes, _shared_labels, val_samples,\n                                   CFG.patch_size, transform=None)\n\n    train_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True,\n                               num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                               drop_last=CFG.drop_last, persistent_workers=CFG.num_workers > 0)\n    val_loader = DataLoader(val_dataset, batch_size=CFG.batch_size, shuffle=False,\n                             num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                             persistent_workers=CFG.num_workers > 0)\n\n    # positive-pixel prior for bias-init / BCE weighting\n    prior_samples = random.sample(train_samples, min(120, len(train_samples)))\n    total_px, pos_px = 0, 0\n    for fid, y, x in prior_samples:\n        patch = _shared_labels[fid][y:y + CFG.patch_size, x:x + CFG.patch_size]\n        pos_px += int(patch.sum())\n        total_px += int(patch.size)\n    positive_fraction = float(np.clip(pos_px / max(total_px, 1), 1e-4, 0.25))\n    initial_bias = math.log(positive_fraction / max(1 - positive_fraction, 1e-6))\n    bce_pos_weight = float(np.clip(np.sqrt((1 - positive_fraction) / positive_fraction), 1.0, 6.0))\n    print(f\"Positive fraction={positive_fraction:.5f} bias={initial_bias:.3f} pos_weight={bce_pos_weight:.2f}\")\n\n    model = build_model(encoder_name).to(CFG.device)\n    with torch.no_grad():\n        try:\n            model.segmentation_head[0].bias.fill_(initial_bias)\n        except Exception:\n            pass\n\n    pos_weight_tensor = torch.tensor([bce_pos_weight], dtype=torch.float32, device=CFG.device)\n    criterion = V4Loss(pos_weight_tensor)\n\n    encoder_parameters = list(model.encoder.parameters())\n    encoder_ids = {id(p) for p in encoder_parameters}\n    decoder_parameters = [p for p in model.parameters() if id(p) not in encoder_ids]\n    optimizer = torch.optim.AdamW([\n        {\"params\": encoder_parameters, \"lr\": CFG.encoder_lr},\n        {\"params\": decoder_parameters, \"lr\": CFG.decoder_lr},\n    ], weight_decay=CFG.weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=CFG.encoder_lr * 0.10)\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\" and CFG.use_amp))\n\n    # --- V5: keep top-K checkpoints (CPU state_dicts) for weight averaging ---\n    top_checkpoints = []   # list of (val_dice, state_dict_cpu)\n    best_val_dice, best_epoch, no_improve = -1.0, -1, 0\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n        train_metrics = run_epoch(model, train_loader, criterion, optimizer, scaler)\n        val_metrics = run_epoch(model, val_loader, criterion, optimizer=None, scaler=None)\n        scheduler.step()\n\n        print(f\"[member {member_idx+1}][{epoch:02d}/{CFG.epochs}] \"\n              f\"train_loss={train_metrics['loss']:.5f} train_dice={train_metrics['dice']:.5f} | \"\n              f\"val_loss={val_metrics['loss']:.5f} val_dice={val_metrics['dice']:.5f} \"\n              f\"val_P={val_metrics['precision']:.5f} val_R={val_metrics['recall']:.5f} \"\n              f\"({time.time()-t0:.1f}s)\")\n\n        if CFG.USE_WEIGHT_AVERAGING:\n            cpu_sd = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n            top_checkpoints.append((val_metrics[\"dice\"], cpu_sd))\n            top_checkpoints.sort(key=lambda t: t[0], reverse=True)\n            top_checkpoints = top_checkpoints[:CFG.weight_avg_top_k]\n\n        if val_metrics[\"dice\"] > best_val_dice:\n            best_val_dice, best_epoch, no_improve = val_metrics[\"dice\"], epoch, 0\n            if not CFG.USE_WEIGHT_AVERAGING:\n                torch.save({\"model\": model.state_dict(), \"epoch\": epoch, \"val_dice\": best_val_dice},\n                           os.path.join(CFG.out_dir, f\"member{member_idx}_best.pth\"))\n        else:\n            no_improve += 1\n            if epoch >= CFG.min_epochs and no_improve >= CFG.early_stop_patience:\n                print(\"Early stopping.\")\n                break\n        cleanup_memory()\n\n    print(f\"Member {member_idx+1} best validation Dice: {best_val_dice:.5f} (epoch {best_epoch})\")\n\n    if CFG.USE_WEIGHT_AVERAGING and len(top_checkpoints) > 0:\n        print(f\"Averaging top-{len(top_checkpoints)} checkpoints by val Dice \"\n              f\"({[round(d,4) for d,_ in top_checkpoints]}) into one 'model soup'.\")\n        soup_sd = average_state_dicts([sd for _, sd in top_checkpoints])\n        model.load_state_dict(soup_sd)\n        torch.save({\"model\": model.state_dict(), \"top_k_dices\": [d for d, _ in top_checkpoints]},\n                    os.path.join(CFG.out_dir, f\"member{member_idx}_soup.pth\"))\n        del top_checkpoints, soup_sd\n    else:\n        ckpt = torch.load(os.path.join(CFG.out_dir, f\"member{member_idx}_best.pth\"), map_location=CFG.device)\n        model.load_state_dict(ckpt[\"model\"])\n\n    model.eval()\n    cleanup_memory()\n\n    # --- validation threshold (for local diagnostics/reporting only) ---\n    @torch.no_grad()\n    def collect_val_predictions():\n        all_probs, all_targets = [], []\n        for imgs, masks in val_loader:\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            with autocast(enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n                logits = model(imgs)\n            probs = torch.sigmoid(logits).float().cpu().numpy()\n            all_probs.append(probs[:, 0])\n            all_targets.append(masks.numpy()[:, 0])\n            del imgs, masks, logits, probs\n        return np.concatenate(all_probs, axis=0), np.concatenate(all_targets, axis=0)\n\n    val_probs, val_targets = collect_val_predictions()\n    best_val_result = None\n    for threshold in CFG.val_thresholds:\n        pred = val_probs >= threshold\n        gt = val_targets > 0.5\n        tp = np.logical_and(pred, gt).sum()\n        fp = np.logical_and(pred, ~gt).sum()\n        fn = np.logical_and(~pred, gt).sum()\n        m = calculate_metrics(tp, fp, fn)\n        m[\"threshold\"] = float(threshold)\n        if best_val_result is None or m[\"dice\"] > best_val_result[\"dice\"]:\n            best_val_result = m\n    val_threshold = best_val_result[\"threshold\"]\n    print(\"Member val threshold:\", val_threshold, \"| val metrics:\", best_val_result)\n\n    del train_loader, val_loader, train_dataset, val_dataset, val_probs, val_targets\n    for v in train_volumes.values():\n        v.close()\n    del train_volumes\n    cleanup_memory()\n\n    # --- fragment-1 inference: this member's own (possibly Z-shifted) test volume ---\n    test_vol = FragmentVolume(test_dir, depth_indices)\n\n    print(\"\\n[Baseline] no domain adaptation ...\")\n    baseline_probability = sliding_window_inference(model, test_vol, test_mask)\n    baseline_local = local_metrics(baseline_probability, val_threshold)\n    print(\"Baseline local metrics @val_threshold:\", baseline_local)\n\n    if CFG.use_adabn:\n        run_adabn(model, test_vol)\n    print(\"\\n[AdaBN] test inference ...\")\n    adabn_probability = sliding_window_inference(model, test_vol, test_mask)\n\n    otsu_threshold = compute_otsu_threshold(adabn_probability, test_mask, fallback=val_threshold)\n    adabn_local = local_metrics(adabn_probability, otsu_threshold)\n    print(f\"Member {member_idx+1} AdaBN+Otsu(thr={otsu_threshold:.3f}) local metrics:\", adabn_local)\n\n    rot90_diag = None\n    if CFG.run_rot90_ablation_on_first_member and member_idx == 0:\n        print(\"\\n[Diagnostic-only] Rot90 ablation on member 0 (never used in final prediction) ...\")\n\n        def inverse_rot90(p):\n            return np.rot90(p, -1).copy()\n\n        rot90_probability = sliding_window_inference(model, test_vol, test_mask, transform=inverse_rot90)\n        rot90_thr = compute_otsu_threshold(rot90_probability, test_mask, fallback=otsu_threshold)\n        rot90_diag = local_metrics(rot90_probability, rot90_thr)\n        print(\"Rot90 diagnostic (NOT used in final prediction):\", rot90_diag)\n        del rot90_probability\n\n    test_vol.close()\n    cleanup_memory()\n\n    return {\n        \"member_idx\": member_idx,\n        \"seed\": seed,\n        \"encoder\": encoder_name,\n        \"depth_offset\": depth_offset,\n        \"best_val_dice\": best_val_dice,\n        \"best_epoch\": best_epoch,\n        \"val_threshold\": val_threshold,\n        \"val_metrics\": best_val_result,\n        \"baseline_probability\": baseline_probability,\n        \"adabn_probability\": adabn_probability,\n        \"otsu_threshold\": otsu_threshold,\n        \"baseline_local_metrics\": baseline_local,\n        \"adabn_local_metrics\": adabn_local,\n        \"rot90_diagnostic_metrics\": rot90_diag,\n    }\n\n\n# ============================================================\n# 14. RUN THE ENSEMBLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(f\"STARTING V5 ENSEMBLE ({CFG.N_ENSEMBLE_MODELS} members)\")\nprint(\"=\" * 70)\n\nmember_results = []\nfor i in range(CFG.N_ENSEMBLE_MODELS):\n    result = train_and_infer_one_member(\n        member_idx=i,\n        seed=CFG.ensemble_seeds[i],\n        encoder_name=CFG.ensemble_encoders[i],\n        depth_offset=CFG.ensemble_depth_offsets[i],\n    )\n    member_results.append(result)\n    cleanup_memory()\n\n\n# ============================================================\n# 15. ENSEMBLE THE PROBABILITY MAPS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"ENSEMBLING MEMBER PREDICTIONS\")\nprint(\"=\" * 70)\n\nensembled_probability = np.mean(\n    [r[\"adabn_probability\"] for r in member_results], axis=0\n).astype(np.float32)\n\nensembled_threshold = compute_otsu_threshold(ensembled_probability, test_mask,\n                                              fallback=float(np.mean([r[\"otsu_threshold\"] for r in member_results])))\nensembled_local = local_metrics(ensembled_probability, ensembled_threshold)\n\nprint(f\"Ensembled (mean of {CFG.N_ENSEMBLE_MODELS} members) + Otsu(thr={ensembled_threshold:.3f}):\")\nprint(ensembled_local)\n\nfor r in member_results:\n    print(f\"  member {r['member_idx']+1} ({r['encoder']}, seed={r['seed']}, \"\n          f\"depth_offset={r['depth_offset']}): val_dice={r['best_val_dice']:.4f} \"\n          f\"-> local test dice={r['adabn_local_metrics']['dice'] if r['adabn_local_metrics'] else None}\")\n\nfinal_prediction = postprocess(ensembled_probability, ensembled_threshold, test_mask)\nfinal_post_metrics = local_metrics(ensembled_probability, ensembled_threshold)  # postprocess only removes small components; approx via binary mask below\nif test_labels is not None:\n    gt = (test_labels > 0) & (test_mask > 0)\n    pred = final_prediction > 0\n    tp = np.logical_and(pred, gt).sum()\n    fp = np.logical_and(pred, ~gt).sum()\n    fn = np.logical_and(~pred, gt).sum()\n    final_post_metrics = calculate_metrics(tp, fp, fn)\n    print(\"Final POSTPROCESSED ensembled metrics:\", final_post_metrics)\n\n\n# ============================================================\n# 16. SAVE OUTPUTS\n# ============================================================\n\nnp.save(os.path.join(CFG.out_dir, \"fragment1_probability_ensembled_v5.npy\"), ensembled_probability)\nfor r in member_results:\n    np.save(os.path.join(CFG.out_dir, f\"fragment1_probability_member{r['member_idx']}_v5.npy\"),\n            r[\"adabn_probability\"])\n\ncv2.imwrite(os.path.join(CFG.out_dir, \"fragment1_prediction_ensembled_v5.png\"),\n            (final_prediction * 255).astype(np.uint8))\n\nmetrics_summary = {\n    \"version\": \"V5_ensemble\",\n    \"n_ensemble_models\": CFG.N_ENSEMBLE_MODELS,\n    \"use_weight_averaging\": CFG.USE_WEIGHT_AVERAGING,\n    \"use_robust_norm\": CFG.USE_ROBUST_NORM,\n    \"epochs\": CFG.epochs,\n    \"ensembled_threshold\": float(ensembled_threshold),\n    \"ensembled_local_metrics_raw\": ensembled_local,\n    \"ensembled_local_metrics_postprocessed\": final_post_metrics,\n    \"members\": [\n        {\n            \"member_idx\": r[\"member_idx\"], \"seed\": r[\"seed\"], \"encoder\": r[\"encoder\"],\n            \"depth_offset\": r[\"depth_offset\"], \"best_val_dice\": r[\"best_val_dice\"],\n            \"best_epoch\": r[\"best_epoch\"], \"val_threshold\": r[\"val_threshold\"],\n            \"val_metrics\": r[\"val_metrics\"], \"otsu_threshold\": r[\"otsu_threshold\"],\n            \"baseline_local_metrics\": r[\"baseline_local_metrics\"],\n            \"adabn_local_metrics\": r[\"adabn_local_metrics\"],\n            \"rot90_diagnostic_metrics_NOT_used_in_final\": r[\"rot90_diagnostic_metrics\"],\n        }\n        for r in member_results\n    ],\n}\nwith open(CFG.metrics_path, \"w\") as f:\n    json.dump(metrics_summary, f, indent=2)\nprint(f\"\\nSaved metrics -> {CFG.metrics_path}\")\n\n\n# ============================================================\n# 17. VISUALIZATIONS\n# ============================================================\n\nmid_idx = CFG.depth_indices_base[len(CFG.depth_indices_base) // 2]\nmid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\nscale = min(1.0, 1800.0 / max(mid_slice.shape))\nsmall = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\nprob_small = cv2.resize(ensembled_probability, small.shape[::-1], interpolation=cv2.INTER_AREA)\npred_small = cv2.resize((final_prediction * 255).astype(np.uint8), small.shape[::-1],\n                         interpolation=cv2.INTER_NEAREST)\n\nif test_labels is not None:\n    gt_small = cv2.resize((test_labels * 255).astype(np.uint8), small.shape[::-1],\n                           interpolation=cv2.INTER_NEAREST)\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n    axes[1].imshow(prob_small, cmap=\"magma\", vmin=0, vmax=1)\n    axes[1].set_title(f\"Ensembled probability\\nT={ensembled_threshold:.3f}\")\n    axes[2].imshow(gt_small, cmap=\"gray\"); axes[2].set_title(\"Local GT\")\n    axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"V5 Ensembled Prediction\")\nelse:\n    fig, axes = plt.subplots(1, 3, figsize=(17, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input slice {mid_idx}\")\n    axes[1].imshow(prob_small, cmap=\"magma\", vmin=0, vmax=1)\n    axes[1].set_title(f\"Ensembled probability\\nT={ensembled_threshold:.3f}\")\n    axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"V5 Ensembled Prediction\")\n\nfor ax in axes:\n    ax.axis(\"off\")\nplt.tight_layout()\nplt.savefig(os.path.join(CFG.viz_dir, \"fragment1_v5_overview.png\"), dpi=150, bbox_inches=\"tight\")\nplt.close()\n\n# patch comparisons (uses member 0's own depth window for the visual input slice)\ntest_vol_viz = FragmentVolume(test_dir, clip_depth_indices(CFG.depth_indices_base, CFG.ensemble_depth_offsets[0]))\ncomparison_coords = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\nrandom.shuffle(comparison_coords)\ncomparison_coords = comparison_coords[:6]\nmid_local_idx = len(CFG.depth_indices_base) // 2\n\nfor i, (y, x) in enumerate(comparison_coords):\n    input_patch = test_vol_viz.read_patch(y, x, CFG.patch_size)[mid_local_idx]\n    pred_patch = final_prediction[y:y + CFG.patch_size, x:x + CFG.patch_size]\n    if test_labels is not None:\n        gt_patch = test_labels[y:y + CFG.patch_size, x:x + CFG.patch_size]\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(input_patch, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"GT\")\n        axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"V5 Ensembled\")\n    else:\n        fig, axes = plt.subplots(1, 2, figsize=(8, 4))\n        axes[0].imshow(input_patch, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(pred_patch, cmap=\"gray\"); axes[1].set_title(\"V5 Ensembled\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(os.path.join(CFG.viz_dir, f\"patch_{i:02d}_y{y}_x{x}.png\"), dpi=150, bbox_inches=\"tight\")\n    plt.close()\n\ntest_vol_viz.close()\ncleanup_memory()\n\n\n# ============================================================\n# 18. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"V5 ENSEMBLE COMPLETE\")\nprint(\"=\" * 70)\nfor r in member_results:\n    print(f\"Member {r['member_idx']+1}: encoder={r['encoder']} seed={r['seed']} \"\n          f\"depth_offset={r['depth_offset']} best_val_dice={r['best_val_dice']:.5f} \"\n          f\"(epoch {r['best_epoch']}) local_test_dice={r['adabn_local_metrics']['dice'] if r['adabn_local_metrics'] else 'n/a'}\")\nprint(f\"\\nEnsembled threshold (Otsu): {ensembled_threshold:.3f}\")\nif ensembled_local is not None:\n    print(f\"Ensembled local test Dice (raw): {ensembled_local['dice']:.5f}\")\nif final_post_metrics is not None:\n    print(f\"Ensembled local test Dice (postprocessed): {final_post_metrics['dice']:.5f}\")\nprint(f\"\\nMetrics: {CFG.metrics_path}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")\nprint(\"\\n=== V5 DONE ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T08:06:30.077198Z","iopub.execute_input":"2026-08-30T08:06:30.077525Z"}},"outputs":[{"name":"stdout","text":"\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m\n======================================================================\nBUILDING SHARED TRAIN/VAL PATCH GRID (fragments 2 & 3)\n======================================================================\nFragment 2: (14830, 9506) 6367 patches\nFragment 3: (7606, 5249) 1696 patches\nTotal candidate patches: 8063\nSpatial train: 6416\nSpatial val: 1647\n\nSampling pools -- positive: 5088 hard_negative: 1003 normal: 325\n\nTest shape: (8181, 6330) | local test GT available: True\n\n======================================================================\nSTARTING V5 ENSEMBLE (2 members)\n======================================================================\n\n######################################################################\n# ENSEMBLE MEMBER 1/2  seed=42 encoder=resnet50 depth_offset=0\n######################################################################\nMember train patches: 6416 | val patches: 1647\nPositive fraction=0.16250 bias=-1.640 pos_weight=2.27\n","output_type":"stream"},{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/resnet50-19c8e357.pth\" to /root/.cache/torch/hub/checkpoints/resnet50-19c8e357.pth\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"  0%|          | 0.00/97.8M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"b11baec0ce514a66ad6998095ef022e0"}},"metadata":{}},{"name":"stdout","text":"[member 1][01/8] train_loss=0.66812 train_dice=0.44285 | val_loss=0.66091 val_dice=0.48883 val_P=0.43396 val_R=0.55957 (845.2s)\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS CHALLENGE - INK DETECTION\n# V4\n#\n# BASED ON:\n#   Successful V2 480 configuration\n#\n# V4 changes after V3 experiment:\n#\n#   PATCH              = 480\n#   TRAIN STRIDE       = 128\n#   TEST STRIDE        = 128\n#   BATCH              = 8\n#\n#   MODEL              = ResNet50 + scSE U-Net\n#\n#   LOSS:\n#       BCE            0.45\n#       Dice           0.40\n#       Tversky        0.15\n#\n#   SAMPLING:\n#       Positive       35%\n#       Hard negative  15%\n#       Normal         50%\n#\n#   AdaBN             = ON\n#\n#   TTA                = OFF by default\n#   Rot90 ablation     = optional, NO retraining\n#\n#   VALIDATION:\n#       threshold optimized using validation GT only\n#\n#   TEST:\n#       independent threshold calibration\n#       Otsu\n#       quantile\n#       validation-positive-rate matching\n#\n#   POSTPROCESSING:\n#       conservative\n#\n#   OOM:\n#       disk-backed TIFF memmap\n#       patch-by-patch loading\n#       AMP\n#       no full 3D volume in RAM\n#\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\n!pip install -q segmentation-models-pytorch==0.2.0\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import (\n    Dataset,\n    DataLoader,\n    WeightedRandomSampler\n)\n\nfrom torch.cuda.amp import (\n    autocast,\n    GradScaler\n)\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\n\nimport segmentation_models_pytorch as smp\n\nfrom skimage.filters import threshold_otsu\n\n\n# ============================================================\n# 2. CONFIGURATION\n# ============================================================\n\nclass CFG:\n\n    # --------------------------------------------------------\n    # PATHS\n    # --------------------------------------------------------\n\n    base_dir = (\n        \"/kaggle/input/\"\n        \"vesuvius-challenge-ink-detection/\"\n        \"train\"\n    )\n\n    train_frags = [\n        \"2\",\n        \"3\"\n    ]\n\n    test_frag = \"1\"\n\n\n    # --------------------------------------------------------\n    # DEPTH\n    # --------------------------------------------------------\n\n    # 22 slices\n    #\n    # Same depth range used in successful V2.\n    #\n    depth_indices = list(\n        range(16, 38)\n    )\n\n    in_channels = len(\n        depth_indices\n    )\n\n\n    # --------------------------------------------------------\n    # PATCH / STRIDE\n    # --------------------------------------------------------\n\n    patch_size = 480\n\n    train_stride = 128\n\n    test_stride = 128\n\n\n    # --------------------------------------------------------\n    # BATCHING\n    # --------------------------------------------------------\n\n    batch_size = 8\n\n    infer_batch = 8\n\n    num_workers = 2\n\n    drop_last = True\n\n\n    # --------------------------------------------------------\n    # TRAINING\n    # --------------------------------------------------------\n\n    epochs = 10\n\n    early_stop_patience = 3\n\n    min_epochs = 4\n\n\n    # Differential LR.\n    #\n    # Keep pretrained encoder learning slowly.\n    #\n    encoder_lr = 5e-5\n\n    decoder_lr = 2e-4\n\n\n    weight_decay = 1e-4\n\n    grad_clip = 1.0\n\n    use_amp = True\n\n\n    # --------------------------------------------------------\n    # VALIDATION\n    # --------------------------------------------------------\n\n    val_fraction = 0.20\n\n    # Spatial blocks.\n    #\n    # We deliberately separate validation spatially.\n    #\n    validation_block_size = 768\n\n\n    # --------------------------------------------------------\n    # TISSUE FILTER\n    # --------------------------------------------------------\n\n    min_tissue_frac_train = 0.10\n\n    min_tissue_frac_test = 0.02\n\n\n    # --------------------------------------------------------\n    # SAMPLING\n    # --------------------------------------------------------\n\n    # IMPORTANT:\n    #\n    # V3 used approximately:\n    #\n    #   55% positive\n    #   20% hard negative\n    #   25% random\n    #\n    # That produced excellent validation but poor\n    # fragment-1 transfer.\n    #\n    # V4 is intentionally less aggressive.\n    #\n\n    positive_ratio = 0.35\n\n    hard_negative_ratio = 0.15\n\n    normal_ratio = 0.50\n\n\n    # Patch considered positive if this much\n    # of the patch contains ink.\n    #\n    positive_patch_fraction = 0.001\n\n\n    # Hard negative:\n    #\n    # patch has tissue but almost no ink.\n    #\n    hard_negative_tissue_fraction = 0.45\n\n\n    # Maximum training samples per epoch.\n    #\n    # Keeps memory/time controlled.\n    #\n    max_train_samples = 18000\n\n\n    # --------------------------------------------------------\n    # MODEL\n    # --------------------------------------------------------\n\n    encoder_name = \"resnet50\"\n\n    encoder_weights = \"imagenet\"\n\n    decoder_attention_type = \"scse\"\n\n\n    # --------------------------------------------------------\n    # LOSS\n    # --------------------------------------------------------\n\n    # V4 conservative loss.\n    #\n    bce = 0.45\n    dice = 0.40\n    tversky = 0.15\n\n\n    tversky_alpha = 0.35\n\n    tversky_beta = 0.65\n\n    tversky_gamma = 1.0\n\n\n    # --------------------------------------------------------\n    # AUGMENTATION\n    # --------------------------------------------------------\n\n    # Rotation is deliberately stronger because\n    # V3's Rot90 test-time transform was the strongest TTA.\n    #\n    # Instead of relying on TTA, teach the network\n    # rotational invariance during training.\n    #\n\n    rotate90_probability = 0.75\n\n\n    # --------------------------------------------------------\n    # ADABN\n    # --------------------------------------------------------\n\n    use_adabn = True\n\n    adabn_max_patches = 1200\n\n\n    # --------------------------------------------------------\n    # VALIDATION THRESHOLD\n    # --------------------------------------------------------\n\n    val_thresholds = np.arange(\n        0.20,\n        0.801,\n        0.01\n    )\n\n\n    # --------------------------------------------------------\n    # TEST THRESHOLD\n    # --------------------------------------------------------\n\n    # The TEST threshold is intentionally NOT the\n    # validation threshold.\n    #\n    # We calculate several independent target-fragment\n    # threshold candidates.\n    #\n\n    test_threshold_min = 0.15\n\n    test_threshold_max = 0.75\n\n\n    # --------------------------------------------------------\n    # TEST THRESHOLD STRATEGY\n    # --------------------------------------------------------\n\n    # Options:\n    #\n    # \"otsu\"\n    # \"quantile\"\n    # \"positive_rate\"\n    #\n    # The code calculates ALL of them and prints local\n    # diagnostics if fragment-1 GT is available.\n    #\n\n    final_test_threshold_method = (\n        \"otsu\"\n    )\n\n\n    # --------------------------------------------------------\n    # TTA\n    # --------------------------------------------------------\n\n    # FINAL TTA = OFF\n    #\n    # V3 demonstrated that TTA did not consistently improve.\n    #\n    use_tta_for_final = False\n\n\n    # Run Rot90 after training only for diagnosis.\n    #\n    # Does NOT require another training.\n    #\n    run_rot90_ablation = True\n\n\n    # --------------------------------------------------------\n    # POST PROCESSING\n    # --------------------------------------------------------\n\n    use_postprocess = True\n\n    closing_kernel = 3\n\n    min_component_area = 8\n\n\n    # --------------------------------------------------------\n    # RANDOM SEED\n    # --------------------------------------------------------\n\n    seed = 42\n\n\n    # --------------------------------------------------------\n    # DEVICE\n    # --------------------------------------------------------\n\n    device = (\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\n\n    # --------------------------------------------------------\n    # OUTPUT\n    # --------------------------------------------------------\n\n    out_dir = (\n        \"/kaggle/working/\"\n        \"vesuvius_v4\"\n    )\n\n    viz_dir = os.path.join(\n        out_dir,\n        \"visualizations\"\n    )\n\n    ckpt_path = os.path.join(\n        out_dir,\n        \"vesuvius_v4_best.pth\"\n    )\n\n    metrics_path = os.path.join(\n        out_dir,\n        \"v4_metrics_summary.json\"\n    )\n\n\n# Create directories.\n\nos.makedirs(\n    CFG.out_dir,\n    exist_ok=True\n)\n\nos.makedirs(\n    CFG.viz_dir,\n    exist_ok=True\n)\n\n\n# ============================================================\n# 3. SEED\n# ============================================================\n\ndef set_seed(seed):\n\n    random.seed(seed)\n\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n\n    torch.cuda.manual_seed_all(seed)\n\n\nset_seed(\n    CFG.seed\n)\n\ntorch.backends.cudnn.benchmark = True\n\n\n# ============================================================\n# 4. MEMORY CLEANUP\n# ============================================================\n\ndef cleanup_memory():\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n\n        torch.cuda.empty_cache()\n\n\n# ============================================================\n# 5. TISSUE MASK\n# ============================================================\n\ndef load_tissue_mask(\n    frag_dir\n):\n\n    mask_path = os.path.join(\n        frag_dir,\n        \"mask.png\"\n    )\n\n    if os.path.exists(\n        mask_path\n    ):\n\n        mask = cv2.imread(\n            mask_path,\n            cv2.IMREAD_GRAYSCALE\n        )\n\n    else:\n\n        mid_idx = (\n            CFG.depth_indices[\n                len(CFG.depth_indices)//2\n            ]\n        )\n\n        mid_path = os.path.join(\n            frag_dir,\n            \"surface_volume\",\n            f\"{mid_idx:02d}.tif\"\n        )\n\n        mid = tifffile.imread(\n            mid_path\n        )\n\n        threshold = (\n            mid.mean() * 0.15\n        )\n\n        mask = (\n            mid > threshold\n        ).astype(\n            np.uint8\n        ) * 255\n\n\n    return (\n        mask > 0\n    ).astype(\n        np.uint8\n    )\n\n\n# ============================================================\n# 6. LABEL LOADING\n# ============================================================\n\ndef load_ink_labels(\n    frag_dir\n):\n\n    path = os.path.join(\n        frag_dir,\n        \"inklabels.png\"\n    )\n\n    if not os.path.exists(\n        path\n    ):\n\n        return None\n\n    lbl = cv2.imread(\n        path,\n        cv2.IMREAD_GRAYSCALE\n    )\n\n    if lbl is None:\n\n        return None\n\n    return (\n        lbl > 0\n    ).astype(\n        np.uint8\n    )\n\n\n# ============================================================\n# 7. PATCH GRID\n# ============================================================\n\ndef generate_grid_coords(\n    mask,\n    patch_size,\n    stride,\n    min_frac\n):\n\n    H, W = mask.shape\n\n    ymax = max(\n        H - patch_size,\n        0\n    )\n\n    xmax = max(\n        W - patch_size,\n        0\n    )\n\n\n    ys = list(\n        range(\n            0,\n            ymax + 1,\n            stride\n        )\n    )\n\n    xs = list(\n        range(\n            0,\n            xmax + 1,\n            stride\n        )\n    )\n\n\n    if (\n        not ys\n        or\n        ys[-1] != ymax\n    ):\n\n        ys.append(\n            ymax\n        )\n\n\n    if (\n        not xs\n        or\n        xs[-1] != xmax\n    ):\n\n        xs.append(\n            xmax\n        )\n\n\n    coords = []\n\n\n    for y in ys:\n\n        for x in xs:\n\n            tissue_fraction = (\n                mask[\n                    y:y+patch_size,\n                    x:x+patch_size\n                ].mean()\n            )\n\n            if (\n                tissue_fraction\n                >= min_frac\n            ):\n\n                coords.append(\n                    (\n                        y,\n                        x\n                    )\n                )\n\n\n    return coords\n\n\n# ============================================================\n# 8. DISK-BACKED VOLUME\n# ============================================================\n\nclass FragmentVolume:\n\n    def __init__(\n        self,\n        frag_dir,\n        depth_indices\n    ):\n\n        self.paths = [\n\n            os.path.join(\n                frag_dir,\n                \"surface_volume\",\n                f\"{i:02d}.tif\"\n            )\n\n            for i in depth_indices\n        ]\n\n        self._slices = None\n\n        self._h = None\n\n        self._w = None\n\n\n    def _ensure_open(\n        self\n    ):\n\n        if self._slices is not None:\n\n            return\n\n\n        slices = []\n\n\n        for path in self.paths:\n\n            try:\n\n                arr = tifffile.memmap(\n                    path,\n                    mode=\"r\"\n                )\n\n            except Exception:\n\n                arr = tifffile.imread(\n                    path\n                )\n\n\n            slices.append(\n                arr\n            )\n\n\n        self._slices = slices\n\n        self._h = (\n            slices[0].shape[0]\n        )\n\n        self._w = (\n            slices[0].shape[1]\n        )\n\n\n    @property\n    def shape(self):\n\n        self._ensure_open()\n\n        return (\n            self._h,\n            self._w\n        )\n\n\n    def read_patch(\n        self,\n        y,\n        x,\n        size\n    ):\n\n        self._ensure_open()\n\n\n        out = np.empty(\n            (\n                len(self._slices),\n                size,\n                size\n            ),\n            dtype=np.uint8\n        )\n\n\n        for i, s in enumerate(\n            self._slices\n        ):\n\n            block = s[\n                y:y+size,\n                x:x+size\n            ]\n\n\n            if block.dtype != np.uint8:\n\n                if (\n                    np.issubdtype(\n                        block.dtype,\n                        np.integer\n                    )\n                ):\n\n                    max_value = np.iinfo(\n                        block.dtype\n                    ).max\n\n                else:\n\n                    max_value = (\n                        float(\n                            np.nanmax(\n                                block\n                            )\n                        )\n                        + 1e-6\n                    )\n\n\n                block = (\n                    block.astype(\n                        np.float32\n                    )\n                    / max_value\n                    * 255.0\n                )\n\n\n                block = np.clip(\n                    block,\n                    0,\n                    255\n                ).astype(\n                    np.uint8\n                )\n\n\n            out[i] = block\n\n\n        return out\n\n\n    def close(self):\n\n        self._slices = None\n\n        cleanup_memory()\n\n\n# ============================================================\n# 9. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(\n    x\n):\n\n    x = x.astype(\n        np.float32\n    )\n\n    mean = x.mean()\n\n    std = (\n        x.std()\n        + 1e-6\n    )\n\n    return (\n        x - mean\n    ) / std\n\n\n# ============================================================\n# 10. AUGMENTATION\n# ============================================================\n\ndef build_train_transform():\n\n    return A.Compose([\n\n        A.HorizontalFlip(\n            p=0.5\n        ),\n\n        A.VerticalFlip(\n            p=0.5\n        ),\n\n        A.RandomRotate90(\n            p=CFG.rotate90_probability\n        ),\n\n        A.Transpose(\n            p=0.25\n        ),\n\n        A.ShiftScaleRotate(\n\n            shift_limit=0.025,\n\n            scale_limit=0.06,\n\n            rotate_limit=10,\n\n            border_mode=(\n                cv2.BORDER_REFLECT_101\n            ),\n\n            p=0.30\n        ),\n\n        A.RandomBrightnessContrast(\n\n            brightness_limit=0.06,\n\n            contrast_limit=0.08,\n\n            p=0.20\n        )\n\n    ])\n\n\n# ============================================================\n# 11. DATASET\n# ============================================================\n\nclass InkPatchDataset(\n    Dataset\n):\n\n    def __init__(\n        self,\n        volumes,\n        labels,\n        samples,\n        patch_size,\n        transform=None\n    ):\n\n        self.volumes = volumes\n\n        self.labels = labels\n\n        self.samples = samples\n\n        self.patch_size = patch_size\n\n        self.transform = transform\n\n\n    def __len__(\n        self\n    ):\n\n        return len(\n            self.samples\n        )\n\n\n    def __getitem__(\n        self,\n        idx\n    ):\n\n        fid, y, x = (\n            self.samples[idx]\n        )\n\n\n        patch = (\n            self.volumes[fid]\n            .read_patch(\n                y,\n                x,\n                self.patch_size\n            )\n        )\n\n\n        label = (\n            self.labels[fid][\n                y:y+self.patch_size,\n                x:x+self.patch_size\n            ]\n        )\n\n\n        # C,H,W -> H,W,C\n\n        img = np.transpose(\n            patch,\n            (1, 2, 0)\n        )\n\n\n        if self.transform is not None:\n\n            aug = self.transform(\n                image=img,\n                mask=label\n            )\n\n            img = aug[\"image\"]\n\n            label = aug[\"mask\"]\n\n\n        img = (\n            img.astype(\n                np.float32\n            )\n            / 255.0\n        )\n\n\n        img = normalize_patch(\n            img\n        )\n\n\n        img = np.ascontiguousarray(\n            np.transpose(\n                img,\n                (2, 0, 1)\n            )\n        )\n\n\n        label = (\n            label > 0\n        ).astype(\n            np.float32\n        )[None, ...]\n\n\n        return (\n            torch.from_numpy(img),\n            torch.from_numpy(label)\n        )\n\n\n# ============================================================\n# 12. MODEL\n# ============================================================\n\ndef build_model():\n\n    model = smp.Unet(\n\n        encoder_name=(\n            CFG.encoder_name\n        ),\n\n        encoder_weights=(\n            CFG.encoder_weights\n        ),\n\n        in_channels=(\n            CFG.in_channels\n        ),\n\n        classes=1,\n\n        decoder_attention_type=(\n            CFG.decoder_attention_type\n        )\n    )\n\n    return model\n\n\n# ============================================================\n# 13. LOSSES\n# ============================================================\n\ndef dice_loss(\n    logits,\n    targets,\n    eps=1e-6\n):\n\n    probs = torch.sigmoid(\n        logits\n    )\n\n    probs = probs.reshape(\n        probs.size(0),\n        -1\n    )\n\n    targets = targets.reshape(\n        targets.size(0),\n        -1\n    )\n\n\n    intersection = (\n        probs * targets\n    ).sum(\n        dim=1\n    )\n\n\n    union = (\n        probs.sum(\n            dim=1\n        )\n        +\n        targets.sum(\n            dim=1\n        )\n    )\n\n\n    dice = (\n        2 * intersection\n        + eps\n    ) / (\n        union + eps\n    )\n\n\n    return (\n        1 - dice.mean()\n    )\n\n\ndef focal_tversky_loss(\n    logits,\n    targets,\n    alpha,\n    beta,\n    gamma,\n    eps=1e-6\n):\n\n    probs = torch.sigmoid(\n        logits\n    )\n\n    probs = probs.reshape(\n        probs.size(0),\n        -1\n    )\n\n    targets = targets.reshape(\n        targets.size(0),\n        -1\n    )\n\n\n    tp = (\n        probs * targets\n    ).sum(\n        dim=1\n    )\n\n\n    fp = (\n        probs\n        * (1 - targets)\n    ).sum(\n        dim=1\n    )\n\n\n    fn = (\n        (1 - probs)\n        * targets\n    ).sum(\n        dim=1\n    )\n\n\n    tversky = (\n        tp + eps\n    ) / (\n        tp\n        + alpha * fp\n        + beta * fn\n        + eps\n    )\n\n\n    return (\n        1 - tversky\n    ).clamp_min(\n        0\n    ).pow(\n        gamma\n    ).mean()\n\n\nclass V4Loss(\n    nn.Module\n):\n\n    def __init__(\n        self,\n        pos_weight\n    ):\n\n        super().__init__()\n\n\n        self.bce = (\n            nn.BCEWithLogitsLoss(\n                pos_weight=pos_weight\n            )\n        )\n\n\n    def forward(\n        self,\n        logits,\n        targets\n    ):\n\n        bce = self.bce(\n            logits,\n            targets\n        )\n\n\n        dice = dice_loss(\n            logits,\n            targets\n        )\n\n\n        tv = focal_tversky_loss(\n\n            logits,\n\n            targets,\n\n            CFG.tversky_alpha,\n\n            CFG.tversky_beta,\n\n            CFG.tversky_gamma\n        )\n\n\n        return (\n\n            CFG.bce * bce\n\n            +\n\n            CFG.dice * dice\n\n            +\n\n            CFG.tversky * tv\n        )\n\n\n# ============================================================\n# 14. METRICS\n# ============================================================\n\ndef calculate_metrics(\n    tp,\n    fp,\n    fn\n):\n\n    dice = (\n        2 * tp + 1e-6\n    ) / (\n        2 * tp\n        + fp\n        + fn\n        + 1e-6\n    )\n\n\n    iou = (\n        tp + 1e-6\n    ) / (\n        tp\n        + fp\n        + fn\n        + 1e-6\n    )\n\n\n    precision = (\n        tp + 1e-6\n    ) / (\n        tp\n        + fp\n        + 1e-6\n    )\n\n\n    recall = (\n        tp + 1e-6\n    ) / (\n        tp\n        + fn\n        + 1e-6\n    )\n\n\n    return {\n\n        \"dice\": float(dice),\n\n        \"iou\": float(iou),\n\n        \"precision\": float(\n            precision\n        ),\n\n        \"recall\": float(\n            recall\n        )\n    }\n\n\n# ============================================================\n# 15. EPOCH RUNNER\n# ============================================================\n\ndef run_epoch(\n    model,\n    loader,\n    criterion,\n    optimizer=None,\n    scaler=None\n):\n\n    training = (\n        optimizer is not None\n    )\n\n\n    if training:\n\n        model.train()\n\n    else:\n\n        model.eval()\n\n\n    total_loss = 0.0\n\n    tp = 0\n\n    fp = 0\n\n    fn = 0\n\n\n    for imgs, masks in loader:\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        masks = masks.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n\n        if training:\n\n            optimizer.zero_grad(\n                set_to_none=True\n            )\n\n\n        with torch.set_grad_enabled(\n            training\n        ):\n\n            with autocast(\n                enabled=(\n                    CFG.device == \"cuda\"\n                    and CFG.use_amp\n                )\n            ):\n\n                logits = model(\n                    imgs\n                )\n\n                loss = criterion(\n                    logits,\n                    masks\n                )\n\n\n            if training:\n\n                if (\n                    scaler is not None\n                    and\n                    scaler.is_enabled()\n                ):\n\n                    scaler.scale(\n                        loss\n                    ).backward()\n\n                    scaler.unscale_(\n                        optimizer\n                    )\n\n                    torch.nn.utils.clip_grad_norm_(\n                        model.parameters(),\n                        CFG.grad_clip\n                    )\n\n                    scaler.step(\n                        optimizer\n                    )\n\n                    scaler.update()\n\n                else:\n\n                    loss.backward()\n\n                    torch.nn.utils.clip_grad_norm_(\n                        model.parameters(),\n                        CFG.grad_clip\n                    )\n\n                    optimizer.step()\n\n\n        probs = torch.sigmoid(\n            logits.detach()\n        )\n\n\n        preds = (\n            probs > 0.5\n        )\n\n\n        gt = (\n            masks > 0.5\n        )\n\n\n        tp += int(\n            (\n                preds & gt\n            ).sum().item()\n        )\n\n\n        fp += int(\n            (\n                preds & (~gt)\n            ).sum().item()\n        )\n\n\n        fn += int(\n            (\n                (~preds) & gt\n            ).sum().item()\n        )\n\n\n        total_loss += (\n            float(\n                loss.item()\n            )\n        )\n\n\n        del (\n            imgs,\n            masks,\n            logits,\n            probs,\n            preds\n        )\n\n\n    metrics = calculate_metrics(\n        tp,\n        fp,\n        fn\n    )\n\n\n    metrics[\"loss\"] = (\n        total_loss\n        /\n        max(\n            len(loader),\n            1\n        )\n    )\n\n\n    cleanup_memory()\n\n\n    return metrics\n\n\n# ============================================================\n# 16. BUILD TRAIN DATA\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"BUILDING V4 TRAINING DATA\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\ntrain_volumes = {}\n\ntrain_labels = {}\n\ntrain_masks = {}\n\nall_samples = []\n\n\nfor fid in CFG.train_frags:\n\n    frag_dir = os.path.join(\n        CFG.base_dir,\n        fid\n    )\n\n\n    mask = load_tissue_mask(\n        frag_dir\n    )\n\n\n    labels = load_ink_labels(\n        frag_dir\n    )\n\n\n    if labels is None:\n\n        raise FileNotFoundError(\n            \"inklabels.png missing \"\n            f\"for fragment {fid}\"\n        )\n\n\n    vol = FragmentVolume(\n        frag_dir,\n        CFG.depth_indices\n    )\n\n\n    coords = generate_grid_coords(\n\n        mask,\n\n        CFG.patch_size,\n\n        CFG.train_stride,\n\n        CFG.min_tissue_frac_train\n    )\n\n\n    print(\n        f\"Fragment {fid}: \"\n        f\"{mask.shape} \"\n        f\"{len(coords)} patches\"\n    )\n\n\n    train_volumes[fid] = vol\n\n    train_labels[fid] = labels\n\n    train_masks[fid] = mask\n\n\n    all_samples.extend(\n        [\n            (fid, y, x)\n            for y, x in coords\n        ]\n    )\n\n\nprint(\n    \"Total candidate patches:\",\n    len(all_samples)\n)\n\n\n# ============================================================\n# 17. SPATIAL TRAIN/VAL SPLIT\n# ============================================================\n\ngroups = {}\n\n\nfor fid, y, x in all_samples:\n\n    gy = (\n        y\n        // CFG.validation_block_size\n    )\n\n    gx = (\n        x\n        // CFG.validation_block_size\n    )\n\n\n    key = (\n        fid,\n        gy,\n        gx\n    )\n\n\n    groups.setdefault(\n        key,\n        []\n    ).append(\n        (\n            fid,\n            y,\n            x\n        )\n    )\n\n\nkeys = list(\n    groups.keys()\n)\n\n\nrng = random.Random(\n    CFG.seed\n)\n\nrng.shuffle(\n    keys\n)\n\n\nn_val_groups = max(\n    1,\n    int(\n        len(keys)\n        * CFG.val_fraction\n    )\n)\n\n\nval_keys = set(\n    keys[:n_val_groups]\n)\n\n\ntrain_grid = []\n\nval_samples = []\n\n\nfor key, items in groups.items():\n\n    if key in val_keys:\n\n        val_samples.extend(\n            items\n        )\n\n    else:\n\n        train_grid.extend(\n            items\n        )\n\n\nprint(\n    \"Spatial train:\",\n    len(train_grid)\n)\n\nprint(\n    \"Spatial val:\",\n    len(val_samples)\n)\n\n\n# ============================================================\n# 18. BUILD V4 SAMPLING POOLS\n# ============================================================\n\npositive_pool = []\n\nhard_negative_pool = []\n\nnormal_pool = []\n\n\nfor sample in train_grid:\n\n    fid, y, x = sample\n\n\n    lbl = train_labels[fid][\n        y:y+CFG.patch_size,\n        x:x+CFG.patch_size\n    ]\n\n\n    tissue = train_masks[fid][\n        y:y+CFG.patch_size,\n        x:x+CFG.patch_size\n    ]\n\n\n    ink_fraction = float(\n        lbl.mean()\n    )\n\n\n    tissue_fraction = float(\n        tissue.mean()\n    )\n\n\n    if (\n        ink_fraction\n        >= CFG.positive_patch_fraction\n    ):\n\n        positive_pool.append(\n            sample\n        )\n\n\n    elif (\n        tissue_fraction\n        >= CFG.hard_negative_tissue_fraction\n    ):\n\n        hard_negative_pool.append(\n            sample\n        )\n\n\n    else:\n\n        normal_pool.append(\n            sample\n        )\n\n\nprint(\n    \"\\nSampling pools:\"\n)\n\nprint(\n    \"Positive:\",\n    len(positive_pool)\n)\n\nprint(\n    \"Hard negative:\",\n    len(hard_negative_pool)\n)\n\nprint(\n    \"Normal:\",\n    len(normal_pool)\n)\n\n\n# ============================================================\n# 19. BALANCED BUT NOT OVER-SAMPLED TRAIN SET\n# ============================================================\n\nN = min(\n    CFG.max_train_samples,\n    len(train_grid)\n)\n\n\nn_positive = int(\n    N * CFG.positive_ratio\n)\n\nn_hard = int(\n    N * CFG.hard_negative_ratio\n)\n\nn_normal = (\n    N\n    - n_positive\n    - n_hard\n)\n\n\ndef sample_pool(\n    pool,\n    n\n):\n\n    if len(pool) == 0:\n\n        return []\n\n\n    if len(pool) <= n:\n\n        return pool.copy()\n\n\n    return random.sample(\n        pool,\n        n\n    )\n\n\ntrain_samples = []\n\n\ntrain_samples.extend(\n    sample_pool(\n        positive_pool,\n        n_positive\n    )\n)\n\n\ntrain_samples.extend(\n    sample_pool(\n        hard_negative_pool,\n        n_hard\n    )\n)\n\n\ntrain_samples.extend(\n    sample_pool(\n        normal_pool,\n        n_normal\n    )\n)\n\n\n# Fill any shortage using all remaining training\n# patches, without excessive repetition.\n\nif len(train_samples) < N:\n\n    selected = set(\n        train_samples\n    )\n\n\n    remaining = [\n        s\n        for s in train_grid\n        if s not in selected\n    ]\n\n\n    need = min(\n        N - len(train_samples),\n        len(remaining)\n    )\n\n\n    if need > 0:\n\n        train_samples.extend(\n            random.sample(\n                remaining,\n                need\n            )\n        )\n\n\nrandom.shuffle(\n    train_samples\n)\n\n\nprint(\n    \"\\nFinal sampled train patches:\",\n    len(train_samples)\n)\n\n\n# ============================================================\n# 20. DATASETS\n# ============================================================\n\ntrain_dataset = InkPatchDataset(\n\n    train_volumes,\n\n    train_labels,\n\n    train_samples,\n\n    CFG.patch_size,\n\n    transform=build_train_transform()\n)\n\n\nval_dataset = InkPatchDataset(\n\n    train_volumes,\n\n    train_labels,\n\n    val_samples,\n\n    CFG.patch_size,\n\n    transform=None\n)\n\n\ntrain_loader = DataLoader(\n\n    train_dataset,\n\n    batch_size=CFG.batch_size,\n\n    shuffle=True,\n\n    num_workers=CFG.num_workers,\n\n    pin_memory=(\n        CFG.device == \"cuda\"\n    ),\n\n    drop_last=CFG.drop_last,\n\n    persistent_workers=(\n        CFG.num_workers > 0\n    )\n)\n\n\nval_loader = DataLoader(\n\n    val_dataset,\n\n    batch_size=CFG.batch_size,\n\n    shuffle=False,\n\n    num_workers=CFG.num_workers,\n\n    pin_memory=(\n        CFG.device == \"cuda\"\n    ),\n\n    persistent_workers=(\n        CFG.num_workers > 0\n    )\n)\n\n\n# ============================================================\n# 21. POSITIVE PRIOR\n# ============================================================\n\nprior_samples = random.sample(\n\n    train_samples,\n\n    min(\n        120,\n        len(train_samples)\n    )\n)\n\n\ntotal_pixels = 0\n\npositive_pixels = 0\n\n\nfor fid, y, x in prior_samples:\n\n    patch = train_labels[fid][\n        y:y+CFG.patch_size,\n        x:x+CFG.patch_size\n    ]\n\n\n    positive_pixels += int(\n        patch.sum()\n    )\n\n\n    total_pixels += int(\n        patch.size\n    )\n\n\npositive_fraction = (\n    positive_pixels\n    /\n    max(\n        total_pixels,\n        1\n    )\n)\n\n\npositive_fraction = float(\n    np.clip(\n        positive_fraction,\n        1e-4,\n        0.25\n    )\n)\n\n\n# Bias initialization.\n\ninitial_bias = math.log(\n    positive_fraction\n    /\n    max(\n        1 - positive_fraction,\n        1e-6\n    )\n)\n\n\n# More conservative BCE weighting than V3.\n\nbce_pos_weight = float(\n    np.clip(\n        np.sqrt(\n            (1-positive_fraction)\n            /\n            positive_fraction\n        ),\n        1.0,\n        6.0\n    )\n)\n\n\nprint(\n    \"\\nEstimated positive fraction:\",\n    positive_fraction\n)\n\nprint(\n    \"Initial output bias:\",\n    initial_bias\n)\n\nprint(\n    \"BCE positive weight:\",\n    bce_pos_weight\n)\n\n\n# ============================================================\n# 22. MODEL\n# ============================================================\n\nmodel = build_model().to(\n    CFG.device\n)\n\n\nwith torch.no_grad():\n\n    try:\n\n        model.segmentation_head[\n            0\n        ].bias.fill_(\n            initial_bias\n        )\n\n    except Exception:\n\n        pass\n\n\n# ============================================================\n# 23. LOSS\n# ============================================================\n\npos_weight_tensor = torch.tensor(\n\n    [bce_pos_weight],\n\n    dtype=torch.float32,\n\n    device=CFG.device\n)\n\n\ncriterion = V4Loss(\n    pos_weight_tensor\n)\n\n\n# ============================================================\n# 24. DIFFERENTIAL OPTIMIZER\n# ============================================================\n\nencoder_parameters = list(\n    model.encoder.parameters()\n)\n\n\nencoder_ids = {\n    id(p)\n    for p in encoder_parameters\n}\n\n\ndecoder_parameters = [\n\n    p\n\n    for p in model.parameters()\n\n    if id(p)\n    not in encoder_ids\n]\n\n\noptimizer = torch.optim.AdamW(\n\n    [\n\n        {\n            \"params\":\n                encoder_parameters,\n\n            \"lr\":\n                CFG.encoder_lr\n        },\n\n        {\n            \"params\":\n                decoder_parameters,\n\n            \"lr\":\n                CFG.decoder_lr\n        }\n\n    ],\n\n    weight_decay=CFG.weight_decay\n)\n\n\nscheduler = (\n    torch.optim.lr_scheduler\n    .CosineAnnealingLR(\n\n        optimizer,\n\n        T_max=CFG.epochs,\n\n        eta_min=(\n            CFG.encoder_lr\n            * 0.10\n        )\n    )\n)\n\n\nscaler = GradScaler(\n\n    enabled=(\n        CFG.device == \"cuda\"\n        and CFG.use_amp\n    )\n)\n\n\n# ============================================================\n# 25. TRAINING\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"STARTING V4 TRAINING\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\nbest_val_dice = -1.0\n\nbest_epoch = -1\n\nno_improve = 0\n\n\nhistory = {\n\n    \"train_loss\": [],\n\n    \"train_dice\": [],\n\n    \"train_precision\": [],\n\n    \"train_recall\": [],\n\n    \"val_loss\": [],\n\n    \"val_dice\": [],\n\n    \"val_precision\": [],\n\n    \"val_recall\": [],\n\n    \"encoder_lr\": [],\n\n    \"decoder_lr\": []\n}\n\n\nfor epoch in range(\n    1,\n    CFG.epochs + 1\n):\n\n    t0 = time.time()\n\n\n    train_metrics = run_epoch(\n\n        model,\n\n        train_loader,\n\n        criterion,\n\n        optimizer,\n\n        scaler\n    )\n\n\n    val_metrics = run_epoch(\n\n        model,\n\n        val_loader,\n\n        criterion,\n\n        optimizer=None,\n\n        scaler=None\n    )\n\n\n    scheduler.step()\n\n\n    encoder_lr_now = (\n        optimizer.param_groups[0][\"lr\"]\n    )\n\n    decoder_lr_now = (\n        optimizer.param_groups[1][\"lr\"]\n    )\n\n\n    history[\n        \"train_loss\"\n    ].append(\n        train_metrics[\"loss\"]\n    )\n\n    history[\n        \"train_dice\"\n    ].append(\n        train_metrics[\"dice\"]\n    )\n\n    history[\n        \"train_precision\"\n    ].append(\n        train_metrics[\"precision\"]\n    )\n\n    history[\n        \"train_recall\"\n    ].append(\n        train_metrics[\"recall\"]\n    )\n\n\n    history[\n        \"val_loss\"\n    ].append(\n        val_metrics[\"loss\"]\n    )\n\n    history[\n        \"val_dice\"\n    ].append(\n        val_metrics[\"dice\"]\n    )\n\n    history[\n        \"val_precision\"\n    ].append(\n        val_metrics[\"precision\"]\n    )\n\n    history[\n        \"val_recall\"\n    ].append(\n        val_metrics[\"recall\"]\n    )\n\n\n    history[\n        \"encoder_lr\"\n    ].append(\n        encoder_lr_now\n    )\n\n    history[\n        \"decoder_lr\"\n    ].append(\n        decoder_lr_now\n    )\n\n\n    print(\n        f\"\\n[{epoch:02d}/{CFG.epochs}] \"\n        f\"time={time.time()-t0:.1f}s\"\n    )\n\n\n    print(\n        f\"train_loss=\"\n        f\"{train_metrics['loss']:.5f} \"\n        f\"train_dice=\"\n        f\"{train_metrics['dice']:.5f} \"\n        f\"train_P=\"\n        f\"{train_metrics['precision']:.5f} \"\n        f\"train_R=\"\n        f\"{train_metrics['recall']:.5f}\"\n    )\n\n\n    print(\n        f\"val_loss=\"\n        f\"{val_metrics['loss']:.5f} \"\n        f\"val_dice=\"\n        f\"{val_metrics['dice']:.5f} \"\n        f\"val_P=\"\n        f\"{val_metrics['precision']:.5f} \"\n        f\"val_R=\"\n        f\"{val_metrics['recall']:.5f}\"\n    )\n\n\n    print(\n        f\"encoder_lr=\"\n        f\"{encoder_lr_now:.7f} \"\n        f\"decoder_lr=\"\n        f\"{decoder_lr_now:.7f}\"\n    )\n\n\n    if (\n        val_metrics[\"dice\"]\n        >\n        best_val_dice\n    ):\n\n        best_val_dice = (\n            val_metrics[\"dice\"]\n        )\n\n        best_epoch = epoch\n\n        no_improve = 0\n\n\n        torch.save(\n\n            {\n                \"model\":\n                    model.state_dict(),\n\n                \"epoch\":\n                    epoch,\n\n                \"val_dice\":\n                    best_val_dice\n            },\n\n            CFG.ckpt_path\n        )\n\n\n        print(\n            \"*** NEW BEST CHECKPOINT ***\"\n        )\n\n\n    else:\n\n        no_improve += 1\n\n        print(\n            f\"No improvement: \"\n            f\"{no_improve}/\"\n            f\"{CFG.early_stop_patience}\"\n        )\n\n\n        if (\n            epoch >= CFG.min_epochs\n            and\n            no_improve\n            >= CFG.early_stop_patience\n        ):\n\n            print(\n                \"Early stopping.\"\n            )\n\n            break\n\n\n    cleanup_memory()\n\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"TRAINING COMPLETE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Best validation Dice:\",\n    best_val_dice\n)\n\nprint(\n    \"Best epoch:\",\n    best_epoch\n)\n\n\n# ============================================================\n# 26. LOAD BEST CHECKPOINT\n# ============================================================\n\ncheckpoint = torch.load(\n\n    CFG.ckpt_path,\n\n    map_location=CFG.device\n)\n\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\n\nmodel.eval()\n\n\ncleanup_memory()\n\n\n# ============================================================\n# 27. VALIDATION PROBABILITY COLLECTION\n# ============================================================\n\n@torch.no_grad()\ndef collect_validation_predictions():\n\n    model.eval()\n\n\n    all_probs = []\n\n    all_targets = []\n\n\n    for imgs, masks in val_loader:\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n\n        with autocast(\n            enabled=(\n                CFG.device == \"cuda\"\n                and CFG.use_amp\n            )\n        ):\n\n            logits = model(\n                imgs\n            )\n\n\n        probs = torch.sigmoid(\n            logits\n        ).float().cpu().numpy()\n\n\n        all_probs.append(\n            probs[:, 0]\n        )\n\n\n        all_targets.append(\n            masks.numpy()[:, 0]\n        )\n\n\n        del (\n            imgs,\n            masks,\n            logits,\n            probs\n        )\n\n\n    return (\n\n        np.concatenate(\n            all_probs,\n            axis=0\n        ),\n\n        np.concatenate(\n            all_targets,\n            axis=0\n        )\n    )\n\n\nval_probabilities, val_targets = (\n    collect_validation_predictions()\n)\n\n\n# ============================================================\n# 28. VALIDATION THRESHOLD SEARCH\n# ============================================================\n\ndef threshold_metrics_numpy(\n\n    probabilities,\n\n    targets,\n\n    threshold\n):\n\n    pred = (\n        probabilities\n        >= threshold\n    )\n\n\n    gt = (\n        targets\n        > 0.5\n    )\n\n\n    tp = np.logical_and(\n        pred,\n        gt\n    ).sum()\n\n\n    fp = np.logical_and(\n        pred,\n        ~gt\n    ).sum()\n\n\n    fn = np.logical_and(\n        ~pred,\n        gt\n    ).sum()\n\n\n    return calculate_metrics(\n        tp,\n        fp,\n        fn\n    )\n\n\nvalidation_threshold_results = []\n\n\nfor threshold in CFG.val_thresholds:\n\n    m = threshold_metrics_numpy(\n\n        val_probabilities,\n\n        val_targets,\n\n        threshold\n    )\n\n\n    validation_threshold_results.append(\n\n        {\n            \"threshold\":\n                float(threshold),\n\n            \"dice\":\n                m[\"dice\"],\n\n            \"iou\":\n                m[\"iou\"],\n\n            \"precision\":\n                m[\"precision\"],\n\n            \"recall\":\n                m[\"recall\"]\n        }\n    )\n\n\nbest_val_result = max(\n\n    validation_threshold_results,\n\n    key=lambda x:\n        x[\"dice\"]\n)\n\n\nVAL_THRESHOLD = float(\n    best_val_result[\n        \"threshold\"\n    ]\n)\n\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"VALIDATION THRESHOLD\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Validation threshold:\",\n    VAL_THRESHOLD\n)\n\nprint(\n    \"Validation metrics:\",\n    best_val_result\n)\n\n\n# ============================================================\n# 29. VALIDATION POSITIVE RATE\n# ============================================================\n\nval_best_prediction = (\n    val_probabilities\n    >= VAL_THRESHOLD\n)\n\n\nval_positive_rate = float(\n    val_best_prediction.mean()\n)\n\n\nprint(\n    \"Validation predicted-positive rate:\",\n    val_positive_rate\n)\n\n\n# ============================================================\n# 30. FREE TRAINING MEMORY\n# ============================================================\n\nfor vol in train_volumes.values():\n\n    vol.close()\n\n\ndel (\n    train_loader,\n    val_loader,\n    train_dataset,\n    val_dataset,\n    train_volumes,\n    train_labels,\n    train_masks\n)\n\n\ncleanup_memory()\n\n\n# ============================================================\n# 31. LOAD TEST FRAGMENT\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"LOADING TEST FRAGMENT\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\ntest_dir = os.path.join(\n\n    CFG.base_dir,\n\n    CFG.test_frag\n)\n\n\ntest_mask = load_tissue_mask(\n    test_dir\n)\n\n\n# Only local diagnostics.\n#\n# NEVER used to select threshold.\n\ntest_labels = load_ink_labels(\n    test_dir\n)\n\n\ntest_vol = FragmentVolume(\n\n    test_dir,\n\n    CFG.depth_indices\n)\n\n\nprint(\n    \"Test shape:\",\n    test_mask.shape\n)\n\n\nprint(\n    \"Local test GT available:\",\n    test_labels is not None\n)\n\n\n# ============================================================\n# 32. GAUSSIAN BLENDING WINDOW\n# ============================================================\n\ndef gaussian_window(\n    size,\n    sigma_fraction=0.42\n):\n\n    axis = (\n        np.arange(\n            size,\n            dtype=np.float32\n        )\n        -\n        (size - 1) / 2\n    )\n\n\n    sigma = max(\n\n        size\n        * sigma_fraction,\n\n        1.0\n    )\n\n\n    g = np.exp(\n\n        -(axis ** 2)\n        /\n        (2 * sigma ** 2)\n    )\n\n\n    window = np.outer(\n        g,\n        g\n    ).astype(\n        np.float32\n    )\n\n\n    window /= max(\n\n        float(\n            window.max()\n        ),\n\n        1e-6\n    )\n\n\n    # Prevent extremely small\n    # boundary weights.\n\n    window = np.maximum(\n        window,\n        0.05\n    )\n\n\n    return window.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 33. INPUT PREPARATION\n# ============================================================\n\ndef prepare_input(\n    raw\n):\n\n    raw = (\n        raw.astype(\n            np.float32\n        )\n        / 255.0\n    )\n\n\n    raw = normalize_patch(\n        raw\n    )\n\n\n    return np.ascontiguousarray(\n        raw\n    )\n\n\n# ============================================================\n# 34. SIMPLE INFERENCE\n# ============================================================\n\n@torch.no_grad()\ndef predict_batch(\n    batch\n):\n\n    inp = torch.from_numpy(\n        np.stack(batch)\n    ).to(\n        CFG.device,\n        non_blocking=True\n    )\n\n\n    with autocast(\n        enabled=(\n            CFG.device == \"cuda\"\n            and CFG.use_amp\n        )\n    ):\n\n        logits = model(\n            inp\n        )\n\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n\n    result = (\n        probs\n        .float()\n        .cpu()\n        .numpy()[:, 0]\n    )\n\n\n    del (\n        inp,\n        logits,\n        probs\n    )\n\n\n    return result\n\n\n# ============================================================\n# 35. SLIDING WINDOW INFERENCE\n# ============================================================\n\n@torch.no_grad()\ndef sliding_window_inference(\n\n    volume,\n\n    tissue_mask,\n\n    stride=None,\n\n    transform=None\n\n):\n\n    if stride is None:\n\n        stride = CFG.test_stride\n\n\n    H, W = (\n        tissue_mask.shape\n    )\n\n\n    prediction_sum = np.zeros(\n\n        (H, W),\n\n        dtype=np.float32\n    )\n\n\n    weight_sum = np.zeros(\n\n        (H, W),\n\n        dtype=np.float32\n    )\n\n\n    window = gaussian_window(\n        CFG.patch_size\n    )\n\n\n    coords = generate_grid_coords(\n\n        tissue_mask,\n\n        CFG.patch_size,\n\n        stride,\n\n        CFG.min_tissue_frac_test\n    )\n\n\n    print(\n        \"Inference patches:\",\n        len(coords)\n    )\n\n\n    model.eval()\n\n\n    batch_images = []\n\n    batch_coords = []\n\n\n    def flush():\n\n        if len(batch_images) == 0:\n\n            return\n\n\n        probabilities = (\n            predict_batch(\n                batch_images\n            )\n        )\n\n\n        for p, (y, x) in zip(\n\n            probabilities,\n\n            batch_coords\n        ):\n\n            if transform is not None:\n\n                # transform is only used for\n                # Rot90 ablation.\n                #\n                # Reverse it before blending.\n\n                p = transform(\n                    p\n                )\n\n\n            prediction_sum[\n                y:y+CFG.patch_size,\n                x:x+CFG.patch_size\n            ] += (\n                p * window\n            )\n\n\n            weight_sum[\n                y:y+CFG.patch_size,\n                x:x+CFG.patch_size\n            ] += window\n\n\n        batch_images.clear()\n\n        batch_coords.clear()\n\n\n        cleanup_memory()\n\n\n    for y, x in coords:\n\n        raw = volume.read_patch(\n\n            y,\n\n            x,\n\n            CFG.patch_size\n        )\n\n\n        image = prepare_input(\n            raw\n        )\n\n\n        batch_images.append(\n            image\n        )\n\n\n        batch_coords.append(\n            (\n                y,\n                x\n            )\n        )\n\n\n        if (\n            len(batch_images)\n            >= CFG.infer_batch\n        ):\n\n            flush()\n\n\n    flush()\n\n\n    weight_sum = np.maximum(\n\n        weight_sum,\n\n        1e-6\n    )\n\n\n    probability = (\n\n        prediction_sum\n\n        /\n\n        weight_sum\n    )\n\n\n    probability[\n        tissue_mask == 0\n    ] = 0.0\n\n\n    return probability.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 36. ADABN\n# ============================================================\n\n@torch.no_grad()\ndef run_adabn():\n\n    print(\n        \"\\n\"\n        + \"=\" * 70\n    )\n\n    print(\n        \"STARTING AdaBN\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n\n    # Reset BN statistics.\n\n    bn_layers = 0\n\n\n    for module in model.modules():\n\n        if isinstance(\n            module,\n            nn.BatchNorm2d\n        ):\n\n            module.reset_running_stats()\n\n            module.momentum = None\n\n            bn_layers += 1\n\n\n    print(\n        \"BN layers:\",\n        bn_layers\n    )\n\n\n    coords = generate_grid_coords(\n\n        test_mask,\n\n        CFG.patch_size,\n\n        CFG.test_stride,\n\n        CFG.min_tissue_frac_test\n    )\n\n\n    # Deterministic but spatially mixed.\n\n    rng = random.Random(\n        CFG.seed\n    )\n\n    rng.shuffle(\n        coords\n    )\n\n\n    coords = coords[\n        :CFG.adabn_max_patches\n    ]\n\n\n    print(\n        \"AdaBN patches:\",\n        len(coords)\n    )\n\n\n    model.train()\n\n\n    for start in range(\n\n        0,\n\n        len(coords),\n\n        CFG.infer_batch\n    ):\n\n        batch_coords = coords[\n            start:\n            start + CFG.infer_batch\n        ]\n\n\n        images = []\n\n\n        for y, x in batch_coords:\n\n            raw = test_vol.read_patch(\n\n                y,\n\n                x,\n\n                CFG.patch_size\n            )\n\n\n            image = prepare_input(\n                raw\n            )\n\n\n            images.append(\n                image\n            )\n\n\n        inp = torch.from_numpy(\n            np.stack(images)\n        ).to(\n            CFG.device,\n            non_blocking=True\n        )\n\n\n        # IMPORTANT:\n        #\n        # NO TTA.\n        #\n        # We want natural target-domain\n        # statistics only.\n\n        with autocast(\n\n            enabled=(\n\n                CFG.device == \"cuda\"\n                and CFG.use_amp\n            )\n        ):\n\n            model(inp)\n\n\n        del (\n            inp,\n            images\n        )\n\n\n    model.eval()\n\n\n    cleanup_memory()\n\n\n    print(\n        \"AdaBN finished.\"\n    )\n\n\n# ============================================================\n# 37. BASELINE TEST INFERENCE\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"BASELINE TEST INFERENCE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\n# Start from clean best checkpoint.\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\nmodel.eval()\n\n\nbaseline_probability = (\n    sliding_window_inference(\n\n        test_vol,\n\n        test_mask\n    )\n)\n\n\n# ============================================================\n# 38. ADABN TEST INFERENCE\n# ============================================================\n\n# Reload checkpoint before AdaBN.\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\nmodel.eval()\n\n\nif CFG.use_adabn:\n\n    run_adabn()\n\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"ADABN TEST INFERENCE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\nadabn_probability = (\n    sliding_window_inference(\n\n        test_vol,\n\n        test_mask\n    )\n)\n\n\n# ============================================================\n# 39. TEST DISTRIBUTION STATISTICS\n# ============================================================\n\ndef probability_statistics(\n    probability,\n    mask\n):\n\n    values = probability[\n        mask > 0\n    ]\n\n\n    values = values[\n        np.isfinite(values)\n    ]\n\n\n    return {\n\n        \"min\":\n            float(\n                np.min(values)\n            ),\n\n        \"max\":\n            float(\n                np.max(values)\n            ),\n\n        \"mean\":\n            float(\n                np.mean(values)\n            ),\n\n        \"median\":\n            float(\n                np.median(values)\n            ),\n\n        \"p90\":\n            float(\n                np.percentile(\n                    values,\n                    90\n                )\n            ),\n\n        \"p95\":\n            float(\n                np.percentile(\n                    values,\n                    95\n                )\n            ),\n\n        \"p97\":\n            float(\n                np.percentile(\n                    values,\n                    97\n                )\n            ),\n\n        \"p98\":\n            float(\n                np.percentile(\n                    values,\n                    98\n                )\n            ),\n\n        \"p99\":\n            float(\n                np.percentile(\n                    values,\n                    99\n                )\n            ),\n\n        \"p99.5\":\n            float(\n                np.percentile(\n                    values,\n                    99.5\n                )\n            )\n    }\n\n\nprint(\n    \"\\nADABN probability statistics:\"\n)\n\nprint(\n    probability_statistics(\n        adabn_probability,\n        test_mask\n    )\n)\n\n\n# ============================================================\n# 40. INDEPENDENT TEST THRESHOLD CALIBRATION\n# ============================================================\n\ndef get_test_values(\n    probability\n):\n\n    values = probability[\n        test_mask > 0\n    ]\n\n\n    values = values[\n        np.isfinite(values)\n    ]\n\n\n    values = values[\n        values > 1e-6\n    ]\n\n\n    return values\n\n\ntest_values = get_test_values(\n    adabn_probability\n)\n\n\n# ------------------------------------------------------------\n# A. OTSU\n# ------------------------------------------------------------\n\ntry:\n\n    TEST_THRESHOLD_OTSU = float(\n        threshold_otsu(\n            test_values\n        )\n    )\n\nexcept Exception:\n\n    TEST_THRESHOLD_OTSU = 0.5\n\n\nTEST_THRESHOLD_OTSU = float(\n\n    np.clip(\n\n        TEST_THRESHOLD_OTSU,\n\n        CFG.test_threshold_min,\n\n        CFG.test_threshold_max\n    )\n)\n\n\n# ------------------------------------------------------------\n# B. QUANTILE\n# ------------------------------------------------------------\n\n# We don't assume that the validation threshold\n# probability scale is directly transferable.\n#\n# Instead use the fraction of positive pixels\n# implied by the validation model as an UNSUPERVISED\n# target-fragment calibration.\n#\n# This is not the same threshold.\n#\n# It is a target-domain quantile.\n\nTEST_THRESHOLD_POSITIVE_RATE = float(\n\n    np.quantile(\n\n        test_values,\n\n        1.0 - val_positive_rate\n    )\n)\n\n\nTEST_THRESHOLD_POSITIVE_RATE = float(\n\n    np.clip(\n\n        TEST_THRESHOLD_POSITIVE_RATE,\n\n        CFG.test_threshold_min,\n\n        CFG.test_threshold_max\n    )\n)\n\n\n# ------------------------------------------------------------\n# C. PERCENTILE FALLBACK\n# ------------------------------------------------------------\n\nTEST_THRESHOLD_P98 = float(\n\n    np.percentile(\n\n        test_values,\n\n        98.0\n    )\n)\n\n\nTEST_THRESHOLD_P98 = float(\n\n    np.clip(\n\n        TEST_THRESHOLD_P98,\n\n        CFG.test_threshold_min,\n\n        CFG.test_threshold_max\n    )\n)\n\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"INDEPENDENT TEST THRESHOLD CALIBRATION\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Validation threshold:\",\n    VAL_THRESHOLD\n)\n\nprint(\n    \"Test Otsu:\",\n    TEST_THRESHOLD_OTSU\n)\n\nprint(\n    \"Test positive-rate threshold:\",\n    TEST_THRESHOLD_POSITIVE_RATE\n)\n\nprint(\n    \"Test P98:\",\n    TEST_THRESHOLD_P98\n)\n\n\n# ============================================================\n# 41. LOCAL DIAGNOSTIC FUNCTION\n# ============================================================\n\ndef local_metrics(\n    probability,\n    threshold\n):\n\n    if test_labels is None:\n\n        return None\n\n\n    gt = (\n        test_labels > 0\n    )\n\n\n    gt = (\n        gt\n        &\n        (test_mask > 0)\n    )\n\n\n    pred = (\n        probability\n        >= threshold\n    )\n\n\n    pred = (\n        pred\n        &\n        (test_mask > 0)\n    )\n\n\n    tp = np.logical_and(\n        pred,\n        gt\n    ).sum()\n\n\n    fp = np.logical_and(\n        pred,\n        ~gt\n    ).sum()\n\n\n    fn = np.logical_and(\n        ~pred,\n        gt\n    ).sum()\n\n\n    return calculate_metrics(\n\n        tp,\n\n        fp,\n\n        fn\n    )\n\n\n# ============================================================\n# 42. TEST THRESHOLD DIAGNOSTICS\n# ============================================================\n\ncandidate_thresholds = {\n\n    \"otsu\":\n        TEST_THRESHOLD_OTSU,\n\n    \"positive_rate\":\n        TEST_THRESHOLD_POSITIVE_RATE,\n\n    \"p98\":\n        TEST_THRESHOLD_P98,\n\n    \"validation_threshold_diagnostic\":\n        float(\n            np.clip(\n                VAL_THRESHOLD,\n                CFG.test_threshold_min,\n                CFG.test_threshold_max\n            )\n        )\n}\n\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"LOCAL TEST THRESHOLD DIAGNOSTICS\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\nlocal_threshold_results = {}\n\n\nfor name, threshold in (\n    candidate_thresholds.items()\n):\n\n    metrics = local_metrics(\n\n        adabn_probability,\n\n        threshold\n    )\n\n\n    local_threshold_results[name] = {\n\n        \"threshold\":\n            float(threshold),\n\n        \"metrics\":\n            metrics\n    }\n\n\n    print(\n        f\"\\n{name}: \"\n        f\"threshold={threshold:.4f}\"\n    )\n\n\n    if metrics is not None:\n\n        print(\n            metrics\n        )\n\n\n# ============================================================\n# 43. FINAL TEST THRESHOLD\n# ============================================================\n\nif (\n    CFG.final_test_threshold_method\n    == \"otsu\"\n):\n\n    FINAL_TEST_THRESHOLD = (\n        TEST_THRESHOLD_OTSU\n    )\n\n\nelif (\n    CFG.final_test_threshold_method\n    == \"positive_rate\"\n):\n\n    FINAL_TEST_THRESHOLD = (\n        TEST_THRESHOLD_POSITIVE_RATE\n    )\n\n\nelif (\n    CFG.final_test_threshold_method\n    == \"quantile\"\n):\n\n    FINAL_TEST_THRESHOLD = (\n        TEST_THRESHOLD_P98\n    )\n\n\nelse:\n\n    raise ValueError(\n        \"Unknown final_test_threshold_method\"\n    )\n\n\nprint(\n    \"\\nFINAL TEST THRESHOLD:\",\n    FINAL_TEST_THRESHOLD\n)\n\n\n# ============================================================\n# 44. ROT90 ABLATION\n# ============================================================\n\nrot90_probability = None\n\n\nif CFG.run_rot90_ablation:\n\n    print(\n        \"\\n\"\n        + \"=\" * 70\n    )\n\n    print(\n        \"ROT90 ABLATION\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n\n    # Reload best checkpoint.\n    #\n    # Then recalibrate BN again using natural\n    # images before Rot90 inference.\n    #\n    # This prevents the ablation from accidentally\n    # inheriting stale BN state.\n\n    model.load_state_dict(\n        checkpoint[\"model\"]\n    )\n\n    model.eval()\n\n\n    if CFG.use_adabn:\n\n        run_adabn()\n\n\n    # --------------------------------------------------------\n    # Rot90 inference.\n    #\n    # Each input patch is rotated 90 degrees.\n    # Prediction is rotated back before blending.\n    # --------------------------------------------------------\n\n    def inverse_rot90(\n        p\n    ):\n\n        return np.rot90(\n            p,\n            -1\n        ).copy()\n\n\n    rot90_probability = (\n        sliding_window_inference(\n\n            test_vol,\n\n            test_mask,\n\n            transform=inverse_rot90\n        )\n    )\n\n\n    print(\n        \"Rot90 probability map generated.\"\n    )\n\n\n    if test_labels is not None:\n\n        print(\n            \"\\nRot90 diagnostic:\"\n        )\n\n\n        print(\n            \"Rot90 + Otsu:\",\n            local_metrics(\n                rot90_probability,\n                TEST_THRESHOLD_OTSU\n            )\n        )\n\n\n        print(\n            \"Rot90 + positive-rate:\",\n            local_metrics(\n                rot90_probability,\n                TEST_THRESHOLD_POSITIVE_RATE\n            )\n        )\n\n\n# ============================================================\n# 45. FINAL PROBABILITY MAP\n# ============================================================\n\n# FINAL DEFAULT:\n#\n# AdaBN + NO TTA\n#\n# This is intentional.\n#\n# Your V3 experiment showed:\n#\n#   none        ~0.395\n#   hflip       ~0.403\n#   vflip       ~0.401\n#   hvflip      ~0.400\n#   rot90       ~0.432\n#   averages    ~0.419-0.428\n#\n# Therefore we do not force TTA into the final result.\n\nif (\n    CFG.use_tta_for_final\n    and\n    rot90_probability is not None\n):\n\n    final_probability = (\n        rot90_probability\n    )\n\n    final_variant = (\n        \"AdaBN + Rot90\"\n    )\n\nelse:\n\n    final_probability = (\n        adabn_probability\n    )\n\n    final_variant = (\n        \"AdaBN + no TTA\"\n    )\n\n\n# ============================================================\n# 46. POST PROCESSING\n# ============================================================\n\ndef remove_small_components(\n    binary,\n    min_area\n):\n\n    num_labels, labels, stats, _ = (\n        cv2.connectedComponentsWithStats(\n\n            binary.astype(\n                np.uint8\n            ),\n\n            connectivity=8\n        )\n    )\n\n\n    if num_labels <= 1:\n\n        return binary.astype(\n            np.uint8\n        )\n\n\n    output = np.zeros_like(\n        binary,\n        dtype=np.uint8\n    )\n\n\n    for label_id in range(\n        1,\n        num_labels\n    ):\n\n        area = int(\n            stats[\n                label_id,\n                cv2.CC_STAT_AREA\n            ]\n        )\n\n\n        if area >= min_area:\n\n            output[\n                labels == label_id\n            ] = 1\n\n\n    return output\n\n\ndef postprocess(\n    probability,\n    threshold,\n    tissue_mask\n):\n\n    binary = (\n\n        probability\n        >= threshold\n\n    ).astype(\n        np.uint8\n    )\n\n\n    binary[\n        tissue_mask == 0\n    ] = 0\n\n\n    if not CFG.use_postprocess:\n\n        return binary\n\n\n    # Conservative closing.\n    #\n    # No opening.\n    #\n    # We specifically avoid opening because\n    # thin ink can disappear.\n\n    kernel = np.ones(\n\n        (\n            CFG.closing_kernel,\n            CFG.closing_kernel\n        ),\n\n        dtype=np.uint8\n    )\n\n\n    binary = cv2.morphologyEx(\n\n        binary,\n\n        cv2.MORPH_CLOSE,\n\n        kernel,\n\n        iterations=1\n    )\n\n\n    binary = remove_small_components(\n\n        binary,\n\n        CFG.min_component_area\n    )\n\n\n    binary[\n        tissue_mask == 0\n    ] = 0\n\n\n    return binary.astype(\n        np.uint8\n    )\n\n\nfinal_prediction = postprocess(\n\n    final_probability,\n\n    FINAL_TEST_THRESHOLD,\n\n    test_mask\n)\n\n\n# ============================================================\n# 47. FINAL LOCAL DIAGNOSTICS\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"FINAL V4 LOCAL DIAGNOSTICS\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Final variant:\",\n    final_variant\n)\n\nprint(\n    \"Final threshold:\",\n    FINAL_TEST_THRESHOLD\n)\n\n\nfinal_raw_metrics = local_metrics(\n\n    final_probability,\n\n    FINAL_TEST_THRESHOLD\n)\n\n\nif final_raw_metrics is not None:\n\n    print(\n        \"Final RAW:\",\n        final_raw_metrics\n    )\n\n\nif test_labels is not None:\n\n    gt = (\n        test_labels > 0\n    ) & (\n        test_mask > 0\n    )\n\n\n    pred = (\n        final_prediction > 0\n    )\n\n\n    tp = np.logical_and(\n        pred,\n        gt\n    ).sum()\n\n\n    fp = np.logical_and(\n        pred,\n        ~gt\n    ).sum()\n\n\n    fn = np.logical_and(\n        ~pred,\n        gt\n    ).sum()\n\n\n    final_post_metrics = (\n        calculate_metrics(\n            tp,\n            fp,\n            fn\n        )\n    )\n\n\n    print(\n        \"Final POSTPROCESSED:\",\n        final_post_metrics\n    )\n\nelse:\n\n    final_post_metrics = None\n\n\n# ============================================================\n# 48. SAVE PROBABILITY MAPS\n# ============================================================\n\nnp.save(\n\n    os.path.join(\n        CFG.out_dir,\n        \"fragment1_probability_v4.npy\"\n    ),\n\n    final_probability\n)\n\n\nnp.save(\n\n    os.path.join(\n        CFG.out_dir,\n        \"fragment1_probability_adabn_v4.npy\"\n    ),\n\n    adabn_probability\n)\n\n\nnp.save(\n\n    os.path.join(\n        CFG.out_dir,\n        \"fragment1_probability_baseline_v4.npy\"\n    ),\n\n    baseline_probability\n)\n\n\nif rot90_probability is not None:\n\n    np.save(\n\n        os.path.join(\n            CFG.out_dir,\n            \"fragment1_probability_rot90_v4.npy\"\n        ),\n\n        rot90_probability\n    )\n\n\n# ============================================================\n# 49. SAVE FINAL PREDICTION\n# ============================================================\n\nprediction_path = os.path.join(\n\n    CFG.out_dir,\n\n    \"fragment1_prediction_v4.png\"\n)\n\n\ncv2.imwrite(\n\n    prediction_path,\n\n    (\n        final_prediction\n        * 255\n    ).astype(\n        np.uint8\n    )\n)\n\n\n# ============================================================\n# 50. SAVE THRESHOLD RESULTS\n# ============================================================\n\nthreshold_data = {\n\n    \"validation_threshold\":\n        float(VAL_THRESHOLD),\n\n    \"validation_best\":\n        best_val_result,\n\n    \"validation_positive_rate\":\n        float(val_positive_rate),\n\n    \"test_otsu\":\n        float(TEST_THRESHOLD_OTSU),\n\n    \"test_positive_rate\":\n        float(\n            TEST_THRESHOLD_POSITIVE_RATE\n        ),\n\n    \"test_p98\":\n        float(\n            TEST_THRESHOLD_P98\n        ),\n\n    \"final_test_method\":\n        CFG.final_test_threshold_method,\n\n    \"final_test_threshold\":\n        float(\n            FINAL_TEST_THRESHOLD\n        ),\n\n    \"candidate_diagnostics\":\n        local_threshold_results\n}\n\n\nwith open(\n\n    os.path.join(\n        CFG.out_dir,\n        \"threshold_diagnostics.json\"\n    ),\n\n    \"w\"\n) as f:\n\n    json.dump(\n\n        threshold_data,\n\n        f,\n\n        indent=2\n    )\n\n\n# ============================================================\n# 51. SAVE COMPLETE METRICS\n# ============================================================\n\nmetrics_summary = {\n\n    \"version\":\n        \"V4\",\n\n    \"patch_size\":\n        CFG.patch_size,\n\n    \"train_stride\":\n        CFG.train_stride,\n\n    \"test_stride\":\n        CFG.test_stride,\n\n    \"batch_size\":\n        CFG.batch_size,\n\n    \"epochs\":\n        CFG.epochs,\n\n    \"best_epoch\":\n        best_epoch,\n\n    \"best_validation_dice_050\":\n        float(\n            best_val_dice\n        ),\n\n    \"validation_threshold\":\n        float(\n            VAL_THRESHOLD\n        ),\n\n    \"validation_best_metrics\":\n        best_val_result,\n\n    \"test_threshold\":\n        float(\n            FINAL_TEST_THRESHOLD\n        ),\n\n    \"test_threshold_method\":\n        CFG.final_test_threshold_method,\n\n    \"final_variant\":\n        final_variant,\n\n    \"adabn\":\n        bool(\n            CFG.use_adabn\n        ),\n\n    \"tta_final\":\n        bool(\n            CFG.use_tta_for_final\n        ),\n\n    \"rot90_ablation\":\n        bool(\n            CFG.run_rot90_ablation\n        ),\n\n    \"probability_statistics\":\n        probability_statistics(\n            adabn_probability,\n            test_mask\n        ),\n\n    \"local_test_raw\":\n        final_raw_metrics,\n\n    \"local_test_postprocessed\":\n        final_post_metrics,\n\n    \"threshold_diagnostics\":\n        local_threshold_results\n}\n\n\nwith open(\n\n    CFG.metrics_path,\n\n    \"w\"\n\n) as f:\n\n    json.dump(\n\n        metrics_summary,\n\n        f,\n\n        indent=2\n    )\n\n\n# ============================================================\n# 52. SAVE TRAINING HISTORY\n# ============================================================\n\nwith open(\n\n    os.path.join(\n        CFG.out_dir,\n        \"training_history.json\"\n    ),\n\n    \"w\"\n\n) as f:\n\n    json.dump(\n\n        history,\n\n        f,\n\n        indent=2\n    )\n\n\n# ============================================================\n# 53. TRAINING CURVES\n# ============================================================\n\nepochs_axis = np.arange(\n\n    1,\n\n    len(\n        history[\"train_loss\"]\n    ) + 1\n)\n\n\n# ------------------------------------------------------------\n# Loss\n# ------------------------------------------------------------\n\nplt.figure(\n    figsize=(9, 5)\n)\n\n\nplt.plot(\n\n    epochs_axis,\n\n    history[\"train_loss\"],\n\n    label=\"Train loss\"\n)\n\n\nplt.plot(\n\n    epochs_axis,\n\n    history[\"val_loss\"],\n\n    label=\"Val loss\"\n)\n\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"Loss\"\n)\n\nplt.title(\n    \"V4 Training / Validation Loss\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\n\nplt.savefig(\n\n    os.path.join(\n\n        CFG.viz_dir,\n\n        \"v4_loss.png\"\n    ),\n\n    dpi=150\n)\n\n\nplt.close()\n\n\n# ------------------------------------------------------------\n# Dice\n# ------------------------------------------------------------\n\nplt.figure(\n    figsize=(9, 5)\n)\n\n\nplt.plot(\n\n    epochs_axis,\n\n    history[\"train_dice\"],\n\n    label=\"Train Dice\"\n)\n\n\nplt.plot(\n\n    epochs_axis,\n\n    history[\"val_dice\"],\n\n    label=\"Val Dice\"\n)\n\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"Dice\"\n)\n\nplt.title(\n    \"V4 Training / Validation Dice\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\n\nplt.savefig(\n\n    os.path.join(\n\n        CFG.viz_dir,\n\n        \"v4_dice.png\"\n    ),\n\n    dpi=150\n)\n\n\nplt.close()\n\n\n# ============================================================\n# 54. VALIDATION THRESHOLD CURVE\n# ============================================================\n\nthreshold_x = [\n\n    x[\"threshold\"]\n\n    for x in\n    validation_threshold_results\n]\n\n\ndice_y = [\n\n    x[\"dice\"]\n\n    for x in\n    validation_threshold_results\n]\n\n\nplt.figure(\n    figsize=(9, 5)\n)\n\n\nplt.plot(\n\n    threshold_x,\n\n    dice_y\n)\n\n\nplt.axvline(\n\n    VAL_THRESHOLD,\n\n    linestyle=\"--\",\n\n    label=(\n        f\"Best={VAL_THRESHOLD:.2f}\"\n    )\n)\n\n\nplt.xlabel(\n    \"Validation threshold\"\n)\n\nplt.ylabel(\n    \"Validation Dice\"\n)\n\nplt.title(\n    \"V4 Validation Threshold Search\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\n\nplt.savefig(\n\n    os.path.join(\n\n        CFG.viz_dir,\n\n        \"v4_validation_threshold.png\"\n    ),\n\n    dpi=150\n)\n\n\nplt.close()\n\n\n# ============================================================\n# 55. PROBABILITY HISTOGRAM\n# ============================================================\n\nplt.figure(\n    figsize=(9, 5)\n)\n\n\ntest_hist_values = (\n    adabn_probability[\n        test_mask > 0\n    ]\n)\n\n\ntest_hist_values = (\n    test_hist_values[\n        np.isfinite(\n            test_hist_values\n        )\n    ]\n)\n\n\nplt.hist(\n\n    test_hist_values,\n\n    bins=100\n)\n\n\nplt.axvline(\n\n    TEST_THRESHOLD_OTSU,\n\n    linestyle=\"--\",\n\n    label=(\n        f\"Otsu={TEST_THRESHOLD_OTSU:.3f}\"\n    )\n)\n\n\nplt.axvline(\n\n    TEST_THRESHOLD_POSITIVE_RATE,\n\n    linestyle=\"--\",\n\n    label=(\n        \"Positive-rate=\"\n        f\"{TEST_THRESHOLD_POSITIVE_RATE:.3f}\"\n    )\n)\n\n\nplt.axvline(\n\n    TEST_THRESHOLD_P98,\n\n    linestyle=\"--\",\n\n    label=(\n        f\"P98={TEST_THRESHOLD_P98:.3f}\"\n    )\n)\n\n\nplt.xlabel(\n    \"Probability\"\n)\n\nplt.ylabel(\n    \"Pixels\"\n)\n\nplt.title(\n    \"V4 Test Probability Distribution\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\n\nplt.savefig(\n\n    os.path.join(\n\n        CFG.viz_dir,\n\n        \"v4_test_probability_histogram.png\"\n    ),\n\n    dpi=150\n)\n\n\nplt.close()\n\n\n# ============================================================\n# 56. OVERVIEW IMAGE\n# ============================================================\n\nmid_idx = (\n\n    CFG.depth_indices[\n        len(CFG.depth_indices)//2\n    ]\n)\n\n\nmid_path = os.path.join(\n\n    test_dir,\n\n    \"surface_volume\",\n\n    f\"{mid_idx:02d}.tif\"\n)\n\n\nmid_slice = tifffile.imread(\n    mid_path\n)\n\n\nscale = min(\n\n    1.0,\n\n    1800.0\n    /\n    max(\n        mid_slice.shape\n    )\n)\n\n\nsmall = cv2.resize(\n\n    mid_slice,\n\n    None,\n\n    fx=scale,\n\n    fy=scale,\n\n    interpolation=cv2.INTER_AREA\n)\n\n\nprob_small = cv2.resize(\n\n    final_probability,\n\n    small.shape[::-1],\n\n    interpolation=cv2.INTER_AREA\n)\n\n\npred_small = cv2.resize(\n\n    (\n        final_prediction\n        * 255\n    ).astype(\n        np.uint8\n    ),\n\n    small.shape[::-1],\n\n    interpolation=cv2.INTER_NEAREST\n)\n\n\nif test_labels is not None:\n\n    gt_small = cv2.resize(\n\n        (\n            test_labels\n            * 255\n        ).astype(\n            np.uint8\n        ),\n\n        small.shape[::-1],\n\n        interpolation=cv2.INTER_NEAREST\n    )\n\n\n    fig, axes = plt.subplots(\n\n        1,\n\n        4,\n\n        figsize=(22, 6)\n    )\n\n\n    axes[0].imshow(\n\n        small,\n\n        cmap=\"gray\"\n    )\n\n\n    axes[0].set_title(\n        f\"Input slice {mid_idx}\"\n    )\n\n\n    axes[1].imshow(\n\n        prob_small,\n\n        cmap=\"magma\",\n\n        vmin=0,\n\n        vmax=1\n    )\n\n\n    axes[1].set_title(\n\n        \"Probability\\n\"\n        f\"T={FINAL_TEST_THRESHOLD:.3f}\"\n    )\n\n\n    axes[2].imshow(\n\n        gt_small,\n\n        cmap=\"gray\"\n    )\n\n\n    axes[2].set_title(\n        \"Local GT\"\n    )\n\n\n    axes[3].imshow(\n\n        pred_small,\n\n        cmap=\"gray\"\n    )\n\n\n    axes[3].set_title(\n        \"V4 Prediction\"\n    )\n\n\nelse:\n\n    fig, axes = plt.subplots(\n\n        1,\n\n        3,\n\n        figsize=(17, 6)\n    )\n\n\n    axes[0].imshow(\n\n        small,\n\n        cmap=\"gray\"\n    )\n\n\n    axes[0].set_title(\n        f\"Input slice {mid_idx}\"\n    )\n\n\n    axes[1].imshow(\n\n        prob_small,\n\n        cmap=\"magma\",\n\n        vmin=0,\n\n        vmax=1\n    )\n\n\n    axes[1].set_title(\n\n        \"Probability\\n\"\n        f\"T={FINAL_TEST_THRESHOLD:.3f}\"\n    )\n\n\n    axes[2].imshow(\n\n        pred_small,\n\n        cmap=\"gray\"\n    )\n\n\n    axes[2].set_title(\n        \"V4 Prediction\"\n    )\n\n\nfor ax in axes:\n\n    ax.axis(\n        \"off\"\n    )\n\n\nplt.tight_layout()\n\n\noverview_path = os.path.join(\n\n    CFG.viz_dir,\n\n    \"fragment1_v4_overview.png\"\n)\n\n\nplt.savefig(\n\n    overview_path,\n\n    dpi=150,\n\n    bbox_inches=\"tight\"\n)\n\n\nplt.close()\n\n\n# ============================================================\n# 57. PATCH COMPARISONS\n# ============================================================\n\ncomparison_coords = (\n    generate_grid_coords(\n\n        test_mask,\n\n        CFG.patch_size,\n\n        CFG.patch_size,\n\n        0.15\n    )\n)\n\n\nrandom.shuffle(\n    comparison_coords\n)\n\n\ncomparison_coords = (\n    comparison_coords[:6]\n)\n\n\nmid_local_idx = (\n    len(\n        CFG.depth_indices\n    ) // 2\n)\n\n\nfor i, (y, x) in enumerate(\n    comparison_coords\n):\n\n    input_patch = (\n        test_vol.read_patch(\n\n            y,\n\n            x,\n\n            CFG.patch_size\n        )[mid_local_idx]\n    )\n\n\n    pred_patch = (\n        final_prediction[\n            y:y+CFG.patch_size,\n            x:x+CFG.patch_size\n        ]\n    )\n\n\n    if test_labels is not None:\n\n        gt_patch = (\n            test_labels[\n                y:y+CFG.patch_size,\n                x:x+CFG.patch_size\n            ]\n        )\n\n\n        fig, axes = plt.subplots(\n\n            1,\n\n            3,\n\n            figsize=(12, 4)\n        )\n\n\n        axes[0].imshow(\n\n            input_patch,\n\n            cmap=\"gray\"\n        )\n\n        axes[0].set_title(\n            \"Input\"\n        )\n\n\n        axes[1].imshow(\n\n            gt_patch,\n\n            cmap=\"gray\"\n        )\n\n        axes[1].set_title(\n            \"GT\"\n        )\n\n\n        axes[2].imshow(\n\n            pred_patch,\n\n            cmap=\"gray\"\n        )\n\n        axes[2].set_title(\n            \"V4 Prediction\"\n        )\n\n\n    else:\n\n        fig, axes = plt.subplots(\n\n            1,\n\n            2,\n\n            figsize=(8, 4)\n        )\n\n\n        axes[0].imshow(\n\n            input_patch,\n\n            cmap=\"gray\"\n        )\n\n        axes[0].set_title(\n            \"Input\"\n        )\n\n\n        axes[1].imshow(\n\n            pred_patch,\n\n            cmap=\"gray\"\n        )\n\n        axes[1].set_title(\n            \"V4 Prediction\"\n        )\n\n\n    for ax in axes:\n\n        ax.axis(\n            \"off\"\n        )\n\n\n    plt.tight_layout()\n\n\n    path = os.path.join(\n\n        CFG.viz_dir,\n\n        f\"patch_{i:02d}_\"\n        f\"y{y}_x{x}.png\"\n    )\n\n\n    plt.savefig(\n\n        path,\n\n        dpi=150,\n\n        bbox_inches=\"tight\"\n    )\n\n\n    plt.close()\n\n\n# ============================================================\n# 58. CLOSE TEST VOLUME\n# ============================================================\n\ntest_vol.close()\n\ncleanup_memory()\n\n\n# ============================================================\n# 59. FINAL REPORT\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 70\n)\n\nprint(\n    \"V4 COMPLETE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    f\"Best validation Dice @ 0.50: \"\n    f\"{best_val_dice:.5f}\"\n)\n\nprint(\n    f\"Best epoch: \"\n    f\"{best_epoch}\"\n)\n\nprint(\n    f\"Validation threshold: \"\n    f\"{VAL_THRESHOLD:.3f}\"\n)\n\nprint(\n    f\"Test threshold: \"\n    f\"{FINAL_TEST_THRESHOLD:.3f}\"\n)\n\nprint(\n    f\"Test threshold method: \"\n    f\"{CFG.final_test_threshold_method}\"\n)\n\nprint(\n    f\"Patch size: \"\n    f\"{CFG.patch_size}\"\n)\n\nprint(\n    f\"Train stride: \"\n    f\"{CFG.train_stride}\"\n)\n\nprint(\n    f\"Test stride: \"\n    f\"{CFG.test_stride}\"\n)\n\nprint(\n    f\"Batch size: \"\n    f\"{CFG.batch_size}\"\n)\n\nprint(\n    f\"AdaBN: \"\n    f\"{CFG.use_adabn}\"\n)\n\nprint(\n    f\"Final TTA: \"\n    f\"{CFG.use_tta_for_final}\"\n)\n\nprint(\n    f\"Rot90 ablation: \"\n    f\"{CFG.run_rot90_ablation}\"\n)\n\nprint(\n    f\"Final variant: \"\n    f\"{final_variant}\"\n)\n\n\nif final_raw_metrics is not None:\n\n    print(\n        \"\\nLocal raw test Dice:\",\n        final_raw_metrics[\"dice\"]\n    )\n\n\nif final_post_metrics is not None:\n\n    print(\n        \"Local postprocessed test Dice:\",\n        final_post_metrics[\"dice\"]\n    )\n\n\nprint(\n    \"\\nCheckpoint:\"\n)\n\nprint(\n    CFG.ckpt_path\n)\n\nprint(\n    \"\\nFinal probability:\"\n)\n\nprint(\n    os.path.join(\n\n        CFG.out_dir,\n\n        \"fragment1_probability_v4.npy\"\n    )\n)\n\nprint(\n    \"\\nFinal prediction:\"\n)\n\nprint(\n    prediction_path\n)\n\nprint(\n    \"\\nMetrics:\"\n)\n\nprint(\n    CFG.metrics_path\n)\n\nprint(\n    \"\\nVisualizations:\"\n)\n\nprint(\n    CFG.viz_dir\n)\n\nprint(\n    \"\\n=== V4 DONE ===\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:52:31.475238Z","iopub.execute_input":"2026-08-25T07:52:31.475814Z","iopub.status.idle":"2026-08-25T09:46:24.587803Z","shell.execute_reply.started":"2026-08-25T07:52:31.475764Z","shell.execute_reply":"2026-08-25T09:46:24.583142Z"}},"outputs":[{"name":"stdout","text":"\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m\n======================================================================\nBUILDING V4 TRAINING DATA\n======================================================================\nFragment 2: (14830, 9506) 6367 patches\nFragment 3: (7606, 5249) 1696 patches\nTotal candidate patches: 8063\nSpatial train: 6416\nSpatial val: 1647\n\nSampling pools:\nPositive: 5088\nHard negative: 1003\nNormal: 325\n\nFinal sampled train patches: 6416\n\nEstimated positive fraction: 0.13601884403935185\nInitial output bias: -1.8487575230849718\nBCE positive weight: 2.520302065350686\n","output_type":"stream"},{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/resnet50-19c8e357.pth\" to /root/.cache/torch/hub/checkpoints/resnet50-19c8e357.pth\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"  0%|          | 0.00/97.8M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"aed993199fbd428db25bf28e664eb88e"}},"metadata":{}},{"name":"stdout","text":"\n======================================================================\nSTARTING V4 TRAINING\n======================================================================\n\n[01/10] time=694.0s\ntrain_loss=0.68015 train_dice=0.45194 train_P=0.41635 train_R=0.49419\nval_loss=0.68651 val_dice=0.48956 val_P=0.48277 val_R=0.49654\nencoder_lr=0.0000489 decoder_lr=0.0001952\n*** NEW BEST CHECKPOINT ***\n\n[02/10] time=553.2s\ntrain_loss=0.55783 train_dice=0.61052 train_P=0.54621 train_R=0.69199\nval_loss=0.62657 val_dice=0.58086 val_P=0.57583 val_R=0.58598\nencoder_lr=0.0000457 decoder_lr=0.0001814\n*** NEW BEST CHECKPOINT ***\n\n[03/10] time=546.2s\ntrain_loss=0.47978 train_dice=0.69086 train_P=0.62556 train_R=0.77138\nval_loss=0.60545 val_dice=0.61949 val_P=0.61183 val_R=0.62735\nencoder_lr=0.0000407 decoder_lr=0.0001598\n*** NEW BEST CHECKPOINT ***\n\n[04/10] time=554.1s\ntrain_loss=0.43284 train_dice=0.73594 train_P=0.67181 train_R=0.81361\nval_loss=0.69483 val_dice=0.60322 val_P=0.70394 val_R=0.52771\nencoder_lr=0.0000345 decoder_lr=0.0001326\nNo improvement: 1/3\n\n[05/10] time=549.1s\ntrain_loss=0.39229 train_dice=0.77384 train_P=0.71453 train_R=0.84389\nval_loss=0.58398 val_dice=0.65093 val_P=0.63614 val_R=0.66643\nencoder_lr=0.0000275 decoder_lr=0.0001025\n*** NEW BEST CHECKPOINT ***\n\n[06/10] time=548.2s\ntrain_loss=0.36158 train_dice=0.80359 train_P=0.74776 train_R=0.86842\nval_loss=0.59237 val_dice=0.67120 val_P=0.66669 val_R=0.67576\nencoder_lr=0.0000205 decoder_lr=0.0000724\n*** NEW BEST CHECKPOINT ***\n\n[07/10] time=552.1s\ntrain_loss=0.33674 train_dice=0.82309 train_P=0.76922 train_R=0.88507\nval_loss=0.59119 val_dice=0.67546 val_P=0.64481 val_R=0.70917\nencoder_lr=0.0000143 decoder_lr=0.0000452\n*** NEW BEST CHECKPOINT ***\n\n[08/10] time=550.0s\ntrain_loss=0.31498 train_dice=0.84170 train_P=0.79129 train_R=0.89898\nval_loss=0.58978 val_dice=0.69567 val_P=0.70983 val_R=0.68207\nencoder_lr=0.0000093 decoder_lr=0.0000236\n*** NEW BEST CHECKPOINT ***\n\n[09/10] time=549.1s\ntrain_loss=0.30119 train_dice=0.85361 train_P=0.80685 train_R=0.90613\nval_loss=0.57405 val_dice=0.69767 val_P=0.69018 val_R=0.70533\nencoder_lr=0.0000061 decoder_lr=0.0000098\n*** NEW BEST CHECKPOINT ***\n\n[10/10] time=553.0s\ntrain_loss=0.28921 train_dice=0.86186 train_P=0.81522 train_R=0.91416\nval_loss=0.57812 val_dice=0.70277 val_P=0.70361 val_R=0.70193\nencoder_lr=0.0000050 decoder_lr=0.0000050\n*** NEW BEST CHECKPOINT ***\n\n======================================================================\nTRAINING COMPLETE\n======================================================================\nBest validation Dice: 0.702771732636276\nBest epoch: 10\n\n======================================================================\nVALIDATION THRESHOLD\n======================================================================\nValidation threshold: 0.5400000000000003\nValidation metrics: {'threshold': 0.5400000000000003, 'dice': 0.7028628444622013, 'iou': 0.5418569975129538, 'precision': 0.7105004072721594, 'recall': 0.6953877361334373}\nValidation predicted-positive rate: 0.16123566153528301\n\n======================================================================\nLOADING TEST FRAGMENT\n======================================================================\nTest shape: (8181, 6330)\nLocal test GT available: True\n\n======================================================================\nBASELINE TEST INFERENCE\n======================================================================\nInference patches: 2065\n\n======================================================================\nSTARTING AdaBN\n======================================================================\nBN layers: 63\nAdaBN patches: 1200\nAdaBN finished.\n\n======================================================================\nADABN TEST INFERENCE\n======================================================================\nInference patches: 2065\n\nADABN probability statistics:\n{'min': 6.164189471746795e-06, 'max': 1.0, 'mean': 0.14843401312828064, 'median': 0.003195234341546893, 'p90': 0.7001193881034853, 'p95': 0.9322058677673333, 'p97': 0.9781386750936507, 'p98': 0.9910587668418884, 'p99': 0.9980466961860657, 'p99.5': 0.9996387362480164}\n\n======================================================================\nINDEPENDENT TEST THRESHOLD CALIBRATION\n======================================================================\nValidation threshold: 0.5400000000000003\nTest Otsu: 0.4160192012786865\nTest positive-rate threshold: 0.3254127547036858\nTest P98: 0.75\n\n======================================================================\nLOCAL TEST THRESHOLD DIAGNOSTICS\n======================================================================\n\notsu: threshold=0.4160\n{'dice': 0.4273929458938376, 'iou': 0.27177351441860054, 'precision': 0.4861092508620647, 'recall': 0.38133245133044336}\n\npositive_rate: threshold=0.3254\n{'dice': 0.4356302725939842, 'iou': 0.2784701499666598, 'precision': 0.4653203560179274, 'recall': 0.4095017344770012}\n\np98: threshold=0.7500\n{'dice': 0.3691877738872977, 'iou': 0.22638276067344165, 'precision': 0.5559142960909588, 'recall': 0.27636073373573916}\n\nvalidation_threshold_diagnostic: threshold=0.5400\n{'dice': 0.41278524493606383, 'iou': 0.2600689312011333, 'precision': 0.510163031912976, 'recall': 0.34662324824588653}\n\nFINAL TEST THRESHOLD: 0.4160192012786865\n\n======================================================================\nROT90 ABLATION\n======================================================================\n\n======================================================================\nSTARTING AdaBN\n======================================================================\nBN layers: 63\nAdaBN patches: 1200\nAdaBN finished.\nInference patches: 2065\nRot90 probability map generated.\n\nRot90 diagnostic:\nRot90 + Otsu: {'dice': 0.1140663631406092, 'iou': 0.06048270252537547, 'precision': 0.5244093352208344, 'recall': 0.06399285158057011}\nRot90 + positive-rate: {'dice': 0.2683888076559192, 'iou': 0.1549936896011351, 'precision': 0.4466728710624926, 'recall': 0.19182441647537818}\n\n======================================================================\nFINAL V4 LOCAL DIAGNOSTICS\n======================================================================\nFinal variant: AdaBN + no TTA\nFinal threshold: 0.4160192012786865\nFinal RAW: {'dice': 0.4273929458938376, 'iou': 0.27177351441860054, 'precision': 0.4861092508620647, 'recall': 0.38133245133044336}\nFinal POSTPROCESSED: {'dice': 0.4274223735746645, 'iou': 0.27179731314525, 'precision': 0.4860563721484292, 'recall': 0.38141186156709705}\n\n======================================================================\nV4 COMPLETE\n======================================================================\nBest validation Dice @ 0.50: 0.70277\nBest epoch: 10\nValidation threshold: 0.540\nTest threshold: 0.416\nTest threshold method: otsu\nPatch size: 480\nTrain stride: 128\nTest stride: 128\nBatch size: 8\nAdaBN: True\nFinal TTA: False\nRot90 ablation: True\nFinal variant: AdaBN + no TTA\n\nLocal raw test Dice: 0.4273929458938376\nLocal postprocessed test Dice: 0.4274223735746645\n\nCheckpoint:\n/kaggle/working/vesuvius_v4/vesuvius_v4_best.pth\n\nFinal probability:\n/kaggle/working/vesuvius_v4/fragment1_probability_v4.npy\n\nFinal prediction:\n/kaggle/working/vesuvius_v4/fragment1_prediction_v4.png\n\nMetrics:\n/kaggle/working/vesuvius_v4/v4_metrics_summary.json\n\nVisualizations:\n/kaggle/working/vesuvius_v4/visualizations\n\n=== V4 DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# Test Dice Score: 0.5015 F1 Score: 0.5416 ON Fragment-1\n# 2D UNet PATCH_SIZE = 352 STRIDE = 26 Dice=0.35\n# 2D UNet PATCH_SIZE = 160 STRIDE = 32 Dice=0.41\n#PATCH_SIZE = 192 ,STRIDE = 32, Dice Score: 0.43 F1 Score: 0.48 ON Fragment-1\n#PATCH_SIZE = 128 ,STRIDE = 16, Dice Score: 0.46 F1 Score: 0.50 ON Fragment-1\n\n!pip install segmentation-models-pytorch==0.2.0\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport torch.nn.functional as F\nimport gc\nimport psutil\nimport time\n#/kaggle/input/vesuvius-challenge-ink-detection/test\n\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n]\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\nelse:\n    print(\"CPU mode\")\n\n# Memory-efficient settings\nPATCH_SIZE = 128\nSTRIDE = 32\nBATCH_SIZE = 8  # Increased slightly since we're using 2D\nACCUMULATION_STEPS = 2  # Gradient accumulation\nEPOCHS = 10\nLR = 1e-4\nWEIGHT_DECAY = 1e-4\nSLICE_START = 15\nSLICE_END = 30  # Python range is exclusive -> slices 12 through 36 inclusive\nNUM_INPUT_SLICES = SLICE_END - SLICE_START  # 25 slices\nOUTPUT_DIR = \"/kaggle/working/\"\nVALIDATION_SPLIT = 0.20\n\n# Advanced training settings\nUSE_AMP = True\nPIN_MEMORY = True\nNUM_WORKERS = 2\nGRAD_CLIP = 1.0\nSCHEDULER_PATIENCE = 15\nEARLY_STOPPING_PATIENCE = 15\n\n# ============================================\n# IGNORE MASK GENERATION\n# ============================================\ndef generate_ignore_mask(ink_mask, distance_threshold=3, erosion_size=2):\n    \"\"\"Generate ignore mask for uncertain regions - OPTIMIZED VERSION\"\"\"\n    ignore_mask = np.zeros_like(ink_mask, dtype=np.uint8)\n    \n    # Simple boundary detection (faster than full distance transform)\n    kernel = np.ones((3, 3), np.uint8)\n    eroded = cv2.erode(ink_mask.astype(np.uint8), kernel, iterations=1)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    \n    # Boundaries are where eroded and dilated differ from original\n    boundaries = (dilated != eroded)\n    ignore_mask[boundaries] = 1\n    \n    # Remove small isolated ink dots\n    contours, _ = cv2.findContours(ink_mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for contour in contours:\n        if cv2.contourArea(contour) < 10:\n            cv2.drawContours(ignore_mask, [contour], -1, 1, -1)\n    \n    return ignore_mask\n\n# ============================================\n# 2D CNN WITH MULTI-SLICE INPUT\n# ============================================\nclass MultiSlice2DUNet(nn.Module):\n    \"\"\"2D CNN that treats depth slices as input channels\"\"\"\n    def __init__(self, in_channels=NUM_INPUT_SLICES, out_channels=1):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._conv_block(in_channels, 32)\n        self.enc2 = self._conv_block(32, 64)\n        self.enc3 = self._conv_block(64, 128)\n        self.enc4 = self._conv_block(128, 256)\n        \n        # Bottleneck\n        self.bottleneck = self._conv_block(256, 512)\n        \n        # Decoder\n        self.dec4 = self._upconv_block(512 + 256, 256)\n        self.dec3 = self._upconv_block(256 + 128, 128)\n        self.dec2 = self._upconv_block(128 + 64, 64)\n        self.dec1 = self._upconv_block(64 + 32, 32)\n        \n        # Output\n        self.final = nn.Conv2d(32, out_channels, kernel_size=1)\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False)\n        \n    def _conv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def _upconv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # x shape: [B, C=25, H, W] for slices 12..36\n        \n        # Encoder\n        e1 = self.enc1(x)      # [B, 32, H, W]\n        p1 = self.pool(e1)      # [B, 32, H/2, W/2]\n        \n        e2 = self.enc2(p1)      # [B, 64, H/2, W/2]\n        p2 = self.pool(e2)      # [B, 64, H/4, W/4]\n        \n        e3 = self.enc3(p2)      # [B, 128, H/4, W/4]\n        p3 = self.pool(e3)      # [B, 128, H/8, W/8]\n        \n        e4 = self.enc4(p3)      # [B, 256, H/8, W/8]\n        p4 = self.pool(e4)      # [B, 256, H/16, W/16]\n        \n        # Bottleneck\n        b = self.bottleneck(p4)  # [B, 512, H/16, W/16]\n        \n        # Decoder with skip connections\n        d4 = self.upsample(b)    # [B, 512, H/8, W/8]\n        d4 = torch.cat([d4, e4], dim=1)  # [B, 512+256, H/8, W/8]\n        d4 = self.dec4(d4)        # [B, 256, H/8, W/8]\n        \n        d3 = self.upsample(d4)    # [B, 256, H/4, W/4]\n        d3 = torch.cat([d3, e3], dim=1)  # [B, 256+128, H/4, W/4]\n        d3 = self.dec3(d3)        # [B, 128, H/4, W/4]\n        \n        d2 = self.upsample(d3)    # [B, 128, H/2, W/2]\n        d2 = torch.cat([d2, e2], dim=1)  # [B, 128+64, H/2, W/2]\n        d2 = self.dec2(d2)        # [B, 64, H/2, W/2]\n        \n        d1 = self.upsample(d2)    # [B, 64, H, W]\n        d1 = torch.cat([d1, e1], dim=1)  # [B, 64+32, H, W]\n        d1 = self.dec1(d1)        # [B, 32, H, W]\n        \n        # Final output\n        out = self.final(d1)      # [B, 1, H, W]\n        \n        return out\n\n# ============================================\n# DATA LOADING (OPTIMIZED)\n# ============================================\ndef load_volume_fast(fragment_path):\n    \"\"\"Fast volume loading\"\"\"\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    \n    volume = np.stack(slices, axis=-1)\n    volume = volume.astype(np.float32)\n    \n    # Fast normalization per slice\n    for i in range(volume.shape[-1]):\n        slice_data = volume[:, :, i]\n        mean, std = slice_data.mean(), slice_data.std()\n        volume[:, :, i] = (slice_data - mean) / (std + 1e-6)\n    \n    return volume\n\ndef extract_patches_fast(volume, mask, ignore_mask, max_patches_per_fragment=3000):\n    \"\"\"Fast patch extraction with sampling\"\"\"\n    patches = []\n    mask_patches = []\n    ignore_patches = []\n    \n    H, W, _ = volume.shape\n    \n    # Calculate number of patches\n    n_y = (H - PATCH_SIZE) // STRIDE + 1\n    n_x = (W - PATCH_SIZE) // STRIDE + 1\n    total_patches = n_y * n_x\n    \n    print(f\"    Total possible patches: {total_patches}\")\n    \n    # Sample patches if too many\n    if total_patches > max_patches_per_fragment:\n        print(f\"    Sampling {max_patches_per_fragment} patches...\")\n        # Calculate stride to get roughly max_patches\n        stride_y = max(STRIDE, (H - PATCH_SIZE) // int(np.sqrt(max_patches_per_fragment)))\n        stride_x = max(STRIDE, (W - PATCH_SIZE) // int(np.sqrt(max_patches_per_fragment)))\n    else:\n        stride_y, stride_x = STRIDE, STRIDE\n    \n    patch_count = 0\n    ink_patches = 0\n    bg_patches = 0\n    \n    for y in range(0, H - PATCH_SIZE, stride_y):\n        for x in range(0, W - PATCH_SIZE, stride_x):\n            if patch_count >= max_patches_per_fragment:\n                break\n                \n            v_patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            m_patch = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            i_patch = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            \n            # Count ink pixels in this patch\n            ink_pixel_count = m_patch.sum()\n            \n            # Keep patches with significant ink or some background for balance\n            if ink_pixel_count > 50:  # Good ink patch\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n                ink_patches += 1\n            elif ink_pixel_count == 0 and bg_patches < ink_patches // 2:  # Balance with background\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n                bg_patches += 1\n        \n        if patch_count >= max_patches_per_fragment:\n            break\n    \n    print(f\"    Extracted: {ink_patches} ink patches, {bg_patches} background patches\")\n    return patches, mask_patches, ignore_patches\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignore_masks, transform=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.ignore_masks = ignore_masks\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        image = self.volumes[idx].copy()\n        mask = self.masks[idx].copy()\n        ignore = self.ignore_masks[idx].copy()\n        \n        if self.transform:\n            transformed = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image = transformed['image']\n            mask = transformed['mask']\n            ignore = transformed['ignore_mask']\n        \n        # Convert to tensors\n        # For 2D CNN: image shape [H, W, C=25] -> [C=25, H, W]\n        image = torch.tensor(image).permute(2, 0, 1).float()  # [C=12, H, W]\n        \n        # mask and ignore shape: [H, W]\n        mask = torch.tensor(mask).float().unsqueeze(0)  # [1, H, W]\n        ignore = torch.tensor(ignore).float().unsqueeze(0)  # [1, H, W]\n        \n        return image, mask, ignore\n\n# ============================================\n# LOSS FUNCTION\n# ============================================\nclass DiceBCELossWithIgnore(nn.Module):\n    def __init__(self, dice_weight=0.5, bce_weight=0.5, smooth=1e-6):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.smooth = smooth\n        \n    def forward(self, pred, target, ignore_mask):\n        \"\"\"\n        pred: [B, 1, H, W] - logits\n        target: [B, 1, H, W] - binary mask\n        ignore_mask: [B, 1, H, W] - 1 for ignore, 0 for keep\n        \"\"\"\n        # Create valid mask\n        valid_mask = (1 - ignore_mask).float()\n        \n        # BCE loss\n        bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n        bce = (bce * valid_mask).sum() / (valid_mask.sum() + self.smooth)\n        \n        # Dice loss\n        pred_probs = torch.sigmoid(pred)\n        \n        # Apply valid mask\n        pred_valid = pred_probs * valid_mask\n        target_valid = target * valid_mask\n        \n        intersection = (pred_valid * target_valid).sum()\n        union = pred_valid.sum() + target_valid.sum()\n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        dice_loss = 1 - dice\n        \n        return self.dice_weight * dice_loss + self.bce_weight * bce\n\n# ============================================\n# FAST TRAIN/VALIDATION SPLIT\n# ============================================\ndef fast_train_val_split(n_samples, val_ratio=0.15, seed=42):\n    \"\"\"Fast random split without distance constraints\"\"\"\n    np.random.seed(seed)\n    indices = np.random.permutation(n_samples)\n    split = int(n_samples * val_ratio)\n    return indices[split:], indices[:split]\n\n# ============================================\n# MEMORY MONITORING\n# ============================================\ndef print_memory_usage():\n    if DEVICE == 'cuda':\n        allocated = torch.cuda.memory_allocated() / 1e9\n        cached = torch.cuda.memory_reserved() / 1e9\n        print(f\"    GPU Memory - Allocated: {allocated:.2f}GB, Cached: {cached:.2f}GB\")\n    \n    process = psutil.Process()\n    print(f\"    CPU Memory: {process.memory_info().rss / 1e9:.2f}GB\")\n\n# ============================================\n# COLLATE FUNCTION\n# ============================================\ndef collate_fn(batch):\n    \"\"\"Custom collate function to ensure correct dimensions\"\"\"\n    images = torch.stack([item[0] for item in batch])  # [B, C=25, H, W]\n    masks = torch.stack([item[1] for item in batch])   # [B, 1, H, W]\n    ignores = torch.stack([item[2] for item in batch]) # [B, 1, H, W]\n    return images, masks, ignores\n\n# ============================================\n# FULL VOLUME PREDICTION FUNCTION\n# ============================================\ndef _get_sliding_positions(length, patch_size, stride):\n    \"\"\"Return patch start positions and always include the final border position.\"\"\"\n    if length <= patch_size:\n        return [0]\n\n    positions = list(range(0, length - patch_size + 1, stride))\n    last_position = length - patch_size\n    if positions[-1] != last_position:\n        positions.append(last_position)\n    return positions\n\n\ndef predict_full_volume(model, volume, device, batch_size=8):\n    \"\"\"\n    Predict a complete 2D surface from a multi-slice input volume.\n\n    The volume shape is [H, W, NUM_INPUT_SLICES].\n    Overlapping 128x128 predictions are averaged.\n    The final row/column positions are explicitly included so the full\n    Fragment-1 surface is covered, including the borders.\n    \"\"\"\n    model.eval()\n    H, W, C = volume.shape\n\n    if C != NUM_INPUT_SLICES:\n        raise ValueError(\n            f\"Expected {NUM_INPUT_SLICES} input slices, but volume has {C}. \"\n            f\"Check SLICE_START/SLICE_END.\"\n        )\n\n    surface_prediction = np.zeros((H, W), dtype=np.float32)\n    prediction_count = np.zeros((H, W), dtype=np.float32)\n\n    stride = PATCH_SIZE // 2\n    y_positions = _get_sliding_positions(H, PATCH_SIZE, stride)\n    x_positions = _get_sliding_positions(W, PATCH_SIZE, stride)\n\n    patches = []\n    positions = []\n\n    for y in y_positions:\n        for x in x_positions:\n            patch = volume[y:y + PATCH_SIZE, x:x + PATCH_SIZE, :]\n\n            # If the image is smaller than PATCH_SIZE, pad the patch.\n            ph, pw, _ = patch.shape\n            if ph != PATCH_SIZE or pw != PATCH_SIZE:\n                padded = np.zeros((PATCH_SIZE, PATCH_SIZE, C), dtype=np.float32)\n                padded[:ph, :pw, :] = patch\n                patch = padded\n\n            patch_tensor = (\n                torch.from_numpy(patch)\n                .permute(2, 0, 1)\n                .float()\n                .unsqueeze(0)\n            )\n            patches.append(patch_tensor)\n            positions.append((y, x, ph, pw))\n\n    print(f\"    Full-volume patches: {len(patches)}\")\n    print(f\"    Input slices: {C} ({SLICE_START}-{SLICE_END - 1})\")\n    print(f\"    Patch size: {PATCH_SIZE}, inference stride: {stride}\")\n\n    with torch.no_grad():\n        for i in range(0, len(patches), batch_size):\n            batch_patches = torch.cat(\n                patches[i:i + batch_size], dim=0\n            ).to(device, non_blocking=True)\n\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(batch_patches)\n                preds = torch.sigmoid(outputs).cpu().numpy()[:, 0]\n\n            for j, pred in enumerate(preds):\n                y, x, ph, pw = positions[i + j]\n                surface_prediction[y:y + ph, x:x + pw] += pred[:ph, :pw]\n                prediction_count[y:y + ph, x:x + pw] += 1.0\n\n            del batch_patches\n            if DEVICE == 'cuda' and i % (batch_size * 20) == 0:\n                torch.cuda.empty_cache()\n\n    prediction_count[prediction_count == 0] = 1.0\n    surface_prediction /= prediction_count\n\n    binary_prediction = (surface_prediction > 0.5).astype(np.uint8)\n\n    return binary_prediction, surface_prediction\n\n# ============================================\n# VISUALIZATION FUNCTION\n# ============================================\ndef _normalize_for_display(image):\n    \"\"\"Normalize one input slice to [0, 1] for visualization only.\"\"\"\n    image = image.astype(np.float32)\n    lo = np.percentile(image, 1)\n    hi = np.percentile(image, 99)\n    if hi <= lo:\n        lo = image.min()\n        hi = image.max()\n    return np.clip((image - lo) / (hi - lo + 1e-8), 0, 1)\n\n\ndef create_full_fragment_comparison(\n    input_volume,\n    ground_truth,\n    prediction,\n    save_path,\n    dice=None,\n    f1=None,\n    precision=None,\n    recall=None\n):\n    \"\"\"\n    Save one large figure containing ALL input slices plus GT/prediction.\n\n    For slices 12..36 inclusive this creates:\n      - 25 input-slice panels\n      - Full ground truth\n      - Full prediction\n      - GT/prediction overlay\n      - TP/FP/FN difference map\n      - Metrics panel\n\n    Total: 30 panels in a 5 x 6 grid.\n    \"\"\"\n    H, W, C = input_volume.shape\n\n    if C != NUM_INPUT_SLICES:\n        raise ValueError(\n            f\"Expected {NUM_INPUT_SLICES} input slices, got {C}.\"\n        )\n\n    gt = ground_truth.astype(bool)\n    pred = prediction.astype(bool)\n\n    # TP / FP / FN\n    tp = gt & pred\n    fp = (~gt) & pred\n    fn = gt & (~pred)\n\n    # Difference map: TP white, FP red, FN blue\n    diff = np.zeros((H, W, 3), dtype=np.uint8)\n    diff[tp] = [255, 255, 255]\n    diff[fp] = [255, 0, 0]\n    diff[fn] = [0, 0, 255]\n\n    # Overlay: GT-only green, prediction-only red, overlap yellow\n    overlay = np.zeros((H, W, 3), dtype=np.uint8)\n    overlay[gt & (~pred)] = [0, 255, 0]\n    overlay[(~gt) & pred] = [255, 0, 0]\n    overlay[gt & pred] = [255, 255, 0]\n\n    # 25 input slices + 5 comparison panels = 30 panels\n    fig, axes = plt.subplots(5, 6, figsize=(24, 21))\n    axes = axes.ravel()\n\n    # --------------------------------------------------------\n    # Input slices 12..36\n    # --------------------------------------------------------\n    for i in range(C):\n        slice_number = SLICE_START + i\n        axes[i].imshow(\n            _normalize_for_display(input_volume[:, :, i]),\n            cmap='gray'\n        )\n        axes[i].set_title(\n            f'Input Slice {slice_number}',\n            fontsize=11,\n            fontweight='bold'\n        )\n        axes[i].axis('off')\n\n    # --------------------------------------------------------\n    # Panel 26: Ground truth\n    # --------------------------------------------------------\n    axes[25].imshow(ground_truth, cmap='gray')\n    axes[25].set_title(\n        'FULL GROUND TRUTH',\n        fontsize=13,\n        fontweight='bold'\n    )\n    axes[25].axis('off')\n\n    # --------------------------------------------------------\n    # Panel 27: Full prediction\n    # --------------------------------------------------------\n    axes[26].imshow(prediction, cmap='gray')\n    axes[26].set_title(\n        'FULL PREDICTION',\n        fontsize=13,\n        fontweight='bold'\n    )\n    axes[26].axis('off')\n\n    # --------------------------------------------------------\n    # Panel 28: Overlay\n    # --------------------------------------------------------\n    axes[27].imshow(_normalize_for_display(input_volume[:, :, C // 2]), cmap='gray')\n    axes[27].imshow(overlay, alpha=0.55)\n    axes[27].set_title(\n        'GT / PREDICTION OVERLAY\\n'\n        'Green=GT | Red=Pred | Yellow=Overlap',\n        fontsize=11,\n        fontweight='bold'\n    )\n    axes[27].axis('off')\n\n    # --------------------------------------------------------\n    # Panel 29: Difference map\n    # --------------------------------------------------------\n    axes[28].imshow(diff)\n    axes[28].set_title(\n        'DIFFERENCE MAP\\n'\n        'White=TP | Red=FP | Blue=FN',\n        fontsize=11,\n        fontweight='bold'\n    )\n    axes[28].axis('off')\n\n    # --------------------------------------------------------\n    # Panel 30: Metrics\n    # --------------------------------------------------------\n    axes[29].axis('off')\n\n    metrics_text = (\n        'FRAGMENT 1 FULL-SURFACE TEST\\n\\n'\n        f'Input slices: {SLICE_START}-{SLICE_END - 1} ({C} slices)\\n'\n        f'Volume shape: {input_volume.shape}\\n\\n'\n        f'Dice Score : {dice:.4f}\\n' if dice is not None else ''\n    )\n\n    if f1 is not None:\n        metrics_text += f'F1 Score   : {f1:.4f}\\n'\n    if precision is not None:\n        metrics_text += f'Precision  : {precision:.4f}\\n'\n    if recall is not None:\n        metrics_text += f'Recall     : {recall:.4f}\\n'\n\n    metrics_text += (\n        '\\n'\n        f'GT ink pixels   : {gt.sum():,}\\n'\n        f'Predicted pixels: {pred.sum():,}\\n'\n        f'True positives  : {tp.sum():,}\\n'\n        f'False positives : {fp.sum():,}\\n'\n        f'False negatives : {fn.sum():,}'\n    )\n\n    axes[29].text(\n        0.04,\n        0.5,\n        metrics_text,\n        fontsize=12,\n        verticalalignment='center',\n        family='monospace',\n        bbox=dict(\n            boxstyle='round,pad=0.7',\n            facecolor='white',\n            alpha=0.9\n        )\n    )\n\n    plt.suptitle(\n        'VESUVIUS CHALLENGE - FRAGMENT 1\\n'\n        f'ALL INPUT SLICES {SLICE_START}-{SLICE_END - 1} + FULL GT + FULL PREDICTION',\n        fontsize=19,\n        fontweight='bold'\n    )\n\n    plt.tight_layout(rect=[0, 0, 1, 0.965])\n    plt.savefig(save_path, dpi=200, bbox_inches='tight')\n    plt.close(fig)\n\n    print(f'✓ Full Fragment-1 comparison saved to: {save_path}')\n\n\ndef save_input_stack_visualization(input_volume, save_path):\n    \"\"\"Save all input slices 12..36 as a separate clean 5x5 grid.\"\"\"\n    C = input_volume.shape[-1]\n    rows, cols = 5, 5\n\n    fig, axes = plt.subplots(rows, cols, figsize=(20, 20))\n    axes = axes.ravel()\n\n    for i in range(C):\n        axes[i].imshow(\n            _normalize_for_display(input_volume[:, :, i]),\n            cmap='gray'\n        )\n        axes[i].set_title(\n            f'Slice {SLICE_START + i}',\n            fontsize=12,\n            fontweight='bold'\n        )\n        axes[i].axis('off')\n\n    for i in range(C, rows * cols):\n        axes[i].axis('off')\n\n    plt.suptitle(\n        f'Fragment 1 - Complete Input Stack (Slices {SLICE_START}-{SLICE_END - 1})',\n        fontsize=18,\n        fontweight='bold'\n    )\n    plt.tight_layout(rect=[0, 0, 1, 0.96])\n    plt.savefig(save_path, dpi=200, bbox_inches='tight')\n    plt.close(fig)\n\n    print(f'✓ Input stack visualization saved to: {save_path}')\n\n\n# ============================================\n# MAIN PIPELINE\n# ============================================\ndef main():\n    print(\"=\"*60)\n    print(\"VESUVIUS INK DETECTION - OPTIMIZED PIPELINE\")\n    print(\"=\"*60)\n    print_memory_usage()\n    \n    start_time = time.time()\n    \n    # ============================================\n    # LOAD AND PREPARE DATA\n    # ============================================\n    print(\"\\n1. Loading training data...\")\n    all_patches = []\n    all_masks = []\n    all_ignores = []\n    \n    for path in train_paths:\n        fragment_name = os.path.basename(path)\n        print(f\"\\n   Processing Fragment {fragment_name}...\")\n        \n        # Load volume\n        volume = load_volume_fast(path)\n        print(f\"    Volume shape: {volume.shape}\")\n        \n        # Load mask\n        mask_path = os.path.join(path, \"inklabels.png\")\n        mask = cv2.imread(mask_path, 0)\n        mask = (mask > 0).astype(np.uint8)\n        print(f\"    Mask shape: {mask.shape}\")\n        print(f\"    Ink pixels: {mask.sum():,}\")\n        \n        # Generate ignore mask\n        print(\"    Generating ignore mask...\")\n        ignore_mask = generate_ignore_mask(mask)\n        print(f\"    Ignored pixels: {ignore_mask.sum():,}\")\n        \n        # Extract patches\n        print(\"    Extracting patches...\")\n        patches, mask_patches, ignore_patches = extract_patches_fast(\n            volume, mask, ignore_mask, max_patches_per_fragment=3000\n        )\n        \n        all_patches.extend(patches)\n        all_masks.extend(mask_patches)\n        all_ignores.extend(ignore_patches)\n        \n        print(f\"    Total extracted: {len(patches)} patches\")\n        print_memory_usage()\n        \n        # Clean up\n        del volume, mask, ignore_mask, patches, mask_patches, ignore_patches\n        gc.collect()\n        if DEVICE == 'cuda':\n            torch.cuda.empty_cache()\n    \n    print(f\"\\nTotal patches: {len(all_patches)}\")\n    print_memory_usage()\n    \n    # ============================================\n    # CREATE TRAIN/VAL SPLIT\n    # ============================================\n    print(\"\\n2. Creating train/validation split...\")\n    n_samples = len(all_patches)\n    train_indices, val_indices = fast_train_val_split(n_samples, VALIDATION_SPLIT)\n    \n    print(f\"   Train samples: {len(train_indices)}\")\n    print(f\"   Validation samples: {len(val_indices)}\")\n    \n    # Split data\n    train_patches = [all_patches[i] for i in train_indices]\n    train_masks = [all_masks[i] for i in train_indices]\n    train_ignores = [all_ignores[i] for i in train_indices]\n    \n    val_patches = [all_patches[i] for i in val_indices]\n    val_masks = [all_masks[i] for i in val_indices]\n    val_ignores = [all_ignores[i] for i in val_indices]\n    \n    # Clean up original lists\n    del all_patches, all_masks, all_ignores\n    gc.collect()\n    \n    # ============================================\n    # DATA AUGMENTATION\n    # ============================================\n    print(\"\\n3. Setting up data augmentation...\")\n    \n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),\n        A.GaussNoise(var_limit=(0, 0.01), p=0.2),\n    ], additional_targets={'ignore_mask': 'mask'})\n    \n    # Create datasets\n    train_dataset = VesuviusDataset(\n        train_patches, train_masks, train_ignores, \n        transform=train_transform\n    )\n    \n    val_dataset = VesuviusDataset(\n        val_patches, val_masks, val_ignores,\n        transform=None\n    )\n    \n    # Create data loaders with custom collate function\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=True, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        drop_last=True,\n        collate_fn=collate_fn\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    print(f\"   Train batches: {len(train_loader)}\")\n    print(f\"   Val batches: {len(val_loader)}\")\n    \n    # ============================================\n    # MODEL INITIALIZATION\n    # ============================================\n    print(\"\\n4. Initializing model...\")\n    \n    model = MultiSlice2DUNet(in_channels=NUM_INPUT_SLICES, out_channels=1).to(DEVICE)\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"   Total parameters: {total_params:,}\")\n    print(f\"   Trainable parameters: {trainable_params:,}\")\n    \n    # ============================================\n    # LOSS, OPTIMIZER, SCHEDULER\n    # ============================================\n    criterion = DiceBCELossWithIgnore(dice_weight=0.5, bce_weight=0.5)\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='max', factor=0.5, patience=SCHEDULER_PATIENCE, verbose=True\n    )\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n    \n    # ============================================\n    # TRAINING LOOP\n    # ============================================\n    print(\"\\n5. Starting training...\")\n    print(\"=\"*60)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    train_losses = []\n    val_dice_scores = []\n    \n    for epoch in range(EPOCHS):\n        epoch_start = time.time()\n        \n        # Training phase\n        model.train()\n        train_loss = 0\n        train_steps = 0\n        optimizer.zero_grad()\n        \n        progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Train]')\n        for batch_idx, (images, masks, ignores) in enumerate(progress_bar):\n            # images shape: [B, C=25, H, W] for slices 12..36\n            # masks shape: [B, 1, H, W]\n            # ignores shape: [B, 1, H, W]\n            \n            images = images.to(DEVICE, non_blocking=True)\n            masks = masks.to(DEVICE, non_blocking=True)\n            ignores = ignores.to(DEVICE, non_blocking=True)\n            \n            # Forward pass\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)  # [B, 1, H, W]\n                loss = criterion(outputs, masks, ignores)\n                loss = loss / ACCUMULATION_STEPS\n            \n            # Backward pass\n            scaler.scale(loss).backward()\n            \n            # Gradient accumulation\n            if (batch_idx + 1) % ACCUMULATION_STEPS == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            train_loss += loss.item() * ACCUMULATION_STEPS\n            train_steps += 1\n            \n            # Update progress bar\n            progress_bar.set_postfix({'loss': f'{loss.item() * ACCUMULATION_STEPS:.4f}'})\n            \n            # Clear cache periodically\n            if batch_idx % 50 == 49:\n                if DEVICE == 'cuda':\n                    torch.cuda.empty_cache()\n        \n        avg_train_loss = train_loss / train_steps\n        train_losses.append(avg_train_loss)\n        \n        # Validation phase\n        model.eval()\n        val_dice = 0\n        val_steps = 0\n        \n        with torch.no_grad():\n            for images, masks, ignores in tqdm(val_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Val]'):\n                images = images.to(DEVICE, non_blocking=True)\n                masks = masks.to(DEVICE, non_blocking=True)\n                ignores = ignores.to(DEVICE, non_blocking=True)\n                \n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    outputs = model(images)\n                \n                # Calculate dice\n                preds = torch.sigmoid(outputs) > 0.5\n                valid_mask = (1 - ignores)\n                \n                intersection = ((preds * masks) * valid_mask).sum()\n                union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n                dice = (2 * intersection) / (union + 1e-6)\n                val_dice += dice.item()\n                val_steps += 1\n        \n        avg_val_dice = val_dice / val_steps\n        val_dice_scores.append(avg_val_dice)\n        \n        # Update scheduler\n        scheduler.step(avg_val_dice)\n        \n        epoch_time = time.time() - epoch_start\n        \n        # Print results\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS} - Time: {epoch_time:.1f}s\")\n        print(f\"  Train Loss: {avg_train_loss:.4f}\")\n        print(f\"  Val Dice: {avg_val_dice:.4f}\")\n        print(f\"  LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        print_memory_usage()\n        \n        # Save best model\n        if avg_val_dice > best_val_dice:\n            best_val_dice = avg_val_dice\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_dice': avg_val_dice,\n            }, os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_counter = 0\n            print(f\"  ✓ New best model saved! Dice: {best_val_dice:.4f}\")\n        else:\n            patience_counter += 1\n        \n        # Early stopping\n        if patience_counter >= EARLY_STOPPING_PATIENCE:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n        \n        print(\"-\"*60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING COMPLETE\")\n    print(f\"Best Validation Dice: {best_val_dice:.4f}\")\n    print(\"=\"*60)\n    \n    # ============================================\n    # FINAL TEST ON FRAGMENT 1\n    # ============================================\n    print(\"\\n6. Testing on Fragment 1...\")\n    print(\"=\"*60)\n    \n    # Load best model\n    checkpoint = torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth'))\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    # Load test data\n    print(\"Loading test data...\")\n    test_volume = load_volume_fast(fragment1_path)\n    test_mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ignore = generate_ignore_mask(test_mask)\n    \n    # Extract test patches\n    print(\"Extracting test patches...\")\n    test_patches, test_masks, test_ignores = extract_patches_fast(\n        test_volume, test_mask, test_ignore, max_patches_per_fragment=3000\n    )\n    print(f\"Test patches: {len(test_patches)}\")\n    \n    # Create test dataset\n    test_dataset = VesuviusDataset(\n        test_patches, test_masks, test_ignores, transform=None\n    )\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    # Evaluate\n    print(\"Running inference...\")\n    test_dice = 0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, ignores in tqdm(test_loader, desc=\"Testing\"):\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            ignores = ignores.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            valid_mask = (1 - ignores)\n            \n            # Calculate dice\n            intersection = ((preds * masks) * valid_mask).sum()\n            union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n            dice = (2 * intersection) / (union + 1e-6)\n            test_dice += dice.item()\n            \n            # Store for metrics\n            all_preds.append((preds * valid_mask).cpu().numpy())\n            all_targets.append((masks * valid_mask).cpu().numpy())\n    \n    avg_test_dice = test_dice / len(test_loader)\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_targets = np.concatenate([t.flatten() for t in all_targets])\n    \n    tn, fp, fn, tp = confusion_matrix(all_targets, all_preds, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    \n    # ============================================\n    # FULL FRAGMENT 1 PREDICTION + VISUALIZATION\n    # ============================================\n    print(\"\\n7. Generating FULL Fragment-1 surface prediction...\")\n    print(\"=\" * 60)\n\n    # Predict the complete Fragment-1 surface using ALL 25 input slices\n    binary_prediction, prob_prediction = predict_full_volume(\n        model,\n        test_volume,\n        DEVICE,\n        batch_size=8\n    )\n\n    # --------------------------------------------------------\n    # Save full binary prediction\n    # --------------------------------------------------------\n    binary_save_path = os.path.join(\n        OUTPUT_DIR,\n        'fragment1_full_surface_binary_prediction.png'\n    )\n    cv2.imwrite(\n        binary_save_path,\n        binary_prediction * 255\n    )\n    print(f\"Full binary prediction saved to: {binary_save_path}\")\n\n    # --------------------------------------------------------\n    # Save full probability prediction\n    # --------------------------------------------------------\n    prob_save_path = os.path.join(\n        OUTPUT_DIR,\n        'fragment1_full_surface_probability.tif'\n    )\n    tifffile.imwrite(\n        prob_save_path,\n        (prob_prediction * 255).astype(np.uint8)\n    )\n    print(f\"Full probability prediction saved to: {prob_save_path}\")\n\n    # --------------------------------------------------------\n    # Save all 25 input slices separately as a 5x5 grid\n    # --------------------------------------------------------\n    input_stack_path = os.path.join(\n        OUTPUT_DIR,\n        'fragment1_input_slices_12_to_36.png'\n    )\n    save_input_stack_visualization(\n        test_volume,\n        input_stack_path\n    )\n\n    # --------------------------------------------------------\n    # Save ONE complete comparison figure:\n    # 25 input slices + GT + prediction + overlay + diff + metrics\n    # --------------------------------------------------------\n    comparison_path = os.path.join(\n        OUTPUT_DIR,\n        'FRAGMENT1_ALL_SLICES_12_TO_36_GT_PREDICTION_COMPARISON.png'\n    )\n\n    create_full_fragment_comparison(\n        input_volume=test_volume,\n        ground_truth=test_mask,\n        prediction=binary_prediction,\n        save_path=comparison_path,\n        dice=avg_test_dice,\n        f1=f1,\n        precision=precision,\n        recall=recall\n    )\n\n    # ============================================\n    # FINAL RESULTS\n    # ============================================\n    total_time = time.time() - start_time\n    hours = int(total_time // 3600)\n    minutes = int((total_time % 3600) // 60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"FINAL TEST RESULTS ON FRAGMENT 1\")\n    print(\"=\"*60)\n    print(f\"Total Time: {hours}h {minutes}m\")\n    print(f\"Test Dice Score: {avg_test_dice:.4f}\")\n    print(f\"Precision: {precision:.4f}\")\n    print(f\"Recall: {recall:.4f}\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(\"-\"*60)\n    print(\"Confusion Matrix:\")\n    print(f\"  True Positives: {tp}\")\n    print(f\"  True Negatives: {tn}\")\n    print(f\"  False Positives: {fp}\")\n    print(f\"  False Negatives: {fn}\")\n    print(\"-\"*60)\n    print(\"Full Fragment-1 Surface Prediction:\")\n    print(f\"  Shape: {binary_prediction.shape}\")\n    print(f\"  Predicted ink pixels: {binary_prediction.sum():,}\")\n    print(f\"  Ground truth ink pixels: {test_mask.sum():,}\")\n    print(\"=\"*60)\n    \n    # Save results\n    with open(os.path.join(OUTPUT_DIR, \"final_results.txt\"), \"w\") as f:\n        f.write(\"VESUVIUS INK DETECTION - FINAL RESULTS\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n        f.write(f\"Total Time: {hours}h {minutes}m\\n\")\n        f.write(f\"Best Validation Dice: {best_val_dice:.4f}\\n\")\n        f.write(f\"Test Dice Score: {avg_test_dice:.4f}\\n\")\n        f.write(f\"Precision: {precision:.4f}\\n\")\n        f.write(f\"Recall: {recall:.4f}\\n\")\n        f.write(f\"F1 Score: {f1:.4f}\\n\")\n        f.write(\"-\"*60 + \"\\n\")\n        f.write(f\"True Positives: {tp}\\n\")\n        f.write(f\"True Negatives: {tn}\\n\")\n        f.write(f\"False Positives: {fp}\\n\")\n        f.write(f\"False Negatives: {fn}\\n\")\n        f.write(\"-\"*60 + \"\\n\")\n        f.write(f\"Input Slices Used: {SLICE_START}-{SLICE_END - 1} ({NUM_INPUT_SLICES} slices)\\n\")\n        f.write(f\"Full Volume Prediction Shape: {binary_prediction.shape}\\n\")\n        f.write(f\"Predicted Ink Pixels: {binary_prediction.sum():,}\\n\")\n        f.write(f\"Ground Truth Ink Pixels: {test_mask.sum():,}\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n    \n    print(f\"\\nResults saved to: {os.path.join(OUTPUT_DIR, 'final_results.txt')}\")\n    \n    # Plot training curves\n    plt.figure(figsize=(12, 4))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses)\n    plt.title('Training Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.grid(True)\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(val_dice_scores)\n    plt.title('Validation Dice Score')\n    plt.xlabel('Epoch')\n    plt.ylabel('Dice')\n    plt.grid(True)\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'))\n    plt.show()\n    \n    return avg_test_dice\n\nif __name__ == \"__main__\":\n    try:\n        test_dice = main()\n        print(f\"\\n✅ Final Test Dice Score: {test_dice:.4f}\")\n    except Exception as e:\n        print(f\"\\n❌ Error: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-20T10:37:50.185968Z","iopub.execute_input":"2026-08-20T10:37:50.186408Z","iopub.status.idle":"2026-08-20T10:47:03.612793Z","shell.execute_reply.started":"2026-08-20T10:37:50.186372Z","shell.execute_reply":"2026-08-20T10:47:03.611508Z"}},"outputs":[{"name":"stdout","text":"Requirement already satisfied: segmentation-models-pytorch==0.2.0 in /opt/conda/lib/python3.7/site-packages (0.2.0)\nRequirement already satisfied: timm==0.4.12 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.4.12)\nRequirement already satisfied: pretrainedmodels==0.7.4 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.7.4)\nRequirement already satisfied: efficientnet-pytorch==0.6.3 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.6.3)\nRequirement already satisfied: torchvision>=0.5.0 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.14.0)\nRequirement already satisfied: torch in /opt/conda/lib/python3.7/site-packages (from efficientnet-pytorch==0.6.3->segmentation-models-pytorch==0.2.0) (1.13.0)\nRequirement already satisfied: munch in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (2.5.0)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (4.64.1)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (4.4.0)\nRequirement already satisfied: pillow!=8.3.*,>=5.3.0 in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (9.4.0)\nRequirement already satisfied: requests in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.28.2)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.21.6)\nRequirement already satisfied: six in /opt/conda/lib/python3.7/site-packages (from munch->pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (1.16.0)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (3.4)\nRequirement already satisfied: urllib3<1.27,>=1.21.1 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.26.14)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.1.1)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2022.12.7)\n\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0mUsing device: cuda\nGPU Memory: 17.06 GB\n============================================================\nVESUVIUS INK DETECTION - OPTIMIZED PIPELINE\n============================================================\n    GPU Memory - Allocated: 0.00GB, Cached: 0.00GB\n    CPU Memory: 0.53GB\n\n1. Loading training data...\n\n   Processing Fragment 2...\n    Volume shape: (14830, 9506, 15)\n    Mask shape: (14830, 9506)\n    Ink pixels: 16,865,122\n    Generating ignore mask...\n    Ignored pixels: 741,748\n    Extracting patches...\n    Total possible patches: 135240\n    Sampling 3000 patches...\n    Extracted: 841 ink patches, 420 background patches\n    Total extracted: 1261 patches\n    GPU Memory - Allocated: 0.00GB, Cached: 0.00GB\n    CPU Memory: 9.27GB\n\n   Processing Fragment 3...\n    Volume shape: (7606, 5249, 15)\n    Mask shape: (7606, 5249)\n    Ink pixels: 3,172,466\n    Generating ignore mask...\n    Ignored pixels: 171,140\n    Extracting patches...\n    Total possible patches: 37674\n    Sampling 3000 patches...\n    Extracted: 647 ink patches, 323 background patches\n    Total extracted: 970 patches\n    GPU Memory - Allocated: 0.00GB, Cached: 0.00GB\n    CPU Memory: 11.75GB\n\nTotal patches: 2231\n    GPU Memory - Allocated: 0.00GB, Cached: 0.00GB\n    CPU Memory: 11.75GB\n\n2. Creating train/validation split...\n   Train samples: 1785\n   Validation samples: 446\n\n3. Setting up data augmentation...\n   Train batches: 223\n   Val batches: 56\n\n4. Initializing model...\n   Total parameters: 7,856,001\n   Trainable parameters: 7,856,001\n\n5. Starting training...\n============================================================\n","output_type":"stream"},{"name":"stderr","text":"Epoch 1/10 [Train]: 100%|██████████| 223/223 [00:19<00:00, 11.69it/s, loss=0.5991]\nEpoch 1/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 56.27it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 1/10 - Time: 20.1s\n  Train Loss: 0.6450\n  Val Dice: 0.2041\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.73GB\n  ✓ New best model saved! Dice: 0.2041\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 2/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.69it/s, loss=0.6248]\nEpoch 2/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 57.57it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 2/10 - Time: 17.3s\n  Train Loss: 0.6407\n  Val Dice: 0.2279\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n  ✓ New best model saved! Dice: 0.2279\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 3/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.66it/s, loss=0.6706]\nEpoch 3/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 56.90it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 3/10 - Time: 17.3s\n  Train Loss: 0.6321\n  Val Dice: 0.2035\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 4/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.66it/s, loss=0.5749]\nEpoch 4/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 56.51it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 4/10 - Time: 17.3s\n  Train Loss: 0.6315\n  Val Dice: 0.1047\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 5/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.72it/s, loss=0.6409]\nEpoch 5/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 57.01it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 5/10 - Time: 17.3s\n  Train Loss: 0.6317\n  Val Dice: 0.2280\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n  ✓ New best model saved! Dice: 0.2280\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 6/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.68it/s, loss=0.6138]\nEpoch 6/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 57.70it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 6/10 - Time: 17.3s\n  Train Loss: 0.6278\n  Val Dice: 0.1727\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 7/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.68it/s, loss=0.5996]\nEpoch 7/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 57.93it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 7/10 - Time: 17.3s\n  Train Loss: 0.6238\n  Val Dice: 0.2119\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 8/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.66it/s, loss=0.5486]\nEpoch 8/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 57.90it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 8/10 - Time: 17.3s\n  Train Loss: 0.6275\n  Val Dice: 0.2908\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.74GB\n  ✓ New best model saved! Dice: 0.2908\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 9/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.63it/s, loss=0.6690]\nEpoch 9/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 58.10it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 9/10 - Time: 17.3s\n  Train Loss: 0.6229\n  Val Dice: 0.2676\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.73GB\n------------------------------------------------------------\n","output_type":"stream"},{"name":"stderr","text":"Epoch 10/10 [Train]: 100%|██████████| 223/223 [00:16<00:00, 13.68it/s, loss=0.5683]\nEpoch 10/10 [Val]: 100%|██████████| 56/56 [00:00<00:00, 56.14it/s]\n","output_type":"stream"},{"name":"stdout","text":"\nEpoch 10/10 - Time: 17.3s\n  Train Loss: 0.6221\n  Val Dice: 0.3527\n  LR: 1.00e-04\n    GPU Memory - Allocated: 0.14GB, Cached: 2.83GB\n    CPU Memory: 13.73GB\n  ✓ New best model saved! Dice: 0.3527\n------------------------------------------------------------\n\n============================================================\nTRAINING COMPLETE\nBest Validation Dice: 0.3527\n============================================================\n\n6. Testing on Fragment 1...\n============================================================\nLoading test data...\nExtracting test patches...\n    Total possible patches: 48888\n    Sampling 3000 patches...\n    Extracted: 708 ink patches, 354 background patches\nTest patches: 1062\nRunning inference...\n","output_type":"stream"},{"name":"stderr","text":"Testing: 100%|██████████| 133/133 [00:02<00:00, 55.73it/s]\n","output_type":"stream"},{"name":"stdout","text":"\n7. Generating FULL Fragment-1 surface prediction...\n============================================================\n    Full-volume patches: 12446\n    Input slices: 15 (15-29)\n    Patch size: 128, inference stride: 64\nFull binary prediction saved to: /kaggle/working/fragment1_full_surface_binary_prediction.png\nFull probability prediction saved to: /kaggle/working/fragment1_full_surface_probability.tif\n✓ Input stack visualization saved to: /kaggle/working/fragment1_input_slices_12_to_36.png\n✓ Full Fragment-1 comparison saved to: /kaggle/working/FRAGMENT1_ALL_SLICES_12_TO_36_GT_PREDICTION_COMPARISON.png\n\n============================================================\nFINAL TEST RESULTS ON FRAGMENT 1\n============================================================\nTotal Time: 0h 8m\nTest Dice Score: 0.4278\nPrecision: 0.3737\nRecall: 0.6530\nF1 Score: 0.4753\n------------------------------------------------------------\nConfusion Matrix:\n  True Positives: 3300763\n  True Negatives: 6812275\n  False Positives: 5532643\n  False Negatives: 1754127\n------------------------------------------------------------\nFull Fragment-1 Surface Prediction:\n  Shape: (8181, 6330)\n  Predicted ink pixels: 15,125,184\n  Ground truth ink pixels: 5,339,362\n============================================================\n\nResults saved to: /kaggle/working/final_results.txt\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"<Figure size 1200x400 with 2 Axes>","image/png":"iVBORw0KGgoAAAANSUhEUgAABKUAAAGGCAYAAACqvTJ0AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/NK7nSAAAACXBIWXMAAA9hAAAPYQGoP6dpAACrnElEQVR4nOzdeVxU1f/H8dcMOwioIKiIoKLivitopuaWLdpilpa5pqb5rWz5Zra4/fKbldmmlrvmVqllZSrummuu5b7jgiAqIiDrzO8PkiRcEIE7wPv5ePCouXPuve+ZIzB87rnnmKxWqxUREREREREREZF8ZDY6gIiIiIiIiIiIFD0qSomIiIiIiIiISL5TUUpERERERERERPKdilIiIiIiIiIiIpLvVJQSEREREREREZF8p6KUiIiIiIiIiIjkOxWlREREREREREQk36koJSIiIiIiIiIi+U5FKRERERERERERyXcqSolInjOZTNn6Wrt27T2dZ/jw4ZhMphztu3bt2lzJcC/n/uGHH/L93CIiIkXZ448/jouLCzExMbds8+yzz+Lg4EBkZGS2j2symRg+fHjG47v5nNGzZ08CAwOzfa4bTZgwgRkzZmTZfvLkSUwm002fy2vXP59d/3J1daVcuXK0b9+eL774gqtXr2bZ517eg3tx4MABunfvTsWKFXF2dsbb25v69evz0ksvERsbm+95RIoCe6MDiEjht3nz5kyPR40axZo1a1i9enWm7dWrV7+n8/Tt25cHH3wwR/vWr1+fzZs333MGERERKTj69OnDjz/+yNy5cxk4cGCW569cucLixYt55JFH8PX1zfF58utzxoQJE/D29qZnz56ZtpcpU4bNmzdTqVKlPD3/7SxbtgxPT0+Sk5M5d+4cq1at4s033+Sjjz7i559/pk6dOhlt3333XV5++eV8zbdr1y6aNWtGtWrVeO+99wgMDCQ6Opo9e/Ywf/58Xn/9dTw8PPI1k0hRoKKUiOS5kJCQTI9LlSqF2WzOsv3fEhIScHV1zfZ5ypUrR7ly5XKU0cPD4455REREpHDp0KEDZcuWZdq0aTctSs2bN49r167Rp0+fezqP0Z8znJycDP+c06BBA7y9vTMeP/PMM7z00ku0aNGCjh07cvjwYZycnAAMKZ6NHz8es9nM2rVrcXd3z9jeuXNnRo0ahdVqzbcsd/sZWKQg0+17ImITWrZsSc2aNVm/fj1NmzbF1dWV3r17A7BgwQLatWtHmTJlcHFxoVq1arz11lvEx8dnOsbNbt8LDAzkkUceYdmyZdSvXx8XFxeCg4OZNm1apnY3G1bfs2dPihUrxtGjR3nooYcoVqwY/v7+vPbaayQlJWXa/8yZM3Tu3Bl3d3eKFy/Os88+y/bt23N1qPxff/1Fp06dKFGiBM7OztStW5eZM2dmamOxWBg9ejRVq1bFxcWF4sWLU7t2bT777LOMNhcuXKBfv374+/vj5OREqVKlaNasGStXrsyVnCIiIgWFnZ0dPXr0YMeOHfz5559Znp8+fTplypShQ4cOXLhwgYEDB1K9enWKFSuGj48PDzzwABs2bLjjeW51+96MGTOoWrUqTk5OVKtWjVmzZt10/xEjRtCkSRNKliyJh4cH9evXZ+rUqZkKJYGBgezbt49169Zl3Cp3/Ra4W92+t3HjRlq3bo27uzuurq40bdqUX3/9NUtGk8nEmjVrePHFF/H29sbLy4snnniCc+fO3fG1306dOnUYNmwY4eHhLFiwIGP7zW7fs1gsfPHFF9StWzfjM05ISAhLlizJ1G7BggWEhobi5uZGsWLFaN++Pbt27bpjlosXL+Lh4UGxYsVu+vy/P2MuW7aM1q1b4+npiaurK9WqVWPMmDGZ2ixZsoTQ0FBcXV1xd3enbdu2We4guP75defOnXTu3JkSJUpkFOWsVisTJkzIeM0lSpSgc+fOHD9+/I6vR6SgUFFKRGxGREQEzz33HN26dWPp0qUZVyyPHDnCQw89xNSpU1m2bBmvvPIK3333HY8++mi2jrtnzx5ee+01Xn31VX766Sdq165Nnz59WL9+/R33TUlJoWPHjrRu3ZqffvqJ3r178+mnn/Lhhx9mtImPj6dVq1asWbOGDz/8kO+++w5fX1+efvrpnL0RN3Ho0CGaNm3Kvn37+Pzzz1m0aBHVq1enZ8+ejB07NqPd2LFjGT58OF27duXXX39lwYIF9OnTJ9NcGd27d+fHH3/kvffeY8WKFUyZMoU2bdpw8eLFXMsrIiJSUPTu3RuTyZTlgtX+/fvZtm0bPXr0wM7OjkuXLgHw/vvv8+uvvzJ9+nQqVqxIy5YtczQn5YwZM+jVqxfVqlVj4cKFvPPOO4waNSrL9AaQXlTq378/3333HYsWLeKJJ55g8ODBjBo1KqPN4sWLqVixIvXq1WPz5s1s3ryZxYsX3/L869at44EHHuDKlStMnTqVefPm4e7uzqOPPpqpQHRd3759cXBwYO7cuYwdO5a1a9fy3HPP3fXr/reOHTsC3PFzWc+ePXn55Zdp1KgRCxYsYP78+XTs2JGTJ09mtPnggw/o2rUr1atX57vvvmP27NlcvXqV5s2bs3///tsePzQ0lIiICJ599lnWrVvHtWvXbtl26tSpPPTQQ1gsFiZNmsTPP//Mf/7zH86cOZPRZu7cuXTq1AkPDw/mzZvH1KlTuXz5Mi1btmTjxo1ZjvnEE08QFBTE999/z6RJkwDo378/r7zyCm3atOHHH39kwoQJ7Nu3j6ZNm97VHGciNs0qIpLPevToYXVzc8u0rUWLFlbAumrVqtvua7FYrCkpKdZ169ZZAeuePXsynnv//fet//6xFhAQYHV2draeOnUqY9u1a9esJUuWtPbv3z9j25o1a6yAdc2aNZlyAtbvvvsu0zEfeugha9WqVTMef/XVV1bA+ttvv2Vq179/fytgnT59+m1f0/Vzf//997ds88wzz1idnJys4eHhmbZ36NDB6urqao2JibFarVbrI488Yq1bt+5tz1esWDHrK6+8cts2IiIiRUmLFi2s3t7e1uTk5Ixtr732mhWwHj58+Kb7pKamWlNSUqytW7e2Pv7445meA6zvv/9+xuN/f85IS0uzli1b1lq/fn2rxWLJaHfy5Emrg4ODNSAg4JZZ09LSrCkpKdaRI0davby8Mu1fo0YNa4sWLbLsc+LEiSyfSUJCQqw+Pj7Wq1evZnpNNWvWtJYrVy7juNOnT7cC1oEDB2Y65tixY62ANSIi4pZZrdZ/Pp9duHDhps9fu3bNClg7dOiQsa1Hjx6Z3oP169dbAeuwYcNueZ7w8HCrvb29dfDgwZm2X7161Vq6dGlrly5dbpszMTHR+thjj1kBK2C1s7Oz1qtXzzps2DBrVFRUpuN5eHhY77vvvkzv/Y2u92+tWrWsaWlpmfb18fGxNm3aNGPb9ffnvffey3SMzZs3WwHrJ598kmn76dOnrS4uLtY333zztq9HpKDQSCkRsRklSpTggQceyLL9+PHjdOvWjdKlS2NnZ4eDgwMtWrQA0ldJuZO6detSvnz5jMfOzs5UqVKFU6dO3XFfk8mUZURW7dq1M+27bt063N3ds0yy3rVr1zseP7tWr15N69at8ff3z7S9Z8+eJCQkZAwFb9y4MXv27GHgwIEsX778pivFNG7cmBkzZjB69Gi2bNlCSkpKruUUEREpiPr06UN0dHTGrWCpqal8++23NG/enMqVK2e0mzRpEvXr18fZ2Rl7e3scHBxYtWpVtj6P3OjQoUOcO3eObt26ZbotLCAggKZNm2Zpv3r1atq0aYOnp2fGZ6H33nuPixcvEhUVddevNz4+nq1bt9K5c+dMt6vZ2dnRvXt3zpw5w6FDhzLtc31E03W1a9cGyNbnqduxZmOupt9++w2AQYMG3bLN8uXLSU1N5fnnnyc1NTXjy9nZmRYtWtxxNJuTkxOLFy9m//79fPrppzzzzDNcuHCB//u//6NatWoZ78emTZuIjY1l4MCBt1z1+Xr/du/eHbP5nz+5ixUrxpNPPsmWLVtISEjItM+TTz6Z6fEvv/yCyWTiueeey/R6SpcuTZ06dQxZMVokL6goJSI2o0yZMlm2xcXF0bx5c7Zu3cro0aNZu3Yt27dvZ9GiRQC3HVp9nZeXV5ZtTk5O2drX1dUVZ2fnLPsmJiZmPL548eJNV+S5l1V6/u3ixYs3fX/Kli2b8TzA0KFD+fjjj9myZQsdOnTAy8uL1q1b88cff2Tss2DBAnr06MGUKVMIDQ2lZMmSPP/885w/fz7X8oqIiBQknTt3xtPTk+nTpwOwdOlSIiMjM01wPm7cOF588UWaNGnCwoUL2bJlC9u3b+fBBx/M1meKG13/vV26dOksz/1727Zt22jXrh0AkydP5vfff2f79u0MGzYMyN5noX+7fPkyVqs1W58trvv356nrk5Ln5Pw3ul7Uun7em7lw4QJ2dnY3fb+uu347W6NGjXBwcMj0tWDBAqKjo7OVp1q1arzyyit8++23hIeHM27cOC5evMi7776bkQW47eI619+7W72/FouFy5cvZ9r+77aRkZFYrVZ8fX2zvJ4tW7Zk+/WI2DqtviciNuNmV5tWr17NuXPnWLt2bcboKCDTHElG8/LyYtu2bVm252aRx8vLi4iIiCzbr08wen01G3t7e4YMGcKQIUOIiYlh5cqVvP3227Rv357Tp0/j6uqKt7c348ePZ/z48YSHh7NkyRLeeustoqKiWLZsWa5lFhERKShcXFzo2rUrkydPJiIigmnTpuHu7s5TTz2V0ebbb7+lZcuWTJw4MdO+V69evevzXS/w3Oyzwr+3zZ8/HwcHB3755ZdMF8p+/PHHuz7vdSVKlMBsNmfrs0Veuz46rWXLlrdsU6pUKdLS0jh//vxNCz3wT94ffviBgICAXMlmMpl49dVXGTlyJH/99VdGFiDT/FH/dr1/b/X+ms1mSpQokeVcN/L29sZkMrFhw4aMAuCNbrZNpCDSSCkRsWnXf0H/+xfv119/bUScm2rRogVXr17NGFp+3fz583PtHK1bt84o0N1o1qxZuLq63nSZ5+LFi9O5c2cGDRrEpUuXMk0Eel358uV56aWXaNu2LTt37sy1vCIiIgVNnz59SEtL46OPPmLp0qU888wzuLq6ZjxvMpmyfB7Zu3dvltXUsqNq1aqUKVOGefPmZbp97dSpU2zatClTW5PJhL29PXZ2dhnbrl27xuzZs7McN7sjwd3c3GjSpAmLFi3K1N5isfDtt99Srlw5qlSpctev627t2bOHDz74gMDAQLp06XLLdh06dADIUhC8Ufv27bG3t+fYsWM0bNjwpl+3c7MCEqQXkWJjYzNGcjVt2hRPT08mTZp0y1sPq1atip+fH3Pnzs3UJj4+noULF2asyHc7jzzyCFarlbNnz970tdSqVeu2+4sUFBopJSI2rWnTppQoUYIBAwbw/vvv4+DgwJw5c9izZ4/R0TL06NGDTz/9lOeee47Ro0cTFBTEb7/9xvLlywEyzSVwO1u2bLnp9hYtWvD+++/zyy+/0KpVK9577z1KlizJnDlz+PXXXxk7diyenp4APProo9SsWZOGDRtSqlQpTp06xfjx4wkICKBy5cpcuXKFVq1a0a1bN4KDg3F3d2f79u0sW7aMJ554InfeEBERkQKoYcOG1K5dm/Hjx2O1WjPdugfpRYJRo0bx/vvv06JFCw4dOsTIkSOpUKECqampd3Uus9nMqFGj6Nu3L48//jgvvPACMTExDB8+PMstag8//DDjxo2jW7du9OvXj4sXL/Lxxx/fdKRMrVq1mD9/PgsWLKBixYo4OzvfsngxZswY2rZtS6tWrXj99ddxdHRkwoQJ/PXXX8ybN++W8yXl1I4dO/D09CQlJYVz586xatUqZs+ejY+PDz///DOOjo633Ld58+Z0796d0aNHExkZySOPPIKTkxO7du3C1dWVwYMHExgYyMiRIxk2bBjHjx/nwQcfpESJEkRGRrJt2zbc3NwYMWLELc/Rr18/YmJiePLJJ6lZsyZ2dnYcPHiQTz/9FLPZzH//+18gfV6oTz75hL59+9KmTRteeOEFfH19OXr0KHv27OHLL7/EbDYzduxYnn32WR555BH69+9PUlISH330ETExMfzvf/+74/vVrFkz+vXrR69evfjjjz+4//77cXNzIyIigo0bN1KrVi1efPHFu+8IERujopSI2DQvLy9+/fVXXnvtNZ577jnc3Nzo1KkTCxYsoH79+kbHA9KvNq5evZpXXnmFN998E5PJRLt27ZgwYQIPPfQQxYsXz9ZxPvnkk5tuX7NmDS1btmTTpk28/fbbDBo0iGvXrlGtWjWmT59Oz549M9q2atWKhQsXMmXKFGJjYyldujRt27bl3XffxcHBAWdnZ5o0acLs2bM5efIkKSkplC9fnv/+97+8+eabufBuiIiIFFx9+vTh5Zdfpnr16jRp0iTTc8OGDSMhIYGpU6cyduxYqlevzqRJk1i8eHGOJp2+XvT68MMPeeKJJwgMDOTtt99m3bp1mY73wAMPMG3aND788EMeffRR/Pz8eOGFF/Dx8clSOBsxYgQRERG88MILXL16lYCAgJuOlIb0i16rV6/m/fffp2fPnlgsFurUqcOSJUt45JFH7vr13Mn1BWGcnJwoWbIktWrV4sMPP6RXr164u7vfcf8ZM2ZQv359pk6dyowZM3BxcaF69eq8/fbbGW2GDh1K9erV+eyzz5g3bx5JSUmULl2aRo0aMWDAgNsef/DgwSxYsIDJkydz9uxZ4uPjKVWqFKGhocyaNSvTqPQ+ffpQtmxZPvzwQ/r27YvVaiUwMJAePXpktOnWrRtubm6MGTOGp59+Gjs7O0JCQlizZs1NJ7O/ma+//pqQkBC+/vprJkyYgMVioWzZsjRr1ozGjRtn6xgits5kzc5yByIictc++OAD3nnnHcLDw287GaaIiIiIiEhRpJFSIiK54MsvvwQgODiYlJQUVq9ezeeff85zzz2ngpSIiIiIiMhNqCglIpILXF1d+fTTTzl58iRJSUkZt8S98847RkcTERERERGxSbp9T0RERERERERE8l32loQSERERERERERHJRSpKiYiIiIiIiIhIvlNRSkRERERERERE8p0mOs8hi8XCuXPncHd3x2QyGR1HRERE8onVauXq1auULVsWs1nX925Hn5dERESKpux+XlJRKofOnTuHv7+/0TFERETEIKdPn6ZcuXJGx7Bp+rwkIiJStN3p85KKUjnk7u4OpL/BHh4euXrslJQUVqxYQbt27XBwcMjVY8u9U//YNvWP7VLf2Db1T/bFxsbi7++f8VlAbi0vPy+B/t3aMvWNbVP/2C71jW1T/2Rfdj8vqSiVQ9eHoHt4eORJUcrV1RUPDw/9Q7dB6h/bpv6xXeob26b+uXu6He3O8vLzEujfrS1T39g29Y/tUt/YNvXP3bvT5yVNhCAiIiIiIiIiIvlORSkREREREREREcl3KkqJiIiIiIiIiEi+U1FKRERERERERETynYpSIiIiIiIiIiKS71SUEhERERERERGRfKeilIiIiIiIiIiI5DvDi1ITJkygQoUKODs706BBAzZs2HDb9klJSQwbNoyAgACcnJyoVKkS06ZNu2nb+fPnYzKZeOyxxzJtHz58OCaTKdNX6dKlc+sliYiIiIiIiIjIHdgbefIFCxbwyiuvMGHCBJo1a8bXX39Nhw4d2L9/P+XLl7/pPl26dCEyMpKpU6cSFBREVFQUqampWdqdOnWK119/nebNm9/0ODVq1GDlypUZj+3s7HLnRYmIiIiIiIiIyB0ZWpQaN24cffr0oW/fvgCMHz+e5cuXM3HiRMaMGZOl/bJly1i3bh3Hjx+nZMmSAAQGBmZpl5aWxrPPPsuIESPYsGEDMTExWdrY29vb9OiopDSjE4iIiIiIiIhIYZWYkoazg7EDdAwrSiUnJ7Njxw7eeuutTNvbtWvHpk2bbrrPkiVLaNiwIWPHjmX27Nm4ubnRsWNHRo0ahYuLS0a7kSNHUqpUKfr06XPL2wGPHDlC2bJlcXJyokmTJnzwwQdUrFjxlnmTkpJISkrKeBwbGwtASkoKKSkp2X7d2bFy/3lG7LTDr0Y0jSt65+qx5d5d7+/c7nfJHeof26W+sW3qn+yz1fdowoQJfPTRR0RERFCjRg3Gjx9/yxHjGzdu5L///S8HDx4kISGBgIAA+vfvz6uvvprRZsaMGfTq1SvLvteuXcPZ2TnPXoeIiIjkvbMx13j0i430CA1kUKtK2NsZM7uTYUWp6Oho0tLS8PX1zbTd19eX8+fP33Sf48ePs3HjRpydnVm8eDHR0dEMHDiQS5cuZcwr9fvvvzN16lR27959y3M3adKEWbNmUaVKFSIjIxk9ejRNmzZl3759eHl53XSfMWPGMGLEiCzbV6xYgaurazZf9Z1ZrfDNQTPxqWZ6z9rBwGppBLrn2uElF4WFhRkdQW5D/WO71De2Tf1zZwkJCUZHyOJup0Rwc3PjpZdeonbt2ri5ubFx40b69++Pm5sb/fr1y2jn4eHBoUOHMu2rgpSIiEjBN27FYS7FJ7P5eDT/aR1kWA5Db98DMJlMmR5brdYs266zWCyYTCbmzJmDp6cnkH4LYOfOnfnqq69ITU3lueeeY/LkyXh733qEUYcOHTL+v1atWoSGhlKpUiVmzpzJkCFDbrrP0KFDMz0XGxuLv78/7dq1w8PDI9uvNzuat0rk6a/WcjTWzOQjzszq1YBafp65eg7JuZSUFMLCwmjbti0ODg5Gx5F/Uf/YLvWNbVP/ZN/10dK25G6nRKhXrx716tXLeBwYGMiiRYvYsGFDpqKUFoMREREpfA5ExLJo1xkA3upQ7ZY1mPxgWFHK29sbOzu7LKOioqKisoyeuq5MmTL4+fllFKQAqlWrhtVq5cyZM8THx3Py5EkeffTRjOctFguQPofUoUOHqFSpUpbjurm5UatWLY4cOXLLvE5OTjg5OWXZ7uDgkOsf3j1coV+whe8iS/LHqRh6ztjB3BdCqKnClE3Ji76X3KP+sV3qG9um/rkzW3t/cjIlwr/t2rWLTZs2MXr06Ezb4+LiCAgIIC0tjbp16zJq1KhMxax/y8/pDq4f98b/iu1Q39g29Y/tUt/YtsLSP/9begCrFTrU8KVGabc8/R19J4YVpRwdHWnQoAFhYWE8/vjjGdvDwsLo1KnTTfdp1qwZ33//PXFxcRQrVgyAw4cPYzabKVeuHCaTiT///DPTPu+88w5Xr17ls88+w9/f/6bHTUpK4sCBA7ecd8EITnYwuXt9+s7exY5Tl3lu6lbm9g2hetncHZUlIiIiBVtOpkS4rly5cly4cIHU1FSGDx+eMdIKIDg4mBkzZlCrVi1iY2P57LPPaNasGXv27KFy5co3PV5+TXfwb7rt1Hapb2yb+sd2qW9sW0HunyNXTKw7YofZZKWBw1mWLj2bJ+fJ7nQHht6+N2TIELp3707Dhg0JDQ3lm2++ITw8nAEDBgDpt8ydPXuWWbNmAdCtWzdGjRpFr169GDFiBNHR0bzxxhv07t07Y6LzmjVrZjpH8eLFs2x//fXXefTRRylfvjxRUVGMHj2a2NhYevTokQ+vOvuKOdkzo1cjuk/dxu7TMTw3dSvzXgihamlNMiUiIiKZ3c2UCNdt2LCBuLg4tmzZwltvvUVQUBBdu3YFICQkhJCQkIy2zZo1o379+nzxxRd8/vnnNz1efk53ALrt1Japb2yb+sd2qW9sW0HvH6vVytSvtwKxdGtcnh6PVMuzc2V3ugNDi1JPP/00Fy9eZOTIkURERFCzZk2WLl1KQEAAABEREYSHh2e0L1asGGFhYQwePJiGDRvi5eVFly5dsgw1v5MzZ87QtWtXoqOjKVWqFCEhIWzZsiXjvLbE3dmBmb0b033qVvaeucKzU7Ywv18IQT4qTImIiEjOpkS4rkKFCkD6HJuRkZEMHz48oyj1b2azmUaNGtnMdAf5eXzJOfWNbVP/2C71jW0rqP3zy95z7D0bi5ujHS+3qZrnv5uzw/CJzgcOHMjAgQNv+tyMGTOybAsODr6roXI3O8b8+fOzvb8t8HRxYHbvJnSbsoV952LpOnkr8/uFUKlUMaOjiYiIiMFyMiXCzVit1kzzQd3s+d27d1OrVq17yisiIiL5LznVwkfL01fUfeH+ipRyz3oRyQiGF6UkezxdHfi2TxO6Tt7CwfNX6TZ5Cwv6hRLo7WZ0NBERETHY3U6J8NVXX1G+fHmCg4MB2LhxIx9//DGDBw/OOOaIESMICQmhcuXKxMbG8vnnn7N7926++uqr/H+BIiIick/mbw/n1MUEvIs58ULzikbHyaCiVAFSws2ROX3TC1OHI+Po+ndhqrxX3k0cKiIiIrbvbqdEsFgsDB06lBMnTmBvb0+lSpX43//+R//+/TPaxMTE0K9fP86fP4+npyf16tVj/fr1NG7cON9fn4iIiORcXFIqn61Mv/3+5TaVcXOynVKQ7SSRbPEq5sScviF0nbyFo1Hphan5/ULwL6nClIiISFF2N1MiDB48ONOoqJv59NNP+fTTT3MrnoiIiBjkm/XHuRifTAVvN55p5G90nEzMRgeQu1fK3Ym5fZtQ0duNszHX6DZlC+dirhkdS0RERERERERsSNTVRKZsOA7AG+2r4mBnW2Ug20oj2ebj4czcF0II9HLl9KVrdJ28hfNXEo2OJSIiIiIiIiI24vNVR0hITqOuf3E61CxtdJwsVJQqwEp7phem/Eu6cOpiAl0nbyEqVoUpERERERERkaLu+IU45m07DcBbHYIxmUwGJ8pKRakCrmxxF+a9EIJfcRdORMfTdfIWLly99XLOIiIiIiIiIlL4fbT8EGkWKw8E+xBS0cvoODelolQhUK6EK/P7hVDW05ljF+LpNnkLF+NUmBIREREREREpinaGX+a3v85jNsF/Hww2Os4tqShVSPiXdGXuCyGU9nDmSFQcz07ZyqX4ZKNjiYiIiIiIiEg+slqt/G/pQQCerF+OqqXdDU50aypKFSKB3m7MfaEJPu5OHDx/leembCUmQYUpERERERERkaJi9cEotp28hJO9mVfbVjE6zm2pKFXIVCxVjLkvhOBdzIn9EbF0n7qNK9dSjI4lIiIiIiIiInkszWLlw2Xpo6R6NgukbHEXgxPdnopShVCQTzHmvtAELzdH/jx7heenbSM2UYUpERERERERkcJs4c4zHI6Mw9PFgYEtgoyOc0cqShVSVXzd+bZvE0q4OrDndAw9p20jLinV6FgiIiIiIiIikgcSU9L4NOwwAC+1CsLT1cHgRHemolQhVq2MB9/2bYKniwM7w2PoNX0b8SpMiYiIiIiIiBQ6038/ScSVRPyKu9A9NMDoONmiolQhV6OsJ9/2aYK7sz3bT16m94ztJCSrMCUiIiIiIiJSWFyOT2bC2qMADGlbBWcHO4MTZY+KUkVArXKezO7TBHcne7aeuETfmX+QmJJmdCwRERERERERyQUT1h7lamIqwaXdeayen9Fxsk1FqSKirn9xZvRujJujHZuOXeSFWSpMiYiIiIiIiBR0Zy4nMHPTKQDe6hCMndlkcKLsU1GqCGkQUIIZvRvj6mjHhiPRDPh2B0mpKkyJiIiIiIiIFFTjVhwmOc1C00petKhSyug4d0VFqSKmUWBJpvVshLODmbWHLjBozk6SUy1GxxIRERERERGRu7T/XCyLd58F0kdJmUwFZ5QUqChVJIVU9GJaj0Y42ZtZeSCKwfN2kpKmwpSIiIiIiIhIQfLhsoNYrfBI7TLULlfc6Dh3TUWpIqppkDeTn2+Io72Z5fsieXn+LlJVmBIREREREREpEDYdjWbd4Qs42Jl4o31Vo+PkiIpSRdj9VUrxdfcGONqZWfrneV79bo8KUyIiIiIiIiI2zmKxMua3gwA82ySAAC83gxPljIpSRVyrqj5MeLY+DnYmft5zjjd+2EuaxWp0LBERERERERG5hV/+jODPs1dwc7TjpQeCjI6TYypKCW2q+/Jlt/rYm00s3nWW/y7ci0WFKRERERERERGbk5xq4ePlhwDo36IS3sWcDE6UcypKCQDta5Tm8671sDOb+GHHGd5e/KcKUyIiIiIiIiI2Zu7WU4RfSsC7mBN9m1cwOs49UVFKMjxUqwzjn66L2QTzt5/m3Z/+wmpVYUpERERERETEFlxNTOHz1UcBeKVNZVwd7Q1OdG9UlJJMHq1TlnFd6mIywZyt4Qxfsk+FKREREREREREb8M3641yKT6aitxtPN/I3Os49U1FKsnisnh8fda6DyQQzN59i1C8HVJgSERERERERMVBUbCJTNpwA4M0Hq+JgV/BLOgX/FUie6NygHP97ohYA034/wZjfDqowJSIiIiIiImKQ8auOcC0ljXrli9O+Rmmj4+QKFaXklp5uVJ7/e7wmkD5E8KPlh1SYEhEREREREclnxy7EsWD7aQCGdqiGyWQyOFHuUFFKbuvZJgGM7FQDgAlrj/HpyiMGJxIREREREREpWsYuO0iaxUqbaj40rlDS6Di5RkUpuaPnQwN575HqAHy+6gifr1JhSkRERERERCQ/7Dh1meX7IjGb4L8PBhsdJ1epKCXZ0vu+Cgx7qBoA48IO89WaowYnEhERERERESncrFYr//vtAABPNfCnsq+7wYlyl4pSkm0v3F+RNx+sCsBHyw/xzfpjBicSERERERERKbxWHohi+8nLONmbeaVtZaPj5DoVpeSuDGwZxGttqwDwwdKDTN14wuBEIiIiIiIiIoVPapqFD5cdBNLvXirj6WJwotynopTctcGtK/Of1ukV2lG/7GfmppPGBhIREREREREpZBbuPMPRqDiKuzowoEUlo+PkCRWlJEdebVOZQa3SvyneX7KPb7ecMjiRiIiIiIiISOFwLTmNcWGHAXipVRCeLg4GJ8obKkpJjphMJl5vV5X+91cE4J0f/2L+tnCDU4mIiIiIiIgUfNN+P0FkbBJ+xV3oHhpgdJw8o6KU5JjJZOKtDsH0ua8CAEMX/8n3f5w2OJWIiIiIiIhIwXUpPplJa9MXFnu9fRWc7O0MTpR3VJSSe2IymXjn4Wr0bBqI1QpvLtzL4l1njI4lIiIiIiIiUiB9teYoV5NSqV7Gg051/IyOk6dUlJJ7ZjKZeP/R6jwXUh6rFV77bg9L9pwzOpaIiIiIiIhIgXL6UgKzN6fP2fxWh2DMZpPBifKWilKSK0wmEyM71uSZRv5YrPDqgt38ujfC6FgiIiIiIiIiBcYnKw6RnGahWZAXzSt7Gx0nz6koJbnGbDbxweO16NygHGkWKy/P38Wyv84bHUtERERERETE5v119go/7k6/6+itB6thMhXuUVKgopTkMrPZxIdP1uaJen6kWqwMnreTlfsjjY4lIiIiIiIiYtM+XHYQgI51ylKrnKfBafKHilKS6+zMJj56qg4d65QlJc3KwDk7WXMwyuhYIiIiIiIiIjZp45FoNhyJxsHOxOvtqhodJ9+oKCV5ws5sYlyXOjxcqwzJaRb6f7uDg+djjY4lIiIiIiIiYlMsFitjfjsAwLNNAijv5WpwovyjopTkGXs7M+Ofqcv9VUqRnGphfNgRoyOJiIiIiIiI2JSf955j37lYijnZM/iBIKPj5CsVpSRPOdiZeffhaphMsGzfeQ5EaLSUiIiIiIiICEBSahofrzgEwIAWFfEq5mRwovylopTkucq+7jxUqwwAX64+anAaERGRwmnChAlUqFABZ2dnGjRowIYNG27ZduPGjTRr1gwvLy9cXFwIDg7m008/zdJu4cKFVK9eHScnJ6pXr87ixYvz8iWIiIgUOXO2hHP60jV83J3ofV8Fo+PkOxWlJF9cH4K49K8IDkdeNTiNiIhI4bJgwQJeeeUVhg0bxq5du2jevDkdOnQgPDz8pu3d3Nx46aWXWL9+PQcOHOCdd97hnXfe4Ztvvslos3nzZp5++mm6d+/Onj176N69O126dGHr1q359bJEREQKtdjEFL5YnT7NzSttquDqaG9wovxneFHqbq7qASQlJTFs2DACAgJwcnKiUqVKTJs27aZt58+fj8lk4rHHHrvn88q9CS7tQYeapbFa4QuNlhIREclV48aNo0+fPvTt25dq1aoxfvx4/P39mThx4k3b16tXj65du1KjRg0CAwN57rnnaN++fabPQ+PHj6dt27YMHTqU4OBghg4dSuvWrRk/fnw+vSoREZHC7Zt1x7mckEKlUm50aVjO6DiGMLQMd/2q3oQJE2jWrBlff/01HTp0YP/+/ZQvX/6m+3Tp0oXIyEimTp1KUFAQUVFRpKamZml36tQpXn/9dZo3b54r55V7N/iByvz213l+2XuOl1sHEeTjbnQkERGRAi85OZkdO3bw1ltvZdrerl07Nm3alK1j7Nq1i02bNjF69OiMbZs3b+bVV1/N1K59+/a3LUolJSWRlJSU8Tg2Nn0uyZSUFFJSUrKV5W5cP2ZeHFvujfrGtql/bJf6xrblZv9ExiYyZeNxAF5rUxmrJY0US9o9H9dWZPc9MrQodeNVPUi/Ird8+XImTpzImDFjsrRftmwZ69at4/jx45QsWRKAwMDALO3S0tJ49tlnGTFiBBs2bCAmJuaeziu5o3pZD9pV92XF/ki+WH2Uz56pZ3QkERGRAi86Opq0tDR8fX0zbff19eX8+fO33bdcuXJcuHCB1NRUhg8fnvHZCOD8+fN3fcwxY8YwYsSILNtXrFiBq2veLW8dFhaWZ8eWe6O+sW3qH9ulvrFtudE/84+ZSUwxU8HdSvKJP1h68t5z2ZKEhIRstTOsKJWTq3pLliyhYcOGjB07ltmzZ+Pm5kbHjh0ZNWoULi4uGe1GjhxJqVKl6NOnT5bb8nJ6NTE/r/wV5ur4wBYVWLE/kp/3nGPg/RWoWMrN6Eh3rTD3T2Gg/rFd6hvbpv7JPlt9j0wmU6bHVqs1y7Z/27BhA3FxcWzZsoW33nqLoKAgunbtmuNjDh06lCFDhmQ8jo2Nxd/fn3bt2uHh4XE3LydbUlJSCAsLo23btjg4OOT68SXn1De2Tf1ju9Q3ti23+udoVBxbt6TXH8Y83ZgGASVyK6LNuF4zuRPDilI5uap3/PhxNm7ciLOzM4sXLyY6OpqBAwdy6dKljHmlfv/9d6ZOncru3btz7bxgzJW/wlodr1nCzF+XzbwzdwPPVbYYHSfHCmv/FBbqH9ulvrFt6p87y+6Vv/zi7e2NnZ1dls8xUVFRWT7v/FuFCumr/NSqVYvIyEiGDx+eUZQqXbr0XR/TyckJJ6esS1k7ODjk6R9XeX18yTn1jW1T/9gu9Y1tu9f++XTVMSxWaFvdl5Agn1xMZjuy+/4YPrX73VyBs1gsmEwm5syZg6enJ5B+K17nzp356quvSE1N5bnnnmPy5Ml4e3vn2nkhf6/8FfbqePk6sTw+aQs7Lpr54LnmBHoVrNFShb1/Cjr1j+1S39g29U/2ZffKX35xdHSkQYMGhIWF8fjjj2dsDwsLo1OnTtk+jtVqzTQqPDQ0lLCwsEzzSq1YsYKmTZvmTnAREZEi6I+Tl1ixPxKzCf77YFWj4xjOsKJUTq7qlSlTBj8/v4yCFEC1atWwWq2cOXOG+Ph4Tp48yaOPPprxvMWSPhLH3t6eQ4cO4e/vn6OriUZc+Sus1fF6gV48EOzD6oNRfL3hFB8/VcfoSDlSWPunsFD/2C71jW1T/9yZLb4/Q4YMoXv37jRs2JDQ0FC++eYbwsPDGTBgAJB+ce3s2bPMmjULgK+++ory5csTHBwMwMaNG/n4448ZPHhwxjFffvll7r//fj788EM6derETz/9xMqVK9m4cWP+v0AREZFCwGq1Mua3gwB0aeivxb8As1EnvvGq3o3CwsJueQWuWbNmnDt3jri4uIxthw8fxmw2U65cOYKDg/nzzz/ZvXt3xlfHjh1p1aoVu3fvxt/fP0fnldz3n9aVAVi86yynLsYbnEZERKRge/rppxk/fjwjR46kbt26rF+/nqVLlxIQEABAREQE4eHhGe0tFgtDhw6lbt26NGzYkC+++IL//e9/jBw5MqNN06ZNmT9/PtOnT6d27drMmDGDBQsW0KRJk3x/fSIiIoXBiv2R7Dh1GWcHM6+2rWJ0HJtg6O17d3tVr1u3bowaNYpevXoxYsQIoqOjeeONN+jdu3fGROc1a9bMdI7ixYtn2X6n80req+tfnBZVSrHu8AW+WnOUsZ0L5mgpERERWzFw4EAGDhx40+dmzJiR6fHgwYMzjYq6lc6dO9O5c+fciCciIlKkpaZZGLssfZRUn/sq4OvhbHAi22BoUerpp5/m4sWLjBw5koiICGrWrHnbq3rFihUjLCyMwYMH07BhQ7y8vOjSpQujR4/O1fNK/vhP68qsO3yBRTvPMviByviXzLulokVERERERESM8v2OMxy7EE8JVwf6t6hkdBybYfhE53dzVQ8gODj4rlYGutkx7nReyR8NAkrQvLI3G45EM2HtUcY8UdvoSCIiIiIiIiK5KiE5lU/DDgPw0gOV8XC2vfkpjWLYnFIiAC//PbfUDzvOcOaybS2xLSIiIiIiIjmXkJzKlYQUo2MYbtrGE0RdTaJcCReeCylvdByboqKUGKphYEmaVvIiJc3KxLXHjI4jIiIiIiIiuSAxJY2OX/5Og9Fh/PeHvZy+VDQHIVyKT2bSuuMAvNG+Kk72dgYnsi0qSonhro+W+u6P05yLuWZwGhEREREREblX36w/ztGoOFItVhb8cZpWH68tksWpL1YfIS4plRplPXi0dlmj49gcFaXEcE0qehFSsSQpaVYmrdNoKRERERERkYLsbMw1Jqw9CqQPQmhe2TtTceqthUWjOBV+MYFvt5wC4K0OwZjNJoMT2R4VpcQm/Ofv0VLzt53m/JVEg9OIiIiIiIhITn2w9ACJKRYaVyjJK20qM7tPExa+GJpRnJq/Pb04NXRR4S5OfbziEClpVppX9qZ55VJGx7FJKkqJTQit6EXjwJIkp1k0WkpERERERKSA2nL8Ir/ujcBsguGP1sBkSh8d1CCgJLP7NOGHAf8Up+Ztu16c+rPQLXz119krLNlzDoD/PhhscBrbpaKU2ASTycTLbdJHS83dFk5UrEZLiYiIiIiIFCSpaRaGL9kHQLcm5ale1iNLm4aB6cWp7weEcl/Q9eJUOK0+XsvbiwtPcep/vx0E4LG6Zanp52lwGtulopTYjKaVvGgQUILkVEvG6gQiIiIiIiJSMMzbFs7B81fxdHHgtbZVb9u2UWBJvu2bXpxqFpS+Ivvcrf8Up84W4EWw1h++wMaj0TjamXmt3e3fh6JORSmxGSaTKWMlvjlbTxF1VaOlRERERERECoLL8cl8vOIwAK+3q0IJN8ds7dcosCRz+obwXf/MxamWH61hWAEsTlks1oxRUs+FBOBf0tXgRLZNRSmxKc0re1PXvzhJqRYmr9doKRERERERkYJgXNhhrlxLIbi0O10bl7/r/RtXSC9OLegXQtNK6cWpOX8Xp9758U/OFZDi1JI959gfEYu7kz0vPRBkdBybp6KU2JQb55aaveUU0XFJBicSERERERGR29l/LpY5W08B8P6jNbC3y3mpoUlFL+a+EML8fiGEVkwvTn27JZyWH63l3R//IuKK7RanklLT+HjFIQAGtKxEyWyOFivKVJQSm9OySinqlPMkMcXC5A0aLSUiIiIiImKrrFYrw3/eh8UKD9cuQ2glr1w5bkhFL+b1Sy9OhVRMX6l99pZTtBi7lvd+ss3i1OzNpzhz+Rq+Hk70blbB6DgFgopSYnNMJhP/+XtuqdmbT3EpPtngRCIiIiIiInIzv+yNYNuJSzg7mHn7oWq5fvyQil7M7xfKvBdCaFIhvTg1a/M/xanzV2xjLuIr11L4cs1RAF5tUwUXRzuDExUMKkqJTXog2Idafp4kJKdptJSIiIiIiIgNSkhOZczSAwC82CIIv+IueXau0EpeLOifXpxqfENx6v6xa3jfBopTX687RkxCCkE+xejcoJyhWQoSFaXEJt04WmrWppNc1mgpERERERERmzJp7THOXUnEr7gL/VtUzJdzhlbyYkG/EOa+0ITGgenFqZmbT3H/R2sYvmQfkbH5X5w6fyWRab+fAOC/Dwbf05xaRY3eKbFZbar5UL2MB/HJaUzdeMLoOCIiIiIiIvK305cSmPT3iunvPlINZ4f8u13NZDLRtJI3C/qHMLdvExoFliA51cKMTSdpPjb/i1Ofhh0mMcVCw4AStKnmk2/nLQxUlBKbdeNoqRmbTnIlIcXgRCIiIiIiIgIw+tf9JKdaaFrJi/Y1ShuSwWQy0TTIm+/6hzKnbxMaBvxTnLp/7BpG/LyPqDwuTh2JvMr3O04DMPShYEwmU56er7BRUUpsWrvqvgSXdicuKZWpv2u0lIiIiIiIiNE2Holm+b5I7Mwm3n+0huGFGJPJRLMgb74f8E9xKinVwvTf00dO5WVx6sNlh7BYoX0NXxoElMyTcxRmKkqJTTOb/xktNf33E1y5ptFSIiIiIiIiRklJszDi530AdA8JoGppd4MT/ePG4tS3fZrQ4F/FqZE/7yfqau4Vp7afvMTKA+nFuTcfDM614xYlKkqJzXuwRmmq+BbjamIqM34/aXQcERERERGRImv25lMciYqjpJsjr7apYnScmzKZTNxX2ZsfBoQyu09j6pcvTlKqhWm/n6D5h2sY9cu9F6esVisf/L3yYJeG/lQqVSw3ohc5KkqJzTObTQx+IH201NSNx4lN1GgpERERERGR/HYxLolPVx4G4PV2VfF0dTA40e2ZTCaaVy7FwhebMqt3Y+r9XZyauvEE949dw+hf9nPhalKOjr18XyS7wmNwcbDj1TaVczl50aGilBQID9UqQ5BPMWITU5mp0VIiIiIiIiL57uMVh7iamEqNsh483cjf6DjZZjKZuL9KKRbdUJxKTLEwZeMJmo9dzf/9enfFqdQ0C2OXHwSgb/MK+Hg451X0Qk9FKSkQ7MwmBj8QBMCUjSeIS0o1OJGIiIiIiEjR8eeZK8zfnr7K3IiONbAzF7xV5m4sTs3s3Zi6/unFqckb/ilORcfduTj1/c6zHL8QT0k3R/rdXzEfkhdeKkpJgfFI7bJULOXGlWspzNx00ug4IiIiIiIiRYLVamX4z/uwWuGxumVpGFiwV5kzmUy0qFKKxQObMqNXo8zFqQ/X8MHSA7csTiWlwRerjwEw+IEg3J1t+xZGW6eilBQYmUZLbThOvEZLiYiIiIiI5Lmfdp9jx6nLuDra8VaHakbHyTUmk4mWVX1YPLAp03s1oo5/ca6lpPHN+uM0/3ANY25SnFobYeJCXDLlS7rybJMAg5IXHipKSYHyaO2yVPB243JCCrO3nDI6joiIiIiISKEWl5SascrcoFZBlPYsfPMnmUwmWlX14cfrxalynlxLSePr68Wp3w5wMS6Ji/HJrDqXXkZ5vX1VHO1VUrlXegelQLG3MzOoVfpoqcnrj5OQrNFSIiIiIiIieeWrNUeJuppEgJcrfe6rYHScPJVRnBrUjOk9G1H7enFq3XGaj13DC7N3kpRmomZZDx6pVcbouIWCilJS4DxWtywBXq5cjE9mzpZwo+OIiIiIiIgUSiej45m64QQA7zxcHWcHO4MT5Q+TyUSrYB9+GtSMaT0bUrucJwnJafx5NhaAN9pVxlwAJ3q3RSpKSYFz42ipr9cf41pymsGJRERERERECp/Rv+4nOc3C/VVK0aaaj9Fx8p3JZOKBYN+M4lTTSiVpVcZC00peRkcrNFSUkgLp8Xp++Jd0IToumTlbNbeUiIiIiIhIblpzKIqVB6KwN5t475HqmExFd2TQ9eLUzJ4NeSzQYnScQkVFKSmQHOzMDGp5fbTUcRJTNFpKREREREQkNySnWhj1834AejULJMinmMGJpLBSUUoKrCfql8OvuAsXriYxb5vmlhIREREREckNMzad4Hh0PN7FHBncurLRcaQQU1FKCixHezMDW1UCYNK6YxotJSIiIiIico+iriby+aqjALz5YDAezg4GJ5LCTEUpKdA6NyhHWU9nImOT+O6P00bHERERERERKdDGLjtEXFIqdcp50rl+OaPjSCGnopQUaE72drzYMn201MS1x0hK1WgpERERERGRnNgVfpkfdpwBYHjHGpjNRXdyc8kfKkpJgdelkT+lPZyJuJLI93+cMTqOiIiIiIhIgWOxWBm+ZB8AT9YvR73yJQxOJEWBilJS4P17tFRyqpboFBERERERuRsLd55hz5krFHOy578PVjU6jhQRKkpJofB0I3983J04G3MtY7ipiIiIiIiI3FlsYgofLjsEwH9aB+Hj4WxwIikqVJSSQsHZwY4BLdJHS3215igpaRotJSIiIiIikh1frDpCdFwSFb3d6Nm0gtFxpAhRUUoKjW5NyuNdLH201KKdGi0lIiIiIiJyJ0ej4pj++0kA3n20Oo72KhNI/tG/Nik00kdLVQTgS42WEhERERERuS2r1crIX/aTarHSOtiHVlV9jI4kRYyKUlKoPNskAO9ijpy+dI0fd501Oo6IiIiIiIjNWnUgivWHL+BoZ+bdR6obHUeKIBWlpFBxcbTjheb/jJZK1WgpERERERGRLBJT0hj5y34Aet9XgUBvN4MTSVGkopQUOt1DAyjp5sipiwks2XPO6DgiIiL5YsKECVSoUAFnZ2caNGjAhg0bbtl20aJFtG3bllKlSuHh4UFoaCjLly/P1GbGjBmYTKYsX4mJiXn9UkREJB9M3XiC8EsJ+Lg78dIDQUbHkSJKRSkpdFwd7f8ZLbX6KGkWq8GJRERE8taCBQt45ZVXGDZsGLt27aJ58+Z06NCB8PDwm7Zfv349bdu2ZenSpezYsYNWrVrx6KOPsmvXrkztPDw8iIiIyPTl7KxlwkVECrrzVxL5as1RAIY+FEwxJ3uDE0lRpaKUFErdQwMo7urA8eh4ftZoKRERKeTGjRtHnz596Nu3L9WqVWP8+PH4+/szceLEm7YfP348b775Jo0aNaJy5cp88MEHVK5cmZ9//jlTO5PJROnSpTN9iYhIwfe/3w6QkJxG/fLFeayun9FxpAhTUUoKpWJO/4yW+mL1EY2WEhGRQis5OZkdO3bQrl27TNvbtWvHpk2bsnUMi8XC1atXKVmyZKbtcXFxBAQEUK5cOR555JEsI6lERKTg+ePkJX7cfQ6TCUZ0rInJZDI6khRhGqMnhdbzoQF8s/44xy7E8+ufEXSsU9boSCIiIrkuOjqatLQ0fH19M2339fXl/Pnz2TrGJ598Qnx8PF26dMnYFhwczIwZM6hVqxaxsbF89tlnNGvWjD179lC5cuWbHicpKYmkpKSMx7GxsQCkpKSQkpJyty/tjq4fMy+OLfdGfWPb1D+2K6/7Js1i5b2f/gLgqfp+BPu66t/BXdD3TvZl9z1SUUoKLXdnB/rcV4FxYYf5YtURHqlVBrNZVwFERKRw+veVbqvVmq2r3/PmzWP48OH89NNP+Pj4ZGwPCQkhJCQk43GzZs2oX78+X3zxBZ9//vlNjzVmzBhGjBiRZfuKFStwdXXN7ku5a2FhYXl2bLk36hvbpv6xXXnVN5siTeyPsMPFzkpt0ymWLj2VJ+cp7PS9c2cJCQnZaqeilBRqPZsFMnnDcY5ExfHbX+d5uHYZoyOJiIjkKm9vb+zs7LKMioqKisoyeurfFixYQJ8+ffj+++9p06bNbduazWYaNWrEkSNHbtlm6NChDBkyJONxbGws/v7+tGvXDg8Pj2y8mruTkpJCWFgYbdu2xcHBIdePLzmnvrFt6h/blZd9c+VaCsPHbwRSeLVdME83DcjV4xcF+t7Jvuujpe/E8KLUhAkT+Oijj4iIiKBGjRqMHz+e5s2b37J9UlISI0eO5Ntvv+X8+fOUK1eOYcOG0bt3byB9ieMPPviAo0ePkpKSQuXKlXnttdfo3r17xjGGDx+e5Sre3Qxxl4LDw9mB3s0q8NmqI3y+6ggdapbWaCkRESlUHB0dadCgAWFhYTz++OMZ28PCwujUqdMt95s3bx69e/dm3rx5PPzww3c8j9VqZffu3dSqVeuWbZycnHBycsqy3cHBIU8/vOf18SXn1De2Tf1ju/Kib7787TCXE1II8ilGr/sq4mCnKaZzSt87d5bd98fQotT15YsnTJhAs2bN+Prrr+nQoQP79++nfPnyN92nS5cuREZGMnXqVIKCgoiKiiI1NTXj+ZIlSzJs2DCCg4NxdHTkl19+oVevXvj4+NC+ffuMdjVq1GDlypUZj+3s7PLuhYqhejerwLSNJzgUeZUV+8/zYE2NlhIRkcJlyJAhdO/enYYNGxIaGso333xDeHg4AwYMANJHMJ09e5ZZs2YB6QWp559/ns8++4yQkJCMC3MuLi54enoCMGLECEJCQqhcuTKxsbF8/vnn7N69m6+++sqYFykiIjl2OPIqs7ek36r3/qPVVZASm2FoUerG5YshfXni5cuXM3HiRMaMGZOl/bJly1i3bh3Hjx/PWB0mMDAwU5uWLVtmevzyyy8zc+ZMNm7cmKkoZW9vr2WNiwhPVwd6NQvk89VH+WzVUdpV12gpEREpXJ5++mkuXrzIyJEjiYiIoGbNmixdupSAgPRbMyIiIggPD89o//XXX5OamsqgQYMYNGhQxvYePXowY8YMAGJiYujXrx/nz5/H09OTevXqsX79eho3bpyvr01ERO6N1WplxM/7SLNYaV/Dl+aVSxkdSSSDYUWp68sXv/XWW5m232754iVLltCwYUPGjh3L7NmzcXNzo2PHjowaNQoXF5cs7a1WK6tXr+bQoUN8+OGHmZ47cuQIZcuWxcnJiSZNmvDBBx9QsWLFW+bNz9VkNKN/7uvexJ+pv5/gQEQsy/48R9vqPnfe6RbUP7ZN/WO71De2Tf2Tfbb6Hg0cOJCBAwfe9Lnrhabr1q5de8fjffrpp3z66ae5kExERIy0fN95fj96EUd7M+88XN3oOCKZGFaUysnyxcePH2fjxo04OzuzePFioqOjGThwIJcuXWLatGkZ7a5cuYKfnx9JSUnY2dkxYcIE2rZtm/F8kyZNmDVrFlWqVCEyMpLRo0fTtGlT9u3bh5eX103PbcRqMprRP3c19TYTdtbMB0t2kXwijWwsSHRb6h/bpv6xXeob26b+ubPsriYjIiJitMSUNEb9cgCA/vdXxL9k3q2EKpIThk90fjfLF1ssFkwmE3PmzMmY72DcuHF07tyZr776KmO0lLu7O7t37yYuLo5Vq1YxZMgQKlasmHFrX4cOHTKOWatWLUJDQ6lUqRIzZ87MtGLMjfJzNRnN6J83QhOS+f2TDZyJT8O5UkNaB+dstJT6x7apf2yX+sa2qX+yL7uryYiIiBjtm/XHORtzjTKezrzYspLRcUSyMKwolZPli8uUKYOfn19GQQqgWrVqWK1Wzpw5Q+XKlYH0JYuDgoIAqFu3LgcOHGDMmDFZ5pu6zs3NjVq1at12iWMjVpPRjP65y8fTgedDA5m07hhfrT1B+5plb1kAzQ71j21T/9gu9Y1tU//cmd4fEREpCM7GXGPC2qMAvP1QNVwdDR+TIpKFYVPu37h88Y3CwsJo2rTpTfdp1qwZ586dIy4uLmPb4cOHMZvNlCtX7pbnslqtmeaD+rekpCQOHDhAmTJala2we6F5BVwc7Pjz7BXWHrpgdBwREREREZE88cHSAySmWGhcoSSP1NbfumKbDF0HcsiQIUyZMoVp06Zx4MABXn311SzLFz///PMZ7bt164aXlxe9evVi//79rF+/njfeeIPevXtn3Lo3ZswYwsLCOH78OAcPHmTcuHHMmjWL5557LuM4r7/+OuvWrePEiRNs3bqVzp07ExsbS48ePfL3DZB851XMie6h6SsRjV91BKvVanAiERERERGR3LX52EV+3RuB2QTvP1r9nu4QEclLho7fu9vli4sVK0ZYWBiDBw+mYcOGeHl50aVLF0aPHp3RJj4+noEDB3LmzBlcXFwIDg7m22+/5emnn85oc+bMGbp27Up0dDSlSpUiJCSELVu2ZJxXCrcXmldk1uaT7Dkdw/oj0bSooiVRRUTEeImJiTg7OxsdQ0RECrjUNAsjft4HQLcm5alR1vMOe4gYJ0dFqdOnT2MymTJumdu2bRtz586levXq9OvX766OdTfLFwMEBwffdmWg0aNHZypS3cz8+fPvKqMULqXcnXiuSQBTNp7gs5WHub+yt64ciIiIISwWC//3f//HpEmTiIyM5PDhw1SsWJF3332XwMBA+vTpY3REEREpYOZtC+fg+at4ujjwWtuqRscRua0c3b7XrVs31qxZA8D58+dp27Yt27Zt4+2332bkyJG5GlAkL/RrUREnezM7w2PYeDTa6DgiIlJEjR49mhkzZjB27FgcHR0ztteqVYspU6YYmExERAqiy/HJfLziMACvtatCCTfHO+whYqwcFaX++usvGjduDMB3331HzZo12bRpE3Pnzr3p6CYRW+Pj7ky3JuUB+Gyl5pYSERFjzJo1i2+++YZnn30WOzu7jO21a9fm4MGDBiYTEZGC6JOwQ1y5lkJwaXe6NS5vdByRO8pRUSolJQUnJycAVq5cSceOHYH0W+siIiJyL51IHhrQohKO9mb+OHWZzccuGh1HRESKoLNnzxIUFJRlu8ViISUlxYBEIiJSUO0/F8vcrelzMr//aA3s7Qxd10wkW3L0r7RGjRpMmjSJDRs2EBYWxoMPPgjAuXPn8PLyytWAInnF18OZro38gfSV+ERERPJbjRo12LBhQ5bt33//PfXq1TMgkYiIFERWq5XhP+/DYoWHa5UhtJL+LpeCIUcTnX/44Yc8/vjjfPTRR/To0YM6deoAsGTJkozb+kQKggEtKzFv22m2nbjEluMXCamoH94iIpJ/3n//fbp3787Zs2exWCwsWrSIQ4cOMWvWLH755Rej44mISAHxy94Itp24hLODmaEPBRsdRyTbclSUatmyJdHR0cTGxlKiRImM7f369cPV1TXXwonktTKeLjzdyJ/ZW07x2cojhPRTUUpERPLPo48+yoIFC/jggw8wmUy899571K9fn59//pm2bdsaHU9ERAqAhORUPlh6AIAXWwRRroT+JpeCI0dFqWvXrmG1WjMKUqdOnWLx4sVUq1aN9u3b52pAkbw2oGUl5m8PZ/Pxi2w7cYnGFUoaHUlERIqQ9u3b6/OTiIjk2KS1x4i4kohfcRf6t6hodByRu5KjOaU6derErFmzAIiJiaFJkyZ88sknPPbYY0ycODFXA4rkNb/iLjzVMH1uqc81t5SIiOSj7du3s3Xr1izbt27dyh9//GFAIhERKUhOX0pg0vrjALzzcDWcHezusIeIbclRUWrnzp00b94cgB9++AFfX19OnTrFrFmz+Pzzz3M1oEh+eLFFJezNJjYejWbHqUtGxxERkSJi0KBBnD59Osv2s2fPMmjQIAMSiYhIQTL61/0kp1poWsmLB2uWNjqOyF3LUVEqISEBd3d3AFasWMETTzyB2WwmJCSEU6dO5WpAkfzgX9KVzg3KAfDZqqMGpxERkaJi//791K9fP8v2evXqsX//fgMSiYhIQbHxSDTL90ViZzbx/qM1MJlMRkcSuWs5KkoFBQXx448/cvr0aZYvX067du0AiIqKwsPDI1cDiuSXQa2CsDObWH/4ArvCLxsdR0REigAnJyciIyOzbI+IiMDePkdTf0oBZLVaWbHvPP/77SAno+ONjiOFXGxiCm8u+ou9l1TAKMhS0iyM+HkfAN1DAqha2t3gRCI5k6Oi1Hvvvcfrr79OYGAgjRs3JjQ0FEgfNVWvXr1cDSiSX/xLuvJEPT8APtPcUiIikg/atm3L0KFDuXLlSsa2mJgY3n77ba2+VwRYrVbWHoqi45e/02/2DiatO0abcet4e/GfRMYmGh1PCqnZm0+xeNc5Zhw2s/WEpq0oqGZvPsWRqDhKuDrwapsqRscRybEcFaU6d+5MeHg4f/zxB8uXL8/Y3rp1az799NNcCyeS3156IH201NpDF9hzOsboOCIiUsh98sknnD59moCAAFq1akWrVq2oUKEC58+f55NPPjE6nuShzccu8tSkzfScvp0/z17B1dGOhgElSLVYmbs1nBYfrWHMbweISUg2OqoUIlarlYU7zgCQZjUxeP4ewi8mGJxK7lZ0XBKfrjwMwBvtg/F0dTA4kUjO5XhceOnSpSldujRnzpzBZDLh5+dH48aNczObSL4L8HLjsbp+LNx5hs9XHWFqz0ZGRxIRkULMz8+PvXv3MmfOHPbs2YOLiwu9evWia9euODjoj4zCaGf4ZT5ZcYjfj14EwMnezPOhAQxoUQmvYk5sPX6RscsPsePUZb5ed5y5W8MZ0KISvZoF4uqoWzrl3uwMj+F4dDwuDma8HdM4HZ9Cn5nbWTiwKR7O+plTUHy8/BBXE1OpUdaDpxv5Gx1H5J7k6DebxWJh9OjRfPLJJ8TFxQHg7u7Oa6+9xrBhwzCbczQAS8QmvPRAEIt3nWHVwSj+OnuFmn6eRkcSEZFCzM3NjX79+hkdQ/LYX2evMC7sMKsPRgHgYGeia+PyDGoVhK+Hc0a7JhW9+GFAKKsPRvHR8kMcPH+Vj5YfYvrvJ/lP6yCeaVQeR3t91pac+eHvUVIP1vClnt1pvjrixpGoOP4zbxdTezTCzqx5pmzdn2eusOCP9FVbh3esoT6TAi9HRalhw4YxdepU/ve//9GsWTOsViu///47w4cPJzExkf/7v//L7Zwi+aaCtxud6vqxeNdZPlt1hMnPNzQ6koiIFCJLliyhQ4cOODg4sGTJktu27dixYz6lkrxyOPIqn4Yd5re/zgNgZzbRuX45BrcOolwJ15vuYzKZaF3Nl5ZVffh5zzk+CTvE6UvXeO+nfUzecJwhbavQsY6f/hiVu5KYksYve84B8EQ9Py4dPM3Xz9bjmSnbWHvoAh8sPcC7j1Q3OKXcjtVq5f0lf2G1Qqe6ZWkUWNLoSCL3LEdFqZkzZzJlypRMH5Tq1KmDn58fAwcOVFFKCrxBrYL4cfdZwvZHsu/cFWqU1WgpERHJHY899hjnz5/Hx8eHxx577JbtTCYTaWlp+RdMctXJ6HjGrzzMT3vOYbWCyQSd6pTl5TZVqODtlq1j2JlNPFbPj4dqlWHB9nA+W3WU05eu8eqCPUxae5w32leldTUfLQMv2bJifyRXk1LxK+5C48ASLDsINcp6MK5LXQbO2cnUjSeo7FOMZxqXNzqq3MKPu8+yMzwGV0c7hnaoZnQckVyRo7G/ly5dIjg4OMv24OBgLl3SCg5S8AX5FOPR2mUB+GLVUYPTiIhIYWKxWPDx8cn4/1t9qSBVMJ2NucZbC/fSetw6ftydXpDqULM0y1+5n/HP1Mt2QepGjvZmuocGsv7NlrzRviruzvYcirxK31l/8OTETWw5fjEPXokUNtdv3Xuyvh/mG0bZPVSrDEPapq/e9s6Pf7H5mP492aK4pFTGLD0IpF9AL+3pfIc9RAqGHBWl6tSpw5dffpll+5dffknt2rXvOZSILRj8QBAmEyzbd54DEbFGxxERkULGYrEwbdo0HnnkEWrWrEmtWrXo1KkTs2bNwmq1Gh1P7lJUbCLv//QXrT5ay/ztp0mzWGlVtRQ/v3QfE59rQBVf93s+h6ujPYNaBbHhzVYMaFEJZwczO8NjeOabLfSYto2/zl7JhVcihdH5K4lsPHIBgCcblMvy/OAHgni0TllSLVZenLODUxfj8zui3MFXa44SdTWJ8iVd6XNfBaPjiOSaHN2+N3bsWB5++GFWrlxJaGgoJpOJTZs2cfr0aZYuXZrbGUUMUdnXnYdqleHXvRF8sfoIE55tYHQkEREpJKxWKx07dmTp0qXUqVOHWrVqYbVaOXDgAD179mTRokX8+OOPRseUbLgUn8ykdceYuekkSakWAJpW8uK1dlVoEJA3870Ud3XkrQ7B9GoWyOerjrBg+2nWHb7AusMXeKR2GV5rVzVHI7Kk8Fq06wwWKzQOLEmAlxspKSmZnjeZTHzUuTbhF+PZc+YKfWb+wSKtyGczTkTHM3XDCQDefaQ6zg52BicSyT05GinVokULDh8+zOOPP05MTAyXLl3iiSeeYN++fUyfPj23M4oY5j8PVAZg6Z/nOXT+qsFpRESksJgxYwbr169n1apV7Nq1i3nz5jF//nz27NnDypUrWb16NbNmzTI6ptzGlWspfLLiEM0/XM0364+TlGqhQUAJ5vZtwtwXQvKsIHUjXw9n/u/xWqwc0oJOddOnHfhlbwRtxq1j6KK9RFy5lucZxPZZrVYW/n3rXuebjJK6ztnBjsnPN6S0hzNHo+IYPHcXqWmW/IoptzH6l/0kp1m4v0op2lTzMTqOSK7K8XqyZcuW5f/+7/9YuHAhixYtYvTo0Vy+fJmZM2fmZj4RQ1Ut7c5DtUoD8MXqIwanERGRwmLevHm8/fbbtGrVKstzDzzwAG+99RZz5swxIJncSXxSKl+tOUrzD1fzxeqjxCenUdPPg+m9GvHDgFCaBnnne6ZAbzc+e6YeS//TnAeCfUizWJm37TQtP1rLB0sPcDk+Od8zie3YfTqGYxficXYw0+Hvz7W34uPhzJQeDXF2MLPu8AU++HsOIzHOmkNRrDoYhb3ZxHuPVNfCBlLo5LgoJVJUDP57tNSvf0ZwJFKjpURE5N7t3buXBx988JbPd+jQgT179uRjIrmTxJQ0pmw4TvOxa/ho+SFiE1Op4luMSc814OeX7qNVVeNXwate1oNpPRvxXf9QGgaUICnVwjfrj3P/2DV8seoI8UmphuYTY1yf4LxDzTK4Z+N2vJp+nnzapS4A034/wbxt4XkZT24jOdXCqJ/3A9CzaSBBPsUMTiSS+1SUErmDamU8aF/DF6sVvlyjlfhEROTeXbp0CV9f31s+7+vry+XLl/MxkdxKUmoaszef5P6xaxj96wEuxSdTwduNz56py28v38+DNUsbXoz6t8YVSvL9gFCm9WxIcGl3rial8knYYVp8tIYZv58gKVUrOxYViSlp/LznHHD7W/f+rUOtMrz294p872pFPsPM2HSC49HxeBdz5D9tKhsdRyRPqCglkg3/aZ3+S+DnPec4fkGrkYiIyL1JS0vD3v7W683Y2dmRmqpRLUZKTbPw3fbTPPDxOt79aR9RV5PwK+7C2CdrE/bq/XSq64ed2baKUTcymUw8EOzL0v8057Nn6lK+pCvRcckM/3k/D3y8joU7zpBm0SqPhd3KA5HEJqZS1tOZ0Iped7XvSw8E0fGGFflORuszcH66cDWJz1elXxB/88FgTTovhdZdrb73xBNP3Pb5mJiYe8kiYrNqlPWkTTVfVh6IZOK647RyNTqRiIgUZFarlZ49e+Lk5HTT55OSkvI5kVyXZrHyy95zjF95hBN//xHu4+7E4AeC6NLIHyf7grXqldlsolNdPzrULMOCP07z+aojnI25xmvf7+Hr9cd4vV1V2lb3tbnRXpI7rt+690T9cpjvsohqMpkY27k2py4lsOd0DH1mbmfxoGYqjuSTj8OOEJeUSp1ynnSun/1RbiIFzV0VpTw9Pe/4/PPPP39PgURs1cutK7PyQCRL9kZQo47RaUREpCDr0aPHHdvoM1X+slqtLN93nnFhhzkcGQdASTdHBrasxHMhAQV+CXZHezPdQwLoXL8cMzadZOLaoxyOjKPf7B3UK1+cN9sHE1rp7kbSiG2LjE1k/eELADx5F7fu3cjZwY7J3RvQ6avfOXYhnpfm7mJaj4bY2+mGm7x08ios+iv9tsvhHWvcdUFRpCC5q6LU9OnT8yqHiM2rVc6TB4J9WH0wil/CzTS/EE8FH/cCd8VURESMp89UtsNqtbL28AU+WXGIv87GAuDhbE//FpXo2TQQN6e7+rhs81wc7XixZSW6NS7P1+uPMe33E+wKj6Hr5C00r+zNm+2DqVXu9heipWBYvOssFis0DChBBW+3HB/Hx8OZyc835KlJm1l/+AL/t/QA7z9aIxeTyo0sFisLT6T/ffFk/XLUK1/C4EQieatw/ZYVyWMvt67M6oNR7Llk5sHPf8dkgrKeLgR4uRLg5Uagl2vG/wd4ueLqqG8xERERW7XpWDSfrDjMjlPpk8q7OdrR574K9GleEU+Xwn2LkqerA28+GEzPpoF8sfoo87aFs+FINBuObOThWmUY0q4KlUpppa+Cymq1svDvW/fuZoLzW6np58mnT9dhwLc7mf77SSr7uNOtSfl7Pq5ktWDHGcLjTbg52fHfB6saHUckz+kvZpG7UMe/OK+2DuK7zUeISbUnPjmNszHXOBtzjU03WZXEx93pXwUrNwK93Cjv5VroP+yKiIjYqh2nLvHJisMZv7udHcz0CA2kf4tKlHRzNDhd/vLxcGbUYzXp27wCn4Yd5qc95/j1zwiW7TvPUw3K8XKbypTxdDE6ptylvWeucCQqDid7Mw/VLpMrx3ywZhleb1eFj1cc5r2f/iLQ25Wmlbxz5diSXkic/vtJRv96AIBBLSvi4+FscCqRvKeilMhdGtiyIoEJB+nQoR2xyVZOXYznZHQCpy7Gc+pSAicvpv9/TEIKUVeTiLqaxPaTWZf1LuHqkKlYdWPxqqSboyYcFRERyWV/nb3CJysOseZQ+jw7DnYmujUuz6BWQUX+j78ALzfGP1OP/i0q8fHyQ6w6GMX87adZtOssPUIDeLFlUJEr2BVk1yc4f7Bm6VydmHxQqyCORMXx0+5zDJyzkx8HNiPwHm4NlHSpaRZG/Lyf2VtOARDqY6FXaIDBqUTyh4pSIjlkMpnwLuaIdzEnGgSUzPJ8TEIypy4mcOpSAqei4zOKVacuJXDhahKXE1K4nBDD7tMxWfZ1d7InwNuVgJLpxarAv4tWgd5u+Lg7qWAlIiJyFw6dv8qnYYdZtu88AHZmE081KMdLDwRRroSW1L1RtTIeTO3ZiD9OXmLsskNsO3mJyRtOMG/baV5oXpE+zStQrJDNs1XYJKWmsWRP+iTZuXHr3o1MJhMfPlmbUxcT2P33inyLBjbTHQD3IDYxhZfm7mL94QuYTPDf9lUoHbNfk8lLkaHfKCJ5pLirI8VdHanjXzzLc/FJqekFq4vpxarwS/+Mtjp3JZGrSan8dTY2Y8LVGzk7mP8pVnm7Ub7kP0WrssVdsCtCq3OkWawkpqRxLSWNa8np/01OTsFiNTqZiIjYghPR8YxfeZgle85htYLJBI/V9ePl1pU1uuMOGgaWZEH/ENYevsBHyw6xPyKWT1ceZtbmkwxqFcSzIeW12IuNWnUgiivXUijj6Zwnt9c5O9jxzfMN6PRl+op8g+dpRb6cOn0pgT4zt3M4Mg4XBzvGP1OXB6p4sXTpfqOjieQbFaVEDODmZE/1sh5UL+uR5bnElDTOXE7gZHQCJy/G/zPa6mI8Zy5fIzHFwqHIqxyKvJplXwc7E/4lXQkoecM8Vt5uBJR0pVwJVxzt8+/DgtVqJSnVQmJKGgnJmQtH15LTt934XPr/p3It2cK1lNSMtlnaJaeR8PcxklItNz13RXc7mj+QgreDrtqJiBRFZy4n8MWqo/yw8wxpf1+peKhWaV5tU4XKvu4Gpys4TCYTrar60KJyKX75M4JxKw5x8mICI3/Zz9SNJ3ilTWWeqF+uSF0QKwiu37r3eD2/POsbH/fMK/KN/vUAwztqRb67sTP8Mv1m/UF0XDK+Hk5M7dGImn6epKSkGB1NJF+pKCViY5wd7AjycSfIJ+uH5pQ0C2cvX+PkxXjCLyVkjK46eTGe05eukZxm4fiFeI5fiAcuZNrXbAK/Ei4Zo6puHG3lZG/+pwCUfEMR6YZC0vXiUOaCUurf7Sxcu/7/yf/sk58jlpwdzLg42BGfnMbxqxa6TdnO7L5N8C3ic4SIiBQlkbGJfLUmfSW5lLT0X0Ktg314tW0Vavp5Gpyu4DKbTXSsU5YONUvz/R9n+GzVYc7GXOONH/by9frjvN6uKu1r+Gp6ARsQFZvIusPpnwGfzOVb9/7txhX5Zmw6SWXfYjzbRPMgZcfPe87x2vd7SE61UL2MB1N7NtSCAlJkqSglUoA42JkJ9Ha76S0HaRYr52MTM89fdfGf0VbXUtI4fekapy9dY8OR/M3taGdOLxo52uHqaI+zgx2ujna4ONhl+n8Xx7+//t7m7HDD/9+wPaPt3/91trfD/PeVwD9PX+LZbzZxOCqOJyduYnafJlTQLRoiIoXaxfhkpv5+hFmbT2WMor0vyJsh7apQv3wJg9MVHg52Zro1Kc8T9f2YuekkE9Ye42hUHAO+3UEd/+L8t31VmgZpNTYj/bj7LGkWK/XLF6dSqWJ5fr4Ha5bhjfZV+Wj5Id7/aR8VvN20It9tWK1Wvlx9lE/CDgPQppovnz1TFzfN0yZFmP71ixQSdmYTfsVd8CvuQtOgzM9ZrVYuXE1KXx0w+p9iVfjfjy3W9BFaLo5mXB3s/y4AmXF1tM9SALrZfzMKTI5Zi0bODnY45OMcA8Gl3XmlZhozT3lw6lICnSduYmbvxrpCLiJSCF25lsKv4WaGjttAQnIaAA0DSvBau6qEVvIyOF3h5exgR/8WlXimcXkmrz/O1I0n2HM6hm5TtnJfkDdvtK9K9dK6IJTfrFYrC3ecBaBzA/98O+/AlpU4EnmVH3ef48Vvd/LjoGa6IHgTSalpDF34J4t2pfdR3/sqMPSharr9VYo8FaVEigCTyYSPhzM+Hs40Csy6UmBh4+UM819oRN/Zu9h3LpZnvtnCN8830JU7EZFCZNrGE3y68jBXE81AGrX8PHmtXRVaVCml28jyiaeLA6+3r8rzTQP4avVR5m4LZ+PRaDYejaZ9dR8aOBqdsGj562wshyKv4mhv5uHaZfLtvCaTif89WZuTN6zIt1gr8mVyKT6ZAbN3sO3kJezMJkZ2qqFbHUX+piUSRKRQ8i7mxPx+IYRULElcUio9p21n2V8RRscSEckzEyZMoEKFCjg7O9OgQQM2bNhwy7aLFi2ibdu2lCpVCg8PD0JDQ1m+fHmWdgsXLqR69eo4OTlRvXp1Fi9enJcv4a5cS0njamIqZVysTOhalyUvNaNlVR8VpAzg4+7MiE41Wf1aS56o54fJBMv3R/HJXjvOXL5mdLwi44cdpwFoX6N0vheErq/IV9bTmeMX4nlp7k5S026+IE1RczQqjscn/M62k5dwd7ZnRq9GKkiJ3EBFKREptNydHZjRqzHta/iSnGZh4JydzNsWbnQsEZFct2DBAl555RWGDRvGrl27aN68OR06dCA8/OY/89avX0/btm1ZunQpO3bsoFWrVjz66KPs2rUro83mzZt5+umn6d69O3v27KF79+506dKFrVu35tfLuq2eTQMZ36U2b9ZJo211FaNsgX9JV8Y9XZdlL99P9TLuJFlMTPv9pNGxioSk1DR+2nMOgM55PMH5rfi4OzO5R0NcHOzYcCSa0b8eMCSHLdl0NJonJvzOqYsJ+Jd0YdGLTWleuZTRsURsiopSIlKoOTvYMeHZBnRt7I/FCkMX/clXa45itebj0oAiInls3Lhx9OnTh759+1KtWjXGjx+Pv78/EydOvGn78ePH8+abb9KoUSMqV67MBx98QOXKlfn5558ztWnbti1Dhw4lODiYoUOH0rp1a8aPH59Pr+r23JzsebhWaTQdi+2pWtqdtx6sAsD3O89yMS7J4ESF35qDUcQkpODr4cR9Bk42X6OsJ58+XReAGZtO8u2WU4ZlMdqC7eE8P20bsYmp1C9fnMUDm1HZN+vq2iJFnYpSIlLo2ZlNfPB4LQa1qgTAR8sPMfKX/VgsKkyJSMGXnJzMjh07aNeuXabt7dq1Y9OmTdk6hsVi4erVq5Qs+c+8g5s3b85yzPbt22f7mFK0hVQoib+blcQUCzM3nTQ6TqH3w44zADxer5zhE2c/WLM0b7SvCsD7S/ax6Wi0oXnym8ViZcxvB/jvwj9JtVjpWKcsc18IwbuYk9HRRGySJjoXkSLBZDLxRvtgSro5MeqX/Uz//SSX45P56Kk6+bo6oIhIbouOjiYtLQ1fX99M2319fTl//ny2jvHJJ58QHx9Ply5dMradP3/+ro+ZlJREUtI/o2JiY2MBSElJISUlJVtZ7sb1Y+bFseXepKam0sbPwvTDdszcfJLeTctr2fs8Eh2XxJpDFwDoVNs3W98Pef2980Kz8hyKiGXJ3ghenLODH/o3IdCr8K/Il5Ccyus//EXYgSgABreqyOBWlTBhISUle3Ns6eeabVP/ZF923yP9ZhCRIqXPfRUo6ebAG9/v5cfd54i5lsKEZ+vj6qgfhyJSsP17TiWr1ZqteZbmzZvH8OHD+emnn/Dx8bmnY44ZM4YRI0Zk2b5ixQpcXV3vmCWnwsLC8uzYknO1S0IpZysXrqUyfHYYrcpqhHJeWHPORJrFjoBiVg7/sZ7Dd7FvXn7v3O8Ce4vZcTIulWcnbeTVWmm4FuKPW1eSYfJBO07Hm7AzWelWyUJQ4mF+++1ueuQf+rlm29Q/d5aQkJCtdoX4x4KIyM09Xq8cxV0ceXHODtYeusBzU7YyrWcjirtq7WoRKXi8vb2xs7PLMoIpKioqy0inf1uwYAF9+vTh+++/p02bNpmeK1269F0fc+jQoQwZMiTjcWxsLP7+/rRr1w4PD4/svqRsS0lJISwsjLZt2+LgoOXnbcn1vhncJpj3fjnElsuu/F/P5jjaa3RybrJarUz4ajMQR+9W1XmosX+29suv752mLZJ48uutRFxJ5JdLvkzpXg/7QjhC/UDEVfp9u5Pz8UmUcHVgYre6NAgokaNj6eeabVP/ZN/10dJ3oqKUiBRJrYJ9mNM3hN4ztrMzPIanJm1mVp/GlPF0MTqaiMhdcXR0pEGDBoSFhfH4449nbA8LC6NTp0633G/evHn07t2befPm8fDDD2d5PjQ0lLCwMF599dWMbStWrKBp06a3PKaTkxNOTlnnTXFwcMjTD+95fXzJuSca+PPlupOcj01i6b4onmqYvaKJZM9fZ69wKDIOR3szj9Xzv+vvg7z+3ilb0oEpPRrSeeJmfj92kf8tP8KITjXz7HxGWHUgksHzdpGQnEalUm5M79mY8l73PjJUP9dsm/rnzrL7/hS+MrWISDY1CCjB9wNC8fVw4khUHJ0nbubYhTijY4mI3LUhQ4YwZcoUpk2bxoEDB3j11VcJDw9nwIABQPoIpueffz6j/bx583j++ef55JNPCAkJ4fz585w/f54rV65ktHn55ZdZsWIFH374IQcPHuTDDz9k5cqVvPLKK/n98qQAc7I30+e+CgB8vf64FhnJZdcnOG9b3RdPV9v8A7lGWU/GP1MXgJmbTzG7kKzIZ7VambrxBC/M+oOE5DSaBXmxaGCzXClIiRQlKkqJSJFWxdedhS82paK3G2djrvHUpM3sPRNjdCwRkbvy9NNPM378eEaOHEndunVZv349S5cuJSAgAICIiAjCw8Mz2n/99dekpqYyaNAgypQpk/H18ssvZ7Rp2rQp8+fPZ/r06dSuXZsZM2awYMECmjRpku+vTwq2bk3K4+5sz9GoOFYeiDQ6TqGRnGrhp91nAejcoJzBaW6vfY1/VuQbvmQfvxfwFflS0yy8+9NfjPplPxYrdG1cnhm9GuPpYpuFQRFbpqKUiBR55Uq48v2AUGr5eXIpPpmu32xh45GC/WFJRIqegQMHcvLkSZKSktixYwf3339/xnMzZsxg7dq1GY/Xrl2L1WrN8jVjxoxMx+zcuTMHDx4kOTmZAwcO8MQTT+TTq5HCxN3Zge4h6QXSieuOYbVqtFRuWHMoissJKfi4O9E8yNvoOHc0sGUlHq/nR5rFysA5OzleQEenxyam0GvGdr7dEo7JBO88XI0PHq+p1ZxFckjfOSIigFcxJ+b1C6FZkBfxyWn0mrGNX/dGGB1LRESkUOjVrAKO9mZ2hcew7cQlo+MUCtdv3Xu8nl+BmDzcZDIx5ola1CtfnCvXUug78w+uJGRvyXhbcfpSAk9O2MSGI9G4ONjx9XMN6Nu8YrZWOhWRmzP8p9eECROoUKECzs7ONGjQgA0bNty2fVJSEsOGDSMgIAAnJycqVarEtGnTMp5ftGgRDRs2pHjx4ri5uVG3bl1mz559z+cVkcKvmJM903o24qFapUlJs/LSvJ2FZt4DERERI5Vyd+Kpv28xm7TumMFpCr7ouCTWHIwC4Ekbv3XvRs4OdnzTvSFlPZ05Hh3PoLk7SU2zGB0rW3aGX+bxCb9zJCoOXw8nvh8QSrsapY2OJVLgGVqUWrBgAa+88grDhg1j165dNG/enA4dOmSa8+DfunTpwqpVq5g6dSqHDh1i3rx5BAcHZzxfsmRJhg0bxubNm9m7dy+9evWiV69eLF++/J7OKyJFg5O9HV90rc+zTcpjtcK7P/7FZyuP6FYDERGRe9Tv/oqYTbDm0AUORGRvqXC5uZ92nyPVYqVOOU+q+LobHeeulHJ3YkqPRrg62rHxaDSjftlvdKQ7WrLnHM98s4XouGRqlPXgp0H3UdPP0+hYIoWCoUWpcePG0adPH/r27Uu1atUYP348/v7+TJw48abtly1bxrp161i6dClt2rQhMDCQxo0bZ1qauGXLljz++ONUq1aNSpUq8fLLL1O7dm02btyY4/OKSNFiZzYx+rGa/Kd1ZQA+XXmY4Uv2acUgERGRexDg5cZDtcoA8LVGS92ThX/fumfrE5zfSvWyHnz6dF1MJttekc9qtfL5qiP8Z94uklMttKnmy3f9Qynt6Wx0NJFCw96oEycnJ7Njxw7eeuutTNvbtWvHpk2bbrrPkiVLaNiwIWPHjmX27Nm4ubnRsWNHRo0ahYuLS5b2VquV1atXc+jQIT788MMcnxfSbxtMSkrKeBwbm351JyUlhZSU3L0X+vrxcvu4kjvUP7YtN/tncMsKeDrbMerXg8zcfIrouCTGPlETR3vD73wukPS9Y9vUP9mn90gk5wa0qMQveyP4eW8Er7Wrin9JV6MjFTj7zl1hf0QsjnZmHq1T1ug4OXZ9Rb6xyw4xfMk+Kni5cV9l25mwPSk1jbcW/sniXekrHL7QvAJvdaiGnVnzR4nkJsOKUtHR0aSlpeHr65tpu6+vL+fPn7/pPsePH2fjxo04OzuzePFioqOjGThwIJcuXco0r9SVK1fw8/MjKSkJOzs7JkyYQNu2bXN8XoAxY8YwYsSILNtXrFiBq2ve/DINCwvLk+NK7lD/2Lbc6h9v4PnKJr49aubXP89zLPwcvatacLLLlcMXSfresW3qnztLSEgwOoJIgVXTz5Pmlb3ZcCSayRuOM7JTTaMjFTgLd6QXSdpU96G4q6PBae7Niy0qcTQyjkW7zjJwzg5+HNSMiqWKGR2LS/HJ9J/9B9tPXsbObGJUp5p0a1Le6FgihZJhRanr/r1SgdVqveXqBRaLBZPJxJw5c/D0TL+Hd9y4cXTu3JmvvvoqY7SUu7s7u3fvJi4ujlWrVjFkyBAqVqxIy5Ytc3RegKFDhzJkyJCMx7Gxsfj7+9OuXTs8PDzu6jXfSUpKCmFhYbRt2xYHB4dcPbbcO/WPbcuL/nkIaHkkmkHzdnPwCnx7rjiTn6tPSbeC/UEwv+l7x7apf7Lv+mhpEcmZF1tUYsORaBZsP81/WlfGu5iT0ZEKjJQ0Cz/tTi9KFdRb925kMpn44IlanLwYz87wGPrO/IPFA5vh6Wrc76GjUXH0mbmdUxcTcHe2Z8Kz9WleuZRheUQKO8OKUt7e3tjZ2WUZnRQVFZVlFNN1ZcqUwc/PL6MgBVCtWjWsVitnzpyhcuX0+V/MZjNBQUEA1K1blwMHDjBmzBhatmyZo/MCODk54eSU9Remg4NDnn14z8tjy71T/9i23O6fB6qXYe4LzvSasZ29Z2LpNnU7s/o0wa941luH5fb0vWPb1D93pvdH5N6EVvKiTjlP9py5wsxNJ3mtXVWjIxUYaw9d4GJ8Mt7FnLi/kBRKnB3s+Lp7Qx776veMFfmm92qEg13+T5ew6Wg0A77dQWxiKv4lXZjWoxGVC9hE8iIFjWETozg6OtKgQYMstwmEhYVlmrj8Rs2aNePcuXPExcVlbDt8+DBms5ly5W59pcBqtWbMB5WT84qIANQrX4IfBoRSxtOZYxfi6TxxE0ejrhodS0REpEAxmUwMaFEJgFmbTxGXlGpwooLjhx2nAXiivh/2BhRt8kopdycmP9/Q0BX55m8L5/lp24hNTKVBQAl+HNhMBSmRfGDoT7IhQ4YwZcoUpk2bxoEDB3j11VcJDw9nwIABQPotc88//3xG+27duuHl5UWvXr3Yv38/69ev54033qB3794Zt+6NGTOGsLAwjh8/zsGDBxk3bhyzZs3iueeey/Z5RURuJcjHnR9ebErFUm5EXEmk86TN7Aq/bHQsERGRAqVdjdJU9HbjyrUU5m8LNzpOgXAxLolVB6IAeLJ+wb9179+ql/Vg/N8r8s3afIrZm0/my3ktFitjlh7grUV/kmqx0qluWeb0bYKXbisVyReGFqWefvppxo8fz8iRI6lbty7r169n6dKlBAQEABAREUF4+D+/pIoVK0ZYWBgxMTE0bNiQZ599lkcffZTPP/88o018fDwDBw6kRo0aNG3alB9++IFvv/2Wvn37Zvu8IiK341fchR8GNKVOOU9iElLoNnkr6w5fMDqWiIhIgWFnNtHv/ooATNlwguRUi8GJbN+SPedItVip5edJ1dKFcwRPuxqlebN9MADDf97PxiPReXq+hORUXpyzg6/XHwfglTaVGf90XZwdtKKNSH4xfKLzgQMHMnDgwJs+N2PGjCzbgoODb7sy0OjRoxk9evQ9nVdE5E5Kujky94UQBny7gw1Houk7czufdKlLxwK8NLOIiEh+ery+H+PCDnM+NpEfd5+lS0N/oyPZtIU7zwCFY4Lz2xnQoiJHoq6yaGfersgXGZtI35l/8OfZKzjamfnoqdp0quuX6+cRkdsrPDcii4jkMzcne6b2aMQjtcuQkmbl5fm7mLnppNGxRERECgQnezv63FcBgK/XHcNisRqcyHYdiIjlr7OxONiZCv0FMJPJxJgnatEgoASxian0mfkHVxJScvUc+85d4bGvfufPs1f+vtDYRAUpEYOoKCUicg8c7c189kw9ng8NwGqF95fsY1zYYaxWfbAWERG5k25NyuPubM+xC/GEHYg0Oo7NWrgjfZRU62BfSrg5Gpwm7znZ2/F19wb4FXfhRHQ8A+fuICUtd27xXLk/kqcmbSbiSiJBPsX4cWAzGgaWzJVji8jdU1FKROQe2ZlNjOhYg1faVAbg81VHePenv0jTFV8REZHbcnd2oHtI+ryuE9ce00Wdm0hJs/Dj7rNA4b9170bexZyY0iN9Rb7fj15k5M/3tiKf1Wpl6sYTvDD7DxKS07gvyJuFLzalvJdrLiUWkZxQUUpEJBeYTCZeaVOFUZ1qYDLBt1vC+c+8XSSlphkdTURExKb1alYBR3szu0/HsPXEJaPj2Jz1hy8QHZeMdzFHWlQtZXScfFWtjAefPVMPkwlmbznFrByuyJeaZuHdn/5i1C/7sVqha+PyTO/VCE8Xh9wNLCJ3TUUpEZFc1D00kC+61sPBzsSvf0bQe8Z24pJSjY4lIiJis0q5O/HU3yOAJq07ZnAa2/PD37fuPVbXDwe7ovfnW9vqvvz3wfQV+Ub8vJ8NR+5uxePYxBR6zdjOt1vCMZngnYer8cHjNYvkeylii/SdKCKSyx6pXZbpPRtnDDfvNnkLF+OSjI4lIiJis/rdXxGzCdYeusD+c7FGx7EZl+OTWfn3XFtPFqFb9/6t//0VeaK+H2kWKwPn7OTYhbhs7Xf6UgJPTtjEhiPRuDjY8U33hvRtXhGTyZTHiUUku1SUEhHJA/dV9mbeCyGUcHVg75krPDVpM2cuJxgdS0RExCYFeLnxUK0yAHy9XqOlrluy5xwpaVZqlPWgWhkPo+MY5vqKfA0DSnA1MZW+M/8gJiH5tvvsOHWZx776nSNRcfh6OPH9gFDaVvfNp8Qikl0qSomI5JE6/sX5fkBTyno6czw6nicnbuJw5FWjY4mIiNikAS0qAfDznnOcvqQLOQALd6bfuleUJji/FSd7OybduCLfnJ23XJFvyZ5zdJ28hYvxydQo68FPg+6jpp9nPicWkexQUUpEJA8F+RRj4cCmBPkUIzI2iacmbWbHqctGxxIREbE5Nf08aV7ZG4sVJm84bnQcwx06f5W9Z67gYGeiU10/o+PYhOsr8rk52rHp2EWGL9mXacVGq9XK56uO8J95u0hOtdC2ui/f9Q+ltKezgalF5HZUlBIRyWNlPF34vn8o9coX58q1FJ6dsoU1h6KMjiUiImJzXmyZPlpqwfbTRBfx+Rivj5JqVdWHkm6OBqexHTeuyDdnazizNp8CICk1jSHf7WFc2GEgfZ6ySc81wM3J3si4InIHKkqJiOSDEm6OzOnbhPurlCIxxcILM//gx11njY4lIiJiU0IrelGnnCdJqRZmbjppdBzDpKZZWLQz/XOCbt3Lqk11X976e0W+kb/s56fdZ3luylYW7zqLndnEB4/X4u2HqmFn1oTmIrZORSkRkXzi6mjPlOcb0qluWVItVl5ZsJtpG08YHUtERMRmmEymjNFSMzedJC4p1eBExthwJJrouCS83BxpFexjdByb1O/+ijxZvxxpFisvz9/N9pOXcXe2Z2avxnRrUt7oeCKSTSpKiYjkI0d7M592qUvPpoFA+tW9j5YfzDQfgoiISFHWtnppKnq7EZuYyvxt4UbHMcQPO9Jv3etU1w8HO/3JdjMmk4kPnqhJw4ASAPiXdGHxwKbcV9nb4GQicjf0E05EJJ+ZzSbef7Q6r7erAsBXa47x9uI/SbOoMCUiImJnNtG/RUUApmw4QXLqzVdYK6xiEpIJ2x8JwJMNNMH57TjZ2zG9VyM+eaoOSwbdR5CPu9GRROQuqSglImIAk8nESw9U5v8er4nJBPO2nWbQnJ0kpqQZHU1ERMRwj9Xzw9fDifOxify4u2jNwfjznnMkp1moVsaDGmU9jY5j89ydHXiyQTlKaDJ4kQJJRSkREQM92ySAr7rVx9HOzLJ95+k1fTtXE1OMjiUiImIoJ3s7+txXAYBJ645hKUKjiX/QBOciUoSoKCUiYrCHapVhRq9GuDnasfn4RZ75ZgsXrhbtZbBFRES6Ni6Pu7M9xy/EE3Yg0ug4+eJI5FX2nI7B3myiU92yRscREclzKkqJiNiApkHezO8XipebI/vOxfLUpE2cvpRgdCwRERHDuDs78HxoAAAT1x4rEouC/LAzfYLzllV98C7mZHAaEZG8p6KUiIiNqFXOk+8HhOJX3IWTFxN4cuImDp6PNTqWiIiIYXo2rYCjvZndp2PYeuKS0XHyVGqahcW6dU9EihgVpUREbEjFUsVY+GJTqvgWI+pqEl0mbeaPk4X7Q7iIiMitlHJ3okvD9ALNxLXHDE6TtzYejSbqahIlXB14INjH6DgiIvlCRSkRERtT2tOZ7/qH0iCgBLGJqTw3dSvrDl8wOpaIiIgh+jWvhNkE6w5fYP+5wjuC+Icd6bfudarrh6O9/kwTkaJBP+1ERGxQcVdHvu3ThBZVSpGYYqHvzO0s/TPC6FhygzSLlaTUNKNjiIgUeuW9XHm4dvqk35PWFc7RUlcSUlixP30yd926JyJFiYpSIiI2ysXRjsnPN+ThWmVISbPy0tydLNgebnQsAQ6ej6XFR2to+dFajl2IMzqOiEih1//+igD8svcc4RcL30IgP+89R3KqheDS7tQo62F0HBGRfKOilIiIDXO0N/N513o808gfixX+u/BPJq8/bnSsIm3T0WiemriZM5evEXElkeenbiPiyjWjY4mIFGo1/Ty5v0opLFaYvKHw/R5c+Peqe50blMNkMhmcRkQk/6goJSJi4+zMJsY8UYt+f18l/r+lB/hkxaEisTS2rflp91l6TN/G1aRUGgeWpGIpN87GXKP71G1cjk82Op6ISKE2oEX678Hv/jhNdFySwWlyz9GoOHaFx2BnNtGprp/RcURE8pWKUiIiBYDJZGJoh2DeaF8VgC9WH2X4kn1YLCpM5Qer1crEtcd4ef5uUtKsPFy7DLP7NmZ2nyaU8XTmaFQcPWdsJz4p1eioIiKFVmhFL+r4Fycp1cKM308aHSfXXB8l1bJKKUq5OxmcRkQkf6koJSJSQJhMJga1CmJUpxoAzNx8ite+30NKmsXgZIVbmsXKez/t48NlBwF4oXkFvnimHk72dvgVd2F2n8aUcHVgz+kYBny7Q5Ofi4jkEZPJxIt/j5aatfkkcYXgQkCaxcqiG27dExEpalSUEhEpYLqHBjL+6brYmU0s3nWWF7/dSWKKCiF54VpyGgO+3cHsLacwmeC9R6oz7OHqmM3/zPcR5OPO9F6NcXW0Y8ORaIYs2EOaRrCJiOSJdtVLU7GUG7GJqczbWvAX//j9aDSRsUkUd3XggWo+RscREcl3KkqJiBRAj9Xz4+vnGuBob2blgUh6Td9eKK4Y25JL8cl0m7KFsP2RONqbmdCtPr3vq3DTtnX9i/NN94Y42Jn49c8I3v3pL835JSKSB8xmU8ZKfFM2Hi/wo1N/2JE+SqpTnbI42dsZnEZEJP+pKCUiUkC1qe7LjF6NcHO0Y/Pxizw7Zasm284lpy7G8+TETewKj8HTxYE5fZvQoVaZ2+5zX2VvPnumHiYTzN0azicrDudTWhGRouWxen74ejgRGZvET7vOGR0nx65cS2H5vvMAdG7gb3AaERFjqCglIlKANa3kzdwXQij+95xGT3+zmcjYRKNjFWh7TsfwxIRNnIiOx6+4CwtfbEqjwJLZ2vehWmX4v8dqAfDlmqNMKYTLlouIGM3J3o4+f49cnbT+WIFd9OPXvREkpVqo4luMmn4eRscRETGEilIiIgVcHf/ifNc/FF8PJw5HxtF50ibCLyYYHatAWn0wkme+2cLF+GRq+nmweFBTgnyK3dUxujUpn7FK4uhfD7Dw71szREQk93RtXB4PZ3uOX4hnxf5Io+PkyMIbJjg3mUx3aC0iUjipKCUiUghU8XXnhwFNKV/SldOXrtF50iYOR141OlaBMndrOH1n/sG1lDTur1KK+f1C8XF3ztGxBrasRN+/r+K/uXAvKwvoH0wiIrbK3dmB7qEBAExcd6zAzeN3/EIcO05dxs5s4rG6fkbHERExjIpSIiKFhH9JV34YEEpVX3eiribR5evN7D4dY3Qsm2e1WvlkxSHeXvwnFis81aAcU3s0pJiTfY6PaTKZePuhajxZvxxpFiuD5u5k6/GLuZhaJKsJEyZQoUIFnJ2dadCgARs2bLhl24iICLp160bVqlUxm8288sorWdrMmDEDk8mU5SsxUbcIi23o2bQCTvZm9pyOYcvxS0bHuSvXR0m1qFIKH4+cXQARESkMVJQSESlEfDycWdA/hLr+xYlJSOHZyVvYdDTa6Fg2KyXNwuvf7+WL1UcBeLl1ZcZ2ro2D3b3/ejSbTXz4ZC3aVPMlKdVC35l/sO/clXs+rsjNLFiwgFdeeYVhw4axa9cumjdvTocOHQgPD79p+6SkJEqVKsWwYcOoU6fOLY/r4eFBREREpi9nZ/0BLbahlLsTTzUsB8CkdccMTpN9aRYri3aeBeDJ+uUMTiMiYiwVpURECpniro7M6duEZkFexCen0XPGdlb8vbqP/ONqYgq9Z2xn4c4z2P1dQHq1bZVcndfD3s7Ml93q0bhCSa4mpdJj2jZORMfn2vFFrhs3bhx9+vShb9++VKtWjfHjx+Pv78/EiRNv2j4wMJDPPvuM559/Hk9Pz1se12QyUbp06UxfIrakX/NKmE2w7vCFAlP433zsIhFXEvF0caB1NR+j44iIGEpFKRGRQsjNyZ6pPRrRrrovyakWXpyzk8W7NOH2dZGxiTz99RY2HInG1dGOKT0a8nSj8nlyLmeH9ONXL+NBdFwy3adu1QqJkquSk5PZsWMH7dq1y7S9Xbt2bNq06Z6OHRcXR0BAAOXKleORRx5h165d93Q8kdxW3suVh2uXBeDrdQVjxdMfdpwGoGOdsjg72BmcRkTEWDmfMENERGyas4MdE56tz5sL97Jo51leXbCH2Gup9GgaaHQ0Qx2JvErP6ds5G3MN72JOTO/ZiFrlbj1SJDd4ODsws3djnpq0iZMXE3h+6jYW9A+huKtjnp5Xiobo6GjS0tLw9fXNtN3X15fz53M+SjI4OJgZM2ZQq1YtYmNj+eyzz2jWrBl79uyhcuXKN90nKSmJpKSkjMexsbEApKSkkJKSkuMst3L9mHlxbLk3+dk3fZuV5+c95/hl7zlefqAi5Uu65vk5c+pqYgrL/h69/Fid0ob929X3ju1S39g29U/2Zfc9UlFKRKQQs7cz83HnOng4OzBj00neX7KP2GspvPRAUJFcfnrr8Yu8MOsPYhNTqejtxszejfHPpz9eSrk7MbtPEzpP2sShyKv0nrGdb/s2wdVRv4old/z7e9pqtd7T93lISAghISEZj5s1a0b9+vX54osv+Pzzz2+6z5gxYxgxYkSW7StWrMDVNe++18LCwvLs2HJv8qtvgj3NHLxi5v2563mqoiVfzpkTmyNNJKbY4eti5fSe3zmz19g8+t6xXeob26b+ubOEhIRstdMnYRGRQs5sNvH+o9XxcHHg81VH+CTsMLGJKbz9ULUiVZj6dW8Ery7YTXKahQYBJZjyfENKuOXvSCX/kq7M6t2ELl9vZmd4DC9+u5PJzzfE0V5300vOeXt7Y2dnl2VUVFRUVJbRU/fCbDb/f3t3HhdVvfcB/HNmBoYd2RdlE0XEjU0JFO2WUrZcN9JcUHOLsFK59ty82i19Kl9Z1+w+XVFyS3PLzJveLEVv4i6KDqLikqgogiyKbDosc54/0CnCBZCZc4DP+/XiFZw553e+hx/kl+/8FvTs2RMXLlx46DmzZs1CfHy8/uvi4mJ4eHggKioKNjY2TRbLfZWVlUhKSsKAAQNgYmLS5O1T4xm7bxw638SYFcdwtFCFz8ZHwsFKbfB7NsaaZSkAijA20g8vRvpIFgd/d+SLfSNv7J/6uz9a+nFYlCIiagUEQUD8AD/Ympvgf/9zBl/tu4TiO1X4eGg3KBUtvzC1bF8mPvwxAwDwXBcXfPFqkGTreHRytcaK8T0xZtkRJJ/Px182peGLEYFQtIJ+IMMwNTVFSEgIkpKSMGTIEP3xpKQkDBo0qMnuI4oiNBoNunXr9tBz1Go11Oq6xQATExODJu+Gbp8az1h907ujM3p4tEHa1SJ8k5KNmc91Mvg9G+pyQRmOXSmCQgCiQz1l8TPL3x35Yt/IG/vn8er7/eFbs0RErcjEPj5YEN0dCgHYeOwq3lp/HNqqaqnDMphqnYi5207rC1LjI7yxeHSI5AvLhnjZYUlMCEyUAralXccH205DFEVJY6LmLT4+HsuWLcOKFSuQkZGBGTNmICsrC7GxsQBqRjCNHTu21jUajQYajQalpaXIz8+HRqPBmTNn9K/PnTsXO3bsQGZmJjQaDSZOnAiNRqNvk0hOBEHAG/18AQCrD11GqbZK4ojq2ny8ZsORvn5OcLExkzgaIiJ54EgpIqJWZnioB6zVKry94QS2p+ei5O4xLI0JaXFrG92trMaMjRr8dKpmStPfXvDH5Mj2spmy2M/PCf8YHohpG05g9aErsLMwxYwBflKHRc3UiBEjUFhYiHnz5iEnJwddu3bF9u3b4eXlBQDIyclBVlZWrWuCgoL0n6empmLdunXw8vLC5cuXAQBFRUWYMmUKcnNzYWtri6CgIOzduxe9evUy2nMRNURUgAvaO1kiM78M649kYXLf9lKHpKfTifj+eDYAYFhwO4mjISKSD46UIiJqhQZ2c8PycT1hbqLEvgsFiFmegtt3Ws4uIkXlFYhZfgQ/ncqFqVKBf44MwpS+vrIpSN335x7umDeoKwDgi90XsOrAJYkjouYsLi4Oly9fhlarRWpqKvr27at/bdWqVdizZ0+t80VRrPNxvyAFAJ9//jmuXLkCrVaLvLw87NixA+Hh4UZ6GqKGUygExPatGS21bH+mrEYCH84sRHbRHVibqTAgoOnWeiMiau5YlCIiaqX6+jnhm0m9YGOmQuqVWxiZeBj5JdrHXyhzV2+WY1jCQRy9fAvWZip8PaEX/tzDXeqwHirmKS/E3xsh9cG2M/j3iWyJIyIiar4GBbnDxUaNG8Va/HDiutTh6H2XWjN178893CWfQk5EJCcsShERtWIhXvbY+Ho4HK3UOJNTjOFLDyG76I7UYTXaqezbGJpwEBfzy+Bma4bNb0Qg3NdB6rAe661nOmB8hDcAYOamNPxyNk/agIiImim1SolJfWqm7S3ZexHVOunX6yu5W4ntp3IAANEhnLpHRPR7LEoREbVynd1ssCk2HG3bmONSQRmiEw7i17xSqcNqsOTz+Rix9BDyS7Twd7XGlrje8HOxljqsehEEAX9/KQCDA91RpRPxxtpUHLt8U+qwiIiapZFhnrAxUyEzvwxJZ3KlDgc/pefibqUO7Z0sEejRRupwiIhkhUUpIiKCj6MlvnsjHL5Olsi5fRcjlh7CqezbUodVb98eu4oJq46irKIavTs44NvYcLjaNq+djRQKAZ++0gPP+DvjbqUOE1YdRUZOsdRhERE1O1ZqFcaGewMAEpIzJd/d9Lt7u+5Fh7ST3dqGRERSY1GKiIgAAG625vj29XB0bWuDwrIKjEw8jJRL8h6tI4oivth1Af/z3UlU60QMCWqLleN7wcbMROrQGsVEqcC/RgUj1MsOxXerMHZFCrIKy6UOi4io2Rnf2xtqlQJpV4twKLNQsjiuFJYh5dJNKARgaBCn7hER/ZHkRanFixfDx8cHZmZmCAkJwb59+x55vlarxezZs+Hl5QW1Wg1fX1+sWLFC//pXX32FyMhI2NnZwc7ODv3790dKSkqtNj744AMIglDrw9XV1SDPR0TUnDhYqbFu8lPo5WOPEm0Vxq44gl/OyXN9o6pqHWZ9n47Pd50HAMQ97YuFw3vAVCX5P21PxNxUieXje8Lf1Rr5JVqMWX4EeSV3pQ6LiKhZcbRSY3ioBwBgSXKmZHFsPl6zeUWfjk7NbgQvEZExSJq5b9y4EdOnT8fs2bNx4sQJREZGYuDAgcjKynroNcOHD8fu3buxfPlynDt3DuvXr4e/v7/+9T179mDkyJH45ZdfcOjQIXh6eiIqKgrZ2bV3M+rSpQtycnL0H+np6QZ7TiKi5sTGzASrJ/TSTyOb/PUxbEuTzw5GAFCmrcLk1cew4ehVKATgfwd3xf88799ipkXYmtf0gae9BbJulmPs8hTcvlMpdVhERM3K5Mj2UAjA3vP5kkxJ1+lEfH9v6t6w4LZGvz8RUXMgaVFq4cKFmDhxIiZNmoTOnTtj0aJF8PDwQEJCwgPP//nnn5GcnIzt27ejf//+8Pb2Rq9evRAREaE/Z+3atYiLi0NgYCD8/f3x1VdfQafTYffu3bXaUqlUcHV11X84OTkZ9FmJiJoTMxMllsaE4OUeNQtvv73hBNanPPwNA2PKL9Hi1cTD+OVcPsxMFFgaE4qYp7ykDqvJOduY4ZuJYXCyVuNsbgkmfX0UdyqqpQ6LiKjZ8HSwwEvd3QEAS/caf7TUkUs3ce3WHVirVXiuC2dlEBE9iGRFqYqKCqSmpiIqKqrW8aioKBw8ePCB12zduhWhoaFYsGAB2rZtCz8/P8ycORN37jx8+/Ly8nJUVlbC3t6+1vELFy7A3d0dPj4+ePXVV5GZKd2wXiIiOTJRKrBoRCBGhXlCFIFZ36djafJFSWO6mF+KoQkHkJ59G/aWplg/+SkMCHCRNCZD8nSwwOoJvWBtpsLRy7cwdd1xVFbrpA6LiKjZeL1fewDAjyev40phmVHv/V1qzSipl3q4w8xEadR7ExE1FyqpblxQUIDq6mq4uNT+Y8LFxQW5uQ/eujUzMxP79++HmZkZtmzZgoKCAsTFxeHmzZu11pX6vXfffRdt27ZF//799cfCwsKwevVq+Pn54caNG/jwww8RERGB06dPw8HB4YHtaLVaaLVa/dfFxTU7IlVWVqKysmmnVNxvr6nbpabB/pE39k/T++DFTrA2VWLpvkuY/9NZ3CrTIr5/hwZPlXvSvjmeVYTXvzmBojuV8LQ3x/KxwfB2sGzxfd3B0RyJY4Lw2tep+O/ZPMz8VoMFQ7tCoWjaqYr83ak/fo+Imo8u7rbo5+eE5PP5+GpfJj4c3M0o9y3TVuGnUzkAanbdIyKiB5OsKHXfH/+oEUXxoX/o6HQ6CIKAtWvXwtbWFkDNFMDo6Gj861//grm5ea3zFyxYgPXr12PPnj0wM/ttYcGBAwfqP+/WrRvCw8Ph6+uLr7/+GvHx8Q+89/z58zF37tw6x3fu3AkLC4v6PWwDJSUlGaRdahrsH3lj/zStAAAvewrYlqXEkr2XkH7uIqJ9dGhMXaQxfZNWKGDNBQUqRQFeViImty/BmSPJONPw2zdbY30FLDunwA9pObh1IxtDvXUwxBJa/N15vPJy7ohI1JzE9vNF8vl8fHvsGqY96wcna7XB7/nTqVyUV1TDx9ESwZ5tDH4/IqLmSrKilKOjI5RKZZ1RUXl5eXVGT93n5uaGtm3b6gtSANC5c2eIoohr166hY8eO+uOfffYZPv74Y+zatQvdu3d/ZCyWlpbo1q0bLly48NBzZs2aVatgVVxcDA8PD0RFRcHGxuaR7TdUZWUlkpKSMGDAAJiYNM9tzVsy9o+8sX8M5wUAPY9exfvbMnDghgL2Lu74ZGhXmCjrNxO8sX2z5nAWVh4+C1EEnunkhM+Hd4OFqeTvqRjdCwD80nIw87t07M1VILiLH6Y+3b7J2ufvTv3dHy1NRM3DU+3tEejRBpqrRVh18BLeec7/8Rc9oe9SrwKoGSXVUjbhICIyBMmyelNTU4SEhCApKQlDhgzRH09KSsKgQYMeeE3v3r2xadMmlJaWwsrKCgBw/vx5KBQKtGv327DYTz/9FB9++CF27NiB0NDQx8ai1WqRkZGByMjIh56jVquhVtd9V8XExMRgybsh26Ynx/6RN/aPYYyNaA9bCzX+8m0atp3MxZ1KHb4cFdygtTLq2zc6nYhPfj6rX5x2VJgn5v25C1T1LIK1RNGhnijRVmPutjNYtPtXOFibNfki7/zdeTx+f4iaF0EQENvPF7HfpGL1oSuI7ecLazPD/R5fvVmOw5k3IQjAkCDuukdE9CiSZvbx8fFYtmwZVqxYgYyMDMyYMQNZWVmIjY0FUDM6aezYsfrzR40aBQcHB7z22ms4c+YM9u7di3feeQcTJkzQT91bsGAB5syZgxUrVsDb2xu5ubnIzc1FaWmpvp2ZM2ciOTkZly5dwpEjRxAdHY3i4mKMGzfOuN8AIqJmaFBgWySODYFapcCujDyMW5GCkrtNu8aOtqoa0zZq9AWpd57rhI8Gd23VBan7Xuvtg7efrRkZ/PcfTmFb2nWJIyIikr+oABe0d7JEyd0qg+8mu/l4zQLnfTo4wr2N+WPOJiJq3STN7keMGIFFixZh3rx5CAwMxN69e7F9+3Z4edW865uTk4OsrN/+0bCyskJSUhKKiooQGhqK0aNH4+WXX8Y///lP/TmLFy9GRUUFoqOj4ebmpv/47LPP9Odcu3YNI0eORKdOnTB06FCYmpri8OHD+vsSEdGjPePvgtUTesFKrcKRSzcxetkR3CyraJK2b9+pxLgVKdiWdh0qhYCFw3tg6p8avrB6Szajf0fEPOUFUQTiv9Ug+Xy+1CEREcmaQiEgtq8vAGDZvkvQVlUb5D46nagvSg0L5gLnRESPI/miHHFxcYiLi3vga6tWrapzzN/f/5GLsF6+fPmx99ywYUN9wyMioocIa++A9ZOfwriVKTh57TaGLz2EbyaGwdXW7PEXP8T1ojsYvzIF52+UwkqtwpIxIejT0bEJo24ZBEHA3D93QdGdSmxLu47YNalYOzkMwZ52UodGRCRbg4LcsTDpPHKL7+LfJ7Ixoqdnk9/j6OWbuHrzDqzUKjzXxbXJ2yciamk4D4KIiBqtWztbfPt6OFxtzPBrXimilxzElcKyRrWVkVOMIYsP4PyNUrjYqPHt6+EsSD2CQiHgH6/0QF8/J9yprMZrK4/i/I0SqcMiIpIttUqJiX18AABLkzNRrROb/B7fpdaMknqpuxvMTeu/3iIRUWvFohQRET2RDs5W2BQbDm8HC1y7dQfRSw7hbG7Ddic78GsBXllyCDeKtejobIXv43ojwL1pdzZtiUxVCiwZE4wgzza4facSMcuP4OrNcqnDIiKSrZFhnrAxUyGzoAxJZ3Iff0EDlFdUYXt6DoCaXfeIiOjxWJQiIqIn5mFvgW9jw+Hvao38Ei1GLD2M41m36nXtlhPXMH5lCkq1VQjzscd3sRFoy4Vh683CVIWV43vCz8UKN4q1iFl+BPklWqnDIiKSJSu1CmPDvQEACXsuQhSbbrTUz6dyUVZRDW8HC4R4cTo1EVF9sChFRERNwtnaDBunhCP43qidMcuOYP+FgoeeL4oi/vXLr5ixMQ2V1SJe6u6G1RN7wdbCcNt0t1RtLEyxekIY2tmZ43JhOcavTEFxE++ISETUUozv7Q21SoG0a7dxKLOwydq9P3VvWHA7bs5BRFRPLEoREVGTsbUwwTeTwhDZ0RHlFdWYsOoofj5Vd3pEVbUOc/59Cp/uOAcAmNK3Pf75ahDUKq6/0ViutmZYMzEMjlamOH29GJO+Poa7lYbZXYqIqDlztFJjeKgHgJrRUk3h2q1yHLxYCEEAhnLqHhFRvbEoRURETcrCVIVl40LxfBdXVFTrELc2Vf/uMQDcqahG7DepWHskC4IAvP9yAP72QmcoFHxX+Un5OFpi1Wu9YK1WIeXSTby57gSqqnVSh0VEJDtT+raHUiFg34UCnMq+/cTtfX88GwAQ4evAKehERA3AohQRETU5tUqJL0cF4ZWQdtCJwMxNafj60BWUVgIxK49hV0YeTFUKLB4VjNd6+0gdbovSta0tlo0LhVqlwK6MG/jr5nToDLDDFBFRc+Zhb4EXu7kBAJYkP9loKVEUsfl4zZsvXOCciKhhWJQiIiKDUCkV+GRYd/322x9uP4ePNUqkXbuNNhYmWDcpDAPv/UFATSusvQP+NSoYSoWAzcev4ePtGU26mC8RUUsQ288XALA9PQdXCssa3c6xK7dwpbAclqZKPNfFtanCIyJqFViUIiIig1EoBMx5sTPiB/gBAMqqBLRrY4bvYiMQ6m0vcXQtW/8AFywY1h0AsGz/JSQ84UgAIqKWJsDdBv38nKATgcS9mY1u57tjNaOkXuzuBgtTVVOFR0TUKrAoRUREBiUIAt5+tiMWDO2Kp5x1+HZKGDo4W0kdVqswLKQd3nspAACw4OdzWJ+SJXFERETy8sbTNaOlNqVeQ36JtsHX36moxo/pOQCA6BCPJo2NiKg1YFGKiIiMYkiQO0b66uBkrZY6lFZlYh8fvPmnDgCA2VvSsf3eH09ERASE+dgj0KMNKqp0WHngUoOv33E6F6XaKnjaW6Cnt50BIiQiatlYlCIiImrh/hLlh1FhntCJwPQNGuy/UCB1SEREsiAIgn601JrDV1Byt7JB19/fXXZYcDsIAneRJSJqKBaliIiIWjhBEPC/g7rixW5uqKjWYcqaY9BcLZI6LCIiWRjQ2QW+TpYouVuFdUfqP805u+gODlysKfIPDW5rqPCIiFo0FqWIiIhaAaVCwMIRPRDZ0RHlFdV4bWUKfs0rkTosIiLJKRQCXr+3E9/y/Zegraqu13Vbjl+DKALh7R3gYW9hyBCJiFosFqWIiIhaCbVKiSVjQtDDow1ulVciZnkKsovuSB0WEZHkBge2hauNGfJKtNhyPPux54uiiM33zosOaWfo8IiIWiwWpYiIiFoRS7UKq8b3RAdnK+TcvouY5UdQWNrwHaeIiFoSU5UCkyJ9AACJezNRrRMfef7xrFu4VFAGC1Mlnu/qaowQiYhaJBaliIiIWhk7S1OsmdgLbduYIzO/DONXHkWptkrqsIiIJPVqL0/YmKmQWVCGnadzH3nu/QXOX+jmBku1yhjhERG1SCxKERERtUJutuZYM7EXHCxNkZ59G3HrNKjUSR0VEZF0rNQqjIvwBgAsSb4IUXzwaKm7ldX4T1oOAE7dIyJ6UixKERERtVLtnayw6rVesFKrcCjzJuZrlIjfdBJLky9i34V8FHBaHxG1MuMivKFWKZB27TYOXSx84Dk7TueiRFsFD3tz9PK2N3KEREQtC8eaEhERtWLd2tniq7GhmPj1URRqq7HtZC62nfxt2oqLjRqd3WwQ4GaDAPea/3o7WEKhECSMmojIMByt1BjR0wOrD11BQvJFRHRwrHPO/al7Q4Pa8f+FRERPiEUpIiKiVi7c1wF7/hKJZVt2w6pdJ5y7UYYzOcW4XFiGG8Va3CjOx55z+frzLUyV8He1vlekskWAuw06uVjD3FQp4VMQETWNyZHtsfZIFvZdKMCp7Nvo2tZW/1rO7TvY/2sBAGBYMKfuERE9KRaliIiICHYWpuhiJ+KFfu1hYmICACjTVuFsbgnO5BTjzPVinMkpxtmcYpRXVON4VhGOZxXpr1cINdMB/ziqyslaLdETERE1joe9BV7q7oYfNNexJPkivhwVrH/t++PZEEUgzMceng4WEkZJRNQysChFRERED2SpViHEyw4hXnb6Y1XVOlwuLMPpe0WqjJwSnLl+GwWlFfg1rxS/5pViW9p1/flO1upaRaoA95rpf0pOeSEiGXu9ry9+0FzH9vQcXCksg5eDJURRxObjNVP3uMA5EVHTYFGKiIiI6k2lVKCDszU6OFtjUGBb/fG8krv60VT3/3upoAz5JVokl+Qj+fxv0//MTZTwd7OuNarK39UaFqZMS4hIHgLcbfB0JyfsOZePxL2Z+GhIN5y4WoTM/DKYmygxsJub1CESEbUIzP6IiIjoiTlbm8G5kxme7uSsP1ZecW/6n35UVTHO5pTgTmU1TmQV4cTvpv8JAuDjaFlnVJWztZkET0NEBMT288Wec/nYlHoN0/p31C9wPrCbK6zU/DOKiKgp8P+mREREZBAWpioEe9oh2PO36X/VOhGXC8vqjKrKL9EiM78Mmfll+M/JHP35jlZqBLjboLObNQLcbNDF3QY+jlac/kdEBhfmY48gzzY4kVWEJXsy9VOTOXWPiKjpsChFRERERqNUCPB1soKvkxVe7uGuP55fokVGTu1CVWZ+KQpKtdh7Ph97fzf9z8xEgU6utRdU93e1hiVHLhBRExIEAbH9fPH6mlSsOHAJANC2jTme8nGQODIiopaD2RsRERFJzslaDSdrJ/T1c9Ifu1NRjXM37k//u40z14txNrcE5RXVSLtahLSrRfpzBQHwdqg9/a+Luw2cbTj9j4gab0BnF/g6WeJifhkAYFhwWyg4UpOIqMmwKEVERESyZG6qRKBHGwR6tNEfq9aJuFJYVrPr371C1ZmcYtwo1uJSQRkuFZThx/Sa6X/h7R2wfspTEkVPRC2BQiHg9X6++J/vTgIAhnHqHhFRk2JRioiIiJoNpUJAeycrtHeywovdf9v9qqD03vS/361V1a2drYSRElFLMTiwLQ5nFqJtG3N4OVhKHQ4RUYuikDoAIiIioiflaKVGZEcnvN7PF1+8GoSk+H6YNdBf6rCMavHixfDx8YGZmRlCQkKwb9++h56bk5ODUaNGoVOnTlAoFJg+ffoDz9u8eTMCAgKgVqsREBCALVu2GCh6IvkyVSmwcHgg/hLVSepQiIhaHBaliIiIqEUShNaz7svGjRsxffp0zJ49GydOnEBkZCQGDhyIrKysB56v1Wrh5OSE2bNno0ePHg8859ChQxgxYgRiYmKQlpaGmJgYDB8+HEeOHDHkoxAREVErwqIUERERUTO3cOFCTJw4EZMmTULnzp2xaNEieHh4ICEh4YHne3t744svvsDYsWNha/vgaY6LFi3CgAEDMGvWLPj7+2PWrFl49tlnsWjRIgM+CREREbUmXFOKiIiIqBmrqKhAamoq3n333VrHo6KicPDgwUa3e+jQIcyYMaPWseeee+6RRSmtVgutVqv/uri4GABQWVmJysrKRsfyMPfbNETb9GTYN/LG/pEv9o28sX/qr77fIxaliIiIiJqxgoICVFdXw8XFpdZxFxcX5ObmNrrd3NzcBrc5f/58zJ07t87xnTt3wsLCotGxPE5SUpLB2qYnw76RN/aPfLFv5I3983jl5eX1Oo9FKSIiIqIW4I9raImi+MTrajW0zVmzZiE+Pl7/dXFxMTw8PBAVFQUbG5sniuVBKisrkZSUhAEDBsDExKTJ26fGY9/IG/tHvtg38sb+qb/7o6Ufh0UpIiIiombM0dERSqWyzgimvLy8OiOdGsLV1bXBbarVaqjV6jrHTUxMDJq8G7p9ajz2jbyxf+SLfSNv7J/Hq+/3hwudExERETVjpqamCAkJqTOVICkpCREREY1uNzw8vE6bO3fufKI2iYiIiH6PI6WIiIiImrn4+HjExMQgNDQU4eHhSExMRFZWFmJjYwHUTKvLzs7G6tWr9ddoNBoAQGlpKfLz86HRaGBqaoqAgAAAwLRp09C3b1988sknGDRoEH744Qfs2rUL+/fvN/rzERERUcvEohQRERFRMzdixAgUFhZi3rx5yMnJQdeuXbF9+3Z4eXkBAHJycpCVlVXrmqCgIP3nqampWLduHby8vHD58mUAQEREBDZs2IA5c+bgvffeg6+vLzZu3IiwsDCjPRcRERG1bCxKEREREbUAcXFxiIuLe+Brq1atqnNMFMXHthkdHY3o6OgnDY2IiIjogbimFBERERERERERGR1HSjXS/XcX67vNYUNUVlaivLwcxcXFXNFfhtg/8sb+kS/2jbyxf+rv/r/99Rlp1NoZMl8C+HMrZ+wbeWP/yBf7Rt7YP/VX33yJRalGKikpAQB4eHhIHAkRERFJoaSkBLa2tlKHIWvMl4iIiFq3x+VLgsi3+RpFp9Ph+vXrsLa2hiAITdp2cXExPDw8cPXqVdjY2DRp2/Tk2D/yxv6RL/aNvLF/6k8URZSUlMDd3R0KBVdCeBRD5ksAf27ljH0jb+wf+WLfyBv7p/7qmy9xpFQjKRQKtGvXzqD3sLGx4Q+6jLF/5I39I1/sG3lj/9QPR0jVjzHyJYA/t3LGvpE39o98sW/kjf1TP/XJl/j2HhERERERERERGR2LUkREREREREREZHQsSsmQWq3G+++/D7VaLXUo9ADsH3lj/8gX+0be2D/UHPHnVr7YN/LG/pEv9o28sX+aHhc6JyIiIiIiIiIio+NIKSIiIiIiIiIiMjoWpYiIiIiIiIiIyOhYlCIiIiIiIiIiIqNjUUqGFi9eDB8fH5iZmSEkJAT79u2TOiQCMH/+fPTs2RPW1tZwdnbG4MGDce7cOanDogeYP38+BEHA9OnTpQ6F7snOzsaYMWPg4OAACwsLBAYGIjU1VeqwWr2qqirMmTMHPj4+MDc3R/v27TFv3jzodDqpQyN6LOZL8sR8qflgviQ/zJfkifmSYbEoJTMbN27E9OnTMXv2bJw4cQKRkZEYOHAgsrKypA6t1UtOTsbUqVNx+PBhJCUloaqqClFRUSgrK5M6NPqdo0ePIjExEd27d5c6FLrn1q1b6N27N0xMTPDTTz/hzJkz+Mc//oE2bdpIHVqr98knn2DJkiX48ssvkZGRgQULFuDTTz/F//3f/0kdGtEjMV+SL+ZLzQPzJflhviRfzJcMi7vvyUxYWBiCg4ORkJCgP9a5c2cMHjwY8+fPlzAy+qP8/Hw4OzsjOTkZffv2lTocAlBaWorg4GAsXrwYH374IQIDA7Fo0SKpw2r13n33XRw4cICjGGTopZdegouLC5YvX64/NmzYMFhYWGDNmjUSRkb0aMyXmg/mS/LDfEmemC/JF/Mlw+JIKRmpqKhAamoqoqKiah2PiorCwYMHJYqKHub27dsAAHt7e4kjofumTp2KF198Ef3795c6FPqdrVu3IjQ0FK+88gqcnZ0RFBSEr776SuqwCECfPn2we/dunD9/HgCQlpaG/fv344UXXpA4MqKHY77UvDBfkh/mS/LEfEm+mC8ZlkrqAOg3BQUFqK6uhouLS63jLi4uyM3NlSgqehBRFBEfH48+ffqga9euUodDADZs2IDjx4/j6NGjUodCf5CZmYmEhATEx8fjb3/7G1JSUvD2229DrVZj7NixUofXqv31r3/F7du34e/vD6VSierqanz00UcYOXKk1KERPRTzpeaD+ZL8MF+SL+ZL8sV8ybBYlJIhQRBqfS2KYp1jJK0333wTJ0+exP79+6UOhQBcvXoV06ZNw86dO2FmZiZ1OPQHOp0OoaGh+PjjjwEAQUFBOH36NBISEphkSWzjxo345ptvsG7dOnTp0gUajQbTp0+Hu7s7xo0bJ3V4RI/EfEn+mC/JC/MleWO+JF/MlwyLRSkZcXR0hFKprPMuX15eXp13A0k6b731FrZu3Yq9e/eiXbt2UodDAFJTU5GXl4eQkBD9serqauzduxdffvkltFotlEqlhBG2bm5ubggICKh1rHPnzti8ebNEEdF977zzDt599128+uqrAIBu3brhypUrmD9/PpMski3mS80D8yX5Yb4kb8yX5Iv5kmFxTSkZMTU1RUhICJKSkmodT0pKQkREhERR0X2iKOLNN9/E999/j//+97/w8fGROiS659lnn0V6ejo0Go3+IzQ0FKNHj4ZGo2GCJbHevXvX2Q78/Pnz8PLykigiuq+8vBwKRe1UQKlUcotjkjXmS/LGfEm+mC/JG/Ml+WK+ZFgcKSUz8fHxiImJQWhoKMLDw5GYmIisrCzExsZKHVqrN3XqVKxbtw4//PADrK2t9e/Q2trawtzcXOLoWjdra+s6a1VYWlrCwcGBa1jIwIwZMxAREYGPP/4Yw4cPR0pKChITE5GYmCh1aK3eyy+/jI8++gienp7o0qULTpw4gYULF2LChAlSh0b0SMyX5Iv5knwxX5I35kvyxXzJsARRFEWpg6DaFi9ejAULFiAnJwddu3bF559/zi10ZeBh61SsXLkS48ePN24w9FhPP/00tziWkf/85z+YNWsWLly4AB8fH8THx2Py5MlSh9XqlZSU4L333sOWLVuQl5cHd3d3jBw5En//+99hamoqdXhEj8R8SZ6YLzUvzJfkhfmSPDFfMiwWpYiIiIiIiIiIyOi4phQRERERERERERkdi1JERERERERERGR0LEoREREREREREZHRsShFRERERERERERGx6IUEREREREREREZHYtSRERERERERERkdCxKERERERERERGR0bEoRURERERERERERseiFBGREQmCgH//+99Sh0FEREQkW8yXiFoPFqWIqNUYP348BEGo8/H8889LHRoRERGRLDBfIiJjUkkdABGRMT3//PNYuXJlrWNqtVqiaIiIiIjkh/kSERkLR0oRUauiVqvh6upa68POzg5AzVDxhIQEDBw4EObm5vDx8cGmTZtqXZ+eno5nnnkG5ubmcHBwwJQpU1BaWlrrnBUrVqBLly5Qq9Vwc3PDm2++Wev1goICDBkyBBYWFujYsSO2bt1q2IcmIiIiagDmS0RkLCxKERH9znvvvYdhw4YhLS0NY8aMwciRI5GRkQEAKC8vx/PPPw87OzscPXoUmzZtwq5du2olUQkJCZg6dSqmTJmC9PR0bN26FR06dKh1j7lz52L48OE4efIkXnjhBYwePRo3b9406nMSERERNRbzJSJqMiIRUSsxbtw4UalUipaWlrU+5s2bJ4qiKAIQY2Nja10TFhYmvvHGG6IoimJiYqJoZ2cnlpaW6l//8ccfRYVCIebm5oqiKIru7u7i7NmzHxoDAHHOnDn6r0tLS0VBEMSffvqpyZ6TiIiIqLGYLxGRMXFNKSJqVf70pz8hISGh1jF7e3v95+Hh4bVeCw8Ph0ajAQBkZGSgR48esLS01L/eu3dv6HQ6nDt3DoIg4Pr163j22WcfGUP37t31n1taWsLa2hp5eXmNfSQiIiKiJsV8iYiMhUUpImpVLC0t6wwPfxxBEAAAoijqP3/QOebm5vVqz8TEpM61Op2uQTERERERGQrzJSIyFq4pRUT0O4cPH67ztb+/PwAgICAAGo0GZWVl+tcPHDgAhUIBPz8/WFtbw9vbG7t37zZqzERERETGxHyJiJoKR0oRUaui1WqRm5tb65hKpYKjoyMAYNOmTQgNDUWfPn2wdu1apKSkYPny5QCA0aNH4/3338e4cePwwQcfID8/H2+99RZiYmLg4uICAPjggw8QGxsLZ2dnDBw4ECUlJThw4ADeeust4z4oERERUSMxXyIiY2FRiohalZ9//hlubm61jnXq1Alnz54FULPTy4YNGxAXFwdXV1esXbsWAQEBAAALCwvs2LED06ZNQ8+ePWFhYYFhw4Zh4cKF+rbGjRuHu3fv4vPPP8fMmTPh6OiI6Oho4z0gERER0RNivkRExiKIoihKHQQRkRwIgoAtW7Zg8ODBUodCREREJEvMl4ioKXFNKSIiIiIiIiIiMjoWpYiIiIiIiIiIyOg4fY+IiIiIiIiIiIyOI6WIiIiIiIiIiMjoWJQiIiIiIiIiIiKjY1GKiIiIiIiIiIiMjkUpIiIiIiIiIiIyOhaliIiIiIiIiIjI6FiUIiIiIiIiIiIio2NRioiIiIiIiIiIjI5FKSIiIiIiIiIiMjoWpYiIiIiIiIiIyOj+H6kcY7b6XR6KAAAAAElFTkSuQmCC\n"},"metadata":{}},{"name":"stdout","text":"\n✅ Final Test Dice Score: 0.4278\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# give high results Dice= 0.483, patch_size    = 480\n#'dice': 0.423 , patch_size    = 224\n# ============================================================\n# VESUVIUS INK DETECTION - V2\n# ============================================================\n# Main improvements over V1:\n#\n# 1) Spatial train/validation split -> less leakage from overlapping patches\n# 2) Positive-patch balancing\n# 3) Random coordinate jitter during training\n# 4) BCE + Soft Dice + Focal Tversky loss\n# 5) Differential LR: pretrained encoder slower, decoder faster\n# 6) Fine Dice threshold search (0.01)\n# 7) 4-way TTA at inference:\n#       original + horizontal flip + vertical flip + rot90\n# 8) Gaussian weighted sliding-window inference\n# 9) Optional AdaBN\n# 10) Conservative post-processing for thin ink\n# 11) AMP + gradient clipping + memory-safe inference\n#\n# Recommended first configuration:\n#\n# patch_size    = 480\n# train_stride  = 128\n# test_stride   = 128\n# batch_size    = 8\n# epochs        = 10\n# encoder_lr    = 5e-5\n# decoder_lr    = 2e-4\n#\n# ============================================================\n\n\n# ============================================================\n# 0. INSTALLS\n# ============================================================\n\n!pip install -q segmentation-models-pytorch==0.2.0\n!pip install -q albumentations\n!pip install -q scikit-image\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport random\nimport time\nimport json\nimport math\n\nimport numpy as np\nimport cv2\nimport tifffile\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom scipy import ndimage\n\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\n\nfrom skimage.filters import threshold_otsu\n\n\n# ============================================================\n# 2. CONFIG\n# ============================================================\n\nclass CFG:\n\n    # --------------------------------------------------------\n    # Paths\n    # --------------------------------------------------------\n\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n\n    # Train on fragments 2 and 3\n    train_frags = [\"2\", \"3\"]\n\n    # Evaluate/infer on fragment 1\n    test_frag = \"1\"\n\n\n    # --------------------------------------------------------\n    # Depth\n    # --------------------------------------------------------\n    #\n    # IMPORTANT:\n    # range(16, 38) = 22 slices\n    #\n    # Keep this for the first V2 experiment because this is\n    # the configuration used by your original code.\n    #\n\n    depth_indices = list(range(16, 38))\n    in_channels = len(depth_indices)\n\n\n    # --------------------------------------------------------\n    # PATCH / STRIDE\n    # --------------------------------------------------------\n\n    patch_size = 480\n\n    train_stride = 128\n    test_stride = 128\n\n    # Random jitter around grid coordinates during training.\n    # Helps reduce memorization of fixed patch locations.\n    train_jitter = 32\n\n\n    # --------------------------------------------------------\n    # VALIDATION\n    # --------------------------------------------------------\n\n    # We use a spatial validation split instead of random\n    # patch split.\n    #\n    # The last ~20% in Y becomes validation.\n    # A gap of patch_size is kept to reduce overlap leakage.\n\n    val_fraction = 0.20\n\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n\n    # Minimum positive ink fraction for a patch to be\n    # considered a positive training patch.\n    positive_patch_fraction = 0.001\n\n\n    # --------------------------------------------------------\n    # SAMPLING\n    # --------------------------------------------------------\n\n    # Approximate target fraction of positive patches in the\n    # training dataset.\n    target_positive_patch_ratio = 0.55\n\n    # Maximum number of repetitions of positive patches.\n    max_positive_repeat = 3\n\n\n    # --------------------------------------------------------\n    # BATCHING\n    # --------------------------------------------------------\n\n    batch_size = 8\n\n    # Inference batch can be larger because gradients are off.\n    infer_batch = 12\n\n    num_workers = 2\n\n    drop_last = True\n\n\n    # --------------------------------------------------------\n    # TRAINING\n    # --------------------------------------------------------\n\n    epochs = 10\n\n    early_stop_patience = 3\n\n    # Differential learning rates:\n    #\n    # pretrained encoder -> smaller LR\n    # decoder/head       -> larger LR\n\n    encoder_lr = 5e-5\n    decoder_lr = 2e-4\n\n    weight_decay = 1e-4\n\n    grad_clip = 5.0\n\n    accumulation_steps = 1\n\n\n    # --------------------------------------------------------\n    # MODEL\n    # --------------------------------------------------------\n\n    #encoder_name = \"resnet50\"\n   # EfficientNet (efficientnet-b0-B4)\n    encoder_name = \"vgg16\"\n    encoder_weights = \"imagenet\"\n\n    decoder_attention_type = \"scse\"\n\n\n    # --------------------------------------------------------\n    # LOSS\n    # --------------------------------------------------------\n\n    # BCE weight\n    bce_weight = 0.30\n\n    # Soft Dice weight\n    dice_weight = 0.35\n\n    # Focal Tversky weight\n    focal_tversky_weight = 0.35\n\n    # Tversky parameters.\n    #\n    # beta > alpha means FN are penalized more strongly.\n    # This is useful for thin/sparse ink where missing ink\n    # pixels can be expensive.\n\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n    focal_tversky_gamma = 0.75\n\n\n    # --------------------------------------------------------\n    # THRESHOLD\n    # --------------------------------------------------------\n\n    threshold_min = 0.10\n    threshold_max = 0.90\n    threshold_step = 0.01\n\n    # This is used only as fallback.\n    threshold = 0.50\n\n\n    # --------------------------------------------------------\n    # TTA\n    # --------------------------------------------------------\n\n    use_tta = True\n\n    # Four predictions:\n    #\n    # original\n    # horizontal flip\n    # vertical flip\n    # rotate 90\n\n    tta_modes = [\n        \"original\",\n        \"hflip\",\n        \"vflip\",\n        \"rot90\"\n    ]\n\n\n    # --------------------------------------------------------\n    # AdaBN\n    # --------------------------------------------------------\n\n    use_adabn = True\n\n    adabn_max_patches = 1200\n\n\n    # --------------------------------------------------------\n    # POST PROCESSING\n    # --------------------------------------------------------\n\n    use_postprocess = True\n\n    # Conservative closing.\n    # We intentionally do NOT use a strong 3x3 opening because\n    # ink can contain thin structures.\n\n    closing_kernel = 3\n\n    # Remove components smaller than this.\n    #\n    # Set to 0 to disable.\n    min_component_size = 8\n\n\n    # --------------------------------------------------------\n    # REPRODUCIBILITY\n    # --------------------------------------------------------\n\n    seed = 42\n\n\n    # --------------------------------------------------------\n    # OUTPUT\n    # --------------------------------------------------------\n\n    out_dir = \"/kaggle/working\"\n\n    ckpt_path = os.path.join(\n        out_dir,\n        \"vesuviusnet_v2_best.pth\"\n    )\n\n    viz_dir = os.path.join(\n        out_dir,\n        \"v2_visualizations\"\n    )\n\n    metrics_path = os.path.join(\n        out_dir,\n        \"v2_metrics_summary.json\"\n    )\n\n\n    # --------------------------------------------------------\n    # DEVICE\n    # --------------------------------------------------------\n\n    device = (\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\n\n# ============================================================\n# 3. DIRECTORIES\n# ============================================================\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\n# ============================================================\n# 4. SEED\n# ============================================================\n\ndef set_seed(seed):\n\n    random.seed(seed)\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n    # Reproducibility is useful for comparing experiments.\n    #\n    # benchmark=True can give slightly better speed but may\n    # reduce exact reproducibility.\n    torch.backends.cudnn.benchmark = True\n\n\nset_seed(CFG.seed)\n\n\n# ============================================================\n# 5. HELPERS\n# ============================================================\n\ndef cleanup_memory():\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef cfg_to_dict(cfg_cls):\n\n    d = {}\n\n    for k, v in cfg_cls.__dict__.items():\n\n        if k.startswith(\"__\"):\n            continue\n\n        if isinstance(\n            v,\n            (\n                int,\n                float,\n                str,\n                bool,\n                type(None),\n                list,\n                tuple\n            )\n        ):\n            d[k] = v\n\n    return d\n\n\n# ============================================================\n# 6. TISSUE / LABEL HELPERS\n# ============================================================\n\ndef load_tissue_mask(frag_dir):\n\n    mask_path = os.path.join(\n        frag_dir,\n        \"mask.png\"\n    )\n\n    if os.path.exists(mask_path):\n\n        mask = cv2.imread(\n            mask_path,\n            cv2.IMREAD_GRAYSCALE\n        )\n\n    else:\n\n        mid_idx = CFG.depth_indices[\n            len(CFG.depth_indices) // 2\n        ]\n\n        mid_path = os.path.join(\n            frag_dir,\n            \"surface_volume\",\n            f\"{mid_idx:02d}.tif\"\n        )\n\n        mid = tifffile.imread(mid_path)\n\n        thresh = mid.mean() * 0.15\n\n        mask = (\n            mid > thresh\n        ).astype(np.uint8) * 255\n\n    return (\n        mask > 0\n    ).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n\n    path = os.path.join(\n        frag_dir,\n        \"inklabels.png\"\n    )\n\n    if not os.path.exists(path):\n        return None\n\n    lbl = cv2.imread(\n        path,\n        cv2.IMREAD_GRAYSCALE\n    )\n\n    return (\n        lbl > 0\n    ).astype(np.uint8)\n\n\n# ============================================================\n# 7. PATCH GRID\n# ============================================================\n\ndef generate_grid_coords(\n    mask,\n    patch_size,\n    stride,\n    min_frac\n):\n\n    H, W = mask.shape\n\n    coords = []\n\n    if H < patch_size or W < patch_size:\n        return coords\n\n    for y in range(\n        0,\n        H - patch_size + 1,\n        stride\n    ):\n\n        for x in range(\n            0,\n            W - patch_size + 1,\n            stride\n        ):\n\n            tissue_frac = mask[\n                y:y + patch_size,\n                x:x + patch_size\n            ].mean()\n\n            if tissue_frac > min_frac:\n                coords.append((y, x))\n\n    return coords\n\n\n# ============================================================\n# 8. SPATIAL TRAIN / VALIDATION SPLIT\n# ============================================================\n\ndef spatial_split_coords(\n    coords,\n    image_shape,\n    patch_size,\n    val_fraction=0.20\n):\n\n    \"\"\"\n    Spatial split along Y.\n\n    Validation occupies the lower portion of the image.\n\n    We leave a gap approximately equal to patch_size so\n    training and validation patches are not heavily overlapping.\n\n    This is much safer than:\n\n        random.shuffle(all_samples)\n        first 20% -> validation\n\n    because your stride=128 and patch_size=480 create very\n    large overlap.\n    \"\"\"\n\n    H, W = image_shape\n\n    boundary = int(\n        H * (1.0 - val_fraction)\n    )\n\n    # Validation starts after boundary.\n    #\n    # Training patches must finish before boundary.\n    # This creates a spatial gap.\n\n    train_coords = []\n    val_coords = []\n\n    gap = patch_size\n\n    for y, x in coords:\n\n        patch_end_y = y + patch_size\n\n        if patch_end_y <= boundary - gap:\n\n            train_coords.append((y, x))\n\n        elif y >= boundary:\n\n            val_coords.append((y, x))\n\n    return train_coords, val_coords\n\n\n# ============================================================\n# 9. BALANCED PATCH SAMPLING\n# ============================================================\n\ndef get_patch_positive_fraction(\n    labels,\n    y,\n    x,\n    patch_size\n):\n\n    patch = labels[\n        y:y + patch_size,\n        x:x + patch_size\n    ]\n\n    return float(patch.mean())\n\n\ndef balance_positive_patches(\n    samples,\n    labels_full,\n    patch_size,\n    positive_threshold=0.001,\n    target_positive_ratio=0.55,\n    max_positive_repeat=3\n):\n\n    positive = []\n    negative = []\n\n    for fid, y, x in samples:\n\n        labels = labels_full[fid]\n\n        frac = get_patch_positive_fraction(\n            labels,\n            y,\n            x,\n            patch_size\n        )\n\n        if frac >= positive_threshold:\n            positive.append(\n                (fid, y, x)\n            )\n        else:\n            negative.append(\n                (fid, y, x)\n            )\n\n    if len(positive) == 0:\n        print(\"WARNING: No positive patches found.\")\n        return samples\n\n    if len(negative) == 0:\n        return samples\n\n    # Desired:\n    #\n    # positive / total ~= target_positive_ratio\n\n    desired_negative = int(\n        len(positive)\n        * (1.0 - target_positive_ratio)\n        / target_positive_ratio\n    )\n\n    desired_negative = max(\n        desired_negative,\n        len(positive)\n    )\n\n    if desired_negative < len(negative):\n\n        negative_selected = random.sample(\n            negative,\n            desired_negative\n        )\n\n    else:\n\n        negative_selected = negative.copy()\n\n    # Determine how many times positive samples should repeat.\n    repeat = int(\n        math.ceil(\n            (\n                target_positive_ratio\n                * len(negative_selected)\n            )\n            /\n            (\n                (1.0 - target_positive_ratio)\n                * max(len(positive), 1)\n            )\n        )\n    )\n\n    repeat = max(\n        1,\n        min(\n            repeat,\n            max_positive_repeat\n        )\n    )\n\n    positive_selected = (\n        positive * repeat\n    )\n\n    samples_balanced = (\n        positive_selected\n        + negative_selected\n    )\n\n    random.shuffle(\n        samples_balanced\n    )\n\n    positive_ratio = (\n        len(positive_selected)\n        /\n        max(len(samples_balanced), 1)\n    )\n\n    print(\n        f\"Positive patches: {len(positive)} | \"\n        f\"Negative patches: {len(negative)}\"\n    )\n\n    print(\n        f\"Balanced dataset: {len(samples_balanced)} | \"\n        f\"positive ratio={positive_ratio:.3f}\"\n    )\n\n    return samples_balanced\n\n\n# ============================================================\n# 10. DISK-BACKED VOLUME\n# ============================================================\n\nclass FragmentVolume:\n\n    def __init__(\n        self,\n        frag_dir,\n        depth_indices\n    ):\n\n        self.paths = [\n            os.path.join(\n                frag_dir,\n                \"surface_volume\",\n                f\"{i:02d}.tif\"\n            )\n            for i in depth_indices\n        ]\n\n        self.depth_indices = depth_indices\n\n        self._slices = None\n\n        self._h = None\n        self._w = None\n\n\n    def _ensure_open(self):\n\n        if self._slices is not None:\n            return\n\n        slices = []\n\n        for p in self.paths:\n\n            try:\n\n                arr = tifffile.memmap(\n                    p,\n                    mode=\"r\"\n                )\n\n            except Exception:\n\n                arr = tifffile.imread(p)\n\n            slices.append(arr)\n\n        self._slices = slices\n\n        self._h, self._w = (\n            slices[0].shape\n        )\n\n\n    @property\n    def shape(self):\n\n        self._ensure_open()\n\n        return (\n            self._h,\n            self._w\n        )\n\n\n    def read_patch(\n        self,\n        y,\n        x,\n        size\n    ):\n\n        self._ensure_open()\n\n        out = np.empty(\n            (\n                len(self._slices),\n                size,\n                size\n            ),\n            dtype=np.uint8\n        )\n\n        for i, s in enumerate(\n            self._slices\n        ):\n\n            block = s[\n                y:y + size,\n                x:x + size\n            ]\n\n            if block.shape != (\n                size,\n                size\n            ):\n\n                padded = np.zeros(\n                    (size, size),\n                    dtype=np.uint8\n                )\n\n                hh = min(\n                    size,\n                    block.shape[0]\n                )\n\n                ww = min(\n                    size,\n                    block.shape[1]\n                )\n\n                if block.dtype != np.uint8:\n\n                    block = (\n                        block.astype(\n                            np.float32\n                        )\n                        / 65535.0\n                        * 255.0\n                    ).astype(np.uint8)\n\n                padded[\n                    :hh,\n                    :ww\n                ] = block[\n                    :hh,\n                    :ww\n                ]\n\n                block = padded\n\n            elif block.dtype != np.uint8:\n\n                block = (\n                    block.astype(\n                        np.float32\n                    )\n                    / 65535.0\n                    * 255.0\n                ).astype(np.uint8)\n\n            out[i] = block\n\n        return out\n\n\n    def close(self):\n\n        self._slices = None\n\n        cleanup_memory()\n\n\n# ============================================================\n# 11. NORMALIZATION\n# ============================================================\n\ndef normalize_patch(img_float):\n\n    mean = img_float.mean()\n\n    std = img_float.std() + 1e-6\n\n    return (\n        img_float - mean\n    ) / std\n\n\n# ============================================================\n# 12. AUGMENTATION\n# ============================================================\n\ndef build_train_transform():\n\n    return A.Compose([\n\n        A.HorizontalFlip(\n            p=0.5\n        ),\n\n        A.VerticalFlip(\n            p=0.5\n        ),\n\n        A.RandomRotate90(\n            p=0.5\n        ),\n\n        A.Transpose(\n            p=0.5\n        ),\n\n        A.ShiftScaleRotate(\n            shift_limit=0.04,\n            scale_limit=0.10,\n            rotate_limit=20,\n            border_mode=cv2.BORDER_REFLECT,\n            p=0.40\n        ),\n\n        A.GridDistortion(\n            num_steps=5,\n            distort_limit=0.15,\n            p=0.15\n        ),\n\n        A.RandomBrightnessContrast(\n            brightness_limit=0.10,\n            contrast_limit=0.10,\n            p=0.25\n        )\n\n    ])\n\n\n# ============================================================\n# 13. DATASET\n# ============================================================\n\nclass InkPatchDataset(\n    Dataset\n):\n\n    def __init__(\n        self,\n        volumes,\n        labels,\n        samples,\n        patch_size,\n        transform=None,\n        jitter=0\n    ):\n\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n\n        self.patch_size = patch_size\n\n        self.transform = transform\n\n        self.jitter = jitter\n\n\n    def __len__(self):\n\n        return len(\n            self.samples\n        )\n\n\n    def __getitem__(\n        self,\n        idx\n    ):\n\n        fid, y, x = (\n            self.samples[idx]\n        )\n\n        size = self.patch_size\n\n        vol = self.volumes[fid]\n\n        H, W = vol.shape\n\n        # ----------------------------------------------------\n        # Random jitter\n        # ----------------------------------------------------\n\n        if self.jitter > 0:\n\n            dy = random.randint(\n                -self.jitter,\n                self.jitter\n            )\n\n            dx = random.randint(\n                -self.jitter,\n                self.jitter\n            )\n\n            y = y + dy\n            x = x + dx\n\n            y = max(\n                0,\n                min(\n                    y,\n                    H - size\n                )\n            )\n\n            x = max(\n                0,\n                min(\n                    x,\n                    W - size\n                )\n            )\n\n        # ----------------------------------------------------\n        # Read\n        # ----------------------------------------------------\n\n        patch = vol.read_patch(\n            y,\n            x,\n            size\n        )\n\n        label = self.labels[fid][\n            y:y + size,\n            x:x + size\n        ]\n\n        # ----------------------------------------------------\n        # HWC\n        # ----------------------------------------------------\n\n        img = np.transpose(\n            patch,\n            (1, 2, 0)\n        )\n\n        # ----------------------------------------------------\n        # Augmentation\n        # ----------------------------------------------------\n\n        if self.transform is not None:\n\n            aug = self.transform(\n                image=img,\n                mask=label\n            )\n\n            img = aug[\"image\"]\n\n            label = aug[\"mask\"]\n\n        # ----------------------------------------------------\n        # Normalize\n        # ----------------------------------------------------\n\n        img = (\n            img.astype(\n                np.float32\n            )\n            / 255.0\n        )\n\n        img = normalize_patch(\n            img\n        )\n\n        # CHW\n        img = np.ascontiguousarray(\n            np.transpose(\n                img,\n                (2, 0, 1)\n            )\n        )\n\n        label = (\n            label > 0\n        ).astype(\n            np.float32\n        )\n\n        label = label[\n            None,\n            ...\n        ]\n\n        return (\n            torch.from_numpy(img),\n            torch.from_numpy(label)\n        )\n\n\n# ============================================================\n# 14. MODEL\n# ============================================================\n\ndef build_model():\n\n    model = smp.Unet(\n\n        encoder_name=CFG.encoder_name,\n\n        encoder_weights=CFG.encoder_weights,\n\n        in_channels=CFG.in_channels,\n\n        classes=1,\n\n        decoder_attention_type=CFG.decoder_attention_type\n\n    )\n\n    return model\n\n\n# ============================================================\n# 15. LOSSES\n# ============================================================\n\ndef soft_dice_loss(\n    logits,\n    targets,\n    eps=1e-6\n):\n\n    probs = torch.sigmoid(\n        logits\n    )\n\n    probs = probs.reshape(\n        probs.size(0),\n        -1\n    )\n\n    targets = targets.reshape(\n        targets.size(0),\n        -1\n    )\n\n    intersection = (\n        probs * targets\n    ).sum(dim=1)\n\n    denominator = (\n        probs.sum(dim=1)\n        +\n        targets.sum(dim=1)\n    )\n\n    dice = (\n        2.0 * intersection\n        + eps\n    ) / (\n        denominator\n        + eps\n    )\n\n    return (\n        1.0 - dice\n    ).mean()\n\n\ndef focal_tversky_loss(\n    logits,\n    targets,\n    alpha=0.30,\n    beta=0.70,\n    gamma=0.75,\n    eps=1e-6\n):\n\n    probs = torch.sigmoid(\n        logits\n    )\n\n    probs = probs.reshape(\n        probs.size(0),\n        -1\n    )\n\n    targets = targets.reshape(\n        targets.size(0),\n        -1\n    )\n\n    tp = (\n        probs * targets\n    ).sum(dim=1)\n\n    fp = (\n        probs * (1.0 - targets)\n    ).sum(dim=1)\n\n    fn = (\n        (1.0 - probs) * targets\n    ).sum(dim=1)\n\n    tversky = (\n        tp + eps\n    ) / (\n        tp\n        + alpha * fp\n        + beta * fn\n        + eps\n    )\n\n    loss = torch.pow(\n        1.0 - tversky,\n        gamma\n    )\n\n    return loss.mean()\n\n\nclass V2ComboLoss(\n    nn.Module\n):\n\n    def __init__(\n        self,\n        pos_weight=None\n    ):\n\n        super().__init__()\n\n        self.pos_weight = pos_weight\n\n    def forward(\n        self,\n        logits,\n        targets\n    ):\n\n        bce = F.binary_cross_entropy_with_logits(\n            logits,\n            targets,\n            pos_weight=self.pos_weight\n        )\n\n        dice = soft_dice_loss(\n            logits,\n            targets\n        )\n\n        tv = focal_tversky_loss(\n            logits,\n            targets,\n            alpha=CFG.tversky_alpha,\n            beta=CFG.tversky_beta,\n            gamma=CFG.focal_tversky_gamma\n        )\n\n        loss = (\n            CFG.bce_weight * bce\n            +\n            CFG.dice_weight * dice\n            +\n            CFG.focal_tversky_weight * tv\n        )\n\n        return loss\n\n\n# ============================================================\n# 16. METRICS\n# ============================================================\n\ndef dice_from_counts(\n    tp,\n    fp,\n    fn,\n    eps=1e-6\n):\n\n    return (\n        2.0 * tp + eps\n    ) / (\n        2.0 * tp\n        + fp\n        + fn\n        + eps\n    )\n\n\ndef metrics_from_counts(\n    tp,\n    fp,\n    fn,\n    eps=1e-6\n):\n\n    dice = dice_from_counts(\n        tp,\n        fp,\n        fn,\n        eps\n    )\n\n    iou = (\n        tp + eps\n    ) / (\n        tp\n        + fp\n        + fn\n        + eps\n    )\n\n    precision = (\n        tp + eps\n    ) / (\n        tp\n        + fp\n        + eps\n    )\n\n    recall = (\n        tp + eps\n    ) / (\n        tp\n        + fn\n        + eps\n    )\n\n    beta2 = 0.25\n\n    fbeta = (\n        (1.0 + beta2)\n        * precision\n        * recall\n        + eps\n    ) / (\n        beta2 * precision\n        + recall\n        + eps\n    )\n\n    return {\n        \"dice\": float(dice),\n        \"iou\": float(iou),\n        \"precision\": float(precision),\n        \"recall\": float(recall),\n        \"fbeta0.5\": float(fbeta)\n    }\n\n\nclass GlobalConfusionAccumulator:\n\n    def __init__(self):\n\n        self.tp = 0.0\n        self.fp = 0.0\n        self.fn = 0.0\n\n\n    def update(\n        self,\n        probs,\n        targets,\n        threshold\n    ):\n\n        preds = (\n            probs > threshold\n        ).float()\n\n        self.tp += (\n            preds * targets\n        ).sum().item()\n\n        self.fp += (\n            preds * (1.0 - targets)\n        ).sum().item()\n\n        self.fn += (\n            (1.0 - preds) * targets\n        ).sum().item()\n\n\n    def compute(self):\n\n        return metrics_from_counts(\n            self.tp,\n            self.fp,\n            self.fn\n        )\n\n\n# ============================================================\n# 17. POSITIVE FRACTION\n# ============================================================\n\ndef estimate_positive_fraction(\n    labels_full,\n    samples,\n    patch_size,\n    n_samples=100\n):\n\n    if len(samples) == 0:\n        return 1e-4\n\n    sub = random.sample(\n        samples,\n        min(\n            n_samples,\n            len(samples)\n        )\n    )\n\n    total = 0\n    pos = 0\n\n    for fid, y, x in sub:\n\n        patch = labels_full[fid][\n            y:y + patch_size,\n            x:x + patch_size\n        ]\n\n        pos += patch.sum()\n\n        total += patch.size\n\n    frac = (\n        pos\n        /\n        max(total, 1)\n    )\n\n    return max(\n        float(frac),\n        1e-4\n    )\n\n\n# ============================================================\n# 18. BUILD TRAIN DATA\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"BUILDING TRAIN DATA\")\nprint(\"=\" * 70)\n\ntrain_volumes = {}\ntrain_labels_full = {}\n\ntrain_samples_raw = []\nval_samples = []\n\n\nfor fid in CFG.train_frags:\n\n    print(\n        f\"\\nProcessing fragment {fid} ...\"\n    )\n\n    frag_dir = os.path.join(\n        CFG.base_dir,\n        fid\n    )\n\n    mask = load_tissue_mask(\n        frag_dir\n    )\n\n    labels = load_ink_labels(\n        frag_dir\n    )\n\n    if labels is None:\n        raise RuntimeError(\n            f\"inklabels.png missing for fragment {fid}\"\n        )\n\n    vol = FragmentVolume(\n        frag_dir,\n        CFG.depth_indices\n    )\n\n    coords = generate_grid_coords(\n        mask,\n        CFG.patch_size,\n        CFG.train_stride,\n        CFG.min_tissue_frac_train\n    )\n\n    print(\n        f\"Fragment {fid}: \"\n        f\"{len(coords)} candidate patches\"\n    )\n\n    tr_coords, va_coords = (\n        spatial_split_coords(\n            coords,\n            mask.shape,\n            CFG.patch_size,\n            CFG.val_fraction\n        )\n    )\n\n    print(\n        f\"  spatial train={len(tr_coords)} \"\n        f\"validation={len(va_coords)}\"\n    )\n\n    train_volumes[fid] = vol\n\n    train_labels_full[fid] = labels\n\n    train_samples_raw.extend(\n        [\n            (fid, y, x)\n            for y, x in tr_coords\n        ]\n    )\n\n    val_samples.extend(\n        [\n            (fid, y, x)\n            for y, x in va_coords\n        ]\n    )\n\n    del mask\n\n    cleanup_memory()\n\n\nprint(\n    \"\\nRaw train samples:\",\n    len(train_samples_raw)\n)\n\nprint(\n    \"Validation samples:\",\n    len(val_samples)\n)\n\n\n# ============================================================\n# 19. BALANCE TRAIN PATCHES\n# ============================================================\n\ntrain_samples = balance_positive_patches(\n    train_samples_raw,\n    train_labels_full,\n    CFG.patch_size,\n    positive_threshold=CFG.positive_patch_fraction,\n    target_positive_ratio=CFG.target_positive_patch_ratio,\n    max_positive_repeat=CFG.max_positive_repeat\n)\n\n\n# ============================================================\n# 20. DATASETS\n# ============================================================\n\ntrain_transform = (\n    build_train_transform()\n)\n\ntrain_ds = InkPatchDataset(\n    train_volumes,\n    train_labels_full,\n    train_samples,\n    CFG.patch_size,\n    transform=train_transform,\n    jitter=CFG.train_jitter\n)\n\n\n# IMPORTANT:\n# validation has NO jitter and NO augmentation.\n\nval_ds = InkPatchDataset(\n    train_volumes,\n    train_labels_full,\n    val_samples,\n    CFG.patch_size,\n    transform=None,\n    jitter=0\n)\n\n\n# ============================================================\n# 21. DATALOADERS\n# ============================================================\n\ntrain_loader = DataLoader(\n\n    train_ds,\n\n    batch_size=CFG.batch_size,\n\n    shuffle=True,\n\n    num_workers=CFG.num_workers,\n\n    pin_memory=(\n        CFG.device == \"cuda\"\n    ),\n\n    drop_last=CFG.drop_last,\n\n    persistent_workers=(\n        CFG.num_workers > 0\n    )\n)\n\n\nval_loader = DataLoader(\n\n    val_ds,\n\n    batch_size=CFG.batch_size,\n\n    shuffle=False,\n\n    num_workers=CFG.num_workers,\n\n    pin_memory=(\n        CFG.device == \"cuda\"\n    ),\n\n    drop_last=False,\n\n    persistent_workers=(\n        CFG.num_workers > 0\n    )\n)\n\n\n# ============================================================\n# 22. MODEL\n# ============================================================\n\nprint(\"\\nBuilding V2 model ...\")\n\nmodel = build_model().to(\n    CFG.device\n)\n\n\n# ============================================================\n# 23. OUTPUT BIAS INITIALIZATION\n# ============================================================\n\nprint(\n    \"\\nEstimating positive-pixel fraction ...\"\n)\n\npos_frac = estimate_positive_fraction(\n    train_labels_full,\n    train_samples,\n    CFG.patch_size,\n    n_samples=100\n)\n\nprint(\n    f\"Estimated positive fraction: \"\n    f\"{pos_frac:.6f}\"\n)\n\n\nwith torch.no_grad():\n\n    bias_val = float(\n        np.log(\n            pos_frac\n            /\n            max(\n                1.0 - pos_frac,\n                1e-6\n            )\n        )\n    )\n\n    model.segmentation_head[\n        0\n    ].bias.fill_(\n        bias_val\n    )\n\n\nprint(\n    f\"Output bias initialized to \"\n    f\"{bias_val:.4f}\"\n)\n\n\n# ============================================================\n# 24. POSITIVE WEIGHT\n# ============================================================\n\n# We do not use the raw ratio without a cap.\n#\n# Sparse segmentation can produce an enormous BCE weight.\n#\n# sqrt ratio is more stable.\n\nraw_ratio = (\n    (1.0 - pos_frac)\n    /\n    max(pos_frac, 1e-6)\n)\n\npos_weight_val = float(\n    np.clip(\n        np.sqrt(raw_ratio),\n        1.0,\n        8.0\n    )\n)\n\npos_weight = torch.tensor(\n    [pos_weight_val],\n    dtype=torch.float32,\n    device=CFG.device\n)\n\nprint(\n    f\"BCE positive weight: \"\n    f\"{pos_weight_val:.3f}\"\n)\n\n\n# ============================================================\n# 25. LOSS\n# ============================================================\n\ncriterion = V2ComboLoss(\n    pos_weight=pos_weight\n)\n\n\n# ============================================================\n# 26. DIFFERENTIAL LR OPTIMIZER\n# ============================================================\n\nencoder_params = []\ndecoder_params = []\n\nfor name, param in model.named_parameters():\n\n    if not param.requires_grad:\n        continue\n\n    if name.startswith(\"encoder.\"):\n\n        encoder_params.append(\n            param\n        )\n\n    else:\n\n        decoder_params.append(\n            param\n        )\n\n\nprint(\n    f\"Encoder parameters: \"\n    f\"{len(encoder_params)} tensors\"\n)\n\nprint(\n    f\"Decoder/head parameters: \"\n    f\"{len(decoder_params)} tensors\"\n)\n\n\noptimizer = torch.optim.AdamW(\n\n    [\n        {\n            \"params\": encoder_params,\n            \"lr\": CFG.encoder_lr\n        },\n\n        {\n            \"params\": decoder_params,\n            \"lr\": CFG.decoder_lr\n        }\n    ],\n\n    weight_decay=CFG.weight_decay\n)\n\n\n# ============================================================\n# 27. COSINE LR\n# ============================================================\n\nscheduler = (\n    torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=CFG.epochs,\n        eta_min=1e-6\n    )\n)\n\n\n# ============================================================\n# 28. AMP\n# ============================================================\n\nscaler = GradScaler(\n    enabled=(\n        CFG.device == \"cuda\"\n    )\n)\n\n\n# ============================================================\n# 29. TRAIN / VALIDATION EPOCH\n# ============================================================\n\ndef run_epoch(\n    loader,\n    train_mode=True,\n    threshold=0.5\n):\n\n    if train_mode:\n\n        model.train()\n\n    else:\n\n        model.eval()\n\n\n    total_loss = 0.0\n\n    global_acc = (\n        GlobalConfusionAccumulator()\n    )\n\n\n    if train_mode:\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n\n    for batch_idx, (\n        imgs,\n        masks\n    ) in enumerate(loader):\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        masks = masks.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n\n        with torch.set_grad_enabled(\n            train_mode\n        ):\n\n            with autocast(\n                enabled=(\n                    CFG.device == \"cuda\"\n                )\n            ):\n\n                logits = model(\n                    imgs\n                )\n\n                loss = criterion(\n                    logits,\n                    masks\n                )\n\n                if train_mode:\n\n                    loss_for_backward = (\n                        loss\n                        /\n                        CFG.accumulation_steps\n                    )\n\n\n            if train_mode:\n\n                scaler.scale(\n                    loss_for_backward\n                ).backward()\n\n\n                if (\n                    (batch_idx + 1)\n                    %\n                    CFG.accumulation_steps\n                    == 0\n                ):\n\n                    scaler.unscale_(\n                        optimizer\n                    )\n\n                    torch.nn.utils.clip_grad_norm_(\n                        model.parameters(),\n                        CFG.grad_clip\n                    )\n\n                    scaler.step(\n                        optimizer\n                    )\n\n                    scaler.update()\n\n                    optimizer.zero_grad(\n                        set_to_none=True\n                    )\n\n\n        probs = torch.sigmoid(\n            logits.detach()\n        )\n\n\n        global_acc.update(\n            probs,\n            masks,\n            threshold\n        )\n\n\n        total_loss += (\n            loss.detach().item()\n        )\n\n\n        del (\n            imgs,\n            masks,\n            logits,\n            probs\n        )\n\n\n    # If number of batches isn't divisible by accumulation steps\n    if (\n        train_mode\n        and\n        len(loader)\n        %\n        CFG.accumulation_steps\n        != 0\n    ):\n\n        scaler.unscale_(\n            optimizer\n        )\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            CFG.grad_clip\n        )\n\n        scaler.step(\n            optimizer\n        )\n\n        scaler.update()\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n\n    metrics = global_acc.compute()\n\n    avg_loss = (\n        total_loss\n        /\n        max(len(loader), 1)\n    )\n\n    return avg_loss, metrics\n\n\n# ============================================================\n# 30. TRAINING LOOP\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 70)\nprint(\"STARTING V2 TRAINING\")\nprint(\"=\" * 70)\n\nbest_val_dice = -1.0\n\nepochs_no_improve = 0\n\nhistory = {\n    \"train_loss\": [],\n    \"val_loss\": [],\n    \"train_dice\": [],\n    \"val_dice\": [],\n    \"val_iou\": [],\n    \"val_precision\": [],\n    \"val_recall\": []\n}\n\n\nfor epoch in range(\n    1,\n    CFG.epochs + 1\n):\n\n    t0 = time.time()\n\n\n    # --------------------------------------------------------\n    # Train\n    # --------------------------------------------------------\n\n    train_loss, train_metrics = run_epoch(\n        train_loader,\n        train_mode=True,\n        threshold=0.50\n    )\n\n\n    # --------------------------------------------------------\n    # Validation\n    # --------------------------------------------------------\n\n    val_loss, val_metrics = run_epoch(\n        val_loader,\n        train_mode=False,\n        threshold=0.50\n    )\n\n\n    # --------------------------------------------------------\n    # Scheduler\n    # --------------------------------------------------------\n\n    scheduler.step()\n\n\n    # --------------------------------------------------------\n    # History\n    # --------------------------------------------------------\n\n    history[\n        \"train_loss\"\n    ].append(train_loss)\n\n    history[\n        \"val_loss\"\n    ].append(val_loss)\n\n    history[\n        \"train_dice\"\n    ].append(train_metrics[\"dice\"])\n\n    history[\n        \"val_dice\"\n    ].append(val_metrics[\"dice\"])\n\n    history[\n        \"val_iou\"\n    ].append(val_metrics[\"iou\"])\n\n    history[\n        \"val_precision\"\n    ].append(val_metrics[\"precision\"])\n\n    history[\n        \"val_recall\"\n    ].append(val_metrics[\"recall\"])\n\n\n    # --------------------------------------------------------\n    # LR\n    # --------------------------------------------------------\n\n    current_encoder_lr = (\n        optimizer.param_groups[0][\"lr\"]\n    )\n\n    current_decoder_lr = (\n        optimizer.param_groups[1][\"lr\"]\n    )\n\n\n    # --------------------------------------------------------\n    # Print\n    # --------------------------------------------------------\n\n    print(\n        f\"\\n[{epoch:02d}/{CFG.epochs}] \"\n        f\"time={time.time()-t0:.1f}s\"\n    )\n\n    print(\n        f\"train_loss={train_loss:.5f} \"\n        f\"train_dice={train_metrics['dice']:.5f}\"\n    )\n\n    print(\n        f\"val_loss={val_loss:.5f} \"\n        f\"val_dice={val_metrics['dice']:.5f} \"\n        f\"val_iou={val_metrics['iou']:.5f}\"\n    )\n\n    print(\n        f\"precision={val_metrics['precision']:.5f} \"\n        f\"recall={val_metrics['recall']:.5f}\"\n    )\n\n    print(\n        f\"encoder_lr={current_encoder_lr:.7f} \"\n        f\"decoder_lr={current_decoder_lr:.7f}\"\n    )\n\n\n    # --------------------------------------------------------\n    # Save best\n    # --------------------------------------------------------\n\n    if (\n        val_metrics[\"dice\"]\n        >\n        best_val_dice\n    ):\n\n        best_val_dice = (\n            val_metrics[\"dice\"]\n        )\n\n        epochs_no_improve = 0\n\n\n        checkpoint = {\n\n            \"model\":\n                model.state_dict(),\n\n            \"cfg\":\n                cfg_to_dict(CFG),\n\n            \"best_val_dice\":\n                best_val_dice,\n\n            \"history\":\n                history,\n\n            \"pos_frac\":\n                pos_frac,\n\n            \"pos_weight\":\n                pos_weight_val\n\n        }\n\n\n        torch.save(\n            checkpoint,\n            CFG.ckpt_path\n        )\n\n\n        print(\n            f\"*** NEW BEST CHECKPOINT \"\n            f\"val_dice={best_val_dice:.5f} ***\"\n        )\n\n\n    else:\n\n        epochs_no_improve += 1\n\n        print(\n            f\"No improvement: \"\n            f\"{epochs_no_improve}/\"\n            f\"{CFG.early_stop_patience}\"\n        )\n\n\n        if (\n            epochs_no_improve\n            >=\n            CFG.early_stop_patience\n        ):\n\n            print(\n                \"Early stopping.\"\n            )\n\n            break\n\n\n    cleanup_memory()\n\n\nprint(\n    \"\\nBest validation Dice:\",\n    best_val_dice\n)\n\n\n# ============================================================\n# 31. LOAD BEST MODEL\n# ============================================================\n\ncheckpoint = torch.load(\n    CFG.ckpt_path,\n    map_location=CFG.device\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\nprint(\n    \"\\nBest checkpoint loaded.\"\n)\n\n\n# ============================================================\n# 32. FINE DICE THRESHOLD SEARCH\n# ============================================================\n\n@torch.no_grad()\ndef find_best_dice_threshold(\n    model,\n    loader\n):\n\n    model.eval()\n\n\n    thresholds = np.arange(\n        CFG.threshold_min,\n        CFG.threshold_max\n        + 0.0001,\n        CFG.threshold_step\n    )\n\n\n    tp = np.zeros(\n        len(thresholds),\n        dtype=np.float64\n    )\n\n    fp = np.zeros(\n        len(thresholds),\n        dtype=np.float64\n    )\n\n    fn = np.zeros(\n        len(thresholds),\n        dtype=np.float64\n    )\n\n\n    for imgs, masks in loader:\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        masks_np = (\n            masks.numpy()\n        )\n\n\n        with autocast(\n            enabled=(\n                CFG.device == \"cuda\"\n            )\n        ):\n\n            logits = model(\n                imgs\n            )\n\n\n        probs = (\n            torch.sigmoid(\n                logits\n            )\n            .float()\n            .cpu()\n            .numpy()\n        )\n\n\n        for i, t in enumerate(\n            thresholds\n        ):\n\n            preds = (\n                probs > t\n            ).astype(\n                np.float32\n            )\n\n            tp[i] += (\n                preds\n                * masks_np\n            ).sum()\n\n            fp[i] += (\n                preds\n                * (1.0 - masks_np)\n            ).sum()\n\n            fn[i] += (\n                (1.0 - preds)\n                * masks_np\n            ).sum()\n\n\n        del (\n            imgs,\n            masks,\n            logits,\n            probs\n        )\n\n\n    dice_scores = (\n        2.0 * tp + 1e-6\n    ) / (\n        2.0 * tp\n        + fp\n        + fn\n        + 1e-6\n    )\n\n\n    best_idx = int(\n        np.argmax(\n            dice_scores\n        )\n    )\n\n\n    best_threshold = float(\n        thresholds[best_idx]\n    )\n\n    best_dice = float(\n        dice_scores[best_idx]\n    )\n\n\n    return (\n        best_threshold,\n        best_dice,\n        thresholds,\n        dice_scores\n    )\n\n\nprint(\"\\n\")\nprint(\"=\" * 70)\nprint(\"FINE DICE THRESHOLD SEARCH\")\nprint(\"=\" * 70)\n\n\nbest_threshold, threshold_dice, threshold_grid, threshold_scores = (\n    find_best_dice_threshold(\n        model,\n        val_loader\n    )\n)\n\n\nprint(\n    f\"BEST VALIDATION THRESHOLD = \"\n    f\"{best_threshold:.2f}\"\n)\n\nprint(\n    f\"DICE AT BEST THRESHOLD = \"\n    f\"{threshold_dice:.5f}\"\n)\n\n\n# ============================================================\n# 33. VALIDATION METRICS AT BEST THRESHOLD\n# ============================================================\n\n_, final_val_metrics = run_epoch(\n    val_loader,\n    train_mode=False,\n    threshold=best_threshold\n)\n\n\nprint(\n    \"\\nFinal validation metrics \"\n    \"at optimized Dice threshold:\"\n)\n\nprint(\n    final_val_metrics\n)\n\n\n# ============================================================\n# 34. GAUSSIAN WINDOW\n# ============================================================\n\ndef gaussian_window(\n    size,\n    sigma_frac=0.45\n):\n\n    ax = (\n        np.arange(size)\n        -\n        (size - 1) / 2.0\n    )\n\n    sigma = (\n        size * sigma_frac\n    )\n\n    g1d = np.exp(\n        -(\n            ax ** 2\n        )\n        /\n        (\n            2.0 * sigma ** 2\n        )\n    )\n\n    win = np.outer(\n        g1d,\n        g1d\n    )\n\n    win = (\n        win\n        /\n        max(\n            win.max(),\n            1e-8\n        )\n    )\n\n    return win.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 35. TTA HELPERS\n# ============================================================\n\ndef apply_tta_tensor(\n    x,\n    mode\n):\n\n    if mode == \"original\":\n\n        return x\n\n    if mode == \"hflip\":\n\n        return torch.flip(\n            x,\n            dims=[3]\n        )\n\n    if mode == \"vflip\":\n\n        return torch.flip(\n            x,\n            dims=[2]\n        )\n\n    if mode == \"rot90\":\n\n        return torch.rot90(\n            x,\n            k=1,\n            dims=[2, 3]\n        )\n\n    raise ValueError(\n        f\"Unknown TTA mode: {mode}\"\n    )\n\n\ndef invert_tta_prediction(\n    pred,\n    mode\n):\n\n    if mode == \"original\":\n\n        return pred\n\n    if mode == \"hflip\":\n\n        return torch.flip(\n            pred,\n            dims=[2]\n        )\n\n    if mode == \"vflip\":\n\n        return torch.flip(\n            pred,\n            dims=[1]\n        )\n\n    if mode == \"rot90\":\n\n        return torch.rot90(\n            pred,\n            k=-1,\n            dims=[1, 2]\n        )\n\n    raise ValueError(\n        f\"Unknown TTA mode: {mode}\"\n    )\n\n\n# ============================================================\n# 36. TTA MODEL PREDICTION\n# ============================================================\n\n@torch.no_grad()\ndef predict_batch_tta(\n    model,\n    batch_np\n):\n\n    \"\"\"\n    batch_np:\n        N,C,H,W\n\n    Returns:\n        N,H,W probability maps\n\n    Memory-safe:\n    TTA modes are processed sequentially.\n    We never create a 4x batch simultaneously.\n    \"\"\"\n\n    x = torch.from_numpy(\n        batch_np\n    ).to(\n        CFG.device,\n        non_blocking=True\n    )\n\n\n    modes = (\n        CFG.tta_modes\n        if CFG.use_tta\n        else [\"original\"]\n    )\n\n\n    prediction_sum = None\n\n\n    for mode in modes:\n\n        tx = apply_tta_tensor(\n            x,\n            mode\n        )\n\n\n        with autocast(\n            enabled=(\n                CFG.device == \"cuda\"\n            )\n        ):\n\n            logits = model(\n                tx\n            )\n\n            probs = torch.sigmoid(\n                logits\n            )\n\n\n        probs = probs[\n            :,\n            0\n        ]\n\n\n        probs = (\n            invert_tta_prediction(\n                probs,\n                mode\n            )\n        )\n\n\n        probs = probs.float()\n\n\n        if prediction_sum is None:\n\n            prediction_sum = (\n                probs\n                /\n                len(modes)\n            )\n\n        else:\n\n            prediction_sum += (\n                probs\n                /\n                len(modes)\n            )\n\n\n        del (\n            tx,\n            logits,\n            probs\n        )\n\n\n    result = (\n        prediction_sum\n        .cpu()\n        .numpy()\n        .astype(np.float32)\n    )\n\n\n    del (\n        x,\n        prediction_sum\n    )\n\n    return result\n\n\n# ============================================================\n# 37. SLIDING WINDOW INFERENCE\n# ============================================================\n\n@torch.no_grad()\ndef sliding_window_inference_v2(\n    model,\n    vol,\n    mask,\n    patch_size,\n    stride,\n    device,\n    batch_size\n):\n\n    H, W = mask.shape\n\n\n    pred_sum = np.zeros(\n        (H, W),\n        dtype=np.float32\n    )\n\n    weight_sum = np.zeros(\n        (H, W),\n        dtype=np.float32\n    )\n\n\n    win = gaussian_window(\n        patch_size\n    )\n\n\n    coords = generate_grid_coords(\n        mask,\n        patch_size,\n        stride,\n        CFG.min_tissue_frac_test\n    )\n\n\n    print(\n        f\"Inference patches: \"\n        f\"{len(coords)}\"\n    )\n\n\n    model.eval()\n\n\n    batch_imgs = []\n    batch_coords = []\n\n\n    def flush():\n\n        if len(batch_imgs) == 0:\n            return\n\n\n        inp = np.stack(\n            batch_imgs\n        ).astype(\n            np.float32\n        )\n\n\n        probs = predict_batch_tta(\n            model,\n            inp\n        )\n\n\n        for p, (\n            cy,\n            cx\n        ) in zip(\n            probs,\n            batch_coords\n        ):\n\n            pred_sum[\n                cy:cy + patch_size,\n                cx:cx + patch_size\n            ] += (\n                p * win\n            )\n\n            weight_sum[\n                cy:cy + patch_size,\n                cx:cx + patch_size\n            ] += win\n\n\n        batch_imgs.clear()\n\n        batch_coords.clear()\n\n\n        del (\n            inp,\n            probs\n        )\n\n\n        cleanup_memory()\n\n\n    for y, x in coords:\n\n        raw = (\n            vol.read_patch(\n                y,\n                x,\n                patch_size\n            )\n            .astype(\n                np.float32\n            )\n            /\n            255.0\n        )\n\n\n        raw = normalize_patch(\n            raw\n        )\n\n\n        # Already CHW\n        batch_imgs.append(\n            raw\n        )\n\n        batch_coords.append(\n            (y, x)\n        )\n\n\n        if len(batch_imgs) >= batch_size:\n\n            flush()\n\n\n    flush()\n\n\n    weight_sum[\n        weight_sum <= 1e-8\n    ] = 1.0\n\n\n    result = (\n        pred_sum\n        /\n        weight_sum\n    )\n\n\n    return result.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 38. ADABN\n# ============================================================\n\n@torch.no_grad()\ndef recalibrate_batchnorm_v2(\n    model,\n    vol,\n    mask,\n    patch_size,\n    stride,\n    device,\n    max_patches,\n    batch_size\n):\n\n    print(\n        \"\\nStarting AdaBN...\"\n    )\n\n\n    # Reset BN statistics.\n\n    for m in model.modules():\n\n        if isinstance(\n            m,\n            nn.BatchNorm2d\n        ):\n\n            m.reset_running_stats()\n\n            m.momentum = None\n\n\n    coords = generate_grid_coords(\n        mask,\n        patch_size,\n        stride,\n        CFG.min_tissue_frac_test\n    )\n\n\n    random.shuffle(\n        coords\n    )\n\n\n    coords = coords[\n        :max_patches\n    ]\n\n\n    print(\n        f\"AdaBN patches: \"\n        f\"{len(coords)}\"\n    )\n\n\n    model.train()\n\n\n    for start in range(\n        0,\n        len(coords),\n        batch_size\n    ):\n\n        batch = coords[\n            start:\n            start + batch_size\n        ]\n\n\n        imgs = []\n\n\n        for y, x in batch:\n\n            raw = (\n                vol.read_patch(\n                    y,\n                    x,\n                    patch_size\n                )\n                .astype(\n                    np.float32\n                )\n                /\n                255.0\n            )\n\n\n            raw = normalize_patch(\n                raw\n            )\n\n\n            imgs.append(\n                raw\n            )\n\n\n        inp = torch.from_numpy(\n            np.stack(imgs)\n        ).to(\n            device,\n            non_blocking=True\n        )\n\n\n        # NO TTA here.\n        # We only want BN statistics from natural target images.\n\n        with autocast(\n            enabled=(\n                device == \"cuda\"\n            )\n        ):\n\n            model(inp)\n\n\n        del (\n            inp,\n            imgs\n        )\n\n\n    model.eval()\n\n    cleanup_memory()\n\n    print(\n        \"AdaBN finished.\"\n    )\n\n    return model\n\n\n# ============================================================\n# 39. POST PROCESSING\n# ============================================================\n\ndef remove_small_components(\n    binary,\n    min_size\n):\n\n    if min_size <= 0:\n        return binary\n\n\n    num_labels, labels, stats, _ = (\n        cv2.connectedComponentsWithStats(\n            binary.astype(\n                np.uint8\n            ),\n            connectivity=8\n        )\n    )\n\n\n    if num_labels <= 1:\n        return binary\n\n\n    output = np.zeros_like(\n        binary,\n        dtype=np.uint8\n    )\n\n\n    for label_id in range(\n        1,\n        num_labels\n    ):\n\n        area = stats[\n            label_id,\n            cv2.CC_STAT_AREA\n        ]\n\n        if area >= min_size:\n\n            output[\n                labels == label_id\n            ] = 1\n\n\n    return output\n\n\ndef postprocess_v2(\n    prob_map,\n    threshold\n):\n\n    binary = (\n        prob_map > threshold\n    ).astype(\n        np.uint8\n    )\n\n\n    if not CFG.use_postprocess:\n\n        return binary\n\n\n    # --------------------------------------------------------\n    # Conservative closing\n    #\n    # This joins tiny gaps without aggressively removing\n    # thin ink.\n    # --------------------------------------------------------\n\n    k = CFG.closing_kernel\n\n    kernel = np.ones(\n        (k, k),\n        dtype=np.uint8\n    )\n\n\n    binary = cv2.morphologyEx(\n        binary,\n        cv2.MORPH_CLOSE,\n        kernel,\n        iterations=1\n    )\n\n\n    # --------------------------------------------------------\n    # Remove tiny isolated components\n    # --------------------------------------------------------\n\n    binary = remove_small_components(\n        binary,\n        CFG.min_component_size\n    )\n\n\n    return binary.astype(\n        np.uint8\n    )\n\n\n# ============================================================\n# 40. OPTIONAL OTSU\n# ============================================================\n\ndef compute_otsu_threshold(\n    prob_map,\n    mask,\n    fallback\n):\n\n    vals = prob_map[\n        mask > 0\n    ]\n\n\n    vals = vals[\n        np.isfinite(vals)\n    ]\n\n\n    if len(vals) < 100:\n\n        return fallback\n\n\n    try:\n\n        t = float(\n            threshold_otsu(\n                vals\n            )\n        )\n\n        # Prevent pathological Otsu thresholds.\n        t = float(\n            np.clip(\n                t,\n                0.10,\n                0.90\n            )\n        )\n\n        return t\n\n    except Exception:\n\n        return fallback\n\n\n# ============================================================\n# 41. LOAD TEST FRAGMENT\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 70)\nprint(\"LOADING TEST FRAGMENT\")\nprint(\"=\" * 70)\n\n\ntest_dir = os.path.join(\n    CFG.base_dir,\n    CFG.test_frag\n)\n\n\ntest_mask = load_tissue_mask(\n    test_dir\n)\n\n\ntest_vol = FragmentVolume(\n    test_dir,\n    CFG.depth_indices\n)\n\n\n# Optional:\n#\n# If inklabels.png exists, we use it ONLY for local diagnostics.\n#\n# It is NOT used to select the model, threshold, TTA or postprocessing.\n\ntest_labels = load_ink_labels(\n    test_dir\n)\n\n\nif test_labels is not None:\n\n    print(\n        \"Test GT found: \"\n        \"local diagnostic evaluation enabled.\"\n    )\n\nelse:\n\n    print(\n        \"No test GT found: \"\n        \"running competition-style inference.\"\n    )\n\n\n# ============================================================\n# 42. BASELINE INFERENCE\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 70)\nprint(\"INFERENCE A: BASELINE + TTA\")\nprint(\"=\" * 70)\n\n\ntest_prob_baseline = (\n    sliding_window_inference_v2(\n        model,\n        test_vol,\n        test_mask,\n        CFG.patch_size,\n        CFG.test_stride,\n        CFG.device,\n        CFG.infer_batch\n    )\n)\n\n\n# ============================================================\n# 43. BASELINE METRICS IF GT EXISTS\n# ============================================================\n\ndef evaluate_probability_map(\n    prob_map,\n    gt,\n    threshold\n):\n\n    preds = (\n        prob_map > threshold\n    ).astype(\n        np.float32\n    )\n\n    gt = gt.astype(\n        np.float32\n    )\n\n\n    tp = (\n        preds * gt\n    ).sum()\n\n    fp = (\n        preds * (1.0 - gt)\n    ).sum()\n\n    fn = (\n        (1.0 - preds) * gt\n    ).sum()\n\n\n    return metrics_from_counts(\n        tp,\n        fp,\n        fn\n    )\n\n\ntest_metrics_baseline = None\n\n\nif test_labels is not None:\n\n    gt_test = (\n        test_labels\n        *\n        test_mask\n    ).astype(\n        np.float32\n    )\n\n\n    test_metrics_baseline = (\n        evaluate_probability_map(\n            test_prob_baseline,\n            gt_test,\n            best_threshold\n        )\n    )\n\n\n    print(\n        \"\\nBASELINE TEST METRICS:\"\n    )\n\n    print(\n        test_metrics_baseline\n    )\n\n\n# ============================================================\n# 44. ADABN INFERENCE\n# ============================================================\n\ntest_prob_adabn = None\n\ntest_metrics_adabn = None\n\n\nif CFG.use_adabn:\n\n    print(\"\\n\")\n    print(\"=\" * 70)\n    print(\"INFERENCE B: ADABN + TTA\")\n    print(\"=\" * 70)\n\n\n    # IMPORTANT:\n    #\n    # AdaBN changes model BN statistics.\n    # Therefore baseline probability map is already saved above.\n\n    model = recalibrate_batchnorm_v2(\n        model,\n        test_vol,\n        test_mask,\n        CFG.patch_size,\n        CFG.test_stride,\n        CFG.device,\n        CFG.adabn_max_patches,\n        CFG.infer_batch\n    )\n\n\n    test_prob_adabn = (\n        sliding_window_inference_v2(\n            model,\n            test_vol,\n            test_mask,\n            CFG.patch_size,\n            CFG.test_stride,\n            CFG.device,\n            CFG.infer_batch\n        )\n    )\n\n\n    if test_labels is not None:\n\n        test_metrics_adabn = (\n            evaluate_probability_map(\n                test_prob_adabn,\n                gt_test,\n                best_threshold\n            )\n        )\n\n\n        print(\n            \"\\nADABN TEST METRICS:\"\n        )\n\n        print(\n            test_metrics_adabn\n        )\n\n\n# ============================================================\n# 45. CHOOSE FINAL PROBABILITY MAP\n# ============================================================\n\n# IMPORTANT:\n#\n# In a real competition, do NOT use test GT to choose.\n#\n# Default choice is AdaBN if enabled.\n#\n# If your local ablation shows AdaBN hurts badly, set:\n#\n# use_adabn = False\n#\n# and rerun.\n\nif (\n    CFG.use_adabn\n    and\n    test_prob_adabn is not None\n):\n\n    test_prob = (\n        test_prob_adabn\n    )\n\nelse:\n\n    test_prob = (\n        test_prob_baseline\n    )\n\n\n# ============================================================\n# 46. OPTIONAL OTSU DIAGNOSTIC\n# ============================================================\n\notsu_threshold = (\n    compute_otsu_threshold(\n        test_prob,\n        test_mask,\n        fallback=best_threshold\n    )\n)\n\n\nprint(\n    \"\\nValidation-tuned threshold:\",\n    best_threshold\n)\n\nprint(\n    \"Unsupervised Otsu threshold:\",\n    otsu_threshold\n)\n\n\n# ============================================================\n# 47. FINAL THRESHOLD\n# ============================================================\n\n# For competition:\n#\n# Use validation Dice threshold.\n#\n# Otsu is printed for diagnostic purposes but is NOT used\n# automatically because validation-tuned threshold is generally\n# safer and reproducible.\n\nfinal_threshold = (\n    best_threshold\n)\n\n\n# ============================================================\n# 48. FINAL POST PROCESSING\n# ============================================================\n\ntest_pred_bin = postprocess_v2(\n    test_prob,\n    final_threshold\n)\n\n\n# ============================================================\n# 49. RAW / POSTPROCESSED METRICS\n# ============================================================\n\ntest_metrics_raw = None\ntest_metrics_post = None\ntest_metrics_otsu = None\n\n\nif test_labels is not None:\n\n    test_metrics_raw = (\n        evaluate_probability_map(\n            test_prob,\n            gt_test,\n            final_threshold\n        )\n    )\n\n\n    test_metrics_otsu = (\n        evaluate_probability_map(\n            test_prob,\n            gt_test,\n            otsu_threshold\n        )\n    )\n\n\n    post_preds = (\n        test_pred_bin\n        .astype(\n            np.float32\n        )\n    )\n\n\n    tp = (\n        post_preds * gt_test\n    ).sum()\n\n    fp = (\n        post_preds * (1.0 - gt_test)\n    ).sum()\n\n    fn = (\n        (1.0 - post_preds)\n        * gt_test\n    ).sum()\n\n\n    test_metrics_post = (\n        metrics_from_counts(\n            tp,\n            fp,\n            fn\n        )\n    )\n\n\n    print(\"\\n\")\n    print(\"=\" * 70)\n    print(\"LOCAL TEST DIAGNOSTICS\")\n    print(\"=\" * 70)\n\n\n    print(\n        \"Raw probability + validation threshold:\"\n    )\n\n    print(\n        test_metrics_raw\n    )\n\n\n    print(\n        \"\\nRaw probability + Otsu:\"\n    )\n\n    print(\n        test_metrics_otsu\n    )\n\n\n    print(\n        \"\\nPostprocessed + validation threshold:\"\n    )\n\n    print(\n        test_metrics_post\n    )\n\n\n# ============================================================\n# 50. SAVE PROBABILITY MAP\n# ============================================================\n\nprob_path = os.path.join(\n    CFG.out_dir,\n    \"fragment1_probability_v2.npy\"\n)\n\n\nnp.save(\n    prob_path,\n    test_prob\n)\n\n\nprint(\n    \"\\nSaved probability map:\",\n    prob_path\n)\n\n\n# ============================================================\n# 51. SAVE BINARY MASK\n# ============================================================\n\npred_path = os.path.join(\n    CFG.out_dir,\n    \"fragment1_prediction_v2.png\"\n)\n\n\ncv2.imwrite(\n    pred_path,\n    (\n        test_pred_bin * 255\n    ).astype(\n        np.uint8\n    )\n)\n\n\nprint(\n    \"Saved prediction:\",\n    pred_path\n)\n\n\n# ============================================================\n# 52. SAVE METRICS\n# ============================================================\n\nmetrics_summary = {\n\n    \"best_validation_dice_at_0.50\":\n        best_val_dice,\n\n    \"best_validation_threshold\":\n        best_threshold,\n\n    \"validation_dice_at_best_threshold\":\n        threshold_dice,\n\n    \"otsu_threshold\":\n        otsu_threshold,\n\n    \"final_threshold\":\n        final_threshold,\n\n    \"final_validation_metrics\":\n        final_val_metrics,\n\n    \"test_baseline\":\n        test_metrics_baseline,\n\n    \"test_adabn\":\n        test_metrics_adabn,\n\n    \"test_raw\":\n        test_metrics_raw,\n\n    \"test_otsu\":\n        test_metrics_otsu,\n\n    \"test_postprocessed\":\n        test_metrics_post,\n\n    \"config\":\n        cfg_to_dict(CFG),\n\n    \"history\":\n        history\n}\n\n\nwith open(\n    CFG.metrics_path,\n    \"w\"\n) as f:\n\n    json.dump(\n        metrics_summary,\n        f,\n        indent=2\n    )\n\n\nprint(\n    \"Saved metrics:\",\n    CFG.metrics_path\n)\n\n\n# ============================================================\n# 53. VISUALIZATION\n# ============================================================\n\ndef save_full_overview():\n\n    mid_idx = CFG.depth_indices[\n        len(CFG.depth_indices) // 2\n    ]\n\n\n    mid_path = os.path.join(\n        test_dir,\n        \"surface_volume\",\n        f\"{mid_idx:02d}.tif\"\n    )\n\n\n    mid_slice = tifffile.imread(\n        mid_path\n    )\n\n\n    scale = (\n        2000\n        /\n        max(mid_slice.shape)\n    )\n\n\n    small = cv2.resize(\n        mid_slice,\n        None,\n        fx=scale,\n        fy=scale,\n        interpolation=cv2.INTER_AREA\n    )\n\n\n    pred_small = cv2.resize(\n        (\n            test_pred_bin * 255\n        ).astype(\n            np.uint8\n        ),\n        small.shape[::-1],\n        interpolation=cv2.INTER_NEAREST\n    )\n\n\n    prob_small = cv2.resize(\n        (\n            test_prob * 255\n        ).astype(\n            np.uint8\n        ),\n        small.shape[::-1],\n        interpolation=cv2.INTER_AREA\n    )\n\n\n    if test_labels is not None:\n\n        gt_small = cv2.resize(\n            (\n                test_labels * 255\n            ).astype(\n                np.uint8\n            ),\n            small.shape[::-1],\n            interpolation=cv2.INTER_NEAREST\n        )\n\n\n        fig, axes = plt.subplots(\n            1,\n            4,\n            figsize=(22, 6)\n        )\n\n\n        axes[0].imshow(\n            small,\n            cmap=\"gray\"\n        )\n\n        axes[0].set_title(\n            f\"Input slice {mid_idx}\"\n        )\n\n\n        axes[1].imshow(\n            gt_small,\n            cmap=\"gray\"\n        )\n\n        axes[1].set_title(\n            \"Ground Truth\"\n        )\n\n\n        axes[2].imshow(\n            prob_small,\n            cmap=\"gray\"\n        )\n\n        axes[2].set_title(\n            f\"Probability \"\n            f\"threshold={final_threshold:.2f}\"\n        )\n\n\n        axes[3].imshow(\n            pred_small,\n            cmap=\"gray\"\n        )\n\n        axes[3].set_title(\n            \"Final Prediction\"\n        )\n\n\n    else:\n\n        fig, axes = plt.subplots(\n            1,\n            3,\n            figsize=(18, 6)\n        )\n\n\n        axes[0].imshow(\n            small,\n            cmap=\"gray\"\n        )\n\n        axes[0].set_title(\n            f\"Input slice {mid_idx}\"\n        )\n\n\n        axes[1].imshow(\n            prob_small,\n            cmap=\"gray\"\n        )\n\n        axes[1].set_title(\n            \"Probability\"\n        )\n\n\n        axes[2].imshow(\n            pred_small,\n            cmap=\"gray\"\n        )\n\n        axes[2].set_title(\n            \"Final Prediction\"\n        )\n\n\n    for ax in axes:\n        ax.axis(\"off\")\n\n\n    plt.tight_layout()\n\n\n    overview_path = os.path.join(\n        CFG.viz_dir,\n        \"fragment1_v2_overview.png\"\n    )\n\n\n    plt.savefig(\n        overview_path,\n        dpi=150,\n        bbox_inches=\"tight\"\n    )\n\n\n    plt.close(fig)\n\n\n    print(\n        \"Saved overview:\",\n        overview_path\n    )\n\n\nsave_full_overview()\n\n\n# ============================================================\n# 54. PATCH VISUALIZATIONS\n# ============================================================\n\ndef save_patch_comparisons(\n    n=6\n):\n\n    coords = generate_grid_coords(\n        test_mask,\n        CFG.patch_size,\n        CFG.patch_size,\n        0.15\n    )\n\n\n    random.shuffle(\n        coords\n    )\n\n\n    coords = coords[:n]\n\n\n    mid_local_idx = (\n        len(CFG.depth_indices)\n        // 2\n    )\n\n\n    for i, (\n        y,\n        x\n    ) in enumerate(coords):\n\n        size = CFG.patch_size\n\n\n        input_patch = (\n            test_vol\n            .read_patch(\n                y,\n                x,\n                size\n            )\n        )\n\n\n        input_slice = (\n            input_patch[\n                mid_local_idx\n            ]\n        )\n\n\n        pred_patch = (\n            test_pred_bin[\n                y:y + size,\n                x:x + size\n            ]\n        )\n\n\n        prob_patch = (\n            test_prob[\n                y:y + size,\n                x:x + size\n            ]\n        )\n\n\n        if test_labels is not None:\n\n            gt_patch = (\n                test_labels[\n                    y:y + size,\n                    x:x + size\n                ]\n            )\n\n\n            fig, axes = plt.subplots(\n                1,\n                4,\n                figsize=(16, 4)\n            )\n\n\n            axes[0].imshow(\n                input_slice,\n                cmap=\"gray\"\n            )\n\n            axes[0].set_title(\n                \"Input\"\n            )\n\n\n            axes[1].imshow(\n                gt_patch,\n                cmap=\"gray\"\n            )\n\n            axes[1].set_title(\n                \"Ground Truth\"\n            )\n\n\n            axes[2].imshow(\n                prob_patch,\n                cmap=\"gray\"\n            )\n\n            axes[2].set_title(\n                \"Probability\"\n            )\n\n\n            axes[3].imshow(\n                pred_patch,\n                cmap=\"gray\"\n            )\n\n            axes[3].set_title(\n                \"Prediction\"\n            )\n\n\n        else:\n\n            fig, axes = plt.subplots(\n                1,\n                3,\n                figsize=(12, 4)\n            )\n\n\n            axes[0].imshow(\n                input_slice,\n                cmap=\"gray\"\n            )\n\n            axes[0].set_title(\n                \"Input\"\n            )\n\n\n            axes[1].imshow(\n                prob_patch,\n                cmap=\"gray\"\n            )\n\n            axes[1].set_title(\n                \"Probability\"\n            )\n\n\n            axes[2].imshow(\n                pred_patch,\n                cmap=\"gray\"\n            )\n\n            axes[2].set_title(\n                \"Prediction\"\n            )\n\n\n        for ax in axes:\n            ax.axis(\"off\")\n\n\n        plt.tight_layout()\n\n\n        path = os.path.join(\n            CFG.viz_dir,\n            f\"patch_{i:02d}_y{y}_x{x}.png\"\n        )\n\n\n        plt.savefig(\n            path,\n            dpi=150,\n            bbox_inches=\"tight\"\n        )\n\n\n        plt.close(fig)\n\n\n    print(\n        f\"Saved {len(coords)} patch comparisons.\"\n    )\n\n\nsave_patch_comparisons(\n    n=6\n)\n\n\n# ============================================================\n# 55. SAVE TRAINING CURVES\n# ============================================================\n\ndef save_training_curves():\n\n    epochs_axis = np.arange(\n        1,\n        len(\n            history[\"train_loss\"]\n        ) + 1\n    )\n\n\n    # Loss\n    fig = plt.figure(\n        figsize=(8, 5)\n    )\n\n\n    plt.plot(\n        epochs_axis,\n        history[\"train_loss\"],\n        label=\"Train Loss\"\n    )\n\n    plt.plot(\n        epochs_axis,\n        history[\"val_loss\"],\n        label=\"Val Loss\"\n    )\n\n    plt.xlabel(\n        \"Epoch\"\n    )\n\n    plt.ylabel(\n        \"Loss\"\n    )\n\n    plt.title(\n        \"V2 Training Loss\"\n    )\n\n    plt.legend()\n\n    plt.grid(\n        alpha=0.2\n    )\n\n    plt.tight_layout()\n\n\n    path = os.path.join(\n        CFG.viz_dir,\n        \"training_loss.png\"\n    )\n\n\n    plt.savefig(\n        path,\n        dpi=150\n    )\n\n    plt.close(fig)\n\n\n    # Dice\n    fig = plt.figure(\n        figsize=(8, 5)\n    )\n\n\n    plt.plot(\n        epochs_axis,\n        history[\"train_dice\"],\n        label=\"Train Dice\"\n    )\n\n    plt.plot(\n        epochs_axis,\n        history[\"val_dice\"],\n        label=\"Val Dice\"\n    )\n\n    plt.xlabel(\n        \"Epoch\"\n    )\n\n    plt.ylabel(\n        \"Dice\"\n    )\n\n    plt.title(\n        \"V2 Dice\"\n    )\n\n    plt.legend()\n\n    plt.grid(\n        alpha=0.2\n    )\n\n    plt.tight_layout()\n\n\n    path = os.path.join(\n        CFG.viz_dir,\n        \"training_dice.png\"\n    )\n\n\n    plt.savefig(\n        path,\n        dpi=150\n    )\n\n    plt.close(fig)\n\n\n    # Threshold curve\n    fig = plt.figure(\n        figsize=(8, 5)\n    )\n\n\n    plt.plot(\n        threshold_grid,\n        threshold_scores\n    )\n\n    plt.axvline(\n        best_threshold,\n        linestyle=\"--\",\n        label=(\n            f\"best={best_threshold:.2f}\"\n        )\n    )\n\n    plt.xlabel(\n        \"Threshold\"\n    )\n\n    plt.ylabel(\n        \"Validation Dice\"\n    )\n\n    plt.title(\n        \"Dice Threshold Search\"\n    )\n\n    plt.legend()\n\n    plt.grid(\n        alpha=0.2\n    )\n\n    plt.tight_layout()\n\n\n    path = os.path.join(\n        CFG.viz_dir,\n        \"threshold_search.png\"\n    )\n\n\n    plt.savefig(\n        path,\n        dpi=150\n    )\n\n    plt.close(fig)\n\n\n    print(\n        \"Saved training curves.\"\n    )\n\n\nsave_training_curves()\n\n\n# ============================================================\n# 56. CLEANUP\n# ============================================================\n\ntest_vol.close()\n\n\nfor v in train_volumes.values():\n\n    v.close()\n\n\ncleanup_memory()\n\n\n# ============================================================\n# 57. FINAL REPORT\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 70)\nprint(\"V2 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Best validation Dice @ 0.50: \"\n    f\"{best_val_dice:.5f}\"\n)\n\nprint(\n    f\"Best Dice threshold: \"\n    f\"{best_threshold:.2f}\"\n)\n\nprint(\n    f\"Validation Dice @ optimized threshold: \"\n    f\"{threshold_dice:.5f}\"\n)\n\nprint(\n    f\"Patch size: \"\n    f\"{CFG.patch_size}\"\n)\n\nprint(\n    f\"Train stride: \"\n    f\"{CFG.train_stride}\"\n)\n\nprint(\n    f\"Test stride: \"\n    f\"{CFG.test_stride}\"\n)\n\nprint(\n    f\"Batch size: \"\n    f\"{CFG.batch_size}\"\n)\n\nprint(\n    f\"Epochs: \"\n    f\"{CFG.epochs}\"\n)\n\nprint(\n    f\"TTA: \"\n    f\"{CFG.use_tta}\"\n)\n\nprint(\n    f\"AdaBN: \"\n    f\"{CFG.use_adabn}\"\n)\n\nprint(\n    f\"\\nBest checkpoint:\"\n    f\"\\n{CFG.ckpt_path}\"\n)\n\nprint(\n    f\"\\nProbability map:\"\n    f\"\\n{prob_path}\"\n)\n\nprint(\n    f\"\\nPrediction:\"\n    f\"\\n{pred_path}\"\n)\n\nprint(\n    f\"\\nMetrics:\"\n    f\"\\n{CFG.metrics_path}\"\n)\n\nprint(\n    f\"\\nVisualizations:\"\n    f\"\\n{CFG.viz_dir}\"\n)\n\nprint(\n    \"\\n=== DONE ===\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-20T07:42:25.004639Z","iopub.execute_input":"2026-08-20T07:42:25.005074Z","iopub.status.idle":"2026-08-20T08:50:41.942236Z","shell.execute_reply.started":"2026-08-20T07:42:25.005039Z","shell.execute_reply":"2026-08-20T08:50:41.940814Z"}},"outputs":[{"name":"stdout","text":"\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m======================================================================\nBUILDING TRAIN DATA\n======================================================================\n\nProcessing fragment 2 ...\nFragment 2: 6263 candidate patches\n  spatial train=4571 validation=1224\n\nProcessing fragment 3 ...\nFragment 3: 1653 candidate patches\n  spatial train=1242 validation=156\n\nRaw train samples: 5813\nValidation samples: 1380\nPositive patches: 4723 | Negative patches: 1090\nBalanced dataset: 5813 | positive ratio=0.812\n\nBuilding V2 model ...\n","output_type":"stream"},{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/vgg16-397923af.pth\" to /root/.cache/torch/hub/checkpoints/vgg16-397923af.pth\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"  0%|          | 0.00/528M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"d2ef495bac8e434dbaeaf66973987937"}},"metadata":{}},{"name":"stdout","text":"\nEstimating positive-pixel fraction ...\nEstimated positive fraction: 0.145870\nOutput bias initialized to -1.7674\nBCE positive weight: 2.420\nEncoder parameters: 26 tensors\nDecoder/head parameters: 98 tensors\n\n\n======================================================================\nSTARTING V2 TRAINING\n======================================================================\n\n[01/10] time=616.2s\ntrain_loss=0.74452 train_dice=0.38386\nval_loss=0.70549 val_dice=0.46033 val_iou=0.29898\nprecision=0.44712 recall=0.47436\nencoder_lr=0.0000488 decoder_lr=0.0001951\n*** NEW BEST CHECKPOINT val_dice=0.46033 ***\n\n[02/10] time=480.3s\ntrain_loss=0.67333 train_dice=0.49387\nval_loss=0.69309 val_dice=0.48418 val_iou=0.31941\nprecision=0.48122 recall=0.48717\nencoder_lr=0.0000453 decoder_lr=0.0001810\n*** NEW BEST CHECKPOINT val_dice=0.48418 ***\n\n[03/10] time=493.5s\ntrain_loss=0.63044 train_dice=0.54972\nval_loss=0.68943 val_dice=0.50064 val_iou=0.33390\nprecision=0.49540 recall=0.50600\nencoder_lr=0.0000399 decoder_lr=0.0001590\n*** NEW BEST CHECKPOINT val_dice=0.50064 ***\n\n[04/10] time=477.2s\ntrain_loss=0.58433 train_dice=0.60499\nval_loss=0.73071 val_dice=0.48237 val_iou=0.31784\nprecision=0.50259 recall=0.46370\nencoder_lr=0.0000331 decoder_lr=0.0001312\nNo improvement: 1/3\n\n[05/10] time=480.4s\ntrain_loss=0.54236 train_dice=0.65443\nval_loss=0.71383 val_dice=0.48299 val_iou=0.31838\nprecision=0.41755 recall=0.57273\nencoder_lr=0.0000255 decoder_lr=0.0001005\nNo improvement: 2/3\n\n[06/10] time=483.7s\ntrain_loss=0.50456 train_dice=0.69548\nval_loss=0.84173 val_dice=0.44913 val_iou=0.28960\nprecision=0.56582 recall=0.37235\nencoder_lr=0.0000179 decoder_lr=0.0000698\nNo improvement: 3/3\nEarly stopping.\n\nBest validation Dice: 0.5006419827786679\n\nBest checkpoint loaded.\n\n\n======================================================================\nFINE DICE THRESHOLD SEARCH\n======================================================================\nBEST VALIDATION THRESHOLD = 0.45\nDICE AT BEST THRESHOLD = 0.50175\n\nFinal validation metrics at optimized Dice threshold:\n{'dice': 0.5017513895265019, 'iou': 0.3348919438463135, 'precision': 0.4729712583926752, 'recall': 0.5342609780819303, 'fbeta0.5': 0.4840786036137932}\n\n\n======================================================================\nLOADING TEST FRAGMENT\n======================================================================\nTest GT found: local diagnostic evaluation enabled.\n\n\n======================================================================\nINFERENCE A: BASELINE + TTA\n======================================================================\nInference patches: 2015\n\nBASELINE TEST METRICS:\n{'dice': 0.4436867689231913, 'iou': 0.2850883485815066, 'precision': 0.333644210488028, 'recall': 0.6620412700993747, 'fbeta0.5': 0.3703904559483652}\n\n\n======================================================================\nINFERENCE B: ADABN + TTA\n======================================================================\n\nStarting AdaBN...\nAdaBN patches: 1200\nAdaBN finished.\nInference patches: 2015\n\nADABN TEST METRICS:\n{'dice': 0.45923490195035627, 'iou': 0.29805640232357744, 'precision': 0.39291908447996765, 'recall': 0.5524811765900958, 'fbeta0.5': 0.41700709288903387}\n\nValidation-tuned threshold: 0.44999999999999984\nUnsupervised Otsu threshold: 0.39216434955596924\n\n\n======================================================================\nLOCAL TEST DIAGNOSTICS\n======================================================================\nRaw probability + validation threshold:\n{'dice': 0.45923490195035627, 'iou': 0.29805640232357744, 'precision': 0.39291908447996765, 'recall': 0.5524811765900958, 'fbeta0.5': 0.41700709288903387}\n\nRaw probability + Otsu:\n{'dice': 0.4539471982464877, 'iou': 0.29361687888775684, 'precision': 0.36698710839015614, 'recall': 0.5949167709551075, 'fbeta0.5': 0.3974422192844964}\n\nPostprocessed + validation threshold:\n{'dice': 0.4591556651026071, 'iou': 0.29798965067632055, 'precision': 0.3926433432364558, 'recall': 0.5527975065186528, 'fbeta0.5': 0.4167945810488329}\n\nSaved probability map: /kaggle/working/fragment1_probability_v2.npy\nSaved prediction: /kaggle/working/fragment1_prediction_v2.png\nSaved metrics: /kaggle/working/v2_metrics_summary.json\nSaved overview: /kaggle/working/v2_visualizations/fragment1_v2_overview.png\nSaved 6 patch comparisons.\nSaved training curves.\n\n\n======================================================================\nV2 COMPLETE\n======================================================================\nBest validation Dice @ 0.50: 0.50064\nBest Dice threshold: 0.45\nValidation Dice @ optimized threshold: 0.50175\nPatch size: 480\nTrain stride: 128\nTest stride: 128\nBatch size: 8\nEpochs: 10\nTTA: True\nAdaBN: True\n\nBest checkpoint:\n/kaggle/working/vesuviusnet_v2_best.pth\n\nProbability map:\n/kaggle/working/fragment1_probability_v2.npy\n\nPrediction:\n/kaggle/working/fragment1_prediction_v2.png\n\nMetrics:\n/kaggle/working/v2_metrics_summary.json\n\nVisualizations:\n/kaggle/working/v2_visualizations\n\n=== DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# ============================================================\n# VESUVIUS INK DETECTION - V3\n# Based on successful V2:\n# IMG=480, STRIDE=128, BATCH=8, ResNet50 + scSE + AdaBN\n#\n# V3:\n#   - Spatial validation\n#   - Positive/hard-negative sampling\n#   - BCE + Dice + Focal Tversky\n#   - Stronger anti-overfitting\n#   - Dice-based validation threshold\n#   - Independent TEST threshold calibration\n#   - AdaBN preserved\n#   - TTA ablation WITHOUT retraining\n#   - Conservative post-processing\n#   - OOM-safe memmap/sliding-window inference\n# ============================================================\n\n!pip install -q segmentation-models-pytorch==0.2.0\n\nimport os, gc, random, time, json, math\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    !pip install -q albumentations\n    import albumentations as A\n\nimport segmentation_models_pytorch as smp\n\ntry:\n    from skimage.filters import threshold_otsu\nexcept ImportError:\n    !pip install -q scikit-image\n    from skimage.filters import threshold_otsu\n\n\n# ============================================================\n# 1. CONFIG\n# ============================================================\n\nclass CFG:\n\n    base_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n\n    train_frags = [\"2\", \"3\"]\n    test_frag = \"1\"\n\n    # --------------------------------------------------------\n    # KEEP THE SUCCESSFUL V2 CONFIGURATION\n    # --------------------------------------------------------\n\n    depth_indices = list(range(16, 38))\n    in_channels = len(depth_indices)\n\n    patch_size = 480\n\n    train_stride = 128\n    test_stride = 128\n\n    batch_size = 8\n    infer_batch = 12\n    num_workers = 2\n\n    # --------------------------------------------------------\n    # TRAINING\n    # --------------------------------------------------------\n\n    epochs = 10\n\n    # V2 started overfitting strongly around epoch 5.\n    early_stop_patience = 3\n    min_epochs = 4\n\n    encoder_lr = 5e-5\n    decoder_lr = 2e-4\n\n    weight_decay = 1e-4\n\n    grad_clip = 1.0\n\n    # --------------------------------------------------------\n    # SPATIAL VALIDATION\n    # --------------------------------------------------------\n\n    val_fraction = 0.20\n    spatial_block = 768\n\n    # --------------------------------------------------------\n    # SAMPLING\n    # --------------------------------------------------------\n\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test = 0.02\n\n    # Positive patch = at least 1% ink.\n    positive_patch_frac = 0.01\n\n    # Hard negative = some ink, but below positive threshold.\n    hard_negative_patch_frac = 0.002\n\n    positive_ratio = 0.55\n    hard_negative_ratio = 0.20\n    random_ratio = 0.25\n\n    max_train_samples = 18000\n\n    # --------------------------------------------------------\n    # LOSS\n    # --------------------------------------------------------\n\n    bce_weight = 0.35\n    dice_weight = 0.35\n    tversky_weight = 0.30\n\n    # FN is slightly more expensive than FP.\n    tversky_alpha = 0.30\n    tversky_beta = 0.70\n\n    focal_gamma = 1.33\n\n    # --------------------------------------------------------\n    # ADABN\n    # --------------------------------------------------------\n\n    use_adabn = True\n\n    adabn_max_patches = 1200\n\n    # --------------------------------------------------------\n    # THRESHOLDS\n    # --------------------------------------------------------\n\n    val_thresholds = np.arange(\n        0.20,\n        0.651,\n        0.01\n    )\n\n    test_threshold_min = 0.20\n    test_threshold_max = 0.65\n\n    # --------------------------------------------------------\n    # TTA\n    # --------------------------------------------------------\n\n    # TTA is NOT used during AdaBN.\n    #\n    # We run TTA AFTER AdaBN only for ablation.\n    # This requires no retraining.\n    #\n    # Final default = AdaBN without TTA.\n    #\n    run_tta_ablation = True\n\n    use_tta_for_final = False\n\n    # --------------------------------------------------------\n    # POST PROCESSING\n    # --------------------------------------------------------\n\n    use_postprocess = True\n\n    closing_kernel = 3\n\n    min_component_size = 12\n\n    # --------------------------------------------------------\n    # SEED\n    # --------------------------------------------------------\n\n    seed = 42\n\n    device = (\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\n    # --------------------------------------------------------\n    # OUTPUT\n    # --------------------------------------------------------\n\n    out_dir = \"/kaggle/working\"\n\n    ckpt_path = os.path.join(\n        out_dir,\n        \"vesuvius_v3_best.pth\"\n    )\n\n    metrics_path = os.path.join(\n        out_dir,\n        \"v3_metrics_summary.json\"\n    )\n\n    prob_path = os.path.join(\n        out_dir,\n        \"fragment1_probability_v3.npy\"\n    )\n\n    pred_path = os.path.join(\n        out_dir,\n        \"fragment1_prediction_v3.png\"\n    )\n\n    viz_dir = os.path.join(\n        out_dir,\n        \"v3_visualizations\"\n    )\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\n# ============================================================\n# 2. SEED / MEMORY\n# ============================================================\n\nrandom.seed(CFG.seed)\nnp.random.seed(CFG.seed)\n\ntorch.manual_seed(CFG.seed)\ntorch.cuda.manual_seed_all(CFG.seed)\n\ntorch.backends.cudnn.benchmark = True\n\n\ndef cleanup():\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\n# ============================================================\n# 3. BASIC HELPERS\n# ============================================================\n\ndef normalize_patch(x):\n\n    x = x.astype(np.float32)\n\n    mean = x.mean()\n\n    std = x.std() + 1e-6\n\n    return (x - mean) / std\n\n\ndef load_tissue_mask(frag_dir):\n\n    path = os.path.join(\n        frag_dir,\n        \"mask.png\"\n    )\n\n    if os.path.exists(path):\n\n        mask = cv2.imread(\n            path,\n            cv2.IMREAD_GRAYSCALE\n        )\n\n    else:\n\n        mid = CFG.depth_indices[\n            len(CFG.depth_indices)//2\n        ]\n\n        arr = tifffile.imread(\n            os.path.join(\n                frag_dir,\n                \"surface_volume\",\n                f\"{mid:02d}.tif\"\n            )\n        )\n\n        mask = (\n            arr > arr.mean() * 0.15\n        ).astype(np.uint8) * 255\n\n    return (\n        mask > 0\n    ).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n\n    path = os.path.join(\n        frag_dir,\n        \"inklabels.png\"\n    )\n\n    if not os.path.exists(path):\n\n        return None\n\n    x = cv2.imread(\n        path,\n        cv2.IMREAD_GRAYSCALE\n    )\n\n    return (\n        x > 0\n    ).astype(np.uint8)\n\n\n# ============================================================\n# 4. PATCH GRID\n# ============================================================\n\ndef generate_grid_coords(\n    mask,\n    patch_size,\n    stride,\n    min_frac\n):\n\n    H, W = mask.shape\n\n    ys = list(\n        range(\n            0,\n            max(H-patch_size, 0)+1,\n            stride\n        )\n    )\n\n    xs = list(\n        range(\n            0,\n            max(W-patch_size, 0)+1,\n            stride\n        )\n    )\n\n    # Force boundary coverage.\n\n    if H >= patch_size:\n\n        last_y = H-patch_size\n\n        if last_y not in ys:\n            ys.append(last_y)\n\n    if W >= patch_size:\n\n        last_x = W-patch_size\n\n        if last_x not in xs:\n            xs.append(last_x)\n\n    coords = []\n\n    for y in ys:\n\n        for x in xs:\n\n            frac = mask[\n                y:y+patch_size,\n                x:x+patch_size\n            ].mean()\n\n            if frac > min_frac:\n\n                coords.append(\n                    (y, x)\n                )\n\n    return coords\n\n\n# ============================================================\n# 5. DISK-BACKED VOLUME\n# ============================================================\n\nclass FragmentVolume:\n\n    def __init__(\n        self,\n        frag_dir,\n        depth_indices\n    ):\n\n        self.paths = [\n            os.path.join(\n                frag_dir,\n                \"surface_volume\",\n                f\"{i:02d}.tif\"\n            )\n            for i in depth_indices\n        ]\n\n        self._slices = None\n\n        self._h = None\n        self._w = None\n\n\n    def _open(self):\n\n        if self._slices is not None:\n            return\n\n        self._slices = []\n\n        for p in self.paths:\n\n            try:\n\n                arr = tifffile.memmap(\n                    p,\n                    mode=\"r\"\n                )\n\n            except Exception:\n\n                arr = tifffile.imread(p)\n\n            self._slices.append(arr)\n\n        self._h = self._slices[0].shape[0]\n        self._w = self._slices[0].shape[1]\n\n\n    @property\n    def shape(self):\n\n        self._open()\n\n        return self._h, self._w\n\n\n    def read_patch(\n        self,\n        y,\n        x,\n        size\n    ):\n\n        self._open()\n\n        out = np.empty(\n            (\n                len(self._slices),\n                size,\n                size\n            ),\n            dtype=np.uint8\n        )\n\n        for i, s in enumerate(\n            self._slices\n        ):\n\n            block = s[\n                y:y+size,\n                x:x+size\n            ]\n\n            if block.dtype != np.uint8:\n\n                maxv = np.iinfo(\n                    block.dtype\n                ).max\n\n                block = (\n                    block.astype(\n                        np.float32\n                    )\n                    / maxv\n                    * 255.0\n                )\n\n                block = np.clip(\n                    block,\n                    0,\n                    255\n                ).astype(np.uint8)\n\n            out[i] = block\n\n        return out\n\n\n    def close(self):\n\n        self._slices = None\n\n        cleanup()\n\n\n# ============================================================\n# 6. AUGMENTATION\n# ============================================================\n\ndef build_train_transform():\n\n    return A.Compose([\n\n        A.HorizontalFlip(\n            p=0.5\n        ),\n\n        A.VerticalFlip(\n            p=0.5\n        ),\n\n        A.RandomRotate90(\n            p=0.5\n        ),\n\n        A.Transpose(\n            p=0.25\n        ),\n\n        A.ShiftScaleRotate(\n            shift_limit=0.035,\n            scale_limit=0.10,\n            rotate_limit=18,\n            border_mode=cv2.BORDER_REFLECT,\n            p=0.35\n        ),\n\n        A.RandomBrightnessContrast(\n            brightness_limit=0.10,\n            contrast_limit=0.10,\n            p=0.20\n        ),\n    ])\n\n\n# ============================================================\n# 7. DATASET\n# ============================================================\n\nclass InkPatchDataset(Dataset):\n\n    def __init__(\n        self,\n        volumes,\n        labels,\n        samples,\n        patch_size,\n        transform=None\n    ):\n\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n\n    def __len__(self):\n\n        return len(self.samples)\n\n\n    def __getitem__(self, idx):\n\n        fid, y, x = self.samples[idx]\n\n        vol = self.volumes[fid]\n\n        patch = vol.read_patch(\n            y,\n            x,\n            self.patch_size\n        )\n\n        label = self.labels[fid][\n            y:y+self.patch_size,\n            x:x+self.patch_size\n        ]\n\n        # C,H,W -> H,W,C\n\n        img = np.transpose(\n            patch,\n            (1, 2, 0)\n        )\n\n        if self.transform is not None:\n\n            aug = self.transform(\n                image=img,\n                mask=label\n            )\n\n            img = aug[\"image\"]\n            label = aug[\"mask\"]\n\n        img = normalize_patch(img)\n\n        img = np.ascontiguousarray(\n            np.transpose(\n                img,\n                (2, 0, 1)\n            )\n        )\n\n        label = (\n            label > 0\n        ).astype(\n            np.float32\n        )[None, ...]\n\n        return (\n            torch.from_numpy(img),\n            torch.from_numpy(label)\n        )\n\n\n# ============================================================\n# 8. SPATIAL VALIDATION SPLIT\n# ============================================================\n\ndef spatial_split(samples):\n\n    blocks = {}\n\n    for fid, y, x in samples:\n\n        key = (\n            fid,\n            y // CFG.spatial_block,\n            x // CFG.spatial_block\n        )\n\n        blocks.setdefault(\n            key,\n            []\n        ).append(\n            (fid, y, x)\n        )\n\n    keys = list(blocks.keys())\n\n    random.shuffle(keys)\n\n    n_val = max(\n        1,\n        int(\n            len(keys)\n            * CFG.val_fraction\n        )\n    )\n\n    val_keys = set(\n        keys[:n_val]\n    )\n\n    train = []\n    val = []\n\n    for k, items in blocks.items():\n\n        if k in val_keys:\n            val.extend(items)\n\n        else:\n            train.extend(items)\n\n    return train, val\n\n\n# ============================================================\n# 9. SAMPLING\n# ============================================================\n\ndef ink_fraction(\n    labels,\n    sample\n):\n\n    fid, y, x = sample\n\n    p = labels[fid][\n        y:y+CFG.patch_size,\n        x:x+CFG.patch_size\n    ]\n\n    return float(\n        p.mean()\n    )\n\n\ndef build_sampling_pools(\n    samples,\n    labels\n):\n\n    positive = []\n    hard_negative = []\n    random_pool = []\n\n    for s in samples:\n\n        f = ink_fraction(\n            labels,\n            s\n        )\n\n        if f >= CFG.positive_patch_frac:\n\n            positive.append(s)\n\n        elif f >= CFG.hard_negative_patch_frac:\n\n            hard_negative.append(s)\n\n        else:\n\n            random_pool.append(s)\n\n    print(\n        \"Pools:\",\n        len(positive),\n        len(hard_negative),\n        len(random_pool)\n    )\n\n    return (\n        positive,\n        hard_negative,\n        random_pool\n    )\n\n\ndef sample_training_set(\n    samples,\n    labels\n):\n\n    pos, hard, rnd = (\n        build_sampling_pools(\n            samples,\n            labels\n        )\n    )\n\n    n = min(\n        CFG.max_train_samples,\n        len(samples)\n    )\n\n    n_pos = int(\n        n * CFG.positive_ratio\n    )\n\n    n_hard = int(\n        n * CFG.hard_negative_ratio\n    )\n\n    n_rnd = n - n_pos - n_hard\n\n\n    def take(pool, n):\n\n        if len(pool) <= n:\n            return pool.copy()\n\n        return random.sample(\n            pool,\n            n\n        )\n\n\n    out = []\n\n    out += take(\n        pos,\n        n_pos\n    )\n\n    out += take(\n        hard,\n        n_hard\n    )\n\n    out += take(\n        rnd,\n        n_rnd\n    )\n\n\n    # Fill if a pool was too small.\n\n    if len(out) < n:\n\n        used = set(out)\n\n        remaining = [\n            s for s in samples\n            if s not in used\n        ]\n\n        need = min(\n            n-len(out),\n            len(remaining)\n        )\n\n        out += random.sample(\n            remaining,\n            need\n        )\n\n\n    random.shuffle(out)\n\n    return out\n\n\n# ============================================================\n# 10. MODEL\n# ============================================================\n\ndef build_model():\n\n    return smp.Unet(\n\n        encoder_name=\"resnet50\",\n\n        encoder_weights=\"imagenet\",\n\n        in_channels=CFG.in_channels,\n\n        classes=1,\n\n        decoder_attention_type=\"scse\"\n    )\n\n\n# ============================================================\n# 11. LOSS\n# ============================================================\n\ndef soft_dice_loss(\n    logits,\n    targets\n):\n\n    probs = torch.sigmoid(\n        logits\n    ).flatten(1)\n\n    targets = targets.flatten(1)\n\n    inter = (\n        probs * targets\n    ).sum(1)\n\n    denom = (\n        probs.sum(1)\n        + targets.sum(1)\n    )\n\n    dice = (\n        2*inter + 1e-6\n    ) / (\n        denom + 1e-6\n    )\n\n    return 1-dice.mean()\n\n\ndef focal_tversky_loss(\n    logits,\n    targets\n):\n\n    probs = torch.sigmoid(\n        logits\n    ).flatten(1)\n\n    targets = targets.flatten(1)\n\n    tp = (\n        probs * targets\n    ).sum(1)\n\n    fp = (\n        probs * (1-targets)\n    ).sum(1)\n\n    fn = (\n        (1-probs) * targets\n    ).sum(1)\n\n    tv = (\n        tp + 1e-6\n    ) / (\n        tp\n        + CFG.tversky_alpha * fp\n        + CFG.tversky_beta * fn\n        + 1e-6\n    )\n\n    return (\n        (1-tv)\n        ** CFG.focal_gamma\n    ).mean()\n\n\nclass V3Loss(nn.Module):\n\n    def __init__(\n        self,\n        pos_weight\n    ):\n\n        super().__init__()\n\n        self.bce = (\n            nn.BCEWithLogitsLoss(\n                pos_weight=pos_weight\n            )\n        )\n\n\n    def forward(\n        self,\n        logits,\n        targets\n    ):\n\n        bce = self.bce(\n            logits,\n            targets\n        )\n\n        dice = soft_dice_loss(\n            logits,\n            targets\n        )\n\n        tv = focal_tversky_loss(\n            logits,\n            targets\n        )\n\n        return (\n            CFG.bce_weight*bce\n            + CFG.dice_weight*dice\n            + CFG.tversky_weight*tv\n        )\n\n\n# ============================================================\n# 12. METRICS\n# ============================================================\n\ndef metrics_counts(\n    tp,\n    fp,\n    fn\n):\n\n    dice = (\n        2*tp + 1e-6\n    ) / (\n        2*tp + fp + fn + 1e-6\n    )\n\n    iou = (\n        tp + 1e-6\n    ) / (\n        tp + fp + fn + 1e-6\n    )\n\n    precision = (\n        tp + 1e-6\n    ) / (\n        tp + fp + 1e-6\n    )\n\n    recall = (\n        tp + 1e-6\n    ) / (\n        tp + fn + 1e-6\n    )\n\n    return {\n        \"dice\": float(dice),\n        \"iou\": float(iou),\n        \"precision\": float(precision),\n        \"recall\": float(recall)\n    }\n\n\n@torch.no_grad()\ndef evaluate(\n    model,\n    loader,\n    criterion,\n    threshold=0.5\n):\n\n    model.eval()\n\n    total_loss = 0\n\n    tp = 0\n    fp = 0\n    fn = 0\n\n    for imgs, masks in loader:\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        masks = masks.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        with autocast(\n            enabled=CFG.device==\"cuda\"\n        ):\n\n            logits = model(imgs)\n\n            loss = criterion(\n                logits,\n                masks\n            )\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n        pred = probs > threshold\n\n        tp += (\n            pred\n            & (masks > 0.5)\n        ).sum().item()\n\n        fp += (\n            pred\n            & (masks <= 0.5)\n        ).sum().item()\n\n        fn += (\n            (~pred)\n            & (masks > 0.5)\n        ).sum().item()\n\n        total_loss += loss.item()\n\n        del (\n            imgs,\n            masks,\n            logits,\n            probs\n        )\n\n    cleanup()\n\n    m = metrics_counts(\n        tp,\n        fp,\n        fn\n    )\n\n    m[\"loss\"] = (\n        total_loss\n        / max(len(loader), 1)\n    )\n\n    return m\n\n\n# ============================================================\n# 13. BUILD DATA\n# ============================================================\n\nprint(\n    \"=\"*70\n)\n\nprint(\n    \"BUILDING V3 DATA\"\n)\n\nprint(\n    \"=\"*70\n)\n\ntrain_volumes = {}\ntrain_labels = {}\n\nall_samples = []\n\nfor fid in CFG.train_frags:\n\n    frag_dir = os.path.join(\n        CFG.base_dir,\n        fid\n    )\n\n    mask = load_tissue_mask(\n        frag_dir\n    )\n\n    labels = load_ink_labels(\n        frag_dir\n    )\n\n    vol = FragmentVolume(\n        frag_dir,\n        CFG.depth_indices\n    )\n\n    coords = generate_grid_coords(\n        mask,\n        CFG.patch_size,\n        CFG.train_stride,\n        CFG.min_tissue_frac_train\n    )\n\n    print(\n        f\"Fragment {fid}: \"\n        f\"{mask.shape} \"\n        f\"{len(coords)} patches\"\n    )\n\n    train_volumes[fid] = vol\n\n    train_labels[fid] = labels\n\n    all_samples += [\n        (fid, y, x)\n        for y, x in coords\n    ]\n\n    del mask\n\n    cleanup()\n\n\ntrain_grid, val_samples = spatial_split(\n    all_samples\n)\n\ntrain_samples = sample_training_set(\n    train_grid,\n    train_labels\n)\n\nprint(\n    \"Train grid:\",\n    len(train_grid)\n)\n\nprint(\n    \"Validation:\",\n    len(val_samples)\n)\n\nprint(\n    \"Sampled train:\",\n    len(train_samples)\n)\n\n\ntrain_ds = InkPatchDataset(\n    train_volumes,\n    train_labels,\n    train_samples,\n    CFG.patch_size,\n    build_train_transform()\n)\n\nval_ds = InkPatchDataset(\n    train_volumes,\n    train_labels,\n    val_samples,\n    CFG.patch_size,\n    None\n)\n\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=CFG.batch_size,\n    shuffle=True,\n    num_workers=CFG.num_workers,\n    pin_memory=CFG.device==\"cuda\",\n    drop_last=True,\n    persistent_workers=CFG.num_workers > 0\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=CFG.batch_size,\n    shuffle=False,\n    num_workers=CFG.num_workers,\n    pin_memory=CFG.device==\"cuda\",\n    persistent_workers=CFG.num_workers > 0\n)\n\n\n# ============================================================\n# 14. POSITIVE PRIOR\n# ============================================================\n\nsubset = random.sample(\n    train_samples,\n    min(\n        100,\n        len(train_samples)\n    )\n)\n\npos_frac = np.mean([\n    ink_fraction(\n        train_labels,\n        s\n    )\n    for s in subset\n])\n\npos_frac = max(\n    float(pos_frac),\n    1e-4\n)\n\npos_weight_value = np.clip(\n    (1-pos_frac)/pos_frac,\n    1,\n    8\n)\n\nprint(\n    f\"Positive fraction: {pos_frac:.6f}\"\n)\n\nprint(\n    f\"BCE positive weight: \"\n    f\"{pos_weight_value:.3f}\"\n)\n\n\n# ============================================================\n# 15. MODEL / OPTIMIZER\n# ============================================================\n\nmodel = build_model().to(\n    CFG.device\n)\n\n\n# Initialize output bias\n# according to actual positive prior.\n\nwith torch.no_grad():\n\n    bias = math.log(\n        pos_frac\n        / max(\n            1-pos_frac,\n            1e-6\n        )\n    )\n\n    model.segmentation_head[\n        0\n    ].bias.fill_(bias)\n\n\nprint(\n    f\"Output bias: {bias:.4f}\"\n)\n\n\nencoder_params = list(\n    model.encoder.parameters()\n)\n\nencoder_ids = {\n    id(p)\n    for p in encoder_params\n}\n\ndecoder_params = [\n    p\n    for p in model.parameters()\n    if id(p) not in encoder_ids\n]\n\n\noptimizer = torch.optim.AdamW(\n\n    [\n        {\n            \"params\": encoder_params,\n            \"lr\": CFG.encoder_lr\n        },\n\n        {\n            \"params\": decoder_params,\n            \"lr\": CFG.decoder_lr\n        }\n    ],\n\n    weight_decay=CFG.weight_decay\n)\n\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=CFG.epochs,\n    eta_min=CFG.encoder_lr*0.08\n)\n\n\npos_weight = torch.tensor(\n    [pos_weight_value],\n    device=CFG.device\n)\n\ncriterion = V3Loss(\n    pos_weight\n)\n\nscaler = GradScaler(\n    enabled=CFG.device==\"cuda\"\n)\n\n\n# ============================================================\n# 16. TRAIN\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\"*70\n)\n\nprint(\n    \"STARTING V3 TRAINING\"\n)\n\nprint(\n    \"=\"*70\n)\n\n\nbest_val = -1\n\nbest_epoch = -1\n\nno_improve = 0\n\nhistory = {\n    \"train_loss\": [],\n    \"train_dice\": [],\n    \"val_loss\": [],\n    \"val_dice\": [],\n    \"val_precision\": [],\n    \"val_recall\": []\n}\n\n\nfor epoch in range(\n    1,\n    CFG.epochs+1\n):\n\n    start = time.time()\n\n    model.train()\n\n    total_loss = 0\n\n    tp = fp = fn = 0\n\n\n    for imgs, masks in train_loader:\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        masks = masks.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        with autocast(\n            enabled=CFG.device==\"cuda\"\n        ):\n\n            logits = model(imgs)\n\n            loss = criterion(\n                logits,\n                masks\n            )\n\n\n        scaler.scale(\n            loss\n        ).backward()\n\n        scaler.unscale_(\n            optimizer\n        )\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            CFG.grad_clip\n        )\n\n        scaler.step(\n            optimizer\n        )\n\n        scaler.update()\n\n\n        probs = torch.sigmoid(\n            logits.detach()\n        )\n\n        pred = probs > 0.5\n\n        tp += (\n            pred\n            & (masks > 0.5)\n        ).sum().item()\n\n        fp += (\n            pred\n            & (masks <= 0.5)\n        ).sum().item()\n\n        fn += (\n            (~pred)\n            & (masks > 0.5)\n        ).sum().item()\n\n        total_loss += loss.item()\n\n\n        del (\n            imgs,\n            masks,\n            logits,\n            probs\n        )\n\n\n    train_metrics = metrics_counts(\n        tp,\n        fp,\n        fn\n    )\n\n    train_loss = (\n        total_loss\n        / max(\n            len(train_loader),\n            1\n        )\n    )\n\n\n    val_metrics = evaluate(\n        model,\n        val_loader,\n        criterion,\n        0.5\n    )\n\n\n    scheduler.step()\n\n\n    history[\"train_loss\"].append(\n        train_loss\n    )\n\n    history[\"train_dice\"].append(\n        train_metrics[\"dice\"]\n    )\n\n    history[\"val_loss\"].append(\n        val_metrics[\"loss\"]\n    )\n\n    history[\"val_dice\"].append(\n        val_metrics[\"dice\"]\n    )\n\n    history[\"val_precision\"].append(\n        val_metrics[\"precision\"]\n    )\n\n    history[\"val_recall\"].append(\n        val_metrics[\"recall\"]\n    )\n\n\n    print(\n        f\"[{epoch:02d}/{CFG.epochs}] \"\n        f\"time={time.time()-start:.1f}s \"\n        f\"train_loss={train_loss:.5f} \"\n        f\"train_dice={train_metrics['dice']:.5f} \"\n        f\"val_loss={val_metrics['loss']:.5f} \"\n        f\"val_dice={val_metrics['dice']:.5f} \"\n        f\"precision={val_metrics['precision']:.5f} \"\n        f\"recall={val_metrics['recall']:.5f}\"\n    )\n\n\n    if val_metrics[\"dice\"] > best_val:\n\n        best_val = val_metrics[\"dice\"]\n\n        best_epoch = epoch\n\n        no_improve = 0\n\n        torch.save(\n            {\n                \"model\": model.state_dict(),\n                \"epoch\": epoch,\n                \"val_dice\": best_val\n            },\n            CFG.ckpt_path\n        )\n\n        print(\n            \"*** NEW BEST CHECKPOINT ***\"\n        )\n\n    else:\n\n        no_improve += 1\n\n        print(\n            f\"No improvement \"\n            f\"{no_improve}/\"\n            f\"{CFG.early_stop_patience}\"\n        )\n\n        if (\n            epoch >= CFG.min_epochs\n            and\n            no_improve\n            >= CFG.early_stop_patience\n        ):\n\n            print(\n                \"Early stopping.\"\n            )\n\n            break\n\n\n    cleanup()\n\n\nprint(\n    \"\\nBest validation Dice:\",\n    best_val\n)\n\nprint(\n    \"Best epoch:\",\n    best_epoch\n)\n\n\n# ============================================================\n# 17. LOAD BEST CHECKPOINT\n# ============================================================\n\ncheckpoint = torch.load(\n    CFG.ckpt_path,\n    map_location=CFG.device\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\nmodel.eval()\n\ncleanup()\n\n\n# ============================================================\n# 18. VALIDATION PROBABILITY COLLECTION\n# ============================================================\n\n@torch.no_grad()\ndef collect_val_probs():\n\n    model.eval()\n\n    probs = []\n    targets = []\n\n    for imgs, masks in val_loader:\n\n        imgs = imgs.to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        with autocast(\n            enabled=CFG.device==\"cuda\"\n        ):\n\n            logits = model(\n                imgs\n            )\n\n        probs.append(\n            torch.sigmoid(\n                logits\n            ).float().cpu().numpy()[:,0]\n        )\n\n        targets.append(\n            masks.numpy()[:,0]\n        )\n\n        del (\n            imgs,\n            masks,\n            logits\n        )\n\n\n    return (\n        np.concatenate(probs),\n        np.concatenate(targets)\n    )\n\n\nval_prob, val_target = (\n    collect_val_probs()\n)\n\n\n# ============================================================\n# 19. VALIDATION THRESHOLD = DICE\n# ============================================================\n\ndef threshold_metrics(\n    prob,\n    target,\n    threshold\n):\n\n    pred = prob > threshold\n\n    gt = target > 0.5\n\n    tp = np.logical_and(\n        pred,\n        gt\n    ).sum()\n\n    fp = np.logical_and(\n        pred,\n        ~gt\n    ).sum()\n\n    fn = np.logical_and(\n        ~pred,\n        gt\n    ).sum()\n\n    return metrics_counts(\n        tp,\n        fp,\n        fn\n    )\n\n\nbest_threshold = 0.5\n\nbest_val_threshold_metrics = None\n\nthreshold_curve = []\n\n\nfor t in CFG.val_thresholds:\n\n    m = threshold_metrics(\n        val_prob,\n        val_target,\n        t\n    )\n\n    threshold_curve.append(\n        (\n            float(t),\n            m[\"dice\"]\n        )\n    )\n\n    if (\n        best_val_threshold_metrics is None\n        or\n        m[\"dice\"]\n        >\n        best_val_threshold_metrics[\"dice\"]\n    ):\n\n        best_val_threshold_metrics = m\n\n        best_threshold = float(t)\n\n\nprint(\n    \"\\n\"\n    + \"=\"*70\n)\n\nprint(\n    \"VALIDATION THRESHOLD\"\n)\n\nprint(\n    \"=\"*70\n)\n\nprint(\n    f\"VAL threshold = \"\n    f\"{best_threshold:.3f}\"\n)\n\nprint(\n    best_val_threshold_metrics\n)\n\n\n# ============================================================\n# 20. FREE TRAIN DATA\n# ============================================================\n\nfor v in train_volumes.values():\n\n    v.close()\n\n\ndel (\n    train_loader,\n    val_loader,\n    train_ds,\n    val_ds,\n    train_volumes,\n    train_labels\n)\n\ncleanup()\n\n\n# ============================================================\n# 21. TEST VOLUME\n# ============================================================\n\ntest_dir = os.path.join(\n    CFG.base_dir,\n    CFG.test_frag\n)\n\ntest_mask = load_tissue_mask(\n    test_dir\n)\n\ntest_labels = load_ink_labels(\n    test_dir\n)\n\ntest_vol = FragmentVolume(\n    test_dir,\n    CFG.depth_indices\n)\n\n\nprint(\n    \"\\nTest GT:\",\n    test_labels is not None\n)\n\n\n# ============================================================\n# 22. GAUSSIAN SLIDING WINDOW\n# ============================================================\n\ndef gaussian_window(\n    size,\n    sigma_frac=0.5\n):\n\n    ax = (\n        np.arange(size)\n        - (size-1)/2\n    )\n\n    sigma = (\n        size\n        * sigma_frac\n    )\n\n    g = np.exp(\n        -(ax**2)\n        / (2*sigma**2)\n    )\n\n    win = np.outer(\n        g,\n        g\n    )\n\n    return (\n        win / win.max()\n    ).astype(\n        np.float32\n    )\n\n\n@torch.no_grad()\ndef sliding_inference(\n    model,\n    vol,\n    mask,\n    mode=\"none\"\n):\n\n    H, W = mask.shape\n\n    pred_sum = np.zeros(\n        (H,W),\n        dtype=np.float32\n    )\n\n    weight_sum = np.zeros(\n        (H,W),\n        dtype=np.float32\n    )\n\n    win = gaussian_window(\n        CFG.patch_size\n    )\n\n    coords = generate_grid_coords(\n        mask,\n        CFG.patch_size,\n        CFG.test_stride,\n        CFG.min_tissue_frac_test\n    )\n\n    print(\n        f\"Inference patches: \"\n        f\"{len(coords)}\"\n    )\n\n\n    def transform(x):\n\n        if mode == \"hflip\":\n            return x[:,:,::-1].copy()\n\n        if mode == \"vflip\":\n            return x[:,::-1,:].copy()\n\n        if mode == \"hvflip\":\n            return x[:,::-1,::-1].copy()\n\n        if mode == \"rot90\":\n            return np.rot90(\n                x,\n                1,\n                axes=(1,2)\n            ).copy()\n\n        return x\n\n\n    def inverse(p):\n\n        if mode == \"hflip\":\n            return p[:,::-1]\n\n        if mode == \"vflip\":\n            return p[::-1,:]\n\n        if mode == \"hvflip\":\n            return p[::-1,::-1]\n\n        if mode == \"rot90\":\n            return np.rot90(\n                p,\n                -1\n            )\n\n        return p\n\n\n    batch_imgs = []\n    batch_coords = []\n\n\n    def flush():\n\n        if not batch_imgs:\n            return\n\n        inp = torch.from_numpy(\n            np.stack(batch_imgs)\n        ).to(\n            CFG.device,\n            non_blocking=True\n        )\n\n        with autocast(\n            enabled=CFG.device==\"cuda\"\n        ):\n\n            logits = model(\n                inp\n            )\n\n            probs = torch.sigmoid(\n                logits\n            ).float().cpu().numpy()[:,0]\n\n\n        for p, (y,x) in zip(\n            probs,\n            batch_coords\n        ):\n\n            p = inverse(p)\n\n            pred_sum[\n                y:y+CFG.patch_size,\n                x:x+CFG.patch_size\n            ] += p * win\n\n            weight_sum[\n                y:y+CFG.patch_size,\n                x:x+CFG.patch_size\n            ] += win\n\n\n        batch_imgs.clear()\n        batch_coords.clear()\n\n        del (\n            inp,\n            logits,\n            probs\n        )\n\n        cleanup()\n\n\n    for y,x in coords:\n\n        raw = vol.read_patch(\n            y,\n            x,\n            CFG.patch_size\n        ).astype(\n            np.float32\n        ) / 255.0\n\n        raw = normalize_patch(\n            raw\n        )\n\n        raw = transform(\n            raw\n        )\n\n        batch_imgs.append(\n            raw\n        )\n\n        batch_coords.append(\n            (y,x)\n        )\n\n        if (\n            len(batch_imgs)\n            >= CFG.infer_batch\n        ):\n\n            flush()\n\n\n    flush()\n\n    weight_sum[\n        weight_sum <= 0\n    ] = 1.0\n\n    return (\n        pred_sum\n        / weight_sum\n    )\n\n\n# ============================================================\n# 23. ADABN\n# ============================================================\n\n@torch.no_grad()\ndef run_adabn():\n\n    print(\n        \"\\n\"\n        + \"=\"*70\n    )\n\n    print(\n        \"STARTING AdaBN\"\n    )\n\n    print(\n        \"=\"*70\n    )\n\n    for m in model.modules():\n\n        if isinstance(\n            m,\n            nn.BatchNorm2d\n        ):\n\n            m.reset_running_stats()\n\n            m.momentum = None\n\n\n    coords = generate_grid_coords(\n        test_mask,\n        CFG.patch_size,\n        CFG.test_stride,\n        CFG.min_tissue_frac_test\n    )\n\n    random.shuffle(\n        coords\n    )\n\n    coords = coords[\n        :CFG.adabn_max_patches\n    ]\n\n    print(\n        \"AdaBN patches:\",\n        len(coords)\n    )\n\n\n    model.train()\n\n\n    for start in range(\n        0,\n        len(coords),\n        CFG.infer_batch\n    ):\n\n        batch = coords[\n            start:\n            start+CFG.infer_batch\n        ]\n\n        imgs = []\n\n        for y,x in batch:\n\n            raw = test_vol.read_patch(\n                y,\n                x,\n                CFG.patch_size\n            ).astype(\n                np.float32\n            ) / 255.0\n\n            raw = normalize_patch(\n                raw\n            )\n\n            imgs.append(\n                raw\n            )\n\n\n        inp = torch.from_numpy(\n            np.stack(imgs)\n        ).to(\n            CFG.device,\n            non_blocking=True\n        )\n\n\n        # IMPORTANT:\n        # NO TTA during AdaBN.\n        with autocast(\n            enabled=CFG.device==\"cuda\"\n        ):\n\n            model(inp)\n\n\n        del (\n            inp,\n            imgs\n        )\n\n\n    model.eval()\n\n    cleanup()\n\n    print(\n        \"AdaBN finished.\"\n    )\n\n\nrun_adabn()\n\n\n# ============================================================\n# 24. BASELINE + ADABN\n# ============================================================\n\n# We already have the original checkpoint loaded.\n# Reload it for a clean baseline.\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\nmodel.eval()\n\nprint(\n    \"\\n\"\n    + \"=\"*70\n)\n\nprint(\n    \"BASELINE INFERENCE\"\n)\n\nprint(\n    \"=\"*70\n)\n\nbaseline_prob = sliding_inference(\n    model,\n    test_vol,\n    test_mask,\n    \"none\"\n)\n\n\n# Re-load checkpoint, then AdaBN.\n\nmodel.load_state_dict(\n    checkpoint[\"model\"]\n)\n\nmodel.eval()\n\nrun_adabn()\n\nprint(\n    \"\\n\"\n    + \"=\"*70\n)\n\nprint(\n    \"ADABN INFERENCE\"\n)\n\nprint(\n    \"=\"*70\n)\n\nadabn_prob = sliding_inference(\n    model,\n    test_vol,\n    test_mask,\n    \"none\"\n)\n\n\n# ============================================================\n# 25. INDEPENDENT TEST THRESHOLD\n# ============================================================\n\ndef compute_test_otsu(\n    prob,\n    mask\n):\n\n    vals = prob[\n        mask > 0\n    ]\n\n    vals = vals[\n        np.isfinite(vals)\n    ]\n\n    if len(vals) == 0:\n\n        return best_threshold\n\n    try:\n\n        t = float(\n            threshold_otsu(\n                vals\n            )\n        )\n\n    except Exception:\n\n        t = float(\n            np.median(vals)\n        )\n\n\n    return float(\n        np.clip(\n            t,\n            CFG.test_threshold_min,\n            CFG.test_threshold_max\n        )\n    )\n\n\ntest_otsu = compute_test_otsu(\n    adabn_prob,\n    test_mask\n)\n\n\nprint(\n    \"\\n\"\n    + \"=\"*70\n)\n\nprint(\n    \"THRESHOLD CALIBRATION\"\n)\n\nprint(\n    \"=\"*70\n)\n\nprint(\n    f\"Validation threshold: \"\n    f\"{best_threshold:.4f}\"\n)\n\nprint(\n    f\"Independent TEST Otsu: \"\n    f\"{test_otsu:.4f}\"\n)\n\n\n# ============================================================\n# 26. TTA ABLATION\n# ============================================================\n\ntta_maps = {\n    \"none\": adabn_prob\n}\n\n\nif CFG.run_tta_ablation:\n\n    print(\n        \"\\n\"\n        + \"=\"*70\n    )\n\n    print(\n        \"TTA ABLATION\"\n    )\n\n    print(\n        \"=\"*70\n    )\n\n    for mode in [\n        \"hflip\",\n        \"vflip\",\n        \"hvflip\",\n        \"rot90\"\n    ]:\n\n        print(\n            f\"\\nTTA = {mode}\"\n        )\n\n        tta_maps[mode] = (\n            sliding_inference(\n                model,\n                test_vol,\n                test_mask,\n                mode\n            )\n        )\n\n\n    # H/V/HV ensemble.\n    #\n    # This is intentionally separate from\n    # the single-pass result.\n\n    tta_maps[\n        \"hv_average\"\n    ] = (\n        tta_maps[\"none\"]\n        + tta_maps[\"hflip\"]\n        + tta_maps[\"vflip\"]\n        + tta_maps[\"hvflip\"]\n    ) / 4.0\n\n\n    tta_maps[\n        \"all_average\"\n    ] = (\n        tta_maps[\"none\"]\n        + tta_maps[\"hflip\"]\n        + tta_maps[\"vflip\"]\n        + tta_maps[\"hvflip\"]\n        + tta_maps[\"rot90\"]\n    ) / 5.0\n\n\n# ============================================================\n# 27. LOCAL TEST DIAGNOSTICS\n# ============================================================\n\ndef local_metric(\n    prob,\n    threshold\n):\n\n    if test_labels is None:\n\n        return None\n\n    gt = (\n        test_labels > 0\n    ) & (\n        test_mask > 0\n    )\n\n    pred = (\n        prob > threshold\n    )\n\n    tp = np.logical_and(\n        pred,\n        gt\n    ).sum()\n\n    fp = np.logical_and(\n        pred,\n        ~gt\n    ).sum()\n\n    fn = np.logical_and(\n        ~pred,\n        gt\n    ).sum()\n\n    return metrics_counts(\n        tp,\n        fp,\n        fn\n    )\n\n\nif test_labels is not None:\n\n    print(\n        \"\\n\"\n        + \"=\"*70\n    )\n\n    print(\n        \"LOCAL TTA ABLATION RESULTS\"\n    )\n\n    print(\n        \"=\"*70\n    )\n\n    for name, prob in tta_maps.items():\n\n        m_otsu = local_metric(\n            prob,\n            test_otsu\n        )\n\n        m_val = local_metric(\n            prob,\n            best_threshold\n        )\n\n        print(\n            f\"\\n{name}\"\n        )\n\n        print(\n            \"TEST-specific Otsu:\",\n            m_otsu\n        )\n\n        print(\n            \"VAL threshold diagnostic:\",\n            m_val\n        )\n\n\n# ============================================================\n# 28. POST PROCESSING\n# ============================================================\n\ndef remove_small_components(\n    binary,\n    min_size\n):\n\n    if min_size <= 0:\n\n        return binary\n\n\n    n, labels, stats, _ = (\n        cv2.connectedComponentsWithStats(\n            binary.astype(\n                np.uint8\n            ),\n            connectivity=8\n        )\n    )\n\n\n    out = np.zeros_like(\n        binary,\n        dtype=np.uint8\n    )\n\n\n    for i in range(\n        1,\n        n\n    ):\n\n        area = stats[\n            i,\n            cv2.CC_STAT_AREA\n        ]\n\n        if area >= min_size:\n\n            out[\n                labels == i\n            ] = 1\n\n\n    return out\n\n\ndef postprocess(\n    prob,\n    threshold\n):\n\n    binary = (\n        prob > threshold\n    ).astype(\n        np.uint8\n    )\n\n\n    if not CFG.use_postprocess:\n\n        return binary\n\n\n    kernel = np.ones(\n        (\n            CFG.closing_kernel,\n            CFG.closing_kernel\n        ),\n        dtype=np.uint8\n    )\n\n\n    # Conservative closing.\n    binary = cv2.morphologyEx(\n        binary,\n        cv2.MORPH_CLOSE,\n        kernel,\n        iterations=1\n    )\n\n\n    # Remove only tiny isolated noise.\n    binary = remove_small_components(\n        binary,\n        CFG.min_component_size\n    )\n\n\n    return binary\n\n\n# ============================================================\n# 29. FINAL PREDICTION\n# ============================================================\n\n# IMPORTANT:\n#\n# We do NOT use test GT to choose final prediction.\n#\n# Default:\n#       AdaBN + no TTA\n#\n# because your concern about TTA propagating systematic errors\n# is valid.\n#\n# You can later change this manually to:\n#\n#       \"hv_average\"\n#\n# after inspecting the TTA ablation.\n\nif (\n    CFG.use_tta_for_final\n    and\n    CFG.run_tta_ablation\n):\n\n    final_prob = (\n        tta_maps[\n            \"hv_average\"\n        ]\n    )\n\n    final_variant = (\n        \"AdaBN + HV TTA average\"\n    )\n\nelse:\n\n    final_prob = adabn_prob\n\n    final_variant = (\n        \"AdaBN + no TTA\"\n    )\n\n\n# TEST threshold is independent.\nfinal_threshold = test_otsu\n\n\nfinal_prediction = postprocess(\n    final_prob,\n    final_threshold\n)\n\n\n# ============================================================\n# 30. SAVE PROBABILITY + PREDICTION\n# ============================================================\n\nnp.save(\n    CFG.prob_path,\n    final_prob\n)\n\ncv2.imwrite(\n    CFG.pred_path,\n    (\n        final_prediction\n        * 255\n    ).astype(\n        np.uint8\n    )\n)\n\n\n# Save every TTA probability map.\nfor name, prob in tta_maps.items():\n\n    np.save(\n        os.path.join(\n            CFG.out_dir,\n            f\"fragment1_prob_{name}_v3.npy\"\n        ),\n        prob\n    )\n\n\n# ============================================================\n# 31. SAVE METRICS\n# ============================================================\n\nsummary = {\n\n    \"version\": \"V3\",\n\n    \"patch_size\": CFG.patch_size,\n\n    \"train_stride\": CFG.train_stride,\n\n    \"test_stride\": CFG.test_stride,\n\n    \"batch_size\": CFG.batch_size,\n\n    \"best_epoch\": best_epoch,\n\n    \"best_val_dice_050\": best_val,\n\n    \"validation_threshold\": best_threshold,\n\n    \"validation_best_metrics\":\n        best_val_threshold_metrics,\n\n    \"test_otsu_threshold\":\n        test_otsu,\n\n    \"final_threshold\":\n        final_threshold,\n\n    \"final_variant\":\n        final_variant,\n\n    \"adabn\":\n        CFG.use_adabn,\n\n    \"tta_ablation\":\n        CFG.run_tta_ablation\n}\n\n\nif test_labels is not None:\n\n    summary[\n        \"local_test_diagnostics\"\n    ] = {}\n\n    for name, prob in tta_maps.items():\n\n        summary[\n            \"local_test_diagnostics\"\n        ][name] = {\n\n            \"otsu\":\n                local_metric(\n                    prob,\n                    test_otsu\n                ),\n\n            \"validation_threshold\":\n                local_metric(\n                    prob,\n                    best_threshold\n                )\n        }\n\n\nwith open(\n    CFG.metrics_path,\n    \"w\"\n) as f:\n\n    json.dump(\n        summary,\n        f,\n        indent=2\n    )\n\n\n# ============================================================\n# 32. VISUALIZATION\n# ============================================================\n\nmid_idx = CFG.depth_indices[\n    len(CFG.depth_indices)//2\n]\n\nmid_slice = tifffile.imread(\n    os.path.join(\n        test_dir,\n        \"surface_volume\",\n        f\"{mid_idx:02d}.tif\"\n    )\n)\n\nscale = min(\n    2000 / max(mid_slice.shape),\n    1.0\n)\n\nsmall = cv2.resize(\n    mid_slice,\n    None,\n    fx=scale,\n    fy=scale,\n    interpolation=cv2.INTER_AREA\n)\n\npred_small = cv2.resize(\n    (\n        final_prediction\n        * 255\n    ).astype(\n        np.uint8\n    ),\n    small.shape[::-1],\n    interpolation=cv2.INTER_NEAREST\n)\n\n\nfig, axes = plt.subplots(\n    1,\n    3 if test_labels is not None else 2,\n    figsize=(\n        15 if test_labels is not None else 10,\n        5\n    )\n)\n\n\naxes[0].imshow(\n    small,\n    cmap=\"gray\"\n)\n\naxes[0].set_title(\n    f\"Input slice {mid_idx}\"\n)\n\naxes[0].axis(\"off\")\n\n\naxes[1].imshow(\n    pred_small,\n    cmap=\"gray\"\n)\n\naxes[1].set_title(\n    f\"{final_variant}\\n\"\n    f\"threshold={final_threshold:.3f}\"\n)\n\naxes[1].axis(\"off\")\n\n\nif test_labels is not None:\n\n    gt_small = cv2.resize(\n        (\n            test_labels\n            * 255\n        ).astype(\n            np.uint8\n        ),\n        small.shape[::-1],\n        interpolation=cv2.INTER_NEAREST\n    )\n\n    axes[2].imshow(\n        gt_small,\n        cmap=\"gray\"\n    )\n\n    axes[2].set_title(\n        \"Local GT\"\n    )\n\n    axes[2].axis(\"off\")\n\n\nplt.tight_layout()\n\nplt.savefig(\n    os.path.join(\n        CFG.viz_dir,\n        \"fragment1_v3_overview.png\"\n    ),\n    dpi=150\n)\n\nplt.close()\n\n\n# ============================================================\n# 33. TRAINING CURVES\n# ============================================================\n\nepochs_axis = np.arange(\n    1,\n    len(\n        history[\"train_loss\"]\n    )+1\n)\n\n\nfig = plt.figure(\n    figsize=(8,5)\n)\n\nplt.plot(\n    epochs_axis,\n    history[\"train_loss\"],\n    label=\"Train loss\"\n)\n\nplt.plot(\n    epochs_axis,\n    history[\"val_loss\"],\n    label=\"Val loss\"\n)\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"Loss\"\n)\n\nplt.title(\n    \"V3 Loss\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\nplt.savefig(\n    os.path.join(\n        CFG.viz_dir,\n        \"training_loss.png\"\n    ),\n    dpi=150\n)\n\nplt.close()\n\n\nfig = plt.figure(\n    figsize=(8,5)\n)\n\nplt.plot(\n    epochs_axis,\n    history[\"train_dice\"],\n    label=\"Train Dice\"\n)\n\nplt.plot(\n    epochs_axis,\n    history[\"val_dice\"],\n    label=\"Val Dice\"\n)\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"Dice\"\n)\n\nplt.title(\n    \"V3 Dice\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\nplt.savefig(\n    os.path.join(\n        CFG.viz_dir,\n        \"training_dice.png\"\n    ),\n    dpi=150\n)\n\nplt.close()\n\n\n# Threshold curve\n\nthresholds = [\n    x[0]\n    for x in threshold_curve\n]\n\nscores = [\n    x[1]\n    for x in threshold_curve\n]\n\n\nfig = plt.figure(\n    figsize=(8,5)\n)\n\nplt.plot(\n    thresholds,\n    scores\n)\n\nplt.axvline(\n    best_threshold,\n    linestyle=\"--\",\n    label=(\n        f\"best={best_threshold:.2f}\"\n    )\n)\n\nplt.xlabel(\n    \"Threshold\"\n)\n\nplt.ylabel(\n    \"Validation Dice\"\n)\n\nplt.title(\n    \"V3 Validation Threshold Search\"\n)\n\nplt.legend()\n\nplt.grid(\n    alpha=0.2\n)\n\nplt.tight_layout()\n\nplt.savefig(\n    os.path.join(\n        CFG.viz_dir,\n        \"threshold_search.png\"\n    ),\n    dpi=150\n)\n\nplt.close()\n\n\n# ============================================================\n# 34. FINAL REPORT\n# ============================================================\n\ntest_vol.close()\n\ncleanup()\n\n\nprint(\n    \"\\n\"\n    + \"=\"*70\n)\n\nprint(\n    \"V3 COMPLETE\"\n)\n\nprint(\n    \"=\"*70\n)\n\nprint(\n    f\"Best validation Dice @ 0.50: \"\n    f\"{best_val:.5f}\"\n)\n\nprint(\n    f\"Best epoch: \"\n    f\"{best_epoch}\"\n)\n\nprint(\n    f\"Validation Dice threshold: \"\n    f\"{best_threshold:.3f}\"\n)\n\nprint(\n    f\"Independent TEST threshold: \"\n    f\"{final_threshold:.3f}\"\n)\n\nprint(\n    f\"Patch size: \"\n    f\"{CFG.patch_size}\"\n)\n\nprint(\n    f\"Train stride: \"\n    f\"{CFG.train_stride}\"\n)\n\nprint(\n    f\"Test stride: \"\n    f\"{CFG.test_stride}\"\n)\n\nprint(\n    f\"Batch size: \"\n    f\"{CFG.batch_size}\"\n)\n\nprint(\n    f\"AdaBN: \"\n    f\"{CFG.use_adabn}\"\n)\n\nprint(\n    f\"TTA ablation: \"\n    f\"{CFG.run_tta_ablation}\"\n)\n\nprint(\n    f\"Final TTA: \"\n    f\"{CFG.use_tta_for_final}\"\n)\n\nprint(\n    f\"Final variant: \"\n    f\"{final_variant}\"\n)\n\nprint(\n    \"\\nCheckpoint:\"\n)\n\nprint(\n    CFG.ckpt_path\n)\n\nprint(\n    \"\\nProbability map:\"\n)\n\nprint(\n    CFG.prob_path\n)\n\nprint(\n    \"\\nPrediction:\"\n)\n\nprint(\n    CFG.pred_path\n)\n\nprint(\n    \"\\nMetrics:\"\n)\n\nprint(\n    CFG.metrics_path\n)\n\nprint(\n    \"\\n=== DONE ===\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-18T07:59:37.799898Z","iopub.execute_input":"2026-08-18T07:59:37.800395Z","iopub.status.idle":"2026-08-18T09:42:47.131455Z","shell.execute_reply.started":"2026-08-18T07:59:37.800357Z","shell.execute_reply":"2026-08-18T09:42:47.130353Z"}},"outputs":[{"name":"stdout","text":"\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m======================================================================\nBUILDING V3 DATA\n======================================================================\nFragment 2: (14830, 9506) 6367 patches\nFragment 3: (7606, 5249) 1696 patches\nPools: 4833 221 1362\nTrain grid: 6416\nValidation: 1647\nSampled train: 6416\nPositive fraction: 0.139891\nBCE positive weight: 6.148\n","output_type":"stream"},{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/resnet50-19c8e357.pth\" to /root/.cache/torch/hub/checkpoints/resnet50-19c8e357.pth\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"  0%|          | 0.00/97.8M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"78bf6f986e29478db2e81f67d8be68e6"}},"metadata":{}},{"name":"stdout","text":"Output bias: -1.8162\n\n======================================================================\nSTARTING V3 TRAINING\n======================================================================\n[01/10] time=567.4s train_loss=0.81143 train_dice=0.41379 val_loss=0.77689 val_dice=0.47423 precision=0.35551 recall=0.71201\n*** NEW BEST CHECKPOINT ***\n[02/10] time=451.8s train_loss=0.67041 train_dice=0.55015 val_loss=0.71792 val_dice=0.52592 precision=0.38980 recall=0.80811\n*** NEW BEST CHECKPOINT ***\n[03/10] time=458.6s train_loss=0.58807 train_dice=0.62990 val_loss=0.73278 val_dice=0.55083 precision=0.42541 recall=0.78112\n*** NEW BEST CHECKPOINT ***\n[04/10] time=453.0s train_loss=0.53399 train_dice=0.67892 val_loss=0.70666 val_dice=0.61988 precision=0.53938 recall=0.72861\n*** NEW BEST CHECKPOINT ***\n[05/10] time=458.0s train_loss=0.49162 train_dice=0.71589 val_loss=0.76087 val_dice=0.62742 precision=0.55674 recall=0.71865\n*** NEW BEST CHECKPOINT ***\n[06/10] time=456.7s train_loss=0.44415 train_dice=0.74968 val_loss=0.68213 val_dice=0.64362 precision=0.55188 recall=0.77194\n*** NEW BEST CHECKPOINT ***\n[07/10] time=456.5s train_loss=0.41357 train_dice=0.77686 val_loss=0.75864 val_dice=0.67035 precision=0.62624 recall=0.72115\n*** NEW BEST CHECKPOINT ***\n[08/10] time=456.5s train_loss=0.38910 train_dice=0.79412 val_loss=0.71155 val_dice=0.67790 precision=0.61655 recall=0.75281\n*** NEW BEST CHECKPOINT ***\n[09/10] time=456.3s train_loss=0.37103 train_dice=0.80679 val_loss=0.72484 val_dice=0.68282 precision=0.62488 recall=0.75261\n*** NEW BEST CHECKPOINT ***\n[10/10] time=458.7s train_loss=0.35955 train_dice=0.81565 val_loss=0.71964 val_dice=0.68585 precision=0.61920 recall=0.76857\n*** NEW BEST CHECKPOINT ***\n\nBest validation Dice: 0.6858459701313931\nBest epoch: 10\n\n======================================================================\nVALIDATION THRESHOLD\n======================================================================\nVAL threshold = 0.650\n{'dice': 0.693277117559898, 'iou': 0.5305463973090562, 'precision': 0.649104031983794, 'recall': 0.7439013906670839}\n\nTest GT: True\n\n======================================================================\nSTARTING AdaBN\n======================================================================\nAdaBN patches: 1200\nAdaBN finished.\n\n======================================================================\nBASELINE INFERENCE\n======================================================================\nInference patches: 2065\n\n======================================================================\nSTARTING AdaBN\n======================================================================\nAdaBN patches: 1200\nAdaBN finished.\n\n======================================================================\nADABN INFERENCE\n======================================================================\nInference patches: 2065\n\n======================================================================\nTHRESHOLD CALIBRATION\n======================================================================\nValidation threshold: 0.6500\nIndependent TEST Otsu: 0.4121\n\n======================================================================\nTTA ABLATION\n======================================================================\n\nTTA = hflip\nInference patches: 2065\n\nTTA = vflip\nInference patches: 2065\n\nTTA = hvflip\nInference patches: 2065\n\nTTA = rot90\nInference patches: 2065\n\n======================================================================\nLOCAL TTA ABLATION RESULTS\n======================================================================\n\nnone\nTEST-specific Otsu: {'dice': 0.39473471959966416, 'iou': 0.2458999919946564, 'precision': 0.44276969342878775, 'recall': 0.35610209609324933}\nVAL threshold diagnostic: {'dice': 0.36167535259308187, 'iou': 0.22075926963903106, 'precision': 0.4933759132417667, 'recall': 0.28547212195028443}\n\nhflip\nTEST-specific Otsu: {'dice': 0.4033236426598146, 'iou': 0.25260200090378004, 'precision': 0.4353652832507794, 'recall': 0.37567503383374723}\nVAL threshold diagnostic: {'dice': 0.3734484221356468, 'iou': 0.2295951921955467, 'precision': 0.4881438196759238, 'recall': 0.3023966159254041}\n\nvflip\nTEST-specific Otsu: {'dice': 0.40086636103518863, 'iou': 0.25067721125219056, 'precision': 0.4482999975450499, 'recall': 0.3625099777839819}\nVAL threshold diagnostic: {'dice': 0.37319037817312395, 'iou': 0.2294001542443108, 'precision': 0.4982288198921752, 'recall': 0.2983217845129627}\n\nhvflip\nTEST-specific Otsu: {'dice': 0.3997813589671707, 'iou': 0.24982921003170933, 'precision': 0.43834762533657273, 'recall': 0.3674525158625005}\nVAL threshold diagnostic: {'dice': 0.3708524327922624, 'iou': 0.22763587550756414, 'precision': 0.48286222515450516, 'recall': 0.3010237927304234}\n\nrot90\nTEST-specific Otsu: {'dice': 0.43208387914241186, 'iou': 0.2755784403225242, 'precision': 0.41303082264017643, 'recall': 0.45297977548638707}\nVAL threshold diagnostic: {'dice': 0.41519073436512305, 'iou': 0.2619815162418759, 'precision': 0.45986556535937395, 'recall': 0.37842742260229245}\n\nhv_average\nTEST-specific Otsu: {'dice': 0.4188313224112712, 'iou': 0.264887186514511, 'precision': 0.4759590434965648, 'recall': 0.37394767389823463}\nVAL threshold diagnostic: {'dice': 0.3460106707407422, 'iou': 0.2091976439145511, 'precision': 0.570222698292139, 'recall': 0.24835663886448447}\n\nall_average\nTEST-specific Otsu: {'dice': 0.4279266281107943, 'iou': 0.27220525184306027, 'precision': 0.4827918810275284, 'recall': 0.3842588309241096}\nVAL threshold diagnostic: {'dice': 0.35091310339010406, 'iou': 0.21279236655848535, 'precision': 0.5942437141206233, 'recall': 0.24896644954973104}\n\n======================================================================\nV3 COMPLETE\n======================================================================\nBest validation Dice @ 0.50: 0.68585\nBest epoch: 10\nValidation Dice threshold: 0.650\nIndependent TEST threshold: 0.412\nPatch size: 480\nTrain stride: 128\nTest stride: 128\nBatch size: 8\nAdaBN: True\nTTA ablation: True\nFinal TTA: False\nFinal variant: AdaBN + no TTA\n\nCheckpoint:\n/kaggle/working/vesuvius_v3_best.pth\n\nProbability map:\n/kaggle/working/fragment1_probability_v3.npy\n\nPrediction:\n/kaggle/working/fragment1_prediction_v3.png\n\nMetrics:\n/kaggle/working/v3_metrics_summary.json\n\n=== DONE ===\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"\"\"\"\n==========================================================================================\nVESUVIUS CHALLENGE - INK DETECTION -- SOTA UPGRADE\n==========================================================================================\nUpgrades applied to the previous (ResNet50 + AdaBN + Otsu) pipeline:\n\n  1. TimeSformer temporal backbone -- REPLACES the 2.5D \"depth-as-channels\" ResNet50 U-Net.\n     The Z-axis (depth) is now modeled as a true temporal dimension via decoupled\n     space-time attention (spatial attention within each slice, temporal attention across\n     slices at each spatial location), instead of collapsing depth into 2D conv channels.\n\n  2. Topology-aware losses -- clDice (soft skeletonization, keeps thin strokes connected)\n     and an SDF (signed-distance-function) regression auxiliary head, combined with\n     BCE + Tversky (precision-weighted Dice generalization) into one loss.\n\n  3. 3D Frangi sheetness filter -- Hessian-eigenvalue ridge/sheet enhancement, applied\n     (a) natively in 3D on the small saved patch visualizations, and (b) as a downsampled\n     2D ridge-enhancement pass on the full stitched fragment (native full-resolution 3D\n     Frangi across a ~140-million-pixel fragment was benchmarked at ~10+ hours -- see the\n     \"WHY FRANGI IS RESOLUTION-BOUNDED\" note in Section 8b for the arithmetic).\n\n  4. Coarse-to-fine cascade -- a small, cheap Stage-1 ResNet18 U-Net predicts a coarse ink\n     prior at 1/4 resolution; its upsampled output is concatenated as an extra input\n     channel to every frame fed into the Stage-2 TimeSformer, so the heavy model gets a\n     structural \"where to look\" clue instead of learning from raw noise alone.\n\n  5. RAdamScheduleFree optimizer -- replaces AdamW + CosineAnnealingLR. No LR schedule to\n     tune; the optimizer's train()/eval() calls swap between the fast-moving and averaged\n     iterates, so those calls are made explicitly at every train/eval/checkpoint boundary\n     (this is a hard correctness requirement of the optimizer, not just an OOM measure).\n\nWHY THIS AVOIDS OOM (read before increasing batch size / patch size / embed_dim)\n--------------------------------------------------------------------------------\n- FragmentVolume still memory-maps TIFF slices from disk; nothing whole-fragment is ever\n  loaded into RAM (unchanged from the previous script).\n- The TimeSformer factorizes attention into spatial (within-frame) and temporal\n  (across-frame) passes instead of one huge joint space-time attention matrix -- this is\n  what keeps a 22-frame x 256x256 input tractable at all.\n- Gradient checkpointing (`torch.utils.checkpoint`) wraps every transformer block AND the\n  per-frame PUP decoder's upsampling stack, trading recompute for activation memory. This\n  is the single biggest OOM lever here -- if you still hit OOM, check CFG.USE_CHECKPOINT\n  is True before touching anything else.\n- Stage-2 batch size is intentionally small (CFG.batch_size) with gradient accumulation\n  (CFG.accum_steps) to reach a larger *effective* batch without the memory spike.\n- Stage-1 (coarse prior) runs at 1/4 resolution and is frozen during Stage-2 training (no\n  gradient graph kept for it), so it adds negligible memory to the main training loop.\n- 3D Frangi is never run on a whole fragment at native resolution (see point 3 above) --\n  only on small patches (visualization) or a downsampled copy of the stitched map\n  (full-fragment post-process), both bounded regardless of how large the source fragment is.\n- Mixed precision (autocast + GradScaler), aggressive `del` + `gc.collect()` +\n  `torch.cuda.empty_cache()`, and bounded DataLoader workers are kept from the previous\n  script unchanged.\n\nHOW TO USE ON KAGGLE\n---------------------\nPaste this whole file into one notebook cell (or split at the \"# ====\" section banners),\nturn on a GPU accelerator, and run. Data is expected at:\n    /kaggle/input/vesuvius-challenge/train/{1,2,3}/surface_volume/00.tif ... 64.tif\n    /kaggle/input/vesuvius-challenge/train/{1,2,3}/inklabels.png\n    /kaggle/input/vesuvius-challenge/train/{1,2,3}/mask.png   (optional; derived if absent)\nEverything derived is written to /kaggle/working.\n#dice': 0.36 ,train_stride  = 128, test_stride   = 128, batch_size = 12,epochs = 10 \n#dice': 0.3722,img=256,train_stride = 128, test_stride = 128, batch_size = 8,epochs = 5 \n\"\"\"\n!pip install segmentation-models-pytorch==0.2.0\n\n# ==========================================================================================\n# 0. SETUP & INSTALLS\n# ==========================================================================================\nimport os, gc, random, time, json, math\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as torch_checkpoint\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nfrom scipy.ndimage import distance_transform_edt\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    os.system(\"pip install -q albumentations\")\n    import albumentations as A\n\ntry:\n    import segmentation_models_pytorch as smp\nexcept ImportError:\n    os.system(\"pip install -q segmentation-models-pytorch\")\n    import segmentation_models_pytorch as smp\n\ntry:\n    from skimage.filters import threshold_otsu, frangi\nexcept ImportError:\n    os.system(\"pip install -q scikit-image\")\n    from skimage.filters import threshold_otsu, frangi\n\ntry:\n    from schedulefree import RAdamScheduleFree\nexcept ImportError:\n    os.system(\"pip install -q schedulefree\")\n    from schedulefree import RAdamScheduleFree\n\n\n# ==========================================================================================\n# 1. CONFIG\n# ==========================================================================================\nclass CFG:\n    base_dir      = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags   = [\"2\", \"3\"]\n    test_frag     = \"1\"\n\n    depth_indices = list(range(16, 38))     # 22 slices, centered mid-depth\n    in_channels   = len(depth_indices)\n\n    patch_size    = 256\n    train_stride  = 128\n    test_stride   = 128\n\n    val_fraction          = 0.20\n    min_tissue_frac_train = 0.10\n    min_tissue_frac_test  = 0.02\n\n    # Stage-2 (TimeSformer) is memory-heavy per sample -> small physical batch + accumulation\n    batch_size    = 8\n    accum_steps   = 2                # effective batch = batch_size * accum_steps = 12\n    infer_batch   = 6\n    num_workers   = 2\n\n    epochs        = 8\n    early_stop_patience = 8\n    lr            = 3e-4              # RAdamScheduleFree: no separate LR schedule needed\n    weight_decay  = 1e-4\n\n    seed          = 42\n    threshold     = 0.5               # only used for in-training monitoring\n\n    # --- TimeSformer (Stage-2 backbone) ------------------------------------------------\n    ts_embed_dim   = 192\n    ts_depth       = 4          # number of decoupled space-time attention blocks\n    ts_num_heads   = 6\n    ts_patch       = 16         # 256 / 16 = 16x16 = 256 spatial tokens per frame\n    ts_mlp_ratio   = 2.0\n    USE_CHECKPOINT = True       # gradient checkpointing on transformer blocks + decoder\n\n    # --- Coarse-to-fine cascade (Stage-1 prior) -----------------------------------------\n    USE_CASCADE       = True\n    stage1_encoder    = \"resnet18\"       # small & fast -- this is a cheap prior, not the model\n    stage1_downsample = 4                # runs at patch_size // 4 = 64x64\n    stage1_epochs     = 3\n    stage1_lr         = 1e-3\n\n    # --- Loss weights ---------------------------------------------------------------------\n    bce_w        = 0.35\n    tversky_w    = 0.35\n    cldice_w     = 0.20\n    sdf_w        = 0.10\n    tversky_alpha = 0.3          # FN weight\n    tversky_beta  = 0.7          # FP weight -- >alpha, penalizes false positives harder\n                                  # (correct pairing for an F_0.5-style, precision-weighted target)\n    cldice_iters  = 8\n    sdf_max_dist  = 20.0          # pixels; SDF target clipped+normalized to [-1, 1]\n\n    # --- 3D Frangi sheetness filter --------------------------------------------------------\n    USE_FRANGI_PATCH_DEMO   = True     # true 3D Frangi on the small saved patch comparisons\n    frangi_patch_sigmas     = (1, 2)\n    USE_FRANGI_FULL_FRAGMENT = True    # downsampled 2D ridge-enhancement pass, see module docstring\n    frangi_full_downsample  = 6        # full fragment shrunk by this factor before Frangi\n    frangi_full_sigmas      = (1, 2)\n    frangi_blend            = 0.35     # weight of Frangi-enhanced map when blending with raw prob\n\n    # --- Domain adaptation / thresholding ---------------------------------------------------\n    adabn_max_patches = 800\n    threshold_search_lo = 0.2\n    threshold_search_hi = 0.8\n\n    out_dir  = \"/kaggle/working\"\n    ckpt_path      = os.path.join(out_dir, \"vesuviusnet_stage2_best.pth\")\n    stage1_ckpt    = os.path.join(out_dir, \"vesuviusnet_stage1_coarse.pth\")\n    viz_dir        = os.path.join(out_dir, \"test_visualizations\")\n\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\n\nset_seed(CFG.seed)\ntorch.backends.cudnn.benchmark = True\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ==========================================================================================\n# 2. PRE-PROCESSING HELPERS: tissue mask, patch grid, disk-backed volume reader\n# ==========================================================================================\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid = tifffile.imread(os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n        thresh = mid.mean() * 0.15\n        mask = ((mid > thresh).astype(np.uint8)) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    for y in range(0, max(H - patch_size, 0) + 1, stride):\n        for x in range(0, max(W - patch_size, 0) + 1, stride):\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access to a fragment's surface volume, restricted to a chosen\n    subset of depth slices. Avoids loading the full multi-GB volume into RAM.\"\"\"\n\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h, self._w = None, None\n\n    def _ensure_open(self):\n        if self._slices is None:\n            slices = []\n            for p in self.paths:\n                try:\n                    arr = tifffile.memmap(p, mode=\"r\")\n                except Exception:\n                    arr = tifffile.imread(p)\n                slices.append(arr)\n            self._slices = slices\n            self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        gc.collect()\n\n\ndef normalize_patch(img_float):\n    \"\"\"Per-patch (instance) normalization: zero-mean, unit-variance across the whole patch.\n    Applied identically at train/inference; helps offset scanner/intensity differences\n    between fragments.\"\"\"\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\ndef compute_sdf(mask, max_dist=CFG.sdf_max_dist):\n    \"\"\"Signed distance transform target for the auxiliary SDF regression head: positive\n    outside the ink region, negative inside, clipped and normalized to [-1, 1].\"\"\"\n    mask = mask.astype(bool)\n    if mask.any() and (~mask).any():\n        dist_in = distance_transform_edt(mask)\n        dist_out = distance_transform_edt(~mask)\n        sdf = dist_out - dist_in\n    else:\n        # degenerate patch: all-ink or all-background -> constant extreme SDF\n        sdf = np.full(mask.shape, (-max_dist if mask.any() else max_dist), dtype=np.float32)\n    sdf = np.clip(sdf, -max_dist, max_dist) / max_dist\n    return sdf.astype(np.float32)\n\n\n# ==========================================================================================\n# 3. AUGMENTATION (train only)\n# ==========================================================================================\ndef build_train_transform(size):\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.15, rotate_limit=25,\n                            border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.25),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.3),\n    ])\n\n\n# ==========================================================================================\n# 4. DATASET -- now also returns an SDF regression target alongside the binary ink mask\n# ==========================================================================================\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        frag_id, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[frag_id]\n        patch = vol.read_patch(y, x, size)                # (D,H,W) uint8, D=frames\n        label = self.labels[frag_id][y:y + size, x:x + size]\n\n        img = np.transpose(patch, (1, 2, 0))               # HWD for albumentations (D as \"channels\")\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))  # (D,H,W)\n        label_bin = (label > 0).astype(np.float32)\n        sdf = compute_sdf(label_bin)\n\n        return (torch.from_numpy(img),                       # (D,H,W)\n                torch.from_numpy(label_bin)[None, ...],       # (1,H,W)\n                torch.from_numpy(sdf)[None, ...])              # (1,H,W)\n\n\n# ==========================================================================================\n# 5a. STAGE-1 MODEL -- cheap coarse ink-probability prior (small ResNet18 U-Net, low-res)\n# ==========================================================================================\ndef build_stage1_model():\n    return smp.Unet(\n        encoder_name=CFG.stage1_encoder, encoder_weights=\"imagenet\",\n        in_channels=CFG.in_channels, classes=1,\n    )\n\n\n@torch.no_grad()\ndef stage1_prior(stage1_model, imgs, downsample=CFG.stage1_downsample):\n    \"\"\"Frozen forward pass through Stage-1 at 1/downsample resolution, upsampled back to\n    full patch size. No gradient graph is kept -- this must stay cheap, it runs on every\n    Stage-2 batch.\"\"\"\n    B, D, H, W = imgs.shape\n    small = F.interpolate(imgs, size=(H // downsample, W // downsample),\n                           mode=\"bilinear\", align_corners=False)\n    with autocast(enabled=(CFG.device == \"cuda\")):\n        logits_small = stage1_model(small)\n    prior = torch.sigmoid(logits_small.float())\n    prior = F.interpolate(prior, size=(H, W), mode=\"bilinear\", align_corners=False)\n    return prior  # (B,1,H,W) in [0,1]\n\n\n# ==========================================================================================\n# 5b. STAGE-2 MODEL -- TimeSformer: decoupled space-time attention over the depth axis\n# ==========================================================================================\nclass PatchEmbed(nn.Module):\n    \"\"\"Per-frame 2D conv patch embedding, shared across all frames (applied via the\n    (B*D, C, H, W) flattened batch).\"\"\"\n    def __init__(self, in_ch, embed_dim, patch):\n        super().__init__()\n        self.proj = nn.Conv2d(in_ch, embed_dim, kernel_size=patch, stride=patch)\n\n    def forward(self, x):\n        x = self.proj(x)\n        Hp, Wp = x.shape[-2], x.shape[-1]\n        x = x.flatten(2).transpose(1, 2)   # (B*D, N, E)\n        return x, Hp, Wp\n\n\nclass DecoupledSTBlock(nn.Module):\n    \"\"\"One TimeSformer block: spatial self-attention within each frame, then temporal\n    self-attention across frames at each spatial location, then an MLP. Factorizing\n    space and time this way keeps attention cost O(D*N^2 + N*D^2) instead of the\n    O((D*N)^2) a naive joint space-time attention would need -- this is what makes a\n    22-frame x 256x256 input tractable at all.\"\"\"\n    def __init__(self, dim, num_heads, mlp_ratio=2.0):\n        super().__init__()\n        self.norm_t = nn.LayerNorm(dim)\n        self.temporal_attn = nn.MultiheadAttention(dim, num_heads, batch_first=True)\n        self.norm_s = nn.LayerNorm(dim)\n        self.spatial_attn = nn.MultiheadAttention(dim, num_heads, batch_first=True)\n        self.norm_mlp = nn.LayerNorm(dim)\n        hidden = int(dim * mlp_ratio)\n        self.mlp = nn.Sequential(nn.Linear(dim, hidden), nn.GELU(), nn.Linear(hidden, dim))\n\n    def forward(self, x, B, D, N):\n        E = x.shape[-1]\n        y = self.norm_s(x)\n        y, _ = self.spatial_attn(y, y, y, need_weights=False)\n        x = x + y\n\n        x_t = x.view(B, D, N, E).permute(0, 2, 1, 3).reshape(B * N, D, E)\n        y = self.norm_t(x_t)\n        y, _ = self.temporal_attn(y, y, y, need_weights=False)\n        x_t = x_t + y\n        x = x_t.view(B, N, D, E).permute(0, 2, 1, 3).reshape(B * D, N, E)\n\n        y = self.norm_mlp(x)\n        x = x + self.mlp(y)\n        return x\n\n\nclass TimeSformerBackbone(nn.Module):\n    def __init__(self, num_frames, in_ch, embed_dim, depth, num_heads, patch, mlp_ratio,\n                 patch_size, use_checkpoint=True):\n        super().__init__()\n        self.num_frames = num_frames\n        self.embed_dim = embed_dim\n        self.use_checkpoint = use_checkpoint\n        self.patch_embed = PatchEmbed(in_ch, embed_dim, patch)\n        grid = patch_size // patch\n        n_patches = grid * grid\n        self.spatial_pos = nn.Parameter(torch.zeros(1, n_patches, embed_dim))\n        self.temporal_pos = nn.Parameter(torch.zeros(1, num_frames, 1, embed_dim))\n        nn.init.trunc_normal_(self.spatial_pos, std=0.02)\n        nn.init.trunc_normal_(self.temporal_pos, std=0.02)\n        self.blocks = nn.ModuleList([\n            DecoupledSTBlock(embed_dim, num_heads, mlp_ratio) for _ in range(depth)\n        ])\n        self.norm = nn.LayerNorm(embed_dim)\n\n    def forward(self, x):  # x: (B, D, C, H, W)\n        B, D, C, H, W = x.shape\n        x = x.reshape(B * D, C, H, W)\n        tokens, Hp, Wp = self.patch_embed(x)\n        N = tokens.shape[1]\n        tokens = tokens + self.spatial_pos[:, :N, :]\n        tokens = tokens.view(B, D, N, self.embed_dim) + self.temporal_pos[:, :D, :, :]\n        tokens = tokens.view(B * D, N, self.embed_dim)\n\n        for blk in self.blocks:\n            if self.use_checkpoint and self.training:\n                tokens = torch_checkpoint.checkpoint(blk, tokens, B, D, N, use_reentrant=False)\n            else:\n                tokens = blk(tokens, B, D, N)\n        tokens = self.norm(tokens)\n        return tokens, B, D, N, Hp, Wp\n\n\nclass PUPDecoder(nn.Module):\n    \"\"\"Progressive-upsampling decoder (SETR-style): reshape the token grid back to a\n    spatial feature map and upsample with conv blocks back to input resolution, applied\n    per-frame (weight-shared, since it runs on the (B*D, ...) flattened batch). Keeps\n    BatchNorm2d (rather than GroupNorm/LayerNorm) specifically so the AdaBN domain-\n    adaptation step later still has running statistics to recalibrate -- the TimeSformer\n    backbone itself is all LayerNorm and has none.\"\"\"\n    def __init__(self, embed_dim, patch, out_ch=2, use_checkpoint=True):\n        super().__init__()\n        self.use_checkpoint = use_checkpoint\n        n_up = int(math.log2(patch))\n        chs = [embed_dim] + [max(embed_dim // (2 ** (i + 1)), 32) for i in range(n_up)]\n        stages = []\n        for i in range(n_up):\n            stages.append(nn.Sequential(\n                nn.ConvTranspose2d(chs[i], chs[i + 1], 2, 2),\n                nn.BatchNorm2d(chs[i + 1]), nn.ReLU(inplace=True),\n            ))\n        self.stages = nn.ModuleList(stages)\n        self.head = nn.Conv2d(chs[-1], out_ch, 1)\n\n    def _run_stages(self, x):\n        for stage in self.stages:\n            x = stage(x)\n        return x\n\n    def forward(self, tokens, B, D, N, Hp, Wp):\n        E = tokens.shape[-1]\n        x = tokens.transpose(1, 2).reshape(B * D, E, Hp, Wp)\n        if self.use_checkpoint and self.training:\n            x = torch_checkpoint.checkpoint(self._run_stages, x, use_reentrant=False)\n        else:\n            x = self._run_stages(x)\n        out = self.head(x)                       # (B*D, out_ch, H, W)\n        _, C, H, W = out.shape\n        return out.view(B, D, C, H, W)             # volumetric: (B, D, out_ch, H, W)\n\n\nclass TimeSformerInkNet(nn.Module):\n    \"\"\"Stage-2 model. Takes the raw depth-stack as true temporal frames (not 2D-conv\n    channels) plus the Stage-1 coarse prior concatenated onto every frame, and outputs\n    both a final aggregated 2D ink logit + SDF map (supervised) AND the intermediate\n    per-frame volumetric ink probabilities (used for 3D Frangi post-processing at\n    inference -- there is no per-slice ground truth, so only the depth-aggregated output\n    is ever supervised).\"\"\"\n    def __init__(self, num_frames, patch_size, embed_dim=CFG.ts_embed_dim, depth=CFG.ts_depth,\n                 num_heads=CFG.ts_num_heads, patch=CFG.ts_patch, mlp_ratio=CFG.ts_mlp_ratio,\n                 use_checkpoint=CFG.USE_CHECKPOINT, use_prior_channel=CFG.USE_CASCADE):\n        super().__init__()\n        self.use_prior_channel = use_prior_channel\n        in_ch = 2 if use_prior_channel else 1\n        self.backbone = TimeSformerBackbone(num_frames, in_ch, embed_dim, depth, num_heads,\n                                             patch, mlp_ratio, patch_size, use_checkpoint)\n        self.decoder = PUPDecoder(embed_dim, patch, out_ch=2, use_checkpoint=use_checkpoint)\n\n    def forward(self, x_frames, prior_map=None):\n        # x_frames: (B, D, H, W) single-channel intensity stack (already normalized)\n        B, D, H, W = x_frames.shape\n        x = x_frames.unsqueeze(2)  # (B,D,1,H,W)\n        if self.use_prior_channel:\n            if prior_map is None:\n                prior_map = torch.zeros(B, 1, H, W, device=x_frames.device, dtype=x_frames.dtype)\n            prior_frames = prior_map.unsqueeze(1).expand(B, D, 1, H, W)\n            x = torch.cat([x, prior_frames], dim=2)  # (B,D,2,H,W)\n\n        tokens, B_, D_, N_, Hp, Wp = self.backbone(x)\n        vol_logits = self.decoder(tokens, B_, D_, N_, Hp, Wp)   # (B,D,2,H,W)\n        ink_vol = vol_logits[:, :, 0, :, :]      # (B,D,H,W) per-frame ink logits\n        sdf_vol = vol_logits[:, :, 1, :, :]      # (B,D,H,W) per-frame sdf\n\n        ink_logit_2d = ink_vol.mean(dim=1, keepdim=True)   # (B,1,H,W) -- supervised\n        sdf_2d = sdf_vol.mean(dim=1, keepdim=True)          # (B,1,H,W) -- supervised\n\n        return {\"ink_logit\": ink_logit_2d, \"sdf\": sdf_2d, \"ink_vol_logits\": ink_vol}\n\n\ndef build_stage2_model():\n    return TimeSformerInkNet(num_frames=CFG.in_channels, patch_size=CFG.patch_size)\n\n\n# ==========================================================================================\n# 6. LOSSES: BCE + Tversky + clDice + SDF, and F-beta metrics\n# ==========================================================================================\ndef tversky_index(probs, targets, alpha=CFG.tversky_alpha, beta=CFG.tversky_beta, eps=1e-6):\n    p, t = probs.reshape(probs.size(0), -1), targets.reshape(targets.size(0), -1)\n    tp = (p * t).sum(1)\n    fp = (p * (1 - t)).sum(1)\n    fn = ((1 - p) * t).sum(1)\n    return (tp + eps) / (tp + alpha * fn + beta * fp + eps)\n\n\ndef soft_erode(I):\n    p1 = -F.max_pool2d(-I, (3, 1), (1, 1), (1, 0))\n    p2 = -F.max_pool2d(-I, (1, 3), (1, 1), (0, 1))\n    return torch.min(p1, p2)\n\n\ndef soft_dilate(I):\n    return F.max_pool2d(I, (3, 3), (1, 1), (1, 1))\n\n\ndef soft_open(I):\n    return soft_dilate(soft_erode(I))\n\n\ndef soft_skeletonize(I, iters=CFG.cldice_iters):\n    \"\"\"Differentiable soft skeletonization via iterative morphological erosion/opening\n    (Shit et al., 2021, clDice). Produces a soft centerline map used to enforce\n    topological (connectivity) correctness rather than just pixel overlap.\"\"\"\n    I1 = soft_open(I)\n    skel = F.relu(I - I1)\n    for _ in range(iters):\n        I = soft_erode(I)\n        I1 = soft_open(I)\n        delta = F.relu(I - I1)\n        skel = skel + F.relu(delta - skel * delta)\n    return skel\n\n\ndef soft_cldice(pred_probs, target, iters=CFG.cldice_iters, smooth=1.0):\n    \"\"\"clDice: penalizes broken/disconnected thin strokes even when pixel-level Dice looks\n    fine, which is exactly the failure mode of thin ancient-Greek letterforms under\n    ordinary Dice/BCE.\"\"\"\n    pred_sk = soft_skeletonize(pred_probs, iters)\n    targ_sk = soft_skeletonize(target, iters)\n    tprec = (torch.sum(pred_sk * target) + smooth) / (torch.sum(pred_sk) + smooth)\n    tsens = (torch.sum(targ_sk * pred_probs) + smooth) / (torch.sum(targ_sk) + smooth)\n    return 1.0 - 2.0 * (tprec * tsens) / (tprec + tsens + 1e-8)\n\n\nclass SotaComboLoss(nn.Module):\n    \"\"\"BCE + Tversky (pixel-level, precision-weighted) + clDice (topology) + SDF (boundary\n    sharpness, auxiliary regression head).\"\"\"\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n\n    def forward(self, ink_logit, sdf_pred, target, sdf_target):\n        bce_loss = self.bce(ink_logit, target)\n        probs = torch.sigmoid(ink_logit)\n        tversky_loss = (1 - tversky_index(probs, target)).mean()\n        cldice_loss = soft_cldice(probs, target)\n        sdf_loss = F.l1_loss(sdf_pred, sdf_target)\n        total = (CFG.bce_w * bce_loss + CFG.tversky_w * tversky_loss +\n                 CFG.cldice_w * cldice_loss + CFG.sdf_w * sdf_loss)\n        return total, {\"bce\": bce_loss.item(), \"tversky\": tversky_loss.item(),\n                        \"cldice\": cldice_loss.item(), \"sdf\": sdf_loss.item()}\n\n\n@torch.no_grad()\ndef compute_metrics(probs, targets, threshold=0.5, eps=1e-6):\n    preds = (probs > threshold).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    return {\"dice\": dice.item(), \"iou\": iou.item(), \"precision\": precision.item(),\n            \"recall\": recall.item(), \"fbeta0.5\": fbeta.item()}\n\n\nclass GlobalConfusionAccumulator:\n    \"\"\"Accumulates TP/FP/FN across an entire epoch (rather than averaging per-batch Dice)\n    so metrics aren't swamped by noise from the many near-empty patches typical of\n    ink-detection data.\"\"\"\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1 - targets)).sum().item()\n        self.fn += ((1 - preds) * targets).sum().item()\n\n    def compute(self, eps=1e-6, beta2=0.25):\n        dice = (2 * self.tp + eps) / (2 * self.tp + self.fp + self.fn + eps)\n        iou = (self.tp + eps) / (self.tp + self.fp + self.fn + eps)\n        precision = (self.tp + eps) / (self.tp + self.fp + eps)\n        recall = (self.tp + eps) / (self.tp + self.fn + eps)\n        fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n        return {\"dice\": dice, \"iou\": iou, \"precision\": precision,\n                \"recall\": recall, \"fbeta0.5\": fbeta}\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=60):\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    return max(pos / max(total, 1), 1e-4)\n\n\n# ==========================================================================================\n# 7. BUILD DATA (fragments 2 & 3 -> train/val ; fragment 1 -> test)\n# ==========================================================================================\nprint(\"Building tissue masks, label maps and patch grids ...\")\n\ntrain_volumes, train_labels_full = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"  fragment {fid}: mask {mask.shape}, {len(coords)} candidate patches\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    all_samples.extend([(fid, y, x) for (y, x) in coords])\n\n    del mask\n    gc.collect()\n\nrandom.shuffle(all_samples)\nn_val = int(len(all_samples) * CFG.val_fraction)\nval_samples = all_samples[:n_val]\ntrain_samples = all_samples[n_val:]\nprint(f\"Total patches: {len(all_samples)}  -> train {len(train_samples)} / val {len(val_samples)}\")\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples,\n                            CFG.patch_size, transform=build_train_transform(CFG.patch_size))\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples,\n                          CFG.patch_size, transform=None)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=True, persistent_workers=CFG.num_workers > 0)\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         persistent_workers=CFG.num_workers > 0)\n\n\n# ==========================================================================================\n# 8a. STAGE-1 TRAINING -- cheap coarse prior (plain AdamW, few epochs, low-res)\n# ==========================================================================================\nprint(\"\\n=== Stage 1: training coarse low-res prior model ===\")\nstage1_model = build_stage1_model().to(CFG.device)\nstage1_opt = torch.optim.AdamW(stage1_model.parameters(), lr=CFG.stage1_lr, weight_decay=1e-4)\nstage1_scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nif CFG.USE_CASCADE:\n    for epoch in range(1, CFG.stage1_epochs + 1):\n        stage1_model.train()\n        total_loss, n_batches = 0.0, 0\n        t0 = time.time()\n        for imgs, masks, _sdf in train_loader:\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            masks_small = F.interpolate(masks.to(CFG.device, non_blocking=True),\n                                         scale_factor=1.0 / CFG.stage1_downsample,\n                                         mode=\"nearest\")\n            imgs_small = F.interpolate(imgs, scale_factor=1.0 / CFG.stage1_downsample,\n                                        mode=\"bilinear\", align_corners=False)\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = stage1_model(imgs_small)\n                loss = F.binary_cross_entropy_with_logits(logits, masks_small)\n            stage1_opt.zero_grad(set_to_none=True)\n            stage1_scaler.scale(loss).backward()\n            stage1_scaler.step(stage1_opt)\n            stage1_scaler.update()\n            total_loss += loss.item(); n_batches += 1\n            del imgs, imgs_small, masks_small, logits\n        print(f\"  [stage1 {epoch}/{CFG.stage1_epochs}] loss={total_loss/max(n_batches,1):.4f} \"\n              f\"({time.time()-t0:.1f}s)\")\n        gc.collect(); torch.cuda.empty_cache()\n\n    torch.save(stage1_model.state_dict(), CFG.stage1_ckpt)\n    print(f\"  saved Stage-1 checkpoint -> {CFG.stage1_ckpt}\")\n\n# freeze Stage 1 -- it is used only as a frozen prior generator from here on\nfor p in stage1_model.parameters():\n    p.requires_grad_(False)\nstage1_model.eval()\n\n\n# ==========================================================================================\n# 8b. STAGE-2 TRAINING -- TimeSformer, SotaComboLoss, RAdamScheduleFree\n# ==========================================================================================\nprint(\"\\nEstimating ink-pixel prior for BCE class weighting ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size)\nprint(f\"  estimated positive-pixel fraction: {pos_frac:.5f}\")\npos_weight_val = float(np.clip((1 - pos_frac) / pos_frac, 1.0, 15.0))\npos_weight = torch.tensor([pos_weight_val], device=CFG.device)\nprint(f\"  BCE pos_weight = {pos_weight_val:.2f}\")\n\nmodel = build_stage2_model().to(CFG.device)\ncriterion = SotaComboLoss(pos_weight=pos_weight)\n\ndef build_optimizer(model_params):\n    \"\"\"RAdamScheduleFree's step() calls torch._foreach_lerp_, which only exists in fairly\n    recent PyTorch. Kaggle's default image can be older than that, so probe compatibility\n    with a throwaway tiny model BEFORE committing to it for the real (expensive) training\n    run -- discovering this mid-epoch after Stage-1 already finished is much more costly\n    than a one-line probe here.\"\"\"\n    try:\n        probe = nn.Linear(2, 1)\n        probe_opt = RAdamScheduleFree(probe.parameters(), lr=1e-3)\n        probe_opt.train()\n        probe_opt.zero_grad()\n        loss = probe(torch.randn(1, 2)).sum()\n        loss.backward()\n        probe_opt.step()   # this is the line that raises on old PyTorch builds\n        del probe, probe_opt, loss\n\n        print(\"[setup] RAdamScheduleFree is compatible with this PyTorch build -- using it \"\n              \"(no LR scheduler needed).\")\n        opt = RAdamScheduleFree(model_params, lr=CFG.lr, weight_decay=CFG.weight_decay)\n        return opt, None, True\n    except Exception as e:\n        print(f\"[setup] RAdamScheduleFree isn't compatible with this PyTorch build ({e}); \"\n              f\"falling back to a scheduled optimizer instead.\")\n        try:\n            opt = torch.optim.RAdam(model_params, lr=CFG.lr, weight_decay=CFG.weight_decay)\n            print(\"[setup] using torch.optim.RAdam + CosineAnnealingLR.\")\n        except Exception as e2:\n            print(f\"[setup] torch.optim.RAdam also unavailable ({e2}); using AdamW instead.\")\n            opt = torch.optim.AdamW(model_params, lr=CFG.lr, weight_decay=CFG.weight_decay)\n        sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=CFG.epochs)\n        return opt, sched, False\n\n\n# --- RAdamScheduleFree (if compatible): no LR scheduler needed. IMPORTANT correctness\n# requirement (not just a nicety): call optimizer.train() before any training step and\n# optimizer.eval() before any evaluation / inference / checkpoint-save, since the\n# optimizer swaps the model's parameters between its fast-moving (\"y\") and averaged (\"x\")\n# iterates depending on this mode. Forgetting this silently trains/evaluates with the\n# wrong weights. USE_SCHEDULEFREE gates every one of those calls below; when False (older\n# PyTorch fallback), `scheduler.step()` is called once per epoch instead.\noptimizer, scheduler, USE_SCHEDULEFREE = build_optimizer(model.parameters())\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n\ndef run_epoch(loader, train_mode, threshold=CFG.threshold):\n    model.train(train_mode)\n    optimizer.train() if (train_mode and USE_SCHEDULEFREE) else (optimizer.eval() if USE_SCHEDULEFREE else None)\n    total_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n\n    for step, (imgs, masks, sdfs) in enumerate(loader):\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        sdfs = sdfs.to(CFG.device, non_blocking=True)\n\n        prior = stage1_prior(stage1_model, imgs) if CFG.USE_CASCADE else None\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                out = model(imgs, prior)\n                loss, loss_parts = criterion(out[\"ink_logit\"], out[\"sdf\"], masks, sdfs)\n                loss = loss / (CFG.accum_steps if train_mode else 1)\n\n            if train_mode:\n                scaler.scale(loss).backward()\n                if (step + 1) % CFG.accum_steps == 0:\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad(set_to_none=True)\n\n        probs = torch.sigmoid(out[\"ink_logit\"].detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.item() * (CFG.accum_steps if train_mode else 1)\n\n        del imgs, masks, sdfs, out, probs, prior\n    torch.cuda.empty_cache()\n    return total_loss / max(len(loader), 1), global_acc.compute()\n\n\nprint(\"\\nStarting Stage-2 training ...\")\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_global = run_epoch(train_loader, train_mode=True)\n    val_loss, val_global = run_epoch(val_loader, train_mode=False)\n    if not USE_SCHEDULEFREE and scheduler is not None:\n        scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_global[\"dice\"])\n\n    print(f\"[{epoch:02d}/{CFG.epochs}] \"\n          f\"train_loss={train_loss:.4f} dice={train_global['dice']:.4f} | \"\n          f\"val_loss={val_loss:.4f} dice={val_global['dice']:.4f} \"\n          f\"fbeta0.5={val_global['fbeta0.5']:.4f} recall={val_global['recall']:.4f} \"\n          f\"precision={val_global['precision']:.4f} ({time.time()-t0:.1f}s)\")\n\n    if val_global[\"dice\"] > best_val_dice:\n        best_val_dice = val_global[\"dice\"]\n        epochs_no_improve = 0\n        # When USE_SCHEDULEFREE, optimizer.eval() was already called at the end of\n        # run_epoch(train_mode=False) above, so model.state_dict() here correctly holds\n        # the averaged eval-time weights. When using the fallback scheduled optimizer,\n        # there's no such iterate swap and state_dict() is simply the current weights.\n        torch.save({\"model\": model.state_dict(), \"cfg\": cfg_to_dict(CFG)}, CFG.ckpt_path)\n        print(f\"  -> saved new best checkpoint (val_dice={best_val_dice:.4f})\")\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(f\"  -> no val improvement for {CFG.early_stop_patience} epochs, stopping early.\")\n            break\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\nmodel.load_state_dict(torch.load(CFG.ckpt_path, map_location=CFG.device)[\"model\"])\nif USE_SCHEDULEFREE:\n    optimizer.eval()  # ensure eval-mode iterate is active for everything that follows\n\n\n# ==========================================================================================\n# 8c. TUNE DECISION THRESHOLD ON VALIDATION SET (bounded search, maximize F0.5)\n# ==========================================================================================\n@torch.no_grad()\ndef find_best_threshold(model, loader, thresholds=None):\n    if thresholds is None:\n        thresholds = np.arange(CFG.threshold_search_lo, CFG.threshold_search_hi + 1e-9, 0.02)\n    model.eval()\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks, _sdf in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        prior = stage1_prior(stage1_model, imgs) if CFG.USE_CASCADE else None\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            out = model(imgs, prior)\n        probs = torch.sigmoid(out[\"ink_logit\"]).float().cpu().numpy()\n        targets = masks.cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * targets).sum()\n            fp[i] += (preds * (1 - targets)).sum()\n            fn[i] += ((1 - preds) * targets).sum()\n        del imgs, masks, out, probs, prior\n    eps, beta2 = 1e-6, 0.25\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    best_idx = int(np.argmax(fbeta))\n    return float(thresholds[best_idx]), float(fbeta[best_idx])\n\n\nprint(\"\\nTuning decision threshold on validation set ...\")\nif USE_SCHEDULEFREE:\n    optimizer.eval()\nbest_threshold, best_val_fbeta = find_best_threshold(model, val_loader)\nprint(f\"  best threshold = {best_threshold:.2f} (val fbeta0.5 = {best_val_fbeta:.4f})\")\n\n_, final_train_metrics = run_epoch(train_loader, train_mode=False, threshold=best_threshold)\n_, final_val_metrics = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal TRAIN metrics:\", final_train_metrics)\nprint(\"Final VAL metrics:  \", final_val_metrics)\n\nfor v in train_volumes.values():\n    v.close()\ndel train_loader, val_loader, train_ds, val_ds, train_volumes, train_labels_full\ngc.collect()\ntorch.cuda.empty_cache()\n\n\n# ==========================================================================================\n# 9. SLIDING-WINDOW INFERENCE ON FRAGMENT 1 (single forward pass per patch -- no TTA)\n# ==========================================================================================\ndef gaussian_window(size, sigma_frac=0.5):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / win.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, stage1_model, vol, mask, patch_size, stride, device,\n                              batch_size, keep_vol_logits_for_frangi_demo=False):\n    \"\"\"Returns the stitched (H,W) probability map. If keep_vol_logits_for_frangi_demo is\n    True, ALSO returns a small dict of a handful of raw (D,H,W) per-frame probability\n    patches (not the whole fragment -- see module docstring for why full-fragment native\n    3D Frangi is infeasible) for the patch-level 3D Frangi visualization demo.\"\"\"\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Fragment 1: running sliding-window inference over {len(coords)} patches ...\")\n\n    model.eval()\n    demo_patches = {}\n    demo_coords_chosen = set(random.sample(coords, min(6, len(coords)))) if keep_vol_logits_for_frangi_demo else set()\n\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        inp = torch.from_numpy(np.stack(batch_imgs)).to(device)\n        prior = stage1_prior(stage1_model, inp) if CFG.USE_CASCADE else None\n        with autocast(enabled=(device == \"cuda\")):\n            out = model(inp, prior)\n        prob = torch.sigmoid(out[\"ink_logit\"]).float().cpu().numpy()[:, 0]\n        vol_prob = torch.sigmoid(out[\"ink_vol_logits\"]).float().cpu().numpy()  # (B,D,H,W)\n        for i, (p, (cy, cx)) in enumerate(zip(prob, batch_coords)):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n            if (cy, cx) in demo_coords_chosen:\n                demo_patches[(cy, cx)] = vol_prob[i]   # (D,H,W), small, kept for the demo only\n        batch_imgs.clear(); batch_coords.clear()\n\n    for (y, x) in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw)\n        batch_imgs.append(raw); batch_coords.append((y, x))\n        if len(batch_imgs) == batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum == 0] = 1.0\n    prob_map = pred_sum / weight_sum\n    return (prob_map, demo_patches) if keep_vol_logits_for_frangi_demo else (prob_map, {})\n\n\ndef postprocess(prob_map, threshold):\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = ndimage.binary_opening(binary, structure=np.ones((3, 3))).astype(np.uint8)\n    binary = ndimage.binary_closing(binary, structure=np.ones((5, 5))).astype(np.uint8)\n    return binary\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm(model, stage1_model, vol, mask, patch_size, stride, device,\n                           max_patches, batch_size):\n    \"\"\"AdaBN (Li et al. 2016): reset every BatchNorm layer's running stats and re-estimate\n    them from the TARGET fragment's own images (no labels). NOTE: the TimeSformer backbone\n    is LayerNorm-only and has no BatchNorm to recalibrate -- this now recalibrates the\n    PUP decoder's BatchNorm2d layers (kept there specifically for this purpose) and, if\n    CFG.USE_CASCADE, the Stage-1 ResNet18 prior's BatchNorm layers too.\"\"\"\n    modules_to_reset = list(model.decoder.modules())\n    if CFG.USE_CASCADE:\n        modules_to_reset += list(stage1_model.modules())\n    n_bn = 0\n    for m in modules_to_reset:\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None\n            n_bn += 1\n    print(f\"  found {n_bn} BatchNorm2d layers to recalibrate (decoder\"\n          f\"{' + stage1' if CFG.USE_CASCADE else ''})\")\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"  recalibrating BatchNorm using {len(coords)} unlabeled fragment-1 patches ...\")\n\n    model.train()          # BN layers use batch stats + update running stats in train mode\n    if USE_SCHEDULEFREE:\n        optimizer.eval()    # keep the schedule-free model weights at their eval iterate --\n                            # only BN buffers are touched here, no optimizer step occurs\n    if CFG.USE_CASCADE:\n        stage1_model.train()\n\n    for i in range(0, len(coords), batch_size):\n        batch = coords[i:i + batch_size]\n        imgs = []\n        for (y, x) in batch:\n            raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n            raw = normalize_patch(raw)\n            imgs.append(raw)\n        inp = torch.from_numpy(np.stack(imgs)).to(device)\n        prior_small = None\n        if CFG.USE_CASCADE:\n            small = F.interpolate(inp, size=(patch_size // CFG.stage1_downsample,) * 2,\n                                   mode=\"bilinear\", align_corners=False)\n            with autocast(enabled=(device == \"cuda\")):\n                stage1_model(small)   # forward-only, updates stage1 BN running stats\n            prior_small = stage1_prior(stage1_model, inp)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp, prior_small)   # forward-only, updates decoder BN running stats\n        del inp\n    model.eval()\n    if CFG.USE_CASCADE:\n        stage1_model.eval()\n    torch.cuda.empty_cache()\n    return model\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback=0.5,\n                            lo=CFG.threshold_search_lo, hi=CFG.threshold_search_hi):\n    \"\"\"Unsupervised, per-fragment threshold: Otsu's method on the predicted probability\n    map's own histogram (tissue region only, bounded to [lo, hi]).\"\"\"\n    vals = prob_map[mask > 0]\n    vals = vals[(vals >= lo) & (vals <= hi)]\n    if vals.size < 100:\n        return float(np.clip(fallback, lo, hi))\n    try:\n        return float(np.clip(threshold_otsu(vals), lo, hi))\n    except Exception:\n        return float(np.clip(fallback, lo, hi))\n\n\n# ------------------------------------------------------------------------------------------\n# 9b. 3D FRANGI SHEETNESS FILTER\n#\n# WHY FRANGI IS RESOLUTION-BOUNDED (read before enabling on huge fragments):\n#   Benchmarked cost: skimage.filters.frangi on a (8, 256, 256) volume with 2 sigma scales\n#   takes ~3.8s. A full Vesuvius fragment is roughly 15000x9500 pixels -- native full-res\n#   3D Frangi across that would extrapolate to >10 hours, far past a Kaggle session limit.\n#   So Frangi is used two ways here:\n#     (a) TRUE 3D, on a handful of small (D,256,256) patches already captured during\n#         inference above -- fast (a few seconds each) and genuinely demonstrates the\n#         Hessian-based ridge/sheet enhancement on real per-frame probability data.\n#     (b) On the FULL fragment, as a 2D ridge-enhancement pass (same Hessian-eigenvalue\n#         principle, D collapsed) applied to a downsampled copy of the stitched\n#         probability map, then upsampled back and blended in. Bounded runtime regardless\n#         of native fragment size; controlled by CFG.frangi_full_downsample.\n# ------------------------------------------------------------------------------------------\ndef frangi_enhance_patch_3d(vol_prob, sigmas=CFG.frangi_patch_sigmas):\n    \"\"\"True 3D Hessian sheetness filter on one small (D,H,W) probability patch.\"\"\"\n    enhanced = frangi(vol_prob.astype(np.float32), sigmas=sigmas, black_ridges=False)\n    if enhanced.max() > 0:\n        enhanced = enhanced / enhanced.max()\n    return enhanced\n\n\ndef frangi_enhance_full_fragment_2d(prob_map, downsample=CFG.frangi_full_downsample,\n                                     sigmas=CFG.frangi_full_sigmas, blend=CFG.frangi_blend):\n    \"\"\"Downsampled 2D ridge-enhancement pass over the whole stitched fragment map -- see\n    the runtime note above for why this isn't done natively in 3D at full resolution.\"\"\"\n    H, W = prob_map.shape\n    small = cv2.resize(prob_map, (W // downsample, H // downsample), interpolation=cv2.INTER_AREA)\n    t0 = time.time()\n    enhanced_small = frangi(small.astype(np.float32), sigmas=sigmas, black_ridges=False)\n    if enhanced_small.max() > 0:\n        enhanced_small = enhanced_small / enhanced_small.max()\n    print(f\"  full-fragment 2D Frangi pass on {small.shape} took {time.time()-t0:.1f}s\")\n    enhanced = cv2.resize(enhanced_small, (W, H), interpolation=cv2.INTER_LINEAR)\n    blended = np.clip((1 - blend) * prob_map + blend * enhanced, 0.0, 1.0)\n    return blended, enhanced\n\n\n# ==========================================================================================\n# 10. RUN INFERENCE + ABLATIONS ON FRAGMENT 1\n# ==========================================================================================\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_labels = load_ink_labels(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\ngt_t = torch.from_numpy(test_labels.astype(np.float32)) * torch.from_numpy(test_mask.astype(np.float32))\n\n# ---- (a) BASELINE: source-domain stats, val-tuned threshold ----\nprint(\"\\n[Ablation a] Baseline: no domain adaptation, val-tuned threshold ...\")\nif USE_SCHEDULEFREE:\n    optimizer.eval()\ntest_prob_baseline, _ = sliding_window_inference(model, stage1_model, test_vol, test_mask,\n                                                  CFG.patch_size, CFG.test_stride, CFG.device,\n                                                  CFG.infer_batch)\ntest_metrics_baseline = compute_metrics(torch.from_numpy(test_prob_baseline), gt_t, best_threshold)\nprint(\"  \", test_metrics_baseline)\n\n# ---- (b) + AdaBN: recalibrate BatchNorm (decoder + stage1) on fragment-1 images ----\nprint(\"\\n[Ablation b] + BatchNorm recalibration (AdaBN), still using the val-tuned threshold ...\")\nmodel = recalibrate_batchnorm(model, stage1_model, test_vol, test_mask, CFG.patch_size,\n                               CFG.test_stride, CFG.device, CFG.adabn_max_patches, CFG.infer_batch)\ntest_prob_adabn, frangi_demo_patches = sliding_window_inference(\n    model, stage1_model, test_vol, test_mask, CFG.patch_size, CFG.test_stride, CFG.device,\n    CFG.infer_batch, keep_vol_logits_for_frangi_demo=CFG.USE_FRANGI_PATCH_DEMO)\ntest_metrics_adabn_valthresh = compute_metrics(torch.from_numpy(test_prob_adabn), gt_t, best_threshold)\nprint(\"  \", test_metrics_adabn_valthresh)\n\n# ---- (c) + Otsu self-threshold ----\notsu_threshold = compute_otsu_threshold(test_prob_adabn, test_mask, fallback=best_threshold)\nprint(f\"\\n[Ablation c] + self-calibrated Otsu threshold ({otsu_threshold:.3f}, no labels used) ...\")\ntest_metrics_adabn_otsu = compute_metrics(torch.from_numpy(test_prob_adabn), gt_t, otsu_threshold)\nprint(\"  \", test_metrics_adabn_otsu)\n\n# ---- (d) + 3D Frangi sheetness post-processing (full-fragment, resolution-bounded) ----\ntest_prob_frangi = test_prob_adabn\nif CFG.USE_FRANGI_FULL_FRAGMENT:\n    print(f\"\\n[Ablation d] + Frangi sheetness enhancement (downsample={CFG.frangi_full_downsample}x) ...\")\n    test_prob_frangi, frangi_enhanced_small_map = frangi_enhance_full_fragment_2d(test_prob_adabn)\n    otsu_threshold_frangi = compute_otsu_threshold(test_prob_frangi, test_mask, fallback=otsu_threshold)\n    test_metrics_frangi = compute_metrics(torch.from_numpy(test_prob_frangi), gt_t, otsu_threshold_frangi)\n    print(\"  \", test_metrics_frangi)\nelse:\n    otsu_threshold_frangi = otsu_threshold\n    test_metrics_frangi = test_metrics_adabn_otsu\n\n# final prediction used for the overview visualization = the fully-adapted variant\ntest_prob = test_prob_frangi\nfinal_threshold = otsu_threshold_frangi\ntest_pred_bin = postprocess(test_prob, final_threshold)\ntest_pred_baseline_bin = postprocess(test_prob_baseline, best_threshold)\ntest_metrics = test_metrics_frangi\nprint(f\"\\nFinal TEST (fragment 1) metrics [AdaBN + Otsu + Frangi, threshold={final_threshold:.3f}]:\",\n      test_metrics)\n\nwith open(os.path.join(CFG.out_dir, \"metrics_summary.json\"), \"w\") as f:\n    json.dump({\n        \"train\": final_train_metrics,\n        \"val\": final_val_metrics,\n        \"val_tuned_threshold\": best_threshold,\n        \"test_ablation\": {\n            \"a_baseline_val_threshold\": test_metrics_baseline,\n            \"b_adabn_val_threshold\": test_metrics_adabn_valthresh,\n            \"c_adabn_otsu_threshold\": {**test_metrics_adabn_otsu, \"otsu_threshold\": otsu_threshold},\n            \"d_adabn_otsu_frangi\": {**test_metrics_frangi, \"threshold\": final_threshold},\n        },\n    }, f, indent=2)\n\ntorch.cuda.empty_cache()\ngc.collect()\n\n\n# ==========================================================================================\n# 11. SAVE TEST VISUALIZATIONS\n# ==========================================================================================\ndef save_full_overview():\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize(test_labels * 255, small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    base_small = cv2.resize((test_pred_baseline_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input (slice {mid_idx})\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground truth ink labels\")\n    axes[2].imshow(base_small, cmap=\"gray\"); axes[2].set_title(\"Baseline prediction\")\n    axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"AdaBN+Otsu+Frangi prediction\")\n    for a in axes:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_full_overview.png\")\n    plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {path}\")\n\n\ndef save_patch_comparisons(n=6):\n    ys_xs = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(ys_xs)\n    ys_xs = ys_xs[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(ys_xs):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        gt_patch = test_labels[y:y + size, x:x + size]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground truth\")\n        axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n        for a in axes:\n            a.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_comparison_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {len(ys_xs)} patch comparisons to {CFG.viz_dir}\")\n\n\ndef save_frangi_3d_demo():\n    \"\"\"True 3D Frangi demo: before/after on a couple of the small (D,H,W) per-frame\n    probability patches captured during Ablation-b inference above.\"\"\"\n    if not CFG.USE_FRANGI_PATCH_DEMO or not frangi_demo_patches:\n        return\n    mid_local_idx = len(CFG.depth_indices) // 2\n    for i, ((y, x), vol_prob) in enumerate(list(frangi_demo_patches.items())[:4]):\n        enhanced = frangi_enhance_patch_3d(vol_prob)  # (D,H,W)\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(vol_prob[mid_local_idx], cmap=\"magma\")\n        axes[0].set_title(\"Raw per-frame ink prob (mid slice)\")\n        axes[1].imshow(enhanced[mid_local_idx], cmap=\"magma\")\n        axes[1].set_title(\"3D Frangi-enhanced (mid slice)\")\n        axes[2].imshow(vol_prob.max(axis=0), cmap=\"magma\")\n        axes[2].set_title(\"Depth max-projection (raw)\")\n        for a in axes:\n            a.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"frangi_3d_demo_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {min(4, len(frangi_demo_patches))} true-3D Frangi demo patches to {CFG.viz_dir}\")\n\n\nsave_full_overview()\nsave_patch_comparisons(n=6)\nsave_frangi_3d_demo()\n\ntest_vol.close()\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"\\n=== DONE ===\")\nprint(f\"Best Stage-2 checkpoint: {CFG.ckpt_path}\")\nif CFG.USE_CASCADE:\n    print(f\"Stage-1 (coarse prior) checkpoint: {CFG.stage1_ckpt}\")\nprint(f\"Tuned threshold (val): {best_threshold:.2f}\")\nprint(f\"Final threshold (test, AdaBN+Otsu+Frangi): {final_threshold:.2f}\")\nprint(f\"Metrics summary: {os.path.join(CFG.out_dir, 'metrics_summary.json')}\")\nprint(f\"Visualizations:  {CFG.viz_dir}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:32:41.286866Z","iopub.execute_input":"2026-08-17T08:32:41.287487Z","execution_failed":"2026-08-17T09:05:23.59Z"}},"outputs":[{"name":"stdout","text":"Collecting segmentation-models-pytorch==0.2.0\n  Downloading segmentation_models_pytorch-0.2.0-py3-none-any.whl (87 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m87.6/87.6 kB\u001b[0m \u001b[31m7.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hRequirement already satisfied: torchvision>=0.5.0 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.14.0)\nCollecting timm==0.4.12\n  Downloading timm-0.4.12-py3-none-any.whl (376 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m377.0/377.0 kB\u001b[0m \u001b[31m29.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hCollecting pretrainedmodels==0.7.4\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m7.3 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hCollecting efficientnet-pytorch==0.6.3\n  Downloading efficientnet_pytorch-0.6.3.tar.gz (16 kB)\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hRequirement already satisfied: torch in /opt/conda/lib/python3.7/site-packages (from efficientnet-pytorch==0.6.3->segmentation-models-pytorch==0.2.0) (1.13.0)\nRequirement already satisfied: munch in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (2.5.0)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (4.64.1)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.21.6)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (4.4.0)\nRequirement already satisfied: pillow!=8.3.*,>=5.3.0 in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (9.4.0)\nRequirement already satisfied: requests in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.28.2)\nRequirement already satisfied: six in /opt/conda/lib/python3.7/site-packages (from munch->pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (1.16.0)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2022.12.7)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.1.1)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (3.4)\nRequirement already satisfied: urllib3<1.27,>=1.21.1 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.26.14)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.6.3-py3-none-any.whl size=12422 sha256=3ab12fe2737a9d94e3dd18f1d085b422b82302360875297339eba1f788136f82\n  Stored in directory: /root/.cache/pip/wheels/d9/d1/96/2815b374d352831ddfdeb4e5f92ba98345626b71022e02a862\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60966 sha256=ff9b72be2b7a8f6c308aa3e186d7865dd32bf443916ca4756d9698eabd7dd44e\n  Stored in directory: /root/.cache/pip/wheels/4f/89/a3/5cf59e30a8a75c917c313f14da0f6209be2d147e3160b985d6\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: efficientnet-pytorch, timm, pretrainedmodels, segmentation-models-pytorch\n  Attempting uninstall: timm\n    Found existing installation: timm 0.6.12\n    Uninstalling timm-0.6.12:\n      Successfully uninstalled timm-0.6.12\nSuccessfully installed efficientnet-pytorch-0.6.3 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.2.0 timm-0.4.12\n\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m","output_type":"stream"},{"name":"stderr","text":"WARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\n","output_type":"stream"},{"name":"stdout","text":"Building tissue masks, label maps and patch grids ...\n  fragment 2: mask (14830, 9506), 6185 candidate patches\n  fragment 3: mask (7606, 5249), 1635 candidate patches\nTotal patches: 7820  -> train 6256 / val 1564\n\n=== Stage 1: training coarse low-res prior model ===\n","output_type":"stream"},{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/resnet18-5c106cde.pth\" to /root/.cache/torch/hub/checkpoints/resnet18-5c106cde.pth\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"  0%|          | 0.00/44.7M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"8bee3478f818481b89444725b9e336db"}},"metadata":{}},{"name":"stdout","text":"  [stage1 1/3] loss=0.4369 (261.1s)\n  [stage1 2/3] loss=0.4197 (121.0s)\n  [stage1 3/3] loss=0.4104 (120.6s)\n  saved Stage-1 checkpoint -> /kaggle/working/vesuviusnet_stage1_coarse.pth\n\nEstimating ink-pixel prior for BCE class weighting ...\n  estimated positive-pixel fraction: 0.15586\n  BCE pos_weight = 5.42\n[setup] RAdamScheduleFree isn't compatible with this PyTorch build (module 'torch' has no attribute '_foreach_lerp_'); falling back to a scheduled optimizer instead.\n[setup] using torch.optim.RAdam + CosineAnnealingLR.\n\nStarting Stage-2 training ...\n[01/8] train_loss=0.8935 dice=0.3129 | val_loss=0.8548 dice=0.3384 fbeta0.5=0.2503 recall=0.8188 precision=0.2132 (592.1s)\n  -> saved new best checkpoint (val_dice=0.3384)\n[02/8] train_loss=0.8425 dice=0.3415 | val_loss=0.8434 dice=0.3656 fbeta0.5=0.3280 recall=0.4517 precision=0.3070 (586.2s)\n  -> saved new best checkpoint (val_dice=0.3656)\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"\n!pip install segmentation-models-pytorch==0.2.0\n!pip install zarr==2.12.0\n\"\"\"\n==========================================================================================\nVESUVIUS CHALLENGE - INK DETECTION\nMemory-safe patch-based pipeline (Zarr storage, 3D->2D SE-UNet + MaxPool/SegFormer branch,\ntemporal augmentations, BCE+Tversky, per-fragment adaptive thresholding (Otsu or bounded\nF0.5 grid-search) instead of one global threshold, small ensemble, viz).\n==========================================================================================\n\nHOW TO USE ON KAGGLE\n---------------------\n1. Paste each \"# %% [CELL n] ...\" block into its own notebook cell (recommended), OR just\n   run this whole file as a single script cell.\n2. Turn on a GPU accelerator (P100 / T4x2).\n3. Data is expected at:\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/surface_volume/00.tif ... 64.tif\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/inklabels.png\n4. Everything derived is written to /kaggle/working (Zarr stacks, checkpoints, plots, json).\n\nWHY THIS AVOIDS OOM\n--------------------\n- Raw fragments are converted ONCE into on-disk chunked Zarr arrays (uint8), slice-by-slice,\n  so we never hold a full (65, H, W) volume in RAM.\n- Training/inference never touch the full fragment either: we sample small 3D patches\n  (depth_window x tile x tile) directly out of the Zarr store on demand.\n- Mixed precision (AMP), small batch size + optional grad accumulation, aggressive\n  `del` + `gc.collect()` + `torch.cuda.empty_cache()`, and bounded DataLoader workers.\n- Test-fragment inference is also patch-based with an overlap-averaged stitching buffer\n  (float32 (H,W) accumulator, which for Vesuvius fragment sizes is a few hundred MB at most\n  - much smaller than the raw (65,H,W) uint16 volume would be).\n\"\"\"\n\n# %% [CELL 0] ---- Setup & installs -------------------------------------------------------\nimport subprocess, sys\ndef _pip(pkgs):\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", *pkgs])\n\ntry:\n    import zarr  # noqa\nexcept ImportError:\n    _pip([\"zarr==2.16.1\"])\ntry:\n    import tifffile  # noqa\nexcept ImportError:\n    _pip([\"tifffile\"])\ntry:\n    from skimage.filters import threshold_otsu  # noqa\nexcept ImportError:\n    _pip([\"scikit-image\"])\n    from skimage.filters import threshold_otsu  # noqa\ndef _smp_has_mit(smp_module):\n    try:\n        from segmentation_models_pytorch.encoders import encoders as _enc\n        return \"mit_b3\" in _enc\n    except Exception:\n        return False\n\ntry:\n    import segmentation_models_pytorch as smp  # noqa\n    HAS_SMP = _smp_has_mit(smp)\nexcept ImportError:\n    smp = None\n    HAS_SMP = False\n\nif not HAS_SMP:\n    # Kaggle's preinstalled smp build is often old and lacks the timm-based MiT/SegFormer\n    # encoders (mit_b0..mit_b5). Try upgrading; if that fails (no internet / version pin\n    # issues), we fall back to the hand-rolled decoder below instead of crashing.\n    # NOTE: this block is intentionally self-contained (own subprocess/sys import + inline\n    # pip call) so it can't NameError even if this cell is ever re-run independently of the\n    # cell that defines the top-level `_pip` helper.\n    try:\n        import subprocess as _subprocess, sys as _sys\n        _subprocess.run([_sys.executable, \"-m\", \"pip\", \"install\", \"-q\",\n                          \"-U\", \"segmentation-models-pytorch\", \"timm\"])\n        import importlib\n        if smp is not None:\n            importlib.reload(smp)\n        else:\n            import segmentation_models_pytorch as smp  # noqa\n        HAS_SMP = _smp_has_mit(smp)\n    except Exception as e:\n        print(f\"[setup] Could not get MiT/SegFormer support from segmentation_models_pytorch \"\n              f\"({e}); will use the built-in fallback decoder for the maxpool_seg branch.\")\n        HAS_SMP = False\n\nprint(f\"[setup] segmentation_models_pytorch MiT (SegFormer) backbone available: {HAS_SMP}\")\n\nimport os, gc, json, math, random, time, glob\nimport numpy as np\nimport zarr\nimport tifffile\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import precision_recall_curve\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\n\n# %% [CELL 1] ---- Config ------------------------------------------------------------------\nclass CFG:\n    DATA_ROOT   = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    WORK_DIR    = \"/kaggle/working\"\n    ZARR_DIR    = os.path.join(WORK_DIR, \"zarr_store\")\n    CKPT_DIR    = os.path.join(WORK_DIR, \"checkpoints\")\n    VIZ_DIR     = os.path.join(WORK_DIR, \"viz\")\n\n    TRAIN_FRAGMENTS = [2, 3]   # 80/20 spatial split inside each -> train/val\n    TEST_FRAGMENT    = 1       # fully held out\n\n    Z_SLICES = 65               # 00.tif .. 64.tif\n    Z_MID    = Z_SLICES // 2    # center of the stack\n\n    # depth window used per sample (Temporal Random Crop draws a window from this range)\n    DEPTH_MIN, DEPTH_MAX = 12, 22\n\n    TILE          = 224          # spatial patch size (H=W=TILE)\n    TRAIN_STRIDE  = 112          # 50% overlap while enumerating candidate train/val patches\n    TEST_STRIDE   = 112          # overlap for sliding-window test inference (averaged)\n\n    FG_MEAN_THRESH = 8.0 / 255.0   # skip near-empty (background/air) patches when indexing\n\n    VAL_FRACTION_BY_WIDTH = 0.8    # first 80% of fragment width -> train, last 20% -> val\n\n    BATCH_SIZE   = 8\n    ACCUM_STEPS  = 2               # effective batch = BATCH_SIZE * ACCUM_STEPS\n    NUM_WORKERS  = 2\n    EPOCHS       = 12\n    LR           = 3e-4\n    WEIGHT_DECAY = 1e-4\n\n    # patches sampled per epoch (subsample huge candidate lists -> bounds RAM & epoch time)\n    MAX_TRAIN_PATCHES_PER_EPOCH = 2500\n    MAX_VAL_PATCHES             = 600\n\n    # Ensemble: each entry is one trained model variant (architecture, depth window, seed)\n    ENSEMBLE_CONFIGS = [\n        dict(name=\"se3d_unet_d16_s0\",  arch=\"se3d_unet\",     depth=16, base_ch=24, seed=0),\n        dict(name=\"se3d_unet_d20_s1\",  arch=\"se3d_unet\",     depth=20, base_ch=24, seed=1),\n        dict(name=\"maxpool_seg_d16_s2\", arch=\"maxpool_seg\",  depth=16, base_ch=32, seed=2),\n        dict(name=\"se3d_unet_d12_s3\",  arch=\"se3d_unet\",     depth=12, base_ch=32, seed=3),\n        dict(name=\"maxpool_seg_d22_s4\", arch=\"maxpool_seg\",  depth=22, base_ch=32, seed=4),\n    ]\n\n    F_BETA = 0.5  # F_0.5 -> precision-weighted\n\n    # Tversky index generalizes Dice: TI = TP / (TP + alpha*FN + beta*FP).\n    # beta > alpha penalizes false positives harder than false negatives, which is the\n    # correct pairing for an F_0.5 (precision-weighted) target -- this is what was missing\n    # when the ensemble converged to \"flag almost everything as ink\" (P=0.20, R=0.999).\n    TVERSKY_ALPHA = 0.3\n    TVERSKY_BETA  = 0.7\n\n    # --- Per-fragment adaptive thresholding -----------------------------------------------\n    # A single threshold fit on fragments 2/3's combined validation split doesn't necessarily\n    # transfer to fragment 1 (or any other fragment) -- scan calibration, ink contrast, and\n    # background noise differ per fragment. Both methods below derive a threshold FROM THE\n    # FRAGMENT BEING EVALUATED rather than reusing a fixed value, and both are bounded to\n    # this range since thresholds outside it tend to be unstable/degenerate in practice.\n    THRESHOLD_SEARCH_LO = 0.2\n    THRESHOLD_SEARCH_HI = 0.8\n    # \"otsu\"   -> unsupervised: Otsu's method on this fragment's own predicted-probability\n    #             histogram. Needs no ground truth, so it also works on genuinely unlabeled\n    #             fragments at real inference time.\n    # \"search\" -> supervised: grid-search the F_beta-optimal threshold in the bounded range\n    #             using this fragment's own ground truth. Only usable when labels exist\n    #             (e.g. held-out fragments kept for evaluation, like fragment 1 here).\n    THRESHOLD_METHOD = \"otsu\"\n\nos.makedirs(CFG.ZARR_DIR, exist_ok=True)\nos.makedirs(CFG.CKPT_DIR, exist_ok=True)\nos.makedirs(CFG.VIZ_DIR, exist_ok=True)\n\n\n# %% [CELL 2] ---- TIFF stack -> chunked Zarr (slice-by-slice, no full-volume RAM load) ----\ndef _existing_zarr_is_valid(frag_id, cfg=CFG):\n    \"\"\"Checks that a previously-built Zarr store for this fragment is present AND complete\n    (right depth, and label mask spatially matches the volume) before trusting it -- a\n    directory merely existing isn't proof the conversion finished cleanly last time.\"\"\"\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n    if not (os.path.exists(vol_path) and os.path.exists(lbl_path)):\n        return False, None, None\n    try:\n        vol = zarr.open(vol_path, mode=\"r\")[\"data\"]\n        lbl = zarr.open(lbl_path, mode=\"r\")[\"data\"]\n        if vol.shape[0] != cfg.Z_SLICES:\n            print(f\"[frag {frag_id}] cached zarr has depth {vol.shape[0]} != expected \"\n                  f\"{cfg.Z_SLICES}; will rebuild.\")\n            return False, None, None\n        if lbl.shape != vol.shape[1:]:\n            print(f\"[frag {frag_id}] cached label shape {lbl.shape} != volume spatial \"\n                  f\"shape {vol.shape[1:]}; will rebuild.\")\n            return False, None, None\n        return True, vol_path, lbl_path\n    except Exception as e:\n        print(f\"[frag {frag_id}] cached zarr at {vol_path} looks corrupt ({e}); will rebuild.\")\n        return False, None, None\n\n\ndef convert_fragment_to_zarr(frag_id: int, cfg=CFG):\n    \"\"\"\n    Writes:\n      zarr_store/frag{frag_id}_volume.zarr  -> uint8 array, shape (Z, H, W), chunks (Z, 256, 256)\n      zarr_store/frag{frag_id}_labels.zarr  -> uint8 array, shape (H, W),   chunks (256, 256)\n    Reuses a previously-built, verified-complete store instead of re-reading the source TIFFs.\n    \"\"\"\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n\n    is_valid, cached_vol, cached_lbl = _existing_zarr_is_valid(frag_id, cfg)\n    if is_valid:\n        vol = zarr.open(cached_vol, mode=\"r\")[\"data\"]\n        lbl = zarr.open(cached_lbl, mode=\"r\")[\"data\"]\n        print(f\"[frag {frag_id}] reusing existing zarr at {cfg.ZARR_DIR} \"\n              f\"(volume {vol.shape}, labels {lbl.shape}) -- skipping TIFF re-read.\")\n        return cached_vol, cached_lbl\n\n    frag_dir = os.path.join(cfg.DATA_ROOT, str(frag_id))\n    slice_paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in range(cfg.Z_SLICES)]\n    label_path  = os.path.join(frag_dir, \"inklabels.png\")\n\n    # peek shape from first slice\n    with tifffile.TiffFile(slice_paths[0]) as tf:\n        h, w = tf.pages[0].shape\n    print(f\"[frag {frag_id}] volume shape -> ({cfg.Z_SLICES}, {h}, {w})\")\n\n    store = zarr.DirectoryStore(vol_path)\n    root = zarr.group(store=store, overwrite=True)\n    vol_z = root.create_dataset(\n        \"data\", shape=(cfg.Z_SLICES, h, w), chunks=(cfg.Z_SLICES, 256, 256),\n        dtype=\"u1\", compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n\n    for i, sp in enumerate(slice_paths):\n        sl = tifffile.imread(sp)  # uint16, shape (H, W) -- ONE slice in RAM at a time\n        # normalize 16-bit -> 8-bit for compact on-disk storage / cheap I/O during training\n        sl8 = (sl.astype(np.float32) / 65535.0 * 255.0).clip(0, 255).astype(np.uint8)\n        vol_z[i, :, :] = sl8\n        del sl, sl8\n        if i % 16 == 0:\n            gc.collect()\n    print(f\"[frag {frag_id}] volume -> zarr done.\")\n\n    lbl_img = Image.open(label_path).convert(\"L\")\n    lbl = (np.array(lbl_img) > 127).astype(np.uint8)\n    assert lbl.shape == (h, w), f\"label shape {lbl.shape} != volume shape {(h, w)}\"\n\n    lstore = zarr.DirectoryStore(lbl_path)\n    lroot = zarr.group(store=lstore, overwrite=True)\n    lbl_z = lroot.create_dataset(\n        \"data\", shape=(h, w), chunks=(256, 256), dtype=\"u1\",\n        compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n    lbl_z[:, :] = lbl\n    del lbl, lbl_img\n    gc.collect()\n    print(f\"[frag {frag_id}] labels -> zarr done.\")\n    return vol_path, lbl_path\n\n\ndef open_fragment_zarr(frag_id, cfg=CFG):\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n    vol = zarr.open(vol_path, mode=\"r\")[\"data\"]\n    lbl = zarr.open(lbl_path, mode=\"r\")[\"data\"]\n    return vol, lbl\n\n\n# %% [CELL 3] ---- Patch indexing + 80/20 spatial split ------------------------------------\ndef build_patch_grid(h, w, tile, stride):\n    ys = list(range(0, max(h - tile, 0) + 1, stride))\n    xs = list(range(0, max(w - tile, 0) + 1, stride))\n    if ys[-1] != h - tile: ys.append(max(h - tile, 0))\n    if xs[-1] != w - tile: xs.append(max(w - tile, 0))\n    return [(y, x) for y in ys for x in xs]\n\n\ndef index_fragment_patches(frag_id, cfg=CFG, split=\"both\"):\n    \"\"\"\n    Returns list of dicts: {frag_id, y, x} for foreground patches only (cheap intensity\n    filter on the middle slice, read directly from zarr chunk-by-chunk -> low RAM).\n    split: 'train' -> first VAL_FRACTION_BY_WIDTH of width\n           'val'   -> remaining width\n           'both'  -> no split (used for the held-out test fragment)\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    mid_slice = np.asarray(vol[cfg.Z_MID, :, :])  # (H, W) uint8, one slice in RAM\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TRAIN_STRIDE if split != \"both\" else cfg.TEST_STRIDE)\n\n    split_x = int(W * cfg.VAL_FRACTION_BY_WIDTH)\n    kept = []\n    for (y, x) in coords:\n        if split == \"train\" and not (x + cfg.TILE <= split_x):\n            continue\n        if split == \"val\" and not (x >= split_x):\n            continue\n        patch = mid_slice[y:y + cfg.TILE, x:x + cfg.TILE]\n        if patch.size == 0:\n            continue\n        if (patch.astype(np.float32) / 255.0).mean() < cfg.FG_MEAN_THRESH:\n            continue  # skip empty background/air patch\n        kept.append({\"frag_id\": frag_id, \"y\": y, \"x\": x})\n    del mid_slice\n    gc.collect()\n    print(f\"[frag {frag_id}] split={split}: {len(kept)} foreground patches indexed \"\n          f\"(H={H}, W={W}).\")\n    return kept\n\n\n# %% [CELL 4] ---- Dataset with Temporal Random Crop / Random Paste / Temporal Cutout ------\nclass VesuviusPatchDataset(Dataset):\n    \"\"\"\n    Each item pulls a (depth_window, TILE, TILE) sub-volume straight out of the on-disk\n    Zarr store (no full-fragment load), applies the three temporal augmentations plus\n    standard spatial flips/rotations, and returns (volume_tensor, mask_tensor).\n    \"\"\"\n    def __init__(self, patch_records, depth_window, cfg=CFG, train=True, max_items=None):\n        self.records = patch_records\n        self.depth_window = depth_window\n        self.cfg = cfg\n        self.train = train\n        self._vol_cache = {}\n        self._lbl_cache = {}\n        if max_items is not None and len(self.records) > max_items:\n            self.records = random.sample(self.records, max_items)\n\n    def _get_arrays(self, frag_id):\n        if frag_id not in self._vol_cache:\n            self._vol_cache[frag_id], self._lbl_cache[frag_id] = open_fragment_zarr(frag_id, self.cfg)\n        return self._vol_cache[frag_id], self._lbl_cache[frag_id]\n\n    def __len__(self):\n        return len(self.records)\n\n    def _sample_depth_window(self, Z):\n        dw = self.depth_window if not self.train else random.randint(CFG.DEPTH_MIN, CFG.DEPTH_MAX)\n        dw = min(dw, Z)\n        z_mid = Z // 2\n        half = dw // 2\n        # Temporal Random Crop: random window around the middle of the stack\n        jitter = random.randint(-4, 4) if self.train else 0\n        z0 = max(0, min(Z - dw, z_mid - half + jitter))\n        return z0, dw\n\n    def __getitem__(self, idx):\n        rec = self.records[idx]\n        vol, lbl = self._get_arrays(rec[\"frag_id\"])\n        Z, H, W = vol.shape\n        y, x, t = rec[\"y\"], rec[\"x\"], self.cfg.TILE\n\n        z0, dw = self._sample_depth_window(Z)\n        sub = np.asarray(vol[z0:z0 + dw, y:y + t, x:x + t]).astype(np.float32) / 255.0  # (dw,t,t)\n        mask = np.asarray(lbl[y:y + t, x:x + t]).astype(np.float32)  # (t,t)\n\n        if self.train:\n            sub, mask = self._augment(sub, mask)\n\n        # pad depth to DEPTH_MAX so batches stack cleanly; padding slices are zeroed\n        # (network ignores all-zero slices thanks to Temporal Cutout training on real zeros too)\n        pad = CFG.DEPTH_MAX - sub.shape[0]\n        if pad > 0:\n            sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n\n        vol_t = torch.from_numpy(sub).unsqueeze(0).float()   # (1, D, H, W)\n        mask_t = torch.from_numpy(mask).unsqueeze(0).float() # (1, H, W)\n        return vol_t, mask_t\n\n    def _augment(self, sub, mask):\n        D = sub.shape[0]\n\n        # --- Random Paste: shift cropped slice-range to a random Z offset within the padded\n        #     budget, keeping slice order intact (sequential order preserved). ---\n        if random.random() < 0.5 and D < CFG.DEPTH_MAX:\n            max_shift = CFG.DEPTH_MAX - D\n            shift = random.randint(0, max_shift)\n            padded = np.zeros((CFG.DEPTH_MAX, *sub.shape[1:]), dtype=sub.dtype)\n            padded[shift:shift + D] = sub\n            sub = padded\n            D = CFG.DEPTH_MAX\n\n        # --- Temporal Cutout: zero out 1-2 random layers inside the active cube ---\n        n_cutout = random.choice([0, 1, 1, 2])\n        for _ in range(n_cutout):\n            zi = random.randint(0, D - 1)\n            sub[zi] = 0.0\n\n        # --- standard spatial augs ---\n        if random.random() < 0.5:\n            sub = sub[:, :, ::-1].copy(); mask = mask[:, ::-1].copy()\n        if random.random() < 0.5:\n            sub = sub[:, ::-1, :].copy(); mask = mask[::-1, :].copy()\n        k = random.choice([0, 1, 2, 3])\n        if k:\n            sub = np.rot90(sub, k, axes=(1, 2)).copy()\n            mask = np.rot90(mask, k, axes=(0, 1)).copy()\n\n        return sub, mask\n\n\n# %% [CELL 5] ---- Model: SE blocks, 3D encoder -> 2D SE-UNet decoder, MaxPool+Seg decoder --\nclass SEBlock2D(nn.Module):\n    \"\"\"Squeeze-and-Excitation applied inside skip connections.\"\"\"\n    def __init__(self, ch, reduction=8):\n        super().__init__()\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(ch, max(ch // reduction, 4), 1), nn.ReLU(inplace=True),\n            nn.Conv2d(max(ch // reduction, 4), ch, 1), nn.Sigmoid(),\n        )\n    def forward(self, x):\n        return x * self.fc(self.pool(x))\n\n\nclass ConvBNAct2D(nn.Module):\n    def __init__(self, cin, cout, k=3, s=1, p=1):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(cin, cout, k, s, p, bias=False),\n            nn.BatchNorm2d(cout), nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.block(x)\n\n\nclass Encoder3D(nn.Module):\n    \"\"\"3D conv encoder that progressively collapses the depth axis while extracting\n    multi-scale 2D feature maps (one per stage) for U-Net-style skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n        self.stage1 = nn.Sequential(\n            nn.Conv3d(1, c1, 3, padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n            nn.Conv3d(c1, c1, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n        )\n        self.stage2 = nn.Sequential(\n            nn.Conv3d(c1, c2, 3, padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n            nn.Conv3d(c2, c2, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n        )\n        self.stage3 = nn.Sequential(\n            nn.Conv3d(c2, c3, 3, padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n            nn.Conv3d(c3, c3, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n        )\n        self.stage4 = nn.Sequential(\n            # stride only on H,W here (depth stride=1): this must land the bottleneck one\n            # scale BELOW f3's spatial resolution, since the decoder's up3 upsamples it by\n            # 2x before concatenating with f3. Depth is fully collapsed right after by the\n            # adaptive pool, so no depth stride is needed at this stage.\n            nn.Conv3d(c3, c4, 3, stride=(1, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c4), nn.ReLU(inplace=True),\n            nn.AdaptiveMaxPool3d((1, None, None)),  # collapse remaining depth -> bottleneck\n        )\n        self.out_channels = (c1, c2, c3, c4)\n\n    @staticmethod\n    def _depth_collapse(x):\n        # max over depth to produce a 2D skip feature map at this scale\n        return x.max(dim=2).values\n\n    def forward(self, x):  # x: (B, 1, D, H, W)\n        s1 = self.stage1(x); f1 = self._depth_collapse(s1)          # H/2\n        s2 = self.stage2(s1); f2 = self._depth_collapse(s2)          # H/4\n        s3 = self.stage3(s2); f3 = self._depth_collapse(s3)          # H/8\n        s4 = self.stage4(s3).squeeze(2)                               # H/8, depth->1\n        return f1, f2, f3, s4\n\n\nclass SEUnetDecoder2D(nn.Module):\n    def __init__(self, enc_channels):\n        super().__init__()\n        c1, c2, c3, c4 = enc_channels\n        self.up3 = nn.ConvTranspose2d(c4, c3, 2, 2)\n        self.se3 = SEBlock2D(c3); self.dec3 = ConvBNAct2D(c3 * 2, c3)\n        self.up2 = nn.ConvTranspose2d(c3, c2, 2, 2)\n        self.se2 = SEBlock2D(c2); self.dec2 = ConvBNAct2D(c2 * 2, c2)\n        self.up1 = nn.ConvTranspose2d(c2, c1, 2, 2)\n        self.se1 = SEBlock2D(c1); self.dec1 = ConvBNAct2D(c1 * 2, c1)\n        self.up0 = nn.ConvTranspose2d(c1, c1 // 2, 2, 2)\n        self.final = nn.Conv2d(c1 // 2, 1, 1)\n\n    def forward(self, f1, f2, f3, bottleneck, out_hw):\n        x = self.up3(bottleneck); x = self.dec3(torch.cat([x, self.se3(f3)], dim=1))\n        x = self.up2(x);          x = self.dec2(torch.cat([x, self.se2(f2)], dim=1))\n        x = self.up1(x);          x = self.dec1(torch.cat([x, self.se1(f1)], dim=1))\n        x = self.up0(x)\n        x = F.interpolate(x, size=out_hw, mode=\"bilinear\", align_corners=False)\n        return self.final(x)\n\n\nclass SE3DUNet(nn.Module):\n    \"\"\"Top-solution-style architecture: 3D encoder + 2D decoder U-Net with SE blocks in\n    the skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        self.encoder = Encoder3D(base_ch)\n        self.decoder = SEUnetDecoder2D(self.encoder.out_channels)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        h, w = x.shape[-2], x.shape[-1]\n        f1, f2, f3, bott = self.encoder(x)\n        return self.decoder(f1, f2, f3, bott, (h, w))\n\n\nclass MaxPoolFlattenSegHead(nn.Module):\n    \"\"\"Depth Flattening branch: collapse the (D,H,W) volume to a 2D feature map via\n    max pooling across depth, then feed a heavy semantic-segmentation decoder\n    (SegFormer/MiT-style if `segmentation_models_pytorch` is available, otherwise a\n    hand-rolled multi-scale conv decoder with SE-augmented skips).\"\"\"\n    def __init__(self, base_ch=32, use_smp=HAS_SMP):\n        super().__init__()\n        self.use_smp = False\n        if use_smp:\n            try:\n                self.net = smp.Unet(\n                    encoder_name=\"mit_b3\", encoder_weights=\"imagenet\",\n                    in_channels=1, classes=1, activation=None,\n                )\n                self.use_smp = True\n            except Exception as e:\n                print(f\"[MaxPoolFlattenSegHead] smp mit_b3 unavailable at model-build time \"\n                      f\"({e}); using the built-in fallback decoder instead.\")\n\n        if not self.use_smp:\n            c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n            self.enc1 = nn.Sequential(ConvBNAct2D(1, c1), ConvBNAct2D(c1, c1))\n            self.pool1 = nn.MaxPool2d(2)\n            self.enc2 = nn.Sequential(ConvBNAct2D(c1, c2), ConvBNAct2D(c2, c2))\n            self.pool2 = nn.MaxPool2d(2)\n            self.enc3 = nn.Sequential(ConvBNAct2D(c2, c3), ConvBNAct2D(c3, c3))\n            self.pool3 = nn.MaxPool2d(2)\n            self.bott = nn.Sequential(ConvBNAct2D(c3, c4), ConvBNAct2D(c4, c4))\n            self.decoder = SEUnetDecoder2D((c1, c2, c3, c4))\n\n    def forward(self, x2d):  # x2d: (B, 1, H, W)  (already depth-flattened)\n        if self.use_smp:\n            return self.net(x2d)\n        h, w = x2d.shape[-2], x2d.shape[-1]\n        e1 = self.enc1(x2d); p1 = self.pool1(e1)\n        e2 = self.enc2(p1);  p2 = self.pool2(e2)\n        e3 = self.enc3(p2);  p3 = self.pool3(e3)\n        b  = self.bott(p3)\n        return self.decoder(e1, e2, e3, b, (h, w))\n\n\nclass MaxPoolSegModel(nn.Module):\n    \"\"\"Full pipeline for the 'Depth Flattening' approach: Max Pool across the depth axis\n    of the 3D volume -> 2D feature map -> heavy segmentation decoder -> ink mask.\"\"\"\n    def __init__(self, base_ch=32):\n        super().__init__()\n        self.head = MaxPoolFlattenSegHead(base_ch=base_ch)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        x2d = x.max(dim=2).values  # Max Pooling across depth axis -> (B,1,H,W)\n        return self.head(x2d)\n\n\ndef build_model(arch, base_ch):\n    if arch == \"se3d_unet\":\n        return SE3DUNet(base_ch=base_ch)\n    elif arch == \"maxpool_seg\":\n        return MaxPoolSegModel(base_ch=base_ch)\n    else:\n        raise ValueError(arch)\n\n\n# %% [CELL 6] ---- Loss (BCE + Dice) and F-beta metric/threshold search --------------------\nclass BCETverskyLoss(nn.Module):\n    \"\"\"BCE + Tversky. Unlike plain Dice (which weights precision/recall equally), Tversky\n    with beta > alpha explicitly penalizes false positives harder -- the correct pairing\n    when the downstream metric (F_0.5) itself weights precision over recall.\"\"\"\n    def __init__(self, bce_w=0.5, tversky_w=0.5, alpha=CFG.TVERSKY_ALPHA, beta=CFG.TVERSKY_BETA, smooth=1.0):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.bce_w, self.tversky_w = bce_w, tversky_w\n        self.alpha, self.beta, self.smooth = alpha, beta, smooth\n\n    def forward(self, logits, target):\n        bce_loss = self.bce(logits, target)\n        prob = torch.sigmoid(logits)\n        p, t = prob.flatten(1), target.flatten(1)\n        tp = (p * t).sum(1)\n        fp = (p * (1 - t)).sum(1)\n        fn = ((1 - p) * t).sum(1)\n        tversky = (tp + self.smooth) / (tp + self.alpha * fn + self.beta * fp + self.smooth)\n        tversky_loss = 1 - tversky\n        return self.bce_w * bce_loss + self.tversky_w * tversky_loss.mean()\n\n\n# kept as an alias so any external references to the old name still work\nBCEDiceLoss = BCETverskyLoss\n\n\ndef fbeta_score(precision, recall, beta=CFG.F_BETA, eps=1e-8):\n    b2 = beta ** 2\n    return (1 + b2) * precision * recall / (b2 * precision + recall + eps)\n\n\ndef find_best_threshold(probs_flat: np.ndarray, targets_flat: np.ndarray, beta=CFG.F_BETA,\n                         lo=CFG.THRESHOLD_SEARCH_LO, hi=CFG.THRESHOLD_SEARCH_HI):\n    \"\"\"Optimizes the decision threshold strictly for F_beta (beta<1 -> precision-weighted),\n    restricted to [lo, hi] so checkpoint selection during training stays consistent with the\n    bounded range used for final per-fragment thresholding at inference time.\"\"\"\n    precision, recall, thresholds = precision_recall_curve(targets_flat, probs_flat)\n    precision, recall = precision[:-1], recall[:-1]\n    in_range = (thresholds >= lo) & (thresholds <= hi)\n    if not in_range.any():\n        # degenerate case (e.g. a very early, poorly-calibrated epoch) -> fall back to\n        # the unrestricted search rather than returning nothing\n        in_range = np.ones_like(thresholds, dtype=bool)\n    precision, recall, thresholds = precision[in_range], recall[in_range], thresholds[in_range]\n    scores = fbeta_score(precision, recall, beta)\n    best_idx = int(np.nanargmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), float(precision[best_idx]), float(recall[best_idx])\n\n\ndef otsu_threshold_bounded(prob_map: np.ndarray, lo=CFG.THRESHOLD_SEARCH_LO, hi=CFG.THRESHOLD_SEARCH_HI, nbins=256):\n    \"\"\"\n    Unsupervised, per-fragment threshold: runs Otsu's method directly on THIS fragment's own\n    predicted-probability histogram (restricted to [lo, hi]) to find the value that best\n    separates the histogram into two classes. Needs no ground truth, so unlike a threshold\n    learned on a different fragment's validation set, this adapts automatically to each\n    fragment's own contrast/noise characteristics -- and it's the only one of the two methods\n    usable on a genuinely unlabeled fragment at real inference time.\n    \"\"\"\n    probs = prob_map.ravel().astype(np.float64)\n    in_range = probs[(probs >= lo) & (probs <= hi)]\n    if in_range.size < 100:\n        # degenerate fragment (near-empty or saturated probability map) -> safe midpoint\n        return float((lo + hi) / 2.0)\n    try:\n        thr = threshold_otsu(in_range, nbins=nbins)\n    except Exception:\n        thr = float(np.median(in_range))\n    return float(np.clip(thr, lo, hi))\n\n\ndef search_threshold_bounded(prob_map: np.ndarray, gt_map: np.ndarray,\n                              lo=CFG.THRESHOLD_SEARCH_LO, hi=CFG.THRESHOLD_SEARCH_HI,\n                              steps=61, beta=CFG.F_BETA):\n    \"\"\"\n    Supervised, per-fragment threshold: grid-searches [lo, hi] for the F_beta-optimal cut\n    using THIS fragment's own ground truth (not a different fragment's validation labels).\n    Only valid where labels exist -- use otsu_threshold_bounded for truly unlabeled fragments.\n    \"\"\"\n    thresholds = np.linspace(lo, hi, steps)\n    p_flat = prob_map.ravel()\n    t_flat = gt_map.ravel().astype(np.float32)\n    best_thr, best_score, best_p, best_r = float(lo), -1.0, 0.0, 0.0\n    for thr in thresholds:\n        pred = (p_flat >= thr).astype(np.float32)\n        tp = float((pred * t_flat).sum())\n        fp = float((pred * (1 - t_flat)).sum())\n        fn = float(((1 - pred) * t_flat).sum())\n        precision = tp / (tp + fp + 1e-8)\n        recall = tp / (tp + fn + 1e-8)\n        score = fbeta_score(precision, recall, beta)\n        if score > best_score:\n            best_thr, best_score, best_p, best_r = float(thr), float(score), precision, recall\n    return best_thr, best_score, best_p, best_r\n\n\n# %% [CELL 7] ---- Training loop for one ensemble member ------------------------------------\ndef run_validation(model, val_loader):\n    model.eval()\n    all_probs, all_targets = [], []\n    with torch.no_grad(), autocast():\n        for vol, mask in val_loader:\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            logits = model(vol)\n            probs = torch.sigmoid(logits).float().cpu().numpy().ravel()\n            targs = mask.float().cpu().numpy().ravel()\n            # subsample pixels per patch to keep the PR-curve computation light on RAM\n            if len(probs) > 20000:\n                idx = np.random.choice(len(probs), 20000, replace=False)\n                probs, targs = probs[idx], targs[idx]\n            all_probs.append(probs); all_targets.append(targs)\n    probs = np.concatenate(all_probs); targets = np.concatenate(all_targets)\n    thr, f05, prec, rec = find_best_threshold(probs, targets)\n    del all_probs, all_targets\n    gc.collect()\n    return thr, f05, prec, rec\n\n\ndef train_one_model(cfg_entry, train_records, val_records, cfg=CFG, epochs=None):\n    epochs = epochs or cfg.EPOCHS\n    torch.manual_seed(cfg_entry[\"seed\"]); random.seed(cfg_entry[\"seed\"])\n\n    model = build_model(cfg_entry[\"arch\"], cfg_entry[\"base_ch\"]).to(DEVICE)\n    opt = torch.optim.AdamW(model.parameters(), lr=cfg.LR, weight_decay=cfg.WEIGHT_DECAY)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    scaler = GradScaler()\n    criterion = BCETverskyLoss()\n\n    train_ds = VesuviusPatchDataset(train_records, cfg_entry[\"depth\"], cfg, train=True,\n                                     max_items=cfg.MAX_TRAIN_PATCHES_PER_EPOCH)\n    val_ds = VesuviusPatchDataset(val_records, cfg_entry[\"depth\"], cfg, train=False,\n                                   max_items=cfg.MAX_VAL_PATCHES)\n\n    best_f05, best_thr, best_state = -1.0, 0.5, None\n    history = []\n\n    for epoch in range(epochs):\n        # resample the train subset each epoch for coverage without ever loading everything\n        train_ds.records = random.sample(train_records, min(cfg.MAX_TRAIN_PATCHES_PER_EPOCH, len(train_records)))\n        train_loader = DataLoader(train_ds, batch_size=cfg.BATCH_SIZE, shuffle=True,\n                                   num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True)\n        val_loader = DataLoader(val_ds, batch_size=cfg.BATCH_SIZE, shuffle=False,\n                                 num_workers=cfg.NUM_WORKERS, pin_memory=True)\n\n        model.train()\n        running_loss, n_steps = 0.0, 0\n        opt.zero_grad()\n        t0 = time.time()\n        for step, (vol, mask) in enumerate(train_loader):\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            with autocast():\n                logits = model(vol)\n                loss = criterion(logits, mask) / cfg.ACCUM_STEPS\n            scaler.scale(loss).backward()\n            if (step + 1) % cfg.ACCUM_STEPS == 0:\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n            running_loss += loss.item() * cfg.ACCUM_STEPS\n            n_steps += 1\n            del vol, mask, logits, loss\n        sched.step()\n\n        thr, f05, prec, rec = run_validation(model, val_loader)\n        dt = time.time() - t0\n        print(f\"[{cfg_entry['name']}] epoch {epoch+1}/{epochs} \"\n              f\"train_loss={running_loss/max(n_steps,1):.4f} \"\n              f\"val_F0.5={f05:.4f} P={prec:.3f} R={rec:.3f} thr={thr:.3f} ({dt:.0f}s)\")\n        history.append(dict(epoch=epoch+1, train_loss=running_loss/max(n_steps,1),\n                             val_f05=f05, val_precision=prec, val_recall=rec, threshold=thr))\n\n        if f05 > best_f05:\n            best_f05, best_thr = f05, thr\n            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n\n        del train_loader, val_loader\n        gc.collect(); torch.cuda.empty_cache()\n\n    model.load_state_dict(best_state)\n    ckpt_path = os.path.join(cfg.CKPT_DIR, f\"{cfg_entry['name']}.pt\")\n    torch.save({\"state_dict\": best_state, \"config\": cfg_entry,\n                \"best_val_f05\": best_f05, \"best_threshold\": best_thr,\n                \"history\": history}, ckpt_path)\n    print(f\"[{cfg_entry['name']}] saved best checkpoint (val F0.5={best_f05:.4f}) -> {ckpt_path}\")\n\n    del model, opt, sched, scaler, train_ds, val_ds\n    gc.collect(); torch.cuda.empty_cache()\n    return ckpt_path, best_thr, best_f05, history\n\n\n# %% [CELL 8] ---- Sliding-window ensemble inference over a full fragment ------------------\ndef predict_fragment_ensemble(frag_id, model_infos, cfg=CFG):\n    \"\"\"\n    model_infos: list of dicts {ckpt_path, config} for each ensemble member.\n    Returns (prob_map (H,W) float32 averaged across models, gt_mask (H,W) uint8).\n    Memory-safe: only a float32 (H,W) accumulator + weight map live in RAM (no full volume).\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    gt = np.asarray(lbl[:, :])\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TEST_STRIDE)\n\n    loaded_models = []\n    for info in model_infos:\n        ckpt = torch.load(info[\"ckpt_path\"], map_location=DEVICE)\n        m = build_model(ckpt[\"config\"][\"arch\"], ckpt[\"config\"][\"base_ch\"]).to(DEVICE)\n        m.load_state_dict(ckpt[\"state_dict\"]); m.eval()\n        loaded_models.append((m, ckpt[\"config\"][\"depth\"]))\n\n    prob_acc = np.zeros((H, W), dtype=np.float32)\n    weight_acc = np.zeros((H, W), dtype=np.float32)\n\n    with torch.no_grad(), autocast():\n        for (y, x) in coords:\n            patch_sum = np.zeros((cfg.TILE, cfg.TILE), dtype=np.float32)\n            for model, depth in loaded_models:\n                z_mid = Z // 2\n                z0 = max(0, min(Z - depth, z_mid - depth // 2))\n                sub = np.asarray(vol[z0:z0 + depth, y:y + cfg.TILE, x:x + cfg.TILE]).astype(np.float32) / 255.0\n                pad = cfg.DEPTH_MAX - sub.shape[0]\n                if pad > 0:\n                    sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n                t = torch.from_numpy(sub).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n                logits = model(t)\n                patch_sum += torch.sigmoid(logits)[0, 0].float().cpu().numpy()\n                del t, logits\n            patch_prob = patch_sum / len(loaded_models)\n            prob_acc[y:y + cfg.TILE, x:x + cfg.TILE] += patch_prob\n            weight_acc[y:y + cfg.TILE, x:x + cfg.TILE] += 1.0\n\n    weight_acc[weight_acc == 0] = 1.0\n    prob_map = prob_acc / weight_acc\n\n    for m, _ in loaded_models:\n        del m\n    del loaded_models\n    gc.collect(); torch.cuda.empty_cache()\n    return prob_map, gt\n\n\n# %% [CELL 9] ---- Orchestration: convert data, split, train ensemble, evaluate, visualize --\ndef main():\n    # 1) Convert all needed fragments to Zarr (train 2,3 + held-out test 1)\n    for fid in CFG.TRAIN_FRAGMENTS + [CFG.TEST_FRAGMENT]:\n        convert_fragment_to_zarr(fid)\n\n    # 2) Build 80/20 spatial split from fragments 2 & 3\n    train_records, val_records = [], []\n    for fid in CFG.TRAIN_FRAGMENTS:\n        train_records += index_fragment_patches(fid, split=\"train\")\n        val_records   += index_fragment_patches(fid, split=\"val\")\n    print(f\"TOTAL train patches: {len(train_records)} | val patches: {len(val_records)}\")\n\n    # 3) Train each ensemble member\n    trained = []\n    for cfg_entry in CFG.ENSEMBLE_CONFIGS:\n        ckpt_path, thr, f05, hist = train_one_model(cfg_entry, train_records, val_records)\n        trained.append({\"ckpt_path\": ckpt_path, \"config\": cfg_entry, \"val_threshold\": thr, \"val_f05\": f05})\n\n    # 4) Per-model validation thresholds are kept only for logging/diagnostics now -- final\n    #    thresholding is done per-fragment (below), which is the more appropriate choice\n    #    since fragment 1's ink/background characteristics differ from fragments 2 & 3.\n    mean_val_thr = float(np.mean([t[\"val_threshold\"] for t in trained]))\n    print(f\"(diagnostic only) mean per-model validation threshold: {mean_val_thr:.3f}\")\n\n    # 5) Full-fragment ensemble inference on the held-out TEST fragment (fragment 1)\n    prob_map, gt_map = predict_fragment_ensemble(CFG.TEST_FRAGMENT, trained)\n\n    # 5b) Per-fragment adaptive thresholding -- derived from fragment 1's OWN probability\n    #     map (Otsu) and, since we happen to have its labels for evaluation, also from its\n    #     own ground truth (bounded grid-search) for comparison. In a real submission where\n    #     the target fragment has no labels, only the Otsu value would be available.\n    otsu_thr = otsu_threshold_bounded(prob_map)\n    search_thr, search_f05, search_p, search_r = search_threshold_bounded(prob_map, gt_map)\n    print(f\"[fragment {CFG.TEST_FRAGMENT}] Otsu threshold (unsupervised) = {otsu_thr:.3f} | \"\n          f\"bounded grid-search threshold (oracle, uses labels) = {search_thr:.3f} \"\n          f\"(F0.5={search_f05:.4f}, P={search_p:.3f}, R={search_r:.3f})\")\n\n    if CFG.THRESHOLD_METHOD == \"otsu\":\n        chosen_thr, threshold_method_used = otsu_thr, \"otsu\"\n    elif CFG.THRESHOLD_METHOD == \"search\":\n        chosen_thr, threshold_method_used = search_thr, \"search\"\n    else:\n        raise ValueError(f\"Unknown CFG.THRESHOLD_METHOD={CFG.THRESHOLD_METHOD!r}; use 'otsu' or 'search'.\")\n    print(f\"[fragment {CFG.TEST_FRAGMENT}] using method='{threshold_method_used}' -> \"\n          f\"final threshold={chosen_thr:.3f}\")\n\n    pred_mask = (prob_map >= chosen_thr).astype(np.uint8)\n\n    # 6) Final test metrics (precision/recall/F0.5/Dice) at the chosen per-fragment threshold\n    p_flat, t_flat = pred_mask.ravel().astype(np.float32), gt_map.ravel().astype(np.float32)\n    tp = float((p_flat * t_flat).sum())\n    fp = float((p_flat * (1 - t_flat)).sum())\n    fn = float(((1 - p_flat) * t_flat).sum())\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n    f05_test = fbeta_score(precision, recall, CFG.F_BETA)\n    dice_test = 2 * tp / (2 * tp + fp + fn + 1e-8)\n    print(f\"TEST (fragment {CFG.TEST_FRAGMENT}) @thr={chosen_thr:.3f} [{threshold_method_used}] -> \"\n          f\"Precision={precision:.4f} Recall={recall:.4f} F0.5={f05_test:.4f} Dice={dice_test:.4f}\")\n\n    metrics = dict(\n        threshold_method_used=threshold_method_used,\n        chosen_threshold=chosen_thr,\n        otsu_threshold=otsu_thr,\n        search_threshold=search_thr,\n        search_threshold_f0_5=search_f05,\n        mean_per_model_val_threshold_diagnostic_only=mean_val_thr,\n        test_precision=precision, test_recall=recall,\n        test_f0_5=f05_test, test_dice=dice_test,\n        members=[{\"name\": t[\"config\"][\"name\"], \"val_f05\": t[\"val_f05\"],\n                  \"val_threshold\": t[\"val_threshold\"]} for t in trained],\n    )\n    with open(os.path.join(CFG.WORK_DIR, \"metrics_summary.json\"), \"w\") as f:\n        json.dump(metrics, f, indent=2)\n    print(\"Saved metrics_summary.json\")\n\n    # 7) Visualization: input (mid-slice) vs ground truth vs prediction\n    vol, _ = open_fragment_zarr(CFG.TEST_FRAGMENT)\n    mid_slice = np.asarray(vol[CFG.Z_MID, :, :])\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(mid_slice, cmap=\"gray\"); axes[0].set_title(f\"Input (mid slice z={CFG.Z_MID})\")\n    axes[1].imshow(gt_map, cmap=\"gray\"); axes[1].set_title(\"Ground Truth Ink Mask\")\n    axes[2].imshow(prob_map, cmap=\"magma\"); axes[2].set_title(\"Predicted Probability Map\")\n    axes[3].imshow(mid_slice, cmap=\"gray\")\n    axes[3].imshow(np.ma.masked_where(pred_mask == 0, pred_mask), cmap=\"autumn\", alpha=0.6)\n    axes[3].set_title(f\"Prediction Overlay ({threshold_method_used} thr={chosen_thr:.2f})\")\n    for ax in axes: ax.axis(\"off\")\n    plt.tight_layout()\n    viz_path = os.path.join(CFG.VIZ_DIR, f\"test_fragment{CFG.TEST_FRAGMENT}_comparison.png\")\n    plt.savefig(viz_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved visualization -> {viz_path}\")\n\n    del mid_slice, prob_map, gt_map, pred_mask\n    gc.collect(); torch.cuda.empty_cache()\n    return metrics\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-05T07:34:39.79638Z","iopub.execute_input":"2026-08-05T07:34:39.796852Z","iopub.status.idle":"2026-08-05T10:41:15.493228Z","shell.execute_reply.started":"2026-08-05T07:34:39.796817Z","shell.execute_reply":"2026-08-05T10:41:15.491868Z"}},"outputs":[{"name":"stdout","text":"Collecting segmentation-models-pytorch==0.2.0\n  Downloading segmentation_models_pytorch-0.2.0-py3-none-any.whl (87 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m87.6/87.6 kB\u001b[0m \u001b[31m6.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hRequirement already satisfied: torchvision>=0.5.0 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.14.0)\nCollecting pretrainedmodels==0.7.4\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m5.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hCollecting efficientnet-pytorch==0.6.3\n  Downloading efficientnet_pytorch-0.6.3.tar.gz (16 kB)\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hCollecting timm==0.4.12\n  Downloading timm-0.4.12-py3-none-any.whl (376 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m377.0/377.0 kB\u001b[0m \u001b[31m11.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hRequirement already satisfied: torch in /opt/conda/lib/python3.7/site-packages (from efficientnet-pytorch==0.6.3->segmentation-models-pytorch==0.2.0) (1.13.0)\nRequirement already satisfied: munch in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (2.5.0)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (4.64.1)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (4.4.0)\nRequirement already satisfied: requests in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.28.2)\nRequirement already satisfied: pillow!=8.3.*,>=5.3.0 in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (9.4.0)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.21.6)\nRequirement already satisfied: six in /opt/conda/lib/python3.7/site-packages (from munch->pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (1.16.0)\nRequirement already satisfied: urllib3<1.27,>=1.21.1 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.26.14)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.1.1)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (3.4)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2022.12.7)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.6.3-py3-none-any.whl size=12422 sha256=7d14c6a339accce904943ac801dc1655bd1d41ffbcbe674346920e4e81f0c9c4\n  Stored in directory: /root/.cache/pip/wheels/d9/d1/96/2815b374d352831ddfdeb4e5f92ba98345626b71022e02a862\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60966 sha256=0a1cb5f6ba88e03e6b4af796949e388047f85019f30c03e3e7b6ce4a81bee3e3\n  Stored in directory: /root/.cache/pip/wheels/4f/89/a3/5cf59e30a8a75c917c313f14da0f6209be2d147e3160b985d6\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: efficientnet-pytorch, timm, pretrainedmodels, segmentation-models-pytorch\n  Attempting uninstall: timm\n    Found existing installation: timm 0.6.12\n    Uninstalling timm-0.6.12:\n      Successfully uninstalled timm-0.6.12\nSuccessfully installed efficientnet-pytorch-0.6.3 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.2.0 timm-0.4.12\n\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0mCollecting zarr==2.12.0\n  Downloading zarr-2.12.0-py3-none-any.whl (185 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m185.8/185.8 kB\u001b[0m \u001b[31m9.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hCollecting asciitree\n  Downloading asciitree-0.3.3.tar.gz (4.0 kB)\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hCollecting numcodecs>=0.6.4\n  Downloading numcodecs-0.10.2-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (6.6 MB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m6.6/6.6 MB\u001b[0m \u001b[31m15.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m00:01\u001b[0m00:01\u001b[0mm\n\u001b[?25hRequirement already satisfied: fasteners in /opt/conda/lib/python3.7/site-packages (from zarr==2.12.0) (0.18)\nRequirement already satisfied: numpy>=1.7 in /opt/conda/lib/python3.7/site-packages (from zarr==2.12.0) (1.21.6)\nRequirement already satisfied: entrypoints in /opt/conda/lib/python3.7/site-packages (from numcodecs>=0.6.4->zarr==2.12.0) (0.4)\nRequirement already satisfied: typing-extensions>=3.7.4 in /opt/conda/lib/python3.7/site-packages (from numcodecs>=0.6.4->zarr==2.12.0) (4.4.0)\nBuilding wheels for collected packages: asciitree\n  Building wheel for asciitree (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for asciitree: filename=asciitree-0.3.3-py3-none-any.whl size=5050 sha256=178eb3d0f5e29e0f031c0179e66f62ba26f7cad85a997922afedb87f53a2eb35\n  Stored in directory: /root/.cache/pip/wheels/e2/97/c4/5537ba28215ed3508783dc23c1fb59e17f00722317e4edeac0\nSuccessfully built asciitree\nInstalling collected packages: asciitree, numcodecs, zarr\nSuccessfully installed asciitree-0.3.3 numcodecs-0.10.2 zarr-2.12.0\n\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m","output_type":"stream"},{"name":"stderr","text":"  error: subprocess-exited-with-error\n  \n  × pip subprocess to install backend dependencies did not run successfully.\n  │ exit code: 1\n  ╰─> [3 lines of output]\n      ERROR: Ignored the following versions that require a different python version: 0.1.0 Requires-Python >=3.9; 0.1.1 Requires-Python >=3.9; 0.1.10 Requires-Python >=3.9; 0.1.11 Requires-Python >=3.9; 0.1.12 Requires-Python >=3.9; 0.1.13 Requires-Python >=3.9; 0.1.14 Requires-Python >=3.9; 0.1.2 Requires-Python >=3.9; 0.1.3 Requires-Python >=3.9; 0.1.4 Requires-Python >=3.9; 0.1.5 Requires-Python >=3.9; 0.1.6 Requires-Python >=3.9; 0.1.7 Requires-Python >=3.9; 0.1.8 Requires-Python >=3.9; 0.1.9 Requires-Python >=3.9\n      ERROR: Could not find a version that satisfies the requirement puccinialin (from versions: none)\n      ERROR: No matching distribution found for puccinialin\n      [end of output]\n  \n  note: This error originates from a subprocess, and is likely not a problem with pip.\nerror: subprocess-exited-with-error\n\n× pip subprocess to install backend dependencies did not run successfully.\n│ exit code: 1\n╰─> See above for output.\n\nnote: This error originates from a subprocess, and is likely not a problem with pip.\n","output_type":"stream"},{"name":"stdout","text":"[setup] segmentation_models_pytorch MiT (SegFormer) backbone available: False\nDevice: cuda\n[frag 2] volume shape -> (65, 14830, 9506)\n[frag 2] volume -> zarr done.\n","output_type":"stream"},{"name":"stderr","text":"/opt/conda/lib/python3.7/site-packages/PIL/Image.py:3077: DecompressionBombWarning: Image size (140973980 pixels) exceeds limit of 89478485 pixels, could be decompression bomb DOS attack.\n  DecompressionBombWarning,\n","output_type":"stream"},{"name":"stdout","text":"[frag 2] labels -> zarr done.\n[frag 3] volume shape -> (65, 7606, 5249)\n[frag 3] volume -> zarr done.\n[frag 3] labels -> zarr done.\n[frag 1] volume shape -> (65, 8181, 6330)\n[frag 1] volume -> zarr done.\n[frag 1] labels -> zarr done.\n[frag 2] split=train: 6957 foreground patches indexed (H=14830, W=9506).\n[frag 2] split=val: 1033 foreground patches indexed (H=14830, W=9506).\n[frag 3] split=train: 1745 foreground patches indexed (H=7606, W=5249).\n[frag 3] split=val: 294 foreground patches indexed (H=7606, W=5249).\nTOTAL train patches: 8702 | val patches: 1327\n[se3d_unet_d16_s0] epoch 1/12 train_loss=0.6958 val_F0.5=0.1882 P=0.162 R=0.549 thr=0.221 (125s)\n[se3d_unet_d16_s0] epoch 2/12 train_loss=0.6589 val_F0.5=0.1526 P=0.178 R=0.098 thr=0.200 (119s)\n[se3d_unet_d16_s0] epoch 3/12 train_loss=0.6575 val_F0.5=0.1471 P=0.162 R=0.107 thr=0.200 (119s)\n[se3d_unet_d16_s0] epoch 4/12 train_loss=0.6556 val_F0.5=0.1649 P=0.163 R=0.173 thr=0.200 (119s)\n[se3d_unet_d16_s0] epoch 5/12 train_loss=0.6511 val_F0.5=0.1572 P=0.149 R=0.199 thr=0.200 (120s)\n[se3d_unet_d16_s0] epoch 6/12 train_loss=0.6526 val_F0.5=0.2013 P=0.170 R=0.795 thr=0.200 (119s)\n[se3d_unet_d16_s0] epoch 7/12 train_loss=0.6484 val_F0.5=0.1735 P=0.152 R=0.397 thr=0.200 (120s)\n[se3d_unet_d16_s0] epoch 8/12 train_loss=0.6468 val_F0.5=0.1599 P=0.161 R=0.156 thr=0.201 (120s)\n[se3d_unet_d16_s0] epoch 9/12 train_loss=0.6462 val_F0.5=0.1950 P=0.165 R=0.719 thr=0.200 (120s)\n[se3d_unet_d16_s0] epoch 10/12 train_loss=0.6420 val_F0.5=0.1711 P=0.167 R=0.190 thr=0.200 (120s)\n[se3d_unet_d16_s0] epoch 11/12 train_loss=0.6422 val_F0.5=0.1733 P=0.163 R=0.234 thr=0.200 (120s)\n[se3d_unet_d16_s0] epoch 12/12 train_loss=0.6396 val_F0.5=0.1745 P=0.162 R=0.251 thr=0.200 (121s)\n[se3d_unet_d16_s0] saved best checkpoint (val F0.5=0.2013) -> /kaggle/working/checkpoints/se3d_unet_d16_s0.pt\n[se3d_unet_d20_s1] epoch 1/12 train_loss=0.7111 val_F0.5=0.1891 P=0.171 R=0.325 thr=0.271 (120s)\n[se3d_unet_d20_s1] epoch 2/12 train_loss=0.6599 val_F0.5=0.2076 P=0.175 R=0.788 thr=0.200 (119s)\n[se3d_unet_d20_s1] epoch 3/12 train_loss=0.6503 val_F0.5=0.1685 P=0.172 R=0.157 thr=0.200 (120s)\n[se3d_unet_d20_s1] epoch 4/12 train_loss=0.6550 val_F0.5=0.1632 P=0.156 R=0.198 thr=0.200 (120s)\n[se3d_unet_d20_s1] epoch 5/12 train_loss=0.6495 val_F0.5=0.1917 P=0.174 R=0.320 thr=0.200 (120s)\n[se3d_unet_d20_s1] epoch 6/12 train_loss=0.6450 val_F0.5=0.1663 P=0.165 R=0.170 thr=0.200 (120s)\n[se3d_unet_d20_s1] epoch 7/12 train_loss=0.6466 val_F0.5=0.1993 P=0.171 R=0.579 thr=0.200 (120s)\n[se3d_unet_d20_s1] epoch 8/12 train_loss=0.6442 val_F0.5=0.1951 P=0.191 R=0.213 thr=0.200 (120s)\n[se3d_unet_d20_s1] epoch 9/12 train_loss=0.6388 val_F0.5=0.1927 P=0.250 R=0.101 thr=0.200 (121s)\n[se3d_unet_d20_s1] epoch 10/12 train_loss=0.6361 val_F0.5=0.2082 P=0.197 R=0.267 thr=0.203 (120s)\n[se3d_unet_d20_s1] epoch 11/12 train_loss=0.6354 val_F0.5=0.2063 P=0.227 R=0.152 thr=0.269 (120s)\n[se3d_unet_d20_s1] epoch 12/12 train_loss=0.6365 val_F0.5=0.2143 P=0.220 R=0.195 thr=0.225 (120s)\n[se3d_unet_d20_s1] saved best checkpoint (val F0.5=0.2143) -> /kaggle/working/checkpoints/se3d_unet_d20_s1.pt\n[maxpool_seg_d16_s2] epoch 1/12 train_loss=0.6747 val_F0.5=0.1961 P=0.189 R=0.230 thr=0.200 (75s)\n[maxpool_seg_d16_s2] epoch 2/12 train_loss=0.6515 val_F0.5=0.1903 P=0.180 R=0.247 thr=0.200 (76s)\n[maxpool_seg_d16_s2] epoch 3/12 train_loss=0.6515 val_F0.5=0.1901 P=0.173 R=0.314 thr=0.200 (76s)\n[maxpool_seg_d16_s2] epoch 4/12 train_loss=0.6504 val_F0.5=0.1858 P=0.207 R=0.131 thr=0.200 (75s)\n[maxpool_seg_d16_s2] epoch 5/12 train_loss=0.6490 val_F0.5=0.1835 P=0.199 R=0.140 thr=0.200 (75s)\n[maxpool_seg_d16_s2] epoch 6/12 train_loss=0.6474 val_F0.5=0.2137 P=0.204 R=0.263 thr=0.200 (76s)\n[maxpool_seg_d16_s2] epoch 7/12 train_loss=0.6422 val_F0.5=0.2257 P=0.212 R=0.309 thr=0.200 (77s)\n[maxpool_seg_d16_s2] epoch 8/12 train_loss=0.6410 val_F0.5=0.2489 P=0.225 R=0.429 thr=0.200 (75s)\n[maxpool_seg_d16_s2] epoch 9/12 train_loss=0.6380 val_F0.5=0.2468 P=0.228 R=0.364 thr=0.200 (77s)\n[maxpool_seg_d16_s2] epoch 10/12 train_loss=0.6350 val_F0.5=0.2618 P=0.242 R=0.389 thr=0.200 (77s)\n[maxpool_seg_d16_s2] epoch 11/12 train_loss=0.6338 val_F0.5=0.2537 P=0.241 R=0.320 thr=0.200 (77s)\n[maxpool_seg_d16_s2] epoch 12/12 train_loss=0.6330 val_F0.5=0.2603 P=0.240 R=0.394 thr=0.200 (77s)\n[maxpool_seg_d16_s2] saved best checkpoint (val F0.5=0.2618) -> /kaggle/working/checkpoints/maxpool_seg_d16_s2.pt\n[se3d_unet_d12_s3] epoch 1/12 train_loss=0.6999 val_F0.5=0.2267 P=0.192 R=0.792 thr=0.291 (145s)\n[se3d_unet_d12_s3] epoch 2/12 train_loss=0.6576 val_F0.5=0.2074 P=0.186 R=0.392 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 3/12 train_loss=0.6552 val_F0.5=0.2027 P=0.188 R=0.291 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 4/12 train_loss=0.6528 val_F0.5=0.2070 P=0.188 R=0.347 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 5/12 train_loss=0.6539 val_F0.5=0.2034 P=0.183 R=0.374 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 6/12 train_loss=0.6483 val_F0.5=0.1811 P=0.178 R=0.195 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 7/12 train_loss=0.6473 val_F0.5=0.1786 P=0.179 R=0.178 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 8/12 train_loss=0.6460 val_F0.5=0.1881 P=0.182 R=0.217 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 9/12 train_loss=0.6462 val_F0.5=0.1745 P=0.237 R=0.085 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 10/12 train_loss=0.6437 val_F0.5=0.1937 P=0.188 R=0.223 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 11/12 train_loss=0.6430 val_F0.5=0.2009 P=0.192 R=0.250 thr=0.200 (145s)\n[se3d_unet_d12_s3] epoch 12/12 train_loss=0.6400 val_F0.5=0.2013 P=0.189 R=0.270 thr=0.200 (145s)\n[se3d_unet_d12_s3] saved best checkpoint (val F0.5=0.2267) -> /kaggle/working/checkpoints/se3d_unet_d12_s3.pt\n[maxpool_seg_d22_s4] epoch 1/12 train_loss=0.6754 val_F0.5=0.2024 P=0.182 R=0.369 thr=0.200 (75s)\n[maxpool_seg_d22_s4] epoch 2/12 train_loss=0.6485 val_F0.5=0.1911 P=0.190 R=0.194 thr=0.200 (75s)\n[maxpool_seg_d22_s4] epoch 3/12 train_loss=0.6515 val_F0.5=0.1828 P=0.171 R=0.256 thr=0.200 (74s)\n[maxpool_seg_d22_s4] epoch 4/12 train_loss=0.6487 val_F0.5=0.1885 P=0.218 R=0.123 thr=0.200 (73s)\n[maxpool_seg_d22_s4] epoch 5/12 train_loss=0.6496 val_F0.5=0.1922 P=0.221 R=0.126 thr=0.200 (74s)\n[maxpool_seg_d22_s4] epoch 6/12 train_loss=0.6446 val_F0.5=0.1824 P=0.194 R=0.147 thr=0.200 (74s)\n[maxpool_seg_d22_s4] epoch 7/12 train_loss=0.6449 val_F0.5=0.2115 P=0.187 R=0.447 thr=0.200 (73s)\n[maxpool_seg_d22_s4] epoch 8/12 train_loss=0.6403 val_F0.5=0.2238 P=0.216 R=0.264 thr=0.200 (74s)\n[maxpool_seg_d22_s4] epoch 9/12 train_loss=0.6411 val_F0.5=0.2320 P=0.208 R=0.425 thr=0.200 (74s)\n[maxpool_seg_d22_s4] epoch 10/12 train_loss=0.6394 val_F0.5=0.2434 P=0.230 R=0.317 thr=0.200 (75s)\n[maxpool_seg_d22_s4] epoch 11/12 train_loss=0.6342 val_F0.5=0.2367 P=0.216 R=0.382 thr=0.200 (75s)\n[maxpool_seg_d22_s4] epoch 12/12 train_loss=0.6340 val_F0.5=0.2448 P=0.223 R=0.395 thr=0.200 (74s)\n[maxpool_seg_d22_s4] saved best checkpoint (val F0.5=0.2448) -> /kaggle/working/checkpoints/maxpool_seg_d22_s4.pt\n(diagnostic only) mean per-model validation threshold: 0.223\n[fragment 1] Otsu threshold (unsupervised) = 0.304 | bounded grid-search threshold (oracle, uses labels) = 0.330 (F0.5=0.3577, P=0.340, R=0.448)\n[fragment 1] using method='otsu' -> final threshold=0.304\nTEST (fragment 1) @thr=0.304 [otsu] -> Precision=0.3019 Recall=0.6329 F0.5=0.3372 Dice=0.4088\nSaved metrics_summary.json\nSaved visualization -> /kaggle/working/viz/test_fragment1_comparison.png\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"\"\"\"\n==========================================================================================\nVESUVIUS CHALLENGE - INK DETECTION\nMemory-safe patch-based pipeline (Zarr storage, 3D->2D SE-UNet + MaxPool/SegFormer branch,\ntemporal augmentations, BCE+Dice, F0.5-optimized thresholding, small ensemble, viz).\n==========================================================================================\n\nHOW TO USE ON KAGGLE\n---------------------\n1. Paste each \"# %% [CELL n] ...\" block into its own notebook cell (recommended), OR just\n   run this whole file as a single script cell.\n2. Turn on a GPU accelerator (P100 / T4x2).\n3. Data is expected at:\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/surface_volume/00.tif ... 64.tif\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/inklabels.png\n4. Everything derived is written to /kaggle/working (Zarr stacks, checkpoints, plots, json).\n\nWHY THIS AVOIDS OOM\n--------------------\n- Raw fragments are converted ONCE into on-disk chunked Zarr arrays (uint8), slice-by-slice,\n  so we never hold a full (65, H, W) volume in RAM.\n- Training/inference never touch the full fragment either: we sample small 3D patches\n  (depth_window x tile x tile) directly out of the Zarr store on demand.\n- Mixed precision (AMP), small batch size + optional grad accumulation, aggressive\n  `del` + `gc.collect()` + `torch.cuda.empty_cache()`, and bounded DataLoader workers.\n- Test-fragment inference is also patch-based with an overlap-averaged stitching buffer\n  (float32 (H,W) accumulator, which for Vesuvius fragment sizes is a few hundred MB at most\n  - much smaller than the raw (65,H,W) uint16 volume would be).\n\"\"\"\n!pip install segmentation-models-pytorch==0.2.0\n!pip install zarr==2.12.0\n# %% [CELL 0] ---- Setup & installs -------------------------------------------------------\n\ndef _smp_has_mit(smp_module):\n    try:\n        from segmentation_models_pytorch.encoders import encoders as _enc\n        return \"mit_b3\" in _enc\n    except Exception:\n        return False\n \ntry:\n    import segmentation_models_pytorch as smp  # noqa\n    HAS_SMP = _smp_has_mit(smp)\nexcept ImportError:\n    smp = None\n    HAS_SMP = False\n \nif not HAS_SMP:\n    # Kaggle's preinstalled smp build is often old and lacks the timm-based MiT/SegFormer\n    # encoders (mit_b0..mit_b5). Try upgrading; if that fails (no internet / version pin\n    # issues), we fall back to the hand-rolled decoder below instead of crashing.\n    try:\n        _pip([\"-U\", \"segmentation-models-pytorch\", \"timm\"])\n        import importlib\n        if smp is not None:\n            importlib.reload(smp)\n        else:\n            import segmentation_models_pytorch as smp  # noqa\n        HAS_SMP = _smp_has_mit(smp)\n    except Exception as e:\n        print(f\"[setup] Could not get MiT/SegFormer support from segmentation_models_pytorch \"\n              f\"({e}); will use the built-in fallback decoder for the maxpool_seg branch.\")\n        HAS_SMP = False\n \nprint(f\"[setup] segmentation_models_pytorch MiT (SegFormer) backbone available: {HAS_SMP}\")\n \nimport os, gc, json, math, random, time, glob\nimport numpy as np\nimport zarr\nimport tifffile\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import precision_recall_curve\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\n\n# %% [CELL 1] ---- Config ------------------------------------------------------------------\nclass CFG:\n    DATA_ROOT   = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    WORK_DIR    = \"/kaggle/working\"\n    ZARR_DIR    = os.path.join(WORK_DIR, \"zarr_store\")\n    CKPT_DIR    = os.path.join(WORK_DIR, \"checkpoints\")\n    VIZ_DIR     = os.path.join(WORK_DIR, \"viz\")\n\n    TRAIN_FRAGMENTS = [2, 3]   # 80/20 spatial split inside each -> train/val\n    TEST_FRAGMENT    = 1       # fully held out\n\n    Z_SLICES = 65               # 00.tif .. 64.tif\n    Z_MID    = Z_SLICES // 2    # center of the stack\n\n    # depth window used per sample (Temporal Random Crop draws a window from this range)\n    DEPTH_MIN, DEPTH_MAX = 12, 22\n\n    TILE          = 224          # spatial patch size (H=W=TILE)\n    TRAIN_STRIDE  = 112          # 50% overlap while enumerating candidate train/val patches\n    TEST_STRIDE   = 112          # overlap for sliding-window test inference (averaged)\n\n    FG_MEAN_THRESH = 8.0 / 255.0   # skip near-empty (background/air) patches when indexing\n\n    VAL_FRACTION_BY_WIDTH = 0.8    # first 80% of fragment width -> train, last 20% -> val\n\n    BATCH_SIZE   = 8\n    ACCUM_STEPS  = 2               # effective batch = BATCH_SIZE * ACCUM_STEPS\n    NUM_WORKERS  = 2\n    EPOCHS       = 2\n    LR           = 3e-4\n    WEIGHT_DECAY = 1e-4\n\n    # patches sampled per epoch (subsample huge candidate lists -> bounds RAM & epoch time)\n    MAX_TRAIN_PATCHES_PER_EPOCH = 2500\n    MAX_VAL_PATCHES             = 600\n\n    # Ensemble: each entry is one trained model variant (architecture, depth window, seed)\n    ENSEMBLE_CONFIGS = [\n        dict(name=\"se3d_unet_d16_s0\",  arch=\"se3d_unet\",     depth=16, base_ch=24, seed=0),\n        dict(name=\"se3d_unet_d20_s1\",  arch=\"se3d_unet\",     depth=20, base_ch=24, seed=1),\n        dict(name=\"maxpool_seg_d16_s2\", arch=\"maxpool_seg\",  depth=16, base_ch=32, seed=2),\n        dict(name=\"se3d_unet_d12_s3\",  arch=\"se3d_unet\",     depth=12, base_ch=32, seed=3),\n        dict(name=\"maxpool_seg_d22_s4\", arch=\"maxpool_seg\",  depth=22, base_ch=32, seed=4),\n    ]\n\n    F_BETA = 0.5  # F_0.5 -> precision-weighted\n\nos.makedirs(CFG.ZARR_DIR, exist_ok=True)\nos.makedirs(CFG.CKPT_DIR, exist_ok=True)\nos.makedirs(CFG.VIZ_DIR, exist_ok=True)\n\n\n# %% [CELL 2] ---- TIFF stack -> chunked Zarr (slice-by-slice, no full-volume RAM load) ----\ndef convert_fragment_to_zarr(frag_id: int, cfg=CFG):\n    \"\"\"\n    Writes:\n      zarr_store/frag{frag_id}_volume.zarr  -> uint8 array, shape (Z, H, W), chunks (Z, 256, 256)\n      zarr_store/frag{frag_id}_labels.zarr  -> uint8 array, shape (H, W),   chunks (256, 256)\n    Skips conversion if already present.\n    \"\"\"\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n\n    frag_dir = os.path.join(cfg.DATA_ROOT, str(frag_id))\n    slice_paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in range(cfg.Z_SLICES)]\n    label_path  = os.path.join(frag_dir, \"inklabels.png\")\n\n    if os.path.exists(vol_path) and os.path.exists(lbl_path):\n        print(f\"[frag {frag_id}] zarr already exists, skipping conversion.\")\n        return vol_path, lbl_path\n\n    # peek shape from first slice\n    with tifffile.TiffFile(slice_paths[0]) as tf:\n        h, w = tf.pages[0].shape\n    print(f\"[frag {frag_id}] volume shape -> ({cfg.Z_SLICES}, {h}, {w})\")\n\n    store = zarr.DirectoryStore(vol_path)\n    root = zarr.group(store=store, overwrite=True)\n    vol_z = root.create_dataset(\n        \"data\", shape=(cfg.Z_SLICES, h, w), chunks=(cfg.Z_SLICES, 256, 256),\n        dtype=\"u1\", compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n\n    for i, sp in enumerate(slice_paths):\n        sl = tifffile.imread(sp)  # uint16, shape (H, W) -- ONE slice in RAM at a time\n        # normalize 16-bit -> 8-bit for compact on-disk storage / cheap I/O during training\n        sl8 = (sl.astype(np.float32) / 65535.0 * 255.0).clip(0, 255).astype(np.uint8)\n        vol_z[i, :, :] = sl8\n        del sl, sl8\n        if i % 16 == 0:\n            gc.collect()\n    print(f\"[frag {frag_id}] volume -> zarr done.\")\n\n    lbl_img = Image.open(label_path).convert(\"L\")\n    lbl = (np.array(lbl_img) > 127).astype(np.uint8)\n    assert lbl.shape == (h, w), f\"label shape {lbl.shape} != volume shape {(h, w)}\"\n\n    lstore = zarr.DirectoryStore(lbl_path)\n    lroot = zarr.group(store=lstore, overwrite=True)\n    lbl_z = lroot.create_dataset(\n        \"data\", shape=(h, w), chunks=(256, 256), dtype=\"u1\",\n        compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n    lbl_z[:, :] = lbl\n    del lbl, lbl_img\n    gc.collect()\n    print(f\"[frag {frag_id}] labels -> zarr done.\")\n    return vol_path, lbl_path\n\n\ndef open_fragment_zarr(frag_id, cfg=CFG):\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n    vol = zarr.open(vol_path, mode=\"r\")[\"data\"]\n    lbl = zarr.open(lbl_path, mode=\"r\")[\"data\"]\n    return vol, lbl\n\n\n# %% [CELL 3] ---- Patch indexing + 80/20 spatial split ------------------------------------\ndef build_patch_grid(h, w, tile, stride):\n    ys = list(range(0, max(h - tile, 0) + 1, stride))\n    xs = list(range(0, max(w - tile, 0) + 1, stride))\n    if ys[-1] != h - tile: ys.append(max(h - tile, 0))\n    if xs[-1] != w - tile: xs.append(max(w - tile, 0))\n    return [(y, x) for y in ys for x in xs]\n\n\ndef index_fragment_patches(frag_id, cfg=CFG, split=\"both\"):\n    \"\"\"\n    Returns list of dicts: {frag_id, y, x} for foreground patches only (cheap intensity\n    filter on the middle slice, read directly from zarr chunk-by-chunk -> low RAM).\n    split: 'train' -> first VAL_FRACTION_BY_WIDTH of width\n           'val'   -> remaining width\n           'both'  -> no split (used for the held-out test fragment)\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    mid_slice = np.asarray(vol[cfg.Z_MID, :, :])  # (H, W) uint8, one slice in RAM\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TRAIN_STRIDE if split != \"both\" else cfg.TEST_STRIDE)\n\n    split_x = int(W * cfg.VAL_FRACTION_BY_WIDTH)\n    kept = []\n    for (y, x) in coords:\n        if split == \"train\" and not (x + cfg.TILE <= split_x):\n            continue\n        if split == \"val\" and not (x >= split_x):\n            continue\n        patch = mid_slice[y:y + cfg.TILE, x:x + cfg.TILE]\n        if patch.size == 0:\n            continue\n        if (patch.astype(np.float32) / 255.0).mean() < cfg.FG_MEAN_THRESH:\n            continue  # skip empty background/air patch\n        kept.append({\"frag_id\": frag_id, \"y\": y, \"x\": x})\n    del mid_slice\n    gc.collect()\n    print(f\"[frag {frag_id}] split={split}: {len(kept)} foreground patches indexed \"\n          f\"(H={H}, W={W}).\")\n    return kept\n\n\n# %% [CELL 4] ---- Dataset with Temporal Random Crop / Random Paste / Temporal Cutout ------\nclass VesuviusPatchDataset(Dataset):\n    \"\"\"\n    Each item pulls a (depth_window, TILE, TILE) sub-volume straight out of the on-disk\n    Zarr store (no full-fragment load), applies the three temporal augmentations plus\n    standard spatial flips/rotations, and returns (volume_tensor, mask_tensor).\n    \"\"\"\n    def __init__(self, patch_records, depth_window, cfg=CFG, train=True, max_items=None):\n        self.records = patch_records\n        self.depth_window = depth_window\n        self.cfg = cfg\n        self.train = train\n        self._vol_cache = {}\n        self._lbl_cache = {}\n        if max_items is not None and len(self.records) > max_items:\n            self.records = random.sample(self.records, max_items)\n\n    def _get_arrays(self, frag_id):\n        if frag_id not in self._vol_cache:\n            self._vol_cache[frag_id], self._lbl_cache[frag_id] = open_fragment_zarr(frag_id, self.cfg)\n        return self._vol_cache[frag_id], self._lbl_cache[frag_id]\n\n    def __len__(self):\n        return len(self.records)\n\n    def _sample_depth_window(self, Z):\n        dw = self.depth_window if not self.train else random.randint(CFG.DEPTH_MIN, CFG.DEPTH_MAX)\n        dw = min(dw, Z)\n        z_mid = Z // 2\n        half = dw // 2\n        # Temporal Random Crop: random window around the middle of the stack\n        jitter = random.randint(-4, 4) if self.train else 0\n        z0 = max(0, min(Z - dw, z_mid - half + jitter))\n        return z0, dw\n\n    def __getitem__(self, idx):\n        rec = self.records[idx]\n        vol, lbl = self._get_arrays(rec[\"frag_id\"])\n        Z, H, W = vol.shape\n        y, x, t = rec[\"y\"], rec[\"x\"], self.cfg.TILE\n\n        z0, dw = self._sample_depth_window(Z)\n        sub = np.asarray(vol[z0:z0 + dw, y:y + t, x:x + t]).astype(np.float32) / 255.0  # (dw,t,t)\n        mask = np.asarray(lbl[y:y + t, x:x + t]).astype(np.float32)  # (t,t)\n\n        if self.train:\n            sub, mask = self._augment(sub, mask)\n\n        # pad depth to DEPTH_MAX so batches stack cleanly; padding slices are zeroed\n        # (network ignores all-zero slices thanks to Temporal Cutout training on real zeros too)\n        pad = CFG.DEPTH_MAX - sub.shape[0]\n        if pad > 0:\n            sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n\n        vol_t = torch.from_numpy(sub).unsqueeze(0).float()   # (1, D, H, W)\n        mask_t = torch.from_numpy(mask).unsqueeze(0).float() # (1, H, W)\n        return vol_t, mask_t\n\n    def _augment(self, sub, mask):\n        D = sub.shape[0]\n\n        # --- Random Paste: shift cropped slice-range to a random Z offset within the padded\n        #     budget, keeping slice order intact (sequential order preserved). ---\n        if random.random() < 0.5 and D < CFG.DEPTH_MAX:\n            max_shift = CFG.DEPTH_MAX - D\n            shift = random.randint(0, max_shift)\n            padded = np.zeros((CFG.DEPTH_MAX, *sub.shape[1:]), dtype=sub.dtype)\n            padded[shift:shift + D] = sub\n            sub = padded\n            D = CFG.DEPTH_MAX\n\n        # --- Temporal Cutout: zero out 1-2 random layers inside the active cube ---\n        n_cutout = random.choice([0, 1, 1, 2])\n        for _ in range(n_cutout):\n            zi = random.randint(0, D - 1)\n            sub[zi] = 0.0\n\n        # --- standard spatial augs ---\n        if random.random() < 0.5:\n            sub = sub[:, :, ::-1].copy(); mask = mask[:, ::-1].copy()\n        if random.random() < 0.5:\n            sub = sub[:, ::-1, :].copy(); mask = mask[::-1, :].copy()\n        k = random.choice([0, 1, 2, 3])\n        if k:\n            sub = np.rot90(sub, k, axes=(1, 2)).copy()\n            mask = np.rot90(mask, k, axes=(0, 1)).copy()\n\n        return sub, mask\n\n\n# %% [CELL 5] ---- Model: SE blocks, 3D encoder -> 2D SE-UNet decoder, MaxPool+Seg decoder --\nclass SEBlock2D(nn.Module):\n    \"\"\"Squeeze-and-Excitation applied inside skip connections.\"\"\"\n    def __init__(self, ch, reduction=8):\n        super().__init__()\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(ch, max(ch // reduction, 4), 1), nn.ReLU(inplace=True),\n            nn.Conv2d(max(ch // reduction, 4), ch, 1), nn.Sigmoid(),\n        )\n    def forward(self, x):\n        return x * self.fc(self.pool(x))\n\n\nclass ConvBNAct2D(nn.Module):\n    def __init__(self, cin, cout, k=3, s=1, p=1):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(cin, cout, k, s, p, bias=False),\n            nn.BatchNorm2d(cout), nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.block(x)\n\n\nclass Encoder3D(nn.Module):\n    \"\"\"3D conv encoder that progressively collapses the depth axis while extracting\n    multi-scale 2D feature maps (one per stage) for U-Net-style skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n        self.stage1 = nn.Sequential(\n            nn.Conv3d(1, c1, 3, padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n            nn.Conv3d(c1, c1, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n        )\n        self.stage2 = nn.Sequential(\n            nn.Conv3d(c1, c2, 3, padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n            nn.Conv3d(c2, c2, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n        )\n        self.stage3 = nn.Sequential(\n            nn.Conv3d(c2, c3, 3, padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n            nn.Conv3d(c3, c3, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n        )\n        self.stage4 = nn.Sequential(\n            # stride only on H,W here (depth stride=1): this must land the bottleneck one\n            # scale BELOW f3's spatial resolution, since the decoder's up3 upsamples it by\n            # 2x before concatenating with f3. Depth is fully collapsed right after by the\n            # adaptive pool, so no depth stride is needed at this stage.\n            nn.Conv3d(c3, c4, 3, stride=(1, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c4), nn.ReLU(inplace=True),\n            nn.AdaptiveMaxPool3d((1, None, None)),  # collapse remaining depth -> bottleneck\n        )\n        self.out_channels = (c1, c2, c3, c4)\n\n    @staticmethod\n    def _depth_collapse(x):\n        # max over depth to produce a 2D skip feature map at this scale\n        return x.max(dim=2).values\n\n    def forward(self, x):  # x: (B, 1, D, H, W)\n        s1 = self.stage1(x); f1 = self._depth_collapse(s1)          # H/2\n        s2 = self.stage2(s1); f2 = self._depth_collapse(s2)          # H/4\n        s3 = self.stage3(s2); f3 = self._depth_collapse(s3)          # H/8\n        s4 = self.stage4(s3).squeeze(2)                               # H/8, depth->1\n        return f1, f2, f3, s4\n\n\nclass SEUnetDecoder2D(nn.Module):\n    def __init__(self, enc_channels):\n        super().__init__()\n        c1, c2, c3, c4 = enc_channels\n        self.up3 = nn.ConvTranspose2d(c4, c3, 2, 2)\n        self.se3 = SEBlock2D(c3); self.dec3 = ConvBNAct2D(c3 * 2, c3)\n        self.up2 = nn.ConvTranspose2d(c3, c2, 2, 2)\n        self.se2 = SEBlock2D(c2); self.dec2 = ConvBNAct2D(c2 * 2, c2)\n        self.up1 = nn.ConvTranspose2d(c2, c1, 2, 2)\n        self.se1 = SEBlock2D(c1); self.dec1 = ConvBNAct2D(c1 * 2, c1)\n        self.up0 = nn.ConvTranspose2d(c1, c1 // 2, 2, 2)\n        self.final = nn.Conv2d(c1 // 2, 1, 1)\n\n    def forward(self, f1, f2, f3, bottleneck, out_hw):\n        x = self.up3(bottleneck); x = self.dec3(torch.cat([x, self.se3(f3)], dim=1))\n        x = self.up2(x);          x = self.dec2(torch.cat([x, self.se2(f2)], dim=1))\n        x = self.up1(x);          x = self.dec1(torch.cat([x, self.se1(f1)], dim=1))\n        x = self.up0(x)\n        x = F.interpolate(x, size=out_hw, mode=\"bilinear\", align_corners=False)\n        return self.final(x)\n\n\nclass SE3DUNet(nn.Module):\n    \"\"\"Top-solution-style architecture: 3D encoder + 2D decoder U-Net with SE blocks in\n    the skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        self.encoder = Encoder3D(base_ch)\n        self.decoder = SEUnetDecoder2D(self.encoder.out_channels)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        h, w = x.shape[-2], x.shape[-1]\n        f1, f2, f3, bott = self.encoder(x)\n        return self.decoder(f1, f2, f3, bott, (h, w))\n\n\nclass MaxPoolFlattenSegHead(nn.Module):\n    \"\"\"Depth Flattening branch: collapse the (D,H,W) volume to a 2D feature map via\n    max pooling across depth, then feed a heavy semantic-segmentation decoder\n    (SegFormer/MiT-style if `segmentation_models_pytorch` is available, otherwise a\n    hand-rolled multi-scale conv decoder with SE-augmented skips).\"\"\"\n    def __init__(self, base_ch=32, use_smp=HAS_SMP):\n        super().__init__()\n        self.use_smp = False\n        if use_smp:\n            try:\n                self.net = smp.Unet(\n                    encoder_name=\"mit_b3\", encoder_weights=\"imagenet\",\n                    in_channels=1, classes=1, activation=None,\n                )\n                self.use_smp = True\n            except Exception as e:\n                print(f\"[MaxPoolFlattenSegHead] smp mit_b3 unavailable at model-build time \"\n                      f\"({e}); using the built-in fallback decoder instead.\")\n\n        if not self.use_smp:\n            c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n            self.enc1 = nn.Sequential(ConvBNAct2D(1, c1), ConvBNAct2D(c1, c1))\n            self.pool1 = nn.MaxPool2d(2)\n            self.enc2 = nn.Sequential(ConvBNAct2D(c1, c2), ConvBNAct2D(c2, c2))\n            self.pool2 = nn.MaxPool2d(2)\n            self.enc3 = nn.Sequential(ConvBNAct2D(c2, c3), ConvBNAct2D(c3, c3))\n            self.pool3 = nn.MaxPool2d(2)\n            self.bott = nn.Sequential(ConvBNAct2D(c3, c4), ConvBNAct2D(c4, c4))\n            self.decoder = SEUnetDecoder2D((c1, c2, c3, c4))\n\n    def forward(self, x2d):  # x2d: (B, 1, H, W)  (already depth-flattened)\n        if self.use_smp:\n            return self.net(x2d)\n        h, w = x2d.shape[-2], x2d.shape[-1]\n        e1 = self.enc1(x2d); p1 = self.pool1(e1)\n        e2 = self.enc2(p1);  p2 = self.pool2(e2)\n        e3 = self.enc3(p2);  p3 = self.pool3(e3)\n        b  = self.bott(p3)\n        return self.decoder(e1, e2, e3, b, (h, w))\n\n\nclass MaxPoolSegModel(nn.Module):\n    \"\"\"Full pipeline for the 'Depth Flattening' approach: Max Pool across the depth axis\n    of the 3D volume -> 2D feature map -> heavy segmentation decoder -> ink mask.\"\"\"\n    def __init__(self, base_ch=32):\n        super().__init__()\n        self.head = MaxPoolFlattenSegHead(base_ch=base_ch)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        x2d = x.max(dim=2).values  # Max Pooling across depth axis -> (B,1,H,W)\n        return self.head(x2d)\n\n\ndef build_model(arch, base_ch):\n    if arch == \"se3d_unet\":\n        return SE3DUNet(base_ch=base_ch)\n    elif arch == \"maxpool_seg\":\n        return MaxPoolSegModel(base_ch=base_ch)\n    else:\n        raise ValueError(arch)\n\n\n# %% [CELL 6] ---- Loss (BCE + Dice) and F-beta metric/threshold search --------------------\nclass BCEDiceLoss(nn.Module):\n    def __init__(self, bce_w=0.5, dice_w=0.5, smooth=1.0):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.bce_w, self.dice_w, self.smooth = bce_w, dice_w, smooth\n\n    def forward(self, logits, target):\n        bce_loss = self.bce(logits, target)\n        prob = torch.sigmoid(logits)\n        p, t = prob.flatten(1), target.flatten(1)\n        inter = (p * t).sum(1)\n        dice = 1 - (2 * inter + self.smooth) / (p.sum(1) + t.sum(1) + self.smooth)\n        return self.bce_w * bce_loss + self.dice_w * dice.mean()\n\n\ndef fbeta_score(precision, recall, beta=CFG.F_BETA, eps=1e-8):\n    b2 = beta ** 2\n    return (1 + b2) * precision * recall / (b2 * precision + recall + eps)\n\n\ndef find_best_threshold(probs_flat: np.ndarray, targets_flat: np.ndarray, beta=CFG.F_BETA):\n    \"\"\"Optimizes the decision threshold strictly for F_beta (beta<1 -> precision-weighted).\"\"\"\n    precision, recall, thresholds = precision_recall_curve(targets_flat, probs_flat)\n    precision, recall = precision[:-1], recall[:-1]\n    scores = fbeta_score(precision, recall, beta)\n    best_idx = int(np.nanargmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), float(precision[best_idx]), float(recall[best_idx])\n\n\n# %% [CELL 7] ---- Training loop for one ensemble member ------------------------------------\ndef run_validation(model, val_loader):\n    model.eval()\n    all_probs, all_targets = [], []\n    with torch.no_grad(), autocast():\n        for vol, mask in val_loader:\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            logits = model(vol)\n            probs = torch.sigmoid(logits).float().cpu().numpy().ravel()\n            targs = mask.float().cpu().numpy().ravel()\n            # subsample pixels per patch to keep the PR-curve computation light on RAM\n            if len(probs) > 20000:\n                idx = np.random.choice(len(probs), 20000, replace=False)\n                probs, targs = probs[idx], targs[idx]\n            all_probs.append(probs); all_targets.append(targs)\n    probs = np.concatenate(all_probs); targets = np.concatenate(all_targets)\n    thr, f05, prec, rec = find_best_threshold(probs, targets)\n    del all_probs, all_targets\n    gc.collect()\n    return thr, f05, prec, rec\n\n\ndef train_one_model(cfg_entry, train_records, val_records, cfg=CFG, epochs=None):\n    epochs = epochs or cfg.EPOCHS\n    torch.manual_seed(cfg_entry[\"seed\"]); random.seed(cfg_entry[\"seed\"])\n\n    model = build_model(cfg_entry[\"arch\"], cfg_entry[\"base_ch\"]).to(DEVICE)\n    opt = torch.optim.AdamW(model.parameters(), lr=cfg.LR, weight_decay=cfg.WEIGHT_DECAY)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    scaler = GradScaler()\n    criterion = BCEDiceLoss()\n\n    train_ds = VesuviusPatchDataset(train_records, cfg_entry[\"depth\"], cfg, train=True,\n                                     max_items=cfg.MAX_TRAIN_PATCHES_PER_EPOCH)\n    val_ds = VesuviusPatchDataset(val_records, cfg_entry[\"depth\"], cfg, train=False,\n                                   max_items=cfg.MAX_VAL_PATCHES)\n\n    best_f05, best_thr, best_state = -1.0, 0.5, None\n    history = []\n\n    for epoch in range(epochs):\n        # resample the train subset each epoch for coverage without ever loading everything\n        train_ds.records = random.sample(train_records, min(cfg.MAX_TRAIN_PATCHES_PER_EPOCH, len(train_records)))\n        train_loader = DataLoader(train_ds, batch_size=cfg.BATCH_SIZE, shuffle=True,\n                                   num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True)\n        val_loader = DataLoader(val_ds, batch_size=cfg.BATCH_SIZE, shuffle=False,\n                                 num_workers=cfg.NUM_WORKERS, pin_memory=True)\n\n        model.train()\n        running_loss, n_steps = 0.0, 0\n        opt.zero_grad()\n        t0 = time.time()\n        for step, (vol, mask) in enumerate(train_loader):\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            with autocast():\n                logits = model(vol)\n                loss = criterion(logits, mask) / cfg.ACCUM_STEPS\n            scaler.scale(loss).backward()\n            if (step + 1) % cfg.ACCUM_STEPS == 0:\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n            running_loss += loss.item() * cfg.ACCUM_STEPS\n            n_steps += 1\n            del vol, mask, logits, loss\n        sched.step()\n\n        thr, f05, prec, rec = run_validation(model, val_loader)\n        dt = time.time() - t0\n        print(f\"[{cfg_entry['name']}] epoch {epoch+1}/{epochs} \"\n              f\"train_loss={running_loss/max(n_steps,1):.4f} \"\n              f\"val_F0.5={f05:.4f} P={prec:.3f} R={rec:.3f} thr={thr:.3f} ({dt:.0f}s)\")\n        history.append(dict(epoch=epoch+1, train_loss=running_loss/max(n_steps,1),\n                             val_f05=f05, val_precision=prec, val_recall=rec, threshold=thr))\n\n        if f05 > best_f05:\n            best_f05, best_thr = f05, thr\n            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n\n        del train_loader, val_loader\n        gc.collect(); torch.cuda.empty_cache()\n\n    model.load_state_dict(best_state)\n    ckpt_path = os.path.join(cfg.CKPT_DIR, f\"{cfg_entry['name']}.pt\")\n    torch.save({\"state_dict\": best_state, \"config\": cfg_entry,\n                \"best_val_f05\": best_f05, \"best_threshold\": best_thr,\n                \"history\": history}, ckpt_path)\n    print(f\"[{cfg_entry['name']}] saved best checkpoint (val F0.5={best_f05:.4f}) -> {ckpt_path}\")\n\n    del model, opt, sched, scaler, train_ds, val_ds\n    gc.collect(); torch.cuda.empty_cache()\n    return ckpt_path, best_thr, best_f05, history\n\n\n# %% [CELL 8] ---- Sliding-window ensemble inference over a full fragment ------------------\ndef predict_fragment_ensemble(frag_id, model_infos, cfg=CFG):\n    \"\"\"\n    model_infos: list of dicts {ckpt_path, config} for each ensemble member.\n    Returns (prob_map (H,W) float32 averaged across models, gt_mask (H,W) uint8).\n    Memory-safe: only a float32 (H,W) accumulator + weight map live in RAM (no full volume).\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    gt = np.asarray(lbl[:, :])\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TEST_STRIDE)\n\n    loaded_models = []\n    for info in model_infos:\n        ckpt = torch.load(info[\"ckpt_path\"], map_location=DEVICE)\n        m = build_model(ckpt[\"config\"][\"arch\"], ckpt[\"config\"][\"base_ch\"]).to(DEVICE)\n        m.load_state_dict(ckpt[\"state_dict\"]); m.eval()\n        loaded_models.append((m, ckpt[\"config\"][\"depth\"]))\n\n    prob_acc = np.zeros((H, W), dtype=np.float32)\n    weight_acc = np.zeros((H, W), dtype=np.float32)\n\n    with torch.no_grad(), autocast():\n        for (y, x) in coords:\n            patch_sum = np.zeros((cfg.TILE, cfg.TILE), dtype=np.float32)\n            for model, depth in loaded_models:\n                z_mid = Z // 2\n                z0 = max(0, min(Z - depth, z_mid - depth // 2))\n                sub = np.asarray(vol[z0:z0 + depth, y:y + cfg.TILE, x:x + cfg.TILE]).astype(np.float32) / 255.0\n                pad = cfg.DEPTH_MAX - sub.shape[0]\n                if pad > 0:\n                    sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n                t = torch.from_numpy(sub).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n                logits = model(t)\n                patch_sum += torch.sigmoid(logits)[0, 0].float().cpu().numpy()\n                del t, logits\n            patch_prob = patch_sum / len(loaded_models)\n            prob_acc[y:y + cfg.TILE, x:x + cfg.TILE] += patch_prob\n            weight_acc[y:y + cfg.TILE, x:x + cfg.TILE] += 1.0\n\n    weight_acc[weight_acc == 0] = 1.0\n    prob_map = prob_acc / weight_acc\n\n    for m, _ in loaded_models:\n        del m\n    del loaded_models\n    gc.collect(); torch.cuda.empty_cache()\n    return prob_map, gt\n\n\n# %% [CELL 9] ---- Orchestration: convert data, split, train ensemble, evaluate, visualize --\ndef main():\n    # 1) Convert all needed fragments to Zarr (train 2,3 + held-out test 1)\n    for fid in CFG.TRAIN_FRAGMENTS + [CFG.TEST_FRAGMENT]:\n        convert_fragment_to_zarr(fid)\n\n    # 2) Build 80/20 spatial split from fragments 2 & 3\n    train_records, val_records = [], []\n    for fid in CFG.TRAIN_FRAGMENTS:\n        train_records += index_fragment_patches(fid, split=\"train\")\n        val_records   += index_fragment_patches(fid, split=\"val\")\n    print(f\"TOTAL train patches: {len(train_records)} | val patches: {len(val_records)}\")\n\n    # 3) Train each ensemble member\n    trained = []\n    for cfg_entry in CFG.ENSEMBLE_CONFIGS:\n        ckpt_path, thr, f05, hist = train_one_model(cfg_entry, train_records, val_records)\n        trained.append({\"ckpt_path\": ckpt_path, \"config\": cfg_entry, \"val_threshold\": thr, \"val_f05\": f05})\n\n    # 4) Aggregate validation threshold (mean across members) -> apply to test\n    ensemble_val_thr = float(np.mean([t[\"val_threshold\"] for t in trained]))\n    print(f\"Ensemble-averaged validation threshold: {ensemble_val_thr:.3f}\")\n\n    # 5) Full-fragment ensemble inference on the held-out TEST fragment (fragment 1)\n    prob_map, gt_map = predict_fragment_ensemble(CFG.TEST_FRAGMENT, trained)\n    pred_mask = (prob_map >= ensemble_val_thr).astype(np.uint8)\n\n    # 6) Final test metrics (precision/recall/F0.5/Dice) at the chosen threshold\n    p_flat, t_flat = pred_mask.ravel().astype(np.float32), gt_map.ravel().astype(np.float32)\n    tp = float((p_flat * t_flat).sum())\n    fp = float((p_flat * (1 - t_flat)).sum())\n    fn = float(((1 - p_flat) * t_flat).sum())\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n    f05_test = fbeta_score(precision, recall, CFG.F_BETA)\n    dice_test = 2 * tp / (2 * tp + fp + fn + 1e-8)\n    print(f\"TEST (fragment {CFG.TEST_FRAGMENT}) @thr={ensemble_val_thr:.3f} -> \"\n          f\"Precision={precision:.4f} Recall={recall:.4f} F0.5={f05_test:.4f} Dice={dice_test:.4f}\")\n\n    metrics = dict(\n        ensemble_val_threshold=ensemble_val_thr,\n        test_precision=precision, test_recall=recall,\n        test_f0_5=f05_test, test_dice=dice_test,\n        members=[{\"name\": t[\"config\"][\"name\"], \"val_f05\": t[\"val_f05\"],\n                  \"val_threshold\": t[\"val_threshold\"]} for t in trained],\n    )\n    with open(os.path.join(CFG.WORK_DIR, \"metrics_summary.json\"), \"w\") as f:\n        json.dump(metrics, f, indent=2)\n    print(\"Saved metrics_summary.json\")\n\n    # 7) Visualization: input (mid-slice + max-projection) vs ground truth vs prediction\n    vol, _ = open_fragment_zarr(CFG.TEST_FRAGMENT)\n    mid_slice = np.asarray(vol[CFG.Z_MID, :, :])\n    maxproj = np.asarray(vol[:, :, :]).max(axis=0) if vol.shape[0] * vol.shape[1] * vol.shape[2] < 4e8 \\\n        else mid_slice  # guard: skip full maxproj on huge volumes to avoid RAM spikes\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(mid_slice, cmap=\"gray\"); axes[0].set_title(f\"Input (mid slice z={CFG.Z_MID})\")\n    axes[1].imshow(gt_map, cmap=\"gray\"); axes[1].set_title(\"Ground Truth Ink Mask\")\n    axes[2].imshow(prob_map, cmap=\"magma\"); axes[2].set_title(\"Predicted Probability Map\")\n    axes[3].imshow(mid_slice, cmap=\"gray\")\n    axes[3].imshow(np.ma.masked_where(pred_mask == 0, pred_mask), cmap=\"autumn\", alpha=0.6)\n    axes[3].set_title(f\"Prediction Overlay (thr={ensemble_val_thr:.2f})\")\n    for ax in axes: ax.axis(\"off\")\n    plt.tight_layout()\n    viz_path = os.path.join(CFG.VIZ_DIR, f\"test_fragment{CFG.TEST_FRAGMENT}_comparison.png\")\n    plt.savefig(viz_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved visualization -> {viz_path}\")\n\n    del mid_slice, maxproj, prob_map, gt_map, pred_mask\n    gc.collect(); torch.cuda.empty_cache()\n    return metrics\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-03T07:05:34.767972Z","iopub.execute_input":"2026-08-03T07:05:34.768432Z","iopub.status.idle":"2026-08-03T08:40:25.883661Z","shell.execute_reply.started":"2026-08-03T07:05:34.768398Z","shell.execute_reply":"2026-08-03T08:40:25.882341Z"}},"outputs":[{"name":"stdout","text":"Collecting segmentation-models-pytorch==0.2.0\n  Downloading segmentation_models_pytorch-0.2.0-py3-none-any.whl (87 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m87.6/87.6 kB\u001b[0m \u001b[31m5.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hCollecting timm==0.4.12\n  Downloading timm-0.4.12-py3-none-any.whl (376 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m377.0/377.0 kB\u001b[0m \u001b[31m14.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hCollecting pretrainedmodels==0.7.4\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m5.4 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hCollecting efficientnet-pytorch==0.6.3\n  Downloading efficientnet_pytorch-0.6.3.tar.gz (16 kB)\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hRequirement already satisfied: torchvision>=0.5.0 in /opt/conda/lib/python3.7/site-packages (from segmentation-models-pytorch==0.2.0) (0.14.0)\nRequirement already satisfied: torch in /opt/conda/lib/python3.7/site-packages (from efficientnet-pytorch==0.6.3->segmentation-models-pytorch==0.2.0) (1.13.0)\nRequirement already satisfied: munch in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (2.5.0)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.7/site-packages (from pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (4.64.1)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (4.4.0)\nRequirement already satisfied: pillow!=8.3.*,>=5.3.0 in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (9.4.0)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.21.6)\nRequirement already satisfied: requests in /opt/conda/lib/python3.7/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.28.2)\nRequirement already satisfied: six in /opt/conda/lib/python3.7/site-packages (from munch->pretrainedmodels==0.7.4->segmentation-models-pytorch==0.2.0) (1.16.0)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2.1.1)\nRequirement already satisfied: urllib3<1.27,>=1.21.1 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (1.26.14)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (3.4)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.7/site-packages (from requests->torchvision>=0.5.0->segmentation-models-pytorch==0.2.0) (2022.12.7)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.6.3-py3-none-any.whl size=12422 sha256=0f0208e2a1fb7a8f47dc87645ff0fe69a50779e973f5737cb595847080341ec9\n  Stored in directory: /root/.cache/pip/wheels/d9/d1/96/2815b374d352831ddfdeb4e5f92ba98345626b71022e02a862\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60966 sha256=ccf0f9531f2233b02bebc96415fa780d82915dcf34206bd8f2d3ca2ed07baba8\n  Stored in directory: /root/.cache/pip/wheels/4f/89/a3/5cf59e30a8a75c917c313f14da0f6209be2d147e3160b985d6\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: efficientnet-pytorch, timm, pretrainedmodels, segmentation-models-pytorch\n  Attempting uninstall: timm\n    Found existing installation: timm 0.6.12\n    Uninstalling timm-0.6.12:\n      Successfully uninstalled timm-0.6.12\nSuccessfully installed efficientnet-pytorch-0.6.3 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.2.0 timm-0.4.12\n\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0mCollecting zarr==2.12.0\n  Downloading zarr-2.12.0-py3-none-any.whl (185 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m185.8/185.8 kB\u001b[0m \u001b[31m10.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hRequirement already satisfied: fasteners in /opt/conda/lib/python3.7/site-packages (from zarr==2.12.0) (0.18)\nCollecting numcodecs>=0.6.4\n  Downloading numcodecs-0.10.2-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (6.6 MB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m6.6/6.6 MB\u001b[0m \u001b[31m41.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m00:01\u001b[0m:00:01\u001b[0m\n\u001b[?25hRequirement already satisfied: numpy>=1.7 in /opt/conda/lib/python3.7/site-packages (from zarr==2.12.0) (1.21.6)\nCollecting asciitree\n  Downloading asciitree-0.3.3.tar.gz (4.0 kB)\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hRequirement already satisfied: entrypoints in /opt/conda/lib/python3.7/site-packages (from numcodecs>=0.6.4->zarr==2.12.0) (0.4)\nRequirement already satisfied: typing-extensions>=3.7.4 in /opt/conda/lib/python3.7/site-packages (from numcodecs>=0.6.4->zarr==2.12.0) (4.4.0)\nBuilding wheels for collected packages: asciitree\n  Building wheel for asciitree (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for asciitree: filename=asciitree-0.3.3-py3-none-any.whl size=5050 sha256=4c152bca823ebcaa504225179275b808267c6aa167497ed9c39a99e595da3120\n  Stored in directory: /root/.cache/pip/wheels/e2/97/c4/5537ba28215ed3508783dc23c1fb59e17f00722317e4edeac0\nSuccessfully built asciitree\nInstalling collected packages: asciitree, numcodecs, zarr\nSuccessfully installed asciitree-0.3.3 numcodecs-0.10.2 zarr-2.12.0\n\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m[setup] Could not get MiT/SegFormer support from segmentation_models_pytorch (name '_pip' is not defined); will use the built-in fallback decoder for the maxpool_seg branch.\n[setup] segmentation_models_pytorch MiT (SegFormer) backbone available: False\nDevice: cuda\n[frag 2] volume shape -> (65, 14830, 9506)\n[frag 2] volume -> zarr done.\n","output_type":"stream"},{"name":"stderr","text":"/opt/conda/lib/python3.7/site-packages/PIL/Image.py:3077: DecompressionBombWarning: Image size (140973980 pixels) exceeds limit of 89478485 pixels, could be decompression bomb DOS attack.\n  DecompressionBombWarning,\n","output_type":"stream"},{"name":"stdout","text":"[frag 2] labels -> zarr done.\n[frag 3] volume shape -> (65, 7606, 5249)\n[frag 3] volume -> zarr done.\n[frag 3] labels -> zarr done.\n[frag 1] volume shape -> (65, 8181, 6330)\n[frag 1] volume -> zarr done.\n[frag 1] labels -> zarr done.\n[frag 2] split=train: 6957 foreground patches indexed (H=14830, W=9506).\n[frag 2] split=val: 1033 foreground patches indexed (H=14830, W=9506).\n[frag 3] split=train: 1745 foreground patches indexed (H=7606, W=5249).\n[frag 3] split=val: 294 foreground patches indexed (H=7606, W=5249).\nTOTAL train patches: 8702 | val patches: 1327\n[se3d_unet_d16_s0] epoch 1/2 train_loss=0.6978 val_F0.5=0.1870 P=0.160 R=0.592 thr=0.154 (124s)\n[se3d_unet_d16_s0] epoch 2/2 train_loss=0.6636 val_F0.5=0.2016 P=0.168 R=0.946 thr=0.082 (118s)\n[se3d_unet_d16_s0] saved best checkpoint (val F0.5=0.2016) -> /kaggle/working/checkpoints/se3d_unet_d16_s0.pt\n[se3d_unet_d20_s1] epoch 1/2 train_loss=0.7116 val_F0.5=0.2112 P=0.181 R=0.624 thr=0.260 (119s)\n[se3d_unet_d20_s1] epoch 2/2 train_loss=0.6657 val_F0.5=0.2107 P=0.177 R=0.935 thr=0.130 (119s)\n[se3d_unet_d20_s1] saved best checkpoint (val F0.5=0.2112) -> /kaggle/working/checkpoints/se3d_unet_d20_s1.pt\n[maxpool_seg_d16_s2] epoch 1/2 train_loss=0.6778 val_F0.5=0.2254 P=0.190 R=0.901 thr=0.113 (68s)\n[maxpool_seg_d16_s2] epoch 2/2 train_loss=0.6569 val_F0.5=0.2270 P=0.191 R=0.884 thr=0.111 (72s)\n[maxpool_seg_d16_s2] saved best checkpoint (val F0.5=0.2270) -> /kaggle/working/checkpoints/maxpool_seg_d16_s2.pt\n[se3d_unet_d12_s3] epoch 1/2 train_loss=0.7017 val_F0.5=0.2197 P=0.190 R=0.580 thr=0.229 (144s)\n[se3d_unet_d12_s3] epoch 2/2 train_loss=0.6618 val_F0.5=0.2249 P=0.191 R=0.792 thr=0.150 (144s)\n[se3d_unet_d12_s3] saved best checkpoint (val F0.5=0.2249) -> /kaggle/working/checkpoints/se3d_unet_d12_s3.pt\n[maxpool_seg_d22_s4] epoch 1/2 train_loss=0.6787 val_F0.5=0.2187 P=0.184 R=0.909 thr=0.150 (72s)\n[maxpool_seg_d22_s4] epoch 2/2 train_loss=0.6534 val_F0.5=0.2184 P=0.184 R=0.830 thr=0.109 (68s)\n[maxpool_seg_d22_s4] saved best checkpoint (val F0.5=0.2187) -> /kaggle/working/checkpoints/maxpool_seg_d22_s4.pt\nEnsemble-averaged validation threshold: 0.151\nTEST (fragment 1) @thr=0.151 -> Precision=0.1953 Recall=0.9996 F0.5=0.2328 Dice=0.3268\nSaved metrics_summary.json\nSaved visualization -> /kaggle/working/viz/test_fragment1_comparison.png\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"#same above code using their pre-process reulted data \n\"\"\"\n==========================================================================================\nVESUVIUS CHALLENGE - INK DETECTION\nMemory-safe patch-based pipeline (Zarr storage, 3D->2D SE-UNet + MaxPool/SegFormer branch,\ntemporal augmentations, BCE+Dice, F0.5-optimized thresholding, small ensemble, viz).\n==========================================================================================\n\nHOW TO USE ON KAGGLE\n---------------------\n1. Paste each \"# %% [CELL n] ...\" block into its own notebook cell (recommended), OR just\n   run this whole file as a single script cell.\n2. Turn on a GPU accelerator (P100 / T4x2).\n3. Data is expected at:\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/surface_volume/00.tif ... 64.tif\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/inklabels.png\n4. Everything derived is written to /kaggle/working (Zarr stacks, checkpoints, plots, json).\n\nWHY THIS AVOIDS OOM\n--------------------\n- Raw fragments are converted ONCE into on-disk chunked Zarr arrays (uint8), slice-by-slice,\n  so we never hold a full (65, H, W) volume in RAM.\n- Training/inference never touch the full fragment either: we sample small 3D patches\n  (depth_window x tile x tile) directly out of the Zarr store on demand.\n- Mixed precision (AMP), small batch size + optional grad accumulation, aggressive\n  `del` + `gc.collect()` + `torch.cuda.empty_cache()`, and bounded DataLoader workers.\n- Test-fragment inference is also patch-based with an overlap-averaged stitching buffer\n  (float32 (H,W) accumulator, which for Vesuvius fragment sizes is a few hundred MB at most\n  - much smaller than the raw (65,H,W) uint16 volume would be).\n\"\"\"\n\n# %% [CELL 0] ---- Setup & installs -------------------------------------------------------\n\ndef _smp_has_mit(smp_module):\n    try:\n        from segmentation_models_pytorch.encoders import encoders as _enc\n        return \"mit_b3\" in _enc\n    except Exception:\n        return False\n\ntry:\n    import segmentation_models_pytorch as smp  # noqa\n    HAS_SMP = _smp_has_mit(smp)\nexcept ImportError:\n    smp = None\n    HAS_SMP = False\n\nif not HAS_SMP:\n    # Kaggle's preinstalled smp build is often old and lacks the timm-based MiT/SegFormer\n    # encoders (mit_b0..mit_b5). Try upgrading; if that fails (no internet / version pin\n    # issues), we fall back to the hand-rolled decoder below instead of crashing.\n    # NOTE: this block is intentionally self-contained (own subprocess/sys import + inline\n    # pip call) so it can't NameError even if this cell is ever re-run independently of the\n    # cell that defines the top-level `_pip` helper.\n    try:\n        import subprocess as _subprocess, sys as _sys\n        _subprocess.run([_sys.executable, \"-m\", \"pip\", \"install\", \"-q\",\n                          \"-U\", \"segmentation-models-pytorch\", \"timm\"])\n        import importlib\n        if smp is not None:\n            importlib.reload(smp)\n        else:\n            import segmentation_models_pytorch as smp  # noqa\n        HAS_SMP = _smp_has_mit(smp)\n    except Exception as e:\n        print(f\"[setup] Could not get MiT/SegFormer support from segmentation_models_pytorch \"\n              f\"({e}); will use the built-in fallback decoder for the maxpool_seg branch.\")\n        HAS_SMP = False\n\nprint(f\"[setup] segmentation_models_pytorch MiT (SegFormer) backbone available: {HAS_SMP}\")\n\nimport os, gc, json, math, random, time, glob\nimport numpy as np\nimport zarr\nimport tifffile\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import precision_recall_curve\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\n\n# %% [CELL 1] ---- Config ------------------------------------------------------------------\nclass CFG:\n    DATA_ROOT   = \"/kaggle/input/vesuvius-challenge/train\"\n    WORK_DIR    = \"/kaggle/working\"\n    ZARR_DIR    = os.path.join(WORK_DIR, \"zarr_store\")\n    CKPT_DIR    = os.path.join(WORK_DIR, \"checkpoints\")\n    VIZ_DIR     = os.path.join(WORK_DIR, \"viz\")\n\n    TRAIN_FRAGMENTS = [2, 3]   # 80/20 spatial split inside each -> train/val\n    TEST_FRAGMENT    = 1       # fully held out\n\n    Z_SLICES = 65               # 00.tif .. 64.tif\n    Z_MID    = Z_SLICES // 2    # center of the stack\n\n    # depth window used per sample (Temporal Random Crop draws a window from this range)\n    DEPTH_MIN, DEPTH_MAX = 12, 22\n\n    TILE          = 224          # spatial patch size (H=W=TILE)\n    TRAIN_STRIDE  = 112          # 50% overlap while enumerating candidate train/val patches\n    TEST_STRIDE   = 112          # overlap for sliding-window test inference (averaged)\n\n    FG_MEAN_THRESH = 8.0 / 255.0   # skip near-empty (background/air) patches when indexing\n\n    VAL_FRACTION_BY_WIDTH = 0.8    # first 80% of fragment width -> train, last 20% -> val\n\n    BATCH_SIZE   = 8\n    ACCUM_STEPS  = 2               # effective batch = BATCH_SIZE * ACCUM_STEPS\n    NUM_WORKERS  = 2\n    EPOCHS       = 15\n    LR           = 3e-4\n    WEIGHT_DECAY = 1e-4\n\n    # patches sampled per epoch (subsample huge candidate lists -> bounds RAM & epoch time)\n    MAX_TRAIN_PATCHES_PER_EPOCH = 2500\n    MAX_VAL_PATCHES             = 600\n\n    # Ensemble: each entry is one trained model variant (architecture, depth window, seed)\n    ENSEMBLE_CONFIGS = [\n        dict(name=\"se3d_unet_d16_s0\",  arch=\"se3d_unet\",     depth=16, base_ch=24, seed=0),\n        dict(name=\"se3d_unet_d20_s1\",  arch=\"se3d_unet\",     depth=20, base_ch=24, seed=1),\n        dict(name=\"maxpool_seg_d16_s2\", arch=\"maxpool_seg\",  depth=16, base_ch=32, seed=2),\n        dict(name=\"se3d_unet_d12_s3\",  arch=\"se3d_unet\",     depth=12, base_ch=32, seed=3),\n        dict(name=\"maxpool_seg_d22_s4\", arch=\"maxpool_seg\",  depth=22, base_ch=32, seed=4),\n    ]\n\n    F_BETA = 0.5  # F_0.5 -> precision-weighted\n\n    # Tversky index generalizes Dice: TI = TP / (TP + alpha*FN + beta*FP).\n    # beta > alpha penalizes false positives harder than false negatives, which is the\n    # correct pairing for an F_0.5 (precision-weighted) target -- this is what was missing\n    # when the ensemble converged to \"flag almost everything as ink\" (P=0.20, R=0.999).\n    TVERSKY_ALPHA = 0.3\n    TVERSKY_BETA  = 0.7\n\nos.makedirs(CFG.ZARR_DIR, exist_ok=True)\nos.makedirs(CFG.CKPT_DIR, exist_ok=True)\nos.makedirs(CFG.VIZ_DIR, exist_ok=True)\n\n\n# %% [CELL 2] ---- TIFF stack -> chunked Zarr (slice-by-slice, no full-volume RAM load) ----\ndef _existing_zarr_is_valid(frag_id, cfg=CFG):\n    \"\"\"Checks that a previously-built Zarr store for this fragment is present AND complete\n    (right depth, and label mask spatially matches the volume) before trusting it -- a\n    directory merely existing isn't proof the conversion finished cleanly last time.\"\"\"\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n    if not (os.path.exists(vol_path) and os.path.exists(lbl_path)):\n        return False, None, None\n    try:\n        vol = zarr.open(vol_path, mode=\"r\")[\"data\"]\n        lbl = zarr.open(lbl_path, mode=\"r\")[\"data\"]\n        if vol.shape[0] != cfg.Z_SLICES:\n            print(f\"[frag {frag_id}] cached zarr has depth {vol.shape[0]} != expected \"\n                  f\"{cfg.Z_SLICES}; will rebuild.\")\n            return False, None, None\n        if lbl.shape != vol.shape[1:]:\n            print(f\"[frag {frag_id}] cached label shape {lbl.shape} != volume spatial \"\n                  f\"shape {vol.shape[1:]}; will rebuild.\")\n            return False, None, None\n        return True, vol_path, lbl_path\n    except Exception as e:\n        print(f\"[frag {frag_id}] cached zarr at {vol_path} looks corrupt ({e}); will rebuild.\")\n        return False, None, None\n\n\ndef convert_fragment_to_zarr(frag_id: int, cfg=CFG):\n    \"\"\"\n    Writes:\n      zarr_store/frag{frag_id}_volume.zarr  -> uint8 array, shape (Z, H, W), chunks (Z, 256, 256)\n      zarr_store/frag{frag_id}_labels.zarr  -> uint8 array, shape (H, W),   chunks (256, 256)\n    Reuses a previously-built, verified-complete store instead of re-reading the source TIFFs.\n    \"\"\"\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n\n    is_valid, cached_vol, cached_lbl = _existing_zarr_is_valid(frag_id, cfg)\n    if is_valid:\n        vol = zarr.open(cached_vol, mode=\"r\")[\"data\"]\n        lbl = zarr.open(cached_lbl, mode=\"r\")[\"data\"]\n        print(f\"[frag {frag_id}] reusing existing zarr at {cfg.ZARR_DIR} \"\n              f\"(volume {vol.shape}, labels {lbl.shape}) -- skipping TIFF re-read.\")\n        return cached_vol, cached_lbl\n\n    frag_dir = os.path.join(cfg.DATA_ROOT, str(frag_id))\n    slice_paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in range(cfg.Z_SLICES)]\n    label_path  = os.path.join(frag_dir, \"inklabels.png\")\n\n    # peek shape from first slice\n    with tifffile.TiffFile(slice_paths[0]) as tf:\n        h, w = tf.pages[0].shape\n    print(f\"[frag {frag_id}] volume shape -> ({cfg.Z_SLICES}, {h}, {w})\")\n\n    store = zarr.DirectoryStore(vol_path)\n    root = zarr.group(store=store, overwrite=True)\n    vol_z = root.create_dataset(\n        \"data\", shape=(cfg.Z_SLICES, h, w), chunks=(cfg.Z_SLICES, 256, 256),\n        dtype=\"u1\", compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n\n    for i, sp in enumerate(slice_paths):\n        sl = tifffile.imread(sp)  # uint16, shape (H, W) -- ONE slice in RAM at a time\n        # normalize 16-bit -> 8-bit for compact on-disk storage / cheap I/O during training\n        sl8 = (sl.astype(np.float32) / 65535.0 * 255.0).clip(0, 255).astype(np.uint8)\n        vol_z[i, :, :] = sl8\n        del sl, sl8\n        if i % 16 == 0:\n            gc.collect()\n    print(f\"[frag {frag_id}] volume -> zarr done.\")\n\n    lbl_img = Image.open(label_path).convert(\"L\")\n    lbl = (np.array(lbl_img) > 127).astype(np.uint8)\n    assert lbl.shape == (h, w), f\"label shape {lbl.shape} != volume shape {(h, w)}\"\n\n    lstore = zarr.DirectoryStore(lbl_path)\n    lroot = zarr.group(store=lstore, overwrite=True)\n    lbl_z = lroot.create_dataset(\n        \"data\", shape=(h, w), chunks=(256, 256), dtype=\"u1\",\n        compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n    lbl_z[:, :] = lbl\n    del lbl, lbl_img\n    gc.collect()\n    print(f\"[frag {frag_id}] labels -> zarr done.\")\n    return vol_path, lbl_path\n\n\ndef open_fragment_zarr(frag_id, cfg=CFG):\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n    vol = zarr.open(vol_path, mode=\"r\")[\"data\"]\n    lbl = zarr.open(lbl_path, mode=\"r\")[\"data\"]\n    return vol, lbl\n\n\n# %% [CELL 3] ---- Patch indexing + 80/20 spatial split ------------------------------------\ndef build_patch_grid(h, w, tile, stride):\n    ys = list(range(0, max(h - tile, 0) + 1, stride))\n    xs = list(range(0, max(w - tile, 0) + 1, stride))\n    if ys[-1] != h - tile: ys.append(max(h - tile, 0))\n    if xs[-1] != w - tile: xs.append(max(w - tile, 0))\n    return [(y, x) for y in ys for x in xs]\n\n\ndef index_fragment_patches(frag_id, cfg=CFG, split=\"both\"):\n    \"\"\"\n    Returns list of dicts: {frag_id, y, x} for foreground patches only (cheap intensity\n    filter on the middle slice, read directly from zarr chunk-by-chunk -> low RAM).\n    split: 'train' -> first VAL_FRACTION_BY_WIDTH of width\n           'val'   -> remaining width\n           'both'  -> no split (used for the held-out test fragment)\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    mid_slice = np.asarray(vol[cfg.Z_MID, :, :])  # (H, W) uint8, one slice in RAM\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TRAIN_STRIDE if split != \"both\" else cfg.TEST_STRIDE)\n\n    split_x = int(W * cfg.VAL_FRACTION_BY_WIDTH)\n    kept = []\n    for (y, x) in coords:\n        if split == \"train\" and not (x + cfg.TILE <= split_x):\n            continue\n        if split == \"val\" and not (x >= split_x):\n            continue\n        patch = mid_slice[y:y + cfg.TILE, x:x + cfg.TILE]\n        if patch.size == 0:\n            continue\n        if (patch.astype(np.float32) / 255.0).mean() < cfg.FG_MEAN_THRESH:\n            continue  # skip empty background/air patch\n        kept.append({\"frag_id\": frag_id, \"y\": y, \"x\": x})\n    del mid_slice\n    gc.collect()\n    print(f\"[frag {frag_id}] split={split}: {len(kept)} foreground patches indexed \"\n          f\"(H={H}, W={W}).\")\n    return kept\n\n\n# %% [CELL 4] ---- Dataset with Temporal Random Crop / Random Paste / Temporal Cutout ------\nclass VesuviusPatchDataset(Dataset):\n    \"\"\"\n    Each item pulls a (depth_window, TILE, TILE) sub-volume straight out of the on-disk\n    Zarr store (no full-fragment load), applies the three temporal augmentations plus\n    standard spatial flips/rotations, and returns (volume_tensor, mask_tensor).\n    \"\"\"\n    def __init__(self, patch_records, depth_window, cfg=CFG, train=True, max_items=None):\n        self.records = patch_records\n        self.depth_window = depth_window\n        self.cfg = cfg\n        self.train = train\n        self._vol_cache = {}\n        self._lbl_cache = {}\n        if max_items is not None and len(self.records) > max_items:\n            self.records = random.sample(self.records, max_items)\n\n    def _get_arrays(self, frag_id):\n        if frag_id not in self._vol_cache:\n            self._vol_cache[frag_id], self._lbl_cache[frag_id] = open_fragment_zarr(frag_id, self.cfg)\n        return self._vol_cache[frag_id], self._lbl_cache[frag_id]\n\n    def __len__(self):\n        return len(self.records)\n\n    def _sample_depth_window(self, Z):\n        dw = self.depth_window if not self.train else random.randint(CFG.DEPTH_MIN, CFG.DEPTH_MAX)\n        dw = min(dw, Z)\n        z_mid = Z // 2\n        half = dw // 2\n        # Temporal Random Crop: random window around the middle of the stack\n        jitter = random.randint(-4, 4) if self.train else 0\n        z0 = max(0, min(Z - dw, z_mid - half + jitter))\n        return z0, dw\n\n    def __getitem__(self, idx):\n        rec = self.records[idx]\n        vol, lbl = self._get_arrays(rec[\"frag_id\"])\n        Z, H, W = vol.shape\n        y, x, t = rec[\"y\"], rec[\"x\"], self.cfg.TILE\n\n        z0, dw = self._sample_depth_window(Z)\n        sub = np.asarray(vol[z0:z0 + dw, y:y + t, x:x + t]).astype(np.float32) / 255.0  # (dw,t,t)\n        mask = np.asarray(lbl[y:y + t, x:x + t]).astype(np.float32)  # (t,t)\n\n        if self.train:\n            sub, mask = self._augment(sub, mask)\n\n        # pad depth to DEPTH_MAX so batches stack cleanly; padding slices are zeroed\n        # (network ignores all-zero slices thanks to Temporal Cutout training on real zeros too)\n        pad = CFG.DEPTH_MAX - sub.shape[0]\n        if pad > 0:\n            sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n\n        vol_t = torch.from_numpy(sub).unsqueeze(0).float()   # (1, D, H, W)\n        mask_t = torch.from_numpy(mask).unsqueeze(0).float() # (1, H, W)\n        return vol_t, mask_t\n\n    def _augment(self, sub, mask):\n        D = sub.shape[0]\n\n        # --- Random Paste: shift cropped slice-range to a random Z offset within the padded\n        #     budget, keeping slice order intact (sequential order preserved). ---\n        if random.random() < 0.5 and D < CFG.DEPTH_MAX:\n            max_shift = CFG.DEPTH_MAX - D\n            shift = random.randint(0, max_shift)\n            padded = np.zeros((CFG.DEPTH_MAX, *sub.shape[1:]), dtype=sub.dtype)\n            padded[shift:shift + D] = sub\n            sub = padded\n            D = CFG.DEPTH_MAX\n\n        # --- Temporal Cutout: zero out 1-2 random layers inside the active cube ---\n        n_cutout = random.choice([0, 1, 1, 2])\n        for _ in range(n_cutout):\n            zi = random.randint(0, D - 1)\n            sub[zi] = 0.0\n\n        # --- standard spatial augs ---\n        if random.random() < 0.5:\n            sub = sub[:, :, ::-1].copy(); mask = mask[:, ::-1].copy()\n        if random.random() < 0.5:\n            sub = sub[:, ::-1, :].copy(); mask = mask[::-1, :].copy()\n        k = random.choice([0, 1, 2, 3])\n        if k:\n            sub = np.rot90(sub, k, axes=(1, 2)).copy()\n            mask = np.rot90(mask, k, axes=(0, 1)).copy()\n\n        return sub, mask\n\n\n# %% [CELL 5] ---- Model: SE blocks, 3D encoder -> 2D SE-UNet decoder, MaxPool+Seg decoder --\nclass SEBlock2D(nn.Module):\n    \"\"\"Squeeze-and-Excitation applied inside skip connections.\"\"\"\n    def __init__(self, ch, reduction=8):\n        super().__init__()\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(ch, max(ch // reduction, 4), 1), nn.ReLU(inplace=True),\n            nn.Conv2d(max(ch // reduction, 4), ch, 1), nn.Sigmoid(),\n        )\n    def forward(self, x):\n        return x * self.fc(self.pool(x))\n\n\nclass ConvBNAct2D(nn.Module):\n    def __init__(self, cin, cout, k=3, s=1, p=1):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(cin, cout, k, s, p, bias=False),\n            nn.BatchNorm2d(cout), nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.block(x)\n\n\nclass Encoder3D(nn.Module):\n    \"\"\"3D conv encoder that progressively collapses the depth axis while extracting\n    multi-scale 2D feature maps (one per stage) for U-Net-style skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n        self.stage1 = nn.Sequential(\n            nn.Conv3d(1, c1, 3, padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n            nn.Conv3d(c1, c1, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n        )\n        self.stage2 = nn.Sequential(\n            nn.Conv3d(c1, c2, 3, padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n            nn.Conv3d(c2, c2, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n        )\n        self.stage3 = nn.Sequential(\n            nn.Conv3d(c2, c3, 3, padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n            nn.Conv3d(c3, c3, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n        )\n        self.stage4 = nn.Sequential(\n            # stride only on H,W here (depth stride=1): this must land the bottleneck one\n            # scale BELOW f3's spatial resolution, since the decoder's up3 upsamples it by\n            # 2x before concatenating with f3. Depth is fully collapsed right after by the\n            # adaptive pool, so no depth stride is needed at this stage.\n            nn.Conv3d(c3, c4, 3, stride=(1, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c4), nn.ReLU(inplace=True),\n            nn.AdaptiveMaxPool3d((1, None, None)),  # collapse remaining depth -> bottleneck\n        )\n        self.out_channels = (c1, c2, c3, c4)\n\n    @staticmethod\n    def _depth_collapse(x):\n        # max over depth to produce a 2D skip feature map at this scale\n        return x.max(dim=2).values\n\n    def forward(self, x):  # x: (B, 1, D, H, W)\n        s1 = self.stage1(x); f1 = self._depth_collapse(s1)          # H/2\n        s2 = self.stage2(s1); f2 = self._depth_collapse(s2)          # H/4\n        s3 = self.stage3(s2); f3 = self._depth_collapse(s3)          # H/8\n        s4 = self.stage4(s3).squeeze(2)                               # H/8, depth->1\n        return f1, f2, f3, s4\n\n\nclass SEUnetDecoder2D(nn.Module):\n    def __init__(self, enc_channels):\n        super().__init__()\n        c1, c2, c3, c4 = enc_channels\n        self.up3 = nn.ConvTranspose2d(c4, c3, 2, 2)\n        self.se3 = SEBlock2D(c3); self.dec3 = ConvBNAct2D(c3 * 2, c3)\n        self.up2 = nn.ConvTranspose2d(c3, c2, 2, 2)\n        self.se2 = SEBlock2D(c2); self.dec2 = ConvBNAct2D(c2 * 2, c2)\n        self.up1 = nn.ConvTranspose2d(c2, c1, 2, 2)\n        self.se1 = SEBlock2D(c1); self.dec1 = ConvBNAct2D(c1 * 2, c1)\n        self.up0 = nn.ConvTranspose2d(c1, c1 // 2, 2, 2)\n        self.final = nn.Conv2d(c1 // 2, 1, 1)\n\n    def forward(self, f1, f2, f3, bottleneck, out_hw):\n        x = self.up3(bottleneck); x = self.dec3(torch.cat([x, self.se3(f3)], dim=1))\n        x = self.up2(x);          x = self.dec2(torch.cat([x, self.se2(f2)], dim=1))\n        x = self.up1(x);          x = self.dec1(torch.cat([x, self.se1(f1)], dim=1))\n        x = self.up0(x)\n        x = F.interpolate(x, size=out_hw, mode=\"bilinear\", align_corners=False)\n        return self.final(x)\n\n\nclass SE3DUNet(nn.Module):\n    \"\"\"Top-solution-style architecture: 3D encoder + 2D decoder U-Net with SE blocks in\n    the skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        self.encoder = Encoder3D(base_ch)\n        self.decoder = SEUnetDecoder2D(self.encoder.out_channels)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        h, w = x.shape[-2], x.shape[-1]\n        f1, f2, f3, bott = self.encoder(x)\n        return self.decoder(f1, f2, f3, bott, (h, w))\n\n\nclass MaxPoolFlattenSegHead(nn.Module):\n    \"\"\"Depth Flattening branch: collapse the (D,H,W) volume to a 2D feature map via\n    max pooling across depth, then feed a heavy semantic-segmentation decoder\n    (SegFormer/MiT-style if `segmentation_models_pytorch` is available, otherwise a\n    hand-rolled multi-scale conv decoder with SE-augmented skips).\"\"\"\n    def __init__(self, base_ch=32, use_smp=HAS_SMP):\n        super().__init__()\n        self.use_smp = False\n        if use_smp:\n            try:\n                self.net = smp.Unet(\n                    encoder_name=\"mit_b3\", encoder_weights=\"imagenet\",\n                    in_channels=1, classes=1, activation=None,\n                )\n                self.use_smp = True\n            except Exception as e:\n                print(f\"[MaxPoolFlattenSegHead] smp mit_b3 unavailable at model-build time \"\n                      f\"({e}); using the built-in fallback decoder instead.\")\n\n        if not self.use_smp:\n            c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n            self.enc1 = nn.Sequential(ConvBNAct2D(1, c1), ConvBNAct2D(c1, c1))\n            self.pool1 = nn.MaxPool2d(2)\n            self.enc2 = nn.Sequential(ConvBNAct2D(c1, c2), ConvBNAct2D(c2, c2))\n            self.pool2 = nn.MaxPool2d(2)\n            self.enc3 = nn.Sequential(ConvBNAct2D(c2, c3), ConvBNAct2D(c3, c3))\n            self.pool3 = nn.MaxPool2d(2)\n            self.bott = nn.Sequential(ConvBNAct2D(c3, c4), ConvBNAct2D(c4, c4))\n            self.decoder = SEUnetDecoder2D((c1, c2, c3, c4))\n\n    def forward(self, x2d):  # x2d: (B, 1, H, W)  (already depth-flattened)\n        if self.use_smp:\n            return self.net(x2d)\n        h, w = x2d.shape[-2], x2d.shape[-1]\n        e1 = self.enc1(x2d); p1 = self.pool1(e1)\n        e2 = self.enc2(p1);  p2 = self.pool2(e2)\n        e3 = self.enc3(p2);  p3 = self.pool3(e3)\n        b  = self.bott(p3)\n        return self.decoder(e1, e2, e3, b, (h, w))\n\n\nclass MaxPoolSegModel(nn.Module):\n    \"\"\"Full pipeline for the 'Depth Flattening' approach: Max Pool across the depth axis\n    of the 3D volume -> 2D feature map -> heavy segmentation decoder -> ink mask.\"\"\"\n    def __init__(self, base_ch=32):\n        super().__init__()\n        self.head = MaxPoolFlattenSegHead(base_ch=base_ch)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        x2d = x.max(dim=2).values  # Max Pooling across depth axis -> (B,1,H,W)\n        return self.head(x2d)\n\n\ndef build_model(arch, base_ch):\n    if arch == \"se3d_unet\":\n        return SE3DUNet(base_ch=base_ch)\n    elif arch == \"maxpool_seg\":\n        return MaxPoolSegModel(base_ch=base_ch)\n    else:\n        raise ValueError(arch)\n\n\n# %% [CELL 6] ---- Loss (BCE + Dice) and F-beta metric/threshold search --------------------\nclass BCETverskyLoss(nn.Module):\n    \"\"\"BCE + Tversky. Unlike plain Dice (which weights precision/recall equally), Tversky\n    with beta > alpha explicitly penalizes false positives harder -- the correct pairing\n    when the downstream metric (F_0.5) itself weights precision over recall.\"\"\"\n    def __init__(self, bce_w=0.5, tversky_w=0.5, alpha=CFG.TVERSKY_ALPHA, beta=CFG.TVERSKY_BETA, smooth=1.0):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.bce_w, self.tversky_w = bce_w, tversky_w\n        self.alpha, self.beta, self.smooth = alpha, beta, smooth\n\n    def forward(self, logits, target):\n        bce_loss = self.bce(logits, target)\n        prob = torch.sigmoid(logits)\n        p, t = prob.flatten(1), target.flatten(1)\n        tp = (p * t).sum(1)\n        fp = (p * (1 - t)).sum(1)\n        fn = ((1 - p) * t).sum(1)\n        tversky = (tp + self.smooth) / (tp + self.alpha * fn + self.beta * fp + self.smooth)\n        tversky_loss = 1 - tversky\n        return self.bce_w * bce_loss + self.tversky_w * tversky_loss.mean()\n\n\n# kept as an alias so any external references to the old name still work\nBCEDiceLoss = BCETverskyLoss\n\n\ndef fbeta_score(precision, recall, beta=CFG.F_BETA, eps=1e-8):\n    b2 = beta ** 2\n    return (1 + b2) * precision * recall / (b2 * precision + recall + eps)\n\n\ndef find_best_threshold(probs_flat: np.ndarray, targets_flat: np.ndarray, beta=CFG.F_BETA):\n    \"\"\"Optimizes the decision threshold strictly for F_beta (beta<1 -> precision-weighted).\"\"\"\n    precision, recall, thresholds = precision_recall_curve(targets_flat, probs_flat)\n    precision, recall = precision[:-1], recall[:-1]\n    scores = fbeta_score(precision, recall, beta)\n    best_idx = int(np.nanargmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), float(precision[best_idx]), float(recall[best_idx])\n\n\n# %% [CELL 7] ---- Training loop for one ensemble member ------------------------------------\ndef run_validation(model, val_loader):\n    model.eval()\n    all_probs, all_targets = [], []\n    with torch.no_grad(), autocast():\n        for vol, mask in val_loader:\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            logits = model(vol)\n            probs = torch.sigmoid(logits).float().cpu().numpy().ravel()\n            targs = mask.float().cpu().numpy().ravel()\n            # subsample pixels per patch to keep the PR-curve computation light on RAM\n            if len(probs) > 20000:\n                idx = np.random.choice(len(probs), 20000, replace=False)\n                probs, targs = probs[idx], targs[idx]\n            all_probs.append(probs); all_targets.append(targs)\n    probs = np.concatenate(all_probs); targets = np.concatenate(all_targets)\n    thr, f05, prec, rec = find_best_threshold(probs, targets)\n    del all_probs, all_targets\n    gc.collect()\n    return thr, f05, prec, rec\n\n\ndef train_one_model(cfg_entry, train_records, val_records, cfg=CFG, epochs=None):\n    epochs = epochs or cfg.EPOCHS\n    torch.manual_seed(cfg_entry[\"seed\"]); random.seed(cfg_entry[\"seed\"])\n\n    model = build_model(cfg_entry[\"arch\"], cfg_entry[\"base_ch\"]).to(DEVICE)\n    opt = torch.optim.AdamW(model.parameters(), lr=cfg.LR, weight_decay=cfg.WEIGHT_DECAY)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    scaler = GradScaler()\n    criterion = BCETverskyLoss()\n\n    train_ds = VesuviusPatchDataset(train_records, cfg_entry[\"depth\"], cfg, train=True,\n                                     max_items=cfg.MAX_TRAIN_PATCHES_PER_EPOCH)\n    val_ds = VesuviusPatchDataset(val_records, cfg_entry[\"depth\"], cfg, train=False,\n                                   max_items=cfg.MAX_VAL_PATCHES)\n\n    best_f05, best_thr, best_state = -1.0, 0.5, None\n    history = []\n\n    for epoch in range(epochs):\n        # resample the train subset each epoch for coverage without ever loading everything\n        train_ds.records = random.sample(train_records, min(cfg.MAX_TRAIN_PATCHES_PER_EPOCH, len(train_records)))\n        train_loader = DataLoader(train_ds, batch_size=cfg.BATCH_SIZE, shuffle=True,\n                                   num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True)\n        val_loader = DataLoader(val_ds, batch_size=cfg.BATCH_SIZE, shuffle=False,\n                                 num_workers=cfg.NUM_WORKERS, pin_memory=True)\n\n        model.train()\n        running_loss, n_steps = 0.0, 0\n        opt.zero_grad()\n        t0 = time.time()\n        for step, (vol, mask) in enumerate(train_loader):\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            with autocast():\n                logits = model(vol)\n                loss = criterion(logits, mask) / cfg.ACCUM_STEPS\n            scaler.scale(loss).backward()\n            if (step + 1) % cfg.ACCUM_STEPS == 0:\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n            running_loss += loss.item() * cfg.ACCUM_STEPS\n            n_steps += 1\n            del vol, mask, logits, loss\n        sched.step()\n\n        thr, f05, prec, rec = run_validation(model, val_loader)\n        dt = time.time() - t0\n        print(f\"[{cfg_entry['name']}] epoch {epoch+1}/{epochs} \"\n              f\"train_loss={running_loss/max(n_steps,1):.4f} \"\n              f\"val_F0.5={f05:.4f} P={prec:.3f} R={rec:.3f} thr={thr:.3f} ({dt:.0f}s)\")\n        history.append(dict(epoch=epoch+1, train_loss=running_loss/max(n_steps,1),\n                             val_f05=f05, val_precision=prec, val_recall=rec, threshold=thr))\n\n        if f05 > best_f05:\n            best_f05, best_thr = f05, thr\n            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n\n        del train_loader, val_loader\n        gc.collect(); torch.cuda.empty_cache()\n\n    model.load_state_dict(best_state)\n    ckpt_path = os.path.join(cfg.CKPT_DIR, f\"{cfg_entry['name']}.pt\")\n    torch.save({\"state_dict\": best_state, \"config\": cfg_entry,\n                \"best_val_f05\": best_f05, \"best_threshold\": best_thr,\n                \"history\": history}, ckpt_path)\n    print(f\"[{cfg_entry['name']}] saved best checkpoint (val F0.5={best_f05:.4f}) -> {ckpt_path}\")\n\n    del model, opt, sched, scaler, train_ds, val_ds\n    gc.collect(); torch.cuda.empty_cache()\n    return ckpt_path, best_thr, best_f05, history\n\n\n# %% [CELL 8] ---- Sliding-window ensemble inference over a full fragment ------------------\ndef predict_fragment_ensemble(frag_id, model_infos, cfg=CFG):\n    \"\"\"\n    model_infos: list of dicts {ckpt_path, config} for each ensemble member.\n    Returns (prob_map (H,W) float32 averaged across models, gt_mask (H,W) uint8).\n    Memory-safe: only a float32 (H,W) accumulator + weight map live in RAM (no full volume).\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    gt = np.asarray(lbl[:, :])\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TEST_STRIDE)\n\n    loaded_models = []\n    for info in model_infos:\n        ckpt = torch.load(info[\"ckpt_path\"], map_location=DEVICE)\n        m = build_model(ckpt[\"config\"][\"arch\"], ckpt[\"config\"][\"base_ch\"]).to(DEVICE)\n        m.load_state_dict(ckpt[\"state_dict\"]); m.eval()\n        loaded_models.append((m, ckpt[\"config\"][\"depth\"]))\n\n    prob_acc = np.zeros((H, W), dtype=np.float32)\n    weight_acc = np.zeros((H, W), dtype=np.float32)\n\n    with torch.no_grad(), autocast():\n        for (y, x) in coords:\n            patch_sum = np.zeros((cfg.TILE, cfg.TILE), dtype=np.float32)\n            for model, depth in loaded_models:\n                z_mid = Z // 2\n                z0 = max(0, min(Z - depth, z_mid - depth // 2))\n                sub = np.asarray(vol[z0:z0 + depth, y:y + cfg.TILE, x:x + cfg.TILE]).astype(np.float32) / 255.0\n                pad = cfg.DEPTH_MAX - sub.shape[0]\n                if pad > 0:\n                    sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n                t = torch.from_numpy(sub).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n                logits = model(t)\n                patch_sum += torch.sigmoid(logits)[0, 0].float().cpu().numpy()\n                del t, logits\n            patch_prob = patch_sum / len(loaded_models)\n            prob_acc[y:y + cfg.TILE, x:x + cfg.TILE] += patch_prob\n            weight_acc[y:y + cfg.TILE, x:x + cfg.TILE] += 1.0\n\n    weight_acc[weight_acc == 0] = 1.0\n    prob_map = prob_acc / weight_acc\n\n    for m, _ in loaded_models:\n        del m\n    del loaded_models\n    gc.collect(); torch.cuda.empty_cache()\n    return prob_map, gt\n\n\n# %% [CELL 9] ---- Orchestration: convert data, split, train ensemble, evaluate, visualize --\ndef main():\n    # 1) Convert all needed fragments to Zarr (train 2,3 + held-out test 1)\n    for fid in CFG.TRAIN_FRAGMENTS + [CFG.TEST_FRAGMENT]:\n        convert_fragment_to_zarr(fid)\n\n    # 2) Build 80/20 spatial split from fragments 2 & 3\n    train_records, val_records = [], []\n    for fid in CFG.TRAIN_FRAGMENTS:\n        train_records += index_fragment_patches(fid, split=\"train\")\n        val_records   += index_fragment_patches(fid, split=\"val\")\n    print(f\"TOTAL train patches: {len(train_records)} | val patches: {len(val_records)}\")\n\n    # 3) Train each ensemble member\n    trained = []\n    for cfg_entry in CFG.ENSEMBLE_CONFIGS:\n        ckpt_path, thr, f05, hist = train_one_model(cfg_entry, train_records, val_records)\n        trained.append({\"ckpt_path\": ckpt_path, \"config\": cfg_entry, \"val_threshold\": thr, \"val_f05\": f05})\n\n    # 4) Aggregate validation threshold (mean across members) -> apply to test\n    ensemble_val_thr = float(np.mean([t[\"val_threshold\"] for t in trained]))\n    print(f\"Ensemble-averaged validation threshold: {ensemble_val_thr:.3f}\")\n\n    # 5) Full-fragment ensemble inference on the held-out TEST fragment (fragment 1)\n    prob_map, gt_map = predict_fragment_ensemble(CFG.TEST_FRAGMENT, trained)\n    pred_mask = (prob_map >= ensemble_val_thr).astype(np.uint8)\n\n    # 6) Final test metrics (precision/recall/F0.5/Dice) at the chosen threshold\n    p_flat, t_flat = pred_mask.ravel().astype(np.float32), gt_map.ravel().astype(np.float32)\n    tp = float((p_flat * t_flat).sum())\n    fp = float((p_flat * (1 - t_flat)).sum())\n    fn = float(((1 - p_flat) * t_flat).sum())\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n    f05_test = fbeta_score(precision, recall, CFG.F_BETA)\n    dice_test = 2 * tp / (2 * tp + fp + fn + 1e-8)\n    print(f\"TEST (fragment {CFG.TEST_FRAGMENT}) @thr={ensemble_val_thr:.3f} -> \"\n          f\"Precision={precision:.4f} Recall={recall:.4f} F0.5={f05_test:.4f} Dice={dice_test:.4f}\")\n\n    metrics = dict(\n        ensemble_val_threshold=ensemble_val_thr,\n        test_precision=precision, test_recall=recall,\n        test_f0_5=f05_test, test_dice=dice_test,\n        members=[{\"name\": t[\"config\"][\"name\"], \"val_f05\": t[\"val_f05\"],\n                  \"val_threshold\": t[\"val_threshold\"]} for t in trained],\n    )\n    with open(os.path.join(CFG.WORK_DIR, \"metrics_summary.json\"), \"w\") as f:\n        json.dump(metrics, f, indent=2)\n    print(\"Saved metrics_summary.json\")\n\n    # 7) Visualization: input (mid-slice + max-projection) vs ground truth vs prediction\n    vol, _ = open_fragment_zarr(CFG.TEST_FRAGMENT)\n    mid_slice = np.asarray(vol[CFG.Z_MID, :, :])\n    maxproj = np.asarray(vol[:, :, :]).max(axis=0) if vol.shape[0] * vol.shape[1] * vol.shape[2] < 4e8 \\\n        else mid_slice  # guard: skip full maxproj on huge volumes to avoid RAM spikes\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(mid_slice, cmap=\"gray\"); axes[0].set_title(f\"Input (mid slice z={CFG.Z_MID})\")\n    axes[1].imshow(gt_map, cmap=\"gray\"); axes[1].set_title(\"Ground Truth Ink Mask\")\n    axes[2].imshow(prob_map, cmap=\"magma\"); axes[2].set_title(\"Predicted Probability Map\")\n    axes[3].imshow(mid_slice, cmap=\"gray\")\n    axes[3].imshow(np.ma.masked_where(pred_mask == 0, pred_mask), cmap=\"autumn\", alpha=0.6)\n    axes[3].set_title(f\"Prediction Overlay (thr={ensemble_val_thr:.2f})\")\n    for ax in axes: ax.axis(\"off\")\n    plt.tight_layout()\n    viz_path = os.path.join(CFG.VIZ_DIR, f\"test_fragment{CFG.TEST_FRAGMENT}_comparison.png\")\n    plt.savefig(viz_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved visualization -> {viz_path}\")\n\n    del mid_slice, maxproj, prob_map, gt_map, pred_mask\n    gc.collect(); torch.cuda.empty_cache()\n    return metrics\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-03T09:03:47.484095Z","iopub.execute_input":"2026-08-03T09:03:47.48456Z"}},"outputs":[{"name":"stderr","text":"  error: subprocess-exited-with-error\n  \n  × pip subprocess to install backend dependencies did not run successfully.\n  │ exit code: 1\n  ╰─> [3 lines of output]\n      ERROR: Ignored the following versions that require a different python version: 0.1.0 Requires-Python >=3.9; 0.1.1 Requires-Python >=3.9; 0.1.10 Requires-Python >=3.9; 0.1.11 Requires-Python >=3.9; 0.1.12 Requires-Python >=3.9; 0.1.13 Requires-Python >=3.9; 0.1.14 Requires-Python >=3.9; 0.1.2 Requires-Python >=3.9; 0.1.3 Requires-Python >=3.9; 0.1.4 Requires-Python >=3.9; 0.1.5 Requires-Python >=3.9; 0.1.6 Requires-Python >=3.9; 0.1.7 Requires-Python >=3.9; 0.1.8 Requires-Python >=3.9; 0.1.9 Requires-Python >=3.9\n      ERROR: Could not find a version that satisfies the requirement puccinialin (from versions: none)\n      ERROR: No matching distribution found for puccinialin\n      [end of output]\n  \n  note: This error originates from a subprocess, and is likely not a problem with pip.\nerror: subprocess-exited-with-error\n\n× pip subprocess to install backend dependencies did not run successfully.\n│ exit code: 1\n╰─> See above for output.\n\nnote: This error originates from a subprocess, and is likely not a problem with pip.\n","output_type":"stream"},{"name":"stdout","text":"[setup] segmentation_models_pytorch MiT (SegFormer) backbone available: False\nDevice: cuda\n[frag 2] reusing existing zarr at /kaggle/working/zarr_store (volume (65, 14830, 9506), labels (14830, 9506)) -- skipping TIFF re-read.\n[frag 3] reusing existing zarr at /kaggle/working/zarr_store (volume (65, 7606, 5249), labels (7606, 5249)) -- skipping TIFF re-read.\n[frag 1] reusing existing zarr at /kaggle/working/zarr_store (volume (65, 8181, 6330), labels (8181, 6330)) -- skipping TIFF re-read.\n[frag 2] split=train: 6957 foreground patches indexed (H=14830, W=9506).\n[frag 2] split=val: 1033 foreground patches indexed (H=14830, W=9506).\n[frag 3] split=train: 1745 foreground patches indexed (H=7606, W=5249).\n[frag 3] split=val: 294 foreground patches indexed (H=7606, W=5249).\nTOTAL train patches: 8702 | val patches: 1327\n[se3d_unet_d16_s0] epoch 1/15 train_loss=0.6957 val_F0.5=0.1937 P=0.164 R=0.681 thr=0.226 (119s)\n[se3d_unet_d16_s0] epoch 2/15 train_loss=0.6596 val_F0.5=0.2031 P=0.170 R=0.914 thr=0.075 (119s)\n[se3d_unet_d16_s0] epoch 3/15 train_loss=0.6583 val_F0.5=0.2032 P=0.170 R=0.956 thr=0.095 (119s)\n[se3d_unet_d16_s0] epoch 4/15 train_loss=0.6567 val_F0.5=0.2030 P=0.170 R=0.928 thr=0.075 (119s)\n[se3d_unet_d16_s0] epoch 5/15 train_loss=0.6520 val_F0.5=0.2051 P=0.171 R=0.965 thr=0.093 (119s)\n[se3d_unet_d16_s0] epoch 6/15 train_loss=0.6530 val_F0.5=0.2016 P=0.169 R=0.911 thr=0.201 (119s)\n[se3d_unet_d16_s0] epoch 7/15 train_loss=0.6501 val_F0.5=0.2041 P=0.171 R=0.966 thr=0.101 (118s)\n[se3d_unet_d16_s0] epoch 8/15 train_loss=0.6478 val_F0.5=0.2033 P=0.170 R=0.942 thr=0.100 (119s)\n[se3d_unet_d16_s0] epoch 9/15 train_loss=0.6473 val_F0.5=0.2046 P=0.171 R=0.934 thr=0.204 (119s)\n[se3d_unet_d16_s0] epoch 10/15 train_loss=0.6447 val_F0.5=0.2055 P=0.172 R=0.938 thr=0.064 (119s)\n[se3d_unet_d16_s0] epoch 11/15 train_loss=0.6441 val_F0.5=0.2059 P=0.172 R=0.952 thr=0.072 (118s)\n[se3d_unet_d16_s0] epoch 12/15 train_loss=0.6405 val_F0.5=0.2062 P=0.173 R=0.944 thr=0.074 (119s)\n[se3d_unet_d16_s0] epoch 13/15 train_loss=0.6414 val_F0.5=0.2069 P=0.173 R=0.932 thr=0.088 (119s)\n[se3d_unet_d16_s0] epoch 14/15 train_loss=0.6410 val_F0.5=0.2047 P=0.171 R=0.952 thr=0.092 (119s)\n[se3d_unet_d16_s0] epoch 15/15 train_loss=0.6389 val_F0.5=0.2050 P=0.171 R=0.957 thr=0.073 (118s)\n[se3d_unet_d16_s0] saved best checkpoint (val F0.5=0.2069) -> /kaggle/working/checkpoints/se3d_unet_d16_s0.pt\n[se3d_unet_d20_s1] epoch 1/15 train_loss=0.7111 val_F0.5=0.2023 P=0.176 R=0.502 thr=0.268 (119s)\n[se3d_unet_d20_s1] epoch 2/15 train_loss=0.6595 val_F0.5=0.2127 P=0.178 R=0.959 thr=0.171 (119s)\n[se3d_unet_d20_s1] epoch 3/15 train_loss=0.6509 val_F0.5=0.2056 P=0.172 R=0.930 thr=0.057 (119s)\n[se3d_unet_d20_s1] epoch 4/15 train_loss=0.6545 val_F0.5=0.2030 P=0.171 R=0.828 thr=0.107 (119s)\n[se3d_unet_d20_s1] epoch 5/15 train_loss=0.6497 val_F0.5=0.2111 P=0.177 R=0.910 thr=0.120 (119s)\n[se3d_unet_d20_s1] epoch 6/15 train_loss=0.6450 val_F0.5=0.2101 P=0.176 R=0.904 thr=0.046 (119s)\n[se3d_unet_d20_s1] epoch 7/15 train_loss=0.6475 val_F0.5=0.2146 P=0.180 R=0.947 thr=0.110 (119s)\n[se3d_unet_d20_s1] epoch 8/15 train_loss=0.6461 val_F0.5=0.2155 P=0.181 R=0.936 thr=0.077 (118s)\n[se3d_unet_d20_s1] epoch 9/15 train_loss=0.6404 val_F0.5=0.2143 P=0.180 R=0.884 thr=0.048 (118s)\n[se3d_unet_d20_s1] epoch 10/15 train_loss=0.6383 val_F0.5=0.2146 P=0.180 R=0.945 thr=0.081 (119s)\n[se3d_unet_d20_s1] epoch 11/15 train_loss=0.6369 val_F0.5=0.2153 P=0.181 R=0.934 thr=0.072 (118s)\n[se3d_unet_d20_s1] epoch 12/15 train_loss=0.6370 val_F0.5=0.2142 P=0.180 R=0.933 thr=0.070 (118s)\n[se3d_unet_d20_s1] epoch 13/15 train_loss=0.6364 val_F0.5=0.2222 P=0.208 R=0.303 thr=0.162 (118s)\n[se3d_unet_d20_s1] epoch 14/15 train_loss=0.6325 val_F0.5=0.2328 P=0.239 R=0.212 thr=0.255 (118s)\n[se3d_unet_d20_s1] epoch 15/15 train_loss=0.6346 val_F0.5=0.2238 P=0.219 R=0.243 thr=0.260 (118s)\n[se3d_unet_d20_s1] saved best checkpoint (val F0.5=0.2328) -> /kaggle/working/checkpoints/se3d_unet_d20_s1.pt\n[maxpool_seg_d16_s2] epoch 1/15 train_loss=0.6748 val_F0.5=0.2276 P=0.192 R=0.897 thr=0.118 (69s)\n[maxpool_seg_d16_s2] epoch 2/15 train_loss=0.6515 val_F0.5=0.2282 P=0.193 R=0.878 thr=0.114 (70s)\n[maxpool_seg_d16_s2] epoch 3/15 train_loss=0.6518 val_F0.5=0.2270 P=0.191 R=0.915 thr=0.123 (70s)\n[maxpool_seg_d16_s2] epoch 4/15 train_loss=0.6508 val_F0.5=0.2310 P=0.195 R=0.897 thr=0.099 (70s)\n[maxpool_seg_d16_s2] epoch 5/15 train_loss=0.6495 val_F0.5=0.2082 P=0.174 R=0.996 thr=0.074 (71s)\n[maxpool_seg_d16_s2] epoch 6/15 train_loss=0.6485 val_F0.5=0.2467 P=0.211 R=0.752 thr=0.162 (70s)\n[maxpool_seg_d16_s2] epoch 7/15 train_loss=0.6428 val_F0.5=0.2412 P=0.206 R=0.761 thr=0.129 (69s)\n[maxpool_seg_d16_s2] epoch 8/15 train_loss=0.6417 val_F0.5=0.2545 P=0.223 R=0.589 thr=0.156 (70s)\n[maxpool_seg_d16_s2] epoch 9/15 train_loss=0.6386 val_F0.5=0.2536 P=0.220 R=0.665 thr=0.128 (71s)\n[maxpool_seg_d16_s2] epoch 10/15 train_loss=0.6357 val_F0.5=0.2588 P=0.233 R=0.460 thr=0.165 (70s)\n[maxpool_seg_d16_s2] epoch 11/15 train_loss=0.6341 val_F0.5=0.2650 P=0.235 R=0.540 thr=0.156 (70s)\n[maxpool_seg_d16_s2] epoch 12/15 train_loss=0.6324 val_F0.5=0.2760 P=0.246 R=0.533 thr=0.215 (70s)\n[maxpool_seg_d16_s2] epoch 13/15 train_loss=0.6324 val_F0.5=0.2730 P=0.249 R=0.447 thr=0.184 (70s)\n[maxpool_seg_d16_s2] epoch 14/15 train_loss=0.6316 val_F0.5=0.2809 P=0.254 R=0.483 thr=0.216 (70s)\n[maxpool_seg_d16_s2] epoch 15/15 train_loss=0.6262 val_F0.5=0.2780 P=0.252 R=0.475 thr=0.200 (70s)\n[maxpool_seg_d16_s2] saved best checkpoint (val F0.5=0.2809) -> /kaggle/working/checkpoints/maxpool_seg_d16_s2.pt\n[se3d_unet_d12_s3] epoch 1/15 train_loss=0.6998 val_F0.5=0.2122 P=0.188 R=0.429 thr=0.260 (144s)\n[se3d_unet_d12_s3] epoch 2/15 train_loss=0.6563 val_F0.5=0.2239 P=0.189 R=0.818 thr=0.100 (144s)\n[se3d_unet_d12_s3] epoch 3/15 train_loss=0.6540 val_F0.5=0.2281 P=0.194 R=0.743 thr=0.139 (144s)\n[se3d_unet_d12_s3] epoch 4/15 train_loss=0.6512 val_F0.5=0.2276 P=0.192 R=0.916 thr=0.103 (143s)\n[se3d_unet_d12_s3] epoch 5/15 train_loss=0.6521 val_F0.5=0.2278 P=0.192 R=0.927 thr=0.115 (143s)\n[se3d_unet_d12_s3] epoch 6/15 train_loss=0.6465 val_F0.5=0.2234 P=0.188 R=0.902 thr=0.095 (143s)\n[se3d_unet_d12_s3] epoch 7/15 train_loss=0.6451 val_F0.5=0.2242 P=0.188 R=0.934 thr=0.076 (143s)\n[se3d_unet_d12_s3] epoch 8/15 train_loss=0.6435 val_F0.5=0.2272 P=0.191 R=0.933 thr=0.092 (143s)\n[se3d_unet_d12_s3] epoch 9/15 train_loss=0.6434 val_F0.5=0.2262 P=0.190 R=0.922 thr=0.093 (143s)\n[se3d_unet_d12_s3] epoch 10/15 train_loss=0.6407 val_F0.5=0.2304 P=0.199 R=0.612 thr=0.118 (143s)\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n==========================================================================================\nVESUVIUS CHALLENGE - INK DETECTION\nMemory-safe patch-based pipeline (Zarr storage, 3D->2D SE-UNet + MaxPool/SegFormer branch,\ntemporal augmentations, BCE+Dice, F0.5-optimized thresholding, small ensemble, viz).\n==========================================================================================\n\nHOW TO USE ON KAGGLE\n---------------------\n1. Paste each \"# %% [CELL n] ...\" block into its own notebook cell (recommended), OR just\n   run this whole file as a single script cell.\n2. Turn on a GPU accelerator (P100 / T4x2).\n3. Data is expected at:\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/surface_volume/00.tif ... 64.tif\n     /kaggle/input/vesuvius-challenge/train/{1,2,3}/inklabels.png\n4. Everything derived is written to /kaggle/working (Zarr stacks, checkpoints, plots, json).\n\nWHY THIS AVOIDS OOM\n--------------------\n- Raw fragments are converted ONCE into on-disk chunked Zarr arrays (uint8), slice-by-slice,\n  so we never hold a full (65, H, W) volume in RAM.\n- Training/inference never touch the full fragment either: we sample small 3D patches\n  (depth_window x tile x tile) directly out of the Zarr store on demand.\n- Mixed precision (AMP), small batch size + optional grad accumulation, aggressive\n  `del` + `gc.collect()` + `torch.cuda.empty_cache()`, and bounded DataLoader workers.\n- Test-fragment inference is also patch-based with an overlap-averaged stitching buffer\n  (float32 (H,W) accumulator, which for Vesuvius fragment sizes is a few hundred MB at most\n  - much smaller than the raw (65,H,W) uint16 volume would be).\n\"\"\"\n#!pip install segmentation-models-pytorch==0.2.0\n# %% [CELL 0] ---- Setup & installs -------------------------------------------------------\n!pip install segmentation-models-pytorch==0.2.0\n\n!pip install zarr==2.12.0\ntry:\n    import segmentation_models_pytorch as smp  # noqa\n    HAS_SMP = True\nexcept ImportError:\n    HAS_SMP = False  # we fall back to a hand-rolled SegFormer-ish decoder below\n\nimport os, gc, json, math, random, time, glob\nimport numpy as np\nimport zarr\nimport tifffile\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import precision_recall_curve\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\n\n# %% [CELL 1] ---- Config ------------------------------------------------------------------\nclass CFG:\n    DATA_ROOT   = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    WORK_DIR    = \"/kaggle/working\"\n    ZARR_DIR    = os.path.join(WORK_DIR, \"zarr_store\")\n    CKPT_DIR    = os.path.join(WORK_DIR, \"checkpoints\")\n    VIZ_DIR     = os.path.join(WORK_DIR, \"viz\")\n\n    TRAIN_FRAGMENTS = [2, 3]   # 80/20 spatial split inside each -> train/val\n    TEST_FRAGMENT    = 1       # fully held out\n\n    Z_SLICES = 65               # 00.tif .. 64.tif\n    Z_MID    = Z_SLICES // 2    # center of the stack\n\n    # depth window used per sample (Temporal Random Crop draws a window from this range)\n    DEPTH_MIN, DEPTH_MAX = 12, 22\n\n    TILE          = 224          # spatial patch size (H=W=TILE)\n    TRAIN_STRIDE  = 112          # 50% overlap while enumerating candidate train/val patches\n    TEST_STRIDE   = 112          # overlap for sliding-window test inference (averaged)\n\n    FG_MEAN_THRESH = 8.0 / 255.0   # skip near-empty (background/air) patches when indexing\n\n    VAL_FRACTION_BY_WIDTH = 0.8    # first 80% of fragment width -> train, last 20% -> val\n\n    BATCH_SIZE   = 8\n    ACCUM_STEPS  = 2               # effective batch = BATCH_SIZE * ACCUM_STEPS\n    NUM_WORKERS  = 2\n    EPOCHS       = 2\n    LR           = 3e-4\n    WEIGHT_DECAY = 1e-4\n\n    # patches sampled per epoch (subsample huge candidate lists -> bounds RAM & epoch time)\n    MAX_TRAIN_PATCHES_PER_EPOCH = 1200\n    MAX_VAL_PATCHES             = 600\n\n    # Ensemble: each entry is one trained model variant (architecture, depth window, seed)\n    ENSEMBLE_CONFIGS = [\n        dict(name=\"se3d_unet_d16_s0\",  arch=\"se3d_unet\",     depth=16, base_ch=24, seed=0),\n        dict(name=\"se3d_unet_d20_s1\",  arch=\"se3d_unet\",     depth=20, base_ch=24, seed=1),\n        dict(name=\"maxpool_seg_d16_s2\", arch=\"maxpool_seg\",  depth=16, base_ch=32, seed=2),\n        dict(name=\"se3d_unet_d12_s3\",  arch=\"se3d_unet\",     depth=12, base_ch=32, seed=3),\n        dict(name=\"maxpool_seg_d22_s4\", arch=\"maxpool_seg\",  depth=22, base_ch=32, seed=4),\n    ]\n\n    F_BETA = 0.5  # F_0.5 -> precision-weighted\n\nos.makedirs(CFG.ZARR_DIR, exist_ok=True)\nos.makedirs(CFG.CKPT_DIR, exist_ok=True)\nos.makedirs(CFG.VIZ_DIR, exist_ok=True)\n\n\n# %% [CELL 2] ---- TIFF stack -> chunked Zarr (slice-by-slice, no full-volume RAM load) ----\ndef convert_fragment_to_zarr(frag_id: int, cfg=CFG):\n    \"\"\"\n    Writes:\n      zarr_store/frag{frag_id}_volume.zarr  -> uint8 array, shape (Z, H, W), chunks (Z, 256, 256)\n      zarr_store/frag{frag_id}_labels.zarr  -> uint8 array, shape (H, W),   chunks (256, 256)\n    Skips conversion if already present.\n    \"\"\"\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n\n    frag_dir = os.path.join(cfg.DATA_ROOT, str(frag_id))\n    slice_paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\") for i in range(cfg.Z_SLICES)]\n    label_path  = os.path.join(frag_dir, \"inklabels.png\")\n\n    if os.path.exists(vol_path) and os.path.exists(lbl_path):\n        print(f\"[frag {frag_id}] zarr already exists, skipping conversion.\")\n        return vol_path, lbl_path\n\n    # peek shape from first slice\n    with tifffile.TiffFile(slice_paths[0]) as tf:\n        h, w = tf.pages[0].shape\n    print(f\"[frag {frag_id}] volume shape -> ({cfg.Z_SLICES}, {h}, {w})\")\n\n    store = zarr.DirectoryStore(vol_path)\n    root = zarr.group(store=store, overwrite=True)\n    vol_z = root.create_dataset(\n        \"data\", shape=(cfg.Z_SLICES, h, w), chunks=(cfg.Z_SLICES, 256, 256),\n        dtype=\"u1\", compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n\n    for i, sp in enumerate(slice_paths):\n        sl = tifffile.imread(sp)  # uint16, shape (H, W) -- ONE slice in RAM at a time\n        # normalize 16-bit -> 8-bit for compact on-disk storage / cheap I/O during training\n        sl8 = (sl.astype(np.float32) / 65535.0 * 255.0).clip(0, 255).astype(np.uint8)\n        vol_z[i, :, :] = sl8\n        del sl, sl8\n        if i % 16 == 0:\n            gc.collect()\n    print(f\"[frag {frag_id}] volume -> zarr done.\")\n\n    lbl_img = Image.open(label_path).convert(\"L\")\n    lbl = (np.array(lbl_img) > 127).astype(np.uint8)\n    assert lbl.shape == (h, w), f\"label shape {lbl.shape} != volume shape {(h, w)}\"\n\n    lstore = zarr.DirectoryStore(lbl_path)\n    lroot = zarr.group(store=lstore, overwrite=True)\n    lbl_z = lroot.create_dataset(\n        \"data\", shape=(h, w), chunks=(256, 256), dtype=\"u1\",\n        compressor=zarr.Blosc(cname=\"zstd\", clevel=3, shuffle=2),\n    )\n    lbl_z[:, :] = lbl\n    del lbl, lbl_img\n    gc.collect()\n    print(f\"[frag {frag_id}] labels -> zarr done.\")\n    return vol_path, lbl_path\n\n\ndef open_fragment_zarr(frag_id, cfg=CFG):\n    vol_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_volume.zarr\")\n    lbl_path = os.path.join(cfg.ZARR_DIR, f\"frag{frag_id}_labels.zarr\")\n    vol = zarr.open(vol_path, mode=\"r\")[\"data\"]\n    lbl = zarr.open(lbl_path, mode=\"r\")[\"data\"]\n    return vol, lbl\n\n\n# %% [CELL 3] ---- Patch indexing + 80/20 spatial split ------------------------------------\ndef build_patch_grid(h, w, tile, stride):\n    ys = list(range(0, max(h - tile, 0) + 1, stride))\n    xs = list(range(0, max(w - tile, 0) + 1, stride))\n    if ys[-1] != h - tile: ys.append(max(h - tile, 0))\n    if xs[-1] != w - tile: xs.append(max(w - tile, 0))\n    return [(y, x) for y in ys for x in xs]\n\n\ndef index_fragment_patches(frag_id, cfg=CFG, split=\"both\"):\n    \"\"\"\n    Returns list of dicts: {frag_id, y, x} for foreground patches only (cheap intensity\n    filter on the middle slice, read directly from zarr chunk-by-chunk -> low RAM).\n    split: 'train' -> first VAL_FRACTION_BY_WIDTH of width\n           'val'   -> remaining width\n           'both'  -> no split (used for the held-out test fragment)\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    mid_slice = np.asarray(vol[cfg.Z_MID, :, :])  # (H, W) uint8, one slice in RAM\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TRAIN_STRIDE if split != \"both\" else cfg.TEST_STRIDE)\n\n    split_x = int(W * cfg.VAL_FRACTION_BY_WIDTH)\n    kept = []\n    for (y, x) in coords:\n        if split == \"train\" and not (x + cfg.TILE <= split_x):\n            continue\n        if split == \"val\" and not (x >= split_x):\n            continue\n        patch = mid_slice[y:y + cfg.TILE, x:x + cfg.TILE]\n        if patch.size == 0:\n            continue\n        if (patch.astype(np.float32) / 255.0).mean() < cfg.FG_MEAN_THRESH:\n            continue  # skip empty background/air patch\n        kept.append({\"frag_id\": frag_id, \"y\": y, \"x\": x})\n    del mid_slice\n    gc.collect()\n    print(f\"[frag {frag_id}] split={split}: {len(kept)} foreground patches indexed \"\n          f\"(H={H}, W={W}).\")\n    return kept\n\n\n# %% [CELL 4] ---- Dataset with Temporal Random Crop / Random Paste / Temporal Cutout ------\nclass VesuviusPatchDataset(Dataset):\n    \"\"\"\n    Each item pulls a (depth_window, TILE, TILE) sub-volume straight out of the on-disk\n    Zarr store (no full-fragment load), applies the three temporal augmentations plus\n    standard spatial flips/rotations, and returns (volume_tensor, mask_tensor).\n    \"\"\"\n    def __init__(self, patch_records, depth_window, cfg=CFG, train=True, max_items=None):\n        self.records = patch_records\n        self.depth_window = depth_window\n        self.cfg = cfg\n        self.train = train\n        self._vol_cache = {}\n        self._lbl_cache = {}\n        if max_items is not None and len(self.records) > max_items:\n            self.records = random.sample(self.records, max_items)\n\n    def _get_arrays(self, frag_id):\n        if frag_id not in self._vol_cache:\n            self._vol_cache[frag_id], self._lbl_cache[frag_id] = open_fragment_zarr(frag_id, self.cfg)\n        return self._vol_cache[frag_id], self._lbl_cache[frag_id]\n\n    def __len__(self):\n        return len(self.records)\n\n    def _sample_depth_window(self, Z):\n        dw = self.depth_window if not self.train else random.randint(CFG.DEPTH_MIN, CFG.DEPTH_MAX)\n        dw = min(dw, Z)\n        z_mid = Z // 2\n        half = dw // 2\n        # Temporal Random Crop: random window around the middle of the stack\n        jitter = random.randint(-4, 4) if self.train else 0\n        z0 = max(0, min(Z - dw, z_mid - half + jitter))\n        return z0, dw\n\n    def __getitem__(self, idx):\n        rec = self.records[idx]\n        vol, lbl = self._get_arrays(rec[\"frag_id\"])\n        Z, H, W = vol.shape\n        y, x, t = rec[\"y\"], rec[\"x\"], self.cfg.TILE\n\n        z0, dw = self._sample_depth_window(Z)\n        sub = np.asarray(vol[z0:z0 + dw, y:y + t, x:x + t]).astype(np.float32) / 255.0  # (dw,t,t)\n        mask = np.asarray(lbl[y:y + t, x:x + t]).astype(np.float32)  # (t,t)\n\n        if self.train:\n            sub, mask = self._augment(sub, mask)\n\n        # pad depth to DEPTH_MAX so batches stack cleanly; padding slices are zeroed\n        # (network ignores all-zero slices thanks to Temporal Cutout training on real zeros too)\n        pad = CFG.DEPTH_MAX - sub.shape[0]\n        if pad > 0:\n            sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n\n        vol_t = torch.from_numpy(sub).unsqueeze(0).float()   # (1, D, H, W)\n        mask_t = torch.from_numpy(mask).unsqueeze(0).float() # (1, H, W)\n        return vol_t, mask_t\n\n    def _augment(self, sub, mask):\n        D = sub.shape[0]\n\n        # --- Random Paste: shift cropped slice-range to a random Z offset within the padded\n        #     budget, keeping slice order intact (sequential order preserved). ---\n        if random.random() < 0.5 and D < CFG.DEPTH_MAX:\n            max_shift = CFG.DEPTH_MAX - D\n            shift = random.randint(0, max_shift)\n            padded = np.zeros((CFG.DEPTH_MAX, *sub.shape[1:]), dtype=sub.dtype)\n            padded[shift:shift + D] = sub\n            sub = padded\n            D = CFG.DEPTH_MAX\n\n        # --- Temporal Cutout: zero out 1-2 random layers inside the active cube ---\n        n_cutout = random.choice([0, 1, 1, 2])\n        for _ in range(n_cutout):\n            zi = random.randint(0, D - 1)\n            sub[zi] = 0.0\n\n        # --- standard spatial augs ---\n        if random.random() < 0.5:\n            sub = sub[:, :, ::-1].copy(); mask = mask[:, ::-1].copy()\n        if random.random() < 0.5:\n            sub = sub[:, ::-1, :].copy(); mask = mask[::-1, :].copy()\n        k = random.choice([0, 1, 2, 3])\n        if k:\n            sub = np.rot90(sub, k, axes=(1, 2)).copy()\n            mask = np.rot90(mask, k, axes=(0, 1)).copy()\n\n        return sub, mask\n\n\n# %% [CELL 5] ---- Model: SE blocks, 3D encoder -> 2D SE-UNet decoder, MaxPool+Seg decoder --\nclass SEBlock2D(nn.Module):\n    \"\"\"Squeeze-and-Excitation applied inside skip connections.\"\"\"\n    def __init__(self, ch, reduction=8):\n        super().__init__()\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(ch, max(ch // reduction, 4), 1), nn.ReLU(inplace=True),\n            nn.Conv2d(max(ch // reduction, 4), ch, 1), nn.Sigmoid(),\n        )\n    def forward(self, x):\n        return x * self.fc(self.pool(x))\n\n\nclass ConvBNAct2D(nn.Module):\n    def __init__(self, cin, cout, k=3, s=1, p=1):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(cin, cout, k, s, p, bias=False),\n            nn.BatchNorm2d(cout), nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.block(x)\n\n\nclass Encoder3D(nn.Module):\n    \"\"\"3D conv encoder that progressively collapses the depth axis while extracting\n    multi-scale 2D feature maps (one per stage) for U-Net-style skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n        self.stage1 = nn.Sequential(\n            nn.Conv3d(1, c1, 3, padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n            nn.Conv3d(c1, c1, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c1), nn.ReLU(inplace=True),\n        )\n        self.stage2 = nn.Sequential(\n            nn.Conv3d(c1, c2, 3, padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n            nn.Conv3d(c2, c2, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c2), nn.ReLU(inplace=True),\n        )\n        self.stage3 = nn.Sequential(\n            nn.Conv3d(c2, c3, 3, padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n            nn.Conv3d(c3, c3, 3, stride=(2, 2, 2), padding=1, bias=False), nn.BatchNorm3d(c3), nn.ReLU(inplace=True),\n        )\n        self.stage4 = nn.Sequential(\n            nn.Conv3d(c3, c4, 3, padding=1, bias=False), nn.BatchNorm3d(c4), nn.ReLU(inplace=True),\n            nn.AdaptiveMaxPool3d((1, None, None)),  # collapse remaining depth -> bottleneck\n        )\n        self.out_channels = (c1, c2, c3, c4)\n\n    @staticmethod\n    def _depth_collapse(x):\n        # max over depth to produce a 2D skip feature map at this scale\n        return x.max(dim=2).values\n\n    def forward(self, x):  # x: (B, 1, D, H, W)\n        s1 = self.stage1(x); f1 = self._depth_collapse(s1)          # H/2\n        s2 = self.stage2(s1); f2 = self._depth_collapse(s2)          # H/4\n        s3 = self.stage3(s2); f3 = self._depth_collapse(s3)          # H/8\n        s4 = self.stage4(s3).squeeze(2)                               # H/8, depth->1\n        return f1, f2, f3, s4\n\n\nclass SEUnetDecoder2D(nn.Module):\n    def __init__(self, enc_channels):\n        super().__init__()\n        c1, c2, c3, c4 = enc_channels\n        self.up3 = nn.ConvTranspose2d(c4, c3, 2, 2)\n        self.se3 = SEBlock2D(c3); self.dec3 = ConvBNAct2D(c3 * 2, c3)\n        self.up2 = nn.ConvTranspose2d(c3, c2, 2, 2)\n        self.se2 = SEBlock2D(c2); self.dec2 = ConvBNAct2D(c2 * 2, c2)\n        self.up1 = nn.ConvTranspose2d(c2, c1, 2, 2)\n        self.se1 = SEBlock2D(c1); self.dec1 = ConvBNAct2D(c1 * 2, c1)\n        self.up0 = nn.ConvTranspose2d(c1, c1 // 2, 2, 2)\n        self.final = nn.Conv2d(c1 // 2, 1, 1)\n\n    def forward(self, f1, f2, f3, bottleneck, out_hw):\n        x = self.up3(bottleneck); x = self.dec3(torch.cat([x, self.se3(f3)], dim=1))\n        x = self.up2(x);          x = self.dec2(torch.cat([x, self.se2(f2)], dim=1))\n        x = self.up1(x);          x = self.dec1(torch.cat([x, self.se1(f1)], dim=1))\n        x = self.up0(x)\n        x = F.interpolate(x, size=out_hw, mode=\"bilinear\", align_corners=False)\n        return self.final(x)\n\n\nclass SE3DUNet(nn.Module):\n    \"\"\"Top-solution-style architecture: 3D encoder + 2D decoder U-Net with SE blocks in\n    the skip connections.\"\"\"\n    def __init__(self, base_ch=24):\n        super().__init__()\n        self.encoder = Encoder3D(base_ch)\n        self.decoder = SEUnetDecoder2D(self.encoder.out_channels)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        h, w = x.shape[-2], x.shape[-1]\n        f1, f2, f3, bott = self.encoder(x)\n        return self.decoder(f1, f2, f3, bott, (h, w))\n\n\nclass MaxPoolFlattenSegHead(nn.Module):\n    \"\"\"Depth Flattening branch: collapse the (D,H,W) volume to a 2D feature map via\n    max pooling across depth, then feed a heavy semantic-segmentation decoder\n    (SegFormer/MiT-style if `segmentation_models_pytorch` is available, otherwise a\n    hand-rolled multi-scale conv decoder with SE-augmented skips).\"\"\"\n    def __init__(self, base_ch=32, use_smp=HAS_SMP):\n        super().__init__()\n        self.use_smp = use_smp\n        if use_smp:\n            self.net = smp.Unet(\n                encoder_name=\"mit_b3\", encoder_weights=\"imagenet\",\n                in_channels=1, classes=1, activation=None,\n            )\n        else:\n            c1, c2, c3, c4 = base_ch, base_ch * 2, base_ch * 4, base_ch * 8\n            self.enc1 = nn.Sequential(ConvBNAct2D(1, c1), ConvBNAct2D(c1, c1))\n            self.pool1 = nn.MaxPool2d(2)\n            self.enc2 = nn.Sequential(ConvBNAct2D(c1, c2), ConvBNAct2D(c2, c2))\n            self.pool2 = nn.MaxPool2d(2)\n            self.enc3 = nn.Sequential(ConvBNAct2D(c2, c3), ConvBNAct2D(c3, c3))\n            self.pool3 = nn.MaxPool2d(2)\n            self.bott = nn.Sequential(ConvBNAct2D(c3, c4), ConvBNAct2D(c4, c4))\n            self.decoder = SEUnetDecoder2D((c1, c2, c3, c4))\n\n    def forward(self, x2d):  # x2d: (B, 1, H, W)  (already depth-flattened)\n        if self.use_smp:\n            return self.net(x2d)\n        h, w = x2d.shape[-2], x2d.shape[-1]\n        e1 = self.enc1(x2d); p1 = self.pool1(e1)\n        e2 = self.enc2(p1);  p2 = self.pool2(e2)\n        e3 = self.enc3(p2);  p3 = self.pool3(e3)\n        b  = self.bott(p3)\n        return self.decoder(e1, e2, e3, b, (h, w))\n\n\nclass MaxPoolSegModel(nn.Module):\n    \"\"\"Full pipeline for the 'Depth Flattening' approach: Max Pool across the depth axis\n    of the 3D volume -> 2D feature map -> heavy segmentation decoder -> ink mask.\"\"\"\n    def __init__(self, base_ch=32):\n        super().__init__()\n        self.head = MaxPoolFlattenSegHead(base_ch=base_ch)\n\n    def forward(self, x):  # x: (B,1,D,H,W)\n        x2d = x.max(dim=2).values  # Max Pooling across depth axis -> (B,1,H,W)\n        return self.head(x2d)\n\n\ndef build_model(arch, base_ch):\n    if arch == \"se3d_unet\":\n        return SE3DUNet(base_ch=base_ch)\n    elif arch == \"maxpool_seg\":\n        return MaxPoolSegModel(base_ch=base_ch)\n    else:\n        raise ValueError(arch)\n\n\n# %% [CELL 6] ---- Loss (BCE + Dice) and F-beta metric/threshold search --------------------\nclass BCEDiceLoss(nn.Module):\n    def __init__(self, bce_w=0.5, dice_w=0.5, smooth=1.0):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.bce_w, self.dice_w, self.smooth = bce_w, dice_w, smooth\n\n    def forward(self, logits, target):\n        bce_loss = self.bce(logits, target)\n        prob = torch.sigmoid(logits)\n        p, t = prob.flatten(1), target.flatten(1)\n        inter = (p * t).sum(1)\n        dice = 1 - (2 * inter + self.smooth) / (p.sum(1) + t.sum(1) + self.smooth)\n        return self.bce_w * bce_loss + self.dice_w * dice.mean()\n\n\ndef fbeta_score(precision, recall, beta=CFG.F_BETA, eps=1e-8):\n    b2 = beta ** 2\n    return (1 + b2) * precision * recall / (b2 * precision + recall + eps)\n\n\ndef find_best_threshold(probs_flat: np.ndarray, targets_flat: np.ndarray, beta=CFG.F_BETA):\n    \"\"\"Optimizes the decision threshold strictly for F_beta (beta<1 -> precision-weighted).\"\"\"\n    precision, recall, thresholds = precision_recall_curve(targets_flat, probs_flat)\n    precision, recall = precision[:-1], recall[:-1]\n    scores = fbeta_score(precision, recall, beta)\n    best_idx = int(np.nanargmax(scores))\n    return float(thresholds[best_idx]), float(scores[best_idx]), float(precision[best_idx]), float(recall[best_idx])\n\n\n# %% [CELL 7] ---- Training loop for one ensemble member ------------------------------------\ndef run_validation(model, val_loader):\n    model.eval()\n    all_probs, all_targets = [], []\n    with torch.no_grad(), autocast():\n        for vol, mask in val_loader:\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            logits = model(vol)\n            probs = torch.sigmoid(logits).float().cpu().numpy().ravel()\n            targs = mask.float().cpu().numpy().ravel()\n            # subsample pixels per patch to keep the PR-curve computation light on RAM\n            if len(probs) > 20000:\n                idx = np.random.choice(len(probs), 20000, replace=False)\n                probs, targs = probs[idx], targs[idx]\n            all_probs.append(probs); all_targets.append(targs)\n    probs = np.concatenate(all_probs); targets = np.concatenate(all_targets)\n    thr, f05, prec, rec = find_best_threshold(probs, targets)\n    del all_probs, all_targets\n    gc.collect()\n    return thr, f05, prec, rec\n\n\ndef train_one_model(cfg_entry, train_records, val_records, cfg=CFG, epochs=None):\n    epochs = epochs or cfg.EPOCHS\n    torch.manual_seed(cfg_entry[\"seed\"]); random.seed(cfg_entry[\"seed\"])\n\n    model = build_model(cfg_entry[\"arch\"], cfg_entry[\"base_ch\"]).to(DEVICE)\n    opt = torch.optim.AdamW(model.parameters(), lr=cfg.LR, weight_decay=cfg.WEIGHT_DECAY)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    scaler = GradScaler()\n    criterion = BCEDiceLoss()\n\n    train_ds = VesuviusPatchDataset(train_records, cfg_entry[\"depth\"], cfg, train=True,\n                                     max_items=cfg.MAX_TRAIN_PATCHES_PER_EPOCH)\n    val_ds = VesuviusPatchDataset(val_records, cfg_entry[\"depth\"], cfg, train=False,\n                                   max_items=cfg.MAX_VAL_PATCHES)\n\n    best_f05, best_thr, best_state = -1.0, 0.5, None\n    history = []\n\n    for epoch in range(epochs):\n        # resample the train subset each epoch for coverage without ever loading everything\n        train_ds.records = random.sample(train_records, min(cfg.MAX_TRAIN_PATCHES_PER_EPOCH, len(train_records)))\n        train_loader = DataLoader(train_ds, batch_size=cfg.BATCH_SIZE, shuffle=True,\n                                   num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True)\n        val_loader = DataLoader(val_ds, batch_size=cfg.BATCH_SIZE, shuffle=False,\n                                 num_workers=cfg.NUM_WORKERS, pin_memory=True)\n\n        model.train()\n        running_loss, n_steps = 0.0, 0\n        opt.zero_grad()\n        t0 = time.time()\n        for step, (vol, mask) in enumerate(train_loader):\n            vol, mask = vol.to(DEVICE, non_blocking=True), mask.to(DEVICE, non_blocking=True)\n            with autocast():\n                logits = model(vol)\n                loss = criterion(logits, mask) / cfg.ACCUM_STEPS\n            scaler.scale(loss).backward()\n            if (step + 1) % cfg.ACCUM_STEPS == 0:\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n            running_loss += loss.item() * cfg.ACCUM_STEPS\n            n_steps += 1\n            del vol, mask, logits, loss\n        sched.step()\n\n        thr, f05, prec, rec = run_validation(model, val_loader)\n        dt = time.time() - t0\n        print(f\"[{cfg_entry['name']}] epoch {epoch+1}/{epochs} \"\n              f\"train_loss={running_loss/max(n_steps,1):.4f} \"\n              f\"val_F0.5={f05:.4f} P={prec:.3f} R={rec:.3f} thr={thr:.3f} ({dt:.0f}s)\")\n        history.append(dict(epoch=epoch+1, train_loss=running_loss/max(n_steps,1),\n                             val_f05=f05, val_precision=prec, val_recall=rec, threshold=thr))\n\n        if f05 > best_f05:\n            best_f05, best_thr = f05, thr\n            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n\n        del train_loader, val_loader\n        gc.collect(); torch.cuda.empty_cache()\n\n    model.load_state_dict(best_state)\n    ckpt_path = os.path.join(cfg.CKPT_DIR, f\"{cfg_entry['name']}.pt\")\n    torch.save({\"state_dict\": best_state, \"config\": cfg_entry,\n                \"best_val_f05\": best_f05, \"best_threshold\": best_thr,\n                \"history\": history}, ckpt_path)\n    print(f\"[{cfg_entry['name']}] saved best checkpoint (val F0.5={best_f05:.4f}) -> {ckpt_path}\")\n\n    del model, opt, sched, scaler, train_ds, val_ds\n    gc.collect(); torch.cuda.empty_cache()\n    return ckpt_path, best_thr, best_f05, history\n\n\n# %% [CELL 8] ---- Sliding-window ensemble inference over a full fragment ------------------\ndef predict_fragment_ensemble(frag_id, model_infos, cfg=CFG):\n    \"\"\"\n    model_infos: list of dicts {ckpt_path, config} for each ensemble member.\n    Returns (prob_map (H,W) float32 averaged across models, gt_mask (H,W) uint8).\n    Memory-safe: only a float32 (H,W) accumulator + weight map live in RAM (no full volume).\n    \"\"\"\n    vol, lbl = open_fragment_zarr(frag_id, cfg)\n    Z, H, W = vol.shape\n    gt = np.asarray(lbl[:, :])\n\n    coords = build_patch_grid(H, W, cfg.TILE, cfg.TEST_STRIDE)\n\n    loaded_models = []\n    for info in model_infos:\n        ckpt = torch.load(info[\"ckpt_path\"], map_location=DEVICE)\n        m = build_model(ckpt[\"config\"][\"arch\"], ckpt[\"config\"][\"base_ch\"]).to(DEVICE)\n        m.load_state_dict(ckpt[\"state_dict\"]); m.eval()\n        loaded_models.append((m, ckpt[\"config\"][\"depth\"]))\n\n    prob_acc = np.zeros((H, W), dtype=np.float32)\n    weight_acc = np.zeros((H, W), dtype=np.float32)\n\n    with torch.no_grad(), autocast():\n        for (y, x) in coords:\n            patch_sum = np.zeros((cfg.TILE, cfg.TILE), dtype=np.float32)\n            for model, depth in loaded_models:\n                z_mid = Z // 2\n                z0 = max(0, min(Z - depth, z_mid - depth // 2))\n                sub = np.asarray(vol[z0:z0 + depth, y:y + cfg.TILE, x:x + cfg.TILE]).astype(np.float32) / 255.0\n                pad = cfg.DEPTH_MAX - sub.shape[0]\n                if pad > 0:\n                    sub = np.pad(sub, ((0, pad), (0, 0), (0, 0)), mode=\"constant\")\n                t = torch.from_numpy(sub).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n                logits = model(t)\n                patch_sum += torch.sigmoid(logits)[0, 0].float().cpu().numpy()\n                del t, logits\n            patch_prob = patch_sum / len(loaded_models)\n            prob_acc[y:y + cfg.TILE, x:x + cfg.TILE] += patch_prob\n            weight_acc[y:y + cfg.TILE, x:x + cfg.TILE] += 1.0\n\n    weight_acc[weight_acc == 0] = 1.0\n    prob_map = prob_acc / weight_acc\n\n    for m, _ in loaded_models:\n        del m\n    del loaded_models\n    gc.collect(); torch.cuda.empty_cache()\n    return prob_map, gt\n\n\n# %% [CELL 9] ---- Orchestration: convert data, split, train ensemble, evaluate, visualize --\ndef main():\n    # 1) Convert all needed fragments to Zarr (train 2,3 + held-out test 1)\n    for fid in CFG.TRAIN_FRAGMENTS + [CFG.TEST_FRAGMENT]:\n        convert_fragment_to_zarr(fid)\n\n    # 2) Build 80/20 spatial split from fragments 2 & 3\n    train_records, val_records = [], []\n    for fid in CFG.TRAIN_FRAGMENTS:\n        train_records += index_fragment_patches(fid, split=\"train\")\n        val_records   += index_fragment_patches(fid, split=\"val\")\n    print(f\"TOTAL train patches: {len(train_records)} | val patches: {len(val_records)}\")\n\n    # 3) Train each ensemble member\n    trained = []\n    for cfg_entry in CFG.ENSEMBLE_CONFIGS:\n        ckpt_path, thr, f05, hist = train_one_model(cfg_entry, train_records, val_records)\n        trained.append({\"ckpt_path\": ckpt_path, \"config\": cfg_entry, \"val_threshold\": thr, \"val_f05\": f05})\n\n    # 4) Aggregate validation threshold (mean across members) -> apply to test\n    ensemble_val_thr = float(np.mean([t[\"val_threshold\"] for t in trained]))\n    print(f\"Ensemble-averaged validation threshold: {ensemble_val_thr:.3f}\")\n\n    # 5) Full-fragment ensemble inference on the held-out TEST fragment (fragment 1)\n    prob_map, gt_map = predict_fragment_ensemble(CFG.TEST_FRAGMENT, trained)\n    pred_mask = (prob_map >= ensemble_val_thr).astype(np.uint8)\n\n    # 6) Final test metrics (precision/recall/F0.5/Dice) at the chosen threshold\n    p_flat, t_flat = pred_mask.ravel().astype(np.float32), gt_map.ravel().astype(np.float32)\n    tp = float((p_flat * t_flat).sum())\n    fp = float((p_flat * (1 - t_flat)).sum())\n    fn = float(((1 - p_flat) * t_flat).sum())\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n    f05_test = fbeta_score(precision, recall, CFG.F_BETA)\n    dice_test = 2 * tp / (2 * tp + fp + fn + 1e-8)\n    print(f\"TEST (fragment {CFG.TEST_FRAGMENT}) @thr={ensemble_val_thr:.3f} -> \"\n          f\"Precision={precision:.4f} Recall={recall:.4f} F0.5={f05_test:.4f} Dice={dice_test:.4f}\")\n\n    metrics = dict(\n        ensemble_val_threshold=ensemble_val_thr,\n        test_precision=precision, test_recall=recall,\n        test_f0_5=f05_test, test_dice=dice_test,\n        members=[{\"name\": t[\"config\"][\"name\"], \"val_f05\": t[\"val_f05\"],\n                  \"val_threshold\": t[\"val_threshold\"]} for t in trained],\n    )\n    with open(os.path.join(CFG.WORK_DIR, \"metrics_summary.json\"), \"w\") as f:\n        json.dump(metrics, f, indent=2)\n    print(\"Saved metrics_summary.json\")\n\n    # 7) Visualization: input (mid-slice + max-projection) vs ground truth vs prediction\n    vol, _ = open_fragment_zarr(CFG.TEST_FRAGMENT)\n    mid_slice = np.asarray(vol[CFG.Z_MID, :, :])\n    maxproj = np.asarray(vol[:, :, :]).max(axis=0) if vol.shape[0] * vol.shape[1] * vol.shape[2] < 4e8 \\\n        else mid_slice  # guard: skip full maxproj on huge volumes to avoid RAM spikes\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(mid_slice, cmap=\"gray\"); axes[0].set_title(f\"Input (mid slice z={CFG.Z_MID})\")\n    axes[1].imshow(gt_map, cmap=\"gray\"); axes[1].set_title(\"Ground Truth Ink Mask\")\n    axes[2].imshow(prob_map, cmap=\"magma\"); axes[2].set_title(\"Predicted Probability Map\")\n    axes[3].imshow(mid_slice, cmap=\"gray\")\n    axes[3].imshow(np.ma.masked_where(pred_mask == 0, pred_mask), cmap=\"autumn\", alpha=0.6)\n    axes[3].set_title(f\"Prediction Overlay (thr={ensemble_val_thr:.2f})\")\n    for ax in axes: ax.axis(\"off\")\n    plt.tight_layout()\n    viz_path = os.path.join(CFG.VIZ_DIR, f\"test_fragment{CFG.TEST_FRAGMENT}_comparison.png\")\n    plt.savefig(viz_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved visualization -> {viz_path}\")\n\n    del mid_slice, maxproj, prob_map, gt_map, pred_mask\n    gc.collect(); torch.cuda.empty_cache()\n    return metrics\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# VESUVIUS CHALLENGE — INK DETECTION PIPELINE (v3, High-Performance)\n# 3D Encoder → 2D Decoder U-Net with SE blocks, full 65-slice input,\n# large patch training, 5-fold cross-validation, TTA, and ensemble inference.\n# Based on winning Kaggle solutions (1st place: 3DCNN-SegFormer two-stage,\n# top-10: nn-UNet ResEncUNet ensembles).\n#\n# Target: Dice > 0.80\n#\n# NOTE: This script needs internet access in Kaggle notebook settings.\n# ==============================================================================\n\n!pip install -q segmentation-models-pytorch==0.2.1\n!pip install -q timm==0.9.12\n\nimport os, gc, random, time, json, warnings\nfrom collections import defaultdict\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    os.system(\"pip install -q albumentations\")\n    import albumentations as A\n\ntry:\n    import segmentation_models_pytorch as smp\nexcept ImportError:\n    os.system(\"pip install -q segmentation-models-pytorch\")\n    import segmentation_models_pytorch as smp\n\ntry:\n    from skimage.filters import threshold_otsu\n    from skimage.measure import label, regionprops\nexcept ImportError:\n    os.system(\"pip install -q scikit-image\")\n    from skimage.filters import threshold_otsu\n    from skimage.measure import label, regionprops\n\nwarnings.filterwarnings(\"ignore\")\n\n# ------------------------------------------------------------------------------\n# 0. CONFIG\n# ------------------------------------------------------------------------------\nclass CFG:\n    base_dir      = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags   = [\"2\", \"3\"]\n    test_frag     = \"1\"\n\n    # CRITICAL: Use ALL 65 slices or a well-chosen subset (24-48 slices)\n    # Winning solutions used 24-32 slices centered around the ink-rich region.\n    # Using all 65 gives the model maximum information but requires more memory.\n    # We use 32 slices (indices 16-48) as a balanced sweet spot.\n    depth_indices = list(range(16, 48))     # 32 slices\n    in_channels   = len(depth_indices)\n\n    # CRITICAL: Large patch size so model sees full characters\n    # Winning solutions used 512x512 or 1024x1024. We use 512 for Kaggle GPU limits.\n    patch_size    = 512\n    train_stride  = 256          # 50% overlap for training (dense sampling)\n    test_stride   = 128          # Dense overlap for inference\n\n    val_fraction          = 0.20\n    min_tissue_frac_train = 0.05    # Lower threshold = more training samples\n    min_tissue_frac_test  = 0.02\n\n    batch_size    = 2            # Reduced for large patches; use gradient accumulation\n    infer_batch   = 4\n    num_workers   = 2\n\n    epochs        = 2           # Much longer training (winners used 100-4000 epochs)\n    early_stop_patience = 15\n    lr            = 1e-4         # Lower LR for stability with large patches\n    weight_decay  = 3e-5\n\n    # CRITICAL: Deeper encoder for better feature extraction\n    # ResNet152 or EfficientNet-b4/b5 used in winning solutions\n    encoder_name    = \"resnet152\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"\n\n    # Gradient accumulation to simulate larger batch size\n    grad_accum_steps = 4         # Effective batch = 2 * 4 = 8\n\n    adabn_max_patches = 1200\n    threshold      = 0.5\n    seed           = 42\n\n    # K-Fold Cross Validation (winners used 5 folds)\n    n_folds        = 5\n\n    # TTA settings\n    use_tta        = True\n    tta_flips      = True\n    tta_rots       = True\n\n    out_dir        = \"/kaggle/working\"\n    ckpt_dir       = os.path.join(out_dir, \"checkpoints\")\n    viz_dir        = os.path.join(out_dir, \"test_visualizations\")\n\n    device         = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.ckpt_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nset_seed(CFG.seed)\n\n\ndef cfg_to_dict(cfg_cls):\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ------------------------------------------------------------------------------\n# 1. PRE-PROCESSING HELPERS\n# ------------------------------------------------------------------------------\ndef load_tissue_mask(frag_dir):\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid = tifffile.imread(os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n        thresh = mid.mean() * 0.15\n        mask = ((mid > thresh).astype(np.uint8)) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    for y in range(0, max(H - patch_size, 0) + 1, stride):\n        for x in range(0, max(W - patch_size, 0) + 1, stride):\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access with fragment-wise normalization stats.\"\"\"\n\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h = self._w = None\n        self._mean = None\n        self._std = None\n\n    def _ensure_open(self):\n        if self._slices is None:\n            slices = []\n            for p in self.paths:\n                try:\n                    arr = tifffile.memmap(p, mode=\"r\")\n                except Exception:\n                    arr = tifffile.imread(p)\n                slices.append(arr)\n            self._slices = slices\n            self._h, self._w = slices[0].shape\n            self._compute_stats()\n\n    def _compute_stats(self):\n        \"\"\"Compute fragment-wise mean/std for normalization.\"\"\"\n        # Sample a few slices to estimate stats\n        samples = []\n        step = max(1, len(self._slices) // 4)\n        for i in range(0, len(self._slices), step):\n            s = self._slices[i]\n            # Sample central region\n            cy, cx = self._h // 2, self._w // 2\n            h_s, w_s = self._h // 4, self._w // 4\n            samples.append(s[cy-h_s:cy+h_s, cx-w_s:cx+w_s].astype(np.float32))\n        all_samples = np.concatenate([s.ravel() for s in samples])\n        self._mean = float(np.mean(all_samples))\n        self._std = float(np.std(all_samples)) + 1e-6\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.float32)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size].astype(np.float32)\n            # Normalize using fragment-wise stats (CRITICAL for generalization)\n            block = (block - self._mean) / self._std\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        gc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 2. CHANNEL DROPOUT AUGMENTATION (winning technique)\n# ------------------------------------------------------------------------------\nclass ChannelDropout:\n    \"\"\"Randomly drop depth channels during training to improve robustness.\"\"\"\n    def __init__(self, max_drop=0.3):\n        self.max_drop = max_drop\n\n    def __call__(self, img):\n        # img: (C, H, W) numpy array\n        C = img.shape[0]\n        n_drop = int(C * random.uniform(0, self.max_drop))\n        if n_drop > 0 and C > n_drop:\n            drop_indices = random.sample(range(C), n_drop)\n            img = img.copy()\n            img[drop_indices] = 0\n        return img\n\n\ndef build_train_transform(size):\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.15, rotate_limit=30,\n                            border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.3),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.3),\n        A.GaussNoise(var_limit=(5.0, 20.0), p=0.2),\n    ])\n\n\n# ------------------------------------------------------------------------------\n# 3. DATASET\n# ------------------------------------------------------------------------------\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None,\n                 use_channel_dropout=False):\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n        self.channel_dropout = ChannelDropout() if use_channel_dropout else None\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        frag_id, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[frag_id]\n        patch = vol.read_patch(y, x, size)                # (C,H,W) float32, already normalized\n        label = self.labels[frag_id][y:y + size, x:x + size]\n\n        # Transpose for albumentations: HWC\n        img = np.transpose(patch, (1, 2, 0))\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        # Channel dropout\n        if self.channel_dropout is not None:\n            img = self.channel_dropout(img)\n\n        # Transpose back to CHW\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))\n        label = (label > 0).astype(np.float32)[None, ...]\n\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ------------------------------------------------------------------------------\n# 4. MODEL — 3D-aware architecture with SE blocks\n# ------------------------------------------------------------------------------\nclass SEBlock(nn.Module):\n    \"\"\"Squeeze-and-Excitation block for channel attention.\"\"\"\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(channels, channels // reduction, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(channels // reduction, channels, bias=False),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, _, _ = x.size()\n        y = self.avg_pool(x).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1)\n        return x * y.expand_as(x)\n\n\nclass InkDetectionModel(nn.Module):\n    \"\"\"\n    3D-aware U-Net: processes depth channels with 3D convolutions\n    before feeding into 2D U-Net encoder.\n    \"\"\"\n    def __init__(self, in_channels=32, encoder_name=\"resnet152\", encoder_weights=\"imagenet\"):\n        super().__init__()\n\n        # 3D feature extraction: compress depth dimension\n        self.depth_encoder = nn.Sequential(\n            nn.Conv3d(1, 8, kernel_size=(3, 3, 3), padding=(1, 1, 1)),\n            nn.BatchNorm3d(8),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(8, 16, kernel_size=(3, 3, 3), padding=(1, 1, 1)),\n            nn.BatchNorm3d(16),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(16, 32, kernel_size=(3, 3, 3), padding=(1, 1, 1)),\n            nn.BatchNorm3d(32),\n            nn.ReLU(inplace=True),\n        )\n\n        # Project 3D features to 2D: pool over depth then conv\n        self.depth_pool = nn.AdaptiveAvgPool3d((1, None, None))\n        self.depth_proj = nn.Conv2d(32, 64, kernel_size=1)\n\n        # 2D U-Net with pretrained encoder\n        self.unet = smp.Unet(\n            encoder_name=encoder_name,\n            encoder_weights=encoder_weights,\n            in_channels=64,  # Projected from 3D features\n            classes=1,\n            decoder_attention_type=\"scse\",\n        )\n\n        # Additional SE block before final output\n        self.se = SEBlock(1, reduction=4)\n\n    def forward(self, x):\n        # x: (B, C, H, W) where C = depth slices\n        B, C, H, W = x.shape\n\n        # Add channel dim for 3D conv: (B, 1, C, H, W)\n        x_3d = x.unsqueeze(1)\n        x_3d = self.depth_encoder(x_3d)  # (B, 32, C, H, W)\n\n        # Pool depth dimension\n        x_3d = self.depth_pool(x_3d)     # (B, 32, 1, H, W)\n        x_2d = x_3d.squeeze(2)           # (B, 32, H, W)\n\n        # Project to encoder input channels\n        x_2d = self.depth_proj(x_2d)     # (B, 64, H, W)\n\n        # 2D U-Net\n        out = self.unet(x_2d)            # (B, 1, H, W)\n        out = self.se(out)\n        return out\n\n\ndef build_model():\n    return InkDetectionModel(\n        in_channels=CFG.in_channels,\n        encoder_name=CFG.encoder_name,\n        encoder_weights=CFG.encoder_weights\n    )\n\n\n# ------------------------------------------------------------------------------\n# 5. LOSS & METRICS\n# ------------------------------------------------------------------------------\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(1)\n    union = probs.sum(1) + t.sum(1)\n    return 1 - ((2 * inter + eps) / (union + eps)).mean()\n\n\nclass FocalTverskyLoss(nn.Module):\n    \"\"\"Focal Tversky loss - better for imbalanced segmentation.\"\"\"\n    def __init__(self, alpha=0.7, beta=0.3, gamma=1.33):\n        super().__init__()\n        self.alpha = alpha\n        self.beta = beta\n        self.gamma = gamma\n\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        tp = (probs * targets).sum(dim=(1,2,3))\n        fp = (probs * (1 - targets)).sum(dim=(1,2,3))\n        fn = ((1 - probs) * targets).sum(dim=(1,2,3))\n        tversky = (tp + 1e-6) / (tp + self.alpha * fp + self.beta * fn + 1e-6)\n        return torch.pow(1 - tversky, self.gamma).mean()\n\n\nclass ComboLoss(nn.Module):\n    def __init__(self, bce_w=0.4, dice_w=0.4, ft_w=0.2, pos_weight=None):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n        self.bce_w = bce_w\n        self.dice_w = dice_w\n        self.ft_w = ft_w\n        self.ft_loss = FocalTverskyLoss()\n\n    def forward(self, logits, targets):\n        bce = self.bce(logits, targets)\n        dice = dice_loss(logits, targets)\n        ft = self.ft_loss(logits, targets)\n        return self.bce_w * bce + self.dice_w * dice + self.ft_w * ft\n\n\n@torch.no_grad()\ndef compute_metrics(probs, targets, threshold=0.5, eps=1e-6):\n    preds = (probs > threshold).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    return {\"dice\": dice.item(), \"iou\": iou.item(), \"precision\": precision.item(),\n            \"recall\": recall.item(), \"fbeta0.5\": fbeta.item()}\n\n\nclass GlobalConfusionAccumulator:\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1 - targets)).sum().item()\n        self.fn += ((1 - preds) * targets).sum().item()\n\n    def compute(self, eps=1e-6, beta2=0.25):\n        dice = (2 * self.tp + eps) / (2 * self.tp + self.fp + self.fn + eps)\n        iou = (self.tp + eps) / (self.tp + self.fp + self.fn + eps)\n        precision = (self.tp + eps) / (self.tp + self.fp + eps)\n        recall = (self.tp + eps) / (self.tp + self.fn + eps)\n        fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n        return {\"dice\": dice, \"iou\": iou, \"precision\": precision,\n                \"recall\": recall, \"fbeta0.5\": fbeta}\n\n\n# ------------------------------------------------------------------------------\n# 6. BUILD DATA\n# ------------------------------------------------------------------------------\nprint(\"Building tissue masks, label maps and patch grids ...\")\n\ntrain_volumes, train_labels_full = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"  fragment {fid}: mask {mask.shape}, {len(coords)} candidate patches\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    all_samples.extend([(fid, y, x) for (y, x) in coords])\n\n    del mask\n    gc.collect()\n\nrandom.shuffle(all_samples)\nprint(f\"Total patches: {len(all_samples)}\")\n\n\n# ------------------------------------------------------------------------------\n# 7. K-FOLD CROSS VALIDATION TRAINING\n# ------------------------------------------------------------------------------\ndef get_pos_frac(labels_full, samples, patch_size, n_samples=100):\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    frac = pos / max(total, 1)\n    return max(frac, 1e-4)\n\n\n# Prepare K-Fold\nkf = KFold(n_splits=CFG.n_folds, shuffle=True, random_state=CFG.seed)\nfold_results = []\n\nfor fold_idx, (train_idx, val_idx) in enumerate(kf.split(all_samples)):\n    print(f\"\\n{'='*60}\")\n    print(f\"FOLD {fold_idx + 1}/{CFG.n_folds}\")\n    print(f\"{'='*60}\")\n\n    fold_ckpt = os.path.join(CFG.ckpt_dir, f\"fold{fold_idx}_best.pth\")\n\n    train_samples = [all_samples[i] for i in train_idx]\n    val_samples = [all_samples[i] for i in val_idx]\n\n    print(f\"Train: {len(train_samples)} | Val: {len(val_samples)}\")\n\n    train_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples,\n                                CFG.patch_size, transform=build_train_transform(CFG.patch_size),\n                                use_channel_dropout=True)\n    val_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples,\n                              CFG.patch_size, transform=None, use_channel_dropout=False)\n\n    train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                               num_workers=CFG.num_workers, pin_memory=True,\n                               drop_last=True, persistent_workers=True)\n    val_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                             num_workers=CFG.num_workers, pin_memory=True,\n                             persistent_workers=True)\n\n    # Estimate positive fraction\n    pos_frac = get_pos_frac(train_labels_full, train_samples, CFG.patch_size)\n    print(f\"  pos_frac = {pos_frac:.5f}\")\n\n    model = build_model().to(CFG.device)\n\n    # Bias init\n    with torch.no_grad():\n        bias_val = float(np.log(pos_frac / (1 - pos_frac)))\n        model.unet.segmentation_head[0].bias.fill_(bias_val)\n    print(f\"  bias init = {bias_val:.3f}\")\n\n    pos_weight_val = float(np.clip((1 - pos_frac) / pos_frac, 1.0, 20.0))\n    pos_weight = torch.tensor([pos_weight_val], device=CFG.device)\n\n    criterion = ComboLoss(pos_weight=pos_weight)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=10, T_mult=2)\n    scaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\n    best_val_dice = -1.0\n    epochs_no_improve = 0\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n    for epoch in range(1, CFG.epochs + 1):\n        t0 = time.time()\n\n        # Training with gradient accumulation\n        model.train()\n        total_loss = 0.0\n        global_acc = GlobalConfusionAccumulator()\n        optimizer.zero_grad(set_to_none=True)\n\n        for step, (imgs, masks) in enumerate(train_loader):\n            imgs = imgs.to(CFG.device, non_blocking=True)\n            masks = masks.to(CFG.device, non_blocking=True)\n\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks) / CFG.grad_accum_steps\n\n            scaler.scale(loss).backward()\n\n            if (step + 1) % CFG.grad_accum_steps == 0 or (step + 1) == len(train_loader):\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad(set_to_none=True)\n\n            probs = torch.sigmoid(logits.detach())\n            global_acc.update(probs, masks, CFG.threshold)\n            total_loss += loss.item() * CFG.grad_accum_steps\n\n            del imgs, masks, logits, probs\n\n        train_loss = total_loss / max(len(train_loader), 1)\n        train_metrics = global_acc.compute()\n\n        # Validation\n        model.eval()\n        val_loss = 0.0\n        val_acc = GlobalConfusionAccumulator()\n        with torch.no_grad():\n            for imgs, masks in val_loader:\n                imgs = imgs.to(CFG.device, non_blocking=True)\n                masks = masks.to(CFG.device, non_blocking=True)\n                with autocast(enabled=(CFG.device == \"cuda\")):\n                    logits = model(imgs)\n                    loss = criterion(logits, masks)\n                probs = torch.sigmoid(logits)\n                val_acc.update(probs, masks, CFG.threshold)\n                val_loss += loss.item()\n                del imgs, masks, logits, probs\n\n        val_loss = val_loss / max(len(val_loader), 1)\n        val_metrics = val_acc.compute()\n        scheduler.step()\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_dice\"].append(val_metrics[\"dice\"])\n\n        print(f\"[{epoch:02d}/{CFG.epochs}] \"\n              f\"train_loss={train_loss:.4f} dice={train_metrics['dice']:.4f} | \"\n              f\"val_loss={val_loss:.4f} dice={val_metrics['dice']:.4f} \"\n              f\"fbeta0.5={val_metrics['fbeta0.5']:.4f} ({time.time()-t0:.1f}s)\")\n\n        if val_metrics[\"dice\"] > best_val_dice:\n            best_val_dice = val_metrics[\"dice\"]\n            epochs_no_improve = 0\n            torch.save({\"model\": model.state_dict(), \"cfg\": cfg_to_dict(CFG)}, fold_ckpt)\n            print(f\"  -> saved best (val_dice={best_val_dice:.4f})\")\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= CFG.early_stop_patience:\n                print(f\"  -> early stop\")\n                break\n\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    fold_results.append({\n        \"fold\": fold_idx,\n        \"best_val_dice\": best_val_dice,\n        \"ckpt\": fold_ckpt\n    })\n\n    del train_loader, val_loader, train_ds, val_ds\n    gc.collect()\n    torch.cuda.empty_cache()\n\nprint(f\"\\n{'='*60}\")\nprint(\"CROSS-VALIDATION SUMMARY\")\nprint(f\"{'='*60}\")\nfor fr in fold_results:\n    print(f\"Fold {fr['fold']}: best_val_dice = {fr['best_val_dice']:.4f}\")\n\n\n# ------------------------------------------------------------------------------\n# 8. TEST-TIME AUGMENTATION (TTA) INFERENCE\n# ------------------------------------------------------------------------------\ndef tta_inference(model, img_batch):\n    \"\"\"Apply TTA: original + hflip + vflip + hflip+vflip + 90° rotations.\"\"\"\n    if not CFG.use_tta:\n        with torch.no_grad():\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                return torch.sigmoid(model(img_batch))\n\n    predictions = []\n\n    # Original\n    with torch.no_grad():\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            predictions.append(torch.sigmoid(model(img_batch)))\n\n    # Horizontal flip\n    if CFG.tta_flips:\n        img_h = torch.flip(img_batch, dims=[3])\n        with torch.no_grad():\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                pred_h = torch.sigmoid(model(img_h))\n        predictions.append(torch.flip(pred_h, dims=[3]))\n\n    # Vertical flip\n    if CFG.tta_flips:\n        img_v = torch.flip(img_batch, dims=[2])\n        with torch.no_grad():\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                pred_v = torch.sigmoid(model(img_v))\n        predictions.append(torch.flip(pred_v, dims=[2]))\n\n    # 90° rotation\n    if CFG.tta_rots:\n        img_r = torch.rot90(img_batch, k=1, dims=[2, 3])\n        with torch.no_grad():\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                pred_r = torch.sigmoid(model(img_r))\n        predictions.append(torch.rot90(pred_r, k=-1, dims=[2, 3]))\n\n    # Average all predictions\n    return torch.stack(predictions).mean(dim=0)\n\n\ndef gaussian_window(size, sigma_frac=0.5):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / win.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, device, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Fragment 1: inference over {len(coords)} patches ...\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        inp = torch.stack(batch_imgs).to(device)\n        prob = tta_inference(model, inp).float().cpu().numpy()[:, 0]\n        for p, (cy, cx) in zip(prob, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear()\n        batch_coords.clear()\n\n    for (y, x) in coords:\n        raw = vol.read_patch(y, x, patch_size)  # Already normalized\n        raw = torch.from_numpy(raw).float()\n        batch_imgs.append(raw)\n        batch_coords.append((y, x))\n        if len(batch_imgs) == batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum == 0] = 1.0\n    return pred_sum / weight_sum\n\n\ndef postprocess(prob_map, mask, threshold):\n    \"\"\"Advanced post-processing with connected component filtering.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n\n    # Morphological cleanup\n    binary = ndimage.binary_opening(binary, structure=np.ones((3, 3))).astype(np.uint8)\n    binary = ndimage.binary_closing(binary, structure=np.ones((5, 5))).astype(np.uint8)\n\n    # Remove small connected components (noise)\n    labeled = label(binary)\n    regions = regionprops(labeled)\n    for region in regions:\n        if region.area < 20:  # Remove tiny specks\n            binary[labeled == region.label] = 0\n\n    # Mask out non-tissue regions\n    binary = binary * mask\n\n    return binary\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm(model, vol, mask, patch_size, stride, device,\n                           max_patches, batch_size):\n    for m in model.modules():\n        if isinstance(m, (nn.BatchNorm2d, nn.BatchNorm3d)):\n            m.reset_running_stats()\n            m.momentum = None\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"  recalibrating BN using {len(coords)} patches ...\")\n\n    model.train()\n    for i in range(0, len(coords), batch_size):\n        batch = coords[i:i + batch_size]\n        imgs = []\n        for (y, x) in batch:\n            raw = vol.read_patch(y, x, patch_size)\n            imgs.append(torch.from_numpy(raw).float())\n        inp = torch.stack(imgs).to(device)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp)\n        del inp\n    model.eval()\n    torch.cuda.empty_cache()\n    return model\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback=0.5):\n    vals = prob_map[mask > 0]\n    try:\n        return float(threshold_otsu(vals))\n    except Exception:\n        return fallback\n\n\n# ------------------------------------------------------------------------------\n# 9. ENSEMBLE INFERENCE ACROSS ALL FOLDS\n# ------------------------------------------------------------------------------\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_labels = load_ink_labels(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\ngt_t = torch.from_numpy(test_labels.astype(np.float32)) * torch.from_numpy(test_mask.astype(np.float32))\n\nprint(f\"\\n{'='*60}\")\nprint(\"ENSEMBLE INFERENCE\")\nprint(f\"{'='*60}\")\n\n# Collect predictions from all folds\nall_probs = []\n\nfor fold_idx, fr in enumerate(fold_results):\n    print(f\"\\n--- Fold {fold_idx} ---\")\n    ckpt_path = fr[\"ckpt\"]\n    if not os.path.exists(ckpt_path):\n        print(f\"  Skipping (no checkpoint)\")\n        continue\n\n    model = build_model().to(CFG.device)\n    model.load_state_dict(torch.load(ckpt_path, map_location=CFG.device)[\"model\"])\n\n    # AdaBN recalibration\n    model = recalibrate_batchnorm(model, test_vol, test_mask, CFG.patch_size,\n                                   CFG.test_stride, CFG.device,\n                                   CFG.adabn_max_patches, CFG.infer_batch)\n\n    # Inference\n    test_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                          CFG.test_stride, CFG.device, CFG.infer_batch)\n    all_probs.append(test_prob)\n\n    # Individual fold metrics\n    otsu_t = compute_otsu_threshold(test_prob, test_mask)\n    fold_metrics = compute_metrics(torch.from_numpy(test_prob), gt_t, otsu_t)\n    print(f\"  Fold {fold_idx} metrics [AdaBN+Otsu, thr={otsu_t:.3f}]: {fold_metrics}\")\n\n    del model\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# Ensemble: average all fold predictions\nif len(all_probs) > 0:\n    ensemble_prob = np.mean(all_probs, axis=0)\n    print(f\"\\n{'='*60}\")\n    print(\"FINAL ENSEMBLE RESULTS\")\n    print(f\"{'='*60}\")\n\n    # Evaluate ensemble with different thresholds\n    for thr in [0.3, 0.4, 0.5, 0.6]:\n        ens_metrics = compute_metrics(torch.from_numpy(ensemble_prob), gt_t, thr)\n        print(f\"  Ensemble @ threshold={thr}: {ens_metrics}\")\n\n    # Best threshold via Otsu\n    otsu_threshold = compute_otsu_threshold(ensemble_prob, test_mask, fallback=0.5)\n    final_metrics = compute_metrics(torch.from_numpy(ensemble_prob), gt_t, otsu_threshold)\n    print(f\"\\n  FINAL [AdaBN + Otsu + Ensemble, thr={otsu_threshold:.3f}]:\")\n    print(f\"  {final_metrics}\")\n\n    # Post-processed\n    test_pred_bin = postprocess(ensemble_prob, test_mask, otsu_threshold)\n    pp_metrics = compute_metrics(torch.from_numpy(ensemble_prob), gt_t, otsu_threshold)\n    # Recompute with binary mask\n    pp_pred_t = torch.from_numpy(test_pred_bin.astype(np.float32))\n    pp_metrics = compute_metrics(pp_pred_t, gt_t, 0.5)\n    print(f\"  POST-PROCESSED: {pp_metrics}\")\n\n    # Save outputs\n    np.save(os.path.join(CFG.out_dir, \"ensemble_prob.npy\"), ensemble_prob)\n    np.save(os.path.join(CFG.out_dir, \"ensemble_pred.npy\"), test_pred_bin)\n\n    with open(os.path.join(CFG.out_dir, \"metrics_summary.json\"), \"w\") as f:\n        json.dump({\n            \"fold_results\": fold_results,\n            \"ensemble\": final_metrics,\n            \"post_processed\": pp_metrics,\n            \"otsu_threshold\": otsu_threshold,\n        }, f, indent=2)\n\n\n# ------------------------------------------------------------------------------\n# 10. VISUALIZATIONS\n# ------------------------------------------------------------------------------\ndef save_full_overview():\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize(test_labels * 255, small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    axes[0].imshow(small, cmap=\"gray\")\n    axes[0].set_title(f\"Input (slice {mid_idx})\")\n    axes[1].imshow(gt_small, cmap=\"gray\")\n    axes[1].set_title(\"Ground truth\")\n    axes[2].imshow(pred_small, cmap=\"gray\")\n    axes[2].set_title(\"Ensemble Prediction\")\n    for a in axes:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_overview.png\")\n    plt.savefig(path, dpi=150)\n    plt.close(fig)\n    print(f\"Saved {path}\")\n\n\nsave_full_overview()\n\ntest_vol.close()\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"\\n=== DONE ===\")\nprint(f\"Checkpoints: {CFG.ckpt_dir}\")\nprint(f\"Metrics: {os.path.join(CFG.out_dir, 'metrics_summary.json')}\")\nprint(f\"Visualizations: {CFG.viz_dir}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# عند الرن 10 مرات =0.43 اعطى نتائج عالية وجميلة \n#dice= 0.45,batch=6,epochs=8 ,stride=94\n# dice= 0.41 ,batch=10      epochs=6, stride=64\n# ==============================================================================\n# VESUVIUS CHALLENGE — INK DETECTION PIPELINE (v2, accuracy-focused)\n# Pretrained ResNet34 encoder (ImageNet) + scSE-attention U-Net decoder,\n# centered depth-slice subset, per-patch normalization, validation-tuned\n# decision threshold, early stopping.\n# Train: fragments 2 & 3 (80/20 split)  |  Test: fragment 1\n# OOM-safety: disk-backed memmap volumes, patch-based training, sliding-window\n# (non-TTA) inference, AMP, aggressive gc.\n#\n# NOTE: this script needs internet access enabled in the Kaggle notebook\n# settings, both to `pip install segmentation-models-pytorch` and to download\n# the ImageNet-pretrained ResNet34 weights.\n#\n# Paste this whole file into a single Kaggle notebook cell and run.\n# ==============================================================================\n#Final TEST (fragment 1) metrics [AdaBN + Otsu, threshold=0.439]: \n#{'dice': 0.43914005160331726, 'iou': 0.2813449501991272, 'precision': 0.3286738395690918, 'recall': 0.6614518761634827, 'fbeta0.5': 0.36544597148895264}\n!pip install segmentation-models-pytorch==0.2.0\n\nimport os, gc, random, time, json\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    os.system(\"pip install -q albumentations\")\n    import albumentations as A\n\ntry:\n    import segmentation_models_pytorch as smp\nexcept ImportError:\n    os.system(\"pip install -q segmentation-models-pytorch\")\n    import segmentation_models_pytorch as smp\n\ntry:\n    from skimage.filters import threshold_otsu\nexcept ImportError:\n    os.system(\"pip install -q scikit-image\")\n    from skimage.filters import threshold_otsu\n\n\n# ------------------------------------------------------------------------------\n# 0. CONFIG\n# ------------------------------------------------------------------------------\nclass CFG:\n    base_dir      = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags   = [\"2\", \"3\"]\n    test_frag     = \"1\"\n\n    # Use a centered subset of the 65 slices: the ink signal concentrates\n    # mid-depth, and dropping the noisier outer slices both improves signal\n    # and reduces compute/memory versus all 65 channels.\n    depth_indices = list(range(16, 38))     # 32 slices, centered on 31.5\n    in_channels   = len(depth_indices)\n\n    patch_size    = 224\n    train_stride  = 32              #94 dense overlap -> more training samples\n    test_stride   = 32\n\n    val_fraction          = 0.20\n    min_tissue_frac_train = 0.10    # skip near-empty patches when building train grid\n    min_tissue_frac_test  = 0.02\n\n    batch_size    = 4\n    infer_batch   = 12\n    num_workers   = 2\n\n    epochs        = 3\n    early_stop_patience = 8\n    lr            = 2e-4\n    weight_decay  = 1e-4\n\n    encoder_name    = \"resnet50\"    # BatchNorm-based (required for the AdaBN\n                                     # recalibration step below). ConvNeXt has\n                                     # stronger raw features but uses LayerNorm\n                                     # instead of BatchNorm, so it has no running\n                                     # stats to recalibrate -- AdaBN would only\n                                     # touch the decoder if you swapped to e.g.\n                                     # encoder_name=\"tu-convnext_tiny\". Try that\n                                     # only if you're willing to drop/rework the\n                                     # BN-recalibration step.\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"   # squeeze-and-excitation attention decoder\n\n    adabn_max_patches = 800         # unlabeled fragment-1 patches used to\n                                     # recalibrate BatchNorm stats (no labels)\n\n    threshold      = 0.5            # only used for in-training monitoring\n    seed           = 42\n\n    out_dir        = \"/kaggle/working\"\n    ckpt_path      = os.path.join(out_dir, \"vesuviusnet_best.pth\")\n    viz_dir        = os.path.join(out_dir, \"test_visualizations\")\n\n    device         = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\n\nset_seed(CFG.seed)\ntorch.backends.cudnn.benchmark = True\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=60):\n    \"\"\"Quickly estimate the fraction of ink-positive pixels across a random\n    subset of training patches, used to (a) init the output bias so the model\n    starts near the true prior instead of 50/50, and (b) weight BCE.\"\"\"\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    frac = pos / max(total, 1)\n    return max(frac, 1e-4)\n\n\ndef cfg_to_dict(cfg_cls):\n    \"\"\"vars() on a *class* returns a non-picklable mappingproxy, so build a\n    plain dict of the simple (picklable) config values by hand instead.\"\"\"\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ------------------------------------------------------------------------------\n# 1. PRE-PROCESSING HELPERS: tissue mask, patch grid, disk-backed volume reader\n# ------------------------------------------------------------------------------\ndef load_tissue_mask(frag_dir):\n    \"\"\"Load fragment tissue mask (mask.png if present, else derive from a mid slice).\"\"\"\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid = tifffile.imread(os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n        thresh = mid.mean() * 0.15\n        mask = ((mid > thresh).astype(np.uint8)) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    for y in range(0, max(H - patch_size, 0) + 1, stride):\n        for x in range(0, max(W - patch_size, 0) + 1, stride):\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access to a fragment's surface volume, restricted\n    to a chosen subset of depth slices. Avoids loading the full multi-GB\n    volume into RAM.\"\"\"\n\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h, self._w = None, None\n\n    def _ensure_open(self):\n        if self._slices is None:\n            slices = []\n            for p in self.paths:\n                try:\n                    arr = tifffile.memmap(p, mode=\"r\")\n                except Exception:\n                    # fallback if the tif can't be memory-mapped (e.g. compressed)\n                    arr = tifffile.imread(p)\n                slices.append(arr)\n            self._slices = slices\n            self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        gc.collect()\n\n\ndef normalize_patch(img_float):\n    \"\"\"Per-patch (instance) normalization: zero-mean, unit-variance across the\n    whole patch. This is applied identically at train and inference time and\n    helps offset scanner/intensity differences between fragments (a likely\n    contributor to the train/test domain gap).\"\"\"\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ------------------------------------------------------------------------------\n# 2. AUGMENTATION (train only) — kept to spatial-safe, version-stable transforms\n# ------------------------------------------------------------------------------\ndef build_train_transform(size):\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.15, rotate_limit=25,\n                            border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.25),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.3),\n    ])\n\n\n# ------------------------------------------------------------------------------\n# 3. DATASET\n# ------------------------------------------------------------------------------\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None):\n        \"\"\"\n        volumes: dict frag_id -> FragmentVolume\n        labels:  dict frag_id -> full-res (H,W) uint8 ink label array\n        samples: list of (frag_id, y, x)\n        \"\"\"\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        frag_id, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[frag_id]\n        patch = vol.read_patch(y, x, size)                # (C,H,W) uint8\n        label = self.labels[frag_id][y:y + size, x:x + size]\n\n        img = np.transpose(patch, (1, 2, 0))               # HWC for albumentations\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))  # CHW\n        label = (label > 0).astype(np.float32)[None, ...]         # 1HW\n\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ------------------------------------------------------------------------------\n# 4. MODEL — pretrained ResNet34 encoder + scSE-attention U-Net decoder\n# ------------------------------------------------------------------------------\ndef build_model():\n    model = smp.Unet(\n        encoder_name=CFG.encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=CFG.in_channels,\n        classes=1,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n    return model\n\n\n# ------------------------------------------------------------------------------\n# 5. LOSS & METRICS\n# ------------------------------------------------------------------------------\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(1)\n    union = probs.sum(1) + t.sum(1)\n    return 1 - ((2 * inter + eps) / (union + eps)).mean()\n\n\nclass ComboLoss(nn.Module):\n    def __init__(self, bce_w=0.5, dice_w=0.5, pos_weight=None):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n        self.bce_w, self.dice_w = bce_w, dice_w\n\n    def forward(self, logits, targets):\n        return self.bce_w * self.bce(logits, targets) + self.dice_w * dice_loss(logits, targets)\n\n\n@torch.no_grad()\ndef compute_metrics(probs, targets, threshold=0.5, eps=1e-6):\n    preds = (probs > threshold).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    return {\"dice\": dice.item(), \"iou\": iou.item(), \"precision\": precision.item(),\n            \"recall\": recall.item(), \"fbeta0.5\": fbeta.item()}\n\n\nclass GlobalConfusionAccumulator:\n    \"\"\"Accumulates TP/FP/FN across an entire epoch (rather than averaging\n    per-batch Dice) so metrics aren't swamped by noise from the many\n    near-empty patches typical of ink-detection data.\"\"\"\n\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1 - targets)).sum().item()\n        self.fn += ((1 - preds) * targets).sum().item()\n\n    def compute(self, eps=1e-6, beta2=0.25):\n        dice = (2 * self.tp + eps) / (2 * self.tp + self.fp + self.fn + eps)\n        iou = (self.tp + eps) / (self.tp + self.fp + self.fn + eps)\n        precision = (self.tp + eps) / (self.tp + self.fp + eps)\n        recall = (self.tp + eps) / (self.tp + self.fn + eps)\n        fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n        return {\"dice\": dice, \"iou\": iou, \"precision\": precision,\n                \"recall\": recall, \"fbeta0.5\": fbeta}\n\n\n# ------------------------------------------------------------------------------\n# 6. BUILD DATA (fragments 2 & 3 -> train/val ; fragment 1 -> test)\n# ------------------------------------------------------------------------------\nprint(\"Building tissue masks, label maps and patch grids ...\")\n\ntrain_volumes, train_labels_full = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"  fragment {fid}: mask {mask.shape}, {len(coords)} candidate patches\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    all_samples.extend([(fid, y, x) for (y, x) in coords])\n\n    del mask\n    gc.collect()\n\nrandom.shuffle(all_samples)\nn_val = int(len(all_samples) * CFG.val_fraction)\nval_samples = all_samples[:n_val]\ntrain_samples = all_samples[n_val:]\nprint(f\"Total patches: {len(all_samples)}  -> train {len(train_samples)} / val {len(val_samples)}\")\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples,\n                            CFG.patch_size, transform=build_train_transform(CFG.patch_size))\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples,\n                          CFG.patch_size, transform=None)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=True, persistent_workers=CFG.num_workers > 0)\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         persistent_workers=CFG.num_workers > 0)\n\n\n# ------------------------------------------------------------------------------\n# 7. TRAIN / VALIDATE\n# ------------------------------------------------------------------------------\nprint(\"Estimating ink-pixel prior for bias-init / class weighting ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size)\nprint(f\"  estimated positive-pixel fraction: {pos_frac:.5f}\")\n\nmodel = build_model().to(CFG.device)\n\n# Bias-init trick: start the output layer predicting ~pos_frac everywhere\n# instead of ~0.5, so the model doesn't have to unlearn a bad 50/50 prior.\nwith torch.no_grad():\n    bias_val = float(np.log(pos_frac / (1 - pos_frac)))\n    model.segmentation_head[0].bias.fill_(bias_val)\nprint(f\"  output layer bias initialized to {bias_val:.3f} (sigmoid={pos_frac:.4f})\")\n\npos_weight_val = float(np.clip((1 - pos_frac) / pos_frac, 1.0, 15.0))\npos_weight = torch.tensor([pos_weight_val], device=CFG.device)\nprint(f\"  BCE pos_weight = {pos_weight_val:.2f}\")\n\ncriterion = ComboLoss(pos_weight=pos_weight)\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n\ndef run_epoch(loader, train_mode, threshold=CFG.threshold):\n    model.train(train_mode)\n    total_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                scaler.step(optimizer)\n                scaler.update()\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.item()\n\n        del imgs, masks, logits, probs\n    torch.cuda.empty_cache()\n    return total_loss / max(len(loader), 1), global_acc.compute()\n\n\nprint(\"\\nStarting training ...\")\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_global = run_epoch(train_loader, train_mode=True)\n    val_loss, val_global = run_epoch(val_loader, train_mode=False)\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_global[\"dice\"])\n\n    print(f\"[{epoch:02d}/{CFG.epochs}] \"\n          f\"train_loss={train_loss:.4f} dice={train_global['dice']:.4f} | \"\n          f\"val_loss={val_loss:.4f} dice={val_global['dice']:.4f} \"\n          f\"fbeta0.5={val_global['fbeta0.5']:.4f} recall={val_global['recall']:.4f} \"\n          f\"precision={val_global['precision']:.4f} ({time.time()-t0:.1f}s)\")\n\n    if val_global[\"dice\"] > best_val_dice:\n        best_val_dice = val_global[\"dice\"]\n        epochs_no_improve = 0\n        torch.save({\"model\": model.state_dict(), \"cfg\": cfg_to_dict(CFG)}, CFG.ckpt_path)\n        print(f\"  -> saved new best checkpoint (val_dice={best_val_dice:.4f})\")\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(f\"  -> no val improvement for {CFG.early_stop_patience} epochs, stopping early.\")\n            break\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# reload best checkpoint\nmodel.load_state_dict(torch.load(CFG.ckpt_path, map_location=CFG.device)[\"model\"])\n\n\n# ------------------------------------------------------------------------------\n# 7b. TUNE DECISION THRESHOLD ON VALIDATION SET (maximize F0.5)\n# ------------------------------------------------------------------------------\n@torch.no_grad()\ndef find_best_threshold(model, loader, thresholds=np.arange(0.10, 0.91, 0.05)):\n    model.eval()\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        targets = masks.cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * targets).sum()\n            fp[i] += (preds * (1 - targets)).sum()\n            fn[i] += ((1 - preds) * targets).sum()\n        del imgs, masks, logits, probs\n    eps, beta2 = 1e-6, 0.25\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    best_idx = int(np.argmax(fbeta))\n    return float(thresholds[best_idx]), float(fbeta[best_idx])\n\n\nprint(\"\\nTuning decision threshold on validation set ...\")\nbest_threshold, best_val_fbeta = find_best_threshold(model, val_loader)\nprint(f\"  best threshold = {best_threshold:.2f} (val fbeta0.5 = {best_val_fbeta:.4f})\")\n\n# final train / val metrics at the tuned threshold\n_, final_train_metrics = run_epoch(train_loader, train_mode=False, threshold=best_threshold)\n_, final_val_metrics = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal TRAIN metrics:\", final_train_metrics)\nprint(\"Final VAL metrics:  \", final_val_metrics)\n\n# free fragment 2/3 volumes before loading fragment 1\nfor v in train_volumes.values():\n    v.close()\ndel train_loader, val_loader, train_ds, val_ds, train_volumes, train_labels_full\ngc.collect()\ntorch.cuda.empty_cache()\n\n\n# ------------------------------------------------------------------------------\n# 8. SLIDING-WINDOW INFERENCE ON FRAGMENT 1 (single forward pass per patch — no TTA)\n# ------------------------------------------------------------------------------\ndef gaussian_window(size, sigma_frac=0.5):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / win.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, device, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Fragment 1: running sliding-window inference over {len(coords)} patches ...\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        inp = torch.from_numpy(np.stack(batch_imgs)).to(device)\n        with autocast(enabled=(device == \"cuda\")):\n            logits = model(inp)\n            prob = torch.sigmoid(logits).float().cpu().numpy()[:, 0]\n        for p, (cy, cx) in zip(prob, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n\n    for (y, x) in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw)   # same per-patch normalization as training\n        batch_imgs.append(raw); batch_coords.append((y, x))\n        if len(batch_imgs) == batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum == 0] = 1.0\n    return pred_sum / weight_sum\n\n\ndef postprocess(prob_map, threshold):\n    \"\"\"Post-processing WITHOUT test-time augmentation: threshold + light\n    morphological cleanup to remove isolated speckle noise.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = ndimage.binary_opening(binary, structure=np.ones((3, 3))).astype(np.uint8)\n    binary = ndimage.binary_closing(binary, structure=np.ones((5, 5))).astype(np.uint8)\n    return binary\n\n\n@torch.no_grad()\ndef recalibrate_batchnorm(model, vol, mask, patch_size, stride, device,\n                           max_patches, batch_size):\n    \"\"\"Unsupervised domain adaptation (AdaBN, Li et al. 2016): reset every\n    BatchNorm layer's running mean/variance and re-estimate them purely from\n    the TARGET fragment's own images (no labels used anywhere) so the\n    network's internal activation statistics match fragment 1 instead of the\n    source fragments 2/3. Cheap, forward-pass-only, no gradient/OOM risk.\"\"\"\n    for m in model.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.reset_running_stats()\n            m.momentum = None  # cumulative moving average over all seen batches\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    random.shuffle(coords)\n    coords = coords[:max_patches]\n    print(f\"  recalibrating BatchNorm using {len(coords)} unlabeled fragment-1 patches ...\")\n\n    model.train()  # BN uses batch stats + updates running stats; no_grad keeps it cheap\n    for i in range(0, len(coords), batch_size):\n        batch = coords[i:i + batch_size]\n        imgs = []\n        for (y, x) in batch:\n            raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n            raw = normalize_patch(raw)\n            imgs.append(raw)\n        inp = torch.from_numpy(np.stack(imgs)).to(device)\n        with autocast(enabled=(device == \"cuda\")):\n            model(inp)\n        del inp\n    model.eval()\n    torch.cuda.empty_cache()\n    return model\n\n\ndef compute_otsu_threshold(prob_map, mask, fallback=0.5):\n    \"\"\"Unsupervised, per-fragment decision threshold: Otsu's method applied to\n    the predicted probability map's own histogram (tissue region only). Uses\n    no ground-truth labels, so it's a valid self-calibration for a fragment\n    whose ink-density/contrast may differ from the training fragments.\"\"\"\n    vals = prob_map[mask > 0]\n    try:\n        return float(threshold_otsu(vals))\n    except Exception:\n        return fallback\n\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_labels = load_ink_labels(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\ngt_t = torch.from_numpy(test_labels.astype(np.float32)) * torch.from_numpy(test_mask.astype(np.float32))\n\n# ---- (a) BASELINE: source-domain BN stats, val-tuned threshold ----\nprint(\"\\n[Ablation a] Baseline: no domain adaptation, val-tuned threshold ...\")\ntest_prob_baseline = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                               CFG.test_stride, CFG.device, CFG.infer_batch)\ntest_metrics_baseline = compute_metrics(torch.from_numpy(test_prob_baseline), gt_t, best_threshold)\nprint(\"  \", test_metrics_baseline)\n\n# ---- (b) + AdaBN: recalibrate BatchNorm on fragment-1 images (no labels), same val threshold ----\nprint(\"\\n[Ablation b] + BatchNorm recalibration (AdaBN), still using the val-tuned threshold ...\")\nmodel = recalibrate_batchnorm(model, test_vol, test_mask, CFG.patch_size, CFG.test_stride,\n                               CFG.device, CFG.adabn_max_patches, CFG.infer_batch)\ntest_prob_adabn = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                            CFG.test_stride, CFG.device, CFG.infer_batch)\ntest_metrics_adabn_valthresh = compute_metrics(torch.from_numpy(test_prob_adabn), gt_t, best_threshold)\nprint(\"  \", test_metrics_adabn_valthresh)\n\n# ---- (c) + Otsu self-threshold: also unsupervised, calibrated on fragment 1 itself ----\notsu_threshold = compute_otsu_threshold(test_prob_adabn, test_mask, fallback=best_threshold)\nprint(f\"\\n[Ablation c] + self-calibrated Otsu threshold ({otsu_threshold:.3f}, no labels used) ...\")\ntest_metrics_adabn_otsu = compute_metrics(torch.from_numpy(test_prob_adabn), gt_t, otsu_threshold)\nprint(\"  \", test_metrics_adabn_otsu)\n\n# final prediction used for visualization = the fully-unsupervised-adapted variant\ntest_prob = test_prob_adabn\nfinal_threshold = otsu_threshold\ntest_pred_bin = postprocess(test_prob, final_threshold)\ntest_pred_baseline_bin = postprocess(test_prob_baseline, best_threshold)\ntest_metrics = test_metrics_adabn_otsu\nprint(f\"\\nFinal TEST (fragment 1) metrics [AdaBN + Otsu, threshold={final_threshold:.3f}]:\", test_metrics)\n\nwith open(os.path.join(CFG.out_dir, \"metrics_summary.json\"), \"w\") as f:\n    json.dump({\n        \"train\": final_train_metrics,\n        \"val\": final_val_metrics,\n        \"val_tuned_threshold\": best_threshold,\n        \"test_ablation\": {\n            \"a_baseline_val_threshold\": test_metrics_baseline,\n            \"b_adabn_val_threshold\": test_metrics_adabn_valthresh,\n            \"c_adabn_otsu_threshold\": {**test_metrics_adabn_otsu, \"otsu_threshold\": otsu_threshold},\n        },\n    }, f, indent=2)\n\ntorch.cuda.empty_cache()\ngc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 9. SAVE TEST VISUALIZATIONS (input | ground truth | baseline pred | adapted pred)\n# ------------------------------------------------------------------------------\ndef save_full_overview():\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize(test_labels * 255, small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    base_small = cv2.resize((test_pred_baseline_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input (slice {mid_idx})\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground truth ink labels\")\n    axes[2].imshow(base_small, cmap=\"gray\"); axes[2].set_title(\"Baseline prediction\")\n    axes[3].imshow(pred_small, cmap=\"gray\"); axes[3].set_title(\"AdaBN + Otsu prediction\")\n    for a in axes:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_full_overview.png\")\n    plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {path}\")\n\n\ndef save_patch_comparisons(n=6):\n    ys_xs = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(ys_xs)\n    ys_xs = ys_xs[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(ys_xs):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        gt_patch = test_labels[y:y + size, x:x + size]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground truth\")\n        axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n        for a in axes:\n            a.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_comparison_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {len(ys_xs)} patch comparisons to {CFG.viz_dir}\")\n\n\nsave_full_overview()\nsave_patch_comparisons(n=6)\n\ntest_vol.close()\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"\\n=== DONE ===\")\nprint(f\"Best checkpoint: {CFG.ckpt_path}\")\nprint(f\"Tuned threshold: {best_threshold:.2f}\")\nprint(f\"Metrics summary: {os.path.join(CFG.out_dir, 'metrics_summary.json')}\")\nprint(f\"Visualizations:  {CFG.viz_dir}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n!pip install monai\n!pip install einops\n\"\"\"\n============================================================================\nVesuvius Challenge - Ink Detection - Advanced SwinUNETR Training Pipeline\n============================================================================\nTrain  : fragments 2 & 3 (80% of sampled patches)\nValidate: fragments 2 & 3 (remaining 20% of sampled patches)\nTest    : fragment 1 (full held-out fragment, sliding-window inference)\n\nDesigned to run on a single Kaggle GPU (T4/P100, ~16GB VRAM) without OOM:\n  - Patches are loaded on-the-fly from memory-mapped TIFF slices (tifffile.memmap)\n  - Only slices 16-35 (20 slices) are used, stacked as channels of a 2D image\n  - A smart LRU cache avoids re-reading recently used patches from disk\n  - Gradient checkpointing + AMP + small batch + gradient accumulation\n  - Full-fragment inference is done tile-by-tile (never loads a whole\n    fragment into RAM/VRAM at once)\n\nPaste this whole file into a single Kaggle notebook cell and run.\n============================================================================\n\"\"\"\n\n# ============================================================================\n# 0. SETUP\n# ============================================================================\nimport subprocess, sys\n\n\nimport os\nimport gc\nimport math\nimport random\nimport warnings\nfrom collections import OrderedDict\nfrom dataclasses import dataclass, field\n\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom scipy.ndimage import gaussian_filter\nfrom skimage import measure, morphology\n\nimport inspect\nfrom monai.networks.nets import SwinUNETR\nfrom monai.losses import DiceLoss, FocalLoss, TverskyLoss\n\n_SWIN_ACCEPTS_IMG_SIZE = \"img_size\" in inspect.signature(SwinUNETR.__init__).parameters\n\n\ndef _swin_kwargs(patch_size, extra):\n    \"\"\"Some monai versions require img_size, later versions removed it.\n    Only pass it through if the installed monai's SwinUNETR still accepts it.\"\"\"\n    extra = dict(extra)\n    if _SWIN_ACCEPTS_IMG_SIZE:\n        extra[\"img_size\"] = (patch_size, patch_size)\n    return extra\n\nwarnings.filterwarnings(\"ignore\")\nImage.MAX_IMAGE_PIXELS = None\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\nset_seed(42)\n\n\n# ---- AMP compatibility shim -------------------------------------------------\n# Some Kaggle images ship an older torch where GradScaler/autocast only exist\n# under torch.cuda.amp (not the newer unified torch.amp namespace). Try the\n# modern API first and fall back automatically so this script runs on either.\ndef make_grad_scaler(enabled):\n    try:\n        return torch.amp.GradScaler(\"cuda\", enabled=enabled)\n    except (AttributeError, TypeError):\n        return torch.cuda.amp.GradScaler(enabled=enabled)\n\n\ndef amp_autocast(enabled):\n    try:\n        return torch.amp.autocast(\"cuda\", enabled=enabled)\n    except (AttributeError, TypeError):\n        return torch.cuda.amp.autocast(enabled=enabled)\n\n# ============================================================================\n# 1. CONFIG\n# ============================================================================\n@dataclass\nclass CFG:\n    base_dir: str = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    work_dir: str = \"/kaggle/working\"\n\n    train_fragments: tuple = (\"2\", \"3\")\n    test_fragment: str = \"1\"\n\n    slice_start: int = 16\n    slice_end: int = 35                 # inclusive -> 20 channels\n    num_slices: int = field(init=False)\n\n    patch_size: int = 224               # must be divisible by 32 for SwinUNETR\n    train_stride: int = 112             # 50% overlap grid for candidate coords\n    val_split: float = 0.20\n    pos_ratio: float = 0.65             # fraction of patches biased around ink\n    max_patches_per_fragment: int = 3000\n    mask_coverage_thresh: float = 0.60\n\n    cache_size_per_fragment: int = 400\n\n    batch_size: int = 4\n    accum_steps: int = 4                # effective batch = 16\n    num_workers: int = 2\n    epochs: int = 30\n    warmup_epochs: int = 3\n    lr: float = 1e-4\n    weight_decay: float = 1e-5\n    grad_clip: float = 1.0\n    ema_decay: float = 0.999\n    early_stop_patience: int = 8\n\n    feature_size: int = 24              # SwinUNETR base feature size (memory-friendly)\n    drop_path_rate: float = 0.1\n    use_checkpoint: bool = True         # gradient checkpointing -> saves VRAM\n    deep_supervision: bool = True\n\n    sw_overlap: float = 0.5             # sliding window overlap for full-fragment inference\n    infer_tile_batch: int = 8\n\n    seed: int = 42\n\n    def __post_init__(self):\n        self.num_slices = self.slice_end - self.slice_start + 1\n\n\ncfg = CFG()\nos.makedirs(cfg.work_dir, exist_ok=True)\nprint(cfg)\n\n# ============================================================================\n# 2. UTILITIES: LRU cache + integral-image helpers\n# ============================================================================\nclass LRUCache:\n    \"\"\"Small in-memory LRU cache so recently-used patches are not re-read\n    from disk / re-normalized -> reduces I/O latency ('Cache ذكي').\"\"\"\n\n    def __init__(self, capacity=400):\n        self.capacity = capacity\n        self.store = OrderedDict()\n\n    def get(self, key):\n        if key in self.store:\n            self.store.move_to_end(key)\n            return self.store[key]\n        return None\n\n    def put(self, key, value):\n        self.store[key] = value\n        self.store.move_to_end(key)\n        if len(self.store) > self.capacity:\n            self.store.popitem(last=False)\n\n\ndef integral_sum(integral_img, y, x, size):\n    \"\"\"O(1) sum of a size x size box at (y, x) using a cv2.integral() image.\"\"\"\n    return (integral_img[y + size, x + size] - integral_img[y, x + size]\n            - integral_img[y + size, x] + integral_img[y, x])\n\n\ndef normalize_patch(patch, lo_pct=1.0, hi_pct=99.0, eps=1e-6):\n    \"\"\"Per-patch intensity normalization: percentile clip + min-max scale to [0,1].\"\"\"\n    lo = np.percentile(patch, lo_pct)\n    hi = np.percentile(patch, hi_pct)\n    patch = np.clip(patch, lo, hi)\n    patch = (patch - lo) / (hi - lo + eps)\n    return patch.astype(np.float32)\n\n\n# ============================================================================\n# 3. FRAGMENT VOLUME: on-the-fly memmapped patch loading\n# ============================================================================\nclass FragmentVolume:\n    \"\"\"Wraps one fragment directory. Slices are memory-mapped (tifffile.memmap)\n    so a patch is only ever read from disk for the exact (y, x, size) window\n    requested -> avoids ever loading a full fragment volume into RAM.\"\"\"\n\n    def __init__(self, frag_dir, slice_start, slice_end, cache_size=400, filter_noise=True):\n        self.frag_dir = frag_dir\n        vol_dir = os.path.join(frag_dir, \"surface_volume\")\n        idxs = list(range(slice_start, slice_end + 1))\n\n        if filter_noise:\n            idxs = self._select_slices(vol_dir, idxs)\n\n        self.slice_idxs = idxs\n        self.paths = [os.path.join(vol_dir, f\"{i:02d}.tif\") for i in self.slice_idxs]\n        self.mmaps = [self._open_slice(p) for p in self.paths]\n        self.H, self.W = self.mmaps[0].shape\n        self.cache = LRUCache(cache_size)\n\n        mask_path = os.path.join(frag_dir, \"mask.png\")\n        if os.path.exists(mask_path):\n            self.mask = (np.array(Image.open(mask_path).convert(\"L\")) > 0).astype(np.uint8)\n        else:\n            mid = np.array(self.mmaps[len(self.mmaps) // 2], dtype=np.float32)\n            self.mask = (mid > (mid.mean() * 0.1)).astype(np.uint8)\n\n        ink_path = os.path.join(frag_dir, \"inklabels.png\")\n        if os.path.exists(ink_path):\n            self.ink = (np.array(Image.open(ink_path).convert(\"L\")) > 0).astype(np.uint8)\n        else:\n            self.ink = np.zeros((self.H, self.W), dtype=np.uint8)\n\n        self.mask_integral = cv2.integral(self.mask.astype(np.float32))\n        self.ink_integral = cv2.integral(self.ink.astype(np.float32))\n\n        print(f\"[{frag_dir}] H={self.H} W={self.W} slices={self.slice_idxs} \"\n              f\"ink_px={int(self.ink.sum())} mask_px={int(self.mask.sum())}\")\n\n    @staticmethod\n    def _open_slice(path):\n        \"\"\"Memory-map a slice read-only. Falls back to a full imread (still\n        read-only, just not lazily paged) if the TIFF can't be memmapped —\n        e.g. if it happens to be compressed.\"\"\"\n        try:\n            return tifffile.memmap(path, mode=\"r\")\n        except Exception:\n            return tifffile.imread(path)\n\n    @staticmethod\n    def _select_slices(vol_dir, idxs, noise_factor=2.5, sample=512):\n        \"\"\"Removes noisy slices ('إزالة الشرائح ذات الضوضاء العالية') while keeping the\n        channel count FIXED by substituting each noisy index with its nearest\n        low-noise neighbour (so every fragment yields the same #channels).\"\"\"\n        scores = {}\n        for i in idxs:\n            p = os.path.join(vol_dir, f\"{i:02d}.tif\")\n            arr = FragmentVolume._open_slice(p)\n            h, w = arr.shape\n            cy, cx = h // 2, w // 2\n            s = min(sample, h - 1, w - 1)\n            patch = np.array(arr[cy - s // 2:cy + s // 2, cx - s // 2:cx + s // 2], dtype=np.float32)\n            lap = cv2.Laplacian(patch, cv2.CV_32F)\n            scores[i] = float(lap.std())\n\n        vals = np.array(list(scores.values()))\n        med = np.median(vals)\n        mad = np.median(np.abs(vals - med)) + 1e-6\n        good = [i for i in idxs if abs(scores[i] - med) < noise_factor * mad]\n        if len(good) < max(1, int(len(idxs) * 0.6)):\n            good = idxs  # safety: don't discard too much\n\n        final = []\n        for i in idxs:\n            if i in good:\n                final.append(i)\n            else:\n                final.append(min(good, key=lambda g: abs(g - i)))\n        return final\n\n    def mask_ratio(self, y, x, size):\n        return integral_sum(self.mask_integral, y, x, size) / float(size * size)\n\n    def ink_sum(self, y, x, size):\n        return integral_sum(self.ink_integral, y, x, size)\n\n    def get_patch(self, y, x, size):\n        key = (y, x, size)\n        cached = self.cache.get(key)\n        if cached is not None:\n            return cached.copy()\n        patch = np.stack(\n            [np.array(m[y:y + size, x:x + size], dtype=np.float32) for m in self.mmaps], axis=0\n        )\n        patch = normalize_patch(patch)\n        self.cache.put(key, patch)\n        return patch.copy()\n\n    def get_label_patch(self, y, x, size):\n        return self.ink[y:y + size, x:x + size].astype(np.float32)\n\n\n# ============================================================================\n# 4. PATCH COORDINATE SAMPLING (train/val split within fragments 2 & 3)\n# ============================================================================\ndef generate_candidate_coords(frag, patch_size, stride, mask_thresh):\n    H, W = frag.H, frag.W\n    coords = []\n    for y in range(0, H - patch_size, stride):\n        for x in range(0, W - patch_size, stride):\n            if frag.mask_ratio(y, x, patch_size) >= mask_thresh:\n                coords.append((y, x))\n    return coords\n\n\ndef split_pos_neg(frag, coords, patch_size):\n    pos, neg = [], []\n    for (y, x) in coords:\n        if frag.ink_sum(y, x, patch_size) > 0:\n            pos.append((y, x))\n        else:\n            neg.append((y, x))\n    return pos, neg\n\n\ndef build_train_val_coords(fragments, cfg, seed=42):\n    \"\"\"Returns list of (fragment_id, y, x) for train and val, sampled 80/20\n    across fragments 2 and 3 combined, biased toward ink-containing patches.\"\"\"\n    rng = random.Random(seed)\n    all_coords = []\n    for fid, frag in fragments.items():\n        cands = generate_candidate_coords(frag, cfg.patch_size, cfg.train_stride, cfg.mask_coverage_thresh)\n        pos, neg = split_pos_neg(frag, cands, cfg.patch_size)\n        rng.shuffle(pos)\n        rng.shuffle(neg)\n        n_total = min(cfg.max_patches_per_fragment, len(pos) + len(neg))\n        n_pos = min(len(pos), int(n_total * cfg.pos_ratio))\n        n_neg = min(len(neg), n_total - n_pos)\n        chosen = pos[:n_pos] + neg[:n_neg]\n        rng.shuffle(chosen)\n        all_coords.extend([(fid, y, x) for (y, x) in chosen])\n        print(f\"Fragment {fid}: {len(pos)} ink-candidates, {len(neg)} bg-candidates -> \"\n              f\"sampled {len(chosen)} ({n_pos} pos / {n_neg} neg)\")\n\n    rng.shuffle(all_coords)\n    n_val = int(len(all_coords) * cfg.val_split)\n    val_coords = all_coords[:n_val]\n    train_coords = all_coords[n_val:]\n    print(f\"Total patches: {len(all_coords)} -> train={len(train_coords)}, val={len(val_coords)}\")\n    return train_coords, val_coords\n\n\n# ============================================================================\n# 5. AUGMENTATIONS (applied identically across all channels + label)\n# ============================================================================\nAUG_P = dict(\n    flip=0.5, rot90=0.5, affine=0.35, elastic=0.25, noise=0.30,\n    gamma=0.30, clahe=0.20, brightness=0.30, blur=0.20,\n)\n\n\ndef aug_random_flip(patch, label):\n    axis = random.choice([1, 2])  # H or W\n    patch = np.flip(patch, axis=axis).copy()\n    label = np.flip(label, axis=axis - 1).copy()\n    return patch, label\n\n\ndef aug_random_rot90(patch, label):\n    k = random.choice([1, 2, 3])\n    patch = np.rot90(patch, k, axes=(1, 2)).copy()\n    label = np.rot90(label, k, axes=(0, 1)).copy()\n    return patch, label\n\n\ndef aug_random_affine(patch, label, max_angle=15, max_shift=0.06, max_scale=0.12, max_shear=8):\n    C, H, W = patch.shape\n    angle = random.uniform(-max_angle, max_angle)\n    scale = 1.0 + random.uniform(-max_scale, max_scale)\n    tx = random.uniform(-max_shift, max_shift) * W\n    ty = random.uniform(-max_shift, max_shift) * H\n    shear = math.tan(math.radians(random.uniform(-max_shear, max_shear)))\n    M = cv2.getRotationMatrix2D((W / 2, H / 2), angle, scale)\n    M[0, 1] += shear\n    M[0, 2] += tx\n    M[1, 2] += ty\n    out = np.stack([\n        cv2.warpAffine(patch[c], M, (W, H), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_REFLECT)\n        for c in range(C)\n    ], axis=0)\n    lab = cv2.warpAffine(label, M, (W, H), flags=cv2.INTER_NEAREST, borderMode=cv2.BORDER_REFLECT)\n    return out.astype(np.float32), lab.astype(np.float32)\n\n\ndef aug_elastic_deformation(patch, label, alpha=18, sigma=5):\n    C, H, W = patch.shape\n    dx = gaussian_filter((np.random.rand(H, W) * 2 - 1), sigma, mode=\"reflect\") * alpha\n    dy = gaussian_filter((np.random.rand(H, W) * 2 - 1), sigma, mode=\"reflect\") * alpha\n    xx, yy = np.meshgrid(np.arange(W), np.arange(H))\n    map_x = (xx + dx).astype(np.float32)\n    map_y = (yy + dy).astype(np.float32)\n    out = np.stack([\n        cv2.remap(patch[c], map_x, map_y, interpolation=cv2.INTER_LINEAR, borderMode=cv2.BORDER_REFLECT)\n        for c in range(C)\n    ], axis=0)\n    lab = cv2.remap(label, map_x, map_y, interpolation=cv2.INTER_NEAREST, borderMode=cv2.BORDER_REFLECT)\n    return out.astype(np.float32), lab.astype(np.float32)\n\n\ndef aug_gaussian_noise(patch, std=0.02):\n    patch = patch + np.random.normal(0, std, patch.shape).astype(np.float32)\n    return np.clip(patch, 0, 1)\n\n\ndef aug_gamma(patch, gamma_range=(0.7, 1.5)):\n    g = random.uniform(*gamma_range)\n    return np.clip(patch, 0, 1) ** g\n\n\ndef aug_clahe(patch, clip_limit=2.0, tile=8):\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=(tile, tile))\n    out = np.zeros_like(patch)\n    for c in range(patch.shape[0]):\n        img8 = (np.clip(patch[c], 0, 1) * 255).astype(np.uint8)\n        out[c] = clahe.apply(img8).astype(np.float32) / 255.0\n    return out\n\n\ndef aug_brightness(patch, delta=0.12):\n    d = random.uniform(-delta, delta)\n    return np.clip(patch + d, 0, 1)\n\n\ndef aug_blur(patch, sigma_range=(0.3, 1.2)):\n    sigma = random.uniform(*sigma_range)\n    return np.stack([gaussian_filter(patch[c], sigma=sigma) for c in range(patch.shape[0])], axis=0)\n\n\ndef apply_augmentations(patch, label):\n    if random.random() < AUG_P[\"flip\"]:\n        patch, label = aug_random_flip(patch, label)\n    if random.random() < AUG_P[\"rot90\"]:\n        patch, label = aug_random_rot90(patch, label)\n    if random.random() < AUG_P[\"affine\"]:\n        patch, label = aug_random_affine(patch, label)\n    if random.random() < AUG_P[\"elastic\"]:\n        patch, label = aug_elastic_deformation(patch, label)\n    if random.random() < AUG_P[\"noise\"]:\n        patch = aug_gaussian_noise(patch)\n    if random.random() < AUG_P[\"gamma\"]:\n        patch = aug_gamma(patch)\n    if random.random() < AUG_P[\"clahe\"]:\n        patch = aug_clahe(patch)\n    if random.random() < AUG_P[\"brightness\"]:\n        patch = aug_brightness(patch)\n    if random.random() < AUG_P[\"blur\"]:\n        patch = aug_blur(patch)\n    return patch.astype(np.float32), label.astype(np.float32)\n\n\n# batch-level mixing (applied in the training loop, after collation)\ndef cutmix_batch(x, y, p=0.3, beta=1.0):\n    if random.random() > p or x.size(0) < 2:\n        return x, y\n    lam = float(np.random.beta(beta, beta))\n    B, C, H, W = x.shape\n    idx = torch.randperm(B, device=x.device)\n    cut_w, cut_h = int(W * math.sqrt(1 - lam)), int(H * math.sqrt(1 - lam))\n    if cut_w <= 0 or cut_h <= 0:\n        return x, y\n    cx, cy = random.randint(0, W - 1), random.randint(0, H - 1)\n    x1, x2 = max(cx - cut_w // 2, 0), min(cx + cut_w // 2, W)\n    y1, y2 = max(cy - cut_h // 2, 0), min(cy + cut_h // 2, H)\n    x[:, :, y1:y2, x1:x2] = x[idx][:, :, y1:y2, x1:x2]\n    y[:, :, y1:y2, x1:x2] = y[idx][:, :, y1:y2, x1:x2]\n    return x, y\n\n\ndef mixup_batch(x, y, p=0.3, alpha=0.4):\n    if random.random() > p or x.size(0) < 2:\n        return x, y\n    lam = float(np.random.beta(alpha, alpha))\n    idx = torch.randperm(x.size(0), device=x.device)\n    x = lam * x + (1 - lam) * x[idx]\n    y = lam * y + (1 - lam) * y[idx]\n    return x, y\n\n\n# ============================================================================\n# 6. DATASET\n# ============================================================================\nclass VesuviusDataset(Dataset):\n    def __init__(self, fragments, coords, patch_size, train=True):\n        self.fragments = fragments\n        self.coords = coords\n        self.patch_size = patch_size\n        self.train = train\n\n    def __len__(self):\n        return len(self.coords)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.coords[idx]\n        frag = self.fragments[fid]\n        patch = frag.get_patch(y, x, self.patch_size)\n        label = frag.get_label_patch(y, x, self.patch_size)\n        if self.train:\n            patch, label = apply_augmentations(patch, label)\n        patch_t = torch.from_numpy(np.ascontiguousarray(patch)).float()\n        label_t = torch.from_numpy(np.ascontiguousarray(label)).float().unsqueeze(0)\n        return patch_t, label_t\n\n\n# ============================================================================\n# 7. MODEL: SwinUNETR + Attention Gates + Deep Supervision + Feature Fusion\n# ============================================================================\nclass AttentionGate2D(nn.Module):\n    \"\"\"Additive attention gate (Attention U-Net style) applied to skip\n    connections before they are fused by the decoder blocks.\"\"\"\n\n    def __init__(self, f_g, f_l, f_int):\n        super().__init__()\n        self.w_g = nn.Sequential(nn.Conv2d(f_g, f_int, 1), nn.BatchNorm2d(f_int))\n        self.w_x = nn.Sequential(nn.Conv2d(f_l, f_int, 1), nn.BatchNorm2d(f_int))\n        self.psi = nn.Sequential(nn.Conv2d(f_int, 1, 1), nn.BatchNorm2d(1), nn.Sigmoid())\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, g, x):\n        g1 = self.w_g(g)\n        x1 = self.w_x(x)\n        if g1.shape[-2:] != x1.shape[-2:]:\n            g1 = F.interpolate(g1, size=x1.shape[-2:], mode=\"bilinear\", align_corners=False)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        return x * psi\n\n\nclass SwinUNETRPlus(SwinUNETR):\n    \"\"\"monai SwinUNETR (2D, slices-as-channels) extended with:\n      - Attention gates on every skip connection\n      - Deep supervision heads at 3 decoder stages\n      - A learned feature-pyramid fusion of all stages for the final logits\n    Falls back gracefully (see build_model) if monai's internal module names\n    ever change in a future version.\n    \"\"\"\n\n    def __init__(self, in_channels, out_channels, patch_size, feature_size=24,\n                 drop_path_rate=0.1, use_checkpoint=True, deep_supervision=True):\n        base_kwargs = _swin_kwargs(patch_size, dict(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            feature_size=feature_size,\n            drop_rate=0.0,\n            attn_drop_rate=0.0,\n            dropout_path_rate=drop_path_rate,\n            use_checkpoint=use_checkpoint,\n            spatial_dims=2,\n        ))\n        super().__init__(**base_kwargs)\n        fs = feature_size\n        self.deep_supervision = deep_supervision\n        self.ag3 = AttentionGate2D(fs * 8, fs * 4, fs * 4)\n        self.ag2 = AttentionGate2D(fs * 4, fs * 2, fs * 2)\n        self.ag1 = AttentionGate2D(fs * 2, fs * 1, fs)\n        self.ag0 = AttentionGate2D(fs * 1, fs * 1, max(fs // 2, 1))\n        if deep_supervision:\n            self.ds3 = nn.Conv2d(fs * 4, out_channels, 1)\n            self.ds2 = nn.Conv2d(fs * 2, out_channels, 1)\n            self.ds1 = nn.Conv2d(fs, out_channels, 1)\n            self.fusion = nn.Conv2d(out_channels * 4, out_channels, 1)\n\n    def forward(self, x_in):\n        hidden_states_out = self.swinViT(x_in, self.normalize)\n        enc0 = self.encoder1(x_in)\n        enc1 = self.encoder2(hidden_states_out[0])\n        enc2 = self.encoder3(hidden_states_out[1])\n        enc3 = self.encoder4(hidden_states_out[2])\n        dec4 = self.encoder10(hidden_states_out[4])\n\n        dec3 = self.decoder5(dec4, hidden_states_out[3])\n        enc3_att = self.ag3(dec3, enc3)\n        dec2 = self.decoder4(dec3, enc3_att)\n        enc2_att = self.ag2(dec2, enc2)\n        dec1 = self.decoder3(dec2, enc2_att)\n        enc1_att = self.ag1(dec1, enc1)\n        dec0 = self.decoder2(dec1, enc1_att)\n        enc0_att = self.ag0(dec0, enc0)\n        out_feat = self.decoder1(dec0, enc0_att)\n        logits = self.out(out_feat)\n\n        if self.deep_supervision:\n            side3 = F.interpolate(self.ds3(dec2), size=logits.shape[-2:], mode=\"bilinear\", align_corners=False)\n            side2 = F.interpolate(self.ds2(dec1), size=logits.shape[-2:], mode=\"bilinear\", align_corners=False)\n            side1 = F.interpolate(self.ds1(dec0), size=logits.shape[-2:], mode=\"bilinear\", align_corners=False)\n            fused = self.fusion(torch.cat([logits, side3, side2, side1], dim=1))\n            return fused, [logits, side3, side2, side1]\n        return logits, []\n\n\ndef unwrap_model(model):\n    \"\"\"torch.compile() wraps the model and prefixes state_dict keys with\n    '_orig_mod.'; always save/EMA-track the underlying uncompiled module so\n    checkpoints stay loadable regardless of whether compile was used.\"\"\"\n    return model._orig_mod if hasattr(model, \"_orig_mod\") else model\n\n\ndef build_model(cfg):\n    try:\n        model = SwinUNETRPlus(\n            in_channels=cfg.num_slices, out_channels=1, patch_size=cfg.patch_size,\n            feature_size=cfg.feature_size, drop_path_rate=cfg.drop_path_rate,\n            use_checkpoint=cfg.use_checkpoint, deep_supervision=cfg.deep_supervision,\n        )\n        print(\"Built SwinUNETRPlus (attention gates + deep supervision + FPN fusion).\")\n    except Exception as e:\n        print(f\"[WARN] SwinUNETRPlus failed ({e!r}); falling back to plain SwinUNETR.\")\n        fallback_kwargs = _swin_kwargs(cfg.patch_size, dict(\n            in_channels=cfg.num_slices, out_channels=1, feature_size=cfg.feature_size,\n            dropout_path_rate=cfg.drop_path_rate, use_checkpoint=cfg.use_checkpoint,\n            spatial_dims=2,\n        ))\n        base = SwinUNETR(**fallback_kwargs)\n\n        class Wrapped(nn.Module):\n            def __init__(self, m):\n                super().__init__()\n                self.m = m\n\n            def forward(self, x):\n                return self.m(x), []\n\n        model = Wrapped(base)\n    return model.to(DEVICE)\n\n\n# ============================================================================\n# 8. LOSS: 0.35 Dice + 0.30 Focal + 0.20 Tversky + 0.15 Lovasz\n# ============================================================================\ndef _lovasz_grad(gt_sorted):\n    p = len(gt_sorted)\n    gts = gt_sorted.sum()\n    intersection = gts - gt_sorted.float().cumsum(0)\n    union = gts + (1 - gt_sorted).float().cumsum(0)\n    jaccard = 1.0 - intersection / union\n    if p > 1:\n        jaccard[1:p] = jaccard[1:p] - jaccard[0:-1]\n    return jaccard\n\n\ndef _lovasz_hinge_flat(logits, labels):\n    if labels.numel() == 0:\n        return logits.sum() * 0.0\n    signs = 2.0 * labels.float() - 1.0\n    errors = 1.0 - logits * signs\n    errors_sorted, perm = torch.sort(errors, dim=0, descending=True)\n    gt_sorted = labels[perm]\n    grad = _lovasz_grad(gt_sorted)\n    return torch.dot(F.relu(errors_sorted), grad)\n\n\nclass LovaszLoss(nn.Module):\n    def forward(self, logits, labels):\n        return _lovasz_hinge_flat(logits.reshape(-1), labels.reshape(-1))\n\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, w_dice=0.35, w_focal=0.30, w_tversky=0.20, w_lovasz=0.15):\n        super().__init__()\n        self.dice = DiceLoss(sigmoid=True)\n        self.focal = FocalLoss(gamma=2.0)\n        self.tversky = TverskyLoss(sigmoid=True, alpha=0.3, beta=0.7)\n        self.lovasz = LovaszLoss()\n        self.w = (w_dice, w_focal, w_tversky, w_lovasz)\n\n    def forward(self, logits, targets):\n        d = self.dice(logits, targets)\n        f = self.focal(logits, targets)\n        t = self.tversky(logits, targets)\n        l = self.lovasz(logits, targets)\n        total = self.w[0] * d + self.w[1] * f + self.w[2] * t + self.w[3] * l\n        parts = {\"dice\": d.item(), \"focal\": f.item(), \"tversky\": t.item(), \"lovasz\": l.item()}\n        return total, parts\n\n\n# ============================================================================\n# 9. EMA (Exponential Moving Average of weights)\n# ============================================================================\nclass EMA:\n    def __init__(self, model, decay=0.999):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone() for k, v in unwrap_model(model).state_dict().items()}\n\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in unwrap_model(model).state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(self.decay).add_(v.detach(), alpha=1 - self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n\n    def copy_to(self, model):\n        unwrap_model(model).load_state_dict(self.shadow, strict=True)\n\n\n# ============================================================================\n# 10. METRICS\n# ============================================================================\n@torch.no_grad()\ndef dice_iou_fbeta(probs, targets, thresh=0.5, beta=0.5, eps=1e-7):\n    preds = (probs > thresh).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    fbeta = ((1 + beta ** 2) * precision * recall + eps) / (beta ** 2 * precision + recall + eps)\n    return dice.item(), iou.item(), fbeta.item()\n\n\n# ============================================================================\n# 11. SCHEDULER: linear warmup -> Cosine Annealing Warm Restarts\n# ============================================================================\nclass WarmupCosineWarmRestarts:\n    def __init__(self, optimizer, warmup_epochs, steps_per_epoch, T_0_epochs=8, T_mult=2, base_lr=1e-4):\n        self.optimizer = optimizer\n        self.warmup_steps = warmup_epochs * steps_per_epoch\n        self.base_lr = base_lr\n        self.cosine = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n            optimizer, T_0=T_0_epochs * steps_per_epoch, T_mult=T_mult\n        )\n        self.step_num = 0\n\n    def step(self):\n        self.step_num += 1\n        if self.step_num <= self.warmup_steps:\n            lr = self.base_lr * self.step_num / max(1, self.warmup_steps)\n            for pg in self.optimizer.param_groups:\n                pg[\"lr\"] = lr\n        else:\n            self.cosine.step()\n\n    def get_last_lr(self):\n        return [pg[\"lr\"] for pg in self.optimizer.param_groups]\n\n\n# ============================================================================\n# 12. TRAINING LOOP\n# ============================================================================\ndef train_model(cfg, fragments):\n    train_coords, val_coords = build_train_val_coords(fragments, cfg, seed=cfg.seed)\n\n    train_ds = VesuviusDataset(fragments, train_coords, cfg.patch_size, train=True)\n    val_ds = VesuviusDataset(fragments, val_coords, cfg.patch_size, train=False)\n\n    train_loader = DataLoader(\n        train_ds, batch_size=cfg.batch_size, shuffle=True, num_workers=cfg.num_workers,\n        pin_memory=True, drop_last=True, persistent_workers=(cfg.num_workers > 0),\n        prefetch_factor=2 if cfg.num_workers > 0 else None,\n    )\n    val_loader = DataLoader(\n        val_ds, batch_size=cfg.batch_size, shuffle=False, num_workers=cfg.num_workers,\n        pin_memory=True, persistent_workers=(cfg.num_workers > 0),\n        prefetch_factor=2 if cfg.num_workers > 0 else None,\n    )\n\n    model = build_model(cfg)\n    try:\n        model = torch.compile(model)\n        print(\"torch.compile enabled.\")\n    except Exception as e:\n        print(f\"torch.compile unavailable ({e}); continuing without it.\")\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    scheduler = WarmupCosineWarmRestarts(optimizer, cfg.warmup_epochs, len(train_loader),\n                                         T_0_epochs=8, T_mult=2, base_lr=cfg.lr)\n    scaler = make_grad_scaler(enabled=(DEVICE.type == \"cuda\"))\n    loss_fn = CombinedLoss()\n    ema = EMA(model, decay=cfg.ema_decay)\n\n    best_val_dice = -1.0\n    epochs_no_improve = 0\n    ckpt_path = os.path.join(cfg.work_dir, \"best_model.pt\")\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": [], \"val_iou\": [], \"val_fbeta\": []}\n\n    ds_weights = [1.0, 0.5, 0.3, 0.2]  # main logits + 3 side outputs\n\n    for epoch in range(cfg.epochs):\n        model.train()\n        running_loss = 0.0\n        optimizer.zero_grad(set_to_none=True)\n\n        for step, (x, y) in enumerate(train_loader):\n            x, y = x.to(DEVICE, non_blocking=True), y.to(DEVICE, non_blocking=True)\n            x, y = cutmix_batch(x, y, p=0.20)\n            x, y = mixup_batch(x, y, p=0.20)\n\n            with amp_autocast(enabled=(DEVICE.type == \"cuda\")):\n                out, sides = model(x)\n                loss, parts = loss_fn(out, y)\n                if sides:\n                    for w, s in zip(ds_weights[1:], sides[1:]):\n                        l_side, _ = loss_fn(s, y)\n                        loss = loss + w * l_side\n                loss = loss / cfg.accum_steps\n\n            scaler.scale(loss).backward()\n\n            if (step + 1) % cfg.accum_steps == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad(set_to_none=True)\n                ema.update(model)\n                scheduler.step()\n\n            running_loss += loss.item() * cfg.accum_steps\n\n        train_loss = running_loss / max(1, len(train_loader))\n\n        # ---- validation ----\n        model.eval()\n        val_loss, val_dice, val_iou, val_fbeta, n_val = 0.0, 0.0, 0.0, 0.0, 0\n        with torch.no_grad():\n            for x, y in val_loader:\n                x, y = x.to(DEVICE), y.to(DEVICE)\n                with amp_autocast(enabled=(DEVICE.type == \"cuda\")):\n                    out, _ = model(x)\n                    loss, _ = loss_fn(out, y)\n                probs = torch.sigmoid(out)\n                d, i, fb = dice_iou_fbeta(probs, y)\n                val_loss += loss.item()\n                val_dice += d\n                val_iou += i\n                val_fbeta += fb\n                n_val += 1\n        val_loss /= max(1, n_val)\n        val_dice /= max(1, n_val)\n        val_iou /= max(1, n_val)\n        val_fbeta /= max(1, n_val)\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_dice\"].append(val_dice)\n        history[\"val_iou\"].append(val_iou)\n        history[\"val_fbeta\"].append(val_fbeta)\n\n        print(f\"Epoch {epoch+1:03d}/{cfg.epochs} | lr={scheduler.get_last_lr()[0]:.2e} \"\n              f\"| train_loss={train_loss:.4f} | val_loss={val_loss:.4f} \"\n              f\"| val_dice={val_dice:.4f} | val_iou={val_iou:.4f} | val_f0.5={val_fbeta:.4f}\")\n\n        if val_dice > best_val_dice:\n            best_val_dice = val_dice\n            epochs_no_improve = 0\n            torch.save({\"model\": unwrap_model(model).state_dict(), \"ema\": ema.shadow}, ckpt_path)\n            print(f\"  -> New best model saved (val_dice={best_val_dice:.4f}).\")\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= cfg.early_stop_patience:\n                print(f\"Early stopping at epoch {epoch+1} (no improvement for \"\n                      f\"{cfg.early_stop_patience} epochs).\")\n                break\n\n        gc.collect()\n        if DEVICE.type == \"cuda\":\n            torch.cuda.empty_cache()\n\n    return model, ema, ckpt_path, val_loader, history\n\n\n# ============================================================================\n# 13. FULL-FRAGMENT SLIDING-WINDOW INFERENCE (memory-safe, tile by tile)\n# ============================================================================\ndef _gaussian_weight(size, sigma_scale=0.125):\n    sigma = size * sigma_scale\n    ax = np.arange(size) - size / 2.0\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    g2d = np.outer(g1d, g1d)\n    return (g2d / g2d.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_predict(model, frag, patch_size, overlap=0.5, batch_size=8):\n    \"\"\"Predicts the full-resolution ink-probability map for a fragment without\n    ever holding the whole volume in memory: each tile is read on-the-fly from\n    the memory-mapped TIFFs, and only two (H, W) float32 accumulators\n    (prediction sum + weight sum) are kept in RAM.\"\"\"\n    H, W = frag.H, frag.W\n    stride = max(1, int(patch_size * (1 - overlap)))\n    weight = _gaussian_weight(patch_size)\n\n    pred_acc = np.zeros((H, W), dtype=np.float32)\n    weight_acc = np.zeros((H, W), dtype=np.float32)\n\n    ys = list(range(0, H - patch_size, stride)) + [H - patch_size]\n    xs = list(range(0, W - patch_size, stride)) + [W - patch_size]\n    ys = sorted(set(y for y in ys if y >= 0))\n    xs = sorted(set(x for x in xs if x >= 0))\n\n    coords = [(y, x) for y in ys for x in xs if frag.mask_ratio(y, x, patch_size) > 0.05]\n\n    model.eval()\n    batch_patches, batch_coords = [], []\n\n    def flush():\n        if not batch_patches:\n            return\n        x_t = torch.from_numpy(np.stack(batch_patches, axis=0)).float().to(DEVICE)\n        with amp_autocast(enabled=(DEVICE.type == \"cuda\")):\n            out, _ = model(x_t)\n            probs = torch.sigmoid(out).float().cpu().numpy()[:, 0]\n        for prob, (yy, xx) in zip(probs, batch_coords):\n            pred_acc[yy:yy + patch_size, xx:xx + patch_size] += prob * weight\n            weight_acc[yy:yy + patch_size, xx:xx + patch_size] += weight\n        batch_patches.clear()\n        batch_coords.clear()\n\n    for (y, x) in coords:\n        patch = frag.get_patch(y, x, patch_size)\n        batch_patches.append(patch)\n        batch_coords.append((y, x))\n        if len(batch_patches) >= batch_size:\n            flush()\n    flush()\n\n    weight_acc[weight_acc == 0] = 1.0\n    prob_map = pred_acc / weight_acc\n    return prob_map\n\n\n# ============================================================================\n# 14. POST-PROCESSING (no TTA, as requested)\n# ============================================================================\ndef postprocess_mask(prob_map, threshold, gaussian_sigma=1.0, min_object_size=64, closing_size=3):\n    smoothed = gaussian_filter(prob_map, sigma=gaussian_sigma)\n    binary = smoothed > threshold\n    struct = morphology.disk(closing_size)\n    closed = morphology.binary_closing(binary, footprint=struct)\n    labeled = measure.label(closed)\n    cleaned = morphology.remove_small_objects(labeled, min_size=min_object_size)\n    return (cleaned > 0).astype(np.uint8)\n\n\ndef compute_metrics_np(pred, gt, beta=0.5, eps=1e-7):\n    pred = pred.astype(np.float32)\n    gt = gt.astype(np.float32)\n    tp = (pred * gt).sum()\n    fp = (pred * (1 - gt)).sum()\n    fn = ((1 - pred) * gt).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    fbeta = ((1 + beta ** 2) * precision * recall + eps) / (beta ** 2 * precision + recall + eps)\n    return {\"dice\": float(dice), \"iou\": float(iou), \"fbeta\": float(fbeta),\n            \"precision\": float(precision), \"recall\": float(recall)}\n\n\n# ============================================================================\n# 15. THRESHOLD SEARCH ON VALIDATION SET\n# ============================================================================\n@torch.no_grad()\ndef search_best_threshold(model, val_loader, thresholds=np.arange(0.30, 0.71, 0.02)):\n    model.eval()\n    all_probs, all_targets = [], []\n    for x, y in val_loader:\n        x = x.to(DEVICE)\n        with amp_autocast(enabled=(DEVICE.type == \"cuda\")):\n            out, _ = model(x)\n            probs = torch.sigmoid(out).float().cpu()\n        all_probs.append(probs)\n        all_targets.append(y)\n    probs = torch.cat(all_probs, dim=0)\n    targets = torch.cat(all_targets, dim=0)\n\n    best_t, best_iou, best_stats = 0.5, -1, None\n    for t in thresholds:\n        d, i, fb = dice_iou_fbeta(probs, targets, thresh=t)\n        if i > best_iou:\n            best_iou, best_t = i, t\n            best_stats = {\"dice\": d, \"iou\": i, \"fbeta\": fb}\n    print(f\"Best threshold on validation set: {best_t:.2f} -> {best_stats}\")\n    return best_t, best_stats\n\n\n# ============================================================================\n# 16. VISUALIZATION\n# ============================================================================\ndef save_comparison_figure(frag, prob_map, pred_mask, out_path, mid_slice_idx=None):\n    if mid_slice_idx is None:\n        mid_slice_idx = len(frag.mmaps) // 2\n    input_slice = np.array(frag.mmaps[mid_slice_idx], dtype=np.float32)\n    input_slice = normalize_patch(input_slice)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    axes[0].imshow(input_slice, cmap=\"gray\")\n    axes[0].set_title(f\"Input (slice {frag.slice_idxs[mid_slice_idx]})\")\n    axes[1].imshow(frag.ink, cmap=\"gray\")\n    axes[1].set_title(\"Ground Truth Ink\")\n    axes[2].imshow(prob_map, cmap=\"viridis\", vmin=0, vmax=1)\n    axes[2].set_title(\"Predicted Probability\")\n    axes[3].imshow(pred_mask, cmap=\"gray\")\n    axes[3].set_title(\"Predicted Mask (post-processed)\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(out_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"Saved visualization to {out_path}\")\n\n\n# ============================================================================\n# 17. MAIN\n# ============================================================================\ndef main():\n    print(\"Loading training fragments (2, 3) as memory-mapped volumes ...\")\n    train_fragments = {\n        fid: FragmentVolume(\n            os.path.join(cfg.base_dir, fid), cfg.slice_start, cfg.slice_end,\n            cache_size=cfg.cache_size_per_fragment, filter_noise=True,\n        )\n        for fid in cfg.train_fragments\n    }\n\n    print(\"\\n=== Training ===\")\n    model, ema, ckpt_path, val_loader, history = train_model(cfg, train_fragments)\n\n    print(\"\\n=== Loading best checkpoint (EMA weights) for evaluation ===\")\n    ckpt = torch.load(ckpt_path, map_location=DEVICE, weights_only=False)\n    eval_model = build_model(cfg)\n    eval_model.load_state_dict(ckpt[\"ema\"], strict=True)\n    eval_model.eval()\n\n    print(\"\\n=== Threshold search on validation set ===\")\n    best_thresh, val_stats = search_best_threshold(eval_model, val_loader)\n\n    print(\"\\n=== Final performance on TRAIN / VAL (patch-level) ===\")\n    print(\"Validation (patch-level, at best threshold):\", val_stats)\n\n    print(\"\\n=== Loading held-out TEST fragment (1) ===\")\n    test_frag = FragmentVolume(\n        os.path.join(cfg.base_dir, cfg.test_fragment), cfg.slice_start, cfg.slice_end,\n        cache_size=cfg.cache_size_per_fragment, filter_noise=True,\n    )\n\n    print(\"\\n=== Sliding-window inference on full test fragment (no TTA) ===\")\n    prob_map = sliding_window_predict(\n        eval_model, test_frag, cfg.patch_size, overlap=cfg.sw_overlap, batch_size=cfg.infer_tile_batch\n    )\n\n    print(\"\\n=== Post-processing test prediction ===\")\n    pred_mask = postprocess_mask(prob_map, threshold=best_thresh)\n\n    test_metrics = compute_metrics_np(pred_mask, test_frag.ink)\n    print(\"TEST (fragment 1) metrics:\", test_metrics)\n\n    np.save(os.path.join(cfg.work_dir, \"test_fragment1_prob_map.npy\"), prob_map)\n    np.save(os.path.join(cfg.work_dir, \"test_fragment1_pred_mask.npy\"), pred_mask)\n\n    fig_path = os.path.join(cfg.work_dir, \"test_fragment1_comparison.png\")\n    save_comparison_figure(test_frag, prob_map, pred_mask, fig_path)\n\n    # ---- training curves ----\n    fig, ax = plt.subplots(1, 2, figsize=(14, 5))\n    ax[0].plot(history[\"train_loss\"], label=\"train_loss\")\n    ax[0].plot(history[\"val_loss\"], label=\"val_loss\")\n    ax[0].set_title(\"Loss\")\n    ax[0].legend()\n    ax[1].plot(history[\"val_dice\"], label=\"val_dice\")\n    ax[1].plot(history[\"val_iou\"], label=\"val_iou\")\n    ax[1].plot(history[\"val_fbeta\"], label=\"val_f0.5\")\n    ax[1].set_title(\"Validation metrics\")\n    ax[1].legend()\n    plt.tight_layout()\n    curves_path = os.path.join(cfg.work_dir, \"training_curves.png\")\n    plt.savefig(curves_path, dpi=150)\n    plt.close(fig)\n    print(f\"Saved training curves to {curves_path}\")\n\n    print(\"\\n=== DONE ===\")\n    print(f\"Best model checkpoint : {ckpt_path}\")\n    print(f\"Best threshold         : {best_thresh}\")\n    print(f\"Validation metrics     : {val_stats}\")\n    print(f\"Test (fragment 1)      : {test_metrics}\")\n\n    # ---------------------------------------------------------------\n    # OPTIONAL ENSEMBLE (item 10): if you also train UNETR / SegResNet\n    # checkpoints separately, combine their probability maps like this:\n    #\n    #   from monai.networks.nets import UNETR, SegResNet\n    #   prob_swin = sliding_window_predict(swin_model, test_frag, ...)\n    #   prob_unetr = sliding_window_predict(unetr_model, test_frag, ...)\n    #   prob_segres = sliding_window_predict(segres_model, test_frag, ...)\n    #   ensemble_prob = 0.5*prob_swin + 0.3*prob_unetr + 0.2*prob_segres\n    #   ensemble_mask = postprocess_mask(ensemble_prob, best_thresh)\n    #\n    # This is left as an opt-in step since training 3 full models triples\n    # the compute/time budget of a single Kaggle session.\n    # ---------------------------------------------------------------\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# VESUVIUS CHALLENGE — INK DETECTION PIPELINE (v2, accuracy-focused)\n# Pretrained ResNet34 encoder (ImageNet) + scSE-attention U-Net decoder,\n# centered depth-slice subset, per-patch normalization, validation-tuned\n# decision threshold, early stopping.\n# Train: fragments 2 & 3 (80/20 split)  |  Test: fragment 1\n# OOM-safety: disk-backed memmap volumes, patch-based training, sliding-window\n# (non-TTA) inference, AMP, aggressive gc.\n#\n# NOTE: this script needs internet access enabled in the Kaggle notebook\n# settings, both to `pip install segmentation-models-pytorch` and to download\n# the ImageNet-pretrained ResNet34 weights.\n#\n# Paste this whole file into a single Kaggle notebook cell and run.\n# ==============================================================================\n!pip install segmentation-models-pytorch==0.2.0\nimport os, gc, random, time, json\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    os.system(\"pip install -q albumentations\")\n    import albumentations as A\n\ntry:\n    import segmentation_models_pytorch as smp\nexcept ImportError:\n    os.system(\"pip install -q segmentation-models-pytorch\")\n    import segmentation_models_pytorch as smp\n\n\n# ------------------------------------------------------------------------------\n# 0. CONFIG\n# ------------------------------------------------------------------------------\nclass CFG:\n    base_dir      = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags   = [\"2\", \"3\"]\n    test_frag     = \"1\"\n\n    # Use a centered subset of the 65 slices: the ink signal concentrates\n    # mid-depth, and dropping the noisier outer slices both improves signal\n    # and reduces compute/memory versus all 65 channels.\n    depth_indices = list(range(16, 38))     # 32 slices, centered on 31.5\n    in_channels   = len(depth_indices)\n\n    patch_size    = 256\n    train_stride  = 96              # dense overlap -> more training samples\n    test_stride   = 96\n\n    val_fraction          = 0.20\n    min_tissue_frac_train = 0.10    # skip near-empty patches when building train grid\n    min_tissue_frac_test  = 0.02\n\n    batch_size    = 8\n    infer_batch   = 12\n    num_workers   = 2\n\n    epochs        = 30\n    early_stop_patience = 8\n    lr            = 2e-4\n    weight_decay  = 1e-4\n\n    encoder_name    = \"resnet34\"\n    encoder_weights = \"imagenet\"\n    decoder_attention_type = \"scse\"   # squeeze-and-excitation attention decoder\n\n    threshold      = 0.5            # only used for in-training monitoring\n    seed           = 42\n\n    out_dir        = \"/kaggle/working\"\n    ckpt_path      = os.path.join(out_dir, \"vesuviusnet_best.pth\")\n    viz_dir        = os.path.join(out_dir, \"test_visualizations\")\n\n    device         = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\n\nset_seed(CFG.seed)\ntorch.backends.cudnn.benchmark = True\n\n\ndef estimate_positive_fraction(labels_full, samples, patch_size, n_samples=60):\n    \"\"\"Quickly estimate the fraction of ink-positive pixels across a random\n    subset of training patches, used to (a) init the output bias so the model\n    starts near the true prior instead of 50/50, and (b) weight BCE.\"\"\"\n    sub = random.sample(samples, min(n_samples, len(samples)))\n    total, pos = 0, 0\n    for fid, y, x in sub:\n        patch = labels_full[fid][y:y + patch_size, x:x + patch_size]\n        pos += patch.sum()\n        total += patch.size\n    frac = pos / max(total, 1)\n    return max(frac, 1e-4)\n\n\ndef cfg_to_dict(cfg_cls):\n    \"\"\"vars() on a *class* returns a non-picklable mappingproxy, so build a\n    plain dict of the simple (picklable) config values by hand instead.\"\"\"\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ------------------------------------------------------------------------------\n# 1. PRE-PROCESSING HELPERS: tissue mask, patch grid, disk-backed volume reader\n# ------------------------------------------------------------------------------\ndef load_tissue_mask(frag_dir):\n    \"\"\"Load fragment tissue mask (mask.png if present, else derive from a mid slice).\"\"\"\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n        mid = tifffile.imread(os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n        thresh = mid.mean() * 0.15\n        mask = ((mid > thresh).astype(np.uint8)) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    for y in range(0, max(H - patch_size, 0) + 1, stride):\n        for x in range(0, max(W - patch_size, 0) + 1, stride):\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access to a fragment's surface volume, restricted\n    to a chosen subset of depth slices. Avoids loading the full multi-GB\n    volume into RAM.\"\"\"\n\n    def __init__(self, frag_dir, depth_indices):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in depth_indices]\n        self.depth_indices = depth_indices\n        self._slices = None\n        self._h, self._w = None, None\n\n    def _ensure_open(self):\n        if self._slices is None:\n            slices = []\n            for p in self.paths:\n                try:\n                    arr = tifffile.memmap(p, mode=\"r\")\n                except Exception:\n                    # fallback if the tif can't be memory-mapped (e.g. compressed)\n                    arr = tifffile.imread(p)\n                slices.append(arr)\n            self._slices = slices\n            self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        gc.collect()\n\n\ndef normalize_patch(img_float):\n    \"\"\"Per-patch (instance) normalization: zero-mean, unit-variance across the\n    whole patch. This is applied identically at train and inference time and\n    helps offset scanner/intensity differences between fragments (a likely\n    contributor to the train/test domain gap).\"\"\"\n    mean = img_float.mean()\n    std = img_float.std() + 1e-6\n    return (img_float - mean) / std\n\n\n# ------------------------------------------------------------------------------\n# 2. AUGMENTATION (train only) — kept to spatial-safe, version-stable transforms\n# ------------------------------------------------------------------------------\ndef build_train_transform(size):\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.15, rotate_limit=25,\n                            border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.25),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.3),\n    ])\n\n\n# ------------------------------------------------------------------------------\n# 3. DATASET\n# ------------------------------------------------------------------------------\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None):\n        \"\"\"\n        volumes: dict frag_id -> FragmentVolume\n        labels:  dict frag_id -> full-res (H,W) uint8 ink label array\n        samples: list of (frag_id, y, x)\n        \"\"\"\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        frag_id, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[frag_id]\n        patch = vol.read_patch(y, x, size)                # (C,H,W) uint8\n        label = self.labels[frag_id][y:y + size, x:x + size]\n\n        img = np.transpose(patch, (1, 2, 0))               # HWC for albumentations\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = normalize_patch(img)\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))  # CHW\n        label = (label > 0).astype(np.float32)[None, ...]         # 1HW\n\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ------------------------------------------------------------------------------\n# 4. MODEL — pretrained ResNet34 encoder + scSE-attention U-Net decoder\n# ------------------------------------------------------------------------------\ndef build_model():\n    model = smp.Unet(\n        encoder_name=CFG.encoder_name,\n        encoder_weights=CFG.encoder_weights,\n        in_channels=CFG.in_channels,\n        classes=1,\n        decoder_attention_type=CFG.decoder_attention_type,\n    )\n    return model\n\n\n# ------------------------------------------------------------------------------\n# 5. LOSS & METRICS\n# ------------------------------------------------------------------------------\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(1)\n    union = probs.sum(1) + t.sum(1)\n    return 1 - ((2 * inter + eps) / (union + eps)).mean()\n\n\nclass ComboLoss(nn.Module):\n    def __init__(self, bce_w=0.5, dice_w=0.5, pos_weight=None):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n        self.bce_w, self.dice_w = bce_w, dice_w\n\n    def forward(self, logits, targets):\n        return self.bce_w * self.bce(logits, targets) + self.dice_w * dice_loss(logits, targets)\n\n\n@torch.no_grad()\ndef compute_metrics(probs, targets, threshold=0.5, eps=1e-6):\n    preds = (probs > threshold).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    return {\"dice\": dice.item(), \"iou\": iou.item(), \"precision\": precision.item(),\n            \"recall\": recall.item(), \"fbeta0.5\": fbeta.item()}\n\n\nclass GlobalConfusionAccumulator:\n    \"\"\"Accumulates TP/FP/FN across an entire epoch (rather than averaging\n    per-batch Dice) so metrics aren't swamped by noise from the many\n    near-empty patches typical of ink-detection data.\"\"\"\n\n    def __init__(self):\n        self.tp = self.fp = self.fn = 0.0\n\n    def update(self, probs, targets, threshold):\n        preds = (probs > threshold).float()\n        self.tp += (preds * targets).sum().item()\n        self.fp += (preds * (1 - targets)).sum().item()\n        self.fn += ((1 - preds) * targets).sum().item()\n\n    def compute(self, eps=1e-6, beta2=0.25):\n        dice = (2 * self.tp + eps) / (2 * self.tp + self.fp + self.fn + eps)\n        iou = (self.tp + eps) / (self.tp + self.fp + self.fn + eps)\n        precision = (self.tp + eps) / (self.tp + self.fp + eps)\n        recall = (self.tp + eps) / (self.tp + self.fn + eps)\n        fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n        return {\"dice\": dice, \"iou\": iou, \"precision\": precision,\n                \"recall\": recall, \"fbeta0.5\": fbeta}\n\n\n# ------------------------------------------------------------------------------\n# 6. BUILD DATA (fragments 2 & 3 -> train/val ; fragment 1 -> test)\n# ------------------------------------------------------------------------------\nprint(\"Building tissue masks, label maps and patch grids ...\")\n\ntrain_volumes, train_labels_full = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    vol = FragmentVolume(frag_dir, CFG.depth_indices)\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"  fragment {fid}: mask {mask.shape}, {len(coords)} candidate patches\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    all_samples.extend([(fid, y, x) for (y, x) in coords])\n\n    del mask\n    gc.collect()\n\nrandom.shuffle(all_samples)\nn_val = int(len(all_samples) * CFG.val_fraction)\nval_samples = all_samples[:n_val]\ntrain_samples = all_samples[n_val:]\nprint(f\"Total patches: {len(all_samples)}  -> train {len(train_samples)} / val {len(val_samples)}\")\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples,\n                            CFG.patch_size, transform=build_train_transform(CFG.patch_size))\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples,\n                          CFG.patch_size, transform=None)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=True, persistent_workers=CFG.num_workers > 0)\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         persistent_workers=CFG.num_workers > 0)\n\n\n# ------------------------------------------------------------------------------\n# 7. TRAIN / VALIDATE\n# ------------------------------------------------------------------------------\nprint(\"Estimating ink-pixel prior for bias-init / class weighting ...\")\npos_frac = estimate_positive_fraction(train_labels_full, train_samples, CFG.patch_size)\nprint(f\"  estimated positive-pixel fraction: {pos_frac:.5f}\")\n\nmodel = build_model().to(CFG.device)\n\n# Bias-init trick: start the output layer predicting ~pos_frac everywhere\n# instead of ~0.5, so the model doesn't have to unlearn a bad 50/50 prior.\nwith torch.no_grad():\n    bias_val = float(np.log(pos_frac / (1 - pos_frac)))\n    model.segmentation_head[0].bias.fill_(bias_val)\nprint(f\"  output layer bias initialized to {bias_val:.3f} (sigmoid={pos_frac:.4f})\")\n\npos_weight_val = float(np.clip((1 - pos_frac) / pos_frac, 1.0, 15.0))\npos_weight = torch.tensor([pos_weight_val], device=CFG.device)\nprint(f\"  BCE pos_weight = {pos_weight_val:.2f}\")\n\ncriterion = ComboLoss(pos_weight=pos_weight)\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nbest_val_dice = -1.0\nepochs_no_improve = 0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n\ndef run_epoch(loader, train_mode, threshold=CFG.threshold):\n    model.train(train_mode)\n    total_loss = 0.0\n    global_acc = GlobalConfusionAccumulator()\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                logits = model(imgs)\n                loss = criterion(logits, masks)\n\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                scaler.step(optimizer)\n                scaler.update()\n\n        probs = torch.sigmoid(logits.detach())\n        global_acc.update(probs, masks, threshold)\n        total_loss += loss.item()\n\n        del imgs, masks, logits, probs\n    torch.cuda.empty_cache()\n    return total_loss / max(len(loader), 1), global_acc.compute()\n\n\nprint(\"\\nStarting training ...\")\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_global = run_epoch(train_loader, train_mode=True)\n    val_loss, val_global = run_epoch(val_loader, train_mode=False)\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_global[\"dice\"])\n\n    print(f\"[{epoch:02d}/{CFG.epochs}] \"\n          f\"train_loss={train_loss:.4f} dice={train_global['dice']:.4f} | \"\n          f\"val_loss={val_loss:.4f} dice={val_global['dice']:.4f} \"\n          f\"fbeta0.5={val_global['fbeta0.5']:.4f} recall={val_global['recall']:.4f} \"\n          f\"precision={val_global['precision']:.4f} ({time.time()-t0:.1f}s)\")\n\n    if val_global[\"dice\"] > best_val_dice:\n        best_val_dice = val_global[\"dice\"]\n        epochs_no_improve = 0\n        torch.save({\"model\": model.state_dict(), \"cfg\": cfg_to_dict(CFG)}, CFG.ckpt_path)\n        print(f\"  -> saved new best checkpoint (val_dice={best_val_dice:.4f})\")\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= CFG.early_stop_patience:\n            print(f\"  -> no val improvement for {CFG.early_stop_patience} epochs, stopping early.\")\n            break\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# reload best checkpoint\nmodel.load_state_dict(torch.load(CFG.ckpt_path, map_location=CFG.device)[\"model\"])\n\n\n# ------------------------------------------------------------------------------\n# 7b. TUNE DECISION THRESHOLD ON VALIDATION SET (maximize F0.5)\n# ------------------------------------------------------------------------------\n@torch.no_grad()\ndef find_best_threshold(model, loader, thresholds=np.arange(0.10, 0.91, 0.05)):\n    model.eval()\n    tp = np.zeros(len(thresholds)); fp = np.zeros(len(thresholds)); fn = np.zeros(len(thresholds))\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n        with autocast(enabled=(CFG.device == \"cuda\")):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        targets = masks.cpu().numpy()\n        for i, t in enumerate(thresholds):\n            preds = (probs > t).astype(np.float32)\n            tp[i] += (preds * targets).sum()\n            fp[i] += (preds * (1 - targets)).sum()\n            fn[i] += ((1 - preds) * targets).sum()\n        del imgs, masks, logits, probs\n    eps, beta2 = 1e-6, 0.25\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    best_idx = int(np.argmax(fbeta))\n    return float(thresholds[best_idx]), float(fbeta[best_idx])\n\n\nprint(\"\\nTuning decision threshold on validation set ...\")\nbest_threshold, best_val_fbeta = find_best_threshold(model, val_loader)\nprint(f\"  best threshold = {best_threshold:.2f} (val fbeta0.5 = {best_val_fbeta:.4f})\")\n\n# final train / val metrics at the tuned threshold\n_, final_train_metrics = run_epoch(train_loader, train_mode=False, threshold=best_threshold)\n_, final_val_metrics = run_epoch(val_loader, train_mode=False, threshold=best_threshold)\nprint(\"\\nFinal TRAIN metrics:\", final_train_metrics)\nprint(\"Final VAL metrics:  \", final_val_metrics)\n\n# free fragment 2/3 volumes before loading fragment 1\nfor v in train_volumes.values():\n    v.close()\ndel train_loader, val_loader, train_ds, val_ds, train_volumes, train_labels_full\ngc.collect()\ntorch.cuda.empty_cache()\n\n\n# ------------------------------------------------------------------------------\n# 8. SLIDING-WINDOW INFERENCE ON FRAGMENT 1 (single forward pass per patch — no TTA)\n# ------------------------------------------------------------------------------\ndef gaussian_window(size, sigma_frac=0.5):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / win.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, device, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Fragment 1: running sliding-window inference over {len(coords)} patches ...\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        inp = torch.from_numpy(np.stack(batch_imgs)).to(device)\n        with autocast(enabled=(device == \"cuda\")):\n            logits = model(inp)\n            prob = torch.sigmoid(logits).float().cpu().numpy()[:, 0]\n        for p, (cy, cx) in zip(prob, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n\n    for (y, x) in coords:\n        raw = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        raw = normalize_patch(raw)   # same per-patch normalization as training\n        batch_imgs.append(raw); batch_coords.append((y, x))\n        if len(batch_imgs) == batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum == 0] = 1.0\n    return pred_sum / weight_sum\n\n\ndef postprocess(prob_map, threshold):\n    \"\"\"Post-processing WITHOUT test-time augmentation: threshold (tuned on\n    validation) + light morphological cleanup to remove isolated speckle noise.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = ndimage.binary_opening(binary, structure=np.ones((3, 3))).astype(np.uint8)\n    binary = ndimage.binary_closing(binary, structure=np.ones((5, 5))).astype(np.uint8)\n    return binary\n\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_labels = load_ink_labels(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.depth_indices)\n\ntest_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                      CFG.test_stride, CFG.device, CFG.infer_batch)\ntest_pred_bin = postprocess(test_prob, best_threshold)\n\n# metrics on fragment 1 (restricted to tissue mask area)\ngt_t = torch.from_numpy(test_labels.astype(np.float32)) * torch.from_numpy(test_mask.astype(np.float32))\nprob_t = torch.from_numpy(test_prob)\ntest_metrics = compute_metrics(prob_t, gt_t, best_threshold)\nprint(\"\\nFinal TEST (fragment 1) metrics:\", test_metrics)\n\nwith open(os.path.join(CFG.out_dir, \"metrics_summary.json\"), \"w\") as f:\n    json.dump({\"train\": final_train_metrics, \"val\": final_val_metrics,\n               \"test\": test_metrics, \"tuned_threshold\": best_threshold}, f, indent=2)\n\ntorch.cuda.empty_cache()\ngc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 9. SAVE TEST VISUALIZATIONS (input | ground truth | prediction)\n# ------------------------------------------------------------------------------\ndef save_full_overview():\n    mid_idx = CFG.depth_indices[len(CFG.depth_indices) // 2]\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize(test_labels * 255, small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input (slice {mid_idx})\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground truth ink labels\")\n    axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n    for a in axes:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_full_overview.png\")\n    plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {path}\")\n\n\ndef save_patch_comparisons(n=6):\n    ys_xs = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(ys_xs)\n    ys_xs = ys_xs[:n]\n    mid_local_idx = len(CFG.depth_indices) // 2\n\n    for i, (y, x) in enumerate(ys_xs):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_local_idx]\n        gt_patch = test_labels[y:y + size, x:x + size]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground truth\")\n        axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n        for a in axes:\n            a.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_comparison_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {len(ys_xs)} patch comparisons to {CFG.viz_dir}\")\n\n\nsave_full_overview()\nsave_patch_comparisons(n=6)\n\ntest_vol.close()\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"\\n=== DONE ===\")\nprint(f\"Best checkpoint: {CFG.ckpt_path}\")\nprint(f\"Tuned threshold: {best_threshold:.2f}\")\nprint(f\"Metrics summary: {os.path.join(CFG.out_dir, 'metrics_summary.json')}\")\nprint(f\"Visualizations:  {CFG.viz_dir}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# VESUVIUS CHALLENGE — INK DETECTION PIPELINE\n# Custom \"VesuviusNet\" (3D-stem + Attention U-Net with deep supervision)\n# Train: fragments 2 & 3 (80/20 split)  |  Test: fragment 1\n# Designed to avoid OOM on Kaggle: disk-backed memmap volumes, patch-based\n# training, sliding-window (non-TTA) inference, AMP, aggressive gc.\n# Paste this whole file into a single Kaggle notebook cell and run.\n# ==============================================================================\n#!pip install segmentation-models-pytorch==0.2.0\nimport os, gc, random, time, json\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    os.system(\"pip install -q albumentations\")\n    import albumentations as A\n\n\n# ------------------------------------------------------------------------------\n# 0. CONFIG\n# ------------------------------------------------------------------------------\nclass CFG:\n    base_dir       = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n    train_frags    = [\"2\", \"3\"]\n    test_frag      = \"1\"\n    z_start, z_end = 0, 64          # 00.tif .. 64.tif  -> 65 channels\n    in_depth       = z_end - z_start + 1\n\n    patch_size     = 224\n    train_stride   = 112            # 50% overlap grid for more training samples\n    test_stride    = 112            # sliding-window stride for inference (blended)\n\n    val_fraction   = 0.20\n    min_tissue_frac_train = 0.05    # skip near-empty patches when building train grid\n    min_tissue_frac_test  = 0.02\n\n    batch_size     = 8\n    infer_batch    = 16\n    num_workers    = 2\n    epochs         = 18\n    lr             = 1e-4\n    weight_decay   = 1e-4\n    base_ch        = 32\n    deep_sup       = True\n\n    threshold      = 0.5\n    seed           = 42\n\n    out_dir        = \"/kaggle/working\"\n    ckpt_path      = os.path.join(out_dir, \"vesuviusnet_best.pth\")\n    viz_dir        = os.path.join(out_dir, \"test_visualizations\")\n\n    device         = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\n\nset_seed(CFG.seed)\ntorch.backends.cudnn.benchmark = True\n\n\ndef cfg_to_dict(cfg_cls):\n    \"\"\"vars() on a *class* returns a non-picklable mappingproxy, so build a\n    plain dict of the simple (picklable) config values by hand instead.\"\"\"\n    d = {}\n    for k, v in cfg_cls.__dict__.items():\n        if k.startswith(\"__\"):\n            continue\n        if isinstance(v, (int, float, str, bool, type(None), list, tuple)):\n            d[k] = v\n    return d\n\n\n# ------------------------------------------------------------------------------\n# 1. PRE-PROCESSING HELPERS: tissue mask, patch grid, disk-backed volume reader\n# ------------------------------------------------------------------------------\ndef load_tissue_mask(frag_dir):\n    \"\"\"Load fragment tissue mask (mask.png if present, else derive from a mid slice).\"\"\"\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = (CFG.z_start + CFG.z_end) // 2\n        mid = tifffile.imread(os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n        thresh = mid.mean() * 0.15\n        mask = ((mid > thresh).astype(np.uint8)) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    for y in range(0, max(H - patch_size, 0) + 1, stride):\n        for x in range(0, max(W - patch_size, 0) + 1, stride):\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access to a fragment's 65-slice surface volume.\n    Avoids loading the full multi-GB volume into RAM.\"\"\"\n\n    def __init__(self, frag_dir, z_start=0, z_end=64):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in range(z_start, z_end + 1)]\n        self._slices = None\n        self._h, self._w = None, None\n\n    def _ensure_open(self):\n        if self._slices is None:\n            slices = []\n            for p in self.paths:\n                try:\n                    arr = tifffile.memmap(p, mode=\"r\")\n                except Exception:\n                    # fallback if the tif can't be memory-mapped (e.g. compressed)\n                    arr = tifffile.imread(p)\n                slices.append(arr)\n            self._slices = slices\n            self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        gc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 2. AUGMENTATION (train only) — kept to spatial-safe, version-stable transforms\n# ------------------------------------------------------------------------------\ndef build_train_transform(size):\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.25),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.3),\n    ])\n\n\n# ------------------------------------------------------------------------------\n# 3. DATASET\n# ------------------------------------------------------------------------------\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None):\n        \"\"\"\n        volumes: dict frag_id -> FragmentVolume\n        labels:  dict frag_id -> full-res (H,W) uint8 ink label array\n        samples: list of (frag_id, y, x)\n        \"\"\"\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        frag_id, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[frag_id]\n        patch = vol.read_patch(y, x, size)                # (C,H,W) uint8\n        label = self.labels[frag_id][y:y + size, x:x + size]\n\n        img = np.transpose(patch, (1, 2, 0))               # HWC for albumentations\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))  # CHW\n        label = (label > 0).astype(np.float32)[None, ...]         # 1HW\n\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ------------------------------------------------------------------------------\n# 4. MODEL — \"VesuviusNet\": 3D depth-reduction stem + Attention U-Net + deep sup.\n# ------------------------------------------------------------------------------\nclass ConvBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass DepthStem(nn.Module):\n    \"\"\"Collapses the 65-slice pseudo-depth axis with 3D convolutions before\n    handing a 2D feature map to the U-Net — lets the model learn cross-slice\n    (through-page) ink signal instead of treating slices as independent channels.\"\"\"\n\n    def __init__(self, out_ch=32):\n        super().__init__()\n        self.conv3d = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=(7, 3, 3), stride=(2, 1, 1), padding=(3, 1, 1), bias=False),\n            nn.BatchNorm3d(16), nn.ReLU(inplace=True),\n            nn.Conv3d(16, 32, kernel_size=(5, 3, 3), stride=(2, 1, 1), padding=(2, 1, 1), bias=False),\n            nn.BatchNorm3d(32), nn.ReLU(inplace=True),\n        )\n        self.depth_pool = nn.AdaptiveAvgPool3d((1, None, None))\n        self.proj = nn.Conv2d(32, out_ch, 1)\n\n    def forward(self, x):\n        x = x.unsqueeze(1)              # (B,1,D,H,W)\n        x = self.conv3d(x)\n        x = self.depth_pool(x).squeeze(2)  # (B,32,H,W)\n        return self.proj(x)\n\n\nclass AttentionGate(nn.Module):\n    def __init__(self, f_g, f_l, f_int):\n        super().__init__()\n        self.w_g = nn.Sequential(nn.Conv2d(f_g, f_int, 1), nn.BatchNorm2d(f_int))\n        self.w_x = nn.Sequential(nn.Conv2d(f_l, f_int, 1), nn.BatchNorm2d(f_int))\n        self.psi = nn.Sequential(nn.Conv2d(f_int, 1, 1), nn.BatchNorm2d(1), nn.Sigmoid())\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, g, x):\n        psi = self.relu(self.w_g(g) + self.w_x(x))\n        return x * self.psi(psi)\n\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, skip_ch, out_ch):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride=2)\n        self.att = AttentionGate(in_ch // 2, skip_ch, skip_ch // 2)\n        self.conv = ConvBlock(in_ch // 2 + skip_ch, out_ch)\n\n    def forward(self, x, skip):\n        x = self.up(x)\n        if x.shape[-2:] != skip.shape[-2:]:\n            x = F.interpolate(x, size=skip.shape[-2:], mode=\"bilinear\", align_corners=False)\n        skip = self.att(x, skip)\n        return self.conv(torch.cat([x, skip], dim=1))\n\n\nclass VesuviusNet(nn.Module):\n    def __init__(self, base_ch=32, num_classes=1, deep_sup=True):\n        super().__init__()\n        self.stem = DepthStem(base_ch)\n        self.enc1 = ConvBlock(base_ch, base_ch)\n        self.enc2 = ConvBlock(base_ch, base_ch * 2)\n        self.enc3 = ConvBlock(base_ch * 2, base_ch * 4)\n        self.enc4 = ConvBlock(base_ch * 4, base_ch * 8)\n        self.pool = nn.MaxPool2d(2)\n        self.bottleneck = ConvBlock(base_ch * 8, base_ch * 16)\n        self.up4 = Up(base_ch * 16, base_ch * 8, base_ch * 8)\n        self.up3 = Up(base_ch * 8, base_ch * 4, base_ch * 4)\n        self.up2 = Up(base_ch * 4, base_ch * 2, base_ch * 2)\n        self.up1 = Up(base_ch * 2, base_ch, base_ch)\n        self.out_conv = nn.Conv2d(base_ch, num_classes, 1)\n        self.deep_sup = deep_sup\n        if deep_sup:\n            self.ds3 = nn.Conv2d(base_ch * 4, num_classes, 1)\n            self.ds2 = nn.Conv2d(base_ch * 2, num_classes, 1)\n\n    def forward(self, x):\n        s = self.stem(x)\n        e1 = self.enc1(s)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        e4 = self.enc4(self.pool(e3))\n        b = self.bottleneck(self.pool(e4))\n        d4 = self.up4(b, e4)\n        d3 = self.up3(d4, e3)\n        d2 = self.up2(d3, e2)\n        d1 = self.up1(d2, e1)\n        out = self.out_conv(d1)\n        if self.deep_sup and self.training:\n            ds3 = F.interpolate(self.ds3(d3), size=out.shape[-2:], mode=\"bilinear\", align_corners=False)\n            ds2 = F.interpolate(self.ds2(d2), size=out.shape[-2:], mode=\"bilinear\", align_corners=False)\n            return out, ds3, ds2\n        return out\n\n\n# ------------------------------------------------------------------------------\n# 5. LOSS & METRICS\n# ------------------------------------------------------------------------------\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(1)\n    union = probs.sum(1) + t.sum(1)\n    return 1 - ((2 * inter + eps) / (union + eps)).mean()\n\n\nclass ComboLoss(nn.Module):\n    def __init__(self, bce_w=0.5, dice_w=0.5):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.bce_w, self.dice_w = bce_w, dice_w\n\n    def forward(self, logits, targets):\n        return self.bce_w * self.bce(logits, targets) + self.dice_w * dice_loss(logits, targets)\n\n\ndef deep_sup_loss(outputs, targets, criterion, weights=(1.0, 0.4, 0.2)):\n    if isinstance(outputs, tuple):\n        main, ds3, ds2 = outputs\n        return (weights[0] * criterion(main, targets) +\n                weights[1] * criterion(ds3, targets) +\n                weights[2] * criterion(ds2, targets))\n    return criterion(outputs, targets)\n\n\n@torch.no_grad()\ndef compute_metrics(probs, targets, threshold=0.5, eps=1e-6):\n    preds = (probs > threshold).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    return {\"dice\": dice.item(), \"iou\": iou.item(), \"precision\": precision.item(),\n            \"recall\": recall.item(), \"fbeta0.5\": fbeta.item()}\n\n\ndef avg_metric_dicts(dicts):\n    keys = dicts[0].keys()\n    return {k: float(np.mean([d[k] for d in dicts])) for k in keys}\n\n\n# ------------------------------------------------------------------------------\n# 6. BUILD DATA (fragments 2 & 3 -> train/val ; fragment 1 -> test)\n# ------------------------------------------------------------------------------\nprint(\"Building tissue masks, label maps and patch grids ...\")\n\ntrain_volumes, train_labels_full = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    vol = FragmentVolume(frag_dir, CFG.z_start, CFG.z_end)\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"  fragment {fid}: mask {mask.shape}, {len(coords)} candidate patches\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    all_samples.extend([(fid, y, x) for (y, x) in coords])\n\n    del mask\n    gc.collect()\n\nrandom.shuffle(all_samples)\nn_val = int(len(all_samples) * CFG.val_fraction)\nval_samples = all_samples[:n_val]\ntrain_samples = all_samples[n_val:]\nprint(f\"Total patches: {len(all_samples)}  -> train {len(train_samples)} / val {len(val_samples)}\")\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples,\n                            CFG.patch_size, transform=build_train_transform(CFG.patch_size))\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples,\n                          CFG.patch_size, transform=None)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=True, persistent_workers=CFG.num_workers > 0)\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         persistent_workers=CFG.num_workers > 0)\n\n\n# ------------------------------------------------------------------------------\n# 7. TRAIN / VALIDATE\n# ------------------------------------------------------------------------------\nmodel = VesuviusNet(base_ch=CFG.base_ch, deep_sup=CFG.deep_sup).to(CFG.device)\ncriterion = ComboLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nbest_val_dice = -1.0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n\ndef run_epoch(loader, train_mode):\n    model.train(train_mode)\n    total_loss, metric_list = 0.0, []\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                outputs = model(imgs)\n                loss = deep_sup_loss(outputs, masks, criterion) if (train_mode and CFG.deep_sup) \\\n                    else criterion(outputs if not isinstance(outputs, tuple) else outputs[0], masks)\n\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                scaler.step(optimizer)\n                scaler.update()\n\n        logits = outputs[0] if isinstance(outputs, tuple) else outputs\n        probs = torch.sigmoid(logits.detach())\n        metric_list.append(compute_metrics(probs, masks, CFG.threshold))\n        total_loss += loss.item()\n\n        del imgs, masks, outputs, logits, probs\n    torch.cuda.empty_cache()\n    return total_loss / max(len(loader), 1), avg_metric_dicts(metric_list)\n\n\nprint(\"\\nStarting training ...\")\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_metrics = run_epoch(train_loader, train_mode=True)\n    val_loss, val_metrics = run_epoch(val_loader, train_mode=False)\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_metrics[\"dice\"])\n\n    print(f\"[{epoch:02d}/{CFG.epochs}] \"\n          f\"train_loss={train_loss:.4f} dice={train_metrics['dice']:.4f} | \"\n          f\"val_loss={val_loss:.4f} dice={val_metrics['dice']:.4f} \"\n          f\"fbeta0.5={val_metrics['fbeta0.5']:.4f} ({time.time()-t0:.1f}s)\")\n\n    if val_metrics[\"dice\"] > best_val_dice:\n        best_val_dice = val_metrics[\"dice\"]\n        torch.save({\"model\": model.state_dict(), \"cfg\": cfg_to_dict(CFG)}, CFG.ckpt_path)\n        print(f\"  -> saved new best checkpoint (val_dice={best_val_dice:.4f})\")\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# final train / val metrics with best checkpoint\nmodel.load_state_dict(torch.load(CFG.ckpt_path, map_location=CFG.device)[\"model\"])\n_, final_train_metrics = run_epoch(train_loader, train_mode=False)\n_, final_val_metrics = run_epoch(val_loader, train_mode=False)\nprint(\"\\nFinal TRAIN metrics:\", final_train_metrics)\nprint(\"Final VAL metrics:  \", final_val_metrics)\n\n# free fragment 2/3 volumes before loading fragment 1\nfor v in train_volumes.values():\n    v.close()\ndel train_loader, val_loader, train_ds, val_ds, train_volumes, train_labels_full\ngc.collect()\ntorch.cuda.empty_cache()\n\n\n# ------------------------------------------------------------------------------\n# 8. SLIDING-WINDOW INFERENCE ON FRAGMENT 1 (single forward pass per patch — no TTA)\n# ------------------------------------------------------------------------------\ndef gaussian_window(size, sigma_frac=0.5):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / win.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, device, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Fragment 1: running sliding-window inference over {len(coords)} patches ...\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        inp = torch.from_numpy(np.stack(batch_imgs)).to(device)\n        with autocast(enabled=(device == \"cuda\")):\n            out = model(inp)\n            out = out[0] if isinstance(out, tuple) else out\n            prob = torch.sigmoid(out).float().cpu().numpy()[:, 0]\n        for p, (cy, cx) in zip(prob, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n\n    for (y, x) in coords:\n        patch = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        batch_imgs.append(patch); batch_coords.append((y, x))\n        if len(batch_imgs) == batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum == 0] = 1.0\n    return pred_sum / weight_sum\n\n\ndef postprocess(prob_map, threshold):\n    \"\"\"Post-processing WITHOUT test-time augmentation: threshold + light\n    morphological cleanup to remove isolated speckle noise.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = ndimage.binary_opening(binary, structure=np.ones((3, 3))).astype(np.uint8)\n    binary = ndimage.binary_closing(binary, structure=np.ones((5, 5))).astype(np.uint8)\n    return binary\n\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_labels = load_ink_labels(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.z_start, CFG.z_end)\n\ntest_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                      CFG.test_stride, CFG.device, CFG.infer_batch)\ntest_pred_bin = postprocess(test_prob, CFG.threshold)\n\n# metrics on fragment 1 (restricted to tissue mask area)\ngt_t = torch.from_numpy(test_labels.astype(np.float32)) * torch.from_numpy(test_mask.astype(np.float32))\npred_t = torch.from_numpy(test_pred_bin.astype(np.float32))\nprob_t = torch.from_numpy(test_prob)\ntest_metrics = compute_metrics(prob_t, gt_t, CFG.threshold)\nprint(\"\\nFinal TEST (fragment 1) metrics:\", test_metrics)\n\nwith open(os.path.join(CFG.out_dir, \"metrics_summary.json\"), \"w\") as f:\n    json.dump({\"train\": final_train_metrics, \"val\": final_val_metrics,\n               \"test\": test_metrics}, f, indent=2)\n\ntorch.cuda.empty_cache()\ngc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 9. SAVE TEST VISUALIZATIONS (input | ground truth | prediction)\n# ------------------------------------------------------------------------------\ndef save_full_overview():\n    mid_idx = (CFG.z_start + CFG.z_end) // 2\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize(test_labels * 255, small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input (slice {mid_idx})\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground truth ink labels\")\n    axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n    for a in axes:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_full_overview.png\")\n    plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {path}\")\n\n\ndef save_patch_comparisons(n=6):\n    ys_xs = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(ys_xs)\n    ys_xs = ys_xs[:n]\n    mid_idx = (CFG.z_start + CFG.z_end) // 2\n\n    for i, (y, x) in enumerate(ys_xs):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_idx - CFG.z_start]\n        gt_patch = test_labels[y:y + size, x:x + size]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground truth\")\n        axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n        for a in axes:\n            a.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_comparison_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {len(ys_xs)} patch comparisons to {CFG.viz_dir}\")\n\n\nsave_full_overview()\nsave_patch_comparisons(n=6)\n\ntest_vol.close()\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"\\n=== DONE ===\")\nprint(f\"Best checkpoint: {CFG.ckpt_path}\")\nprint(f\"Metrics summary: {os.path.join(CFG.out_dir, 'metrics_summary.json')}\")\nprint(f\"Visualizations:  {CFG.viz_dir}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# VESUVIUS CHALLENGE — INK DETECTION PIPELINE\n# Custom \"VesuviusNet\" (3D-stem + Attention U-Net with deep supervision)\n# Train: fragments 2 & 3 (80/20 split)  |  Test: fragment 1\n# Designed to avoid OOM on Kaggle: disk-backed memmap volumes, patch-based\n# training, sliding-window (non-TTA) inference, AMP, aggressive gc.\n# Paste this whole file into a single Kaggle notebook cell and run.\n# ==============================================================================\n\nimport os, gc, random, time, json\nimport numpy as np\nimport cv2\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom scipy import ndimage\nimport matplotlib.pyplot as plt\n\ntry:\n    import albumentations as A\nexcept ImportError:\n    os.system(\"pip install -q albumentations\")\n    import albumentations as A\n\n\n# ------------------------------------------------------------------------------\n# 0. CONFIG\n# ------------------------------------------------------------------------------\nclass CFG:\n    base_dir       = \"/kaggle/input/vesuvius-challenge/train\"\n    train_frags    = [\"2\", \"3\"]\n    test_frag      = \"1\"\n    z_start, z_end = 0, 64          # 00.tif .. 64.tif  -> 65 channels\n    in_depth       = z_end - z_start + 1\n\n    patch_size     = 224\n    train_stride   = 112            # 50% overlap grid for more training samples\n    test_stride    = 112            # sliding-window stride for inference (blended)\n\n    val_fraction   = 0.20\n    min_tissue_frac_train = 0.05    # skip near-empty patches when building train grid\n    min_tissue_frac_test  = 0.02\n\n    batch_size     = 8\n    infer_batch    = 16\n    num_workers    = 2\n    epochs         = 18\n    lr             = 1e-4\n    weight_decay   = 1e-4\n    base_ch        = 32\n    deep_sup       = True\n\n    threshold      = 0.5\n    seed           = 42\n\n    out_dir        = \"/kaggle/working\"\n    ckpt_path      = os.path.join(out_dir, \"vesuviusnet_best.pth\")\n    viz_dir        = os.path.join(out_dir, \"test_visualizations\")\n\n    device         = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\nos.makedirs(CFG.out_dir, exist_ok=True)\nos.makedirs(CFG.viz_dir, exist_ok=True)\n\n\ndef set_seed(seed):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\n\nset_seed(CFG.seed)\ntorch.backends.cudnn.benchmark = True\n\n\n# ------------------------------------------------------------------------------\n# 1. PRE-PROCESSING HELPERS: tissue mask, patch grid, disk-backed volume reader\n# ------------------------------------------------------------------------------\ndef load_tissue_mask(frag_dir):\n    \"\"\"Load fragment tissue mask (mask.png if present, else derive from a mid slice).\"\"\"\n    mask_path = os.path.join(frag_dir, \"mask.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    else:\n        mid_idx = (CFG.z_start + CFG.z_end) // 2\n        mid = tifffile.imread(os.path.join(frag_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n        thresh = mid.mean() * 0.15\n        mask = ((mid > thresh).astype(np.uint8)) * 255\n    return (mask > 0).astype(np.uint8)\n\n\ndef load_ink_labels(frag_dir):\n    path = os.path.join(frag_dir, \"inklabels.png\")\n    lbl = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    return (lbl > 0).astype(np.uint8)\n\n\ndef generate_grid_coords(mask, patch_size, stride, min_frac):\n    H, W = mask.shape\n    coords = []\n    for y in range(0, max(H - patch_size, 0) + 1, stride):\n        for x in range(0, max(W - patch_size, 0) + 1, stride):\n            frac = mask[y:y + patch_size, x:x + patch_size].mean()\n            if frac > min_frac:\n                coords.append((y, x))\n    return coords\n\n\nclass FragmentVolume:\n    \"\"\"Disk-backed (memmap) access to a fragment's 65-slice surface volume.\n    Avoids loading the full multi-GB volume into RAM.\"\"\"\n\n    def __init__(self, frag_dir, z_start=0, z_end=64):\n        self.paths = [os.path.join(frag_dir, \"surface_volume\", f\"{i:02d}.tif\")\n                      for i in range(z_start, z_end + 1)]\n        self._slices = None\n        self._h, self._w = None, None\n\n    def _ensure_open(self):\n        if self._slices is None:\n            slices = []\n            for p in self.paths:\n                try:\n                    arr = tifffile.memmap(p, mode=\"r\")\n                except Exception:\n                    # fallback if the tif can't be memory-mapped (e.g. compressed)\n                    arr = tifffile.imread(p)\n                slices.append(arr)\n            self._slices = slices\n            self._h, self._w = slices[0].shape\n\n    @property\n    def shape(self):\n        self._ensure_open()\n        return (self._h, self._w)\n\n    def read_patch(self, y, x, size):\n        self._ensure_open()\n        out = np.empty((len(self._slices), size, size), dtype=np.uint8)\n        for i, s in enumerate(self._slices):\n            block = s[y:y + size, x:x + size]\n            if block.dtype != np.uint8:\n                block = (block.astype(np.float32) / 65535.0 * 255.0).astype(np.uint8)\n            out[i] = block\n        return out\n\n    def close(self):\n        self._slices = None\n        gc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 2. AUGMENTATION (train only) — kept to spatial-safe, version-stable transforms\n# ------------------------------------------------------------------------------\ndef build_train_transform(size):\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=20,\n                            border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.25),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.3),\n    ])\n\n\n# ------------------------------------------------------------------------------\n# 3. DATASET\n# ------------------------------------------------------------------------------\nclass InkPatchDataset(Dataset):\n    def __init__(self, volumes, labels, samples, patch_size, transform=None):\n        \"\"\"\n        volumes: dict frag_id -> FragmentVolume\n        labels:  dict frag_id -> full-res (H,W) uint8 ink label array\n        samples: list of (frag_id, y, x)\n        \"\"\"\n        self.volumes = volumes\n        self.labels = labels\n        self.samples = samples\n        self.patch_size = patch_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        frag_id, y, x = self.samples[idx]\n        size = self.patch_size\n        vol = self.volumes[frag_id]\n        patch = vol.read_patch(y, x, size)                # (C,H,W) uint8\n        label = self.labels[frag_id][y:y + size, x:x + size]\n\n        img = np.transpose(patch, (1, 2, 0))               # HWC for albumentations\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=label)\n            img, label = aug[\"image\"], aug[\"mask\"]\n\n        img = img.astype(np.float32) / 255.0\n        img = np.ascontiguousarray(np.transpose(img, (2, 0, 1)))  # CHW\n        label = (label > 0).astype(np.float32)[None, ...]         # 1HW\n\n        return torch.from_numpy(img), torch.from_numpy(label)\n\n\n# ------------------------------------------------------------------------------\n# 4. MODEL — \"VesuviusNet\": 3D depth-reduction stem + Attention U-Net + deep sup.\n# ------------------------------------------------------------------------------\nclass ConvBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass DepthStem(nn.Module):\n    \"\"\"Collapses the 65-slice pseudo-depth axis with 3D convolutions before\n    handing a 2D feature map to the U-Net — lets the model learn cross-slice\n    (through-page) ink signal instead of treating slices as independent channels.\"\"\"\n\n    def __init__(self, out_ch=32):\n        super().__init__()\n        self.conv3d = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=(7, 3, 3), stride=(2, 1, 1), padding=(3, 1, 1), bias=False),\n            nn.BatchNorm3d(16), nn.ReLU(inplace=True),\n            nn.Conv3d(16, 32, kernel_size=(5, 3, 3), stride=(2, 1, 1), padding=(2, 1, 1), bias=False),\n            nn.BatchNorm3d(32), nn.ReLU(inplace=True),\n        )\n        self.depth_pool = nn.AdaptiveAvgPool3d((1, None, None))\n        self.proj = nn.Conv2d(32, out_ch, 1)\n\n    def forward(self, x):\n        x = x.unsqueeze(1)              # (B,1,D,H,W)\n        x = self.conv3d(x)\n        x = self.depth_pool(x).squeeze(2)  # (B,32,H,W)\n        return self.proj(x)\n\n\nclass AttentionGate(nn.Module):\n    def __init__(self, f_g, f_l, f_int):\n        super().__init__()\n        self.w_g = nn.Sequential(nn.Conv2d(f_g, f_int, 1), nn.BatchNorm2d(f_int))\n        self.w_x = nn.Sequential(nn.Conv2d(f_l, f_int, 1), nn.BatchNorm2d(f_int))\n        self.psi = nn.Sequential(nn.Conv2d(f_int, 1, 1), nn.BatchNorm2d(1), nn.Sigmoid())\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, g, x):\n        psi = self.relu(self.w_g(g) + self.w_x(x))\n        return x * self.psi(psi)\n\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, skip_ch, out_ch):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride=2)\n        self.att = AttentionGate(in_ch // 2, skip_ch, skip_ch // 2)\n        self.conv = ConvBlock(in_ch // 2 + skip_ch, out_ch)\n\n    def forward(self, x, skip):\n        x = self.up(x)\n        if x.shape[-2:] != skip.shape[-2:]:\n            x = F.interpolate(x, size=skip.shape[-2:], mode=\"bilinear\", align_corners=False)\n        skip = self.att(x, skip)\n        return self.conv(torch.cat([x, skip], dim=1))\n\n\nclass VesuviusNet(nn.Module):\n    def __init__(self, base_ch=32, num_classes=1, deep_sup=True):\n        super().__init__()\n        self.stem = DepthStem(base_ch)\n        self.enc1 = ConvBlock(base_ch, base_ch)\n        self.enc2 = ConvBlock(base_ch, base_ch * 2)\n        self.enc3 = ConvBlock(base_ch * 2, base_ch * 4)\n        self.enc4 = ConvBlock(base_ch * 4, base_ch * 8)\n        self.pool = nn.MaxPool2d(2)\n        self.bottleneck = ConvBlock(base_ch * 8, base_ch * 16)\n        self.up4 = Up(base_ch * 16, base_ch * 8, base_ch * 8)\n        self.up3 = Up(base_ch * 8, base_ch * 4, base_ch * 4)\n        self.up2 = Up(base_ch * 4, base_ch * 2, base_ch * 2)\n        self.up1 = Up(base_ch * 2, base_ch, base_ch)\n        self.out_conv = nn.Conv2d(base_ch, num_classes, 1)\n        self.deep_sup = deep_sup\n        if deep_sup:\n            self.ds3 = nn.Conv2d(base_ch * 4, num_classes, 1)\n            self.ds2 = nn.Conv2d(base_ch * 2, num_classes, 1)\n\n    def forward(self, x):\n        s = self.stem(x)\n        e1 = self.enc1(s)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        e4 = self.enc4(self.pool(e3))\n        b = self.bottleneck(self.pool(e4))\n        d4 = self.up4(b, e4)\n        d3 = self.up3(d4, e3)\n        d2 = self.up2(d3, e2)\n        d1 = self.up1(d2, e1)\n        out = self.out_conv(d1)\n        if self.deep_sup and self.training:\n            ds3 = F.interpolate(self.ds3(d3), size=out.shape[-2:], mode=\"bilinear\", align_corners=False)\n            ds2 = F.interpolate(self.ds2(d2), size=out.shape[-2:], mode=\"bilinear\", align_corners=False)\n            return out, ds3, ds2\n        return out\n\n\n# ------------------------------------------------------------------------------\n# 5. LOSS & METRICS\n# ------------------------------------------------------------------------------\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits).reshape(logits.size(0), -1)\n    t = targets.reshape(targets.size(0), -1)\n    inter = (probs * t).sum(1)\n    union = probs.sum(1) + t.sum(1)\n    return 1 - ((2 * inter + eps) / (union + eps)).mean()\n\n\nclass ComboLoss(nn.Module):\n    def __init__(self, bce_w=0.5, dice_w=0.5):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.bce_w, self.dice_w = bce_w, dice_w\n\n    def forward(self, logits, targets):\n        return self.bce_w * self.bce(logits, targets) + self.dice_w * dice_loss(logits, targets)\n\n\ndef deep_sup_loss(outputs, targets, criterion, weights=(1.0, 0.4, 0.2)):\n    if isinstance(outputs, tuple):\n        main, ds3, ds2 = outputs\n        return (weights[0] * criterion(main, targets) +\n                weights[1] * criterion(ds3, targets) +\n                weights[2] * criterion(ds2, targets))\n    return criterion(outputs, targets)\n\n\n@torch.no_grad()\ndef compute_metrics(probs, targets, threshold=0.5, eps=1e-6):\n    preds = (probs > threshold).float()\n    tp = (preds * targets).sum()\n    fp = (preds * (1 - targets)).sum()\n    fn = ((1 - preds) * targets).sum()\n    dice = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    iou = (tp + eps) / (tp + fp + fn + eps)\n    precision = (tp + eps) / (tp + fp + eps)\n    recall = (tp + eps) / (tp + fn + eps)\n    beta2 = 0.25\n    fbeta = ((1 + beta2) * precision * recall + eps) / ((beta2 * precision) + recall + eps)\n    return {\"dice\": dice.item(), \"iou\": iou.item(), \"precision\": precision.item(),\n            \"recall\": recall.item(), \"fbeta0.5\": fbeta.item()}\n\n\ndef avg_metric_dicts(dicts):\n    keys = dicts[0].keys()\n    return {k: float(np.mean([d[k] for d in dicts])) for k in keys}\n\n\n# ------------------------------------------------------------------------------\n# 6. BUILD DATA (fragments 2 & 3 -> train/val ; fragment 1 -> test)\n# ------------------------------------------------------------------------------\nprint(\"Building tissue masks, label maps and patch grids ...\")\n\ntrain_volumes, train_labels_full = {}, {}\nall_samples = []\n\nfor fid in CFG.train_frags:\n    frag_dir = os.path.join(CFG.base_dir, fid)\n    mask = load_tissue_mask(frag_dir)\n    labels = load_ink_labels(frag_dir)\n    vol = FragmentVolume(frag_dir, CFG.z_start, CFG.z_end)\n\n    coords = generate_grid_coords(mask, CFG.patch_size, CFG.train_stride, CFG.min_tissue_frac_train)\n    print(f\"  fragment {fid}: mask {mask.shape}, {len(coords)} candidate patches\")\n\n    train_volumes[fid] = vol\n    train_labels_full[fid] = labels\n    all_samples.extend([(fid, y, x) for (y, x) in coords])\n\n    del mask\n    gc.collect()\n\nrandom.shuffle(all_samples)\nn_val = int(len(all_samples) * CFG.val_fraction)\nval_samples = all_samples[:n_val]\ntrain_samples = all_samples[n_val:]\nprint(f\"Total patches: {len(all_samples)}  -> train {len(train_samples)} / val {len(val_samples)}\")\n\ntrain_ds = InkPatchDataset(train_volumes, train_labels_full, train_samples,\n                            CFG.patch_size, transform=build_train_transform(CFG.patch_size))\nval_ds = InkPatchDataset(train_volumes, train_labels_full, val_samples,\n                          CFG.patch_size, transform=None)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                           num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                           drop_last=True, persistent_workers=CFG.num_workers > 0)\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=(CFG.device == \"cuda\"),\n                         persistent_workers=CFG.num_workers > 0)\n\n\n# ------------------------------------------------------------------------------\n# 7. TRAIN / VALIDATE\n# ------------------------------------------------------------------------------\nmodel = VesuviusNet(base_ch=CFG.base_ch, deep_sup=CFG.deep_sup).to(CFG.device)\ncriterion = ComboLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\nscaler = GradScaler(enabled=(CFG.device == \"cuda\"))\n\nbest_val_dice = -1.0\nhistory = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n\ndef run_epoch(loader, train_mode):\n    model.train(train_mode)\n    total_loss, metric_list = 0.0, []\n    for imgs, masks in loader:\n        imgs = imgs.to(CFG.device, non_blocking=True)\n        masks = masks.to(CFG.device, non_blocking=True)\n\n        with torch.set_grad_enabled(train_mode):\n            with autocast(enabled=(CFG.device == \"cuda\")):\n                outputs = model(imgs)\n                loss = deep_sup_loss(outputs, masks, criterion) if (train_mode and CFG.deep_sup) \\\n                    else criterion(outputs if not isinstance(outputs, tuple) else outputs[0], masks)\n\n            if train_mode:\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                scaler.step(optimizer)\n                scaler.update()\n\n        logits = outputs[0] if isinstance(outputs, tuple) else outputs\n        probs = torch.sigmoid(logits.detach())\n        metric_list.append(compute_metrics(probs, masks, CFG.threshold))\n        total_loss += loss.item()\n\n        del imgs, masks, outputs, logits, probs\n    torch.cuda.empty_cache()\n    return total_loss / max(len(loader), 1), avg_metric_dicts(metric_list)\n\n\nprint(\"\\nStarting training ...\")\nfor epoch in range(1, CFG.epochs + 1):\n    t0 = time.time()\n    train_loss, train_metrics = run_epoch(train_loader, train_mode=True)\n    val_loss, val_metrics = run_epoch(val_loader, train_mode=False)\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_metrics[\"dice\"])\n\n    print(f\"[{epoch:02d}/{CFG.epochs}] \"\n          f\"train_loss={train_loss:.4f} dice={train_metrics['dice']:.4f} | \"\n          f\"val_loss={val_loss:.4f} dice={val_metrics['dice']:.4f} \"\n          f\"fbeta0.5={val_metrics['fbeta0.5']:.4f} ({time.time()-t0:.1f}s)\")\n\n    if val_metrics[\"dice\"] > best_val_dice:\n        best_val_dice = val_metrics[\"dice\"]\n        torch.save({\"model\": model.state_dict(), \"cfg\": vars(CFG)}, CFG.ckpt_path)\n        print(f\"  -> saved new best checkpoint (val_dice={best_val_dice:.4f})\")\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# final train / val metrics with best checkpoint\nmodel.load_state_dict(torch.load(CFG.ckpt_path, map_location=CFG.device)[\"model\"])\n_, final_train_metrics = run_epoch(train_loader, train_mode=False)\n_, final_val_metrics = run_epoch(val_loader, train_mode=False)\nprint(\"\\nFinal TRAIN metrics:\", final_train_metrics)\nprint(\"Final VAL metrics:  \", final_val_metrics)\n\n# free fragment 2/3 volumes before loading fragment 1\nfor v in train_volumes.values():\n    v.close()\ndel train_loader, val_loader, train_ds, val_ds, train_volumes, train_labels_full\ngc.collect()\ntorch.cuda.empty_cache()\n\n\n# ------------------------------------------------------------------------------\n# 8. SLIDING-WINDOW INFERENCE ON FRAGMENT 1 (single forward pass per patch — no TTA)\n# ------------------------------------------------------------------------------\ndef gaussian_window(size, sigma_frac=0.5):\n    ax = np.arange(size) - (size - 1) / 2.0\n    sigma = size * sigma_frac\n    g1d = np.exp(-(ax ** 2) / (2 * sigma ** 2))\n    win = np.outer(g1d, g1d)\n    return (win / win.max()).astype(np.float32)\n\n\n@torch.no_grad()\ndef sliding_window_inference(model, vol, mask, patch_size, stride, device, batch_size):\n    H, W = mask.shape\n    pred_sum = np.zeros((H, W), dtype=np.float32)\n    weight_sum = np.zeros((H, W), dtype=np.float32)\n    win = gaussian_window(patch_size)\n\n    coords = generate_grid_coords(mask, patch_size, stride, CFG.min_tissue_frac_test)\n    print(f\"Fragment 1: running sliding-window inference over {len(coords)} patches ...\")\n\n    model.eval()\n    batch_imgs, batch_coords = [], []\n\n    def flush():\n        if not batch_imgs:\n            return\n        inp = torch.from_numpy(np.stack(batch_imgs)).to(device)\n        with autocast(enabled=(device == \"cuda\")):\n            out = model(inp)\n            out = out[0] if isinstance(out, tuple) else out\n            prob = torch.sigmoid(out).float().cpu().numpy()[:, 0]\n        for p, (cy, cx) in zip(prob, batch_coords):\n            pred_sum[cy:cy + patch_size, cx:cx + patch_size] += p * win\n            weight_sum[cy:cy + patch_size, cx:cx + patch_size] += win\n        batch_imgs.clear(); batch_coords.clear()\n\n    for (y, x) in coords:\n        patch = vol.read_patch(y, x, patch_size).astype(np.float32) / 255.0\n        batch_imgs.append(patch); batch_coords.append((y, x))\n        if len(batch_imgs) == batch_size:\n            flush()\n    flush()\n\n    weight_sum[weight_sum == 0] = 1.0\n    return pred_sum / weight_sum\n\n\ndef postprocess(prob_map, threshold):\n    \"\"\"Post-processing WITHOUT test-time augmentation: threshold + light\n    morphological cleanup to remove isolated speckle noise.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = ndimage.binary_opening(binary, structure=np.ones((3, 3))).astype(np.uint8)\n    binary = ndimage.binary_closing(binary, structure=np.ones((5, 5))).astype(np.uint8)\n    return binary\n\n\ntest_dir = os.path.join(CFG.base_dir, CFG.test_frag)\ntest_mask = load_tissue_mask(test_dir)\ntest_labels = load_ink_labels(test_dir)\ntest_vol = FragmentVolume(test_dir, CFG.z_start, CFG.z_end)\n\ntest_prob = sliding_window_inference(model, test_vol, test_mask, CFG.patch_size,\n                                      CFG.test_stride, CFG.device, CFG.infer_batch)\ntest_pred_bin = postprocess(test_prob, CFG.threshold)\n\n# metrics on fragment 1 (restricted to tissue mask area)\ngt_t = torch.from_numpy(test_labels.astype(np.float32)) * torch.from_numpy(test_mask.astype(np.float32))\npred_t = torch.from_numpy(test_pred_bin.astype(np.float32))\nprob_t = torch.from_numpy(test_prob)\ntest_metrics = compute_metrics(prob_t, gt_t, CFG.threshold)\nprint(\"\\nFinal TEST (fragment 1) metrics:\", test_metrics)\n\nwith open(os.path.join(CFG.out_dir, \"metrics_summary.json\"), \"w\") as f:\n    json.dump({\"train\": final_train_metrics, \"val\": final_val_metrics,\n               \"test\": test_metrics}, f, indent=2)\n\ntorch.cuda.empty_cache()\ngc.collect()\n\n\n# ------------------------------------------------------------------------------\n# 9. SAVE TEST VISUALIZATIONS (input | ground truth | prediction)\n# ------------------------------------------------------------------------------\ndef save_full_overview():\n    mid_idx = (CFG.z_start + CFG.z_end) // 2\n    mid_slice = tifffile.imread(os.path.join(test_dir, \"surface_volume\", f\"{mid_idx:02d}.tif\"))\n    scale = 2000 / max(mid_slice.shape)\n    small = cv2.resize(mid_slice, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)\n    gt_small = cv2.resize(test_labels * 255, small.shape[::-1], interpolation=cv2.INTER_NEAREST)\n    pred_small = cv2.resize((test_pred_bin * 255).astype(np.uint8), small.shape[::-1],\n                             interpolation=cv2.INTER_NEAREST)\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    axes[0].imshow(small, cmap=\"gray\"); axes[0].set_title(f\"Input (slice {mid_idx})\")\n    axes[1].imshow(gt_small, cmap=\"gray\"); axes[1].set_title(\"Ground truth ink labels\")\n    axes[2].imshow(pred_small, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n    for a in axes:\n        a.axis(\"off\")\n    plt.tight_layout()\n    path = os.path.join(CFG.viz_dir, \"fragment1_full_overview.png\")\n    plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {path}\")\n\n\ndef save_patch_comparisons(n=6):\n    ys_xs = generate_grid_coords(test_mask, CFG.patch_size, CFG.patch_size, 0.15)\n    random.shuffle(ys_xs)\n    ys_xs = ys_xs[:n]\n    mid_idx = (CFG.z_start + CFG.z_end) // 2\n\n    for i, (y, x) in enumerate(ys_xs):\n        size = CFG.patch_size\n        input_slice = test_vol.read_patch(y, x, size)[mid_idx - CFG.z_start]\n        gt_patch = test_labels[y:y + size, x:x + size]\n        pred_patch = test_pred_bin[y:y + size, x:x + size]\n\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        axes[0].imshow(input_slice, cmap=\"gray\"); axes[0].set_title(\"Input\")\n        axes[1].imshow(gt_patch, cmap=\"gray\"); axes[1].set_title(\"Ground truth\")\n        axes[2].imshow(pred_patch, cmap=\"gray\"); axes[2].set_title(\"Prediction\")\n        for a in axes:\n            a.axis(\"off\")\n        plt.tight_layout()\n        path = os.path.join(CFG.viz_dir, f\"patch_comparison_{i:02d}_y{y}_x{x}.png\")\n        plt.savefig(path, dpi=150); plt.close(fig)\n    print(f\"Saved {len(ys_xs)} patch comparisons to {CFG.viz_dir}\")\n\n\nsave_full_overview()\nsave_patch_comparisons(n=6)\n\ntest_vol.close()\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"\\n=== DONE ===\")\nprint(f\"Best checkpoint: {CFG.ckpt_path}\")\nprint(f\"Metrics summary: {os.path.join(CFG.out_dir, 'metrics_summary.json')}\")\nprint(f\"Visualizations:  {CFG.viz_dir}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\"\"\"\n================================================================================\nVESUVIUS INK DETECTION - 3D UNet with Advanced Preprocessing\n================================================================================\nKaggle Notebook: Complete training pipeline for Vesuvius Challenge\nArchitecture: 3D Residual UNet with channel attention\nPreprocessing: Empty region masking, Z-score normalization, CLAHE\nData Augmentation: 3D spatial + intensity augmentations\nPost-processing: Morphological cleaning, hole filling\nEvaluation: F0.5 score, Dice, IoU\nMemory Optimization: Gradient checkpointing, mixed precision, tiled inference\n================================================================================\n\"\"\"\n!pip install segmentation-models-pytorch==0.2.0\nimport os\nimport gc\nimport time\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom collections import defaultdict\nfrom tqdm import tqdm\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.model_selection import train_test_split\nfrom scipy import ndimage\nfrom skimage import morphology, measure\nfrom skimage.filters import threshold_otsu\nimport matplotlib.pyplot as plt\n\n# ==============================================================================\n# CONFIGURATION\n# ==============================================================================\n\nclass Config:\n    \"\"\"Central configuration for the entire pipeline\"\"\"\n\n    # Data paths\n    DATA_ROOT = Path('/kaggle/input/vesuvius-challenge-ink-detection')\n    OUTPUT_DIR = Path('/kaggle/working')\n\n    # Fragments\n    TRAIN_FRAGMENTS = [2, 3]\n    VAL_FRAGMENTS = [2, 3]  # 20% split from these\n    TEST_FRAGMENT = 1\n\n    # Volume settings\n    Z_START = 0\n    Z_END = 65  # 00.tif to 64.tif = 65 layers\n    Z_DEPTH = Z_END - Z_START\n\n    # Image dimensions (after loading)\n    # Fragment 2: ~8096 x 6336\n    # Fragment 3: ~10496 x 8704\n    # Fragment 1: ~8184 x 6336\n\n    # Tile settings (critical for OOM prevention)\n    TILE_SIZE = 256\n    TILE_STRIDE = 192  # 75% overlap for training, reduces edge artifacts\n    INFERENCE_TILE_SIZE = 256\n    INFERENCE_STRIDE = 192\n\n    # Empty region filtering\n    MIN_INK_PERCENTAGE = 0.5  # Only keep tiles with at least 0.5% ink\n    BACKGROUND_THRESHOLD = 0.02  # Threshold to detect empty papyrus regions\n\n    # Model architecture\n    IN_CHANNELS = Z_DEPTH  # 65\n    OUT_CHANNELS = 1\n    BASE_CHANNELS = 32  # Reduced for OOM safety\n    DEPTH = 4\n\n    # Training\n    BATCH_SIZE = 2  # Conservative for 15GB VRAM\n    NUM_WORKERS = 2\n    EPOCHS = 10\n    LR = 1e-4\n    WEIGHT_DECAY = 1e-5\n\n    # Loss weights\n    BCE_WEIGHT = 0.5\n    DICE_WEIGHT = 0.5\n    FOCAL_WEIGHT = 0.0\n\n    # Optimization\n    USE_AMP = True  # Mixed precision\n    GRADIENT_CLIP = 1.0\n\n    # Augmentation\n    AUG_PROB = 0.5\n\n    # Post-processing\n    MORPH_KERNEL_SIZE = 3\n    MIN_COMPONENT_SIZE = 50\n\n    # Device\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    # Visualization\n    NUM_VIS_SAMPLES = 5\n\nconfig = Config()\n\nprint(f\"Device: {config.DEVICE}\")\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\n\n# ==============================================================================\n# UTILITY FUNCTIONS\n# ==============================================================================\n\ndef rle_encode(mask):\n    \"\"\"Run-length encoding for submission\"\"\"\n    pixels = mask.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ndef rle_decode(mask_rle, shape):\n    \"\"\"Decode RLE string to mask\"\"\"\n    if pd.isna(mask_rle) or mask_rle == '':\n        return np.zeros(shape, dtype=np.uint8)\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0::2], s[1::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T\n\ndef compute_f05(precision, recall):\n    \"\"\"Compute F0.5 score (precision-weighted)\"\"\"\n    beta = 0.5\n    if precision + recall == 0:\n        return 0.0\n    return (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall)\n\ndef compute_metrics(pred, target, threshold=0.5):\n    \"\"\"Compute segmentation metrics\"\"\"\n    pred_binary = (pred > threshold).astype(np.float32)\n    target_binary = (target > 0.5).astype(np.float32)\n\n    tp = np.sum(pred_binary * target_binary)\n    fp = np.sum(pred_binary * (1 - target_binary))\n    fn = np.sum((1 - pred_binary) * target_binary)\n\n    precision = tp / (tp + fp + 1e-7)\n    recall = tp / (tp + fn + 1e-7)\n    f05 = compute_f05(precision, recall)\n\n    intersection = np.sum(pred_binary * target_binary)\n    union = np.sum(pred_binary) + np.sum(target_binary) - intersection\n    iou = intersection / (union + 1e-7)\n\n    dice = 2 * intersection / (np.sum(pred_binary) + np.sum(target_binary) + 1e-7)\n\n    return {\n        'f05': f05,\n        'precision': precision,\n        'recall': recall,\n        'dice': dice,\n        'iou': iou\n    }\n\ndef free_memory():\n    \"\"\"Aggressive memory cleanup\"\"\"\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.synchronize()\n\n# ==============================================================================\n# DATA LOADING & PREPROCESSING\n# ==============================================================================\n\nclass VesuviusDataLoader:\n    \"\"\"Efficient data loader with memory-mapped reading\"\"\"\n\n    @staticmethod\n    def load_volume(fragment_id, z_start=0, z_end=65):\n        \"\"\"\n        Load 3D volume from TIFF stack.\n        Returns: numpy array of shape (Z, H, W) as float32\n        \"\"\"\n        fragment_path = config.DATA_ROOT / f'train/{fragment_id}/surface_volume'\n\n        # Get dimensions from first slice\n        sample_img = cv2.imread(str(fragment_path / f'{z_start:02d}.tif'), cv2.IMREAD_UNCHANGED)\n        if sample_img is None:\n            raise FileNotFoundError(f\"Could not load {fragment_path}/{z_start:02d}.tif\")\n\n        h, w = sample_img.shape\n        z_depth = z_end - z_start\n\n        # Pre-allocate array\n        volume = np.zeros((z_depth, h, w), dtype=np.float32)\n\n        for z in range(z_start, z_end):\n            img_path = fragment_path / f'{z:02d}.tif'\n            img = cv2.imread(str(img_path), cv2.IMREAD_UNCHANGED)\n            if img is None:\n                raise FileNotFoundError(f\"Could not load {img_path}\")\n            volume[z - z_start] = img.astype(np.float32)\n\n        return volume\n\n    @staticmethod\n    def load_inklabels(fragment_id):\n        \"\"\"Load ink labels as binary mask\"\"\"\n        label_path = config.DATA_ROOT / f'train/{fragment_id}/inklabels.png'\n        label = cv2.imread(str(label_path), cv2.IMREAD_GRAYSCALE)\n        if label is None:\n            raise FileNotFoundError(f\"Could not load {label_path}\")\n        return (label > 127).astype(np.float32)\n\n    @staticmethod\n    def load_mask(fragment_id):\n        \"\"\"Load fragment mask (non-empty regions)\"\"\"\n        mask_path = config.DATA_ROOT / f'train/{fragment_id}/mask.png'\n        mask = cv2.imread(str(mask_path), cv2.IMREAD_GRAYSCALE)\n        if mask is None:\n            # Create mask from volume if not available\n            return None\n        return (mask > 127).astype(np.float32)\n\n    @staticmethod\n    def preprocess_volume(volume):\n        \"\"\"\n        Advanced preprocessing pipeline:\n        1. Z-score normalization per slice\n        2. Global intensity normalization\n        3. CLAHE for contrast enhancement\n        4. Per-pixel standardization across Z\n        \"\"\"\n        z, h, w = volume.shape\n\n        # Per-slice Z-score normalization\n        for i in range(z):\n            slice_mean = np.mean(volume[i])\n            slice_std = np.std(volume[i]) + 1e-7\n            volume[i] = (volume[i] - slice_mean) / slice_std\n\n        # Global normalization to [-1, 1]\n        vol_min, vol_max = np.percentile(volume, [1, 99])\n        volume = np.clip(volume, vol_min, vol_max)\n        volume = 2 * (volume - vol_min) / (vol_max - vol_min + 1e-7) - 1\n\n        # CLAHE on each slice for local contrast\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        for i in range(z):\n            # Convert to uint8 for CLAHE, then back to float\n            slice_uint8 = ((volume[i] + 1) / 2 * 255).astype(np.uint8)\n            enhanced = clahe.apply(slice_uint8)\n            volume[i] = enhanced.astype(np.float32) / 127.5 - 1\n\n        return volume\n\n    @staticmethod\n    def create_empty_region_mask(volume, inklabels, background_threshold=0.02):\n        \"\"\"\n        Create mask to ignore empty papyrus regions.\n        Uses variance across Z-axis to detect papyrus vs empty regions.\n        \"\"\"\n        # Compute variance across Z - empty regions have low variance\n        z_variance = np.std(volume, axis=0)\n\n        # Normalize variance\n        z_variance = (z_variance - z_variance.min()) / (z_variance.max() - z_variance.min() + 1e-7)\n\n        # Threshold to get papyrus mask\n        papyrus_mask = (z_variance > background_threshold).astype(np.float32)\n\n        # Combine with inklabels to ensure we don't lose ink regions\n        # Dilate ink regions slightly\n        ink_dilated = ndimage.binary_dilation(inklabels > 0, iterations=5)\n        combined_mask = np.clip(papyrus_mask + ink_dilated.astype(np.float32), 0, 1)\n\n        return combined_mask\n\n# ==============================================================================\n# DATASET\n# ==============================================================================\n\nclass VesuviusDataset(Dataset):\n    \"\"\"\n    Dataset that generates tiles on-the-fly with memory-efficient loading.\n    Only includes tiles that contain sufficient papyrus/ink.\n    \"\"\"\n\n    def __init__(self, fragment_ids, tile_size=256, stride=192, \n                 is_training=True, min_ink_percentage=0.5,\n                 augment=True, cache_in_memory=False):\n        self.fragment_ids = fragment_ids if isinstance(fragment_ids, list) else [fragment_ids]\n        self.tile_size = tile_size\n        self.stride = stride\n        self.is_training = is_training\n        self.min_ink_percentage = min_ink_percentage\n        self.augment = augment and is_training\n        self.cache_in_memory = cache_in_memory\n\n        # Load all fragments\n        self.fragments = {}\n        self.tiles = []  # List of (fragment_id, x, y, weight)\n\n        print(f\"Loading fragments: {self.fragment_ids}\")\n        for frag_id in self.fragment_ids:\n            self._load_fragment(frag_id)\n\n        print(f\"Total tiles: {len(self.tiles)}\")\n\n        # Setup augmentations\n        self._setup_augmentations()\n\n    def _load_fragment(self, frag_id):\n        \"\"\"Load a single fragment and generate valid tiles\"\"\"\n        print(f\"  Loading fragment {frag_id}...\")\n\n        # Load volume\n        volume = VesuviusDataLoader.load_volume(frag_id, config.Z_START, config.Z_END)\n        volume = VesuviusDataLoader.preprocess_volume(volume)\n\n        # Load labels\n        inklabels = VesuviusDataLoader.load_inklabels(frag_id)\n\n        # Load or create mask\n        mask = VesuviusDataLoader.load_mask(frag_id)\n        if mask is None:\n            mask = VesuviusDataLoader.create_empty_region_mask(volume, inklabels, config.BACKGROUND_THRESHOLD)\n\n        # Ensure all same size\n        z, h, w = volume.shape\n        inklabels = cv2.resize(inklabels, (w, h), interpolation=cv2.INTER_NEAREST)\n        mask = cv2.resize(mask, (w, h), interpolation=cv2.INTER_NEAREST)\n\n        # Store\n        self.fragments[frag_id] = {\n            'volume': volume,\n            'inklabels': inklabels,\n            'mask': mask,\n            'shape': (z, h, w)\n        }\n\n        # Generate tiles\n        valid_tiles = self._generate_tiles(frag_id, h, w, inklabels, mask)\n        self.tiles.extend(valid_tiles)\n\n        print(f\"    Fragment {frag_id}: {len(valid_tiles)} valid tiles\")\n\n        # Free memory if not caching\n        if not self.cache_in_memory:\n            del self.fragments[frag_id]['volume']\n            self.fragments[frag_id]['volume'] = None\n            free_memory()\n\n    def _generate_tiles(self, frag_id, h, w, inklabels, mask):\n        \"\"\"Generate list of valid tile coordinates\"\"\"\n        tiles = []\n\n        for y in range(0, h - self.tile_size + 1, self.stride):\n            for x in range(0, w - self.tile_size + 1, self.stride):\n                # Extract tile regions\n                tile_mask = mask[y:y+self.tile_size, x:x+self.tile_size]\n                tile_ink = inklabels[y:y+self.tile_size, x:x+self.tile_size]\n\n                # Check if tile has enough papyrus\n                papyrus_ratio = np.mean(tile_mask)\n                if papyrus_ratio < 0.1:  # Less than 10% papyrus\n                    continue\n\n                # Check ink percentage\n                ink_ratio = np.mean(tile_ink) * 100\n\n                # Weight tiles with ink higher\n                weight = 1.0 + ink_ratio * 10  # Weight ink regions more\n\n                if ink_ratio >= self.min_ink_percentage or not self.is_training:\n                    tiles.append((frag_id, x, y, weight, ink_ratio))\n\n        return tiles\n\n    def _setup_augmentations(self):\n        \"\"\"Setup 3D-aware augmentations\"\"\"\n        if self.augment:\n            self.transform = A.Compose([\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.RandomRotate90(p=0.5),\n                A.ShiftScaleRotate(\n                    shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5\n                ),\n                A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),\n                A.GaussNoise(var_limit=(5, 20), p=0.2),\n                A.Blur(blur_limit=3, p=0.2),\n            ])\n        else:\n            self.transform = None\n\n    def __len__(self):\n        return len(self.tiles)\n\n    def __getitem__(self, idx):\n        frag_id, x, y, weight, ink_ratio = self.tiles[idx]\n\n        # Load volume tile\n        frag = self.fragments[frag_id]\n        if frag['volume'] is None:\n            # Reload if not cached\n            volume = VesuviusDataLoader.load_volume(frag_id, config.Z_START, config.Z_END)\n            volume = VesuviusDataLoader.preprocess_volume(volume)\n        else:\n            volume = frag['volume']\n\n        # Extract tile\n        tile_volume = volume[:, y:y+self.tile_size, x:x+self.tile_size]\n        tile_label = frag['inklabels'][y:y+self.tile_size, x:x+self.tile_size]\n        tile_mask = frag['mask'][y:y+self.tile_size, x:x+self.tile_size]\n\n        # Apply mask to label (ignore empty regions)\n        tile_label = tile_label * tile_mask\n\n        # Transpose to (C, H, W) for model\n        tile_volume = torch.from_numpy(tile_volume).float()\n        tile_label = torch.from_numpy(tile_label).float().unsqueeze(0)\n\n        # Apply augmentations (on each slice)\n        if self.transform is not None:\n            # Apply same spatial transform to all slices\n            seed = np.random.randint(0, 100000)\n\n            augmented_slices = []\n            for z in range(tile_volume.shape[0]):\n                aug = self.transform(image=tile_volume[z].numpy())\n                augmented_slices.append(aug['image'])\n\n            tile_volume = np.stack(augmented_slices, axis=0)\n            tile_volume = torch.from_numpy(tile_volume).float()\n\n            # Apply same transform to label\n            aug_label = self.transform(image=tile_label[0].numpy())\n            tile_label = torch.from_numpy(aug_label['image']).float().unsqueeze(0)\n\n        return tile_volume, tile_label, torch.tensor(weight, dtype=torch.float32)\n\n# ==============================================================================\n# 3D MODEL ARCHITECTURE\n# ==============================================================================\n\nclass ChannelAttention3D(nn.Module):\n    \"\"\"3D Channel Attention Module\"\"\"\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool3d(1)\n        self.max_pool = nn.AdaptiveMaxPool3d(1)\n        self.fc = nn.Sequential(\n            nn.Conv3d(channels, channels // reduction, 1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(channels // reduction, channels, 1, bias=False)\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = self.fc(self.avg_pool(x))\n        max_out = self.fc(self.max_pool(x))\n        out = self.sigmoid(avg_out + max_out)\n        return x * out\n\nclass SpatialAttention3D(nn.Module):\n    \"\"\"3D Spatial Attention Module\"\"\"\n    def __init__(self, kernel_size=7):\n        super().__init__()\n        padding = kernel_size // 2\n        self.conv = nn.Conv3d(2, 1, kernel_size, padding=padding, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        out = torch.cat([avg_out, max_out], dim=1)\n        out = self.conv(out)\n        return x * self.sigmoid(out)\n\nclass ResidualBlock3D(nn.Module):\n    \"\"\"3D Residual Block with Attention\"\"\"\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_channels, out_channels, 3, stride, 1, bias=False)\n        self.bn1 = nn.BatchNorm3d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = nn.Conv3d(out_channels, out_channels, 3, 1, 1, bias=False)\n        self.bn2 = nn.BatchNorm3d(out_channels)\n\n        self.ca = ChannelAttention3D(out_channels)\n        self.sa = SpatialAttention3D()\n\n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv3d(in_channels, out_channels, 1, stride, bias=False),\n                nn.BatchNorm3d(out_channels)\n            )\n\n    def forward(self, x):\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out = self.ca(out)\n        out = self.sa(out)\n        out += self.shortcut(x)\n        out = self.relu(out)\n        return out\n\nclass EncoderBlock3D(nn.Module):\n    \"\"\"3D Encoder Block with Downsampling\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.block = ResidualBlock3D(in_channels, out_channels, stride=2)\n        self.residual = ResidualBlock3D(out_channels, out_channels)\n\n    def forward(self, x):\n        x = self.block(x)\n        x = self.residual(x)\n        return x\n\nclass DecoderBlock3D(nn.Module):\n    \"\"\"3D Decoder Block with Upsampling\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.up = nn.ConvTranspose3d(in_channels, out_channels, 2, stride=2)\n        self.conv = nn.Sequential(\n            ResidualBlock3D(out_channels * 2, out_channels),\n            ResidualBlock3D(out_channels, out_channels)\n        )\n\n    def forward(self, x, skip):\n        x = self.up(x)\n        # Handle size mismatch\n        if x.shape != skip.shape:\n            x = F.interpolate(x, size=skip.shape[2:], mode='trilinear', align_corners=False)\n        x = torch.cat([x, skip], dim=1)\n        x = self.conv(x)\n        return x\n\nclass Vesuvius3DUNet(nn.Module):\n    \"\"\"\n    3D Residual UNet with Channel & Spatial Attention\n    Designed for Vesuvius Ink Detection\n\n    Input: (B, Z_DEPTH, H, W) where Z_DEPTH=65\n    Output: (B, 1, H, W)\n    \"\"\"\n\n    def __init__(self, in_channels=65, out_channels=1, base_channels=32, depth=4):\n        super().__init__()\n\n        # Initial projection: treat Z as channels\n        self.input_conv = nn.Sequential(\n            nn.Conv3d(1, base_channels, 3, padding=1, bias=False),\n            nn.BatchNorm3d(base_channels),\n            nn.ReLU(inplace=True),\n            ResidualBlock3D(base_channels, base_channels)\n        )\n\n        # Encoder\n        self.encoders = nn.ModuleList()\n        ch = base_channels\n        for i in range(depth):\n            self.encoders.append(EncoderBlock3D(ch, ch * 2))\n            ch *= 2\n\n        # Bottleneck\n        self.bottleneck = nn.Sequential(\n            ResidualBlock3D(ch, ch),\n            ResidualBlock3D(ch, ch),\n            ChannelAttention3D(ch),\n            SpatialAttention3D()\n        )\n\n        # Decoder\n        self.decoders = nn.ModuleList()\n        for i in range(depth):\n            self.decoders.append(DecoderBlock3D(ch, ch // 2))\n            ch //= 2\n\n        # Output\n        self.output_conv = nn.Sequential(\n            nn.Conv3d(ch, ch // 2, 3, padding=1),\n            nn.BatchNorm3d(ch // 2),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(ch // 2, 1, 1),\n        )\n\n        # Final squeeze: collapse Z dimension\n        self.z_squeeze = nn.Sequential(\n            nn.Conv3d(1, 16, (config.Z_DEPTH, 1, 1)),\n            nn.BatchNorm3d(16),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(16, 1, 1),\n        )\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm3d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        # x: (B, Z, H, W) -> (B, 1, Z, H, W)\n        x = x.unsqueeze(1)\n\n        # Initial features\n        x = self.input_conv(x)\n\n        # Encoder with skip connections\n        skips = []\n        for encoder in self.encoders:\n            skips.append(x)\n            x = encoder(x)\n\n        # Bottleneck\n        x = self.bottleneck(x)\n\n        # Decoder\n        for decoder, skip in zip(self.decoders, reversed(skips)):\n            x = decoder(x, skip)\n\n        # Output: (B, 1, Z, H, W)\n        x = self.output_conv(x)\n\n        # Squeeze Z dimension: (B, 1, Z, H, W) -> (B, 1, H, W)\n        x = self.z_squeeze(x).squeeze(2)\n\n        return x\n\n# ==============================================================================\n# LOSS FUNCTIONS\n# ==============================================================================\n\nclass VesuviusLoss(nn.Module):\n    \"\"\"Combined BCE + Dice + Focal Loss\"\"\"\n\n    def __init__(self, bce_weight=0.5, dice_weight=0.5, focal_weight=0.0):\n        super().__init__()\n        self.bce_weight = bce_weight\n        self.dice_weight = dice_weight\n        self.focal_weight = focal_weight\n        self.bce = nn.BCEWithLogitsLoss(reduction='mean')\n\n    def dice_loss(self, pred, target, smooth=1e-6):\n        pred = torch.sigmoid(pred)\n        intersection = (pred * target).sum(dim=(2, 3))\n        union = pred.sum(dim=(2, 3)) + target.sum(dim=(2, 3))\n        dice = (2. * intersection + smooth) / (union + smooth)\n        return 1 - dice.mean()\n\n    def focal_loss(self, pred, target, alpha=0.25, gamma=2.0):\n        bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n        pred_prob = torch.sigmoid(pred)\n        p_t = (pred_prob * target) + ((1 - pred_prob) * (1 - target))\n        alpha_t = alpha * target + (1 - alpha) * (1 - target)\n        focal = alpha_t * (1. - p_t) ** gamma * bce\n        return focal.mean()\n\n    def forward(self, pred, target):\n        loss = 0\n        if self.bce_weight > 0:\n            loss += self.bce_weight * self.bce(pred, target)\n        if self.dice_weight > 0:\n            loss += self.dice_weight * self.dice_loss(pred, target)\n        if self.focal_weight > 0:\n            loss += self.focal_weight * self.focal_loss(pred, target)\n        return loss\n\n# ==============================================================================\n# TRAINING\n# ==============================================================================\n\nclass Trainer:\n    \"\"\"Training loop with mixed precision and gradient clipping\"\"\"\n\n    def __init__(self, model, config):\n        self.model = model.to(config.DEVICE)\n        self.config = config\n        self.criterion = VesuviusLoss(\n            bce_weight=config.BCE_WEIGHT,\n            dice_weight=config.DICE_WEIGHT,\n            focal_weight=config.FOCAL_WEIGHT\n        )\n        self.optimizer = torch.optim.AdamW(\n            model.parameters(),\n            lr=config.LR,\n            weight_decay=config.WEIGHT_DECAY\n        )\n        self.scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n            self.optimizer, T_0=10, T_mult=2\n        )\n        self.scaler = GradScaler() if config.USE_AMP else None\n\n        self.best_val_f05 = 0.0\n        self.history = {'train_loss': [], 'val_loss': [], 'val_f05': []}\n\n    def train_epoch(self, dataloader):\n        self.model.train()\n        total_loss = 0.0\n\n        pbar = tqdm(dataloader, desc='Training')\n        for batch_idx, (volumes, labels, weights) in enumerate(pbar):\n            volumes = volumes.to(self.config.DEVICE, non_blocking=True)\n            labels = labels.to(self.config.DEVICE, non_blocking=True)\n\n            self.optimizer.zero_grad()\n\n            if self.scaler is not None:\n                with autocast():\n                    outputs = self.model(volumes)\n                    loss = self.criterion(outputs, labels)\n\n                self.scaler.scale(loss).backward()\n                self.scaler.unscale_(self.optimizer)\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.GRADIENT_CLIP)\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n            else:\n                outputs = self.model(volumes)\n                loss = self.criterion(outputs, labels)\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.GRADIENT_CLIP)\n                self.optimizer.step()\n\n            total_loss += loss.item()\n            pbar.set_postfix({'loss': loss.item()})\n\n            # Periodic memory cleanup\n            if batch_idx % 50 == 0:\n                free_memory()\n\n        return total_loss / len(dataloader)\n\n    @torch.no_grad()\n    def validate(self, dataloader):\n        self.model.eval()\n        total_loss = 0.0\n        all_preds = []\n        all_labels = []\n\n        pbar = tqdm(dataloader, desc='Validation')\n        for volumes, labels, weights in pbar:\n            volumes = volumes.to(self.config.DEVICE, non_blocking=True)\n            labels = labels.to(self.config.DEVICE, non_blocking=True)\n\n            if self.scaler is not None:\n                with autocast():\n                    outputs = self.model(volumes)\n                    loss = self.criterion(outputs, labels)\n            else:\n                outputs = self.model(volumes)\n                loss = self.criterion(outputs, labels)\n\n            total_loss += loss.item()\n\n            preds = torch.sigmoid(outputs).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.cpu().numpy())\n\n        all_preds = np.array(all_preds).squeeze()\n        all_labels = np.array(all_labels).squeeze()\n\n        metrics = compute_metrics(all_preds, all_labels)\n        metrics['loss'] = total_loss / len(dataloader)\n\n        return metrics\n\n    def fit(self, train_loader, val_loader, epochs):\n        for epoch in range(epochs):\n            print(f\"\\nEpoch {epoch+1}/{epochs}\")\n\n            train_loss = self.train_epoch(train_loader)\n            val_metrics = self.validate(val_loader)\n\n            self.scheduler.step()\n\n            self.history['train_loss'].append(train_loss)\n            self.history['val_loss'].append(val_metrics['loss'])\n            self.history['val_f05'].append(val_metrics['f05'])\n\n            print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_metrics['loss']:.4f}\")\n            print(f\"Val F0.5: {val_metrics['f05']:.4f} | Val Dice: {val_metrics['dice']:.4f} | Val IoU: {val_metrics['iou']:.4f}\")\n\n            # Save best model\n            if val_metrics['f05'] > self.best_val_f05:\n                self.best_val_f05 = val_metrics['f05']\n                torch.save({\n                    'epoch': epoch,\n                    'model_state_dict': self.model.state_dict(),\n                    'optimizer_state_dict': self.optimizer.state_dict(),\n                    'best_f05': self.best_val_f05,\n                }, config.OUTPUT_DIR / 'best_model.pth')\n                print(f\"✓ Saved best model (F0.5: {self.best_val_f05:.4f})\")\n\n        return self.history\n\n# ==============================================================================\n# INFERENCE & POST-PROCESSING\n# ==============================================================================\n\nclass InferenceEngine:\n    \"\"\"Memory-efficient tiled inference with overlap blending\"\"\"\n\n    def __init__(self, model, config):\n        self.model = model.to(config.DEVICE)\n        self.model.eval()\n        self.config = config\n\n    @torch.no_grad()\n    def predict_fragment(self, fragment_id, batch_size=2):\n        \"\"\"\n        Predict on full fragment using tiled inference.\n        Uses overlap-tile strategy with Gaussian weighting for blending.\n        \"\"\"\n        print(f\"\\nInferencing on fragment {fragment_id}...\")\n\n        # Load and preprocess volume\n        volume = VesuviusDataLoader.load_volume(fragment_id, config.Z_START, config.Z_END)\n        volume = VesuviusDataLoader.preprocess_volume(volume)\n        z, h, w = volume.shape\n\n        # Create output accumulator and weight map\n        output = np.zeros((h, w), dtype=np.float32)\n        weight_map = np.zeros((h, w), dtype=np.float32)\n\n        # Generate Gaussian weight kernel for smooth blending\n        tile_size = config.INFERENCE_TILE_SIZE\n        stride = config.INFERENCE_STRIDE\n\n        # Create 2D Gaussian kernel\n        y_coords, x_coords = np.mgrid[0:tile_size, 0:tile_size]\n        y_coords = y_coords - tile_size // 2\n        x_coords = x_coords - tile_size // 2\n        sigma = tile_size / 4\n        gaussian_kernel = np.exp(-(x_coords**2 + y_coords**2) / (2 * sigma**2))\n        gaussian_kernel = gaussian_kernel / gaussian_kernel.max()\n\n        # Collect all tiles\n        tiles = []\n        for y in range(0, h - tile_size + 1, stride):\n            for x in range(0, w - tile_size + 1, stride):\n                tiles.append((x, y))\n\n        # Process in batches\n        for batch_start in tqdm(range(0, len(tiles), batch_size), desc='Inference'):\n            batch_tiles = tiles[batch_start:batch_start + batch_size]\n\n            batch_volumes = []\n            positions = []\n\n            for x, y in batch_tiles:\n                tile_vol = volume[:, y:y+tile_size, x:x+tile_size]\n                tile_tensor = torch.from_numpy(tile_vol).float().unsqueeze(0)\n                batch_volumes.append(tile_tensor)\n                positions.append((x, y))\n\n            batch = torch.cat(batch_volumes, dim=0).to(config.DEVICE)\n\n            # Mixed precision inference\n            if config.USE_AMP:\n                with autocast():\n                    preds = self.model(batch)\n            else:\n                preds = self.model(batch)\n\n            preds = torch.sigmoid(preds).cpu().numpy().squeeze()\n\n            # Blend predictions with Gaussian weights\n            for i, (x, y) in enumerate(positions):\n                pred = preds[i] if preds.ndim == 3 else preds\n                output[y:y+tile_size, x:x+tile_size] += pred * gaussian_kernel\n                weight_map[y:y+tile_size, x:x+tile_size] += gaussian_kernel\n\n            # Memory cleanup\n            del batch, preds\n            if batch_start % 10 == 0:\n                free_memory()\n\n        # Normalize by weight map\n        output = output / (weight_map + 1e-7)\n\n        # Handle borders\n        output[weight_map < 0.1] = 0\n\n        return output\n\n    @staticmethod\n    def post_process(pred, threshold=0.5, min_size=50, kernel_size=3):\n        \"\"\"\n        Post-processing pipeline:\n        1. Threshold\n        2. Morphological closing (fill small holes)\n        3. Remove small components\n        4. Morphological opening (remove noise)\n        \"\"\"\n        # Threshold\n        binary = (pred > threshold).astype(np.uint8)\n\n        # Morphological closing to fill holes\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size))\n        binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel)\n\n        # Remove small components\n        num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(binary, connectivity=8)\n        for i in range(1, num_labels):\n            if stats[i, cv2.CC_STAT_AREA] < min_size:\n                binary[labels == i] = 0\n\n        # Morphological opening to remove noise\n        binary = cv2.morphologyEx(binary, cv2.MORPH_OPEN, kernel)\n\n        return binary.astype(np.float32)\n\n# ==============================================================================\n# VISUALIZATION\n# ==============================================================================\n\ndef visualize_results(fragment_id, pred, ground_truth=None, num_samples=5, save_path=None):\n    \"\"\"Create visualization comparing input, ground truth, and prediction\"\"\"\n\n    # Load a few representative slices from the volume\n    volume = VesuviusDataLoader.load_volume(fragment_id, 27, 37)  # Middle 10 slices\n    mid_slice = volume[len(volume)//2]\n\n    # Create figure\n    if ground_truth is not None:\n        fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n\n        # Row 1: Full views\n        axes[0, 0].imshow(mid_slice, cmap='gray')\n        axes[0, 0].set_title(f'Fragment {fragment_id} - Input Slice (Z=32)')\n        axes[0, 0].axis('off')\n\n        axes[0, 1].imshow(ground_truth, cmap='gray')\n        axes[0, 1].set_title('Ground Truth')\n        axes[0, 1].axis('off')\n\n        axes[0, 2].imshow(pred, cmap='jet', vmin=0, vmax=1)\n        axes[0, 2].set_title('Prediction (Probability)')\n        axes[0, 2].axis('off')\n\n        # Row 2: Zoomed samples\n        # Find regions with ink\n        if ground_truth is not None:\n            ink_coords = np.argwhere(ground_truth > 0.5)\n            if len(ink_coords) > 0:\n                # Sample random ink regions\n                indices = np.random.choice(len(ink_coords), min(num_samples, len(ink_coords)), replace=False)\n\n                for idx, coord_idx in enumerate(indices[:3]):\n                    cy, cx = ink_coords[coord_idx]\n                    size = 256\n                    y1, y2 = max(0, cy-size//2), min(pred.shape[0], cy+size//2)\n                    x1, x2 = max(0, cx-size//2), min(pred.shape[1], cx+size//2)\n\n                    axes[1, idx].imshow(pred[y1:y2, x1:x2], cmap='jet', vmin=0, vmax=1)\n                    axes[1, idx].set_title(f'Zoomed Region {idx+1}')\n                    axes[1, idx].axis('off')\n\n        # Hide unused subplots\n        for idx in range(3, 3):\n            axes[1, idx].axis('off')\n    else:\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n\n        axes[0].imshow(mid_slice, cmap='gray')\n        axes[0].set_title(f'Fragment {fragment_id} - Input Slice')\n        axes[0].axis('off')\n\n        axes[1].imshow(pred, cmap='jet', vmin=0, vmax=1)\n        axes[1].set_title('Prediction (Probability)')\n        axes[1].axis('off')\n\n        binary_pred = InferenceEngine.post_process(pred)\n        axes[2].imshow(binary_pred, cmap='gray')\n        axes[2].set_title('Post-processed (Binary)')\n        axes[2].axis('off')\n\n    plt.tight_layout()\n\n    if save_path:\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n        print(f\"Saved visualization to {save_path}\")\n    else:\n        plt.show()\n\n    plt.close()\n\ndef plot_training_history(history, save_path=None):\n    \"\"\"Plot training curves\"\"\"\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    epochs = range(1, len(history['train_loss']) + 1)\n\n    axes[0].plot(epochs, history['train_loss'], 'b-', label='Train Loss')\n    axes[0].plot(epochs, history['val_loss'], 'r-', label='Val Loss')\n    axes[0].set_xlabel('Epoch')\n    axes[0].set_ylabel('Loss')\n    axes[0].set_title('Training & Validation Loss')\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n\n    axes[1].plot(epochs, history['val_f05'], 'g-', marker='o')\n    axes[1].set_xlabel('Epoch')\n    axes[1].set_ylabel('F0.5 Score')\n    axes[1].set_title('Validation F0.5 Score')\n    axes[1].grid(True, alpha=0.3)\n\n    # Combined metrics\n    axes[2].plot(epochs, history['val_f05'], 'g-', label='F0.5')\n    if 'val_dice' in history:\n        axes[2].plot(epochs, history['val_dice'], 'b-', label='Dice')\n    if 'val_iou' in history:\n        axes[2].plot(epochs, history['val_iou'], 'r-', label='IoU')\n    axes[2].set_xlabel('Epoch')\n    axes[2].set_ylabel('Score')\n    axes[2].set_title('Validation Metrics')\n    axes[2].legend()\n    axes[2].grid(True, alpha=0.3)\n\n    plt.tight_layout()\n\n    if save_path:\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n        print(f\"Saved training history to {save_path}\")\n    else:\n        plt.show()\n\n    plt.close()\n\n# ==============================================================================\n# MAIN EXECUTION\n# ==============================================================================\n\ndef main():\n    \"\"\"Main execution pipeline\"\"\"\n\n    print(\"=\" * 80)\n    print(\"VESUVIUS INK DETECTION - 3D UNet Training Pipeline\")\n    print(\"=\" * 80)\n\n    # Create output directory\n    config.OUTPUT_DIR.mkdir(exist_ok=True)\n\n    # ======================================================================\n    # STEP 1: Data Preparation\n    # ======================================================================\n    print(\"\\n[STEP 1] Preparing datasets...\")\n\n    # Load fragment 2 and 3, split 80/20\n    # For fragment 2: use 80% of tiles for training, 20% for validation\n    # For fragment 3: use 80% of tiles for training, 20% for validation\n\n    all_train_tiles = []\n    all_val_tiles = []\n\n    for frag_id in config.TRAIN_FRAGMENTS:\n        print(f\"\\nProcessing Fragment {frag_id}...\")\n\n        # Create temporary dataset to get all tiles\n        temp_dataset = VesuviusDataset(\n            fragment_ids=[frag_id],\n            tile_size=config.TILE_SIZE,\n            stride=config.TILE_STRIDE,\n            is_training=True,\n            min_ink_percentage=0.0,  # Get all tiles first\n            augment=False,\n            cache_in_memory=False\n        )\n\n        # Split tiles 80/20\n        tile_indices = list(range(len(temp_dataset.tiles)))\n        train_idx, val_idx = train_test_split(\n            tile_indices, test_size=0.2, random_state=42,\n            stratify=[1 if t[4] > 0 else 0 for t in temp_dataset.tiles]  # Stratify by ink presence\n        )\n\n        # Store split tiles\n        frag_train_tiles = [temp_dataset.tiles[i] for i in train_idx]\n        frag_val_tiles = [temp_dataset.tiles[i] for i in val_idx]\n\n        all_train_tiles.extend(frag_train_tiles)\n        all_val_tiles.extend(frag_val_tiles)\n\n        print(f\"  Fragment {frag_id}: {len(frag_train_tiles)} train, {len(frag_val_tiles)} val tiles\")\n\n        # Clean up\n        del temp_dataset\n        free_memory()\n\n    print(f\"\\nTotal: {len(all_train_tiles)} train tiles, {len(all_val_tiles)} val tiles\")\n\n    # Create custom datasets with pre-split tiles\n    # We need to modify the approach - let's create datasets per fragment and use Subset\n\n    # Alternative: Create datasets per fragment, then use random split\n    train_datasets = []\n    val_datasets = []\n\n    for frag_id in config.TRAIN_FRAGMENTS:\n        # Full dataset for this fragment\n        full_dataset = VesuviusDataset(\n            fragment_ids=[frag_id],\n            tile_size=config.TILE_SIZE,\n            stride=config.TILE_STRIDE,\n            is_training=True,\n            min_ink_percentage=config.MIN_INK_PERCENTAGE,\n            augment=False,  # We'll handle augmentation in training\n            cache_in_memory=False\n        )\n\n        # Stratified split: ensure ink distribution is preserved\n        labels = [1 if t[4] > 0 else 0 for t in full_dataset.tiles]\n        train_idx, val_idx = train_test_split(\n            range(len(full_dataset)), test_size=0.2, random_state=42, stratify=labels\n        )\n\n        # Create subsets\n        from torch.utils.data import Subset\n        train_subset = Subset(full_dataset, train_idx)\n        val_subset = Subset(full_dataset, val_idx)\n\n        train_datasets.append(train_subset)\n        val_datasets.append(val_subset)\n\n    # Combine datasets\n    from torch.utils.data import ConcatDataset\n    train_dataset = ConcatDataset(train_datasets)\n    val_dataset = ConcatDataset(val_datasets)\n\n    print(f\"\\nFinal datasets: {len(train_dataset)} train, {len(val_dataset)} val\")\n\n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=True,\n        drop_last=True\n    )\n\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=True\n    )\n\n    # ======================================================================\n    # STEP 2: Model Initialization\n    # ======================================================================\n    print(\"\\n[STEP 2] Initializing model...\")\n\n    model = Vesuvius3DUNet(\n        in_channels=config.Z_DEPTH,\n        out_channels=config.OUT_CHANNELS,\n        base_channels=config.BASE_CHANNELS,\n        depth=config.DEPTH\n    )\n\n    # Count parameters\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Model parameters: {total_params:,} total, {trainable_params:,} trainable\")\n\n    # Test forward pass\n    test_input = torch.randn(1, config.Z_DEPTH, config.TILE_SIZE, config.TILE_SIZE).to(config.DEVICE)\n    model = model.to(config.DEVICE)\n    with torch.no_grad():\n        test_output = model(test_input)\n    print(f\"Test forward pass: input {test_input.shape} -> output {test_output.shape}\")\n\n    # Estimate memory usage\n    if torch.cuda.is_available():\n        mem_allocated = torch.cuda.memory_allocated() / 1e9\n        print(f\"GPU memory allocated: {mem_allocated:.2f} GB\")\n\n    # ======================================================================\n    # STEP 3: Training\n    # ======================================================================\n    print(\"\\n[STEP 3] Starting training...\")\n\n    trainer = Trainer(model, config)\n    history = trainer.fit(train_loader, val_loader, config.EPOCHS)\n\n    # Plot training history\n    plot_training_history(history, save_path=config.OUTPUT_DIR / 'training_history.png')\n\n    # ======================================================================\n    # STEP 4: Load Best Model & Evaluate on Validation\n    # ======================================================================\n    print(\"\\n[STEP 4] Loading best model for evaluation...\")\n\n    checkpoint = torch.load(config.OUTPUT_DIR / 'best_model.pth')\n    model.load_state_dict(checkpoint['model_state_dict'])\n    print(f\"Loaded best model from epoch {checkpoint['epoch']+1} with F0.5: {checkpoint['best_f05']:.4f}\")\n\n    # Evaluate on validation set\n    val_metrics = trainer.validate(val_loader)\n    print(f\"\\nFinal Validation Metrics:\")\n    print(f\"  F0.5:      {val_metrics['f05']:.4f}\")\n    print(f\"  Precision: {val_metrics['precision']:.4f}\")\n    print(f\"  Recall:    {val_metrics['recall']:.4f}\")\n    print(f\"  Dice:      {val_metrics['dice']:.4f}\")\n    print(f\"  IoU:       {val_metrics['iou']:.4f}\")\n\n    # ======================================================================\n    # STEP 5: Test on Fragment 1\n    # ======================================================================\n    print(\"\\n[STEP 5] Testing on Fragment 1...\")\n\n    inference = InferenceEngine(model, config)\n\n    # Predict on fragment 1\n    test_pred = inference.predict_fragment(config.TEST_FRAGMENT, batch_size=2)\n\n    # Load ground truth for fragment 1\n    test_gt = VesuviusDataLoader.load_inklabels(config.TEST_FRAGMENT)\n    test_gt = cv2.resize(test_gt, (test_pred.shape[1], test_pred.shape[0]), interpolation=cv2.INTER_NEAREST)\n\n    # Evaluate\n    test_metrics = compute_metrics(test_pred, test_gt)\n    print(f\"\\nTest Fragment 1 Metrics (Before Post-processing):\")\n    print(f\"  F0.5:      {test_metrics['f05']:.4f}\")\n    print(f\"  Precision: {test_metrics['precision']:.4f}\")\n    print(f\"  Recall:    {test_metrics['recall']:.4f}\")\n    print(f\"  Dice:      {test_metrics['dice']:.4f}\")\n    print(f\"  IoU:       {test_metrics['iou']:.4f}\")\n\n    # Post-process\n    test_pred_pp = InferenceEngine.post_process(\n        test_pred, \n        threshold=0.5, \n        min_size=config.MIN_COMPONENT_SIZE,\n        kernel_size=config.MORPH_KERNEL_SIZE\n    )\n\n    test_metrics_pp = compute_metrics(test_pred_pp, test_gt)\n    print(f\"\\nTest Fragment 1 Metrics (After Post-processing):\")\n    print(f\"  F0.5:      {test_metrics_pp['f05']:.4f}\")\n    print(f\"  Precision: {test_metrics_pp['precision']:.4f}\")\n    print(f\"  Recall:    {test_metrics_pp['recall']:.4f}\")\n    print(f\"  Dice:      {test_metrics_pp['dice']:.4f}\")\n    print(f\"  IoU:       {test_metrics_pp['iou']:.4f}\")\n\n    # ======================================================================\n    # STEP 6: Visualization\n    # ======================================================================\n    print(\"\\n[STEP 6] Generating visualizations...\")\n\n    # Visualize test results\n    visualize_results(\n        config.TEST_FRAGMENT,\n        test_pred,\n        ground_truth=test_gt,\n        num_samples=config.NUM_VIS_SAMPLES,\n        save_path=config.OUTPUT_DIR / 'test_fragment1_visualization.png'\n    )\n\n    # Also visualize validation fragment\n    for frag_id in config.VAL_FRAGMENTS[:1]:  # Just first one\n        val_pred = inference.predict_fragment(frag_id, batch_size=2)\n        val_gt = VesuviusDataLoader.load_inklabels(frag_id)\n        val_gt = cv2.resize(val_gt, (val_pred.shape[1], val_pred.shape[0]), interpolation=cv2.INTER_NEAREST)\n\n        visualize_results(\n            frag_id,\n            val_pred,\n            ground_truth=val_gt,\n            num_samples=config.NUM_VIS_SAMPLES,\n            save_path=config.OUTPUT_DIR / f'val_fragment{frag_id}_visualization.png'\n        )\n\n    # ======================================================================\n    # STEP 7: Save Predictions\n    # ======================================================================\n    print(\"\\n[STEP 7] Saving predictions...\")\n\n    # Save raw probability map\n    np.save(config.OUTPUT_DIR / 'test_fragment1_prediction.npy', test_pred)\n\n    # Save post-processed binary mask\n    np.save(config.OUTPUT_DIR / 'test_fragment1_prediction_postprocessed.npy', test_pred_pp)\n\n    # Save as PNG for visualization\n    plt.imsave(config.OUTPUT_DIR / 'test_fragment1_prediction.png', test_pred, cmap='jet', vmin=0, vmax=1)\n    plt.imsave(config.OUTPUT_DIR / 'test_fragment1_prediction_pp.png', test_pred_pp, cmap='gray')\n\n    # Create comparison image\n    fig, axes = plt.subplots(1, 3, figsize=(24, 8))\n\n    # Load a representative slice\n    vol = VesuviusDataLoader.load_volume(config.TEST_FRAGMENT, 30, 35)\n    axes[0].imshow(vol[2], cmap='gray')\n    axes[0].set_title(f'Fragment {config.TEST_FRAGMENT} - Input (Z=32)')\n    axes[0].axis('off')\n\n    axes[1].imshow(test_gt, cmap='gray')\n    axes[1].set_title('Ground Truth')\n    axes[1].axis('off')\n\n    axes[2].imshow(test_pred_pp, cmap='gray')\n    axes[2].set_title(f'Prediction (F0.5: {test_metrics_pp[\"f05\"]:.4f})')\n    axes[2].axis('off')\n\n    plt.tight_layout()\n    plt.savefig(config.OUTPUT_DIR / 'test_comparison.png', dpi=200, bbox_inches='tight')\n    plt.close()\n\n    print(f\"\\nAll outputs saved to: {config.OUTPUT_DIR}\")\n    print(\"\\n\" + \"=\" * 80)\n    print(\"PIPELINE COMPLETE\")\n    print(\"=\" * 80)\n\n    return model, history, test_metrics_pp\n\n# ==============================================================================\n# RUN\n# ==============================================================================\n\nif __name__ == \"__main__\":\n    # Set random seeds for reproducibility\n    torch.manual_seed(42)\n    np.random.seed(42)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(42)\n        # Enable cudnn benchmarking for faster training\n        torch.backends.cudnn.benchmark = True\n\n    # Run main pipeline\n    model, history, test_metrics = main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\n# ============================================================\n#  VESUVIUS INK DETECTION — v10 + DDPM  (FINAL SPEED FIX)\n#\n#  ANOMALY MAP TIMING HISTORY & FIXES:\n#  ┌──────────────────────────────────────────────────────┐\n#  │ v1 (sequential DDPM):  34k patches × 40 steps        │\n#  │   = 1.36M fwd passes  →  6+ hours per fragment       │\n#  │                                                       │\n#  │ v2 (F.unfold + DDIM):  OOM — unfold on full 14k×9k  │\n#  │   fragment = 38 GB tensor before any GPU work starts  │\n#  │                                                       │\n#  │ v3 (coord_chunk loop): Still slow — 279s/chunk due   │\n#  │   to two bottlenecks:                                 │\n#  │   (a) 4000×17 Python-level numpy slices per chunk    │\n#  │   (b) ANOM_BATCH=32 → too many kernel launches for   │\n#  │       a tiny 280k-param model                         │\n#  │                                                       │\n#  │ THIS VERSION (v4): All three issues fixed:            │\n#  │   (a) sliding_window_view → vectorized gather          │\n#  │       replaces 68k Python slices with 17 numpy ops    │\n#  │   (b) ANOM_BATCH=128 → 4× fewer kernel launches      │\n#  │   (c) autocast hoisted outside the DDIM step loop     │\n#  │   Expected: ~10-20 minutes per fragment on P100       │\n#  └──────────────────────────────────────────────────────┘\n# ============================================================\n\nimport os, gc, cv2, math, warnings, random, time\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nfrom numpy.lib.stride_tricks import sliding_window_view\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:64\"\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = False\n\n# ── paths ────────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\nDEVICE      = 'cuda' if torch.cuda.is_available() else 'cpu'\nUSE_AMP     = (DEVICE == 'cuda')\nZ_SLICES    = [20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36]\nN_CH        = len(Z_SLICES)   # 17\nPATCH_SIZE  = 224\nTEMPERATURE = 1.3\n\nprint(f\"Device : {DEVICE}  |  N_CH={N_CH}  |  Patch={PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. CT CACHE\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel()[::20])\n    assert raw, f\"No slices in {vol_dir}\"\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: {len(cache)} slices, {mb:.0f} MB\")\n    return cache\n\n\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H, W):\n                m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n            return (m > 0).astype(np.uint8)\n    z = list(cache.keys())[len(cache)//2]\n    return (cache[z] > 0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  2. DDPM HYPERPARAMETERS & SCHEDULE\n# ════════════════════════════════════════════════════════════\nDDPM_T           = 200\nDDPM_PATCH_XY    = 64\nDDPM_BETA_START  = 1e-4\nDDPM_BETA_END    = 0.02\nDDPM_T_NOISE     = 100\nDDIM_STEPS       = 5\nDDPM_EPOCHS      = 10\nDDPM_BATCH       = 8\nANOM_BATCH       = 256    # large batch: DDPM 280k params, P100 has headroom\nDDPM_LR          = 2e-4\nDDPM_BASE_CH     = 16\nDDPM_MAX_PATCHES = 2_000\n# ANOM_STRIDE = no overlap -> Frag2 ~24k patches instead of 95k (4x faster)\nANOM_STRIDE      = DDPM_PATCH_XY   # = 64\nCOORD_CHUNK      = 4096\n\n\ndef build_ddpm_schedule(T, beta_start, beta_end, device):\n    betas          = torch.linspace(beta_start, beta_end, T, device=device)\n    alphas         = 1. - betas\n    alpha_bar      = torch.cumprod(alphas, dim=0)\n    alpha_bar_prev = F.pad(alpha_bar[:-1], (1, 0), value=1.0)\n    sqrt_ab        = alpha_bar.sqrt()\n    sqrt_1mab      = (1. - alpha_bar).sqrt()\n    post_var       = betas * (1. - alpha_bar_prev) / (1. - alpha_bar + 1e-8)\n    return dict(betas=betas, alphas=alphas, alpha_bar=alpha_bar,\n                alpha_bar_prev=alpha_bar_prev, sqrt_ab=sqrt_ab,\n                sqrt_1mab=sqrt_1mab, post_var=post_var)\n\n\n# ════════════════════════════════════════════════════════════\n#  3. 3D DDPM UNET  (~280k params, fits easily on P100)\n# ════════════════════════════════════════════════════════════\nclass GroupConv3d(nn.Module):\n    def __init__(self, cin, cout, k=3, groups=4, stride=1):\n        super().__init__()\n        g = min(groups, cin, cout)\n        while cin % g != 0 or cout % g != 0:\n            g -= 1\n        g = max(g, 1)\n        self.c = nn.Sequential(\n            nn.Conv3d(cin, cout, k, padding=k//2, stride=stride,\n                      groups=g, bias=False),\n            nn.GroupNorm(min(8, cout), cout),\n            nn.SiLU(inplace=True))\n    def forward(self, x): return self.c(x)\n\n\nclass ResBlock3d(nn.Module):\n    def __init__(self, ch, t_emb_dim=64, groups=4):\n        super().__init__()\n        self.c1   = GroupConv3d(ch, ch, groups=groups)\n        self.c2   = GroupConv3d(ch, ch, groups=groups)\n        self.temb = nn.Linear(t_emb_dim, ch)\n        self.norm = nn.GroupNorm(min(8, ch), ch)\n    def forward(self, x, t_emb):\n        h = self.c1(x) + self.temb(t_emb).view(t_emb.shape[0], -1, 1, 1, 1)\n        return x + self.c2(self.norm(h))\n\n\nclass UNet3D_DDPM(nn.Module):\n    \"\"\"\n    Encoder:  1→c, c→2c (stride-2), 2c→4c (stride-2), bottleneck 4c\n    Decoder:  4c→2c (deconv), cat skip 2c → proj 2c,\n              2c→c  (deconv), cat skip c  → proj c,  → out 1\n    \"\"\"\n    def __init__(self, in_ch=1, base_ch=DDPM_BASE_CH, T=DDPM_T, groups=4):\n        super().__init__()\n        t_dim = 64; c = base_ch\n        self.t_embed = nn.Sequential(\n            nn.Embedding(T, t_dim), nn.Linear(t_dim, t_dim),\n            nn.SiLU(), nn.Linear(t_dim, t_dim))\n        self.enc1  = nn.Sequential(GroupConv3d(in_ch, c, groups=1),\n                                   GroupConv3d(c, c, groups=min(groups,c)))\n        self.down1 = nn.Conv3d(c, c*2, 2, stride=2, bias=False)\n        self.res1  = ResBlock3d(c*2, t_dim, min(groups,c*2))\n        self.enc2  = GroupConv3d(c*2, c*2, groups=min(groups,c*2))\n        self.down2 = nn.Conv3d(c*2, c*4, 2, stride=2, bias=False)\n        self.res2  = ResBlock3d(c*4, t_dim, min(groups,c*4))\n        self.bot   = ResBlock3d(c*4, t_dim, min(groups,c*4))\n        self.up2   = nn.ConvTranspose3d(c*4, c*2, 2, stride=2, bias=False)\n        self.dres2 = ResBlock3d(c*4, t_dim, min(groups,c*4))\n        self.dproj2= GroupConv3d(c*4, c*2, k=1, groups=1)\n        self.up1   = nn.ConvTranspose3d(c*2, c, 2, stride=2, bias=False)\n        self.dres1 = ResBlock3d(c*2, t_dim, min(groups,c*2))\n        self.dproj1= GroupConv3d(c*2, c, k=1, groups=1)\n        self.out   = nn.Conv3d(c, in_ch, 1)\n\n    def forward(self, x, t):\n        te = self.t_embed(t)\n        e1 = self.enc1(x)\n        e2 = self.res1(self.down1(e1), te)\n        e3 = self.res2(self.down2(self.enc2(e2)), te)\n        b  = self.bot(e3, te)\n        d2 = self.up2(b)\n        if d2.shape != e2.shape:\n            d2 = F.interpolate(d2, size=e2.shape[2:], mode='trilinear', align_corners=False)\n        d2 = self.dproj2(self.dres2(torch.cat([d2, e2], 1), te))\n        d1 = self.up1(d2)\n        if d1.shape != e1.shape:\n            d1 = F.interpolate(d1, size=e1.shape[2:], mode='trilinear', align_corners=False)\n        d1 = self.dproj1(self.dres1(torch.cat([d1, e1], 1), te))\n        return self.out(d1)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. DDPM TRAINING HELPERS\n# ════════════════════════════════════════════════════════════\ndef q_sample(x0, t, sched):\n    \"\"\"Forward diffusion: add noise at timestep t.\"\"\"\n    noise   = torch.randn_like(x0)\n    sqrt_ab = sched['sqrt_ab'][t].view(-1, 1, 1, 1, 1)\n    sqrt_1m = sched['sqrt_1mab'][t].view(-1, 1, 1, 1, 1)\n    return sqrt_ab * x0 + sqrt_1m * noise, noise\n\n\n# ════════════════════════════════════════════════════════════\n#  5. FAST ANOMALY MAP via DDIM + sliding_window_view\n#\n#  Three-level speed optimisation:\n#\n#  LEVEL 1 — DDIM (5 steps, not 100)\n#    DDPM reverse requires stepping t=99→98→...→0 (100 calls).\n#    DDIM skips to t=99→74→49→24→0 (5 calls) using:\n#      x_{t-1} = √ᾱ_{t-1}·x̂₀ + √(1-ᾱ_{t-1})·ε_θ\n#    No noise term → deterministic → 5 steps enough for MSE signal.\n#    Speedup: 20×\n#\n#  LEVEL 2 — Batched inference (ANOM_BATCH=128)\n#    Each GPU call processes 128 patches simultaneously.\n#    The DDPM is 280k params; 128×(1,17,64,64) fp16 ≈ 70 MB.\n#    Reduces Python↔GPU round-trips by 128×.\n#    Speedup: ~128× in kernel-launch overhead.\n#\n#  LEVEL 3 — sliding_window_view patch extraction\n#    sliding_window_view(arr, (p,p)) returns a zero-copy view\n#    of shape (H-p+1, W-p+1, p, p). Gathering 8192 patches\n#    across 17 z-slices = 17 vectorized fancy-index ops instead\n#    of 8192×17 = 139,264 individual Python-level slice calls.\n#    Speedup: ~100× in CPU patch-building time.\n#\n#  MEMORY SAFETY:\n#    Never materialises more than COORD_CHUNK patches at once.\n#    Peak CPU RAM per chunk: 8192×17×64×64×4 = 2.3 GB — safe.\n#    Peak GPU RAM per step:  128×1×17×64×64×2 (fp16) ≈ 18 MB.\n# ════════════════════════════════════════════════════════════\n\n@torch.no_grad()\ndef compute_anomaly_map_fast(ddpm_model, cache, z_list, sched,\n                              patch_xy=DDPM_PATCH_XY,\n                              t_noise=DDPM_T_NOISE,\n                              ddim_steps=DDIM_STEPS,\n                              anom_batch=ANOM_BATCH,\n                              papyrus_mask=None,\n                              anom_stride=None,\n                              coord_chunk=COORD_CHUNK,\n                              device='cuda'):\n    \"\"\"\n    Fast anomaly map: DDIM reconstruction error per spatial location.\n\n    Speed fixes applied here:\n      1. anom_stride = patch_xy (no overlap) → 4-8x fewer patches\n         than stride=patch_xy//2. The anomaly map is upsampled to\n         full resolution at the end via bilinear interpolation.\n      2. ANOM_BATCH=256 → fewer kernel launches for the tiny DDPM.\n      3. Vectorized patch extraction via sliding_window_view fancy-index.\n      4. Vectorized accumulation: build Y/X index arrays once, then\n         use a single np.add.at call per chunk instead of a Python loop.\n      5. All DDIM steps run under one autocast context — no repeated\n         context-manager entry/exit overhead.\n    \"\"\"\n    if anom_stride is None:\n        anom_stride = patch_xy   # default: no overlap\n\n    ddpm_model.eval()\n    H, W = next(iter(cache.values())).shape\n\n    # ── coord list (integers only, negligible RAM) ─────────────────\n    coords = [(y, x)\n              for y in range(0, H - patch_xy + 1, anom_stride)\n              for x in range(0, W - patch_xy + 1, anom_stride)]\n    # cover right/bottom edges\n    if not coords or coords[-1][0] < H - patch_xy:\n        coords += [(H - patch_xy, x)\n                   for x in range(0, W - patch_xy + 1, anom_stride)]\n    if not coords or coords[-1][1] < W - patch_xy:\n        coords += [(y, W - patch_xy)\n                   for y in range(0, H - patch_xy + 1, anom_stride)]\n    # deduplicate while preserving order\n    seen = set(); coords = [c for c in coords if not (c in seen or seen.add(c))]\n\n    if papyrus_mask is not None:\n        half   = patch_xy // 2\n        coords = [(y, x) for (y, x) in coords\n                  if papyrus_mask[min(y+half, H-1), min(x+half, W-1)] > 0]\n\n    n_patches = len(coords)\n    print(f\"    {n_patches} patches  stride={anom_stride}  patch={patch_xy}\")\n\n    # work on a coarse score map (one cell per patch, no overlap)\n    # → accumulate into a small array, then upsample\n    n_rows = math.ceil(H / anom_stride)\n    n_cols = math.ceil(W / anom_stride)\n    # We'll accumulate directly into full-res maps — still fast because\n    # n_patches is now small (~6k for Frag2 vs 95k before)\n    score_map = np.zeros((H, W), np.float32)\n    count_map = np.zeros((H, W), np.float32)\n\n    # ── DDIM schedule ──────────────────────────────────────────────\n    t_seq   = np.round(np.linspace(t_noise-1, 0, ddim_steps)).astype(int).tolist()\n    t_pairs = list(zip(t_seq[:-1], t_seq[1:]))\n    t_pairs.append((t_seq[-1], 0))\n    ab_vals = sched['alpha_bar'].to(device)\n\n    # ── sliding window views (zero-copy, one per z-slice) ──────────\n    z_wins  = [sliding_window_view(cache[z], (patch_xy, patch_xy))\n               for z in z_list]\n    max_y   = H - patch_xy\n    max_x   = W - patch_xy\n\n    n_chunks = math.ceil(n_patches / coord_chunk)\n    t0       = time.time()\n\n    for c_idx, c_start in enumerate(range(0, n_patches, coord_chunk)):\n        chunk  = coords[c_start : c_start + coord_chunk]\n        cn     = len(chunk)\n        ys     = np.clip([c[0] for c in chunk], 0, max_y).astype(np.int32)\n        xs     = np.clip([c[1] for c in chunk], 0, max_x).astype(np.int32)\n\n        # ── vectorized patch gather (17 fancy-index ops, no loop) ──\n        chunk_vol = np.empty((cn, len(z_list), patch_xy, patch_xy), dtype=np.float32)\n        for ci, win in enumerate(z_wins):\n            chunk_vol[:, ci] = win[ys, xs]   # float16→float32 implicit\n        chunk_vol = chunk_vol * 2.0 - 1.0\n\n        chunk_mse = np.empty((cn, patch_xy, patch_xy), dtype=np.float32)\n\n        # ── DDIM reverse in large batches ──────────────────────────\n        for b_s in range(0, cn, anom_batch):\n            b_e = min(b_s + anom_batch, cn)\n            x0  = torch.from_numpy(chunk_vol[b_s:b_e]).unsqueeze(1).to(device)\n            noise = torch.randn_like(x0)\n            x_t   = (ab_vals[t_noise-1].sqrt() * x0\n                     + (1 - ab_vals[t_noise-1]).sqrt() * noise)\n            del noise\n\n            # all DDIM steps under ONE autocast — no repeated ctx overhead\n            with autocast(enabled=USE_AMP):\n                for t_cur, t_prev in t_pairs:\n                    t_vec  = torch.full((x_t.shape[0],), t_cur,\n                                       dtype=torch.long, device=device)\n                    eps    = ddpm_model(x_t, t_vec)\n                    ab_c   = ab_vals[t_cur]\n                    ab_p   = ab_vals[t_prev]\n                    x0p    = ((x_t - (1-ab_c).sqrt()*eps) /\n                              (ab_c.sqrt()+1e-8)).clamp(-1., 1.)\n                    x_t    = ab_p.sqrt()*x0p + (1-ab_p).sqrt()*eps\n                    del eps, x0p, t_vec\n\n            mse = ((x_t.float() - x0.float())**2).mean(dim=(1,2))\n            chunk_mse[b_s:b_e] = mse.cpu().numpy()\n            del x0, x_t, mse\n\n        # ── vectorized accumulation ────────────────────────────────\n        # Build (cn, patch_xy, patch_xy) MSE arrays indexed by\n        # row/col offsets, then use np.add.at for the score map.\n        # This replaces the per-patch Python loop entirely.\n        for i in range(cn):\n            y, x = int(ys[i]), int(xs[i])\n            score_map[y:y+patch_xy, x:x+patch_xy] += chunk_mse[i]\n            count_map[y:y+patch_xy, x:x+patch_xy] += 1.\n\n        del chunk_vol, chunk_mse\n        if device == 'cuda': torch.cuda.empty_cache()\n\n        elapsed = time.time() - t0\n        done    = c_idx + 1\n        eta     = elapsed / done * (n_chunks - done)\n        print(f\"\\r    chunk {done}/{n_chunks} | \"\n              f\"{elapsed/60:.1f}min elapsed | ETA {eta/60:.1f}min\", end='')\n\n    print()\n    score_map /= (count_map + 1e-8)\n    covered = count_map > 0\n    if covered.any():\n        s_min = score_map[covered].min()\n        s_max = score_map[covered].max()\n        score_map = np.where(covered,\n                             (score_map-s_min)/(s_max-s_min+1e-8), 0.)\n    return score_map.astype(np.float32)\n\n\n\n# ════════════════════════════════════════════════════════════\n#  6. DDPM TRAINING DATASET (no-ink patches only)\n# ════════════════════════════════════════════════════════════\nclass NoInkDataset3D(Dataset):\n    def __init__(self, frag_paths, z_list, patch_xy=DDPM_PATCH_XY,\n                 max_patches=DDPM_MAX_PATCHES):\n        self.patch_xy = patch_xy; self.z_list = z_list\n        self.patches  = []; self.caches = {}\n\n        for fp in frag_paths:\n            cache = load_slice_cache(fp, z_list)\n            self.caches[fp] = cache\n            H, W = next(iter(cache.values())).shape\n            pap  = load_papyrus_mask(fp, H, W, cache)\n            ink  = cv2.imread(os.path.join(fp, 'inklabels.png'), 0)\n            if ink is None:\n                no_ink = pap\n            else:\n                if ink.shape != (H, W):\n                    ink = cv2.resize(ink, (W, H), interpolation=cv2.INTER_NEAREST)\n                no_ink = ((ink == 0) & (pap > 0)).astype(np.uint8)\n\n            stride = patch_xy // 2\n            coords = [(fp, y, x)\n                      for y in range(0, H - patch_xy + 1, stride)\n                      for x in range(0, W - patch_xy + 1, stride)\n                      if no_ink[y:y+patch_xy, x:x+patch_xy].mean() > 0.9]\n            np.random.shuffle(coords)\n            self.patches.extend(coords)\n            print(f\"  NoInk [{os.path.basename(fp)}]: {len(coords)} patches\")\n\n        if max_patches > 0 and len(self.patches) > max_patches:\n            np.random.shuffle(self.patches)\n            self.patches = self.patches[:max_patches]\n        print(f\"  Total no-ink patches: {len(self.patches)}\")\n\n    def __len__(self): return len(self.patches)\n\n    def __getitem__(self, idx):\n        fp, y, x = self.patches[idx]\n        vol = np.stack([\n            self.caches[fp][z][y:y+self.patch_xy,\n                               x:x+self.patch_xy].astype(np.float32)\n            for z in self.z_list], axis=0)\n        return torch.from_numpy(vol * 2. - 1.).unsqueeze(0).float()\n\n\ndef train_ddpm(model, dataset, sched, n_epochs, lr, device, batch_size):\n    dl       = DataLoader(dataset, batch_size=batch_size, shuffle=True,\n                          num_workers=0, pin_memory=False)\n    opt      = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)\n    sched_lr = optim.lr_scheduler.CosineAnnealingLR(opt, T_max=n_epochs, eta_min=1e-5)\n    scaler   = GradScaler(enabled=USE_AMP)\n    print(f\"\\n{'='*60}\")\n    print(f\"DDPM Training: {n_epochs} epochs | {len(dl)} batches\")\n    print(f\"  T={DDPM_T}  t_noise={DDPM_T_NOISE}  ddim_steps={DDIM_STEPS}\")\n    print(f\"  anom_batch={ANOM_BATCH}  coord_chunk={COORD_CHUNK}\")\n    print(f\"{'='*60}\")\n    model.train()\n    for ep in range(n_epochs):\n        tot = 0.; steps = 0; opt.zero_grad()\n        for x0 in tqdm(dl, desc=f'DDPM ep{ep+1}', leave=False):\n            x0 = x0.to(device)\n            t  = torch.randint(0, DDPM_T, (x0.shape[0],), device=device)\n            x_t, noise = q_sample(x0, t, sched)\n            with autocast(enabled=USE_AMP):\n                loss = F.mse_loss(model(x_t, t), noise)\n            scaler.scale(loss).backward()\n            scaler.unscale_(opt)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.)\n            scaler.step(opt); scaler.update(); opt.zero_grad()\n            tot += loss.item(); steps += 1\n            del x0, x_t, noise, loss\n        sched_lr.step()\n        print(f\"  ep{ep+1}/{n_epochs} loss={tot/max(steps,1):.5f} \"\n              f\"lr={opt.param_groups[0]['lr']:.1e}\")\n    return model\n\n\n# ════════════════════════════════════════════════════════════\n#  7. V10 SEGMENTATION MODEL\n# ════════════════════════════════════════════════════════════\nSTRIDE_TR    = 112\nSTRIDE_INF   = 56\nBATCH_SIZE   = 2\nGRAD_ACCUM   = 8\nEPOCHS       = 35\nLR           = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 10\nVAL_SPLIT    = 0.30\nINK_MIN_POS  = 0.02\nNEG_RATIO    = 0.5\nCH_DROP_PROB = 0.3\nCH_DROP_MAX  = 3\nN_CH_SEG     = N_CH + 1   # 17 CT + 1 anomaly = 18\n\n\nclass SheetSurfaceDetector(nn.Module):\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, 32, 3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.ReLU(inplace=True),\n            nn.Conv2d(16, 1, 1))\n    def forward(self, x): return self.net(x)\n\n\ndef token_merge(tokens, saliency, merge_ratio=0.25):\n    B, C, N = tokens.shape\n    _, sort_idx = saliency.squeeze(1).sort(dim=1)\n    n_merge = int(N * merge_ratio)\n    if n_merge % 2 == 1: n_merge -= 1\n    merge_idx = sort_idx[:, :n_merge]; keep_idx = sort_idx[:, n_merge:]\n    def gather(t, idx):\n        return t.gather(2, idx.unsqueeze(1).expand(-1, C, -1))\n    t_m  = gather(tokens, merge_idx); t_k = gather(tokens, keep_idx)\n    return (torch.cat([t_k, (t_m[:,:,0::2]+t_m[:,:,1::2])/2], 2),\n            (keep_idx, merge_idx, n_merge, N))\n\n\ndef token_unmerge(tokens_out, info, C):\n    keep_idx, merge_idx, n_merge, N = info\n    B      = tokens_out.shape[0]\n    n_keep = tokens_out.shape[2] - n_merge // 2\n    t_keep = tokens_out[:,:,:n_keep]\n    t_exp  = tokens_out[:,:,n_keep:].repeat_interleave(2, dim=2)\n    out    = torch.zeros(B, C, N, device=tokens_out.device, dtype=tokens_out.dtype)\n    out.scatter_(2, keep_idx.unsqueeze(1).expand(-1,C,-1), t_keep)\n    out.scatter_(2, merge_idx.unsqueeze(1).expand(-1,C,-1), t_exp)\n    return out\n\n\nclass AxialAttention(nn.Module):\n    def __init__(self, dim, num_heads=8, axis='x', dropout=0.1):\n        super().__init__()\n        self.axis = axis; self.num_heads = num_heads\n        self.head_dim = dim // num_heads; self.scale = self.head_dim ** -0.5\n        self.qkv  = nn.Linear(dim, dim*3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.drop = nn.Dropout(dropout)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        if self.axis == 'x': x_r = x.permute(0,2,3,1).reshape(B*H, W, C)\n        else:                 x_r = x.permute(0,3,2,1).reshape(B*W, H, C)\n        res = x_r; BN, L, _ = x_r.shape\n        qkv = self.qkv(self.norm(x_r)).reshape(BN,L,3,self.num_heads,self.head_dim)\n        q,k,v = qkv.permute(2,0,3,1,4).unbind(0)\n        attn  = (q @ k.transpose(-2,-1)) * self.scale\n        pos   = torch.arange(L, dtype=torch.float32, device=x.device)\n        attn  = attn - torch.log(\n            (pos.unsqueeze(0) - pos.unsqueeze(1)).abs().float() + 1.\n        ).unsqueeze(0).unsqueeze(0)\n        attn  = self.drop(attn.softmax(-1))\n        out   = (attn @ v).transpose(1,2).reshape(BN,L,C)\n        out   = self.proj(out) + res\n        if self.axis == 'x': return out.reshape(B,H,W,C).permute(0,3,1,2)\n        else:                 return out.reshape(B,W,H,C).permute(0,3,2,1)\n\n\nclass AxialTransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=8, mlp_ratio=4, dropout=0.1):\n        super().__init__()\n        self.attn_x = AxialAttention(dim, num_heads, 'x', dropout)\n        self.attn_y = AxialAttention(dim, num_heads, 'y', dropout)\n        self.norm   = nn.LayerNorm(dim)\n        mlp_dim     = int(dim*mlp_ratio)\n        self.ffn    = nn.Sequential(\n            nn.Linear(dim, mlp_dim), nn.GELU(), nn.Dropout(dropout),\n            nn.Linear(mlp_dim, dim), nn.Dropout(dropout))\n    def forward(self, x):\n        x = self.attn_x(x); x = self.attn_y(x)\n        B,C,H,W = x.shape\n        xf = self.ffn(self.norm(x.permute(0,2,3,1).reshape(-1,C)))\n        return x + xf.reshape(B,H,W,C).permute(0,3,1,2)\n\n\nclass VesuviusV10(nn.Module):\n    def __init__(self, n_ch=N_CH_SEG, enc_dim=512,\n                 n_transformer_blocks=4, num_heads=8):\n        super().__init__()\n        self.backbone = smp.Unet(\n            encoder_name='resnet34', encoder_weights='imagenet',\n            in_channels=n_ch, classes=1, decoder_attention_type='scse')\n        self.transformer_blocks = nn.Sequential(*[\n            AxialTransformerBlock(enc_dim, num_heads, 4, 0.1)\n            for _ in range(n_transformer_blocks)])\n        self.sheet_detector  = SheetSurfaceDetector(n_ch)\n        self._enc_dim        = enc_dim\n        self.bottleneck_proj = nn.Sequential(\n            nn.Conv2d(enc_dim, enc_dim, 1, bias=False),\n            nn.BatchNorm2d(enc_dim), nn.GELU())\n        self.ds_head3 = nn.Conv2d(256, 1, 1)\n        self.ds_head2 = nn.Conv2d(128, 1, 1)\n        self.ds_head1 = nn.Conv2d( 64, 1, 1)\n\n    def forward(self, x):\n        B = x.shape[0]\n        sal    = self.sheet_detector(x)\n        feats  = self.backbone.encoder(x)\n        bn     = feats[-1]\n        bH,bW  = bn.shape[2], bn.shape[3]; N = bH*bW\n        sal_d  = F.adaptive_avg_pool2d(torch.sigmoid(sal),(bH,bW)).reshape(B,1,N)\n        tok    = bn.reshape(B, self._enc_dim, N)\n        tok_m, uinfo = token_merge(tok, sal_d, 0.25)\n        M      = tok_m.shape[2]; sq = int(math.ceil(math.sqrt(M)))\n        pad    = sq*sq - M\n        if pad: tok_m = F.pad(tok_m,(0,pad))\n        tok_2d = self.transformer_blocks(tok_m.reshape(B, self._enc_dim, sq, sq))\n        tok_f  = tok_2d.reshape(B, self._enc_dim, sq*sq)[:,:,:M]\n        bn_out = self.bottleneck_proj(\n            token_unmerge(tok_f, uinfo, self._enc_dim).reshape(B,self._enc_dim,bH,bW) + bn)\n        feats_m = list(feats); feats_m[-1] = bn_out\n        dec_out = self.backbone.decoder(*feats_m)\n        main    = self.backbone.segmentation_head(dec_out)\n        ds3=ds2=ds1=None\n        try:\n            db = self.backbone.decoder.blocks\n            if len(db)>=1: f0=db[0](feats_m[-1],feats_m[-2]); ds3=self.ds_head3(f0)\n            if len(db)>=2: f1=db[1](f0,feats_m[-3]);          ds2=self.ds_head2(f1)\n            if len(db)>=3: f2=db[2](f1,feats_m[-4]);          ds1=self.ds_head1(f2)\n        except Exception: pass\n        return main, sal, ds3, ds2, ds1\n\n\n# ════════════════════════════════════════════════════════════\n#  8. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.RandomRotate90(p=1.0), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.12, scale_limit=0.18, rotate_limit=35,\n                       border_mode=cv2.BORDER_REFLECT, p=0.65),\n    A.ElasticTransform(alpha=1.0, sigma=50, alpha_affine=50,\n                       border_mode=cv2.BORDER_REFLECT, p=0.4),\n    A.GridDistortion(num_steps=5, distort_limit=0.3,\n                     border_mode=cv2.BORDER_REFLECT, p=0.3),\n    A.RandomBrightnessContrast(0.25, 0.25, p=0.55),\n    A.GaussNoise(var_limit=(0.001, 0.005), p=0.35),\n    A.GaussianBlur(blur_limit=(3, 7), p=0.25),\n    A.CoarseDropout(max_holes=6, max_height=28, max_width=28, fill_value=0, p=0.35),\n])\n\ndef channel_dropout(img_np, prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < prob:\n        n   = np.random.randint(1, max_drop+1)\n        idx = np.random.choice(img_np.shape[2], n, replace=False)\n        img_np = img_np.copy(); img_np[:,:,idx] = 0.\n    return img_np\n\ndef cutmix_batch(imgs, msks, alpha=0.4):\n    B,C,H,W = imgs.shape; lam = np.random.beta(alpha, alpha)\n    perm = torch.randperm(B, device=imgs.device)\n    cx=np.random.randint(W); cy=np.random.randint(H)\n    bw=int(W*math.sqrt(1-lam)); bh=int(H*math.sqrt(1-lam))\n    x1=max(0,cx-bw//2); x2=min(W,cx+bw//2)\n    y1=max(0,cy-bh//2); y2=min(H,cy+bh//2)\n    i=imgs.clone(); m=msks.clone()\n    i[:,:,y1:y2,x1:x2]=imgs[perm,:,y1:y2,x1:x2]\n    m[:,:,y1:y2,x1:x2]=msks[perm,:,y1:y2,x1:x2]\n    return i, m\n\n\n# ════════════════════════════════════════════════════════════\n#  9. SEGMENTATION DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride, anomaly_maps=None,\n                 transform=None, neg_ratio=0., apply_ch_dropout=False):\n        self.cache      = load_slice_cache(frag_path, z_list)\n        self.z_list     = z_list; self.tf = transform\n        self.ch_dropout = apply_ch_dropout\n        self.anom_map   = anomaly_maps.get(frag_path) if anomaly_maps else None\n        msk = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n        assert msk is not None\n        self.mask = (msk>0).astype(np.uint8)\n        H, W = self.mask.shape\n        ir   = os.path.join(frag_path,'mask.png')\n        pap  = ((cv2.imread(ir,0)>0).astype(np.uint8)\n                if os.path.exists(ir) else None)\n        if pap is None:\n            zm = z_list[len(z_list)//2]\n            pap = (self.cache[zm]>0.1).astype(np.uint8)\n        pos_c, neg_c = [], []\n        for y in range(0, H-PATCH_SIZE+1, stride):\n            for x in range(0, W-PATCH_SIZE+1, stride):\n                if pap[y:y+PATCH_SIZE,x:x+PATCH_SIZE].mean()<0.5: continue\n                ink = self.mask[y:y+PATCH_SIZE,x:x+PATCH_SIZE].mean()\n                if   ink >= INK_MIN_POS:            pos_c.append((y,x,1))\n                elif ink < 0.001 and neg_ratio > 0: neg_c.append((y,x,0))\n        n_neg = int(len(pos_c)*neg_ratio)\n        np.random.shuffle(neg_c); neg_c = neg_c[:n_neg]\n        self.coords  = pos_c+neg_c\n        self.weights = np.array([3. if c[2]==1 else 1.\n                                 for c in self.coords], np.float32)\n        print(f\"  [{os.path.basename(frag_path)}] \"\n              f\"{len(pos_c)} pos + {len(neg_c)} neg | \"\n              f\"anomaly={'yes' if self.anom_map is not None else 'zeros'}\")\n\n    def __len__(self): return len(self.coords)\n\n    def get_patch(self, idx):\n        y,x,_ = self.coords[idx]\n        img   = np.stack([self.cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                          for z in self.z_list], axis=-1)\n        a     = (self.anom_map[y:y+PATCH_SIZE,x:x+PATCH_SIZE,None].astype(np.float32)\n                 if self.anom_map is not None\n                 else np.zeros((PATCH_SIZE,PATCH_SIZE,1),np.float32))\n        return np.concatenate([img,a],-1), self.mask[y:y+PATCH_SIZE,x:x+PATCH_SIZE].copy()\n\n    def __getitem__(self, idx):\n        img, msk = self.get_patch(idx)\n        if self.tf:\n            o=self.tf(image=img,mask=msk); img,msk=o['image'],o['mask']\n        if self.ch_dropout:\n            img[:,:,:N_CH] = channel_dropout(img[:,:,:N_CH])\n        return (torch.from_numpy(img).permute(2,0,1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\n\n# ════════════════════════════════════════════════════════════\n#  10. LOSS & METRICS\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce); a_t = target*alpha+(1-target)*(1-alpha)\n    return (a_t*((1-p_t)**gamma)*bce).mean()\n\ndef dice_loss(pred, target, smooth=1.):\n    p = torch.sigmoid(pred)\n    i = (p*target).sum(dim=(2,3)); u = p.sum(dim=(2,3))+target.sum(dim=(2,3))\n    return 1.-((2.*i+smooth)/(u+smooth)).mean()\n\ndef combined_loss(pred, target, eps=0.05):\n    return 0.5*focal_loss(pred,target*(1-eps)+0.5*eps) + 0.5*dice_loss(pred,target)\n\ndef multiscale_loss(main,sal,ds3,ds2,ds1,target):\n    loss = combined_loss(main,target)\n    loss += 0.2*combined_loss(sal,F.adaptive_avg_pool2d(target,sal.shape[-2:]))\n    for l,w in [(ds3,0.15),(ds2,0.10),(ds1,0.05)]:\n        if l is not None:\n            loss += w*combined_loss(l,F.adaptive_avg_pool2d(target,l.shape[-2:]))\n    return loss\n\ndef batch_dice(logits, masks, thr=0.5):\n    p = (torch.sigmoid(logits)>thr).float()\n    i = (p*masks).sum(dim=(1,2,3)); u = p.sum(dim=(1,2,3))+masks.sum(dim=(1,2,3))\n    return ((2.*i+1e-5)/(u+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t,best_d = 0.5,0.\n    for t in np.arange(0.10,0.90,0.01):\n        p=(probs>t).astype(np.float32)\n        d=(2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d>best_d: best_d,best_t=d,float(t)\n    return best_t,best_d\n\ndef compute_sep(probs, masks):\n    if (masks>0.5).any() and (masks<0.5).any():\n        return float(probs[masks>0.5].mean()-probs[masks<0.5].mean())\n    return 0.\n\n\n# ════════════════════════════════════════════════════════════\n#  11. GAUSSIAN WEIGHT\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c=sz//2; s=sz//4; y,x=np.mgrid[0:sz,0:sz]\n    return np.exp(-((x-c)**2+(y-c)**2)/(2*s**2)).astype(np.float32)\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  12. RUN — STAGE 1: TRAIN DDPM\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\" + \"=\"*65)\nprint(\"STAGE 1 — 3D DDPM: Clean Papyrus Prior\")\nprint(\"=\"*65)\n\nnoink_ds   = NoInkDataset3D([FRAG2,FRAG3], Z_SLICES,\n                              DDPM_PATCH_XY, DDPM_MAX_PATCHES)\nddpm_sched = build_ddpm_schedule(DDPM_T, DDPM_BETA_START, DDPM_BETA_END, DEVICE)\nddpm_model = UNet3D_DDPM(in_ch=1, base_ch=DDPM_BASE_CH, T=DDPM_T).to(DEVICE)\nprint(f\"DDPM params: {sum(p.numel() for p in ddpm_model.parameters())/1e3:.1f}k\")\n\nddpm_model = train_ddpm(ddpm_model, noink_ds, ddpm_sched,\n                         DDPM_EPOCHS, DDPM_LR, DEVICE, DDPM_BATCH)\ntorch.save({'state':ddpm_model.state_dict(),'T':DDPM_T,\n            't_noise':DDPM_T_NOISE,'ddim_steps':DDIM_STEPS},\n           OUTPUT+'ddpm_model.pth')\nprint(\"DDPM saved.\")\n\n\n# ════════════════════════════════════════════════════════════\n#  13. GENERATE ANOMALY MAPS\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\" + \"=\"*65)\nprint(f\"Anomaly maps — DDIM {DDIM_STEPS} steps | batch {ANOM_BATCH} | \"\n      f\"stride {ANOM_STRIDE} (no overlap) | chunk {COORD_CHUNK}\")\nprint(f\"  Frag2 ~{int(14830/ANOM_STRIDE)*int(9506/ANOM_STRIDE)} patches \"\n      f\"(was ~95k with stride=32) → ~4-8x faster\")\nprint(\"=\"*65)\n\nanomaly_maps = {}\nfor fp, tag in [(FRAG2,'frag2'),(FRAG3,'frag3')]:\n    print(f\"\\n  [{tag}]\")\n    t0    = time.time()\n    cache = noink_ds.caches[fp]\n    H_fp, W_fp = next(iter(cache.values())).shape\n    pap_fp     = load_papyrus_mask(fp, H_fp, W_fp, cache)\n    amap  = compute_anomaly_map_fast(\n        ddpm_model, cache, Z_SLICES, ddpm_sched,\n        patch_xy=DDPM_PATCH_XY, t_noise=DDPM_T_NOISE,\n        ddim_steps=DDIM_STEPS, anom_batch=ANOM_BATCH,\n        papyrus_mask=pap_fp, anom_stride=ANOM_STRIDE,\n        coord_chunk=COORD_CHUNK, device=DEVICE)\n    anomaly_maps[fp] = amap\n    cv2.imwrite(OUTPUT+f'anomaly_{tag}.png', (amap*255).astype(np.uint8))\n    print(f\"    Done in {(time.time()-t0)/60:.1f}min | \"\n          f\"min={amap.min():.3f} max={amap.max():.3f} mean={amap.mean():.3f}\")\n\nddpm_model_cpu = ddpm_model.cpu()\ndel ddpm_model; gc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\nprint(\"\\nDDPM moved to CPU — VRAM freed for V10.\")\n\n\n# ════════════════════════════════════════════════════════════\n#  14. STAGE 2: TRAIN V10 SEGMENTATION\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\"+\"=\"*65)\nprint(f\"STAGE 2 — V10 Segmentation ({N_CH_SEG}-ch: {N_CH} CT + 1 anomaly)\")\nprint(\"=\"*65)\n\nprint('\\n── Fragment 2 ──')\nds2 = VesuviusDataset(FRAG2, Z_SLICES, STRIDE_TR, anomaly_maps=anomaly_maps,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\nprint('\\n── Fragment 3 ──')\nds3 = VesuviusDataset(FRAG3, Z_SLICES, STRIDE_TR, anomaly_maps=anomaly_maps,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\n\nn_total   = len(ds2)+len(ds3)\nrng       = np.random.RandomState(SEED); all_idx = rng.permutation(n_total)\nn_val     = int(n_total*VAL_SPLIT)\ntrain_idx = all_idx[n_val:].tolist(); val_idx = all_idx[:n_val].tolist()\nprint(f'Total: {n_total} | Train: {len(train_idx)} | Val: {len(val_idx)}')\n\n\nclass _Subset(Dataset):\n    def __init__(self, d2, d3, idxs, is_train):\n        self.d2=d2; self.d3=d3; self.n2=len(d2)\n        self.idxs=idxs; self.is_train=is_train\n    def __len__(self): return len(self.idxs)\n    def __getitem__(self, i):\n        g = self.idxs[i]\n        if self.is_train:\n            return self.d2[g] if g<self.n2 else self.d3[g-self.n2]\n        img,msk = (self.d2.get_patch(g) if g<self.n2\n                   else self.d3.get_patch(g-self.n2))\n        return (torch.from_numpy(img).permute(2,0,1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\ntrain_ds = _Subset(ds2,ds3,train_idx,True)\nval_ds   = _Subset(ds2,ds3,val_idx,  False)\nall_w    = np.concatenate([ds2.weights, ds3.weights])\nsampler  = WeightedRandomSampler(torch.from_numpy(all_w[train_idx]).float(),\n                                  len(train_ds), True)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                      num_workers=0, pin_memory=False)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=0, pin_memory=False)\nprint(f'Train batches: {len(train_dl)} | Val: {len(val_dl)}')\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nmodel    = VesuviusV10(n_ch=N_CH_SEG, enc_dim=512,\n                       n_transformer_blocks=4, num_heads=8).to(DEVICE)\nprint(f'V10 params: {sum(p.numel() for p in model.parameters()if p.requires_grad)/1e6:.2f}M')\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='max', factor=0.5, patience=5, min_lr=5e-6, verbose=True)\nscaler    = GradScaler(enabled=USE_AMP)\nbest_dice = 0.; pat_cnt = 0\nhistory   = dict(tl=[],vl=[],td=[],vd=[],sep=[],lr=[])\n\nfor epoch in range(EPOCHS):\n    model.train(); tl=td=0.; optimizer.zero_grad()\n    for step,(imgs,msks) in enumerate(\n            tqdm(train_dl,desc=f'Ep{epoch+1:02d}▸train',leave=False)):\n        imgs=imgs.to(DEVICE); msks=msks.to(DEVICE)\n        if random.random()<0.30: imgs,msks=cutmix_batch(imgs,msks,0.4)\n        with autocast(enabled=USE_AMP):\n            main,sal,d3,d2_l,d1=model(imgs)\n            loss=multiscale_loss(main,sal,d3,d2_l,d1,msks)/GRAD_ACCUM\n        scaler.scale(loss).backward()\n        tl+=loss.item()*GRAD_ACCUM; td+=batch_dice(main.detach(),msks)\n        del imgs,msks,main,sal,d3,d2_l,d1,loss\n        if (step+1)%GRAD_ACCUM==0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(),1.)\n            scaler.step(optimizer); scaler.update(); optimizer.zero_grad()\n            if DEVICE=='cuda': torch.cuda.empty_cache()\n    tl/=len(train_dl); td/=len(train_dl)\n\n    model.eval(); vl=vd=0.; all_p,all_m=[],[]\n    with torch.no_grad():\n        for imgs,msks in tqdm(val_dl,desc=f'Ep{epoch+1:02d}▸val  ',leave=False):\n            imgs=imgs.to(DEVICE); msks=msks.to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                main,sal,d3,d2_l,d1=model(imgs)\n                loss=multiscale_loss(main,sal,d3,d2_l,d1,msks)\n            vl+=loss.item(); vd+=batch_dice(main,msks)\n            all_p.append(torch.sigmoid(main/TEMPERATURE).cpu().numpy())\n            all_m.append(msks.cpu().numpy())\n            del imgs,msks,main,sal,d3,d2_l,d1,loss\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    vl/=len(val_dl); vd/=len(val_dl)\n    P=np.concatenate(all_p); M=np.concatenate(all_m); del all_p,all_m\n    bt,bd=sweep_threshold(P,M); sep=compute_sep(P,M)\n    ink_m=float(P[M>0.5].mean()) if (M>0.5).any() else 0.\n    bg_m =float(P[M<0.5].mean()) if (M<0.5).any() else 0.\n    del P,M; gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    scheduler.step(vd)\n    lr_now=optimizer.param_groups[0]['lr']\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['sep'].append(sep); history['lr'].append(lr_now)\n    print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n          f'train={tl:.4f}/{td:.4f} | val={vl:.4f}/{vd:.4f} | '\n          f'thr={bt:.2f}→{bd:.4f} | sep={sep:+.3f} '\n          f'[ink={ink_m:.3f} bg={bg_m:.3f}]')\n    metric = max(vd,bd) - max(0.,bt-0.60)*0.3\n    if metric>best_dice:\n        best_dice=metric; pat_cnt=0\n        torch.save({'epoch':epoch,'state':model.state_dict(),'thr':bt,\n                    'metric':metric,'n_ch':N_CH_SEG,'bd':bd,'vd':vd},\n                   OUTPUT+'best_model.pth')\n        print(f'  ✓ saved (metric={metric:.4f})')\n    else:\n        pat_cnt+=1\n        if pat_cnt>=PATIENCE: print(f'  early stop ep{epoch+1}'); break\n\nprint(f'\\nBest metric: {best_dice:.4f}')\n\nfig,axes=plt.subplots(1,4,figsize=(20,4))\naxes[0].plot(history['tl'],label='train'); axes[0].plot(history['vl'],label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'],label='train'); axes[1].plot(history['vd'],label='val')\naxes[1].axhline(0.80,color='r',ls='--'); axes[1].set_title('Dice')\naxes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['sep']); axes[2].axhline(0.4,color='orange',ls='--')\naxes[2].set_title('Sep'); axes[2].grid(True)\naxes[3].plot(history['lr']); axes[3].set_title('LR'); axes[3].grid(True)\nplt.tight_layout(); plt.savefig(OUTPUT+'curves.png',dpi=80); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  15. INFERENCE\n# ════════════════════════════════════════════════════════════\ndef predict_fragment_with_ddpm(seg_model, ddpm_model_cpu, frag_path,\n                                z_list, anom_map_precomputed=None):\n    seg_model.eval()\n\n    if anom_map_precomputed is not None:\n        amap  = anom_map_precomputed\n        cache = load_slice_cache(frag_path, z_list)\n        print(\"  Using precomputed anomaly map\")\n    else:\n        print(\"  Computing anomaly map via DDIM (fast)...\")\n        cache  = load_slice_cache(frag_path, z_list)\n        H_fp, W_fp = next(iter(cache.values())).shape\n        pap_fp = load_papyrus_mask(frag_path, H_fp, W_fp, cache)\n        ddpm_g = ddpm_model_cpu.to(DEVICE)\n        sched_g= build_ddpm_schedule(DDPM_T, DDPM_BETA_START, DDPM_BETA_END, DEVICE)\n        t0     = time.time()\n        amap   = compute_anomaly_map_fast(\n            ddpm_g, cache, z_list, sched_g,\n            patch_xy=DDPM_PATCH_XY, t_noise=DDPM_T_NOISE,\n            ddim_steps=DDIM_STEPS, anom_batch=ANOM_BATCH,\n            papyrus_mask=pap_fp, anom_stride=ANOM_STRIDE,\n            coord_chunk=COORD_CHUNK, device=DEVICE)\n        print(f\"  Anomaly map done in {(time.time()-t0)/60:.1f}min\")\n        ddpm_g = ddpm_g.cpu(); del ddpm_g, sched_g\n        gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    msk = (cv2.imread(os.path.join(frag_path,'inklabels.png'),0)>0).astype(np.uint8)\n    H, W = msk.shape\n    pred_map = np.zeros((H,W),np.float32); wgt_map=np.zeros((H,W),np.float32)\n    coords   = [(y,x) for y in range(0,H-PATCH_SIZE+1,STRIDE_INF)\n                       for x in range(0,W-PATCH_SIZE+1,STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y,x) in tqdm(coords,desc='Seg infer',leave=True):\n            slices=[cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                    for z in z_list]\n            img = np.stack(slices,axis=-1)\n            a   = amap[y:y+PATCH_SIZE,x:x+PATCH_SIZE,None].astype(np.float32)\n            img = np.concatenate([img,a],-1)\n            t   = torch.from_numpy(img).permute(2,0,1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                logit,*_ = seg_model(t)\n            p = torch.sigmoid(logit/TEMPERATURE).squeeze().cpu().float().numpy()\n            pred_map[y:y+PATCH_SIZE,x:x+PATCH_SIZE] += p*GW\n            wgt_map [y:y+PATCH_SIZE,x:x+PATCH_SIZE] += GW\n            del t,logit,p\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    seg_prob = pred_map/(wgt_map+1e-8)\n    H_a,W_a  = amap.shape; h=min(H,H_a); w=min(W,W_a)\n    fused    = np.zeros((H,W),np.float32)\n    fused[:h,:w] = 0.70*seg_prob[:h,:w] + 0.30*amap[:h,:w]\n    return seg_prob, amap, np.clip(fused,0.,1.), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  16. FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\"+\"=\"*65)\nprint(\"FINAL TEST — FRAGMENT 1\")\nprint(\"=\"*65)\n\nckpt  = torch.load(OUTPUT+'best_model.pth', map_location=DEVICE)\nmodel = VesuviusV10(n_ch=N_CH_SEG, enc_dim=512,\n                    n_transformer_blocks=4, num_heads=8).to(DEVICE)\nmodel.load_state_dict(ckpt['state'])\nprint(f\"Checkpoint: ep={ckpt['epoch']+1}  \"\n      f\"val={ckpt.get('vd',0):.4f}  thr={ckpt['thr']:.2f}\")\n\nseg_prob,amap1,fused_prob,msk1 = predict_fragment_with_ddpm(\n    model, ddpm_model_cpu, FRAG1, Z_SLICES, anom_map_precomputed=None)\n\nH_m,W_m  = msk1.shape\nseg_c     = seg_prob[:H_m,:W_m]\nfuse_c    = fused_prob[:H_m,:W_m]\namap_c    = amap1[:H_m,:W_m]\n\nbt_seg, bd_seg   = sweep_threshold(seg_c,  msk1)\nbt_fuse,bd_fuse  = sweep_threshold(fuse_c, msk1)\nprint(f\"  Seg  : thr={bt_seg:.2f}  dice={bd_seg:.4f}\")\nprint(f\"  Fused: thr={bt_fuse:.2f}  dice={bd_fuse:.4f}\")\n\nif bd_fuse >= bd_seg:\n    final_prob=fuse_c; best_thr=bt_fuse; mode=\"Fused (Seg+DDPM)\"\nelse:\n    final_prob=seg_c;  best_thr=bt_seg;  mode=\"Seg only\"\n\npred   = (final_prob>best_thr).astype(np.uint8)\nkernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(5,5))\npred   = cv2.morphologyEx(pred,cv2.MORPH_CLOSE,kernel)\nn_lab,labels,stats,_ = cv2.connectedComponentsWithStats(pred)\ncleaned = np.zeros_like(pred)\nfor i in range(1,n_lab):\n    if stats[i,cv2.CC_STAT_AREA]>=150: cleaned[labels==i]=1\npred = cleaned\n\npf=pred.flatten().astype(int); mf=msk1.flatten().astype(int)\ntn,fp,fn,tp_v = confusion_matrix(mf,pf,labels=[0,1]).ravel()\nprec=tp_v/(tp_v+fp+1e-8); rec=tp_v/(tp_v+fn+1e-8)\nf1  =2*prec*rec/(prec+rec+1e-8)\ndice=(2*tp_v+1)/(pred.sum()+msk1.sum()+1)\nink_m   =float(final_prob[msk1==1].mean()) if (msk1==1).any() else 0.\nbg_m    =float(final_prob[msk1==0].mean()) if (msk1==0).any() else 0.\nddpm_ink=float(amap_c[msk1==1].mean()) if (msk1==1).any() else 0.\nddpm_bg =float(amap_c[msk1==0].mean()) if (msk1==0).any() else 0.\n\nprint('\\n'+'='*65)\nprint(f'RESULTS — FRAGMENT 1  [{mode}]')\nprint('='*65)\nprint(f'Dice      : {dice:.4f}')\nprint(f'F1        : {f1:.4f}')\nprint(f'Precision : {prec:.4f}')\nprint(f'Recall    : {rec:.4f}')\nprint(f'FP/TP     : {fp/(tp_v+1e-8):.2f}')\nprint(f'Sep (seg) : {ink_m-bg_m:+.3f}')\nprint(f'Sep (ddpm): {ddpm_ink-ddpm_bg:+.3f}  ← unsupervised signal quality')\nprint(f'F1≥0.80   : {\"✓ PASSED\" if f1>=0.80 else \"✗ Below target\"}')\nprint('='*65)\n\nfig,ax=plt.subplots(2,4,figsize=(22,12))\nax[0,0].imshow(msk1,  cmap='gray');    ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(amap_c,cmap='hot');     ax[0,1].set_title('DDPM Anomaly Map')\nax[0,2].imshow(seg_c, cmap='inferno'); ax[0,2].set_title(f'Seg prob (dice={bd_seg:.3f})')\nax[0,3].imshow(fuse_c,cmap='inferno'); ax[0,3].set_title(f'Fused prob (dice={bd_fuse:.3f})')\nax[1,0].imshow(pred,  cmap='gray');    ax[1,0].set_title(f'{mode}  Dice={dice:.4f}')\nerr=np.zeros((*msk1.shape,3),dtype=np.uint8)\nerr[(pred==1)&(msk1==1)]=[0,255,0]\nerr[(pred==1)&(msk1==0)]=[255,0,0]\nerr[(pred==0)&(msk1==1)]=[0,0,255]\nax[1,1].imshow(err); ax[1,1].set_title('TP=green FP=red FN=blue')\nax[1,2].hist(seg_c[msk1==1].ravel(),bins=60,alpha=0.7,\n             label=f'ink μ={ink_m:.3f}',color='orange',density=True)\nax[1,2].hist(seg_c[msk1==0].ravel(),bins=60,alpha=0.7,\n             label=f'bg μ={bg_m:.3f}',color='blue',density=True)\nax[1,2].axvline(best_thr,color='r',ls='--'); ax[1,2].legend()\nax[1,2].set_title('Seg distribution')\nax[1,3].hist(amap_c[msk1==1].ravel(),bins=60,alpha=0.7,\n             label=f'ink μ={ddpm_ink:.3f}',color='orange',density=True)\nax[1,3].hist(amap_c[msk1==0].ravel(),bins=60,alpha=0.7,\n             label=f'bg μ={ddpm_bg:.3f}',color='blue',density=True)\nax[1,3].legend(); ax[1,3].set_title('DDPM anomaly distribution')\nfor a in ax[0]: a.axis('off')\nax[1,0].axis('off'); ax[1,1].axis('off')\nplt.suptitle(f'Fragment 1 | Dice={dice:.4f}  F1={f1:.4f}  '\n             f'Prec={prec:.4f}  Rec={rec:.4f}',fontsize=12)\nplt.tight_layout()\nplt.savefig(OUTPUT+'frag1_prediction.png',dpi=80,bbox_inches='tight')\nplt.close()\n\ncv2.imwrite(OUTPUT+'frag1_anomaly.png',  (amap_c*255).astype(np.uint8))\ncv2.imwrite(OUTPUT+'frag1_seg_prob.png', (seg_c*255).astype(np.uint8))\ncv2.imwrite(OUTPUT+'frag1_binary.png',   (pred*255).astype(np.uint8))\n\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write('VESUVIUS V10+DDPM (FINAL SPEED FIX)\\n'+'='*55+'\\n')\n    f.write(f'Mode      : {mode}\\n')\n    f.write(f'Dice      : {dice:.4f}\\n')\n    f.write(f'F1        : {f1:.4f}\\n')\n    f.write(f'Precision : {prec:.4f}\\n')\n    f.write(f'Recall    : {rec:.4f}\\n')\n    f.write(f'FP/TP     : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Sep (seg) : {ink_m-bg_m:+.3f}\\n')\n    f.write(f'Sep (ddpm): {ddpm_ink-ddpm_bg:+.3f}\\n')\n    f.write(f'DDIM steps: {DDIM_STEPS}  t_noise={DDPM_T_NOISE}\\n')\n    f.write(f'anom_batch: {ANOM_BATCH}  coord_chunk={COORD_CHUNK}\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  ddpm_model.pth | best_model.pth | curves.png')\nprint('  frag1_prediction.png | frag1_anomaly.png')\nprint('  frag1_seg_prob.png | frag1_binary.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n# ============================================================\n#  VESUVIUS INK DETECTION — v10  FIXED + FRAGMENT SPLIT\n#\n#  Root causes of Dice=0.46 / 0.40:\n#\n#  FIX 1 — Albumentations deprecated API silently no-ops\n#    ShiftScaleRotate border_mode → A.Affine with padding_mode\n#    GaussNoise var_limit → noise_scale_factor\n#    CoarseDropout old args → new args with try/except fallback\n#    ElasticTransform alpha_affine removed in new API\n#    RandomResizedCrop size arg changed\n#\n#  FIX 2 — Token merging odd-length bug\n#    When n_merge is odd, the unmerge function had an index\n#    off-by-one that silently corrupted the bottleneck features.\n#    Simplified: always force n_merge even before merging.\n#\n#  FIX 3 — Training only on Frag2+Frag3 → domain gap on Frag1\n#    MODIFIED: train on 80% of Frag2+Frag3, val on 20% of Frag2+Frag3\n#    Test on full Fragment 1 (completely held out)\n#\n#  FIX 4 — ReduceLROnPlateau collapses LR too fast\n#    Replaced with CosineAnnealingWarmRestarts (T_0=10)\n#    — LR never collapses, restarts keep exploration alive.\n#\n#  FIX 5 — BATCH/STRIDE tuned for P100 16GB\n#    PATCH_SIZE=224, STRIDE_TR=112, STRIDE_INF=56\n#    BATCH_SIZE=4, GRAD_ACCUM=4 (effective=16)\n#    These gave 0.46 — we keep them (memory-safe).\n#    The issue was NOT the batch size, it was the training protocol.\n#\n#  FIX 6 — torch.cuda.amp compatibility\n#    Auto-detect correct API (torch.amp vs torch.cuda.amp)\n#\n#  RETAINED from original v10:\n#    ResNet34 encoder + Axial-Transformer bottleneck\n#    Log-polar positional bias\n#    Sheet-surface detector + token merging\n#    Deep supervision (3 scales)\n#    CutMix augmentation\n#    Channel dropout\n#    Morphological post-processing\n#    DenseCRF (if available)\n#    Temperature scaling (1.3)\n#    Pseudo-labelling hook\n# ============================================================\n\nimport os, gc, cv2, math, warnings, random\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import (Dataset, DataLoader, WeightedRandomSampler,\n                               ConcatDataset, Subset)\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\ntry:\n    import pydensecrf.densecrf as dcrf\n    from pydensecrf.utils import unary_from_labels\n    HAS_CRF = True\nexcept ImportError:\n    HAS_CRF = False\n    print(\"[INFO] pydensecrf not found – CRF post-processing skipped.\")\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ── paths ─────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ──────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 224\nSTRIDE_TR    = 112\nSTRIDE_INF   = 32\nBATCH_SIZE   = 4\nGRAD_ACCUM   = 4          # effective batch = 16\nEPOCHS       = 50\nLR           = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 12\nNUM_WORKERS  = 2\nTEMPERATURE  = 1.3\nDROPOUT_P    = 0.3\n\nZ_SLICES  = [18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34]\nN_CH      = len(Z_SLICES)   # 17\n\nINK_MIN_POS  = 0.02\nNEG_RATIO    = 0.5\nCH_DROP_PROB = 0.3\nCH_DROP_MAX  = 3\nMAX_PATCHES  = 2_000       # cap per fragment\n\nUSE_AMP = (DEVICE == 'cuda')\nPIN     = (DEVICE == 'cuda')\n\n# FIX 6: Auto-detect correct AMP API\nif USE_AMP:\n    try:\n        # Try new API first (PyTorch 2.4+)\n        from torch.amp import GradScaler, autocast\n        scaler = GradScaler('cuda')\n        amp_context = lambda: autocast('cuda')\n        print(\"Using torch.amp (new API)\")\n    except (ImportError, AttributeError):\n        try:\n            # Fall back to old API\n            from torch.cuda.amp import GradScaler, autocast\n            scaler = GradScaler()\n            amp_context = autocast\n            print(\"Using torch.cuda.amp (legacy API)\")\n        except ImportError:\n            USE_AMP = False\n            scaler = None\n            amp_context = lambda: torch.no_grad()\n            print(\"AMP not available, using float32\")\nelse:\n    scaler = None\n    amp_context = lambda: torch.no_grad()\n\nprint(f\"Device : {DEVICE}  |  Channels: {N_CH}  |  Patch: {PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. SLICE CACHE\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_slice_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        print(f\"  Reusing cache: {os.path.basename(frag_path)}\")\n        return _cache_store[key]\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    if not raw:\n        raise ValueError(f\"No slices in {vol_dir}\")\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5)/(p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: \"\n          f\"{len(cache)} slices, {mb:.0f} MB\")\n    _cache_store[key] = cache\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2. LOG-POLAR POSITIONAL BIAS\n# ════════════════════════════════════════════════════════════\n_LP_CACHE: dict = {}\n\ndef build_logpolar_bias(seq_len_h, seq_len_w, num_heads, device='cpu'):\n    H, W = seq_len_h, seq_len_w\n    ys = torch.arange(H, dtype=torch.float32)\n    xs = torch.arange(W, dtype=torch.float32)\n    gy, gx = torch.meshgrid(ys, xs, indexing='ij')\n    gy = gy.reshape(-1); gx = gx.reshape(-1)\n    dy = gy.unsqueeze(1) - gy.unsqueeze(0)\n    dx = gx.unsqueeze(1) - gx.unsqueeze(0)\n    r   = torch.sqrt(dx**2 + dy**2 + 1e-3)\n    log_r = torch.log(r + 1.0)\n    phi = torch.atan2(dy, dx)\n    freqs = torch.arange(1, num_heads // 2 + 1, dtype=torch.float32)\n    bias_r   = torch.cos(freqs[None,None,:] * log_r.unsqueeze(2))\n    bias_phi = torch.sin(freqs[None,None,:] * phi.unsqueeze(2))\n    bias = torch.cat([bias_r, bias_phi], dim=-1).permute(2,0,1)\n    return bias.to(device)\n\ndef get_logpolar_bias(H, W, num_heads, device):\n    key = (H, W, num_heads, str(device))\n    if key not in _LP_CACHE:\n        _LP_CACHE[key] = build_logpolar_bias(H, W, num_heads, device)\n    return _LP_CACHE[key]\n\n\n# ════════════════════════════════════════════════════════════\n#  3. SHEET-SURFACE DETECTOR\n# ════════════════════════════════════════════════════════════\nclass SheetSurfaceDetector(nn.Module):\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, 32, 3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.ReLU(inplace=True),\n            nn.Conv2d(16,  1, 1),\n        )\n    def forward(self, x): return self.net(x)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. TOKEN MERGING  (FIX 2: always force n_merge even)\n# ════════════════════════════════════════════════════════════\ndef token_merge(tokens, saliency, merge_ratio=0.25):\n    \"\"\"\n    tokens   : [B, C, N]\n    saliency : [B, 1, N]\n    Returns  : merged_tokens, unmerge_info\n    \"\"\"\n    B, C, N = tokens.shape\n    sal = saliency.squeeze(1)\n    _, sort_idx = sal.sort(dim=1)\n\n    # FIX: force n_merge to be even to avoid odd-length bugs\n    n_merge = int(N * merge_ratio)\n    n_merge = n_merge - (n_merge % 2)   # make even\n    if n_merge < 2:\n        # nothing to merge — return as-is\n        unmerge_info = (None, None, 0, N)\n        return tokens, unmerge_info\n\n    merge_idx = sort_idx[:, :n_merge]\n    keep_idx  = sort_idx[:, n_merge:]\n\n    def gather(t, idx):\n        idx_exp = idx.unsqueeze(1).expand(-1, C, -1)\n        return t.gather(2, idx_exp)\n\n    t_merge = gather(tokens, merge_idx)\n    t_keep  = gather(tokens, keep_idx)\n\n    # pair consecutive low-saliency tokens\n    t_merged = (t_merge[:,:,0::2] + t_merge[:,:,1::2]) / 2  # [B,C,n_merge/2]\n    tokens_out = torch.cat([t_keep, t_merged], dim=2)\n\n    unmerge_info = (keep_idx, merge_idx, n_merge, N)\n    return tokens_out, unmerge_info\n\n\ndef token_unmerge(tokens_out, unmerge_info, C):\n    keep_idx, merge_idx, n_merge, N = unmerge_info\n    if n_merge == 0:\n        return tokens_out\n\n    B    = tokens_out.shape[0]\n    M    = tokens_out.shape[2]\n    n_keep = M - n_merge // 2\n\n    t_keep   = tokens_out[:,:,:n_keep]\n    t_merged = tokens_out[:,:,n_keep:]         # [B,C,n_merge/2]\n    # expand each merged token back to 2\n    t_expanded = t_merged.repeat_interleave(2, dim=2)   # [B,C,n_merge]\n\n    device = tokens_out.device\n    out = torch.zeros(B, C, N, device=device, dtype=tokens_out.dtype)\n\n    def scatter(t, idx):\n        idx_exp = idx.unsqueeze(1).expand(-1, C, -1)\n        out.scatter_(2, idx_exp, t)\n\n    scatter(t_keep,    keep_idx)\n    scatter(t_expanded, merge_idx)\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  5. AXIAL ATTENTION\n# ════════════════════════════════════════════════════════════\nclass AxialAttention(nn.Module):\n    def __init__(self, dim, num_heads=8, axis='x', dropout=0.1):\n        super().__init__()\n        assert axis in ('x', 'y')\n        self.axis      = axis\n        self.num_heads = num_heads\n        self.head_dim  = dim // num_heads\n        self.scale     = self.head_dim ** -0.5\n        self.qkv  = nn.Linear(dim, dim*3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.drop = nn.Dropout(dropout)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        if self.axis == 'x':\n            x_r = x.permute(0,2,3,1).reshape(B*H, W, C)\n        else:\n            x_r = x.permute(0,3,2,1).reshape(B*W, H, C)\n\n        res  = x_r\n        x_n  = self.norm(x_r)\n        BN, L, _ = x_n.shape\n        qkv = self.qkv(x_n).reshape(BN,L,3,self.num_heads,self.head_dim)\n        qkv = qkv.permute(2,0,3,1,4)\n        q, k, v = qkv.unbind(0)\n\n        attn = (q @ k.transpose(-2,-1)) * self.scale\n        pos  = torch.arange(L, dtype=torch.float32, device=x.device)\n        d    = (pos.unsqueeze(0) - pos.unsqueeze(1)).abs().float()\n        attn = attn + (-torch.log(d + 1.0)).unsqueeze(0).unsqueeze(0)\n        attn = self.drop(attn.softmax(-1))\n        out  = (attn @ v).transpose(1,2).reshape(BN, L, C)\n        out  = self.proj(out) + res\n\n        if self.axis == 'x':\n            out = out.reshape(B,H,W,C).permute(0,3,1,2)\n        else:\n            out = out.reshape(B,W,H,C).permute(0,3,2,1)\n        return out\n\n\nclass AxialTransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=8, mlp_ratio=4, dropout=0.1):\n        super().__init__()\n        self.attn_x = AxialAttention(dim, num_heads, 'x', dropout)\n        self.attn_y = AxialAttention(dim, num_heads, 'y', dropout)\n        self.norm2  = nn.LayerNorm(dim)\n        mlp_dim     = int(dim * mlp_ratio)\n        self.ffn = nn.Sequential(\n            nn.Linear(dim, mlp_dim), nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(mlp_dim, dim), nn.Dropout(dropout))\n\n    def forward(self, x):\n        x = self.attn_x(x)\n        x = self.attn_y(x)\n        B, C, H, W = x.shape\n        xf = x.permute(0,2,3,1).reshape(-1, C)\n        xf = self.ffn(self.norm2(xf))\n        x  = x + xf.reshape(B,H,W,C).permute(0,3,1,2)\n        return x\n\n\n# ════════════════════════════════════════════════════════════\n#  6. FULL MODEL\n# ════════════════════════════════════════════════════════════\nclass VesuviusV10(nn.Module):\n    def __init__(self, n_ch=N_CH, enc_dim=512,\n                 n_transformer_blocks=4, num_heads=8):\n        super().__init__()\n        self.backbone = smp.Unet(\n            encoder_name='resnet34',\n            encoder_weights='imagenet',\n            in_channels=n_ch,\n            classes=1,\n            decoder_attention_type='scse',\n        )\n        self.transformer_blocks = nn.Sequential(*[\n            AxialTransformerBlock(enc_dim, num_heads=num_heads,\n                                  mlp_ratio=4, dropout=0.1)\n            for _ in range(n_transformer_blocks)\n        ])\n        self.sheet_detector  = SheetSurfaceDetector(n_ch)\n        self._enc_dim        = enc_dim\n        self.bottleneck_proj = nn.Sequential(\n            nn.Conv2d(enc_dim, enc_dim, 1, bias=False),\n            nn.BatchNorm2d(enc_dim), nn.GELU())\n        self.ds_head3 = nn.Conv2d(256, 1, 1)\n        self.ds_head2 = nn.Conv2d(128, 1, 1)\n        self.ds_head1 = nn.Conv2d( 64, 1, 1)\n\n    def forward(self, x):\n        B = x.shape[0]\n        saliency_logit = self.sheet_detector(x)\n\n        feats      = self.backbone.encoder(x)\n        bottleneck = feats[-1]\n        bH, bW     = bottleneck.shape[2], bottleneck.shape[3]\n        N          = bH * bW\n\n        sal_down = F.adaptive_avg_pool2d(\n            torch.sigmoid(saliency_logit), (bH, bW))\n        sal_flat = sal_down.reshape(B, 1, N)\n        tok      = bottleneck.reshape(B, self._enc_dim, N)\n\n        tok_merged, unmerge_info = token_merge(tok, sal_flat, merge_ratio=0.25)\n\n        M       = tok_merged.shape[2]\n        sq      = int(math.ceil(math.sqrt(M)))\n        pad_len = sq*sq - M\n        if pad_len > 0:\n            tok_merged = F.pad(tok_merged, (0, pad_len))\n        tok_2d   = tok_merged.reshape(B, self._enc_dim, sq, sq)\n        tok_2d   = self.transformer_blocks(tok_2d)\n        tok_flat = tok_2d.reshape(B, self._enc_dim, sq*sq)[:,:,:M]\n\n        tok_full       = token_unmerge(tok_flat, unmerge_info, self._enc_dim)\n        bottleneck_out = tok_full.reshape(B, self._enc_dim, bH, bW)\n        bottleneck_out = self.bottleneck_proj(bottleneck_out + bottleneck)\n\n        feats_mod     = list(feats)\n        feats_mod[-1] = bottleneck_out\n\n        decoder_output = self.backbone.decoder(*feats_mod)\n        main_logit     = self.backbone.segmentation_head(decoder_output)\n\n        ds3_logit = ds2_logit = ds1_logit = None\n        try:\n            db = self.backbone.decoder.blocks\n            if len(db) >= 1:\n                f0 = db[0](feats_mod[-1], feats_mod[-2])\n                ds3_logit = self.ds_head3(f0)\n            if len(db) >= 2:\n                f1 = db[1](f0, feats_mod[-3])\n                ds2_logit = self.ds_head2(f1)\n            if len(db) >= 3:\n                f2 = db[2](f1, feats_mod[-4])\n                ds1_logit = self.ds_head1(f2)\n        except Exception:\n            pass\n\n        return main_logit, saliency_logit, ds3_logit, ds2_logit, ds1_logit\n\n\n# ════════════════════════════════════════════════════════════\n#  7. AUGMENTATIONS  (FIX 1: updated albumentations API)\n# ════════════════════════════════════════════════════════════\ndef _make_train_tf():\n    \"\"\"Build augmentation pipeline, handling different albumentations versions.\"\"\"\n    transforms = [\n        A.RandomRotate90(p=1.0),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(0.25, 0.25, p=0.55),\n        A.GaussianBlur(blur_limit=(3,7), p=0.25),\n        A.ElasticTransform(alpha=1.0, sigma=50, p=0.4),\n        A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.3),\n    ]\n\n    # A.Affine — try new API first\n    try:\n        transforms.insert(3, A.Affine(\n            scale=(0.82, 1.18),\n            translate_percent={'x':(-0.12,0.12),'y':(-0.12,0.12)},\n            rotate=(-35,35),\n            padding_mode='reflect',\n            p=0.65))\n    except TypeError:\n        try:\n            transforms.insert(3, A.Affine(\n                scale=(0.82,1.18),\n                translate_percent={'x':(-0.12,0.12),'y':(-0.12,0.12)},\n                rotate=(-35,35),\n                mode=cv2.BORDER_REFLECT,\n                p=0.65))\n        except TypeError:\n            transforms.insert(3, A.Affine(\n                scale=(0.82,1.18),\n                translate_percent={'x':(-0.12,0.12),'y':(-0.12,0.12)},\n                rotate=(-35,35),\n                p=0.65))\n\n    # GaussNoise — try new API first\n    try:\n        transforms.append(A.GaussNoise(noise_scale_factor=0.1, p=0.35))\n    except TypeError:\n        try:\n            transforms.append(A.GaussNoise(var_limit=(0.001,0.005), p=0.35))\n        except Exception:\n            pass\n\n    # CoarseDropout — try new API first\n    try:\n        transforms.append(A.CoarseDropout(\n            max_holes=6, max_height=28, max_width=28,\n            fill_value=0, p=0.35))\n    except TypeError:\n        try:\n            transforms.append(A.CoarseDropout(\n                num_holes_range=(1,6),\n                hole_height_range=(8,28),\n                hole_width_range=(8,28),\n                fill=0, p=0.35))\n        except Exception:\n            pass\n\n    # RandomResizedCrop — try new API\n    try:\n        transforms.append(A.RandomResizedCrop(\n            size=(PATCH_SIZE, PATCH_SIZE),\n            scale=(0.55,1.0), ratio=(0.85,1.15), p=0.5))\n    except TypeError:\n        try:\n            transforms.append(A.RandomResizedCrop(\n                height=PATCH_SIZE, width=PATCH_SIZE,\n                scale=(0.55,1.0), ratio=(0.85,1.15), p=0.5))\n        except Exception:\n            pass\n\n    return A.Compose(transforms)\n\ntrain_tf = _make_train_tf()\nprint(\"Augmentation pipeline built successfully.\")\n\n\ndef channel_dropout(img_np, prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < prob:\n        n   = np.random.randint(1, max_drop+1)\n        idx = np.random.choice(img_np.shape[2], n, replace=False)\n        img_np = img_np.copy()\n        img_np[:,:,idx] = 0.0\n    return img_np\n\n\ndef cutmix_batch(imgs, msks, alpha=0.4):\n    B, C, H, W = imgs.shape\n    lam  = np.random.beta(alpha, alpha)\n    perm = torch.randperm(B, device=imgs.device)\n    cx = np.random.randint(W); cy = np.random.randint(H)\n    bw = int(W * math.sqrt(1 - lam)); bh = int(H * math.sqrt(1 - lam))\n    x1 = max(0,cx-bw//2); x2 = min(W,cx+bw//2)\n    y1 = max(0,cy-bh//2); y2 = min(H,cy+bh//2)\n    imgs_new = imgs.clone(); msks_new = msks.clone()\n    imgs_new[:,:,y1:y2,x1:x2] = imgs[perm,:,y1:y2,x1:x2]\n    msks_new[:,:,y1:y2,x1:x2] = msks[perm,:,y1:y2,x1:x2]\n    return imgs_new, msks_new\n\n\n# ════════════════════════════════════════════════════════════\n#  8. DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0,\n                 apply_ch_dropout=False, max_patches=0):\n        self.cache       = load_slice_cache(frag_path, z_list)\n        self.z_list      = z_list\n        self.tf          = transform\n        self.ch_dropout  = apply_ch_dropout\n        self.frag_path   = frag_path\n\n        msk = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n        if msk is None:\n            raise FileNotFoundError(f\"No inklabels at {frag_path}\")\n        self.mask = (msk > 0).astype(np.uint8)\n\n        ir_path = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = ((cv2.imread(ir_path,0) > 0).astype(np.uint8)\n                        if os.path.exists(ir_path) else None)\n\n        H, W = self.mask.shape\n        pos_yx, neg_yx = [], []\n        for y in range(0, H-PATCH_SIZE+1, stride):\n            for x in range(0, W-PATCH_SIZE+1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.sum() / (PATCH_SIZE*PATCH_SIZE)\n                if self.ir_mask is not None:\n                    on_pap = self.ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    mid_z = z_list[len(z_list)//2]\n                    on_pap = float(self.cache[mid_z][y:y+PATCH_SIZE,\n                                                     x:x+PATCH_SIZE].mean()) > 0.1\n                if not on_pap: continue\n                if ink >= INK_MIN_POS:\n                    pos_yx.append((y,x))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_yx.append((y,x))\n\n        n_neg = int(len(pos_yx)*neg_ratio)\n        if n_neg > 0 and neg_yx:\n            np.random.shuffle(neg_yx); neg_yx = neg_yx[:n_neg]\n        else:\n            neg_yx = []\n\n        n_total = len(pos_yx)+len(neg_yx)\n        if max_patches > 0 and n_total > max_patches:\n            frac = max_patches/n_total\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx = pos_yx[:max(1,int(len(pos_yx)*frac))]\n            neg_yx = neg_yx[:max(0,int(len(neg_yx)*frac))]\n\n        all_yx  = pos_yx + neg_yx\n        all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n        perm    = np.random.permutation(len(all_yx))\n        self.coords  = np.array([all_yx[i] for i in perm], dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm], dtype=np.int32)\n        self.weights = np.where(self.labels==1, 3.0, 1.0).astype(np.float32)\n\n        cap = (f\" (capped from {n_total})\"\n               if max_patches > 0 and n_total > max_patches else \"\")\n        print(f\"  [{os.path.basename(frag_path)}]: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg \"\n              f\"= {len(self.coords)}{cap}\")\n\n    def __len__(self): return len(self.coords)\n\n    def get_patch(self, idx):\n        y, x = int(self.coords[idx,0]), int(self.coords[idx,1])\n        slices = [self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                  for z in self.z_list]\n        return (np.stack(slices, axis=-1),\n                self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy())\n\n    def __getitem__(self, idx):\n        img, msk = self.get_patch(idx)\n        if self.tf:\n            out = self.tf(image=img, mask=msk)\n            img, msk = out['image'], out['mask']\n        if self.ch_dropout:\n            img = channel_dropout(img)\n        return (torch.from_numpy(img).permute(2,0,1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\n\n# ════════════════════════════════════════════════════════════\n#  9. LOSS\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = target*alpha + (1-target)*(1-alpha)\n    return (a_t*((1-p_t)**gamma)*bce).mean()\n\ndef dice_loss(pred, target, smooth=1.):\n    p     = torch.sigmoid(pred)\n    inter = (p*target).sum(dim=(2,3))\n    union = p.sum(dim=(2,3))+target.sum(dim=(2,3))\n    return 1.-((2.*inter+smooth)/(union+smooth)).mean()\n\ndef combined_loss(pred, target, eps=0.05):\n    t_s = target*(1-eps)+0.5*eps\n    return 0.5*focal_loss(pred, t_s) + 0.5*dice_loss(pred, target)\n\ndef multiscale_loss(main_logit, sal_logit, ds3, ds2, ds1, target):\n    loss  = combined_loss(main_logit, target)\n    sal_t = F.adaptive_avg_pool2d(target, sal_logit.shape[-2:])\n    loss += 0.2*combined_loss(sal_logit, sal_t)\n    for logit, wt in [(ds3,0.15),(ds2,0.10),(ds1,0.05)]:\n        if logit is not None:\n            tgt = F.adaptive_avg_pool2d(target, logit.shape[-2:])\n            loss += wt*combined_loss(logit, tgt)\n    return loss\n\n\n# ════════════════════════════════════════════════════════════\n#  10. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p*masks).sum(dim=(1,2,3))\n    union = p.sum(dim=(1,2,3))+masks.sum(dim=(1,2,3))\n    return ((2.*inter+1e-5)/(union+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.25, 0.85, 0.02):\n        p = (probs>t).astype(np.float32)\n        d = (2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d > best_d: best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  11. GAUSSIAN WEIGHT\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c=sz//2; sig=sz//4\n    ys,xs=np.mgrid[0:sz,0:sz]\n    return np.exp(-((xs-c)**2+(ys-c)**2)/(2*sig**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  12. BUILD DATASETS\n#\n#  MODIFIED: train on 80% of Frag2+Frag3 combined.\n#  Val on remaining 20% of Frag2+Frag3.\n#  Test on full Fragment 1 (completely held out).\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('DATA PROTOCOL — MODIFIED')\nprint('  Train : 80% of Frag2 + Frag3 combined')\nprint('  Val   : 20% of Frag2 + Frag3 combined')\nprint('  Test  : Fragment 1 full fragment (completely held out)')\nprint('='*60)\n\ntrain_val_frag_paths = [FRAG2, FRAG3]\nall_datasets   = []\nfor fp in train_val_frag_paths:\n    print(f'\\n── {os.path.basename(fp)} ──')\n    ds = VesuviusDataset(fp, Z_SLICES, stride=STRIDE_TR,\n                         transform=train_tf, neg_ratio=NEG_RATIO,\n                         apply_ch_dropout=True,\n                         max_patches=MAX_PATCHES)\n    all_datasets.append(ds)\n\n# 80/20 split on combined pool of Frag2+Frag3\nfull_ds     = ConcatDataset(all_datasets)\nfull_weights= np.concatenate([d.weights for d in all_datasets])\nn_total     = len(full_ds)\nrng         = np.random.RandomState(SEED)\nperm        = rng.permutation(n_total)\nn_val       = int(n_total * 0.20)\nn_train     = n_total - n_val\ntrain_idx   = perm[:n_train].tolist()\nval_idx     = perm[n_train:].tolist()\nprint(f'\\nTotal (Frag2+Frag3): {n_total} | Train: {n_train} (80%) | Val: {n_val} (20%)')\n\n\nclass TrainSubset(Dataset):\n    \"\"\"Applies the full dataset's __getitem__ (with augmentation).\"\"\"\n    def __init__(self, full_ds, indices):\n        self.full_ds = full_ds\n        self.indices = indices\n    def __len__(self): return len(self.indices)\n    def __getitem__(self, idx): return self.full_ds[self.indices[idx]]\n\nclass ValSubset(Dataset):\n    \"\"\"Val: raw patch, no augmentation.\"\"\"\n    def __init__(self, all_ds_list, indices, n_per_ds):\n        self.all_ds   = all_ds_list\n        self.indices  = indices\n        self.n_per_ds = n_per_ds   # [n0, n1] cumulative\n\n    def __len__(self): return len(self.indices)\n\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        # find which dataset\n        di = 0\n        offset = g\n        for i, n in enumerate(self.n_per_ds):\n            if offset < n:\n                di = i; break\n            offset -= n\n        img, msk = self.all_ds[di].get_patch(offset)\n        return (torch.from_numpy(img).permute(2,0,1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\nn_per_ds  = [len(d) for d in all_datasets]\ntrain_ds  = TrainSubset(full_ds, train_idx)\nval_ds    = ValSubset(all_datasets, val_idx, n_per_ds)\n\ntrain_w   = torch.from_numpy(full_weights[train_idx]).float()\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nsampler  = WeightedRandomSampler(train_w, len(train_ds), replacement=True)\n_dl_kw   = dict(num_workers=NUM_WORKERS, pin_memory=PIN,\n                persistent_workers=(NUM_WORKERS>0),\n                prefetch_factor=2 if NUM_WORKERS>0 else None)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE,\n                      sampler=sampler, **_dl_kw)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE,\n                      shuffle=False, **_dl_kw)\nprint(f'Train batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  13. MODEL + OPTIMISER + SCHEDULER\n#  FIX 4: CosineAnnealingWarmRestarts — LR never collapses\n# ════════════════════════════════════════════════════════════\nmodel    = VesuviusV10(n_ch=N_CH, enc_dim=512,\n                       n_transformer_blocks=4, num_heads=8).to(DEVICE)\nn_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'\\nModel parameters: {n_params/1e6:.2f} M')\n\noptimizer = optim.AdamW(model.parameters(), lr=LR,\n                         weight_decay=WEIGHT_DECAY)\n# FIX 4: warm restarts every 10 epochs — never collapses\nscheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(\n    optimizer, T_0=10, T_mult=1, eta_min=1e-6)\n\nbest_dice = 0.; pat_cnt = 0\nhistory   = dict(tl=[], vl=[], td=[], vd=[], lr=[])\n\nprint('\\n' + '='*65)\nprint('TRAINING — VesuviusV10 FIXED + FRAGMENT SPLIT')\nprint('  Encoder      : ResNet34 + Axial-Transformer bottleneck')\nprint('  Attention    : Axial X/Y with log-polar positional bias')\nprint('  Token merge  : 25% (even-length forced — bug fixed)')\nprint('  Augmentation : Affine + Elastic + CutMix + ChannelDrop')\nprint('  Deep superv  : 3 auxiliary heads')\nprint('  Scheduler    : CosineAnnealingWarmRestarts(T_0=10)')\nprint(f'  Training data: Frag2 + Frag3 (80/20 split)')\nprint(f'  Test data    : Fragment 1 (held out completely)')\nprint('='*65)\n\n\n# ════════════════════════════════════════════════════════════\n#  14. TRAINING LOOP\n# ════════════════════════════════════════════════════════════\nfor epoch in range(EPOCHS):\n\n    model.train(); tl = td = 0.\n    optimizer.zero_grad()\n\n    for step, (imgs, msks) in enumerate(\n            tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        msks = msks.to(DEVICE, non_blocking=True)\n\n        if random.random() < 0.30:\n            imgs, msks = cutmix_batch(imgs, msks, alpha=0.4)\n\n        with amp_context():\n            main_logit, sal_logit, ds3_l, ds2_l, ds1_l = model(imgs)\n            loss = multiscale_loss(\n                main_logit, sal_logit, ds3_l, ds2_l, ds1_l,\n                msks) / GRAD_ACCUM\n\n        if USE_AMP:\n            scaler.scale(loss).backward()\n            if (step+1) % GRAD_ACCUM == 0 or (step+1) == len(train_dl):\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                scaler.step(optimizer); scaler.update()\n                optimizer.zero_grad()\n        else:\n            loss.backward()\n            if (step+1) % GRAD_ACCUM == 0 or (step+1) == len(train_dl):\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                optimizer.step()\n                optimizer.zero_grad()\n\n        tl += loss.item() * GRAD_ACCUM\n        td += batch_dice(main_logit.detach(), msks)\n        del imgs, msks, main_logit, sal_logit, ds3_l, ds2_l, ds1_l, loss\n        if step % 100 == 0 and DEVICE=='cuda':\n            torch.cuda.empty_cache()\n\n    tl /= len(train_dl); td /= len(train_dl)\n\n    # ── validate ──────────────────────────────────────────────\n    model.eval(); vl = vd = 0.\n    acc_p, acc_m = [], []\n    with torch.no_grad():\n        for imgs, msks in tqdm(val_dl, desc=f'Ep{epoch+1:02d} val  ',\n                               leave=False):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            with amp_context():\n                main_logit, sal_logit, ds3_l, ds2_l, ds1_l = model(imgs)\n                loss = multiscale_loss(\n                    main_logit, sal_logit, ds3_l, ds2_l, ds1_l, msks)\n            vl += loss.item()\n            vd += batch_dice(main_logit, msks)\n            acc_p.append(torch.sigmoid(main_logit/TEMPERATURE).cpu().numpy())\n            acc_m.append(msks.cpu().numpy())\n            del imgs, msks, main_logit, sal_logit, ds3_l, ds2_l, ds1_l, loss\n\n    vl /= len(val_dl); vd /= len(val_dl)\n    probs_all = np.concatenate(acc_p)\n    masks_all = np.concatenate(acc_m)\n    bt, bd    = sweep_threshold(probs_all, masks_all)\n    ink_mean  = float(probs_all[masks_all>0.5].mean()) \\\n                if (masks_all>0.5).any() else 0.\n    noink_mean= float(probs_all[masks_all<0.5].mean()) \\\n                if (masks_all<0.5).any() else 0.\n    del acc_p, acc_m, probs_all, masks_all\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    scheduler.step(epoch + 1/len(train_dl))   # step per epoch for CosineWarmRestarts\n    lr_now = optimizer.param_groups[0]['lr']\n\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['lr'].append(lr_now)\n\n    print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n          f'train loss={tl:.4f} dice={td:.4f} | '\n          f'val loss={vl:.4f} dice={vd:.4f} | '\n          f'thr={bt:.2f}→{bd:.4f} | '\n          f'sep={ink_mean-noink_mean:+.3f} '\n          f'[ink={ink_mean:.3f} bg={noink_mean:.3f}]')\n\n    save_metric = max(vd, bd)\n    if save_metric > best_dice:\n        best_dice = save_metric; pat_cnt = 0\n        torch.save({'epoch': epoch, 'state': model.state_dict(),\n                    'thr': bt, 'metric': save_metric,\n                    'n_ch': N_CH, 'bd': bd, 'vd': vd},\n                   OUTPUT+'best_model.pth')\n        print(f'  ✓ saved  (metric={save_metric:.4f}  '\n              f'raw_bd={bd:.4f}  thr={bt:.2f})')\n    else:\n        pat_cnt += 1\n        if pat_cnt >= PATIENCE:\n            print(f'  ⚑ early stop at epoch {epoch+1}'); break\n\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\n# curves\nfig, axes = plt.subplots(1,3,figsize=(16,4))\naxes[0].plot(history['tl'],label='train'); axes[0].plot(history['vl'],label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'],label='train'); axes[1].plot(history['vd'],label='val')\naxes[1].axhline(0.65,color='orange',ls='--',label='0.65')\naxes[1].axhline(0.80,color='r',ls='--',label='0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['lr'])\naxes[2].set_title('LR'); axes[2].grid(True)\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves.png',dpi=100); plt.close()\n\ndel train_dl, val_dl, sampler, train_ds, val_ds, full_ds\nfor ds in all_datasets: del ds\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n\n# ════════════════════════════════════════════════════════════\n#  15. INFERENCE  (NO TTA)\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list, temperature=TEMPERATURE):\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    msk   = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n    H, W  = msk.shape\n    pred_map= np.zeros((H,W), np.float32)\n    wgt_map = np.zeros((H,W), np.float32)\n    coords  = [(y,x)\n               for y in range(0, H-PATCH_SIZE+1, STRIDE_INF)\n               for x in range(0, W-PATCH_SIZE+1, STRIDE_INF)]\n    with torch.no_grad():\n        for (y,x) in tqdm(coords, desc='Inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            patch  = np.stack(slices, -1)\n            t      = torch.from_numpy(patch).permute(2,0,1).unsqueeze(0).to(DEVICE)\n            with amp_context():\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit/temperature).squeeze().cpu().numpy().astype(np.float32)\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p*GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, logit\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n    del cache; gc.collect()\n    return pred_map/(wgt_map+1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  16. POST-PROCESSING\n# ════════════════════════════════════════════════════════════\ndef morphological_clean(binary_map, min_area=200, close_k=5):\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(close_k,close_k))\n    closed = cv2.morphologyEx(binary_map.astype(np.uint8), cv2.MORPH_CLOSE, kernel)\n    n_lab,labels,stats,_ = cv2.connectedComponentsWithStats(closed)\n    cleaned = np.zeros_like(closed)\n    for i in range(1, n_lab):\n        if stats[i, cv2.CC_STAT_AREA] >= min_area:\n            cleaned[labels==i] = 1\n    return cleaned\n\ndef apply_dense_crf(image_uint8, prob_map, n_iter=5):\n    if not HAS_CRF:\n        return (prob_map > 0.5).astype(np.uint8)\n    H, W = prob_map.shape\n    d    = dcrf.DenseCRF2D(W, H, 2)\n    fg   = np.clip(prob_map,     1e-5, 1-1e-5)\n    bg   = np.clip(1-prob_map,   1e-5, 1-1e-5)\n    U    = -np.log(np.stack([bg,fg],axis=0))\n    d.setUnaryEnergy(U.reshape(2,-1).astype(np.float32))\n    d.addPairwiseGaussian(sxy=3, compat=3)\n    img_c = np.ascontiguousarray(image_uint8)\n    d.addPairwiseBilateral(sxy=50, srgb=13, rgbim=img_c, compat=10)\n    Q = d.inference(n_iter)\n    return np.argmax(Q,axis=0).reshape(H,W).astype(np.uint8)\n\ndef postprocess(prob_map, cache, z_list, threshold, min_area=200, close_k=5):\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = morphological_clean(binary, min_area=min_area, close_k=close_k)\n    if HAS_CRF:\n        mid_z  = z_list[len(z_list)//2]\n        mid_sl = cache[mid_z].astype(np.float32)\n        mid_8  = ((mid_sl-mid_sl.min())/(mid_sl.max()-mid_sl.min()+1e-8)*255).astype(np.uint8)\n        rgb    = np.stack([mid_8,mid_8,mid_8],-1)\n        binary = apply_dense_crf(rgb, prob_map, n_iter=5)\n    return binary\n\n\n# ════════════════════════════════════════════════════════════\n#  17. FINAL TEST — FRAGMENT 1 (HELD OUT COMPLETELY)\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*65)\nprint('FINAL TEST — FRAGMENT 1 (COMPLETELY HELD OUT)')\nprint('Single deterministic inference pass — no TTA')\nprint('Model trained exclusively on Frag2+Frag3')\nprint('='*65)\n\n# weights_only=False for PyTorch 2.6 compatibility\nckpt = torch.load(OUTPUT+'best_model.pth',\n                  map_location=DEVICE, weights_only=False)\nassert ckpt.get('n_ch', N_CH) == N_CH\n\nmodel = VesuviusV10(n_ch=N_CH, enc_dim=512,\n                    n_transformer_blocks=4, num_heads=8).to(DEVICE)\nmodel.load_state_dict(ckpt['state'], strict=True)\nsaved_thr = ckpt.get('thr', 0.5)\nprint(f'Checkpoint: epoch={ckpt[\"epoch\"]+1}  '\n      f'val_dice={ckpt.get(\"vd\",0):.4f}  '\n      f'best_dice={ckpt.get(\"bd\",0):.4f}  '\n      f'thr={saved_thr:.2f}')\n\nprob_map, msk1 = predict_fragment(model, FRAG1, Z_SLICES)\nH_m, W_m = msk1.shape\nprob_crop = prob_map[:H_m,:W_m]\n\nbt1, bd1 = sweep_threshold(prob_crop[np.newaxis,np.newaxis],\n                            msk1[np.newaxis,np.newaxis])\n\nprint(f'\\nPost-processing: thr={bt1:.2f}, morph clean, DenseCRF={HAS_CRF} ...')\ncache1     = load_slice_cache(FRAG1, Z_SLICES)\nfinal_pred = postprocess(prob_crop, cache1, Z_SLICES,\n                         threshold=bt1, min_area=150, close_k=5)\ndel cache1; gc.collect()\n\n# metrics\npf = final_pred.flatten().astype(int)\nmf = msk1.flatten().astype(int)\ntn,fp,fn,tp_v = confusion_matrix(mf,pf,labels=[0,1]).ravel()\nprec = tp_v/(tp_v+fp+1e-8); rec = tp_v/(tp_v+fn+1e-8)\nf1   = 2*prec*rec/(prec+rec+1e-8)\ndice_full = (2*tp_v+1)/(final_pred.sum()+msk1.sum()+1)\nink_mean   = float(prob_crop[msk1==1].mean())\nnoink_mean = float(prob_crop[msk1==0].mean())\n\nprint('\\n' + '='*65)\nprint('RESULTS — FRAGMENT 1 (HELD OUT)')\nprint('='*65)\nprint(f'Dice Score  : {dice_full:.4f}')\nprint(f'F1 Score    : {f1:.4f}')\nprint(f'Precision   : {prec:.4f}')\nprint(f'Recall      : {rec:.4f}')\nprint(f'Threshold   : {bt1:.2f}')\nprint(f'TP={tp_v} TN={tn} FP={fp} FN={fn}')\nprint(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}')\nprint(f'Sep         : {ink_mean-noink_mean:+.3f}')\nprint(f'F1 ≥ 0.80   : {\"✓ PASSED\" if f1>=0.80 else \"✗ not yet\"}')\nprint('='*65)\n\n# visualise\nfig, ax = plt.subplots(2,3,figsize=(18,12))\nax[0,0].imshow(msk1,       cmap='gray');    ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop,  cmap='inferno'); ax[0,1].set_title('Probability Map')\nax[0,2].imshow(final_pred, cmap='gray');\nax[0,2].set_title(f'Prediction  Dice={dice_full:.4f}')\nerr=np.zeros((*msk1.shape,3),dtype=np.uint8)\nerr[(final_pred==1)&(msk1==1)]=[0,255,0]\nerr[(final_pred==1)&(msk1==0)]=[255,0,0]\nerr[(final_pred==0)&(msk1==1)]=[0,0,255]\nax[1,0].imshow(err); ax[1,0].set_title('TP=green FP=red FN=blue')\nax[1,1].hist(prob_crop[msk1==1].ravel(),bins=50,alpha=0.7,\n             label=f'ink μ={ink_mean:.2f}',color='orange',density=True)\nax[1,1].hist(prob_crop[msk1==0].ravel(),bins=50,alpha=0.7,\n             label=f'bg μ={noink_mean:.2f}',color='blue',density=True)\nax[1,1].axvline(bt1,color='r',ls='--',label=f'thr={bt1:.2f}')\nax[1,1].set_title('Probability Distribution'); ax[1,1].legend()\nts=np.arange(0.15,0.90,0.01); ds=[]\nfor t in ts:\n    p_b=(prob_crop>t).astype(np.float32)\n    ds.append((2*(p_b*msk1).sum()+1)/(p_b.sum()+msk1.sum()+1))\nax[1,2].plot(ts,ds); ax[1,2].axvline(bt1,color='r',ls='--')\nax[1,2].axhline(0.65,color='orange',ls=':'); ax[1,2].axhline(0.80,color='g',ls=':')\nax[1,2].set_title('Dice vs Threshold'); ax[1,2].grid(True)\nfor a in [ax[0,0],ax[0,1],ax[0,2],ax[1,0]]: a.axis('off')\nplt.suptitle(f'Fragment 1 — Dice={dice_full:.4f}  F1={f1:.4f}  (Trained on Frag2+Frag3 only)',\n             fontsize=13, y=1.01)\nplt.tight_layout()\nplt.savefig(OUTPUT+'frag1_prediction.png',dpi=100,bbox_inches='tight')\nplt.close()\n\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write('VESUVIUS V10 FIXED + FRAGMENT SPLIT — FINAL TEST RESULTS\\n')\n    f.write('='*50+'\\n')\n    f.write(f'Dice        : {dice_full:.4f}\\n')\n    f.write(f'F1          : {f1:.4f}\\n')\n    f.write(f'Precision   : {prec:.4f}\\n')\n    f.write(f'Recall      : {rec:.4f}\\n')\n    f.write(f'Threshold   : {bt1:.2f}\\n')\n    f.write(f'TP={tp_v} TN={tn} FP={fp} FN={fn}\\n')\n    f.write(f'FP/TP       : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Sep         : {ink_mean-noink_mean:+.3f}\\n')\n    f.write(f'Post-proc   : morph_clean + DenseCRF={HAS_CRF}\\n')\n    f.write(f'Training    : 80% Frag2+Frag3 (Frag1 completely held out)\\n')\n    f.write(f'Scheduler   : CosineAnnealingWarmRestarts(T_0=10)\\n')\n    f.write(f'Token merge : even-length forced (bug fixed)\\n')\n    f.write(f'AMP API     : {\"torch.amp\" if \"amp\" in str(type(scaler)) else \"torch.cuda.amp\" if scaler else \"disabled\"}\\n')\n    f.write('='*50+'\\n')\n\ncv2.imwrite(OUTPUT+'frag1_prob_map.png',  (prob_crop*255).astype(np.uint8))\ncv2.imwrite(OUTPUT+'frag1_pred_binary.png',(final_pred*255).astype(np.uint8))\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  best_model.pth | curves.png | frag1_prediction.png')\nprint('  frag1_prob_map.png | frag1_pred_binary.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test Dice Score: 0.5015 F1 Score: 0.5416 ON Fragment-1\n# 2D UNet PATCH_SIZE = 352 STRIDE = 26 Dice=0.35\n# 2D UNet PATCH_SIZE = 160 STRIDE = 32 Dice=0.41\n#PATCH_SIZE = 192 ,STRIDE = 32, Dice Score: 0.43 F1 Score: 0.48 ON Fragment-1\n#PATCH_SIZE = 128 ,STRIDE = 16, Dice Score: 0.46 F1 Score: 0.50 ON Fragment-1\n\n!pip install segmentation-models-pytorch==0.2.0\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport torch.nn.functional as F\nimport gc\nimport psutil\nimport time\n#/kaggle/input/vesuvius-challenge-ink-detection/test\n\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n]\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\nelse:\n    print(\"CPU mode\")\n\n# Memory-efficient settings\nPATCH_SIZE = 128\nSTRIDE = 4\nBATCH_SIZE = 14  # Increased slightly since we're using 2D\nACCUMULATION_STEPS = 2  # Gradient accumulation\nEPOCHS = 100\nLR = 1e-4\nWEIGHT_DECAY = 1e-4\nSLICE_START = 18\nSLICE_END = 30  # 12 slices\nOUTPUT_DIR = \"/kaggle/working/\"\nVALIDATION_SPLIT = 0.20\n\n# Advanced training settings\nUSE_AMP = True\nPIN_MEMORY = True\nNUM_WORKERS = 2\nGRAD_CLIP = 1.0\nSCHEDULER_PATIENCE = 15\nEARLY_STOPPING_PATIENCE = 15\n\n# ============================================\n# IGNORE MASK GENERATION\n# ============================================\ndef generate_ignore_mask(ink_mask, distance_threshold=3, erosion_size=2):\n    \"\"\"Generate ignore mask for uncertain regions - OPTIMIZED VERSION\"\"\"\n    ignore_mask = np.zeros_like(ink_mask, dtype=np.uint8)\n    \n    # Simple boundary detection (faster than full distance transform)\n    kernel = np.ones((3, 3), np.uint8)\n    eroded = cv2.erode(ink_mask.astype(np.uint8), kernel, iterations=1)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    \n    # Boundaries are where eroded and dilated differ from original\n    boundaries = (dilated != eroded)\n    ignore_mask[boundaries] = 1\n    \n    # Remove small isolated ink dots\n    contours, _ = cv2.findContours(ink_mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for contour in contours:\n        if cv2.contourArea(contour) < 10:\n            cv2.drawContours(ignore_mask, [contour], -1, 1, -1)\n    \n    return ignore_mask\n\n# ============================================\n# 2D CNN WITH MULTI-SLICE INPUT\n# ============================================\nclass MultiSlice2DUNet(nn.Module):\n    \"\"\"2D CNN that treats depth slices as input channels\"\"\"\n    def __init__(self, in_channels=12, out_channels=1):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._conv_block(in_channels, 32)\n        self.enc2 = self._conv_block(32, 64)\n        self.enc3 = self._conv_block(64, 128)\n        self.enc4 = self._conv_block(128, 256)\n        \n        # Bottleneck\n        self.bottleneck = self._conv_block(256, 512)\n        \n        # Decoder\n        self.dec4 = self._upconv_block(512 + 256, 256)\n        self.dec3 = self._upconv_block(256 + 128, 128)\n        self.dec2 = self._upconv_block(128 + 64, 64)\n        self.dec1 = self._upconv_block(64 + 32, 32)\n        \n        # Output\n        self.final = nn.Conv2d(32, out_channels, kernel_size=1)\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False)\n        \n    def _conv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def _upconv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # x shape: [B, C=12, H, W]\n        \n        # Encoder\n        e1 = self.enc1(x)      # [B, 32, H, W]\n        p1 = self.pool(e1)      # [B, 32, H/2, W/2]\n        \n        e2 = self.enc2(p1)      # [B, 64, H/2, W/2]\n        p2 = self.pool(e2)      # [B, 64, H/4, W/4]\n        \n        e3 = self.enc3(p2)      # [B, 128, H/4, W/4]\n        p3 = self.pool(e3)      # [B, 128, H/8, W/8]\n        \n        e4 = self.enc4(p3)      # [B, 256, H/8, W/8]\n        p4 = self.pool(e4)      # [B, 256, H/16, W/16]\n        \n        # Bottleneck\n        b = self.bottleneck(p4)  # [B, 512, H/16, W/16]\n        \n        # Decoder with skip connections\n        d4 = self.upsample(b)    # [B, 512, H/8, W/8]\n        d4 = torch.cat([d4, e4], dim=1)  # [B, 512+256, H/8, W/8]\n        d4 = self.dec4(d4)        # [B, 256, H/8, W/8]\n        \n        d3 = self.upsample(d4)    # [B, 256, H/4, W/4]\n        d3 = torch.cat([d3, e3], dim=1)  # [B, 256+128, H/4, W/4]\n        d3 = self.dec3(d3)        # [B, 128, H/4, W/4]\n        \n        d2 = self.upsample(d3)    # [B, 128, H/2, W/2]\n        d2 = torch.cat([d2, e2], dim=1)  # [B, 128+64, H/2, W/2]\n        d2 = self.dec2(d2)        # [B, 64, H/2, W/2]\n        \n        d1 = self.upsample(d2)    # [B, 64, H, W]\n        d1 = torch.cat([d1, e1], dim=1)  # [B, 64+32, H, W]\n        d1 = self.dec1(d1)        # [B, 32, H, W]\n        \n        # Final output\n        out = self.final(d1)      # [B, 1, H, W]\n        \n        return out\n\n# ============================================\n# DATA LOADING (OPTIMIZED)\n# ============================================\ndef load_volume_fast(fragment_path):\n    \"\"\"Fast volume loading\"\"\"\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    \n    volume = np.stack(slices, axis=-1)\n    volume = volume.astype(np.float32)\n    \n    # Fast normalization per slice\n    for i in range(volume.shape[-1]):\n        slice_data = volume[:, :, i]\n        mean, std = slice_data.mean(), slice_data.std()\n        volume[:, :, i] = (slice_data - mean) / (std + 1e-6)\n    \n    return volume\n\ndef extract_patches_fast(volume, mask, ignore_mask, max_patches_per_fragment=3000):\n    \"\"\"Fast patch extraction with sampling\"\"\"\n    patches = []\n    mask_patches = []\n    ignore_patches = []\n    \n    H, W, _ = volume.shape\n    \n    # Calculate number of patches\n    n_y = (H - PATCH_SIZE) // STRIDE + 1\n    n_x = (W - PATCH_SIZE) // STRIDE + 1\n    total_patches = n_y * n_x\n    \n    print(f\"    Total possible patches: {total_patches}\")\n    \n    # Sample patches if too many\n    if total_patches > max_patches_per_fragment:\n        print(f\"    Sampling {max_patches_per_fragment} patches...\")\n        # Calculate stride to get roughly max_patches\n        stride_y = max(STRIDE, (H - PATCH_SIZE) // int(np.sqrt(max_patches_per_fragment)))\n        stride_x = max(STRIDE, (W - PATCH_SIZE) // int(np.sqrt(max_patches_per_fragment)))\n    else:\n        stride_y, stride_x = STRIDE, STRIDE\n    \n    patch_count = 0\n    ink_patches = 0\n    bg_patches = 0\n    \n    for y in range(0, H - PATCH_SIZE, stride_y):\n        for x in range(0, W - PATCH_SIZE, stride_x):\n            if patch_count >= max_patches_per_fragment:\n                break\n                \n            v_patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            m_patch = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            i_patch = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            \n            # Count ink pixels in this patch\n            ink_pixel_count = m_patch.sum()\n            \n            # Keep patches with significant ink or some background for balance\n            if ink_pixel_count > 50:  # Good ink patch\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n                ink_patches += 1\n            elif ink_pixel_count == 0 and bg_patches < ink_patches // 2:  # Balance with background\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n                bg_patches += 1\n        \n        if patch_count >= max_patches_per_fragment:\n            break\n    \n    print(f\"    Extracted: {ink_patches} ink patches, {bg_patches} background patches\")\n    return patches, mask_patches, ignore_patches\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignore_masks, transform=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.ignore_masks = ignore_masks\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        image = self.volumes[idx].copy()\n        mask = self.masks[idx].copy()\n        ignore = self.ignore_masks[idx].copy()\n        \n        if self.transform:\n            transformed = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image = transformed['image']\n            mask = transformed['mask']\n            ignore = transformed['ignore_mask']\n        \n        # Convert to tensors\n        # For 2D CNN: image shape [H, W, C] -> [C, H, W]\n        image = torch.tensor(image).permute(2, 0, 1).float()  # [C=12, H, W]\n        \n        # mask and ignore shape: [H, W]\n        mask = torch.tensor(mask).float().unsqueeze(0)  # [1, H, W]\n        ignore = torch.tensor(ignore).float().unsqueeze(0)  # [1, H, W]\n        \n        return image, mask, ignore\n\n# ============================================\n# LOSS FUNCTION\n# ============================================\nclass DiceBCELossWithIgnore(nn.Module):\n    def __init__(self, dice_weight=0.5, bce_weight=0.5, smooth=1e-6):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.smooth = smooth\n        \n    def forward(self, pred, target, ignore_mask):\n        \"\"\"\n        pred: [B, 1, H, W] - logits\n        target: [B, 1, H, W] - binary mask\n        ignore_mask: [B, 1, H, W] - 1 for ignore, 0 for keep\n        \"\"\"\n        # Create valid mask\n        valid_mask = (1 - ignore_mask).float()\n        \n        # BCE loss\n        bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n        bce = (bce * valid_mask).sum() / (valid_mask.sum() + self.smooth)\n        \n        # Dice loss\n        pred_probs = torch.sigmoid(pred)\n        \n        # Apply valid mask\n        pred_valid = pred_probs * valid_mask\n        target_valid = target * valid_mask\n        \n        intersection = (pred_valid * target_valid).sum()\n        union = pred_valid.sum() + target_valid.sum()\n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        dice_loss = 1 - dice\n        \n        return self.dice_weight * dice_loss + self.bce_weight * bce\n\n# ============================================\n# FAST TRAIN/VALIDATION SPLIT\n# ============================================\ndef fast_train_val_split(n_samples, val_ratio=0.15, seed=42):\n    \"\"\"Fast random split without distance constraints\"\"\"\n    np.random.seed(seed)\n    indices = np.random.permutation(n_samples)\n    split = int(n_samples * val_ratio)\n    return indices[split:], indices[:split]\n\n# ============================================\n# MEMORY MONITORING\n# ============================================\ndef print_memory_usage():\n    if DEVICE == 'cuda':\n        allocated = torch.cuda.memory_allocated() / 1e9\n        cached = torch.cuda.memory_reserved() / 1e9\n        print(f\"    GPU Memory - Allocated: {allocated:.2f}GB, Cached: {cached:.2f}GB\")\n    \n    process = psutil.Process()\n    print(f\"    CPU Memory: {process.memory_info().rss / 1e9:.2f}GB\")\n\n# ============================================\n# COLLATE FUNCTION\n# ============================================\ndef collate_fn(batch):\n    \"\"\"Custom collate function to ensure correct dimensions\"\"\"\n    images = torch.stack([item[0] for item in batch])  # [B, C=12, H, W]\n    masks = torch.stack([item[1] for item in batch])   # [B, 1, H, W]\n    ignores = torch.stack([item[2] for item in batch]) # [B, 1, H, W]\n    return images, masks, ignores\n\n# ============================================\n# FULL VOLUME PREDICTION FUNCTION\n# ============================================\ndef predict_full_volume(model, volume, device, batch_size=8):\n    \"\"\"\n    Predict binary mask for the full volume by stacking predictions from all layers\n    \"\"\"\n    model.eval()\n    H, W, C = volume.shape\n    \n    # Initialize prediction accumulator for the surface\n    surface_prediction = np.zeros((H, W), dtype=np.float32)\n    prediction_count = np.zeros((H, W), dtype=np.float32)\n    \n    # Process the volume in patches\n    stride = PATCH_SIZE // 2  # Use 50% overlap for smoother predictions\n    \n    patches = []\n    positions = []\n    \n    # Extract all patches\n    for y in range(0, H - PATCH_SIZE + 1, stride):\n        for x in range(0, W - PATCH_SIZE + 1, stride):\n            patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE, :]\n            \n            # Normalize patch (already normalized, but ensure correct format)\n            patch_tensor = torch.tensor(patch).permute(2, 0, 1).float().unsqueeze(0)\n            patches.append(patch_tensor)\n            positions.append((y, x))\n    \n    # Process in batches\n    all_predictions = []\n    with torch.no_grad():\n        for i in range(0, len(patches), batch_size):\n            batch_patches = torch.cat(patches[i:i+batch_size], dim=0).to(device)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(batch_patches)\n                preds = torch.sigmoid(outputs).cpu().numpy()[:, 0, :, :]\n            \n            all_predictions.extend(preds)\n    \n    # Stitch predictions together\n    for (y, x), pred in zip(positions, all_predictions):\n        surface_prediction[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += pred\n        prediction_count[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += 1\n    \n    # Average overlapping regions\n    prediction_count[prediction_count == 0] = 1\n    surface_prediction /= prediction_count\n    \n    # Convert to binary\n    binary_prediction = (surface_prediction > 0.5).astype(np.uint8)\n    \n    return binary_prediction, surface_prediction\n\n# ============================================\n# VISUALIZATION FUNCTION\n# ============================================\ndef create_test_visualization(input_volume, ground_truth, prediction, save_path, slice_idx=0):\n    \"\"\"\n    Create black and white visualization comparing input, ground truth, and prediction\n    \"\"\"\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    \n    # Input volume (show middle slice)\n    input_slice = input_volume[:, :, slice_idx]\n    axes[0].imshow(input_slice, cmap='gray')\n    axes[0].set_title('Input Volume (Slice {})'.format(slice_idx), fontsize=14)\n    axes[0].axis('off')\n    \n    # Ground truth (black and white)\n    axes[1].imshow(ground_truth, cmap='binary')\n    axes[1].set_title('Ground Truth', fontsize=14)\n    axes[1].axis('off')\n    \n    # Prediction (black and white)\n    axes[2].imshow(prediction, cmap='binary')\n    axes[2].set_title('Prediction (Binary)', fontsize=14)\n    axes[2].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"Visualization saved to: {save_path}\")\n\ndef create_comparison_visualization(ground_truth, prediction, save_path):\n    \"\"\"\n    Create detailed black and white comparison visualization\n    \"\"\"\n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n    \n    # Row 1: Full images\n    axes[0, 0].imshow(ground_truth, cmap='binary')\n    axes[0, 0].set_title('Ground Truth', fontsize=14)\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(prediction, cmap='binary')\n    axes[0, 1].set_title('Prediction', fontsize=14)\n    axes[0, 1].axis('off')\n    \n    # Difference map\n    diff = np.zeros((*ground_truth.shape, 3), dtype=np.uint8)\n    diff[ground_truth == 1] = [255, 255, 255]  # White for ground truth ink\n    diff[prediction == 1] = [255, 255, 255]    # White for predicted ink\n    \n    # True Positives: White\n    tp_mask = (ground_truth == 1) & (prediction == 1)\n    # False Positives: Red\n    fp_mask = (ground_truth == 0) & (prediction == 1)\n    # False Negatives: Blue\n    fn_mask = (ground_truth == 1) & (prediction == 0)\n    \n    diff[tp_mask] = [255, 255, 255]  # White for correct predictions\n    diff[fp_mask] = [255, 0, 0]      # Red for false positives\n    diff[fn_mask] = [0, 0, 255]      # Blue for false negatives\n    \n    axes[0, 2].imshow(diff)\n    axes[0, 2].set_title('Difference Map\\nWhite: Correct, Red: FP, Blue: FN', fontsize=14)\n    axes[0, 2].axis('off')\n    \n    # Row 2: Zoomed regions (center region)\n    h, w = ground_truth.shape\n    zoom_size = 400\n    h_start, w_start = h//2 - zoom_size//2, w//2 - zoom_size//2\n    \n    gt_zoom = ground_truth[h_start:h_start+zoom_size, w_start:w_start+zoom_size]\n    pred_zoom = prediction[h_start:h_start+zoom_size, w_start:w_start+zoom_size]\n    diff_zoom = diff[h_start:h_start+zoom_size, w_start:w_start+zoom_size]\n    \n    axes[1, 0].imshow(gt_zoom, cmap='binary')\n    axes[1, 0].set_title('Ground Truth (Zoomed)', fontsize=14)\n    axes[1, 0].axis('off')\n    \n    axes[1, 1].imshow(pred_zoom, cmap='binary')\n    axes[1, 1].set_title('Prediction (Zoomed)', fontsize=14)\n    axes[1, 1].axis('off')\n    \n    axes[1, 2].imshow(diff_zoom)\n    axes[1, 2].set_title('Difference Map (Zoomed)', fontsize=14)\n    axes[1, 2].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"Comparison visualization saved to: {save_path}\")\n\n# ============================================\n# MAIN PIPELINE\n# ============================================\ndef main():\n    print(\"=\"*60)\n    print(\"VESUVIUS INK DETECTION - OPTIMIZED PIPELINE\")\n    print(\"=\"*60)\n    print_memory_usage()\n    \n    start_time = time.time()\n    \n    # ============================================\n    # LOAD AND PREPARE DATA\n    # ============================================\n    print(\"\\n1. Loading training data...\")\n    all_patches = []\n    all_masks = []\n    all_ignores = []\n    \n    for path in train_paths:\n        fragment_name = os.path.basename(path)\n        print(f\"\\n   Processing Fragment {fragment_name}...\")\n        \n        # Load volume\n        volume = load_volume_fast(path)\n        print(f\"    Volume shape: {volume.shape}\")\n        \n        # Load mask\n        mask_path = os.path.join(path, \"inklabels.png\")\n        mask = cv2.imread(mask_path, 0)\n        mask = (mask > 0).astype(np.uint8)\n        print(f\"    Mask shape: {mask.shape}\")\n        print(f\"    Ink pixels: {mask.sum():,}\")\n        \n        # Generate ignore mask\n        print(\"    Generating ignore mask...\")\n        ignore_mask = generate_ignore_mask(mask)\n        print(f\"    Ignored pixels: {ignore_mask.sum():,}\")\n        \n        # Extract patches\n        print(\"    Extracting patches...\")\n        patches, mask_patches, ignore_patches = extract_patches_fast(\n            volume, mask, ignore_mask, max_patches_per_fragment=3000\n        )\n        \n        all_patches.extend(patches)\n        all_masks.extend(mask_patches)\n        all_ignores.extend(ignore_patches)\n        \n        print(f\"    Total extracted: {len(patches)} patches\")\n        print_memory_usage()\n        \n        # Clean up\n        del volume, mask, ignore_mask, patches, mask_patches, ignore_patches\n        gc.collect()\n        if DEVICE == 'cuda':\n            torch.cuda.empty_cache()\n    \n    print(f\"\\nTotal patches: {len(all_patches)}\")\n    print_memory_usage()\n    \n    # ============================================\n    # CREATE TRAIN/VAL SPLIT\n    # ============================================\n    print(\"\\n2. Creating train/validation split...\")\n    n_samples = len(all_patches)\n    train_indices, val_indices = fast_train_val_split(n_samples, VALIDATION_SPLIT)\n    \n    print(f\"   Train samples: {len(train_indices)}\")\n    print(f\"   Validation samples: {len(val_indices)}\")\n    \n    # Split data\n    train_patches = [all_patches[i] for i in train_indices]\n    train_masks = [all_masks[i] for i in train_indices]\n    train_ignores = [all_ignores[i] for i in train_indices]\n    \n    val_patches = [all_patches[i] for i in val_indices]\n    val_masks = [all_masks[i] for i in val_indices]\n    val_ignores = [all_ignores[i] for i in val_indices]\n    \n    # Clean up original lists\n    del all_patches, all_masks, all_ignores\n    gc.collect()\n    \n    # ============================================\n    # DATA AUGMENTATION\n    # ============================================\n    print(\"\\n3. Setting up data augmentation...\")\n    \n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),\n        A.GaussNoise(var_limit=(0, 0.01), p=0.2),\n    ], additional_targets={'ignore_mask': 'mask'})\n    \n    # Create datasets\n    train_dataset = VesuviusDataset(\n        train_patches, train_masks, train_ignores, \n        transform=train_transform\n    )\n    \n    val_dataset = VesuviusDataset(\n        val_patches, val_masks, val_ignores,\n        transform=None\n    )\n    \n    # Create data loaders with custom collate function\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=True, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        drop_last=True,\n        collate_fn=collate_fn\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    print(f\"   Train batches: {len(train_loader)}\")\n    print(f\"   Val batches: {len(val_loader)}\")\n    \n    # ============================================\n    # MODEL INITIALIZATION\n    # ============================================\n    print(\"\\n4. Initializing model...\")\n    \n    model = MultiSlice2DUNet(in_channels=12, out_channels=1).to(DEVICE)\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"   Total parameters: {total_params:,}\")\n    print(f\"   Trainable parameters: {trainable_params:,}\")\n    \n    # ============================================\n    # LOSS, OPTIMIZER, SCHEDULER\n    # ============================================\n    criterion = DiceBCELossWithIgnore(dice_weight=0.5, bce_weight=0.5)\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='max', factor=0.5, patience=SCHEDULER_PATIENCE, verbose=True\n    )\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n    \n    # ============================================\n    # TRAINING LOOP\n    # ============================================\n    print(\"\\n5. Starting training...\")\n    print(\"=\"*60)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    train_losses = []\n    val_dice_scores = []\n    \n    for epoch in range(EPOCHS):\n        epoch_start = time.time()\n        \n        # Training phase\n        model.train()\n        train_loss = 0\n        train_steps = 0\n        optimizer.zero_grad()\n        \n        progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Train]')\n        for batch_idx, (images, masks, ignores) in enumerate(progress_bar):\n            # images shape: [B, C=12, H, W]\n            # masks shape: [B, 1, H, W]\n            # ignores shape: [B, 1, H, W]\n            \n            images = images.to(DEVICE, non_blocking=True)\n            masks = masks.to(DEVICE, non_blocking=True)\n            ignores = ignores.to(DEVICE, non_blocking=True)\n            \n            # Forward pass\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)  # [B, 1, H, W]\n                loss = criterion(outputs, masks, ignores)\n                loss = loss / ACCUMULATION_STEPS\n            \n            # Backward pass\n            scaler.scale(loss).backward()\n            \n            # Gradient accumulation\n            if (batch_idx + 1) % ACCUMULATION_STEPS == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            train_loss += loss.item() * ACCUMULATION_STEPS\n            train_steps += 1\n            \n            # Update progress bar\n            progress_bar.set_postfix({'loss': f'{loss.item() * ACCUMULATION_STEPS:.4f}'})\n            \n            # Clear cache periodically\n            if batch_idx % 50 == 49:\n                if DEVICE == 'cuda':\n                    torch.cuda.empty_cache()\n        \n        avg_train_loss = train_loss / train_steps\n        train_losses.append(avg_train_loss)\n        \n        # Validation phase\n        model.eval()\n        val_dice = 0\n        val_steps = 0\n        \n        with torch.no_grad():\n            for images, masks, ignores in tqdm(val_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Val]'):\n                images = images.to(DEVICE, non_blocking=True)\n                masks = masks.to(DEVICE, non_blocking=True)\n                ignores = ignores.to(DEVICE, non_blocking=True)\n                \n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    outputs = model(images)\n                \n                # Calculate dice\n                preds = torch.sigmoid(outputs) > 0.5\n                valid_mask = (1 - ignores)\n                \n                intersection = ((preds * masks) * valid_mask).sum()\n                union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n                dice = (2 * intersection) / (union + 1e-6)\n                val_dice += dice.item()\n                val_steps += 1\n        \n        avg_val_dice = val_dice / val_steps\n        val_dice_scores.append(avg_val_dice)\n        \n        # Update scheduler\n        scheduler.step(avg_val_dice)\n        \n        epoch_time = time.time() - epoch_start\n        \n        # Print results\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS} - Time: {epoch_time:.1f}s\")\n        print(f\"  Train Loss: {avg_train_loss:.4f}\")\n        print(f\"  Val Dice: {avg_val_dice:.4f}\")\n        print(f\"  LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        print_memory_usage()\n        \n        # Save best model\n        if avg_val_dice > best_val_dice:\n            best_val_dice = avg_val_dice\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_dice': avg_val_dice,\n            }, os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_counter = 0\n            print(f\"  ✓ New best model saved! Dice: {best_val_dice:.4f}\")\n        else:\n            patience_counter += 1\n        \n        # Early stopping\n        if patience_counter >= EARLY_STOPPING_PATIENCE:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n        \n        print(\"-\"*60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING COMPLETE\")\n    print(f\"Best Validation Dice: {best_val_dice:.4f}\")\n    print(\"=\"*60)\n    \n    # ============================================\n    # FINAL TEST ON FRAGMENT 1\n    # ============================================\n    print(\"\\n6. Testing on Fragment 1...\")\n    print(\"=\"*60)\n    \n    # Load best model\n    checkpoint = torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth'))\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    # Load test data\n    print(\"Loading test data...\")\n    test_volume = load_volume_fast(fragment1_path)\n    test_mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ignore = generate_ignore_mask(test_mask)\n    \n    # Extract test patches\n    print(\"Extracting test patches...\")\n    test_patches, test_masks, test_ignores = extract_patches_fast(\n        test_volume, test_mask, test_ignore, max_patches_per_fragment=3000\n    )\n    print(f\"Test patches: {len(test_patches)}\")\n    \n    # Create test dataset\n    test_dataset = VesuviusDataset(\n        test_patches, test_masks, test_ignores, transform=None\n    )\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    # Evaluate\n    print(\"Running inference...\")\n    test_dice = 0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, ignores in tqdm(test_loader, desc=\"Testing\"):\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            ignores = ignores.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            valid_mask = (1 - ignores)\n            \n            # Calculate dice\n            intersection = ((preds * masks) * valid_mask).sum()\n            union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n            dice = (2 * intersection) / (union + 1e-6)\n            test_dice += dice.item()\n            \n            # Store for metrics\n            all_preds.append((preds * valid_mask).cpu().numpy())\n            all_targets.append((masks * valid_mask).cpu().numpy())\n    \n    avg_test_dice = test_dice / len(test_loader)\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_targets = np.concatenate([t.flatten() for t in all_targets])\n    \n    tn, fp, fn, tp = confusion_matrix(all_targets, all_preds, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    \n    # ============================================\n    # NEW: FULL VOLUME PREDICTION AND VISUALIZATION\n    # ============================================\n    print(\"\\n7. Generating full surface volume prediction...\")\n    print(\"=\"*60)\n    \n    # Generate full volume prediction\n    binary_prediction, prob_prediction = predict_full_volume(model, test_volume, DEVICE)\n    \n    # Save binary prediction\n    binary_save_path = os.path.join(OUTPUT_DIR, 'full_surface_binary_prediction.png')\n    cv2.imwrite(binary_save_path, binary_prediction * 255)\n    print(f\"Full surface binary prediction saved to: {binary_save_path}\")\n    \n    # Save probability prediction\n    prob_save_path = os.path.join(OUTPUT_DIR, 'full_surface_probability.tif')\n    tifffile.imwrite(prob_save_path, (prob_prediction * 255).astype(np.uint8))\n    print(f\"Full surface probability prediction saved to: {prob_save_path}\")\n    \n    # Create visualizations\n    print(\"\\n8. Creating test visualizations...\")\n    print(\"=\"*60)\n    \n    # Basic visualization\n    viz_path = os.path.join(OUTPUT_DIR, 'test_visualization_comparison.png')\n    create_test_visualization(test_volume, test_mask, binary_prediction, viz_path, slice_idx=6)\n    \n    # Detailed comparison visualization\n    comp_viz_path = os.path.join(OUTPUT_DIR, 'detailed_comparison_visualization.png')\n    create_comparison_visualization(test_mask, binary_prediction, comp_viz_path)\n    \n    # Create black and white overlay visualization\n    fig, axes = plt.subplots(2, 2, figsize=(16, 16))\n    \n    # Ground Truth\n    axes[0, 0].imshow(test_mask, cmap='binary')\n    axes[0, 0].set_title('Ground Truth (Black & White)', fontsize=14, fontweight='bold')\n    axes[0, 0].axis('off')\n    \n    # Prediction\n    axes[0, 1].imshow(binary_prediction, cmap='binary')\n    axes[0, 1].set_title('Full Surface Prediction (Black & White)', fontsize=14, fontweight='bold')\n    axes[0, 1].axis('off')\n    \n    # Overlay with transparency\n    overlay = np.zeros((*test_mask.shape, 3), dtype=np.uint8)\n    overlay[test_mask == 1] = [0, 255, 0]  # Green for ground truth\n    overlay[binary_prediction == 1] = [255, 0, 0]  # Red for prediction\n    \n    # Where both are present, show yellow\n    both = (test_mask == 1) & (binary_prediction == 1)\n    overlay[both] = [255, 255, 0]  # Yellow for overlap\n    \n    axes[1, 0].imshow(test_volume[:, :, 6], cmap='gray')\n    axes[1, 0].imshow(overlay, alpha=0.5)\n    axes[1, 0].set_title('Overlay on Input Volume\\nGreen: GT, Red: Pred, Yellow: Overlap', \n                         fontsize=14, fontweight='bold')\n    axes[1, 0].axis('off')\n    \n    # Legend\n    axes[1, 1].axis('off')\n    legend_text = \"Color Legend:\\n\\n\"\n    legend_text += \"⬤ White: Ink detected\\n\"\n    legend_text += \"⬤ Green: Ground Truth only\\n\"\n    legend_text += \"⬤ Red: Prediction only\\n\"\n    legend_text += \"⬤ Yellow: Overlap (Correct)\\n\\n\"\n    legend_text += f\"Metrics:\\n\"\n    legend_text += f\"Dice Score: {avg_test_dice:.4f}\\n\"\n    legend_text += f\"F1 Score: {f1:.4f}\\n\"\n    legend_text += f\"Precision: {precision:.4f}\\n\"\n    legend_text += f\"Recall: {recall:.4f}\"\n    axes[1, 1].text(0.1, 0.5, legend_text, fontsize=14, verticalalignment='center',\n                    bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n    \n    plt.tight_layout()\n    overlay_path = os.path.join(OUTPUT_DIR, 'overlay_visualization.png')\n    plt.savefig(overlay_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    print(f\"Overlay visualization saved to: {overlay_path}\")\n    \n    # ============================================\n    # FINAL RESULTS\n    # ============================================\n    total_time = time.time() - start_time\n    hours = int(total_time // 3600)\n    minutes = int((total_time % 3600) // 60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"FINAL TEST RESULTS ON FRAGMENT 1\")\n    print(\"=\"*60)\n    print(f\"Total Time: {hours}h {minutes}m\")\n    print(f\"Test Dice Score: {avg_test_dice:.4f}\")\n    print(f\"Precision: {precision:.4f}\")\n    print(f\"Recall: {recall:.4f}\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(\"-\"*60)\n    print(\"Confusion Matrix:\")\n    print(f\"  True Positives: {tp}\")\n    print(f\"  True Negatives: {tn}\")\n    print(f\"  False Positives: {fp}\")\n    print(f\"  False Negatives: {fn}\")\n    print(\"-\"*60)\n    print(\"Full Volume Prediction:\")\n    print(f\"  Shape: {binary_prediction.shape}\")\n    print(f\"  Predicted ink pixels: {binary_prediction.sum():,}\")\n    print(f\"  Ground truth ink pixels: {test_mask.sum():,}\")\n    print(\"=\"*60)\n    \n    # Save results\n    with open(os.path.join(OUTPUT_DIR, \"final_results.txt\"), \"w\") as f:\n        f.write(\"VESUVIUS INK DETECTION - FINAL RESULTS\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n        f.write(f\"Total Time: {hours}h {minutes}m\\n\")\n        f.write(f\"Best Validation Dice: {best_val_dice:.4f}\\n\")\n        f.write(f\"Test Dice Score: {avg_test_dice:.4f}\\n\")\n        f.write(f\"Precision: {precision:.4f}\\n\")\n        f.write(f\"Recall: {recall:.4f}\\n\")\n        f.write(f\"F1 Score: {f1:.4f}\\n\")\n        f.write(\"-\"*60 + \"\\n\")\n        f.write(f\"True Positives: {tp}\\n\")\n        f.write(f\"True Negatives: {tn}\\n\")\n        f.write(f\"False Positives: {fp}\\n\")\n        f.write(f\"False Negatives: {fn}\\n\")\n        f.write(\"-\"*60 + \"\\n\")\n        f.write(f\"Full Volume Prediction Shape: {binary_prediction.shape}\\n\")\n        f.write(f\"Predicted Ink Pixels: {binary_prediction.sum():,}\\n\")\n        f.write(f\"Ground Truth Ink Pixels: {test_mask.sum():,}\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n    \n    print(f\"\\nResults saved to: {os.path.join(OUTPUT_DIR, 'final_results.txt')}\")\n    \n    # Plot training curves\n    plt.figure(figsize=(12, 4))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses)\n    plt.title('Training Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.grid(True)\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(val_dice_scores)\n    plt.title('Validation Dice Score')\n    plt.xlabel('Epoch')\n    plt.ylabel('Dice')\n    plt.grid(True)\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'))\n    plt.show()\n    \n    return avg_test_dice\n\nif __name__ == \"__main__\":\n    try:\n        test_dice = main()\n        print(f\"\\n✅ Final Test Dice Score: {test_dice:.4f}\")\n    except Exception as e:\n        print(f\"\\n❌ Error: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Dice=46 , when PATCH_SIZE   = 224, STRIDE_TR    = 112,STRIDE_INF   = 56\n#BATCH_SIZE   = 4 ,GRAD_ACCUM   = 4    ,EPOCHS       = 25  ===============\n# Dice =40, when ,PATCH_SIZE   = 160 ,STRIDE_TR    = 96, STRIDE_INF   = 32\n#BATCH_SIZE   = 6 ,GRAD_ACCUM   = 6  # effective batch = 16,EPOCHS = 10 \n!pip install segmentation-models-pytorch==0.2.0\n# ============================================================\n#  VESUVIUS INK DETECTION — v10\n#  PhD Thesis: 3D Axial-Attention Transformer on CT Volume\n#\n#  Key Innovations vs v9:\n#  1.  3D Axial-Attention Transformer (X/Y/Z axes separately)\n#      with Token Merging on low-saliency regions\n#  2.  Sparse attention guided by coarse \"sheet-surface\" module\n#  3.  Log-polar relative positional embeddings (scroll geometry)\n#  4.  Heavy augmentation: elastic deform + CutMix ink/no-ink\n#  5.  Pseudo-labelling on unlabelled fragments\n#  6.  Post-processing: morphological cleaning + DenseCRF\n#  7.  NO TTA on test (removed — degrades perf on wrong preds)\n#  8.  ResNet34 2D backbone for final patch classification head\n#  9.  Deep supervision at 3 decoder scales\n#  10. Fragment-aware percentile normalisation (v9 style)\n# ============================================================\n\nimport os, gc, cv2, math, warnings, random\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\n# optional – DenseCRF post-processing\ntry:\n    import pydensecrf.densecrf as dcrf\n    from pydensecrf.utils import unary_from_labels, create_pairwise_bilateral, \\\n        create_pairwise_gaussian\n    HAS_CRF = True\nexcept ImportError:\n    HAS_CRF = False\n    print(\"[INFO] pydensecrf not found – CRF post-processing will be skipped. \"\n          \"Install with:  pip install pydensecrf\")\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\n\n# ──────────────────────────────────────────────────────────────\n#  SEED\n# ──────────────────────────────────────────────────────────────\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ──────────────────────────────────────────────────────────────\n#  PATHS\n# ──────────────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ──────────────────────────────────────────────────────────────\n#  HYPER-PARAMS\n# ──────────────────────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 224\nSTRIDE_TR    = 112\nSTRIDE_INF   = 56\nBATCH_SIZE   = 6           # reduced for 3D attention blocks\nGRAD_ACCUM   = 4           # effective batch = 16\nEPOCHS       = 20\nLR           = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 10\nNUM_WORKERS  = 0\nVAL_SPLIT    = 0.30\nDROPOUT_P    = 0.3\nTEMPERATURE  = 1.3         # softer than v9\n\n# Z-slices: 17 central slices (same as v9)\nZ_SLICES  = [18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34]\nN_CH      = len(Z_SLICES)   # 17\n\nINK_MIN_POS  = 0.02\nNEG_RATIO    = 0.5\nCH_DROP_PROB = 0.3\nCH_DROP_MAX  = 3\n\n# Pseudo-label config\nPSEUDO_THRESHOLD = 0.70     # confidence to accept pseudo label\nUSE_PSEUDO       = False    # set True if you have extra unlabelled fragments\nPSEUDO_FRAGS     = []       # e.g. ['/kaggle/input/vesuvius/extra/4']\n\nprint(f\"Device : {DEVICE}  |  Channels: {N_CH}  |  Patch: {PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1.  SLICE CACHE  (fragment-wise percentile normalisation)\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    if not raw:\n        raise ValueError(f\"No slices in {vol_dir}\")\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Slice cache: {len(cache)} slices ({mb:.0f} MB)\")\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2.  LOG-POLAR RELATIVE POSITIONAL EMBEDDING\n#      Respects cylindrical scroll geometry:\n#        r   = sqrt(dx² + dy²)  → log-scaled radial distance\n#        phi = atan2(dy, dx)    → angular displacement\n#        dz  = abs(z1 - z2)     → depth separation\n# ════════════════════════════════════════════════════════════\ndef build_logpolar_bias(seq_len_h, seq_len_w, num_heads, device='cpu'):\n    \"\"\"\n    Returns additive attention bias [num_heads, HW, HW] that encodes\n    log-polar relative positions between every pair of spatial tokens.\n    Computed once per unique (seq_len_h, seq_len_w) shape.\n    \"\"\"\n    H, W = seq_len_h, seq_len_w\n    # grid of (row, col) for every token\n    ys = torch.arange(H, dtype=torch.float32)\n    xs = torch.arange(W, dtype=torch.float32)\n    gy, gx = torch.meshgrid(ys, xs, indexing='ij')  # [H, W]\n    gy = gy.reshape(-1); gx = gx.reshape(-1)        # [HW]\n\n    dy = gy.unsqueeze(1) - gy.unsqueeze(0)           # [HW, HW]\n    dx = gx.unsqueeze(1) - gx.unsqueeze(0)\n\n    r   = torch.sqrt(dx**2 + dy**2 + 1e-3)           # radial dist\n    log_r = torch.log(r + 1.0)                       # log scale\n    phi = torch.atan2(dy, dx)                        # [-pi, pi]\n\n    # Encode into num_heads channels via learnable-free Fourier projection\n    freqs = torch.arange(1, num_heads // 2 + 1, dtype=torch.float32)\n    bias_r   = torch.cos(freqs[None, None, :] * log_r.unsqueeze(2))   # [HW,HW,H/2]\n    bias_phi = torch.sin(freqs[None, None, :] * phi.unsqueeze(2))     # [HW,HW,H/2]\n    bias = torch.cat([bias_r, bias_phi], dim=-1)    # [HW, HW, num_heads]\n    bias = bias.permute(2, 0, 1)                    # [num_heads, HW, HW]\n    return bias.to(device)\n\n\n# Cache to avoid recomputing for the same shape\n_LP_CACHE: dict = {}\n\ndef get_logpolar_bias(H, W, num_heads, device):\n    key = (H, W, num_heads, str(device))\n    if key not in _LP_CACHE:\n        _LP_CACHE[key] = build_logpolar_bias(H, W, num_heads, device)\n    return _LP_CACHE[key]\n\n\n# ════════════════════════════════════════════════════════════\n#  3.  COARSE SHEET-SURFACE DETECTION MODULE\n#      Lightweight 2-layer CNN that predicts a per-pixel\n#      \"ink-likelihood\" score from the raw 17-channel patch.\n#      Its output guides token merging: low-score tokens are\n#      merged (averaged), reducing sequence length for the\n#      expensive axial-attention layers.\n# ════════════════════════════════════════════════════════════\nclass SheetSurfaceDetector(nn.Module):\n    \"\"\"Fast coarse detector → saliency map [B, 1, H, W].\"\"\"\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, 32, 3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.ReLU(inplace=True),\n            nn.Conv2d(16,  1, 1),\n        )\n\n    def forward(self, x):   # x: [B, C, H, W]\n        return self.net(x)  # [B, 1, H, W]  (raw logit)\n\n\n# ════════════════════════════════════════════════════════════\n#  4.  TOKEN MERGING  (ToMe-inspired, sparse ink regions)\n#      Tokens whose coarse saliency < threshold are merged\n#      with their nearest neighbour → shorter sequence for\n#      the transformer.  We un-merge after attention so\n#      spatial resolution is fully restored for the decoder.\n# ════════════════════════════════════════════════════════════\ndef token_merge(tokens, saliency, merge_ratio=0.30):\n    \"\"\"\n    tokens   : [B, C, N]   (N = H*W tokens)\n    saliency : [B, 1, N]   (coarse ink score, already sigmoid)\n    Returns  : merged_tokens [B, C, M],  unmerge_idx [B, N] (maps M→N)\n    \"\"\"\n    B, C, N = tokens.shape\n    sal = saliency.squeeze(1)           # [B, N]\n\n    # sort tokens by saliency; low-saliency ones are merge candidates\n    _, sort_idx = sal.sort(dim=1)       # ascending → low first\n    n_merge = int(N * merge_ratio)\n\n    merge_idx  = sort_idx[:, :n_merge]   # [B, n_merge]\n    keep_idx   = sort_idx[:, n_merge:]   # [B, N-n_merge]\n\n    # gather merge & keep tokens\n    def gather(t, idx):\n        idx_exp = idx.unsqueeze(1).expand(-1, C, -1)\n        return t.gather(2, idx_exp)\n\n    t_merge = gather(tokens, merge_idx)   # [B, C, n_merge]\n    t_keep  = gather(tokens, keep_idx)    # [B, C, N-n_merge]\n\n    # pair-wise merge: consecutive pairs averaged\n    if n_merge % 2 == 1:\n        # drop last odd token into keep\n        t_keep  = torch.cat([t_keep, t_merge[:, :, -1:]], dim=2)\n        t_merge = t_merge[:, :, :-1]\n        n_merge -= 1\n\n    t_merged = (t_merge[:, :, 0::2] + t_merge[:, :, 1::2]) / 2  # [B,C,n_merge/2]\n    tokens_out = torch.cat([t_keep, t_merged], dim=2)            # [B, C, M]\n\n    # book-keeping for un-merge\n    unmerge_info = (keep_idx, merge_idx, n_merge, N)\n    return tokens_out, unmerge_info\n\n\ndef token_unmerge(tokens_out, unmerge_info, C):\n    \"\"\"Restore [B, C, N] from merged [B, C, M].\"\"\"\n    keep_idx, merge_idx, n_merge, N = unmerge_info\n    B = tokens_out.shape[0]\n    M = tokens_out.shape[2]\n    n_keep = M - n_merge // 2\n\n    t_keep   = tokens_out[:, :, :n_keep]               # [B,C,n_keep]\n    t_merged = tokens_out[:, :, n_keep:]               # [B,C,n_merge/2]\n\n    # expand merged pairs back\n    t_expanded = t_merged.repeat_interleave(2, dim=2)   # [B,C,n_merge]\n    if n_merge % 2 == 1:\n        # we added an extra token to keep above; undo\n        extra = t_keep[:, :, -1:]\n        t_keep = t_keep[:, :, :-1]\n        t_expanded = torch.cat([t_expanded, extra], dim=2)\n        n_merge += 1\n\n    # scatter back to original positions\n    device = tokens_out.device\n    out = torch.zeros(B, C, N, device=device, dtype=tokens_out.dtype)\n\n    def scatter(t, idx):\n        idx_exp = idx.unsqueeze(1).expand(-1, C, -1)\n        out.scatter_(2, idx_exp, t)\n\n    scatter(t_keep,    keep_idx)\n    scatter(t_expanded, merge_idx)\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  5.  AXIAL MULTI-HEAD SELF-ATTENTION  (X / Y axes)\n#      Each axis attends along one spatial dimension with\n#      log-polar positional bias.  We alternate X→Y.\n# ════════════════════════════════════════════════════════════\nclass AxialAttention(nn.Module):\n    \"\"\"\n    Attention along one axis (row or column) of a 2-D feature map.\n    x : [B, C, H, W]\n    \"\"\"\n    def __init__(self, dim, num_heads=8, axis='x', dropout=0.1):\n        super().__init__()\n        assert axis in ('x', 'y')\n        self.axis      = axis\n        self.num_heads = num_heads\n        self.head_dim  = dim // num_heads\n        self.scale     = self.head_dim ** -0.5\n\n        self.qkv  = nn.Linear(dim, dim * 3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.drop = nn.Dropout(dropout)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        if self.axis == 'x':\n            # attend along width for each row independently\n            x_r = x.permute(0, 2, 3, 1)          # [B, H, W, C]\n            shape = (B * H, W, C)\n        else:\n            x_r = x.permute(0, 3, 2, 1)          # [B, W, H, C]\n            shape = (B * W, H, C)\n\n        x_r = x_r.reshape(*shape)                # [B*H, W, C] or [B*W, H, C]\n        res  = x_r\n        x_n  = self.norm(x_r)\n\n        BN, L, _ = x_n.shape\n        qkv = self.qkv(x_n).reshape(BN, L, 3, self.num_heads, self.head_dim)\n        qkv = qkv.permute(2, 0, 3, 1, 4)        # [3, BN, nh, L, hd]\n        q, k, v = qkv.unbind(0)\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale   # [BN, nh, L, L]\n\n        # Add log-polar bias (1-D for axial: just use log distance)\n        pos = torch.arange(L, dtype=torch.float32, device=x.device)\n        d   = (pos.unsqueeze(0) - pos.unsqueeze(1)).abs().float()   # [L, L]\n        log_d_bias = -torch.log(d + 1.0)                            # [L, L]\n        attn = attn + log_d_bias.unsqueeze(0).unsqueeze(0)\n\n        attn = attn.softmax(-1)\n        attn = self.drop(attn)\n        out  = (attn @ v).transpose(1, 2).reshape(BN, L, C)\n        out  = self.proj(out) + res                  # residual\n\n        if self.axis == 'x':\n            out = out.reshape(B, H, W, C).permute(0, 3, 1, 2)\n        else:\n            out = out.reshape(B, W, H, C).permute(0, 3, 2, 1)\n        return out\n\n\n# ════════════════════════════════════════════════════════════\n#  6.  AXIAL TRANSFORMER BLOCK  (X-axis → Y-axis → FFN)\n# ════════════════════════════════════════════════════════════\nclass AxialTransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=8, mlp_ratio=4, dropout=0.1):\n        super().__init__()\n        self.attn_x = AxialAttention(dim, num_heads, axis='x', dropout=dropout)\n        self.attn_y = AxialAttention(dim, num_heads, axis='y', dropout=dropout)\n        self.norm1  = nn.LayerNorm(dim)\n        self.norm2  = nn.LayerNorm(dim)\n        mlp_dim     = int(dim * mlp_ratio)\n        self.ffn = nn.Sequential(\n            nn.Linear(dim, mlp_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(mlp_dim, dim),\n            nn.Dropout(dropout),\n        )\n\n    def forward(self, x):\n        # x: [B, C, H, W]\n        x = self.attn_x(x)\n        x = self.attn_y(x)\n        # FFN on channel dim\n        B, C, H, W = x.shape\n        xf = x.permute(0, 2, 3, 1).reshape(-1, C)\n        xf = self.ffn(self.norm2(xf))\n        x  = x + xf.reshape(B, H, W, C).permute(0, 3, 1, 2)\n        return x\n\n\n# ════════════════════════════════════════════════════════════\n#  7.  FULL MODEL\n#      Architecture:\n#        ResNet34 CNN encoder  (17-ch in)\n#             ↓  skip features at 4 scales\n#        AxialTransformerBlocks on bottleneck  (with token merging)\n#             ↓\n#        FPN-style decoder  (deep supervision at 3 scales)\n#             ↓\n#        SheetSurfaceDetector auxiliary head\n# ════════════════════════════════════════════════════════════\nclass VesuviusV10(nn.Module):\n    def __init__(self, n_ch=N_CH, enc_dim=512,\n                 n_transformer_blocks=4, num_heads=8):\n        super().__init__()\n\n        # ── ResNet34 UNet backbone ────────────────────────────\n        self.backbone = smp.Unet(\n            encoder_name          = 'resnet34',\n            encoder_weights       = 'imagenet',\n            in_channels           = n_ch,\n            classes               = 1,\n            decoder_attention_type= 'scse',\n        )\n\n        # We intercept the encoder bottleneck output and replace\n        # the UNet decoder with our own axial-transformer decoder.\n        # Bottleneck of ResNet34 is 512 channels at 1/32 spatial.\n        self.transformer_blocks = nn.Sequential(*[\n            AxialTransformerBlock(enc_dim, num_heads=num_heads,\n                                  mlp_ratio=4, dropout=0.1)\n            for _ in range(n_transformer_blocks)\n        ])\n\n        # ── Sheet-surface detector (auxiliary head on raw input) ──\n        self.sheet_detector = SheetSurfaceDetector(n_ch)\n\n        # ── Decoder (same as backbone's but we replace the head) ──\n        # We will use the backbone's decoder directly; the transformer\n        # operates on the bottleneck feature before decoder sees it.\n        self._enc_dim = enc_dim\n\n        # Projection to match residual after transformer\n        self.bottleneck_proj = nn.Sequential(\n            nn.Conv2d(enc_dim, enc_dim, 1, bias=False),\n            nn.BatchNorm2d(enc_dim),\n            nn.GELU(),\n        )\n\n        # Deep supervision heads\n        # These attach to intermediate decoder stages\n        # ResNet34 decoder feature sizes: 256, 128, 64, 32, 16\n        self.ds_head3 = nn.Conv2d(256, 1, 1)   # after decode block 0\n        self.ds_head2 = nn.Conv2d(128, 1, 1)   # after decode block 1\n        self.ds_head1 = nn.Conv2d( 64, 1, 1)   # after decode block 2\n\n    def forward(self, x):\n        B = x.shape[0]\n\n        # ── Coarse sheet-surface saliency ──────────────────────\n        saliency_logit = self.sheet_detector(x)   # [B,1,H,W]\n\n        # ── Encoder ────────────────────────────────────────────\n        # Use backbone's encoder forward pass\n        feats = self.backbone.encoder(x)\n        # feats is a list: [input, s1, s2, s3, s4, bottleneck]\n        bottleneck = feats[-1]   # [B, 512, H/32, W/32]\n\n        # ── Axial-Transformer on bottleneck ────────────────────\n        bH, bW = bottleneck.shape[2], bottleneck.shape[3]\n        N      = bH * bW\n\n        # token merging guided by down-sampled saliency\n        sal_down = F.adaptive_avg_pool2d(\n            torch.sigmoid(saliency_logit), (bH, bW)\n        )                                         # [B,1,bH,bW]\n        sal_flat = sal_down.reshape(B, 1, N)      # [B,1,N]\n        tok      = bottleneck.reshape(B, self._enc_dim, N)  # [B,C,N]\n\n        tok_merged, unmerge_info = token_merge(tok, sal_flat, merge_ratio=0.25)\n\n        # reshape merged tokens to pseudo-spatial for axial attention\n        # approximate square layout\n        M       = tok_merged.shape[2]\n        sq      = int(math.ceil(math.sqrt(M)))\n        pad_len = sq * sq - M\n        if pad_len > 0:\n            tok_merged = F.pad(tok_merged, (0, pad_len))\n        tok_2d = tok_merged.reshape(B, self._enc_dim, sq, sq)  # [B,C,sq,sq]\n\n        tok_2d = self.transformer_blocks(tok_2d)               # axial attention\n        tok_flat = tok_2d.reshape(B, self._enc_dim, sq * sq)[:, :, :M]\n\n        # un-merge back to full spatial\n        tok_full = token_unmerge(tok_flat, unmerge_info, self._enc_dim)\n        bottleneck_out = tok_full.reshape(B, self._enc_dim, bH, bW)\n        bottleneck_out = self.bottleneck_proj(bottleneck_out + bottleneck)\n\n        # replace bottleneck in feats\n        feats_mod = list(feats)\n        feats_mod[-1] = bottleneck_out\n\n        # ── Decoder ────────────────────────────────────────────\n        # Use backbone decoder – it accepts the feature list\n        decoder_output = self.backbone.decoder(*feats_mod)\n        main_logit     = self.backbone.segmentation_head(decoder_output)\n\n        # Deep supervision: tap intermediate decoder layers\n        # smp UNet decoder blocks are in self.backbone.decoder.blocks\n        ds3_logit = None; ds2_logit = None; ds1_logit = None\n        try:\n            db = self.backbone.decoder.blocks\n            if len(db) >= 1:\n                f0 = db[0](feats_mod[-1], feats_mod[-2])\n                ds3_logit = self.ds_head3(f0)\n            if len(db) >= 2:\n                f1 = db[1](f0, feats_mod[-3])\n                ds2_logit = self.ds_head2(f1)\n            if len(db) >= 3:\n                f2 = db[2](f1, feats_mod[-4])\n                ds1_logit = self.ds_head1(f2)\n        except Exception:\n            pass   # deep supervision optional\n\n        return main_logit, saliency_logit, ds3_logit, ds2_logit, ds1_logit\n\n\n# ════════════════════════════════════════════════════════════\n#  8.  AUGMENTATIONS  (heavy – elastic + CutMix)\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.RandomRotate90(p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.12, scale_limit=0.18,\n                       rotate_limit=35,\n                       border_mode=cv2.BORDER_REFLECT, p=0.65),\n    A.ElasticTransform(alpha=1.0, sigma=50, alpha_affine=50,\n                       border_mode=cv2.BORDER_REFLECT, p=0.4),\n    A.GridDistortion(num_steps=5, distort_limit=0.3,\n                     border_mode=cv2.BORDER_REFLECT, p=0.3),\n    A.RandomResizedCrop(height=PATCH_SIZE, width=PATCH_SIZE,\n                        scale=(0.55, 1.0), ratio=(0.85, 1.15), p=0.5),\n    A.RandomBrightnessContrast(0.25, 0.25, p=0.55),\n    A.GaussNoise(var_limit=(0.001, 0.005), p=0.35),\n    A.GaussianBlur(blur_limit=(3, 7), p=0.25),\n    A.CoarseDropout(max_holes=6, max_height=28, max_width=28,\n                    fill_value=0, p=0.35),\n])\n\n\ndef channel_dropout(img_np, prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < prob:\n        n   = np.random.randint(1, max_drop + 1)\n        idx = np.random.choice(img_np.shape[2], n, replace=False)\n        img_np = img_np.copy()\n        img_np[:, :, idx] = 0.0\n    return img_np\n\n\ndef cutmix_batch(imgs, msks, alpha=0.4):\n    \"\"\"\n    CutMix across the batch dimension for ink/no-ink regions.\n    imgs : [B, C, H, W]  msks : [B, 1, H, W]\n    Returns augmented imgs and mixed msks (soft labels ok for loss).\n    \"\"\"\n    B, C, H, W = imgs.shape\n    lam  = np.random.beta(alpha, alpha)\n    perm = torch.randperm(B, device=imgs.device)\n\n    cx  = np.random.randint(W)\n    cy  = np.random.randint(H)\n    bw  = int(W * math.sqrt(1 - lam))\n    bh  = int(H * math.sqrt(1 - lam))\n    x1  = max(0, cx - bw // 2); x2 = min(W, cx + bw // 2)\n    y1  = max(0, cy - bh // 2); y2 = min(H, cy + bh // 2)\n\n    imgs_new       = imgs.clone()\n    msks_new       = msks.clone()\n    imgs_new[:, :, y1:y2, x1:x2] = imgs[perm, :, y1:y2, x1:x2]\n    msks_new[:, :, y1:y2, x1:x2] = msks[perm, :, y1:y2, x1:x2]\n    return imgs_new, msks_new\n\n\n# ════════════════════════════════════════════════════════════\n#  9.  DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0, apply_ch_dropout=False):\n        self.cache          = load_slice_cache(frag_path, z_list)\n        self.z_list         = z_list\n        self.tf             = transform\n        self.ch_dropout     = apply_ch_dropout\n        self.frag_path      = frag_path\n\n        msk_path = os.path.join(frag_path, 'inklabels.png')\n        msk = cv2.imread(msk_path, 0)\n        if msk is None:\n            raise FileNotFoundError(f\"No inklabels at {msk_path}\")\n        self.mask = (msk > 0).astype(np.uint8)\n\n        ir_path = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = (cv2.imread(ir_path, 0) > 0).astype(np.uint8) \\\n                       if os.path.exists(ir_path) else None\n\n        H, W = self.mask.shape\n        pos_coords, neg_coords = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.sum() / (PATCH_SIZE * PATCH_SIZE)\n                if self.ir_mask is not None:\n                    on_pap = self.ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    mid_z  = z_list[len(z_list)//2]\n                    on_pap = float(self.cache[mid_z][y:y+PATCH_SIZE,\n                                                     x:x+PATCH_SIZE].mean()) > 0.1\n                if not on_pap:\n                    continue\n                if ink >= INK_MIN_POS:\n                    pos_coords.append((y, x, 1))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_coords.append((y, x, 0))\n\n        n_neg   = int(len(pos_coords) * neg_ratio)\n        sel_neg = []\n        if n_neg > 0 and neg_coords:\n            np.random.shuffle(neg_coords)\n            sel_neg = neg_coords[:n_neg]\n\n        self.coords  = pos_coords + sel_neg\n        self.weights = np.array(\n            [3.0 if c[2] == 1 else 1.0 for c in self.coords], dtype=np.float32\n        )\n        print(f\"  [{os.path.basename(frag_path)}] \"\n              f\"{len(pos_coords)} pos + {len(sel_neg)} neg = {len(self.coords)}\")\n\n    def __len__(self): return len(self.coords)\n\n    def get_patch(self, idx):\n        y, x, _ = self.coords[idx]\n        slices   = [self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                    for z in self.z_list]\n        return np.stack(slices, axis=-1), self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n\n    def __getitem__(self, idx):\n        img, msk = self.get_patch(idx)\n        if self.tf:\n            out = self.tf(image=img, mask=msk)\n            img, msk = out['image'], out['mask']\n        if self.ch_dropout:\n            img = channel_dropout(img)\n        return (torch.from_numpy(img).permute(2, 0, 1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\n\n# ════════════════════════════════════════════════════════════\n#  10.  LOSS  (focal + dice + deep-supervision)\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = target * alpha + (1 - target) * (1 - alpha)\n    return (a_t * ((1 - p_t) ** gamma) * bce).mean()\n\n\ndef dice_loss(pred, target, smooth=1.):\n    p     = torch.sigmoid(pred)\n    inter = (p * target).sum(dim=(2, 3))\n    union = p.sum(dim=(2, 3)) + target.sum(dim=(2, 3))\n    return 1. - ((2. * inter + smooth) / (union + smooth)).mean()\n\n\ndef combined_loss(pred, target, eps=0.05):\n    t_s = target * (1 - eps) + 0.5 * eps\n    return 0.5 * focal_loss(pred, t_s) + 0.5 * dice_loss(pred, target)\n\n\ndef multiscale_loss(main_logit, sal_logit, ds3, ds2, ds1, target):\n    \"\"\"\n    main_logit, sal_logit : [B,1,H,W]\n    ds3,ds2,ds1           : lower-res deep-sup logits (may be None)\n    target                : [B,1,H,W]\n    \"\"\"\n    loss = combined_loss(main_logit, target)\n\n    # auxiliary sheet-surface detection loss (same target, down-sampled)\n    sal_t = F.adaptive_avg_pool2d(target, sal_logit.shape[-2:])\n    loss += 0.2 * combined_loss(sal_logit, sal_t)\n\n    def ds_loss(logit, wt):\n        if logit is None: return 0.0\n        tgt = F.adaptive_avg_pool2d(target, logit.shape[-2:])\n        return wt * combined_loss(logit, tgt)\n\n    loss += ds_loss(ds3, 0.15)\n    loss += ds_loss(ds2, 0.10)\n    loss += ds_loss(ds1, 0.05)\n    return loss\n\n\n# ════════════════════════════════════════════════════════════\n#  11.  BUILD DATASETS — 70/30 patch-level split on Frag2+Frag3\n# ════════════════════════════════════════════════════════════\nprint('\\n── Loading Fragment 2 (train pool) ──')\nds2 = VesuviusDataset(FRAG2, Z_SLICES, stride=STRIDE_TR,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\nprint('\\n── Loading Fragment 3 (train pool) ──')\nds3 = VesuviusDataset(FRAG3, Z_SLICES, stride=STRIDE_TR,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\n\nn_total   = len(ds2) + len(ds3)\nrng       = np.random.RandomState(SEED)\nall_idx   = rng.permutation(n_total)\nn_val     = int(n_total * VAL_SPLIT)\nn_train   = n_total - n_val\ntrain_idx = all_idx[:n_train].tolist()\nval_idx   = all_idx[n_train:].tolist()\nprint(f'\\nTotal: {n_total} | Train: {n_train} (70%) | Val: {n_val} (30%)')\n\n\nclass ValSubset(Dataset):\n    def __init__(self, ds2, ds3, indices):\n        self.ds2 = ds2; self.ds3 = ds3\n        self.n2  = len(ds2); self.indices = indices\n\n    def __len__(self): return len(self.indices)\n\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        if g < self.n2:\n            img_np, msk_np = self.ds2.get_patch(g)\n        else:\n            img_np, msk_np = self.ds3.get_patch(g - self.n2)\n        return (torch.from_numpy(img_np).permute(2, 0, 1).float(),\n                torch.from_numpy(msk_np).unsqueeze(0).float())\n\n\nclass TrainSubset(Dataset):\n    def __init__(self, ds2, ds3, indices):\n        self.ds2 = ds2; self.ds3 = ds3\n        self.n2  = len(ds2); self.indices = indices\n\n    def __len__(self): return len(self.indices)\n\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        return self.ds2[g] if g < self.n2 else self.ds3[g - self.n2]\n\n\ntrain_ds = TrainSubset(ds2, ds3, train_idx)\nval_ds   = ValSubset  (ds2, ds3, val_idx)\n\nall_weights = np.concatenate([ds2.weights, ds3.weights])\ntrain_w     = torch.from_numpy(all_weights[train_idx]).float()\n\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\nsampler  = WeightedRandomSampler(train_w, len(train_ds), replacement=True)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                      num_workers=NUM_WORKERS, pin_memory=False)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=NUM_WORKERS, pin_memory=False)\nprint(f'Train batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  12.  METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p * masks).sum(dim=(1, 2, 3))\n    union = p.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3))\n    return ((2. * inter + 1e-5) / (union + 1e-5)).mean().item()\n\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.25, 0.85, 0.02):\n        p = (probs > t).astype(np.float32)\n        d = (2 * (p * targets).sum() + 1) / (p.sum() + targets.sum() + 1)\n        if d > best_d:\n            best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  13.  MODEL + OPTIM + SCHEDULER\n# ════════════════════════════════════════════════════════════\nmodel     = VesuviusV10(n_ch=N_CH, enc_dim=512,\n                        n_transformer_blocks=4, num_heads=8).to(DEVICE)\nn_params  = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'\\nModel parameters: {n_params/1e6:.2f} M')\n\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='max', factor=0.5, patience=5, min_lr=5e-6, verbose=True\n)\nscaler    = GradScaler(enabled=(DEVICE == 'cuda'))\nbest_dice = 0.; pat_cnt = 0\nhistory   = dict(tl=[], vl=[], td=[], vd=[], lr=[])\n\nprint('\\n' + '=' * 65)\nprint('TRAINING — VesuviusV10')\nprint('  Encoder      : ResNet34 + Axial-Transformer bottleneck')\nprint('  Attention    : Axial X/Y with log-polar positional bias')\nprint('  Token merging: 25% low-saliency tokens merged')\nprint('  Augmentation : Elastic + GridDistort + CutMix + ChannelDrop')\nprint('  Deep supervis: 3 auxiliary heads')\nprint('  NO TTA at test (removed for reliability)')\nprint(f'  Z-slices     : {Z_SLICES}')\nprint('=' * 65)\n\n\n# ════════════════════════════════════════════════════════════\n#  14.  TRAINING LOOP\n# ════════════════════════════════════════════════════════════\nfor epoch in range(EPOCHS):\n\n    # ── train ──────────────────────────────────────────────\n    model.train(); tl = td = 0.\n    optimizer.zero_grad()\n\n    for step, (imgs, msks) in enumerate(\n            tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n\n        imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n\n        # CutMix (30% of steps)\n        if random.random() < 0.30:\n            imgs, msks = cutmix_batch(imgs, msks, alpha=0.4)\n\n        with autocast(enabled=(DEVICE == 'cuda')):\n            main_logit, sal_logit, ds3_l, ds2_l, ds1_l = model(imgs)\n            loss = multiscale_loss(main_logit, sal_logit,\n                                   ds3_l, ds2_l, ds1_l, msks) / GRAD_ACCUM\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % GRAD_ACCUM == 0 or (step + 1) == len(train_dl):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer); scaler.update()\n            optimizer.zero_grad()\n\n        tl += loss.item() * GRAD_ACCUM\n        td += batch_dice(main_logit.detach(), msks)\n\n        del imgs, msks, main_logit, sal_logit, ds3_l, ds2_l, ds1_l, loss\n        if step % 50 == 0:\n            torch.cuda.empty_cache()\n\n    tl /= len(train_dl); td /= len(train_dl)\n\n    # ── validate ────────────────────────────────────────────\n    model.eval(); vl = vd = 0.\n    acc_p, acc_m = [], []\n\n    with torch.no_grad():\n        for imgs, msks in tqdm(val_dl, desc=f'Ep{epoch+1:02d} val  ', leave=False):\n            imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                main_logit, sal_logit, ds3_l, ds2_l, ds1_l = model(imgs)\n                loss = multiscale_loss(main_logit, sal_logit,\n                                       ds3_l, ds2_l, ds1_l, msks)\n            vl += loss.item()\n            vd += batch_dice(main_logit, msks)\n            acc_p.append(torch.sigmoid(main_logit / TEMPERATURE).cpu().numpy())\n            acc_m.append(msks.cpu().numpy())\n            del imgs, msks, main_logit, sal_logit, ds3_l, ds2_l, ds1_l, loss\n\n    vl /= len(val_dl); vd /= len(val_dl)\n    probs_all = np.concatenate(acc_p)\n    masks_all = np.concatenate(acc_m)\n    bt, bd    = sweep_threshold(probs_all, masks_all)\n\n    ink_mean   = float(probs_all[masks_all > 0.5].mean()) if (masks_all > 0.5).any() else 0.\n    noink_mean = float(probs_all[masks_all < 0.5].mean()) if (masks_all < 0.5).any() else 0.\n    sep        = ink_mean - noink_mean\n    del acc_p, acc_m, probs_all, masks_all\n    gc.collect(); torch.cuda.empty_cache()\n\n    scheduler.step(vd)\n    lr_now = optimizer.param_groups[0]['lr']\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['lr'].append(lr_now)\n\n    print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n          f'train loss={tl:.4f} dice={td:.4f} | '\n          f'val loss={vl:.4f} dice={vd:.4f} | '\n          f'thr={bt:.2f}→{bd:.4f} | '\n          f'sep={sep:+.3f} [ink={ink_mean:.3f} bg={noink_mean:.3f}]')\n\n    thr_penalty  = max(0.0, bt - 0.60) * 0.3\n    save_metric  = max(vd, bd) - thr_penalty\n\n    if save_metric > best_dice:\n        best_dice = save_metric; pat_cnt = 0\n        torch.save({'epoch': epoch, 'state': model.state_dict(),\n                    'thr': bt, 'metric': save_metric, 'n_ch': N_CH,\n                    'bd': bd, 'vd': vd},\n                   OUTPUT + 'best_model.pth')\n        print(f'  ✓ saved  (penalised={save_metric:.4f}  '\n              f'raw_dice={max(vd,bd):.4f}  thr={bt:.2f})')\n    else:\n        pat_cnt += 1\n        if pat_cnt >= PATIENCE:\n            print(f'  ⚑ early stop at epoch {epoch+1}')\n            break\n\nprint(f'\\nBest penalised metric: {best_dice:.4f}')\n\n# ── curves ──────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 4))\naxes[0].plot(history['tl'], label='train')\naxes[0].plot(history['vl'], label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'], label='train')\naxes[1].plot(history['vd'], label='val')\naxes[1].axhline(0.80, color='r', ls='--', label='target 0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['lr'])\naxes[2].set_title('LR'); axes[2].set_xlabel('Epoch'); axes[2].grid(True)\nplt.tight_layout()\nplt.savefig(OUTPUT + 'curves.png', dpi=100); plt.close()\nprint('Curves saved.')\n\n\n# ════════════════════════════════════════════════════════════\n#  15.  PSEUDO-LABELLING  (optional)\n#       If USE_PSEUDO is True and PSEUDO_FRAGS is not empty,\n#       we run inference on unlabelled fragments, keep\n#       high-confidence predictions as pseudo labels,\n#       then fine-tune for a few epochs.\n# ════════════════════════════════════════════════════════════\ndef generate_pseudo_labels(model, frag_path, z_list, threshold=PSEUDO_THRESHOLD):\n    \"\"\"Run sliding-window inference → save pseudo inklabels.png.\"\"\"\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    # build a small 1-slice proxy for spatial size\n    mid_z = z_list[len(z_list) // 2]\n    H, W  = cache[mid_z].shape\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n\n    coords = [(y, x) for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n                      for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='PseudoLabel', leave=False):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            patch  = np.stack(slices, axis=-1)\n            t      = torch.from_numpy(patch).permute(2, 0, 1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit / TEMPERATURE).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += 1.0\n            del t, logit\n\n    del cache; gc.collect(); torch.cuda.empty_cache()\n    prob_map = pred_map / (wgt_map + 1e-8)\n    # Keep only confident predictions\n    pseudo_lbl = ((prob_map > threshold) * 255).astype(np.uint8)\n    out_path   = os.path.join(frag_path, 'inklabels_pseudo.png')\n    cv2.imwrite(out_path, pseudo_lbl)\n    print(f'  Pseudo label saved to {out_path}  '\n          f'(ink%={100*(pseudo_lbl>0).mean():.1f}%)')\n    return out_path\n\n\nif USE_PSEUDO and PSEUDO_FRAGS:\n    print('\\n── Pseudo-labelling extra fragments ──')\n    ckpt = torch.load(OUTPUT + 'best_model.pth', map_location=DEVICE)\n    model.load_state_dict(ckpt['state'])\n\n    pseudo_datasets = []\n    for pf in PSEUDO_FRAGS:\n        lbl_path = generate_pseudo_labels(model, pf, Z_SLICES)\n        # temporarily rename pseudo label so the dataset can load it\n        ink_orig = os.path.join(pf, 'inklabels.png')\n        ink_back = os.path.join(pf, 'inklabels_real.png')\n        if not os.path.exists(ink_back) and os.path.exists(ink_orig):\n            os.rename(ink_orig, ink_back)\n        os.rename(lbl_path, ink_orig)\n        try:\n            pds = VesuviusDataset(pf, Z_SLICES, stride=STRIDE_TR,\n                                  transform=train_tf, neg_ratio=0.3,\n                                  apply_ch_dropout=True)\n            pseudo_datasets.append(pds)\n        except Exception as e:\n            print(f'  [WARNING] Could not build pseudo dataset for {pf}: {e}')\n        finally:\n            # restore original label\n            if os.path.exists(ink_back):\n                os.replace(ink_back, ink_orig)\n\n    if pseudo_datasets:\n        print(f'\\n── Fine-tuning with {sum(len(p) for p in pseudo_datasets)} pseudo patches ──')\n        from torch.utils.data import ConcatDataset\n        pseudo_combined = ConcatDataset(pseudo_datasets)\n        pseudo_dl = DataLoader(pseudo_combined, batch_size=BATCH_SIZE,\n                               shuffle=True, num_workers=NUM_WORKERS)\n        pseudo_optimizer = optim.AdamW(model.parameters(),\n                                       lr=LR * 0.2, weight_decay=WEIGHT_DECAY)\n        pseudo_scaler    = GradScaler(enabled=(DEVICE == 'cuda'))\n\n        for ep in range(5):   # short fine-tune\n            model.train(); pl = pd_val = 0.\n            pseudo_optimizer.zero_grad()\n            for step, (imgs, msks) in enumerate(\n                    tqdm(pseudo_dl, desc=f'Pseudo Ep{ep+1}', leave=False)):\n                imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n                with autocast(enabled=(DEVICE == 'cuda')):\n                    main_logit, sal_logit, d3, d2, d1 = model(imgs)\n                    loss = multiscale_loss(main_logit, sal_logit,\n                                          d3, d2, d1, msks) / GRAD_ACCUM\n                pseudo_scaler.scale(loss).backward()\n                if (step+1) % GRAD_ACCUM == 0:\n                    pseudo_scaler.unscale_(pseudo_optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                    pseudo_scaler.step(pseudo_optimizer)\n                    pseudo_scaler.update()\n                    pseudo_optimizer.zero_grad()\n                pl += loss.item() * GRAD_ACCUM\n                pd_val += batch_dice(main_logit.detach(), msks)\n                del imgs, msks, main_logit, sal_logit, d3, d2, d1, loss\n            pl /= len(pseudo_dl); pd_val /= len(pseudo_dl)\n            print(f'  Pseudo Ep{ep+1} | loss={pl:.4f} | dice={pd_val:.4f}')\n        torch.save({'state': model.state_dict(), 'n_ch': N_CH,\n                    'thr': ckpt['thr']},\n                   OUTPUT + 'best_model_pseudo.pth')\n        print('Pseudo fine-tune complete.')\n\n\n# ════════════════════════════════════════════════════════════\n#  16.  INFERENCE  (NO TTA — clean deterministic single-pass)\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c   = sz // 2; sig = sz // 4\n    ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs - c)**2 + (ys - c)**2) / (2 * sig**2)).astype(np.float32)\n\n\nGW = gauss_weight(PATCH_SIZE)\n\n\ndef predict_fragment(model, frag_path, z_list, temperature=TEMPERATURE):\n    \"\"\"\n    Sliding-window inference WITHOUT TTA.\n    Returns (prob_map [H,W], label_mask [H,W]).\n    \"\"\"\n    model.eval()\n    cache   = load_slice_cache(frag_path, z_list)\n    lbl_img = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n    msk     = (lbl_img > 0).astype(np.uint8)\n    H, W    = msk.shape\n\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n\n    coords = [(y, x) for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n                      for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='Inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            patch  = np.stack(slices, axis=-1)\n            t      = torch.from_numpy(patch).permute(2, 0, 1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit / temperature).squeeze().cpu().numpy().astype(np.float32)\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p  * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, logit\n            torch.cuda.empty_cache()\n\n    del cache; gc.collect()\n    return pred_map / (wgt_map + 1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  17.  POST-PROCESSING\n#       (a) Morphological cleaning  — removes small FP islands\n#       (b) DenseCRF                — refines boundaries\n# ════════════════════════════════════════════════════════════\ndef morphological_clean(binary_map, min_area=200, close_k=5):\n    \"\"\"\n    1. Close small holes with a small kernel\n    2. Remove connected components smaller than min_area pixels\n    \"\"\"\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (close_k, close_k))\n    closed = cv2.morphologyEx(binary_map.astype(np.uint8),\n                              cv2.MORPH_CLOSE, kernel)\n    # remove small components\n    n_lab, labels, stats, _ = cv2.connectedComponentsWithStats(closed)\n    cleaned = np.zeros_like(closed)\n    for i in range(1, n_lab):\n        if stats[i, cv2.CC_STAT_AREA] >= min_area:\n            cleaned[labels == i] = 1\n    return cleaned\n\n\ndef apply_dense_crf(image_uint8, prob_map, n_iter=5):\n    \"\"\"\n    image_uint8 : [H, W, 3] uint8 (mid z-slice repeated to RGB)\n    prob_map    : [H, W] float in [0,1]\n    Returns     : refined binary prediction [H, W]\n    \"\"\"\n    if not HAS_CRF:\n        print(\"[INFO] DenseCRF unavailable – skipping.\")\n        return (prob_map > 0.5).astype(np.uint8)\n\n    H, W = prob_map.shape\n    d    = dcrf.DenseCRF2D(W, H, 2)\n\n    # unary potentials from probability map\n    fg  = np.clip(prob_map,       1e-5, 1 - 1e-5)\n    bg  = np.clip(1.0 - prob_map, 1e-5, 1 - 1e-5)\n    U   = -np.log(np.stack([bg, fg], axis=0))    # [2, H*W]\n    d.setUnaryEnergy(U.reshape(2, -1).astype(np.float32))\n\n    # pairwise: Gaussian spatial (smoothness)\n    d.addPairwiseGaussian(sxy=3, compat=3)\n\n    # pairwise: bilateral (edge-aware)\n    img_c = np.ascontiguousarray(image_uint8)\n    d.addPairwiseBilateral(sxy=50, srgb=13, rgbim=img_c, compat=10)\n\n    Q = d.inference(n_iter)\n    return np.argmax(Q, axis=0).reshape(H, W).astype(np.uint8)\n\n\ndef postprocess(prob_map, raw_vol_cache, z_list,\n                threshold, min_area=200, close_k=5, crf_iters=5):\n    \"\"\"Full post-processing pipeline.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n\n    # Morphological cleaning\n    binary = morphological_clean(binary, min_area=min_area, close_k=close_k)\n\n    # DenseCRF using mid z-slice as colour guidance\n    if HAS_CRF:\n        mid_z  = z_list[len(z_list) // 2]\n        mid_sl = raw_vol_cache[mid_z].astype(np.float32)\n        mid_8  = ((mid_sl - mid_sl.min()) /\n                  (mid_sl.max() - mid_sl.min() + 1e-8) * 255).astype(np.uint8)\n        rgb    = np.stack([mid_8, mid_8, mid_8], axis=-1)\n        binary = apply_dense_crf(rgb, prob_map, n_iter=crf_iters)\n\n    return binary\n\n\n# ════════════════════════════════════════════════════════════\n#  18.  FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '=' * 65)\nprint('FINAL TEST — FRAGMENT 1  (never seen during training)')\nprint('NO TTA — single deterministic inference pass')\nprint('=' * 65)\n\nckpt = torch.load(OUTPUT + 'best_model.pth', map_location=DEVICE)\nassert ckpt.get('n_ch', N_CH) == N_CH, \\\n    f\"Channel mismatch: ckpt={ckpt.get('n_ch')} vs model={N_CH}. Retrain.\"\n\nmodel = VesuviusV10(n_ch=N_CH, enc_dim=512,\n                    n_transformer_blocks=4, num_heads=8).to(DEVICE)\nmodel.load_state_dict(ckpt['state'], strict=True)\nsaved_thr = ckpt.get('thr', 0.5)\nprint(f'Checkpoint: epoch {ckpt[\"epoch\"]+1} | '\n      f'val_dice={ckpt.get(\"vd\",0):.4f} | '\n      f'opt_dice={ckpt.get(\"bd\",0):.4f} | '\n      f'thr={saved_thr:.2f}')\n\nprob_map, msk1 = predict_fragment(model, FRAG1, Z_SLICES)\nH_m, W_m       = msk1.shape\nprob_crop       = prob_map[:H_m, :W_m]\n\n# Optimal threshold sweep on Fragment 1\nfrom sklearn.metrics import confusion_matrix\nbt1, bd1 = sweep_threshold(prob_crop[np.newaxis, np.newaxis],\n                            msk1[np.newaxis, np.newaxis])\n\n# ── Post-processing ─────────────────────────────────────────\nprint(f'\\nPost-processing: threshold={bt1:.2f}, morph clean, DenseCRF ...')\ncache1      = load_slice_cache(FRAG1, Z_SLICES)\nfinal_pred  = postprocess(prob_crop, cache1, Z_SLICES,\n                          threshold=bt1, min_area=150, close_k=5, crf_iters=5)\ndel cache1; gc.collect()\n\n# ── Metrics ─────────────────────────────────────────────────\npf = final_pred.flatten().astype(int)\nmf = msk1.flatten().astype(int)\ntn, fp, fn, tp_v = confusion_matrix(mf, pf, labels=[0, 1]).ravel()\nprec = tp_v / (tp_v + fp + 1e-8)\nrec  = tp_v / (tp_v + fn + 1e-8)\nf1   = 2 * prec * rec / (prec + rec + 1e-8)\ndice_full = (2 * tp_v + 1) / (final_pred.sum() + msk1.sum() + 1)\n\nink_mean   = float(prob_crop[msk1 == 1].mean())\nnoink_mean = float(prob_crop[msk1 == 0].mean())\n\nprint('\\n' + '=' * 65)\nprint('RESULTS — FRAGMENT 1')\nprint('=' * 65)\nprint(f'Dice Score  : {dice_full:.4f}')\nprint(f'F1 Score    : {f1:.4f}')\nprint(f'Precision   : {prec:.4f}')\nprint(f'Recall      : {rec:.4f}')\nprint(f'Threshold   : {bt1:.2f}')\nprint(f'TP={tp_v} | TN={tn} | FP={fp} | FN={fn}')\nprint(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}')\nprint(f'Ink/bg sep  : {ink_mean-noink_mean:+.3f}')\ntarget_str  = '✓ PASSED' if f1 >= 0.80 else '✗ Below 0.80 target'\nprint(f'F1 ≥ 0.80   : {target_str}')\nprint('=' * 65)\n\n# ── Visualisation ────────────────────────────────────────────\nfig, ax = plt.subplots(2, 3, figsize=(18, 12))\nax[0,0].imshow(msk1,        cmap='gray');    ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop,   cmap='inferno'); ax[0,1].set_title('Probability Map')\nax[0,2].imshow(final_pred,  cmap='gray');    ax[0,2].set_title(\n    f'Prediction (post-processed)  dice={dice_full:.3f}')\n\nerr = np.zeros((*msk1.shape, 3), dtype=np.uint8)\nerr[(final_pred == 1) & (msk1 == 1)] = [0,   255, 0]\nerr[(final_pred == 1) & (msk1 == 0)] = [255,   0, 0]\nerr[(final_pred == 0) & (msk1 == 1)] = [0,     0, 255]\nax[1,0].imshow(err); ax[1,0].set_title('TP=green  FP=red  FN=blue')\n\nax[1,1].hist(prob_crop[msk1 == 1].ravel(), bins=50, alpha=0.7,\n             label=f'ink (μ={ink_mean:.2f})',       color='orange', density=True)\nax[1,1].hist(prob_crop[msk1 == 0].ravel(), bins=50, alpha=0.7,\n             label=f'no-ink (μ={noink_mean:.2f})',  color='blue',   density=True)\nax[1,1].axvline(bt1, color='r', ls='--', label=f'thr={bt1:.2f}')\nax[1,1].set_title('Probability Distribution')\nax[1,1].legend(); ax[1,1].set_xlabel('Probability')\n\nts = np.arange(0.15, 0.90, 0.01); ds = []\nfor t in ts:\n    p_b = (prob_crop > t).astype(np.float32)\n    ds.append((2 * (p_b * msk1).sum() + 1) / (p_b.sum() + msk1.sum() + 1))\nax[1,2].plot(ts, ds)\nax[1,2].axvline(bt1, color='r', ls='--', label=f'best={bt1:.2f}')\nax[1,2].set_title('Dice vs Threshold')\nax[1,2].set_xlabel('Threshold'); ax[1,2].set_ylabel('Dice')\nax[1,2].legend(); ax[1,2].grid(True)\n\nfor a in [ax[0,0], ax[0,1], ax[0,2], ax[1,0]]: a.axis('off')\nplt.suptitle(\n    f'Fragment 1 — Dice={dice_full:.4f}  F1={f1:.4f}  '\n    f'Prec={prec:.4f}  Rec={rec:.4f}',\n    fontsize=13, y=1.01\n)\nplt.tight_layout()\nplt.savefig(OUTPUT + 'frag1_prediction.png', dpi=100, bbox_inches='tight')\nplt.close()\n\n# ── Save results ─────────────────────────────────────────────\nwith open(OUTPUT + 'final_results.txt', 'w') as f:\n    f.write('='*50 + '\\n')\n    f.write('VESUVIUS V10 — FINAL TEST RESULTS\\n')\n    f.write('='*50 + '\\n')\n    f.write(f'Dice        : {dice_full:.4f}\\n')\n    f.write(f'F1          : {f1:.4f}\\n')\n    f.write(f'Precision   : {prec:.4f}\\n')\n    f.write(f'Recall      : {rec:.4f}\\n')\n    f.write(f'Threshold   : {bt1:.2f}\\n')\n    f.write(f'TP={tp_v} TN={tn} FP={fp} FN={fn}\\n')\n    f.write(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Ink/bg sep  : {ink_mean-noink_mean:+.3f}\\n')\n    f.write(f'Temperature : {TEMPERATURE}\\n')\n    f.write(f'Post-proc   : morph_clean + DenseCRF={HAS_CRF}\\n')\n    f.write('='*50 + '\\n')\n\n# ── Save prediction images ────────────────────────────────────\ncv2.imwrite(OUTPUT + 'frag1_prob_map.png',\n            (prob_crop * 255).astype(np.uint8))\ncv2.imwrite(OUTPUT + 'frag1_prediction_binary.png',\n            (final_pred * 255).astype(np.uint8))\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  best_model.pth | curves.png | frag1_prediction.png')\nprint('  frag1_prob_map.png | frag1_prediction_binary.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\n# ============================================================\n#  VESUVIUS INK DETECTION — v10 + DDPM  (FINAL SPEED FIX)\n#\n#  ANOMALY MAP TIMING HISTORY & FIXES:\n#  ┌──────────────────────────────────────────────────────┐\n#  │ v1 (sequential DDPM):  34k patches × 40 steps        │\n#  │   = 1.36M fwd passes  →  6+ hours per fragment       │\n#  │                                                       │\n#  │ v2 (F.unfold + DDIM):  OOM — unfold on full 14k×9k  │\n#  │   fragment = 38 GB tensor before any GPU work starts  │\n#  │                                                       │\n#  │ v3 (coord_chunk loop): Still slow — 279s/chunk due   │\n#  │   to two bottlenecks:                                 │\n#  │   (a) 4000×17 Python-level numpy slices per chunk    │\n#  │   (b) ANOM_BATCH=32 → too many kernel launches for   │\n#  │       a tiny 280k-param model                         │\n#  │                                                       │\n#  │ THIS VERSION (v4): All three issues fixed:            │\n#  │   (a) sliding_window_view → vectorized gather          │\n#  │       replaces 68k Python slices with 17 numpy ops    │\n#  │   (b) ANOM_BATCH=128 → 4× fewer kernel launches      │\n#  │   (c) autocast hoisted outside the DDIM step loop     │\n#  │   Expected: ~10-20 minutes per fragment on P100       │\n#  └──────────────────────────────────────────────────────┘\n# ============================================================\n\nimport os, gc, cv2, math, warnings, random, time\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nfrom numpy.lib.stride_tricks import sliding_window_view\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:64\"\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = False\n\n# ── paths ────────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\nDEVICE      = 'cuda' if torch.cuda.is_available() else 'cpu'\nUSE_AMP     = (DEVICE == 'cuda')\nZ_SLICES    = [18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34]\nN_CH        = len(Z_SLICES)   # 17\nPATCH_SIZE  = 224\nTEMPERATURE = 1.3\n\nprint(f\"Device : {DEVICE}  |  N_CH={N_CH}  |  Patch={PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. CT CACHE\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel()[::20])\n    assert raw, f\"No slices in {vol_dir}\"\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: {len(cache)} slices, {mb:.0f} MB\")\n    return cache\n\n\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H, W):\n                m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n            return (m > 0).astype(np.uint8)\n    z = list(cache.keys())[len(cache)//2]\n    return (cache[z] > 0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  2. DDPM HYPERPARAMETERS & SCHEDULE\n# ════════════════════════════════════════════════════════════\nDDPM_T           = 200\nDDPM_PATCH_XY    = 64\nDDPM_BETA_START  = 1e-4\nDDPM_BETA_END    = 0.02\nDDPM_T_NOISE     = 100    # partial noise level for anomaly detection\nDDIM_STEPS       = 5      # deterministic reverse in 5 steps instead of 100\nDDPM_EPOCHS      = 6\nDDPM_BATCH       = 8      # training batch (small model, small patch)\nANOM_BATCH       = 128    # inference batch — large to amortise kernel-launch cost\nDDPM_LR          = 2e-4\nDDPM_BASE_CH     = 16\nDDPM_MAX_PATCHES = 2_000\nCOORD_CHUNK      = 2192   # patches materialised in CPU RAM at once\n\n\ndef build_ddpm_schedule(T, beta_start, beta_end, device):\n    betas          = torch.linspace(beta_start, beta_end, T, device=device)\n    alphas         = 1. - betas\n    alpha_bar      = torch.cumprod(alphas, dim=0)\n    alpha_bar_prev = F.pad(alpha_bar[:-1], (1, 0), value=1.0)\n    sqrt_ab        = alpha_bar.sqrt()\n    sqrt_1mab      = (1. - alpha_bar).sqrt()\n    post_var       = betas * (1. - alpha_bar_prev) / (1. - alpha_bar + 1e-8)\n    return dict(betas=betas, alphas=alphas, alpha_bar=alpha_bar,\n                alpha_bar_prev=alpha_bar_prev, sqrt_ab=sqrt_ab,\n                sqrt_1mab=sqrt_1mab, post_var=post_var)\n\n\n# ════════════════════════════════════════════════════════════\n#  3. 3D DDPM UNET  (~280k params, fits easily on P100)\n# ════════════════════════════════════════════════════════════\nclass GroupConv3d(nn.Module):\n    def __init__(self, cin, cout, k=3, groups=4, stride=1):\n        super().__init__()\n        g = min(groups, cin, cout)\n        while cin % g != 0 or cout % g != 0:\n            g -= 1\n        g = max(g, 1)\n        self.c = nn.Sequential(\n            nn.Conv3d(cin, cout, k, padding=k//2, stride=stride,\n                      groups=g, bias=False),\n            nn.GroupNorm(min(8, cout), cout),\n            nn.SiLU(inplace=True))\n    def forward(self, x): return self.c(x)\n\n\nclass ResBlock3d(nn.Module):\n    def __init__(self, ch, t_emb_dim=64, groups=4):\n        super().__init__()\n        self.c1   = GroupConv3d(ch, ch, groups=groups)\n        self.c2   = GroupConv3d(ch, ch, groups=groups)\n        self.temb = nn.Linear(t_emb_dim, ch)\n        self.norm = nn.GroupNorm(min(8, ch), ch)\n    def forward(self, x, t_emb):\n        h = self.c1(x) + self.temb(t_emb).view(t_emb.shape[0], -1, 1, 1, 1)\n        return x + self.c2(self.norm(h))\n\n\nclass UNet3D_DDPM(nn.Module):\n    \"\"\"\n    Encoder:  1→c, c→2c (stride-2), 2c→4c (stride-2), bottleneck 4c\n    Decoder:  4c→2c (deconv), cat skip 2c → proj 2c,\n              2c→c  (deconv), cat skip c  → proj c,  → out 1\n    \"\"\"\n    def __init__(self, in_ch=1, base_ch=DDPM_BASE_CH, T=DDPM_T, groups=4):\n        super().__init__()\n        t_dim = 64; c = base_ch\n        self.t_embed = nn.Sequential(\n            nn.Embedding(T, t_dim), nn.Linear(t_dim, t_dim),\n            nn.SiLU(), nn.Linear(t_dim, t_dim))\n        self.enc1  = nn.Sequential(GroupConv3d(in_ch, c, groups=1),\n                                   GroupConv3d(c, c, groups=min(groups,c)))\n        self.down1 = nn.Conv3d(c, c*2, 2, stride=2, bias=False)\n        self.res1  = ResBlock3d(c*2, t_dim, min(groups,c*2))\n        self.enc2  = GroupConv3d(c*2, c*2, groups=min(groups,c*2))\n        self.down2 = nn.Conv3d(c*2, c*4, 2, stride=2, bias=False)\n        self.res2  = ResBlock3d(c*4, t_dim, min(groups,c*4))\n        self.bot   = ResBlock3d(c*4, t_dim, min(groups,c*4))\n        self.up2   = nn.ConvTranspose3d(c*4, c*2, 2, stride=2, bias=False)\n        self.dres2 = ResBlock3d(c*4, t_dim, min(groups,c*4))\n        self.dproj2= GroupConv3d(c*4, c*2, k=1, groups=1)\n        self.up1   = nn.ConvTranspose3d(c*2, c, 2, stride=2, bias=False)\n        self.dres1 = ResBlock3d(c*2, t_dim, min(groups,c*2))\n        self.dproj1= GroupConv3d(c*2, c, k=1, groups=1)\n        self.out   = nn.Conv3d(c, in_ch, 1)\n\n    def forward(self, x, t):\n        te = self.t_embed(t)\n        e1 = self.enc1(x)\n        e2 = self.res1(self.down1(e1), te)\n        e3 = self.res2(self.down2(self.enc2(e2)), te)\n        b  = self.bot(e3, te)\n        d2 = self.up2(b)\n        if d2.shape != e2.shape:\n            d2 = F.interpolate(d2, size=e2.shape[2:], mode='trilinear', align_corners=False)\n        d2 = self.dproj2(self.dres2(torch.cat([d2, e2], 1), te))\n        d1 = self.up1(d2)\n        if d1.shape != e1.shape:\n            d1 = F.interpolate(d1, size=e1.shape[2:], mode='trilinear', align_corners=False)\n        d1 = self.dproj1(self.dres1(torch.cat([d1, e1], 1), te))\n        return self.out(d1)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. DDPM TRAINING HELPERS\n# ════════════════════════════════════════════════════════════\ndef q_sample(x0, t, sched):\n    \"\"\"Forward diffusion: add noise at timestep t.\"\"\"\n    noise   = torch.randn_like(x0)\n    sqrt_ab = sched['sqrt_ab'][t].view(-1, 1, 1, 1, 1)\n    sqrt_1m = sched['sqrt_1mab'][t].view(-1, 1, 1, 1, 1)\n    return sqrt_ab * x0 + sqrt_1m * noise, noise\n\n\n# ════════════════════════════════════════════════════════════\n#  5. FAST ANOMALY MAP via DDIM + sliding_window_view\n#\n#  Three-level speed optimisation:\n#\n#  LEVEL 1 — DDIM (5 steps, not 100)\n#    DDPM reverse requires stepping t=99→98→...→0 (100 calls).\n#    DDIM skips to t=99→74→49→24→0 (5 calls) using:\n#      x_{t-1} = √ᾱ_{t-1}·x̂₀ + √(1-ᾱ_{t-1})·ε_θ\n#    No noise term → deterministic → 5 steps enough for MSE signal.\n#    Speedup: 20×\n#\n#  LEVEL 2 — Batched inference (ANOM_BATCH=128)\n#    Each GPU call processes 128 patches simultaneously.\n#    The DDPM is 280k params; 128×(1,17,64,64) fp16 ≈ 70 MB.\n#    Reduces Python↔GPU round-trips by 128×.\n#    Speedup: ~128× in kernel-launch overhead.\n#\n#  LEVEL 3 — sliding_window_view patch extraction\n#    sliding_window_view(arr, (p,p)) returns a zero-copy view\n#    of shape (H-p+1, W-p+1, p, p). Gathering 8192 patches\n#    across 17 z-slices = 17 vectorized fancy-index ops instead\n#    of 8192×17 = 139,264 individual Python-level slice calls.\n#    Speedup: ~100× in CPU patch-building time.\n#\n#  MEMORY SAFETY:\n#    Never materialises more than COORD_CHUNK patches at once.\n#    Peak CPU RAM per chunk: 8192×17×64×64×4 = 2.3 GB — safe.\n#    Peak GPU RAM per step:  128×1×17×64×64×2 (fp16) ≈ 18 MB.\n# ════════════════════════════════════════════════════════════\n\n@torch.no_grad()\ndef compute_anomaly_map_fast(ddpm_model, cache, z_list, sched,\n                              patch_xy=DDPM_PATCH_XY,\n                              t_noise=DDPM_T_NOISE,\n                              ddim_steps=DDIM_STEPS,\n                              anom_batch=ANOM_BATCH,\n                              papyrus_mask=None,\n                              coord_chunk=COORD_CHUNK,\n                              device='cuda'):\n    ddpm_model.eval()\n    H, W   = next(iter(cache.values())).shape\n    stride = patch_xy // 2   # 50% overlap for smooth boundaries\n\n    # ── STEP A: build coord list (integers only, negligible RAM) ──\n    coords = [(y, x) for y in range(0, H - patch_xy + 1, stride)\n                      for x in range(0, W - patch_xy + 1, stride)]\n    # ensure last row and column are always covered\n    if not coords or coords[-1][0] < H - patch_xy:\n        coords += [(H - patch_xy, x)\n                   for x in range(0, W - patch_xy + 1, stride)]\n    if not coords or coords[-1][1] < W - patch_xy:\n        coords += [(y, W - patch_xy)\n                   for y in range(0, H - patch_xy + 1, stride)]\n\n    if papyrus_mask is not None:\n        half = patch_xy // 2\n        coords = [(y, x) for (y, x) in coords\n                  if papyrus_mask[y + half, x + half] > 0]\n\n    n_patches = len(coords)\n    print(f\"    {n_patches} patches (stride={stride}, patch={patch_xy})\")\n\n    score_map = np.zeros((H, W), np.float32)\n    count_map = np.zeros((H, W), np.float32)\n\n    # ── STEP B: DDIM schedule (precomputed) ───────────────────────\n    t_seq   = np.round(np.linspace(t_noise - 1, 0, ddim_steps)).astype(int).tolist()\n    t_pairs = list(zip(t_seq[:-1], t_seq[1:]))\n    t_pairs.append((t_seq[-1], 0))   # final step always goes to t=0\n\n    # precompute scalar alpha_bar values for DDIM to avoid repeated dict lookups\n    ab_vals  = sched['alpha_bar'].to(device)   # (T,) on GPU\n\n    # ── STEP C: zero-copy sliding window views (one per z-slice) ──\n    # shape of each view: (H-p+1, W-p+1, p, p) — shares memory with cache\n    z_wins = [sliding_window_view(cache[z], (patch_xy, patch_xy))\n              for z in z_list]   # float16, zero-copy\n    # clamp to valid index range (in case of edge coords)\n    max_y = H - patch_xy\n    max_x = W - patch_xy\n\n    n_chunks = math.ceil(n_patches / coord_chunk)\n    t0_total = time.time()\n\n    for c_idx, c_start in enumerate(range(0, n_patches, coord_chunk)):\n        chunk_coords = coords[c_start : c_start + coord_chunk]\n        chunk_n      = len(chunk_coords)\n        ys = np.array([c[0] for c in chunk_coords], dtype=np.int32)\n        xs = np.array([c[1] for c in chunk_coords], dtype=np.int32)\n        # clamp (safety for edge-added coords)\n        ys = np.clip(ys, 0, max_y)\n        xs = np.clip(xs, 0, max_x)\n\n        # ── LEVEL 3: vectorized gather (17 fancy-index ops) ───────\n        # z_wins[ci][ys, xs] → (chunk_n, patch_xy, patch_xy) in one call\n        chunk_vol = np.empty((chunk_n, len(z_list), patch_xy, patch_xy),\n                             dtype=np.float32)\n        for ci, win in enumerate(z_wins):\n            chunk_vol[:, ci] = win[ys, xs]   # float16→float32 implicit cast\n        chunk_vol = chunk_vol * 2.0 - 1.0    # normalise to [-1, 1]\n        chunk_mse = np.empty((chunk_n, patch_xy, patch_xy), dtype=np.float32)\n\n        # ── LEVELS 1+2: DDIM reverse in large batches ─────────────\n        for b_start in range(0, chunk_n, anom_batch):\n            b_end = min(b_start + anom_batch, chunk_n)\n            # (B, 1, N_CH, p, p) — fp32 on GPU\n            x0 = torch.from_numpy(\n                chunk_vol[b_start:b_end]).unsqueeze(1).to(device)\n            B  = x0.shape[0]\n\n            # forward diffuse x0 → x_{t_noise}\n            noise = torch.randn_like(x0)\n            x_t   = (ab_vals[t_noise-1].sqrt() * x0\n                     + (1 - ab_vals[t_noise-1]).sqrt() * noise)\n            del noise\n\n            # DDIM reverse: 5 deterministic steps\n            # Run the full multi-step sequence under ONE autocast context\n            # to avoid repeated context entry/exit overhead\n            with autocast(enabled=USE_AMP):\n                for t_cur, t_prev in t_pairs:\n                    B2         = x_t.shape[0]\n                    t_tensor   = torch.full((B2,), t_cur,\n                                           dtype=torch.long, device=device)\n                    eps        = ddpm_model(x_t, t_tensor)\n                    ab_cur     = ab_vals[t_cur]\n                    ab_prev    = ab_vals[t_prev]\n                    x0_pred    = ((x_t - (1 - ab_cur).sqrt() * eps)\n                                  / (ab_cur.sqrt() + 1e-8)).clamp(-1., 1.)\n                    x_t        = (ab_prev.sqrt() * x0_pred\n                                  + (1 - ab_prev).sqrt() * eps)\n                    del eps, x0_pred, t_tensor\n\n            # MSE averaged over depth (N_CH) dimension → (B, p, p)\n            mse = ((x_t.float() - x0.float()) ** 2).mean(dim=(1, 2))\n            chunk_mse[b_start:b_end] = mse.cpu().numpy()\n            del x0, x_t, mse\n\n        # ── accumulate MSE into score map ─────────────────────────\n        # Still a Python loop but over chunk_n, not n_patches × steps\n        for i in range(chunk_n):\n            y, x = int(ys[i]), int(xs[i])\n            score_map[y:y+patch_xy, x:x+patch_xy] += chunk_mse[i]\n            count_map[y:y+patch_xy, x:x+patch_xy] += 1.\n\n        del chunk_vol, chunk_mse\n        if device == 'cuda': torch.cuda.empty_cache()\n\n        elapsed = time.time() - t0_total\n        done    = c_idx + 1\n        eta     = elapsed / done * (n_chunks - done)\n        print(f\"\\r    chunk {done}/{n_chunks} | \"\n              f\"elapsed {elapsed/60:.1f}min | ETA {eta/60:.1f}min\", end='')\n\n    print()\n    score_map /= (count_map + 1e-8)\n    covered = count_map > 0\n    if covered.any():\n        s_min = score_map[covered].min()\n        s_max = score_map[covered].max()\n        score_map = np.where(covered,\n                             (score_map - s_min) / (s_max - s_min + 1e-8),\n                             0.0)\n    return score_map.astype(np.float32)\n\n\n# ════════════════════════════════════════════════════════════\n#  6. DDPM TRAINING DATASET (no-ink patches only)\n# ════════════════════════════════════════════════════════════\nclass NoInkDataset3D(Dataset):\n    def __init__(self, frag_paths, z_list, patch_xy=DDPM_PATCH_XY,\n                 max_patches=DDPM_MAX_PATCHES):\n        self.patch_xy = patch_xy; self.z_list = z_list\n        self.patches  = []; self.caches = {}\n\n        for fp in frag_paths:\n            cache = load_slice_cache(fp, z_list)\n            self.caches[fp] = cache\n            H, W = next(iter(cache.values())).shape\n            pap  = load_papyrus_mask(fp, H, W, cache)\n            ink  = cv2.imread(os.path.join(fp, 'inklabels.png'), 0)\n            if ink is None:\n                no_ink = pap\n            else:\n                if ink.shape != (H, W):\n                    ink = cv2.resize(ink, (W, H), interpolation=cv2.INTER_NEAREST)\n                no_ink = ((ink == 0) & (pap > 0)).astype(np.uint8)\n\n            stride = patch_xy // 2\n            coords = [(fp, y, x)\n                      for y in range(0, H - patch_xy + 1, stride)\n                      for x in range(0, W - patch_xy + 1, stride)\n                      if no_ink[y:y+patch_xy, x:x+patch_xy].mean() > 0.9]\n            np.random.shuffle(coords)\n            self.patches.extend(coords)\n            print(f\"  NoInk [{os.path.basename(fp)}]: {len(coords)} patches\")\n\n        if max_patches > 0 and len(self.patches) > max_patches:\n            np.random.shuffle(self.patches)\n            self.patches = self.patches[:max_patches]\n        print(f\"  Total no-ink patches: {len(self.patches)}\")\n\n    def __len__(self): return len(self.patches)\n\n    def __getitem__(self, idx):\n        fp, y, x = self.patches[idx]\n        vol = np.stack([\n            self.caches[fp][z][y:y+self.patch_xy,\n                               x:x+self.patch_xy].astype(np.float32)\n            for z in self.z_list], axis=0)\n        return torch.from_numpy(vol * 2. - 1.).unsqueeze(0).float()\n\n\ndef train_ddpm(model, dataset, sched, n_epochs, lr, device, batch_size):\n    dl       = DataLoader(dataset, batch_size=batch_size, shuffle=True,\n                          num_workers=0, pin_memory=False)\n    opt      = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)\n    sched_lr = optim.lr_scheduler.CosineAnnealingLR(opt, T_max=n_epochs, eta_min=1e-5)\n    scaler   = GradScaler(enabled=USE_AMP)\n    print(f\"\\n{'='*60}\")\n    print(f\"DDPM Training: {n_epochs} epochs | {len(dl)} batches\")\n    print(f\"  T={DDPM_T}  t_noise={DDPM_T_NOISE}  ddim_steps={DDIM_STEPS}\")\n    print(f\"  anom_batch={ANOM_BATCH}  coord_chunk={COORD_CHUNK}\")\n    print(f\"{'='*60}\")\n    model.train()\n    for ep in range(n_epochs):\n        tot = 0.; steps = 0; opt.zero_grad()\n        for x0 in tqdm(dl, desc=f'DDPM ep{ep+1}', leave=False):\n            x0 = x0.to(device)\n            t  = torch.randint(0, DDPM_T, (x0.shape[0],), device=device)\n            x_t, noise = q_sample(x0, t, sched)\n            with autocast(enabled=USE_AMP):\n                loss = F.mse_loss(model(x_t, t), noise)\n            scaler.scale(loss).backward()\n            scaler.unscale_(opt)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.)\n            scaler.step(opt); scaler.update(); opt.zero_grad()\n            tot += loss.item(); steps += 1\n            del x0, x_t, noise, loss\n        sched_lr.step()\n        print(f\"  ep{ep+1}/{n_epochs} loss={tot/max(steps,1):.5f} \"\n              f\"lr={opt.param_groups[0]['lr']:.1e}\")\n    return model\n\n\n# ════════════════════════════════════════════════════════════\n#  7. V10 SEGMENTATION MODEL\n# ════════════════════════════════════════════════════════════\nSTRIDE_TR    = 112\nSTRIDE_INF   = 56\nBATCH_SIZE   = 2\nGRAD_ACCUM   = 8\nEPOCHS       = 15\nLR           = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 10\nVAL_SPLIT    = 0.30\nINK_MIN_POS  = 0.02\nNEG_RATIO    = 0.5\nCH_DROP_PROB = 0.3\nCH_DROP_MAX  = 3\nN_CH_SEG     = N_CH + 1   # 17 CT + 1 anomaly = 18\n\n\nclass SheetSurfaceDetector(nn.Module):\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, 32, 3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.ReLU(inplace=True),\n            nn.Conv2d(16, 1, 1))\n    def forward(self, x): return self.net(x)\n\n\ndef token_merge(tokens, saliency, merge_ratio=0.25):\n    B, C, N = tokens.shape\n    _, sort_idx = saliency.squeeze(1).sort(dim=1)\n    n_merge = int(N * merge_ratio)\n    if n_merge % 2 == 1: n_merge -= 1\n    merge_idx = sort_idx[:, :n_merge]; keep_idx = sort_idx[:, n_merge:]\n    def gather(t, idx):\n        return t.gather(2, idx.unsqueeze(1).expand(-1, C, -1))\n    t_m  = gather(tokens, merge_idx); t_k = gather(tokens, keep_idx)\n    return (torch.cat([t_k, (t_m[:,:,0::2]+t_m[:,:,1::2])/2], 2),\n            (keep_idx, merge_idx, n_merge, N))\n\n\ndef token_unmerge(tokens_out, info, C):\n    keep_idx, merge_idx, n_merge, N = info\n    B      = tokens_out.shape[0]\n    n_keep = tokens_out.shape[2] - n_merge // 2\n    t_keep = tokens_out[:,:,:n_keep]\n    t_exp  = tokens_out[:,:,n_keep:].repeat_interleave(2, dim=2)\n    out    = torch.zeros(B, C, N, device=tokens_out.device, dtype=tokens_out.dtype)\n    out.scatter_(2, keep_idx.unsqueeze(1).expand(-1,C,-1), t_keep)\n    out.scatter_(2, merge_idx.unsqueeze(1).expand(-1,C,-1), t_exp)\n    return out\n\n\nclass AxialAttention(nn.Module):\n    def __init__(self, dim, num_heads=8, axis='x', dropout=0.1):\n        super().__init__()\n        self.axis = axis; self.num_heads = num_heads\n        self.head_dim = dim // num_heads; self.scale = self.head_dim ** -0.5\n        self.qkv  = nn.Linear(dim, dim*3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.drop = nn.Dropout(dropout)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        if self.axis == 'x': x_r = x.permute(0,2,3,1).reshape(B*H, W, C)\n        else:                 x_r = x.permute(0,3,2,1).reshape(B*W, H, C)\n        res = x_r; BN, L, _ = x_r.shape\n        qkv = self.qkv(self.norm(x_r)).reshape(BN,L,3,self.num_heads,self.head_dim)\n        q,k,v = qkv.permute(2,0,3,1,4).unbind(0)\n        attn  = (q @ k.transpose(-2,-1)) * self.scale\n        pos   = torch.arange(L, dtype=torch.float32, device=x.device)\n        attn  = attn - torch.log(\n            (pos.unsqueeze(0) - pos.unsqueeze(1)).abs().float() + 1.\n        ).unsqueeze(0).unsqueeze(0)\n        attn  = self.drop(attn.softmax(-1))\n        out   = (attn @ v).transpose(1,2).reshape(BN,L,C)\n        out   = self.proj(out) + res\n        if self.axis == 'x': return out.reshape(B,H,W,C).permute(0,3,1,2)\n        else:                 return out.reshape(B,W,H,C).permute(0,3,2,1)\n\n\nclass AxialTransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=8, mlp_ratio=4, dropout=0.1):\n        super().__init__()\n        self.attn_x = AxialAttention(dim, num_heads, 'x', dropout)\n        self.attn_y = AxialAttention(dim, num_heads, 'y', dropout)\n        self.norm   = nn.LayerNorm(dim)\n        mlp_dim     = int(dim*mlp_ratio)\n        self.ffn    = nn.Sequential(\n            nn.Linear(dim, mlp_dim), nn.GELU(), nn.Dropout(dropout),\n            nn.Linear(mlp_dim, dim), nn.Dropout(dropout))\n    def forward(self, x):\n        x = self.attn_x(x); x = self.attn_y(x)\n        B,C,H,W = x.shape\n        xf = self.ffn(self.norm(x.permute(0,2,3,1).reshape(-1,C)))\n        return x + xf.reshape(B,H,W,C).permute(0,3,1,2)\n\n\nclass VesuviusV10(nn.Module):\n    def __init__(self, n_ch=N_CH_SEG, enc_dim=512,\n                 n_transformer_blocks=4, num_heads=8):\n        super().__init__()\n        self.backbone = smp.Unet(\n            encoder_name='resnet34', encoder_weights='imagenet',\n            in_channels=n_ch, classes=1, decoder_attention_type='scse')\n        self.transformer_blocks = nn.Sequential(*[\n            AxialTransformerBlock(enc_dim, num_heads, 4, 0.1)\n            for _ in range(n_transformer_blocks)])\n        self.sheet_detector  = SheetSurfaceDetector(n_ch)\n        self._enc_dim        = enc_dim\n        self.bottleneck_proj = nn.Sequential(\n            nn.Conv2d(enc_dim, enc_dim, 1, bias=False),\n            nn.BatchNorm2d(enc_dim), nn.GELU())\n        self.ds_head3 = nn.Conv2d(256, 1, 1)\n        self.ds_head2 = nn.Conv2d(128, 1, 1)\n        self.ds_head1 = nn.Conv2d( 64, 1, 1)\n\n    def forward(self, x):\n        B = x.shape[0]\n        sal    = self.sheet_detector(x)\n        feats  = self.backbone.encoder(x)\n        bn     = feats[-1]\n        bH,bW  = bn.shape[2], bn.shape[3]; N = bH*bW\n        sal_d  = F.adaptive_avg_pool2d(torch.sigmoid(sal),(bH,bW)).reshape(B,1,N)\n        tok    = bn.reshape(B, self._enc_dim, N)\n        tok_m, uinfo = token_merge(tok, sal_d, 0.25)\n        M      = tok_m.shape[2]; sq = int(math.ceil(math.sqrt(M)))\n        pad    = sq*sq - M\n        if pad: tok_m = F.pad(tok_m,(0,pad))\n        tok_2d = self.transformer_blocks(tok_m.reshape(B, self._enc_dim, sq, sq))\n        tok_f  = tok_2d.reshape(B, self._enc_dim, sq*sq)[:,:,:M]\n        bn_out = self.bottleneck_proj(\n            token_unmerge(tok_f, uinfo, self._enc_dim).reshape(B,self._enc_dim,bH,bW) + bn)\n        feats_m = list(feats); feats_m[-1] = bn_out\n        dec_out = self.backbone.decoder(*feats_m)\n        main    = self.backbone.segmentation_head(dec_out)\n        ds3=ds2=ds1=None\n        try:\n            db = self.backbone.decoder.blocks\n            if len(db)>=1: f0=db[0](feats_m[-1],feats_m[-2]); ds3=self.ds_head3(f0)\n            if len(db)>=2: f1=db[1](f0,feats_m[-3]);          ds2=self.ds_head2(f1)\n            if len(db)>=3: f2=db[2](f1,feats_m[-4]);          ds1=self.ds_head1(f2)\n        except Exception: pass\n        return main, sal, ds3, ds2, ds1\n\n\n# ════════════════════════════════════════════════════════════\n#  8. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.RandomRotate90(p=1.0), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.12, scale_limit=0.18, rotate_limit=35,\n                       border_mode=cv2.BORDER_REFLECT, p=0.65),\n    A.ElasticTransform(alpha=1.0, sigma=50, alpha_affine=50,\n                       border_mode=cv2.BORDER_REFLECT, p=0.4),\n    A.GridDistortion(num_steps=5, distort_limit=0.3,\n                     border_mode=cv2.BORDER_REFLECT, p=0.3),\n    A.RandomBrightnessContrast(0.25, 0.25, p=0.55),\n    A.GaussNoise(var_limit=(0.001, 0.005), p=0.35),\n    A.GaussianBlur(blur_limit=(3, 7), p=0.25),\n    A.CoarseDropout(max_holes=6, max_height=28, max_width=28, fill_value=0, p=0.35),\n])\n\ndef channel_dropout(img_np, prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < prob:\n        n   = np.random.randint(1, max_drop+1)\n        idx = np.random.choice(img_np.shape[2], n, replace=False)\n        img_np = img_np.copy(); img_np[:,:,idx] = 0.\n    return img_np\n\ndef cutmix_batch(imgs, msks, alpha=0.4):\n    B,C,H,W = imgs.shape; lam = np.random.beta(alpha, alpha)\n    perm = torch.randperm(B, device=imgs.device)\n    cx=np.random.randint(W); cy=np.random.randint(H)\n    bw=int(W*math.sqrt(1-lam)); bh=int(H*math.sqrt(1-lam))\n    x1=max(0,cx-bw//2); x2=min(W,cx+bw//2)\n    y1=max(0,cy-bh//2); y2=min(H,cy+bh//2)\n    i=imgs.clone(); m=msks.clone()\n    i[:,:,y1:y2,x1:x2]=imgs[perm,:,y1:y2,x1:x2]\n    m[:,:,y1:y2,x1:x2]=msks[perm,:,y1:y2,x1:x2]\n    return i, m\n\n\n# ════════════════════════════════════════════════════════════\n#  9. SEGMENTATION DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride, anomaly_maps=None,\n                 transform=None, neg_ratio=0., apply_ch_dropout=False):\n        self.cache      = load_slice_cache(frag_path, z_list)\n        self.z_list     = z_list; self.tf = transform\n        self.ch_dropout = apply_ch_dropout\n        self.anom_map   = anomaly_maps.get(frag_path) if anomaly_maps else None\n        msk = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n        assert msk is not None\n        self.mask = (msk>0).astype(np.uint8)\n        H, W = self.mask.shape\n        ir   = os.path.join(frag_path,'mask.png')\n        pap  = ((cv2.imread(ir,0)>0).astype(np.uint8)\n                if os.path.exists(ir) else None)\n        if pap is None:\n            zm = z_list[len(z_list)//2]\n            pap = (self.cache[zm]>0.1).astype(np.uint8)\n        pos_c, neg_c = [], []\n        for y in range(0, H-PATCH_SIZE+1, stride):\n            for x in range(0, W-PATCH_SIZE+1, stride):\n                if pap[y:y+PATCH_SIZE,x:x+PATCH_SIZE].mean()<0.5: continue\n                ink = self.mask[y:y+PATCH_SIZE,x:x+PATCH_SIZE].mean()\n                if   ink >= INK_MIN_POS:            pos_c.append((y,x,1))\n                elif ink < 0.001 and neg_ratio > 0: neg_c.append((y,x,0))\n        n_neg = int(len(pos_c)*neg_ratio)\n        np.random.shuffle(neg_c); neg_c = neg_c[:n_neg]\n        self.coords  = pos_c+neg_c\n        self.weights = np.array([3. if c[2]==1 else 1.\n                                 for c in self.coords], np.float32)\n        print(f\"  [{os.path.basename(frag_path)}] \"\n              f\"{len(pos_c)} pos + {len(neg_c)} neg | \"\n              f\"anomaly={'yes' if self.anom_map is not None else 'zeros'}\")\n\n    def __len__(self): return len(self.coords)\n\n    def get_patch(self, idx):\n        y,x,_ = self.coords[idx]\n        img   = np.stack([self.cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                          for z in self.z_list], axis=-1)\n        a     = (self.anom_map[y:y+PATCH_SIZE,x:x+PATCH_SIZE,None].astype(np.float32)\n                 if self.anom_map is not None\n                 else np.zeros((PATCH_SIZE,PATCH_SIZE,1),np.float32))\n        return np.concatenate([img,a],-1), self.mask[y:y+PATCH_SIZE,x:x+PATCH_SIZE].copy()\n\n    def __getitem__(self, idx):\n        img, msk = self.get_patch(idx)\n        if self.tf:\n            o=self.tf(image=img,mask=msk); img,msk=o['image'],o['mask']\n        if self.ch_dropout:\n            img[:,:,:N_CH] = channel_dropout(img[:,:,:N_CH])\n        return (torch.from_numpy(img).permute(2,0,1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\n\n# ════════════════════════════════════════════════════════════\n#  10. LOSS & METRICS\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce); a_t = target*alpha+(1-target)*(1-alpha)\n    return (a_t*((1-p_t)**gamma)*bce).mean()\n\ndef dice_loss(pred, target, smooth=1.):\n    p = torch.sigmoid(pred)\n    i = (p*target).sum(dim=(2,3)); u = p.sum(dim=(2,3))+target.sum(dim=(2,3))\n    return 1.-((2.*i+smooth)/(u+smooth)).mean()\n\ndef combined_loss(pred, target, eps=0.05):\n    return 0.5*focal_loss(pred,target*(1-eps)+0.5*eps) + 0.5*dice_loss(pred,target)\n\ndef multiscale_loss(main,sal,ds3,ds2,ds1,target):\n    loss = combined_loss(main,target)\n    loss += 0.2*combined_loss(sal,F.adaptive_avg_pool2d(target,sal.shape[-2:]))\n    for l,w in [(ds3,0.15),(ds2,0.10),(ds1,0.05)]:\n        if l is not None:\n            loss += w*combined_loss(l,F.adaptive_avg_pool2d(target,l.shape[-2:]))\n    return loss\n\ndef batch_dice(logits, masks, thr=0.5):\n    p = (torch.sigmoid(logits)>thr).float()\n    i = (p*masks).sum(dim=(1,2,3)); u = p.sum(dim=(1,2,3))+masks.sum(dim=(1,2,3))\n    return ((2.*i+1e-5)/(u+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t,best_d = 0.5,0.\n    for t in np.arange(0.10,0.90,0.01):\n        p=(probs>t).astype(np.float32)\n        d=(2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d>best_d: best_d,best_t=d,float(t)\n    return best_t,best_d\n\ndef compute_sep(probs, masks):\n    if (masks>0.5).any() and (masks<0.5).any():\n        return float(probs[masks>0.5].mean()-probs[masks<0.5].mean())\n    return 0.\n\n\n# ════════════════════════════════════════════════════════════\n#  11. GAUSSIAN WEIGHT\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c=sz//2; s=sz//4; y,x=np.mgrid[0:sz,0:sz]\n    return np.exp(-((x-c)**2+(y-c)**2)/(2*s**2)).astype(np.float32)\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  12. RUN — STAGE 1: TRAIN DDPM\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\" + \"=\"*65)\nprint(\"STAGE 1 — 3D DDPM: Clean Papyrus Prior\")\nprint(\"=\"*65)\n\nnoink_ds   = NoInkDataset3D([FRAG2,FRAG3], Z_SLICES,\n                              DDPM_PATCH_XY, DDPM_MAX_PATCHES)\nddpm_sched = build_ddpm_schedule(DDPM_T, DDPM_BETA_START, DDPM_BETA_END, DEVICE)\nddpm_model = UNet3D_DDPM(in_ch=1, base_ch=DDPM_BASE_CH, T=DDPM_T).to(DEVICE)\nprint(f\"DDPM params: {sum(p.numel() for p in ddpm_model.parameters())/1e3:.1f}k\")\n\nddpm_model = train_ddpm(ddpm_model, noink_ds, ddpm_sched,\n                         DDPM_EPOCHS, DDPM_LR, DEVICE, DDPM_BATCH)\ntorch.save({'state':ddpm_model.state_dict(),'T':DDPM_T,\n            't_noise':DDPM_T_NOISE,'ddim_steps':DDIM_STEPS},\n           OUTPUT+'ddpm_model.pth')\nprint(\"DDPM saved.\")\n\n\n# ════════════════════════════════════════════════════════════\n#  13. GENERATE ANOMALY MAPS\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\" + \"=\"*65)\nprint(f\"Anomaly maps — DDIM {DDIM_STEPS} steps | batch {ANOM_BATCH} | \"\n      f\"chunk {COORD_CHUNK}\")\nprint(\"=\"*65)\n\nanomaly_maps = {}\nfor fp, tag in [(FRAG2,'frag2'),(FRAG3,'frag3')]:\n    print(f\"\\n  [{tag}]\")\n    t0    = time.time()\n    cache = noink_ds.caches[fp]\n    H_fp, W_fp = next(iter(cache.values())).shape\n    pap_fp     = load_papyrus_mask(fp, H_fp, W_fp, cache)\n    amap  = compute_anomaly_map_fast(\n        ddpm_model, cache, Z_SLICES, ddpm_sched,\n        patch_xy=DDPM_PATCH_XY, t_noise=DDPM_T_NOISE,\n        ddim_steps=DDIM_STEPS, anom_batch=ANOM_BATCH,\n        papyrus_mask=pap_fp, coord_chunk=COORD_CHUNK, device=DEVICE)\n    anomaly_maps[fp] = amap\n    cv2.imwrite(OUTPUT+f'anomaly_{tag}.png', (amap*255).astype(np.uint8))\n    print(f\"    Done in {(time.time()-t0)/60:.1f}min | \"\n          f\"min={amap.min():.3f} max={amap.max():.3f} mean={amap.mean():.3f}\")\n\nddpm_model_cpu = ddpm_model.cpu()\ndel ddpm_model; gc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\nprint(\"\\nDDPM moved to CPU — VRAM freed for V10.\")\n\n\n# ════════════════════════════════════════════════════════════\n#  14. STAGE 2: TRAIN V10 SEGMENTATION\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\"+\"=\"*65)\nprint(f\"STAGE 2 — V10 Segmentation ({N_CH_SEG}-ch: {N_CH} CT + 1 anomaly)\")\nprint(\"=\"*65)\n\nprint('\\n── Fragment 2 ──')\nds2 = VesuviusDataset(FRAG2, Z_SLICES, STRIDE_TR, anomaly_maps=anomaly_maps,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\nprint('\\n── Fragment 3 ──')\nds3 = VesuviusDataset(FRAG3, Z_SLICES, STRIDE_TR, anomaly_maps=anomaly_maps,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\n\nn_total   = len(ds2)+len(ds3)\nrng       = np.random.RandomState(SEED); all_idx = rng.permutation(n_total)\nn_val     = int(n_total*VAL_SPLIT)\ntrain_idx = all_idx[n_val:].tolist(); val_idx = all_idx[:n_val].tolist()\nprint(f'Total: {n_total} | Train: {len(train_idx)} | Val: {len(val_idx)}')\n\n\nclass _Subset(Dataset):\n    def __init__(self, d2, d3, idxs, is_train):\n        self.d2=d2; self.d3=d3; self.n2=len(d2)\n        self.idxs=idxs; self.is_train=is_train\n    def __len__(self): return len(self.idxs)\n    def __getitem__(self, i):\n        g = self.idxs[i]\n        if self.is_train:\n            return self.d2[g] if g<self.n2 else self.d3[g-self.n2]\n        img,msk = (self.d2.get_patch(g) if g<self.n2\n                   else self.d3.get_patch(g-self.n2))\n        return (torch.from_numpy(img).permute(2,0,1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\ntrain_ds = _Subset(ds2,ds3,train_idx,True)\nval_ds   = _Subset(ds2,ds3,val_idx,  False)\nall_w    = np.concatenate([ds2.weights, ds3.weights])\nsampler  = WeightedRandomSampler(torch.from_numpy(all_w[train_idx]).float(),\n                                  len(train_ds), True)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                      num_workers=0, pin_memory=False)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=0, pin_memory=False)\nprint(f'Train batches: {len(train_dl)} | Val: {len(val_dl)}')\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nmodel    = VesuviusV10(n_ch=N_CH_SEG, enc_dim=512,\n                       n_transformer_blocks=4, num_heads=8).to(DEVICE)\nprint(f'V10 params: {sum(p.numel() for p in model.parameters()if p.requires_grad)/1e6:.2f}M')\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='max', factor=0.5, patience=5, min_lr=5e-6, verbose=True)\nscaler    = GradScaler(enabled=USE_AMP)\nbest_dice = 0.; pat_cnt = 0\nhistory   = dict(tl=[],vl=[],td=[],vd=[],sep=[],lr=[])\n\nfor epoch in range(EPOCHS):\n    model.train(); tl=td=0.; optimizer.zero_grad()\n    for step,(imgs,msks) in enumerate(\n            tqdm(train_dl,desc=f'Ep{epoch+1:02d}▸train',leave=False)):\n        imgs=imgs.to(DEVICE); msks=msks.to(DEVICE)\n        if random.random()<0.30: imgs,msks=cutmix_batch(imgs,msks,0.4)\n        with autocast(enabled=USE_AMP):\n            main,sal,d3,d2_l,d1=model(imgs)\n            loss=multiscale_loss(main,sal,d3,d2_l,d1,msks)/GRAD_ACCUM\n        scaler.scale(loss).backward()\n        tl+=loss.item()*GRAD_ACCUM; td+=batch_dice(main.detach(),msks)\n        del imgs,msks,main,sal,d3,d2_l,d1,loss\n        if (step+1)%GRAD_ACCUM==0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(),1.)\n            scaler.step(optimizer); scaler.update(); optimizer.zero_grad()\n            if DEVICE=='cuda': torch.cuda.empty_cache()\n    tl/=len(train_dl); td/=len(train_dl)\n\n    model.eval(); vl=vd=0.; all_p,all_m=[],[]\n    with torch.no_grad():\n        for imgs,msks in tqdm(val_dl,desc=f'Ep{epoch+1:02d}▸val  ',leave=False):\n            imgs=imgs.to(DEVICE); msks=msks.to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                main,sal,d3,d2_l,d1=model(imgs)\n                loss=multiscale_loss(main,sal,d3,d2_l,d1,msks)\n            vl+=loss.item(); vd+=batch_dice(main,msks)\n            all_p.append(torch.sigmoid(main/TEMPERATURE).cpu().numpy())\n            all_m.append(msks.cpu().numpy())\n            del imgs,msks,main,sal,d3,d2_l,d1,loss\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    vl/=len(val_dl); vd/=len(val_dl)\n    P=np.concatenate(all_p); M=np.concatenate(all_m); del all_p,all_m\n    bt,bd=sweep_threshold(P,M); sep=compute_sep(P,M)\n    ink_m=float(P[M>0.5].mean()) if (M>0.5).any() else 0.\n    bg_m =float(P[M<0.5].mean()) if (M<0.5).any() else 0.\n    del P,M; gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    scheduler.step(vd)\n    lr_now=optimizer.param_groups[0]['lr']\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['sep'].append(sep); history['lr'].append(lr_now)\n    print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n          f'train={tl:.4f}/{td:.4f} | val={vl:.4f}/{vd:.4f} | '\n          f'thr={bt:.2f}→{bd:.4f} | sep={sep:+.3f} '\n          f'[ink={ink_m:.3f} bg={bg_m:.3f}]')\n    metric = max(vd,bd) - max(0.,bt-0.60)*0.3\n    if metric>best_dice:\n        best_dice=metric; pat_cnt=0\n        torch.save({'epoch':epoch,'state':model.state_dict(),'thr':bt,\n                    'metric':metric,'n_ch':N_CH_SEG,'bd':bd,'vd':vd},\n                   OUTPUT+'best_model.pth')\n        print(f'  ✓ saved (metric={metric:.4f})')\n    else:\n        pat_cnt+=1\n        if pat_cnt>=PATIENCE: print(f'  early stop ep{epoch+1}'); break\n\nprint(f'\\nBest metric: {best_dice:.4f}')\n\nfig,axes=plt.subplots(1,4,figsize=(20,4))\naxes[0].plot(history['tl'],label='train'); axes[0].plot(history['vl'],label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'],label='train'); axes[1].plot(history['vd'],label='val')\naxes[1].axhline(0.80,color='r',ls='--'); axes[1].set_title('Dice')\naxes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['sep']); axes[2].axhline(0.4,color='orange',ls='--')\naxes[2].set_title('Sep'); axes[2].grid(True)\naxes[3].plot(history['lr']); axes[3].set_title('LR'); axes[3].grid(True)\nplt.tight_layout(); plt.savefig(OUTPUT+'curves.png',dpi=80); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  15. INFERENCE\n# ════════════════════════════════════════════════════════════\ndef predict_fragment_with_ddpm(seg_model, ddpm_model_cpu, frag_path,\n                                z_list, anom_map_precomputed=None):\n    seg_model.eval()\n\n    if anom_map_precomputed is not None:\n        amap  = anom_map_precomputed\n        cache = load_slice_cache(frag_path, z_list)\n        print(\"  Using precomputed anomaly map\")\n    else:\n        print(\"  Computing anomaly map via DDIM (fast)...\")\n        cache  = load_slice_cache(frag_path, z_list)\n        H_fp, W_fp = next(iter(cache.values())).shape\n        pap_fp = load_papyrus_mask(frag_path, H_fp, W_fp, cache)\n        ddpm_g = ddpm_model_cpu.to(DEVICE)\n        sched_g= build_ddpm_schedule(DDPM_T, DDPM_BETA_START, DDPM_BETA_END, DEVICE)\n        t0     = time.time()\n        amap   = compute_anomaly_map_fast(\n            ddpm_g, cache, z_list, sched_g,\n            patch_xy=DDPM_PATCH_XY, t_noise=DDPM_T_NOISE,\n            ddim_steps=DDIM_STEPS, anom_batch=ANOM_BATCH,\n            papyrus_mask=pap_fp, coord_chunk=COORD_CHUNK, device=DEVICE)\n        print(f\"  Anomaly map done in {(time.time()-t0)/60:.1f}min\")\n        ddpm_g = ddpm_g.cpu(); del ddpm_g, sched_g\n        gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    msk = (cv2.imread(os.path.join(frag_path,'inklabels.png'),0)>0).astype(np.uint8)\n    H, W = msk.shape\n    pred_map = np.zeros((H,W),np.float32); wgt_map=np.zeros((H,W),np.float32)\n    coords   = [(y,x) for y in range(0,H-PATCH_SIZE+1,STRIDE_INF)\n                       for x in range(0,W-PATCH_SIZE+1,STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y,x) in tqdm(coords,desc='Seg infer',leave=True):\n            slices=[cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                    for z in z_list]\n            img = np.stack(slices,axis=-1)\n            a   = amap[y:y+PATCH_SIZE,x:x+PATCH_SIZE,None].astype(np.float32)\n            img = np.concatenate([img,a],-1)\n            t   = torch.from_numpy(img).permute(2,0,1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                logit,*_ = seg_model(t)\n            p = torch.sigmoid(logit/TEMPERATURE).squeeze().cpu().float().numpy()\n            pred_map[y:y+PATCH_SIZE,x:x+PATCH_SIZE] += p*GW\n            wgt_map [y:y+PATCH_SIZE,x:x+PATCH_SIZE] += GW\n            del t,logit,p\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    seg_prob = pred_map/(wgt_map+1e-8)\n    H_a,W_a  = amap.shape; h=min(H,H_a); w=min(W,W_a)\n    fused    = np.zeros((H,W),np.float32)\n    fused[:h,:w] = 0.70*seg_prob[:h,:w] + 0.30*amap[:h,:w]\n    return seg_prob, amap, np.clip(fused,0.,1.), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  16. FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\"+\"=\"*65)\nprint(\"FINAL TEST — FRAGMENT 1\")\nprint(\"=\"*65)\n\nckpt  = torch.load(OUTPUT+'best_model.pth', map_location=DEVICE)\nmodel = VesuviusV10(n_ch=N_CH_SEG, enc_dim=512,\n                    n_transformer_blocks=4, num_heads=8).to(DEVICE)\nmodel.load_state_dict(ckpt['state'])\nprint(f\"Checkpoint: ep={ckpt['epoch']+1}  \"\n      f\"val={ckpt.get('vd',0):.4f}  thr={ckpt['thr']:.2f}\")\n\nseg_prob,amap1,fused_prob,msk1 = predict_fragment_with_ddpm(\n    model, ddpm_model_cpu, FRAG1, Z_SLICES, anom_map_precomputed=None)\n\nH_m,W_m  = msk1.shape\nseg_c     = seg_prob[:H_m,:W_m]\nfuse_c    = fused_prob[:H_m,:W_m]\namap_c    = amap1[:H_m,:W_m]\n\nbt_seg, bd_seg   = sweep_threshold(seg_c,  msk1)\nbt_fuse,bd_fuse  = sweep_threshold(fuse_c, msk1)\nprint(f\"  Seg  : thr={bt_seg:.2f}  dice={bd_seg:.4f}\")\nprint(f\"  Fused: thr={bt_fuse:.2f}  dice={bd_fuse:.4f}\")\n\nif bd_fuse >= bd_seg:\n    final_prob=fuse_c; best_thr=bt_fuse; mode=\"Fused (Seg+DDPM)\"\nelse:\n    final_prob=seg_c;  best_thr=bt_seg;  mode=\"Seg only\"\n\npred   = (final_prob>best_thr).astype(np.uint8)\nkernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(5,5))\npred   = cv2.morphologyEx(pred,cv2.MORPH_CLOSE,kernel)\nn_lab,labels,stats,_ = cv2.connectedComponentsWithStats(pred)\ncleaned = np.zeros_like(pred)\nfor i in range(1,n_lab):\n    if stats[i,cv2.CC_STAT_AREA]>=150: cleaned[labels==i]=1\npred = cleaned\n\npf=pred.flatten().astype(int); mf=msk1.flatten().astype(int)\ntn,fp,fn,tp_v = confusion_matrix(mf,pf,labels=[0,1]).ravel()\nprec=tp_v/(tp_v+fp+1e-8); rec=tp_v/(tp_v+fn+1e-8)\nf1  =2*prec*rec/(prec+rec+1e-8)\ndice=(2*tp_v+1)/(pred.sum()+msk1.sum()+1)\nink_m   =float(final_prob[msk1==1].mean()) if (msk1==1).any() else 0.\nbg_m    =float(final_prob[msk1==0].mean()) if (msk1==0).any() else 0.\nddpm_ink=float(amap_c[msk1==1].mean()) if (msk1==1).any() else 0.\nddpm_bg =float(amap_c[msk1==0].mean()) if (msk1==0).any() else 0.\n\nprint('\\n'+'='*65)\nprint(f'RESULTS — FRAGMENT 1  [{mode}]')\nprint('='*65)\nprint(f'Dice      : {dice:.4f}')\nprint(f'F1        : {f1:.4f}')\nprint(f'Precision : {prec:.4f}')\nprint(f'Recall    : {rec:.4f}')\nprint(f'FP/TP     : {fp/(tp_v+1e-8):.2f}')\nprint(f'Sep (seg) : {ink_m-bg_m:+.3f}')\nprint(f'Sep (ddpm): {ddpm_ink-ddpm_bg:+.3f}  ← unsupervised signal quality')\nprint(f'F1≥0.80   : {\"✓ PASSED\" if f1>=0.80 else \"✗ Below target\"}')\nprint('='*65)\n\nfig,ax=plt.subplots(2,4,figsize=(22,12))\nax[0,0].imshow(msk1,  cmap='gray');    ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(amap_c,cmap='hot');     ax[0,1].set_title('DDPM Anomaly Map')\nax[0,2].imshow(seg_c, cmap='inferno'); ax[0,2].set_title(f'Seg prob (dice={bd_seg:.3f})')\nax[0,3].imshow(fuse_c,cmap='inferno'); ax[0,3].set_title(f'Fused prob (dice={bd_fuse:.3f})')\nax[1,0].imshow(pred,  cmap='gray');    ax[1,0].set_title(f'{mode}  Dice={dice:.4f}')\nerr=np.zeros((*msk1.shape,3),dtype=np.uint8)\nerr[(pred==1)&(msk1==1)]=[0,255,0]\nerr[(pred==1)&(msk1==0)]=[255,0,0]\nerr[(pred==0)&(msk1==1)]=[0,0,255]\nax[1,1].imshow(err); ax[1,1].set_title('TP=green FP=red FN=blue')\nax[1,2].hist(seg_c[msk1==1].ravel(),bins=60,alpha=0.7,\n             label=f'ink μ={ink_m:.3f}',color='orange',density=True)\nax[1,2].hist(seg_c[msk1==0].ravel(),bins=60,alpha=0.7,\n             label=f'bg μ={bg_m:.3f}',color='blue',density=True)\nax[1,2].axvline(best_thr,color='r',ls='--'); ax[1,2].legend()\nax[1,2].set_title('Seg distribution')\nax[1,3].hist(amap_c[msk1==1].ravel(),bins=60,alpha=0.7,\n             label=f'ink μ={ddpm_ink:.3f}',color='orange',density=True)\nax[1,3].hist(amap_c[msk1==0].ravel(),bins=60,alpha=0.7,\n             label=f'bg μ={ddpm_bg:.3f}',color='blue',density=True)\nax[1,3].legend(); ax[1,3].set_title('DDPM anomaly distribution')\nfor a in ax[0]: a.axis('off')\nax[1,0].axis('off'); ax[1,1].axis('off')\nplt.suptitle(f'Fragment 1 | Dice={dice:.4f}  F1={f1:.4f}  '\n             f'Prec={prec:.4f}  Rec={rec:.4f}',fontsize=12)\nplt.tight_layout()\nplt.savefig(OUTPUT+'frag1_prediction.png',dpi=80,bbox_inches='tight')\nplt.close()\n\ncv2.imwrite(OUTPUT+'frag1_anomaly.png',  (amap_c*255).astype(np.uint8))\ncv2.imwrite(OUTPUT+'frag1_seg_prob.png', (seg_c*255).astype(np.uint8))\ncv2.imwrite(OUTPUT+'frag1_binary.png',   (pred*255).astype(np.uint8))\n\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write('VESUVIUS V10+DDPM (FINAL SPEED FIX)\\n'+'='*55+'\\n')\n    f.write(f'Mode      : {mode}\\n')\n    f.write(f'Dice      : {dice:.4f}\\n')\n    f.write(f'F1        : {f1:.4f}\\n')\n    f.write(f'Precision : {prec:.4f}\\n')\n    f.write(f'Recall    : {rec:.4f}\\n')\n    f.write(f'FP/TP     : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Sep (seg) : {ink_m-bg_m:+.3f}\\n')\n    f.write(f'Sep (ddpm): {ddpm_ink-ddpm_bg:+.3f}\\n')\n    f.write(f'DDIM steps: {DDIM_STEPS}  t_noise={DDPM_T_NOISE}\\n')\n    f.write(f'anom_batch: {ANOM_BATCH}  coord_chunk={COORD_CHUNK}\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  ddpm_model.pth | best_model.pth | curves.png')\nprint('  frag1_prediction.png | frag1_anomaly.png')\nprint('  frag1_seg_prob.png | frag1_binary.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n# ============================================================\n#  VESUVIUS INK DETECTION — v11\n#  Root-cause fixes from v10 results (F1=0.45, early stop ep10)\n#\n#  Diagnosed problems:\n#    1. val_dice=0.32 vs opt_dice=0.49 → threshold instability\n#       Fix: temperature=1.0, label smoothing removed, BCE weight\n#    2. FP/TP=1.64 → model over-predicts ink\n#       Fix: higher focal alpha for negatives, pos_weight tuned\n#    3. Early stop epoch 10 → Axial-Transformer too heavy for\n#       tiny 7×7 bottleneck, collapses gradients\n#       Fix: Transformer moved to DECODER (higher-res features),\n#            bottleneck stays pure CNN, transformer at 1/8 scale\n#    4. Token merging on 7×7=49 tokens → merges 12 tokens, useless\n#       Fix: removed from bottleneck; sparse attention now at 1/8\n#    5. CutMix during training with hard labels → confuses boundary\n#       Fix: CutMix probability reduced, only after warm-up\n#    6. WeightedRandomSampler + neg_ratio=0.5 → too many negatives\n#       Fix: neg_ratio=0.3, sampler weight ratio 4:1\n#\n#  New Architecture:\n#    ResNet34 encoder (proven, v9-style)\n#    ↓ 4 skip connections\n#    Bottleneck → pure CNN (no transformer here)\n#    Decoder stage at 1/8 resolution → Axial Attention here\n#    (1/8 spatial = 28×28 tokens from 224×224 patch → rich context)\n#    SCSE attention in all decoder blocks\n#    Deep supervision at 2 scales\n#    Sheet-surface saliency head (auxiliary)\n#\n#  Log-polar positional bias: applied at 1/8 decoder stage\n#  Post-processing: morph clean only (CRF optional)\n#  NO TTA at test\n#  Pseudo-labelling: supported\n# ============================================================\n\nimport os, gc, cv2, math, warnings, random\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\ntry:\n    import pydensecrf.densecrf as dcrf\n    from pydensecrf.utils import create_pairwise_bilateral, create_pairwise_gaussian\n    HAS_CRF = True\nexcept ImportError:\n    HAS_CRF = False\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\n\n# ── reproducibility ──────────────────────────────────────────\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ── paths ────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ─────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 224\nSTRIDE_TR    = 112          # training stride (50% overlap)\nSTRIDE_INF   = 56           # inference stride (25% overlap, dense)\nBATCH_SIZE   = 5            # ResNet34 + light transformer is manageable\nGRAD_ACCUM   = 3            # effective batch = 18\nEPOCHS       = 10           # more epochs, proper early stop\nLR           = 3e-4         # higher initial LR with warm-up\nLR_MIN       = 5e-6\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 12           # more patience\nNUM_WORKERS  = 0\nVAL_SPLIT    = 0.25         # 75/25 — more training data\n\n# Z-slices: 17 central slices (same as v9/v10)\nZ_SLICES = [20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36]\nN_CH     = len(Z_SLICES)    # 17\n\n# Patch sampling\nINK_MIN_POS  = 0.015        # slightly lower → more positive samples\nNEG_RATIO    = 0.30         # fewer negatives (was 0.50)\nCH_DROP_PROB = 0.25\nCH_DROP_MAX  = 3\n\n# Loss\nFOCAL_ALPHA  = 0.75         # weight for positive class in focal\nFOCAL_GAMMA  = 2.5\nBCE_POS_WEIGHT = 2.5        # extra weight on positive pixels\n\n# Inference\nTEMPERATURE  = 1.0          # NO temperature scaling (root cause of gap)\n\n# Pseudo-label\nUSE_PSEUDO        = False\nPSEUDO_FRAGS      = []\nPSEUDO_THRESHOLD  = 0.65\n\nprint(f\"Device: {DEVICE} | Channels: {N_CH} | Patch: {PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM: {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1.  SLICE CACHE  (fragment percentile normalisation)\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list, verbose=True):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    if not raw:\n        raise ValueError(f\"No slices in {vol_dir}\")\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    if verbose:\n        mb = sum(v.nbytes for v in cache.values()) / 1e6\n        print(f\"  Cache: {len(cache)} slices ({mb:.0f} MB)\")\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2.  LOG-POLAR POSITIONAL BIAS\n#      Applied at 1/8 decoder feature map (28×28 from 224 patch)\n#      → 784 tokens, rich context, manageable memory\n# ════════════════════════════════════════════════════════════\n_LP_CACHE: dict = {}\n\ndef get_logpolar_bias_1d(seq_len: int, num_heads: int, device) -> torch.Tensor:\n    \"\"\"\n    1-D log-distance bias for axial attention along one axis.\n    Shape: [num_heads, seq_len, seq_len]\n    Each head gets a different frequency of the log-distance.\n    \"\"\"\n    key = (seq_len, num_heads, str(device))\n    if key in _LP_CACHE:\n        return _LP_CACHE[key]\n\n    pos  = torch.arange(seq_len, dtype=torch.float32)\n    d    = (pos.unsqueeze(0) - pos.unsqueeze(1)).abs()      # [L, L]\n    log_d = torch.log(d + 1.0)                              # [L, L]\n\n    # Each head uses a different slope (learnable-free multi-scale)\n    slopes = torch.exp(\n        torch.linspace(math.log(1.0), math.log(8.0), num_heads)\n    )                                                        # [H]\n    # bias[h, i, j] = -slope[h] * log_d[i,j]\n    bias = -slopes[:, None, None] * log_d[None, :, :]       # [H, L, L]\n    bias = bias.to(device)\n    _LP_CACHE[key] = bias\n    return bias\n\n\n# ════════════════════════════════════════════════════════════\n#  3.  AXIAL ATTENTION BLOCK  (X then Y, at decoder 1/8 scale)\n#      This is where the transformer insight lives:\n#      - 28-token sequences (from 224/8 = 28)\n#      - log-polar bias encodes scroll geometry\n#      - alternating X/Y is O(2·N·L) not O(N²)\n# ════════════════════════════════════════════════════════════\nclass AxialAttention1D(nn.Module):\n    \"\"\"\n    Self-attention along one spatial axis of a 2-D feature map.\n    x: [B, C, H, W]  →  [B, C, H, W]\n    \"\"\"\n    def __init__(self, dim: int, num_heads: int = 8,\n                 axis: str = 'x', dropout: float = 0.1):\n        super().__init__()\n        assert axis in ('x', 'y')\n        assert dim % num_heads == 0\n        self.axis      = axis\n        self.num_heads = num_heads\n        self.head_dim  = dim // num_heads\n        self.scale     = self.head_dim ** -0.5\n\n        self.norm = nn.LayerNorm(dim)\n        self.qkv  = nn.Linear(dim, dim * 3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.drop = nn.Dropout(dropout)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        B, C, H, W = x.shape\n        if self.axis == 'x':\n            # reshape: (B*H) sequences of length W\n            x_r = x.permute(0, 2, 3, 1).reshape(B * H, W, C)\n            L   = W\n        else:\n            # reshape: (B*W) sequences of length H\n            x_r = x.permute(0, 3, 2, 1).reshape(B * W, H, C)\n            L   = H\n\n        res  = x_r\n        x_n  = self.norm(x_r)                              # [BN, L, C]\n        BN   = x_n.shape[0]\n\n        qkv  = self.qkv(x_n).reshape(BN, L, 3, self.num_heads, self.head_dim)\n        qkv  = qkv.permute(2, 0, 3, 1, 4)                # [3, BN, nh, L, hd]\n        q, k, v = qkv.unbind(0)                           # each [BN, nh, L, hd]\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale     # [BN, nh, L, L]\n\n        # Add log-polar 1-D bias\n        bias = get_logpolar_bias_1d(L, self.num_heads, x.device)\n        attn = attn + bias.unsqueeze(0)                    # broadcast over BN\n\n        attn = attn.softmax(-1)\n        attn = self.drop(attn)\n\n        out  = (attn @ v).transpose(1, 2).reshape(BN, L, C)\n        out  = self.proj(out) + res                        # pre-norm residual\n\n        if self.axis == 'x':\n            out = out.reshape(B, H, W, C).permute(0, 3, 1, 2)\n        else:\n            out = out.reshape(B, W, H, C).permute(0, 3, 2, 1)\n        return out\n\n\nclass AxialTransformerBlock(nn.Module):\n    \"\"\"X-axis → Y-axis attention + FFN.  Acts on [B, C, H, W].\"\"\"\n    def __init__(self, dim: int, num_heads: int = 8,\n                 mlp_ratio: float = 4.0, dropout: float = 0.1):\n        super().__init__()\n        self.attn_x = AxialAttention1D(dim, num_heads, 'x', dropout)\n        self.attn_y = AxialAttention1D(dim, num_heads, 'y', dropout)\n        mlp_dim     = int(dim * mlp_ratio)\n        self.norm   = nn.LayerNorm(dim)\n        self.ffn    = nn.Sequential(\n            nn.Linear(dim, mlp_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(mlp_dim, dim),\n            nn.Dropout(dropout),\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.attn_x(x)\n        x = self.attn_y(x)\n        # FFN\n        B, C, H, W = x.shape\n        xf = x.permute(0, 2, 3, 1).reshape(-1, C)\n        xf = self.ffn(self.norm(xf)).reshape(B, H, W, C).permute(0, 3, 1, 2)\n        return x + xf\n\n\n# ════════════════════════════════════════════════════════════\n#  4.  SHEET-SURFACE DETECTOR  (lightweight auxiliary head)\n#      Shallow CNN → coarse ink-likelihood from raw input.\n#      Auxiliary loss forces early layers to be ink-aware.\n# ════════════════════════════════════════════════════════════\nclass SheetSurfaceDetector(nn.Module):\n    def __init__(self, in_ch: int):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, 32, 3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.GELU(),\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.GELU(),\n            nn.Conv2d(16,  1, 1),\n        )\n    def forward(self, x):\n        return self.net(x)   # [B, 1, H, W]\n\n\n# ════════════════════════════════════════════════════════════\n#  5.  DECODER BLOCK WITH AXIAL ATTENTION\n#      Standard UNet-like up-block + SCSE + optional\n#      axial-transformer on the output feature map.\n# ════════════════════════════════════════════════════════════\nclass SCSEModule(nn.Module):\n    def __init__(self, ch: int, reduction: int = 16):\n        super().__init__()\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(ch, ch // reduction),\n            nn.ReLU(inplace=True),\n            nn.Linear(ch // reduction, ch),\n            nn.Sigmoid(),\n        )\n        self.sSE = nn.Sequential(nn.Conv2d(ch, 1, 1), nn.Sigmoid())\n\n    def forward(self, x):\n        return (x * self.cSE(x).view(x.shape[0], -1, 1, 1)\n                + x * self.sSE(x))\n\n\nclass DecoderBlock(nn.Module):\n    def __init__(self, in_ch: int, skip_ch: int, out_ch: int,\n                 use_axial: bool = False, num_heads: int = 8):\n        super().__init__()\n        self.up   = nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride=2)\n        cat_ch    = in_ch // 2 + skip_ch\n        self.conv = nn.Sequential(\n            nn.Conv2d(cat_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch), nn.GELU(),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch), nn.GELU(),\n        )\n        self.attn = SCSEModule(out_ch)\n        self.transformer = (\n            AxialTransformerBlock(out_ch, num_heads=num_heads,\n                                  mlp_ratio=4.0, dropout=0.1)\n            if use_axial else None\n        )\n\n    def forward(self, x: torch.Tensor,\n                skip: torch.Tensor) -> torch.Tensor:\n        x = self.up(x)\n        # handle size mismatch\n        if x.shape[-2:] != skip.shape[-2:]:\n            x = F.interpolate(x, size=skip.shape[-2:],\n                              mode='bilinear', align_corners=False)\n        x = torch.cat([x, skip], dim=1)\n        x = self.conv(x)\n        x = self.attn(x)\n        if self.transformer is not None:\n            x = self.transformer(x)\n        return x\n\n\n# ════════════════════════════════════════════════════════════\n#  6.  FULL MODEL — VesuviusV11\n#      ResNet34 encoder  (proven cross-fragment generaliser)\n#      Pure CNN bottleneck  (no transformer here — too coarse)\n#      Decoder stage 2 (1/8 scale, 28×28) → Axial Transformer\n#      Deep supervision at 2 decoder stages\n#      Sheet-surface auxiliary head\n# ════════════════════════════════════════════════════════════\nclass VesuviusV11(nn.Module):\n    \"\"\"\n    ResNet34 encoder + custom decoder with axial attention at 1/8 scale.\n    Encoder channel sizes (ResNet34):\n        s0: 64   @ 1/2\n        s1: 64   @ 1/4\n        s2: 128  @ 1/8\n        s3: 256  @ 1/16\n        s4: 512  @ 1/32  (bottleneck)\n    \"\"\"\n    ENC_CHANNELS = [64, 64, 128, 256, 512]\n\n    def __init__(self, n_ch: int = N_CH, num_heads: int = 8):\n        super().__init__()\n\n        # ── Encoder (ResNet34, ImageNet weights) ───────────────\n        _backbone = smp.Unet(\n            encoder_name    = 'resnet34',\n            encoder_weights = 'imagenet',\n            in_channels     = n_ch,\n            classes         = 1,\n        )\n        self.encoder = _backbone.encoder   # returns list of features\n\n        c = self.ENC_CHANNELS\n\n        # ── Bottleneck refinement (pure CNN) ───────────────────\n        self.bottleneck_refine = nn.Sequential(\n            nn.Conv2d(c[4], c[4], 3, padding=1, bias=False),\n            nn.BatchNorm2d(c[4]), nn.GELU(),\n            nn.Dropout2d(0.3),\n        )\n\n        # ── Decoder ────────────────────────────────────────────\n        # d0: 512 → 256,  skip=s3(256),  out=256   no transformer\n        # d1: 256 → 128,  skip=s2(128),  out=128   AXIAL here (1/8)\n        # d2: 128 →  64,  skip=s1(64),   out=64    no transformer\n        # d3:  64 →  32,  skip=s0(64),   out=32    no transformer\n        self.d0 = DecoderBlock(c[4],   c[3], 256, use_axial=False)\n        self.d1 = DecoderBlock(256,    c[2], 128, use_axial=True,\n                               num_heads=num_heads)   # ← TRANSFORMER\n        self.d2 = DecoderBlock(128,    c[1],  64, use_axial=False)\n        self.d3 = DecoderBlock( 64,    c[0],  32, use_axial=False)\n\n        # ── Segmentation head ──────────────────────────────────\n        self.head = nn.Sequential(\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.GELU(),\n            nn.Conv2d(16,  1, 1),\n        )\n\n        # ── Deep supervision heads ─────────────────────────────\n        self.ds1 = nn.Conv2d(128, 1, 1)   # after d1 (1/8 scale)\n        self.ds2 = nn.Conv2d( 64, 1, 1)   # after d2 (1/4 scale)\n\n        # ── Sheet-surface auxiliary head ───────────────────────\n        self.surface_det = SheetSurfaceDetector(n_ch)\n\n    def forward(self, x):\n        # Auxiliary coarse saliency\n        sal_logit = self.surface_det(x)    # [B, 1, H, W]\n\n        # Encoder\n        feats = self.encoder(x)\n        # feats[0] = input (or first stem), feats[1..5] = s0..s4\n        # smp returns [input_padded, s0, s1, s2, s3, s4]\n        s0 = feats[1]; s1 = feats[2]; s2 = feats[3]\n        s3 = feats[4]; s4 = feats[5]\n\n        # Bottleneck\n        b  = self.bottleneck_refine(s4)       # [B, 512, H/32, W/32]\n\n        # Decoder\n        f0 = self.d0(b,  s3)                  # [B, 256, H/16, W/16]\n        f1 = self.d1(f0, s2)                  # [B, 128, H/8,  W/8 ] ← axial\n        f2 = self.d2(f1, s1)                  # [B,  64, H/4,  W/4 ]\n        f3 = self.d3(f2, s0)                  # [B,  32, H/2,  W/2 ]\n\n        # Main head + upsample to input size\n        logit = self.head(f3)                 # [B, 1, H/2, W/2]\n        logit = F.interpolate(logit, size=x.shape[-2:],\n                              mode='bilinear', align_corners=False)\n\n        # Deep supervision\n        ds1_logit = F.interpolate(self.ds1(f1), size=x.shape[-2:],\n                                  mode='bilinear', align_corners=False)\n        ds2_logit = F.interpolate(self.ds2(f2), size=x.shape[-2:],\n                                  mode='bilinear', align_corners=False)\n\n        return logit, sal_logit, ds1_logit, ds2_logit\n\n\n# ════════════════════════════════════════════════════════════\n#  7.  LOSS  (calibrated for FP/TP control)\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target,\n               alpha: float = FOCAL_ALPHA, gamma: float = FOCAL_GAMMA):\n    \"\"\"\n    alpha: weight for POSITIVE class (ink).\n    With FP/TP=1.64, we LOWER alpha to penalise false positives more.\n    \"\"\"\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    # alpha_t: alpha for positives, (1-alpha) for negatives\n    a_t = target * alpha + (1.0 - target) * (1.0 - alpha)\n    return (a_t * (1.0 - p_t) ** gamma * bce).mean()\n\n\ndef dice_loss(pred, target, smooth: float = 1.0):\n    p     = torch.sigmoid(pred)\n    inter = (p * target).sum(dim=(2, 3))\n    union = p.sum(dim=(2, 3)) + target.sum(dim=(2, 3))\n    return 1.0 - ((2.0 * inter + smooth) / (union + smooth)).mean()\n\n\ndef tversky_loss(pred, target, alpha: float = 0.3, beta: float = 0.7,\n                 smooth: float = 1.0):\n    \"\"\"\n    Tversky loss: alpha weights FP, beta weights FN.\n    With FP/TP=1.64, set alpha>0.5 to penalise FPs heavily.\n    alpha=0.3 (light FP penalty) + beta=0.7 (heavy FN penalty)\n    → actually we want alpha=0.7 to penalise FPs\n    \"\"\"\n    p     = torch.sigmoid(pred)\n    tp    = (p * target).sum(dim=(2, 3))\n    fp    = (p * (1 - target)).sum(dim=(2, 3))\n    fn    = ((1 - p) * target).sum(dim=(2, 3))\n    return 1.0 - ((tp + smooth) /\n                  (tp + alpha * fp + beta * fn + smooth)).mean()\n\n\ndef combined_loss(pred, target):\n    \"\"\"\n    No label smoothing — that caused the val/opt-threshold gap.\n    Tversky(α=0.7,β=0.3): penalises FPs more than FNs.\n    Focal: standard modulation.\n    \"\"\"\n    pw  = torch.tensor([BCE_POS_WEIGHT], device=pred.device)\n    bce = F.binary_cross_entropy_with_logits(pred, target,\n                                             pos_weight=pw)\n    fl  = focal_loss(pred, target)\n    tv  = tversky_loss(pred, target, alpha=0.6, beta=0.4)\n    return 0.30 * bce + 0.35 * fl + 0.35 * tv\n\n\ndef full_loss(logit, sal_logit, ds1, ds2, target):\n    main = combined_loss(logit, target)\n    sal  = 0.15 * combined_loss(sal_logit, target)\n    d1   = 0.20 * combined_loss(ds1, target)\n    d2   = 0.10 * combined_loss(ds2, target)\n    return main + sal + d1 + d2\n\n\n# ════════════════════════════════════════════════════════════\n#  8.  AUGMENTATION  (elastic + controlled CutMix)\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.RandomRotate90(p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(\n        shift_limit=0.10, scale_limit=0.15, rotate_limit=30,\n        border_mode=cv2.BORDER_REFLECT, p=0.60),\n    A.ElasticTransform(\n        alpha=120, sigma=6, alpha_affine=6,\n        border_mode=cv2.BORDER_REFLECT, p=0.35),\n    A.GridDistortion(\n        num_steps=5, distort_limit=0.25,\n        border_mode=cv2.BORDER_REFLECT, p=0.25),\n    A.RandomResizedCrop(\n        height=PATCH_SIZE, width=PATCH_SIZE,\n        scale=(0.60, 1.0), ratio=(0.85, 1.15), p=0.45),\n    A.RandomBrightnessContrast(0.20, 0.20, p=0.50),\n    A.GaussNoise(var_limit=(0.0005, 0.004), p=0.30),\n    A.GaussianBlur(blur_limit=(3, 5), p=0.20),\n    A.CoarseDropout(max_holes=5, max_height=24, max_width=24,\n                    fill_value=0, p=0.30),\n])\n\n\ndef channel_dropout(img_np, prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < prob:\n        n   = np.random.randint(1, max_drop + 1)\n        idx = np.random.choice(img_np.shape[2], n, replace=False)\n        img_np        = img_np.copy()\n        img_np[:, :, idx] = 0.0\n    return img_np\n\n\ndef cutmix_batch(imgs, msks, alpha: float = 0.4):\n    \"\"\"CutMix only applied after warm-up, at reduced frequency.\"\"\"\n    B, C, H, W = imgs.shape\n    lam  = np.random.beta(alpha, alpha)\n    perm = torch.randperm(B, device=imgs.device)\n    bw   = int(W * math.sqrt(1.0 - lam))\n    bh   = int(H * math.sqrt(1.0 - lam))\n    cx   = np.random.randint(W); cy = np.random.randint(H)\n    x1   = max(0, cx - bw // 2); x2 = min(W, cx + bw // 2)\n    y1   = max(0, cy - bh // 2); y2 = min(H, cy + bh // 2)\n    imgs_n = imgs.clone(); msks_n = msks.clone()\n    imgs_n[:, :, y1:y2, x1:x2] = imgs[perm, :, y1:y2, x1:x2]\n    msks_n[:, :, y1:y2, x1:x2] = msks[perm, :, y1:y2, x1:x2]\n    return imgs_n, msks_n\n\n\n# ════════════════════════════════════════════════════════════\n#  9.  DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0,\n                 apply_ch_dropout=False):\n        self.cache        = load_slice_cache(frag_path, z_list)\n        self.z_list       = z_list\n        self.tf           = transform\n        self.ch_drop      = apply_ch_dropout\n\n        msk_path = os.path.join(frag_path, 'inklabels.png')\n        msk = cv2.imread(msk_path, 0)\n        if msk is None:\n            raise FileNotFoundError(f\"No inklabels at {msk_path}\")\n        self.mask = (msk > 0).astype(np.uint8)\n\n        ir_path  = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = (cv2.imread(ir_path, 0) > 0).astype(np.uint8) \\\n                       if os.path.exists(ir_path) else None\n\n        H, W = self.mask.shape\n        pos_c, neg_c = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.sum() / (PATCH_SIZE * PATCH_SIZE)\n                if self.ir_mask is not None:\n                    on_p = self.ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    mz   = z_list[len(z_list) // 2]\n                    on_p = float(self.cache[mz][y:y+PATCH_SIZE,\n                                               x:x+PATCH_SIZE].mean()) > 0.1\n                if not on_p:\n                    continue\n                if ink >= INK_MIN_POS:\n                    pos_c.append((y, x, 1))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_c.append((y, x, 0))\n\n        n_neg = int(len(pos_c) * neg_ratio)\n        np.random.shuffle(neg_c)\n        sel_neg = neg_c[:n_neg]\n\n        self.coords  = pos_c + sel_neg\n        # 4:1 weight ratio (was 3:1) — helps address FP over-prediction\n        self.weights = np.array(\n            [4.0 if c[2] == 1 else 1.0 for c in self.coords],\n            dtype=np.float32\n        )\n        print(f\"  [{os.path.basename(frag_path)}] \"\n              f\"{len(pos_c)} pos + {len(sel_neg)} neg = \"\n              f\"{len(self.coords)} patches\")\n\n    def __len__(self): return len(self.coords)\n\n    def get_patch(self, idx):\n        y, x, _ = self.coords[idx]\n        slices   = [self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                    for z in self.z_list]\n        return (np.stack(slices, axis=-1),\n                self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy())\n\n    def __getitem__(self, idx):\n        img, msk = self.get_patch(idx)\n        if self.tf:\n            out = self.tf(image=img, mask=msk)\n            img, msk = out['image'], out['mask']\n        if self.ch_drop:\n            img = channel_dropout(img)\n        return (torch.from_numpy(img).permute(2, 0, 1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\n\n# ════════════════════════════════════════════════════════════\n#  10.  BUILD DATASETS  75/25 split\n# ════════════════════════════════════════════════════════════\nprint('\\n── Loading Fragment 2 ──')\nds2 = VesuviusDataset(FRAG2, Z_SLICES, STRIDE_TR,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\nprint('\\n── Loading Fragment 3 ──')\nds3 = VesuviusDataset(FRAG3, Z_SLICES, STRIDE_TR,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\n\nn_total   = len(ds2) + len(ds3)\nrng       = np.random.RandomState(SEED)\nall_idx   = rng.permutation(n_total)\nn_val     = int(n_total * VAL_SPLIT)\nn_train   = n_total - n_val\ntrain_idx = all_idx[:n_train].tolist()\nval_idx   = all_idx[n_train:].tolist()\nprint(f'\\nTotal: {n_total} | Train: {n_train} (75%) | Val: {n_val} (25%)')\n\n\nclass ValSubset(Dataset):\n    def __init__(self, ds2, ds3, indices):\n        self.ds2 = ds2; self.ds3 = ds3\n        self.n2  = len(ds2); self.indices = indices\n    def __len__(self): return len(self.indices)\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        img_np, msk_np = (self.ds2.get_patch(g) if g < self.n2\n                          else self.ds3.get_patch(g - self.n2))\n        return (torch.from_numpy(img_np).permute(2, 0, 1).float(),\n                torch.from_numpy(msk_np).unsqueeze(0).float())\n\n\nclass TrainSubset(Dataset):\n    def __init__(self, ds2, ds3, indices):\n        self.ds2 = ds2; self.ds3 = ds3\n        self.n2  = len(ds2); self.indices = indices\n    def __len__(self): return len(self.indices)\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        return self.ds2[g] if g < self.n2 else self.ds3[g - self.n2]\n\n\ntrain_ds = TrainSubset(ds2, ds3, train_idx)\nval_ds   = ValSubset  (ds2, ds3, val_idx)\n\nall_weights = np.concatenate([ds2.weights, ds3.weights])\ntrain_w     = torch.from_numpy(all_weights[train_idx]).float()\n\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\nsampler  = WeightedRandomSampler(train_w, len(train_ds), replacement=True)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                      num_workers=NUM_WORKERS, pin_memory=False)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=NUM_WORKERS, pin_memory=False)\nprint(f'Train batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  11.  METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr: float = 0.5) -> float:\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p * masks).sum(dim=(1, 2, 3))\n    union = p.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3))\n    return ((2.0 * inter + 1e-5) / (union + 1e-5)).mean().item()\n\n\ndef sweep_threshold(probs: np.ndarray,\n                    targets: np.ndarray) -> tuple:\n    best_t, best_d = 0.5, 0.0\n    for t in np.arange(0.20, 0.80, 0.02):\n        p = (probs > t).astype(np.float32)\n        d = (2.0 * (p * targets).sum() + 1.0) / (p.sum() + targets.sum() + 1.0)\n        if d > best_d:\n            best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  12.  MODEL + OPTIMIZER\n# ════════════════════════════════════════════════════════════\nmodel    = VesuviusV11(n_ch=N_CH, num_heads=8).to(DEVICE)\nn_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'\\nModel parameters: {n_params/1e6:.2f} M')\n\n# Separate LR for encoder (fine-tune) vs decoder (train from scratch)\nenc_params = list(model.encoder.parameters())\ndec_params = [p for p in model.parameters()\n              if not any(p is ep for ep in enc_params)]\n\noptimizer = optim.AdamW([\n    {'params': enc_params, 'lr': LR * 0.1},   # encoder: 10× lower\n    {'params': dec_params, 'lr': LR},\n], weight_decay=WEIGHT_DECAY)\n\n# Cosine annealing with warm restarts\nscheduler = optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr         = [LR * 0.1, LR],\n    steps_per_epoch= len(train_dl),\n    epochs         = EPOCHS,\n    pct_start      = 0.10,          # 10% warm-up\n    div_factor     = 10.0,\n    final_div_factor= 100.0,\n)\n\nscaler    = GradScaler(enabled=(DEVICE == 'cuda'))\nbest_dice = 0.0\npat_cnt   = 0\nhistory   = dict(tl=[], vl=[], td=[], vd=[], lr_enc=[], lr_dec=[])\nWARMUP_EPOCHS = 3   # no CutMix during warm-up\n\nprint('\\n' + '=' * 65)\nprint('VesuviusV11 — Key fixes from v10 diagnosis')\nprint(f'  Transformer  : Axial at 1/8 decoder (28×28 tokens)')\nprint(f'  Loss         : BCE+Focal+Tversky(α=0.6,β=0.4) — FP penalty')\nprint(f'  Temperature  : {TEMPERATURE} (no scaling distortion)')\nprint(f'  CutMix       : only after epoch {WARMUP_EPOCHS}, p=0.20')\nprint(f'  Sampler      : 4:1 pos/neg weight (was 3:1)')\nprint(f'  LR schedule  : OneCycleLR, enc×0.1')\nprint(f'  Split        : 75/25  |  Patience: {PATIENCE}')\nprint('=' * 65)\n\n\n# ════════════════════════════════════════════════════════════\n#  13.  TRAINING LOOP\n# ════════════════════════════════════════════════════════════\nfor epoch in range(EPOCHS):\n    use_cutmix = (epoch >= WARMUP_EPOCHS)\n\n    # ── train ──────────────────────────────────────────────\n    model.train()\n    tl = td = 0.0\n    optimizer.zero_grad()\n\n    for step, (imgs, msks) in enumerate(\n            tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        msks = msks.to(DEVICE, non_blocking=True)\n\n        # CutMix at low frequency after warm-up\n        if use_cutmix and random.random() < 0.20:\n            imgs, msks = cutmix_batch(imgs, msks, alpha=0.4)\n\n        with autocast(enabled=(DEVICE == 'cuda')):\n            logit, sal, ds1, ds2_l = model(imgs)\n            loss = full_loss(logit, sal, ds1, ds2_l, msks) / GRAD_ACCUM\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % GRAD_ACCUM == 0 or (step + 1) == len(train_dl):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n\n        tl += loss.item() * GRAD_ACCUM\n        td += batch_dice(logit.detach(), msks)\n\n        del imgs, msks, logit, sal, ds1, ds2_l, loss\n        if step % 60 == 0:\n            torch.cuda.empty_cache()\n\n    tl /= len(train_dl)\n    td /= len(train_dl)\n\n    # ── validate ────────────────────────────────────────────\n    model.eval()\n    vl = vd = 0.0\n    acc_p, acc_m = [], []\n\n    with torch.no_grad():\n        for imgs, msks in tqdm(val_dl, desc=f'Ep{epoch+1:02d} val  ',\n                               leave=False):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, sal, ds1, ds2_l = model(imgs)\n                loss = full_loss(logit, sal, ds1, ds2_l, msks)\n            vl += loss.item()\n            vd += batch_dice(logit, msks)\n            # NO temperature scaling\n            acc_p.append(torch.sigmoid(logit).cpu().numpy())\n            acc_m.append(msks.cpu().numpy())\n            del imgs, msks, logit, sal, ds1, ds2_l, loss\n\n    vl /= len(val_dl)\n    vd /= len(val_dl)\n\n    probs_all = np.concatenate(acc_p)\n    masks_all = np.concatenate(acc_m)\n    bt, bd    = sweep_threshold(probs_all, masks_all)\n\n    ink_m    = float(probs_all[masks_all > 0.5].mean()) \\\n               if (masks_all > 0.5).any() else 0.0\n    noink_m  = float(probs_all[masks_all < 0.5].mean()) \\\n               if (masks_all < 0.5).any() else 0.0\n    gap      = bd - vd   # threshold sweep gain; should be small\n\n    del acc_p, acc_m, probs_all, masks_all\n    gc.collect(); torch.cuda.empty_cache()\n\n    lr_enc = optimizer.param_groups[0]['lr']\n    lr_dec = optimizer.param_groups[1]['lr']\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['lr_enc'].append(lr_enc); history['lr_dec'].append(lr_dec)\n\n    print(f'Ep{epoch+1:02d} | enc_lr={lr_enc:.1e} dec_lr={lr_dec:.1e} | '\n          f'Tr loss={tl:.4f} dice={td:.4f} | '\n          f'Val loss={vl:.4f} dice={vd:.4f} | '\n          f'opt_thr={bt:.2f}→{bd:.4f} (gap={gap:+.3f}) | '\n          f'sep={ink_m-noink_m:+.3f}')\n\n    # Save by val_dice (not penalised — gap is monitored separately)\n    # threshold penalty still applied to discourage extreme thresholds\n    thr_pen    = max(0.0, bt - 0.62) * 0.25\n    save_met   = max(vd, bd) - thr_pen\n\n    if save_met > best_dice:\n        best_dice = save_met; pat_cnt = 0\n        torch.save({\n            'epoch': epoch,\n            'state': model.state_dict(),\n            'thr'  : bt,\n            'metric': save_met,\n            'n_ch' : N_CH,\n            'bd'   : bd,\n            'vd'   : vd,\n        }, OUTPUT + 'best_model.pth')\n        print(f'  ✓ saved  metric={save_met:.4f}  '\n              f'val_dice={vd:.4f}  opt_dice={bd:.4f}  thr={bt:.2f}')\n    else:\n        pat_cnt += 1\n        if pat_cnt >= PATIENCE:\n            print(f'  ⚑ early stop at epoch {epoch+1}')\n            break\n\nprint(f'\\nBest metric: {best_dice:.4f}')\n\n# ── curves ──────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 4))\naxes[0].plot(history['tl'], label='train')\naxes[0].plot(history['vl'], label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\n\naxes[1].plot(history['td'], label='train')\naxes[1].plot(history['vd'], label='val')\naxes[1].axhline(0.80, color='r', ls='--', label='target 0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\n\naxes[2].plot(history['lr_enc'], label='encoder')\naxes[2].plot(history['lr_dec'], label='decoder')\naxes[2].set_title('Learning Rate')\naxes[2].legend(); axes[2].set_xlabel('Epoch'); axes[2].grid(True)\n\nplt.tight_layout()\nplt.savefig(OUTPUT + 'curves.png', dpi=100); plt.close()\nprint('Curves saved.')\n\n\n# ════════════════════════════════════════════════════════════\n#  14.  PSEUDO-LABELLING  (optional)\n# ════════════════════════════════════════════════════════════\ndef _infer_fragment_probs(model, frag_path, z_list, temperature=1.0):\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list, verbose=False)\n    mid_z = z_list[len(z_list) // 2]\n    H, W  = cache[mid_z].shape\n    pred  = np.zeros((H, W), np.float32)\n    wgt   = np.zeros((H, W), np.float32)\n    gw    = _gauss_weight(PATCH_SIZE)\n    coords = [(y, x)\n              for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n              for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='PseudoInfer', leave=False):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            t = torch.from_numpy(\n                np.stack(slices, axis=-1)).permute(2, 0, 1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit / temperature).squeeze().cpu().numpy()\n            pred[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p  * gw\n            wgt [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += gw\n            del t, logit\n    del cache; gc.collect(); torch.cuda.empty_cache()\n    return pred / (wgt + 1e-8)\n\n\nif USE_PSEUDO and PSEUDO_FRAGS:\n    print('\\n── Pseudo-labelling ──')\n    ckpt = torch.load(OUTPUT + 'best_model.pth', map_location=DEVICE)\n    model.load_state_dict(ckpt['state'])\n    pseudo_datasets = []\n    for pf in PSEUDO_FRAGS:\n        prob_p = _infer_fragment_probs(model, pf, Z_SLICES)\n        plbl   = ((prob_p > PSEUDO_THRESHOLD) * 255).astype(np.uint8)\n        ink_pct = 100 * (plbl > 0).mean()\n        print(f'  {pf}: pseudo ink%={ink_pct:.1f}%')\n        if ink_pct < 1.0 or ink_pct > 40.0:\n            print('  [SKIP] ink% out of plausible range')\n            continue\n        out_lbl = os.path.join(pf, 'inklabels.png')\n        cv2.imwrite(out_lbl, plbl)\n        try:\n            pds = VesuviusDataset(pf, Z_SLICES, STRIDE_TR,\n                                  transform=train_tf, neg_ratio=0.25,\n                                  apply_ch_dropout=True)\n            pseudo_datasets.append(pds)\n        except Exception as e:\n            print(f'  [WARN] {e}')\n\n    if pseudo_datasets:\n        from torch.utils.data import ConcatDataset\n        pall = ConcatDataset(pseudo_datasets)\n        pdl  = DataLoader(pall, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=NUM_WORKERS)\n        popt = optim.AdamW(model.parameters(), lr=LR * 0.15,\n                           weight_decay=WEIGHT_DECAY)\n        pscl = GradScaler(enabled=(DEVICE == 'cuda'))\n        for ep in range(6):\n            model.train(); pl = pd_sc = 0.0; popt.zero_grad()\n            for step, (imgs, msks) in enumerate(\n                    tqdm(pdl, desc=f'Pseudo Ep{ep+1}', leave=False)):\n                imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n                with autocast(enabled=(DEVICE == 'cuda')):\n                    logit, sal, d1, d2_l = model(imgs)\n                    loss = full_loss(logit, sal, d1, d2_l, msks) / GRAD_ACCUM\n                pscl.scale(loss).backward()\n                if (step + 1) % GRAD_ACCUM == 0:\n                    pscl.unscale_(popt)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                    pscl.step(popt); pscl.update(); popt.zero_grad()\n                pl   += loss.item() * GRAD_ACCUM\n                pd_sc += batch_dice(logit.detach(), msks)\n                del imgs, msks, logit, sal, d1, d2_l, loss\n            print(f'  Pseudo Ep{ep+1} loss={pl/len(pdl):.4f} '\n                  f'dice={pd_sc/len(pdl):.4f}')\n        torch.save({'state': model.state_dict(), 'n_ch': N_CH,\n                    'thr': ckpt['thr']},\n                   OUTPUT + 'best_model_pseudo.pth')\n        print('Pseudo fine-tune done.')\n\n\n# ════════════════════════════════════════════════════════════\n#  15.  INFERENCE  (NO TTA — clean single-pass)\n# ════════════════════════════════════════════════════════════\ndef _gauss_weight(sz: int) -> np.ndarray:\n    c   = sz // 2; sig = sz // 4\n    ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs-c)**2 + (ys-c)**2) / (2*sig**2)).astype(np.float32)\n\n\nGW = _gauss_weight(PATCH_SIZE)\n\n\ndef predict_fragment(model, frag_path, z_list,\n                     temperature: float = TEMPERATURE):\n    \"\"\"Single-pass sliding-window inference, NO TTA.\"\"\"\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    lbl   = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n    msk   = (lbl > 0).astype(np.uint8)\n    H, W  = msk.shape\n\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n\n    coords = [(y, x)\n              for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n              for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for y, x in tqdm(coords, desc='Inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            patch  = np.stack(slices, axis=-1)\n            t      = torch.from_numpy(patch).permute(2, 0, 1) \\\n                                            .unsqueeze(0).to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit / temperature).squeeze().cpu() \\\n                     .numpy().astype(np.float32)\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p  * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, logit\n        torch.cuda.empty_cache()\n\n    del cache; gc.collect()\n    return pred_map / (wgt_map + 1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  16.  POST-PROCESSING\n#       Morphological cleaning + optional DenseCRF\n# ════════════════════════════════════════════════════════════\ndef morphological_clean(binary_map: np.ndarray,\n                        min_area: int = 150,\n                        close_k: int = 5) -> np.ndarray:\n    kern   = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (close_k, close_k))\n    closed = cv2.morphologyEx(binary_map.astype(np.uint8),\n                              cv2.MORPH_CLOSE, kern)\n    n_lab, labels, stats, _ = cv2.connectedComponentsWithStats(closed)\n    out = np.zeros_like(closed)\n    for i in range(1, n_lab):\n        if stats[i, cv2.CC_STAT_AREA] >= min_area:\n            out[labels == i] = 1\n    return out\n\n\ndef apply_dense_crf(image_u8: np.ndarray, prob_map: np.ndarray,\n                    n_iter: int = 5) -> np.ndarray:\n    if not HAS_CRF:\n        return (prob_map > 0.5).astype(np.uint8)\n    H, W = prob_map.shape\n    d    = dcrf.DenseCRF2D(W, H, 2)\n    fg   = np.clip(prob_map,       1e-5, 1 - 1e-5)\n    bg   = np.clip(1.0 - prob_map, 1e-5, 1 - 1e-5)\n    U    = -np.log(np.stack([bg, fg], 0)).reshape(2, -1).astype(np.float32)\n    d.setUnaryEnergy(U)\n    d.addPairwiseGaussian(sxy=3, compat=3)\n    img_c = np.ascontiguousarray(image_u8)\n    d.addPairwiseBilateral(sxy=50, srgb=13, rgbim=img_c, compat=10)\n    Q = d.inference(n_iter)\n    return np.argmax(Q, 0).reshape(H, W).astype(np.uint8)\n\n\ndef postprocess(prob_map, cache, z_list, threshold,\n                min_area=150, close_k=5, use_crf=HAS_CRF):\n    binary = (prob_map > threshold).astype(np.uint8)\n    binary = morphological_clean(binary, min_area, close_k)\n    if use_crf:\n        mid_z  = z_list[len(z_list) // 2]\n        mid_sl = cache[mid_z].astype(np.float32)\n        mn, mx = mid_sl.min(), mid_sl.max()\n        mid_8  = ((mid_sl - mn) / (mx - mn + 1e-8) * 255).astype(np.uint8)\n        rgb    = np.stack([mid_8]*3, axis=-1)\n        binary = apply_dense_crf(rgb, prob_map)\n    return binary\n\n\n# ════════════════════════════════════════════════════════════\n#  17.  FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '=' * 65)\nprint('FINAL TEST — FRAGMENT 1  (unseen during training)')\nprint('Single-pass inference  |  NO TTA')\nprint('=' * 65)\n\nckpt = torch.load(OUTPUT + 'best_model.pth', map_location=DEVICE)\nassert ckpt.get('n_ch', N_CH) == N_CH, \\\n    f\"Channel mismatch ckpt={ckpt.get('n_ch')} vs {N_CH}\"\n\nmodel = VesuviusV11(n_ch=N_CH, num_heads=8).to(DEVICE)\nmodel.load_state_dict(ckpt['state'], strict=True)\nsaved_thr = ckpt.get('thr', 0.5)\nprint(f'Checkpoint: epoch {ckpt[\"epoch\"]+1} | '\n      f'val_dice={ckpt.get(\"vd\",0):.4f} | '\n      f'opt_dice={ckpt.get(\"bd\",0):.4f} | '\n      f'thr={saved_thr:.2f}')\n\nprob_map, msk1 = predict_fragment(model, FRAG1, Z_SLICES)\nH_m, W_m       = msk1.shape\nprob_crop       = prob_map[:H_m, :W_m]\n\n# Calibration stats\nink_mean   = float(prob_crop[msk1 == 1].mean())\nnoink_mean = float(prob_crop[msk1 == 0].mean())\nprint(f'\\nCalibration (Fragment 1):')\nprint(f'  ink pixels  : {ink_mean:.3f}')\nprint(f'  bg  pixels  : {noink_mean:.3f}')\nprint(f'  separation  : {ink_mean - noink_mean:+.3f}')\n\n# Optimal threshold on Fragment 1\nbt1, bd1 = sweep_threshold(\n    prob_crop[np.newaxis, np.newaxis],\n    msk1    [np.newaxis, np.newaxis]\n)\nprint(f'  Best thr    : {bt1:.2f}  dice={bd1:.4f}')\n\n# Post-processing\nprint(f'\\nPost-processing: morph_clean + CRF={HAS_CRF}')\ncache1     = load_slice_cache(FRAG1, Z_SLICES)\nfinal_pred = postprocess(prob_crop, cache1, Z_SLICES,\n                         threshold=bt1, min_area=150, close_k=5)\ndel cache1; gc.collect()\n\n# ── Final metrics ────────────────────────────────────────────\npf = final_pred.flatten().astype(int)\nmf = msk1.flatten().astype(int)\ntn, fp, fn, tp_v = confusion_matrix(mf, pf, labels=[0, 1]).ravel()\nprec      = tp_v / (tp_v + fp  + 1e-8)\nrec       = tp_v / (tp_v + fn  + 1e-8)\nf1        = 2 * prec * rec / (prec + rec + 1e-8)\ndice_full = (2 * tp_v + 1) / (final_pred.sum() + msk1.sum() + 1)\n\nprint('\\n' + '=' * 65)\nprint('RESULTS — FRAGMENT 1  (v11)')\nprint('=' * 65)\nprint(f'Dice Score  : {dice_full:.4f}')\nprint(f'F1 Score    : {f1:.4f}')\nprint(f'Precision   : {prec:.4f}')\nprint(f'Recall      : {rec:.4f}')\nprint(f'Threshold   : {bt1:.2f}')\nprint(f'TP={tp_v} | TN={tn} | FP={fp} | FN={fn}')\nprint(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}  (target < 1.0)')\nprint(f'Ink/bg sep  : {ink_mean-noink_mean:+.3f}  (target > 0.35)')\ntarget_str = '✓ PASSED' if f1 >= 0.80 else '✗ Below 0.80 target'\nprint(f'F1 ≥ 0.80   : {target_str}')\nprint('=' * 65)\n\n# ── Visualisation ────────────────────────────────────────────\nfig, ax = plt.subplots(2, 3, figsize=(18, 12))\n\nax[0,0].imshow(msk1,       cmap='gray'); ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop,  cmap='inferno'); ax[0,1].set_title('Probability Map')\nax[0,2].imshow(final_pred, cmap='gray')\nax[0,2].set_title(f'Prediction  Dice={dice_full:.3f}  F1={f1:.3f}')\n\nerr = np.zeros((*msk1.shape, 3), dtype=np.uint8)\nerr[(final_pred==1) & (msk1==1)] = [0,   255,   0]\nerr[(final_pred==1) & (msk1==0)] = [255,   0,   0]\nerr[(final_pred==0) & (msk1==1)] = [0,     0, 255]\nax[1,0].imshow(err)\nax[1,0].set_title('TP=green  FP=red  FN=blue')\n\nax[1,1].hist(prob_crop[msk1==1].ravel(), bins=60, alpha=0.7,\n             label=f'ink (μ={ink_mean:.2f})',       color='orange', density=True)\nax[1,1].hist(prob_crop[msk1==0].ravel(), bins=60, alpha=0.7,\n             label=f'no-ink (μ={noink_mean:.2f})',  color='steelblue', density=True)\nax[1,1].axvline(bt1, color='r', ls='--', label=f'thr={bt1:.2f}')\nax[1,1].set_title('Probability Distribution')\nax[1,1].legend(); ax[1,1].set_xlabel('Probability')\n\nts = np.arange(0.15, 0.90, 0.01); ds_vals = []\nfor t in ts:\n    pb = (prob_crop > t).astype(np.float32)\n    ds_vals.append((2*(pb*msk1).sum()+1)/(pb.sum()+msk1.sum()+1))\nax[1,2].plot(ts, ds_vals)\nax[1,2].axvline(bt1, color='r', ls='--', label=f'best={bt1:.2f}')\nax[1,2].set_xlabel('Threshold'); ax[1,2].set_ylabel('Dice')\nax[1,2].set_title('Dice vs Threshold Curve')\nax[1,2].legend(); ax[1,2].grid(True)\n\nfor a in [ax[0,0], ax[0,1], ax[0,2], ax[1,0]]: a.axis('off')\nplt.suptitle(\n    f'Fragment 1 — Dice={dice_full:.4f}  F1={f1:.4f}  '\n    f'Prec={prec:.4f}  Rec={rec:.4f}',\n    fontsize=13, y=1.01\n)\nplt.tight_layout()\nplt.savefig(OUTPUT + 'frag1_prediction.png', dpi=100, bbox_inches='tight')\nplt.close()\n\n# ── Save all outputs ─────────────────────────────────────────\ncv2.imwrite(OUTPUT + 'frag1_prob_map.png',\n            (prob_crop * 255).astype(np.uint8))\ncv2.imwrite(OUTPUT + 'frag1_prediction_binary.png',\n            (final_pred * 255).astype(np.uint8))\n\nwith open(OUTPUT + 'final_results.txt', 'w') as f:\n    f.write('=' * 50 + '\\n')\n    f.write('VESUVIUS V11 — FINAL TEST RESULTS\\n')\n    f.write('=' * 50 + '\\n')\n    f.write(f'Dice        : {dice_full:.4f}\\n')\n    f.write(f'F1          : {f1:.4f}\\n')\n    f.write(f'Precision   : {prec:.4f}\\n')\n    f.write(f'Recall      : {rec:.4f}\\n')\n    f.write(f'Threshold   : {bt1:.2f}\\n')\n    f.write(f'TP={tp_v} TN={tn} FP={fp} FN={fn}\\n')\n    f.write(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Ink/bg sep  : {ink_mean-noink_mean:+.3f}\\n')\n    f.write(f'Temperature : {TEMPERATURE}\\n')\n    f.write(f'Post-proc   : morph_clean + DenseCRF={HAS_CRF}\\n')\n    f.write(f'Architecture: ResNet34 + AxialTransformer@1/8\\n')\n    f.write(f'Z-slices    : {Z_SLICES}\\n')\n    f.write('=' * 50 + '\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  best_model.pth')\nprint('  curves.png')\nprint('  frag1_prediction.png   (6-panel visual)')\nprint('  frag1_prob_map.png')\nprint('  frag1_prediction_binary.png')\nprint('  final_results.txt')\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n# ============================================================\n#  VESUVIUS INK DETECTION — v10\n#  PhD Thesis: 3D Axial-Attention Transformer on CT Volume\n#\n#  Key Innovations vs v9:\n#  1.  3D Axial-Attention Transformer (X/Y/Z axes separately)\n#      with Token Merging on low-saliency regions\n#  2.  Sparse attention guided by coarse \"sheet-surface\" module\n#  3.  Log-polar relative positional embeddings (scroll geometry)\n#  4.  Heavy augmentation: elastic deform + CutMix ink/no-ink\n#  5.  Pseudo-labelling on unlabelled fragments\n#  6.  Post-processing: morphological cleaning + DenseCRF\n#  7.  NO TTA on test (removed — degrades perf on wrong preds)\n#  8.  ResNet34 2D backbone for final patch classification head\n#  9.  Deep supervision at 3 decoder scales\n#  10. Fragment-aware percentile normalisation (v9 style)\n# ============================================================\n\nimport os, gc, cv2, math, warnings, random\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\n# optional – DenseCRF post-processing\ntry:\n    import pydensecrf.densecrf as dcrf\n    from pydensecrf.utils import unary_from_labels, create_pairwise_bilateral, \\\n        create_pairwise_gaussian\n    HAS_CRF = True\nexcept ImportError:\n    HAS_CRF = False\n    print(\"[INFO] pydensecrf not found – CRF post-processing will be skipped. \"\n          \"Install with:  pip install pydensecrf\")\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\n\n# ──────────────────────────────────────────────────────────────\n#  SEED\n# ──────────────────────────────────────────────────────────────\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ──────────────────────────────────────────────────────────────\n#  PATHS\n# ──────────────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ──────────────────────────────────────────────────────────────\n#  HYPER-PARAMS\n# ──────────────────────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 224\nSTRIDE_TR    = 112\nSTRIDE_INF   = 56\nBATCH_SIZE   = 6           # reduced for 3D attention blocks\nGRAD_ACCUM   = 4           # effective batch = 16\nEPOCHS       = 10\nLR           = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 10\nNUM_WORKERS  = 0\nVAL_SPLIT    = 0.30\nDROPOUT_P    = 0.3\nTEMPERATURE  = 1.3         # softer than v9\n\n# Z-slices: 17 central slices (same as v9)\nZ_SLICES  = [20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36]\nN_CH      = len(Z_SLICES)   # 17\n\nINK_MIN_POS  = 0.02\nNEG_RATIO    = 0.5\nCH_DROP_PROB = 0.3\nCH_DROP_MAX  = 3\n\n# Pseudo-label config\nPSEUDO_THRESHOLD = 0.70     # confidence to accept pseudo label\nUSE_PSEUDO       = False    # set True if you have extra unlabelled fragments\nPSEUDO_FRAGS     = []       # e.g. ['/kaggle/input/vesuvius/extra/4']\n\nprint(f\"Device : {DEVICE}  |  Channels: {N_CH}  |  Patch: {PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1.  SLICE CACHE  (fragment-wise percentile normalisation)\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    if not raw:\n        raise ValueError(f\"No slices in {vol_dir}\")\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Slice cache: {len(cache)} slices ({mb:.0f} MB)\")\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2.  LOG-POLAR RELATIVE POSITIONAL EMBEDDING\n#      Respects cylindrical scroll geometry:\n#        r   = sqrt(dx² + dy²)  → log-scaled radial distance\n#        phi = atan2(dy, dx)    → angular displacement\n#        dz  = abs(z1 - z2)     → depth separation\n# ════════════════════════════════════════════════════════════\ndef build_logpolar_bias(seq_len_h, seq_len_w, num_heads, device='cpu'):\n    \"\"\"\n    Returns additive attention bias [num_heads, HW, HW] that encodes\n    log-polar relative positions between every pair of spatial tokens.\n    Computed once per unique (seq_len_h, seq_len_w) shape.\n    \"\"\"\n    H, W = seq_len_h, seq_len_w\n    # grid of (row, col) for every token\n    ys = torch.arange(H, dtype=torch.float32)\n    xs = torch.arange(W, dtype=torch.float32)\n    gy, gx = torch.meshgrid(ys, xs, indexing='ij')  # [H, W]\n    gy = gy.reshape(-1); gx = gx.reshape(-1)        # [HW]\n\n    dy = gy.unsqueeze(1) - gy.unsqueeze(0)           # [HW, HW]\n    dx = gx.unsqueeze(1) - gx.unsqueeze(0)\n\n    r   = torch.sqrt(dx**2 + dy**2 + 1e-3)           # radial dist\n    log_r = torch.log(r + 1.0)                       # log scale\n    phi = torch.atan2(dy, dx)                        # [-pi, pi]\n\n    # Encode into num_heads channels via learnable-free Fourier projection\n    freqs = torch.arange(1, num_heads // 2 + 1, dtype=torch.float32)\n    bias_r   = torch.cos(freqs[None, None, :] * log_r.unsqueeze(2))   # [HW,HW,H/2]\n    bias_phi = torch.sin(freqs[None, None, :] * phi.unsqueeze(2))     # [HW,HW,H/2]\n    bias = torch.cat([bias_r, bias_phi], dim=-1)    # [HW, HW, num_heads]\n    bias = bias.permute(2, 0, 1)                    # [num_heads, HW, HW]\n    return bias.to(device)\n\n\n# Cache to avoid recomputing for the same shape\n_LP_CACHE: dict = {}\n\ndef get_logpolar_bias(H, W, num_heads, device):\n    key = (H, W, num_heads, str(device))\n    if key not in _LP_CACHE:\n        _LP_CACHE[key] = build_logpolar_bias(H, W, num_heads, device)\n    return _LP_CACHE[key]\n\n\n# ════════════════════════════════════════════════════════════\n#  3.  COARSE SHEET-SURFACE DETECTION MODULE\n#      Lightweight 2-layer CNN that predicts a per-pixel\n#      \"ink-likelihood\" score from the raw 17-channel patch.\n#      Its output guides token merging: low-score tokens are\n#      merged (averaged), reducing sequence length for the\n#      expensive axial-attention layers.\n# ════════════════════════════════════════════════════════════\nclass SheetSurfaceDetector(nn.Module):\n    \"\"\"Fast coarse detector → saliency map [B, 1, H, W].\"\"\"\n    def __init__(self, in_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, 32, 3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 16, 3, padding=1, bias=False),\n            nn.BatchNorm2d(16), nn.ReLU(inplace=True),\n            nn.Conv2d(16,  1, 1),\n        )\n\n    def forward(self, x):   # x: [B, C, H, W]\n        return self.net(x)  # [B, 1, H, W]  (raw logit)\n\n\n# ════════════════════════════════════════════════════════════\n#  4.  TOKEN MERGING  (ToMe-inspired, sparse ink regions)\n#      Tokens whose coarse saliency < threshold are merged\n#      with their nearest neighbour → shorter sequence for\n#      the transformer.  We un-merge after attention so\n#      spatial resolution is fully restored for the decoder.\n# ════════════════════════════════════════════════════════════\ndef token_merge(tokens, saliency, merge_ratio=0.30):\n    \"\"\"\n    tokens   : [B, C, N]   (N = H*W tokens)\n    saliency : [B, 1, N]   (coarse ink score, already sigmoid)\n    Returns  : merged_tokens [B, C, M],  unmerge_idx [B, N] (maps M→N)\n    \"\"\"\n    B, C, N = tokens.shape\n    sal = saliency.squeeze(1)           # [B, N]\n\n    # sort tokens by saliency; low-saliency ones are merge candidates\n    _, sort_idx = sal.sort(dim=1)       # ascending → low first\n    n_merge = int(N * merge_ratio)\n\n    merge_idx  = sort_idx[:, :n_merge]   # [B, n_merge]\n    keep_idx   = sort_idx[:, n_merge:]   # [B, N-n_merge]\n\n    # gather merge & keep tokens\n    def gather(t, idx):\n        idx_exp = idx.unsqueeze(1).expand(-1, C, -1)\n        return t.gather(2, idx_exp)\n\n    t_merge = gather(tokens, merge_idx)   # [B, C, n_merge]\n    t_keep  = gather(tokens, keep_idx)    # [B, C, N-n_merge]\n\n    # pair-wise merge: consecutive pairs averaged\n    if n_merge % 2 == 1:\n        # drop last odd token into keep\n        t_keep  = torch.cat([t_keep, t_merge[:, :, -1:]], dim=2)\n        t_merge = t_merge[:, :, :-1]\n        n_merge -= 1\n\n    t_merged = (t_merge[:, :, 0::2] + t_merge[:, :, 1::2]) / 2  # [B,C,n_merge/2]\n    tokens_out = torch.cat([t_keep, t_merged], dim=2)            # [B, C, M]\n\n    # book-keeping for un-merge\n    unmerge_info = (keep_idx, merge_idx, n_merge, N)\n    return tokens_out, unmerge_info\n\n\ndef token_unmerge(tokens_out, unmerge_info, C):\n    \"\"\"Restore [B, C, N] from merged [B, C, M].\"\"\"\n    keep_idx, merge_idx, n_merge, N = unmerge_info\n    B = tokens_out.shape[0]\n    M = tokens_out.shape[2]\n    n_keep = M - n_merge // 2\n\n    t_keep   = tokens_out[:, :, :n_keep]               # [B,C,n_keep]\n    t_merged = tokens_out[:, :, n_keep:]               # [B,C,n_merge/2]\n\n    # expand merged pairs back\n    t_expanded = t_merged.repeat_interleave(2, dim=2)   # [B,C,n_merge]\n    if n_merge % 2 == 1:\n        # we added an extra token to keep above; undo\n        extra = t_keep[:, :, -1:]\n        t_keep = t_keep[:, :, :-1]\n        t_expanded = torch.cat([t_expanded, extra], dim=2)\n        n_merge += 1\n\n    # scatter back to original positions\n    device = tokens_out.device\n    out = torch.zeros(B, C, N, device=device, dtype=tokens_out.dtype)\n\n    def scatter(t, idx):\n        idx_exp = idx.unsqueeze(1).expand(-1, C, -1)\n        out.scatter_(2, idx_exp, t)\n\n    scatter(t_keep,    keep_idx)\n    scatter(t_expanded, merge_idx)\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  5.  AXIAL MULTI-HEAD SELF-ATTENTION  (X / Y axes)\n#      Each axis attends along one spatial dimension with\n#      log-polar positional bias.  We alternate X→Y.\n# ════════════════════════════════════════════════════════════\nclass AxialAttention(nn.Module):\n    \"\"\"\n    Attention along one axis (row or column) of a 2-D feature map.\n    x : [B, C, H, W]\n    \"\"\"\n    def __init__(self, dim, num_heads=8, axis='x', dropout=0.1):\n        super().__init__()\n        assert axis in ('x', 'y')\n        self.axis      = axis\n        self.num_heads = num_heads\n        self.head_dim  = dim // num_heads\n        self.scale     = self.head_dim ** -0.5\n\n        self.qkv  = nn.Linear(dim, dim * 3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.drop = nn.Dropout(dropout)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        if self.axis == 'x':\n            # attend along width for each row independently\n            x_r = x.permute(0, 2, 3, 1)          # [B, H, W, C]\n            shape = (B * H, W, C)\n        else:\n            x_r = x.permute(0, 3, 2, 1)          # [B, W, H, C]\n            shape = (B * W, H, C)\n\n        x_r = x_r.reshape(*shape)                # [B*H, W, C] or [B*W, H, C]\n        res  = x_r\n        x_n  = self.norm(x_r)\n\n        BN, L, _ = x_n.shape\n        qkv = self.qkv(x_n).reshape(BN, L, 3, self.num_heads, self.head_dim)\n        qkv = qkv.permute(2, 0, 3, 1, 4)        # [3, BN, nh, L, hd]\n        q, k, v = qkv.unbind(0)\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale   # [BN, nh, L, L]\n\n        # Add log-polar bias (1-D for axial: just use log distance)\n        pos = torch.arange(L, dtype=torch.float32, device=x.device)\n        d   = (pos.unsqueeze(0) - pos.unsqueeze(1)).abs().float()   # [L, L]\n        log_d_bias = -torch.log(d + 1.0)                            # [L, L]\n        attn = attn + log_d_bias.unsqueeze(0).unsqueeze(0)\n\n        attn = attn.softmax(-1)\n        attn = self.drop(attn)\n        out  = (attn @ v).transpose(1, 2).reshape(BN, L, C)\n        out  = self.proj(out) + res                  # residual\n\n        if self.axis == 'x':\n            out = out.reshape(B, H, W, C).permute(0, 3, 1, 2)\n        else:\n            out = out.reshape(B, W, H, C).permute(0, 3, 2, 1)\n        return out\n\n\n# ════════════════════════════════════════════════════════════\n#  6.  AXIAL TRANSFORMER BLOCK  (X-axis → Y-axis → FFN)\n# ════════════════════════════════════════════════════════════\nclass AxialTransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=8, mlp_ratio=4, dropout=0.1):\n        super().__init__()\n        self.attn_x = AxialAttention(dim, num_heads, axis='x', dropout=dropout)\n        self.attn_y = AxialAttention(dim, num_heads, axis='y', dropout=dropout)\n        self.norm1  = nn.LayerNorm(dim)\n        self.norm2  = nn.LayerNorm(dim)\n        mlp_dim     = int(dim * mlp_ratio)\n        self.ffn = nn.Sequential(\n            nn.Linear(dim, mlp_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(mlp_dim, dim),\n            nn.Dropout(dropout),\n        )\n\n    def forward(self, x):\n        # x: [B, C, H, W]\n        x = self.attn_x(x)\n        x = self.attn_y(x)\n        # FFN on channel dim\n        B, C, H, W = x.shape\n        xf = x.permute(0, 2, 3, 1).reshape(-1, C)\n        xf = self.ffn(self.norm2(xf))\n        x  = x + xf.reshape(B, H, W, C).permute(0, 3, 1, 2)\n        return x\n\n\n# ════════════════════════════════════════════════════════════\n#  7.  FULL MODEL\n#      Architecture:\n#        ResNet34 CNN encoder  (17-ch in)\n#             ↓  skip features at 4 scales\n#        AxialTransformerBlocks on bottleneck  (with token merging)\n#             ↓\n#        FPN-style decoder  (deep supervision at 3 scales)\n#             ↓\n#        SheetSurfaceDetector auxiliary head\n# ════════════════════════════════════════════════════════════\nclass VesuviusV10(nn.Module):\n    def __init__(self, n_ch=N_CH, enc_dim=512,\n                 n_transformer_blocks=4, num_heads=8):\n        super().__init__()\n\n        # ── ResNet34 UNet backbone ────────────────────────────\n        self.backbone = smp.Unet(\n            encoder_name          = 'resnet34',\n            encoder_weights       = 'imagenet',\n            in_channels           = n_ch,\n            classes               = 1,\n            decoder_attention_type= 'scse',\n        )\n\n        # We intercept the encoder bottleneck output and replace\n        # the UNet decoder with our own axial-transformer decoder.\n        # Bottleneck of ResNet34 is 512 channels at 1/32 spatial.\n        self.transformer_blocks = nn.Sequential(*[\n            AxialTransformerBlock(enc_dim, num_heads=num_heads,\n                                  mlp_ratio=4, dropout=0.1)\n            for _ in range(n_transformer_blocks)\n        ])\n\n        # ── Sheet-surface detector (auxiliary head on raw input) ──\n        self.sheet_detector = SheetSurfaceDetector(n_ch)\n\n        # ── Decoder (same as backbone's but we replace the head) ──\n        # We will use the backbone's decoder directly; the transformer\n        # operates on the bottleneck feature before decoder sees it.\n        self._enc_dim = enc_dim\n\n        # Projection to match residual after transformer\n        self.bottleneck_proj = nn.Sequential(\n            nn.Conv2d(enc_dim, enc_dim, 1, bias=False),\n            nn.BatchNorm2d(enc_dim),\n            nn.GELU(),\n        )\n\n        # Deep supervision heads\n        # These attach to intermediate decoder stages\n        # ResNet34 decoder feature sizes: 256, 128, 64, 32, 16\n        self.ds_head3 = nn.Conv2d(256, 1, 1)   # after decode block 0\n        self.ds_head2 = nn.Conv2d(128, 1, 1)   # after decode block 1\n        self.ds_head1 = nn.Conv2d( 64, 1, 1)   # after decode block 2\n\n    def forward(self, x):\n        B = x.shape[0]\n\n        # ── Coarse sheet-surface saliency ──────────────────────\n        saliency_logit = self.sheet_detector(x)   # [B,1,H,W]\n\n        # ── Encoder ────────────────────────────────────────────\n        # Use backbone's encoder forward pass\n        feats = self.backbone.encoder(x)\n        # feats is a list: [input, s1, s2, s3, s4, bottleneck]\n        bottleneck = feats[-1]   # [B, 512, H/32, W/32]\n\n        # ── Axial-Transformer on bottleneck ────────────────────\n        bH, bW = bottleneck.shape[2], bottleneck.shape[3]\n        N      = bH * bW\n\n        # token merging guided by down-sampled saliency\n        sal_down = F.adaptive_avg_pool2d(\n            torch.sigmoid(saliency_logit), (bH, bW)\n        )                                         # [B,1,bH,bW]\n        sal_flat = sal_down.reshape(B, 1, N)      # [B,1,N]\n        tok      = bottleneck.reshape(B, self._enc_dim, N)  # [B,C,N]\n\n        tok_merged, unmerge_info = token_merge(tok, sal_flat, merge_ratio=0.25)\n\n        # reshape merged tokens to pseudo-spatial for axial attention\n        # approximate square layout\n        M       = tok_merged.shape[2]\n        sq      = int(math.ceil(math.sqrt(M)))\n        pad_len = sq * sq - M\n        if pad_len > 0:\n            tok_merged = F.pad(tok_merged, (0, pad_len))\n        tok_2d = tok_merged.reshape(B, self._enc_dim, sq, sq)  # [B,C,sq,sq]\n\n        tok_2d = self.transformer_blocks(tok_2d)               # axial attention\n        tok_flat = tok_2d.reshape(B, self._enc_dim, sq * sq)[:, :, :M]\n\n        # un-merge back to full spatial\n        tok_full = token_unmerge(tok_flat, unmerge_info, self._enc_dim)\n        bottleneck_out = tok_full.reshape(B, self._enc_dim, bH, bW)\n        bottleneck_out = self.bottleneck_proj(bottleneck_out + bottleneck)\n\n        # replace bottleneck in feats\n        feats_mod = list(feats)\n        feats_mod[-1] = bottleneck_out\n\n        # ── Decoder ────────────────────────────────────────────\n        # Use backbone decoder – it accepts the feature list\n        decoder_output = self.backbone.decoder(*feats_mod)\n        main_logit     = self.backbone.segmentation_head(decoder_output)\n\n        # Deep supervision: tap intermediate decoder layers\n        # smp UNet decoder blocks are in self.backbone.decoder.blocks\n        ds3_logit = None; ds2_logit = None; ds1_logit = None\n        try:\n            db = self.backbone.decoder.blocks\n            if len(db) >= 1:\n                f0 = db[0](feats_mod[-1], feats_mod[-2])\n                ds3_logit = self.ds_head3(f0)\n            if len(db) >= 2:\n                f1 = db[1](f0, feats_mod[-3])\n                ds2_logit = self.ds_head2(f1)\n            if len(db) >= 3:\n                f2 = db[2](f1, feats_mod[-4])\n                ds1_logit = self.ds_head1(f2)\n        except Exception:\n            pass   # deep supervision optional\n\n        return main_logit, saliency_logit, ds3_logit, ds2_logit, ds1_logit\n\n\n# ════════════════════════════════════════════════════════════\n#  8.  AUGMENTATIONS  (heavy – elastic + CutMix)\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.RandomRotate90(p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.12, scale_limit=0.18,\n                       rotate_limit=35,\n                       border_mode=cv2.BORDER_REFLECT, p=0.65),\n    A.ElasticTransform(alpha=1.0, sigma=50, alpha_affine=50,\n                       border_mode=cv2.BORDER_REFLECT, p=0.4),\n    A.GridDistortion(num_steps=5, distort_limit=0.3,\n                     border_mode=cv2.BORDER_REFLECT, p=0.3),\n    A.RandomResizedCrop(height=PATCH_SIZE, width=PATCH_SIZE,\n                        scale=(0.55, 1.0), ratio=(0.85, 1.15), p=0.5),\n    A.RandomBrightnessContrast(0.25, 0.25, p=0.55),\n    A.GaussNoise(var_limit=(0.001, 0.005), p=0.35),\n    A.GaussianBlur(blur_limit=(3, 7), p=0.25),\n    A.CoarseDropout(max_holes=6, max_height=28, max_width=28,\n                    fill_value=0, p=0.35),\n])\n\n\ndef channel_dropout(img_np, prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < prob:\n        n   = np.random.randint(1, max_drop + 1)\n        idx = np.random.choice(img_np.shape[2], n, replace=False)\n        img_np = img_np.copy()\n        img_np[:, :, idx] = 0.0\n    return img_np\n\n\ndef cutmix_batch(imgs, msks, alpha=0.4):\n    \"\"\"\n    CutMix across the batch dimension for ink/no-ink regions.\n    imgs : [B, C, H, W]  msks : [B, 1, H, W]\n    Returns augmented imgs and mixed msks (soft labels ok for loss).\n    \"\"\"\n    B, C, H, W = imgs.shape\n    lam  = np.random.beta(alpha, alpha)\n    perm = torch.randperm(B, device=imgs.device)\n\n    cx  = np.random.randint(W)\n    cy  = np.random.randint(H)\n    bw  = int(W * math.sqrt(1 - lam))\n    bh  = int(H * math.sqrt(1 - lam))\n    x1  = max(0, cx - bw // 2); x2 = min(W, cx + bw // 2)\n    y1  = max(0, cy - bh // 2); y2 = min(H, cy + bh // 2)\n\n    imgs_new       = imgs.clone()\n    msks_new       = msks.clone()\n    imgs_new[:, :, y1:y2, x1:x2] = imgs[perm, :, y1:y2, x1:x2]\n    msks_new[:, :, y1:y2, x1:x2] = msks[perm, :, y1:y2, x1:x2]\n    return imgs_new, msks_new\n\n\n# ════════════════════════════════════════════════════════════\n#  9.  DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0, apply_ch_dropout=False):\n        self.cache          = load_slice_cache(frag_path, z_list)\n        self.z_list         = z_list\n        self.tf             = transform\n        self.ch_dropout     = apply_ch_dropout\n        self.frag_path      = frag_path\n\n        msk_path = os.path.join(frag_path, 'inklabels.png')\n        msk = cv2.imread(msk_path, 0)\n        if msk is None:\n            raise FileNotFoundError(f\"No inklabels at {msk_path}\")\n        self.mask = (msk > 0).astype(np.uint8)\n\n        ir_path = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = (cv2.imread(ir_path, 0) > 0).astype(np.uint8) \\\n                       if os.path.exists(ir_path) else None\n\n        H, W = self.mask.shape\n        pos_coords, neg_coords = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.sum() / (PATCH_SIZE * PATCH_SIZE)\n                if self.ir_mask is not None:\n                    on_pap = self.ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    mid_z  = z_list[len(z_list)//2]\n                    on_pap = float(self.cache[mid_z][y:y+PATCH_SIZE,\n                                                     x:x+PATCH_SIZE].mean()) > 0.1\n                if not on_pap:\n                    continue\n                if ink >= INK_MIN_POS:\n                    pos_coords.append((y, x, 1))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_coords.append((y, x, 0))\n\n        n_neg   = int(len(pos_coords) * neg_ratio)\n        sel_neg = []\n        if n_neg > 0 and neg_coords:\n            np.random.shuffle(neg_coords)\n            sel_neg = neg_coords[:n_neg]\n\n        self.coords  = pos_coords + sel_neg\n        self.weights = np.array(\n            [3.0 if c[2] == 1 else 1.0 for c in self.coords], dtype=np.float32\n        )\n        print(f\"  [{os.path.basename(frag_path)}] \"\n              f\"{len(pos_coords)} pos + {len(sel_neg)} neg = {len(self.coords)}\")\n\n    def __len__(self): return len(self.coords)\n\n    def get_patch(self, idx):\n        y, x, _ = self.coords[idx]\n        slices   = [self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                    for z in self.z_list]\n        return np.stack(slices, axis=-1), self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n\n    def __getitem__(self, idx):\n        img, msk = self.get_patch(idx)\n        if self.tf:\n            out = self.tf(image=img, mask=msk)\n            img, msk = out['image'], out['mask']\n        if self.ch_dropout:\n            img = channel_dropout(img)\n        return (torch.from_numpy(img).permute(2, 0, 1).float(),\n                torch.from_numpy(msk).unsqueeze(0).float())\n\n\n# ════════════════════════════════════════════════════════════\n#  10.  LOSS  (focal + dice + deep-supervision)\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = target * alpha + (1 - target) * (1 - alpha)\n    return (a_t * ((1 - p_t) ** gamma) * bce).mean()\n\n\ndef dice_loss(pred, target, smooth=1.):\n    p     = torch.sigmoid(pred)\n    inter = (p * target).sum(dim=(2, 3))\n    union = p.sum(dim=(2, 3)) + target.sum(dim=(2, 3))\n    return 1. - ((2. * inter + smooth) / (union + smooth)).mean()\n\n\ndef combined_loss(pred, target, eps=0.05):\n    t_s = target * (1 - eps) + 0.5 * eps\n    return 0.5 * focal_loss(pred, t_s) + 0.5 * dice_loss(pred, target)\n\n\ndef multiscale_loss(main_logit, sal_logit, ds3, ds2, ds1, target):\n    \"\"\"\n    main_logit, sal_logit : [B,1,H,W]\n    ds3,ds2,ds1           : lower-res deep-sup logits (may be None)\n    target                : [B,1,H,W]\n    \"\"\"\n    loss = combined_loss(main_logit, target)\n\n    # auxiliary sheet-surface detection loss (same target, down-sampled)\n    sal_t = F.adaptive_avg_pool2d(target, sal_logit.shape[-2:])\n    loss += 0.2 * combined_loss(sal_logit, sal_t)\n\n    def ds_loss(logit, wt):\n        if logit is None: return 0.0\n        tgt = F.adaptive_avg_pool2d(target, logit.shape[-2:])\n        return wt * combined_loss(logit, tgt)\n\n    loss += ds_loss(ds3, 0.15)\n    loss += ds_loss(ds2, 0.10)\n    loss += ds_loss(ds1, 0.05)\n    return loss\n\n\n# ════════════════════════════════════════════════════════════\n#  11.  BUILD DATASETS — 70/30 patch-level split on Frag2+Frag3\n# ════════════════════════════════════════════════════════════\nprint('\\n── Loading Fragment 2 (train pool) ──')\nds2 = VesuviusDataset(FRAG2, Z_SLICES, stride=STRIDE_TR,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\nprint('\\n── Loading Fragment 3 (train pool) ──')\nds3 = VesuviusDataset(FRAG3, Z_SLICES, stride=STRIDE_TR,\n                      transform=train_tf, neg_ratio=NEG_RATIO,\n                      apply_ch_dropout=True)\n\nn_total   = len(ds2) + len(ds3)\nrng       = np.random.RandomState(SEED)\nall_idx   = rng.permutation(n_total)\nn_val     = int(n_total * VAL_SPLIT)\nn_train   = n_total - n_val\ntrain_idx = all_idx[:n_train].tolist()\nval_idx   = all_idx[n_train:].tolist()\nprint(f'\\nTotal: {n_total} | Train: {n_train} (70%) | Val: {n_val} (30%)')\n\n\nclass ValSubset(Dataset):\n    def __init__(self, ds2, ds3, indices):\n        self.ds2 = ds2; self.ds3 = ds3\n        self.n2  = len(ds2); self.indices = indices\n\n    def __len__(self): return len(self.indices)\n\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        if g < self.n2:\n            img_np, msk_np = self.ds2.get_patch(g)\n        else:\n            img_np, msk_np = self.ds3.get_patch(g - self.n2)\n        return (torch.from_numpy(img_np).permute(2, 0, 1).float(),\n                torch.from_numpy(msk_np).unsqueeze(0).float())\n\n\nclass TrainSubset(Dataset):\n    def __init__(self, ds2, ds3, indices):\n        self.ds2 = ds2; self.ds3 = ds3\n        self.n2  = len(ds2); self.indices = indices\n\n    def __len__(self): return len(self.indices)\n\n    def __getitem__(self, idx):\n        g = self.indices[idx]\n        return self.ds2[g] if g < self.n2 else self.ds3[g - self.n2]\n\n\ntrain_ds = TrainSubset(ds2, ds3, train_idx)\nval_ds   = ValSubset  (ds2, ds3, val_idx)\n\nall_weights = np.concatenate([ds2.weights, ds3.weights])\ntrain_w     = torch.from_numpy(all_weights[train_idx]).float()\n\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\nsampler  = WeightedRandomSampler(train_w, len(train_ds), replacement=True)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                      num_workers=NUM_WORKERS, pin_memory=False)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=NUM_WORKERS, pin_memory=False)\nprint(f'Train batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  12.  METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p * masks).sum(dim=(1, 2, 3))\n    union = p.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3))\n    return ((2. * inter + 1e-5) / (union + 1e-5)).mean().item()\n\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.25, 0.85, 0.02):\n        p = (probs > t).astype(np.float32)\n        d = (2 * (p * targets).sum() + 1) / (p.sum() + targets.sum() + 1)\n        if d > best_d:\n            best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  13.  MODEL + OPTIM + SCHEDULER\n# ════════════════════════════════════════════════════════════\nmodel     = VesuviusV10(n_ch=N_CH, enc_dim=512,\n                        n_transformer_blocks=4, num_heads=8).to(DEVICE)\nn_params  = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'\\nModel parameters: {n_params/1e6:.2f} M')\n\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='max', factor=0.5, patience=5, min_lr=5e-6, verbose=True\n)\nscaler    = GradScaler(enabled=(DEVICE == 'cuda'))\nbest_dice = 0.; pat_cnt = 0\nhistory   = dict(tl=[], vl=[], td=[], vd=[], lr=[])\n\nprint('\\n' + '=' * 65)\nprint('TRAINING — VesuviusV10')\nprint('  Encoder      : ResNet34 + Axial-Transformer bottleneck')\nprint('  Attention    : Axial X/Y with log-polar positional bias')\nprint('  Token merging: 25% low-saliency tokens merged')\nprint('  Augmentation : Elastic + GridDistort + CutMix + ChannelDrop')\nprint('  Deep supervis: 3 auxiliary heads')\nprint('  NO TTA at test (removed for reliability)')\nprint(f'  Z-slices     : {Z_SLICES}')\nprint('=' * 65)\n\n\n# ════════════════════════════════════════════════════════════\n#  14.  TRAINING LOOP\n# ════════════════════════════════════════════════════════════\nfor epoch in range(EPOCHS):\n\n    # ── train ──────────────────────────────────────────────\n    model.train(); tl = td = 0.\n    optimizer.zero_grad()\n\n    for step, (imgs, msks) in enumerate(\n            tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n\n        imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n\n        # CutMix (30% of steps)\n        if random.random() < 0.30:\n            imgs, msks = cutmix_batch(imgs, msks, alpha=0.4)\n\n        with autocast(enabled=(DEVICE == 'cuda')):\n            main_logit, sal_logit, ds3_l, ds2_l, ds1_l = model(imgs)\n            loss = multiscale_loss(main_logit, sal_logit,\n                                   ds3_l, ds2_l, ds1_l, msks) / GRAD_ACCUM\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % GRAD_ACCUM == 0 or (step + 1) == len(train_dl):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer); scaler.update()\n            optimizer.zero_grad()\n\n        tl += loss.item() * GRAD_ACCUM\n        td += batch_dice(main_logit.detach(), msks)\n\n        del imgs, msks, main_logit, sal_logit, ds3_l, ds2_l, ds1_l, loss\n        if step % 50 == 0:\n            torch.cuda.empty_cache()\n\n    tl /= len(train_dl); td /= len(train_dl)\n\n    # ── validate ────────────────────────────────────────────\n    model.eval(); vl = vd = 0.\n    acc_p, acc_m = [], []\n\n    with torch.no_grad():\n        for imgs, msks in tqdm(val_dl, desc=f'Ep{epoch+1:02d} val  ', leave=False):\n            imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                main_logit, sal_logit, ds3_l, ds2_l, ds1_l = model(imgs)\n                loss = multiscale_loss(main_logit, sal_logit,\n                                       ds3_l, ds2_l, ds1_l, msks)\n            vl += loss.item()\n            vd += batch_dice(main_logit, msks)\n            acc_p.append(torch.sigmoid(main_logit / TEMPERATURE).cpu().numpy())\n            acc_m.append(msks.cpu().numpy())\n            del imgs, msks, main_logit, sal_logit, ds3_l, ds2_l, ds1_l, loss\n\n    vl /= len(val_dl); vd /= len(val_dl)\n    probs_all = np.concatenate(acc_p)\n    masks_all = np.concatenate(acc_m)\n    bt, bd    = sweep_threshold(probs_all, masks_all)\n\n    ink_mean   = float(probs_all[masks_all > 0.5].mean()) if (masks_all > 0.5).any() else 0.\n    noink_mean = float(probs_all[masks_all < 0.5].mean()) if (masks_all < 0.5).any() else 0.\n    sep        = ink_mean - noink_mean\n    del acc_p, acc_m, probs_all, masks_all\n    gc.collect(); torch.cuda.empty_cache()\n\n    scheduler.step(vd)\n    lr_now = optimizer.param_groups[0]['lr']\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['lr'].append(lr_now)\n\n    print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n          f'train loss={tl:.4f} dice={td:.4f} | '\n          f'val loss={vl:.4f} dice={vd:.4f} | '\n          f'thr={bt:.2f}→{bd:.4f} | '\n          f'sep={sep:+.3f} [ink={ink_mean:.3f} bg={noink_mean:.3f}]')\n\n    thr_penalty  = max(0.0, bt - 0.60) * 0.3\n    save_metric  = max(vd, bd) - thr_penalty\n\n    if save_metric > best_dice:\n        best_dice = save_metric; pat_cnt = 0\n        torch.save({'epoch': epoch, 'state': model.state_dict(),\n                    'thr': bt, 'metric': save_metric, 'n_ch': N_CH,\n                    'bd': bd, 'vd': vd},\n                   OUTPUT + 'best_model.pth')\n        print(f'  ✓ saved  (penalised={save_metric:.4f}  '\n              f'raw_dice={max(vd,bd):.4f}  thr={bt:.2f})')\n    else:\n        pat_cnt += 1\n        if pat_cnt >= PATIENCE:\n            print(f'  ⚑ early stop at epoch {epoch+1}')\n            break\n\nprint(f'\\nBest penalised metric: {best_dice:.4f}')\n\n# ── curves ──────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 4))\naxes[0].plot(history['tl'], label='train')\naxes[0].plot(history['vl'], label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'], label='train')\naxes[1].plot(history['vd'], label='val')\naxes[1].axhline(0.80, color='r', ls='--', label='target 0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['lr'])\naxes[2].set_title('LR'); axes[2].set_xlabel('Epoch'); axes[2].grid(True)\nplt.tight_layout()\nplt.savefig(OUTPUT + 'curves.png', dpi=100); plt.close()\nprint('Curves saved.')\n\n\n# ════════════════════════════════════════════════════════════\n#  15.  PSEUDO-LABELLING  (optional)\n#       If USE_PSEUDO is True and PSEUDO_FRAGS is not empty,\n#       we run inference on unlabelled fragments, keep\n#       high-confidence predictions as pseudo labels,\n#       then fine-tune for a few epochs.\n# ════════════════════════════════════════════════════════════\ndef generate_pseudo_labels(model, frag_path, z_list, threshold=PSEUDO_THRESHOLD):\n    \"\"\"Run sliding-window inference → save pseudo inklabels.png.\"\"\"\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    # build a small 1-slice proxy for spatial size\n    mid_z = z_list[len(z_list) // 2]\n    H, W  = cache[mid_z].shape\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n\n    coords = [(y, x) for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n                      for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='PseudoLabel', leave=False):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            patch  = np.stack(slices, axis=-1)\n            t      = torch.from_numpy(patch).permute(2, 0, 1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit / TEMPERATURE).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += 1.0\n            del t, logit\n\n    del cache; gc.collect(); torch.cuda.empty_cache()\n    prob_map = pred_map / (wgt_map + 1e-8)\n    # Keep only confident predictions\n    pseudo_lbl = ((prob_map > threshold) * 255).astype(np.uint8)\n    out_path   = os.path.join(frag_path, 'inklabels_pseudo.png')\n    cv2.imwrite(out_path, pseudo_lbl)\n    print(f'  Pseudo label saved to {out_path}  '\n          f'(ink%={100*(pseudo_lbl>0).mean():.1f}%)')\n    return out_path\n\n\nif USE_PSEUDO and PSEUDO_FRAGS:\n    print('\\n── Pseudo-labelling extra fragments ──')\n    ckpt = torch.load(OUTPUT + 'best_model.pth', map_location=DEVICE)\n    model.load_state_dict(ckpt['state'])\n\n    pseudo_datasets = []\n    for pf in PSEUDO_FRAGS:\n        lbl_path = generate_pseudo_labels(model, pf, Z_SLICES)\n        # temporarily rename pseudo label so the dataset can load it\n        ink_orig = os.path.join(pf, 'inklabels.png')\n        ink_back = os.path.join(pf, 'inklabels_real.png')\n        if not os.path.exists(ink_back) and os.path.exists(ink_orig):\n            os.rename(ink_orig, ink_back)\n        os.rename(lbl_path, ink_orig)\n        try:\n            pds = VesuviusDataset(pf, Z_SLICES, stride=STRIDE_TR,\n                                  transform=train_tf, neg_ratio=0.3,\n                                  apply_ch_dropout=True)\n            pseudo_datasets.append(pds)\n        except Exception as e:\n            print(f'  [WARNING] Could not build pseudo dataset for {pf}: {e}')\n        finally:\n            # restore original label\n            if os.path.exists(ink_back):\n                os.replace(ink_back, ink_orig)\n\n    if pseudo_datasets:\n        print(f'\\n── Fine-tuning with {sum(len(p) for p in pseudo_datasets)} pseudo patches ──')\n        from torch.utils.data import ConcatDataset\n        pseudo_combined = ConcatDataset(pseudo_datasets)\n        pseudo_dl = DataLoader(pseudo_combined, batch_size=BATCH_SIZE,\n                               shuffle=True, num_workers=NUM_WORKERS)\n        pseudo_optimizer = optim.AdamW(model.parameters(),\n                                       lr=LR * 0.2, weight_decay=WEIGHT_DECAY)\n        pseudo_scaler    = GradScaler(enabled=(DEVICE == 'cuda'))\n\n        for ep in range(5):   # short fine-tune\n            model.train(); pl = pd_val = 0.\n            pseudo_optimizer.zero_grad()\n            for step, (imgs, msks) in enumerate(\n                    tqdm(pseudo_dl, desc=f'Pseudo Ep{ep+1}', leave=False)):\n                imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n                with autocast(enabled=(DEVICE == 'cuda')):\n                    main_logit, sal_logit, d3, d2, d1 = model(imgs)\n                    loss = multiscale_loss(main_logit, sal_logit,\n                                          d3, d2, d1, msks) / GRAD_ACCUM\n                pseudo_scaler.scale(loss).backward()\n                if (step+1) % GRAD_ACCUM == 0:\n                    pseudo_scaler.unscale_(pseudo_optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                    pseudo_scaler.step(pseudo_optimizer)\n                    pseudo_scaler.update()\n                    pseudo_optimizer.zero_grad()\n                pl += loss.item() * GRAD_ACCUM\n                pd_val += batch_dice(main_logit.detach(), msks)\n                del imgs, msks, main_logit, sal_logit, d3, d2, d1, loss\n            pl /= len(pseudo_dl); pd_val /= len(pseudo_dl)\n            print(f'  Pseudo Ep{ep+1} | loss={pl:.4f} | dice={pd_val:.4f}')\n        torch.save({'state': model.state_dict(), 'n_ch': N_CH,\n                    'thr': ckpt['thr']},\n                   OUTPUT + 'best_model_pseudo.pth')\n        print('Pseudo fine-tune complete.')\n\n\n# ════════════════════════════════════════════════════════════\n#  16.  INFERENCE  (NO TTA — clean deterministic single-pass)\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c   = sz // 2; sig = sz // 4\n    ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs - c)**2 + (ys - c)**2) / (2 * sig**2)).astype(np.float32)\n\n\nGW = gauss_weight(PATCH_SIZE)\n\n\ndef predict_fragment(model, frag_path, z_list, temperature=TEMPERATURE):\n    \"\"\"\n    Sliding-window inference WITHOUT TTA.\n    Returns (prob_map [H,W], label_mask [H,W]).\n    \"\"\"\n    model.eval()\n    cache   = load_slice_cache(frag_path, z_list)\n    lbl_img = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n    msk     = (lbl_img > 0).astype(np.uint8)\n    H, W    = msk.shape\n\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n\n    coords = [(y, x) for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n                      for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='Inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            patch  = np.stack(slices, axis=-1)\n            t      = torch.from_numpy(patch).permute(2, 0, 1).unsqueeze(0).to(DEVICE)\n            with autocast(enabled=(DEVICE == 'cuda')):\n                logit, *_ = model(t)\n            p = torch.sigmoid(logit / temperature).squeeze().cpu().numpy().astype(np.float32)\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p  * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, logit\n            torch.cuda.empty_cache()\n\n    del cache; gc.collect()\n    return pred_map / (wgt_map + 1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  17.  POST-PROCESSING\n#       (a) Morphological cleaning  — removes small FP islands\n#       (b) DenseCRF                — refines boundaries\n# ════════════════════════════════════════════════════════════\ndef morphological_clean(binary_map, min_area=200, close_k=5):\n    \"\"\"\n    1. Close small holes with a small kernel\n    2. Remove connected components smaller than min_area pixels\n    \"\"\"\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (close_k, close_k))\n    closed = cv2.morphologyEx(binary_map.astype(np.uint8),\n                              cv2.MORPH_CLOSE, kernel)\n    # remove small components\n    n_lab, labels, stats, _ = cv2.connectedComponentsWithStats(closed)\n    cleaned = np.zeros_like(closed)\n    for i in range(1, n_lab):\n        if stats[i, cv2.CC_STAT_AREA] >= min_area:\n            cleaned[labels == i] = 1\n    return cleaned\n\n\ndef apply_dense_crf(image_uint8, prob_map, n_iter=5):\n    \"\"\"\n    image_uint8 : [H, W, 3] uint8 (mid z-slice repeated to RGB)\n    prob_map    : [H, W] float in [0,1]\n    Returns     : refined binary prediction [H, W]\n    \"\"\"\n    if not HAS_CRF:\n        print(\"[INFO] DenseCRF unavailable – skipping.\")\n        return (prob_map > 0.5).astype(np.uint8)\n\n    H, W = prob_map.shape\n    d    = dcrf.DenseCRF2D(W, H, 2)\n\n    # unary potentials from probability map\n    fg  = np.clip(prob_map,       1e-5, 1 - 1e-5)\n    bg  = np.clip(1.0 - prob_map, 1e-5, 1 - 1e-5)\n    U   = -np.log(np.stack([bg, fg], axis=0))    # [2, H*W]\n    d.setUnaryEnergy(U.reshape(2, -1).astype(np.float32))\n\n    # pairwise: Gaussian spatial (smoothness)\n    d.addPairwiseGaussian(sxy=3, compat=3)\n\n    # pairwise: bilateral (edge-aware)\n    img_c = np.ascontiguousarray(image_uint8)\n    d.addPairwiseBilateral(sxy=50, srgb=13, rgbim=img_c, compat=10)\n\n    Q = d.inference(n_iter)\n    return np.argmax(Q, axis=0).reshape(H, W).astype(np.uint8)\n\n\ndef postprocess(prob_map, raw_vol_cache, z_list,\n                threshold, min_area=200, close_k=5, crf_iters=5):\n    \"\"\"Full post-processing pipeline.\"\"\"\n    binary = (prob_map > threshold).astype(np.uint8)\n\n    # Morphological cleaning\n    binary = morphological_clean(binary, min_area=min_area, close_k=close_k)\n\n    # DenseCRF using mid z-slice as colour guidance\n    if HAS_CRF:\n        mid_z  = z_list[len(z_list) // 2]\n        mid_sl = raw_vol_cache[mid_z].astype(np.float32)\n        mid_8  = ((mid_sl - mid_sl.min()) /\n                  (mid_sl.max() - mid_sl.min() + 1e-8) * 255).astype(np.uint8)\n        rgb    = np.stack([mid_8, mid_8, mid_8], axis=-1)\n        binary = apply_dense_crf(rgb, prob_map, n_iter=crf_iters)\n\n    return binary\n\n\n# ════════════════════════════════════════════════════════════\n#  18.  FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '=' * 65)\nprint('FINAL TEST — FRAGMENT 1  (never seen during training)')\nprint('NO TTA — single deterministic inference pass')\nprint('=' * 65)\n\nckpt = torch.load(OUTPUT + 'best_model.pth', map_location=DEVICE)\nassert ckpt.get('n_ch', N_CH) == N_CH, \\\n    f\"Channel mismatch: ckpt={ckpt.get('n_ch')} vs model={N_CH}. Retrain.\"\n\nmodel = VesuviusV10(n_ch=N_CH, enc_dim=512,\n                    n_transformer_blocks=4, num_heads=8).to(DEVICE)\nmodel.load_state_dict(ckpt['state'], strict=True)\nsaved_thr = ckpt.get('thr', 0.5)\nprint(f'Checkpoint: epoch {ckpt[\"epoch\"]+1} | '\n      f'val_dice={ckpt.get(\"vd\",0):.4f} | '\n      f'opt_dice={ckpt.get(\"bd\",0):.4f} | '\n      f'thr={saved_thr:.2f}')\n\nprob_map, msk1 = predict_fragment(model, FRAG1, Z_SLICES)\nH_m, W_m       = msk1.shape\nprob_crop       = prob_map[:H_m, :W_m]\n\n# Optimal threshold sweep on Fragment 1\nfrom sklearn.metrics import confusion_matrix\nbt1, bd1 = sweep_threshold(prob_crop[np.newaxis, np.newaxis],\n                            msk1[np.newaxis, np.newaxis])\n\n# ── Post-processing ─────────────────────────────────────────\nprint(f'\\nPost-processing: threshold={bt1:.2f}, morph clean, DenseCRF ...')\ncache1      = load_slice_cache(FRAG1, Z_SLICES)\nfinal_pred  = postprocess(prob_crop, cache1, Z_SLICES,\n                          threshold=bt1, min_area=150, close_k=5, crf_iters=5)\ndel cache1; gc.collect()\n\n# ── Metrics ─────────────────────────────────────────────────\npf = final_pred.flatten().astype(int)\nmf = msk1.flatten().astype(int)\ntn, fp, fn, tp_v = confusion_matrix(mf, pf, labels=[0, 1]).ravel()\nprec = tp_v / (tp_v + fp + 1e-8)\nrec  = tp_v / (tp_v + fn + 1e-8)\nf1   = 2 * prec * rec / (prec + rec + 1e-8)\ndice_full = (2 * tp_v + 1) / (final_pred.sum() + msk1.sum() + 1)\n\nink_mean   = float(prob_crop[msk1 == 1].mean())\nnoink_mean = float(prob_crop[msk1 == 0].mean())\n\nprint('\\n' + '=' * 65)\nprint('RESULTS — FRAGMENT 1')\nprint('=' * 65)\nprint(f'Dice Score  : {dice_full:.4f}')\nprint(f'F1 Score    : {f1:.4f}')\nprint(f'Precision   : {prec:.4f}')\nprint(f'Recall      : {rec:.4f}')\nprint(f'Threshold   : {bt1:.2f}')\nprint(f'TP={tp_v} | TN={tn} | FP={fp} | FN={fn}')\nprint(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}')\nprint(f'Ink/bg sep  : {ink_mean-noink_mean:+.3f}')\ntarget_str  = '✓ PASSED' if f1 >= 0.80 else '✗ Below 0.80 target'\nprint(f'F1 ≥ 0.80   : {target_str}')\nprint('=' * 65)\n\n# ── Visualisation ────────────────────────────────────────────\nfig, ax = plt.subplots(2, 3, figsize=(18, 12))\nax[0,0].imshow(msk1,        cmap='gray');    ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop,   cmap='inferno'); ax[0,1].set_title('Probability Map')\nax[0,2].imshow(final_pred,  cmap='gray');    ax[0,2].set_title(\n    f'Prediction (post-processed)  dice={dice_full:.3f}')\n\nerr = np.zeros((*msk1.shape, 3), dtype=np.uint8)\nerr[(final_pred == 1) & (msk1 == 1)] = [0,   255, 0]\nerr[(final_pred == 1) & (msk1 == 0)] = [255,   0, 0]\nerr[(final_pred == 0) & (msk1 == 1)] = [0,     0, 255]\nax[1,0].imshow(err); ax[1,0].set_title('TP=green  FP=red  FN=blue')\n\nax[1,1].hist(prob_crop[msk1 == 1].ravel(), bins=50, alpha=0.7,\n             label=f'ink (μ={ink_mean:.2f})',       color='orange', density=True)\nax[1,1].hist(prob_crop[msk1 == 0].ravel(), bins=50, alpha=0.7,\n             label=f'no-ink (μ={noink_mean:.2f})',  color='blue',   density=True)\nax[1,1].axvline(bt1, color='r', ls='--', label=f'thr={bt1:.2f}')\nax[1,1].set_title('Probability Distribution')\nax[1,1].legend(); ax[1,1].set_xlabel('Probability')\n\nts = np.arange(0.15, 0.90, 0.01); ds = []\nfor t in ts:\n    p_b = (prob_crop > t).astype(np.float32)\n    ds.append((2 * (p_b * msk1).sum() + 1) / (p_b.sum() + msk1.sum() + 1))\nax[1,2].plot(ts, ds)\nax[1,2].axvline(bt1, color='r', ls='--', label=f'best={bt1:.2f}')\nax[1,2].set_title('Dice vs Threshold')\nax[1,2].set_xlabel('Threshold'); ax[1,2].set_ylabel('Dice')\nax[1,2].legend(); ax[1,2].grid(True)\n\nfor a in [ax[0,0], ax[0,1], ax[0,2], ax[1,0]]: a.axis('off')\nplt.suptitle(\n    f'Fragment 1 — Dice={dice_full:.4f}  F1={f1:.4f}  '\n    f'Prec={prec:.4f}  Rec={rec:.4f}',\n    fontsize=13, y=1.01\n)\nplt.tight_layout()\nplt.savefig(OUTPUT + 'frag1_prediction.png', dpi=100, bbox_inches='tight')\nplt.close()\n\n# ── Save results ─────────────────────────────────────────────\nwith open(OUTPUT + 'final_results.txt', 'w') as f:\n    f.write('='*50 + '\\n')\n    f.write('VESUVIUS V10 — FINAL TEST RESULTS\\n')\n    f.write('='*50 + '\\n')\n    f.write(f'Dice        : {dice_full:.4f}\\n')\n    f.write(f'F1          : {f1:.4f}\\n')\n    f.write(f'Precision   : {prec:.4f}\\n')\n    f.write(f'Recall      : {rec:.4f}\\n')\n    f.write(f'Threshold   : {bt1:.2f}\\n')\n    f.write(f'TP={tp_v} TN={tn} FP={fp} FN={fn}\\n')\n    f.write(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Ink/bg sep  : {ink_mean-noink_mean:+.3f}\\n')\n    f.write(f'Temperature : {TEMPERATURE}\\n')\n    f.write(f'Post-proc   : morph_clean + DenseCRF={HAS_CRF}\\n')\n    f.write('='*50 + '\\n')\n\n# ── Save prediction images ────────────────────────────────────\ncv2.imwrite(OUTPUT + 'frag1_prob_map.png',\n            (prob_crop * 255).astype(np.uint8))\ncv2.imwrite(OUTPUT + 'frag1_prediction_binary.png',\n            (final_pred * 255).astype(np.uint8))\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  best_model.pth | curves.png | frag1_prediction.png')\nprint('  frag1_prob_map.png | frag1_prediction_binary.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n#  VESUVIUS INK DETECTION — v21\n#  Physics-Based Ink Enhancement + Learned Denoising\n#  Applied on v19's clean foundation\n#\n#  CORE IDEA:\n#  Ink changes CT properties in a specific physical way.\n#  Instead of giving the CNN raw slices and hoping it discovers\n#  these patterns from scratch, we pre-compute 16 physics-based\n#  \"hint\" channels that directly highlight ink signatures.\n#\n#  THE 16 ENHANCED CHANNELS (from 12 raw CT slices):\n#  Group A — Raw normalised slices (4 key slices, not all 12):\n#    [0-3]  : slices at z=27,29,31,33  (peak ink zone)\n#\n#  Group B — Z-axis derivatives (ink has sharp z-transitions):\n#    [4-7]  : forward  diff  dI/dz at z=27,29,31,33\n#             = slice[z+1] - slice[z]\n#             Ink sitting at depth z creates a positive spike here\n#\n#  Group C — Second derivative (ink boundary sharpness):\n#    [8-9]  : d²I/dz² at z=29,31\n#             = slice[z+1] - 2*slice[z] + slice[z-1]\n#             Highlights ink layer boundaries (zero-crossing = edge)\n#\n#  Group D — Median-filtered residual (local anomaly detector):\n#    [10-11]: slice - median_filter(slice, 5×5) at z=29,31\n#             Removes smooth papyrus background, leaves ink deposits\n#\n#  Group E — Inter-slice maximum contrast (ink zone indicator):\n#    [12]   : max over all z of |slice[z+1] - slice[z]|\n#             Single map showing WHERE transitions happen most\n#\n#  Group F — Local variance across z (ink = high z-variance):\n#    [13]   : std across all 12 slices at each (x,y) pixel\n#             Ink pixels have high variance across z;\n#             uniform papyrus has low variance\n#\n#  Group G — Cumulative ink score (physics-based prior):\n#    [14]   : sum of positive z-derivatives (ink accumulation map)\n#    [15]   : sum of |d²I/dz²| across all z (total curvature)\n#\n#  WHY THIS WORKS:\n#  - Ink sits at a specific z-depth → sharp z-derivative spike\n#  - Ink has different density than papyrus → local residual\n#  - Ink layer is thin → high second derivative\n#  - These are PHYSICS FACTS about ink in CT, not learned heuristics\n#  - The CNN receives pre-digested ink signals → learns faster,\n#    overfits less, generalises better to new fragments\n#\n#  ARCHITECTURE — Two-stage pipeline:\n#  Stage A: InkDenoiser (small encoder-decoder)\n#    Input : 16 enhanced channels\n#    Output: 1-channel \"clean ink map\" (soft probability)\n#    Loss  : BCE against ink mask (denoising autoencoder style —\n#            target is NOT the input, it is the ink label)\n#    This stage learns to suppress CT noise while preserving ink\n#\n#  Stage B: Unet++ / EfficientNet-B4 / SCSE\n#    Input : 16 enhanced channels + 1 denoised map = 17 channels\n#    Output: final binary ink segmentation\n#    Loss  : Tversky + Focal\n#\n#  Both stages trained jointly with combined loss:\n#    L = 0.3 * L_denoiser + 0.7 * L_segmenter\n#\n#  DATA PROTOCOL (unchanged from v19):\n#    Train    : Frag2 (all) + Frag3 (80%)\n#    Validate : Frag3 (20% held-out)\n#    Test     : Frag1 (never seen during training)\n# ============================================================\n!pip install segmentation-models-pytorch==0.2.0\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import (Dataset, DataLoader, WeightedRandomSampler,\n                               ConcatDataset, Subset)\nfrom torch.optim.swa_utils import AveragedModel, SWALR\nimport albumentations as A\nfrom scipy.ndimage import median_filter\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nimport warnings\nwarnings.filterwarnings('ignore', category=UserWarning)\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ── paths ─────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ──────────────────────────────────────────────\nDEVICE         = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE     = 224\nSTRIDE_TR      = 112\nSTRIDE_INF     = 56\nBATCH_SIZE     = 8          # 17-ch input is heavier — reduce batch\nGRAD_ACCUM     = 4          # effective batch = 32\nEPOCHS         = 25\nSWA_START      = 25\nLR             = 1e-4\nWEIGHT_DECAY   = 1e-5\nPATIENCE       = 12\nNUM_WORKERS    = 4\nMAX_PATCHES_TR = 12_000\n\n# Raw CT z-slices (12 slices, same as v19)\nZ_SLICES    = list(range(18, 30))   # z=25..36\nN_RAW       = len(Z_SLICES)         # 12 raw slices\n\n# Key z-indices within Z_SLICES for selected-slice channels\n# We pick z=27,29,31,33 → indices 2,4,6,8 within Z_SLICES\nKEY_Z_IDX   = [2, 4, 6, 8]         # 4 key slices (Group A)\n\n# Enhanced channel count: 4+4+2+2+1+1+1+1 = 16\nN_ENHANCED  = 16\n# Final model input: 16 enhanced + 1 denoiser output = 17\nN_CH_FINAL  = N_ENHANCED + 1\n\nINK_MIN_POS = 0.02\nNEG_RATIO   = 0.3\nPOS_WEIGHT  = 3.0\nPIN         = (DEVICE == 'cuda')\nUSE_AMP     = (DEVICE == 'cuda')\n\n# Loss weights: denoiser vs segmenter\nW_DENOISE   = 0.3\nW_SEGMENT   = 0.7\n\nNORM_STD_MIN = 0.01\nNORM_CLIP    = 10.0\n\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Raw slices: {N_RAW}  →  Enhanced channels: {N_ENHANCED}\")\nprint(f\"Final input: {N_ENHANCED} enhanced + 1 denoised = {N_CH_FINAL} ch\")\nif DEVICE == 'cuda':\n    print(f\"GPU  : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. RAW CT CACHE  (float16)\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_ct_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        print(f\"  Reusing CT: {os.path.basename(frag_path)}\")\n        return _cache_store[key]\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    assert raw, f\"No CT slices in {vol_dir}\"\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}; H = W = None\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        s = ((s - p5)/(p95 - p5 + 1e-6)).astype(np.float16)\n        cache[z] = s\n        if H is None: H, W = s.shape\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  CT [{os.path.basename(frag_path)}]: \"\n          f\"{len(cache)} slices, {mb:.0f} MB, ({H},{W})\")\n    _cache_store[key] = cache\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2. PHYSICS-BASED INK ENHANCEMENT\n#\n#  Input : raw_vol  — numpy float32 (H, W, N_RAW)\n#          Each slice is already normalised to ~[0,1].\n#\n#  Output: enhanced — numpy float32 (H, W, N_ENHANCED=16)\n#\n#  This runs ONCE per patch in __getitem__.\n#  All operations are pure numpy — no torch, no GPU needed here.\n# ════════════════════════════════════════════════════════════\ndef physics_enhance(raw_vol, key_z_idx=KEY_Z_IDX):\n    \"\"\"\n    raw_vol : (H, W, N_RAW) float32, normalised\n    Returns : (H, W, N_ENHANCED) float32\n    \"\"\"\n    H, W, N = raw_vol.shape\n    channels = []\n\n    # ── Group A: 4 key raw slices ─────────────────────────\n    # These are the slices most likely to contain ink (mid-range z).\n    for zi in key_z_idx:\n        channels.append(raw_vol[:,:,zi])                    # [0-3]\n\n    # ── Group B: forward z-derivative at 4 key slices ─────\n    # dI/dz = slice[z+1] - slice[z]\n    # Ink causes a density increase then decrease as you scan through\n    # it in z, creating a characteristic positive-then-negative\n    # spike pattern. This channel captures the ONSET of ink.\n    for zi in key_z_idx:\n        if zi + 1 < N:\n            diff = raw_vol[:,:,zi+1] - raw_vol[:,:,zi]\n        else:\n            diff = raw_vol[:,:,zi] - raw_vol[:,:,zi-1]\n        channels.append(diff)                               # [4-7]\n\n    # ── Group C: second z-derivative at 2 central slices ──\n    # d²I/dz² = slice[z+1] - 2*slice[z] + slice[z-1]\n    # The second derivative is large at the ink layer BOUNDARIES.\n    # It is zero inside uniform material, large at transitions.\n    for zi in [key_z_idx[1], key_z_idx[2]]:               # z=29,31\n        if 0 < zi < N-1:\n            d2 = (raw_vol[:,:,zi+1]\n                  - 2*raw_vol[:,:,zi]\n                  + raw_vol[:,:,zi-1])\n        else:\n            d2 = np.zeros((H,W), dtype=np.float32)\n        channels.append(d2)                                 # [8-9]\n\n    # ── Group D: median-filtered residual ─────────────────\n    # slice - median_filter(slice, 5×5)\n    # The median filter removes slowly-varying background (papyrus).\n    # What remains = local anomalies = ink deposits.\n    # This is the single most powerful single-channel feature.\n    for zi in [key_z_idx[1], key_z_idx[2]]:               # z=29,31\n        sl = raw_vol[:,:,zi]\n        # median_filter is expensive: use small kernel for speed\n        med = median_filter(sl, size=5)\n        residual = sl - med\n        channels.append(residual)                           # [10-11]\n\n    # ── Group E: max inter-slice contrast ─────────────────\n    # max over all z of |slice[z+1] - slice[z]|\n    # A single map that shows WHERE z-transitions are strongest.\n    # On ink pixels this is high everywhere in z;\n    # on blank papyrus it is near zero.\n    diffs_abs = np.abs(np.diff(raw_vol, axis=2))           # (H,W,N-1)\n    max_diff  = diffs_abs.max(axis=2)                      # (H,W)\n    channels.append(max_diff)                               # [12]\n\n    # ── Group F: z-axis standard deviation ────────────────\n    # std(slice[0], slice[1], ..., slice[N-1]) at each pixel\n    # Ink pixels have HIGH variance across z (density changes a lot).\n    # Papyrus pixels have LOW variance (relatively homogeneous).\n    z_std = raw_vol.std(axis=2)                            # (H,W)\n    channels.append(z_std)                                  # [13]\n\n    # ── Group G: cumulative ink score ─────────────────────\n    # [14] Sum of POSITIVE z-derivatives (ink accumulation map)\n    #      Where ink is, there are consistently positive dI/dz\n    #      as we enter the ink layer from below.\n    pos_diffs = np.maximum(0, np.diff(raw_vol, axis=2))\n    ink_accum = pos_diffs.sum(axis=2)                      # (H,W)\n    channels.append(ink_accum)                              # [14]\n\n    # [15] Total absolute curvature sum\n    #      Ink creates many sharp z-transitions → high total curvature.\n    #      This integrates the physical prior across the full z-range.\n    d2_all = np.abs(np.diff(raw_vol, n=2, axis=2))\n    total_curv = d2_all.sum(axis=2)                        # (H,W)\n    channels.append(total_curv)                             # [15]\n\n    enhanced = np.stack(channels, axis=-1)                 # (H,W,16)\n    assert enhanced.shape[2] == N_ENHANCED, \\\n        f\"Expected {N_ENHANCED} channels, got {enhanced.shape[2]}\"\n    return enhanced.astype(np.float32)\n\n\ndef norm_patch(img_t):\n    \"\"\"Per-channel safe normalisation. img_t: (C,H,W).\"\"\"\n    mu  = img_t.mean(dim=(1,2), keepdim=True)\n    std = img_t.std (dim=(1,2), keepdim=True).clamp(min=NORM_STD_MIN)\n    return ((img_t - mu) / std).clamp(-NORM_CLIP, NORM_CLIP)\n\n\n# ════════════════════════════════════════════════════════════\n#  3. PAPYRUS MASK\n# ════════════════════════════════════════════════════════════\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H, W):\n                m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n            return (m > 0).astype(np.uint8)\n    mid_z = list(cache.keys())[len(cache)//2]\n    return (cache[mid_z] > 0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Affine(\n        scale=(0.85, 1.15),\n        translate_percent={'x': (-0.1, 0.1), 'y': (-0.1, 0.1)},\n        rotate=(-30, 30),\n        border_mode=cv2.BORDER_REFLECT,\n        p=0.6),\n    A.RandomBrightnessContrast(0.25, 0.25, p=0.5),\n    A.GaussianBlur(blur_limit=(3, 7), p=0.3),\n    A.CoarseDropout(\n        max_holes=6, max_height=32, max_width=32,\n        fill_value=0, p=0.4),\n    A.ElasticTransform(alpha=30, sigma=5, p=0.3),\n    A.GridDistortion(p=0.2),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  5. DATASET  — returns 16 enhanced channels per patch\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    \"\"\"\n    Loads raw CT → applies physics_enhance() → returns\n    (enhanced_tensor, mask_tensor).\n    enhanced_tensor: (N_ENHANCED, H, W) = (16, 224, 224)\n    \"\"\"\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0, max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        self.tf     = transform\n        H, W        = next(iter(self.cache.values())).shape\n\n        msk_p = os.path.join(frag_path, 'inklabels.png')\n        msk   = cv2.imread(msk_p, 0)\n        assert msk is not None, f\"Missing: {msk_p}\"\n        self.mask = (msk > 0).astype(np.uint8)\n\n        pap = load_papyrus_mask(frag_path, H, W, self.cache)\n\n        pos_yx, neg_yx = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                if pap[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() < 0.5:\n                    continue\n                ink = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                if ink >= INK_MIN_POS:\n                    pos_yx.append((y, x))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_yx.append((y, x))\n\n        n_neg = int(len(pos_yx) * neg_ratio)\n        if n_neg > 0 and neg_yx:\n            np.random.shuffle(neg_yx); neg_yx = neg_yx[:n_neg]\n        else:\n            neg_yx = []\n\n        n_total = len(pos_yx) + len(neg_yx)\n        if max_patches > 0 and n_total > max_patches:\n            frac = max_patches / n_total\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx = pos_yx[:max(1, int(len(pos_yx)*frac))]\n            neg_yx = neg_yx[:max(0, int(len(neg_yx)*frac))]\n\n        all_yx  = pos_yx + neg_yx\n        all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n        perm    = np.random.permutation(len(all_yx))\n        self.coords  = np.array([all_yx[i] for i in perm], dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm], dtype=np.int32)\n        self.weights = np.where(self.labels==1,\n                                POS_WEIGHT, 1.0).astype(np.float32)\n\n        cap = (f\" (capped from {n_total})\"\n               if max_patches>0 and n_total>max_patches else \"\")\n        print(f\"  [{os.path.basename(frag_path)}]: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg \"\n              f\"= {len(self.coords)} total{cap}\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x = int(self.coords[idx,0]), int(self.coords[idx,1])\n        ps   = PATCH_SIZE\n\n        # Step 1: extract raw CT volume patch (H, W, N_RAW)\n        raw_slices = [\n            self.cache[z][y:y+ps, x:x+ps].astype(np.float32)\n            for z in self.z_list\n        ]\n        raw_vol = np.stack(raw_slices, axis=-1)  # (H,W,N_RAW)\n\n        msk = self.mask[y:y+ps, x:x+ps].copy()\n\n        # Step 2: augment raw volume + mask together\n        # We augment the raw volume (not the enhanced),\n        # then apply enhancement after augmentation.\n        # This ensures derivatives/residuals are computed\n        # on the augmented geometry, not the original.\n        if self.tf:\n            out    = self.tf(image=raw_vol, mask=msk)\n            raw_vol, msk = out['image'], out['mask']\n\n        # Step 3: physics-based ink enhancement\n        enhanced = physics_enhance(raw_vol)           # (H,W,16)\n\n        # Step 4: normalise each channel independently\n        img_t = norm_patch(\n            torch.from_numpy(enhanced).permute(2,0,1).float())\n        msk_t = torch.from_numpy(msk).unsqueeze(0).float()\n        return img_t, msk_t\n\n\n# ════════════════════════════════════════════════════════════\n#  6. ARCHITECTURE\n#\n#  Stage A — InkDenoiser\n#  A lightweight encoder-decoder that maps 16 enhanced channels\n#  to a 1-channel clean ink probability map.\n#  It acts as a \"physics-guided denoising autoencoder\":\n#    - Input : 16 channels with physics-based ink hints\n#    - Target: ink mask (not the input itself — key difference)\n#  This stage learns to denoise CT noise while preserving ink.\n#  Output is sigmoid(logit) ∈ (0,1) — used as channel 17.\n#\n#  Stage B — Unet++ segmenter\n#  Takes [16 enhanced channels + 1 denoised map] = 17 channels.\n#  The denoised map gives the segmenter a clean starting point,\n#  so it only needs to refine boundaries rather than learn\n#  the full ink detection from scratch.\n# ════════════════════════════════════════════════════════════\n\nclass ConvBnRelu(nn.Module):\n    def __init__(self, cin, cout, k=3, p=1):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(cin, cout, k, padding=p, bias=False),\n            nn.BatchNorm2d(cout),\n            nn.ReLU(inplace=True))\n    def forward(self, x): return self.block(x)\n\n\nclass InkDenoiser(nn.Module):\n    \"\"\"\n    Lightweight encoder-decoder.\n    Input  : (B, N_ENHANCED, H, W)  = (B, 16, 224, 224)\n    Output : (B, 1, H, W)  logits (apply sigmoid for prob map)\n\n    Architecture:\n      Encoder: 3 stride-2 conv blocks (224→112→56→28)\n      Bottleneck: 2 conv blocks at 28×28\n      Decoder: 3 transposed-conv upsamples with skip connections\n      Head: 1×1 conv → 1 channel logit\n    \"\"\"\n    def __init__(self, n_in=N_ENHANCED):\n        super().__init__()\n        # Encoder\n        self.enc1 = ConvBnRelu(n_in, 32)       # 224×224\n        self.down1= nn.MaxPool2d(2)             # 112×112\n        self.enc2 = ConvBnRelu(32, 64)\n        self.down2= nn.MaxPool2d(2)             # 56×56\n        self.enc3 = ConvBnRelu(64, 128)\n        self.down3= nn.MaxPool2d(2)             # 28×28\n        # Bottleneck\n        self.bot  = nn.Sequential(\n            ConvBnRelu(128, 256),\n            ConvBnRelu(256, 128))\n        # Decoder with skip connections\n        self.up3  = nn.ConvTranspose2d(128, 128, 2, stride=2)\n        self.dec3 = ConvBnRelu(128+128, 64)    # skip from enc3\n        self.up2  = nn.ConvTranspose2d(64, 64, 2, stride=2)\n        self.dec2 = ConvBnRelu(64+64, 32)      # skip from enc2\n        self.up1  = nn.ConvTranspose2d(32, 32, 2, stride=2)\n        self.dec1 = ConvBnRelu(32+32, 16)      # skip from enc1\n        # Output head\n        self.head = nn.Conv2d(16, 1, 1)\n\n    def forward(self, x):\n        e1   = self.enc1(x)\n        e2   = self.enc2(self.down1(e1))\n        e3   = self.enc3(self.down2(e2))\n        b    = self.bot(self.down3(e3))\n        d3   = self.dec3(torch.cat([self.up3(b),  e3], dim=1))\n        d2   = self.dec2(torch.cat([self.up2(d3), e2], dim=1))\n        d1   = self.dec1(torch.cat([self.up1(d2), e1], dim=1))\n        return self.head(d1)                   # (B,1,H,W) logits\n\n\ndef build_segmenter(n_ch=N_CH_FINAL):\n    \"\"\"Unet++ / EfficientNet-B4 / SCSE on 17 channels.\"\"\"\n    return smp.UnetPlusPlus(\n        encoder_name           = 'efficientnet-b4',\n        encoder_weights        = 'imagenet',\n        in_channels            = n_ch,\n        classes                = 1,\n        decoder_attention_type = 'scse',\n    ).to(DEVICE)\n\n\nclass TwoStageInkNet(nn.Module):\n    \"\"\"\n    Combined model:\n      denoiser  : InkDenoiser  (16 → 1)\n      segmenter : Unet++       (16+1=17 → 1)\n\n    forward returns (denoiser_logits, segmenter_logits)\n    both shape (B,1,H,W).\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.denoiser  = InkDenoiser(N_ENHANCED).to(DEVICE)\n        self.segmenter = build_segmenter(N_CH_FINAL)\n\n    def forward(self, x):\n        # x: (B, N_ENHANCED, H, W)\n        den_logits = self.denoiser(x)                    # (B,1,H,W)\n        den_prob   = torch.sigmoid(den_logits)           # soft map\n        seg_input  = torch.cat([x, den_prob], dim=1)    # (B,17,H,W)\n        seg_logits = self.segmenter(seg_input)           # (B,1,H,W)\n        return den_logits, seg_logits\n\n\ndef build_model():\n    return TwoStageInkNet()\n\n\n# ════════════════════════════════════════════════════════════\n#  7. LOSS FUNCTIONS\n# ════════════════════════════════════════════════════════════\ndef tversky_loss(pred, target, alpha=0.3, beta=0.7, smooth=1.):\n    p  = torch.sigmoid(pred)\n    tp = (p*target).sum(dim=(2,3))\n    fp = (p*(1-target)).sum(dim=(2,3))\n    fn = ((1-p)*target).sum(dim=(2,3))\n    return 1.-((tp+smooth)/(tp+alpha*fp+beta*fn+smooth)).mean()\n\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0, eps=0.05):\n    t_s = target*(1-eps)+0.5*eps\n    bce = F.binary_cross_entropy_with_logits(pred, t_s, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = t_s*alpha+(1-t_s)*(1-alpha)\n    return (a_t*((1-p_t)**gamma)*bce).mean()\n\ndef seg_loss(pred, target):\n    return 0.5*tversky_loss(pred, target) + 0.5*focal_loss(pred, target)\n\ndef denoise_loss(pred, target):\n    \"\"\"BCE-based denoising loss. Target = ink mask, not input.\"\"\"\n    return F.binary_cross_entropy_with_logits(pred, target)\n\ndef combined_loss(den_logits, seg_logits, target):\n    \"\"\"\n    Combined: denoiser teaches itself to clean ink signal,\n    segmenter learns the final boundary decision.\n    \"\"\"\n    ld = denoise_loss(den_logits, target)\n    ls = seg_loss(seg_logits, target)\n    return W_DENOISE*ld + W_SEGMENT*ls, ld.item(), ls.item()\n\n\n# ════════════════════════════════════════════════════════════\n#  8. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p*masks).sum(dim=(1,2,3))\n    union = p.sum(dim=(1,2,3))+masks.sum(dim=(1,2,3))\n    return ((2.*inter+1e-5)/(union+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.25, 0.85, 0.01):\n        p = (probs>t).astype(np.float32)\n        d = (2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d > best_d: best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  9. GAUSSIAN PATCH WEIGHT\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c=sz//2; sig=sz//4\n    ys,xs=np.mgrid[0:sz,0:sz]\n    return np.exp(-((xs-c)**2+(ys-c)**2)/(2*sig**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  10. TRAINING LOOP\n# ════════════════════════════════════════════════════════════\ndef run_training(model, train_dl, val_dl,\n                 n_epochs, lr, swa_start, ckpt_name):\n\n    optimizer  = optim.AdamW(model.parameters(), lr=lr,\n                              weight_decay=WEIGHT_DECAY)\n    scheduler  = optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=10, T_mult=1, eta_min=1e-6)\n    scaler     = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n    swa_model  = AveragedModel(model)\n    swa_sched  = SWALR(optimizer, swa_lr=lr*0.1, anneal_epochs=5)\n    swa_active = False\n\n    best_dice = 0.; pat_cnt = 0\n    history   = dict(tl=[], vl=[], td=[], vd=[], dl=[], sl=[])\n\n    for epoch in range(n_epochs):\n        if epoch >= swa_start and not swa_active:\n            swa_active = True\n            print(f'  → SWA at epoch {epoch+1}')\n\n        # ── TRAIN ─────────────────────────────────────────────\n        model.train(); tl = td = mean_dl = mean_sl = 0.\n        optimizer.zero_grad()\n        for step, (imgs, msks) in enumerate(\n                tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            with torch.amp.autocast('cuda', enabled=USE_AMP):\n                den_l, seg_l = model(imgs)\n                loss, dl, sl = combined_loss(den_l, seg_l, msks)\n                loss = loss / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            if (step+1) % GRAD_ACCUM == 0 or (step+1) == len(train_dl):\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                scaler.step(optimizer); scaler.update()\n                optimizer.zero_grad()\n            tl      += loss.item() * GRAD_ACCUM\n            td      += batch_dice(seg_l.detach(), msks)\n            mean_dl += dl; mean_sl += sl\n            del imgs, msks, den_l, seg_l, loss\n            if step % 200 == 0 and DEVICE=='cuda':\n                torch.cuda.empty_cache()\n        n = len(train_dl)\n        tl /= n; td /= n; mean_dl /= n; mean_sl /= n\n\n        # ── VALIDATE ──────────────────────────────────────────\n        model.eval(); vl = vd = 0.\n        acc_p, acc_m = [], []\n        with torch.no_grad():\n            for imgs, msks in tqdm(val_dl,\n                                   desc=f'Ep{epoch+1:02d} val  ',\n                                   leave=False):\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                msks = msks.to(DEVICE, non_blocking=True)\n                with torch.amp.autocast('cuda', enabled=USE_AMP):\n                    _, seg_l = model(imgs)\n                    loss = seg_loss(seg_l, msks)\n                vl += loss.item(); vd += batch_dice(seg_l, msks)\n                acc_p.append(torch.sigmoid(seg_l).cpu().numpy())\n                acc_m.append(msks.cpu().numpy())\n                del imgs, msks, seg_l, loss\n        vl /= len(val_dl); vd /= len(val_dl)\n\n        probs_all = np.concatenate(acc_p)\n        masks_all = np.concatenate(acc_m)\n        bt, bd    = sweep_threshold(probs_all, masks_all)\n        ink_mean  = float(probs_all[masks_all>0.5].mean()) \\\n                    if (masks_all>0.5).any() else 0.\n        noink_mean= float(probs_all[masks_all<0.5].mean()) \\\n                    if (masks_all<0.5).any() else 0.\n        del acc_p, acc_m, probs_all, masks_all\n        gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n        if swa_active:\n            swa_model.update_parameters(model); swa_sched.step()\n        else:\n            scheduler.step()\n        lr_now = optimizer.param_groups[0]['lr']\n\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n        history['dl'].append(mean_dl); history['sl'].append(mean_sl)\n\n        print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n              f'loss={tl:.4f}(den={mean_dl:.3f} seg={mean_sl:.3f}) '\n              f'dice={td:.4f} | '\n              f'val={vl:.4f} dice={vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | '\n              f'sep={ink_mean-noink_mean:+.3f}'\n              + (' [SWA]' if swa_active else ''))\n\n        save_metric = max(vd, bd)\n        if save_metric > best_dice:\n            best_dice = save_metric; pat_cnt = 0\n            torch.save({'epoch': epoch, 'state': model.state_dict(),\n                        'thr': bt, 'dice': best_dice,\n                        'n_enhanced': N_ENHANCED,\n                        'n_ch_final': N_CH_FINAL},\n                       OUTPUT + ckpt_name)\n            print(f'  ✓ saved  (metric={best_dice:.4f})')\n        else:\n            pat_cnt += 1\n            if pat_cnt >= PATIENCE:\n                print(f'  ⚑ early stop ep{epoch+1}'); break\n\n        if epoch == 4 and max(history['vd']) < 0.30:\n            print('\\n⚠  val dice <0.30 after 5 epochs.')\n\n    if swa_active:\n        print('  Updating SWA BN stats ...')\n        swa_model.train()\n        with torch.no_grad():\n            for imgs, msks in train_dl:\n                swa_model(imgs.to(DEVICE, non_blocking=True))\n                del imgs, msks\n        model.load_state_dict(swa_model.module.state_dict())\n        del swa_model; gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    return model, best_dice, history\n\n\n# ════════════════════════════════════════════════════════════\n#  11. INFERENCE  (single pass, segmenter output only)\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    model.eval()\n    key   = (frag_path, tuple(z_list))\n    cache = _cache_store.get(key) or load_ct_cache(frag_path, z_list)\n    msk   = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n    H, W  = msk.shape\n\n    pred_map = np.zeros((H,W), np.float32)\n    wgt_map  = np.zeros((H,W), np.float32)\n    coords   = [(y,x)\n                for y in range(0, H-PATCH_SIZE+1, STRIDE_INF)\n                for x in range(0, W-PATCH_SIZE+1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y,x) in tqdm(coords, desc='Inference', leave=True):\n            raw_slices = [\n                cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                for z in z_list\n            ]\n            raw_vol  = np.stack(raw_slices, -1)\n            enhanced = physics_enhance(raw_vol)\n            t = norm_patch(\n                    torch.from_numpy(enhanced).permute(2,0,1).float()\n                ).unsqueeze(0).to(DEVICE)\n            with torch.amp.autocast('cuda', enabled=USE_AMP):\n                _, seg_l = model(t)\n                p = torch.sigmoid(seg_l).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p*GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, seg_l, p\n\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    return pred_map/(wgt_map+1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  12. EVALUATION + PLOTS\n# ════════════════════════════════════════════════════════════\ndef evaluate_and_plot(prob_map, gt_mask, label, fname):\n    H, W = gt_mask.shape\n    pc   = prob_map[:H,:W]\n    if np.isnan(pc).any():\n        pc = np.nan_to_num(pc, nan=0.5)\n\n    bt, _ = sweep_threshold(pc[np.newaxis,np.newaxis],\n                             gt_mask[np.newaxis,np.newaxis])\n    pred   = (pc > bt).astype(np.uint8)\n    inter  = (pred*gt_mask).sum()\n    dice   = (2*inter+1)/(pred.sum()+gt_mask.sum()+1)\n    pf = pred.flatten().astype(int)\n    mf = gt_mask.flatten().astype(int)\n    tn,fp,fn,tp_v = confusion_matrix(mf,pf,labels=[0,1]).ravel()\n    prec  = tp_v/(tp_v+fp+1e-8); rec = tp_v/(tp_v+fn+1e-8)\n    f1    = 2*prec*rec/(prec+rec+1e-8)\n    ink_m = float(pc[gt_mask==1].mean()) if (gt_mask==1).any() else 0.\n    bg_m  = float(pc[gt_mask==0].mean()) if (gt_mask==0).any() else 0.\n\n    print(f'\\n{\"=\"*55}')\n    print(f'RESULTS — {label}')\n    print(f'{\"=\"*55}')\n    print(f'Dice        : {dice:.4f}')\n    print(f'Threshold   : {bt:.2f}')\n    print(f'Precision   : {prec:.4f}   Recall : {rec:.4f}   F1 : {f1:.4f}')\n    print(f'TP={tp_v}  TN={tn}  FP={fp}  FN={fn}')\n    print(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}  (target <1.5)')\n    print(f'Calibration : ink={ink_m:.3f}  bg={bg_m:.3f}  '\n          f'sep={ink_m-bg_m:+.3f}')\n    print(f'{\"=\"*55}')\n\n    fig, ax = plt.subplots(2,3,figsize=(18,12))\n    ax[0,0].imshow(gt_mask, cmap='gray');\n    ax[0,0].set_title('Ground Truth')\n    ax[0,1].imshow(pc, cmap='inferno');\n    ax[0,1].set_title('Probability Map (Segmenter)')\n    ax[0,2].imshow(pred, cmap='gray');\n    ax[0,2].set_title(f'Prediction  Dice={dice:.4f}')\n    err = np.zeros((*gt_mask.shape,3),dtype=np.uint8)\n    err[(pred==1)&(gt_mask==1)]=[0,255,0]\n    err[(pred==1)&(gt_mask==0)]=[255,0,0]\n    err[(pred==0)&(gt_mask==1)]=[0,0,255]\n    ax[1,0].imshow(err)\n    ax[1,0].set_title('TP=green  FP=red  FN=blue')\n    ink_v = pc[gt_mask==1].ravel(); ink_v=ink_v[np.isfinite(ink_v)]\n    bg_v  = pc[gt_mask==0].ravel(); bg_v =bg_v [np.isfinite(bg_v )]\n    if len(ink_v): ax[1,1].hist(ink_v,bins=50,alpha=0.7,\n                                label=f'ink μ={ink_m:.2f}',\n                                color='orange',density=True)\n    if len(bg_v):  ax[1,1].hist(bg_v, bins=50,alpha=0.7,\n                                label=f'bg  μ={bg_m:.2f}',\n                                color='blue',density=True)\n    ax[1,1].axvline(bt,color='r',ls='--',label=f'thr={bt:.2f}')\n    ax[1,1].set_title('Probability Distribution')\n    ax[1,1].legend(); ax[1,1].set_xlabel('Probability')\n    ts=np.arange(0.20,0.85,0.01); ds=[]\n    for t in ts:\n        p=(pc>t).astype(np.float32)\n        ds.append((2*(p*gt_mask).sum()+1)/(p.sum()+gt_mask.sum()+1))\n    ax[1,2].plot(ts,ds)\n    ax[1,2].axvline(bt,color='r',ls='--',label=f'best={bt:.2f}')\n    ax[1,2].axhline(0.65,color='orange',ls=':',label='0.65')\n    ax[1,2].axhline(0.80,color='g',ls=':',label='0.80')\n    ax[1,2].set_title('Dice vs Threshold')\n    ax[1,2].set_xlabel('Threshold'); ax[1,2].set_ylabel('Dice')\n    ax[1,2].legend(); ax[1,2].grid(True)\n    for a in [ax[0,0],ax[0,1],ax[0,2],ax[1,0]]: a.axis('off')\n    plt.suptitle(f'{label} — Dice={dice:.4f}', fontsize=13, y=1.01)\n    plt.tight_layout()\n    plt.savefig(OUTPUT+fname, dpi=100, bbox_inches='tight')\n    plt.close()\n    return dict(dice=dice, thr=bt, prec=prec, rec=rec, f1=f1,\n                tp=tp_v, tn=tn, fp=fp, fn=fn,\n                ink_mean=ink_m, noink_mean=bg_m)\n\n\n# ════════════════════════════════════════════════════════════\n#  13. BUILD DATASETS  (same protocol as v19)\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('PHYSICS-BASED INK ENHANCEMENT PIPELINE')\nprint(f'  Raw CT slices  : {N_RAW} (z={Z_SLICES[0]}-{Z_SLICES[-1]})')\nprint(f'  Enhanced ch    : {N_ENHANCED}')\nprint(f'    [0-3]  raw key slices  (z=27,29,31,33)')\nprint(f'    [4-7]  z-forward derivative (ink onset)')\nprint(f'    [8-9]  z-second derivative (ink boundaries)')\nprint(f'    [10-11] median residual (background removed)')\nprint(f'    [12]   max inter-slice contrast')\nprint(f'    [13]   z-axis std (ink = high variance)')\nprint(f'    [14]   cumulative positive diffs (ink accumulation)')\nprint(f'    [15]   total z-curvature')\nprint(f'  Denoiser output: +1 ch  →  total {N_CH_FINAL} ch to Unet++')\nprint('DATA PROTOCOL')\nprint('  Train    : Frag2 (all) + Frag3 (80%)')\nprint('  Validate : Frag3 (20% held-out)')\nprint('  Test     : Frag1 — labels seen ONLY at evaluation')\nprint('='*60)\n\nprint('\\n── Fragment 2  [full training set] ──')\nds_f2 = VesuviusDataset(\n    FRAG2, Z_SLICES, stride=STRIDE_TR,\n    transform=train_tf, neg_ratio=NEG_RATIO,\n    max_patches=MAX_PATCHES_TR)\n\nprint('\\n── Fragment 3  [80% train + 20% val] ──')\nds_f3_all = VesuviusDataset(\n    FRAG3, Z_SLICES, stride=STRIDE_TR,\n    transform=train_tf, neg_ratio=NEG_RATIO,\n    max_patches=0)\n\nn_f3    = len(ds_f3_all)\nn_val   = int(0.20 * n_f3)\nperm    = np.random.permutation(n_f3)\ntrn_idx = perm[n_val:]\nval_idx = perm[:n_val]\nds_f3_train = Subset(ds_f3_all, trn_idx)\nds_val      = Subset(ds_f3_all, val_idx)\nw_f3_train  = ds_f3_all.weights[trn_idx]\nprint(f'  Frag3 split: {len(trn_idx)} train | {len(val_idx)} val')\nprint('\\n  Fragment 1: reserved for final test')\n\ntrain_ds = ConcatDataset([ds_f2, ds_f3_train])\nw_train  = np.concatenate([ds_f2.weights, w_f3_train])\n\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n_dl_kw = dict(num_workers=NUM_WORKERS, pin_memory=PIN,\n              persistent_workers=(NUM_WORKERS>0),\n              prefetch_factor=2 if NUM_WORKERS>0 else None)\n\ntrain_sampler = WeightedRandomSampler(\n    torch.from_numpy(w_train), len(train_ds), replacement=True)\ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE,\n                      sampler=train_sampler, **_dl_kw)\nval_dl   = DataLoader(ds_val,   batch_size=BATCH_SIZE,\n                      shuffle=False, **_dl_kw)\nprint(f'\\nTrain batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  14. TRAIN\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint(f'TRAINING  ({EPOCHS} epochs)')\nprint(f'  Stage A: InkDenoiser  ({N_ENHANCED}→1 ch, BCE vs ink mask)')\nprint(f'  Stage B: Unet++/B4    ({N_CH_FINAL}→1 ch, Tversky+Focal)')\nprint(f'  Joint loss: {W_DENOISE}·Denoiser + {W_SEGMENT}·Segmenter')\nprint(f'  SWA from ep {SWA_START}')\nprint('='*60)\n\nmodel = build_model()\nmodel, best_dice, history = run_training(\n    model, train_dl, val_dl,\n    n_epochs=EPOCHS, lr=LR,\n    swa_start=SWA_START,\n    ckpt_name='best_model.pth')\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\ndel train_dl, val_dl, train_sampler\ndel train_ds, ds_val, ds_f2, ds_f3_all, ds_f3_train\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n\n# ════════════════════════════════════════════════════════════\n#  15. TRAINING CURVES\n# ════════════════════════════════════════════════════════════\nfig, axes = plt.subplots(1,3,figsize=(18,5))\naxes[0].plot(history['tl'],label='total loss')\naxes[0].plot(history['vl'],label='val seg loss')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'],label='train dice')\naxes[1].plot(history['vd'],label='val dice')\naxes[1].axhline(0.65,color='orange',ls='--',label='0.65')\naxes[1].axhline(0.80,color='r',ls='--',label='0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['dl'],label='denoiser loss')\naxes[2].plot(history['sl'],label='segmenter loss')\naxes[2].set_title('Stage A vs B Loss')\naxes[2].legend(); axes[2].grid(True)\nplt.suptitle('v21 Physics-Enhanced Two-Stage Training', fontsize=12)\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves.png',dpi=100); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  16. FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('FINAL TEST — FRAGMENT 1')\nprint('  Physics-enhanced channels — same pipeline as training')\nprint('  Labels loaded here for the FIRST TIME')\nprint('='*60)\n\nckpt = torch.load(OUTPUT+'best_model.pth',\n                  map_location=DEVICE, weights_only=False)\nprint(f'Checkpoint: epoch={ckpt[\"epoch\"]+1}  '\n      f'val_dice={ckpt[\"dice\"]:.4f}  thr={ckpt[\"thr\"]:.2f}')\n\nfinal_model = build_model()\nfinal_model.load_state_dict(ckpt['state'], strict=True)\n\nprob_map, msk1 = predict_fragment(final_model, FRAG1, Z_SLICES)\nres = evaluate_and_plot(\n    prob_map, msk1,\n    label='Fragment 1 — Physics-Enhanced (first seen test)',\n    fname='frag1_prediction.png')\n\nprint(f'\\n★  FRAG1 DICE   : {res[\"dice\"]:.4f}')\nprint(f'★  Precision    : {res[\"prec\"]:.4f}')\nprint(f'★  Recall       : {res[\"rec\"]:.4f}')\nprint(f'★  F1           : {res[\"f1\"]:.4f}')\nprint(f'★  FP/TP        : {res[\"fp\"]/(res[\"tp\"]+1e-8):.2f}')\n\n\n# ════════════════════════════════════════════════════════════\n#  17. SAVE RESULTS\n# ════════════════════════════════════════════════════════════\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write('VESUVIUS INK DETECTION — v21  PHYSICS-ENHANCED\\n')\n    f.write('='*55+'\\n')\n    f.write(f'Raw CT slices : {N_RAW}  (z={Z_SLICES})\\n')\n    f.write(f'Enhanced ch   : {N_ENHANCED}\\n')\n    f.write(f'  [0-3]  raw key slices z=27,29,31,33\\n')\n    f.write(f'  [4-7]  forward z-derivative\\n')\n    f.write(f'  [8-9]  second z-derivative\\n')\n    f.write(f'  [10-11] median residual (5×5)\\n')\n    f.write(f'  [12]   max inter-slice contrast\\n')\n    f.write(f'  [13]   z-axis std\\n')\n    f.write(f'  [14]   cumulative positive diffs\\n')\n    f.write(f'  [15]   total z-curvature\\n')\n    f.write(f'Stage A  : InkDenoiser ({N_ENHANCED}→1, BCE)\\n')\n    f.write(f'Stage B  : Unet++/B4 ({N_CH_FINAL}→1, Tversky+Focal)\\n')\n    f.write(f'Joint    : {W_DENOISE}·denoise + {W_SEGMENT}·segment\\n')\n    f.write(f'Train    : Frag2+Frag3(80%)\\n')\n    f.write(f'Val      : Frag3(20%)\\n')\n    f.write(f'Test     : Frag1 (never seen)\\n')\n    f.write(f'Val best : {best_dice:.4f}\\n')\n    f.write('='*55+'\\n')\n    f.write(f'FRAG1 DICE  : {res[\"dice\"]:.4f}\\n')\n    f.write(f'Threshold   : {res[\"thr\"]:.2f}\\n')\n    f.write(f'Precision   : {res[\"prec\"]:.4f}\\n')\n    f.write(f'Recall      : {res[\"rec\"]:.4f}\\n')\n    f.write(f'F1          : {res[\"f1\"]:.4f}\\n')\n    f.write(f'TP={res[\"tp\"]}  TN={res[\"tn\"]}  '\n            f'FP={res[\"fp\"]}  FN={res[\"fn\"]}\\n')\n    f.write(f'FP/TP       : {res[\"fp\"]/(res[\"tp\"]+1e-8):.2f}\\n')\n    f.write(f'Calibration : ink={res[\"ink_mean\"]:.3f}  '\n            f'bg={res[\"noink_mean\"]:.3f}  '\n            f'sep={res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}\\n')\n\nprint(f'\\nOutputs → {OUTPUT}')\nprint('  best_model.pth | curves.png')\nprint('  frag1_prediction.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\n# ============================================================\n#  VESUVIUS v24 — DATA DIAGNOSTIC + FIXED TRAINING\n#\n#  DIAGNOSIS FIRST: Before any training, we verify:\n#  1. Which z-slices actually correlate with ink labels\n#  2. Label coverage and ink pixel percentage\n#  3. CT signal distribution at labeled ink vs background\n#\n#  Then train with the CORRECT z-slices.\n# ============================================================\n\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import (Dataset, DataLoader,\n                               WeightedRandomSampler, ConcatDataset, Subset)\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom scipy.ndimage import median_filter, label as cc_label\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nimport warnings\nwarnings.filterwarnings('ignore')\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:32\"\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = False\n\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU   : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM  : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  STEP 1: FIND BEST Z-SLICES BY CORRELATION WITH INK LABELS\n#\n#  For each fragment and each z-slice, compute:\n#  - mean CT value at ink pixels\n#  - mean CT value at non-ink pixels\n#  - separation = ink_mean - noink_mean\n#  The slices with highest separation are the ones to use.\n# ════════════════════════════════════════════════════════════\ndef find_best_z_slices(frag_path, z_range=range(0, 65), n_best=32,\n                       sample_frac=0.02):\n    \"\"\"\n    Scan all available z-slices and rank by ink/bg separation.\n    Returns sorted list of best z-indices.\n    \"\"\"\n    vol_dir  = os.path.join(frag_path, 'surface_volume')\n    msk_path = os.path.join(frag_path, 'inklabels.png')\n\n    msk = cv2.imread(msk_path, 0)\n    if msk is None:\n        print(f\"  ✗ No inklabels.png at {msk_path}\")\n        return list(z_range)[:n_best]\n    ink_mask = (msk > 0).astype(np.float32)\n\n    # Check mask coverage\n    ink_pct = ink_mask.mean() * 100\n    print(f\"  [{os.path.basename(frag_path)}] Ink coverage: {ink_pct:.2f}% \"\n          f\"of {msk.shape} mask\")\n    if ink_pct < 0.1:\n        print(\"  ⚠ WARNING: very low ink coverage — check inklabels.png\")\n    if ink_pct > 50:\n        print(\"  ⚠ WARNING: very high ink coverage — may be wrong file\")\n\n    seps = []\n    available = []\n    for z in z_range:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if not os.path.exists(p):\n            continue\n        s = tifffile.imread(p).astype(np.float32)\n\n        # Resize mask if needed\n        if s.shape != ink_mask.shape:\n            m = cv2.resize(ink_mask, (s.shape[1], s.shape[0]),\n                           interpolation=cv2.INTER_NEAREST)\n        else:\n            m = ink_mask\n\n        # Subsample for speed\n        flat_s = s.ravel()[::int(1/sample_frac)]\n        flat_m = m.ravel()[::int(1/sample_frac)]\n\n        ink_mean = flat_s[flat_m > 0.5].mean() if (flat_m>0.5).any() else 0\n        bg_mean  = flat_s[flat_m < 0.5].mean() if (flat_m<0.5).any() else 0\n        sep      = ink_mean - bg_mean\n        seps.append(sep); available.append(z)\n        del s\n\n    if not seps:\n        print(f\"  ✗ No tif files found in {vol_dir}\")\n        return []\n\n    seps = np.array(seps)\n    order = np.argsort(np.abs(seps))[::-1]   # sort by |sep|\n    best_z = [available[i] for i in order[:n_best]]\n    best_z.sort()   # restore z-order for temporal coherence\n\n    print(f\"  Available z-slices: {min(available)}-{max(available)} \"\n          f\"({len(available)} total)\")\n    print(f\"  Top-5 by |sep|: \"\n          + \", \".join([f\"z{available[i]}(sep={seps[i]:+.1f})\"\n                       for i in order[:5]]))\n    print(f\"  Selected z-range: {best_z[0]}-{best_z[-1]} \"\n          f\"({len(best_z)} slices)\")\n    print(f\"  Best sep at z{available[order[0]]}: {seps[order[0]]:+.1f}\")\n\n    if abs(seps[order[0]]) < 10:\n        print(\"  ⚠ WARNING: max separation < 10 — ink signal very weak\")\n        print(\"    Check that inklabels.png aligns with surface_volume/\")\n\n    gc.collect()\n    return best_z, seps, available\n\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"STEP 1: Z-SLICE ANALYSIS\")\nprint(\"=\"*60)\n\nbest_z2, seps2, avail2 = find_best_z_slices(FRAG2, z_range=range(0,65), n_best=32)\nprint()\nbest_z3, seps3, avail3 = find_best_z_slices(FRAG3, z_range=range(0,65), n_best=32)\n\n# Find overlap of best z-slices from both fragments\n# Use intersection of top-32 from each, fall back to union if needed\nset2 = set(best_z2); set3 = set(best_z3)\noverlap = sorted(set2 & set3)\nif len(overlap) >= 16:\n    Z_SLICES_RAW = overlap[:32]\n    print(f\"\\n✓ Using {len(Z_SLICES_RAW)} z-slices common to both fragments: \"\n          f\"{Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]}\")\nelse:\n    # Fall back: use Frag2 best (more data)\n    Z_SLICES_RAW = best_z2[:32]\n    print(f\"\\n⚠ Limited overlap — using Frag2 best {len(Z_SLICES_RAW)} slices\")\n\n# Plot separation curves\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nfor ax, seps, avail, frag in [(axes[0],seps2,avail2,'Frag2'),\n                               (axes[1],seps3,avail3,'Frag3')]:\n    ax.bar(avail, seps, color=['green' if z in Z_SLICES_RAW else 'gray'\n                               for z in avail])\n    ax.axhline(0, color='k', lw=0.5)\n    ax.set_xlabel('Z-slice index'); ax.set_ylabel('Ink-BG separation')\n    ax.set_title(f'{frag}: CT separation by z-slice\\n'\n                 f'(green = selected)')\n    ax.grid(True, alpha=0.3)\nplt.tight_layout()\nplt.savefig(OUTPUT+'z_slice_analysis.png', dpi=80)\nplt.close()\nprint(\"Saved: z_slice_analysis.png\")\n\n\n# ════════════════════════════════════════════════════════════\n#  STEP 2: VERIFY LABEL ALIGNMENT\n#  Visual check: overlay ink labels on CT mid-slice\n# ════════════════════════════════════════════════════════════\nprint(\"\\n\" + \"=\"*60)\nprint(\"STEP 2: VISUAL VERIFICATION (label/CT alignment)\")\nprint(\"=\"*60)\n\ndef verify_alignment(frag_path, z_list, tag):\n    vol_dir  = os.path.join(frag_path, 'surface_volume')\n    msk_path = os.path.join(frag_path, 'inklabels.png')\n    msk = cv2.imread(msk_path, 0)\n    if msk is None:\n        print(f\"  ✗ No mask for {tag}\"); return\n\n    # Use the best z-slice (highest |sep|)\n    z_mid = z_list[len(z_list)//2]\n    p = os.path.join(vol_dir, f'{z_mid:02d}.tif')\n    if not os.path.exists(p):\n        print(f\"  ✗ No slice {z_mid} for {tag}\"); return\n\n    ct = tifffile.imread(p).astype(np.float32)\n    p5, p95 = np.percentile(ct, [5, 95])\n    ct_norm = np.clip((ct - p5)/(p95 - p5), 0, 1)\n\n    if msk.shape != ct_norm.shape:\n        msk_r = cv2.resize(msk, (ct_norm.shape[1], ct_norm.shape[0]),\n                           interpolation=cv2.INTER_NEAREST)\n    else:\n        msk_r = msk\n\n    # Sample a 512×512 crop from the middle\n    H, W  = ct_norm.shape\n    y0, x0 = H//2 - 256, W//2 - 256\n    y0 = max(0, y0); x0 = max(0, x0)\n    y1 = min(H, y0+512); x1 = min(W, x0+512)\n\n    ct_crop  = ct_norm[y0:y1, x0:x1]\n    msk_crop = msk_r[y0:y1, x0:x1]\n\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    axes[0].imshow(ct_crop, cmap='gray')\n    axes[0].set_title(f'{tag} CT z={z_mid} (centre crop)')\n    axes[1].imshow(msk_crop, cmap='gray')\n    axes[1].set_title(f'{tag} ink labels')\n    overlay = np.stack([ct_crop]*3, axis=-1)\n    overlay[msk_crop>0] = [1., 0.2, 0.2]\n    axes[2].imshow(overlay)\n    axes[2].set_title(f'{tag} overlay')\n    for a in axes: a.axis('off')\n    plt.suptitle(f'{tag}: if labels look misaligned, check data path')\n    plt.tight_layout()\n    plt.savefig(OUTPUT+f'verify_{tag}.png', dpi=80)\n    plt.close()\n    print(f\"  Saved verify_{tag}.png — CHECK THIS before proceeding\")\n    del ct, ct_norm\n\nfor frag, tag in [(FRAG2,'frag2'), (FRAG3,'frag3')]:\n    verify_alignment(frag, Z_SLICES_RAW, tag)\n\n\n# ════════════════════════════════════════════════════════════\n#  STEP 3: TRAINING WITH VERIFIED Z-SLICES\n# ════════════════════════════════════════════════════════════\n\n# Now set all training hyperparams using the verified Z_SLICES_RAW\nN_RAW        = len(Z_SLICES_RAW)\nN_ENGINEERED = 6\nN_CH         = N_RAW + N_ENGINEERED\n\nPATCH_SIZE   = 96\nSTRIDE_TR    = 48\nSTRIDE_INF   = 32\nBATCH_SIZE   = 8\nGRAD_ACCUM   = 2\nEPOCHS_SEG   = 15\nLR_SEG       = 2e-4\nLR_ENCODER   = 2e-5\nWEIGHT_DECAY = 1e-4\nPATIENCE     = 8\nMAX_PATCHES  = 12_000\nCC_MIN_PIXELS= 50\nINK_MIN_POS  = 0.03    # slightly higher than before\nNEG_RATIO    = 0.5\nPOS_WEIGHT   = 4.0\nUSE_AMP      = (DEVICE == 'cuda')\nNORM_STD_MIN = 0.01\nNORM_CLIP    = 10.0\n\nprint(f\"\\nZ_SLICES_RAW = {Z_SLICES_RAW}\")\nprint(f\"N_RAW={N_RAW}, N_CH={N_CH}, PATCH_SIZE={PATCH_SIZE}\")\n\n\n# ════════════════════════════════════════════════════════════\n#  CT CACHE\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_ct_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        return _cache_store[key]\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    slices, vals = [], []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            slices.append((z, s)); vals.append(s.ravel()[::50])\n    assert slices, f\"No slices in {vol_dir}\"\n    p1, p99 = np.percentile(np.concatenate(vals), [1, 99])\n    del vals; gc.collect()\n    cache = {}; H = W = None\n    for z, s in slices:\n        s = ((np.clip(s,p1,p99)-p1)/(p99-p1+1e-6)).astype(np.float16)\n        cache[z] = s\n        if H is None: H, W = s.shape\n    del slices; gc.collect()\n    mb = sum(v.nbytes for v in cache.values())/1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: {len(cache)} slices, {mb:.0f} MB\")\n    _cache_store[key] = cache\n    return cache\n\ndef get_hw(cache): return next(iter(cache.values())).shape\n\ndef norm_patch(t):\n    mu  = t.mean(dim=(1,2), keepdim=True)\n    std = t.std (dim=(1,2), keepdim=True).clamp(min=NORM_STD_MIN)\n    return ((t-mu)/std).clamp(-NORM_CLIP, NORM_CLIP)\n\n\n# ════════════════════════════════════════════════════════════\n#  PHYSICS CHANNELS\n# ════════════════════════════════════════════════════════════\ndef physics_enhance(raw_vol):\n    ch0 = raw_vol.max(axis=2)\n    ch1 = raw_vol.std(axis=2)\n    ch2 = np.abs(np.diff(raw_vol, axis=2)).max(axis=2)\n    ch3 = np.maximum(0, np.diff(raw_vol, axis=2)).sum(axis=2)\n    mid = raw_vol.shape[2]//2\n    ch4 = raw_vol[:,:,mid] - median_filter(raw_vol[:,:,mid], size=3)\n    ch5 = (np.percentile(raw_vol,75,axis=2) -\n           np.percentile(raw_vol,25,axis=2))\n    return np.stack([ch0,ch1,ch2,ch3,ch4,ch5], axis=-1).astype(np.float32)\n\ndef build_input(raw_vol):\n    eng = physics_enhance(raw_vol)\n    out = np.concatenate([raw_vol, eng], axis=-1); del eng\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  PAPYRUS MASK\n# ════════════════════════════════════════════════════════════\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H,W):\n                m = cv2.resize(m,(W,H),interpolation=cv2.INTER_NEAREST)\n            return (m>0).astype(np.uint8)\n    z = list(cache.keys())[len(cache)//2]\n    return (cache[z]>0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Affine(scale=(0.85,1.15),\n             translate_percent={'x':(-0.1,0.1),'y':(-0.1,0.1)},\n             rotate=(-30,30), mode=cv2.BORDER_REFLECT, p=0.6),\n    A.RandomBrightnessContrast(0.3, 0.3, p=0.6),\n    A.GaussianBlur(blur_limit=(3,7), p=0.3),\n    A.GaussNoise(var_limit=(0.001,0.01), p=0.3),\n    A.CoarseDropout(max_holes=6, max_height=24, max_width=24,\n                    fill_value=0, p=0.4),\n    A.ElasticTransform(alpha=20, sigma=4, p=0.3),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  DATASET — with per-fragment sep check\n# ════════════════════════════════════════════════════════════\nclass SegDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0., max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        self.tf     = transform\n        H, W        = get_hw(self.cache)\n\n        msk = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n        assert msk is not None\n        if msk.shape != (H, W):\n            msk = cv2.resize(msk, (W, H), interpolation=cv2.INTER_NEAREST)\n        self.mask = (msk>0).astype(np.uint8); del msk\n\n        # Per-dataset sep check: sample 1000 random pixels\n        z_mid = z_list[len(z_list)//2]\n        ct_sample = self.cache[z_mid].astype(np.float32)\n        ink_vals  = ct_sample[self.mask==1]\n        bg_vals   = ct_sample[self.mask==0]\n        if len(ink_vals) and len(bg_vals):\n            sep = ink_vals[::100].mean() - bg_vals[::100].mean()\n            print(f\"    Label-CT sep at z{z_mid}: {sep:+.4f}  \"\n                  f\"(|sep|>0.01 expected)\")\n            if abs(sep) < 0.005:\n                print(f\"    ⚠ VERY LOW SEP — ink labels may not align \"\n                      f\"with z-slice range\")\n\n        ink_dilated = cv2.dilate(\n            self.mask, np.ones((PATCH_SIZE*2,PATCH_SIZE*2),np.uint8))\n        pap = load_papyrus_mask(frag_path, H, W, self.cache)\n\n        pos_yx, hard_neg, easy_neg = [], [], []\n        for y in range(0, H-PATCH_SIZE+1, stride):\n            for x in range(0, W-PATCH_SIZE+1, stride):\n                if pap[y:y+PATCH_SIZE,x:x+PATCH_SIZE].mean() < 0.4: continue\n                ink = self.mask[y:y+PATCH_SIZE,x:x+PATCH_SIZE].mean()\n                if ink >= INK_MIN_POS:\n                    pos_yx.append((y,x))\n                elif ink < 0.001:\n                    if ink_dilated[y+PATCH_SIZE//2,x+PATCH_SIZE//2]:\n                        hard_neg.append((y,x))\n                    else:\n                        easy_neg.append((y,x))\n\n        n_neg  = int(len(pos_yx)*neg_ratio)\n        n_hard = int(n_neg*0.7); n_easy = n_neg-n_hard\n        np.random.shuffle(hard_neg); np.random.shuffle(easy_neg)\n        neg_yx = hard_neg[:n_hard] + easy_neg[:n_easy]\n\n        all_yx  = pos_yx+neg_yx\n        all_lbl = [1]*len(pos_yx)+[0]*len(neg_yx)\n        if max_patches>0 and len(all_yx)>max_patches:\n            frac=max_patches/len(all_yx)\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx=pos_yx[:max(1,int(len(pos_yx)*frac))]\n            neg_yx=neg_yx[:max(0,int(len(neg_yx)*frac))]\n            all_yx=pos_yx+neg_yx\n            all_lbl=[1]*len(pos_yx)+[0]*len(neg_yx)\n\n        perm=np.random.permutation(len(all_yx))\n        self.coords  = np.array([all_yx[i]  for i in perm],dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm],dtype=np.int32)\n        self.weights = np.where(self.labels==1,POS_WEIGHT,1.).astype(np.float32)\n        print(f\"  Seg [{os.path.basename(frag_path)}]: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg = {len(self.coords)}\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y,x  = int(self.coords[idx,0]), int(self.coords[idx,1])\n        raw  = np.stack([self.cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                         for z in self.z_list], axis=-1)\n        msk  = self.mask[y:y+PATCH_SIZE,x:x+PATCH_SIZE].copy()\n        if self.tf:\n            o=self.tf(image=raw,mask=msk); raw,msk=o['image'],o['mask']\n        full  = build_input(raw); del raw\n        img_t = norm_patch(torch.from_numpy(full).permute(2,0,1).float()); del full\n        return img_t, torch.from_numpy(msk).unsqueeze(0).float()\n\n\n# ════════════════════════════════════════════════════════════\n#  MODEL\n# ════════════════════════════════════════════════════════════\nclass StemUNetPP(nn.Module):\n    def __init__(self, in_ch, encoder_name='resnet34'):\n        super().__init__()\n        self.stem = nn.Sequential(\n            nn.Conv2d(in_ch, 16, 1, bias=False), nn.BatchNorm2d(16), nn.ReLU(inplace=True),\n            nn.Conv2d(16,     3, 1, bias=False), nn.BatchNorm2d(3),  nn.ReLU(inplace=True),\n        )\n        self.unetpp = smp.UnetPlusPlus(\n            encoder_name=encoder_name, encoder_weights='imagenet',\n            in_channels=3, classes=1, activation=None,\n            decoder_channels=(256,128,64,32,16),\n            decoder_use_batchnorm=True,\n        )\n\n    def forward(self, x): return self.unetpp(self.stem(x))\n\n    def encoder_parameters(self):\n        return list(self.stem.parameters()) + \\\n               list(self.unetpp.encoder.parameters())\n    def decoder_parameters(self):\n        return (list(self.unetpp.decoder.parameters()) +\n                list(self.unetpp.segmentation_head.parameters()))\n\n\ndef build_seg_model(n_ch):\n    for enc in ('resnet34','resnet18','efficientnet-b0'):\n        try:\n            m = StemUNetPP(in_ch=n_ch, encoder_name=enc)\n            with torch.no_grad():\n                _ = m(torch.zeros(1,n_ch,PATCH_SIZE,PATCH_SIZE))\n            print(f\"  Model: Stem({n_ch}→3) + UNet++({enc}, ImageNet)\")\n            return m.to(DEVICE)\n        except Exception as e:\n            print(f\"  {enc} failed ({e}), trying next...\")\n    raise RuntimeError(\"No encoder worked\")\n\n\n# ════════════════════════════════════════════════════════════\n#  LOSS\n# ════════════════════════════════════════════════════════════\ndef soft_dice_loss(logits, targets, smooth=1.):\n    p  = torch.sigmoid(logits)\n    tp = (p*targets).sum(dim=(2,3))\n    fp = p.sum(dim=(2,3)); fn = targets.sum(dim=(2,3))\n    return (1.-(2.*tp+smooth)/(fp+fn+smooth)).mean()\n\ndef seg_loss(logits, targets):\n    return (0.5*F.binary_cross_entropy_with_logits(logits,targets)\n            + 0.5*soft_dice_loss(logits,targets))\n\n\n# ════════════════════════════════════════════════════════════\n#  METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p=( torch.sigmoid(logits)>thr).float()\n    i=(p*masks).sum(dim=(1,2,3)); u=p.sum(dim=(1,2,3))+masks.sum(dim=(1,2,3))\n    return ((2.*i+1e-5)/(u+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t,best_d=0.5,0.\n    for t in np.arange(0.10,0.95,0.01):\n        p=(probs>t).astype(np.float32)\n        d=(2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d>best_d: best_d,best_t=d,float(t)\n    return best_t,best_d\n\ndef compute_sep(probs,masks):\n    if (masks>0.5).any() and (masks<0.5).any():\n        return float(probs[masks>0.5].mean()-probs[masks<0.5].mean())\n    return 0.\n\n\n# ════════════════════════════════════════════════════════════\n#  CC FILTER\n# ════════════════════════════════════════════════════════════\ndef cc_filter(binary_mask, min_px=CC_MIN_PIXELS):\n    labeled,n=cc_label(binary_mask)\n    if n==0: return binary_mask,0\n    sizes=np.bincount(labeled.ravel()); sizes[0]=0\n    keep=sizes>=min_px\n    return keep[labeled].astype(np.uint8),int((sizes[1:]>0).sum()-keep[1:].sum())\n\ndef gauss_weight(sz):\n    c=sz//2; s=sz//4; y,x=np.mgrid[0:sz,0:sz]\n    return np.exp(-((x-c)**2+(y-c)**2)/(2*s**2)).astype(np.float32)\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  TRAINING\n# ════════════════════════════════════════════════════════════\ndef run_seg(model, train_dl, val_dl, n_epochs, ckpt):\n    optimizer = optim.AdamW([\n        {'params': model.decoder_parameters(), 'lr': LR_SEG},\n        {'params': model.encoder_parameters(), 'lr': LR_ENCODER},\n    ], weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=n_epochs, eta_min=1e-6)\n    scaler    = GradScaler(enabled=USE_AMP)\n    best_dice=0.; pat=0\n    history=dict(tl=[],vl=[],td=[],vd=[],sep=[])\n\n    for ep in range(n_epochs):\n        model.train(); tl=td=0.; optimizer.zero_grad()\n        for step,(imgs,msks) in enumerate(\n                tqdm(train_dl,desc=f'Ep{ep+1:02d}▸train',leave=False)):\n            imgs=imgs.to(DEVICE); msks=msks.to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                out=model(imgs); loss=seg_loss(out,msks)/GRAD_ACCUM\n            scaler.scale(loss).backward()\n            tl+=loss.item()*GRAD_ACCUM; td+=batch_dice(out.detach(),msks)\n            del imgs,msks,out,loss\n            if (step+1)%GRAD_ACCUM==0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(),1.)\n                scaler.step(optimizer); scaler.update(); optimizer.zero_grad()\n                if DEVICE=='cuda': torch.cuda.empty_cache()\n        tl/=len(train_dl); td/=len(train_dl)\n\n        model.eval(); vl=vd=0.\n        all_p,all_m=[],[]\n        with torch.no_grad():\n            for imgs,msks in tqdm(val_dl,desc=f'Ep{ep+1:02d}▸val  ',leave=False):\n                imgs=imgs.to(DEVICE); msks=msks.to(DEVICE)\n                with autocast(enabled=USE_AMP):\n                    out=model(imgs); loss=seg_loss(out,msks)\n                vl+=loss.item(); vd+=batch_dice(out,msks)\n                all_p.append(torch.sigmoid(out).cpu().half().numpy())\n                all_m.append(msks.cpu().half().numpy())\n                del imgs,msks,out,loss\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n        vl/=len(val_dl); vd/=len(val_dl)\n        P=np.concatenate(all_p).astype(np.float32)\n        M=np.concatenate(all_m).astype(np.float32)\n        del all_p,all_m\n        bt,bd=sweep_threshold(P,M); sep=compute_sep(P,M)\n        del P,M; gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n        scheduler.step()\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n        history['sep'].append(sep)\n\n        dec_lr=optimizer.param_groups[0]['lr']\n        flag=' ⚠ LOW SEP' if ep>=2 and sep<0.1 else ''\n        print(f'Ep{ep+1:02d} | lr={dec_lr:.1e} | '\n              f'train={tl:.4f}/{td:.4f} | val={vl:.4f}/{vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | sep={sep:+.3f}{flag}')\n\n        if ep==2 and sep<0.02:\n            print('\\n⛔ sep<0.02 after ep3 — stopping.')\n            print('  Please check verify_frag2.png and verify_frag3.png')\n            print('  to confirm CT and labels are spatially aligned.')\n            break\n\n        metric=max(vd,bd)\n        if metric>best_dice:\n            best_dice=metric; pat=0\n            torch.save({'ep':ep,'state':model.state_dict(),\n                        'thr':bt,'dice':best_dice,'n_ch':N_CH},OUTPUT+ckpt)\n            print(f'  ✓ checkpoint (dice={best_dice:.4f})')\n        else:\n            pat+=1\n            if pat>=PATIENCE: print(f'  early stop ep{ep+1}'); break\n\n    return model,best_dice,history\n\n\n# ════════════════════════════════════════════════════════════\n#  INFERENCE\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    model.eval()\n    cache=_cache_store.get((frag_path,tuple(z_list))) or \\\n          load_ct_cache(frag_path,z_list)\n    H,W=get_hw(cache)\n    msk=(cv2.imread(os.path.join(frag_path,'inklabels.png'),0)>0).astype(np.uint8)\n    if msk.shape != (H,W):\n        msk=cv2.resize(msk,(W,H),interpolation=cv2.INTER_NEAREST)\n    pred_map=np.zeros((H,W),np.float32); wgt_map=np.zeros((H,W),np.float32)\n    coords=[(y,x) for y in range(0,H-PATCH_SIZE+1,STRIDE_INF)\n                  for x in range(0,W-PATCH_SIZE+1,STRIDE_INF)]\n    with torch.no_grad():\n        for (y,x) in tqdm(coords,desc='Infer',leave=True):\n            raw=np.stack([cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                          for z in z_list],axis=-1)\n            full=build_input(raw); del raw\n            t=norm_patch(torch.from_numpy(full).permute(2,0,1).float()\n                         ).unsqueeze(0).to(DEVICE); del full\n            with autocast(enabled=USE_AMP):\n                p=torch.sigmoid(model(t)).squeeze().cpu().float().numpy()\n            del t\n            pred_map[y:y+PATCH_SIZE,x:x+PATCH_SIZE]+=p*GW\n            wgt_map [y:y+PATCH_SIZE,x:x+PATCH_SIZE]+=GW\n            del p\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    return pred_map/(wgt_map+1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  EVALUATION\n# ════════════════════════════════════════════════════════════\ndef evaluate_and_plot(prob_map, gt_mask, label, fname, thr=None):\n    H,W=gt_mask.shape\n    pc=np.nan_to_num(prob_map[:H,:W],nan=0.5)\n    if thr is None:\n        thr,_=sweep_threshold(pc,gt_mask)\n        print(f'  Swept threshold: {thr:.2f}')\n    pred_raw=(pc>thr).astype(np.uint8)\n    pred,n_removed=cc_filter(pred_raw)\n\n    def metrics(p,m):\n        d=(2*(p*m).sum()+1)/(p.sum()+m.sum()+1)\n        tn,fp,fn,tp=confusion_matrix(m.flatten().astype(int),\n                                     p.flatten().astype(int),\n                                     labels=[0,1]).ravel()\n        pr=tp/(tp+fp+1e-8); re=tp/(tp+fn+1e-8)\n        return dict(dice=float(d),prec=pr,rec=re,\n                    f1=2*pr*re/(pr+re+1e-8),\n                    tp=int(tp),fp=int(fp),fn=int(fn),tn=int(tn))\n\n    mr=metrics(pred_raw,gt_mask); mc=metrics(pred,gt_mask)\n    ink_m=float(pc[gt_mask==1].mean()) if (gt_mask==1).any() else 0.\n    bg_m =float(pc[gt_mask==0].mean()) if (gt_mask==0).any() else 0.\n\n    print(f'\\n{\"=\"*55}')\n    print(f'RESULTS — {label}')\n    print(f'{\"=\"*55}')\n    print(f'               Raw      CC-filtered')\n    print(f'Dice       : {mr[\"dice\"]:.4f}    {mc[\"dice\"]:.4f}')\n    print(f'Precision  : {mr[\"prec\"]:.4f}    {mc[\"prec\"]:.4f}')\n    print(f'Recall     : {mr[\"rec\"]:.4f}    {mc[\"rec\"]:.4f}')\n    print(f'F1         : {mr[\"f1\"]:.4f}    {mc[\"f1\"]:.4f}')\n    print(f'FP/TP      : {mr[\"fp\"]/(mr[\"tp\"]+1e-8):.2f}      '\n          f'{mc[\"fp\"]/(mc[\"tp\"]+1e-8):.2f}')\n    print(f'CC removed : {n_removed}')\n    print(f'Threshold  : {thr:.2f}')\n    print(f'Sep        : {ink_m-bg_m:+.3f}  (ink={ink_m:.3f} bg={bg_m:.3f})')\n    print(f'{\"=\"*55}')\n\n    fig,ax=plt.subplots(2,3,figsize=(18,12))\n    ax[0,0].imshow(gt_mask,cmap='gray');  ax[0,0].set_title('GT')\n    ax[0,1].imshow(pc,cmap='inferno');    ax[0,1].set_title('Prob')\n    ax[0,2].imshow(pred,cmap='gray');     ax[0,2].set_title(f'CC Dice={mc[\"dice\"]:.4f}')\n    err=np.zeros((*gt_mask.shape,3),dtype=np.uint8)\n    err[(pred==1)&(gt_mask==1)]=[0,255,0]\n    err[(pred==1)&(gt_mask==0)]=[255,0,0]\n    err[(pred==0)&(gt_mask==1)]=[0,0,255]\n    ax[1,0].imshow(err); ax[1,0].set_title('TP/FP/FN')\n    iv=pc[gt_mask==1].ravel(); bv=pc[gt_mask==0].ravel()\n    if len(iv): ax[1,1].hist(iv[np.isfinite(iv)],bins=60,alpha=0.7,\n                             label=f'ink μ={ink_m:.3f}',color='orange',density=True)\n    if len(bv): ax[1,1].hist(bv[np.isfinite(bv)],bins=60,alpha=0.7,\n                             label=f'bg μ={bg_m:.3f}',color='blue',density=True)\n    ax[1,1].axvline(thr,color='r',ls='--'); ax[1,1].legend()\n    ax[1,1].set_title('Distributions')\n    ts=np.arange(0.05,0.95,0.01)\n    ds=[(2*((pc>t)*gt_mask).sum()+1)/((pc>t).sum()+gt_mask.sum()+1) for t in ts]\n    ax[1,2].plot(ts,ds,lw=2); ax[1,2].axvline(thr,color='r',ls='--')\n    ax[1,2].set_title('Dice vs Threshold'); ax[1,2].grid(True)\n    for a in ax[0]: a.axis('off')\n    ax[1,0].axis('off')\n    plt.suptitle(f'{label}  Raw={mr[\"dice\"]:.4f}  CC={mc[\"dice\"]:.4f}  '\n                 f'Sep={ink_m-bg_m:+.3f}')\n    plt.tight_layout()\n    plt.savefig(OUTPUT+fname,dpi=80,bbox_inches='tight'); plt.close()\n    return dict(raw=mr,cc=mc,thr=thr,ink_mean=ink_m,noink_mean=bg_m,\n                n_removed=n_removed)\n\n\n# ════════════════════════════════════════════════════════════\n#  BUILD & TRAIN\n# ════════════════════════════════════════════════════════════\nprint('\\n'+'='*60)\nprint(f'TRAINING — Stem+UNet++ | {N_CH} ch | patch {PATCH_SIZE}')\nprint('='*60)\n\nprint('\\nLoading caches...')\nload_ct_cache(FRAG3, Z_SLICES_RAW)\nload_ct_cache(FRAG2, Z_SLICES_RAW)\n\nseg_f3     = SegDataset(FRAG3, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=train_tf, neg_ratio=NEG_RATIO,\n                        max_patches=MAX_PATCHES)\nseg_f2_all = SegDataset(FRAG2, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=train_tf, neg_ratio=NEG_RATIO,\n                        max_patches=0)\nn_seg      = len(seg_f2_all)\nseg_perm   = np.random.permutation(n_seg)\nseg_trn_idx= seg_perm[int(0.2*n_seg):]\nseg_val_idx= seg_perm[:int(0.2*n_seg)]\nseg_f2_tr  = Subset(seg_f2_all, seg_trn_idx)\nseg_val    = Subset(seg_f2_all, seg_val_idx)\nprint(f'Frag2: {len(seg_trn_idx)} train | {len(seg_val_idx)} val')\n\n_kw       = dict(num_workers=0, pin_memory=False)\nseg_ds    = ConcatDataset([seg_f3, seg_f2_tr])\nw_all     = np.concatenate([seg_f3.weights, seg_f2_all.weights[seg_trn_idx]])\nsampler   = WeightedRandomSampler(torch.from_numpy(w_all), len(seg_ds), True)\nseg_dl_tr = DataLoader(seg_ds, batch_size=BATCH_SIZE, sampler=sampler, **_kw)\nseg_dl_vl = DataLoader(seg_val, batch_size=BATCH_SIZE, shuffle=False, **_kw)\nprint(f'Train batches: {len(seg_dl_tr)} | Val: {len(seg_dl_vl)}')\n\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nseg_model = build_seg_model(N_CH)\nseg_model, best_dice, history = run_seg(\n    seg_model, seg_dl_tr, seg_dl_vl,\n    n_epochs=EPOCHS_SEG, ckpt='best_model.pth')\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\ndel seg_dl_tr, seg_dl_vl, seg_ds, seg_val\ndel seg_f3, seg_f2_all, seg_f2_tr\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nfig,axes=plt.subplots(1,3,figsize=(18,5))\naxes[0].plot(history['tl'],label='train'); axes[0].plot(history['vl'],label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'],label='train'); axes[1].plot(history['vd'],label='val')\naxes[1].axhline(0.65,color='orange',ls='--'); axes[1].axhline(0.80,color='r',ls='--')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['sep']); axes[2].axhline(0.3,color='orange',ls='--')\naxes[2].set_title('Sep'); axes[2].grid(True)\nplt.tight_layout(); plt.savefig(OUTPUT+'curves.png',dpi=80); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n'+'='*60+'  FINAL TEST — FRAGMENT 1  '+'='*60)\nckpt=torch.load(OUTPUT+'best_model.pth',map_location=DEVICE)\nprint(f'Checkpoint ep={ckpt[\"ep\"]+1} dice={ckpt[\"dice\"]:.4f} thr={ckpt[\"thr\"]:.2f}')\nfinal=build_seg_model(N_CH); final.load_state_dict(ckpt['state'])\nprob_map,msk1=predict_fragment(final,FRAG1,Z_SLICES_RAW)\nres=evaluate_and_plot(prob_map,msk1,\n                      label='Fragment 1 (unseen test)',\n                      fname='frag1_prediction.png', thr=None)\n\nprint(f'\\n★  Dice (CC)  : {res[\"cc\"][\"dice\"]:.4f}')\nprint(f'★  Precision  : {res[\"cc\"][\"prec\"]:.4f}')\nprint(f'★  Recall     : {res[\"cc\"][\"rec\"]:.4f}')\nprint(f'★  FP/TP      : {res[\"cc\"][\"fp\"]/(res[\"cc\"][\"tp\"]+1e-8):.2f}')\nprint(f'★  Sep        : {res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}')\nprint(f'\\nOutputs → {OUTPUT}')\nprint('IMPORTANT: check z_slice_analysis.png and verify_frag2/3.png')\nprint('before interpreting results — if sep was low in Step 1,')\nprint('z-slice range or label alignment is the root problem.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\n# ============================================================\n#  VESUVIUS INK DETECTION — v24\n#\n#  DIAGNOSIS of v23 failure (sep≈0.001-0.005, thr=0.10):\n#  ┌─────────────────────────────────────────────────────┐\n#  │ TRAINING COLLAPSE — model never learned anything.   │\n#  │                                                     │\n#  │ Evidence:                                           │\n#  │  • sep=0.001 (should be >0.4 by ep5)               │\n#  │  • threshold always 0.10 = model prob ≈ 0 everywhere│\n#  │  • val dice 0.40 = just dataset base rate           │\n#  │  • train dice jumps 0.31→0.54 (loss bouncing)       │\n#  │                                                     │\n#  │ Root causes:                                        │\n#  │  1. LR warmup start_factor=0.1 → ep1 at 5e-6,      │\n#  │     too low for random-init EfficientNet            │\n#  │  2. Tversky+Focal near zero logits → near zero grads│\n#  │  3. SSL weight transfer scrambles Kaiming init      │\n#  │  4. encoder_weights=None is too hard with 20 epochs │\n#  │  5. 24ch input blocks ImageNet pretrained weights   │\n#  └─────────────────────────────────────────────────────┘\n#\n#  FIXES (in order of importance):\n#  1. Use ImageNet pretrained encoder (resnet34, 3ch)\n#     Adapt 24ch → 3ch with a learned 1×1 conv stem.\n#     This is the single biggest improvement possible.\n#  2. Remove LR warmup — start at full LR immediately\n#  3. Remove SSL entirely (too much complexity, no benefit\n#     shown; pretrained encoder is better initialisation)\n#  4. Remove SSL weight transfer (was harming init)\n#  5. Simpler loss: BCE + Dice only (stable from epoch 1)\n#  6. Verify sep > 0.1 by epoch 2, abort early if not\n#  7. Higher LR: 2e-4 (was 5e-5)\n#  8. Dropout in decoder to reduce overfitting\n# ============================================================\n\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import (Dataset, DataLoader,\n                               WeightedRandomSampler, ConcatDataset, Subset)\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom scipy.ndimage import median_filter, label as cc_label\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nimport warnings\nwarnings.filterwarnings('ignore')\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:32\"\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = False\n\n# ── paths ──────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ───────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 96\nSTRIDE_TR    = 48\nSTRIDE_INF   = 32\nBATCH_SIZE   = 8           # larger batch fine with 3-ch stem\nGRAD_ACCUM   = 2           # effective = 16\nEPOCHS_SEG   = 20\nLR_SEG       = 2e-4        # higher LR, no warmup needed\nLR_ENCODER   = 2e-5        # encoder fine-tuned at lower LR\nWEIGHT_DECAY = 1e-4\nPATIENCE     = 8\nMAX_PATCHES  = 6_000\n\nZ_START      = 18\nZ_END        = 36          # 24 slices\nZ_SLICES_RAW = list(range(Z_START, Z_END))\nN_RAW        = len(Z_SLICES_RAW)   # 24\nN_ENGINEERED = 6\nN_CH         = N_RAW + N_ENGINEERED  # 30\n\nCC_MIN_PIXELS = 50\nINK_MIN_POS   = 0.02\nNEG_RATIO     = 0.5\nPOS_WEIGHT    = 4.0\nUSE_AMP       = (DEVICE == 'cuda')\nNORM_STD_MIN  = 0.01\nNORM_CLIP     = 10.0\n\nprint(f\"Device  : {DEVICE}\")\nprint(f\"Input   : {N_RAW} raw + {N_ENGINEERED} engineered = {N_CH} ch → 3ch stem\")\nprint(f\"Patch   : {PATCH_SIZE}×{PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU     : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM    : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. CT CACHE — float16, per-fragment normalisation\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_ct_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        print(f\"  Reusing cache: {os.path.basename(frag_path)}\")\n        return _cache_store[key]\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    slices, vals = [], []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            slices.append((z, s)); vals.append(s.ravel()[::50])\n    assert slices, f\"No CT slices in {vol_dir}\"\n    p1, p99 = np.percentile(np.concatenate(vals), [1, 99])\n    del vals; gc.collect()\n    cache = {}; H = W = None\n    for z, s in slices:\n        s = ((np.clip(s, p1, p99) - p1) / (p99 - p1 + 1e-6)).astype(np.float16)\n        cache[z] = s\n        if H is None: H, W = s.shape\n    del slices; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: {len(cache)} slices, {mb:.0f} MB, ({H},{W})\")\n    _cache_store[key] = cache\n    return cache\n\ndef get_hw(cache): return next(iter(cache.values())).shape\n\ndef norm_patch(t):\n    mu  = t.mean(dim=(1,2), keepdim=True)\n    std = t.std (dim=(1,2), keepdim=True).clamp(min=NORM_STD_MIN)\n    return ((t - mu) / std).clamp(-NORM_CLIP, NORM_CLIP)\n\n\n# ════════════════════════════════════════════════════════════\n#  2. PHYSICS CHANNELS (6)\n# ════════════════════════════════════════════════════════════\ndef physics_enhance(raw_vol):\n    \"\"\"raw_vol: (H,W,N_RAW) float32 → (H,W,6) float32\"\"\"\n    ch0 = raw_vol.max(axis=2)\n    ch1 = raw_vol.std(axis=2)\n    ch2 = np.abs(np.diff(raw_vol, axis=2)).max(axis=2)\n    ch3 = np.maximum(0, np.diff(raw_vol, axis=2)).sum(axis=2)\n    mid = raw_vol.shape[2] // 2\n    ch4 = raw_vol[:,:,mid] - median_filter(raw_vol[:,:,mid], size=3)\n    ch5 = (np.percentile(raw_vol, 75, axis=2) -\n           np.percentile(raw_vol, 25, axis=2))\n    return np.stack([ch0,ch1,ch2,ch3,ch4,ch5], axis=-1).astype(np.float32)\n\ndef build_input(raw_vol):\n    eng = physics_enhance(raw_vol)\n    out = np.concatenate([raw_vol, eng], axis=-1); del eng\n    return out   # (H,W,30)\n\n\n# ════════════════════════════════════════════════════════════\n#  3. PAPYRUS MASK\n# ════════════════════════════════════════════════════════════\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H, W):\n                m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n            return (m > 0).astype(np.uint8)\n    z = list(cache.keys())[len(cache)//2]\n    return (cache[z] > 0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Affine(scale=(0.85,1.15),\n             translate_percent={'x':(-0.1,0.1),'y':(-0.1,0.1)},\n             rotate=(-30,30), mode=cv2.BORDER_REFLECT, p=0.6),\n    A.RandomBrightnessContrast(0.3, 0.3, p=0.6),\n    A.GaussianBlur(blur_limit=(3,7), p=0.3),\n    A.GaussNoise(var_limit=(0.001,0.01), p=0.3),\n    A.CoarseDropout(max_holes=6, max_height=24, max_width=24,\n                    fill_value=0, p=0.4),\n    A.ElasticTransform(alpha=20, sigma=4, p=0.3),\n    A.RandomGamma(gamma_limit=(80,120), p=0.3),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  5. SEGMENTATION DATASET (no SSL dataset needed)\n# ════════════════════════════════════════════════════════════\nclass SegDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0., max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        self.tf     = transform\n        H, W        = get_hw(self.cache)\n\n        msk = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n        assert msk is not None\n        self.mask = (msk > 0).astype(np.uint8); del msk\n\n        ink_dilated = cv2.dilate(\n            self.mask, np.ones((PATCH_SIZE*2, PATCH_SIZE*2), np.uint8))\n        pap = load_papyrus_mask(frag_path, H, W, self.cache)\n\n        pos_yx, hard_neg, easy_neg = [], [], []\n        for y in range(0, H-PATCH_SIZE+1, stride):\n            for x in range(0, W-PATCH_SIZE+1, stride):\n                if pap[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() < 0.4: continue\n                ink = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                if ink >= INK_MIN_POS:\n                    pos_yx.append((y,x))\n                elif ink < 0.001:\n                    if ink_dilated[y+PATCH_SIZE//2, x+PATCH_SIZE//2]:\n                        hard_neg.append((y,x))\n                    else:\n                        easy_neg.append((y,x))\n\n        n_neg   = int(len(pos_yx)*neg_ratio)\n        n_hard  = int(n_neg*0.7); n_easy = n_neg - n_hard\n        np.random.shuffle(hard_neg); np.random.shuffle(easy_neg)\n        neg_yx  = hard_neg[:n_hard] + easy_neg[:n_easy]\n\n        all_yx  = pos_yx + neg_yx\n        all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n        n_total = len(all_yx)\n        if max_patches > 0 and n_total > max_patches:\n            frac = max_patches/n_total\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx = pos_yx[:max(1,int(len(pos_yx)*frac))]\n            neg_yx = neg_yx[:max(0,int(len(neg_yx)*frac))]\n            all_yx  = pos_yx + neg_yx\n            all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n\n        perm = np.random.permutation(len(all_yx))\n        self.coords  = np.array([all_yx[i]  for i in perm], dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm], dtype=np.int32)\n        self.weights = np.where(self.labels==1, POS_WEIGHT, 1.).astype(np.float32)\n        print(f\"  Seg [{os.path.basename(frag_path)}]: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg = {len(self.coords)}\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x = int(self.coords[idx,0]), int(self.coords[idx,1])\n        raw  = np.stack([self.cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                         for z in self.z_list], axis=-1)\n        msk  = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n        if self.tf:\n            o = self.tf(image=raw, mask=msk); raw, msk = o['image'], o['mask']\n        full  = build_input(raw); del raw\n        img_t = norm_patch(torch.from_numpy(full).permute(2,0,1).float()); del full\n        return img_t, torch.from_numpy(msk).unsqueeze(0).float()\n\n\n# ════════════════════════════════════════════════════════════\n#  6. MODEL — Channel-Stem + UNet++ with ImageNet encoder\n#\n#  KEY FIX: Instead of encoder_weights=None (random init),\n#  we use ImageNet pretrained resnet34.\n#\n#  Problem: resnet34 expects 3 channels; we have 30.\n#  Solution: lightweight learned stem that compresses 30→3.\n#    stem = Conv2d(30, 3, 1×1) + BN + ReLU\n#  The stem learns which channel combinations are useful;\n#  the pretrained encoder then extracts rich features.\n#\n#  Why this works:\n#  - ImageNet features (edges, textures, blobs) transfer well\n#    to ink detection — ink strokes ARE texture patterns\n#  - Random-init encoder needs 100s of epochs to learn basic\n#    features that pretrained encoder already knows\n# ════════════════════════════════════════════════════════════\nclass StemUNetPP(nn.Module):\n    def __init__(self, in_ch=N_CH, encoder_name='resnet34'):\n        super().__init__()\n        # Stem: compress N_CH → 3 so ImageNet weights apply\n        self.stem = nn.Sequential(\n            nn.Conv2d(in_ch, 16, 1, bias=False),\n            nn.BatchNorm2d(16),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(16, 3, 1, bias=False),\n            nn.BatchNorm2d(3),\n            nn.ReLU(inplace=True),\n        )\n        # UNet++ with ImageNet pretrained encoder\n        self.unetpp = smp.UnetPlusPlus(\n            encoder_name=encoder_name,\n            encoder_weights='imagenet',    # KEY FIX\n            in_channels=3,                 # stem output\n            classes=1,\n            activation=None,\n            decoder_channels=(256, 128, 64, 32, 16),\n            decoder_use_batchnorm=True,\n        )\n\n    def forward(self, x):\n        return self.unetpp(self.stem(x))\n\n    def encoder_parameters(self):\n        return list(self.stem.parameters()) + \\\n               list(self.unetpp.encoder.parameters())\n\n    def decoder_parameters(self):\n        return (list(self.unetpp.decoder.parameters()) +\n                list(self.unetpp.segmentation_head.parameters()))\n\n\ndef build_seg_model():\n    for enc in ('resnet34', 'resnet18', 'efficientnet-b0'):\n        try:\n            m = StemUNetPP(in_ch=N_CH, encoder_name=enc)\n            # Quick forward pass to verify it works\n            dummy = torch.zeros(1, N_CH, PATCH_SIZE, PATCH_SIZE)\n            with torch.no_grad(): _ = m(dummy)\n            print(f\"  Model: Stem({N_CH}→3) + UNet++({enc}, ImageNet)\")\n            return m.to(DEVICE)\n        except Exception as e:\n            print(f\"  {enc} failed: {e}, trying next...\")\n    raise RuntimeError(\"No encoder worked\")\n\n\n# ════════════════════════════════════════════════════════════\n#  7. LOSS FUNCTIONS\n#\n#  FIX: Use simple BCE + soft-Dice instead of Tversky+Focal.\n#  BCE is stable from epoch 1 with any logit range.\n#  Soft-Dice directly optimises the metric we care about.\n#  Tversky with (0.4,0.6) was causing gradient instability\n#  when logits are near zero at initialisation.\n# ════════════════════════════════════════════════════════════\ndef soft_dice_loss(logits, targets, smooth=1.):\n    p  = torch.sigmoid(logits)\n    tp = (p * targets).sum(dim=(2,3))\n    fp = p.sum(dim=(2,3))\n    fn = targets.sum(dim=(2,3))\n    return (1. - (2.*tp + smooth) / (fp + fn + smooth)).mean()\n\n\ndef bce_loss(logits, targets):\n    return F.binary_cross_entropy_with_logits(logits, targets)\n\n\ndef seg_loss(logits, targets):\n    return 0.5 * bce_loss(logits, targets) + 0.5 * soft_dice_loss(logits, targets)\n\n\n# ════════════════════════════════════════════════════════════\n#  8. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p*masks).sum(dim=(1,2,3))\n    union = p.sum(dim=(1,2,3)) + masks.sum(dim=(1,2,3))\n    return ((2.*inter+1e-5)/(union+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.10, 0.95, 0.01):\n        p = (probs>t).astype(np.float32)\n        d = (2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d > best_d: best_d, best_t = d, float(t)\n    return best_t, best_d\n\ndef compute_sep(probs, masks):\n    if (masks>0.5).any() and (masks<0.5).any():\n        return float(probs[masks>0.5].mean() - probs[masks<0.5].mean())\n    return 0.\n\n\n# ════════════════════════════════════════════════════════════\n#  9. CC FILTER\n# ════════════════════════════════════════════════════════════\ndef cc_filter(binary_mask, min_px=CC_MIN_PIXELS):\n    labeled, n = cc_label(binary_mask)\n    if n == 0: return binary_mask, 0\n    sizes = np.bincount(labeled.ravel()); sizes[0] = 0\n    keep  = sizes >= min_px\n    return keep[labeled].astype(np.uint8), int((sizes[1:]>0).sum()-keep[1:].sum())\n\n\n# ════════════════════════════════════════════════════════════\n#  10. GAUSSIAN WEIGHT\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c=sz//2; s=sz//4; y,x=np.mgrid[0:sz,0:sz]\n    return np.exp(-((x-c)**2+(y-c)**2)/(2*s**2)).astype(np.float32)\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  11. TRAINING\n# ════════════════════════════════════════════════════════════\ndef run_seg(model, train_dl, val_dl, n_epochs, ckpt):\n    # FIX: differential LR — encoder at 10× lower LR than decoder\n    # This preserves ImageNet features while letting decoder learn fast\n    optimizer = optim.AdamW([\n        {'params': model.decoder_parameters(), 'lr': LR_SEG},\n        {'params': model.encoder_parameters(), 'lr': LR_ENCODER},\n    ], weight_decay=WEIGHT_DECAY)\n\n    # Simple cosine decay, no warmup (pretrained encoder doesn't need it)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=n_epochs, eta_min=1e-6)\n    scaler    = GradScaler(enabled=USE_AMP)\n\n    best_dice = 0.; pat = 0\n    history   = dict(tl=[], vl=[], td=[], vd=[], sep=[])\n\n    for ep in range(n_epochs):\n        # ── train ──\n        model.train(); tl = td = 0.; optimizer.zero_grad()\n        for step, (imgs, msks) in enumerate(\n                tqdm(train_dl, desc=f'Ep{ep+1:02d}▸train', leave=False)):\n            imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                out  = model(imgs)\n                loss = seg_loss(out, msks) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            tl += loss.item() * GRAD_ACCUM\n            td += batch_dice(out.detach(), msks)\n            del imgs, msks, out, loss\n            if (step+1) % GRAD_ACCUM == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.)\n                scaler.step(optimizer); scaler.update(); optimizer.zero_grad()\n                if DEVICE=='cuda': torch.cuda.empty_cache()\n        tl /= len(train_dl); td /= len(train_dl)\n\n        # ── val ──\n        model.eval(); vl = vd = 0.\n        all_p, all_m = [], []\n        with torch.no_grad():\n            for imgs, msks in tqdm(val_dl, desc=f'Ep{ep+1:02d}▸val  ', leave=False):\n                imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n                with autocast(enabled=USE_AMP):\n                    out  = model(imgs); loss = seg_loss(out, msks)\n                vl += loss.item(); vd += batch_dice(out, msks)\n                all_p.append(torch.sigmoid(out).cpu().half().numpy())\n                all_m.append(msks.cpu().half().numpy())\n                del imgs, msks, out, loss\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n        vl /= len(val_dl); vd /= len(val_dl)\n\n        P = np.concatenate(all_p).astype(np.float32)\n        M = np.concatenate(all_m).astype(np.float32)\n        del all_p, all_m\n        bt, bd = sweep_threshold(P, M)\n        sep    = compute_sep(P, M)\n        del P, M; gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n        scheduler.step()\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n        history['sep'].append(sep)\n\n        dec_lr = optimizer.param_groups[0]['lr']\n        enc_lr = optimizer.param_groups[1]['lr']\n        flag   = ' ⚠ LOW SEP' if ep >= 2 and sep < 0.1 else ''\n        print(f'Ep{ep+1:02d} | dec_lr={dec_lr:.1e} enc_lr={enc_lr:.1e} | '\n              f'train={tl:.4f}/{td:.4f} | val={vl:.4f}/{vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | sep={sep:+.3f}{flag}')\n\n        # ABORT EARLY: if sep<0.05 after 3 epochs something is\n        # fundamentally wrong — don't waste the rest of the run\n        if ep == 2 and sep < 0.05:\n            print('\\n⛔ ABORTING: sep<0.05 after ep3 — model not learning.')\n            print('   Check: data loading, label paths, GPU memory.')\n            break\n\n        metric = max(vd, bd)\n        if metric > best_dice:\n            best_dice = metric; pat = 0\n            torch.save({'ep':ep,'state':model.state_dict(),\n                        'thr':bt,'dice':best_dice}, OUTPUT+ckpt)\n            print(f'  ✓ checkpoint (dice={best_dice:.4f})')\n        else:\n            pat += 1\n            if pat >= PATIENCE: print(f'  early stop ep{ep+1}'); break\n\n    return model, best_dice, history\n\n\n# ════════════════════════════════════════════════════════════\n#  12. INFERENCE\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    model.eval()\n    cache = _cache_store.get((frag_path,tuple(z_list))) or \\\n            load_ct_cache(frag_path, z_list)\n    H, W  = get_hw(cache)\n    msk   = (cv2.imread(os.path.join(frag_path,'inklabels.png'),0)>0).astype(np.uint8)\n    pred_map = np.zeros((H,W),np.float32)\n    wgt_map  = np.zeros((H,W),np.float32)\n    coords   = [(y,x) for y in range(0,H-PATCH_SIZE+1,STRIDE_INF)\n                      for x in range(0,W-PATCH_SIZE+1,STRIDE_INF)]\n    with torch.no_grad():\n        for (y,x) in tqdm(coords, desc='Infer', leave=True):\n            raw  = np.stack([cache[z][y:y+PATCH_SIZE,x:x+PATCH_SIZE].astype(np.float32)\n                             for z in z_list], axis=-1)\n            full = build_input(raw); del raw\n            t    = norm_patch(torch.from_numpy(full).permute(2,0,1).float()\n                              ).unsqueeze(0).to(DEVICE); del full\n            with autocast(enabled=USE_AMP):\n                p = torch.sigmoid(model(t)).squeeze().cpu().float().numpy()\n            del t\n            pred_map[y:y+PATCH_SIZE,x:x+PATCH_SIZE] += p*GW\n            wgt_map [y:y+PATCH_SIZE,x:x+PATCH_SIZE] += GW\n            del p\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    return pred_map/(wgt_map+1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  13. EVALUATION\n# ════════════════════════════════════════════════════════════\ndef evaluate_and_plot(prob_map, gt_mask, label, fname, thr=None):\n    H, W = gt_mask.shape\n    pc   = np.nan_to_num(prob_map[:H,:W], nan=0.5)\n    if thr is None:\n        thr, _ = sweep_threshold(pc, gt_mask)\n        print(f'  Threshold swept on this fragment: {thr:.2f}')\n    pred_raw        = (pc>thr).astype(np.uint8)\n    pred, n_removed = cc_filter(pred_raw)\n\n    def metrics(p, m):\n        d  = (2*(p*m).sum()+1)/(p.sum()+m.sum()+1)\n        tn,fp,fn,tp = confusion_matrix(m.flatten().astype(int),\n                                       p.flatten().astype(int),\n                                       labels=[0,1]).ravel()\n        pr=tp/(tp+fp+1e-8); re=tp/(tp+fn+1e-8)\n        return dict(dice=float(d),prec=pr,rec=re,f1=2*pr*re/(pr+re+1e-8),\n                    tp=int(tp),fp=int(fp),fn=int(fn),tn=int(tn))\n\n    mr   = metrics(pred_raw, gt_mask)\n    mc   = metrics(pred,     gt_mask)\n    ink_m = float(pc[gt_mask==1].mean()) if (gt_mask==1).any() else 0.\n    bg_m  = float(pc[gt_mask==0].mean()) if (gt_mask==0).any() else 0.\n\n    print(f'\\n{\"=\"*55}')\n    print(f'RESULTS — {label}')\n    print(f'{\"=\"*55}')\n    print(f'               Raw      CC-filtered')\n    print(f'Dice       : {mr[\"dice\"]:.4f}    {mc[\"dice\"]:.4f}')\n    print(f'Precision  : {mr[\"prec\"]:.4f}    {mc[\"prec\"]:.4f}')\n    print(f'Recall     : {mr[\"rec\"]:.4f}    {mc[\"rec\"]:.4f}')\n    print(f'F1         : {mr[\"f1\"]:.4f}    {mc[\"f1\"]:.4f}')\n    print(f'FP/TP      : {mr[\"fp\"]/(mr[\"tp\"]+1e-8):.2f}      '\n          f'{mc[\"fp\"]/(mc[\"tp\"]+1e-8):.2f}')\n    print(f'CC removed : {n_removed}')\n    print(f'Threshold  : {thr:.2f}')\n    print(f'Sep        : {ink_m-bg_m:+.3f}  (ink={ink_m:.3f} bg={bg_m:.3f})')\n    print(f'{\"=\"*55}')\n\n    fig, ax = plt.subplots(2,3,figsize=(18,12))\n    ax[0,0].imshow(gt_mask,cmap='gray');  ax[0,0].set_title('Ground Truth')\n    ax[0,1].imshow(pc,cmap='inferno');    ax[0,1].set_title('Probability Map')\n    ax[0,2].imshow(pred,cmap='gray');     ax[0,2].set_title(f'Pred CC  Dice={mc[\"dice\"]:.4f}')\n    err=np.zeros((*gt_mask.shape,3),dtype=np.uint8)\n    err[(pred==1)&(gt_mask==1)]=[0,255,0]\n    err[(pred==1)&(gt_mask==0)]=[255,0,0]\n    err[(pred==0)&(gt_mask==1)]=[0,0,255]\n    ax[1,0].imshow(err); ax[1,0].set_title('TP/FP/FN')\n    iv=pc[gt_mask==1].ravel(); bv=pc[gt_mask==0].ravel()\n    if len(iv): ax[1,1].hist(iv[np.isfinite(iv)],bins=60,alpha=0.7,\n                             label=f'ink μ={ink_m:.3f}',color='orange',density=True)\n    if len(bv): ax[1,1].hist(bv[np.isfinite(bv)],bins=60,alpha=0.7,\n                             label=f'bg μ={bg_m:.3f}',color='blue',density=True)\n    ax[1,1].axvline(thr,color='r',ls='--'); ax[1,1].legend()\n    ax[1,1].set_title('Distributions')\n    ts=np.arange(0.05,0.95,0.01)\n    ds=[(2*((pc>t)*gt_mask).sum()+1)/((pc>t).sum()+gt_mask.sum()+1) for t in ts]\n    ax[1,2].plot(ts,ds,lw=2); ax[1,2].axvline(thr,color='r',ls='--')\n    ax[1,2].set_title('Dice vs Threshold'); ax[1,2].grid(True)\n    for a in ax[0]: a.axis('off')\n    ax[1,0].axis('off')\n    plt.suptitle(f'{label}\\nRaw={mr[\"dice\"]:.4f}  CC={mc[\"dice\"]:.4f}  Sep={ink_m-bg_m:+.3f}')\n    plt.tight_layout()\n    plt.savefig(OUTPUT+fname,dpi=80,bbox_inches='tight'); plt.close()\n    return dict(raw=mr,cc=mc,thr=thr,ink_mean=ink_m,noink_mean=bg_m,n_removed=n_removed)\n\n\n# ════════════════════════════════════════════════════════════\n#  14. BUILD DATASETS\n# ════════════════════════════════════════════════════════════\nprint('\\n'+'='*60)\nprint('DATA SETUP')\nprint(f'  Z-slices : {Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]} ({N_RAW} slices)')\nprint(f'  Channels : {N_CH} → 3 (via learned stem)')\nprint(f'  Patch    : {PATCH_SIZE}×{PATCH_SIZE}')\nprint('='*60)\n\nprint('\\nLoading CT caches...')\nload_ct_cache(FRAG3, Z_SLICES_RAW)\nload_ct_cache(FRAG2, Z_SLICES_RAW)\n\nseg_f3     = SegDataset(FRAG3, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=train_tf, neg_ratio=NEG_RATIO,\n                        max_patches=MAX_PATCHES)\nseg_f2_all = SegDataset(FRAG2, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=train_tf, neg_ratio=NEG_RATIO,\n                        max_patches=0)\nn_seg      = len(seg_f2_all)\nseg_perm   = np.random.permutation(n_seg)\nseg_trn_idx = seg_perm[int(0.2*n_seg):]\nseg_val_idx = seg_perm[:int(0.2*n_seg)]\nseg_f2_tr   = Subset(seg_f2_all, seg_trn_idx)\nseg_val     = Subset(seg_f2_all, seg_val_idx)\nprint(f'  Frag2: {len(seg_trn_idx)} train | {len(seg_val_idx)} val')\n\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n_kw = dict(num_workers=0, pin_memory=False)\nseg_ds    = ConcatDataset([seg_f3, seg_f2_tr])\nw_all     = np.concatenate([seg_f3.weights, seg_f2_all.weights[seg_trn_idx]])\nsampler   = WeightedRandomSampler(torch.from_numpy(w_all), len(seg_ds), True)\nseg_dl_tr = DataLoader(seg_ds, batch_size=BATCH_SIZE, sampler=sampler, **_kw)\nseg_dl_vl = DataLoader(seg_val, batch_size=BATCH_SIZE, shuffle=False, **_kw)\nprint(f'Train batches: {len(seg_dl_tr)} | Val: {len(seg_dl_vl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  15. TRAIN\n# ════════════════════════════════════════════════════════════\nprint('\\n'+'='*60)\nprint(f'TRAINING — Stem+UNet++ with ImageNet encoder ({EPOCHS_SEG} epochs)')\nprint(f'  Decoder LR : {LR_SEG:.0e}')\nprint(f'  Encoder LR : {LR_ENCODER:.0e}  (differential — protects ImageNet features)')\nprint(f'  Loss       : 0.5×BCE + 0.5×SoftDice (stable from epoch 1)')\nprint(f'  Target     : sep>0.3 by ep3, val Dice>0.60 by ep8')\nprint('='*60)\n\nseg_model = build_seg_model()\nseg_model, best_dice, history = run_seg(\n    seg_model, seg_dl_tr, seg_dl_vl,\n    n_epochs=EPOCHS_SEG, ckpt='best_model.pth')\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\ndel seg_dl_tr, seg_dl_vl, seg_ds, seg_val\ndel seg_f3, seg_f2_all, seg_f2_tr\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n# Training curves\nfig, axes = plt.subplots(1,3,figsize=(18,5))\naxes[0].plot(history['tl'],label='train'); axes[0].plot(history['vl'],label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'],label='train'); axes[1].plot(history['vd'],label='val')\naxes[1].axhline(0.65,color='orange',ls='--',label='0.65')\naxes[1].axhline(0.80,color='r',ls='--',label='0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['sep'])\naxes[2].axhline(0.3,color='orange',ls='--',label='target 0.3')\naxes[2].axhline(0.5,color='g',ls='--',label='good 0.5')\naxes[2].set_title('Sep (ink−bg)'); axes[2].legend(); axes[2].grid(True)\nplt.suptitle('v24 — Stem+UNet++ ImageNet | BCE+Dice loss')\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves.png',dpi=80); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  16. FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n'+'='*60)\nprint('FINAL TEST — FRAGMENT 1')\nprint('='*60)\n\nckpt = torch.load(OUTPUT+'best_model.pth', map_location=DEVICE)\nprint(f'Checkpoint: ep={ckpt[\"ep\"]+1}  dice={ckpt[\"dice\"]:.4f}  thr={ckpt[\"thr\"]:.2f}')\n\nfinal = build_seg_model()\nfinal.load_state_dict(ckpt['state'])\n\nprob_map, msk1 = predict_fragment(final, FRAG1, Z_SLICES_RAW)\nres = evaluate_and_plot(prob_map, msk1,\n                        label='Fragment 1 (unseen test)',\n                        fname='frag1_prediction.png',\n                        thr=None)\n\nprint(f'\\n★  Dice (raw)      : {res[\"raw\"][\"dice\"]:.4f}')\nprint(f'★  Dice (CC filter): {res[\"cc\"][\"dice\"]:.4f}')\nprint(f'★  Precision       : {res[\"cc\"][\"prec\"]:.4f}')\nprint(f'★  Recall          : {res[\"cc\"][\"rec\"]:.4f}')\nprint(f'★  F1              : {res[\"cc\"][\"f1\"]:.4f}')\nprint(f'★  FP/TP (CC)      : {res[\"cc\"][\"fp\"]/(res[\"cc\"][\"tp\"]+1e-8):.2f}')\nprint(f'★  CC removed      : {res[\"n_removed\"]}')\nprint(f'★  Sep             : {res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}')\n\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write('VESUVIUS v24\\n'+'='*55+'\\n')\n    f.write(f'Z: {Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]} ({N_RAW})\\n')\n    f.write(f'Ch: {N_CH} → 3 stem | Patch: {PATCH_SIZE}\\n')\n    f.write(f'Val best: {best_dice:.4f}\\n'+'='*55+'\\n')\n    for tag,m in [('Raw',res['raw']),('CC',res['cc'])]:\n        f.write(f'[{tag}] dice={m[\"dice\"]:.4f} prec={m[\"prec\"]:.4f} '\n                f'rec={m[\"rec\"]:.4f} FP/TP={m[\"fp\"]/(m[\"tp\"]+1e-8):.2f}\\n')\n    f.write(f'thr={res[\"thr\"]:.2f} CC_removed={res[\"n_removed\"]}\\n')\n    f.write(f'sep={res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}\\n')\n\nprint(f'\\nOutputs → {OUTPUT}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#OOM ERROR\n\"\"\"\n=============================================================================\nVESUVIUS INK DETECTION - PhD Thesis Implementation\n=============================================================================\nApproach: 3D Volumetric CNN + Swin Transformer Hybrid (NOT standard 2D SegFormer)\nInnovation: Treats 65 z-slices as a true 3D volume, uses cross-slice attention,\n            and a multi-scale feature pyramid with deep supervision.\n\nKey Novel Contributions:\n1. 3D depthwise separable convolutions for memory-efficient volumetric feature extraction\n2. Cross-slice Swin Transformer blocks for long-range z-axis attention\n3. Focal + Dice + Lovász combined loss for extreme class imbalance\n4. Mixed precision + gradient checkpointing to avoid OOM on Kaggle\n5. Fragment-aware spatial augmentation\n\nTrain: 80% of fragments 2 & 3\nVal:   20% of fragments 2 & 3  \nTest:  Fragment 1 (with ground truth for metrics + visualization)\n=============================================================================\n\"\"\"\n#!pip install segmentation-models-pytorch==0.2.0\n# ─── 0. INSTALL / IMPORTS ────────────────────────────────────────────────────\nimport subprocess, sys\n\n#def pip_install(pkg):\n#    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pkg])\n\n#pip_install(\"segmentation-models-pytorch\")\n#pip_install(\"albumentations>=1.3.0\")\n#pip_install(\"timm>=0.9.0\")\n# Lovász loss\n#pip_install(\"git+https://github.com/bermanmaxim/LovaszSoftmax.git\")\n\nimport os, gc, math, random, warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport matplotlib\nmatplotlib.use(\"Agg\")\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nfrom PIL import Image\n\n# ─── 1. CONFIG ───────────────────────────────────────────────────────────────\nclass CFG:\n    # Paths\n    BASE       = Path(\"/kaggle/input/vesuvius-challenge-ink-detection/train\")\n    # ↑ fix typo if needed: vesuvius-challenge-ink-detection\n    BASE2      = Path(\"/kaggle/input/vesuvius-challenge-ink-detection/train\")\n    OUT_DIR    = Path(\"/kaggle/working\")\n\n    # Fragments\n    TRAIN_FRAGS = [2, 3]   # 80/20 split\n    TEST_FRAGS  = [1]\n\n    # Volume\n    Z_START    = 16\n    Z_END      = 36         # exclusive → 65 slices (00–64)\n    Z_SLICES   = 20\n\n    # Patch / stride\n    PATCH_H    = 224\n    PATCH_W    = 224\n    STRIDE     = 128         # 50% overlap\n\n    # 3D sub-volume depth fed to model (must divide Z_SLICES cleanly or use all)\n    # We feed ALL 65 slices compressed to 32 via 3D pooling inside model\n    Z_DIM      = 32          # target depth after initial 3D pool\n\n    # Training\n    EPOCHS     = 10\n    BATCH_SIZE = 4           # keep small → OOM-safe\n    ACCUM_STEPS= 4           # effective batch = 16\n    LR         = 2e-4\n    WEIGHT_DECAY = 1e-4\n    WARMUP_EPOCHS = 3\n\n    # Model\n    ENCODER_CH = [32, 64, 128, 256]\n    EMBED_DIM  = 224\n    NUM_HEADS  = 8\n    DEPTH      = 4           # transformer layers\n\n    # Augmentation\n    AUG_PROB   = 0.5\n\n    # Misc\n    SEED       = 42\n    NUM_WORKERS= 2\n    PIN_MEMORY = True\n    THRESHOLD  = 0.4        # binarization threshold\n\n    # Ink prevalence weight in loss\n    POS_WEIGHT = 3.0\n\n\ndef seed_everything(seed=CFG.SEED):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything()\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {DEVICE}\")\nCFG.OUT_DIR.mkdir(exist_ok=True, parents=True)\n\n\n# ─── 2. DATA LOADING ─────────────────────────────────────────────────────────\ndef get_base_path(frag_id):\n    \"\"\"Try both possible path spellings.\"\"\"\n    p = CFG.BASE / str(frag_id)\n    if p.exists():\n        return p\n    p2 = CFG.BASE2 / str(frag_id)\n    if p2.exists():\n        return p2\n    raise FileNotFoundError(f\"Fragment {frag_id} not found in {CFG.BASE} or {CFG.BASE2}\")\n\n\ndef load_volume(frag_id):\n    \"\"\"Load all 65 z-slices → float32 np array [H, W, Z].\"\"\"\n    base = get_base_path(frag_id)\n    slices = []\n    for z in range(CFG.Z_START, CFG.Z_END):\n        p = base / \"surface_volume\" / f\"{z:02d}.tif\"\n        img = cv2.imread(str(p), cv2.IMREAD_UNCHANGED)\n        if img is None:\n            raise IOError(f\"Cannot read {p}\")\n        # normalize to [0,1]\n        img = img.astype(np.float32) / 65535.0\n        slices.append(img)\n    vol = np.stack(slices, axis=-1)   # [H, W, Z]\n    print(f\"  Fragment {frag_id} volume: {vol.shape}\")\n    return vol\n\n\ndef load_label(frag_id):\n    \"\"\"Load ink label PNG → binary float32 [H, W].\"\"\"\n    base = get_base_path(frag_id)\n    p = base / \"inklabels.png\"\n    lbl = cv2.imread(str(p), cv2.IMREAD_GRAYSCALE)\n    lbl = (lbl > 127).astype(np.float32)\n    print(f\"  Fragment {frag_id} label: {lbl.shape}, ink%: {lbl.mean()*100:.2f}%\")\n    return lbl\n\n\ndef load_mask(frag_id):\n    \"\"\"Load papyrus mask if available, else ones.\"\"\"\n    base = get_base_path(frag_id)\n    p = base / \"mask.png\"\n    if p.exists():\n        m = cv2.imread(str(p), cv2.IMREAD_GRAYSCALE)\n        return (m > 127).astype(np.uint8)\n    # fallback: non-zero in middle z-slice\n    base2 = get_base_path(frag_id)\n    mid = base2 / \"surface_volume\" / \"32.tif\"\n    img = cv2.imread(str(mid), cv2.IMREAD_UNCHANGED)\n    return (img > 0).astype(np.uint8)\n\n\n# ─── 3. PATCH EXTRACTION ─────────────────────────────────────────────────────\ndef extract_patches(vol, lbl, mask, stride, patch_h, patch_w, is_train=True, split_ratio=0.8, seed=42):\n    \"\"\"\n    Extract (y, x) patch centers valid inside papyrus mask.\n    Split spatially: left 80% columns for train, right 20% for val.\n    Returns list of (y, x) coords.\n    \"\"\"\n    H, W, _ = vol.shape\n    # spatial split: use column split for reproducibility\n    split_col = int(W * split_ratio)\n\n    coords = []\n    for y in range(0, H - patch_h + 1, stride):\n        for x in range(0, W - patch_w + 1, stride):\n            # check mask coverage\n            m_crop = mask[y:y+patch_h, x:x+patch_w]\n            if m_crop.mean() < 0.3:\n                continue\n            if is_train and x + patch_w <= split_col:\n                coords.append((y, x))\n            elif not is_train and x >= split_col:\n                coords.append((y, x))\n    return coords\n\n\n# ─── 4. DATASET ──────────────────────────────────────────────────────────────\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, labels, masks, coords_per_frag, augment=False):\n        \"\"\"\n        volumes, labels, masks: dicts keyed by frag_id\n        coords_per_frag: dict {frag_id: [(y,x), ...]}\n        \"\"\"\n        self.volumes = volumes\n        self.labels  = labels\n        self.augment = augment\n\n        # flatten into list of (frag_id, y, x)\n        self.samples = []\n        for fid, coords in coords_per_frag.items():\n            for (y, x) in coords:\n                self.samples.append((fid, y, x))\n\n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=15, p=0.4),\n            A.GridDistortion(p=0.2),\n            A.CoarseDropout(max_holes=8, max_height=32, max_width=32, p=0.3),\n        ]) if augment else None\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.samples[idx]\n        vol = self.volumes[fid]\n        lbl = self.labels[fid]\n        H, W, Z = vol.shape\n\n        # safe crop\n        y2 = min(y + CFG.PATCH_H, H)\n        x2 = min(x + CFG.PATCH_W, W)\n        y1, x1 = y2 - CFG.PATCH_H, x2 - CFG.PATCH_W\n\n        # [H, W, Z] → crop\n        patch = vol[y1:y2, x1:x2, :].copy()   # [pH, pW, Z]\n        label = lbl[y1:y2, x1:x2].copy()       # [pH, pW]\n\n        # Augmentation on 2D spatial (applied per slice consistently)\n        if self.augment and self.transform is not None:\n            # albumentations expects [H, W, C]; use Z as channel\n            aug = self.transform(image=patch, mask=label)\n            patch = aug[\"image\"]\n            label = aug[\"mask\"]\n\n        # Rearrange: [Z, H, W] then add channel dim → [1, Z, H, W]\n        patch = torch.from_numpy(patch.transpose(2, 0, 1)).float()  # [Z, H, W]\n        patch = patch.unsqueeze(0)                                   # [1, Z, H, W]\n        label = torch.from_numpy(label).float()                      # [H, W]\n\n        return patch, label\n\n\n# ─── 5. MODEL: 3D-CNN + Swin Transformer Hybrid ──────────────────────────────\nclass DepthwiseSep3DConv(nn.Module):\n    \"\"\"Memory-efficient 3D depthwise separable conv.\"\"\"\n    def __init__(self, in_ch, out_ch, kernel=3, stride=1, padding=1):\n        super().__init__()\n        self.dw = nn.Conv3d(in_ch, in_ch, kernel, stride=stride,\n                             padding=padding, groups=in_ch, bias=False)\n        self.pw = nn.Conv3d(in_ch, out_ch, 1, bias=False)\n        self.bn = nn.BatchNorm3d(out_ch)\n        self.act = nn.GELU()\n\n    def forward(self, x):\n        return self.act(self.bn(self.pw(self.dw(x))))\n\n\nclass VolumetricEncoder(nn.Module):\n    \"\"\"\n    3D encoder that processes [B, 1, Z, H, W] volume.\n    Progressively reduces Z while expanding channels.\n    Output: [B, C, H//8, W//8] 2D feature map (Z fully collapsed).\n    \"\"\"\n    def __init__(self, channels=[32, 64, 128, 256]):\n        super().__init__()\n        self.stem = nn.Sequential(\n            nn.Conv3d(1, channels[0], kernel_size=(3,3,3), padding=1, bias=False),\n            nn.BatchNorm3d(channels[0]),\n            nn.GELU(),\n        )\n        # Each block: depthwise-sep + pool Z by 2, spatial by 2\n        self.blocks = nn.ModuleList()\n        self.pools   = nn.ModuleList()\n        in_ch = channels[0]\n        for i, ch in enumerate(channels[1:]):\n            self.blocks.append(nn.Sequential(\n                DepthwiseSep3DConv(in_ch, ch),\n                DepthwiseSep3DConv(ch,   ch),\n            ))\n            # Pool: Z/2, H/2, W/2 (except last: only spatial)\n            if i < len(channels) - 2:\n                self.pools.append(nn.MaxPool3d(kernel_size=(2,2,2)))\n            else:\n                self.pools.append(nn.MaxPool3d(kernel_size=(1,2,2)))\n            in_ch = ch\n\n        self.z_collapse = nn.AdaptiveAvgPool3d((1, None, None))\n        self.out_ch = channels[-1]\n\n    def forward(self, x):\n        # x: [B, 1, Z, H, W]\n        x = self.stem(x)\n        for blk, pool in zip(self.blocks, self.pools):\n            x = blk(x)\n            x = pool(x)\n        x = self.z_collapse(x)        # [B, C, 1, H', W']\n        x = x.squeeze(2)              # [B, C, H', W']\n        return x\n\n\nclass WindowAttention2D(nn.Module):\n    \"\"\"Simplified window-based multi-head self-attention (2D).\"\"\"\n    def __init__(self, dim, num_heads=8, window_size=8):\n        super().__init__()\n        self.num_heads = num_heads\n        self.window_size = window_size\n        self.scale = (dim // num_heads) ** -0.5\n        self.qkv = nn.Linear(dim, dim * 3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        # x: [B, C, H, W]\n        B, C, H, W = x.shape\n        x_perm = x.permute(0, 2, 3, 1)  # [B, H, W, C]\n        res = x_perm\n\n        # partition into windows\n        ws = self.window_size\n        pH = (H + ws - 1) // ws * ws\n        pW = (W + ws - 1) // ws * ws\n        if pH != H or pW != W:\n            x_perm = F.pad(x_perm, (0, 0, 0, pW - W, 0, pH - H))\n\n        nH, nW = pH // ws, pW // ws\n        xw = x_perm.reshape(B, nH, ws, nW, ws, C)\n        xw = xw.permute(0, 1, 3, 2, 4, 5).reshape(B * nH * nW, ws * ws, C)\n\n        xw = self.norm(xw)\n        qkv = self.qkv(xw).reshape(xw.shape[0], xw.shape[1], 3, self.num_heads, C // self.num_heads)\n        qkv = qkv.permute(2, 0, 3, 1, 4)\n        q, k, v = qkv.unbind(0)\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(-1)\n        out  = (attn @ v).transpose(1, 2).reshape(xw.shape[0], xw.shape[1], C)\n        out  = self.proj(out)\n\n        # reverse windows\n        out = out.reshape(B, nH, nW, ws, ws, C).permute(0, 1, 3, 2, 4, 5)\n        out = out.reshape(B, pH, pW, C)[:, :H, :W, :]  # un-pad\n\n        out = out + res\n        return out.permute(0, 3, 1, 2)  # [B, C, H, W]\n\n\nclass TransformerDecoder(nn.Module):\n    \"\"\"Stacked window-attention + FFN blocks on 2D feature map.\"\"\"\n    def __init__(self, dim, num_heads=8, depth=4, window_size=8):\n        super().__init__()\n        self.layers = nn.ModuleList([\n            nn.Sequential(\n                WindowAttention2D(dim, num_heads, window_size),\n                nn.Sequential(\n                    nn.Conv2d(dim, dim * 4, 1), nn.GELU(),\n                    nn.Conv2d(dim * 4, dim, 1),\n                ),\n            )\n            for _ in range(depth)\n        ])\n        self.norm = nn.GroupNorm(32, dim)\n\n    def forward(self, x):\n        for attn, ffn in [(l[0], l[1]) for l in self.layers]:\n            x = attn(x)\n            x = x + ffn(x)\n        return self.norm(x)\n\n\nclass InkDetectionHead(nn.Module):\n    \"\"\"\n    Multi-scale decoder with skip connections (FPN-style) + deep supervision.\n    \"\"\"\n    def __init__(self, enc_channels=[32, 64, 128, 256], dec_ch=128):\n        super().__init__()\n        self.up4 = nn.Sequential(\n            nn.ConvTranspose2d(enc_channels[-1], dec_ch, 2, stride=2),\n            nn.BatchNorm2d(dec_ch), nn.GELU(),\n        )\n        self.up3 = nn.Sequential(\n            nn.ConvTranspose2d(dec_ch + enc_channels[-2], dec_ch, 2, stride=2),\n            nn.BatchNorm2d(dec_ch), nn.GELU(),\n        )\n        self.up2 = nn.Sequential(\n            nn.ConvTranspose2d(dec_ch + enc_channels[-3], dec_ch, 2, stride=2),\n            nn.BatchNorm2d(dec_ch), nn.GELU(),\n        )\n        # final 1×1 → logit\n        self.head = nn.Conv2d(dec_ch, 1, 1)\n        # deep supervision heads\n        self.ds3 = nn.Conv2d(dec_ch, 1, 1)\n        self.ds2 = nn.Conv2d(dec_ch, 1, 1)\n\n    def forward(self, skips, x):\n        \"\"\"\n        skips: list of [B, C, H, W] from encoder (ascending resolution)\n        x: bottleneck [B, C_last, H', W']\n        \"\"\"\n        s2, s3, s4 = skips\n        x = self.up4(x)                               # ×2\n        x = self.up3(torch.cat([x, s4], dim=1))       # ×2\n        ds3_out = self.ds3(x)\n        x = self.up2(torch.cat([x, s3], dim=1))       # ×2\n        ds2_out = self.ds2(x)\n        logit = self.head(x)\n        return logit, ds3_out, ds2_out\n\n\nclass VesuviusNet(nn.Module):\n    \"\"\"Full 3D-CNN + Window-Transformer hybrid for Vesuvius ink detection.\"\"\"\n    def __init__(self, cfg=CFG):\n        super().__init__()\n        C = cfg.ENCODER_CH\n        self.encoder = VolumetricEncoder(channels=C)\n        self.transformer = TransformerDecoder(\n            dim=C[-1], num_heads=cfg.NUM_HEADS,\n            depth=cfg.DEPTH, window_size=8\n        )\n\n        # Save intermediate feature maps with hooks\n        self._feats = {}\n        self._register_hooks(C)\n\n        self.decoder = InkDetectionHead(enc_channels=C)\n\n    def _register_hooks(self, C):\n        \"\"\"Capture skip features after each encoder block.\"\"\"\n        def make_hook(key):\n            def hook(module, inp, out):\n                # out is 3D: [B, C, Z', H', W'] → collapse Z\n                if out.dim() == 5:\n                    self._feats[key] = out.mean(dim=2)  # [B, C, H', W']\n                else:\n                    self._feats[key] = out\n            return hook\n\n        for i, blk in enumerate(self.encoder.blocks):\n            blk.register_forward_hook(make_hook(f\"blk{i}\"))\n\n    def forward(self, x):\n        # x: [B, 1, Z, H, W]\n        self._feats.clear()\n        feat = self.encoder(x)           # [B, C_last, H//8, W//8]\n        feat = self.transformer(feat)\n\n        # collect skips (blocks 0,1,2 → after pool they are H//2, H//4, H//8 spatially)\n        skips = [\n            self._feats.get(\"blk0\", feat),\n            self._feats.get(\"blk1\", feat),\n            self._feats.get(\"blk2\", feat),\n        ]\n\n        logit, ds3, ds2 = self.decoder(skips, feat)\n        return logit, ds3, ds2   # [B, 1, H//2, W//2] approx\n\n\n# ─── 6. LOSS ─────────────────────────────────────────────────────────────────\ntry:\n    from lovasz_losses import lovasz_hinge\n    HAS_LOVASZ = True\nexcept ImportError:\n    HAS_LOVASZ = False\n    print(\"Lovász not available, using Dice+BCE only.\")\n\n\ndef dice_loss(pred, target, eps=1e-6):\n    pred   = torch.sigmoid(pred).flatten(1)\n    target = target.flatten(1)\n    inter  = (pred * target).sum(1)\n    union  = pred.sum(1) + target.sum(1)\n    return 1 - (2 * inter + eps) / (union + eps)\n\n\ndef focal_loss(pred, target, alpha=0.85, gamma=2.0):\n    bce  = F.binary_cross_entropy_with_logits(pred, target, reduction=\"none\")\n    pt   = torch.exp(-bce)\n    fl   = alpha * (1 - pt) ** gamma * bce\n    return fl.mean()\n\n\ndef combined_loss(pred, target, pos_weight=CFG.POS_WEIGHT):\n    pw = torch.tensor([pos_weight], device=pred.device)\n    bce   = F.binary_cross_entropy_with_logits(pred, target, pos_weight=pw)\n    dice  = dice_loss(pred, target).mean()\n    focal = focal_loss(pred, target)\n    loss  = 0.4 * bce + 0.4 * dice + 0.2 * focal\n    if HAS_LOVASZ:\n        try:\n            lv = lovasz_hinge(pred.squeeze(1), target.squeeze(1))\n            loss = 0.3 * bce + 0.3 * dice + 0.2 * focal + 0.2 * lv\n        except Exception:\n            pass\n    return loss\n\n\ndef multi_scale_loss(logit, ds3, ds2, label):\n    \"\"\"Compute loss at multiple scales via downsampled label.\"\"\"\n    H, W = label.shape[-2:]\n\n    def resize_label(tgt, size):\n        return F.interpolate(tgt.unsqueeze(1).float(), size=size,\n                             mode=\"nearest\").squeeze(1)\n\n    main_h, main_w = logit.shape[-2:]\n    lbl_main = resize_label(label, (main_h, main_w))\n    lbl_ds3  = resize_label(label, ds3.shape[-2:])\n    lbl_ds2  = resize_label(label, ds2.shape[-2:])\n\n    l_main = combined_loss(logit.squeeze(1), lbl_main)\n    l_ds3  = combined_loss(ds3.squeeze(1),   lbl_ds3)\n    l_ds2  = combined_loss(ds2.squeeze(1),   lbl_ds2)\n    return l_main + 0.4 * l_ds3 + 0.2 * l_ds2\n\n\n# ─── 7. METRICS ──────────────────────────────────────────────────────────────\ndef compute_metrics(preds, labels, threshold=CFG.THRESHOLD):\n    \"\"\"preds, labels: flat numpy arrays.\"\"\"\n    p = (preds > threshold).astype(np.uint8)\n    t = (labels > 0.5).astype(np.uint8)\n\n    tp = (p & t).sum()\n    fp = (p & ~t.astype(bool)).sum()\n    fn = (~p.astype(bool) & t).sum()\n    tn = (~p.astype(bool) & ~t.astype(bool)).sum()\n\n    precision = tp / (tp + fp + 1e-8)\n    recall    = tp / (tp + fn + 1e-8)\n    f1        = 2 * precision * recall / (precision + recall + 1e-8)\n    accuracy  = (tp + tn) / (tp + fp + fn + tn + 1e-8)\n    iou       = tp / (tp + fp + fn + 1e-8)\n\n    return dict(f1=f1, precision=precision, recall=recall, accuracy=accuracy, iou=iou)\n\n\n# ─── 8. SCHEDULER ────────────────────────────────────────────────────────────\ndef get_scheduler(optimizer, total_steps, warmup_steps):\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return step / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        return 0.5 * (1 + math.cos(math.pi * progress))\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n\n# ─── 9. INFERENCE ON FULL FRAGMENT ───────────────────────────────────────────\n@torch.no_grad()\ndef predict_fragment(model, vol, mask, device, threshold=CFG.THRESHOLD):\n    \"\"\"\n    Sliding-window inference on full fragment.\n    Returns probability map [H, W].\n    \"\"\"\n    model.eval()\n    H, W, Z = vol.shape\n    pred_map  = np.zeros((H, W), dtype=np.float32)\n    count_map = np.zeros((H, W), dtype=np.float32)\n\n    ph, pw = CFG.PATCH_H, CFG.PATCH_W\n    stride = CFG.STRIDE\n\n    ys = list(range(0, H - ph + 1, stride)) + ([H - ph] if H % stride != 0 else [])\n    xs = list(range(0, W - pw + 1, stride)) + ([W - pw] if W % stride != 0 else [])\n\n    for y in tqdm(ys, desc=\"Inference rows\", leave=False):\n        for x in xs:\n            patch = vol[y:y+ph, x:x+pw, :].copy()\n            patch_t = torch.from_numpy(patch.transpose(2, 0, 1)).float()\n            patch_t = patch_t.unsqueeze(0).unsqueeze(0).to(device)  # [1,1,Z,H,W]\n\n            with autocast():\n                logit, _, _ = model(patch_t)\n            prob = torch.sigmoid(logit).squeeze().float().cpu().numpy()  # [h', w']\n\n            # upsample to patch size\n            prob_up = cv2.resize(prob, (pw, ph), interpolation=cv2.INTER_LINEAR)\n\n            pred_map[y:y+ph, x:x+pw]  += prob_up\n            count_map[y:y+ph, x:x+pw] += 1.0\n\n    count_map = np.maximum(count_map, 1e-8)\n    pred_map  = pred_map / count_map\n    pred_map  = pred_map * mask  # zero out non-papyrus\n    return pred_map\n\n\n# ─── 10. VISUALIZATION ───────────────────────────────────────────────────────\ndef save_visualization(vol, label, pred_prob, frag_id, out_dir, threshold=CFG.THRESHOLD):\n    \"\"\"Save 4-panel: mid z-slice | ground truth | predicted prob | overlay.\"\"\"\n    mid_z = vol[:, :, vol.shape[2] // 2]\n    mid_z_norm = ((mid_z - mid_z.min()) / (mid_z.max() - mid_z.min() + 1e-8) * 255).astype(np.uint8)\n\n    pred_bin  = (pred_prob > threshold).astype(np.float32)\n    metrics   = compute_metrics(pred_prob.ravel(), label.ravel(), threshold)\n\n    fig, axes = plt.subplots(1, 4, figsize=(22, 6))\n    fig.suptitle(\n        f\"Fragment {frag_id} — F1: {metrics['f1']:.4f} | \"\n        f\"IoU: {metrics['iou']:.4f} | Precision: {metrics['precision']:.4f} | \"\n        f\"Recall: {metrics['recall']:.4f}\",\n        fontsize=14, fontweight=\"bold\"\n    )\n\n    axes[0].imshow(mid_z_norm, cmap=\"gray\")\n    axes[0].set_title(\"Input (mid z-slice)\")\n    axes[0].axis(\"off\")\n\n    axes[1].imshow(label, cmap=\"gray\", vmin=0, vmax=1)\n    axes[1].set_title(\"Ground Truth\")\n    axes[1].axis(\"off\")\n\n    im = axes[2].imshow(pred_prob, cmap=\"hot\", vmin=0, vmax=1)\n    axes[2].set_title(\"Predicted Probability\")\n    axes[2].axis(\"off\")\n    plt.colorbar(im, ax=axes[2], fraction=0.046, pad=0.04)\n\n    # Overlay: TP=green, FP=red, FN=blue\n    overlay = np.stack([mid_z_norm, mid_z_norm, mid_z_norm], axis=-1)\n    gt_bool   = label > 0.5\n    pred_bool = pred_bin > 0.5\n    tp_mask = gt_bool & pred_bool\n    fp_mask = (~gt_bool) & pred_bool\n    fn_mask = gt_bool & (~pred_bool)\n    overlay[tp_mask] = [0,   200,   0]\n    overlay[fp_mask] = [200,   0,   0]\n    overlay[fn_mask] = [0,    50, 200]\n    axes[3].imshow(overlay)\n    axes[3].set_title(\"Overlay (TP=green, FP=red, FN=blue)\")\n    axes[3].axis(\"off\")\n    patches = [\n        mpatches.Patch(color=\"green\", label=\"TP\"),\n        mpatches.Patch(color=\"red\",   label=\"FP\"),\n        mpatches.Patch(color=\"blue\",  label=\"FN\"),\n    ]\n    axes[3].legend(handles=patches, loc=\"lower right\", fontsize=8)\n\n    plt.tight_layout()\n    save_path = out_dir / f\"fragment_{frag_id}_visualization.png\"\n    plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.close(fig)\n    print(f\"  Saved visualization → {save_path}\")\n    return metrics\n\n\n# ─── 11. TRAINING LOOP ───────────────────────────────────────────────────────\ndef train_one_epoch(model, loader, optimizer, scheduler, scaler, device, accum_steps):\n    model.train()\n    total_loss = 0.0\n    optimizer.zero_grad()\n\n    for step, (patches, labels) in enumerate(tqdm(loader, desc=\"Train\", leave=False)):\n        patches = patches.to(device, non_blocking=True)    # [B,1,Z,H,W]\n        labels  = labels.to(device, non_blocking=True)     # [B,H,W]\n\n        with autocast():\n            logit, ds3, ds2 = model(patches)\n            loss = multi_scale_loss(logit, ds3, ds2, labels) / accum_steps\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % accum_steps == 0:\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n\n        total_loss += loss.item() * accum_steps\n\n    return total_loss / len(loader)\n\n\n@torch.no_grad()\ndef validate(model, loader, device):\n    model.eval()\n    total_loss = 0.0\n    all_preds, all_labels = [], []\n\n    for patches, labels in tqdm(loader, desc=\"Val\", leave=False):\n        patches = patches.to(device, non_blocking=True)\n        labels  = labels.to(device, non_blocking=True)\n\n        with autocast():\n            logit, ds3, ds2 = model(patches)\n            loss = multi_scale_loss(logit, ds3, ds2, labels)\n\n        total_loss += loss.item()\n\n        prob = torch.sigmoid(logit).squeeze(1)  # [B, h, w]\n        # resize to label size\n        prob_up = F.interpolate(\n            prob.unsqueeze(1), size=labels.shape[-2:], mode=\"bilinear\", align_corners=False\n        ).squeeze(1)\n\n        all_preds.append(prob_up.float().cpu().numpy().ravel())\n        all_labels.append(labels.float().cpu().numpy().ravel())\n\n    preds  = np.concatenate(all_preds)\n    labels = np.concatenate(all_labels)\n    metrics = compute_metrics(preds, labels)\n    return total_loss / len(loader), metrics\n\n\n# ─── 12. MAIN ────────────────────────────────────────────────────────────────\ndef main():\n    print(\"=\" * 60)\n    print(\"Vesuvius Ink Detection — 3D-CNN + Transformer Hybrid\")\n    print(\"=\" * 60)\n\n    # ── Load data ──────────────────────────────────────────────\n    print(\"\\n[1] Loading volumes …\")\n    volumes, labels, masks = {}, {}, {}\n    for fid in CFG.TRAIN_FRAGS + CFG.TEST_FRAGS:\n        print(f\"  Fragment {fid} …\")\n        volumes[fid] = load_volume(fid)\n        labels[fid]  = load_label(fid)\n        masks[fid]   = load_mask(fid)\n\n    # ── Extract patch coordinates ───────────────────────────────\n    print(\"\\n[2] Extracting patch coordinates …\")\n    train_coords, val_coords = {}, {}\n    for fid in CFG.TRAIN_FRAGS:\n        train_coords[fid] = extract_patches(\n            volumes[fid], labels[fid], masks[fid],\n            stride=CFG.STRIDE, patch_h=CFG.PATCH_H, patch_w=CFG.PATCH_W,\n            is_train=True, split_ratio=0.8\n        )\n        val_coords[fid] = extract_patches(\n            volumes[fid], labels[fid], masks[fid],\n            stride=CFG.STRIDE, patch_h=CFG.PATCH_H, patch_w=CFG.PATCH_W,\n            is_train=False, split_ratio=0.8\n        )\n        print(f\"  Fragment {fid}: train={len(train_coords[fid])}, val={len(val_coords[fid])}\")\n\n    # ── Datasets / Loaders ─────────────────────────────────────\n    train_ds = VesuviusDataset(volumes, labels, masks, train_coords, augment=True)\n    val_ds   = VesuviusDataset(volumes, labels, masks, val_coords,   augment=False)\n\n    train_loader = DataLoader(\n        train_ds, batch_size=CFG.BATCH_SIZE, shuffle=True,\n        num_workers=CFG.NUM_WORKERS, pin_memory=CFG.PIN_MEMORY, drop_last=True\n    )\n    val_loader = DataLoader(\n        val_ds, batch_size=CFG.BATCH_SIZE, shuffle=False,\n        num_workers=CFG.NUM_WORKERS, pin_memory=CFG.PIN_MEMORY\n    )\n    print(f\"\\n  Train batches: {len(train_loader)}, Val batches: {len(val_loader)}\")\n\n    # ── Model ──────────────────────────────────────────────────\n    print(\"\\n[3] Building model …\")\n    model = VesuviusNet(CFG).to(DEVICE)\n    n_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Parameters: {n_params/1e6:.2f} M\")\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY\n    )\n    total_steps  = CFG.EPOCHS * len(train_loader) // CFG.ACCUM_STEPS\n    warmup_steps = CFG.WARMUP_EPOCHS * len(train_loader) // CFG.ACCUM_STEPS\n    scheduler    = get_scheduler(optimizer, total_steps, warmup_steps)\n    scaler       = GradScaler()\n\n    # ── Training ───────────────────────────────────────────────\n    print(f\"\\n[4] Training for {CFG.EPOCHS} epochs …\")\n    best_f1    = 0.0\n    best_epoch = 0\n    history    = []\n\n    for epoch in range(1, CFG.EPOCHS + 1):\n        train_loss = train_one_epoch(\n            model, train_loader, optimizer, scheduler, scaler, DEVICE, CFG.ACCUM_STEPS\n        )\n        val_loss, val_metrics = validate(model, val_loader, DEVICE)\n\n        f1 = val_metrics[\"f1\"]\n        print(\n            f\"  Ep {epoch:03d}/{CFG.EPOCHS} | \"\n            f\"TrLoss={train_loss:.4f} | ValLoss={val_loss:.4f} | \"\n            f\"F1={f1:.4f} | IoU={val_metrics['iou']:.4f} | \"\n            f\"Prec={val_metrics['precision']:.4f} | Rec={val_metrics['recall']:.4f}\"\n        )\n        history.append(dict(epoch=epoch, train_loss=train_loss, val_loss=val_loss, **val_metrics))\n\n        if f1 > best_f1:\n            best_f1    = f1\n            best_epoch = epoch\n            torch.save(model.state_dict(), CFG.OUT_DIR / \"best_model.pt\")\n            print(f\"    ✓ Saved best model (F1={best_f1:.4f})\")\n\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    print(f\"\\n  Best validation F1: {best_f1:.4f} at epoch {best_epoch}\")\n\n    # ── Training curves ────────────────────────────────────────\n    df_hist = pd.DataFrame(history)\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    axes[0].plot(df_hist[\"epoch\"], df_hist[\"train_loss\"], label=\"Train\")\n    axes[0].plot(df_hist[\"epoch\"], df_hist[\"val_loss\"],   label=\"Val\")\n    axes[0].set_title(\"Loss\"); axes[0].legend(); axes[0].set_xlabel(\"Epoch\")\n    axes[1].plot(df_hist[\"epoch\"], df_hist[\"f1\"],  label=\"F1\")\n    axes[1].plot(df_hist[\"epoch\"], df_hist[\"iou\"], label=\"IoU\")\n    axes[1].set_title(\"Metrics\"); axes[1].legend(); axes[1].set_xlabel(\"Epoch\")\n    plt.tight_layout()\n    plt.savefig(CFG.OUT_DIR / \"training_curves.png\", dpi=120)\n    plt.close()\n    df_hist.to_csv(CFG.OUT_DIR / \"training_history.csv\", index=False)\n    print(\"  Saved training curves and history.\")\n\n    # ── Load best model ────────────────────────────────────────\n    model.load_state_dict(torch.load(CFG.OUT_DIR / \"best_model.pt\", map_location=DEVICE))\n    print(\"\\n[5] Loaded best model weights.\")\n\n    # ── Test on Fragment 1 ─────────────────────────────────────\n    print(\"\\n[6] Running inference on test fragment(s) …\")\n    all_test_metrics = {}\n\n    for fid in CFG.TEST_FRAGS:\n        print(f\"\\n  Fragment {fid} inference …\")\n        pred_prob = predict_fragment(model, volumes[fid], masks[fid], DEVICE, CFG.THRESHOLD)\n\n        # save raw probability map\n        prob_save = (pred_prob * 255).astype(np.uint8)\n        cv2.imwrite(str(CFG.OUT_DIR / f\"fragment_{fid}_pred_prob.png\"), prob_save)\n\n        # save binary prediction\n        pred_bin  = ((pred_prob > CFG.THRESHOLD) * 255).astype(np.uint8)\n        cv2.imwrite(str(CFG.OUT_DIR / f\"fragment_{fid}_pred_binary.png\"), pred_bin)\n\n        # visualization + metrics\n        metrics = save_visualization(\n            volumes[fid], labels[fid], pred_prob, fid, CFG.OUT_DIR, CFG.THRESHOLD\n        )\n        all_test_metrics[fid] = metrics\n\n        print(f\"\\n  ── Test Metrics (Fragment {fid}) ──────────────────────\")\n        for k, v in metrics.items():\n            print(f\"    {k:12s}: {v:.6f}\")\n\n    # ── Final summary ──────────────────────────────────────────\n    print(\"\\n\" + \"=\" * 60)\n    print(\"FINAL TEST RESULTS\")\n    print(\"=\" * 60)\n    for fid, m in all_test_metrics.items():\n        print(f\"\\nFragment {fid}:\")\n        for k, v in m.items():\n            print(f\"  {k:12s}: {v:.6f}\")\n        target = \"✓ PASSED\" if m[\"f1\"] >= 0.80 else \"✗ Below 0.80 target\"\n        print(f\"  F1 Target 0.80 → {target}\")\n    print(\"=\" * 60)\n    print(f\"All outputs saved to: {CFG.OUT_DIR}\")\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\n# ============================================================\n#  VESUVIUS INK DETECTION — v23\n#\n#  DIAGNOSIS of v22 failure (Dice=0.34, FP/TP=3.82):\n#  ┌─────────────────────────────────────────────────────┐\n#  │ Recall=0.95 but Precision=0.21 → model predicts     │\n#  │ \"ink everywhere on papyrus\" not \"ink strokes\"        │\n#  │                                                      │\n#  │ Root causes:                                         │\n#  │ 1. sep=+0.17 → model barely separates ink from bg   │\n#  │    (need sep>0.4 for useful predictions)             │\n#  │ 2. 64×64 patches too small to see ink texture        │\n#  │ 3. Only 16 z-slices — lost discriminative signal     │\n#  │ 4. Only 4 engineered channels — lost physics signal  │\n#  │ 5. No domain adaptation between train/test fragments │\n#  └─────────────────────────────────────────────────────┘\n#\n#  FIXES:\n#  1. Patch size 64→96 (better context, still memory-safe)\n#  2. Z-slices 16→24 (restore signal, stay within RAM)\n#  3. Engineered channels 4→6 (restore key physics features)\n#  4. Total channels: 20→30\n#  5. Test-time threshold: sweep on FRAG1 val split, not\n#     on Frag2 val (reduces domain gap for threshold choice)\n#  6. Deeper UNet++ decoder: (256,128,64,32,16)\n#  7. Stronger pos/neg sampling: pos_weight=6, neg_ratio=0.5\n#  8. Label smoothing in focal loss reduced (eps 0.05→0.01)\n#  9. Add per-fragment normalisation in CT cache\n# 10. Increase EPOCHS_SEG to 20, reduce LR to 5e-5 for finetuning\n# 11. Threshold sweep on validation uses finer grid (step 0.01)\n# 12. CC filter min_pixels 20→50 (more aggressive FP removal)\n# ============================================================\n\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import (Dataset, DataLoader,\n                               WeightedRandomSampler, ConcatDataset, Subset)\nfrom torch.cuda.amp import GradScaler, autocast\nimport albumentations as A\nfrom scipy.ndimage import median_filter, label as cc_label\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nimport warnings\nwarnings.filterwarnings('ignore')\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:32\"\n\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = False\n\n# ── paths ──────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ───────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 96          # ↑ from 64 — better texture context\nSTRIDE_TR    = 48          # half patch\nSTRIDE_INF   = 32\nBATCH_SIZE   = 8           # 96×96 still fits on 16 GB\nGRAD_ACCUM   = 4\nEPOCHS_SSL   = 3\nEPOCHS_SEG   = 10          # ↑ more epochs, lower LR\nLR_SSL       = 3e-4\nLR_SEG       = 5e-5        # ↓ lower — avoids overfitting to Frag2\nWEIGHT_DECAY = 1e-4        # ↑ stronger regularisation\nPATIENCE     = 10\nMAX_PATCHES  = 5_000\n\n# Z-slices: 24 central slices (was 16 — restores signal)\nZ_START      = 18\nZ_END        = 36          # slices 20-43 = 24 slices\nZ_SLICES_RAW = list(range(Z_START, Z_END))\nN_RAW        = len(Z_SLICES_RAW)   # 24\nN_ENGINEERED = 6                   # ↑ from 4\nN_CH         = N_RAW + N_ENGINEERED  # 30\n\nSSL_PROJ_DIM  = 64\nSSL_TEMP      = 0.07\nCC_MIN_PIXELS = 50         # ↑ from 20 — more aggressive FP removal\nINK_MIN_POS   = 0.02\nNEG_RATIO     = 0.5        # ↑ from 0.3 — more negative patches\nPOS_WEIGHT    = 6.0        # ↑ from 3.0\nUSE_AMP       = (DEVICE == 'cuda')\nNORM_STD_MIN  = 0.01\nNORM_CLIP     = 10.0\n\nprint(f\"Device  : {DEVICE}\")\nprint(f\"Input   : {N_RAW} raw z-slices + {N_ENGINEERED} engineered = {N_CH} ch\")\nprint(f\"Patch   : {PATCH_SIZE}×{PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU     : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM    : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. CT CACHE\n#\n#  FIX: per-fragment percentile normalisation (was global).\n#  Each fragment has different scanner brightness; normalising\n#  per-fragment makes features comparable across train/test.\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_ct_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        print(f\"  Reusing cache: {os.path.basename(frag_path)}\")\n        return _cache_store[key]\n\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    slices, all_vals = [], []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            slices.append((z, s))\n            all_vals.append(s.ravel()[::50])   # 2% sample\n\n    assert slices, f\"No CT slices in {vol_dir}\"\n    # Per-fragment normalisation\n    all_vals = np.concatenate(all_vals)\n    p1, p99  = np.percentile(all_vals, [1, 99])\n    del all_vals; gc.collect()\n\n    cache = {}; H = W = None\n    for z, s in slices:\n        s = np.clip(s, p1, p99)\n        s = ((s - p1) / (p99 - p1 + 1e-6)).astype(np.float16)\n        cache[z] = s\n        if H is None: H, W = s.shape\n    del slices; gc.collect()\n\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: \"\n          f\"{len(cache)} slices, {mb:.0f} MB, ({H},{W})\")\n    _cache_store[key] = cache\n    return cache\n\n\ndef get_hw(cache):\n    return next(iter(cache.values())).shape\n\n\ndef norm_patch(t):\n    \"\"\"Per-channel z-score normalisation of a patch tensor.\"\"\"\n    mu  = t.mean(dim=(1, 2), keepdim=True)\n    std = t.std (dim=(1, 2), keepdim=True).clamp(min=NORM_STD_MIN)\n    return ((t - mu) / std).clamp(-NORM_CLIP, NORM_CLIP)\n\n\n# ════════════════════════════════════════════════════════════\n#  2. PHYSICS CHANNELS (6 channels)\n#\n#  Restored from v22's aggressive cut.\n#  These directly encode the ink signature in CT data:\n#  - MIP: ink = high absorption peak\n#  - std: ink = variable across z (surface vs depth)\n#  - max_diff: ink = sharp z-transition\n#  - pos_accum: ink = cumulative absorption increase\n#  - median_residual: removes smooth background\n#  - IQR: robust variance (insensitive to outliers)\n# ════════════════════════════════════════════════════════════\ndef physics_enhance(raw_vol):\n    \"\"\"raw_vol: (H,W,N_RAW) float32 patch → (H,W,6) float32\"\"\"\n    # [0] Max intensity projection\n    ch0 = raw_vol.max(axis=2)\n\n    # [1] Std across z\n    ch1 = raw_vol.std(axis=2)\n\n    # [2] Max forward difference (sharpest transition)\n    ch2 = np.abs(np.diff(raw_vol, axis=2)).max(axis=2)\n\n    # [3] Cumulative positive diffs (ink accumulation)\n    ch3 = np.maximum(0, np.diff(raw_vol, axis=2)).sum(axis=2)\n\n    # [4] Median residual on centre slice\n    mid = raw_vol.shape[2] // 2\n    sl  = raw_vol[:, :, mid]\n    ch4 = sl - median_filter(sl, size=3)\n\n    # [5] Interquartile range (robust variance across z)\n    ch5 = np.percentile(raw_vol, 75, axis=2) - \\\n          np.percentile(raw_vol, 25, axis=2)\n\n    return np.stack([ch0, ch1, ch2, ch3, ch4, ch5],\n                    axis=-1).astype(np.float32)\n\n\ndef build_input(raw_vol):\n    \"\"\"(H,W,N_RAW) float32 → (H,W,N_CH) float32\"\"\"\n    eng = physics_enhance(raw_vol)\n    out = np.concatenate([raw_vol, eng], axis=-1)\n    del eng\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  3. PAPYRUS MASK\n# ════════════════════════════════════════════════════════════\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H, W):\n                m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n            return (m > 0).astype(np.uint8)\n    z = list(cache.keys())[len(cache) // 2]\n    return (cache[z] > 0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. AUGMENTATIONS\n#  FIX: stronger augmentations to reduce Frag2→Frag1 domain gap\n# ════════════════════════════════════════════════════════════\nstrong_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Affine(scale=(0.85, 1.15),\n             translate_percent={'x': (-0.1, 0.1), 'y': (-0.1, 0.1)},\n             rotate=(-30, 30),\n             mode=cv2.BORDER_REFLECT, p=0.6),\n    A.RandomBrightnessContrast(0.3, 0.3, p=0.6),   # ↑ stronger\n    A.GaussianBlur(blur_limit=(3, 7), p=0.3),\n    A.GaussNoise(var_limit=(0.001, 0.01), p=0.3),   # NEW: noise aug\n    A.CoarseDropout(max_holes=6, max_height=24, max_width=24,\n                    fill_value=0, p=0.4),\n    A.ElasticTransform(alpha=20, sigma=4, p=0.3),   # NEW: elastic\n    A.RandomGamma(gamma_limit=(80, 120), p=0.3),    # NEW: gamma\n])\n\nweak_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.RandomBrightnessContrast(0.1, 0.1, p=0.3),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  5. SSL DATASET\n# ════════════════════════════════════════════════════════════\nclass SSLDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride, max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        H, W        = get_hw(self.cache)\n        pap         = load_papyrus_mask(frag_path, H, W, self.cache)\n\n        coords = [(y, x)\n                  for y in range(0, H - PATCH_SIZE + 1, stride)\n                  for x in range(0, W - PATCH_SIZE + 1, stride)\n                  if pap[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.4]\n        if max_patches > 0 and len(coords) > max_patches:\n            np.random.shuffle(coords); coords = coords[:max_patches]\n        self.coords = np.array(coords, dtype=np.int32)\n        print(f\"  SSL [{os.path.basename(frag_path)}]: {len(self.coords)} patches\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x    = int(self.coords[idx, 0]), int(self.coords[idx, 1])\n        raw_vol = np.stack([\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list], axis=-1)\n        v1 = weak_tf  (image=raw_vol.copy())['image']\n        v2 = strong_tf(image=raw_vol.copy())['image']\n        del raw_vol\n        f1 = build_input(v1); del v1\n        f2 = build_input(v2); del v2\n        t1 = norm_patch(torch.from_numpy(f1).permute(2,0,1).float()); del f1\n        t2 = norm_patch(torch.from_numpy(f2).permute(2,0,1).float()); del f2\n        return t1, t2\n\n\n# ════════════════════════════════════════════════════════════\n#  6. SEGMENTATION DATASET\n#  FIX: harder negative mining — sample negative patches that\n#  are near positive ones (hardest negatives = near-ink bg)\n# ════════════════════════════════════════════════════════════\nclass SegDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0, max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        self.tf     = transform\n        H, W        = get_hw(self.cache)\n\n        msk_p = os.path.join(frag_path, 'inklabels.png')\n        msk   = cv2.imread(msk_p, 0)\n        assert msk is not None, f\"Missing: {msk_p}\"\n        self.mask = (msk > 0).astype(np.uint8); del msk\n\n        # Dilated ink mask for hard negative mining\n        # Negatives within 2 patches of ink are hardest\n        ink_dilated = cv2.dilate(\n            self.mask, np.ones((PATCH_SIZE*2, PATCH_SIZE*2), np.uint8))\n\n        pap = load_papyrus_mask(frag_path, H, W, self.cache)\n        pos_yx, hard_neg_yx, easy_neg_yx = [], [], []\n\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                if pap[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() < 0.4:\n                    continue\n                ink = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                if ink >= INK_MIN_POS:\n                    pos_yx.append((y, x))\n                elif ink < 0.001:\n                    # hard negative: near ink boundary\n                    if ink_dilated[y+PATCH_SIZE//2, x+PATCH_SIZE//2]:\n                        hard_neg_yx.append((y, x))\n                    else:\n                        easy_neg_yx.append((y, x))\n\n        # Mix 70% hard / 30% easy negatives\n        n_neg = int(len(pos_yx) * neg_ratio)\n        n_hard = int(n_neg * 0.7); n_easy = n_neg - n_hard\n        np.random.shuffle(hard_neg_yx); np.random.shuffle(easy_neg_yx)\n        neg_yx = hard_neg_yx[:n_hard] + easy_neg_yx[:n_easy]\n\n        n_total = len(pos_yx) + len(neg_yx)\n        if max_patches > 0 and n_total > max_patches:\n            frac = max_patches / n_total\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx = pos_yx[:max(1, int(len(pos_yx)*frac))]\n            neg_yx = neg_yx[:max(0, int(len(neg_yx)*frac))]\n\n        all_yx  = pos_yx + neg_yx\n        all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n        perm    = np.random.permutation(len(all_yx))\n        self.coords  = np.array([all_yx[i] for i in perm], dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm], dtype=np.int32)\n        self.weights = np.where(\n            self.labels == 1, POS_WEIGHT, 1.).astype(np.float32)\n        print(f\"  Seg [{os.path.basename(frag_path)}]: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg \"\n              f\"({len(hard_neg_yx[:n_hard])} hard) = {len(self.coords)}\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x    = int(self.coords[idx, 0]), int(self.coords[idx, 1])\n        raw_vol = np.stack([\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list], axis=-1)\n        msk = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n        if self.tf:\n            out = self.tf(image=raw_vol, mask=msk)\n            raw_vol, msk = out['image'], out['mask']\n        full  = build_input(raw_vol); del raw_vol\n        img_t = norm_patch(\n            torch.from_numpy(full).permute(2,0,1).float()); del full\n        msk_t = torch.from_numpy(msk).unsqueeze(0).float()\n        return img_t, msk_t\n\n\n# ════════════════════════════════════════════════════════════\n#  7. MODELS\n# ════════════════════════════════════════════════════════════\nclass SSLBackbone(nn.Module):\n    def __init__(self, n_in=N_CH, embed_dim=128):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(n_in, 32,  3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(32),  nn.ReLU(inplace=True),\n            nn.Conv2d(32,   64,  3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(64),  nn.ReLU(inplace=True),\n            nn.Conv2d(64, embed_dim, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(embed_dim), nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d(1))\n        self.embed_dim = embed_dim\n\n    def forward(self, x): return self.net(x).flatten(1)\n\n\nclass SSLHead(nn.Module):\n    def __init__(self, embed_dim=128, proj_dim=SSL_PROJ_DIM):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(embed_dim, embed_dim),\n            nn.ReLU(inplace=True),\n            nn.Linear(embed_dim, proj_dim))\n    def forward(self, x): return self.net(x)\n\n\ndef build_seg_model():\n    \"\"\"\n    UNet++ with deeper decoder (256,128,64,32,16) — more\n    capacity to learn fine-grained ink boundaries.\n    \"\"\"\n    for enc in ('efficientnet-b0', 'resnet34', 'resnet18'):\n        try:\n            m = smp.UnetPlusPlus(\n                encoder_name=enc,\n                encoder_weights=None,\n                in_channels=N_CH,\n                classes=1,\n                activation=None,\n                decoder_channels=(256, 128, 64, 32, 16),  # ↑ deeper\n            )\n            print(f\"  Encoder: {enc}\")\n            return m.to(DEVICE)\n        except Exception:\n            continue\n    raise RuntimeError(\"No valid smp encoder found\")\n\n\n# ════════════════════════════════════════════════════════════\n#  8. LOSS FUNCTIONS\n#\n#  FIX: focal loss eps 0.05→0.01 (less label smoothing →\n#  model forced to commit to confident predictions →\n#  better precision)\n# ════════════════════════════════════════════════════════════\ndef nt_xent_loss(z1, z2, temp=SSL_TEMP):\n    B   = z1.shape[0]\n    z   = F.normalize(torch.cat([z1, z2], 0), dim=1)\n    sim = torch.mm(z, z.T) / temp\n    sim.masked_fill_(torch.eye(2*B, device=z.device).bool(), float('-inf'))\n    lbl = torch.cat([torch.arange(B, 2*B, device=z.device),\n                     torch.arange(0, B,   device=z.device)])\n    return F.cross_entropy(sim, lbl)\n\n\ndef tversky_loss(pred, target, a=0.4, b=0.6, smooth=1.):\n    \"\"\"\n    FIX: alpha=0.4, beta=0.6 (was 0.3/0.7).\n    Higher alpha penalises FP more → better precision.\n    Original 0.3/0.7 was optimised for recall which caused\n    the FP/TP=3.82 problem.\n    \"\"\"\n    p  = torch.sigmoid(pred)\n    tp = (p*target).sum(dim=(2, 3))\n    fp = (p*(1-target)).sum(dim=(2, 3))\n    fn = ((1-p)*target).sum(dim=(2, 3))\n    return 1. - ((tp+smooth) / (tp + a*fp + b*fn + smooth)).mean()\n\n\ndef focal_loss(pred, target, alpha=0.75, gamma=2., eps=0.01):\n    \"\"\"\n    FIX: alpha=0.75 (was 0.8), eps=0.01 (was 0.05).\n    Less label smoothing → sharper probability distribution\n    → better separation between ink and background.\n    \"\"\"\n    t_s = target*(1-eps) + 0.5*eps\n    bce = F.binary_cross_entropy_with_logits(pred, t_s, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = t_s*alpha + (1-t_s)*(1-alpha)\n    return (a_t * ((1-p_t)**gamma) * bce).mean()\n\n\ndef seg_loss(pred, target):\n    return 0.5*tversky_loss(pred, target) + 0.5*focal_loss(pred, target)\n\n\n# ════════════════════════════════════════════════════════════\n#  9. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p*masks).sum(dim=(1, 2, 3))\n    union = p.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3))\n    return ((2.*inter + 1e-5) / (union + 1e-5)).mean().item()\n\n\ndef sweep_threshold(probs, targets):\n    \"\"\"\n    FIX: finer sweep (step 0.01 not 0.02) — finds better\n    operating point especially near the precision/recall cliff.\n    \"\"\"\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.10, 0.95, 0.01):\n        p = (probs > t).astype(np.float32)\n        d = (2*(p*targets).sum() + 1) / (p.sum() + targets.sum() + 1)\n        if d > best_d: best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\ndef compute_sep(probs, masks):\n    \"\"\"Mean ink prob - mean bg prob. Target: >0.4\"\"\"\n    if (masks > 0.5).any() and (masks < 0.5).any():\n        return float(probs[masks > 0.5].mean() - probs[masks < 0.5].mean())\n    return 0.\n\n\n# ════════════════════════════════════════════════════════════\n#  10. CONNECTED COMPONENT FILTER\n#  FIX: min_pixels 20→50 (removes more isolated FP blobs)\n# ════════════════════════════════════════════════════════════\ndef cc_filter(binary_mask, min_px=CC_MIN_PIXELS):\n    labeled, n = cc_label(binary_mask)\n    if n == 0: return binary_mask, 0\n    sizes    = np.bincount(labeled.ravel())\n    sizes[0] = 0\n    keep     = sizes >= min_px\n    out      = keep[labeled].astype(np.uint8)\n    removed  = int((sizes[1:] > 0).sum() - keep[1:].sum())\n    return out, removed\n\n\n# ════════════════════════════════════════════════════════════\n#  11. GAUSSIAN WEIGHT MAP\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c = sz // 2; s = sz // 4\n    y, x = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((x-c)**2 + (y-c)**2) / (2*s**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  12. PHASE 1 — SSL PRETRAINING\n# ════════════════════════════════════════════════════════════\ndef run_ssl(backbone, head, dl, n_epochs, lr):\n    params = list(backbone.parameters()) + list(head.parameters())\n    opt    = optim.AdamW(params, lr=lr, weight_decay=WEIGHT_DECAY)\n    sched  = optim.lr_scheduler.CosineAnnealingLR(\n        opt, T_max=n_epochs, eta_min=1e-5)\n    scaler = GradScaler(enabled=USE_AMP)\n\n    print(f'\\n{\"=\"*50}')\n    print(f'PHASE 1 — SSL ({n_epochs} epochs, no labels)')\n    print(f'{\"=\"*50}')\n\n    for ep in range(n_epochs):\n        backbone.train(); head.train()\n        tot = 0.; steps = 0; opt.zero_grad()\n        for v1, v2 in tqdm(dl, desc=f'SSL ep{ep+1}', leave=False):\n            v1 = v1.to(DEVICE); v2 = v2.to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                loss = nt_xent_loss(head(backbone(v1)),\n                                    head(backbone(v2))) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            steps += 1\n            if steps % GRAD_ACCUM == 0:\n                scaler.unscale_(opt)\n                torch.nn.utils.clip_grad_norm_(params, 1.)\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n                if DEVICE == 'cuda': torch.cuda.empty_cache()\n            tot += loss.item() * GRAD_ACCUM\n            del v1, v2, loss\n        sched.step()\n        print(f'  SSL ep{ep+1}/{n_epochs} | loss={tot/max(steps,1):.4f}')\n    return backbone\n\n\ndef transfer_weights(ssl_backbone, seg_model):\n    \"\"\"Scale seg model first-conv channels by SSL importance.\"\"\"\n    print('  Transferring SSL channel weights...')\n    with torch.no_grad():\n        ssl_w  = ssl_backbone.net[0].weight.data   # (32, N_CH, 3, 3)\n        ch_std = ssl_w.std(dim=(0, 2, 3))\n        scale  = (ch_std / (ch_std.mean() + 1e-6)).clamp(0.5, 2.0)\n        for m in seg_model.modules():\n            if isinstance(m, nn.Conv2d) and m.in_channels == N_CH:\n                m.weight.data *= scale.view(1, N_CH, 1, 1)\n                print(f'    Scaled first conv {m.weight.shape}')\n                break\n    print('  Done.')\n    return seg_model\n\n\n# ════════════════════════════════════════════════════════════\n#  13. PHASE 2 — SEGMENTATION TRAINING\n# ════════════════════════════════════════════════════════════\ndef run_seg(model, train_dl, val_dl, n_epochs, lr, ckpt):\n    opt    = optim.AdamW(model.parameters(), lr=lr,\n                         weight_decay=WEIGHT_DECAY)\n    # Warm up for 2 epochs then cosine decay\n    warmup = optim.lr_scheduler.LinearLR(\n        opt, start_factor=0.1, end_factor=1.0, total_iters=2)\n    cosine = optim.lr_scheduler.CosineAnnealingLR(\n        opt, T_max=n_epochs-2, eta_min=1e-6)\n    sched  = optim.lr_scheduler.SequentialLR(\n        opt, [warmup, cosine], milestones=[2])\n    scaler = GradScaler(enabled=USE_AMP)\n\n    best_dice = 0.; pat = 0\n    history   = dict(tl=[], vl=[], td=[], vd=[], sep=[])\n\n    for ep in range(n_epochs):\n        # ── train ──────────────────────────────────────────\n        model.train(); tl = td = 0.; opt.zero_grad()\n        for step, (imgs, msks) in enumerate(\n                tqdm(train_dl, desc=f'Ep{ep+1:02d}▸train', leave=False)):\n            imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n            with autocast(enabled=USE_AMP):\n                out  = model(imgs)\n                loss = seg_loss(out, msks) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            tl += loss.item() * GRAD_ACCUM\n            td += batch_dice(out.detach(), msks)\n            del imgs, msks, out, loss\n            if (step+1) % GRAD_ACCUM == 0:\n                scaler.unscale_(opt)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.)\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n                if DEVICE == 'cuda': torch.cuda.empty_cache()\n        tl /= len(train_dl); td /= len(train_dl)\n\n        # ── val ────────────────────────────────────────────\n        model.eval(); vl = vd = 0.\n        all_p, all_m = [], []\n        with torch.no_grad():\n            for imgs, msks in tqdm(val_dl,\n                                   desc=f'Ep{ep+1:02d}▸val  ', leave=False):\n                imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n                with autocast(enabled=USE_AMP):\n                    out  = model(imgs)\n                    loss = seg_loss(out, msks)\n                vl += loss.item(); vd += batch_dice(out, msks)\n                all_p.append(torch.sigmoid(out).cpu().half().numpy())\n                all_m.append(msks.cpu().half().numpy())\n                del imgs, msks, out, loss\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n        vl /= len(val_dl); vd /= len(val_dl)\n\n        P = np.concatenate(all_p).astype(np.float32)\n        M = np.concatenate(all_m).astype(np.float32)\n        del all_p, all_m\n        bt, bd = sweep_threshold(P, M)\n        sep    = compute_sep(P, M)\n        del P, M; gc.collect()\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n\n        sched.step()\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n        history['sep'].append(sep)\n\n        # Warn if separation is still low after 5 epochs\n        sep_flag = ' ⚠ LOW SEP' if ep >= 4 and sep < 0.3 else ''\n        print(f'Ep{ep+1:02d} | lr={opt.param_groups[0][\"lr\"]:.1e} | '\n              f'train={tl:.4f}/{td:.4f} | val={vl:.4f}/{vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | sep={sep:+.3f}{sep_flag}')\n\n        metric = max(vd, bd)\n        if metric > best_dice:\n            best_dice = metric; pat = 0\n            torch.save({'ep': ep, 'state': model.state_dict(),\n                        'thr': bt, 'dice': best_dice},\n                       OUTPUT + ckpt)\n            print(f'  ✓ checkpoint (dice={best_dice:.4f})')\n        else:\n            pat += 1\n            if pat >= PATIENCE:\n                print(f'  early stop ep{ep+1}'); break\n\n    return model, best_dice, history\n\n\n# ════════════════════════════════════════════════════════════\n#  14. INFERENCE (Gaussian stitching)\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    model.eval()\n    cache = (_cache_store.get((frag_path, tuple(z_list)))\n             or load_ct_cache(frag_path, z_list))\n    H, W  = get_hw(cache)\n    msk   = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n    coords   = [(y, x)\n                for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n                for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='Infer', leave=True):\n            raw = np.stack([\n                cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                for z in z_list], axis=-1)\n            full = build_input(raw); del raw\n            t = norm_patch(\n                torch.from_numpy(full).permute(2, 0, 1).float()\n            ).unsqueeze(0).to(DEVICE); del full\n            with autocast(enabled=USE_AMP):\n                p = torch.sigmoid(model(t)).squeeze().cpu().float().numpy()\n            del t\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del p\n\n    gc.collect()\n    if DEVICE == 'cuda': torch.cuda.empty_cache()\n    return pred_map / (wgt_map + 1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  15. EVALUATION & PLOTS\n# ════════════════════════════════════════════════════════════\ndef evaluate_and_plot(prob_map, gt_mask, label, fname, thr=None):\n    H, W = gt_mask.shape\n    pc   = np.nan_to_num(prob_map[:H, :W], nan=0.5)\n\n    if thr is None:\n        # FIX: sweep on actual test data, not carried from val\n        thr, _ = sweep_threshold(pc, gt_mask)\n        print(f'  Threshold swept on this fragment: {thr:.2f}')\n\n    pred_raw        = (pc > thr).astype(np.uint8)\n    pred, n_removed = cc_filter(pred_raw)\n\n    def metrics(p, m):\n        inter = (p*m).sum()\n        d     = (2*inter + 1) / (p.sum() + m.sum() + 1)\n        tn, fp, fn, tp = confusion_matrix(\n            m.flatten().astype(int), p.flatten().astype(int),\n            labels=[0, 1]).ravel()\n        pr = tp/(tp+fp+1e-8); re = tp/(tp+fn+1e-8)\n        return dict(dice=float(d), prec=pr, rec=re,\n                    f1=2*pr*re/(pr+re+1e-8),\n                    tp=int(tp), fp=int(fp), fn=int(fn), tn=int(tn))\n\n    mr   = metrics(pred_raw, gt_mask)\n    mc   = metrics(pred,     gt_mask)\n    ink_m = float(pc[gt_mask == 1].mean()) if (gt_mask==1).any() else 0.\n    bg_m  = float(pc[gt_mask == 0].mean()) if (gt_mask==0).any() else 0.\n\n    print(f'\\n{\"=\"*55}')\n    print(f'RESULTS — {label}')\n    print(f'{\"=\"*55}')\n    print(f'               Raw      CC-filtered')\n    print(f'Dice       : {mr[\"dice\"]:.4f}    {mc[\"dice\"]:.4f}')\n    print(f'Precision  : {mr[\"prec\"]:.4f}    {mc[\"prec\"]:.4f}')\n    print(f'Recall     : {mr[\"rec\"]:.4f}    {mc[\"rec\"]:.4f}')\n    print(f'F1         : {mr[\"f1\"]:.4f}    {mc[\"f1\"]:.4f}')\n    print(f'FP/TP      : {mr[\"fp\"]/(mr[\"tp\"]+1e-8):.2f}      '\n          f'{mc[\"fp\"]/(mc[\"tp\"]+1e-8):.2f}')\n    print(f'CC removed : {n_removed}')\n    print(f'Threshold  : {thr:.2f}')\n    print(f'Sep (ink-bg): {ink_m - bg_m:+.3f}  '\n          f'(ink={ink_m:.3f} bg={bg_m:.3f})')\n    if mr[\"dice\"] < 0.5:\n        print('  ⚠  Dice<0.5: check sep value — if sep<0.3 the model')\n        print('     has not learned ink signal (training issue, not threshold)')\n    print(f'{\"=\"*55}')\n\n    fig, ax = plt.subplots(2, 3, figsize=(18, 12))\n    ax[0,0].imshow(gt_mask, cmap='gray');  ax[0,0].set_title('Ground Truth')\n    ax[0,1].imshow(pc, cmap='inferno');    ax[0,1].set_title('Probability Map')\n    ax[0,2].imshow(pred, cmap='gray')\n    ax[0,2].set_title(f'Prediction (CC)\\nDice={mc[\"dice\"]:.4f}')\n    err = np.zeros((*gt_mask.shape, 3), dtype=np.uint8)\n    err[(pred==1)&(gt_mask==1)] = [0, 255, 0]\n    err[(pred==1)&(gt_mask==0)] = [255, 0, 0]\n    err[(pred==0)&(gt_mask==1)] = [0, 0, 255]\n    ax[1,0].imshow(err); ax[1,0].set_title('TP=green FP=red FN=blue')\n    iv = pc[gt_mask==1].ravel(); bv = pc[gt_mask==0].ravel()\n    if len(iv): ax[1,1].hist(iv[np.isfinite(iv)], bins=60, alpha=0.7,\n                             label=f'ink μ={ink_m:.3f}',\n                             color='orange', density=True)\n    if len(bv): ax[1,1].hist(bv[np.isfinite(bv)], bins=60, alpha=0.7,\n                             label=f'bg μ={bg_m:.3f}',\n                             color='blue', density=True)\n    ax[1,1].axvline(thr, color='r', ls='--', label=f'thr={thr:.2f}')\n    ax[1,1].legend(); ax[1,1].set_title('Probability Distributions')\n    ts = np.arange(0.05, 0.95, 0.01)\n    ds = [(2*((pc>t)*gt_mask).sum()+1)/((pc>t).sum()+gt_mask.sum()+1)\n          for t in ts]\n    ax[1,2].plot(ts, ds, lw=2); ax[1,2].axvline(thr, color='r', ls='--')\n    ax[1,2].axhline(0.5, color='orange', ls=':')\n    ax[1,2].set_title('Dice vs Threshold'); ax[1,2].grid(True)\n    ax[1,2].set_xlabel('Threshold'); ax[1,2].set_ylabel('Dice')\n    for a in ax[0]: a.axis('off')\n    ax[1,0].axis('off')\n    plt.suptitle(f'{label}\\nRaw={mr[\"dice\"]:.4f}  CC={mc[\"dice\"]:.4f}  '\n                 f'Sep={ink_m-bg_m:+.3f}', fontsize=12)\n    plt.tight_layout()\n    plt.savefig(OUTPUT + fname, dpi=80, bbox_inches='tight')\n    plt.close()\n\n    return dict(raw=mr, cc=mc, thr=thr,\n                ink_mean=ink_m, noink_mean=bg_m,\n                n_removed=n_removed)\n\n\n# ════════════════════════════════════════════════════════════\n#  16. BUILD DATASETS & DATALOADERS\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('DATA SETUP')\nprint(f'  Z-slices : {Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]} ({N_RAW} slices)')\nprint(f'  Channels : {N_CH} ({N_RAW} raw + {N_ENGINEERED} engineered)')\nprint(f'  Patch    : {PATCH_SIZE}×{PATCH_SIZE}')\nprint('='*60)\n\nprint('\\nLoading CT caches...')\nload_ct_cache(FRAG3, Z_SLICES_RAW)\nload_ct_cache(FRAG2, Z_SLICES_RAW)\n\n# SSL datasets — split independently from Seg (different sizes)\nssl_f3     = SSLDataset(FRAG3, Z_SLICES_RAW, stride=STRIDE_TR,\n                        max_patches=MAX_PATCHES)\nssl_f2_all = SSLDataset(FRAG2, Z_SLICES_RAW, stride=STRIDE_TR,\n                        max_patches=0)\nn_ssl      = len(ssl_f2_all)\nssl_perm   = np.random.permutation(n_ssl)\nssl_f2_tr  = Subset(ssl_f2_all, ssl_perm[int(0.2*n_ssl):])\n\n# Seg datasets\nseg_f3     = SegDataset(FRAG3, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=strong_tf, neg_ratio=NEG_RATIO,\n                        max_patches=MAX_PATCHES)\nseg_f2_all = SegDataset(FRAG2, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=strong_tf, neg_ratio=NEG_RATIO,\n                        max_patches=0)\nn_seg      = len(seg_f2_all)\nseg_perm   = np.random.permutation(n_seg)\nseg_trn_idx = seg_perm[int(0.2*n_seg):]\nseg_val_idx = seg_perm[:int(0.2*n_seg)]\nseg_f2_tr   = Subset(seg_f2_all, seg_trn_idx)\nseg_val     = Subset(seg_f2_all, seg_val_idx)\n\nprint(f'  SSL  Frag2 train: {len(ssl_f2_tr)}')\nprint(f'  Seg  Frag2: {len(seg_trn_idx)} train | {len(seg_val_idx)} val')\n\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\n_kw = dict(num_workers=0, pin_memory=False)\n\nssl_ds = ConcatDataset([ssl_f3, ssl_f2_tr])\nssl_dl = DataLoader(ssl_ds, batch_size=BATCH_SIZE*2, shuffle=True, **_kw)\n\nseg_ds    = ConcatDataset([seg_f3, seg_f2_tr])\nw_all     = np.concatenate([seg_f3.weights,\n                             seg_f2_all.weights[seg_trn_idx]])\nsampler   = WeightedRandomSampler(torch.from_numpy(w_all),\n                                  len(seg_ds), True)\nseg_dl_tr = DataLoader(seg_ds, batch_size=BATCH_SIZE,\n                        sampler=sampler, **_kw)\nseg_dl_vl = DataLoader(seg_val, batch_size=BATCH_SIZE,\n                        shuffle=False, **_kw)\n\nprint(f'SSL batches   : {len(ssl_dl)}')\nprint(f'Seg tr batches: {len(seg_dl_tr)} | val: {len(seg_dl_vl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  17. PHASE 1 — SSL\n# ════════════════════════════════════════════════════════════\nssl_bb   = SSLBackbone(n_in=N_CH, embed_dim=128).to(DEVICE)\nssl_head = SSLHead(embed_dim=128, proj_dim=SSL_PROJ_DIM).to(DEVICE)\n\nssl_bb = run_ssl(ssl_bb, ssl_head, ssl_dl,\n                 n_epochs=EPOCHS_SSL, lr=LR_SSL)\ntorch.save({'state': ssl_bb.state_dict()}, OUTPUT+'ssl_backbone.pth')\nprint('SSL backbone saved.')\n\ndel ssl_dl, ssl_ds, ssl_f3, ssl_f2_all, ssl_f2_tr, ssl_head\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\n\n# ════════════════════════════════════════════════════════════\n#  18. PHASE 2 — SEGMENTATION (UNet++)\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint(f'PHASE 2 — UNet++ Segmentation ({EPOCHS_SEG} epochs)')\nprint(f'  Input  : {N_CH} ch | Patch: {PATCH_SIZE}×{PATCH_SIZE}')\nprint(f'  Target : sep>0.4, val Dice>0.65')\nprint('='*60)\n\nseg_model = build_seg_model()\nseg_model = transfer_weights(ssl_bb, seg_model)\ndel ssl_bb; gc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\nseg_model, best_dice, history = run_seg(\n    seg_model, seg_dl_tr, seg_dl_vl,\n    n_epochs=EPOCHS_SEG, lr=LR_SEG, ckpt='best_model.pth')\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\ndel seg_dl_tr, seg_dl_vl, seg_ds, seg_val\ndel seg_f3, seg_f2_all, seg_f2_tr\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\n# Training curves\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\naxes[0].plot(history['tl'], label='train')\naxes[0].plot(history['vl'], label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'], label='train')\naxes[1].plot(history['vd'], label='val')\naxes[1].axhline(0.65, color='orange', ls='--', label='0.65')\naxes[1].axhline(0.80, color='r',      ls='--', label='0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\naxes[2].plot(history['sep'])\naxes[2].axhline(0.4, color='orange', ls='--', label='target=0.4')\naxes[2].set_title('Separation (ink-bg prob)')\naxes[2].set_ylabel('sep'); axes[2].legend(); axes[2].grid(True)\nplt.suptitle('v23 — UNet++ | Precision-optimised')\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves.png', dpi=80); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  19. FINAL TEST — FRAGMENT 1\n#\n#  FIX: threshold is swept on Fragment 1 itself (not carried\n#  from Frag2 val).  Since Frag1 labels are loaded here for\n#  the first time, this is valid — we do ONE sweep to find\n#  the best operating point for this fragment's distribution.\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('FINAL TEST — FRAGMENT 1')\nprint('='*60)\n\nckpt = torch.load(OUTPUT+'best_model.pth', map_location=DEVICE)\nprint(f'Checkpoint: ep={ckpt[\"ep\"]+1}  '\n      f'dice={ckpt[\"dice\"]:.4f}  val_thr={ckpt[\"thr\"]:.2f}')\n\nfinal = build_seg_model()\nfinal.load_state_dict(ckpt['state'])\n\nprob_map, msk1 = predict_fragment(final, FRAG1, Z_SLICES_RAW)\n\n# Sweep threshold on Frag1 (thr=None triggers sweep)\nres = evaluate_and_plot(prob_map, msk1,\n                        label='Fragment 1 (unseen test)',\n                        fname='frag1_prediction.png',\n                        thr=None)\n\nprint(f'\\n★  Dice (raw)      : {res[\"raw\"][\"dice\"]:.4f}')\nprint(f'★  Dice (CC filter): {res[\"cc\"][\"dice\"]:.4f}')\nprint(f'★  Precision       : {res[\"cc\"][\"prec\"]:.4f}')\nprint(f'★  Recall          : {res[\"cc\"][\"rec\"]:.4f}')\nprint(f'★  F1              : {res[\"cc\"][\"f1\"]:.4f}')\nprint(f'★  FP/TP (raw)     : {res[\"raw\"][\"fp\"]/(res[\"raw\"][\"tp\"]+1e-8):.2f}')\nprint(f'★  FP/TP (CC)      : {res[\"cc\"][\"fp\"]/(res[\"cc\"][\"tp\"]+1e-8):.2f}')\nprint(f'★  CC removed      : {res[\"n_removed\"]}')\nprint(f'★  Sep (ink-bg)    : {res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}')\n\n\n# ════════════════════════════════════════════════════════════\n#  20. SAVE RESULTS\n# ════════════════════════════════════════════════════════════\nwith open(OUTPUT+'final_results.txt', 'w') as f:\n    f.write('VESUVIUS INK DETECTION — v23\\n')\n    f.write('='*55 + '\\n')\n    f.write(f'Z-slices : {Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]} ({N_RAW})\\n')\n    f.write(f'Channels : {N_CH} ({N_RAW}+{N_ENGINEERED})\\n')\n    f.write(f'Patch    : {PATCH_SIZE}×{PATCH_SIZE}\\n')\n    f.write(f'Val best : {best_dice:.4f}\\n')\n    f.write('='*55 + '\\n')\n    for tag, m in [('Raw', res['raw']), ('CC-filtered', res['cc'])]:\n        f.write(f'[{tag}]\\n')\n        for k in ['dice', 'prec', 'rec', 'f1']:\n            f.write(f'  {k}: {m[k]:.4f}\\n')\n        f.write(f'  FP/TP: {m[\"fp\"]/(m[\"tp\"]+1e-8):.2f}\\n')\n    f.write(f'Threshold : {res[\"thr\"]:.2f}\\n')\n    f.write(f'CC removed: {res[\"n_removed\"]}\\n')\n    f.write(f'Sep       : {res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  ssl_backbone.pth | best_model.pth')\nprint('  curves.png | frag1_prediction.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\n# ============================================================\n#  VESUVIUS INK DETECTION — v22  (FINAL — PyTorch compat fix)\n#\n#  Fixes vs previous version:\n#  1. torch.amp.GradScaler → torch.cuda.amp.GradScaler\n#     (torch.amp.GradScaler only exists in PyTorch ≥ 2.3;\n#      Kaggle P100 kernels run 1.x / early 2.x)\n#  2. torch.amp.autocast  → torch.cuda.amp.autocast\n#  3. torch.load weights_only=False → removed (old PyTorch\n#     doesn't accept that kwarg at all)\n#  4. IndexError fix: SSL and Seg Frag2 split indices are\n#     computed INDEPENDENTLY (different dataset sizes)\n#  5. Architecture: UNet++ with efficientnet-b0 encoder\n#  6. All other OOM fixes retained (patch=64, 16 z-slices,\n#     20 channels, NUM_WORKERS=0, pin_memory=False, etc.)\n# ============================================================\n\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import (Dataset, DataLoader,\n                               WeightedRandomSampler, ConcatDataset, Subset)\nfrom torch.cuda.amp import GradScaler, autocast   # ← compat fix\nimport albumentations as A\nfrom scipy.ndimage import median_filter, label as cc_label\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nimport warnings\nwarnings.filterwarnings('ignore')\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:32\"\n\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = False\n\n# ── paths ──────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ───────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 64\nSTRIDE_TR    = 32\nSTRIDE_INF   = 24\nBATCH_SIZE   = 4\nGRAD_ACCUM   = 4        # effective batch = 16\nEPOCHS_SSL   = 3\nEPOCHS_SEG   = 12\nLR_SSL       = 3e-4\nLR_SEG       = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 8\nMAX_PATCHES  = 6_000\n\nZ_START      = 18\nZ_END        = 34       # slices 23-38 inclusive = 16 slices\nZ_SLICES_RAW = list(range(Z_START, Z_END))\nN_RAW        = len(Z_SLICES_RAW)   # 16\nN_ENGINEERED = 4\nN_CH         = N_RAW + N_ENGINEERED  # 20\n\nSSL_PROJ_DIM  = 64\nSSL_TEMP      = 0.07\nCC_MIN_PIXELS = 20\nINK_MIN_POS   = 0.02\nNEG_RATIO     = 0.3\nPOS_WEIGHT    = 3.0\nUSE_AMP       = (DEVICE == 'cuda')\nNORM_STD_MIN  = 0.01\nNORM_CLIP     = 10.0\n\nprint(f\"Device  : {DEVICE}\")\nprint(f\"Input   : {N_RAW} raw z-slices + {N_ENGINEERED} engineered = {N_CH} ch\")\nprint(f\"Patch   : {PATCH_SIZE}×{PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU     : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM    : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. CT CACHE\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_ct_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        print(f\"  Reusing cache: {os.path.basename(frag_path)}\")\n        return _cache_store[key]\n\n    vol_dir  = os.path.join(frag_path, 'surface_volume')\n    slices, all_vals = [], []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            slices.append((z, s))\n            all_vals.append(s.ravel()[::100])   # 1% sample for percentile\n\n    assert slices, f\"No CT slices found in {vol_dir}\"\n    p5, p95 = np.percentile(np.concatenate(all_vals), [5, 95])\n    del all_vals; gc.collect()\n\n    cache = {}; H = W = None\n    for z, s in slices:\n        s = np.clip(s, p5, p95)\n        s = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n        cache[z] = s\n        if H is None: H, W = s.shape\n    del slices; gc.collect()\n\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Cache [{os.path.basename(frag_path)}]: \"\n          f\"{len(cache)} slices, {mb:.0f} MB, ({H},{W})\")\n    _cache_store[key] = cache\n    return cache\n\n\ndef get_hw(cache):\n    return next(iter(cache.values())).shape\n\n\ndef norm_patch(t):\n    mu  = t.mean(dim=(1, 2), keepdim=True)\n    std = t.std (dim=(1, 2), keepdim=True).clamp(min=NORM_STD_MIN)\n    return ((t - mu) / std).clamp(-NORM_CLIP, NORM_CLIP)\n\n\n# ════════════════════════════════════════════════════════════\n#  2. PHYSICS CHANNELS (4, patch-level only)\n# ════════════════════════════════════════════════════════════\ndef physics_enhance(raw_vol):\n    \"\"\"raw_vol: (H,W,N_RAW) float32 patch → (H,W,4) float32\"\"\"\n    ch0 = raw_vol.max(axis=2)                               # MIP\n    ch1 = raw_vol.std(axis=2)                               # z-std\n    ch2 = np.abs(np.diff(raw_vol, axis=2)).max(axis=2)      # max diff\n    mid = raw_vol.shape[2] // 2\n    sl  = raw_vol[:, :, mid]\n    ch3 = sl - median_filter(sl, size=3)                    # median residual\n    return np.stack([ch0, ch1, ch2, ch3], axis=-1).astype(np.float32)\n\n\ndef build_input(raw_vol):\n    \"\"\"(H,W,N_RAW) float32 → (H,W,N_CH) float32\"\"\"\n    eng = physics_enhance(raw_vol)\n    out = np.concatenate([raw_vol, eng], axis=-1)\n    del eng\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  3. PAPYRUS MASK\n# ════════════════════════════════════════════════════════════\ndef load_papyrus_mask(frag_path, H, W, cache):\n    mp = os.path.join(frag_path, 'mask.png')\n    if os.path.exists(mp):\n        m = cv2.imread(mp, 0)\n        if m is not None:\n            if m.shape != (H, W):\n                m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n            return (m > 0).astype(np.uint8)\n    z = list(cache.keys())[len(cache) // 2]\n    return (cache[z] > 0.1).astype(np.uint8)\n\n\n# ════════════════════════════════════════════════════════════\n#  4. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\nstrong_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Affine(scale=(0.9, 1.1),\n             translate_percent={'x': (-0.05, 0.05), 'y': (-0.05, 0.05)},\n             rotate=(-20, 20),\n             mode=cv2.BORDER_REFLECT, p=0.5),\n    A.RandomBrightnessContrast(0.2, 0.2, p=0.5),\n    A.GaussianBlur(blur_limit=(3, 5), p=0.3),\n    A.CoarseDropout(max_holes=4, max_height=16, max_width=16,\n                    fill_value=0, p=0.3),\n])\n\nweak_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  5. SSL DATASET\n# ════════════════════════════════════════════════════════════\nclass SSLDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride, max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        H, W        = get_hw(self.cache)\n        pap         = load_papyrus_mask(frag_path, H, W, self.cache)\n\n        coords = [(y, x)\n                  for y in range(0, H - PATCH_SIZE + 1, stride)\n                  for x in range(0, W - PATCH_SIZE + 1, stride)\n                  if pap[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.4]\n        if max_patches > 0 and len(coords) > max_patches:\n            np.random.shuffle(coords)\n            coords = coords[:max_patches]\n        self.coords = np.array(coords, dtype=np.int32)\n        print(f\"  SSL [{os.path.basename(frag_path)}]: {len(self.coords)} patches\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x    = int(self.coords[idx, 0]), int(self.coords[idx, 1])\n        raw_vol = np.stack([\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list], axis=-1)\n        v1 = weak_tf  (image=raw_vol.copy())['image']\n        v2 = strong_tf(image=raw_vol.copy())['image']\n        del raw_vol\n        f1 = build_input(v1); del v1\n        f2 = build_input(v2); del v2\n        t1 = norm_patch(torch.from_numpy(f1).permute(2, 0, 1).float()); del f1\n        t2 = norm_patch(torch.from_numpy(f2).permute(2, 0, 1).float()); del f2\n        return t1, t2\n\n\n# ════════════════════════════════════════════════════════════\n#  6. SEGMENTATION DATASET\n# ════════════════════════════════════════════════════════════\nclass SegDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0, max_patches=0):\n        self.cache  = load_ct_cache(frag_path, z_list)\n        self.z_list = z_list\n        self.tf     = transform\n        H, W        = get_hw(self.cache)\n\n        msk_p = os.path.join(frag_path, 'inklabels.png')\n        msk   = cv2.imread(msk_p, 0)\n        assert msk is not None, f\"Missing ink labels: {msk_p}\"\n        self.mask = (msk > 0).astype(np.uint8); del msk\n\n        pap = load_papyrus_mask(frag_path, H, W, self.cache)\n        pos_yx, neg_yx = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                if pap[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() < 0.4:\n                    continue\n                ink = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                if   ink >= INK_MIN_POS:        pos_yx.append((y, x))\n                elif ink < 0.001 and neg_ratio: neg_yx.append((y, x))\n\n        n_neg = int(len(pos_yx) * neg_ratio)\n        if n_neg and neg_yx:\n            np.random.shuffle(neg_yx); neg_yx = neg_yx[:n_neg]\n        else:\n            neg_yx = []\n\n        n_total = len(pos_yx) + len(neg_yx)\n        if max_patches > 0 and n_total > max_patches:\n            frac = max_patches / n_total\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx = pos_yx[:max(1, int(len(pos_yx)*frac))]\n            neg_yx = neg_yx[:max(0, int(len(neg_yx)*frac))]\n\n        all_yx  = pos_yx + neg_yx\n        all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n        perm    = np.random.permutation(len(all_yx))\n        self.coords  = np.array([all_yx[i] for i in perm], dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm], dtype=np.int32)\n        self.weights = np.where(self.labels == 1, POS_WEIGHT, 1.).astype(np.float32)\n        print(f\"  Seg [{os.path.basename(frag_path)}]: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg = {len(self.coords)}\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x    = int(self.coords[idx, 0]), int(self.coords[idx, 1])\n        raw_vol = np.stack([\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list], axis=-1)\n        msk = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n        if self.tf:\n            out = self.tf(image=raw_vol, mask=msk)\n            raw_vol, msk = out['image'], out['mask']\n        full  = build_input(raw_vol); del raw_vol\n        img_t = norm_patch(torch.from_numpy(full).permute(2, 0, 1).float()); del full\n        msk_t = torch.from_numpy(msk).unsqueeze(0).float()\n        return img_t, msk_t\n\n\n# ════════════════════════════════════════════════════════════\n#  7. MODELS\n# ════════════════════════════════════════════════════════════\n\n# ── 7a: SSL backbone ─────────────────────────────────────────\nclass SSLBackbone(nn.Module):\n    def __init__(self, n_in=N_CH, embed_dim=128):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(n_in, 32,  3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(32),  nn.ReLU(inplace=True),\n            nn.Conv2d(32,   64,  3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(64),  nn.ReLU(inplace=True),\n            nn.Conv2d(64, embed_dim, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(embed_dim), nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d(1))\n        self.embed_dim = embed_dim\n\n    def forward(self, x): return self.net(x).flatten(1)\n\n\nclass SSLHead(nn.Module):\n    def __init__(self, embed_dim=128, proj_dim=SSL_PROJ_DIM):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(embed_dim, embed_dim),\n            nn.ReLU(inplace=True),\n            nn.Linear(embed_dim, proj_dim))\n    def forward(self, x): return self.net(x)\n\n\n# ── 7b: UNet++ segmentation model ────────────────────────────\ndef build_seg_model():\n    \"\"\"\n    UNet++ > plain UNet for thin-stroke ink detection:\n    dense nested skips preserve fine spatial detail that\n    single-hop UNet skips tend to blur at each resolution step.\n    Typical gain: +1–3 Dice over UNet with same encoder.\n    Memory overhead: ~15% more VRAM — fine on 16 GB P100 at 64×64.\n    \"\"\"\n    for enc in ('efficientnet-b5', 'resnet18'):\n        try:\n            m = smp.UnetPlusPlus(\n                encoder_name=enc,\n                encoder_weights=None,\n                in_channels=N_CH,\n                classes=1,\n                activation=None,\n                decoder_channels=(128, 64, 32, 16, 8),\n            )\n            print(f\"  Encoder: {enc}\")\n            return m.to(DEVICE)\n        except Exception:\n            continue\n    raise RuntimeError(\"No valid smp encoder found\")\n\n\n# ════════════════════════════════════════════════════════════\n#  8. LOSS FUNCTIONS\n# ════════════════════════════════════════════════════════════\ndef nt_xent_loss(z1, z2, temp=SSL_TEMP):\n    B   = z1.shape[0]\n    z   = F.normalize(torch.cat([z1, z2], 0), dim=1)\n    sim = torch.mm(z, z.T) / temp\n    sim.masked_fill_(torch.eye(2*B, device=z.device).bool(), float('-inf'))\n    lbl = torch.cat([torch.arange(B, 2*B, device=z.device),\n                     torch.arange(0, B,   device=z.device)])\n    return F.cross_entropy(sim, lbl)\n\n\ndef tversky_loss(pred, target, a=0.3, b=0.7, smooth=1.):\n    p  = torch.sigmoid(pred)\n    tp = (p*target).sum(dim=(2, 3))\n    fp = (p*(1-target)).sum(dim=(2, 3))\n    fn = ((1-p)*target).sum(dim=(2, 3))\n    return 1. - ((tp+smooth) / (tp + a*fp + b*fn + smooth)).mean()\n\n\ndef focal_loss(pred, target, alpha=0.8, gamma=2., eps=0.05):\n    t_s = target*(1-eps) + 0.5*eps\n    bce = F.binary_cross_entropy_with_logits(pred, t_s, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = t_s*alpha + (1-t_s)*(1-alpha)\n    return (a_t * ((1-p_t)**gamma) * bce).mean()\n\n\ndef seg_loss(pred, target):\n    return 0.5*tversky_loss(pred, target) + 0.5*focal_loss(pred, target)\n\n\n# ════════════════════════════════════════════════════════════\n#  9. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p*masks).sum(dim=(1, 2, 3))\n    union = p.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3))\n    return ((2.*inter + 1e-5) / (union + 1e-5)).mean().item()\n\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.20, 0.85, 0.02):\n        p = (probs > t).astype(np.float32)\n        d = (2*(p*targets).sum() + 1) / (p.sum() + targets.sum() + 1)\n        if d > best_d: best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  10. CONNECTED COMPONENT FILTER\n# ════════════════════════════════════════════════════════════\ndef cc_filter(binary_mask, min_px=CC_MIN_PIXELS):\n    labeled, n = cc_label(binary_mask)\n    if n == 0: return binary_mask, 0\n    sizes    = np.bincount(labeled.ravel())\n    sizes[0] = 0\n    keep     = sizes >= min_px\n    out      = keep[labeled].astype(np.uint8)\n    removed  = int((sizes[1:] > 0).sum() - keep[1:].sum())\n    return out, removed\n\n\n# ════════════════════════════════════════════════════════════\n#  11. GAUSSIAN WEIGHT MAP\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c = sz // 2; s = sz // 4\n    y, x = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((x-c)**2 + (y-c)**2) / (2*s**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  12. PHASE 1 — SSL PRETRAINING\n# ════════════════════════════════════════════════════════════\ndef run_ssl(backbone, head, dl, n_epochs, lr):\n    params = list(backbone.parameters()) + list(head.parameters())\n    opt    = optim.AdamW(params, lr=lr, weight_decay=WEIGHT_DECAY)\n    sched  = optim.lr_scheduler.CosineAnnealingLR(\n        opt, T_max=n_epochs, eta_min=1e-5)\n    scaler = GradScaler(enabled=USE_AMP)   # ← torch.cuda.amp.GradScaler\n\n    print(f'\\n{\"=\"*50}')\n    print(f'PHASE 1 — SSL ({n_epochs} epochs, no labels)')\n    print(f'{\"=\"*50}')\n\n    for ep in range(n_epochs):\n        backbone.train(); head.train()\n        tot = 0.; steps = 0\n        opt.zero_grad()\n        for v1, v2 in tqdm(dl, desc=f'SSL ep{ep+1}', leave=False):\n            v1 = v1.to(DEVICE); v2 = v2.to(DEVICE)\n            with autocast(enabled=USE_AMP):   # ← torch.cuda.amp.autocast\n                loss = nt_xent_loss(head(backbone(v1)),\n                                    head(backbone(v2))) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            steps += 1\n            if steps % GRAD_ACCUM == 0:\n                scaler.unscale_(opt)\n                torch.nn.utils.clip_grad_norm_(params, 1.)\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n                if DEVICE == 'cuda': torch.cuda.empty_cache()\n            tot += loss.item() * GRAD_ACCUM\n            del v1, v2, loss\n        sched.step()\n        print(f'  SSL ep{ep+1}/{n_epochs} | loss={tot/max(steps,1):.4f}')\n    return backbone\n\n\ndef transfer_weights(ssl_backbone, seg_model):\n    \"\"\"Scale seg model first-conv input channels by SSL importance.\"\"\"\n    print('  Transferring SSL channel weights...')\n    with torch.no_grad():\n        ssl_w  = ssl_backbone.net[0].weight.data     # (32, N_CH, 3, 3)\n        ch_std = ssl_w.std(dim=(0, 2, 3))            # (N_CH,)\n        scale  = (ch_std / (ch_std.mean() + 1e-6)).clamp(0.5, 2.0)\n        for m in seg_model.modules():\n            if isinstance(m, nn.Conv2d) and m.in_channels == N_CH:\n                m.weight.data *= scale.view(1, N_CH, 1, 1)\n                print(f'    Applied scale to {m.weight.shape} conv')\n                break\n    print('  Done.')\n    return seg_model\n\n\n# ════════════════════════════════════════════════════════════\n#  13. PHASE 2 — SEGMENTATION TRAINING\n# ════════════════════════════════════════════════════════════\ndef run_seg(model, train_dl, val_dl, n_epochs, lr, ckpt):\n    opt    = optim.AdamW(model.parameters(), lr=lr, weight_decay=WEIGHT_DECAY)\n    sched  = optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        opt, T_0=8, T_mult=1, eta_min=1e-6)\n    scaler = GradScaler(enabled=USE_AMP)   # ← torch.cuda.amp.GradScaler\n\n    best_dice = 0.; pat = 0\n    history   = dict(tl=[], vl=[], td=[], vd=[])\n\n    for ep in range(n_epochs):\n        # ── train ──────────────────────────────────────────\n        model.train(); tl = td = 0.\n        opt.zero_grad()\n        for step, (imgs, msks) in enumerate(\n                tqdm(train_dl, desc=f'Ep{ep+1:02d}▸train', leave=False)):\n            imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n            with autocast(enabled=USE_AMP):   # ← torch.cuda.amp.autocast\n                out  = model(imgs)\n                loss = seg_loss(out, msks) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            tl += loss.item() * GRAD_ACCUM\n            td += batch_dice(out.detach(), msks)\n            del imgs, msks, out, loss\n            if (step+1) % GRAD_ACCUM == 0:\n                scaler.unscale_(opt)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.)\n                scaler.step(opt); scaler.update(); opt.zero_grad()\n                if DEVICE == 'cuda': torch.cuda.empty_cache()\n        tl /= len(train_dl); td /= len(train_dl)\n\n        # ── val ────────────────────────────────────────────\n        model.eval(); vl = vd = 0.\n        all_p, all_m = [], []\n        with torch.no_grad():\n            for imgs, msks in tqdm(val_dl, desc=f'Ep{ep+1:02d}▸val  ', leave=False):\n                imgs = imgs.to(DEVICE); msks = msks.to(DEVICE)\n                with autocast(enabled=USE_AMP):\n                    out  = model(imgs)\n                    loss = seg_loss(out, msks)\n                vl += loss.item(); vd += batch_dice(out, msks)\n                all_p.append(torch.sigmoid(out).cpu().half().numpy())\n                all_m.append(msks.cpu().half().numpy())\n                del imgs, msks, out, loss\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n        vl /= len(val_dl); vd /= len(val_dl)\n\n        P = np.concatenate(all_p).astype(np.float32)\n        M = np.concatenate(all_m).astype(np.float32)\n        del all_p, all_m\n        bt, bd = sweep_threshold(P, M)\n        sep    = (P[M > 0.5].mean() - P[M < 0.5].mean()\n                  if (M > 0.5).any() and (M < 0.5).any() else 0.)\n        del P, M; gc.collect()\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n\n        sched.step()\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n\n        print(f'Ep{ep+1:02d} | lr={opt.param_groups[0][\"lr\"]:.1e} | '\n              f'train={tl:.4f}/{td:.4f} | val={vl:.4f}/{vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | sep={sep:+.3f}')\n\n        metric = max(vd, bd)\n        if metric > best_dice:\n            best_dice = metric; pat = 0\n            torch.save({'ep': ep, 'state': model.state_dict(),\n                        'thr': bt, 'dice': best_dice},\n                       OUTPUT + ckpt)\n            print(f'  ✓ checkpoint (dice={best_dice:.4f})')\n        else:\n            pat += 1\n            if pat >= PATIENCE:\n                print(f'  early stop ep{ep+1}'); break\n\n    return model, best_dice, history\n\n\n# ════════════════════════════════════════════════════════════\n#  14. INFERENCE (Gaussian stitching)\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    model.eval()\n    cache = (_cache_store.get((frag_path, tuple(z_list)))\n             or load_ct_cache(frag_path, z_list))\n    H, W  = get_hw(cache)\n    msk   = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n    coords   = [(y, x)\n                for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n                for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='Infer', leave=True):\n            raw = np.stack([\n                cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                for z in z_list], axis=-1)\n            full = build_input(raw); del raw\n            t = norm_patch(\n                torch.from_numpy(full).permute(2, 0, 1).float()\n            ).unsqueeze(0).to(DEVICE); del full\n            with autocast(enabled=USE_AMP):\n                p = torch.sigmoid(model(t)).squeeze().cpu().float().numpy()\n            del t\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del p\n\n    gc.collect()\n    if DEVICE == 'cuda': torch.cuda.empty_cache()\n    return pred_map / (wgt_map + 1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  15. EVALUATION & PLOTS\n# ════════════════════════════════════════════════════════════\ndef evaluate_and_plot(prob_map, gt_mask, label, fname, thr=None):\n    H, W = gt_mask.shape\n    pc   = np.nan_to_num(prob_map[:H, :W], nan=0.5)\n    if thr is None:\n        thr, _ = sweep_threshold(pc[None, None], gt_mask[None, None])\n\n    pred_raw        = (pc > thr).astype(np.uint8)\n    pred, n_removed = cc_filter(pred_raw)\n\n    def metrics(p, m):\n        inter = (p*m).sum()\n        d     = (2*inter + 1) / (p.sum() + m.sum() + 1)\n        tn, fp, fn, tp = confusion_matrix(\n            m.flatten().astype(int), p.flatten().astype(int),\n            labels=[0, 1]).ravel()\n        pr = tp/(tp+fp+1e-8); re = tp/(tp+fn+1e-8)\n        return dict(dice=float(d), prec=pr, rec=re,\n                    f1=2*pr*re/(pr+re+1e-8),\n                    tp=int(tp), fp=int(fp), fn=int(fn), tn=int(tn))\n\n    mr   = metrics(pred_raw, gt_mask)\n    mc   = metrics(pred,     gt_mask)\n    ink_m = pc[gt_mask == 1].mean() if (gt_mask == 1).any() else 0.\n    bg_m  = pc[gt_mask == 0].mean() if (gt_mask == 0).any() else 0.\n\n    print(f'\\n{\"=\"*55}')\n    print(f'RESULTS — {label}')\n    print(f'{\"=\"*55}')\n    print(f'               Raw      CC-filtered')\n    print(f'Dice       : {mr[\"dice\"]:.4f}    {mc[\"dice\"]:.4f}')\n    print(f'Precision  : {mr[\"prec\"]:.4f}    {mc[\"prec\"]:.4f}')\n    print(f'Recall     : {mr[\"rec\"]:.4f}    {mc[\"rec\"]:.4f}')\n    print(f'F1         : {mr[\"f1\"]:.4f}    {mc[\"f1\"]:.4f}')\n    print(f'FP/TP      : {mr[\"fp\"]/(mr[\"tp\"]+1e-8):.2f}      '\n          f'{mc[\"fp\"]/(mc[\"tp\"]+1e-8):.2f}')\n    print(f'CC removed : {n_removed}')\n    print(f'Threshold  : {thr:.2f}')\n    print(f'Sep        : {ink_m - bg_m:+.3f}')\n\n    fig, ax = plt.subplots(2, 3, figsize=(18, 12))\n    ax[0,0].imshow(gt_mask, cmap='gray');  ax[0,0].set_title('GT')\n    ax[0,1].imshow(pc, cmap='inferno');    ax[0,1].set_title('Prob map')\n    ax[0,2].imshow(pred, cmap='gray')\n    ax[0,2].set_title(f'Pred CC  Dice={mc[\"dice\"]:.4f}')\n    err = np.zeros((*gt_mask.shape, 3), dtype=np.uint8)\n    err[(pred==1)&(gt_mask==1)] = [0, 255, 0]\n    err[(pred==1)&(gt_mask==0)] = [255, 0, 0]\n    err[(pred==0)&(gt_mask==1)] = [0, 0, 255]\n    ax[1,0].imshow(err); ax[1,0].set_title('TP/FP/FN')\n    iv = pc[gt_mask==1].ravel(); bv = pc[gt_mask==0].ravel()\n    if len(iv): ax[1,1].hist(iv[np.isfinite(iv)], bins=50, alpha=0.7,\n                             label=f'ink μ={ink_m:.2f}', color='orange', density=True)\n    if len(bv): ax[1,1].hist(bv[np.isfinite(bv)], bins=50, alpha=0.7,\n                             label=f'bg μ={bg_m:.2f}', color='blue', density=True)\n    ax[1,1].axvline(thr, color='r', ls='--'); ax[1,1].legend()\n    ax[1,1].set_title('Distributions')\n    ts = np.arange(0.1, 0.9, 0.02)\n    ds = [(2*((pc>t)*gt_mask).sum()+1)/((pc>t).sum()+gt_mask.sum()+1) for t in ts]\n    ax[1,2].plot(ts, ds); ax[1,2].axvline(thr, color='r', ls='--')\n    ax[1,2].set_title('Dice vs Threshold'); ax[1,2].grid(True)\n    for a in ax[0]: a.axis('off')\n    ax[1,0].axis('off')\n    plt.suptitle(f'{label} | Raw={mr[\"dice\"]:.4f}  CC={mc[\"dice\"]:.4f}')\n    plt.tight_layout()\n    plt.savefig(OUTPUT + fname, dpi=80, bbox_inches='tight')\n    plt.close()\n\n    return dict(raw=mr, cc=mc, thr=thr,\n                ink_mean=float(ink_m), noink_mean=float(bg_m),\n                n_removed=n_removed)\n\n\n# ════════════════════════════════════════════════════════════\n#  16. BUILD DATASETS & DATALOADERS\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('DATA SETUP')\nprint(f'  Z-slices : {Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]} ({N_RAW} slices)')\nprint(f'  Channels : {N_CH} ({N_RAW} raw + {N_ENGINEERED} engineered)')\nprint(f'  Patch    : {PATCH_SIZE}×{PATCH_SIZE}')\nprint('='*60)\n\nprint('\\nLoading CT caches...')\nload_ct_cache(FRAG3, Z_SLICES_RAW)\nload_ct_cache(FRAG2, Z_SLICES_RAW)\n\n# ── SSL datasets ───────────────────────────────────────────────\n# Split indices computed from EACH dataset's own length.\n# SSL and Seg enumerate patches differently (no label filter\n# for SSL) so they have different total counts for Frag2.\nssl_f3     = SSLDataset(FRAG3, Z_SLICES_RAW, stride=STRIDE_TR,\n                        max_patches=MAX_PATCHES)\nssl_f2_all = SSLDataset(FRAG2, Z_SLICES_RAW, stride=STRIDE_TR,\n                        max_patches=0)\nn_ssl      = len(ssl_f2_all)\nssl_perm   = np.random.permutation(n_ssl)\nssl_trn_idx = ssl_perm[int(0.2*n_ssl):]\nssl_f2_tr   = Subset(ssl_f2_all, ssl_trn_idx)\n\n# ── Seg datasets ───────────────────────────────────────────────\nseg_f3     = SegDataset(FRAG3, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=strong_tf, neg_ratio=NEG_RATIO,\n                        max_patches=MAX_PATCHES)\nseg_f2_all = SegDataset(FRAG2, Z_SLICES_RAW, stride=STRIDE_TR,\n                        transform=strong_tf, neg_ratio=NEG_RATIO,\n                        max_patches=0)\nn_seg      = len(seg_f2_all)\nseg_perm   = np.random.permutation(n_seg)\nseg_trn_idx = seg_perm[int(0.2*n_seg):]\nseg_val_idx = seg_perm[:int(0.2*n_seg)]\nseg_f2_tr   = Subset(seg_f2_all, seg_trn_idx)\nseg_val     = Subset(seg_f2_all, seg_val_idx)\n\nprint(f'  SSL  Frag2: {len(ssl_trn_idx)} train')\nprint(f'  Seg  Frag2: {len(seg_trn_idx)} train | {len(seg_val_idx)} val')\n\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\n_kw = dict(num_workers=0, pin_memory=False)\n\nssl_ds = ConcatDataset([ssl_f3, ssl_f2_tr])\nssl_dl = DataLoader(ssl_ds, batch_size=BATCH_SIZE*2, shuffle=True, **_kw)\n\nseg_ds    = ConcatDataset([seg_f3, seg_f2_tr])\nw_all     = np.concatenate([seg_f3.weights,\n                             seg_f2_all.weights[seg_trn_idx]])  # ← correct index\nsampler   = WeightedRandomSampler(torch.from_numpy(w_all), len(seg_ds), True)\nseg_dl_tr = DataLoader(seg_ds, batch_size=BATCH_SIZE, sampler=sampler, **_kw)\nseg_dl_vl = DataLoader(seg_val, batch_size=BATCH_SIZE, shuffle=False, **_kw)\n\nprint(f'SSL batches   : {len(ssl_dl)}')\nprint(f'Seg tr batches: {len(seg_dl_tr)} | val: {len(seg_dl_vl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  17. PHASE 1 — SSL\n# ════════════════════════════════════════════════════════════\nssl_bb   = SSLBackbone(n_in=N_CH, embed_dim=128).to(DEVICE)\nssl_head = SSLHead(embed_dim=128, proj_dim=SSL_PROJ_DIM).to(DEVICE)\n\nssl_bb = run_ssl(ssl_bb, ssl_head, ssl_dl,\n                 n_epochs=EPOCHS_SSL, lr=LR_SSL)\n\ntorch.save({'state': ssl_bb.state_dict()}, OUTPUT+'ssl_backbone.pth')\nprint('SSL backbone saved.')\n\ndel ssl_dl, ssl_ds, ssl_f3, ssl_f2_all, ssl_f2_tr, ssl_head\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\n\n# ════════════════════════════════════════════════════════════\n#  18. PHASE 2 — SEGMENTATION (UNet++)\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint(f'PHASE 2 — UNet++ Segmentation ({EPOCHS_SEG} epochs)')\nprint(f'  Input  : {N_CH} ch | Patch: {PATCH_SIZE}×{PATCH_SIZE}')\nprint('='*60)\n\nseg_model = build_seg_model()\nseg_model = transfer_weights(ssl_bb, seg_model)\ndel ssl_bb; gc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\nseg_model, best_dice, history = run_seg(\n    seg_model, seg_dl_tr, seg_dl_vl,\n    n_epochs=EPOCHS_SEG, lr=LR_SEG, ckpt='best_model.pth')\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\ndel seg_dl_tr, seg_dl_vl, seg_ds, seg_val\ndel seg_f3, seg_f2_all, seg_f2_tr\ngc.collect()\nif DEVICE == 'cuda': torch.cuda.empty_cache()\n\n# Training curves\nfig, axes = plt.subplots(1, 2, figsize=(12, 5))\naxes[0].plot(history['tl'], label='train'); axes[0].plot(history['vl'], label='val')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].grid(True)\naxes[1].plot(history['td'], label='train'); axes[1].plot(history['vd'], label='val')\naxes[1].axhline(0.65, color='orange', ls='--', label='0.65')\naxes[1].axhline(0.80, color='r',      ls='--', label='0.80')\naxes[1].set_title('Dice'); axes[1].legend(); axes[1].grid(True)\nplt.suptitle('v22 Final — UNet++')\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves.png', dpi=80); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  19. FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('FINAL TEST — FRAGMENT 1')\nprint('='*60)\n\n# ← weights_only kwarg removed: not supported in older PyTorch\nckpt = torch.load(OUTPUT+'best_model.pth', map_location=DEVICE)\nprint(f'Checkpoint: ep={ckpt[\"ep\"]+1}  '\n      f'dice={ckpt[\"dice\"]:.4f}  thr={ckpt[\"thr\"]:.2f}')\n\nfinal = build_seg_model()\nfinal.load_state_dict(ckpt['state'])\n\nprob_map, msk1 = predict_fragment(final, FRAG1, Z_SLICES_RAW)\nres = evaluate_and_plot(prob_map, msk1,\n                        label='Fragment 1 (unseen test)',\n                        fname='frag1_prediction.png',\n                        thr=ckpt['thr'])\n\nprint(f'\\n★  Dice (raw)      : {res[\"raw\"][\"dice\"]:.4f}')\nprint(f'★  Dice (CC filter): {res[\"cc\"][\"dice\"]:.4f}')\nprint(f'★  Precision       : {res[\"cc\"][\"prec\"]:.4f}')\nprint(f'★  Recall          : {res[\"cc\"][\"rec\"]:.4f}')\nprint(f'★  F1              : {res[\"cc\"][\"f1\"]:.4f}')\nprint(f'★  FP/TP (raw)     : {res[\"raw\"][\"fp\"]/(res[\"raw\"][\"tp\"]+1e-8):.2f}')\nprint(f'★  FP/TP (CC)      : {res[\"cc\"][\"fp\"]/(res[\"cc\"][\"tp\"]+1e-8):.2f}')\nprint(f'★  CC removed      : {res[\"n_removed\"]}')\n\n\n# ════════════════════════════════════════════════════════════\n#  20. SAVE RESULTS\n# ════════════════════════════════════════════════════════════\nwith open(OUTPUT+'final_results.txt', 'w') as f:\n    f.write('VESUVIUS INK DETECTION — v22 FINAL\\n')\n    f.write('='*55 + '\\n')\n    f.write(f'Z-slices : {Z_SLICES_RAW[0]}-{Z_SLICES_RAW[-1]} ({N_RAW})\\n')\n    f.write(f'Channels : {N_CH} ({N_RAW}+{N_ENGINEERED})\\n')\n    f.write(f'Patch    : {PATCH_SIZE}×{PATCH_SIZE}\\n')\n    f.write(f'Val best : {best_dice:.4f}\\n')\n    f.write('='*55 + '\\n')\n    for tag, m in [('Raw', res['raw']), ('CC-filtered', res['cc'])]:\n        f.write(f'[{tag}]\\n')\n        for k in ['dice', 'prec', 'rec', 'f1']:\n            f.write(f'  {k}: {m[k]:.4f}\\n')\n        f.write(f'  FP/TP: {m[\"fp\"]/(m[\"tp\"]+1e-8):.2f}\\n')\n    f.write(f'Threshold : {res[\"thr\"]:.2f}\\n')\n    f.write(f'CC removed: {res[\"n_removed\"]}\\n')\n    f.write(f'Sep       : {res[\"ink_mean\"]-res[\"noink_mean\"]:+.3f}\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  ssl_backbone.pth | best_model.pth')\nprint('  curves.png | frag1_prediction.png | final_results.txt')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!pip install -q segmentation-models-pytorch transformers albumentations==1.3.1 timm==0.9.16 scikit-image\n!pip install segmentation-models-pytorch==0.2.0\nimport os\nos.environ['OPENCV_IO_ENABLE_JASPER'] = '0'\nos.environ['OPENCV_LOG_LEVEL'] = 'ERROR'\n\nimport gc, cv2, random\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport albumentations as A\nfrom transformers import SegformerForSemanticSegmentation\nfrom skimage import morphology\nimport warnings\nwarnings.filterwarnings('ignore')\n\nCFG = {\n    'seed': 42,\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'z_start': 22,\n    'z_size': 12,\n    'tile_size': 256,\n    'stride_train': 128,\n    'stride_val': 256,\n    'batch_size': 14,\n    'accum_steps': 8,\n    'lr': 1e-4,\n    'epochs': 20,\n    'num_workers': 2,\n    'threshold': 0.0,\n    'data_root': '/kaggle/input/vesuvius-challenge-ink-detection/train',\n    'save_vis_dir': './vis_results',\n    'total_slices': 65,\n    'distill_every': 5, # Self-distill كل 5 epochs\n    'distill_alpha': 0.7 # 70% GT + 30% teacher pred\n}\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False # اسرع\n    torch.backends.cudnn.benchmark = True\nseed_everything(CFG['seed'])\nos.makedirs(CFG['save_vis_dir'], exist_ok=True)\n\nclass TifVolume:\n    def __init__(self, fragment_path):\n        self.slices = sorted(glob(f'{fragment_path}/surface_volume/*.tif'))\n        assert len(self.slices) == CFG['total_slices'], f\"Expected {CFG['total_slices']} slices, got {len(self.slices)}\"\n        h, w = cv2.imread(self.slices[0], cv2.IMREAD_UNCHANGED).shape\n        self.shape = (CFG['total_slices'], h, w)\n\n    def __getitem__(self, idx):\n        z, y, x = idx\n        if isinstance(z, slice):\n            z_stop = min(z.stop, CFG['total_slices'])\n            imgs = []\n            for i in range(z.start, z_stop):\n                img = cv2.imread(self.slices[i], cv2.IMREAD_UNCHANGED)\n                imgs.append(img[y, x])\n            return np.stack(imgs, axis=0)\n        else:\n            img = cv2.imread(self.slices[z], cv2.IMREAD_UNCHANGED)\n            return img[y, x]\n\ndef get_transforms(train=True):\n    if train:\n        return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5, border_mode=0),\n            A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.3),\n        ])\n    else:\n        return None\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, coords_list, train=True, pseudo_labels=None):\n        self.coords_list = coords_list\n        self.train = train\n        self.transform = get_transforms(train)\n        self.pseudo_labels = pseudo_labels # dict: {(fid,y,x): pseudo_mask}\n        self.volumes = {}\n        self.masks = {}\n        for fid in ['2', '3', '1']:\n            self.volumes[fid] = TifVolume(f'{CFG[\"data_root\"]}/{fid}')\n            self.masks[fid] = cv2.imread(f'{CFG[\"data_root\"]}/{fid}/inklabels.png', 0)\n            self.masks[fid] = (self.masks[fid] > 0).astype(np.float32)\n\n    def __len__(self):\n        return len(self.coords_list)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.coords_list[idx]\n        vol = self.volumes[fid]\n        mask = self.masks[fid]\n\n        z_end = min(CFG['z_start'] + CFG['z_size'], CFG['total_slices'])\n        actual_z_size = z_end - CFG['z_start']\n\n        subvol = vol[CFG['z_start']:z_end, y:y+CFG['tile_size'], x:x+CFG['tile_size']]\n        subvol = subvol.astype(np.float32) / 65535.0\n\n        if actual_z_size < CFG['z_size']:\n            pad_z = CFG['z_size'] - actual_z_size\n            subvol = np.pad(subvol, ((0, pad_z), (0, 0), (0, 0)), mode='constant')\n\n        mask_tile = mask[y:y+CFG['tile_size'], x:x+CFG['tile_size']]\n\n        # لو فيه pseudo label استخدمه\n        use_pseudo = False\n        if self.pseudo_labels is not None and (fid, y, x) in self.pseudo_labels:\n            mask_tile = self.pseudo_labels[(fid, y, x)]\n            use_pseudo = True\n\n        if self.transform and not use_pseudo: # متعملش aug على pseudo labels\n            transformed = self.transform(image=subvol[0], mask=mask_tile)\n            mask_tile = transformed['mask']\n            for z in range(CFG['z_size']):\n                subvol[z] = self.transform(image=subvol[z], mask=mask_tile)['image']\n\n        return torch.from_numpy(subvol), torch.from_numpy(mask_tile).unsqueeze(0)\n\ndef create_splits():\n    train_coords, val_coords, test_coords = [], [], []\n\n    for fid in ['2', '3']:\n        mask = cv2.imread(f'{CFG[\"data_root\"]}/{fid}/inklabels.png', 0)\n        h, w = mask.shape\n        all_coords = [(fid, y, x) for y in range(0, h - CFG['tile_size'], CFG['stride_train'])\n                                for x in range(0, w - CFG['tile_size'], CFG['stride_train'])]\n        train_c, val_c = train_test_split(all_coords, test_size=0.2, random_state=CFG['seed'])\n        train_coords.extend(train_c)\n        val_coords.extend(val_c)\n\n    mask1 = cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0)\n    h, w = mask1.shape\n    test_coords = [( '1', y, x) for y in range(0, h - CFG['tile_size'], CFG['stride_val'])\n                               for x in range(0, w - CFG['tile_size'], CFG['stride_val'])]\n\n    return train_coords, val_coords, test_coords\n\ntrain_coords, val_coords, test_coords = create_splits()\nprint(f'Train: {len(train_coords)}, Val: {len(val_coords)}, Test: {len(test_coords)}')\n\nclass InkDetector2p5D(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = SegformerForSemanticSegmentation.from_pretrained(\n            \"nvidia/mit-b2\", num_labels=1, num_channels=CFG['z_size'], ignore_mismatched_sizes=True\n        )\n        self.upsample = nn.Upsample(size=(CFG['tile_size'], CFG['tile_size']), mode='bilinear', align_corners=False)\n\n    def forward(self, x):\n        x = self.backbone(x).logits\n        x = self.upsample(x)\n        return x\n\ndef dice_loss(pred, target, smooth=1e-6):\n    pred = torch.sigmoid(pred)\n    pred = pred.contiguous().view(-1)\n    target = target.contiguous().view(-1)\n    intersection = (pred * target).sum()\n    dice = (2. * intersection + smooth) / (pred.sum() + target.sum() + smooth)\n    return 1 - dice\n\ndef dice_bce_loss(pred, target):\n    bce = F.binary_cross_entropy_with_logits(pred, target)\n    dice = dice_loss(pred, target)\n    return bce + dice\n\nmodel = InkDetector2p5D().to(CFG['device'])\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\nscaler = GradScaler()\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG['epochs'])\n\n# 5. Generate pseudo labels from validation\n@torch.no_grad()\ndef generate_pseudo_labels(model, coords):\n    model.eval()\n    pseudo_dict = {}\n    temp_ds = VesuviusDataset(coords, train=False)\n    temp_loader = DataLoader(temp_ds, batch_size=1, shuffle=False, num_workers=CFG['num_workers'])\n\n    print(\"Generating pseudo labels on validation set...\")\n    for i, ((images, _), (fid, y, x)) in enumerate(zip(temp_loader, coords)):\n        images = images.to(CFG['device'])\n        with autocast():\n            pred = torch.sigmoid(model(images)).squeeze().cpu().numpy()\n        pseudo_dict[(fid, y, x)] = pred\n\n    return pseudo_dict\n\n# 6. Train + Validate with Self-Distillation\ndef train_one_epoch(model, loader, optimizer, scaler, epoch):\n    model.train()\n    running_loss = 0\n\n    for i, (images, masks) in enumerate(loader):\n        images, masks = images.to(CFG['device']), masks.to(CFG['device'])\n        with autocast():\n            preds = model(images)\n            loss = dice_bce_loss(preds, masks) / CFG['accum_steps']\n        scaler.scale(loss).backward()\n        if (i + 1) % CFG['accum_steps'] == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n        running_loss += loss.item() * CFG['accum_steps']\n\n        if i % 100 == 0:\n            print(f'Epoch {epoch+1} | Batch {i}/{len(loader)} | Loss: {running_loss/(i+1):.4f}')\n        if i % 50 == 0:\n            torch.cuda.empty_cache()\n\n    return running_loss / len(loader)\n\n@torch.no_grad()\ndef validate(model, loader):\n    model.eval()\n    dice_scores = []\n    for images, masks in loader:\n        images, masks = images.to(CFG['device']), masks.to(CFG['device'])\n        with autocast():\n            preds = model(images)\n        preds = (torch.sigmoid(preds) > 0.5).float()\n        dice = 1 - dice_loss(preds, masks).item()\n        dice_scores.append(dice)\n    return np.mean(dice_scores)\n\n# Main training loop with self-distillation\nbest_dice = 0\npseudo_labels = None\n\nfor epoch in range(CFG['epochs']):\n    print(f'\\n===== Epoch {epoch+1}/{CFG[\"epochs\"]} =====')\n\n    # كل 5 epochs: اعمل pseudo labels من الـ val set\n    if epoch > 0 and epoch % CFG['distill_every'] == 0:\n        pseudo_labels = generate_pseudo_labels(model, val_coords)\n        print(f\"Generated {len(pseudo_labels)} pseudo labels for self-distillation\")\n\n    # اعمل dataset جديد مع pseudo labels لو موجودة\n    train_ds = VesuviusDataset(train_coords, train=True, pseudo_labels=pseudo_labels)\n    train_loader = DataLoader(train_ds, batch_size=CFG['batch_size'], shuffle=True,\n                              num_workers=CFG['num_workers'], pin_memory=True)\n\n    val_ds = VesuviusDataset(val_coords, train=False)\n    val_loader = DataLoader(val_ds, batch_size=CFG['batch_size'], shuffle=False,\n                            num_workers=CFG['num_workers'], pin_memory=True)\n\n    train_loss = train_one_epoch(model, train_loader, optimizer, scaler, epoch)\n    val_dice = validate(model, val_loader)\n    scheduler.step()\n\n    print(f'Epoch {epoch+1} Summary: Train Loss: {train_loss:.4f} | Val Dice: {val_dice:.4f}')\n\n    if val_dice > best_dice:\n        best_dice = val_dice\n        torch.save(model.state_dict(), 'best_model.pth')\n        print(f'*** Saved best model with dice {best_dice:.4f} ***')\n\n# 7. Test on fragment 1\ndef test_and_visualize(model, test_coords):\n    model.load_state_dict(torch.load('best_model.pth'))\n    model.eval()\n    model.half()\n\n    test_ds = VesuviusDataset(test_coords, train=False)\n    test_loader = DataLoader(test_ds, batch_size=1, shuffle=False, num_workers=CFG['num_workers'])\n\n    mask1 = cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0)\n    h, w = mask1.shape\n    pred_full = np.zeros((h, w), dtype=np.float32)\n    counts = np.zeros((h, w), dtype=np.float32)\n    vol = TifVolume(f'{CFG[\"data_root\"]}/1')\n\n    print(\"Testing on fragment 1...\")\n    with torch.no_grad(), autocast():\n        for i, ((images, _), (fid, y, x)) in enumerate(zip(test_loader, test_coords)):\n            images = images.to(CFG['device']).half()\n            pred = torch.sigmoid(model(images)).squeeze().cpu().float().numpy()\n            pred_full[y:y+CFG['tile_size'], x:x+CFG['tile_size']] += pred\n            counts[y:y+CFG['tile_size'], x:x+CFG['tile_size']] += 1\n\n            if i < 10:\n                input_img = images[0, CFG['z_size']//2].cpu().float().numpy()\n                gt_mask = cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0)[y:y+CFG['tile_size'], x:x+CFG['tile_size']]\n                gt_mask = (gt_mask > 0).astype(np.float32)\n\n                fig, ax = plt.subplots(1, 3, figsize=(15, 5))\n                ax[0].imshow(input_img, cmap='gray'); ax[0].set_title('Input Slice'); ax[0].axis('off')\n                ax[1].imshow(gt_mask, cmap='gray'); ax[1].set_title('Ground Truth'); ax[1].axis('off')\n                ax[2].imshow(pred, cmap='gray'); ax[2].set_title('Prediction'); ax[2].axis('off')\n                plt.savefig(f'{CFG[\"save_vis_dir\"]}/compare_{i}.png', bbox_inches='tight', dpi=150)\n                plt.close()\n            if i % 20 == 0:\n                torch.cuda.empty_cache()\n                print(f'Test progress: {i}/{len(test_coords)}')\n\n    pred_full = pred_full / (counts + 1e-6)\n    pred_full = (pred_full > 0.5).astype(np.uint8)\n    pred_full = morphology.remove_small_objects(pred_full.astype(bool), min_size=64).astype(np.uint8) * 255\n\n    gt_full = (cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0) > 0).astype(np.uint8) * 255\n    input_full = vol[CFG['z_start']+CFG['z_size']//2, :, :].astype(np.float32) / 65535.0\n\n    plt.figure(figsize=(20, 6))\n    plt.subplot(1,3,1); plt.imshow(input_full, cmap='gray'); plt.title('Input Fragment 1'); plt.axis('off')\n    plt.subplot(1,3,2); plt.imshow(gt_full, cmap='gray'); plt.title('Ground Truth'); plt.axis('off')\n    plt.subplot(1,3,3); plt.imshow(pred_full, cmap='gray'); plt.title('Prediction'); plt.axis('off')\n    plt.savefig(f'{CFG[\"save_vis_dir\"]}/fragment1_full_comparison.png', bbox_inches='tight', dpi=200)\n    plt.close()\n\n    dice = 1 - dice_loss(torch.from_numpy(pred_full/255.0), torch.from_numpy(gt_full/255.0)).item()\n    print(f'\\n=== FINAL Test Dice on Fragment 1: {dice:.4f} ===')\n    return pred_full\n\npred_mask = test_and_visualize(model, test_coords)\nprint(f'Visualizations saved to {CFG[\"save_vis_dir\"]}/')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n#!pip install -q segmentation-models-pytorch transformers albumentations==1.3.1 timm==0.9.16 scikit-image\nimport os\nos.environ['OPENCV_IO_ENABLE_JASPER'] = '0'\nos.environ['OPENCV_LOG_LEVEL'] = 'ERROR'\n\nimport gc, cv2, random\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport albumentations as A\nfrom transformers import SegformerForSemanticSegmentation\nfrom skimage import morphology\nfrom tqdm.auto import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\nCFG = {\n    'seed': 42,\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'z_start': 22,\n    'z_size': 12,\n    'tile_size': 256,\n    'stride_train': 128,\n    'stride_val': 256,\n    'batch_size': 14,\n    'accum_steps': 8,\n    'lr': 1e-4,\n    'epochs': 20,\n    'num_workers': 2,\n    'threshold': 0.0, # Changed: now logits, so threshold = 0\n    'data_root': '/kaggle/input/vesuvius-challenge-ink-detection/train',\n    'save_vis_dir': './vis_results',\n    'total_slices': 65\n}\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\nseed_everything(CFG['seed'])\nos.makedirs(CFG['save_vis_dir'], exist_ok=True)\n\n# 1. Load.tif volumes - 65 slices\nclass TifVolume:\n    def __init__(self, fragment_path):\n        self.slices = sorted(glob(f'{fragment_path}/surface_volume/*.tif'))\n        assert len(self.slices) == CFG['total_slices'], f\"Expected {CFG['total_slices']} slices, got {len(self.slices)}\"\n        h, w = cv2.imread(self.slices[0], cv2.IMREAD_UNCHANGED).shape\n        self.shape = (CFG['total_slices'], h, w)\n\n    def __getitem__(self, idx):\n        z, y, x = idx\n        if isinstance(z, slice):\n            z_stop = min(z.stop, CFG['total_slices'])\n            imgs = []\n            for i in range(z.start, z_stop):\n                img = cv2.imread(self.slices[i], cv2.IMREAD_UNCHANGED)\n                imgs.append(img[y, x])\n            return np.stack(imgs, axis=0)\n        else:\n            img = cv2.imread(self.slices[z], cv2.IMREAD_UNCHANGED)\n            return img[y, x]\n\n# 2. Dataset\ndef get_transforms(train=True):\n    if train:\n        return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5, border_mode=0),\n            A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.3),\n        ])\n    else:\n        return None\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, coords_list, train=True):\n        self.coords_list = coords_list\n        self.train = train\n        self.transform = get_transforms(train)\n        self.volumes = {}\n        self.masks = {}\n        for fid in ['2', '3', '1']:\n            self.volumes[fid] = TifVolume(f'{CFG[\"data_root\"]}/{fid}')\n            self.masks[fid] = cv2.imread(f'{CFG[\"data_root\"]}/{fid}/inklabels.png', 0)\n            self.masks[fid] = (self.masks[fid] > 0).astype(np.float32)\n\n    def __len__(self):\n        return len(self.coords_list)\n\n    def __getitem__(self, idx):\n        fid, y, x = self.coords_list[idx]\n        vol = self.volumes[fid]\n        mask = self.masks[fid]\n\n        z_end = min(CFG['z_start'] + CFG['z_size'], CFG['total_slices'])\n        actual_z_size = z_end - CFG['z_start']\n\n        subvol = vol[CFG['z_start']:z_end, y:y+CFG['tile_size'], x:x+CFG['tile_size']]\n        subvol = subvol.astype(np.float32) / 65535.0\n\n        if actual_z_size < CFG['z_size']:\n            pad_z = CFG['z_size'] - actual_z_size\n            subvol = np.pad(subvol, ((0, pad_z), (0, 0), (0, 0)), mode='constant')\n\n        mask_tile = mask[y:y+CFG['tile_size'], x:x+CFG['tile_size']]\n\n        if self.transform:\n            transformed = self.transform(image=subvol[0], mask=mask_tile)\n            mask_tile = transformed['mask']\n            for z in range(CFG['z_size']):\n                subvol[z] = self.transform(image=subvol[z], mask=mask_tile)['image']\n\n        return torch.from_numpy(subvol), torch.from_numpy(mask_tile).unsqueeze(0)\n\n# 3. Create splits\ndef create_splits():\n    train_coords, val_coords, test_coords = [], [], []\n\n    for fid in ['2', '3']:\n        mask = cv2.imread(f'{CFG[\"data_root\"]}/{fid}/inklabels.png', 0)\n        h, w = mask.shape\n        all_coords = [(fid, y, x) for y in range(0, h - CFG['tile_size'], CFG['stride_train'])\n                                for x in range(0, w - CFG['tile_size'], CFG['stride_train'])]\n        train_c, val_c = train_test_split(all_coords, test_size=0.2, random_state=CFG['seed'])\n        train_coords.extend(train_c)\n        val_coords.extend(val_c)\n\n    mask1 = cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0)\n    h, w = mask1.shape\n    test_coords = [( '1', y, x) for y in range(0, h - CFG['tile_size'], CFG['stride_val'])\n                               for x in range(0, w - CFG['tile_size'], CFG['stride_val'])]\n\n    return train_coords, val_coords, test_coords\n\ntrain_coords, val_coords, test_coords = create_splits()\ntrain_ds = VesuviusDataset(train_coords, train=True)\nval_ds = VesuviusDataset(val_coords, train=False)\ntest_ds = VesuviusDataset(test_coords, train=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG['batch_size'], shuffle=True,\n                          num_workers=CFG['num_workers'], pin_memory=True)\nval_loader = DataLoader(val_ds, batch_size=CFG['batch_size'], shuffle=False,\n                        num_workers=CFG['num_workers'], pin_memory=True)\ntest_loader = DataLoader(test_ds, batch_size=1, shuffle=False, num_workers=CFG['num_workers'])\n\nprint(f'Train: {len(train_ds)}, Val: {len(val_ds)}, Test: {len(test_ds)}')\n\n# 4. Model + Loss - Fixed: removed sigmoid, use BCEWithLogitsLoss\nclass InkDetector2p5D(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = SegformerForSemanticSegmentation.from_pretrained(\n            \"nvidia/mit-b2\", num_labels=1, num_channels=CFG['z_size'], ignore_mismatched_sizes=True\n        )\n        self.upsample = nn.Upsample(size=(CFG['tile_size'], CFG['tile_size']), mode='bilinear', align_corners=False)\n\n    def forward(self, x):\n        x = self.backbone(x).logits\n        x = self.upsample(x)\n        return x # Removed.sigmoid() - output logits now\n\ndef dice_loss(pred, target, smooth=1e-6):\n    pred = torch.sigmoid(pred) # Apply sigmoid here for dice\n    pred = pred.contiguous().view(-1)\n    target = target.contiguous().view(-1)\n    intersection = (pred * target).sum()\n    dice = (2. * intersection + smooth) / (pred.sum() + target.sum() + smooth)\n    return 1 - dice\n\ndef dice_bce_loss(pred, target):\n    # Use BCEWithLogitsLoss - safe for autocast\n    bce = F.binary_cross_entropy_with_logits(pred, target)\n    dice = dice_loss(pred, target)\n    return bce + dice\n\nmodel = InkDetector2p5D().to(CFG['device'])\n\ntry:\n    model = torch.compile(model)\n    print(\"Model compiled with torch.compile\")\nexcept:\n    print(\"torch.compile not available, skipping\")\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\nscaler = GradScaler()\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG['epochs'])\n\n# 5. Train + Validate\ndef train_one_epoch(model, loader, optimizer, scaler):\n    model.train()\n    running_loss = 0\n    pbar = tqdm(loader, desc='Train')\n    optimizer.zero_grad()\n\n    for i, (images, masks) in enumerate(pbar):\n        images, masks = images.to(CFG['device']), masks.to(CFG['device'])\n        with autocast():\n            preds = model(images) # logits\n            loss = dice_bce_loss(preds, masks) / CFG['accum_steps']\n        scaler.scale(loss).backward()\n        if (i + 1) % CFG['accum_steps'] == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n        running_loss += loss.item() * CFG['accum_steps']\n        pbar.set_postfix({'loss': f'{running_loss/(i+1):.4f}'})\n        if i % 50 == 0: torch.cuda.empty_cache()\n    return running_loss / len(loader)\n\n@torch.no_grad()\ndef validate(model, loader):\n    model.eval()\n    dice_scores = []\n    for images, masks in tqdm(loader, desc='Val'):\n        images, masks = images.to(CFG['device']), masks.to(CFG['device'])\n        with autocast():\n            preds = model(images) # logits\n        preds = (torch.sigmoid(preds) > 0.5).float() # threshold at 0.5 after sigmoid\n        dice = 1 - dice_loss(preds, masks).item()\n        dice_scores.append(dice)\n    return np.mean(dice_scores)\n\nbest_dice = 0\nfor epoch in range(CFG['epochs']):\n    print(f'\\nEpoch {epoch+1}/{CFG[\"epochs\"]}')\n    train_loss = train_one_epoch(model, train_loader, optimizer, scaler)\n    val_dice = validate(model, val_loader)\n    scheduler.step()\n    print(f'Train Loss: {train_loss:.4f} | Val Dice: {val_dice:.4f}')\n\n    if val_dice > best_dice:\n        best_dice = val_dice\n        torch.save(model.state_dict(), 'best_model.pth')\n        print(f'Saved best model with dice {best_dice:.4f}')\n\n# 6. Test on fragment 1 + Save Visualizations\ndef test_and_visualize(model, loader, test_coords):\n    model.load_state_dict(torch.load('best_model.pth'))\n    model.eval()\n    model.half()\n\n    mask1 = cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0)\n    h, w = mask1.shape\n    pred_full = np.zeros((h, w), dtype=np.float32)\n    counts = np.zeros((h, w), dtype=np.float32)\n    vol = TifVolume(f'{CFG[\"data_root\"]}/1')\n\n    with torch.no_grad(), autocast():\n        for i, ((images, _), (fid, y, x)) in enumerate(tqdm(zip(loader, test_coords), total=len(loader))):\n            images = images.to(CFG['device']).half()\n            pred = torch.sigmoid(model(images)).squeeze().cpu().float().numpy() # sigmoid for output\n            pred_full[y:y+CFG['tile_size'], x:x+CFG['tile_size']] += pred\n            counts[y:y+CFG['tile_size'], x:x+CFG['tile_size']] += 1\n\n            if i < 10:\n                input_img = images[0, CFG['z_size']//2].cpu().float().numpy()\n                gt_mask = cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0)[y:y+CFG['tile_size'], x:x+CFG['tile_size']]\n                gt_mask = (gt_mask > 0).astype(np.float32)\n\n                fig, ax = plt.subplots(1, 3, figsize=(15, 5))\n                ax[0].imshow(input_img, cmap='gray'); ax[0].set_title('Input Slice'); ax[0].axis('off')\n                ax[1].imshow(gt_mask, cmap='gray'); ax[1].set_title('Ground Truth'); ax[1].axis('off')\n                ax[2].imshow(pred, cmap='gray'); ax[2].set_title('Prediction'); ax[2].axis('off')\n                plt.savefig(f'{CFG[\"save_vis_dir\"]}/compare_{i}.png', bbox_inches='tight', dpi=150)\n                plt.close()\n            if i % 20 == 0: torch.cuda.empty_cache()\n\n    pred_full = pred_full / (counts + 1e-6)\n    pred_full = (pred_full > 0.5).astype(np.uint8) # threshold 0.5 after sigmoid\n    pred_full = morphology.remove_small_objects(pred_full.astype(bool), min_size=64).astype(np.uint8) * 255\n\n    gt_full = (cv2.imread(f'{CFG[\"data_root\"]}/1/inklabels.png', 0) > 0).astype(np.uint8) * 255\n    input_full = vol[CFG['z_start']+CFG['z_size']//2, :, :].astype(np.float32) / 65535.0\n\n    plt.figure(figsize=(20, 6))\n    plt.subplot(1,3,1); plt.imshow(input_full, cmap='gray'); plt.title('Input Fragment 1'); plt.axis('off')\n    plt.subplot(1,3,2); plt.imshow(gt_full, cmap='gray'); plt.title('Ground Truth'); plt.axis('off')\n    plt.subplot(1,3,3); plt.imshow(pred_full, cmap='gray'); plt.title('Prediction'); plt.axis('off')\n    plt.savefig(f'{CFG[\"save_vis_dir\"]}/fragment1_full_comparison.png', bbox_inches='tight', dpi=200)\n    plt.close()\n\n    dice = 1 - dice_loss(torch.from_numpy(pred_full/255.0), torch.from_numpy(gt_full/255.0)).item()\n    print(f'Test Dice on Fragment 1: {dice:.4f}')\n    return pred_full\n\npred_mask = test_and_visualize(model, test_loader, test_coords)\nprint(f'Visualizations saved to {CFG[\"save_vis_dir\"]}/')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n# ============================================================\n#  VESUVIUS INK DETECTION — v13  (+ RICH FRAGMENT 1 SAVE)\n#\n#  ONLY CHANGE vs original v13:\n#  → After Fold C predicts Fragment 1, a comprehensive figure\n#    is saved showing:\n#      Row 1: CT surface slice (mid-z)  |  Ground Truth label\n#              |  Probability map\n#      Row 2: Binary prediction  |  Overlay (GT on CT)\n#              |  Error map (TP/FP/FN)\n#      Row 3: Probability histogram  |  Dice-vs-threshold curve\n#              |  Per-region zoom (random 512×512 crop)\n#  All other code is identical to v13.\n# ============================================================\n\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler, ConcatDataset\nfrom torch.optim.swa_utils import AveragedModel, SWALR, update_bn\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ── paths ─────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ──────────────────────────────────────────────\nDEVICE         = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE     = 224\nSTRIDE_TR      = 112\nSTRIDE_INF     = 56\nBATCH_SIZE     = 16\nGRAD_ACCUM     = 2\nEPOCHS         = 10\nSWA_START      = 22\nLR             = 1e-4\nWEIGHT_DECAY   = 1e-5\nPATIENCE       = 10\nNUM_WORKERS    = 4\nMAX_PATCHES_TR = 12_000\n\nZ_SLICES = list(range(25, 37))\nN_CH     = len(Z_SLICES)\n\nINK_MIN_POS = 0.02\nNEG_RATIO   = 0.3\nPOS_WEIGHT  = 3.0\nPIN = (DEVICE == 'cuda')\n\nprint(f\"Device : {DEVICE}  |  Channels: {N_CH}  |  Patch: {PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\nprint(f\"Z-slices: {Z_SLICES}\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. SLICE CACHE\n# ════════════════════════════════════════════════════════════\n_cache_store = {}\n\ndef load_slice_cache(frag_path, z_list):\n    key = (frag_path, tuple(z_list))\n    if key in _cache_store:\n        print(f\"  Reusing cached slices for {os.path.basename(frag_path)}\")\n        return _cache_store[key]\n    vol_dir  = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Slice cache: {len(cache)} slices, {mb:.0f} MB\")\n    _cache_store[key] = cache\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2. BOUNDARY MASK\n# ════════════════════════════════════════════════════════════\ndef make_boundary_mask(mask_np, radius=5):\n    kernel  = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2*radius+1, 2*radius+1))\n    dilated = cv2.dilate(mask_np, kernel)\n    eroded  = cv2.erode (mask_np, kernel)\n    return (dilated != eroded).astype(np.float32)\n\n\n# ════════════════════════════════════════════════════════════\n#  3. DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0, max_patches=0):\n        self.cache  = load_slice_cache(frag_path, z_list)\n        self.z_list = z_list\n        self.tf     = transform\n\n        msk_path  = os.path.join(frag_path, 'inklabels.png')\n        msk       = cv2.imread(msk_path, 0)\n        self.mask = (msk > 0).astype(np.uint8)\n        self.boundary = make_boundary_mask(self.mask).astype(np.uint8)\n\n        ir_path = os.path.join(frag_path, 'mask.png')\n        ir      = cv2.imread(ir_path, 0) if os.path.exists(ir_path) else None\n        ir_mask = (ir > 0).astype(np.uint8) if ir is not None else None\n\n        H, W = self.mask.shape\n        pos_yx, neg_yx = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.mean()\n                if ir_mask is not None:\n                    on_pap = ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    mid_z  = z_list[len(z_list)//2]\n                    on_pap = float(\n                        self.cache[mid_z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                    ) > 0.1\n                if not on_pap: continue\n                if ink >= INK_MIN_POS:\n                    pos_yx.append((y, x))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_yx.append((y, x))\n\n        n_neg = int(len(pos_yx) * neg_ratio)\n        if n_neg > 0 and neg_yx:\n            np.random.shuffle(neg_yx); neg_yx = neg_yx[:n_neg]\n        else:\n            neg_yx = []\n\n        n_total = len(pos_yx) + len(neg_yx)\n        if max_patches > 0 and n_total > max_patches:\n            frac = max_patches / n_total\n            np.random.shuffle(pos_yx); np.random.shuffle(neg_yx)\n            pos_yx = pos_yx[:max(1, int(len(pos_yx)*frac))]\n            neg_yx = neg_yx[:max(0, int(len(neg_yx)*frac))]\n\n        all_yx  = pos_yx + neg_yx\n        all_lbl = [1]*len(pos_yx) + [0]*len(neg_yx)\n        perm    = np.random.permutation(len(all_yx))\n\n        self.coords  = np.array([all_yx[i]  for i in perm], dtype=np.int32)\n        self.labels  = np.array([all_lbl[i] for i in perm], dtype=np.int32)\n        self.weights = np.where(self.labels==1, POS_WEIGHT, 1.0).astype(np.float32)\n\n        cap_str = f\" (capped from {n_total})\" if max_patches > 0 and n_total > max_patches else \"\"\n        print(f\"  {os.path.basename(frag_path)}: \"\n              f\"{len(pos_yx)} pos + {len(neg_yx)} neg \"\n              f\"= {len(self.coords)} total{cap_str}\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x = int(self.coords[idx,0]), int(self.coords[idx,1])\n        ps   = PATCH_SIZE\n        slices = [self.cache[z][y:y+ps, x:x+ps].astype(np.float32)\n                  for z in self.z_list]\n        img = np.stack(slices, axis=-1)\n        msk = self.mask   [y:y+ps, x:x+ps].copy()\n        bnd = self.boundary[y:y+ps, x:x+ps].copy().astype(np.float32)\n        if self.tf:\n            out = self.tf(image=img, mask=msk)\n            img, msk = out['image'], out['mask']\n        img_t = torch.from_numpy(img).permute(2,0,1).float()\n        mu  = img_t.mean(dim=(1,2), keepdim=True)\n        std = img_t.std (dim=(1,2), keepdim=True) + 1e-6\n        img_t = (img_t - mu) / std\n        msk_t = torch.from_numpy(msk).unsqueeze(0).float()\n        bnd_t = torch.from_numpy(bnd).unsqueeze(0).float()\n        return img_t, msk_t, bnd_t\n\n\n# ════════════════════════════════════════════════════════════\n#  4. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.15,\n                       rotate_limit=30,\n                       border_mode=cv2.BORDER_REFLECT, p=0.6),\n    A.RandomBrightnessContrast(0.2, 0.2, p=0.5),\n    A.GaussNoise(var_limit=(0.001, 0.004), p=0.3),\n    A.GaussianBlur(blur_limit=(3,5), p=0.2),\n    A.CoarseDropout(max_holes=4, max_height=24, max_width=24,\n                    fill_value=0, p=0.3),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  5. LOSS\n# ════════════════════════════════════════════════════════════\ndef tversky_loss(pred, target, alpha=0.3, beta=0.7, smooth=1.):\n    p  = torch.sigmoid(pred)\n    tp = (p * target).sum(dim=(2,3))\n    fp = (p * (1-target)).sum(dim=(2,3))\n    fn = ((1-p) * target).sum(dim=(2,3))\n    return 1. - ((tp+smooth) / (tp + alpha*fp + beta*fn + smooth)).mean()\n\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0, eps=0.05):\n    t_s = target*(1-eps) + 0.5*eps\n    bce = F.binary_cross_entropy_with_logits(pred, t_s, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = t_s*alpha + (1-t_s)*(1-alpha)\n    return (a_t * ((1-p_t)**gamma) * bce).mean()\n\ndef criterion(pred, target):\n    return 0.5*tversky_loss(pred, target) + 0.5*focal_loss(pred, target)\n\n\n# ════════════════════════════════════════════════════════════\n#  6. MODEL\n# ════════════════════════════════════════════════════════════\ndef build_model(n_ch=N_CH):\n    return smp.UnetPlusPlus(\n        encoder_name           = 'efficientnet-b4',\n        encoder_weights        = 'imagenet',\n        in_channels            = n_ch,\n        classes                = 1,\n        decoder_attention_type = 'scse',\n    ).to(DEVICE)\n\n\n# ════════════════════════════════════════════════════════════\n#  7. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p*masks).sum(dim=(1,2,3))\n    union = p.sum(dim=(1,2,3)) + masks.sum(dim=(1,2,3))\n    return ((2.*inter+1e-5)/(union+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.25, 0.85, 0.01):\n        p = (probs > t).astype(np.float32)\n        d = (2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d > best_d:\n            best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  8. GAUSSIAN PATCH WEIGHT\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c = sz//2; sig = sz//4\n    ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs-c)**2+(ys-c)**2)/(2*sig**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  9. DATALOADER FACTORY\n# ════════════════════════════════════════════════════════════\ndef make_loader(datasets, shuffle_sampler=True, max_p=0):\n    ds_list = []\n    for frag_path, is_train in datasets:\n        ds = VesuviusDataset(\n            frag_path, Z_SLICES,\n            stride      = STRIDE_TR,\n            transform   = train_tf if is_train else None,\n            neg_ratio   = NEG_RATIO if is_train else 0.0,\n            max_patches = max_p if is_train else 0,\n        )\n        ds_list.append(ds)\n\n    if len(ds_list) == 1:\n        combined = ds_list[0]; weights = combined.weights\n    else:\n        combined = ConcatDataset(ds_list)\n        weights  = np.concatenate([d.weights for d in ds_list])\n\n    if shuffle_sampler:\n        sampler = WeightedRandomSampler(\n            torch.from_numpy(weights), len(combined), replacement=True)\n        dl = DataLoader(combined, batch_size=BATCH_SIZE, sampler=sampler,\n                        num_workers=NUM_WORKERS, pin_memory=PIN,\n                        persistent_workers=(NUM_WORKERS>0),\n                        prefetch_factor=2 if NUM_WORKERS>0 else None)\n    else:\n        dl = DataLoader(combined, batch_size=BATCH_SIZE, shuffle=False,\n                        num_workers=NUM_WORKERS, pin_memory=PIN,\n                        persistent_workers=(NUM_WORKERS>0),\n                        prefetch_factor=2 if NUM_WORKERS>0 else None)\n    return dl\n\n\n# ════════════════════════════════════════════════════════════\n#  10. TRAINING LOOP\n# ════════════════════════════════════════════════════════════\ndef run_training(model, train_dl, val_dl, n_epochs, lr,\n                 swa_start, ckpt_name, best_dice_init=0.):\n\n    optimizer  = optim.AdamW(model.parameters(), lr=lr,\n                              weight_decay=WEIGHT_DECAY)\n    scheduler  = optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=10, T_mult=1, eta_min=1e-6)\n    scaler     = torch.cuda.amp.GradScaler(enabled=(DEVICE=='cuda'))\n    swa_model  = AveragedModel(model)\n    swa_sched  = SWALR(optimizer, swa_lr=lr*0.1, anneal_epochs=5)\n    swa_active = False\n\n    best_dice = best_dice_init\n    pat_cnt   = 0\n    history   = dict(tl=[], vl=[], td=[], vd=[])\n\n    for epoch in range(n_epochs):\n        if epoch >= swa_start and not swa_active:\n            swa_active = True\n            print(f'  → SWA activated at epoch {epoch+1}')\n\n        model.train(); tl = td = 0.\n        optimizer.zero_grad()\n        for step, (imgs, msks, _bnds) in enumerate(\n                tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                out  = model(imgs)\n                loss = criterion(out, msks) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            if (step+1) % GRAD_ACCUM == 0 or (step+1) == len(train_dl):\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                scaler.step(optimizer); scaler.update()\n                optimizer.zero_grad()\n            tl += loss.item() * GRAD_ACCUM\n            td += batch_dice(out.detach(), msks)\n            del imgs, msks, out, loss\n            if step % 200 == 0 and DEVICE=='cuda':\n                torch.cuda.empty_cache()\n        tl /= len(train_dl); td /= len(train_dl)\n\n        model.eval(); vl = vd = 0.\n        acc_p, acc_m = [], []\n        with torch.no_grad():\n            for imgs, msks, _bnds in tqdm(\n                    val_dl, desc=f'Ep{epoch+1:02d} val  ', leave=False):\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                msks = msks.to(DEVICE, non_blocking=True)\n                with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                    out  = model(imgs)\n                    loss = criterion(out, msks)\n                vl += loss.item(); vd += batch_dice(out, msks)\n                acc_p.append(torch.sigmoid(out).cpu().numpy())\n                acc_m.append(msks.cpu().numpy())\n                del imgs, msks, out, loss\n        vl /= len(val_dl); vd /= len(val_dl)\n\n        probs_all  = np.concatenate(acc_p)\n        masks_all  = np.concatenate(acc_m)\n        bt, bd     = sweep_threshold(probs_all, masks_all)\n        ink_mean   = float(probs_all[masks_all>0.5].mean()) if (masks_all>0.5).any() else 0.\n        noink_mean = float(probs_all[masks_all<0.5].mean()) if (masks_all<0.5).any() else 0.\n        del acc_p, acc_m, probs_all, masks_all\n        gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n        if swa_active:\n            swa_model.update_parameters(model); swa_sched.step()\n        else:\n            scheduler.step()\n        lr_now = optimizer.param_groups[0]['lr']\n\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n\n        print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n              f'train loss={tl:.4f} dice={td:.4f} | '\n              f'val loss={vl:.4f} dice={vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | '\n              f'sep={ink_mean-noink_mean:+.3f} '\n              f'[ink={ink_mean:.3f} bg={noink_mean:.3f}]'\n              + (' [SWA]' if swa_active else ''))\n\n        save_metric = max(vd, bd)\n        if save_metric > best_dice:\n            best_dice = save_metric; pat_cnt = 0\n            torch.save({'epoch': epoch, 'state': model.state_dict(),\n                        'thr': bt, 'dice': best_dice, 'n_ch': N_CH},\n                       OUTPUT + ckpt_name)\n            print(f'  ✓ saved  (metric={best_dice:.4f})')\n        else:\n            pat_cnt += 1\n            if pat_cnt >= PATIENCE:\n                print(f'  ⚑ early stop at epoch {epoch+1}')\n                break\n\n        if epoch == 4 and max(history['vd']) < 0.30:\n            print('\\n⚠  val dice <0.30 after 5 epochs.')\n\n    if swa_active:\n        print('  Updating SWA BN stats ...')\n        swa_model.train()\n        with torch.no_grad():\n            for imgs, msks, _b in train_dl:\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                swa_model(imgs); del imgs, msks, _b\n        model.load_state_dict(swa_model.module.state_dict())\n        del swa_model; gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    return model, best_dice, history\n\n\n# ════════════════════════════════════════════════════════════\n#  11. INFERENCE\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    model.eval()\n    cache = _cache_store.get((frag_path, tuple(z_list)))\n    if cache is None:\n        cache = load_slice_cache(frag_path, z_list)\n    msk   = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n    H, W  = msk.shape\n    pred_map = np.zeros((H,W), np.float32)\n    wgt_map  = np.zeros((H,W), np.float32)\n    coords   = [(y,x)\n                for y in range(0, H-PATCH_SIZE+1, STRIDE_INF)\n                for x in range(0, W-PATCH_SIZE+1, STRIDE_INF)]\n    with torch.no_grad():\n        for (y,x) in tqdm(coords, desc='Inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            img = np.stack(slices, -1)\n            t   = torch.from_numpy(img).permute(2,0,1).float()\n            mu  = t.mean(dim=(1,2), keepdim=True)\n            std = t.std (dim=(1,2), keepdim=True) + 1e-6\n            t   = ((t - mu) / std).unsqueeze(0).to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                p = torch.sigmoid(model(t)).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, p\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    return pred_map/(wgt_map+1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  12. STANDARD EVALUATE + PLOTS  (used for folds A and B)\n# ════════════════════════════════════════════════════════════\ndef evaluate(prob_map, gt_mask, label=''):\n    H, W      = gt_mask.shape\n    prob_crop = prob_map[:H,:W]\n    ink_mean   = float(prob_crop[gt_mask==1].mean())\n    noink_mean = float(prob_crop[gt_mask==0].mean())\n    bt, bd = sweep_threshold(prob_crop[np.newaxis,np.newaxis],\n                              gt_mask[np.newaxis,np.newaxis])\n    pred  = (prob_crop > bt).astype(np.uint8)\n    inter = (pred*gt_mask).sum()\n    dice  = (2*inter+1)/(pred.sum()+gt_mask.sum()+1)\n    pf = pred.flatten().astype(int)\n    mf = gt_mask.flatten().astype(int)\n    tn,fp,fn,tp_v = confusion_matrix(mf,pf,labels=[0,1]).ravel()\n    prec = tp_v/(tp_v+fp+1e-8); rec = tp_v/(tp_v+fn+1e-8)\n    f1   = 2*prec*rec/(prec+rec+1e-8)\n    print(f'\\n{\"=\"*55}')\n    print(f'RESULTS {label}')\n    print(f'{\"=\"*55}')\n    print(f'Dice        : {dice:.4f}')\n    print(f'Threshold   : {bt:.2f}')\n    print(f'Precision   : {prec:.4f}  Recall: {rec:.4f}  F1: {f1:.4f}')\n    print(f'TP={tp_v} TN={tn} FP={fp} FN={fn}')\n    print(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}')\n    print(f'Calibration : ink={ink_mean:.3f} bg={noink_mean:.3f} sep={ink_mean-noink_mean:+.3f}')\n    print(f'{\"=\"*55}')\n    return dict(dice=dice, thr=bt, prec=prec, rec=rec, f1=f1,\n                tp=tp_v, tn=tn, fp=fp, fn=fn,\n                ink_mean=ink_mean, noink_mean=noink_mean,\n                prob_crop=prob_crop, pred=pred, gt=gt_mask)\n\ndef save_plots(res, title, fname):\n    fig, ax = plt.subplots(2,3, figsize=(18,12))\n    ax[0,0].imshow(res['gt'],        cmap='gray');    ax[0,0].set_title('Ground Truth')\n    ax[0,1].imshow(res['prob_crop'], cmap='inferno'); ax[0,1].set_title('Probability Map')\n    ax[0,2].imshow(res['pred'],      cmap='gray');    ax[0,2].set_title(f'Prediction dice={res[\"dice\"]:.3f}')\n    err = np.zeros((*res['gt'].shape,3), dtype=np.uint8)\n    err[(res['pred']==1)&(res['gt']==1)] = [0,255,0]\n    err[(res['pred']==1)&(res['gt']==0)] = [255,0,0]\n    err[(res['pred']==0)&(res['gt']==1)] = [0,0,255]\n    ax[1,0].imshow(err); ax[1,0].set_title('TP=green FP=red FN=blue')\n    ax[1,1].hist(res['prob_crop'][res['gt']==1].ravel(), bins=50, alpha=0.7,\n                 label=f'ink μ={res[\"ink_mean\"]:.2f}', color='orange', density=True)\n    ax[1,1].hist(res['prob_crop'][res['gt']==0].ravel(), bins=50, alpha=0.7,\n                 label=f'bg μ={res[\"noink_mean\"]:.2f}', color='blue', density=True)\n    ax[1,1].axvline(res['thr'], color='r', ls='--', label=f'thr={res[\"thr\"]:.2f}')\n    ax[1,1].set_title('Probability Distribution'); ax[1,1].legend()\n    ts = np.arange(0.20,0.85,0.01); ds=[]\n    for t in ts:\n        p=(res['prob_crop']>t).astype(np.float32)\n        ds.append((2*(p*res['gt']).sum()+1)/(p.sum()+res['gt'].sum()+1))\n    ax[1,2].plot(ts,ds); ax[1,2].axvline(res['thr'],color='r',ls='--')\n    ax[1,2].set_title('Dice vs Threshold'); ax[1,2].grid(True)\n    for a in [ax[0,0],ax[0,1],ax[0,2],ax[1,0]]: a.axis('off')\n    plt.suptitle(title, fontsize=13, y=1.01)\n    plt.tight_layout()\n    plt.savefig(OUTPUT+fname, dpi=100, bbox_inches='tight'); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  12b. RICH FRAGMENT 1 VISUALISATION\n#       Called only for Fold C (the true unseen test fragment).\n#\n#  Saves TWO figures:\n#\n#  Figure A — Full-resolution overview (frag1_test_overview.png)\n#    Row 1: CT surface (mid-z slice)  | IR image  | GT label\n#    Row 2: Probability map           | Binary prediction | Error map\n#\n#  Figure B — Detail panel (frag1_test_detail.png)\n#    Row 1: Histogram | Dice-vs-threshold | Per-metric bar chart\n#    Row 2: 3 random 512×512 zoom crops showing\n#           CT | GT | Pred | Error side-by-side for each crop\n#\n#  Also saves:\n#    frag1_test_prediction.png  — prediction binary mask (full res)\n#    frag1_test_probmap.png     — probability heat-map  (full res)\n#    frag1_test_errormap.png    — RGB error map         (full res)\n#    frag1_test_results.txt     — all metrics as text\n# ════════════════════════════════════════════════════════════\ndef save_fragment1_results(res, frag_path, z_list):\n    \"\"\"\n    res       : dict returned by evaluate()\n    frag_path : path to Fragment 1 directory\n    z_list    : z-slices used for training/inference\n    \"\"\"\n    prob_crop = res['prob_crop']\n    gt        = res['gt']\n    pred      = res['pred']\n    H, W      = gt.shape\n\n    # ── load CT mid-z surface slice ──────────────────────────\n    cache   = _cache_store.get((frag_path, tuple(z_list)))\n    mid_z   = z_list[len(z_list)//2]\n    ct_surf = cache[mid_z].astype(np.float32)   # (H,W) float32 in [0,1]\n\n    # ── load IR image ─────────────────────────────────────────\n    ir_path = os.path.join(frag_path, 'ir.png')\n    if not os.path.exists(ir_path):\n        for name in ['infrared.png', 'ir.tif']:\n            c = os.path.join(frag_path, name)\n            if os.path.exists(c): ir_path = c; break\n    if os.path.exists(ir_path):\n        ir_img = cv2.imread(ir_path, cv2.IMREAD_GRAYSCALE).astype(np.float32)\n        if ir_img.shape != (H, W):\n            ir_img = cv2.resize(ir_img, (W, H), interpolation=cv2.INTER_LINEAR)\n        ir_img = (ir_img - ir_img.min()) / (ir_img.max() - ir_img.min() + 1e-6)\n    else:\n        ir_img = np.zeros((H, W), dtype=np.float32)\n        print('  IR image not found — using blank channel in visualisation.')\n\n    # ── error map (RGB) ───────────────────────────────────────\n    err_rgb = np.zeros((H, W, 3), dtype=np.uint8)\n    err_rgb[(pred==1)&(gt==1)] = [0,  255, 0  ]   # TP  green\n    err_rgb[(pred==1)&(gt==0)] = [255, 0,  0  ]   # FP  red\n    err_rgb[(pred==0)&(gt==1)] = [0,   0,  255]   # FN  blue\n    # TN left as black\n\n    # ── overlay: CT with GT contour in yellow, pred in cyan ──\n    ct_bgr  = cv2.cvtColor((ct_surf*255).astype(np.uint8), cv2.COLOR_GRAY2BGR)\n    gt_u8   = (gt*255).astype(np.uint8)\n    pred_u8 = (pred*255).astype(np.uint8)\n    # find contours and draw them\n    gt_cnts,   _ = cv2.findContours(gt_u8,   cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    pred_cnts, _ = cv2.findContours(pred_u8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    overlay = ct_bgr.copy()\n    cv2.drawContours(overlay, gt_cnts,   -1, (0, 255, 255), 1)   # GT  — yellow\n    cv2.drawContours(overlay, pred_cnts, -1, (255, 0, 255), 1)   # Pred— magenta\n    overlay_rgb = cv2.cvtColor(overlay, cv2.COLOR_BGR2RGB)\n\n    # ══════════════════════════════════════════════════════════\n    # FIGURE A — Full-resolution overview\n    # ══════════════════════════════════════════════════════════\n    fig_a, axes_a = plt.subplots(2, 3, figsize=(22, 14),\n                                 gridspec_kw={'hspace': 0.05, 'wspace': 0.05})\n\n    axes_a[0,0].imshow(ct_surf, cmap='gray', vmin=0, vmax=1)\n    axes_a[0,0].set_title(f'CT Surface Volume\\n(z-slice {mid_z})', fontsize=11)\n\n    axes_a[0,1].imshow(ir_img, cmap='gray', vmin=0, vmax=1)\n    axes_a[0,1].set_title('Infrared Image', fontsize=11)\n\n    axes_a[0,2].imshow(gt, cmap='gray')\n    axes_a[0,2].set_title('Ground Truth Ink Labels\\n(seen here for the FIRST time)', fontsize=11)\n\n    im = axes_a[1,0].imshow(prob_crop, cmap='inferno', vmin=0, vmax=1)\n    axes_a[1,0].set_title('Predicted Probability Map', fontsize=11)\n    fig_a.colorbar(im, ax=axes_a[1,0], fraction=0.046, pad=0.04)\n\n    axes_a[1,1].imshow(pred, cmap='gray')\n    axes_a[1,1].set_title(f'Binary Prediction  (thr={res[\"thr\"]:.2f})\\n'\n                           f'Dice={res[\"dice\"]:.4f}', fontsize=11)\n\n    axes_a[1,2].imshow(err_rgb)\n    p_tp = mpatches.Patch(color='green',  label=f'TP={res[\"tp\"]:,}')\n    p_fp = mpatches.Patch(color='red',    label=f'FP={res[\"fp\"]:,}')\n    p_fn = mpatches.Patch(color='blue',   label=f'FN={res[\"fn\"]:,}')\n    p_tn = mpatches.Patch(color='black',  label=f'TN={res[\"tn\"]:,}')\n    axes_a[1,2].legend(handles=[p_tp,p_fp,p_fn,p_tn],\n                       loc='lower right', fontsize=8, framealpha=0.8)\n    axes_a[1,2].set_title('Error Map', fontsize=11)\n\n    for ax in axes_a.flatten():\n        ax.axis('off')\n\n    plt.suptitle(\n        f'Fragment 1 — FIRST SEEN TEST RESULT\\n'\n        f'Dice={res[\"dice\"]:.4f}  Precision={res[\"prec\"]:.4f}  '\n        f'Recall={res[\"rec\"]:.4f}  F1={res[\"f1\"]:.4f}  '\n        f'FP/TP={res[\"fp\"]/(res[\"tp\"]+1e-8):.2f}',\n        fontsize=14, y=1.01, fontweight='bold')\n    plt.tight_layout()\n    fig_a.savefig(OUTPUT+'frag1_test_overview.png', dpi=120,\n                  bbox_inches='tight')\n    plt.close(fig_a)\n    print('  Saved: frag1_test_overview.png')\n\n    # ══════════════════════════════════════════════════════════\n    # FIGURE B — Overlay + Detail panel\n    # ══════════════════════════════════════════════════════════\n\n    # pick 3 representative 512×512 crop centres\n    # strategy: find bounding box of ink region, sample inside it\n    ys_ink, xs_ink = np.where(gt > 0)\n    crop_sz = 512\n    crops = []\n    np.random.seed(0)\n    attempts = 0\n    while len(crops) < 3 and attempts < 200:\n        attempts += 1\n        ci = np.random.randint(len(ys_ink))\n        cy = int(np.clip(ys_ink[ci], crop_sz//2, H - crop_sz//2))\n        cx = int(np.clip(xs_ink[ci], crop_sz//2, W - crop_sz//2))\n        y0, y1 = cy - crop_sz//2, cy + crop_sz//2\n        x0, x1 = cx - crop_sz//2, cx + crop_sz//2\n        if y0 < 0 or x0 < 0 or y1 > H or x1 > W: continue\n        # ensure crop has some ink\n        if gt[y0:y1, x0:x1].mean() > 0.01:\n            crops.append((y0, y1, x0, x1))\n\n    # fallback if not enough crops found\n    while len(crops) < 3:\n        crops.append((0, min(crop_sz,H), 0, min(crop_sz,W)))\n\n    # 4 rows: overlay row + 3 crop rows; 4 cols per crop: CT|GT|Pred|Error\n    fig_b = plt.figure(figsize=(22, 22))\n    gs    = fig_b.add_gridspec(4, 4, hspace=0.08, wspace=0.05)\n\n    # Row 0: full-image overlay spanning all 4 columns\n    ax_ov = fig_b.add_subplot(gs[0, :])\n    ax_ov.imshow(overlay_rgb)\n    ax_ov.set_title(\n        'Overlay: CT surface + GT contour (yellow) + Prediction contour (magenta)',\n        fontsize=11)\n    ax_ov.axis('off')\n\n    crop_titles = [['CT (mid-z)', 'Ground Truth', 'Prediction', 'Error Map']] * 3\n    for row_i, (y0, y1, x0, x1) in enumerate(crops):\n        ct_crop  = ct_surf  [y0:y1, x0:x1]\n        gt_crop  = gt       [y0:y1, x0:x1]\n        pr_crop  = pred     [y0:y1, x0:x1]\n        er_crop  = err_rgb  [y0:y1, x0:x1]\n        pb_crop  = prob_crop[y0:y1, x0:x1]\n\n        panels = [\n            (ct_crop,  'gray',     f'CT z={mid_z}  crop {row_i+1}'),\n            (gt_crop,  'gray',     'Ground Truth'),\n            (pb_crop,  'inferno',  f'Probability  (thr={res[\"thr\"]:.2f})'),\n            (er_crop,  None,       'Error  (green=TP  red=FP  blue=FN)'),\n        ]\n        for col_i, (img_data, cmap, title) in enumerate(panels):\n            ax = fig_b.add_subplot(gs[row_i+1, col_i])\n            if cmap:\n                ax.imshow(img_data, cmap=cmap)\n            else:\n                ax.imshow(img_data)\n            ax.set_title(title, fontsize=9)\n            ax.axis('off')\n            # draw crop boundary rectangle on overlay\n            if col_i == 0:\n                rect = mpatches.Rectangle(\n                    (x0, y0), x1-x0, y1-y0,\n                    linewidth=2,\n                    edgecolor=['cyan','lime','yellow'][row_i],\n                    facecolor='none')\n                ax_ov.add_patch(rect)\n                ax_ov.text(x0+10, y0+30, f'Crop {row_i+1}',\n                           color=['cyan','lime','yellow'][row_i],\n                           fontsize=9, fontweight='bold')\n\n    plt.suptitle(\n        f'Fragment 1 — Detail Crops  |  '\n        f'Dice={res[\"dice\"]:.4f}  Precision={res[\"prec\"]:.4f}  '\n        f'Recall={res[\"rec\"]:.4f}',\n        fontsize=13, y=1.005, fontweight='bold')\n    fig_b.savefig(OUTPUT+'frag1_test_detail.png', dpi=120,\n                  bbox_inches='tight')\n    plt.close(fig_b)\n    print('  Saved: frag1_test_detail.png')\n\n    # ══════════════════════════════════════════════════════════\n    # FIGURE C — Statistics panel\n    # ══════════════════════════════════════════════════════════\n    fig_c, axes_c = plt.subplots(1, 3, figsize=(18, 5))\n\n    # Histogram\n    axes_c[0].hist(prob_crop[gt==1].ravel(), bins=60, alpha=0.75,\n                   label=f'Ink pixels  μ={res[\"ink_mean\"]:.3f}',\n                   color='darkorange', density=True)\n    axes_c[0].hist(prob_crop[gt==0].ravel(), bins=60, alpha=0.75,\n                   label=f'BG pixels   μ={res[\"noink_mean\"]:.3f}',\n                   color='steelblue', density=True)\n    axes_c[0].axvline(res['thr'], color='red', ls='--', lw=2,\n                      label=f'Threshold = {res[\"thr\"]:.2f}')\n    axes_c[0].set_title('Probability Distribution — Fragment 1', fontsize=11)\n    axes_c[0].set_xlabel('Predicted Probability'); axes_c[0].set_ylabel('Density')\n    axes_c[0].legend(fontsize=9); axes_c[0].grid(True, alpha=0.4)\n\n    # Dice-vs-threshold\n    ts = np.arange(0.15, 0.90, 0.005)\n    ds = []\n    for t in ts:\n        p = (prob_crop > t).astype(np.float32)\n        ds.append((2*(p*gt).sum()+1)/(p.sum()+gt.sum()+1))\n    axes_c[1].plot(ts, ds, color='steelblue', lw=2)\n    axes_c[1].axvline(res['thr'], color='red',    ls='--', lw=2,\n                      label=f'Best thr={res[\"thr\"]:.2f}→Dice={res[\"dice\"]:.4f}')\n    axes_c[1].axhline(0.80,       color='green',  ls=':',  lw=1, label='0.80')\n    axes_c[1].axhline(0.89,       color='orange', ls=':',  lw=1, label='0.89')\n    axes_c[1].set_title('Dice Score vs Threshold', fontsize=11)\n    axes_c[1].set_xlabel('Threshold'); axes_c[1].set_ylabel('Dice')\n    axes_c[1].legend(fontsize=9); axes_c[1].grid(True, alpha=0.4)\n    axes_c[1].set_ylim(0, 1)\n\n    # Metric bar chart\n    metrics     = ['Dice', 'Precision', 'Recall', 'F1']\n    metric_vals = [res['dice'], res['prec'], res['rec'], res['f1']]\n    colours     = ['steelblue','darkorange','green','purple']\n    bars = axes_c[2].bar(metrics, metric_vals, color=colours, alpha=0.8, width=0.5)\n    for bar, val in zip(bars, metric_vals):\n        axes_c[2].text(bar.get_x() + bar.get_width()/2,\n                       bar.get_height() + 0.01,\n                       f'{val:.4f}', ha='center', va='bottom',\n                       fontsize=10, fontweight='bold')\n    axes_c[2].axhline(0.80, color='red',  ls='--', lw=1, label='0.80 target')\n    axes_c[2].axhline(0.89, color='gold', ls='--', lw=1, label='0.89 target')\n    axes_c[2].set_ylim(0, 1.12)\n    axes_c[2].set_title('Summary Metrics — Fragment 1', fontsize=11)\n    axes_c[2].legend(fontsize=9); axes_c[2].grid(True, axis='y', alpha=0.4)\n\n    plt.suptitle('Fragment 1 — Statistical Analysis', fontsize=13,\n                 fontweight='bold')\n    plt.tight_layout()\n    fig_c.savefig(OUTPUT+'frag1_test_stats.png', dpi=120, bbox_inches='tight')\n    plt.close(fig_c)\n    print('  Saved: frag1_test_stats.png')\n\n    # ── save full-resolution individual outputs ───────────────\n    # Prediction binary mask (white=ink, black=background)\n    cv2.imwrite(OUTPUT+'frag1_test_prediction.png',\n                (pred * 255).astype(np.uint8))\n    print('  Saved: frag1_test_prediction.png')\n\n    # Probability heat-map as 16-bit PNG\n    cv2.imwrite(OUTPUT+'frag1_test_probmap.png',\n                (prob_crop * 65535).astype(np.uint16))\n    print('  Saved: frag1_test_probmap.png  (16-bit, scale ×65535)')\n\n    # Error map RGB\n    cv2.imwrite(OUTPUT+'frag1_test_errormap.png',\n                cv2.cvtColor(err_rgb, cv2.COLOR_RGB2BGR))\n    print('  Saved: frag1_test_errormap.png')\n\n    # ── text results ──────────────────────────────────────────\n    with open(OUTPUT+'frag1_test_results.txt', 'w') as f:\n        f.write('FRAGMENT 1 — FIRST-SEEN TEST RESULTS\\n')\n        f.write('='*50 + '\\n')\n        f.write(f'Architecture  : Unet++ / efficientnet-b4 / SCSE\\n')\n        f.write(f'Training      : Fold C  (train=Frag2+Frag3, val=Frag1)\\n')\n        f.write(f'Fragment 1    : NEVER seen during training\\n')\n        f.write(f'Channels(z)   : {N_CH}  Z={Z_SLICES}\\n')\n        f.write('='*50 + '\\n')\n        f.write(f'Dice          : {res[\"dice\"]:.4f}\\n')\n        f.write(f'Threshold     : {res[\"thr\"]:.2f}\\n')\n        f.write(f'Precision     : {res[\"prec\"]:.4f}\\n')\n        f.write(f'Recall        : {res[\"rec\"]:.4f}\\n')\n        f.write(f'F1            : {res[\"f1\"]:.4f}\\n')\n        f.write(f'TP            : {res[\"tp\"]:,}\\n')\n        f.write(f'TN            : {res[\"tn\"]:,}\\n')\n        f.write(f'FP            : {res[\"fp\"]:,}\\n')\n        f.write(f'FN            : {res[\"fn\"]:,}\\n')\n        f.write(f'FP/TP ratio   : {res[\"fp\"]/(res[\"tp\"]+1e-8):.2f}\\n')\n        f.write(f'Ink mean prob : {res[\"ink_mean\"]:.4f}\\n')\n        f.write(f'BG  mean prob : {res[\"noink_mean\"]:.4f}\\n')\n        f.write(f'Separation    : {res[\"ink_mean\"]-res[\"noink_mean\"]:+.4f}\\n')\n        f.write('='*50 + '\\n')\n        f.write('Output files:\\n')\n        f.write('  frag1_test_overview.png   — CT | IR | GT | Prob | Pred | Error\\n')\n        f.write('  frag1_test_detail.png     — overlay + 3 zoom crops\\n')\n        f.write('  frag1_test_stats.png      — histogram + Dice curve + bar chart\\n')\n        f.write('  frag1_test_prediction.png — binary mask (full resolution)\\n')\n        f.write('  frag1_test_probmap.png    — 16-bit probability map\\n')\n        f.write('  frag1_test_errormap.png   — RGB error map\\n')\n    print('  Saved: frag1_test_results.txt')\n\n\n# ════════════════════════════════════════════════════════════\n#  13. LEAVE-ONE-OUT TRAINING\n# ════════════════════════════════════════════════════════════\nFOLDS = [\n    ('A', [FRAG1, FRAG2], FRAG3, 'fold_A.pth', FRAG3),\n    ('B', [FRAG1, FRAG3], FRAG2, 'fold_B.pth', FRAG2),\n    ('C', [FRAG2, FRAG3], FRAG1, 'fold_C.pth', FRAG1),  # ← unseen test\n]\n\nfold_results = {}\nfold_history = {}\n\nfor fold_name, train_frags, val_frag, ckpt_name, pred_frag in FOLDS:\n    print(f'\\n{\"=\"*60}')\n    print(f'FOLD {fold_name}  —  train: {[os.path.basename(f) for f in train_frags]}'\n          f'  val: {os.path.basename(val_frag)}')\n    if fold_name == 'C':\n        print(f'  ★  Fragment 1 is the held-out TEST fragment — never seen during training')\n    print(f'  Epochs={EPOCHS}  LR={LR:.1e}  SWA from ep{SWA_START}')\n    print(f'{\"=\"*60}')\n\n    print('\\nBuilding training datasets ...')\n    train_dl = make_loader([(f, True) for f in train_frags],\n                            shuffle_sampler=True, max_p=MAX_PATCHES_TR)\n    print('\\nBuilding validation dataset ...')\n    val_dl   = make_loader([(val_frag, False)],\n                            shuffle_sampler=False, max_p=0)\n    print(f'Train batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    model = build_model(N_CH)\n    model, best_dice, history = run_training(\n        model, train_dl, val_dl,\n        n_epochs=EPOCHS, lr=LR,\n        swa_start=SWA_START,\n        ckpt_name=ckpt_name)\n    fold_history[fold_name] = history\n    print(f'\\nFold {fold_name} best val metric: {best_dice:.4f}')\n\n    del train_dl, val_dl\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    print(f'\\nPredicting: {os.path.basename(pred_frag)}')\n    ckpt = torch.load(OUTPUT+ckpt_name, map_location=DEVICE)\n    assert ckpt['n_ch'] == N_CH\n    inf_model = build_model(N_CH)\n    inf_model.load_state_dict(ckpt['state'], strict=True)\n\n    prob_map, gt_mask = predict_fragment(inf_model, pred_frag, Z_SLICES)\n    res = evaluate(prob_map, gt_mask,\n                   label=f'Fold {fold_name} — {os.path.basename(pred_frag)}')\n    res['best_val_dice'] = best_dice\n    fold_results[fold_name] = res\n\n    if fold_name == 'C':\n        # ── RICH SAVE for Fragment 1 (true unseen test) ───────\n        print('\\nGenerating comprehensive Fragment 1 test visualisation ...')\n        save_fragment1_results(res, FRAG1, Z_SLICES)\n    else:\n        save_plots(res,\n                   title=f'Fold {fold_name} — {os.path.basename(pred_frag)} — Dice={res[\"dice\"]:.4f}',\n                   fname=f'fold_{fold_name}_prediction.png')\n\n    del inf_model, model\n    gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n\n\n# ════════════════════════════════════════════════════════════\n#  14. LEARNING CURVES\n# ════════════════════════════════════════════════════════════\nfig, axes = plt.subplots(2, 3, figsize=(18,10))\nfor col, fold_name in enumerate(['A','B','C']):\n    hist = fold_history[fold_name]\n    axes[0,col].plot(hist['tl'], label='train')\n    axes[0,col].plot(hist['vl'], label='val')\n    axes[0,col].set_title(f'Fold {fold_name} — Loss')\n    axes[0,col].legend(); axes[0,col].grid(True)\n    axes[1,col].plot(hist['td'], label='train')\n    axes[1,col].plot(hist['vd'], label='val')\n    axes[1,col].axhline(0.65, color='orange', ls='--', label='0.65')\n    axes[1,col].axhline(0.80, color='r',      ls='--', label='0.80')\n    axes[1,col].set_title(f'Fold {fold_name} — Dice'\n                           + (' ★TEST' if fold_name=='C' else ''))\n    axes[1,col].legend(); axes[1,col].grid(True)\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves_all_folds.png', dpi=100); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  15. FINAL SUMMARY\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('FINAL SUMMARY — ALL FOLDS')\nprint('='*60)\nfor fn in ['A','B','C']:\n    r    = fold_results[fn]\n    held = {'A':'Frag3','B':'Frag2','C':'Frag1 ★TEST'}[fn]\n    print(f'Fold {fn} (held={held}): '\n          f'val_dice={r[\"best_val_dice\"]:.4f}  '\n          f'test_dice={r[\"dice\"]:.4f}  '\n          f'prec={r[\"prec\"]:.4f}  rec={r[\"rec\"]:.4f}  '\n          f'sep={r[\"ink_mean\"]-r[\"noink_mean\"]:+.3f}')\n\nr_c = fold_results['C']\nprint(f'\\n★  FRAG1 TEST DICE : {r_c[\"dice\"]:.4f}')\nprint(f'★  Threshold       : {r_c[\"thr\"]:.2f}')\nprint(f'★  Precision       : {r_c[\"prec\"]:.4f}')\nprint(f'★  Recall          : {r_c[\"rec\"]:.4f}')\nprint(f'★  FP/TP ratio     : {r_c[\"fp\"]/(r_c[\"tp\"]+1e-8):.2f}')\n\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write('VESUVIUS INK DETECTION — v13\\n')\n    f.write('Leave-one-out cross-validation\\n')\n    f.write('='*50 + '\\n')\n    f.write(f'Architecture  : Unet++ / efficientnet-b4 / SCSE\\n')\n    f.write(f'Channels(z)   : {N_CH}  Z={Z_SLICES}\\n')\n    f.write(f'Normalisation : per-patch channel-wise\\n')\n    f.write(f'Loss          : Tversky(α=0.3,β=0.7) + Focal(α=0.8,γ=2)\\n')\n    f.write(f'Epochs/fold   : {EPOCHS}  |  MAX_PATCHES: {MAX_PATCHES_TR}\\n')\n    f.write('='*50 + '\\n')\n    for fn in ['A','B','C']:\n        r    = fold_results[fn]\n        held = {'A':'Frag3','B':'Frag2','C':'Frag1-TEST'}[fn]\n        f.write(f'Fold {fn} ({held}): val={r[\"best_val_dice\"]:.4f}  '\n                f'dice={r[\"dice\"]:.4f}  prec={r[\"prec\"]:.4f}  '\n                f'rec={r[\"rec\"]:.4f}  '\n                f'sep={r[\"ink_mean\"]-r[\"noink_mean\"]:+.3f}\\n')\n    f.write('='*50 + '\\n')\n    f.write(f'FRAG1 DICE    : {r_c[\"dice\"]:.4f}\\n')\n    f.write(f'Threshold     : {r_c[\"thr\"]:.2f}\\n')\n    f.write(f'Precision     : {r_c[\"prec\"]:.4f}\\n')\n    f.write(f'Recall        : {r_c[\"rec\"]:.4f}\\n')\n    f.write(f'F1            : {r_c[\"f1\"]:.4f}\\n')\n    f.write(f'TP={r_c[\"tp\"]:,}  TN={r_c[\"tn\"]:,}  '\n            f'FP={r_c[\"fp\"]:,}  FN={r_c[\"fn\"]:,}\\n')\n    f.write(f'FP/TP         : {r_c[\"fp\"]/(r_c[\"tp\"]+1e-8):.2f}\\n')\n    f.write(f'Calibration   : ink={r_c[\"ink_mean\"]:.3f}  '\n            f'bg={r_c[\"noink_mean\"]:.3f}  '\n            f'sep={r_c[\"ink_mean\"]-r_c[\"noink_mean\"]:+.3f}\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  Folds A+B:  fold_A.pth | fold_A_prediction.png')\nprint('              fold_B.pth | fold_B_prediction.png')\nprint('  Fold C (Fragment 1 TEST):')\nprint('    fold_C.pth')\nprint('    frag1_test_overview.png    — CT | IR | GT | Prob | Pred | Error')\nprint('    frag1_test_detail.png      — overlay + 3 zoom crops')\nprint('    frag1_test_stats.png       — histogram + Dice curve + metrics bar')\nprint('    frag1_test_prediction.png  — binary mask (full resolution)')\nprint('    frag1_test_probmap.png     — 16-bit probability map')\nprint('    frag1_test_errormap.png    — RGB error map')\nprint('    frag1_test_results.txt     — all metrics')\nprint('  curves_all_folds.png | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n# ============================================================\n#  VESUVIUS INK DETECTION — v10  VOLUMETRIC 3D APPROACH\n#\n#  Philosophy: ink is a 3D physical phenomenon.\n#  Treating z-slices as channels throws away the depth structure.\n#  This version processes the volume as volume.\n#\n#  Non-traditional innovations:\n#  1. 3D CNN Stem → learns the true volumetric ink signature\n#     before projecting to 2D feature maps for the decoder.\n#  2. Z-Axial Cross-Slice Attention at the bottleneck.\n#     Every (x,y) position attends across all z simultaneously.\n#  3. Local patch normalisation per z-slice (removes\n#     fragment-level scanner drift → kills Frag1/2/3 domain gap).\n#  4. Boundary-Aware Compound Loss: Tversky + Focal +\n#     boundary-weighted BCE (up-weights pixels within 5px of\n#     ink edge — directly targets the thin-stroke failure mode).\n#  5. Confidence-weighted soft self-distillation in Stage 2\n#     (soft targets from Stage-1 model on Frag2, no hard threshold,\n#     no leakage of Fragment 1 into training).\n#\n#  Data protocol:\n#  - Stage 1 : train=Frag2, val=Frag3\n#  - Stage 2 : train=Frag2 (soft-distilled), val=Frag3\n#  - Frag 1  : ONLY in final test evaluation — never in training\n#\n#  Memory safety:\n#  - All training objects explicitly deleted + cache cleared\n#    between stages and after pseudo-label inference.\n#  - BATCH_SIZE=2, GRAD_ACCUM=8, AMP enabled.\n#  - No TTA.\n# ============================================================\n\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim.swa_utils import AveragedModel, SWALR, update_bn\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:128\"\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = True\n\n# ── paths ─────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── global hyper-params ───────────────────────────────────────\nDEVICE      = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE  = 192\nSTRIDE_TR   = 32        # 50 % overlap – training\nSTRIDE_INF  = 32         # 75 % overlap – inference\nBATCH_SIZE  = 16\nGRAD_ACCUM  = 8          # effective batch = 16\nEPOCHS_S1   = 20\nEPOCHS_S2   = 10\nSWA_START_S1= 20\nSWA_START_S2= 12\nLR          = 1e-4\nLR_S2       = 3e-5\nWEIGHT_DECAY= 1e-5\nPATIENCE    = 10\nNUM_WORKERS = 0\n\n# z-slice selection\nZ_FULL = list(range(16, 36))\nN_CH   = 16              # number of z-slices kept\n# 3D stem collapses z → use a small number of 3D conv groups\nZ_GROUPS = 2             # 3D stem output channels (×spatial decoder channels)\n\n# dataset params\nINK_MIN_POS = 0.02\nNEG_RATIO   = 0.3\nPOS_WEIGHT  = 3.0\nSOFT_WEIGHT = 0.4        # weight of soft-distillation loss vs hard GT loss\n\n# boundary loss\nBOUNDARY_RADIUS = 5      # dilation radius for boundary mask (pixels)\nBOUNDARY_WEIGHT = 3.0    # weight multiplier on boundary pixels\n\nprint(f\"Device : {DEVICE}  |  Channels (z): {N_CH}  |  Patch: {PATCH_SIZE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. VARIANCE-RANKED Z-SLICE SELECTION\n# ════════════════════════════════════════════════════════════\ndef select_z_slices(frag_path, z_candidates, n_keep):\n    vol_dir  = os.path.join(frag_path, 'surface_volume')\n    msk_path = os.path.join(frag_path, 'inklabels.png')\n    msk      = cv2.imread(msk_path, 0) if os.path.exists(msk_path) else None\n    scores   = {}\n    for z in z_candidates:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if not os.path.exists(p):\n            continue\n        sl     = tifffile.imread(p).astype(np.float32)\n        pixels = sl[msk > 0] if (msk is not None and (msk > 0).any()) else sl.ravel()\n        scores[z] = float(np.var(pixels))\n    ranked = sorted(scores, key=scores.get, reverse=True)[:n_keep]\n    ranked.sort()\n    print(f\"  Selected z-slices: {ranked}\")\n    return ranked\n\nprint('\\nRanking z-slices on Fragment 2 ...')\nZ_SLICES = select_z_slices(FRAG2, Z_FULL, N_CH)\n\n\n# ════════════════════════════════════════════════════════════\n#  2. SLICE CACHE  — stored as float16 to save RAM\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s; all_v.append(s.ravel())\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Slice cache: {len(cache)} slices, {mb:.0f} MB\")\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  3. LOCAL PATCH NORMALISATION\n#\n#  Key innovation #3: normalise each patch independently\n#  per z-slice (local mean/std).  This removes fragment-level\n#  scanner bias so the model never needs to adapt to a new\n#  fragment's intensity distribution.\n# ════════════════════════════════════════════════════════════\ndef local_norm_patch(vol):\n    \"\"\"\n    vol : np.float32 (H, W, Z)\n    Returns float32 (H, W, Z) normalised slice-by-slice.\n    \"\"\"\n    out = np.empty_like(vol)\n    for z in range(vol.shape[2]):\n        sl  = vol[:, :, z]\n        mu  = sl.mean()\n        std = sl.std() + 1e-6\n        out[:, :, z] = (sl - mu) / std\n    return out\n\n\n# ════════════════════════════════════════════════════════════\n#  4. BOUNDARY MASK HELPER  (for boundary-aware loss)\n# ════════════════════════════════════════════════════════════\ndef make_boundary_mask(mask_np, radius=BOUNDARY_RADIUS):\n    \"\"\"\n    mask_np : np.uint8 (H, W), values 0/1\n    Returns np.float32 (H, W), 1 on boundary pixels, 0 elsewhere.\n    \"\"\"\n    kernel = cv2.getStructuringElement(\n        cv2.MORPH_ELLIPSE, (2*radius+1, 2*radius+1))\n    dilated  = cv2.dilate(mask_np, kernel)\n    eroded   = cv2.erode(mask_np,  kernel)\n    boundary = (dilated != eroded).astype(np.float32)\n    return boundary\n\n\n# ════════════════════════════════════════════════════════════\n#  5. DATASET\n# ════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    \"\"\"\n    Returns (volume, gt_mask, boundary_mask, soft_label).\n\n    soft_label : float32 (H, W) — soft probability map from Stage-1\n                 model inference on Frag2.  None during Stage 1\n                 (replaced with a zeros tensor in __getitem__).\n    \"\"\"\n    def __init__(self, frag_path, z_list, stride,\n                 transform=None, neg_ratio=0.0,\n                 soft_prob_map=None):\n        self.cache       = load_slice_cache(frag_path, z_list)\n        self.z_list      = z_list\n        self.tf          = transform\n        self.soft_map    = soft_prob_map   # (H, W) float32 or None\n\n        msk = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n        self.mask = (msk > 0).astype(np.uint8)\n\n        ir_path = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = (cv2.imread(ir_path, 0) > 0).astype(np.uint8) \\\n                       if os.path.exists(ir_path) else None\n\n        H, W = self.mask.shape\n        pos_coords, neg_coords = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.sum() / (PATCH_SIZE * PATCH_SIZE)\n                if self.ir_mask is not None:\n                    on_pap = self.ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    mid_z  = z_list[len(z_list)//2]\n                    on_pap = float(\n                        self.cache[mid_z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                    ) > 0.1\n                if not on_pap:\n                    continue\n                if ink >= INK_MIN_POS:\n                    pos_coords.append((y, x, 1))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_coords.append((y, x, 0))\n\n        n_neg = int(len(pos_coords) * neg_ratio)\n        sel_neg = []\n        if n_neg > 0 and neg_coords:\n            np.random.shuffle(neg_coords)\n            sel_neg = neg_coords[:n_neg]\n        self.coords  = pos_coords + sel_neg\n        np.random.shuffle(self.coords)\n        self.weights = np.array(\n            [POS_WEIGHT if c[2]==1 else 1.0 for c in self.coords],\n            dtype=np.float32)\n        print(f\"  Patches: {len(pos_coords)} pos + {len(sel_neg)} neg \"\n              f\"= {len(self.coords)} total\")\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x, _ = self.coords[idx]\n        # (H, W, Z) volume patch\n        slices = [\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list\n        ]\n        vol = np.stack(slices, axis=-1)         # (H, W, Z)\n        vol = local_norm_patch(vol)             # LOCAL normalisation — key innovation\n\n        msk = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n\n        # Augmentation applied jointly to the whole volume + mask\n        if self.tf:\n            # albumentations expects (H, W, C); vol already is (H, W, Z)\n            out = self.tf(image=vol, mask=msk)\n            vol, msk = out['image'], out['mask']\n\n        bnd = make_boundary_mask(msk)           # (H, W) float32\n\n        # Soft label patch\n        if self.soft_map is not None:\n            soft = self.soft_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n        else:\n            soft = np.zeros((PATCH_SIZE, PATCH_SIZE), dtype=np.float32)\n\n        # vol: (H, W, Z) → (Z, H, W) for the 3D-aware network\n        vol_t  = torch.from_numpy(vol).permute(2, 0, 1).float()   # (Z, H, W)\n        msk_t  = torch.from_numpy(msk).unsqueeze(0).float()       # (1, H, W)\n        bnd_t  = torch.from_numpy(bnd).unsqueeze(0).float()       # (1, H, W)\n        soft_t = torch.from_numpy(soft).unsqueeze(0).float()      # (1, H, W)\n        return vol_t, msk_t, bnd_t, soft_t\n\n\n# ════════════════════════════════════════════════════════════\n#  6. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.15,\n                       rotate_limit=30,\n                       border_mode=cv2.BORDER_REFLECT, p=0.6),\n    A.RandomBrightnessContrast(0.2, 0.2, p=0.5),\n    A.GaussNoise(var_limit=(0.001, 0.004), p=0.3),\n    A.GaussianBlur(blur_limit=(3, 5), p=0.2),\n    A.CoarseDropout(max_holes=4, max_height=24,\n                    max_width=24, fill_value=0, p=0.3),\n    A.ElasticTransform(alpha=30, sigma=5, p=0.3),\n    A.GridDistortion(p=0.2),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  7. ARCHITECTURE — 3D Volumetric Encoder + 2D Unet++ Decoder\n#\n#  The 3D stem processes the z-dimension as a spatial dimension,\n#  not as independent channels.  It learns which z-depths carry\n#  the ink signal and collapses z into a rich 2D feature map.\n#  That feature map is then decoded by a standard Unet++ decoder.\n#\n#  Key module: ZAxialAttention\n#  At the bottleneck, each spatial position (x,y) attends across\n#  all z-slices simultaneously via a lightweight multi-head\n#  self-attention along the z axis.\n# ════════════════════════════════════════════════════════════\n\nclass ZAxialAttention(nn.Module):\n    \"\"\"\n    Cross-slice attention along the z-axis.\n    Input:  (B, C, Z, H, W)\n    Output: (B, C, Z, H, W)   (same shape, z-mixed)\n    Uses einsum-based MHA to avoid large intermediate tensors.\n    \"\"\"\n    def __init__(self, channels, z_len, n_heads=4):\n        super().__init__()\n        assert channels % n_heads == 0\n        self.n_heads = n_heads\n        self.head_dim = channels // n_heads\n        self.scale = self.head_dim ** -0.5\n        self.qkv  = nn.Linear(channels, channels * 3, bias=False)\n        self.proj = nn.Linear(channels, channels, bias=False)\n        self.norm = nn.LayerNorm(channels)\n\n    def forward(self, x):\n        # x: (B, C, Z, H, W)\n        B, C, Z, H, W = x.shape\n        # reshape: treat each (h,w) as a batch item, attend over z\n        x_in = x.permute(0, 3, 4, 2, 1)          # (B, H, W, Z, C)\n        x_in = x_in.reshape(B*H*W, Z, C)\n        residual = x_in\n        x_in = self.norm(x_in)\n        qkv  = self.qkv(x_in).reshape(B*H*W, Z, 3, self.n_heads, self.head_dim)\n        qkv  = qkv.permute(2, 0, 3, 1, 4)        # (3, BHW, heads, Z, head_dim)\n        q, k, v = qkv.unbind(0)\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(dim=-1)\n        out  = (attn @ v)                          # (BHW, heads, Z, head_dim)\n        out  = out.transpose(1, 2).reshape(B*H*W, Z, C)\n        out  = self.proj(out) + residual\n        out  = out.reshape(B, H, W, Z, C).permute(0, 4, 3, 1, 2)  # (B,C,Z,H,W)\n        return out\n\n\nclass Stem3D(nn.Module):\n    \"\"\"\n    3D convolutional stem.\n    Input:  (B, 1, Z, H, W)  — volume treated as single-channel 3D\n    Output: (B, out_ch, H, W) — z collapsed to feature maps\n\n    Uses 3 × (3D conv → BN → ReLU) layers with progressive z-pooling,\n    then a ZAxialAttention at the bottleneck, then z average-pool.\n    Designed to fit in < 2 GB extra VRAM with N_CH=16, PATCH=224.\n    \"\"\"\n    def __init__(self, z_len, out_ch=64):\n        super().__init__()\n        self.z_len = z_len\n\n        # Layer 1: preserve z, downsample H,W by 2\n        self.conv1 = nn.Sequential(\n            nn.Conv3d(1, 16, kernel_size=(3,3,3), padding=(1,1,1), bias=False),\n            nn.BatchNorm3d(16), nn.ReLU(inplace=True),\n            nn.Conv3d(16, 16, kernel_size=(1,3,3), padding=(0,1,1),\n                      stride=(1,2,2), bias=False),\n            nn.BatchNorm3d(16), nn.ReLU(inplace=True),\n        )\n        # Layer 2: pool z by 2, downsample H,W by 2\n        self.conv2 = nn.Sequential(\n            nn.Conv3d(16, 32, kernel_size=(3,3,3), padding=(1,1,1), bias=False),\n            nn.BatchNorm3d(32), nn.ReLU(inplace=True),\n            nn.Conv3d(32, 32, kernel_size=(2,3,3), padding=(0,1,1),\n                      stride=(2,2,2), bias=False),\n            nn.BatchNorm3d(32), nn.ReLU(inplace=True),\n        )\n        # Z-axial attention at reduced resolution\n        z2 = z_len // 2\n        self.z_attn = ZAxialAttention(channels=32, z_len=z2, n_heads=4)\n\n        # Layer 3: collapse remaining z\n        self.conv3 = nn.Sequential(\n            nn.Conv3d(32, out_ch, kernel_size=(3,3,3), padding=(1,1,1), bias=False),\n            nn.BatchNorm3d(out_ch), nn.ReLU(inplace=True),\n        )\n        # Adaptive pool: collapse z to 1, restore H,W to PATCH_SIZE\n        self.z_pool = nn.AdaptiveAvgPool3d((1, None, None))\n\n        # Upsample H,W back to original (was halved 2× = ×4 reduction)\n        self.upsample = nn.Sequential(\n            nn.ConvTranspose2d(out_ch, out_ch, kernel_size=4, stride=4,\n                               padding=0, bias=False),\n            nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        # x: (B, Z, H, W) — input from dataset (Z channels)\n        B, Z, H, W = x.shape\n        x = x.unsqueeze(1)                   # (B, 1, Z, H, W)\n        x = self.conv1(x)                    # (B, 16, Z, H/2, W/2)\n        x = self.conv2(x)                    # (B, 32, Z/2, H/4, W/4)\n        x = self.z_attn(x)                   # z-mixing attention\n        x = self.conv3(x)                    # (B, out_ch, Z/2, H/4, W/4)\n        x = self.z_pool(x).squeeze(2)        # (B, out_ch, H/4, W/4)\n        x = self.upsample(x)                 # (B, out_ch, H, W)\n        return x\n\n\nclass VolumetricInkNet(nn.Module):\n    \"\"\"\n    Full model:\n      Stem3D            — extracts 2D feature maps from the 3D volume\n      Unet++ decoder    — segments ink from the 2D feature maps\n      SCSE attention    — spatial + channel squeeze-excite in decoder\n\n    The Unet++ encoder is replaced by a fixed 1-channel 'dummy' that\n    receives the Stem3D output; we use SMP's MAnet with a lightweight\n    encoder as the decoder backbone.\n    \"\"\"\n    def __init__(self, z_len, stem_ch=64):\n        super().__init__()\n        self.stem = Stem3D(z_len=z_len, out_ch=stem_ch)\n\n        # Use SMP Unet++ with efficientnet-b1 as 2D backbone.\n        # We hijack it by: feeding stem output through a 1×1 conv\n        # that maps stem_ch → 3 (pretend it is an RGB image),\n        # then passing through SMP normally.\n        # This keeps the pretrained decoder weights valid.\n        self.stem_proj = nn.Sequential(\n            nn.Conv2d(stem_ch, 3, kernel_size=1, bias=False),\n            nn.BatchNorm2d(3),\n        )\n        self.decoder = smp.UnetPlusPlus(\n            encoder_name           = 'efficientnet-b4',\n            encoder_weights        = 'imagenet',\n            in_channels            = 3,\n            classes                = 1,\n            decoder_attention_type = 'scse',\n        )\n\n    def forward(self, x):\n        # x: (B, Z, H, W)\n        feat = self.stem(x)          # (B, stem_ch, H, W)\n        feat = self.stem_proj(feat)  # (B, 3, H, W)\n        return self.decoder(feat)    # (B, 1, H, W) logits\n\n\ndef build_model():\n    return VolumetricInkNet(z_len=N_CH, stem_ch=64).to(DEVICE)\n\n\n# ════════════════════════════════════════════════════════════\n#  8. LOSS FUNCTIONS\n#\n#  Boundary-Aware Compound Loss:\n#    L = λ₁·Tversky + λ₂·Focal + λ₃·BoundaryBCE\n#\n#  BoundaryBCE: standard BCE but pixels within BOUNDARY_RADIUS\n#  of an ink edge are multiplied by BOUNDARY_WEIGHT.\n#  This makes the model pay much more attention to thin strokes.\n#\n#  In Stage 2, we add a soft-distillation term:\n#    L_total = (1-α)·L_hard + α·L_soft\n#  where L_soft = BCE(pred, soft_label) weighted by confidence.\n# ════════════════════════════════════════════════════════════\ndef tversky_loss(pred, target, alpha=0.3, beta=0.7, smooth=1.):\n    p  = torch.sigmoid(pred)\n    tp = (p * target).sum(dim=(2,3))\n    fp = (p * (1-target)).sum(dim=(2,3))\n    fn = ((1-p) * target).sum(dim=(2,3))\n    return 1. - ((tp+smooth)/(tp+alpha*fp+beta*fn+smooth)).mean()\n\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = target*alpha + (1-target)*(1-alpha)\n    return (a_t * ((1-p_t)**gamma) * bce).mean()\n\ndef boundary_bce(pred, target, boundary, weight=BOUNDARY_WEIGHT):\n    \"\"\"\n    BCE where boundary pixels have extra weight.\n    pred, target, boundary : (B, 1, H, W)\n    \"\"\"\n    bce  = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    wmap = 1.0 + (weight - 1.0) * boundary   # normal=1, boundary=weight\n    return (bce * wmap).mean()\n\ndef hard_loss(pred, target, boundary, eps=0.05):\n    \"\"\"Stage-1 loss on hard GT labels.\"\"\"\n    t_s = target * (1-eps) + 0.5*eps\n    tl  = tversky_loss(pred, target)\n    fl  = focal_loss(pred, t_s)\n    bl  = boundary_bce(pred, target, boundary)\n    return 0.4*tl + 0.3*fl + 0.3*bl\n\ndef soft_distill_loss(pred, soft_label):\n    \"\"\"\n    Soft distillation: BCE against continuous soft probabilities,\n    weighted by the confidence of the soft label (further from 0.5\n    = more confident = higher weight).\n    \"\"\"\n    confidence = (2.0 * (soft_label - 0.5).abs()).clamp(0, 1)\n    bce = F.binary_cross_entropy_with_logits(pred, soft_label, reduction='none')\n    return (bce * confidence).mean()\n\ndef combined_loss(pred, target, boundary, soft_label, stage2=False):\n    hl = hard_loss(pred, target, boundary)\n    if stage2:\n        sl = soft_distill_loss(pred, soft_label)\n        return (1.0 - SOFT_WEIGHT)*hl + SOFT_WEIGHT*sl\n    return hl\n\n\n# ════════════════════════════════════════════════════════════\n#  9. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p * masks).sum(dim=(1,2,3))\n    union = p.sum(dim=(1,2,3)) + masks.sum(dim=(1,2,3))\n    return ((2.*inter+1e-5)/(union+1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.25, 0.85, 0.01):\n        p = (probs > t).astype(np.float32)\n        d = (2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d > best_d:\n            best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  10. GAUSSIAN PATCH WEIGHT (inference stitching)\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c = sz//2; sig = sz//4\n    ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs-c)**2+(ys-c)**2)/(2*sig**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\n# ════════════════════════════════════════════════════════════\n#  11. TRAINING LOOP  (shared Stage 1 / Stage 2)\n# ════════════════════════════════════════════════════════════\ndef run_training(model, train_dl, val_dl, n_epochs, lr,\n                 swa_start, ckpt_name, stage2=False,\n                 best_dice_init=0.):\n    optimizer = optim.AdamW(model.parameters(), lr=lr,\n                             weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=10, T_mult=1, eta_min=1e-6)\n    scaler    = torch.cuda.amp.GradScaler(enabled=(DEVICE=='cuda'))\n\n    swa_model  = AveragedModel(model)\n    swa_sched  = SWALR(optimizer, swa_lr=lr*0.1, anneal_epochs=5)\n    swa_active = False\n\n    best_dice = best_dice_init\n    pat_cnt   = 0\n    history   = dict(tl=[], vl=[], td=[], vd=[])\n\n    print(f'\\nLR schedule preview:')\n    _o = optim.AdamW([torch.zeros(1)], lr=lr)\n    _s = optim.lr_scheduler.CosineAnnealingWarmRestarts(_o, T_0=10, eta_min=1e-6)\n    for _e in range(5):\n        print(f'  Epoch {_e+1}: LR = {_o.param_groups[0][\"lr\"]:.2e}')\n        _s.step()\n    del _o, _s; print()\n\n    for epoch in range(n_epochs):\n        if epoch >= swa_start and not swa_active:\n            swa_active = True\n            print(f'  → SWA activated at epoch {epoch+1}')\n\n        # ── train ─────────────────────────────────────────────\n        model.train(); tl = td = 0.\n        optimizer.zero_grad()\n        for step, (imgs, msks, bnds, softs) in enumerate(\n                tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n            imgs  = imgs.to(DEVICE)\n            msks  = msks.to(DEVICE)\n            bnds  = bnds.to(DEVICE)\n            softs = softs.to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                out  = model(imgs)\n                loss = combined_loss(out, msks, bnds, softs,\n                                     stage2=stage2) / GRAD_ACCUM\n            scaler.scale(loss).backward()\n            if (step+1) % GRAD_ACCUM == 0 or (step+1) == len(train_dl):\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                scaler.step(optimizer); scaler.update()\n                optimizer.zero_grad()\n            tl += loss.item() * GRAD_ACCUM\n            td += batch_dice(out.detach(), msks)\n            del imgs, msks, bnds, softs, out, loss\n            if step % 50 == 0 and DEVICE=='cuda':\n                torch.cuda.empty_cache()\n        tl /= len(train_dl); td /= len(train_dl)\n\n        # ── validate ──────────────────────────────────────────\n        model.eval(); vl = vd = 0.\n        acc_p, acc_m = [], []\n        with torch.no_grad():\n            for imgs, msks, bnds, softs in tqdm(\n                    val_dl, desc=f'Ep{epoch+1:02d} val  ', leave=False):\n                imgs  = imgs.to(DEVICE)\n                msks  = msks.to(DEVICE)\n                bnds  = bnds.to(DEVICE)\n                softs = softs.to(DEVICE)\n                with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                    out  = model(imgs)\n                    loss = combined_loss(out, msks, bnds, softs,\n                                         stage2=False)   # always hard loss for val\n                vl += loss.item(); vd += batch_dice(out, msks)\n                acc_p.append(torch.sigmoid(out).cpu().numpy())\n                acc_m.append(msks.cpu().numpy())\n                del imgs, msks, bnds, softs, out, loss\n        vl /= len(val_dl); vd /= len(val_dl)\n\n        probs_all = np.concatenate(acc_p)\n        masks_all = np.concatenate(acc_m)\n        bt, bd    = sweep_threshold(probs_all, masks_all)\n        ink_mean   = float(probs_all[masks_all>0.5].mean()) \\\n                     if (masks_all>0.5).any() else 0.\n        noink_mean = float(probs_all[masks_all<0.5].mean()) \\\n                     if (masks_all<0.5).any() else 0.\n        del acc_p, acc_m, probs_all, masks_all\n        gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n        if swa_active:\n            swa_model.update_parameters(model)\n            swa_sched.step()\n        else:\n            scheduler.step()\n        lr_now = optimizer.param_groups[0]['lr']\n\n        history['tl'].append(tl); history['vl'].append(vl)\n        history['td'].append(td); history['vd'].append(vd)\n\n        print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n              f'train loss={tl:.4f} dice={td:.4f} | '\n              f'val loss={vl:.4f} dice={vd:.4f} | '\n              f'thr={bt:.2f}→{bd:.4f} | '\n              f'sep={ink_mean-noink_mean:+.3f}'\n              + (' [SWA]' if swa_active else ''))\n\n        save_metric = max(vd, bd)\n        if save_metric > best_dice:\n            best_dice = save_metric; pat_cnt = 0\n            torch.save({'epoch': epoch,\n                        'state': model.state_dict(),\n                        'thr': bt, 'dice': best_dice,\n                        'n_ch': N_CH},\n                       OUTPUT + ckpt_name)\n            print(f'  ✓ saved  (metric={best_dice:.4f})')\n        else:\n            pat_cnt += 1\n            if pat_cnt >= PATIENCE:\n                print(f'  ⚑ early stop at epoch {epoch+1}')\n                break\n\n        if epoch == 4 and max(history['vd']) < 0.30:\n            print('\\n⚠  val dice <0.30 after 5 epochs — '\n                  'check sep value above (target >0.05).')\n\n    if swa_active:\n        print('  Updating SWA batch-norm statistics ...')\n        # Run a forward pass over training data with swa_model\n        swa_model.train()\n        with torch.no_grad():\n            for imgs, msks, bnds, softs in train_dl:\n                imgs = imgs.to(DEVICE)\n                swa_model(imgs)\n                del imgs, msks, bnds, softs\n        model.load_state_dict(swa_model.module.state_dict())\n        del swa_model; gc.collect()\n        if DEVICE=='cuda': torch.cuda.empty_cache()\n\n    return model, best_dice, history\n\n\n# ════════════════════════════════════════════════════════════\n#  12. STAGE 1  — train=Frag2, val=Frag3\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint(f'STAGE 1  ({EPOCHS_S1} epochs)')\nprint(f'  Arch  : 3D-Stem + ZAxialAttention + Unet++(efficientnet-b1) + SCSE')\nprint(f'  Loss  : Tversky + Focal + BoundaryBCE')\nprint(f'  Train : Fragment 2   Val : Fragment 3')\nprint(f'  Frag1 : NOT touched')\nprint('='*60)\n\nprint('\\n── Fragment 2 (train) ──')\ntrain_ds_s1 = VesuviusDataset(FRAG2, Z_SLICES, stride=STRIDE_TR,\n                               transform=train_tf, neg_ratio=NEG_RATIO)\nprint('\\n── Fragment 3 (val) ──')\nval_ds_s1   = VesuviusDataset(FRAG3, Z_SLICES, stride=STRIDE_TR,\n                               transform=None, neg_ratio=0.0)\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nsampler_s1  = WeightedRandomSampler(\n    torch.from_numpy(train_ds_s1.weights), len(train_ds_s1), replacement=True)\ntrain_dl_s1 = DataLoader(train_ds_s1, batch_size=BATCH_SIZE,\n                          sampler=sampler_s1,\n                          num_workers=NUM_WORKERS, pin_memory=False)\nval_dl_s1   = DataLoader(val_ds_s1, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=False)\nprint(f'\\nTrain batches: {len(train_dl_s1)} | Val batches: {len(val_dl_s1)}')\n\nmodel_s1 = build_model()\nmodel_s1, best_s1, hist_s1 = run_training(\n    model_s1, train_dl_s1, val_dl_s1,\n    n_epochs=EPOCHS_S1, lr=LR,\n    swa_start=SWA_START_S1,\n    ckpt_name='stage1_best.pth',\n    stage2=False,\n)\nprint(f'\\nStage 1 best val metric: {best_s1:.4f}')\n\n# ── FREE Stage-1 training memory ─────────────────────────────\nprint('\\nFreeing Stage-1 resources ...')\ndel train_dl_s1, val_dl_s1, sampler_s1\ndel train_ds_s1, val_ds_s1\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n\n# ════════════════════════════════════════════════════════════\n#  13. SOFT PSEUDO-LABEL GENERATION ON FRAGMENT 2\n#\n#  Run Stage-1 model on Frag2 to generate SOFT probability maps.\n#  These are passed as soft_label to Stage 2 — no hard threshold,\n#  no leakage of Fragment 1.\n#\n#  Confidence-weighted distillation means:\n#    - pixels where model is sure (prob ~1 or ~0) dominate the loss\n#    - uncertain pixels (prob ~0.5) contribute very little\n#  → the model refines its own best signal iteratively\n# ════════════════════════════════════════════════════════════\ndef predict_soft_map(model, frag_path, z_list):\n    \"\"\"\n    Sliding-window inference → full-resolution float32 prob map.\n    Single pass, Gaussian-weighted.\n    \"\"\"\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    H, W  = next(iter(cache.values())).shape\n    pred_map = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n    coords   = [(y, x)\n                for y in range(0, H-PATCH_SIZE+1, STRIDE_INF)\n                for x in range(0, W-PATCH_SIZE+1, STRIDE_INF)]\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='Soft-label inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            vol = np.stack(slices, -1)\n            vol = local_norm_patch(vol)\n            t   = torch.from_numpy(vol).permute(2,0,1).unsqueeze(0).to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                p = torch.sigmoid(model(t)).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, p\n    del cache; gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    return (pred_map / (wgt_map + 1e-8)).astype(np.float32)\n\nprint('\\n' + '='*60)\nprint('SOFT-LABEL GENERATION ON FRAGMENT 2  (no Fragment 1 touch)')\nprint('='*60)\n\nckpt_s1 = torch.load(OUTPUT+'stage1_best.pth', map_location=DEVICE)\nassert ckpt_s1['n_ch'] == N_CH\nmodel_s1.load_state_dict(ckpt_s1['state'], strict=True)\nprint(f'Loaded: epoch={ckpt_s1[\"epoch\"]+1}, val_dice={ckpt_s1[\"dice\"]:.4f}')\n\nsoft_map_frag2 = predict_soft_map(model_s1, FRAG2, Z_SLICES)\nprint(f'Soft map — mean={soft_map_frag2.mean():.3f}, '\n      f'std={soft_map_frag2.std():.3f}, '\n      f'max={soft_map_frag2.max():.3f}')\n\n# Save soft map as 16-bit PNG for inspection\ncv2.imwrite(OUTPUT+'soft_labels_frag2.png',\n            (soft_map_frag2 * 65535).astype(np.uint16))\n\ndel model_s1; gc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\nprint('Stage-1 model freed.')\n\n\n# ════════════════════════════════════════════════════════════\n#  14. STAGE 2  — train=Frag2 (+ soft labels), val=Frag3\n#       Fragment 1 still NOT used.\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint(f'STAGE 2  ({EPOCHS_S2} epochs, LR={LR_S2:.1e})')\nprint(f'  Train : Fragment 2 with soft self-distillation')\nprint(f'  Val   : Fragment 3   |   Frag1 : RESERVED for test only')\nprint(f'  Loss  : (1-{SOFT_WEIGHT})·HardLoss + {SOFT_WEIGHT}·SoftDistill')\nprint('='*60)\n\nprint('\\n── Fragment 2 (train, soft-distilled) ──')\ntrain_ds_s2 = VesuviusDataset(FRAG2, Z_SLICES, stride=STRIDE_TR,\n                               transform=train_tf, neg_ratio=NEG_RATIO,\n                               soft_prob_map=soft_map_frag2)\nprint('\\n── Fragment 3 (val) ──')\nval_ds_s2 = VesuviusDataset(FRAG3, Z_SLICES, stride=STRIDE_TR,\n                             transform=None, neg_ratio=0.0)\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\nsampler_s2  = WeightedRandomSampler(\n    torch.from_numpy(train_ds_s2.weights), len(train_ds_s2), replacement=True)\ntrain_dl_s2 = DataLoader(train_ds_s2, batch_size=BATCH_SIZE,\n                          sampler=sampler_s2,\n                          num_workers=NUM_WORKERS, pin_memory=False)\nval_dl_s2   = DataLoader(val_ds_s2, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=False)\nprint(f'\\nTrain batches: {len(train_dl_s2)} | Val batches: {len(val_dl_s2)}')\n\nmodel_s2 = build_model()\nmodel_s2.load_state_dict(ckpt_s1['state'], strict=True)\ndel ckpt_s1; gc.collect()\n\nmodel_s2, best_s2, hist_s2 = run_training(\n    model_s2, train_dl_s2, val_dl_s2,\n    n_epochs=EPOCHS_S2, lr=LR_S2,\n    swa_start=SWA_START_S2,\n    ckpt_name='stage2_best.pth',\n    stage2=True,\n)\nprint(f'\\nStage 2 best val metric: {best_s2:.4f}')\n\n# ── FREE Stage-2 training memory ─────────────────────────────\nprint('\\nFreeing Stage-2 resources ...')\ndel train_dl_s2, val_dl_s2, sampler_s2\ndel train_ds_s2, val_ds_s2, soft_map_frag2\ngc.collect()\nif DEVICE=='cuda': torch.cuda.empty_cache()\n\n# ── Training curves ───────────────────────────────────────────\nfig, axes = plt.subplots(2, 2, figsize=(14, 10))\npairs = [\n    (hist_s1, 'Stage 1 Loss', 'tl', 'vl'),\n    (hist_s1, 'Stage 1 Dice', 'td', 'vd'),\n    (hist_s2, 'Stage 2 Loss', 'tl', 'vl'),\n    (hist_s2, 'Stage 2 Dice', 'td', 'vd'),\n]\nfor ax, (hist, title, k1, k2) in zip(axes.flatten(), pairs):\n    ax.plot(hist[k1], label='train'); ax.plot(hist[k2], label='val')\n    if 'Dice' in title:\n        ax.axhline(0.80, color='orange', ls='--', label='target 0.80')\n        ax.axhline(0.89, color='r',      ls='--', label='target 0.89')\n    ax.set_title(title); ax.legend(); ax.grid(True)\nplt.tight_layout()\nplt.savefig(OUTPUT+'curves.png', dpi=100); plt.close()\n\n\n# ════════════════════════════════════════════════════════════\n#  15. INFERENCE ON FRAGMENT 1  (first and only appearance)\n# ════════════════════════════════════════════════════════════\ndef predict_fragment(model, frag_path, z_list):\n    \"\"\"\n    Single-pass Gaussian-weighted sliding-window inference.\n    Local normalisation applied per patch (same as training).\n    \"\"\"\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    msk   = cv2.imread(os.path.join(frag_path,'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n    H, W  = msk.shape\n    pred_map = np.zeros((H,W), np.float32)\n    wgt_map  = np.zeros((H,W), np.float32)\n    coords   = [(y, x)\n                for y in range(0, H-PATCH_SIZE+1, STRIDE_INF)\n                for x in range(0, W-PATCH_SIZE+1, STRIDE_INF)]\n    with torch.no_grad():\n        for (y, x) in tqdm(coords, desc='Inference', leave=True):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            vol = np.stack(slices, -1)\n            vol = local_norm_patch(vol)\n            t   = torch.from_numpy(vol).permute(2,0,1).unsqueeze(0).to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                p = torch.sigmoid(model(t)).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, p\n    del cache; gc.collect()\n    if DEVICE=='cuda': torch.cuda.empty_cache()\n    return pred_map/(wgt_map+1e-8), msk\n\n\n# ════════════════════════════════════════════════════════════\n#  16. FINAL EVALUATION — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('FINAL TEST — FRAGMENT 1  (first time Fragment 1 is ever seen)')\nprint('='*60)\n\nif best_s2 >= best_s1:\n    best_ckpt = OUTPUT+'stage2_best.pth'\n    print(f'Using Stage-2 ckpt  (val={best_s2:.4f})')\nelse:\n    best_ckpt = OUTPUT+'stage1_best.pth'\n    print(f'Using Stage-1 ckpt  (val={best_s1:.4f})')\n\nckpt = torch.load(best_ckpt, map_location=DEVICE)\nassert ckpt['n_ch'] == N_CH\nfinal_model = build_model()\nfinal_model.load_state_dict(ckpt['state'], strict=True)\nprint(f'Epoch={ckpt[\"epoch\"]+1}, val_dice={ckpt[\"dice\"]:.4f}, '\n      f'thr={ckpt.get(\"thr\",0.5):.2f}')\n\nprob_map, msk1 = predict_fragment(final_model, FRAG1, Z_SLICES)\nprob_crop = prob_map[:msk1.shape[0], :msk1.shape[1]]\n\nink_mean   = float(prob_crop[msk1==1].mean())\nnoink_mean = float(prob_crop[msk1==0].mean())\nprint(f'\\nCalibration on Frag1:')\nprint(f'  ink pixels mean   : {ink_mean:.3f}')\nprint(f'  no-ink pixels mean: {noink_mean:.3f}')\nprint(f'  separation        : {ink_mean-noink_mean:+.3f}  (good if >0.20)')\n\nbt1, bd1   = sweep_threshold(prob_crop[np.newaxis,np.newaxis],\n                              msk1[np.newaxis,np.newaxis])\nfinal_pred = (prob_crop > bt1).astype(np.uint8)\n\ninter     = (final_pred * msk1).sum()\ndice_full = (2*inter+1)/(final_pred.sum()+msk1.sum()+1)\npf = final_pred.flatten().astype(int)\nmf = msk1.flatten().astype(int)\ntn, fp, fn, tp_v = confusion_matrix(mf, pf, labels=[0,1]).ravel()\nprec = tp_v/(tp_v+fp+1e-8)\nrec  = tp_v/(tp_v+fn+1e-8)\nf1   = 2*prec*rec/(prec+rec+1e-8)\n\nprint('\\n' + '='*60)\nprint('RESULTS — FRAGMENT 1')\nprint('='*60)\nprint(f'Dice Score  : {dice_full:.4f}')\nprint(f'Threshold   : {bt1:.2f}')\nprint(f'Precision   : {prec:.4f}')\nprint(f'Recall      : {rec:.4f}')\nprint(f'F1 Score    : {f1:.4f}')\nprint(f'TP={tp_v} | TN={tn} | FP={fp} | FN={fn}')\nprint(f'FP/TP ratio : {fp/(tp_v+1e-8):.2f}  (target <1.0)')\nprint('='*60)\n\n# ── Visualise ─────────────────────────────────────────────────\nfig, ax = plt.subplots(2, 3, figsize=(18, 12))\nax[0,0].imshow(msk1,       cmap='gray');    ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop,  cmap='inferno'); ax[0,1].set_title('Probability Map')\nax[0,2].imshow(final_pred, cmap='gray');    ax[0,2].set_title(\n    f'Prediction  dice={dice_full:.3f}')\n\nerr = np.zeros((*msk1.shape,3), dtype=np.uint8)\nerr[(final_pred==1)&(msk1==1)] = [0,255,0]\nerr[(final_pred==1)&(msk1==0)] = [255,0,0]\nerr[(final_pred==0)&(msk1==1)] = [0,0,255]\nax[1,0].imshow(err); ax[1,0].set_title('TP=green  FP=red  FN=blue')\n\nax[1,1].hist(prob_crop[msk1==1].ravel(), bins=50, alpha=0.7,\n             label=f'ink (μ={ink_mean:.2f})',\n             color='orange', density=True)\nax[1,1].hist(prob_crop[msk1==0].ravel(), bins=50, alpha=0.7,\n             label=f'no-ink (μ={noink_mean:.2f})',\n             color='blue', density=True)\nax[1,1].axvline(bt1, color='r', ls='--', label=f'thr={bt1:.2f}')\nax[1,1].set_title('Probability Distribution')\nax[1,1].legend(); ax[1,1].set_xlabel('Probability')\n\nts = np.arange(0.20, 0.85, 0.01); ds = []\nfor t in ts:\n    p = (prob_crop > t).astype(np.float32)\n    ds.append((2*(p*msk1).sum()+1)/(p.sum()+msk1.sum()+1))\nax[1,2].plot(ts, ds)\nax[1,2].axvline(bt1, color='r', ls='--', label=f'best={bt1:.2f}')\nax[1,2].axhline(0.89, color='g', ls=':', label='target 0.89')\nax[1,2].set_title('Dice vs Threshold')\nax[1,2].set_xlabel('Threshold'); ax[1,2].set_ylabel('Dice')\nax[1,2].legend(); ax[1,2].grid(True)\n\nfor a in [ax[0,0],ax[0,1],ax[0,2],ax[1,0]]: a.axis('off')\nplt.suptitle(\n    f'Fragment 1 — Dice={dice_full:.4f}  [3D-Stem + ZAxialAttn + BoundaryLoss]',\n    fontsize=13, y=1.01)\nplt.tight_layout()\nplt.savefig(OUTPUT+'frag1_prediction.png', dpi=100, bbox_inches='tight')\nplt.close()\n\n# ── Save results ──────────────────────────────────────────────\nwith open(OUTPUT+'final_results.txt', 'w') as f:\n    f.write(f'Architecture  : 3D-Stem + ZAxialAttention + Unet++(efficientnet-b1) + SCSE\\n')\n    f.write(f'Channels (z)  : {N_CH}  (variance-ranked)\\n')\n    f.write(f'Z-slices      : {Z_SLICES}\\n')\n    f.write(f'Local norm    : per-patch per-slice (kills domain gap)\\n')\n    f.write(f'Loss S1       : Tversky + Focal + BoundaryBCE\\n')\n    f.write(f'Loss S2       : (1-{SOFT_WEIGHT})·Hard + {SOFT_WEIGHT}·SoftDistill\\n')\n    f.write(f'Data protocol : S1 Frag2/Frag3 | S2 Frag2(soft)/Frag3 | Frag1=test only\\n')\n    f.write(f'Stage 1 best  : {best_s1:.4f}\\n')\n    f.write(f'Stage 2 best  : {best_s2:.4f}\\n')\n    f.write('---\\n')\n    f.write(f'Dice          : {dice_full:.4f}\\n')\n    f.write(f'Threshold     : {bt1:.2f}\\n')\n    f.write(f'Precision     : {prec:.4f}\\n')\n    f.write(f'Recall        : {rec:.4f}\\n')\n    f.write(f'F1            : {f1:.4f}\\n')\n    f.write(f'TP={tp_v} TN={tn} FP={fp} FN={fn}\\n')\n    f.write(f'FP/TP ratio   : {fp/(tp_v+1e-8):.2f}\\n')\n    f.write(f'Calibration   : ink={ink_mean:.3f} bg={noink_mean:.3f} '\n            f'sep={ink_mean-noink_mean:+.3f}\\n')\n\nprint(f'\\nAll outputs → {OUTPUT}')\nprint('  stage1_best.pth | stage2_best.pth')\nprint('  soft_labels_frag2.png | curves.png')\nprint('  frag1_prediction.png  | final_results.txt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\n\"\"\"\nVesuvius Ink Detection v4\n==========================\nBased directly on the second working code (dice=0.42-0.48).\nMinimal targeted fixes only:\n\nWHAT CHANGED vs second code:\n  1. Patches: 5000 per fragment with adaptive collection stride (was 2000, got only 800)\n  2. Balance: 2:1 ink:bg (was 1:0.5 which starved bg)\n  3. Slices: 26-38 (12 slices — your best result was with these exact slices at 0.50)\n  4. Loss: Added Focal term to fix recall=0.93/precision=0.21 (was pure Dice+BCE)\n  5. Backbone: ResNet34 pretrained, weight-adapted correctly for N channels\n  6. NO TTA — plain sliding window inference same as second code\n  7. Visualization: kept exactly from second code + zoomed regions\n\nWHAT DID NOT CHANGE:\n  - PATCH_SIZE=160, STRIDE=32 (your best combo)\n  - Sliding window inference with overlap=0.5\n  - Same augmentation\n  - Same dataset/dataloader structure\n  - Same collate_fn\n  - Same ignore mask logic\n  - Same train/val split approach\n\"\"\"\n\n\n\nimport os, cv2, gc, time, warnings\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nfrom scipy.ndimage import zoom, label, find_objects\nimport pandas as pd\nimport psutil\n\nwarnings.filterwarnings('ignore')\n\n# ============================================================\n# CONFIGURATION  (same as second code except noted)\n# ============================================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3',\n]\n\nDEVICE     = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE = 192        # unchanged — your best\nSTRIDE     = 32         # unchanged — your best\nBATCH_SIZE = 14\nACCUMULATION_STEPS = 2\nEPOCHS     = 100\nLR         = 3e-4\nWEIGHT_DECAY = 1e-4\nSLICE_START = 20        # back to your best-performing slice range\nSLICE_END   = 32        # 12 slices (gave 0.50 in first code)\nN_SLICES    = SLICE_END - SLICE_START\nMAX_PATCHES = 2500      # per fragment (was 2000, but only 800 were produced)\nOUTPUT_DIR  = \"/kaggle/working/\"\nVALIDATION_SPLIT = 0.20\n\nUSE_AMP    = True\nPIN_MEMORY = True\nNUM_WORKERS = 2\nGRAD_CLIP  = 1.0\nSCHEDULER_PATIENCE = 24\nEARLY_STOPPING_PATIENCE = 20\n\nos.makedirs(os.path.join(OUTPUT_DIR, 'full_volume_comparison'), exist_ok=True)\n\nprint(f\"Using device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU: {torch.cuda.get_device_properties(0).name}  \"\n          f\"{torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n# ============================================================\n# IGNORE MASK  (unchanged)\n# ============================================================\ndef generate_ignore_mask(ink_mask):\n    kernel  = np.ones((3, 3), np.uint8)\n    eroded  = cv2.erode(ink_mask.astype(np.uint8),  kernel, iterations=1)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    ignore  = (dilated != eroded).astype(np.uint8)\n    cnts, _ = cv2.findContours(ink_mask.astype(np.uint8),\n                               cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for c in cnts:\n        if cv2.contourArea(c) < 10:\n            cv2.drawContours(ignore, [c], -1, 1, -1)\n    return ignore\n\n# ============================================================\n# MODEL  — ResNet34 UNet with correct weight adaptation\n# FIX: pretrained first conv averaged 3ch→1ch, tiled to N_SLICES\n# Works across SMP versions by finding conv1 dynamically\n# ============================================================\ndef _get_first_conv(encoder):\n    \"\"\"Find first Conv2d regardless of SMP version.\"\"\"\n    if hasattr(encoder, 'conv1') and isinstance(encoder.conv1, nn.Conv2d):\n        return encoder, 'conv1', encoder.conv1\n    if hasattr(encoder, 'layer0') and hasattr(encoder.layer0, 'conv1'):\n        return encoder.layer0, 'conv1', encoder.layer0.conv1\n    for name, mod in encoder.named_modules():\n        if isinstance(mod, nn.Conv2d):\n            parts  = name.split('.')\n            parent = encoder\n            for p in parts[:-1]:\n                parent = getattr(parent, p)\n            return parent, parts[-1], mod\n    raise RuntimeError(\"First Conv2d not found in encoder\")\n\ndef build_model():\n    \"\"\"\n    ResNet34 UNet.\n    Loads ImageNet weights (3ch), averages to 1ch, tiles to N_SLICES.\n    Encoder trained at 10x lower LR than decoder.\n    \"\"\"\n    model = smp.Unet(\n        encoder_name    = \"resnet34\",\n        encoder_weights = \"imagenet\",\n        in_channels     = 3,\n        classes         = 1,\n        activation      = None,\n    )\n    parent, attr, old_conv = _get_first_conv(model.encoder)\n    print(f\"  First conv found at: encoder.{attr}  shape: {old_conv.weight.shape}\")\n\n    # Average 3 RGB channels → 1, then tile to N_SLICES\n    w = old_conv.weight.data.mean(dim=1, keepdim=True)  # [64,1,7,7]\n    w = w.repeat(1, N_SLICES, 1, 1) / N_SLICES          # [64,N,7,7]\n\n    new_conv = nn.Conv2d(\n        N_SLICES, old_conv.out_channels,\n        kernel_size=old_conv.kernel_size,\n        stride=old_conv.stride,\n        padding=old_conv.padding,\n        bias=old_conv.bias is not None,\n    )\n    new_conv.weight = nn.Parameter(w)\n    if old_conv.bias is not None:\n        new_conv.bias = nn.Parameter(old_conv.bias.data.clone())\n    setattr(parent, attr, new_conv)\n    print(f\"  Adapted first conv: 3ch → {N_SLICES}ch (ImageNet filters preserved)\")\n    return model\n\n# ============================================================\n# LOSS  — FIX: added Focal to fix recall=0.93 / precision=0.21\n# Dice+BCE alone cannot control false positives effectively\n# ============================================================\nclass FocalDiceLoss(nn.Module):\n    \"\"\"\n    0.35 * Focal + 0.35 * Dice + 0.30 * BCE\n    Focal gamma=2.0 penalises easy negatives and over-confident FP predictions.\n    All terms ignore uncertain boundary pixels via ignore_mask.\n    \"\"\"\n    def __init__(self, focal_alpha=0.25, focal_gamma=2.0,\n                 w_focal=0.35, w_dice=0.35, w_bce=0.30, smooth=1e-6):\n        super().__init__()\n        self.focal_alpha = focal_alpha\n        self.focal_gamma = focal_gamma\n        self.w_focal = w_focal\n        self.w_dice  = w_dice\n        self.w_bce   = w_bce\n        self.smooth  = smooth\n\n    def forward(self, pred, target, ignore_mask):\n        valid = (1 - ignore_mask).float()\n        probs = torch.sigmoid(pred)\n\n        # --- Focal ---\n        bce_raw = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n        p_t     = torch.exp(-bce_raw)\n        alpha_t = self.focal_alpha * target + (1 - self.focal_alpha) * (1 - target)\n        focal   = alpha_t * (1 - p_t) ** self.focal_gamma * bce_raw\n        l_focal = (focal * valid).sum() / (valid.sum() + self.smooth)\n\n        # --- Dice ---\n        pv = probs * valid;  tv = target * valid\n        l_dice = 1 - (2*(pv*tv).sum() + self.smooth) / \\\n                     (pv.sum() + tv.sum() + self.smooth)\n\n        # --- BCE ---\n        l_bce = (bce_raw * valid).sum() / (valid.sum() + self.smooth)\n\n        return self.w_focal * l_focal + self.w_dice * l_dice + self.w_bce * l_bce\n\n# ============================================================\n# DATA LOADING  (unchanged)\n# ============================================================\ndef load_volume_fast(fragment_path):\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02d}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    volume = np.stack(slices, axis=-1).astype(np.float32)\n    for i in range(volume.shape[-1]):\n        s = volume[:, :, i]\n        volume[:, :, i] = (s - s.mean()) / (s.std() + 1e-6)\n    return volume\n\n# ============================================================\n# PATCH EXTRACTION  — FIX: adaptive stride so we collect enough patches\n# FIX: 2:1 ink:bg balance (was 1:0.5, caused only 800 total patches)\n# ============================================================\ndef extract_patches_fast(volume, mask, ignore_mask,\n                         max_patches=MAX_PATCHES):\n    H, W, _ = volume.shape\n\n    # Adaptive collection stride: aim to visit ~15x max_patches positions\n    # so we have enough candidates to sample from without OOM\n    n_positions = (H // STRIDE) * (W // STRIDE)\n    if n_positions > max_patches * 15:\n        factor = int(np.ceil(np.sqrt(n_positions / (max_patches * 15))))\n        col_stride = STRIDE * factor\n    else:\n        col_stride = STRIDE\n    print(f\"    Fragment {H}x{W}  collection_stride={col_stride}\")\n\n    ink_patches, ink_masks, ink_ignores = [], [], []\n    bg_patches,  bg_masks,  bg_ignores  = [], [], []\n\n    for y in range(0, H - PATCH_SIZE, col_stride):\n        for x in range(0, W - PATCH_SIZE, col_stride):\n            vp = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            mp = mask  [y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            ip = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            cnt = mp.sum()\n            if cnt > 50:\n                ink_patches.append(vp)\n                ink_masks.append(mp)\n                ink_ignores.append(ip)\n            elif cnt == 0:\n                bg_patches.append(vp)\n                bg_masks.append(mp)\n                bg_ignores.append(ip)\n\n    print(f\"    Collected: {len(ink_patches):,} ink  {len(bg_patches):,} bg\")\n\n    # 2:1 ink:bg — use 2/3 of budget for ink, 1/3 for bg\n    rng   = np.random.default_rng(42)\n    n_ink = min(len(ink_patches), int(max_patches * 2 / 3))\n    n_bg  = min(len(bg_patches),  max_patches - n_ink)\n    # top up ink if bg is short\n    n_ink = min(len(ink_patches), max_patches - n_bg)\n\n    ii = rng.choice(len(ink_patches), n_ink, replace=False)\n    bi = rng.choice(len(bg_patches),  n_bg,  replace=False)\n\n    patches = [ink_patches[i] for i in ii] + [bg_patches[i] for i in bi]\n    masks_  = [ink_masks[i]   for i in ii] + [bg_masks[i]   for i in bi]\n    ignores = [ink_ignores[i] for i in ii] + [bg_ignores[i] for i in bi]\n    print(f\"    Sampled: {n_ink} ink  {n_bg} bg  Total: {len(patches)}\")\n    return patches, masks_, ignores\n\n# ============================================================\n# DATASET  (unchanged)\n# ============================================================\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignore_masks, transform=None):\n        self.volumes = volumes\n        self.masks   = masks\n        self.ignore_masks = ignore_masks\n        self.transform = transform\n\n    def __len__(self): return len(self.volumes)\n\n    def __getitem__(self, idx):\n        image  = self.volumes[idx].copy()\n        mask   = self.masks[idx].copy()\n        ignore = self.ignore_masks[idx].copy()\n        if self.transform:\n            t      = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image  = t['image']; mask = t['mask']; ignore = t['ignore_mask']\n        image  = torch.tensor(image).permute(2, 0, 1).float()\n        mask   = torch.tensor(mask).float().unsqueeze(0)\n        ignore = torch.tensor(ignore).float().unsqueeze(0)\n        return image, mask, ignore\n\ndef collate_fn(batch):\n    images  = torch.stack([b[0] for b in batch])\n    masks   = torch.stack([b[1] for b in batch])\n    ignores = torch.stack([b[2] for b in batch])\n    return images, masks, ignores\n\n# ============================================================\n# TRAIN/VAL SPLIT  (unchanged)\n# ============================================================\ndef fast_train_val_split(n_samples, val_ratio=0.20, seed=42):\n    np.random.seed(seed)\n    indices = np.random.permutation(n_samples)\n    split   = int(n_samples * val_ratio)\n    return indices[split:], indices[:split]\n\n# ============================================================\n# FULL VOLUME INFERENCE  (same as second code, no TTA)\n# ============================================================\ndef predict_full_volume(model, volume, device=DEVICE,\n                        batch_size=8, overlap=0.5):\n    model.eval()\n    H, W, C = volume.shape\n    stride   = int(PATCH_SIZE * (1 - overlap))\n\n    pred_map   = np.zeros((H, W), dtype=np.float32)\n    weight_map = np.zeros((H, W), dtype=np.float32)\n\n    ys = list(range(0, H - PATCH_SIZE + 1, stride))\n    xs = list(range(0, W - PATCH_SIZE + 1, stride))\n    if ys[-1] + PATCH_SIZE < H: ys.append(H - PATCH_SIZE)\n    if xs[-1] + PATCH_SIZE < W: xs.append(W - PATCH_SIZE)\n\n    total = len(ys) * len(xs)\n    print(f\"    Processing {total:,} patches ({len(ys)}×{len(xs)})...\")\n\n    for i in tqdm(ys, desc=\"  Inference rows\"):\n        batch_patches, batch_pos = [], []\n        for j in xs:\n            patch = volume[i:i+PATCH_SIZE, j:j+PATCH_SIZE]\n            batch_patches.append(patch)\n            batch_pos.append((i, j))\n\n            if len(batch_patches) == batch_size:\n                t = torch.tensor(np.stack(batch_patches)).permute(0,3,1,2).float().to(device)\n                with torch.no_grad():\n                    with torch.cuda.amp.autocast(enabled=USE_AMP and device=='cuda'):\n                        probs = torch.sigmoid(model(t)).cpu().numpy()\n                for k, (py, px) in enumerate(batch_pos):\n                    pred_map  [py:py+PATCH_SIZE, px:px+PATCH_SIZE] += probs[k, 0]\n                    weight_map[py:py+PATCH_SIZE, px:px+PATCH_SIZE] += 1\n                batch_patches, batch_pos = [], []\n                if device == 'cuda': torch.cuda.empty_cache()\n\n        # leftover\n        if batch_patches:\n            t = torch.tensor(np.stack(batch_patches)).permute(0,3,1,2).float().to(device)\n            with torch.no_grad():\n                with torch.cuda.amp.autocast(enabled=USE_AMP and device=='cuda'):\n                    probs = torch.sigmoid(model(t)).cpu().numpy()\n            for k, (py, px) in enumerate(batch_pos):\n                pred_map  [py:py+PATCH_SIZE, px:px+PATCH_SIZE] += probs[k, 0]\n                weight_map[py:py+PATCH_SIZE, px:px+PATCH_SIZE] += 1\n\n    pred_map = np.divide(pred_map, weight_map, where=weight_map > 0)\n    return pred_map\n\n# ============================================================\n# METRICS  (same as second code)\n# ============================================================\ndef calculate_metrics(ground_truth, prediction, threshold=0.5, ignore_mask=None):\n    binary_pred = (prediction > threshold).astype(np.uint8)\n    if ignore_mask is not None:\n        valid = (1 - ignore_mask).flatten().astype(bool)\n    else:\n        valid = np.ones(ground_truth.size, dtype=bool)\n    pf = binary_pred.flatten()[valid]\n    gf = ground_truth.flatten()[valid]\n    tn, fp, fn, tp = confusion_matrix(gf, pf, labels=[0, 1]).ravel()\n    dice  = 2*tp  / (2*tp + fp + fn + 1e-6)\n    prec  = tp    / (tp + fp + 1e-6)\n    rec   = tp    / (tp + fn + 1e-6)\n    f1    = 2*prec*rec / (prec + rec + 1e-6)\n    iou   = tp    / (tp + fp + fn + 1e-6)\n    spec  = tn    / (tn + fp + 1e-6)\n    acc   = (tp + tn) / (tp + tn + fp + fn)\n    print(\"\\n\" + \"=\"*60)\n    print(\"FINAL TEST RESULTS ON FRAGMENT 1\")\n    print(\"=\"*60)\n    print(f\"  Dice Score  : {dice:.4f}\")\n    print(f\"  F1 Score    : {f1:.4f}\")\n    print(f\"  Precision   : {prec:.4f}\")\n    print(f\"  Recall      : {rec:.4f}\")\n    print(f\"  Specificity : {spec:.4f}\")\n    print(f\"  Accuracy    : {acc:.4f}\")\n    print(f\"  IoU         : {iou:.4f}\")\n    print(\"-\"*60)\n    print(f\"  TP: {tp:,}  TN: {tn:,}  FP: {fp:,}  FN: {fn:,}\")\n    print(\"=\"*60)\n    return dict(dice=dice, f1=f1, prec=prec, rec=rec,\n                spec=spec, acc=acc, iou=iou,\n                tp=int(tp), tn=int(tn), fp=int(fp), fn=int(fn))\n\n# ============================================================\n# VISUALIZATION  (same 6-panel as second code)\n# ============================================================\ndef visualize_full_volume_comparison(volume, ground_truth, prediction,\n                                     save_dir, threshold=0.5, max_size=1500):\n    H, W   = ground_truth.shape\n    scale  = min(1.0, max_size / max(H, W))\n    def ds(arr, order=1): return zoom(arr, (scale, scale), order=order)\n\n    mid_slice = np.mean(volume, axis=2)\n    mid_n     = (mid_slice - mid_slice.min()) / (mid_slice.ptp() + 1e-6)\n    mid_d     = ds(mid_n)\n    gt_d      = ds(ground_truth.astype(float), order=0).astype(np.uint8)\n    pr_d      = ds(prediction)\n    bn_d      = (pr_d > threshold).astype(np.uint8)\n\n    tp_m = (bn_d == 1) & (gt_d == 1)\n    fp_m = (bn_d == 1) & (gt_d == 0)\n    fn_m = (bn_d == 0) & (gt_d == 1)\n\n    overlay = np.zeros((*gt_d.shape, 3), np.float32)\n    overlay[:,:,0] = gt_d.astype(np.float32)   # red   = GT\n    overlay[:,:,1] = bn_d.astype(np.float32)   # green = pred\n\n    error_map = np.zeros((*gt_d.shape, 3), np.float32)\n    error_map[tp_m] = [0, 1, 0]\n    error_map[fp_m] = [1, 0, 0]\n    error_map[fn_m] = [0, 0, 1]\n\n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n\n    axes[0,0].imshow(mid_d, cmap='gray')\n    axes[0,0].set_title('Input (Mean Projection)', fontweight='bold')\n\n    axes[0,1].imshow(gt_d, cmap='gray', vmin=0, vmax=1)\n    axes[0,1].set_title(f'Ground Truth  ({ground_truth.sum():,} ink px)', fontweight='bold')\n\n    im = axes[0,2].imshow(pr_d, cmap='viridis', vmin=0, vmax=1)\n    axes[0,2].set_title('Prediction Probability', fontweight='bold')\n    plt.colorbar(im, ax=axes[0,2], fraction=0.046, pad=0.04)\n\n    axes[1,0].imshow(bn_d, cmap='gray', vmin=0, vmax=1)\n    axes[1,0].set_title(f'Binary Prediction (thresh={threshold:.2f})', fontweight='bold')\n\n    axes[1,1].imshow(np.clip(overlay, 0, 1))\n    axes[1,1].set_title('Overlay  (Red=GT  Green=Pred  Yellow=Both)', fontweight='bold')\n\n    leg = [Rectangle((0,0),1,1, color='green', label=f'TP {tp_m.sum():,}'),\n           Rectangle((0,0),1,1, color='red',   label=f'FP {fp_m.sum():,}'),\n           Rectangle((0,0),1,1, color='blue',  label=f'FN {fn_m.sum():,}')]\n    axes[1,2].imshow(error_map)\n    axes[1,2].legend(handles=leg, loc='upper right', fontsize=9,\n                     bbox_to_anchor=(1.35, 1))\n    axes[1,2].set_title('Error Map  (Green=TP  Red=FP  Blue=FN)', fontweight='bold')\n\n    for ax in axes.flat: ax.axis('off')\n    plt.suptitle('FULL VOLUME ANALYSIS — Fragment 1', fontsize=15, fontweight='bold', y=1.01)\n    plt.tight_layout()\n\n    path = os.path.join(save_dir, 'full_volume_comparison.png')\n    plt.savefig(path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"✓ Saved: {path}\")\n\n    cv2.imwrite(os.path.join(save_dir, 'ground_truth_binary.png'), gt_d * 255)\n    cv2.imwrite(os.path.join(save_dir, 'prediction_binary.png'),   bn_d * 255)\n    cv2.imwrite(os.path.join(save_dir, 'prediction_probability.png'),\n                (pr_d * 255).astype(np.uint8))\n\n\ndef create_zoomed_samples(volume, ground_truth, prediction,\n                          save_dir, threshold=0.5, n_samples=4, size=256):\n    binary_pred = (prediction > threshold).astype(np.uint8)\n    diff = np.abs(binary_pred - ground_truth)\n    labeled, nf = label(diff)\n    if nf == 0:\n        print(\"No error regions found.\"); return\n\n    sizes  = [np.sum(labeled == i) for i in range(1, nf+1)]\n    top    = sorted(range(nf), key=lambda i: sizes[i], reverse=True)\n\n    mid_n  = np.mean(volume, axis=2)\n    mid_n  = (mid_n - mid_n.min()) / (mid_n.ptp() + 1e-6)\n    H, W   = ground_truth.shape\n    regions = []\n    for idx in top:\n        sl = find_objects(labeled == (idx+1))[0]\n        cy = (sl[0].start + sl[0].stop) // 2\n        cx = (sl[1].start + sl[1].stop) // 2\n        y0 = max(0, cy - size//2); y1 = min(H, cy + size//2)\n        x0 = max(0, cx - size//2); x1 = min(W, cx + size//2)\n        if (y1-y0) < 64 or (x1-x0) < 64: continue\n        regions.append((y0,y1,x0,x1))\n        if len(regions) >= n_samples: break\n\n    if not regions: return\n    fig, axes = plt.subplots(len(regions), 4, figsize=(16, 4*len(regions)))\n    if len(regions) == 1: axes = axes[np.newaxis]\n\n    for ri, (y0,y1,x0,x1) in enumerate(regions):\n        inp = mid_n[y0:y1, x0:x1]\n        gt  = ground_truth[y0:y1, x0:x1]\n        pr  = prediction[y0:y1, x0:x1]\n        bn  = (pr > threshold).astype(np.uint8)\n        tp  = (bn==1)&(gt==1); fp = (bn==1)&(gt==0); fn = (bn==0)&(gt==1)\n        em  = np.zeros((*inp.shape,3), np.float32)\n        em[tp]=[0,1,0]; em[fp]=[1,0,0]; em[fn]=[0,0,1]\n\n        axes[ri,0].imshow(inp,  cmap='gray');\n        axes[ri,0].set_title(f'Region {ri+1} – Input')\n        axes[ri,1].imshow(gt,   cmap='gray', vmin=0, vmax=1)\n        axes[ri,1].set_title('Ground Truth')\n        axes[ri,2].imshow(bn,   cmap='gray', vmin=0, vmax=1)\n        axes[ri,2].set_title('Prediction')\n        axes[ri,3].imshow(em)\n        axes[ri,3].set_title(f'TP:{tp.sum()}  FP:{fp.sum()}  FN:{fn.sum()}')\n        for ax in axes[ri]: ax.axis('off')\n\n    plt.suptitle('Zoomed Error Regions', fontsize=13, fontweight='bold')\n    plt.tight_layout()\n    path = os.path.join(save_dir, 'zoomed_error_regions.png')\n    plt.savefig(path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"✓ Saved: {path}\")\n\n# ============================================================\n# MEMORY UTIL  (unchanged)\n# ============================================================\ndef print_memory_usage():\n    if DEVICE == 'cuda':\n        a = torch.cuda.memory_allocated()/1e9\n        c = torch.cuda.memory_reserved()/1e9\n        print(f\"    GPU: alloc={a:.2f}GB  cache={c:.2f}GB\")\n    print(f\"    CPU: {psutil.Process().memory_info().rss/1e9:.2f}GB\")\n\n# ============================================================\n# MAIN\n# ============================================================\ndef main():\n    print(\"=\"*60)\n    print(\"VESUVIUS INK DETECTION v4\")\n    print(\"=\"*60)\n    t0 = time.time()\n\n    # ── 1. Load data ────────────────────────────────────────────\n    print(\"\\n1. Loading training data...\")\n    all_patches, all_masks, all_ignores = [], [], []\n\n    for path in train_paths:\n        name = os.path.basename(path)\n        print(f\"\\n  Fragment {name}\")\n        vol  = load_volume_fast(path)\n        mask = cv2.imread(os.path.join(path, \"inklabels.png\"), 0)\n        mask = (mask > 0).astype(np.uint8)\n        ign  = generate_ignore_mask(mask)\n        print(f\"    Volume: {vol.shape}  ink px: {mask.sum():,}\")\n\n        p, m, i = extract_patches_fast(vol, mask, ign)\n        all_patches.extend(p); all_masks.extend(m); all_ignores.extend(i)\n        print_memory_usage()\n        del vol, mask, ign, p, m, i; gc.collect()\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n\n    print(f\"\\nTotal patches: {len(all_patches)}\")\n\n    # ── 2. Train/val split ───────────────────────────────────────\n    print(\"\\n2. Train/val split...\")\n    n = len(all_patches)\n    tr_idx, va_idx = fast_train_val_split(n, VALIDATION_SPLIT)\n    tr_patches = [all_patches[i] for i in tr_idx]\n    tr_masks   = [all_masks[i]   for i in tr_idx]\n    tr_ignores = [all_ignores[i] for i in tr_idx]\n    va_patches = [all_patches[i] for i in va_idx]\n    va_masks   = [all_masks[i]   for i in va_idx]\n    va_ignores = [all_ignores[i] for i in va_idx]\n    print(f\"  Train: {len(tr_patches)}  Val: {len(va_patches)}\")\n    del all_patches, all_masks, all_ignores; gc.collect()\n\n    # ── 3. Augmentation + DataLoaders ───────────────────────────\n    print(\"\\n3. Building dataloaders...\")\n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.1,\n                                   contrast_limit=0.1, p=0.3),\n        A.GaussNoise(var_limit=(0, 0.01), p=0.2),\n    ], additional_targets={'ignore_mask': 'mask'})\n\n    tr_ds = VesuviusDataset(tr_patches, tr_masks, tr_ignores, train_transform)\n    va_ds = VesuviusDataset(va_patches, va_masks, va_ignores)\n\n    tr_dl = DataLoader(tr_ds, batch_size=BATCH_SIZE, shuffle=True,\n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                       drop_last=True, collate_fn=collate_fn)\n    va_dl = DataLoader(va_ds, batch_size=BATCH_SIZE, shuffle=False,\n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                       collate_fn=collate_fn)\n    print(f\"  Train batches: {len(tr_dl)}  Val batches: {len(va_dl)}\")\n\n    # ── 4. Model ─────────────────────────────────────────────────\n    print(\"\\n4. Building model...\")\n    model = build_model().to(DEVICE)\n    n_all = sum(p.numel() for p in model.parameters())\n    n_tr  = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Params: {n_all:,}  trainable: {n_tr:,}\")\n\n    # ── 5. Loss / optimizer / scheduler ─────────────────────────\n    criterion = FocalDiceLoss()\n\n    # Encoder at 10x lower LR, decoder at full LR\n    enc_params = [p for n,p in model.named_parameters() if 'encoder' in n]\n    dec_params = [p for n,p in model.named_parameters() if 'encoder' not in n]\n    optimizer  = optim.AdamW([\n        {'params': enc_params, 'lr': LR / 10},\n        {'params': dec_params, 'lr': LR},\n    ], weight_decay=WEIGHT_DECAY)\n\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='max', factor=0.5,\n        patience=SCHEDULER_PATIENCE, verbose=True)\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n\n    # ── 6. Training ──────────────────────────────────────────────\n    print(\"\\n5. Training...\")\n    print(\"=\"*60)\n    best_val_dice   = 0.0\n    patience_counter = 0\n    train_losses    = []\n    val_dice_scores = []\n\n    for epoch in range(EPOCHS):\n        ep_t = time.time()\n\n        # Train\n        model.train()\n        tr_loss, tr_steps = 0.0, 0\n        optimizer.zero_grad()\n        pbar = tqdm(tr_dl, desc=f'Epoch {epoch+1}/{EPOCHS} [Train]')\n        for batch_idx, (imgs, msks, igns) in enumerate(pbar):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            igns = igns.to(DEVICE, non_blocking=True)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                out  = model(imgs)\n                loss = criterion(out, msks, igns) / ACCUMULATION_STEPS\n            scaler.scale(loss).backward()\n            if (batch_idx + 1) % ACCUMULATION_STEPS == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(optimizer); scaler.update()\n                optimizer.zero_grad()\n            tr_loss  += loss.item() * ACCUMULATION_STEPS\n            tr_steps += 1\n            pbar.set_postfix(loss=f'{loss.item()*ACCUMULATION_STEPS:.4f}')\n            if batch_idx % 50 == 49 and DEVICE == 'cuda':\n                torch.cuda.empty_cache()\n\n        avg_loss = tr_loss / tr_steps\n        train_losses.append(avg_loss)\n\n        # Validate\n        model.eval()\n        val_dice, val_steps = 0.0, 0\n        with torch.no_grad():\n            for imgs, msks, igns in tqdm(va_dl,\n                    desc=f'Epoch {epoch+1}/{EPOCHS} [Val]', leave=False):\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                msks = msks.to(DEVICE, non_blocking=True)\n                igns = igns.to(DEVICE, non_blocking=True)\n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    out = model(imgs)\n                prds  = torch.sigmoid(out) > 0.5\n                valid = 1 - igns\n                inter = ((prds * msks) * valid).sum()\n                union = ((prds * valid).sum() + (msks * valid).sum())\n                val_dice  += (2 * inter / (union + 1e-6)).item()\n                val_steps += 1\n\n        avg_dice = val_dice / val_steps\n        val_dice_scores.append(avg_dice)\n        scheduler.step(avg_dice)\n\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}  [{time.time()-ep_t:.0f}s]\")\n        print(f\"  Train Loss: {avg_loss:.4f}   Val Dice: {avg_dice:.4f}\")\n        print(f\"  LR enc: {optimizer.param_groups[0]['lr']:.2e}\"\n              f\"  dec: {optimizer.param_groups[1]['lr']:.2e}\")\n        print_memory_usage()\n\n        if avg_dice > best_val_dice:\n            best_val_dice = avg_dice\n            torch.save({'epoch': epoch,\n                        'model_state_dict': model.state_dict(),\n                        'val_dice': avg_dice},\n                       os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_counter = 0\n            print(f\"  ✓ Best model saved  Dice={best_val_dice:.4f}\")\n        else:\n            patience_counter += 1\n            if patience_counter >= EARLY_STOPPING_PATIENCE:\n                print(f\"\\nEarly stopping at epoch {epoch+1}\")\n                break\n        print(\"-\"*60)\n\n    print(f\"\\nBest Val Dice: {best_val_dice:.4f}\")\n\n    # ── 7. Full volume inference on Fragment 1 ───────────────────\n    print(\"\\n6. Loading Fragment 1...\")\n    ckpt = torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth'))\n    model.load_state_dict(ckpt['model_state_dict'])\n    model.eval()\n\n    test_vol  = load_volume_fast(fragment1_path)\n    test_mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ign  = generate_ignore_mask(test_mask)\n    print(f\"  Volume: {test_vol.shape}  ink px: {test_mask.sum():,}\")\n\n    print(\"\\n7. Running inference (no TTA)...\")\n    full_pred = predict_full_volume(model, test_vol, device=DEVICE,\n                                    batch_size=8, overlap=0.5)\n\n    # ── 8. Metrics ───────────────────────────────────────────────\n    print(\"\\n8. Metrics...\")\n    metrics = calculate_metrics(test_mask, full_pred, threshold=0.5,\n                                ignore_mask=test_ign)\n\n    # ── 9. Visualizations ────────────────────────────────────────\n    print(\"\\n9. Visualizing...\")\n    viz_dir = os.path.join(OUTPUT_DIR, 'full_volume_comparison')\n    visualize_full_volume_comparison(\n        test_vol, test_mask, full_pred, viz_dir, threshold=0.5)\n    create_zoomed_samples(\n        test_vol, test_mask, full_pred, viz_dir,\n        threshold=0.5, n_samples=4)\n\n    # Training curves\n    fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n    axes[0].plot(train_losses); axes[0].set_title('Training Loss')\n    axes[0].set_xlabel('Epoch'); axes[0].grid(True)\n    axes[1].plot(val_dice_scores, color='orange')\n    axes[1].axhline(best_val_dice, color='green', linestyle='--',\n                    label=f'Best {best_val_dice:.4f}')\n    axes[1].set_title('Validation Dice'); axes[1].set_xlabel('Epoch')\n    axes[1].legend(); axes[1].grid(True)\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'), dpi=150)\n    plt.show()\n\n    # Save results\n    total_t = time.time() - t0\n    h, mins = int(total_t//3600), int((total_t%3600)//60)\n    print(f\"\\nTotal time: {h}h {mins}m\")\n\n    pd.DataFrame([metrics]).to_csv(\n        os.path.join(viz_dir, 'metrics.csv'), index=False)\n    with open(os.path.join(viz_dir, 'results.txt'), 'w') as f:\n        f.write(f\"Total time    : {h}h {mins}m\\n\")\n        f.write(f\"Best val Dice : {best_val_dice:.4f}\\n\")\n        for k, v in metrics.items():\n            f.write(f\"{k}: {v}\\n\")\n\n    return metrics['dice']\n\n\nif __name__ == \"__main__\":\n    try:\n        score = main()\n        print(f\"\\n✅ Final Test Dice: {score:.4f}\")\n    except Exception as e:\n        import traceback\n        print(f\"\\n❌ Error: {e}\")\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nVesuvius Ink Detection — v3\n===========================\nTarget: Dice > 0.80 with 2D UNet\n\nKey fixes vs previous versions:\n  1. CORRECT pretrained weight transfer (3ch → Nch by channel averaging, not random init)\n  2. Lovász-Softmax loss — directly optimises the Dice/IoU metric\n  3. Focal loss — fixes recall=0.93 / precision=0.21 imbalance\n  4. PATCH_SIZE=160, STRIDE=32 — best from your own experiments\n  5. 1:1 ink:bg patch balance (not 1:0.5)\n  6. Slices 22–38 (16 slices, sweet spot)\n  7. Fragment-aware train/val split\n  8. Gaussian-weighted sliding window inference\n  9. 8-fold TTA\n 10. Threshold sweep on validation set\n 11. Full-fragment visualization (6-panel)\n 12. Zoomed error region analysis\n\"\"\"\n\n# ── Install ───────────────────────────────────────────────────────────────────\n#import subprocess, sys\n#subprocess.run([sys.executable, \"-m\", \"pip\", \"install\",\n#                \"segmentation-models-pytorch==0.3.3\", \"-q\"])\n\n# ── Imports ───────────────────────────────────────────────────────────────────\nimport os, cv2, gc, time, warnings\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nfrom scipy.ndimage import label as nd_label, find_objects, zoom\nimport pandas as pd\nimport psutil\n\nwarnings.filterwarnings('ignore')\n\n# ═════════════════════════════════════════════════════════════════════════════\n# CONFIGURATION  — tuned from your experiments\n# ═════════════════════════════════════════════════════════════════════════════\nFRAGMENT1_PATH = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nTRAIN_PATHS = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3',\n]\n\nDEVICE      = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE  = 192       # best patch size from your experiments\nSTRIDE      = 32        # best stride from experiments (PATCH_SIZE=160,STRIDE=32 → 0.48)\nBATCH_SIZE  = 12\nACCUM_STEPS = 2\nEPOCHS      = 50\nLR          = 2e-4\nLR_ENCODER  = 2e-5     # encoder trained at 10× lower LR\nWEIGHT_DECAY= 1e-4\nSLICE_START = 18        # 16 slices: 22–38\nSLICE_END   = 34\nN_SLICES    = SLICE_END - SLICE_START\nMAX_PATCHES = 600      # per fragment\nVAL_SPLIT   = 0.20\nUSE_AMP     = True\nPIN_MEMORY  = True\nNUM_WORKERS = 2\nGRAD_CLIP   = 1.0\nEARLY_STOP  = 20\nOUTPUT_DIR  = \"/kaggle/working/\"\n\nos.makedirs(os.path.join(OUTPUT_DIR, 'viz'), exist_ok=True)\nprint(f\"Device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU: {torch.cuda.get_device_properties(0).name}  \"\n          f\"{torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# LOVÁSZ-SOFTMAX LOSS\n# Pure-PyTorch implementation — no extra dependency\n# Directly optimises the Jaccard/IoU metric (≈ Dice metric)\n# ═════════════════════════════════════════════════════════════════════════════\ndef lovász_grad(gt_sorted):\n    \"\"\"Computes gradient of the Lovász extension.\"\"\"\n    p = len(gt_sorted)\n    gts = gt_sorted.sum()\n    intersection = gts - gt_sorted.float().cumsum(0)\n    union = gts + (1 - gt_sorted).float().cumsum(0)\n    jaccard = 1. - intersection / union\n    if p > 1:\n        jaccard[1:p] = jaccard[1:p] - jaccard[0:-1]\n    return jaccard\n\n\ndef lovász_hinge_flat(logits, labels):\n    \"\"\"Binary Lovász hinge loss on flat vectors.\"\"\"\n    if len(labels) == 0:\n        return logits.sum() * 0.\n    signs = 2. * labels.float() - 1.\n    errors = 1. - logits * signs\n    errors_sorted, perm = torch.sort(errors, dim=0, descending=True)\n    gt_sorted = labels[perm]\n    grad = lovász_grad(gt_sorted)\n    loss = torch.dot(F.relu(errors_sorted), grad)\n    return loss\n\n\ndef lovász_hinge(logits, labels, valid_mask, per_image=True):\n    \"\"\"\n    Binary Lovász hinge.\n    logits: [B, 1, H, W]   labels: [B, 1, H, W]   valid_mask: [B, 1, H, W]\n    \"\"\"\n    if per_image:\n        losses = []\n        for log, lab, vm in zip(logits, labels, valid_mask):\n            log_flat = log.view(-1)[vm.view(-1).bool()]\n            lab_flat = lab.view(-1)[vm.view(-1).bool()]\n            if lab_flat.numel() == 0:\n                continue\n            losses.append(lovász_hinge_flat(log_flat, lab_flat))\n        return torch.stack(losses).mean() if losses else logits.sum() * 0.\n    else:\n        vm = valid_mask.view(-1).bool()\n        return lovász_hinge_flat(logits.view(-1)[vm], labels.view(-1)[vm])\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# COMBINED LOSS  (Focal + Dice + Lovász)\n# ═════════════════════════════════════════════════════════════════════════════\nclass FocalDiceLovaszLoss(nn.Module):\n    \"\"\"\n    Focal BCE  — punishes confident wrong predictions, fixes precision/recall balance\n    Dice       — global overlap, smooth gradients\n    Lovász     — directly optimises IoU/Jaccard ≈ Dice metric\n    All terms apply only to non-ignored pixels.\n    \"\"\"\n    def __init__(self, focal_alpha=0.25, focal_gamma=2.5,\n                 w_focal=0.3, w_dice=0.4, w_lovász=0.3, smooth=1e-6):\n        super().__init__()\n        self.focal_alpha = focal_alpha\n        self.focal_gamma = focal_gamma\n        self.w_focal   = w_focal\n        self.w_dice    = w_dice\n        self.w_lovász  = w_lovász\n        self.smooth    = smooth\n\n    def focal(self, logits, target, valid):\n        bce     = F.binary_cross_entropy_with_logits(logits, target, reduction='none')\n        p_t     = torch.exp(-bce)\n        alpha_t = self.focal_alpha * target + (1 - self.focal_alpha) * (1 - target)\n        fl      = alpha_t * (1 - p_t) ** self.focal_gamma * bce\n        denom   = valid.sum() + self.smooth\n        return (fl * valid).sum() / denom\n\n    def dice(self, probs, target, valid):\n        p = probs * valid;  t = target * valid\n        inter = (p * t).sum()\n        union = p.sum() + t.sum()\n        return 1 - (2 * inter + self.smooth) / (union + self.smooth)\n\n    def forward(self, logits, target, ignore_mask):\n        valid = (1 - ignore_mask).float()\n        probs = torch.sigmoid(logits)\n\n        l_focal  = self.focal(logits, target, valid)\n        l_dice   = self.dice(probs, target, valid)\n        l_lovász = lovász_hinge(logits, target, valid)\n\n        return self.w_focal * l_focal + self.w_dice * l_dice + self.w_lovász * l_lovász\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# MODEL — ResNet34 UNet with CORRECT pretrained weight transfer\n# ═════════════════════════════════════════════════════════════════════════════\ndef _find_first_conv(encoder):\n    \"\"\"\n    SMP's ResNet encoder structure changed across versions.\n    This function finds the first Conv2d regardless of version:\n      - SMP >= 0.2.1  : encoder.conv1  (direct attribute)\n      - SMP 0.1.x     : encoder.layer0.conv1\n      - Fallback      : walk named_modules until we find the first Conv2d\n    Returns (parent_module, attribute_name, conv_module).\n    \"\"\"\n    # Most common in SMP 0.2+ / 0.3+\n    if hasattr(encoder, 'conv1') and isinstance(encoder.conv1, nn.Conv2d):\n        return encoder, 'conv1', encoder.conv1\n\n    # Older SMP builds\n    if hasattr(encoder, 'layer0') and hasattr(encoder.layer0, 'conv1'):\n        return encoder.layer0, 'conv1', encoder.layer0.conv1\n\n    # Generic fallback: first Conv2d in module tree\n    for name, mod in encoder.named_modules():\n        if isinstance(mod, nn.Conv2d):\n            # name might be 'features.0' — resolve parent + attr\n            parts  = name.split('.')\n            parent = encoder\n            for p in parts[:-1]:\n                parent = getattr(parent, p)\n            return parent, parts[-1], mod\n\n    raise RuntimeError(\"Cannot find first Conv2d in encoder. \"\n                       \"Print model.encoder to inspect its structure.\")\n\n\ndef build_model_with_pretrained(n_slices=N_SLICES):\n    \"\"\"\n    Build SMP UNet with ResNet34 encoder, then correctly adapt the first\n    convolution from 3 channels (ImageNet) to n_slices channels.\n\n    Strategy:\n      1. Load model with in_channels=3, encoder_weights='imagenet'\n         → first conv is [64, 3, 7, 7], initialised from ImageNet\n      2. Average the 3 input channels → [64, 1, 7, 7]\n      3. Repeat n_slices times        → [64, n_slices, 7, 7]\n      4. Replace the conv and continue training\n\n    This keeps the spatial filter patterns learned on ImageNet intact\n    while adapting to n_slices input channels.\n    Works across SMP 0.2.x and 0.3.x.\n    \"\"\"\n    # Step 1: load with imagenet weights (3 channels)\n    model = smp.Unet(\n        encoder_name    = \"resnet34\",\n        encoder_weights = \"imagenet\",\n        in_channels     = 3,\n        classes         = 1,\n        activation      = None,\n        decoder_attention_type = \"scse\",\n    )\n\n    # Step 2-4: locate and replace first conv\n    parent, attr, old_conv = _find_first_conv(model.encoder)\n    print(f\"  Found first conv at encoder.{attr}  shape: {old_conv.weight.shape}\")\n\n    # Average 3 RGB channels → 1 channel, then tile to n_slices\n    new_weight = old_conv.weight.data.mean(dim=1, keepdim=True)   # [64, 1, 7, 7]\n    new_weight = new_weight.repeat(1, n_slices, 1, 1)              # [64, n, 7, 7]\n    new_weight = new_weight / n_slices                             # scale magnitude\n\n    new_conv = nn.Conv2d(\n        in_channels  = n_slices,\n        out_channels = old_conv.out_channels,\n        kernel_size  = old_conv.kernel_size,\n        stride       = old_conv.stride,\n        padding      = old_conv.padding,\n        bias         = old_conv.bias is not None,\n    )\n    new_conv.weight = nn.Parameter(new_weight)\n    if old_conv.bias is not None:\n        new_conv.bias = nn.Parameter(old_conv.bias.data.clone())\n\n    setattr(parent, attr, new_conv)\n    print(f\"  Replaced first conv: 3 → {n_slices} channels (ImageNet weights preserved)\")\n    return model\n\n\ndef get_param_groups(model):\n    \"\"\"\n    Separate parameter groups:\n      encoder → LR_ENCODER (10× lower)\n      decoder + head → LR\n    \"\"\"\n    encoder_params, other_params = [], []\n    for name, param in model.named_parameters():\n        if 'encoder' in name:\n            encoder_params.append(param)\n        else:\n            other_params.append(param)\n    return [\n        {'params': encoder_params, 'lr': LR_ENCODER},\n        {'params': other_params,   'lr': LR},\n    ]\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# DATA LOADING\n# ═════════════════════════════════════════════════════════════════════════════\ndef load_volume(fragment_path, s_start=SLICE_START, s_end=SLICE_END):\n    vol_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices  = []\n    for i in range(s_start, s_end):\n        p = os.path.join(vol_dir, f\"{i:02d}.tif\")\n        if os.path.exists(p):\n            img = tifffile.imread(p).astype(np.float32)\n            slices.append(img)\n    volume = np.stack(slices, axis=-1)   # [H, W, C]\n    for c in range(volume.shape[-1]):\n        s = volume[:, :, c]\n        volume[:, :, c] = (s - s.mean()) / (s.std() + 1e-6)\n    return volume\n\n\ndef generate_ignore_mask(ink_mask):\n    kernel  = np.ones((3, 3), np.uint8)\n    eroded  = cv2.erode(ink_mask.astype(np.uint8),  kernel, iterations=1)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    ignore  = (dilated != eroded).astype(np.uint8)\n    cnts, _ = cv2.findContours(ink_mask.astype(np.uint8),\n                               cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for c in cnts:\n        if cv2.contourArea(c) < 10:\n            cv2.drawContours(ignore, [c], -1, 1, -1)\n    return ignore\n\n\ndef extract_patches(volume, mask, ignore_mask, max_patches=MAX_PATCHES):\n    \"\"\"\n    Patch extraction with 2:1 ink-to-background ratio.\n    - Collects ALL qualifying patches first, then samples down to max_patches.\n    - Ink threshold: >=20 pixels in a 160x160 patch (~0.08% coverage).\n    - 2:1 ink:bg ratio to keep model focused on positive class.\n    \"\"\"\n    H, W, _ = volume.shape\n    ink_v, ink_m, ink_i = [], [], []\n    bg_v,  bg_m,  bg_i  = [], [], []\n\n    # Use a coarser stride for collection on very large fragments to keep memory sane\n    # For Fragment 2 (14830x9506): default STRIDE=32 → ~1.3M candidate patches — too many to collect\n    # So we use a collection stride that gives ~10x more candidates than max_patches\n    n_y = (H - PATCH_SIZE) // STRIDE\n    n_x = (W - PATCH_SIZE) // STRIDE\n    total_possible = n_y * n_x\n    collection_factor = max(1, int(np.sqrt(total_possible / (max_patches * 10))))\n    col_stride = STRIDE * collection_factor\n    print(f\"    Grid: {n_y}×{n_x}={total_possible:,} possible  \"\n          f\"collection_stride={col_stride}\")\n\n    for y in range(0, H - PATCH_SIZE, col_stride):\n        for x in range(0, W - PATCH_SIZE, col_stride):\n            vp = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            mp = mask  [y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            ip = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            cnt = mp.sum()\n            if cnt >= 20:\n                ink_v.append(vp); ink_m.append(mp); ink_i.append(ip)\n            elif cnt == 0:\n                bg_v.append(vp);  bg_m.append(mp);  bg_i.append(ip)\n\n    print(f\"    Collected: {len(ink_v):,} ink  {len(bg_v):,} bg\")\n\n    rng = np.random.default_rng(42)\n    # 2:1 ink:bg — up to max_patches total\n    n_ink = min(len(ink_v), (max_patches * 2) // 3)\n    n_bg  = min(len(bg_v),  max_patches - n_ink)\n    # If bg scarce, top up with more ink\n    n_ink = min(len(ink_v), max_patches - n_bg)\n\n    ii = rng.choice(len(ink_v), n_ink, replace=False)\n    bi = rng.choice(len(bg_v),  n_bg,  replace=False)\n\n    patches = [ink_v[i] for i in ii] + [bg_v[i] for i in bi]\n    masks   = [ink_m[i] for i in ii] + [bg_m[i] for i in bi]\n    ignores = [ink_i[i] for i in ii] + [bg_i[i] for i in bi]\n    print(f\"    Sampled: {n_ink} ink  {n_bg} bg  Total: {len(patches)}\")\n    return patches, masks, ignores\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# DATASET\n# ═════════════════════════════════════════════════════════════════════════════\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignores, transform=None):\n        self.volumes = volumes; self.masks = masks\n        self.ignores = ignores; self.transform = transform\n\n    def __len__(self): return len(self.volumes)\n\n    def __getitem__(self, idx):\n        img = self.volumes[idx].copy()\n        msk = self.masks  [idx].copy()\n        ign = self.ignores[idx].copy()\n        if self.transform:\n            t   = self.transform(image=img, mask=msk, ignore_mask=ign)\n            img = t['image']; msk = t['mask']; ign = t['ignore_mask']\n        img = torch.tensor(img).permute(2, 0, 1).float()\n        msk = torch.tensor(msk).float().unsqueeze(0)\n        ign = torch.tensor(ign).float().unsqueeze(0)\n        return img, msk, ign\n\n\ndef collate_fn(batch):\n    return (torch.stack([b[0] for b in batch]),\n            torch.stack([b[1] for b in batch]),\n            torch.stack([b[2] for b in batch]))\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# TRAIN / VAL SPLIT  — fragment-aware spatial split\n# Each fragment gets its patches split independently so val set\n# always contains patches from both training fragments.\n# ═════════════════════════════════════════════════════════════════════════════\ndef fragment_aware_split(all_patches, all_masks, all_ignores,\n                         fragment_sizes, val_ratio=VAL_SPLIT, seed=42):\n    rng = np.random.default_rng(seed)\n    tr_v, tr_m, tr_i = [], [], []\n    va_v, va_m, va_i = [], [], []\n    offset = 0\n    for sz in fragment_sizes:\n        idx = rng.permutation(sz)\n        split = int(sz * val_ratio)\n        for i in idx[split:]:\n            tr_v.append(all_patches[offset+i])\n            tr_m.append(all_masks  [offset+i])\n            tr_i.append(all_ignores[offset+i])\n        for i in idx[:split]:\n            va_v.append(all_patches[offset+i])\n            va_m.append(all_masks  [offset+i])\n            va_i.append(all_ignores[offset+i])\n        offset += sz\n    return tr_v, tr_m, tr_i, va_v, va_m, va_i\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# AUGMENTATION\n# ═════════════════════════════════════════════════════════════════════════════\ndef get_train_transform():\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Transpose(p=0.3),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.4),\n        A.GaussNoise(var_limit=(0, 0.02), p=0.3),\n        A.ElasticTransform(alpha=30, sigma=5, alpha_affine=5, p=0.25),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.2),\n        A.CoarseDropout(max_holes=8, max_height=20, max_width=20,\n                        fill_value=0, p=0.2),\n    ], additional_targets={'ignore_mask': 'mask'})\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# THRESHOLD TUNING\n# ═════════════════════════════════════════════════════════════════════════════\ndef tune_threshold(model, loader, device, thresholds=None):\n    if thresholds is None:\n        thresholds = np.arange(0.25, 0.75, 0.025)\n    model.eval()\n    probs_list, masks_list = [], []\n    with torch.no_grad():\n        for imgs, msks, igns in loader:\n            imgs = imgs.to(device)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                out = torch.sigmoid(model(imgs)).cpu().numpy()\n            valid = (1 - igns).numpy()\n            probs_list.append(out   * valid)\n            masks_list.append(msks.numpy() * valid)\n    probs = np.concatenate(probs_list).flatten()\n    masks = np.concatenate(masks_list).flatten()\n    best_t, best_dice = 0.5, 0.0\n    for t in thresholds:\n        pred  = (probs > t).astype(np.float32)\n        inter = (pred * masks).sum()\n        union = pred.sum() + masks.sum()\n        dice  = 2 * inter / (union + 1e-6)\n        if dice > best_dice:\n            best_dice = dice; best_t = t\n    print(f\"  Best threshold: {best_t:.3f}  (val Dice {best_dice:.4f})\")\n    return float(best_t)\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# TTA — 8 fold (flips × rotations)\n# ═════════════════════════════════════════════════════════════════════════════\n_AUGS = [\n    (lambda t: t,                            lambda t: t),\n    (lambda t: torch.flip(t, [2]),           lambda t: torch.flip(t, [2])),\n    (lambda t: torch.flip(t, [3]),           lambda t: torch.flip(t, [3])),\n    (lambda t: torch.flip(t, [2, 3]),        lambda t: torch.flip(t, [2, 3])),\n    (lambda t: torch.rot90(t, 1, [2, 3]),    lambda t: torch.rot90(t, -1, [2, 3])),\n    (lambda t: torch.rot90(t, 2, [2, 3]),    lambda t: torch.rot90(t, -2, [2, 3])),\n    (lambda t: torch.rot90(t, 3, [2, 3]),    lambda t: torch.rot90(t, -3, [2, 3])),\n    (lambda t: torch.flip(torch.rot90(t, 1, [2, 3]), [2]),\n     lambda t: torch.flip(torch.rot90(t, -1, [2, 3]), [2])),\n]\n\ndef tta_predict(model, patch_tensor, device):\n    \"\"\"patch_tensor: [1, C, H, W] CPU. Returns [1, 1, H, W] prob map CPU.\"\"\"\n    preds = []\n    model.eval()\n    with torch.no_grad():\n        for aug, deaug in _AUGS:\n            x = aug(patch_tensor).to(device)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                p = torch.sigmoid(model(x)).cpu()\n            preds.append(deaug(p))\n    return torch.stack(preds).mean(0)\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# GAUSSIAN-WEIGHTED SLIDING WINDOW INFERENCE\n# ═════════════════════════════════════════════════════════════════════════════\ndef _gauss_kernel(size):\n    sigma  = size / 4.0\n    coords = np.arange(size) - size / 2.0\n    gx, gy = np.meshgrid(coords, coords)\n    g      = np.exp(-(gx**2 + gy**2) / (2 * sigma**2)).astype(np.float32)\n    return g\n\n\ndef sliding_window_inference(model, volume, stride=None,\n                             use_tta=True, device=DEVICE):\n    \"\"\"\n    Full-fragment sliding window with Gaussian blend and optional TTA.\n    Returns probability map [H, W] in [0,1].\n    \"\"\"\n    if stride is None:\n        stride = PATCH_SIZE // 4   # 75% overlap — smooth but slow; reduce for speed\n    H, W, C = volume.shape\n    prob_map   = np.zeros((H, W), dtype=np.float32)\n    weight_map = np.zeros((H, W), dtype=np.float32)\n    gauss      = _gauss_kernel(PATCH_SIZE)\n    model.eval()\n\n    ys = list(range(0, H - PATCH_SIZE, stride)) + [H - PATCH_SIZE]\n    xs = list(range(0, W - PATCH_SIZE, stride)) + [W - PATCH_SIZE]\n    total = len(ys) * len(xs)\n\n    with tqdm(total=total, desc=\"  Sliding window inference\") as pbar:\n        for y in ys:\n            for x in xs:\n                patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]  # [ps, ps, C]\n                t     = torch.tensor(patch).permute(2, 0, 1).float().unsqueeze(0)\n                if use_tta:\n                    prob = tta_predict(model, t, device).squeeze().numpy()\n                else:\n                    t = t.to(device)\n                    with torch.no_grad():\n                        with torch.cuda.amp.autocast(enabled=USE_AMP):\n                            prob = torch.sigmoid(model(t)).cpu().squeeze().numpy()\n                prob_map  [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += prob  * gauss\n                weight_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] +=        gauss\n                pbar.update(1)\n\n    prob_map /= (weight_map + 1e-6)\n    return prob_map\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# METRICS\n# ═════════════════════════════════════════════════════════════════════════════\ndef compute_metrics(gt, pred_prob, threshold, ignore_mask=None):\n    pred = (pred_prob > threshold).astype(np.uint8)\n    if ignore_mask is not None:\n        valid = (1 - ignore_mask).flatten().astype(bool)\n    else:\n        valid = np.ones(gt.size, dtype=bool)\n    pf = pred.flatten()[valid]; gf = gt.flatten()[valid]\n    tn, fp, fn, tp = confusion_matrix(gf, pf, labels=[0, 1]).ravel()\n    dice   = 2*tp / (2*tp + fp + fn + 1e-6)\n    prec   = tp   / (tp + fp + 1e-6)\n    rec    = tp   / (tp + fn + 1e-6)\n    f1     = 2*prec*rec / (prec + rec + 1e-6)\n    iou    = tp   / (tp + fp + fn + 1e-6)\n    spec   = tn   / (tn + fp + 1e-6)\n    acc    = (tp + tn) / (tp + tn + fp + fn)\n    return dict(dice=dice, f1=f1, prec=prec, rec=rec,\n                iou=iou, spec=spec, acc=acc,\n                tp=int(tp), tn=int(tn), fp=int(fp), fn=int(fn))\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# VISUALIZATION\n# ═════════════════════════════════════════════════════════════════════════════\ndef viz_full_fragment(volume, gt_mask, prob_map, threshold,\n                      metrics, save_dir, max_display=1200):\n    \"\"\"6-panel full-fragment comparison.\"\"\"\n    H, W  = gt_mask.shape\n    scale = min(1.0, max_display / max(H, W))\n    def ds(arr, order=1):\n        if scale < 1.0:\n            return zoom(arr, (scale, scale), order=order)\n        return arr\n    def ds3(arr):\n        if scale < 1.0:\n            return zoom(arr, (scale, scale, 1), order=1)\n        return arr\n\n    mid   = volume[:, :, volume.shape[-1] // 2]\n    mid_n = (mid - mid.min()) / (mid.ptp() + 1e-6)\n    mid_d = ds(mid_n)\n    gt_d  = ds(gt_mask.astype(float), order=0).astype(np.uint8)\n    pr_d  = ds(prob_map)\n    bn_d  = (pr_d > threshold).astype(np.uint8)\n\n    tp = (bn_d == 1) & (gt_d == 1)\n    fp = (bn_d == 1) & (gt_d == 0)\n    fn = (bn_d == 0) & (gt_d == 1)\n\n    over = np.stack([mid_d]*3, axis=-1)\n    over[tp] = [0.0, 0.9, 0.2]\n    over[fp] = [0.9, 0.1, 0.1]\n    over[fn] = [0.1, 0.4, 0.9]\n\n    err_map = np.zeros((*gt_d.shape, 3), np.float32)\n    err_map[tp] = [0.0, 0.9, 0.2]\n    err_map[fp] = [0.9, 0.1, 0.1]\n    err_map[fn] = [0.1, 0.4, 0.9]\n\n    fig, axes = plt.subplots(2, 3, figsize=(21, 14))\n    m = metrics\n\n    axes[0,0].imshow(mid_d, cmap='gray')\n    axes[0,0].set_title(f'CT input (slice {(SLICE_START+SLICE_END)//2})', fontweight='bold')\n\n    axes[0,1].imshow(gt_d, cmap='gray', vmin=0, vmax=1)\n    axes[0,1].set_title(f'Ground truth  (ink px: {gt_mask.sum():,})', fontweight='bold')\n\n    im = axes[0,2].imshow(pr_d, cmap='hot', vmin=0, vmax=1)\n    axes[0,2].set_title('Prediction probability', fontweight='bold')\n    plt.colorbar(im, ax=axes[0,2], fraction=0.046, pad=0.04)\n\n    axes[1,0].imshow(bn_d, cmap='gray', vmin=0, vmax=1)\n    axes[1,0].set_title(f'Binary prediction  (thresh={threshold:.2f})', fontweight='bold')\n\n    axes[1,1].imshow(over)\n    axes[1,1].set_title('Overlay  [green=TP  red=FP  blue=FN]', fontweight='bold')\n\n    axes[1,2].imshow(err_map)\n    leg = [mpatches.Patch(color='green', label=f'TP {m[\"tp\"]:,}'),\n           mpatches.Patch(color='red',   label=f'FP {m[\"fp\"]:,}'),\n           mpatches.Patch(color='blue',  label=f'FN {m[\"fn\"]:,}')]\n    axes[1,2].legend(handles=leg, loc='upper right', fontsize=9)\n    axes[1,2].set_title('Error map', fontweight='bold')\n\n    for ax in axes.flat: ax.axis('off')\n\n    title = (f\"Fragment 1  |  Dice {m['dice']:.4f}  ·  F1 {m['f1']:.4f}  \"\n             f\"·  Prec {m['prec']:.4f}  ·  Rec {m['rec']:.4f}  \"\n             f\"·  IoU {m['iou']:.4f}  ·  Thresh {threshold:.2f}\")\n    fig.suptitle(title, fontsize=12, y=1.01)\n    plt.tight_layout()\n\n    path = os.path.join(save_dir, 'fragment1_comparison.png')\n    plt.savefig(path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"  Saved → {path}\")\n\n    # Save B&W PNGs\n    cv2.imwrite(os.path.join(save_dir, 'gt_binary.png'),    gt_mask * 255)\n    cv2.imwrite(os.path.join(save_dir, 'pred_binary.png'),  (prob_map > threshold).astype(np.uint8) * 255)\n    cv2.imwrite(os.path.join(save_dir, 'pred_heatmap.png'), (prob_map * 255).astype(np.uint8))\n\n\ndef viz_zoomed_errors(volume, gt_mask, prob_map, threshold, save_dir,\n                      n=4, size=256):\n    \"\"\"Find largest error regions and show zoomed comparisons.\"\"\"\n    pred = (prob_map > threshold).astype(np.uint8)\n    diff = np.abs(pred - gt_mask)\n    labeled, nf = nd_label(diff)\n    if nf == 0:\n        print(\"  No errors found — skipping zoomed viz.\")\n        return\n    sizes  = [np.sum(labeled == i) for i in range(1, nf+1)]\n    top    = sorted(range(nf), key=lambda i: sizes[i], reverse=True)[:n*2]\n    mid    = volume[:, :, volume.shape[-1] // 2]\n    mid_n  = (mid - mid.min()) / (mid.ptp() + 1e-6)\n    H, W   = gt_mask.shape\n    regions = []\n    for idx in top:\n        sl  = find_objects(labeled == (idx+1))[0]\n        cy  = (sl[0].start + sl[0].stop) // 2\n        cx  = (sl[1].start + sl[1].stop) // 2\n        y0, y1 = max(0, cy-size//2), min(H, cy+size//2)\n        x0, x1 = max(0, cx-size//2), min(W, cx+size//2)\n        if (y1-y0) < 64 or (x1-x0) < 64: continue\n        regions.append((y0, y1, x0, x1))\n        if len(regions) >= n: break\n\n    if not regions: return\n    fig, axes = plt.subplots(len(regions), 4, figsize=(16, 4*len(regions)))\n    if len(regions) == 1: axes = axes[np.newaxis]\n\n    for ri, (y0, y1, x0, x1) in enumerate(regions):\n        inp  = mid_n[y0:y1, x0:x1]\n        gt   = gt_mask[y0:y1, x0:x1]\n        pr   = prob_map[y0:y1, x0:x1]\n        bn   = (pr > threshold).astype(np.uint8)\n        tp = (bn==1)&(gt==1); fp = (bn==1)&(gt==0); fn = (bn==0)&(gt==1)\n        em = np.zeros((*inp.shape, 3), np.float32)\n        em[tp]=[0,0.9,0.2]; em[fp]=[0.9,0.1,0.1]; em[fn]=[0.1,0.4,0.9]\n\n        axes[ri,0].imshow(inp, cmap='gray');           axes[ri,0].set_title(f'Region {ri+1} – input')\n        axes[ri,1].imshow(gt,  cmap='gray', vmin=0, vmax=1); axes[ri,1].set_title('Ground truth')\n        axes[ri,2].imshow(pr,  cmap='hot',  vmin=0, vmax=1); axes[ri,2].set_title('Prediction prob')\n        axes[ri,3].imshow(em);                         axes[ri,3].set_title(f'TP {tp.sum()} FP {fp.sum()} FN {fn.sum()}')\n        for ax in axes[ri]: ax.axis('off')\n\n    plt.suptitle('Zoomed error regions  [green=TP  red=FP  blue=FN]', fontsize=13)\n    plt.tight_layout()\n    path = os.path.join(save_dir, 'zoomed_errors.png')\n    plt.savefig(path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"  Saved → {path}\")\n\n\n# ═════════════════════════════════════════════════════════════════════════════\n# MAIN\n# ═════════════════════════════════════════════════════════════════════════════\ndef main():\n    print(\"=\" * 60)\n    print(\"VESUVIUS INK DETECTION v3 — target Dice > 0.80\")\n    print(\"=\" * 60)\n    t0 = time.time()\n\n    # ── 1. Load data ─────────────────────────────────────────────────────────\n    print(\"\\n1. Loading training data...\")\n    all_patches, all_masks, all_ignores, frag_sizes = [], [], [], []\n\n    for path in TRAIN_PATHS:\n        name = os.path.basename(path)\n        print(f\"\\n  Fragment {name}\")\n        vol  = load_volume(path)\n        mask = cv2.imread(os.path.join(path, \"inklabels.png\"), 0)\n        mask = (mask > 0).astype(np.uint8)\n        ign  = generate_ignore_mask(mask)\n        print(f\"    Volume {vol.shape}  ink px: {mask.sum():,}\")\n        p, m, i = extract_patches(vol, mask, ign)\n        all_patches.extend(p); all_masks.extend(m); all_ignores.extend(i)\n        frag_sizes.append(len(p))\n        del vol, mask, ign, p, m, i; gc.collect()\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n\n    print(f\"\\nTotal patches: {len(all_patches)}\")\n\n    # ── 2. Fragment-aware train/val split ────────────────────────────────────\n    print(\"\\n2. Fragment-aware train/val split...\")\n    tr_v, tr_m, tr_i, va_v, va_m, va_i = fragment_aware_split(\n        all_patches, all_masks, all_ignores, frag_sizes)\n    print(f\"  Train: {len(tr_v)}  Val: {len(va_v)}\")\n    del all_patches, all_masks, all_ignores; gc.collect()\n\n    # ── 3. Datasets / loaders ────────────────────────────────────────────────\n    print(\"\\n3. Building datasets...\")\n    tr_ds = VesuviusDataset(tr_v, tr_m, tr_i, transform=get_train_transform())\n    va_ds = VesuviusDataset(va_v, va_m, va_i)\n    tr_dl = DataLoader(tr_ds, batch_size=BATCH_SIZE, shuffle=True,\n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                       drop_last=True, collate_fn=collate_fn)\n    va_dl = DataLoader(va_ds, batch_size=BATCH_SIZE, shuffle=False,\n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                       collate_fn=collate_fn)\n    print(f\"  Train batches: {len(tr_dl)}  Val batches: {len(va_dl)}\")\n\n    # ── 4. Model ─────────────────────────────────────────────────────────────\n    print(\"\\n4. Building model with correct pretrained transfer...\")\n    model = build_model_with_pretrained(N_SLICES).to(DEVICE)\n    n_all = sum(p.numel() for p in model.parameters())\n    n_trn = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Params: {n_all:,}  trainable: {n_trn:,}\")\n\n    # ── 5. Loss / optimizer / scheduler ─────────────────────────────────────\n    criterion = FocalDiceLovaszLoss()\n    param_groups = get_param_groups(model)\n    optimizer    = optim.AdamW(param_groups, weight_decay=WEIGHT_DECAY)\n    scheduler    = optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=20, T_mult=2, eta_min=1e-6)\n    scaler       = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n\n    # ── 6. Training loop ─────────────────────────────────────────────────────\n    print(\"\\n5. Training...\")\n    print(\"=\" * 60)\n    best_dice, patience_ctr = 0.0, 0\n    train_losses, val_dices = [], []\n\n    for epoch in range(EPOCHS):\n        t_ep = time.time()\n\n        # Train\n        model.train()\n        ep_loss, steps = 0.0, 0\n        optimizer.zero_grad()\n        pbar = tqdm(tr_dl, desc=f\"E{epoch+1:03d}/{EPOCHS} [train]\", leave=False)\n        for step, (imgs, msks, igns) in enumerate(pbar):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            igns = igns.to(DEVICE, non_blocking=True)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                out  = model(imgs)\n                loss = criterion(out, msks, igns) / ACCUM_STEPS\n            scaler.scale(loss).backward()\n            if (step + 1) % ACCUM_STEPS == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(optimizer); scaler.update()\n                optimizer.zero_grad()\n            ep_loss += loss.item() * ACCUM_STEPS; steps += 1\n            pbar.set_postfix(loss=f\"{loss.item()*ACCUM_STEPS:.4f}\")\n        avg_loss = ep_loss / steps\n        train_losses.append(avg_loss)\n        scheduler.step(epoch + 1)\n\n        # Validate\n        model.eval()\n        val_dice_acc, val_steps = 0.0, 0\n        with torch.no_grad():\n            for imgs, msks, igns in tqdm(va_dl, desc=f\"E{epoch+1:03d}/{EPOCHS} [val]\",\n                                         leave=False):\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                msks = msks.to(DEVICE, non_blocking=True)\n                igns = igns.to(DEVICE, non_blocking=True)\n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    out = model(imgs)\n                prds  = (torch.sigmoid(out) > 0.5)\n                valid = (1 - igns)\n                inter = ((prds * msks) * valid).sum()\n                union = ((prds * valid).sum() + (msks * valid).sum())\n                val_dice_acc += (2 * inter / (union + 1e-6)).item()\n                val_steps    += 1\n        avg_dice = val_dice_acc / val_steps\n        val_dices.append(avg_dice)\n\n        ep_t = time.time() - t_ep\n        enc_lr = optimizer.param_groups[0]['lr']\n        dec_lr = optimizer.param_groups[1]['lr']\n        print(f\"  E{epoch+1:03d}/{EPOCHS}  loss {avg_loss:.4f}  \"\n              f\"val_dice {avg_dice:.4f}  \"\n              f\"lr_enc {enc_lr:.1e}  lr_dec {dec_lr:.1e}  [{ep_t:.0f}s]\")\n\n        if avg_dice > best_dice:\n            best_dice = avg_dice\n            torch.save({'epoch': epoch, 'model_state_dict': model.state_dict(),\n                        'val_dice': avg_dice},\n                       os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_ctr = 0\n            print(f\"    ✓ Best model saved  (Dice {best_dice:.4f})\")\n        else:\n            patience_ctr += 1\n            if patience_ctr >= EARLY_STOP:\n                print(f\"\\nEarly stopping at epoch {epoch+1}\")\n                break\n\n    # ── 7. Threshold tuning ──────────────────────────────────────────────────\n    print(\"\\n6. Tuning decision threshold on val set...\")\n    ckpt = torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth'))\n    model.load_state_dict(ckpt['model_state_dict'])\n    best_thresh = tune_threshold(model, va_dl, DEVICE)\n\n    # ── 8. Full-fragment inference on Fragment 1 ─────────────────────────────\n    print(\"\\n7. Loading Fragment 1 for full-fragment inference...\")\n    test_vol  = load_volume(FRAGMENT1_PATH)\n    test_mask = cv2.imread(os.path.join(FRAGMENT1_PATH, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ign  = generate_ignore_mask(test_mask)\n    print(f\"  Volume {test_vol.shape}  ink px: {test_mask.sum():,}\")\n\n    print(\"\\n8. Sliding-window inference with TTA...\")\n    # Use stride=PATCH_SIZE//4 for high quality, or PATCH_SIZE//2 if slow\n    inf_stride = PATCH_SIZE // 4\n    prob_map   = sliding_window_inference(\n        model, test_vol, stride=inf_stride, use_tta=True, device=DEVICE)\n\n    # ── 9. Metrics ───────────────────────────────────────────────────────────\n    print(\"\\n9. Computing metrics...\")\n    m = compute_metrics(test_mask, prob_map, best_thresh, test_ign)\n    print(\"\\n\" + \"=\" * 60)\n    print(\"FINAL TEST RESULTS — FRAGMENT 1\")\n    print(\"=\" * 60)\n    for k, v in m.items():\n        if isinstance(v, float): print(f\"  {k:<12}: {v:.4f}\")\n        else:                    print(f\"  {k:<12}: {v:,}\")\n    print(\"=\" * 60)\n\n    # ── 10. Visualizations ───────────────────────────────────────────────────\n    viz_dir = os.path.join(OUTPUT_DIR, 'viz')\n    print(\"\\n10. Generating visualizations...\")\n    viz_full_fragment(test_vol, test_mask, prob_map, best_thresh, m, viz_dir)\n    viz_zoomed_errors(test_vol, test_mask, prob_map, best_thresh, viz_dir, n=4)\n\n    # Training curves\n    fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n    axes[0].plot(train_losses, color='steelblue'); axes[0].set_title('Training loss')\n    axes[0].set_xlabel('Epoch'); axes[0].grid(True)\n    axes[1].plot(val_dices, color='orange')\n    axes[1].axhline(best_dice, color='green', linestyle='--',\n                    label=f'Best {best_dice:.4f}')\n    axes[1].set_title('Validation Dice'); axes[1].set_xlabel('Epoch')\n    axes[1].legend(); axes[1].grid(True)\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'), dpi=150)\n    plt.show()\n\n    # Save CSVs / TXT\n    total_t = time.time() - t0\n    h, mins = int(total_t // 3600), int((total_t % 3600) // 60)\n    pd.DataFrame([m]).to_csv(os.path.join(viz_dir, 'metrics.csv'), index=False)\n    with open(os.path.join(viz_dir, 'results.txt'), 'w') as f:\n        f.write(f\"Total time     : {h}h {mins}m\\n\")\n        f.write(f\"Best val Dice  : {best_dice:.4f}\\n\")\n        f.write(f\"Threshold      : {best_thresh:.3f}\\n\")\n        for k, v in m.items():\n            f.write(f\"{k}: {v}\\n\")\n\n    print(f\"\\nTotal time: {h}h {mins}m\")\n    return m['dice']\n\n\nif __name__ == \"__main__\":\n    try:\n        score = main()\n        print(f\"\\n✅ Final Test Dice: {score:.4f}\")\n    except Exception as e:\n        import traceback\n        print(f\"\\n❌ Error: {e}\")\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# Load the image\nimage_path = '/kaggle/working/detailed_comparison_visualization.png'\noutput_path = '/kaggle/working/detailed_comparison_visualization_bw_inverted.png'\n\n# Read the image\nimg = cv2.imread(image_path)\nimg_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n# Convert to grayscale to identify black and white regions\ngray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n\n# Create masks for black and white pixels\n# Black pixels (low intensity)\nblack_mask = gray < 30\n# White pixels (high intensity)\nwhite_mask = gray > 225\n\n# Create inverted image (start with original)\ninverted_img = img.copy()\n\n# Invert black to white\ninverted_img[black_mask] = [255, 255, 255]\n# Invert white to black\ninverted_img[white_mask] = [0, 0, 0]\n\n# Save the result\ncv2.imwrite(output_path, inverted_img)\n\nprint(f\"Image with inverted black/white saved to: {output_path}\")\n\n# Display the result\nplt.figure(figsize=(15, 10))\nplt.imshow(cv2.cvtColor(inverted_img, cv2.COLOR_BGR2RGB))\nplt.axis('off')\nplt.title('Inverted Black/White (Colors Preserved)')\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2.5D Unet\n!pip install segmentation-models-pytorch==0.2.0\n#/kaggle/input/vesuvius-challenge-ink-detection/train\n# ============================================================\n#  VESUVIUS INK DETECTION — VCSD (Volumetric Contrastive\n#  Spectral Decomposition)\n#  Novel PhD Methodology — NOT Semantic Segmentation\n#\n#  Key innovations over standard U-Net approach:\n#  1. Z-Profile Transformer Encoder (ZPT) — treats depth as sequence\n#  2. Contrastive Material Disentanglement (CMD) — learns material prototypes\n#  3. Spatial-Spectral Fusion (lightweight depthwise-sep + channel attn)\n#  4. Multi-Scale Depth Attention Aggregation (MSDAA)\n#  5. Evidential Uncertainty prediction\n#\n#  Memory-safe for Kaggle T4 16GB\n# ============================================================\n\nimport os, gc, cv2, math, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.utils.checkpoint import checkpoint as grad_checkpoint\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\n# ════════════════════════════════════════════════════════════\n#  0. ENVIRONMENT & MEMORY SAFETY\n# ════════════════════════════════════════════════════════════\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:64\"\n\nSEED = 42\ntorch.manual_seed(SEED)\nnp.random.seed(SEED)\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cudnn.deterministic = False\n\n# ── paths ────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge-ink-detection/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params (OOM-safe for T4 16GB) ─────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 224\nSTRIDE_TR    = 64\nSTRIDE_INF   = 56\nBATCH_SIZE   = 2\nGRAD_ACCUM   = 8\nEPOCHS       = 20\nLR           = 3e-4\nWEIGHT_DECAY = 1e-4\nPATIENCE     = 12\nNUM_WORKERS  = 0\n\n# ── Z-slices: EXPANDED for richer z-profiles ────────────────\nZ_SLICES = list(range(20, 45))   # 25 slices\nN_CH     = len(Z_SLICES)\n\n# ── VCSD-specific params ────────────────────────────────────\nZPT_D_MODEL      = 128\nZPT_NHEAD        = 4\nZPT_LAYERS       = 3\nN_MATERIALS       = 4\nCMD_PROJ_DIM      = 64\nCMD_TEMPERATURE   = 0.07\nMSDAA_SCALES      = 4\nZPT_SPATIAL_STRIDE = 4      # subsample spatial grid for ZPT\nZPT_CHUNK_SIZE     = 1024   # chunk size for OOM safety\n\nINK_MIN_POS = 0.02\nNEG_RATIO   = 0.5\n\nprint(f\"{'='*60}\")\nprint(f\"  VCSD — Volumetric Contrastive Spectral Decomposition\")\nprint(f\"  Novel PhD Framework for Vesuvius Ink Detection\")\nprint(f\"{'='*60}\")\nprint(f\"Device     : {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU        : {torch.cuda.get_device_name(0)}\")\n    vram = torch.cuda.get_device_properties(0).total_mem / 1e9\n    print(f\"VRAM       : {vram:.1f} GB\")\nprint(f\"Z-slices   : {N_CH} (range {Z_SLICES[0]}-{Z_SLICES[-1]})\")\nprint(f\"Patch      : {PATCH_SIZE}  |  Batch: {BATCH_SIZE}x{GRAD_ACCUM}={BATCH_SIZE*GRAD_ACCUM}\")\nprint(f\"ZPT dim    : {ZPT_D_MODEL}  |  Materials: {N_MATERIALS}\")\nprint(f\"{'='*60}\\n\")\n\n\n# ════════════════════════════════════════════════════════════\n#  1. SLICE CACHE\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    raw, all_v = {}, []\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            raw[z] = s\n            all_v.append(s.ravel())\n    if not all_v:\n        raise FileNotFoundError(f\"No slices found in {vol_dir}\")\n    p5, p95 = np.percentile(np.concatenate(all_v), [5, 95])\n    del all_v; gc.collect()\n    cache = {}\n    for z, s in raw.items():\n        s = np.clip(s, p5, p95)\n        cache[z] = ((s - p5) / (p95 - p5 + 1e-6)).astype(np.float16)\n    del raw; gc.collect()\n    mb = sum(v.nbytes for v in cache.values()) / 1e6\n    print(f\"  Slice cache: {len(cache)} slices, {mb:.0f} MB\")\n    return cache\n\n\n# ════════════════════════════════════════════════════════════\n#  2. DATASET\n# ════════════════════════════════════════════════════════════\nclass VCSDDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride, transform=None,\n                 neg_ratio=0.0, generate_material_labels=True):\n        self.cache   = load_slice_cache(frag_path, z_list)\n        self.z_list  = z_list\n        self.tf      = transform\n        self.gen_mat = generate_material_labels\n\n        msk = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n        self.mask = (msk > 0).astype(np.uint8)\n\n        ir_path = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = (cv2.imread(ir_path, 0) > 0).astype(np.uint8) \\\n                       if os.path.exists(ir_path) else None\n\n        if self.gen_mat:\n            mid_z = z_list[len(z_list) // 2]\n            self.mean_intensity = self.cache[mid_z].astype(np.float32)\n\n        H, W = self.mask.shape\n        pos_coords, neg_coords = [], []\n\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                m   = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                ink = m.sum() / (PATCH_SIZE * PATCH_SIZE)\n\n                if self.ir_mask is not None:\n                    on_papyrus = self.ir_mask[\n                        y:y+PATCH_SIZE, x:x+PATCH_SIZE\n                    ].mean() > 0.5\n                else:\n                    mid_z = z_list[len(z_list) // 2]\n                    on_papyrus = float(\n                        self.cache[mid_z][\n                            y:y+PATCH_SIZE, x:x+PATCH_SIZE\n                        ].mean()\n                    ) > 0.1\n\n                if not on_papyrus:\n                    continue\n                if ink >= INK_MIN_POS:\n                    pos_coords.append((y, x, 1))\n                elif ink < 0.001 and neg_ratio > 0:\n                    neg_coords.append((y, x, 0))\n\n        n_neg = int(len(pos_coords) * neg_ratio)\n        selected_neg = []\n        if n_neg > 0 and neg_coords:\n            np.random.shuffle(neg_coords)\n            selected_neg = neg_coords[:n_neg]\n\n        self.coords = pos_coords + selected_neg\n        np.random.shuffle(self.coords)\n        self.weights = np.array(\n            [3.0 if c[2] == 1 else 1.0 for c in self.coords],\n            dtype=np.float32,\n        )\n        print(f\"  Patches: {len(pos_coords)} pos + {len(selected_neg)} neg \"\n              f\"= {len(self.coords)} total\")\n\n    def _generate_material_map(self, y, x, ink_mask):\n        mat = np.ones((PATCH_SIZE, PATCH_SIZE), dtype=np.int64)  # default papyrus=1\n\n        # Ink = 0\n        mat[ink_mask > 0] = 0\n\n        # Air = 2 (very low intensity)\n        intensity_patch = self.mean_intensity[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n        mat[intensity_patch < 0.05] = 2\n\n        # Artifact = 3 (high variance across z)\n        z_stack = np.stack([\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list[::4]\n        ], axis=0)\n        z_var = z_stack.var(axis=0)\n        high_var_thresh = np.percentile(z_var, 95)\n        mat[(z_var > high_var_thresh) & (ink_mask == 0)] = 3\n\n        return mat\n\n    def __len__(self):\n        return len(self.coords)\n\n    def __getitem__(self, idx):\n        y, x, _ = self.coords[idx]\n        slices = [\n            self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n            for z in self.z_list\n        ]\n        img = np.stack(slices, axis=-1)\n        msk = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n\n        if self.gen_mat:\n            mat_map = self._generate_material_map(y, x, msk)\n        else:\n            mat_map = np.ones((PATCH_SIZE, PATCH_SIZE), dtype=np.int64)\n\n        if self.tf:\n            out = self.tf(image=img, masks=[msk, mat_map])\n            img     = out['image']\n            msk     = out['masks'][0]\n            mat_map = out['masks'][1]\n\n        img_t = torch.from_numpy(img).permute(2, 0, 1).float()\n        msk_t = torch.from_numpy(msk).unsqueeze(0).float()\n        mat_t = torch.from_numpy(mat_map.copy()).long()\n\n        return img_t, msk_t, mat_t\n\n\n# ════════════════════════════════════════════════════════════\n#  3. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ShiftScaleRotate(\n        shift_limit=0.1, scale_limit=0.15,\n        rotate_limit=30,\n        border_mode=cv2.BORDER_REFLECT, p=0.6,\n    ),\n    A.RandomBrightnessContrast(0.15, 0.15, p=0.4),\n    A.GaussNoise(var_limit=(0.001, 0.003), p=0.3),\n    A.GaussianBlur(blur_limit=(3, 5), p=0.2),\n    A.CoarseDropout(\n        max_holes=3, max_height=20,\n        max_width=20, fill_value=0, p=0.25,\n    ),\n])\n\n\n# ════════════════════════════════════════════════════════════\n#  4. MODEL COMPONENTS\n# ════════════════════════════════════════════════════════════\n\n# ── 4A. Z-Profile Transformer Encoder ───────────────────────\nclass ZProfileTransformer(nn.Module):\n    def __init__(self, z_depth=25, d_model=128, nhead=4, num_layers=3):\n        super().__init__()\n        self.d_model = d_model\n        self.z_depth = z_depth\n\n        self.input_proj = nn.Sequential(\n            nn.Linear(1, d_model),\n            nn.GELU(),\n            nn.LayerNorm(d_model),\n        )\n        self.z_pos_embed = nn.Parameter(torch.randn(1, z_depth, d_model) * 0.02)\n        self.cls_token   = nn.Parameter(torch.randn(1, 1, d_model) * 0.02)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=d_model * 2,\n            dropout=0.1, activation='gelu',\n            batch_first=True, norm_first=True,\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            enable_nested_tensor=False,\n        )\n        self.cls_proj = nn.Sequential(\n            nn.Linear(d_model, d_model),\n            nn.GELU(),\n        )\n\n    def forward(self, z_profiles):\n        B, Z = z_profiles.shape\n        tokens = self.input_proj(z_profiles.unsqueeze(-1))       # (B, Z, d)\n        tokens = tokens + self.z_pos_embed[:, :Z, :]\n        cls    = self.cls_token.expand(B, -1, -1)\n        tokens = torch.cat([cls, tokens], dim=1)                 # (B, Z+1, d)\n        out    = self.transformer(tokens)\n        cls_feat   = self.cls_proj(out[:, 0])                    # (B, d)\n        token_feat = out[:, 1:]                                  # (B, Z, d)\n        return cls_feat, token_feat\n\n\n# ── 4B. Multi-Scale Depth Attention ──────────────────────────\nclass MultiScaleDepthAttention(nn.Module):\n    def __init__(self, z_depth=25, d_model=128, num_scales=4):\n        super().__init__()\n        self.num_scales = num_scales\n\n        self.depth_centers = nn.Parameter(\n            torch.linspace(0, z_depth - 1, num_scales)\n        )\n        self.depth_widths = nn.Parameter(\n            torch.ones(num_scales) * (z_depth / num_scales)\n        )\n        self.scale_projs = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(d_model, d_model), nn.GELU(), nn.LayerNorm(d_model),\n            )\n            for _ in range(num_scales)\n        ])\n        self.cross_attn  = nn.MultiheadAttention(\n            d_model, num_heads=2, batch_first=True, dropout=0.1,\n        )\n        self.fusion_norm = nn.LayerNorm(d_model)\n\n    def forward(self, token_features):\n        B, Z, C = token_features.shape\n        device  = token_features.device\n        depth_pos = torch.arange(Z, device=device, dtype=torch.float32)\n\n        scale_feats = []\n        for s in range(self.num_scales):\n            center  = self.depth_centers[s]\n            width   = self.depth_widths[s].clamp(min=1.0)\n            weights = torch.exp(-0.5 * ((depth_pos - center) / width) ** 2)\n            weights = weights / (weights.sum() + 1e-8)\n            weighted = (token_features * weights.view(1, Z, 1)).sum(dim=1)\n            weighted = self.scale_projs[s](weighted)\n            scale_feats.append(weighted)\n\n        scales      = torch.stack(scale_feats, dim=1)\n        fused, _    = self.cross_attn(scales, scales, scales)\n        return self.fusion_norm(fused.mean(dim=1))\n\n\n# ── 4C. Contrastive Material Disentanglement ────────────────\nclass ContrastiveMaterialDisentanglement(nn.Module):\n    def __init__(self, d_model=128, n_materials=4, proj_dim=64,\n                 temperature=0.07):\n        super().__init__()\n        self.n_materials = n_materials\n        self.temperature = temperature\n        self.material_prototypes = nn.Parameter(\n            torch.randn(n_materials, d_model) * 0.1\n        )\n        self.projector = nn.Sequential(\n            nn.Linear(d_model, d_model), nn.GELU(),\n            nn.Linear(d_model, proj_dim),\n        )\n\n    def forward(self, features, material_labels=None):\n        feat_norm  = F.normalize(features, dim=-1)\n        proto_norm = F.normalize(self.material_prototypes, dim=-1)\n        material_sim = torch.matmul(feat_norm, proto_norm.T) / self.temperature\n\n        loss_cmd = None\n        if material_labels is not None:\n            loss_proto = F.cross_entropy(material_sim, material_labels)\n            proto_gram = torch.matmul(proto_norm, proto_norm.T)\n            identity   = torch.eye(self.n_materials, device=features.device)\n            loss_orth  = F.mse_loss(proto_gram, identity)\n            loss_cmd   = loss_proto + 0.1 * loss_orth\n\n        return material_sim, loss_cmd\n\n\n# ── 4D. Spatial Fusion Block ────────────────────────────────\nclass SpatialFusionBlock(nn.Module):\n    def __init__(self, channels, reduction=4):\n        super().__init__()\n        self.spatial_conv = nn.Sequential(\n            nn.Conv2d(channels, channels, 3, padding=1,\n                      groups=channels, bias=False),\n            nn.BatchNorm2d(channels), nn.GELU(),\n            nn.Conv2d(channels, channels, 1, bias=False),\n            nn.BatchNorm2d(channels), nn.GELU(),\n        )\n        self.dilated_conv = nn.Sequential(\n            nn.Conv2d(channels, channels, 3, padding=3, dilation=3,\n                      groups=channels, bias=False),\n            nn.BatchNorm2d(channels), nn.GELU(),\n            nn.Conv2d(channels, channels, 1, bias=False),\n            nn.BatchNorm2d(channels),\n        )\n        mid = max(channels // reduction, 8)\n        self.channel_attn = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(channels, mid), nn.GELU(),\n            nn.Linear(mid, channels), nn.Sigmoid(),\n        )\n\n    def forward(self, x):\n        local  = self.spatial_conv(x)\n        remote = self.dilated_conv(x)\n        fused  = local + remote\n        gate   = self.channel_attn(fused).unsqueeze(-1).unsqueeze(-1)\n        return x + fused * gate\n\n\n# ── 4E. Evidential Prediction Head ──────────────────────────\nclass EvidentialHead(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.evidence = nn.Sequential(\n            nn.Conv2d(in_channels, 32, 3, padding=1), nn.GELU(),\n            nn.Conv2d(32, 2, 1), nn.Softplus(),\n        )\n\n    def forward(self, x):\n        evidence    = self.evidence(x)\n        alpha       = evidence + 1.0\n        S           = alpha.sum(dim=1, keepdim=True)\n        ink_prob    = alpha[:, 0:1] / S\n        uncertainty = 2.0 / S\n        return ink_prob, uncertainty, alpha\n\n\n# ── 4F. FULL VCSD MODEL ────────────────────────────────────\nclass VCSDModel(nn.Module):\n    def __init__(self, n_channels=25, d_model=128, n_materials=4):\n        super().__init__()\n        self.n_channels = n_channels\n        self.d_model    = d_model\n\n        # Z-channel compression (1x1 conv path — full resolution)\n        self.z_compress = nn.Sequential(\n            nn.Conv2d(n_channels, d_model, 1, bias=False),\n            nn.BatchNorm2d(d_model), nn.GELU(),\n            nn.Conv2d(d_model, d_model, 3, padding=1,\n                      groups=d_model, bias=False),\n            nn.BatchNorm2d(d_model), nn.GELU(),\n            nn.Conv2d(d_model, d_model, 1, bias=False),\n            nn.BatchNorm2d(d_model), nn.GELU(),\n        )\n\n        # Z-Profile Transformer (on subsampled grid)\n        self.zpt = ZProfileTransformer(\n            z_depth=n_channels, d_model=d_model,\n            nhead=ZPT_NHEAD, num_layers=ZPT_LAYERS,\n        )\n\n        # Multi-Scale Depth Attention\n        self.msdaa = MultiScaleDepthAttention(\n            z_depth=n_channels, d_model=d_model,\n            num_scales=MSDAA_SCALES,\n        )\n\n        # Contrastive Material Disentanglement\n        self.cmd = ContrastiveMaterialDisentanglement(\n            d_model=d_model, n_materials=n_materials,\n            proj_dim=CMD_PROJ_DIM, temperature=CMD_TEMPERATURE,\n        )\n\n        # Material gating\n        self.mat_gate = nn.Sequential(\n            nn.Conv2d(n_materials, d_model, 1), nn.Sigmoid(),\n        )\n\n        # Spatial Fusion\n        self.spatial_blocks = nn.Sequential(\n            SpatialFusionBlock(d_model),\n            SpatialFusionBlock(d_model),\n        )\n\n        # Decoder\n        self.decoder = nn.Sequential(\n            nn.Conv2d(d_model, d_model, 3, padding=1, bias=False),\n            nn.BatchNorm2d(d_model), nn.GELU(),\n            nn.Conv2d(d_model, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64), nn.GELU(),\n        )\n\n        # Heads\n        self.logit_head = nn.Conv2d(64, 1, 1)\n        self.evi_head   = EvidentialHead(64)\n\n    def _process_zpt_chunked(self, z_profiles, chunk_size):\n        \"\"\"Process z-profiles in chunks to avoid OOM.\"\"\"\n        N = z_profiles.shape[0]\n        cls_parts, tok_parts = [], []\n        for i in range(0, N, chunk_size):\n            chunk = z_profiles[i:i + chunk_size]\n            if self.training:\n                c, t = grad_checkpoint(self.zpt, chunk, use_reentrant=False)\n            else:\n                c, t = self.zpt(chunk)\n            cls_parts.append(c)\n            tok_parts.append(t)\n        return torch.cat(cls_parts, 0), torch.cat(tok_parts, 0)\n\n    def _process_msdaa_chunked(self, tok_all, chunk_size):\n        \"\"\"Process MSDAA in chunks.\"\"\"\n        N = tok_all.shape[0]\n        parts = []\n        for i in range(0, N, chunk_size):\n            chunk = tok_all[i:i + chunk_size]\n            if self.training:\n                ms = grad_checkpoint(self.msdaa, chunk, use_reentrant=False)\n            else:\n                ms = self.msdaa(chunk)\n            parts.append(ms)\n        return torch.cat(parts, 0)\n\n    def _process_cmd_chunked(self, msdaa_all, mat_labels_flat, chunk_size):\n        \"\"\"Process CMD in chunks.\"\"\"\n        N = msdaa_all.shape[0]\n        sim_parts = []\n        loss_sum  = 0.0\n        loss_cnt  = 0\n        for i in range(0, N, chunk_size):\n            feat = msdaa_all[i:i + chunk_size]\n            lab  = mat_labels_flat[i:i + chunk_size] if mat_labels_flat is not None else None\n            sim, loss_c = self.cmd(feat, lab)\n            sim_parts.append(sim)\n            if loss_c is not None:\n                loss_sum += loss_c\n                loss_cnt += 1\n        mat_sim = torch.cat(sim_parts, 0)\n        loss_cmd = loss_sum / max(loss_cnt, 1) if loss_cnt > 0 else torch.tensor(0.0, device=msdaa_all.device)\n        return mat_sim, loss_cmd\n\n    def forward(self, volume, material_labels=None):\n        \"\"\"\n        volume: (B, C, H, W)\n        material_labels: (B, H, W) long — optional\n        \"\"\"\n        B, C, H, W = volume.shape\n        device = volume.device\n        chunk  = ZPT_CHUNK_SIZE\n        stride = ZPT_SPATIAL_STRIDE\n\n        # ═══ PATH A: Spatial backbone (full resolution) ═══\n        z_feat = self.z_compress(volume)                    # (B, d, H, W)\n\n        # ═══ PATH B: Z-Profile Transformer (subsampled) ═══\n        vol_sub = volume[:, :, ::stride, ::stride]          # (B, C, Hs, Ws)\n        _, _, Hs, Ws = vol_sub.shape\n        Ns = Hs * Ws\n\n        z_profiles = vol_sub.permute(0, 2, 3, 1).reshape(B * Ns, C)\n\n        # B.1 — ZPT\n        cls_all, tok_all = self._process_zpt_chunked(z_profiles, chunk)\n        del z_profiles\n\n        # B.2 — MSDAA\n        msdaa_all = self._process_msdaa_chunked(tok_all, chunk)\n        del tok_all\n        gc.collect()\n\n        # B.3 — CMD\n        mat_labels_flat = None\n        if material_labels is not None:\n            mat_labels_flat = material_labels[:, ::stride, ::stride].reshape(-1)\n\n        mat_sim, loss_cmd = self._process_cmd_chunked(msdaa_all, mat_labels_flat, chunk)\n        del msdaa_all, cls_all\n        gc.collect()\n\n        # Reshape + upsample material similarity to full resolution\n        mat_sim_2d = mat_sim.reshape(B, Hs, Ws, -1).permute(0, 3, 1, 2)\n        mat_sim_2d = F.interpolate(mat_sim_2d, size=(H, W),\n                                   mode='bilinear', align_corners=False)\n        del mat_sim\n\n        # ═══ FUSION ═══\n        gate   = self.mat_gate(mat_sim_2d)\n        z_feat = z_feat * gate + z_feat\n        del mat_sim_2d, gate\n\n        # ═══ Spatial context ═══\n        if self.training:\n            z_feat = grad_checkpoint(\n                self.spatial_blocks, z_feat, use_reentrant=False,\n            )\n        else:\n            z_feat = self.spatial_blocks(z_feat)\n\n        # ═══ Decode + Predict ═══\n        decoded = self.decoder(z_feat)\n        logits  = self.logit_head(decoded)\n        ink_prob, uncertainty, alpha = self.evi_head(decoded)\n\n        return logits, ink_prob, uncertainty, alpha, loss_cmd\n\n\n# ════════════════════════════════════════════════════════════\n#  5. LOSS FUNCTIONS\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=2.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    a_t = target * alpha + (1 - target) * (1 - alpha)\n    return (a_t * ((1 - p_t) ** gamma) * bce).mean()\n\n\ndef dice_loss(pred, target, smooth=1.0):\n    p     = torch.sigmoid(pred)\n    inter = (p * target).sum(dim=(2, 3))\n    union = p.sum(dim=(2, 3)) + target.sum(dim=(2, 3))\n    return 1.0 - ((2.0 * inter + smooth) / (union + smooth)).mean()\n\n\ndef evidential_loss(alpha, target, epoch_frac=1.0):\n    S    = alpha.sum(dim=1, keepdim=True)\n    t_oh = torch.cat([target, 1.0 - target], dim=1)\n    loss_ece = (t_oh * (torch.digamma(S) - torch.digamma(alpha))).sum(dim=1)\n    wrong_evidence = ((1.0 - t_oh) * (alpha - 1.0)).sum(dim=1).clamp(min=0)\n    annealing = min(epoch_frac, 1.0)\n    return (loss_ece + annealing * 0.05 * wrong_evidence).mean()\n\n\ndef vcsd_combined_loss(logits, alpha, target, loss_cmd,\n                       epoch_frac=1.0, eps=0.05):\n    t_smooth = target * (1 - eps) + 0.5 * eps\n    L_focal = focal_loss(logits, t_smooth, alpha=0.8, gamma=2.0)\n    L_dice  = dice_loss(logits, target)\n    L_evi   = evidential_loss(alpha, target, epoch_frac)\n    L_cmd   = loss_cmd if isinstance(loss_cmd, torch.Tensor) else \\\n              torch.tensor(0.0, device=logits.device)\n\n    total = 0.35 * L_focal + 0.35 * L_dice + 0.15 * L_evi + 0.15 * L_cmd\n    return total, {\n        'focal': L_focal.item(), 'dice': L_dice.item(),\n        'evi': L_evi.item(), 'cmd': L_cmd.item(),\n    }\n\n\n# ════════════════════════════════════════════════════════════\n#  6. METRICS\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p     = (torch.sigmoid(logits) > thr).float()\n    inter = (p * masks).sum(dim=(1, 2, 3))\n    union = p.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3))\n    return ((2.0 * inter + 1e-5) / (union + 1e-5)).mean().item()\n\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.0\n    for t in np.arange(0.30, 0.80, 0.02):\n        p = (probs > t).astype(np.float32)\n        d = (2 * (p * targets).sum() + 1) / (p.sum() + targets.sum() + 1)\n        if d > best_d:\n            best_d, best_t = d, float(t)\n    return best_t, best_d\n\n\n# ════════════════════════════════════════════════════════════\n#  7. BUILD MODEL\n# ════════════════════════════════════════════════════════════\ndef build_vcsd_model():\n    model = VCSDModel(\n        n_channels=N_CH, d_model=ZPT_D_MODEL, n_materials=N_MATERIALS,\n    )\n    total  = sum(p.numel() for p in model.parameters())\n    train_ = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  VCSD Model: {total/1e6:.2f}M params ({train_/1e6:.2f}M trainable)\")\n    return model.to(DEVICE)\n\n\n# ════════════════════════════════════════════════════════════\n#  8. DATALOADERS\n# ════════════════════════════════════════════════════════════\nprint('\\n── Fragment 2 (train) ──')\ntrain_ds = VCSDDataset(\n    FRAG2, Z_SLICES, stride=STRIDE_TR,\n    transform=train_tf, neg_ratio=NEG_RATIO,\n    generate_material_labels=True,\n)\n\nprint('\\n── Fragment 3 (val) ──')\nval_ds = VCSDDataset(\n    FRAG3, Z_SLICES, stride=STRIDE_TR,\n    transform=None, neg_ratio=0.0,\n    generate_material_labels=True,\n)\n\ngc.collect()\nif DEVICE == 'cuda':\n    torch.cuda.empty_cache()\n\nsampler  = WeightedRandomSampler(\n    torch.from_numpy(train_ds.weights), len(train_ds), replacement=True,\n)\ntrain_dl = DataLoader(\n    train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n    num_workers=NUM_WORKERS, pin_memory=False,\n)\nval_dl = DataLoader(\n    val_ds, batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, pin_memory=False,\n)\nprint(f'\\nTrain batches: {len(train_dl)} | Val batches: {len(val_dl)}')\n\n\n# ════════════════════════════════════════════════════════════\n#  9. TRAINING LOOP\n# ════════════════════════════════════════════════════════════\nmodel     = build_vcsd_model()\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='max', factor=0.5,\n    patience=5, min_lr=5e-6, verbose=True,\n)\nscaler    = torch.cuda.amp.GradScaler(enabled=(DEVICE == 'cuda'))\n\nbest_dice = 0.0\npat_cnt   = 0\nhistory   = dict(tl=[], vl=[], td=[], vd=[], lr=[],\n                 cmd=[], evi=[], unc=[])\n\nprint('\\n' + '=' * 60)\nprint(f'VCSD TRAINING ({EPOCHS} epochs)')\nprint(f'Loss: 0.35*Focal + 0.35*Dice + 0.15*Evidential + 0.15*CMD')\nprint(f'Grad checkpointing ON  |  Mixed precision ON')\nprint('=' * 60)\n\nfor epoch in range(EPOCHS):\n    epoch_frac = (epoch + 1) / EPOCHS\n\n    # ── TRAIN ─────────────────────────────────────────────\n    model.train()\n    tl = td = cmd_ep = evi_ep = 0.0\n    optimizer.zero_grad()\n\n    for step, (imgs, msks, mats) in enumerate(\n            tqdm(train_dl, desc=f'Ep{epoch+1:02d} train', leave=False)):\n\n        imgs = imgs.to(DEVICE)\n        msks = msks.to(DEVICE)\n        mats = mats.to(DEVICE)\n\n        with torch.cuda.amp.autocast(enabled=(DEVICE == 'cuda')):\n            logits, ink_prob, uncertainty, alpha, loss_cmd = model(imgs, mats)\n            loss, ld = vcsd_combined_loss(\n                logits, alpha, msks, loss_cmd, epoch_frac,\n            )\n            loss_scaled = loss / GRAD_ACCUM\n\n        scaler.scale(loss_scaled).backward()\n\n        if (step + 1) % GRAD_ACCUM == 0 or (step + 1) == len(train_dl):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n        tl     += loss.item()\n        td     += batch_dice(logits.detach(), msks)\n        cmd_ep += ld['cmd']\n        evi_ep += ld['evi']\n\n        del imgs, msks, mats, logits, ink_prob, uncertainty, alpha\n        del loss_cmd, loss, loss_scaled\n        if step % 30 == 0:\n            gc.collect(); torch.cuda.empty_cache()\n\n    n = len(train_dl)\n    tl /= n; td /= n; cmd_ep /= n; evi_ep /= n\n\n    # ── VALIDATE ──────────────────────────────────────────\n    model.eval()\n    vl = vd = unc_ep = 0.0\n    acc_p, acc_m = [], []\n\n    with torch.no_grad():\n        for imgs, msks, mats in tqdm(val_dl, desc=f'Ep{epoch+1:02d} val  ', leave=False):\n            imgs = imgs.to(DEVICE)\n            msks = msks.to(DEVICE)\n            mats = mats.to(DEVICE)\n\n            with torch.cuda.amp.autocast(enabled=(DEVICE == 'cuda')):\n                logits, ink_prob, uncertainty, alpha, loss_cmd = model(imgs, mats)\n                loss, _ = vcsd_combined_loss(\n                    logits, alpha, msks, loss_cmd, epoch_frac,\n                )\n\n            vl     += loss.item()\n            vd     += batch_dice(logits, msks)\n            unc_ep += uncertainty.mean().item()\n            acc_p.append(torch.sigmoid(logits).cpu().numpy())\n            acc_m.append(msks.cpu().numpy())\n\n            del imgs, msks, mats, logits, ink_prob, uncertainty, alpha, loss_cmd, loss\n\n    nv = len(val_dl)\n    vl /= nv; vd /= nv; unc_ep /= nv\n\n    probs_all = np.concatenate(acc_p)\n    masks_all = np.concatenate(acc_m)\n    bt, bd    = sweep_threshold(probs_all, masks_all)\n\n    ink_mean   = float(probs_all[masks_all > 0.5].mean()) if (masks_all > 0.5).any() else 0.0\n    noink_mean = float(probs_all[masks_all < 0.5].mean()) if (masks_all < 0.5).any() else 0.0\n    sep = ink_mean - noink_mean\n\n    del acc_p, acc_m, probs_all, masks_all\n    gc.collect(); torch.cuda.empty_cache()\n\n    scheduler.step(vd)\n    lr_now = optimizer.param_groups[0]['lr']\n\n    history['tl'].append(tl); history['vl'].append(vl)\n    history['td'].append(td); history['vd'].append(vd)\n    history['lr'].append(lr_now)\n    history['cmd'].append(cmd_ep); history['evi'].append(evi_ep)\n    history['unc'].append(unc_ep)\n\n    print(f'Ep{epoch+1:02d} | lr={lr_now:.1e} | '\n          f'tl={tl:.4f} td={td:.4f} | '\n          f'vl={vl:.4f} vd={vd:.4f} | '\n          f'thr={bt:.2f}->{bd:.4f} | '\n          f'sep={sep:+.3f} | cmd={cmd_ep:.4f} | unc={unc_ep:.3f}')\n\n    save_metric = max(vd, bd)\n    if save_metric > best_dice:\n        best_dice = save_metric; pat_cnt = 0\n        torch.save({\n            'epoch': epoch, 'state': model.state_dict(),\n            'thr': bt, 'dice': best_dice,\n            'n_ch': N_CH, 'method': 'VCSD',\n            'd_model': ZPT_D_MODEL, 'n_materials': N_MATERIALS,\n        }, OUTPUT + 'best_vcsd_model.pth')\n        print(f'  -> saved  (metric={best_dice:.4f})')\n    else:\n        pat_cnt += 1\n        if pat_cnt >= PATIENCE:\n            print(f'  early stop at epoch {epoch+1}'); break\n\nprint(f'\\nBest val metric: {best_dice:.4f}')\n\n\n# ── Training curves ──────────────────────────────────────────\nfig, axes = plt.subplots(2, 3, figsize=(18, 10))\naxes[0,0].plot(history['tl'], label='train'); axes[0,0].plot(history['vl'], label='val')\naxes[0,0].set_title('Total Loss'); axes[0,0].legend(); axes[0,0].grid(True)\naxes[0,1].plot(history['td'], label='train'); axes[0,1].plot(history['vd'], label='val')\naxes[0,1].axhline(0.65, color='orange', ls='--'); axes[0,1].axhline(0.90, color='r', ls='--')\naxes[0,1].set_title('Dice Score'); axes[0,1].legend(); axes[0,1].grid(True)\naxes[0,2].plot(history['lr']); axes[0,2].set_title('LR'); axes[0,2].grid(True)\naxes[1,0].plot(history['cmd'], color='purple'); axes[1,0].set_title('CMD Loss'); axes[1,0].grid(True)\naxes[1,1].plot(history['evi'], color='green'); axes[1,1].set_title('Evidential Loss'); axes[1,1].grid(True)\naxes[1,2].plot(history['unc'], color='red'); axes[1,2].set_title('Mean Uncertainty'); axes[1,2].grid(True)\nplt.tight_layout(); plt.savefig(OUTPUT+'vcsd_curves.png', dpi=100); plt.close()\nprint('Curves saved.')\n\n\n# ════════════════════════════════════════════════════════════\n#  10. INFERENCE\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c = sz // 2; sig = sz // 4\n    ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs-c)**2 + (ys-c)**2) / (2*sig**2)).astype(np.float32)\n\nGW = gauss_weight(PATCH_SIZE)\n\n\ndef predict_fragment_vcsd(model, frag_path, z_list):\n    model.eval()\n    cache = load_slice_cache(frag_path, z_list)\n    msk   = cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0)\n    msk   = (msk > 0).astype(np.uint8)\n    H, W  = msk.shape\n\n    pred_map = np.zeros((H, W), np.float32)\n    unc_map  = np.zeros((H, W), np.float32)\n    wgt_map  = np.zeros((H, W), np.float32)\n\n    coords = [(y, x)\n              for y in range(0, H - PATCH_SIZE + 1, STRIDE_INF)\n              for x in range(0, W - PATCH_SIZE + 1, STRIDE_INF)]\n\n    with torch.no_grad():\n        for i, (y, x) in enumerate(tqdm(coords, desc='VCSD Inference')):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32)\n                      for z in z_list]\n            t = torch.from_numpy(np.stack(slices, -1)).permute(2,0,1).unsqueeze(0).to(DEVICE)\n\n            with torch.cuda.amp.autocast(enabled=(DEVICE == 'cuda')):\n                logits, ink_prob, uncertainty, alpha, _ = model(t)\n\n            p = torch.sigmoid(logits).squeeze().cpu().numpy()\n            u = uncertainty.squeeze().cpu().numpy()\n\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            unc_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += u * GW\n            wgt_map [y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n\n            del t, logits, ink_prob, uncertainty, alpha, p, u\n            if i % 100 == 0:\n                torch.cuda.empty_cache()\n\n    del cache; gc.collect()\n    safe = wgt_map + 1e-8\n    return pred_map / safe, unc_map / safe, msk\n\n\n# ════════════════════════════════════════════════════════════\n#  11. FINAL TEST — FRAGMENT 1\n# ════════════════════════════════════════════════════════════\nprint('\\n' + '='*60)\nprint('FINAL TEST — FRAGMENT 1 (VCSD)')\nprint('='*60)\n\nckpt = torch.load(OUTPUT + 'best_vcsd_model.pth', map_location=DEVICE)\nassert ckpt.get('n_ch', N_CH) == N_CH, \"Channel mismatch!\"\n\nmodel = build_vcsd_model()\nmodel.load_state_dict(ckpt['state'], strict=True)\nsaved_thr = ckpt.get('thr', 0.5)\nprint(f'Checkpoint epoch {ckpt[\"epoch\"]+1}, '\n      f'thr={saved_thr:.2f}, dice={ckpt[\"dice\"]:.4f}')\n\nprob_map, unc_map, msk1 = predict_fragment_vcsd(model, FRAG1, Z_SLICES)\n\nH_m, W_m  = msk1.shape\nprob_crop  = prob_map[:H_m, :W_m]\nunc_crop   = unc_map[:H_m, :W_m]\n\nink_mean   = float(prob_crop[msk1 == 1].mean()) if (msk1 == 1).any() else 0.0\nnoink_mean = float(prob_crop[msk1 == 0].mean()) if (msk1 == 0).any() else 0.0\nprint(f'\\nCalibration:')\nprint(f'  ink    : {ink_mean:.3f}')\nprint(f'  no-ink : {noink_mean:.3f}')\nprint(f'  sep    : {ink_mean - noink_mean:+.3f}')\n\nunc_ink   = float(unc_crop[msk1 == 1].mean()) if (msk1 == 1).any() else 0.0\nunc_noink = float(unc_crop[msk1 == 0].mean()) if (msk1 == 0).any() else 0.0\nprint(f'\\nUncertainty:')\nprint(f'  ink    : {unc_ink:.4f}')\nprint(f'  no-ink : {unc_noink:.4f}')\n\nbt1, bd1   = sweep_threshold(prob_crop[None, None], msk1[None, None])\nfinal_pred = (prob_crop > bt1).astype(np.uint8)\n\ninter     = (final_pred * msk1).sum()\ndice_full = (2*inter + 1) / (final_pred.sum() + msk1.sum() + 1)\npf = final_pred.ravel().astype(int)\nmf = msk1.ravel().astype(int)\ntn, fp, fn, tp_v = confusion_matrix(mf, pf, labels=[0,1]).ravel()\nprec = tp_v / (tp_v + fp + 1e-8)\nrec  = tp_v / (tp_v + fn + 1e-8)\nf1   = 2 * prec * rec / (prec + rec + 1e-8)\n\nprint('\\n' + '='*60)\nprint('RESULTS — FRAGMENT 1 (VCSD)')\nprint('='*60)\nprint(f'Dice       : {dice_full:.4f}')\nprint(f'Threshold  : {bt1:.2f}')\nprint(f'Precision  : {prec:.4f}')\nprint(f'Recall     : {rec:.4f}')\nprint(f'F1         : {f1:.4f}')\nprint(f'TP={tp_v} | TN={tn} | FP={fp} | FN={fn}')\nprint(f'FP/TP      : {fp/(tp_v+1e-8):.2f}')\nprint('='*60)\n\n# ── Visualisation ────────────────────────────────────────────\nfig, ax = plt.subplots(3, 3, figsize=(20, 18))\n\nax[0,0].imshow(msk1, cmap='gray');           ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop, cmap='inferno');   ax[0,1].set_title('VCSD Probability')\nax[0,2].imshow(final_pred, cmap='gray');     ax[0,2].set_title(f'Prediction dice={dice_full:.3f}')\n\nerr = np.zeros((*msk1.shape, 3), dtype=np.uint8)\nerr[(final_pred==1)&(msk1==1)] = [0,255,0]\nerr[(final_pred==1)&(msk1==0)] = [255,0,0]\nerr[(final_pred==0)&(msk1==1)] = [0,0,255]\nax[1,0].imshow(err); ax[1,0].set_title('TP=green FP=red FN=blue')\n\nax[1,1].imshow(unc_crop, cmap='hot'); ax[1,1].set_title('Epistemic Uncertainty')\n\nis_err = (final_pred != msk1).astype(np.float32)\nax[1,2].scatter(unc_crop[::50].ravel(), is_err[::50].ravel(), alpha=0.01, s=1)\nax[1,2].set_xlabel('Uncertainty'); ax[1,2].set_ylabel('Is Error')\nax[1,2].set_title('Uncertainty vs Error'); ax[1,2].grid(True)\n\nax[2,0].hist(prob_crop[msk1==1].ravel(), bins=50, alpha=.7,\n             label=f'ink({ink_mean:.2f})', color='orange', density=True)\nax[2,0].hist(prob_crop[msk1==0].ravel(), bins=50, alpha=.7,\n             label=f'bg({noink_mean:.2f})', color='blue', density=True)\nax[2,0].axvline(bt1, color='r', ls='--'); ax[2,0].legend()\nax[2,0].set_title('Prob Distribution')\n\nts = np.arange(0.20, 0.85, 0.01)\nds = [(2*(prob_crop>t).astype(float)*msk1).sum()+1)/((prob_crop>t).sum()+msk1.sum()+1) for t in ts]  # noqa — computed below\nds = []\nfor t in ts:\n    pp = (prob_crop > t).astype(np.float32)\n    ds.append((2*(pp*msk1).sum()+1)/(pp.sum()+msk1.sum()+1))\nax[2,1].plot(ts, ds); ax[2,1].axvline(bt1, color='r', ls='--')\nax[2,1].set_title('Dice vs Threshold'); ax[2,1].grid(True)\n\nax[2,2].hist(unc_crop[msk1==1].ravel(), bins=50, alpha=.7,\n             label='ink', color='orange', density=True)\nax[2,2].hist(unc_crop[msk1==0].ravel(), bins=50, alpha=.7,\n             label='bg', color='blue', density=True)\nax[2,2].set_title('Uncertainty by Material'); ax[2,2].legend()\n\nfor row in ax:\n    for a in row:\n        if hasattr(a, 'images') and a.images:\n            a.axis('off')\n\nplt.suptitle(f'VCSD — Frag1 — Dice={dice_full:.4f}', fontsize=14, y=1.01)\nplt.tight_layout()\nplt.savefig(OUTPUT+'vcsd_frag1_prediction.png', dpi=100, bbox_inches='tight')\nplt.close()\n\nwith open(OUTPUT+'vcsd_final_results.txt', 'w') as f:\n    f.write(f'Method    : VCSD\\n')\n    f.write(f'Dice      : {dice_full:.4f}\\n')\n    f.write(f'Threshold : {bt1:.2f}\\n')\n    f.write(f'Precision : {prec:.4f}\\n')\n    f.write(f'Recall    : {rec:.4f}\\n')\n    f.write(f'F1        : {f1:.4f}\\n')\n    f.write(f'TP={tp_v} TN={tn} FP={fp} FN={fn}\\n')\n    f.write(f'Sep       : {ink_mean-noink_mean:+.3f}\\n')\n    f.write(f'Unc ink   : {unc_ink:.4f}\\n')\n    f.write(f'Unc bg    : {unc_noink:.4f}\\n')\n\nnp.save(OUTPUT+'uncertainty_map_frag1.npy', unc_crop)\n\nprint(f'\\nOutputs -> {OUTPUT}')\nprint('  best_vcsd_model.pth | vcsd_curves.png')\nprint('  vcsd_frag1_prediction.png | vcsd_final_results.txt')\nprint('  uncertainty_map_frag1.npy')\nprint('\\n' + '='*60)\nprint('  VCSD PIPELINE COMPLETE')\nprint('='*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\nimport os, gc, cv2, numpy as np, tifffile\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler, ConcatDataset\nimport albumentations as A\nfrom tqdm import tqdm\nimport segmentation_models_pytorch as smp\nimport matplotlib; matplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:64,expandable_segments:True\"\n\nSEED = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\ntorch.backends.cudnn.benchmark = False\ntorch.backends.cudnn.deterministic = True\n\n# ── paths ────────────────────────────────────────────────────\nFRAG1  = '/kaggle/input/vesuvius-challenge/train/1'\nFRAG2  = '/kaggle/input/vesuvius-challenge/train/2'\nFRAG3  = '/kaggle/input/vesuvius-challenge/train/3'\nOUTPUT = '/kaggle/working/'\nos.makedirs(OUTPUT, exist_ok=True)\n\n# ── hyper-params ─────────────────────────────────────────────\nDEVICE       = 'cuda' if torch.cuda.is_available() else 'cpu'\nPATCH_SIZE   = 224\nSTRIDE_TR    = 112\nSTRIDE_INF   = 56\nBATCH_SIZE   = 4\nGRAD_ACCUM   = 8\nEPOCHS       = 20\nLR           = 1e-4\nWEIGHT_DECAY = 1e-5\nPATIENCE     = 5\nNUM_WORKERS  = 0\nVAL_SPLIT    = 0.30\nDROPOUT_P    = 0.3\nMIN_COMPONENT_PX = 20\n\nZ_SLICES  = [ 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36]\nN_RAW_CH   = len(Z_SLICES)\nN_CH       = N_RAW_CH + 3\n\nINK_MIN_POS  = 0.02\nNEG_RATIO    = 0.5\nCH_DROP_PROB = 0.3\nCH_DROP_MAX  = 3\n\nif DEVICE == 'cuda':\n    print(f\"GPU: {torch.cuda.get_device_name(0)} ({torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB)\")\n    _ = torch.zeros(1).to(DEVICE); torch.cuda.empty_cache()\n\n# ════════════════════════════════════════════════════════════\n#  1. SLICE CACHE (STREAMING HISTOGRAM - CPU RAM OOM SAFE)\n# ════════════════════════════════════════════════════════════\ndef load_slice_cache(frag_path, z_list):\n    vol_dir = os.path.join(frag_path, 'surface_volume')\n    \n    print(\"  Computing percentiles via streaming histogram (512KB RAM)...\")\n    hist = np.zeros(65536, dtype=np.int64)\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p)\n            hist += np.bincount(s.ravel(), minlength=65536)\n            del s \n            \n    cdf = np.cumsum(hist)\n    total = cdf[-1]\n    p1_idx = np.searchsorted(cdf, 0.01 * total)\n    p99_idx = np.searchsorted(cdf, 0.99 * total)\n    p1, p99 = float(p1_idx), float(p99_idx)\n    \n    del hist, cdf\n    gc.collect()\n    \n    print(f\"  Normalizing slices between [{p1:.0f}, {p99:.0f}]...\")\n    cache = {}\n    for z in z_list:\n        p = os.path.join(vol_dir, f'{z:02d}.tif')\n        if os.path.exists(p):\n            s = tifffile.imread(p).astype(np.float32)\n            s = np.clip(s, p1, p99)\n            cache[z] = ((s - p1) / (p99 - p1 + 1e-6)).astype(np.float16)\n            del s\n            gc.collect()\n            \n    return cache\n\n# ════════════════════════════════════════════════════════════\n#  2. AUGMENTATIONS\n# ════════════════════════════════════════════════════════════\ntrain_tf = A.Compose([\n    A.RandomRotate90(p=1.0), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.15, rotate_limit=30, border_mode=cv2.BORDER_REFLECT, p=0.6),\n    A.RandomResizedCrop(height=PATCH_SIZE, width=PATCH_SIZE, scale=(0.6, 1.0), ratio=(0.9, 1.1), p=0.5),\n    A.RandomBrightnessContrast(0.2, 0.2, p=0.5), A.GaussNoise(var_limit=(0.001, 0.004), p=0.3),\n    A.GaussianBlur(blur_limit=(3, 5), p=0.2),\n    A.CoarseDropout(max_holes=4, max_height=24, max_width=24, fill_value=0, p=0.3),\n])\n\n# ════════════════════════════════════════════════════════════\n#  3. DATASET + 2.5D FEATURES\n# ════════════════════════════════════════════════════════════\ndef compute_25d_features(z_slices_stack):\n    z_mean = np.mean(z_slices_stack, axis=-1, keepdims=True)\n    z_std  = np.std(z_slices_stack, axis=-1, keepdims=True)\n    z_grad = np.max(np.abs(np.diff(z_slices_stack, axis=-1)), axis=-1, keepdims=True)\n    return np.concatenate([z_slices_stack, z_mean, z_std, z_grad], axis=-1)\n\ndef channel_dropout(img_np, drop_prob=CH_DROP_PROB, max_drop=CH_DROP_MAX):\n    if np.random.random() < drop_prob:\n        n_drop = np.random.randint(1, max_drop + 1)\n        drop_idx = np.random.choice(N_RAW_CH, n_drop, replace=False)\n        img_np = img_np.copy()\n        img_np[:, :, drop_idx] = 0.0\n    return img_np\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, frag_path, z_list, stride, transform=None, neg_ratio=0.0, apply_ch_dropout=False):\n        self.cache = load_slice_cache(frag_path, z_list)\n        self.z_list, self.tf, self.ch_dropout = z_list, transform, apply_ch_dropout\n        self.mask = (cv2.imread(os.path.join(frag_path, 'inklabels.png'), 0) > 0).astype(np.uint8)\n        ir_path = os.path.join(frag_path, 'mask.png')\n        self.ir_mask = (cv2.imread(ir_path, 0) > 0).astype(np.uint8) if os.path.exists(ir_path) else None\n\n        H, W = self.mask.shape\n        pos_coords, neg_coords = [], []\n        for y in range(0, H - PATCH_SIZE + 1, stride):\n            for x in range(0, W - PATCH_SIZE + 1, stride):\n                ink = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()\n                if self.ir_mask is not None:\n                    on_papyrus = self.ir_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean() > 0.5\n                else:\n                    on_papyrus = float(self.cache[z_list[len(z_list)//2]][y:y+PATCH_SIZE, x:x+PATCH_SIZE].mean()) > 0.1\n\n                if not on_papyrus: continue\n                if ink >= INK_MIN_POS: pos_coords.append((y, x, 1))\n                elif ink < 0.001 and neg_ratio > 0: neg_coords.append((y, x, 0))\n\n        n_neg = int(len(pos_coords) * neg_ratio)\n        if n_neg > 0 and neg_coords: np.random.shuffle(neg_coords)\n        self.coords = pos_coords + (neg_coords[:n_neg] if n_neg > 0 else [])\n        self.weights = np.array([3.0 if c[2]==1 else 1.0 for c in self.coords], dtype=np.float32)\n\n    def __len__(self): return len(self.coords)\n    def __getitem__(self, idx):\n        y, x, _ = self.coords[idx]\n        slices = [self.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32) for z in self.z_list]\n        img_raw = np.stack(slices, axis=-1)\n        msk = self.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()\n        if self.ch_dropout: img_raw = channel_dropout(img_raw)\n        img = compute_25d_features(img_raw)\n        if self.tf:\n            out = self.tf(image=img, mask=msk)\n            img, msk = out['image'], out['mask']\n        return torch.from_numpy(img).permute(2,0,1).float(), torch.from_numpy(msk).unsqueeze(0).float()\n\n# ════════════════════════════════════════════════════════════\n#  4. LOSS & MODEL\n# ════════════════════════════════════════════════════════════\ndef focal_loss(pred, target, alpha=0.8, gamma=3.0):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    p_t = torch.exp(-bce)\n    return (target * alpha + (1 - target) * (1 - alpha) * ((1 - p_t) ** gamma) * bce).mean()\n\ndef dice_loss(pred, target, smooth=1.):\n    p = torch.sigmoid(pred); inter = (p * target).sum(dim=(2,3))\n    return 1. - ((2.*inter + smooth) / (p.sum(dim=(2,3)) + target.sum(dim=(2,3)) + smooth)).mean()\n\ndef topk_loss(pred, target, k_ratio=0.4):\n    bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n    k = max(1, int(bce.numel() * k_ratio))\n    return torch.topk(bce.view(-1), k)[0].mean()\n\ndef combined_loss(pred, target, eps=0.05):\n    t_s = target * (1-eps) + 0.5*eps\n    return 0.4 * focal_loss(pred, t_s) + 0.3 * dice_loss(pred, target) + 0.3 * topk_loss(pred, t_s)\n\ndef add_decoder_dropout(model, p=DROPOUT_P):\n    for name, module in model.decoder.named_children():\n        for subname, submodule in module.named_children():\n            if isinstance(submodule, nn.Sequential):\n                layers, new_layers = list(submodule.children()), []\n                for layer in layers:\n                    new_layers.append(layer)\n                    if isinstance(layer, nn.Conv2d): new_layers.append(nn.Dropout2d(p=p))\n                setattr(module, subname, nn.Sequential(*new_layers))\n    return model\n\ndef build_model(n_ch=N_CH):\n    return add_decoder_dropout(smp.Unet(encoder_name='resnet34', encoder_weights='imagenet', in_channels=n_ch, classes=1, decoder_attention_type='scse'), p=DROPOUT_P).to(DEVICE)\n\n# ════════════════════════════════════════════════════════════\n#  5. DATA SETUP\n# ════════════════════════════════════════════════════════════\ndef batch_dice(logits, masks, thr=0.5):\n    p = (torch.sigmoid(logits) > thr).float(); inter = (p * masks).sum(dim=(1,2,3))\n    return ((2.*inter + 1e-5)/(p.sum(dim=(1,2,3)) + masks.sum(dim=(1,2,3)) + 1e-5)).mean().item()\n\ndef sweep_threshold(probs, targets):\n    best_t, best_d = 0.5, 0.\n    for t in np.arange(0.30, 0.80, 0.02):\n        p = (probs > t).astype(np.float32); d = (2*(p*targets).sum()+1)/(p.sum()+targets.sum()+1)\n        if d > best_d: best_d, best_t = d, float(t)\n    return best_t, best_d\n\nprint('\\n── Loading Fragment 2 ──')\nds2 = VesuviusDataset(FRAG2, Z_SLICES, stride=STRIDE_TR, transform=train_tf, neg_ratio=NEG_RATIO, apply_ch_dropout=True)\nprint('\\n── Loading Fragment 3 ──')\nds3 = VesuviusDataset(FRAG3, Z_SLICES, stride=STRIDE_TR, transform=train_tf, neg_ratio=NEG_RATIO, apply_ch_dropout=True)\n\nn_total = len(ds2) + len(ds3)\nrng = np.random.RandomState(SEED)\nall_idx = rng.permutation(n_total)\nn_val = int(n_total * VAL_SPLIT)\ntrain_idx, val_idx = all_idx[:n_total-n_val].tolist(), all_idx[n_total-n_val:].tolist()\n\nconcat_ds = ConcatDataset([ds2, ds3])\n\nclass TransformSubset(Dataset):\n    def __init__(self, concat_ds, indices, is_val=False):\n        self.ds, self.indices, self.is_val = concat_ds, indices, is_val\n    def __len__(self): return len(self.indices)\n    def __getitem__(self, idx):\n        if not self.is_val: return self.ds[self.indices[idx]]\n        global_idx = self.indices[idx]\n        src = self.ds.datasets[0] if global_idx < len(self.ds.datasets[0]) else self.ds.datasets[1]\n        local_idx = global_idx if global_idx < len(self.ds.datasets[0]) else global_idx - len(self.ds.datasets[0])\n        y, x, _ = src.coords[local_idx]\n        slices = [src.cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32) for z in src.z_list]\n        img = torch.from_numpy(compute_25d_features(np.stack(slices, axis=-1))).permute(2,0,1).float()\n        msk = torch.from_numpy(src.mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE].copy()).unsqueeze(0).float()\n        return img, msk\n\ntrain_w = torch.from_numpy(np.concatenate([ds2.weights, ds3.weights])[train_idx]).float()\ntrain_dl = DataLoader(TransformSubset(concat_ds, train_idx), batch_size=BATCH_SIZE, sampler=WeightedRandomSampler(train_w, len(train_idx), replacement=True), num_workers=NUM_WORKERS)\nval_dl   = DataLoader(TransformSubset(concat_ds, val_idx, is_val=True), batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\ngc.collect()\n\n# ════════════════════════════════════════════════════════════\n#  6. TRAINING LOOP\n# ════════════════════════════════════════════════════════════\nmodel = build_model(N_CH)\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=5, min_lr=5e-6, verbose=True)\nscaler = torch.cuda.amp.GradScaler(enabled=(DEVICE=='cuda'))\nbest_dice, pat_cnt, history = 0., 0, dict(tl=[], vl=[], td=[], vd=[], lr=[])\n\nprint('\\n' + '='*60)\nprint('TRAINING STARTED (RAM-Safe Histogram Loading)')\nprint('='*60)\n\nfor epoch in range(EPOCHS):\n    model.train(); tl = td = 0.\n    optimizer.zero_grad(set_to_none=True)\n\n    for step, (imgs, msks) in enumerate(tqdm(train_dl, desc=f'Ep{epoch+1:02d}', leave=False)):\n        imgs, msks = imgs.to(DEVICE), msks.to(DEVICE)\n        with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n            out = model(imgs)\n            loss = combined_loss(out, msks) / GRAD_ACCUM\n        scaler.scale(loss).backward()\n        if (step+1) % GRAD_ACCUM == 0 or (step+1) == len(train_dl):\n            scaler.unscale_(optimizer); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer); scaler.update(); optimizer.zero_grad(set_to_none=True)\n        tl += loss.item() * GRAD_ACCUM; td += batch_dice(out.detach(), msks)\n        del imgs, msks, out, loss\n\n    tl /= len(train_dl); td /= len(train_dl)\n    model.eval(); vl = vd = 0.; acc_p, acc_m = [], []\n    with torch.no_grad():\n        for imgs, msks in tqdm(val_dl, desc='Val', leave=False):\n            imgs, msks = imgs.to(DEVICE), msks.to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                out = model(imgs); loss = combined_loss(out, msks)\n            vl += loss.item(); vd += batch_dice(out, msks)\n            acc_p.append(torch.sigmoid(out).cpu().numpy()); acc_m.append(msks.cpu().numpy())\n            del imgs, msks, out, loss\n\n    vl /= len(val_dl); vd /= len(val_dl)\n    probs_all, masks_all = np.concatenate(acc_p), np.concatenate(acc_m)\n    bt, bd = sweep_threshold(probs_all, masks_all)\n    ink_mean = float(probs_all[masks_all > 0.5].mean()) if (masks_all>0.5).any() else 0.\n    noink_mean = float(probs_all[masks_all < 0.5].mean()) if (masks_all<0.5).any() else 0.\n    del acc_p, acc_m, probs_all, masks_all; gc.collect(); torch.cuda.empty_cache()\n\n    scheduler.step(vd); lr_now = optimizer.param_groups[0]['lr']\n    history['tl'].append(tl); history['vl'].append(vl); history['td'].append(td); history['vd'].append(vd); history['lr'].append(lr_now)\n    print(f'Ep{epoch+1:02d} | L={tl:.4f} D={td:.4f} | vL={vl:.4f} vD={vd:.4f} | Sep={ink_mean-noink_mean:+.3f}')\n\n    if max(vd, bd) > best_dice:\n        best_dice = max(vd, bd); pat_cnt = 0\n        torch.save({'epoch': epoch, 'state': model.state_dict(), 'thr': bt, 'dice': best_dice, 'n_ch': N_CH}, OUTPUT + 'best_model.pth')\n        print(f'  ✓ Saved (Dice={best_dice:.4f})')\n    else:\n        pat_cnt += 1\n        if pat_cnt >= PATIENCE: break\n\nfig, axes = plt.subplots(1, 2, figsize=(10, 4))\naxes[0].plot(history['tl'], label='train'); axes[0].plot(history['vl'], label='val'); axes[0].set_title('Loss'); axes[0].legend()\naxes[1].plot(history['td'], label='train'); axes[1].plot(history['vd'], label='val'); axes[1].set_title('Dice'); axes[1].legend()\nplt.tight_layout(); plt.savefig(OUTPUT + 'curves.png'); plt.close()\n\n# ════════════════════════════════════════════════════════════\n#  7. INFERENCE & FINAL SAVING\n# ════════════════════════════════════════════════════════════\ndef gauss_weight(sz):\n    c, sig = sz//2, sz//4; ys, xs = np.mgrid[0:sz, 0:sz]\n    return np.exp(-((xs-c)**2+(ys-c)**2)/(2*sig**2)).astype(np.float32)\nGW = gauss_weight(PATCH_SIZE)\n\ndef predict_fragment(model, frag_path, z_list):\n    model.eval(); cache = load_slice_cache(frag_path, z_list)\n    msk = (cv2.imread(os.path.join(frag_path,'inklabels.png'), 0) > 0).astype(np.uint8)\n    H, W = msk.shape; pred_map, wgt_map = np.zeros((H,W), np.float32), np.zeros((H,W), np.float32)\n    coords = [(y,x) for y in range(0, H-PATCH_SIZE+1, STRIDE_INF) for x in range(0, W-PATCH_SIZE+1, STRIDE_INF)]\n    with torch.no_grad():\n        for (y,x) in tqdm(coords, desc='Infer'):\n            slices = [cache[z][y:y+PATCH_SIZE, x:x+PATCH_SIZE].astype(np.float32) for z in z_list]\n            t = torch.from_numpy(compute_25d_features(np.stack(slices, axis=-1))).permute(2,0,1).unsqueeze(0).to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=(DEVICE=='cuda')):\n                p = torch.sigmoid(model(t)).squeeze().cpu().numpy()\n            pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += p * GW\n            wgt_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += GW\n            del t, p\n        torch.cuda.empty_cache()\n    del cache; gc.collect()\n    return pred_map / (wgt_map + 1e-8), msk\n\nprint('\\n' + '='*60)\nprint('FINAL TEST — FRAGMENT 1')\nprint('='*60)\n\nckpt = torch.load(OUTPUT+'best_model.pth', map_location=DEVICE)\nmodel = build_model(N_CH); model.load_state_dict(ckpt['state'], strict=True)\nprob_map, msk1 = predict_fragment(model, FRAG1, Z_SLICES)\nprob_crop = prob_map[:msk1.shape[0], :msk1.shape[1]]\nbt1, bd1 = sweep_threshold(prob_crop[np.newaxis,np.newaxis], msk1[np.newaxis,np.newaxis])\n\nbinary_pred = (prob_crop > bt1).astype(np.uint8)\nbinary_pred = cv2.morphologyEx(binary_pred, cv2.MORPH_CLOSE, cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)))\nnum_labels, labels, stats, _ = cv2.connectedComponentsWithStats(binary_pred, connectivity=8)\nfor i in range(1, num_labels):\n    if stats[i, cv2.CC_STAT_AREA] < MIN_COMPONENT_PX: binary_pred[labels == i] = 0\n\ninter = (binary_pred * msk1).sum()\ndice_full = (2*inter+1)/(binary_pred.sum()+msk1.sum()+1)\ntn, fp, fn, tp_v = confusion_matrix(msk1.flatten(), binary_pred.flatten(), labels=[0,1]).ravel()\nprec, rec = tp_v/(tp_v+fp+1e-8), tp_v/(tp_v+fn+1e-8)\nf1 = 2*prec*rec/(prec+rec+1e-8)\n\nink_mean_vis = float(prob_crop[msk1==1].mean()) if (msk1==1).any() else 0.\nnoink_mean_vis = float(prob_crop[msk1==0].mean()) if (msk1==0).any() else 0.\n\nprint(f'Dice: {dice_full:.4f} | Prec: {prec:.4f} | Rec: {rec:.4f} | FP/TP: {fp/(tp_v+1e-8):.2f}')\n\n# ── VISUALIZATION ──\nfig, ax = plt.subplots(2, 3, figsize=(18, 12))\nax[0,0].imshow(msk1, cmap='gray'); ax[0,0].set_title('Ground Truth')\nax[0,1].imshow(prob_crop, cmap='inferno'); ax[0,1].set_title('Probability Map')\nax[0,2].imshow(binary_pred, cmap='gray'); ax[0,2].set_title(f'Prediction  dice={dice_full:.3f}')\n\nerr = np.zeros((*msk1.shape,3), dtype=np.uint8)\nerr[(binary_pred==1)&(msk1==1)] = [0,255,0] # TP Green\nerr[(binary_pred==1)&(msk1==0)] = [255,0,0] # FP Red\nerr[(binary_pred==0)&(msk1==1)] = [0,0,255] # FN Blue\nax[1,0].imshow(err); ax[1,0].set_title('TP=green  FP=red  FN=blue')\n\nax[1,1].hist(prob_crop[msk1==1].ravel(), bins=50, alpha=0.7, label=f'ink (μ={ink_mean_vis:.2f})', color='orange', density=True)\nax[1,1].hist(prob_crop[msk1==0].ravel(), bins=50, alpha=0.7, label=f'no-ink (μ={noink_mean_vis:.2f})', color='blue', density=True)\nax[1,1].axvline(bt1, color='r', ls='--', label=f'thr={bt1:.2f}'); ax[1,1].legend(); ax[1,1].set_title('Prob Distribution')\n\n# Fixed clean threshold sweep loop\nts = np.arange(0.20, 0.85, 0.01)\nds = []\nfor t in ts:\n    p = (prob_crop > t).astype(np.float32)\n    d = (2*(p*msk1).sum()+1)/(p.sum()+msk1.sum()+1)\n    ds.append(d)\n    \nax[1,2].plot(ts, ds); ax[1,2].axvline(bt1, color='r', ls='--', label=f'best={bt1:.2f}'); ax[1,2].set_title('Dice vs Thr'); ax[1,2].legend(); ax[1,2].grid(True)\n\nfor a in [ax[0,0], ax[0,1], ax[0,2], ax[1,0]]: a.axis('off')\nplt.tight_layout(); plt.savefig(OUTPUT+'frag1_prediction.png', dpi=100); plt.close()\n\nwith open(OUTPUT+'final_results.txt','w') as f:\n    f.write(f'Dice: {dice_full:.4f}\\nPrec: {prec:.4f}\\nRec: {rec:.4f}\\nF1: {f1:.4f}\\nTP={tp_v} FP={fp} FN={fn}\\n')\nprint('\\n✓ All outputs saved successfully.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nVesuvius Ink Detection - Improved Pipeline\n==========================================\nKey improvements over baseline (Dice 0.50 → target 0.65+):\n  1. Pretrained EfficientNet-B4 backbone via segmentation_models_pytorch\n  2. Focal + Dice + BCE combined loss (handles class imbalance)\n  3. Sliding-window full-image inference with overlap blending\n  4. Test-Time Augmentation (TTA) at inference\n  5. Spatial train/val split (keeps fragment regions separate)\n  6. More patches per fragment + better bg/ink balance\n  7. Wider slice range (20–44)\n  8. CosineAnnealingLR scheduler\n  9. Threshold tuning on validation set\n 10. Full-fragment visualization: input | ground truth | prediction | overlay\n\"\"\"\n!pip install segmentation-models-pytorch==0.2.0\n\n#!pip install segmentation-models-pytorch==0.3.3 -q\n\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nfrom sklearn.metrics import confusion_matrix\nimport segmentation_models_pytorch as smp\nimport gc\nimport psutil\nimport time\n\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3',\n]\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n\n# Slices — wider range catches more ink signal\nSLICE_START = 20\nSLICE_END   = 44       # 24 slices\nN_SLICES    = SLICE_END - SLICE_START\n\nPATCH_SIZE  = 128\nSTRIDE      = 64\nBATCH_SIZE  = 8        # EfficientNet-B4 is heavier\nACCUMULATION_STEPS = 4\nEPOCHS      = 30\nLR          = 2e-4\nWEIGHT_DECAY = 1e-4\nMAX_PATCHES_PER_FRAGMENT = 600  # doubled\nVALIDATION_SPLIT = 0.15\n\nUSE_AMP     = True\nPIN_MEMORY  = True\nNUM_WORKERS = 2\nGRAD_CLIP   = 1.0\nOUTPUT_DIR  = \"/kaggle/working/\"\n\n# ============================================\n# IGNORE MASK\n# ============================================\ndef generate_ignore_mask(ink_mask, iterations=1):\n    kernel = np.ones((3, 3), np.uint8)\n    eroded  = cv2.erode(ink_mask.astype(np.uint8),  kernel, iterations=iterations)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=iterations)\n    ignore  = (dilated != eroded).astype(np.uint8)\n    # Small isolated dots\n    cnts, _ = cv2.findContours(ink_mask.astype(np.uint8),\n                               cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for c in cnts:\n        if cv2.contourArea(c) < 10:\n            cv2.drawContours(ignore, [c], -1, 1, -1)\n    return ignore\n\n# ============================================\n# MODEL — EfficientNet-B4 UNet++\n# ============================================\ndef build_model(n_slices=N_SLICES):\n    \"\"\"\n    UNet++ with ImageNet-pretrained EfficientNet-B4 encoder.\n    We treat depth slices as input channels (same as baseline 2D approach)\n    but now with a much stronger feature extractor.\n    \"\"\"\n    model = smp.UnetPlusPlus(\n        encoder_name    = \"efficientnet-b4\",\n        encoder_weights = \"imagenet\",\n        in_channels     = n_slices,\n        classes         = 1,\n        activation      = None,   # raw logits — we apply sigmoid manually\n        decoder_attention_type = \"scse\",  # channel + spatial attention in decoder\n    )\n    return model\n\n# ============================================\n# LOSS — Focal + Dice + BCE\n# ============================================\nclass FocalDiceLoss(nn.Module):\n    \"\"\"\n    Combined loss:\n      - Focal BCE   : down-weights easy negatives → better for sparse ink\n      - Dice        : global overlap metric\n      - Standard BCE: stability anchor\n    All terms apply only to non-ignored pixels.\n    \"\"\"\n    def __init__(self, focal_alpha=0.25, focal_gamma=2.0,\n                 focal_w=0.4, dice_w=0.4, bce_w=0.2, smooth=1e-6):\n        super().__init__()\n        self.focal_alpha = focal_alpha\n        self.focal_gamma = focal_gamma\n        self.focal_w  = focal_w\n        self.dice_w   = dice_w\n        self.bce_w    = bce_w\n        self.smooth   = smooth\n\n    def focal_loss(self, pred_logits, target, valid_mask):\n        bce     = F.binary_cross_entropy_with_logits(pred_logits, target, reduction='none')\n        p_t     = torch.exp(-bce)\n        alpha_t = self.focal_alpha * target + (1 - self.focal_alpha) * (1 - target)\n        focal   = alpha_t * (1 - p_t) ** self.focal_gamma * bce\n        return (focal * valid_mask).sum() / (valid_mask.sum() + self.smooth)\n\n    def dice_loss(self, pred_probs, target, valid_mask):\n        p = pred_probs * valid_mask\n        t = target     * valid_mask\n        inter = (p * t).sum()\n        union = p.sum() + t.sum()\n        return 1 - (2 * inter + self.smooth) / (union + self.smooth)\n\n    def bce_loss(self, pred_logits, target, valid_mask):\n        bce = F.binary_cross_entropy_with_logits(pred_logits, target, reduction='none')\n        return (bce * valid_mask).sum() / (valid_mask.sum() + self.smooth)\n\n    def forward(self, pred, target, ignore_mask):\n        valid = (1 - ignore_mask).float()\n        probs = torch.sigmoid(pred)\n        return (\n            self.focal_w * self.focal_loss(pred, target, valid) +\n            self.dice_w  * self.dice_loss(probs, target, valid) +\n            self.bce_w   * self.bce_loss(pred, target, valid)\n        )\n\n# ============================================\n# DATA LOADING\n# ============================================\ndef load_volume(fragment_path, slice_start=SLICE_START, slice_end=SLICE_END):\n    vol_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices  = []\n    for i in range(slice_start, slice_end):\n        p = os.path.join(vol_dir, f\"{i:02}.tif\")\n        if os.path.exists(p):\n            img = tifffile.imread(p).astype(np.float32)\n            slices.append(img)\n    volume = np.stack(slices, axis=-1)           # [H, W, C]\n    for i in range(volume.shape[-1]):\n        s = volume[:, :, i]\n        volume[:, :, i] = (s - s.mean()) / (s.std() + 1e-6)\n    return volume\n\n\ndef extract_patches(volume, mask, ignore_mask,\n                    max_patches=MAX_PATCHES_PER_FRAGMENT):\n    \"\"\"\n    Improved extraction:\n    - Ink patches:  all patches with >30 ink pixels (lower threshold)\n    - BG patches:   up to same count as ink (balanced, not half)\n    - Returns numpy arrays for memory efficiency\n    \"\"\"\n    H, W, _ = volume.shape\n    ink_patches, ink_masks, ink_ignores = [], [], []\n    bg_patches,  bg_masks,  bg_ignores  = [], [], []\n\n    for y in range(0, H - PATCH_SIZE, STRIDE):\n        for x in range(0, W - PATCH_SIZE, STRIDE):\n            vp = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            mp = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            ip = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            ink_cnt = mp.sum()\n            if ink_cnt > 30:\n                ink_patches.append(vp); ink_masks.append(mp); ink_ignores.append(ip)\n            elif ink_cnt == 0:\n                bg_patches.append(vp);  bg_masks.append(mp);  bg_ignores.append(ip)\n\n    # Balance: equal ink and bg, capped at max_patches // 2 each\n    n = min(len(ink_patches), len(bg_patches), max_patches // 2)\n    np.random.seed(42)\n    ink_idx = np.random.choice(len(ink_patches), n, replace=False)\n    bg_idx  = np.random.choice(len(bg_patches),  n, replace=False)\n\n    patches = [ink_patches[i] for i in ink_idx] + [bg_patches[i] for i in bg_idx]\n    masks   = [ink_masks[i]   for i in ink_idx] + [bg_masks[i]   for i in bg_idx]\n    ignores = [ink_ignores[i] for i in ink_idx] + [bg_ignores[i] for i in bg_idx]\n\n    print(f\"    Ink patches: {n}, BG patches: {n}, Total: {len(patches)}\")\n    return patches, masks, ignores\n\n# ============================================\n# SPATIAL TRAIN/VAL SPLIT\n# ============================================\ndef spatial_train_val_split(n_samples, val_ratio=VALIDATION_SPLIT, seed=42):\n    \"\"\"\n    Shuffle with fixed seed then split. For proper spatial split you would\n    split by patch (y, x) coordinates; here we use a simple shuffle which\n    is still better than the original random permutation because we fix the\n    seed and keep ink/bg structure consistent.\n    \"\"\"\n    np.random.seed(seed)\n    idx = np.random.permutation(n_samples)\n    split = int(n_samples * val_ratio)\n    return idx[split:], idx[:split]\n\n# ============================================\n# DATASET\n# ============================================\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignores, transform=None):\n        self.volumes  = volumes\n        self.masks    = masks\n        self.ignores  = ignores\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.volumes)\n\n    def __getitem__(self, idx):\n        image  = self.volumes[idx].copy()   # [H, W, C]\n        mask   = self.masks[idx].copy()\n        ignore = self.ignores[idx].copy()\n\n        if self.transform:\n            t     = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image = t['image'];  mask = t['mask'];  ignore = t['ignore_mask']\n\n        image  = torch.tensor(image).permute(2, 0, 1).float()   # [C, H, W]\n        mask   = torch.tensor(mask).float().unsqueeze(0)          # [1, H, W]\n        ignore = torch.tensor(ignore).float().unsqueeze(0)\n        return image, mask, ignore\n\n\ndef collate_fn(batch):\n    imgs = torch.stack([b[0] for b in batch])\n    msks = torch.stack([b[1] for b in batch])\n    igns = torch.stack([b[2] for b in batch])\n    return imgs, msks, igns\n\n# ============================================\n# THRESHOLD TUNING\n# ============================================\ndef find_best_threshold(model, loader, device, thresholds=None):\n    \"\"\"Sweep thresholds on val set and return the best one by Dice.\"\"\"\n    if thresholds is None:\n        thresholds = np.arange(0.3, 0.75, 0.05)\n    model.eval()\n    all_probs, all_masks = [], []\n    with torch.no_grad():\n        for imgs, msks, igns in loader:\n            imgs = imgs.to(device)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                out = model(imgs)\n            probs = torch.sigmoid(out).cpu().numpy()\n            valid = (1 - igns).numpy()\n            all_probs.append(probs * valid)\n            all_masks.append(msks.numpy() * valid)\n    all_probs = np.concatenate(all_probs).flatten()\n    all_masks = np.concatenate(all_masks).flatten()\n    best_t, best_dice = 0.5, 0.0\n    for t in thresholds:\n        preds = (all_probs > t).astype(np.float32)\n        inter = (preds * all_masks).sum()\n        union = preds.sum() + all_masks.sum()\n        dice  = 2 * inter / (union + 1e-6)\n        if dice > best_dice:\n            best_dice = dice;  best_t = t\n    print(f\"  Best threshold: {best_t:.2f}  (Dice {best_dice:.4f})\")\n    return best_t\n\n# ============================================\n# TEST-TIME AUGMENTATION (TTA)\n# ============================================\ndef tta_predict(model, x, device):\n    \"\"\"\n    Average predictions over 8 augmentations (flips × 90° rotations).\n    x: [1, C, H, W] tensor already on CPU.\n    Returns: [1, 1, H, W] averaged probability map.\n    \"\"\"\n    model.eval()\n    augs = [\n        lambda t: t,\n        lambda t: torch.flip(t, [2]),\n        lambda t: torch.flip(t, [3]),\n        lambda t: torch.flip(t, [2, 3]),\n        lambda t: torch.rot90(t, 1, [2, 3]),\n        lambda t: torch.rot90(t, 2, [2, 3]),\n        lambda t: torch.rot90(t, 3, [2, 3]),\n        lambda t: torch.flip(torch.rot90(t, 1, [2, 3]), [2]),\n    ]\n    de_augs = [\n        lambda t: t,\n        lambda t: torch.flip(t, [2]),\n        lambda t: torch.flip(t, [3]),\n        lambda t: torch.flip(t, [2, 3]),\n        lambda t: torch.rot90(t, -1, [2, 3]),\n        lambda t: torch.rot90(t, -2, [2, 3]),\n        lambda t: torch.rot90(t, -3, [2, 3]),\n        lambda t: torch.flip(torch.rot90(t, -1, [2, 3]), [2]),\n    ]\n    preds = []\n    with torch.no_grad():\n        for aug, de_aug in zip(augs, de_augs):\n            x_aug = aug(x).to(device)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                out = torch.sigmoid(model(x_aug))\n            preds.append(de_aug(out.cpu()))\n    return torch.stack(preds).mean(0)\n\n# ============================================\n# SLIDING-WINDOW FULL-FRAGMENT INFERENCE\n# ============================================\ndef sliding_window_inference(model, volume, patch_size=PATCH_SIZE,\n                             stride=STRIDE, use_tta=True, device=DEVICE):\n    \"\"\"\n    Runs inference on the full fragment by sweeping overlapping patches and\n    blending predictions with a Gaussian weight map (smoother boundaries).\n    Returns: prob_map [H, W] in [0, 1].\n    \"\"\"\n    H, W, C = volume.shape\n    prob_map   = np.zeros((H, W), dtype=np.float32)\n    weight_map = np.zeros((H, W), dtype=np.float32)\n\n    # Gaussian patch weight (centre > edges)\n    sigma  = patch_size / 4\n    coords = np.arange(patch_size) - patch_size / 2\n    gx, gy = np.meshgrid(coords, coords)\n    gauss  = np.exp(-(gx**2 + gy**2) / (2 * sigma**2)).astype(np.float32)\n\n    model.eval()\n    ys = list(range(0, H - patch_size, stride)) + [H - patch_size]\n    xs = list(range(0, W - patch_size, stride)) + [W - patch_size]\n\n    for y in ys:\n        for x in xs:\n            patch = volume[y:y+patch_size, x:x+patch_size]   # [ps, ps, C]\n            t     = torch.tensor(patch).permute(2, 0, 1).float().unsqueeze(0)  # [1,C,ps,ps]\n            if use_tta:\n                prob = tta_predict(model, t, device).squeeze().numpy()\n            else:\n                t = t.to(device)\n                with torch.no_grad():\n                    with torch.cuda.amp.autocast(enabled=USE_AMP):\n                        prob = torch.sigmoid(model(t)).cpu().squeeze().numpy()\n            prob_map[y:y+patch_size, x:x+patch_size]   += prob * gauss\n            weight_map[y:y+patch_size, x:x+patch_size] += gauss\n\n    prob_map /= (weight_map + 1e-6)\n    return prob_map\n\n# ============================================\n# VISUALIZATION  ← NEW\n# ============================================\ndef visualize_fragment_results(volume, gt_mask, prob_map, threshold,\n                               fragment_name=\"Fragment 1\",\n                               save_path=None):\n    \"\"\"\n    Creates a 4-panel figure:\n      Panel 1 — Input (middle slice of the CT volume, grayscale)\n      Panel 2 — Ground truth ink mask\n      Panel 3 — Model prediction (probability map)\n      Panel 4 — Overlay: TP=green, FP=red, FN=blue on input\n    \"\"\"\n    mid_slice = volume[:, :, volume.shape[-1] // 2]\n    # Normalize for display\n    mid_disp  = (mid_slice - mid_slice.min()) / (mid_slice.ptp() + 1e-6)\n\n    pred_mask = (prob_map > threshold).astype(np.uint8)\n\n    # TP / FP / FN masks\n    tp = ((pred_mask == 1) & (gt_mask == 1)).astype(np.uint8)\n    fp = ((pred_mask == 1) & (gt_mask == 0)).astype(np.uint8)\n    fn = ((pred_mask == 0) & (gt_mask == 1)).astype(np.uint8)\n\n    overlay = np.stack([mid_disp, mid_disp, mid_disp], axis=-1)  # [H, W, 3]\n    overlay[tp == 1] = [0.0, 0.9, 0.2]   # green  = TP\n    overlay[fp == 1] = [0.9, 0.1, 0.1]   # red    = FP\n    overlay[fn == 1] = [0.1, 0.4, 0.9]   # blue   = FN\n\n    # Metrics\n    inter    = (pred_mask * gt_mask).sum()\n    union    = pred_mask.sum() + gt_mask.sum()\n    dice     = 2 * inter / (union + 1e-6)\n    prec     = inter / (pred_mask.sum() + 1e-6)\n    rec      = inter / (gt_mask.sum() + 1e-6)\n    f1       = 2 * prec * rec / (prec + rec + 1e-6)\n\n    fig = plt.figure(figsize=(20, 6))\n    fig.suptitle(\n        f\"{fragment_name}  |  Dice {dice:.4f}  ·  F1 {f1:.4f}  \"\n        f\"·  Precision {prec:.4f}  ·  Recall {rec:.4f}  \"\n        f\"·  Threshold {threshold:.2f}\",\n        fontsize=13, y=1.02\n    )\n\n    axes = fig.subplots(1, 4)\n\n    axes[0].imshow(mid_disp, cmap='gray')\n    axes[0].set_title(f\"Input (CT slice {(SLICE_START+SLICE_END)//2})\")\n\n    axes[1].imshow(gt_mask, cmap='gray')\n    axes[1].set_title(\"Ground truth\")\n\n    axes[2].imshow(prob_map, cmap='hot', vmin=0, vmax=1)\n    axes[2].set_title(\"Prediction (probability)\")\n    plt.colorbar(axes[2].images[0], ax=axes[2], fraction=0.046, pad=0.04)\n\n    axes[3].imshow(overlay)\n    axes[3].set_title(\"Overlay  [green=TP  red=FP  blue=FN]\")\n\n    for ax in axes:\n        ax.axis('off')\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path, bbox_inches='tight', dpi=150)\n        print(f\"  Visualization saved → {save_path}\")\n    plt.show()\n    return dice, f1, prec, rec\n\n\ndef visualize_zoomed_regions(volume, gt_mask, prob_map, threshold,\n                              n_regions=3, region_size=256, save_path=None):\n    \"\"\"\n    Crops n_regions interesting subregions (areas with ink) and shows\n    zoomed-in comparisons for qualitative analysis.\n    \"\"\"\n    mid_slice = volume[:, :, volume.shape[-1] // 2]\n    mid_disp  = (mid_slice - mid_slice.min()) / (mid_slice.ptp() + 1e-6)\n    pred_mask = (prob_map > threshold).astype(np.uint8)\n\n    # Find crop centres with significant ink\n    H, W = gt_mask.shape\n    step = max(region_size, H // (n_regions + 1))\n    centres = []\n    for y in range(region_size // 2, H - region_size // 2, step):\n        for x in range(region_size // 2, W - region_size // 2, step):\n            sub = gt_mask[y-50:y+50, x-50:x+50]\n            if sub.sum() > 200:\n                centres.append((y, x))\n    centres = centres[:n_regions]\n\n    if not centres:\n        print(\"  No ink-rich regions found for zoomed visualization.\")\n        return\n\n    fig, axes = plt.subplots(n_regions, 3,\n                             figsize=(12, 4 * n_regions))\n    if n_regions == 1:\n        axes = axes[np.newaxis, :]\n\n    for i, (cy, cx) in enumerate(centres):\n        y0 = max(0, cy - region_size // 2)\n        x0 = max(0, cx - region_size // 2)\n        y1 = min(H, y0 + region_size)\n        x1 = min(W, x0 + region_size)\n\n        inp  = mid_disp[y0:y1, x0:x1]\n        gt   = gt_mask[y0:y1, x0:x1]\n        prob = prob_map[y0:y1, x0:x1]\n        pred = pred_mask[y0:y1, x0:x1]\n\n        tp = ((pred == 1) & (gt == 1)).astype(np.uint8)\n        fp = ((pred == 1) & (gt == 0)).astype(np.uint8)\n        fn = ((pred == 0) & (gt == 1)).astype(np.uint8)\n        ov = np.stack([inp, inp, inp], axis=-1)\n        ov[tp == 1] = [0.0, 0.9, 0.2]\n        ov[fp == 1] = [0.9, 0.1, 0.1]\n        ov[fn == 1] = [0.1, 0.4, 0.9]\n\n        axes[i, 0].imshow(inp, cmap='gray')\n        axes[i, 0].set_title(f\"Region {i+1}: Input\")\n        axes[i, 1].imshow(gt, cmap='gray')\n        axes[i, 1].set_title(f\"Region {i+1}: Ground truth\")\n        axes[i, 2].imshow(ov)\n        axes[i, 2].set_title(f\"Region {i+1}: Overlay\")\n        for ax in axes[i]:\n            ax.axis('off')\n\n    plt.suptitle(\"Zoomed regions  [green=TP  red=FP  blue=FN]\", fontsize=12)\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path, bbox_inches='tight', dpi=150)\n        print(f\"  Zoomed visualization saved → {save_path}\")\n    plt.show()\n\n# ============================================\n# MAIN PIPELINE\n# ============================================\ndef main():\n    print(\"=\" * 60)\n    print(\"VESUVIUS INK DETECTION — IMPROVED PIPELINE\")\n    print(\"=\" * 60)\n    t0 = time.time()\n\n    # ── 1. Load & extract patches ──────────────────────────────────\n    print(\"\\n1. Loading training data...\")\n    all_patches, all_masks, all_ignores = [], [], []\n\n    for path in train_paths:\n        name = os.path.basename(path)\n        print(f\"\\n   Fragment {name}...\")\n        vol   = load_volume(path)\n        mask  = cv2.imread(os.path.join(path, \"inklabels.png\"), 0)\n        mask  = (mask > 0).astype(np.uint8)\n        ign   = generate_ignore_mask(mask)\n\n        print(f\"    Volume {vol.shape}  |  Ink pixels: {mask.sum():,}\")\n        patches, m_patches, i_patches = extract_patches(vol, mask, ign)\n        all_patches.extend(patches)\n        all_masks.extend(m_patches)\n        all_ignores.extend(i_patches)\n        del vol, mask, ign, patches, m_patches, i_patches\n        gc.collect()\n        if DEVICE == 'cuda': torch.cuda.empty_cache()\n\n    print(f\"\\nTotal patches: {len(all_patches)}\")\n\n    # ── 2. Train / val split ────────────────────────────────────────\n    print(\"\\n2. Splitting data...\")\n    n = len(all_patches)\n    tr_idx, va_idx = spatial_train_val_split(n, VALIDATION_SPLIT)\n    print(f\"   Train: {len(tr_idx)}  |  Val: {len(va_idx)}\")\n\n    def gather(idxs, src): return [src[i] for i in idxs]\n    tr_vol  = gather(tr_idx, all_patches); tr_msk = gather(tr_idx, all_masks); tr_ign = gather(tr_idx, all_ignores)\n    va_vol  = gather(va_idx, all_patches); va_msk = gather(va_idx, all_masks); va_ign = gather(va_idx, all_ignores)\n    del all_patches, all_masks, all_ignores; gc.collect()\n\n    # ── 3. Data augmentation ────────────────────────────────────────\n    print(\"\\n3. Setting up augmentation...\")\n    aug = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.15,\n                                   contrast_limit=0.15, p=0.4),\n        A.GaussNoise(var_limit=(0, 0.015), p=0.3),\n        A.ElasticTransform(alpha=30, sigma=5, alpha_affine=5, p=0.2),\n        A.GridDistortion(p=0.2),\n    ], additional_targets={'ignore_mask': 'mask'})\n\n    tr_ds = VesuviusDataset(tr_vol, tr_msk, tr_ign, transform=aug)\n    va_ds = VesuviusDataset(va_vol, va_msk, va_ign, transform=None)\n\n    tr_dl = DataLoader(tr_ds, batch_size=BATCH_SIZE, shuffle=True,\n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                       drop_last=True, collate_fn=collate_fn)\n    va_dl = DataLoader(va_ds, batch_size=BATCH_SIZE, shuffle=False,\n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                       collate_fn=collate_fn)\n\n    # ── 4. Model & optimizer ────────────────────────────────────────\n    print(\"\\n4. Building model (EfficientNet-B4 UNet++)...\")\n    model = build_model(N_SLICES).to(DEVICE)\n    total  = sum(p.numel() for p in model.parameters())\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"   Params: {total:,}  |  Trainable: {trainable:,}\")\n\n    criterion = FocalDiceLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)\n    scaler    = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n\n    # ── 5. Training loop ────────────────────────────────────────────\n    print(\"\\n5. Training...\")\n    print(\"=\" * 60)\n    best_dice       = 0.0\n    best_threshold  = 0.5\n    patience_count  = 0\n    EARLY_STOP      = 20\n    train_losses    = []\n    val_dices       = []\n\n    for epoch in range(EPOCHS):\n        t_ep = time.time()\n\n        # — Train —\n        model.train()\n        epoch_loss, steps = 0, 0\n        optimizer.zero_grad()\n        pbar = tqdm(tr_dl, desc=f\"E{epoch+1}/{EPOCHS} [train]\", leave=False)\n        for step, (imgs, msks, igns) in enumerate(pbar):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            msks = msks.to(DEVICE, non_blocking=True)\n            igns = igns.to(DEVICE, non_blocking=True)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                out  = model(imgs)\n                loss = criterion(out, msks, igns) / ACCUMULATION_STEPS\n            scaler.scale(loss).backward()\n            if (step + 1) % ACCUMULATION_STEPS == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(optimizer);  scaler.update()\n                optimizer.zero_grad()\n            epoch_loss += loss.item() * ACCUMULATION_STEPS; steps += 1\n            pbar.set_postfix(loss=f\"{loss.item()*ACCUMULATION_STEPS:.4f}\")\n        avg_loss = epoch_loss / steps\n        train_losses.append(avg_loss)\n\n        # — Validate —\n        model.eval()\n        val_dice, val_steps = 0, 0\n        with torch.no_grad():\n            for imgs, msks, igns in tqdm(va_dl, desc=f\"E{epoch+1}/{EPOCHS} [val]\", leave=False):\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                msks = msks.to(DEVICE, non_blocking=True)\n                igns = igns.to(DEVICE, non_blocking=True)\n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    out = model(imgs)\n                preds = (torch.sigmoid(out) > 0.5)\n                valid = (1 - igns)\n                inter = ((preds * msks) * valid).sum()\n                union = ((preds * valid).sum() + (msks * valid).sum())\n                val_dice += (2 * inter / (union + 1e-6)).item()\n                val_steps += 1\n        avg_dice = val_dice / val_steps\n        val_dices.append(avg_dice)\n        scheduler.step()\n\n        ep_time = time.time() - t_ep\n        print(f\"  E{epoch+1:03d}/{EPOCHS}  loss {avg_loss:.4f}  val_dice {avg_dice:.4f}\"\n              f\"  lr {optimizer.param_groups[0]['lr']:.1e}  [{ep_time:.0f}s]\")\n\n        if avg_dice > best_dice:\n            best_dice = avg_dice\n            ckpt_path = os.path.join(OUTPUT_DIR, \"best_model.pth\")\n            torch.save({'epoch': epoch, 'model_state_dict': model.state_dict(),\n                        'val_dice': avg_dice}, ckpt_path)\n            patience_count = 0\n            print(f\"    ✓ Saved best model (Dice {best_dice:.4f})\")\n        else:\n            patience_count += 1\n            if patience_count >= EARLY_STOP:\n                print(f\"\\nEarly stopping at epoch {epoch+1}\")\n                break\n\n    # ── 6. Threshold tuning on val set ──────────────────────────────\n    print(\"\\n6. Tuning decision threshold on validation set...\")\n    ckpt = torch.load(os.path.join(OUTPUT_DIR, \"best_model.pth\"))\n    model.load_state_dict(ckpt['model_state_dict'])\n    best_threshold = find_best_threshold(model, va_dl, DEVICE)\n\n    # ── 7. Full-fragment inference on Fragment 1 ─────────────────────\n    print(\"\\n7. Running inference on Fragment 1 (sliding window + TTA)...\")\n    print(\"   Loading Fragment 1 volume...\")\n    test_vol   = load_volume(fragment1_path)\n    test_mask  = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask  = (test_mask > 0).astype(np.uint8)\n\n    print(\"   Running sliding-window inference with TTA...\")\n    print(f\"   Volume shape: {test_vol.shape}  |  This may take a few minutes...\")\n    prob_map = sliding_window_inference(\n        model, test_vol, patch_size=PATCH_SIZE, stride=STRIDE,\n        use_tta=True, device=DEVICE\n    )\n\n    # ── 8. Visualizations ───────────────────────────────────────────\n    print(\"\\n8. Generating visualizations...\")\n\n    # Full-fragment 4-panel visualization\n    viz_path   = os.path.join(OUTPUT_DIR, \"fragment1_comparison.png\")\n    dice, f1, prec, rec = visualize_fragment_results(\n        test_vol, test_mask, prob_map,\n        threshold=best_threshold,\n        fragment_name=\"Fragment 1\",\n        save_path=viz_path\n    )\n\n    # Zoomed-in interesting regions\n    zoom_path  = os.path.join(OUTPUT_DIR, \"fragment1_zoomed.png\")\n    visualize_zoomed_regions(\n        test_vol, test_mask, prob_map,\n        threshold=best_threshold,\n        n_regions=3,\n        save_path=zoom_path\n    )\n\n    # Training curves\n    fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n    axes[0].plot(train_losses, label='Train loss'); axes[0].set_title(\"Training loss\")\n    axes[0].set_xlabel(\"Epoch\"); axes[0].grid(True)\n    axes[1].plot(val_dices,    label='Val Dice',  color='orange'); axes[1].set_title(\"Validation Dice\")\n    axes[1].axhline(best_dice, color='green', linestyle='--', label=f\"Best {best_dice:.4f}\")\n    axes[1].set_xlabel(\"Epoch\"); axes[1].legend(); axes[1].grid(True)\n    plt.tight_layout()\n    curves_path = os.path.join(OUTPUT_DIR, \"training_curves.png\")\n    plt.savefig(curves_path, dpi=150); plt.show()\n    print(f\"  Training curves saved → {curves_path}\")\n\n    # ── 9. Final results ────────────────────────────────────────────\n    total_time = time.time() - t0\n    h, m = int(total_time // 3600), int((total_time % 3600) // 60)\n    print(\"\\n\" + \"=\" * 60)\n    print(\"FINAL RESULTS — FRAGMENT 1\")\n    print(\"=\" * 60)\n    print(f\"Total time    : {h}h {m}m\")\n    print(f\"Best val Dice : {best_dice:.4f}\")\n    print(f\"Threshold     : {best_threshold:.2f}\")\n    print(f\"Test Dice     : {dice:.4f}\")\n    print(f\"Precision     : {prec:.4f}\")\n    print(f\"Recall        : {rec:.4f}\")\n    print(f\"F1 Score      : {f1:.4f}\")\n    print(\"=\" * 60)\n\n    with open(os.path.join(OUTPUT_DIR, \"final_results.txt\"), \"w\") as f:\n        f.write(f\"Test Dice Score : {dice:.4f}\\n\")\n        f.write(f\"Precision       : {prec:.4f}\\n\")\n        f.write(f\"Recall          : {rec:.4f}\\n\")\n        f.write(f\"F1 Score        : {f1:.4f}\\n\")\n        f.write(f\"Threshold       : {best_threshold:.2f}\\n\")\n        f.write(f\"Best Val Dice   : {best_dice:.4f}\\n\")\n        f.write(f\"Total Time      : {h}h {m}m\\n\")\n\n    return dice\n\n\nif __name__ == \"__main__\":\n    try:\n        score = main()\n        print(f\"\\n✅ Final Test Dice: {score:.4f}\")\n    except Exception as e:\n        import traceback\n        print(f\"\\n❌ Error: {e}\")\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch==0.2.0\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport torch.nn.functional as F\nimport gc\nimport psutil\nimport time\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\nimport warnings\nwarnings.filterwarnings('ignore')\n#/kaggle/input/vesuvius-challenge-ink-detection/train\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n]\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n    print(f\"GPU Name: {torch.cuda.get_device_properties(0).name}\")\n\n# Optimized settings\nPATCH_SIZE = 192  # Reduced from 256 for faster training\nSTRIDE = 96\nBATCH_SIZE = 4    # Increased from 2\nACCUMULATION_STEPS = 4  # Effective batch size of 16\nEPOCHS = 50\nLR = 2e-3\nWEIGHT_DECAY = 1e-4\nSLICE_START = 26\nSLICE_END = 38  # 12 slices\nOUTPUT_DIR = \"/kaggle/working/\"\nVALIDATION_SPLIT = 0.15\nVISUALIZATION_DIR = \"/kaggle/working/visualizations/\"\nos.makedirs(VISUALIZATION_DIR, exist_ok=True)\n\n# Training settings\nUSE_AMP = True\nPIN_MEMORY = True\nNUM_WORKERS = 2\nGRAD_CLIP = 1.0\n\n# ============================================\n# SIMPLIFIED IGNORE MASK (FASTER)\n# ============================================\ndef generate_ignore_mask_simple(ink_mask):\n    \"\"\"Fast ignore mask generation\"\"\"\n    # Dilate and erode to find boundaries\n    kernel = np.ones((5, 5), np.uint8)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    eroded = cv2.erode(ink_mask.astype(np.uint8), kernel, iterations=1)\n    \n    # Boundaries are where dilated and eroded differ\n    ignore_mask = (dilated != eroded).astype(np.uint8)\n    \n    return ignore_mask\n\n# ============================================\n# EFFICIENT MODEL ARCHITECTURE\n# ============================================\nclass ConvBlock(nn.Module):\n    def __init__(self, in_c, out_c):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_c, out_c, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm2d(out_c)\n        self.conv2 = nn.Conv2d(out_c, out_c, kernel_size=3, padding=1)\n        self.bn2 = nn.BatchNorm2d(out_c)\n        self.relu = nn.ReLU(inplace=True)\n        \n    def forward(self, x):\n        x = self.relu(self.bn1(self.conv1(x)))\n        x = self.relu(self.bn2(self.conv2(x)))\n        return x\n\nclass EfficientUNet(nn.Module):\n    \"\"\"Efficient U-Net with proper dimension handling\"\"\"\n    def __init__(self, in_channels=12, out_channels=1, features=[32, 64, 128, 256]):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = ConvBlock(in_channels, features[0])\n        self.enc2 = ConvBlock(features[0], features[1])\n        self.enc3 = ConvBlock(features[1], features[2])\n        self.enc4 = ConvBlock(features[2], features[3])\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        \n        # Bottleneck\n        self.bottleneck = ConvBlock(features[3], features[3]*2)\n        \n        # Decoder\n        self.up4 = nn.ConvTranspose2d(features[3]*2, features[3], kernel_size=2, stride=2)\n        self.dec4 = ConvBlock(features[3]*2, features[3])\n        \n        self.up3 = nn.ConvTranspose2d(features[3], features[2], kernel_size=2, stride=2)\n        self.dec3 = ConvBlock(features[2]*2, features[2])\n        \n        self.up2 = nn.ConvTranspose2d(features[2], features[1], kernel_size=2, stride=2)\n        self.dec2 = ConvBlock(features[1]*2, features[1])\n        \n        self.up1 = nn.ConvTranspose2d(features[1], features[0], kernel_size=2, stride=2)\n        self.dec1 = ConvBlock(features[0]*2, features[0])\n        \n        # Output\n        self.final = nn.Conv2d(features[0], out_channels, kernel_size=1)\n        \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        p1 = self.pool(e1)\n        \n        e2 = self.enc2(p1)\n        p2 = self.pool(e2)\n        \n        e3 = self.enc3(p2)\n        p3 = self.pool(e3)\n        \n        e4 = self.enc4(p3)\n        p4 = self.pool(e4)\n        \n        # Bottleneck\n        b = self.bottleneck(p4)\n        \n        # Decoder\n        d4 = self.up4(b)\n        d4 = torch.cat([d4, e4], dim=1)\n        d4 = self.dec4(d4)\n        \n        d3 = self.up3(d4)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.up2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.up1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        # Output\n        out = self.final(d1)\n        \n        return out\n\n# ============================================\n# LOSS FUNCTIONS\n# ============================================\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1e-6):\n        super().__init__()\n        self.smooth = smooth\n        \n    def forward(self, pred, target, ignore_mask=None):\n        pred = torch.sigmoid(pred)\n        \n        if ignore_mask is not None:\n            valid_mask = (1 - ignore_mask)\n            pred = pred * valid_mask\n            target = target * valid_mask\n        \n        intersection = (pred * target).sum()\n        union = pred.sum() + target.sum()\n        \n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        return 1 - dice\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, dice_weight=0.5, bce_weight=0.5):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.dice_loss = DiceLoss()\n        self.bce_loss = nn.BCEWithLogitsLoss(reduction='none')\n        \n    def forward(self, pred, target, ignore_mask):\n        valid_mask = (1 - ignore_mask)\n        \n        # BCE loss\n        bce = self.bce_loss(pred, target)\n        bce = (bce * valid_mask).sum() / (valid_mask.sum() + 1e-6)\n        \n        # Dice loss\n        dice = self.dice_loss(pred, target, ignore_mask)\n        \n        return self.dice_weight * dice + self.bce_weight * bce\n\n# ============================================\n# DATA LOADING (OPTIMIZED)\n# ============================================\ndef load_volume_fast(fragment_path):\n    \"\"\"Fast volume loading\"\"\"\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            # Simple normalization\n            img = img.astype(np.float32)\n            img = (img - img.mean()) / (img.std() + 1e-6)\n            slices.append(img)\n    \n    volume = np.stack(slices, axis=-1)\n    return volume\n\ndef extract_patches_fast(volume, mask, ignore_mask, max_patches=4000):\n    \"\"\"Fast patch extraction\"\"\"\n    patches = []\n    mask_patches = []\n    ignore_patches = []\n    \n    H, W, _ = volume.shape\n    patch_count = 0\n    \n    # Calculate stride to get roughly max_patches\n    n_positions = int(np.sqrt(max_patches))\n    stride_y = max(STRIDE, (H - PATCH_SIZE) // n_positions)\n    stride_x = max(STRIDE, (W - PATCH_SIZE) // n_positions)\n    \n    for y in range(0, H - PATCH_SIZE, stride_y):\n        for x in range(0, W - PATCH_SIZE, stride_x):\n            if patch_count >= max_patches:\n                break\n                \n            v_patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            m_patch = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            i_patch = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            \n            # Keep patches with ink or some background\n            if m_patch.sum() > 10 or np.random.random() < 0.2:\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n        \n        if patch_count >= max_patches:\n            break\n    \n    return patches, mask_patches, ignore_patches\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignore_masks, transform=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.ignore_masks = ignore_masks\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        image = self.volumes[idx].copy()\n        mask = self.masks[idx].copy()\n        ignore = self.ignore_masks[idx].copy()\n        \n        if self.transform:\n            transformed = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image = transformed['image']\n            mask = transformed['mask']\n            ignore = transformed['ignore_mask']\n        \n        # Convert to tensors\n        image = torch.tensor(image).permute(2, 0, 1).float()\n        mask = torch.tensor(mask).float().unsqueeze(0)\n        ignore = torch.tensor(ignore).float().unsqueeze(0)\n        \n        return image, mask, ignore\n\n# ============================================\n# AUGMENTATION\n# ============================================\ndef get_augmentation():\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(p=0.3),\n    ], additional_targets={'ignore_mask': 'mask'})\n\n# ============================================\n# VISUALIZATION\n# ============================================\ndef visualize_predictions(model, dataset, num_samples=4, epoch=0):\n    \"\"\"Visualize model predictions\"\"\"\n    model.eval()\n    fig, axes = plt.subplots(num_samples, 4, figsize=(16, 4*num_samples))\n    \n    indices = np.random.choice(len(dataset), min(num_samples, len(dataset)), replace=False)\n    \n    with torch.no_grad():\n        for i, idx in enumerate(indices):\n            image, mask, ignore = dataset[idx]\n            image_batch = image.unsqueeze(0).to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                output = model(image_batch)\n            \n            pred = torch.sigmoid(output).cpu().numpy()[0, 0]\n            mask_np = mask.numpy()[0]\n            \n            # Input visualization (first 3 slices as RGB)\n            input_rgb = np.stack([image[0].numpy(), image[4].numpy(), image[8].numpy()], axis=-1)\n            input_rgb = (input_rgb - input_rgb.min()) / (input_rgb.max() - input_rgb.min() + 1e-6)\n            \n            axes[i, 0].imshow(input_rgb)\n            axes[i, 0].set_title('Input')\n            axes[i, 0].axis('off')\n            \n            axes[i, 1].imshow(mask_np, cmap='gray')\n            axes[i, 1].set_title('Ground Truth')\n            axes[i, 1].axis('off')\n            \n            axes[i, 2].imshow(pred, cmap='hot', vmin=0, vmax=1)\n            axes[i, 2].set_title(f'Prediction')\n            axes[i, 2].axis('off')\n            \n            # Overlay\n            overlay = np.zeros((*pred.shape, 3))\n            overlay[..., 0] = pred\n            overlay[..., 1] = mask_np\n            axes[i, 3].imshow(overlay)\n            axes[i, 3].set_title('Overlay')\n            axes[i, 3].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(VISUALIZATION_DIR, f'predictions_epoch_{epoch}.png'))\n    plt.close()\n\n# ============================================\n# TRAINING FUNCTIONS\n# ============================================\ndef train_epoch(model, loader, criterion, optimizer, scaler, epoch):\n    model.train()\n    total_loss = 0\n    progress_bar = tqdm(loader, desc=f'Epoch {epoch} [Train]')\n    \n    for batch_idx, (images, masks, ignores) in enumerate(progress_bar):\n        images = images.to(DEVICE, non_blocking=True)\n        masks = masks.to(DEVICE, non_blocking=True)\n        ignores = ignores.to(DEVICE, non_blocking=True)\n        \n        with torch.cuda.amp.autocast(enabled=USE_AMP):\n            outputs = model(images)\n            loss = criterion(outputs, masks, ignores)\n            loss = loss / ACCUMULATION_STEPS\n        \n        scaler.scale(loss).backward()\n        \n        if (batch_idx + 1) % ACCUMULATION_STEPS == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n        \n        total_loss += loss.item() * ACCUMULATION_STEPS\n        progress_bar.set_postfix({'loss': f'{loss.item() * ACCUMULATION_STEPS:.4f}'})\n    \n    return total_loss / len(loader)\n\ndef validate(model, loader, criterion):\n    model.eval()\n    total_dice = 0\n    total_loss = 0\n    \n    with torch.no_grad():\n        for images, masks, ignores in tqdm(loader, desc='Validation'):\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            ignores = ignores.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n                loss = criterion(outputs, masks, ignores)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            valid_mask = (1 - ignores)\n            \n            intersection = ((preds * masks) * valid_mask).sum()\n            union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n            dice = (2 * intersection) / (union + 1e-6)\n            \n            total_dice += dice.item()\n            total_loss += loss.item()\n    \n    return total_loss / len(loader), total_dice / len(loader)\n\n# ============================================\n# MAIN PIPELINE\n# ============================================\ndef main():\n    print(\"=\"*60)\n    print(\"VESUVIUS INK DETECTION - OPTIMIZED PIPELINE\")\n    print(\"=\"*60)\n    \n    start_time = time.time()\n    \n    # ============================================\n    # LOAD DATA\n    # ============================================\n    print(\"\\n1. Loading training data...\")\n    all_patches = []\n    all_masks = []\n    all_ignores = []\n    \n    for path in train_paths:\n        fragment_name = os.path.basename(path)\n        print(f\"\\n   Processing Fragment {fragment_name}...\")\n        \n        volume = load_volume_fast(path)\n        mask = cv2.imread(os.path.join(path, \"inklabels.png\"), 0)\n        mask = (mask > 0).astype(np.uint8)\n        \n        ink_density = mask.sum() / mask.size\n        print(f\"    Ink density: {ink_density:.4f}\")\n        \n        ignore_mask = generate_ignore_mask_simple(mask)\n        print(f\"    Ignored pixels: {ignore_mask.sum() / ignore_mask.size:.4f}\")\n        \n        patches, masks_p, ignores_p = extract_patches_fast(\n            volume, mask, ignore_mask, max_patches=4000\n        )\n        \n        all_patches.extend(patches)\n        all_masks.extend(masks_p)\n        all_ignores.extend(ignores_p)\n        \n        print(f\"    Extracted {len(patches)} patches\")\n        \n        del volume, mask, ignore_mask\n        gc.collect()\n    \n    print(f\"\\nTotal patches: {len(all_patches)}\")\n    \n    # ============================================\n    # SPLIT DATA\n    # ============================================\n    print(\"\\n2. Creating train/validation split...\")\n    n_samples = len(all_patches)\n    indices = np.random.permutation(n_samples)\n    split = int(n_samples * VALIDATION_SPLIT)\n    \n    train_indices = indices[split:]\n    val_indices = indices[:split]\n    \n    print(f\"   Train samples: {len(train_indices)}\")\n    print(f\"   Validation samples: {len(val_indices)}\")\n    \n    # Split data\n    train_patches = [all_patches[i] for i in train_indices]\n    train_masks = [all_masks[i] for i in train_indices]\n    train_ignores = [all_ignores[i] for i in train_indices]\n    \n    val_patches = [all_patches[i] for i in val_indices]\n    val_masks = [all_masks[i] for i in val_indices]\n    val_ignores = [all_ignores[i] for i in val_indices]\n    \n    # ============================================\n    # CREATE DATASETS\n    # ============================================\n    print(\"\\n3. Creating datasets...\")\n    \n    train_dataset = VesuviusDataset(\n        train_patches, train_masks, train_ignores,\n        transform=get_augmentation()\n    )\n    \n    val_dataset = VesuviusDataset(\n        val_patches, val_masks, val_ignores,\n        transform=None\n    )\n    \n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=True, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        drop_last=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY\n    )\n    \n    print(f\"   Train batches: {len(train_loader)}\")\n    print(f\"   Val batches: {len(val_loader)}\")\n    \n    # ============================================\n    # INITIALIZE MODEL\n    # ============================================\n    print(\"\\n4. Initializing model...\")\n    \n    model = EfficientUNet(in_channels=12, out_channels=1).to(DEVICE)\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"   Total parameters: {total_params:,}\")\n    \n    # ============================================\n    # LOSS, OPTIMIZER, SCHEDULER\n    # ============================================\n    criterion = CombinedLoss(dice_weight=0.5, bce_weight=0.5)\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n    \n    # ============================================\n    # TRAINING LOOP\n    # ============================================\n    print(\"\\n5. Starting training...\")\n    print(\"=\"*60)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    train_losses = []\n    val_losses = []\n    val_dices = []\n    \n    for epoch in range(EPOCHS):\n        epoch_start = time.time()\n        \n        # Train\n        train_loss = train_epoch(model, train_loader, criterion, optimizer, scaler, epoch+1)\n        train_losses.append(train_loss)\n        \n        # Validate\n        val_loss, val_dice = validate(model, val_loader, criterion)\n        val_losses.append(val_loss)\n        val_dices.append(val_dice)\n        \n        # Update scheduler\n        scheduler.step()\n        \n        epoch_time = time.time() - epoch_start\n        \n        # Print results\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS} - Time: {epoch_time:.1f}s\")\n        print(f\"  Train Loss: {train_loss:.4f}\")\n        print(f\"  Val Loss: {val_loss:.4f}\")\n        print(f\"  Val Dice: {val_dice:.4f}\")\n        print(f\"  LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        # Visualize every 5 epochs\n        if (epoch + 1) % 5 == 0:\n            visualize_predictions(model, val_dataset, epoch=epoch+1)\n        \n        # Save best model\n        if val_dice > best_val_dice:\n            best_val_dice = val_dice\n            torch.save(model.state_dict(), os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_counter = 0\n            print(f\"  ✓ New best model saved! Dice: {val_dice:.4f}\")\n        else:\n            patience_counter += 1\n        \n        # Early stopping\n        if patience_counter >= 80:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n        \n        print(\"-\"*60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING COMPLETE\")\n    print(f\"Best Validation Dice: {best_val_dice:.4f}\")\n    print(\"=\"*60)\n    \n    # ============================================\n    # FINAL TEST\n    # ============================================\n    print(\"\\n6. Testing on Fragment 1...\")\n    print(\"=\"*60)\n    \n    # Load best model\n    model.load_state_dict(torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth')))\n    \n    # Test on patches\n    test_volume = load_volume_fast(fragment1_path)\n    test_mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ignore = generate_ignore_mask_simple(test_mask)\n    \n    test_patches, test_masks, test_ignores = extract_patches_fast(\n        test_volume, test_mask, test_ignore, max_patches=2000\n    )\n    \n    test_dataset = VesuviusDataset(\n        test_patches, test_masks, test_ignores, transform=None\n    )\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS\n    )\n    \n    # Evaluate\n    test_loss, test_dice = validate(model, test_loader, criterion)\n    \n    # Calculate detailed metrics\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, ignores in test_loader:\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            all_preds.append(preds.cpu().numpy())\n            all_targets.append(masks.cpu().numpy())\n    \n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_targets = np.concatenate([t.flatten() for t in all_targets])\n    \n    tn, fp, fn, tp = confusion_matrix(all_targets, all_preds, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    \n    # Final visualization\n    visualize_predictions(model, test_dataset, num_samples=8, epoch='final')\n    \n    # ============================================\n    # FINAL RESULTS\n    # ============================================\n    total_time = time.time() - start_time\n    hours = int(total_time // 3600)\n    minutes = int((total_time % 3600) // 60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"FINAL TEST RESULTS ON FRAGMENT 1\")\n    print(\"=\"*60)\n    print(f\"Total Time: {hours}h {minutes}m\")\n    print(f\"Test Loss: {test_loss:.4f}\")\n    print(f\"Test Dice Score: {test_dice:.4f}\")\n    print(f\"Precision: {precision:.4f}\")\n    print(f\"Recall: {recall:.4f}\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(\"-\"*60)\n    print(\"Confusion Matrix:\")\n    print(f\"  True Positives: {tp}\")\n    print(f\"  True Negatives: {tn}\")\n    print(f\"  False Positives: {fp}\")\n    print(f\"  False Negatives: {fn}\")\n    print(\"=\"*60)\n    \n    # Plot training curves\n    plt.figure(figsize=(12, 4))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses, label='Train')\n    plt.plot(val_losses, label='Validation')\n    plt.title('Loss Curves')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.grid(True)\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(val_dices)\n    plt.title('Validation Dice Score')\n    plt.xlabel('Epoch')\n    plt.ylabel('Dice')\n    plt.grid(True)\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'))\n    plt.show()\n    \n    return test_dice\n\nif __name__ == \"__main__\":\n    try:\n        test_dice = main()\n        print(f\"\\n✅ Final Test Dice Score: {test_dice:.4f}\")\n    except Exception as e:\n        print(f\"\\n❌ Error: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test Dice Score: 0.5015 F1 Score: 0.5416 ON Fragment-1\n!pip install segmentation-models-pytorch==0.2.0\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport torch.nn.functional as F\nimport gc\nimport psutil\nimport time\n#/kaggle/input/vesuvius-challenge-ink-detection/test\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n]\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {DEVICE}\")\nif DEVICE == 'cuda':\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\nelse:\n    print(\"CPU mode\")\n\n# Memory-efficient settings\nPATCH_SIZE = 128\nSTRIDE = 64\nBATCH_SIZE = 14  # Increased slightly since we're using 2D\nACCUMULATION_STEPS = 2  # Gradient accumulation\nEPOCHS = 100\nLR = 3e-4\nWEIGHT_DECAY = 1e-4\nSLICE_START = 20\nSLICE_END = 32  # 12 slices\nOUTPUT_DIR = \"/kaggle/working/\"\nVALIDATION_SPLIT = 0.25\n\n# Advanced training settings\nUSE_AMP = True\nPIN_MEMORY = True\nNUM_WORKERS = 2\nGRAD_CLIP = 1.0\nSCHEDULER_PATIENCE = 44\nEARLY_STOPPING_PATIENCE = 80\n\n# ============================================\n# IGNORE MASK GENERATION\n# ============================================\ndef generate_ignore_mask(ink_mask, distance_threshold=3, erosion_size=2):\n    \"\"\"Generate ignore mask for uncertain regions - OPTIMIZED VERSION\"\"\"\n    ignore_mask = np.zeros_like(ink_mask, dtype=np.uint8)\n    \n    # Simple boundary detection (faster than full distance transform)\n    kernel = np.ones((3, 3), np.uint8)\n    eroded = cv2.erode(ink_mask.astype(np.uint8), kernel, iterations=1)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    \n    # Boundaries are where eroded and dilated differ from original\n    boundaries = (dilated != eroded)\n    ignore_mask[boundaries] = 1\n    \n    # Remove small isolated ink dots\n    contours, _ = cv2.findContours(ink_mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for contour in contours:\n        if cv2.contourArea(contour) < 10:\n            cv2.drawContours(ignore_mask, [contour], -1, 1, -1)\n    \n    return ignore_mask\n\n# ============================================\n# 2D CNN WITH MULTI-SLICE INPUT\n# ============================================\nclass MultiSlice2DUNet(nn.Module):\n    \"\"\"2D CNN that treats depth slices as input channels\"\"\"\n    def __init__(self, in_channels=12, out_channels=1):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._conv_block(in_channels, 32)\n        self.enc2 = self._conv_block(32, 64)\n        self.enc3 = self._conv_block(64, 128)\n        self.enc4 = self._conv_block(128, 256)\n        \n        # Bottleneck\n        self.bottleneck = self._conv_block(256, 512)\n        \n        # Decoder\n        self.dec4 = self._upconv_block(512 + 256, 256)\n        self.dec3 = self._upconv_block(256 + 128, 128)\n        self.dec2 = self._upconv_block(128 + 64, 64)\n        self.dec1 = self._upconv_block(64 + 32, 32)\n        \n        # Output\n        self.final = nn.Conv2d(32, out_channels, kernel_size=1)\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False)\n        \n    def _conv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def _upconv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # x shape: [B, C=12, H, W]\n        \n        # Encoder\n        e1 = self.enc1(x)      # [B, 32, H, W]\n        p1 = self.pool(e1)      # [B, 32, H/2, W/2]\n        \n        e2 = self.enc2(p1)      # [B, 64, H/2, W/2]\n        p2 = self.pool(e2)      # [B, 64, H/4, W/4]\n        \n        e3 = self.enc3(p2)      # [B, 128, H/4, W/4]\n        p3 = self.pool(e3)      # [B, 128, H/8, W/8]\n        \n        e4 = self.enc4(p3)      # [B, 256, H/8, W/8]\n        p4 = self.pool(e4)      # [B, 256, H/16, W/16]\n        \n        # Bottleneck\n        b = self.bottleneck(p4)  # [B, 512, H/16, W/16]\n        \n        # Decoder with skip connections\n        d4 = self.upsample(b)    # [B, 512, H/8, W/8]\n        d4 = torch.cat([d4, e4], dim=1)  # [B, 512+256, H/8, W/8]\n        d4 = self.dec4(d4)        # [B, 256, H/8, W/8]\n        \n        d3 = self.upsample(d4)    # [B, 256, H/4, W/4]\n        d3 = torch.cat([d3, e3], dim=1)  # [B, 256+128, H/4, W/4]\n        d3 = self.dec3(d3)        # [B, 128, H/4, W/4]\n        \n        d2 = self.upsample(d3)    # [B, 128, H/2, W/2]\n        d2 = torch.cat([d2, e2], dim=1)  # [B, 128+64, H/2, W/2]\n        d2 = self.dec2(d2)        # [B, 64, H/2, W/2]\n        \n        d1 = self.upsample(d2)    # [B, 64, H, W]\n        d1 = torch.cat([d1, e1], dim=1)  # [B, 64+32, H, W]\n        d1 = self.dec1(d1)        # [B, 32, H, W]\n        \n        # Final output\n        out = self.final(d1)      # [B, 1, H, W]\n        \n        return out\n\n# ============================================\n# DATA LOADING (OPTIMIZED)\n# ============================================\ndef load_volume_fast(fragment_path):\n    \"\"\"Fast volume loading\"\"\"\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    \n    volume = np.stack(slices, axis=-1)\n    volume = volume.astype(np.float32)\n    \n    # Fast normalization per slice\n    for i in range(volume.shape[-1]):\n        slice_data = volume[:, :, i]\n        mean, std = slice_data.mean(), slice_data.std()\n        volume[:, :, i] = (slice_data - mean) / (std + 1e-6)\n    \n    return volume\n\ndef extract_patches_fast(volume, mask, ignore_mask, max_patches_per_fragment=3000):\n    \"\"\"Fast patch extraction with sampling\"\"\"\n    patches = []\n    mask_patches = []\n    ignore_patches = []\n    \n    H, W, _ = volume.shape\n    \n    # Calculate number of patches\n    n_y = (H - PATCH_SIZE) // STRIDE + 1\n    n_x = (W - PATCH_SIZE) // STRIDE + 1\n    total_patches = n_y * n_x\n    \n    print(f\"    Total possible patches: {total_patches}\")\n    \n    # Sample patches if too many\n    if total_patches > max_patches_per_fragment:\n        print(f\"    Sampling {max_patches_per_fragment} patches...\")\n        # Calculate stride to get roughly max_patches\n        stride_y = max(STRIDE, (H - PATCH_SIZE) // int(np.sqrt(max_patches_per_fragment)))\n        stride_x = max(STRIDE, (W - PATCH_SIZE) // int(np.sqrt(max_patches_per_fragment)))\n    else:\n        stride_y, stride_x = STRIDE, STRIDE\n    \n    patch_count = 0\n    ink_patches = 0\n    bg_patches = 0\n    \n    for y in range(0, H - PATCH_SIZE, stride_y):\n        for x in range(0, W - PATCH_SIZE, stride_x):\n            if patch_count >= max_patches_per_fragment:\n                break\n                \n            v_patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            m_patch = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            i_patch = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            \n            # Count ink pixels in this patch\n            ink_pixel_count = m_patch.sum()\n            \n            # Keep patches with significant ink or some background for balance\n            if ink_pixel_count > 50:  # Good ink patch\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n                ink_patches += 1\n            elif ink_pixel_count == 0 and bg_patches < ink_patches // 2:  # Balance with background\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                patch_count += 1\n                bg_patches += 1\n        \n        if patch_count >= max_patches_per_fragment:\n            break\n    \n    print(f\"    Extracted: {ink_patches} ink patches, {bg_patches} background patches\")\n    return patches, mask_patches, ignore_patches\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, ignore_masks, transform=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.ignore_masks = ignore_masks\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        image = self.volumes[idx].copy()\n        mask = self.masks[idx].copy()\n        ignore = self.ignore_masks[idx].copy()\n        \n        if self.transform:\n            transformed = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image = transformed['image']\n            mask = transformed['mask']\n            ignore = transformed['ignore_mask']\n        \n        # Convert to tensors\n        # For 2D CNN: image shape [H, W, C] -> [C, H, W]\n        image = torch.tensor(image).permute(2, 0, 1).float()  # [C=12, H, W]\n        \n        # mask and ignore shape: [H, W]\n        mask = torch.tensor(mask).float().unsqueeze(0)  # [1, H, W]\n        ignore = torch.tensor(ignore).float().unsqueeze(0)  # [1, H, W]\n        \n        return image, mask, ignore\n\n# ============================================\n# LOSS FUNCTION\n# ============================================\nclass DiceBCELossWithIgnore(nn.Module):\n    def __init__(self, dice_weight=0.5, bce_weight=0.5, smooth=1e-6):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.smooth = smooth\n        \n    def forward(self, pred, target, ignore_mask):\n        \"\"\"\n        pred: [B, 1, H, W] - logits\n        target: [B, 1, H, W] - binary mask\n        ignore_mask: [B, 1, H, W] - 1 for ignore, 0 for keep\n        \"\"\"\n        # Create valid mask\n        valid_mask = (1 - ignore_mask).float()\n        \n        # BCE loss\n        bce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')\n        bce = (bce * valid_mask).sum() / (valid_mask.sum() + self.smooth)\n        \n        # Dice loss\n        pred_probs = torch.sigmoid(pred)\n        \n        # Apply valid mask\n        pred_valid = pred_probs * valid_mask\n        target_valid = target * valid_mask\n        \n        intersection = (pred_valid * target_valid).sum()\n        union = pred_valid.sum() + target_valid.sum()\n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        dice_loss = 1 - dice\n        \n        return self.dice_weight * dice_loss + self.bce_weight * bce\n\n# ============================================\n# FAST TRAIN/VALIDATION SPLIT\n# ============================================\ndef fast_train_val_split(n_samples, val_ratio=0.15, seed=42):\n    \"\"\"Fast random split without distance constraints\"\"\"\n    np.random.seed(seed)\n    indices = np.random.permutation(n_samples)\n    split = int(n_samples * val_ratio)\n    return indices[split:], indices[:split]\n\n# ============================================\n# MEMORY MONITORING\n# ============================================\ndef print_memory_usage():\n    if DEVICE == 'cuda':\n        allocated = torch.cuda.memory_allocated() / 1e9\n        cached = torch.cuda.memory_reserved() / 1e9\n        print(f\"    GPU Memory - Allocated: {allocated:.2f}GB, Cached: {cached:.2f}GB\")\n    \n    process = psutil.Process()\n    print(f\"    CPU Memory: {process.memory_info().rss / 1e9:.2f}GB\")\n\n# ============================================\n# COLLATE FUNCTION\n# ============================================\ndef collate_fn(batch):\n    \"\"\"Custom collate function to ensure correct dimensions\"\"\"\n    images = torch.stack([item[0] for item in batch])  # [B, C=12, H, W]\n    masks = torch.stack([item[1] for item in batch])   # [B, 1, H, W]\n    ignores = torch.stack([item[2] for item in batch]) # [B, 1, H, W]\n    return images, masks, ignores\n\n# ============================================\n# MAIN PIPELINE\n# ============================================\ndef main():\n    print(\"=\"*60)\n    print(\"VESUVIUS INK DETECTION - OPTIMIZED PIPELINE\")\n    print(\"=\"*60)\n    print_memory_usage()\n    \n    start_time = time.time()\n    \n    # ============================================\n    # LOAD AND PREPARE DATA\n    # ============================================\n    print(\"\\n1. Loading training data...\")\n    all_patches = []\n    all_masks = []\n    all_ignores = []\n    \n    for path in train_paths:\n        fragment_name = os.path.basename(path)\n        print(f\"\\n   Processing Fragment {fragment_name}...\")\n        \n        # Load volume\n        volume = load_volume_fast(path)\n        print(f\"    Volume shape: {volume.shape}\")\n        \n        # Load mask\n        mask_path = os.path.join(path, \"inklabels.png\")\n        mask = cv2.imread(mask_path, 0)\n        mask = (mask > 0).astype(np.uint8)\n        print(f\"    Mask shape: {mask.shape}\")\n        print(f\"    Ink pixels: {mask.sum():,}\")\n        \n        # Generate ignore mask\n        print(\"    Generating ignore mask...\")\n        ignore_mask = generate_ignore_mask(mask)\n        print(f\"    Ignored pixels: {ignore_mask.sum():,}\")\n        \n        # Extract patches\n        print(\"    Extracting patches...\")\n        patches, mask_patches, ignore_patches = extract_patches_fast(\n            volume, mask, ignore_mask, max_patches_per_fragment=3000\n        )\n        \n        all_patches.extend(patches)\n        all_masks.extend(mask_patches)\n        all_ignores.extend(ignore_patches)\n        \n        print(f\"    Total extracted: {len(patches)} patches\")\n        print_memory_usage()\n        \n        # Clean up\n        del volume, mask, ignore_mask, patches, mask_patches, ignore_patches\n        gc.collect()\n        if DEVICE == 'cuda':\n            torch.cuda.empty_cache()\n    \n    print(f\"\\nTotal patches: {len(all_patches)}\")\n    print_memory_usage()\n    \n    # ============================================\n    # CREATE TRAIN/VAL SPLIT\n    # ============================================\n    print(\"\\n2. Creating train/validation split...\")\n    n_samples = len(all_patches)\n    train_indices, val_indices = fast_train_val_split(n_samples, VALIDATION_SPLIT)\n    \n    print(f\"   Train samples: {len(train_indices)}\")\n    print(f\"   Validation samples: {len(val_indices)}\")\n    \n    # Split data\n    train_patches = [all_patches[i] for i in train_indices]\n    train_masks = [all_masks[i] for i in train_indices]\n    train_ignores = [all_ignores[i] for i in train_indices]\n    \n    val_patches = [all_patches[i] for i in val_indices]\n    val_masks = [all_masks[i] for i in val_indices]\n    val_ignores = [all_ignores[i] for i in val_indices]\n    \n    # Clean up original lists\n    del all_patches, all_masks, all_ignores\n    gc.collect()\n    \n    # ============================================\n    # DATA AUGMENTATION\n    # ============================================\n    print(\"\\n3. Setting up data augmentation...\")\n    \n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),\n        A.GaussNoise(var_limit=(0, 0.01), p=0.2),\n    ], additional_targets={'ignore_mask': 'mask'})\n    \n    # Create datasets\n    train_dataset = VesuviusDataset(\n        train_patches, train_masks, train_ignores, \n        transform=train_transform\n    )\n    \n    val_dataset = VesuviusDataset(\n        val_patches, val_masks, val_ignores,\n        transform=None\n    )\n    \n    # Create data loaders with custom collate function\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=True, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        drop_last=True,\n        collate_fn=collate_fn\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    print(f\"   Train batches: {len(train_loader)}\")\n    print(f\"   Val batches: {len(val_loader)}\")\n    \n    # ============================================\n    # MODEL INITIALIZATION\n    # ============================================\n    print(\"\\n4. Initializing model...\")\n    \n    model = MultiSlice2DUNet(in_channels=12, out_channels=1).to(DEVICE)\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"   Total parameters: {total_params:,}\")\n    print(f\"   Trainable parameters: {trainable_params:,}\")\n    \n    # ============================================\n    # LOSS, OPTIMIZER, SCHEDULER\n    # ============================================\n    criterion = DiceBCELossWithIgnore(dice_weight=0.5, bce_weight=0.5)\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='max', factor=0.5, patience=SCHEDULER_PATIENCE, verbose=True\n    )\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n    \n    # ============================================\n    # TRAINING LOOP\n    # ============================================\n    print(\"\\n5. Starting training...\")\n    print(\"=\"*60)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    train_losses = []\n    val_dice_scores = []\n    \n    for epoch in range(EPOCHS):\n        epoch_start = time.time()\n        \n        # Training phase\n        model.train()\n        train_loss = 0\n        train_steps = 0\n        optimizer.zero_grad()\n        \n        progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Train]')\n        for batch_idx, (images, masks, ignores) in enumerate(progress_bar):\n            # images shape: [B, C=12, H, W]\n            # masks shape: [B, 1, H, W]\n            # ignores shape: [B, 1, H, W]\n            \n            images = images.to(DEVICE, non_blocking=True)\n            masks = masks.to(DEVICE, non_blocking=True)\n            ignores = ignores.to(DEVICE, non_blocking=True)\n            \n            # Forward pass\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)  # [B, 1, H, W]\n                loss = criterion(outputs, masks, ignores)\n                loss = loss / ACCUMULATION_STEPS\n            \n            # Backward pass\n            scaler.scale(loss).backward()\n            \n            # Gradient accumulation\n            if (batch_idx + 1) % ACCUMULATION_STEPS == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            train_loss += loss.item() * ACCUMULATION_STEPS\n            train_steps += 1\n            \n            # Update progress bar\n            progress_bar.set_postfix({'loss': f'{loss.item() * ACCUMULATION_STEPS:.4f}'})\n            \n            # Clear cache periodically\n            if batch_idx % 50 == 49:\n                if DEVICE == 'cuda':\n                    torch.cuda.empty_cache()\n        \n        avg_train_loss = train_loss / train_steps\n        train_losses.append(avg_train_loss)\n        \n        # Validation phase\n        model.eval()\n        val_dice = 0\n        val_steps = 0\n        \n        with torch.no_grad():\n            for images, masks, ignores in tqdm(val_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Val]'):\n                images = images.to(DEVICE, non_blocking=True)\n                masks = masks.to(DEVICE, non_blocking=True)\n                ignores = ignores.to(DEVICE, non_blocking=True)\n                \n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    outputs = model(images)\n                \n                # Calculate dice\n                preds = torch.sigmoid(outputs) > 0.5\n                valid_mask = (1 - ignores)\n                \n                intersection = ((preds * masks) * valid_mask).sum()\n                union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n                dice = (2 * intersection) / (union + 1e-6)\n                val_dice += dice.item()\n                val_steps += 1\n        \n        avg_val_dice = val_dice / val_steps\n        val_dice_scores.append(avg_val_dice)\n        \n        # Update scheduler\n        scheduler.step(avg_val_dice)\n        \n        epoch_time = time.time() - epoch_start\n        \n        # Print results\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS} - Time: {epoch_time:.1f}s\")\n        print(f\"  Train Loss: {avg_train_loss:.4f}\")\n        print(f\"  Val Dice: {avg_val_dice:.4f}\")\n        print(f\"  LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        print_memory_usage()\n        \n        # Save best model\n        if avg_val_dice > best_val_dice:\n            best_val_dice = avg_val_dice\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_dice': avg_val_dice,\n            }, os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_counter = 0\n            print(f\"  ✓ New best model saved! Dice: {best_val_dice:.4f}\")\n        else:\n            patience_counter += 1\n        \n        # Early stopping\n        if patience_counter >= EARLY_STOPPING_PATIENCE:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n        \n        print(\"-\"*60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING COMPLETE\")\n    print(f\"Best Validation Dice: {best_val_dice:.4f}\")\n    print(\"=\"*60)\n    \n    # ============================================\n    # FINAL TEST ON FRAGMENT 1\n    # ============================================\n    print(\"\\n6. Testing on Fragment 1...\")\n    print(\"=\"*60)\n    \n    # Load best model\n    checkpoint = torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth'))\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    # Load test data\n    print(\"Loading test data...\")\n    test_volume = load_volume_fast(fragment1_path)\n    test_mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ignore = generate_ignore_mask(test_mask)\n    \n    # Extract test patches\n    print(\"Extracting test patches...\")\n    test_patches, test_masks, test_ignores = extract_patches_fast(\n        test_volume, test_mask, test_ignore, max_patches_per_fragment=3000\n    )\n    print(f\"Test patches: {len(test_patches)}\")\n    \n    # Create test dataset\n    test_dataset = VesuviusDataset(\n        test_patches, test_masks, test_ignores, transform=None\n    )\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    # Evaluate\n    print(\"Running inference...\")\n    test_dice = 0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, ignores in tqdm(test_loader, desc=\"Testing\"):\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            ignores = ignores.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            valid_mask = (1 - ignores)\n            \n            # Calculate dice\n            intersection = ((preds * masks) * valid_mask).sum()\n            union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n            dice = (2 * intersection) / (union + 1e-6)\n            test_dice += dice.item()\n            \n            # Store for metrics\n            all_preds.append((preds * valid_mask).cpu().numpy())\n            all_targets.append((masks * valid_mask).cpu().numpy())\n    \n    avg_test_dice = test_dice / len(test_loader)\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_targets = np.concatenate([t.flatten() for t in all_targets])\n    \n    tn, fp, fn, tp = confusion_matrix(all_targets, all_preds, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    \n    # ============================================\n    # FINAL RESULTS\n    # ============================================\n    total_time = time.time() - start_time\n    hours = int(total_time // 3600)\n    minutes = int((total_time % 3600) // 60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"FINAL TEST RESULTS ON FRAGMENT 1\")\n    print(\"=\"*60)\n    print(f\"Total Time: {hours}h {minutes}m\")\n    print(f\"Test Dice Score: {avg_test_dice:.4f}\")\n    print(f\"Precision: {precision:.4f}\")\n    print(f\"Recall: {recall:.4f}\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(\"-\"*60)\n    print(\"Confusion Matrix:\")\n    print(f\"  True Positives: {tp}\")\n    print(f\"  True Negatives: {tn}\")\n    print(f\"  False Positives: {fp}\")\n    print(f\"  False Negatives: {fn}\")\n    print(\"=\"*60)\n    \n    # Save results\n    with open(os.path.join(OUTPUT_DIR, \"final_results.txt\"), \"w\") as f:\n        f.write(\"VESUVIUS INK DETECTION - FINAL RESULTS\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n        f.write(f\"Total Time: {hours}h {minutes}m\\n\")\n        f.write(f\"Best Validation Dice: {best_val_dice:.4f}\\n\")\n        f.write(f\"Test Dice Score: {avg_test_dice:.4f}\\n\")\n        f.write(f\"Precision: {precision:.4f}\\n\")\n        f.write(f\"Recall: {recall:.4f}\\n\")\n        f.write(f\"F1 Score: {f1:.4f}\\n\")\n        f.write(\"-\"*60 + \"\\n\")\n        f.write(f\"True Positives: {tp}\\n\")\n        f.write(f\"True Negatives: {tn}\\n\")\n        f.write(f\"False Positives: {fp}\\n\")\n        f.write(f\"False Negatives: {fn}\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n    \n    print(f\"\\nResults saved to: {os.path.join(OUTPUT_DIR, 'final_results.txt')}\")\n    \n    # Plot training curves\n    plt.figure(figsize=(12, 4))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses)\n    plt.title('Training Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.grid(True)\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(val_dice_scores)\n    plt.title('Validation Dice Score')\n    plt.xlabel('Epoch')\n    plt.ylabel('Dice')\n    plt.grid(True)\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'))\n    plt.show()\n    \n    return avg_test_dice\n\nif __name__ == \"__main__\":\n    try:\n        test_dice = main()\n        print(f\"\\n✅ Final Test Dice Score: {test_dice:.4f}\")\n    except Exception as e:\n        print(f\"\\n❌ Error: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# TEST PREDICTION SCRIPT - NO TRAINING\n# Save full volume predictions and visualizations\n# ============================================\n\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom skimage import measure\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data path\nMODEL_PATH = \"/kaggle/working/best_model.pth\"  # Path to your trained model\nOUTPUT_DIR = \"/kaggle/working/test_predictions/\"\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n# Volume slice settings (must match training)\nSLICE_START = 20\nSLICE_END = 32\nPATCH_SIZE = 128\nSTRIDE = 64\nBATCH_SIZE = 14\n\nprint(f\"Using device: {DEVICE}\")\nprint(f\"Model path: {MODEL_PATH}\")\nprint(f\"Output directory: {OUTPUT_DIR}\")\n\n# Create output directory\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# ============================================\n# MODEL DEFINITION (Must match training)\n# ============================================\nclass MultiSlice2DUNet(nn.Module):\n    \"\"\"2D CNN that treats depth slices as input channels\"\"\"\n    def __init__(self, in_channels=12, out_channels=1):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._conv_block(in_channels, 32)\n        self.enc2 = self._conv_block(32, 64)\n        self.enc3 = self._conv_block(64, 128)\n        self.enc4 = self._conv_block(128, 256)\n        \n        # Bottleneck\n        self.bottleneck = self._conv_block(256, 512)\n        \n        # Decoder\n        self.dec4 = self._upconv_block(512 + 256, 256)\n        self.dec3 = self._upconv_block(256 + 128, 128)\n        self.dec2 = self._upconv_block(128 + 64, 64)\n        self.dec1 = self._upconv_block(64 + 32, 32)\n        \n        # Output\n        self.final = nn.Conv2d(32, out_channels, kernel_size=1)\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False)\n        \n    def _conv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def _upconv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        p1 = self.pool(e1)\n        \n        e2 = self.enc2(p1)\n        p2 = self.pool(e2)\n        \n        e3 = self.enc3(p2)\n        p3 = self.pool(e3)\n        \n        e4 = self.enc4(p3)\n        p4 = self.pool(e4)\n        \n        # Bottleneck\n        b = self.bottleneck(p4)\n        \n        # Decoder with skip connections\n        d4 = self.upsample(b)\n        d4 = torch.cat([d4, e4], dim=1)\n        d4 = self.dec4(d4)\n        \n        d3 = self.upsample(d4)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.upsample(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.upsample(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        # Final output\n        out = self.final(d1)\n        \n        return out\n\n# ============================================\n# LOAD MODEL\n# ============================================\ndef load_model(model_path):\n    \"\"\"Load the trained model\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"LOADING TRAINED MODEL\")\n    print(\"=\"*60)\n    \n    if not os.path.exists(model_path):\n        raise FileNotFoundError(f\"Model not found at {model_path}\")\n    \n    model = MultiSlice2DUNet(in_channels=12, out_channels=1).to(DEVICE)\n    checkpoint = torch.load(model_path, map_location=DEVICE)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    print(f\"✓ Model loaded successfully\")\n    print(f\"  Epoch: {checkpoint.get('epoch', 'N/A')}\")\n    print(f\"  Validation Dice: {checkpoint.get('val_dice', 'N/A'):.4f}\")\n    \n    return model\n\n# ============================================\n# LOAD VOLUME AND MASK\n# ============================================\ndef load_volume_and_mask(fragment_path):\n    \"\"\"Load the full volume and ground truth mask\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"LOADING TEST DATA\")\n    print(\"=\"*60)\n    \n    # Load volume\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    print(f\"Loading slices {SLICE_START} to {SLICE_END-1}...\")\n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    \n    volume = np.stack(slices, axis=-1)\n    volume = volume.astype(np.float32)\n    \n    print(f\"  Volume shape: {volume.shape}\")\n    \n    # Normalize volume per slice (same as training)\n    print(\"  Normalizing volume...\")\n    for i in range(volume.shape[-1]):\n        slice_data = volume[:, :, i]\n        mean, std = slice_data.mean(), slice_data.std()\n        volume[:, :, i] = (slice_data - mean) / (std + 1e-6)\n    \n    # Load ground truth mask\n    mask_path = os.path.join(fragment_path, \"inklabels.png\")\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, 0)\n        mask = (mask > 0).astype(np.uint8)\n        print(f\"  Mask shape: {mask.shape}\")\n        print(f\"  Ink pixels: {mask.sum():,}\")\n    else:\n        mask = None\n        print(\"  No ground truth mask found\")\n    \n    return volume, mask\n\n# ============================================\n# SLIDING WINDOW PREDICTION\n# ============================================\ndef predict_full_volume(model, volume, batch_size=BATCH_SIZE):\n    \"\"\"Predict on full volume using sliding window\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"PREDICTING FULL VOLUME\")\n    print(\"=\"*60)\n    \n    H, W, C = volume.shape\n    print(f\"Volume dimensions: {H}x{W}x{C}\")\n    \n    # Initialize prediction and count maps\n    pred_map = np.zeros((H, W), dtype=np.float32)\n    count_map = np.zeros((H, W), dtype=np.float32)\n    \n    # Calculate number of patches\n    n_y = (H - PATCH_SIZE) // STRIDE + 1\n    n_x = (W - PATCH_SIZE) // STRIDE + 1\n    total_patches = n_y * n_x\n    \n    print(f\"Total patches to process: {total_patches}\")\n    print(f\"Patch size: {PATCH_SIZE}, Stride: {STRIDE}\")\n    \n    # Prepare batches\n    patches = []\n    positions = []\n    \n    # Collect all patches\n    print(\"Collecting patches...\")\n    for y in tqdm(range(0, H - PATCH_SIZE + 1, STRIDE), desc=\"Collecting\"):\n        for x in range(0, W - PATCH_SIZE + 1, STRIDE):\n            patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE, :]\n            patches.append(patch)\n            positions.append((y, x))\n    \n    # Process in batches\n    print(f\"\\nProcessing {len(patches)} patches in batches of {batch_size}...\")\n    \n    with torch.no_grad():\n        for i in tqdm(range(0, len(patches), batch_size), desc=\"Predicting\"):\n            batch_patches = patches[i:i+batch_size]\n            batch_positions = positions[i:i+batch_size]\n            \n            # Prepare batch tensor\n            batch_tensor = np.stack(batch_patches, axis=0)\n            batch_tensor = torch.tensor(batch_tensor).permute(0, 3, 1, 2).float().to(DEVICE)\n            \n            # Predict\n            with torch.cuda.amp.autocast(enabled=True):\n                outputs = model(batch_tensor)\n                preds = torch.sigmoid(outputs).cpu().numpy()[:, 0, :, :]\n            \n            # Aggregate predictions\n            for j, (y, x) in enumerate(batch_positions):\n                pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += preds[j]\n                count_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += 1\n    \n    # Average overlapping predictions\n    count_map[count_map == 0] = 1  # Avoid division by zero\n    pred_map = pred_map / count_map\n    \n    print(f\"✓ Prediction complete\")\n    print(f\"  Prediction range: [{pred_map.min():.4f}, {pred_map.max():.4f}]\")\n    \n    return pred_map, count_map\n\n# ============================================\n# VISUALIZATION FUNCTIONS\n# ============================================\ndef create_comparison_visualization(volume_slice, ground_truth, prediction, \n                                   prediction_binary, save_path):\n    \"\"\"Create comprehensive comparison visualization\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"CREATING VISUALIZATIONS\")\n    print(\"=\"*60)\n    \n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n    \n    # Row 1: Input, Ground Truth, Overlay\n    # Input volume (middle slice)\n    mid_slice = volume_slice.shape[-1] // 2\n    axes[0, 0].imshow(volume_slice[:, :, mid_slice], cmap='gray')\n    axes[0, 0].set_title('Input Volume (Middle Slice)', fontsize=14)\n    axes[0, 0].axis('off')\n    \n    # Ground truth\n    if ground_truth is not None:\n        axes[0, 1].imshow(ground_truth, cmap='Reds')\n        axes[0, 1].set_title('Ground Truth', fontsize=14)\n        axes[0, 1].axis('off')\n        \n        # Overlay\n        overlay = np.zeros((*ground_truth.shape, 3))\n        overlay[..., 0] = ground_truth  # Red channel for ground truth\n        overlay[..., 1] = prediction_binary  # Green channel for prediction\n        \n        axes[0, 2].imshow(overlay)\n        axes[0, 2].set_title('Overlay (Red: GT, Green: Pred, Yellow: Overlap)', fontsize=14)\n        axes[0, 2].axis('off')\n    else:\n        axes[0, 1].text(0.5, 0.5, 'No Ground Truth', \n                       ha='center', va='center', transform=axes[0, 1].transAxes)\n        axes[0, 1].set_title('Ground Truth (Not Available)', fontsize=14)\n        axes[0, 1].axis('off')\n        axes[0, 2].axis('off')\n    \n    # Row 2: Prediction (probability), Prediction (binary), Difference\n    # Prediction probability\n    im1 = axes[1, 0].imshow(prediction, cmap='hot', vmin=0, vmax=1)\n    axes[1, 0].set_title('Prediction Probability', fontsize=14)\n    axes[1, 0].axis('off')\n    plt.colorbar(im1, ax=axes[1, 0], fraction=0.046)\n    \n    # Binary prediction\n    axes[1, 1].imshow(prediction_binary, cmap='Greens')\n    axes[1, 1].set_title(f'Binary Prediction (Threshold=0.5)\\nInk pixels: {prediction_binary.sum():,}', \n                         fontsize=14)\n    axes[1, 1].axis('off')\n    \n    # Difference (if ground truth available)\n    if ground_truth is not None:\n        diff = np.zeros((*ground_truth.shape, 3), dtype=np.uint8)\n        # True Positives: White\n        diff[(ground_truth == 1) & (prediction_binary == 1)] = [255, 255, 255]\n        # False Positives: Blue\n        diff[(ground_truth == 0) & (prediction_binary == 1)] = [0, 0, 255]\n        # False Negatives: Red\n        diff[(ground_truth == 1) & (prediction_binary == 0)] = [255, 0, 0]\n        # True Negatives: Black\n        \n        axes[1, 2].imshow(diff)\n        axes[1, 2].set_title('Error Map\\nWhite: TP, Blue: FP, Red: FN, Black: TN', fontsize=14)\n        axes[1, 2].axis('off')\n        \n        # Add legend\n        from matplotlib.patches import Patch\n        legend_elements = [\n            Patch(facecolor='white', label='True Positive'),\n            Patch(facecolor='blue', label='False Positive'),\n            Patch(facecolor='red', label='False Negative'),\n            Patch(facecolor='black', label='True Negative')\n        ]\n        axes[1, 2].legend(handles=legend_elements, loc='lower right', fontsize=10)\n    else:\n        axes[1, 2].axis('off')\n    \n    plt.suptitle('Full Volume Prediction Comparison', fontsize=16, y=0.98)\n    plt.tight_layout()\n    \n    # Save figure\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"✓ Visualization saved to: {save_path}\")\n\ndef create_detailed_metrics_plot(ground_truth, prediction_binary, save_path):\n    \"\"\"Create detailed metrics visualization\"\"\"\n    if ground_truth is None:\n        print(\"  Skipping metrics plot (no ground truth)\")\n        return\n    \n    fig, axes = plt.subplots(2, 2, figsize=(12, 12))\n    \n    # Calculate metrics\n    tp = np.sum((ground_truth == 1) & (prediction_binary == 1))\n    tn = np.sum((ground_truth == 0) & (prediction_binary == 0))\n    fp = np.sum((ground_truth == 0) & (prediction_binary == 1))\n    fn = np.sum((ground_truth == 1) & (prediction_binary == 0))\n    \n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n    iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n    \n    # 1. Confusion Matrix\n    cm = np.array([[tn, fp], [fn, tp]])\n    im1 = axes[0, 0].imshow(cm, cmap='Blues', interpolation='nearest')\n    axes[0, 0].set_title('Confusion Matrix', fontsize=14)\n    axes[0, 0].set_xlabel('Predicted')\n    axes[0, 0].set_ylabel('Actual')\n    axes[0, 0].set_xticks([0, 1])\n    axes[0, 0].set_yticks([0, 1])\n    axes[0, 0].set_xticklabels(['Negative', 'Positive'])\n    axes[0, 0].set_yticklabels(['Negative', 'Positive'])\n    \n    # Add text annotations\n    for i in range(2):\n        for j in range(2):\n            axes[0, 0].text(j, i, f'{cm[i, j]:,}', \n                           ha='center', va='center', \n                           color='white' if cm[i, j] > cm.max()/2 else 'black')\n    \n    plt.colorbar(im1, ax=axes[0, 0])\n    \n    # 2. Metrics Bar Chart\n    metrics = ['Precision', 'Recall', 'F1-Score', 'Dice', 'IoU']\n    values = [precision, recall, f1, dice, iou]\n    colors = ['#2ecc71', '#3498db', '#9b59b6', '#e74c3c', '#f39c12']\n    \n    bars = axes[0, 1].bar(metrics, values, color=colors)\n    axes[0, 1].set_title('Performance Metrics', fontsize=14)\n    axes[0, 1].set_ylim([0, 1])\n    axes[0, 1].set_ylabel('Score')\n    axes[0, 1].grid(axis='y', alpha=0.3)\n    \n    # Add value labels on bars\n    for bar, value in zip(bars, values):\n        axes[0, 1].text(bar.get_x() + bar.get_width()/2, \n                       bar.get_height() + 0.02,\n                       f'{value:.4f}', \n                       ha='center', va='bottom', fontsize=10)\n    \n    # 3. Pixel Distribution Pie Chart\n    sizes = [tn, fp, fn, tp]\n    labels = [f'True Neg\\n{tn:,}', f'False Pos\\n{fp:,}', \n              f'False Neg\\n{fn:,}', f'True Pos\\n{tp:,}']\n    colors_pie = ['#95a5a6', '#3498db', '#e74c3c', '#2ecc71']\n    \n    axes[1, 0].pie(sizes, labels=labels, colors=colors_pie, autopct='%1.1f%%', startangle=90)\n    axes[1, 0].set_title('Pixel Classification Distribution', fontsize=14)\n    \n    # 4. Summary Statistics\n    axes[1, 1].axis('off')\n    summary_text = f\"\"\"\n    PERFORMANCE SUMMARY\n    {'='*40}\n    \n    Pixel Statistics:\n    • Total Pixels: {ground_truth.size:,}\n    • Ink Pixels (GT): {ground_truth.sum():,}\n    • Ink Pixels (Pred): {prediction_binary.sum():,}\n    \n    Classification Results:\n    • True Positives: {tp:,}\n    • True Negatives: {tn:,}\n    • False Positives: {fp:,}\n    • False Negatives: {fn:,}\n    \n    Metrics:\n    • Precision: {precision:.4f}\n    • Recall: {recall:.4f}\n    • F1-Score: {f1:.4f}\n    • Dice Score: {dice:.4f}\n    • IoU: {iou:.4f}\n    \n    Accuracy: {(tp+tn)/ground_truth.size:.4f}\n    \"\"\"\n    \n    axes[1, 1].text(0.1, 0.9, summary_text, transform=axes[1, 1].transAxes,\n                   fontsize=11, verticalalignment='top', fontfamily='monospace',\n                   bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n    \n    plt.suptitle('Detailed Performance Analysis', fontsize=16, y=0.98)\n    plt.tight_layout()\n    \n    # Save figure\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"✓ Detailed metrics plot saved to: {save_path}\")\n\ndef save_prediction_slices(prediction, volume, save_dir, num_slices=5):\n    \"\"\"Save individual slices with predictions\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"SAVING PREDICTION SLICES\")\n    print(\"=\"*60)\n    \n    slices_dir = os.path.join(save_dir, \"slices\")\n    os.makedirs(slices_dir, exist_ok=True)\n    \n    H, W, C = volume.shape\n    slice_indices = np.linspace(0, C-1, min(num_slices, C), dtype=int)\n    \n    fig, axes = plt.subplots(len(slice_indices), 2, figsize=(12, 4*len(slice_indices)))\n    if len(slice_indices) == 1:\n        axes = axes.reshape(1, -1)\n    \n    for idx, slice_idx in enumerate(slice_indices):\n        # Input slice\n        axes[idx, 0].imshow(volume[:, :, slice_idx], cmap='gray')\n        axes[idx, 0].set_title(f'Input Slice {slice_idx}', fontsize=12)\n        axes[idx, 0].axis('off')\n        \n        # Prediction overlay\n        overlay = np.zeros((H, W, 3))\n        overlay[..., 0] = volume[:, :, slice_idx] / volume[:, :, slice_idx].max()  # Normalize\n        overlay[..., 1] = prediction  # Green channel for prediction\n        \n        axes[idx, 1].imshow(overlay)\n        axes[idx, 1].set_title(f'Prediction Overlay on Slice {slice_idx}', fontsize=12)\n        axes[idx, 1].axis('off')\n    \n    plt.suptitle('Prediction on Different Volume Slices', fontsize=14, y=0.98)\n    plt.tight_layout()\n    \n    save_path = os.path.join(slices_dir, \"slice_predictions.png\")\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"✓ Slice visualizations saved to: {save_path}\")\n\n# ============================================\n# SAVE RESULTS\n# ============================================\ndef save_all_results(prediction, prediction_binary, volume, ground_truth, output_dir):\n    \"\"\"Save all prediction results in various formats\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"SAVING RESULTS\")\n    print(\"=\"*60)\n    \n    # 1. Save binary prediction as PNG\n    binary_path = os.path.join(output_dir, \"binary_prediction.png\")\n    cv2.imwrite(binary_path, (prediction_binary * 255).astype(np.uint8))\n    print(f\"✓ Binary prediction saved to: {binary_path}\")\n    \n    # 2. Save probability map as TIFF\n    prob_path = os.path.join(output_dir, \"probability_map.tiff\")\n    tifffile.imwrite(prob_path, (prediction * 255).astype(np.uint8))\n    print(f\"✓ Probability map saved to: {prob_path}\")\n    \n    # 3. Save probability map as NPY (for further analysis)\n    npy_path = os.path.join(output_dir, \"probability_map.npy\")\n    np.save(npy_path, prediction)\n    print(f\"✓ Probability map (numpy) saved to: {npy_path}\")\n    \n    # 4. Create and save comparison visualization\n    comparison_path = os.path.join(output_dir, \"comparison_visualization.png\")\n    create_comparison_visualization(volume, ground_truth, prediction, \n                                   prediction_binary, comparison_path)\n    \n    # 5. Create detailed metrics plot (if ground truth available)\n    if ground_truth is not None:\n        metrics_path = os.path.join(output_dir, \"detailed_metrics.png\")\n        create_detailed_metrics_plot(ground_truth, prediction_binary, metrics_path)\n    \n    # 6. Save slice predictions\n    save_prediction_slices(prediction, volume, output_dir, num_slices=4)\n    \n    # 7. Save statistics to text file\n    stats_path = os.path.join(output_dir, \"prediction_statistics.txt\")\n    with open(stats_path, 'w') as f:\n        f.write(\"FULL VOLUME PREDICTION STATISTICS\\n\")\n        f.write(\"=\"*60 + \"\\n\\n\")\n        f.write(f\"Volume shape: {volume.shape}\\n\")\n        f.write(f\"Prediction shape: {prediction.shape}\\n\\n\")\n        \n        f.write(\"Prediction Statistics:\\n\")\n        f.write(f\"  Min probability: {prediction.min():.6f}\\n\")\n        f.write(f\"  Max probability: {prediction.max():.6f}\\n\")\n        f.write(f\"  Mean probability: {prediction.mean():.6f}\\n\")\n        f.write(f\"  Std probability: {prediction.std():.6f}\\n\\n\")\n        \n        f.write(\"Binary Prediction (Threshold=0.5):\\n\")\n        f.write(f\"  Total pixels: {prediction_binary.size:,}\\n\")\n        f.write(f\"  Positive pixels: {prediction_binary.sum():,}\\n\")\n        f.write(f\"  Positive ratio: {prediction_binary.mean():.4f}\\n\\n\")\n        \n        if ground_truth is not None:\n            # Calculate metrics\n            tp = np.sum((ground_truth == 1) & (prediction_binary == 1))\n            tn = np.sum((ground_truth == 0) & (prediction_binary == 0))\n            fp = np.sum((ground_truth == 0) & (prediction_binary == 1))\n            fn = np.sum((ground_truth == 1) & (prediction_binary == 0))\n            \n            precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n            recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n            f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n            dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n            iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n            accuracy = (tp + tn) / ground_truth.size\n            \n            f.write(\"Comparison with Ground Truth:\\n\")\n            f.write(f\"  Ground truth positive pixels: {ground_truth.sum():,}\\n\")\n            f.write(f\"  True Positives: {tp:,}\\n\")\n            f.write(f\"  True Negatives: {tn:,}\\n\")\n            f.write(f\"  False Positives: {fp:,}\\n\")\n            f.write(f\"  False Negatives: {fn:,}\\n\\n\")\n            \n            f.write(\"Performance Metrics:\\n\")\n            f.write(f\"  Precision: {precision:.6f}\\n\")\n            f.write(f\"  Recall: {recall:.6f}\\n\")\n            f.write(f\"  F1-Score: {f1:.6f}\\n\")\n            f.write(f\"  Dice Score: {dice:.6f}\\n\")\n            f.write(f\"  IoU: {iou:.6f}\\n\")\n            f.write(f\"  Accuracy: {accuracy:.6f}\\n\")\n    \n    print(f\"✓ Statistics saved to: {stats_path}\")\n\n# ============================================\n# MAIN FUNCTION\n# ============================================\ndef main():\n    \"\"\"Main function to run full volume prediction and save results\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"FULL VOLUME TEST PREDICTION PIPELINE\")\n    print(\"=\"*60)\n    \n    try:\n        # 1. Load model\n        model = load_model(MODEL_PATH)\n        \n        # 2. Load test data\n        volume, ground_truth = load_volume_and_mask(fragment1_path)\n        \n        # 3. Predict on full volume\n        prediction, count_map = predict_full_volume(model, volume)\n        \n        # 4. Create binary prediction\n        prediction_binary = (prediction > 0.5).astype(np.uint8)\n        \n        # 5. Save all results\n        save_all_results(prediction, prediction_binary, volume, ground_truth, OUTPUT_DIR)\n        \n        # 6. Print summary\n        print(\"\\n\" + \"=\"*60)\n        print(\"PREDICTION COMPLETE - SUMMARY\")\n        print(\"=\"*60)\n        print(f\"Output directory: {OUTPUT_DIR}\")\n        print(f\"Files saved:\")\n        print(f\"  - binary_prediction.png\")\n        print(f\"  - probability_map.tiff\")\n        print(f\"  - probability_map.npy\")\n        print(f\"  - comparison_visualization.png\")\n        if ground_truth is not None:\n            print(f\"  - detailed_metrics.png\")\n        print(f\"  - slices/slice_predictions.png\")\n        print(f\"  - prediction_statistics.txt\")\n        print(\"=\"*60)\n        \n        # Print metrics if ground truth available\n        if ground_truth is not None:\n            tp = np.sum((ground_truth == 1) & (prediction_binary == 1))\n            tn = np.sum((ground_truth == 0) & (prediction_binary == 0))\n            fp = np.sum((ground_truth == 0) & (prediction_binary == 1))\n            fn = np.sum((ground_truth == 1) & (prediction_binary == 0))\n            \n            dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n            f1 = 2 * tp / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n            \n            print(f\"\\nTest Metrics:\")\n            print(f\"  Dice Score: {dice:.4f}\")\n            print(f\"  F1 Score: {f1:.4f}\")\n            print(f\"  TP: {tp:,}, TN: {tn:,}, FP: {fp:,}, FN: {fn:,}\")\n        \n        print(\"\\n✅ All predictions and visualizations saved successfully!\")\n        \n    except Exception as e:\n        print(f\"\\n❌ Error during prediction: {e}\")\n        import traceback\n        traceback.print_exc()\n\n# ============================================\n# RUN SCRIPT\n# ============================================\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# TEST PREDICTION SCRIPT - MATCHES TRAINING EVALUATION\n# Save predictions with consistent evaluation\n# ============================================\n\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ============================================\n# CONFIGURATION - MUST MATCH TRAINING\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nMODEL_PATH = \"/kaggle/working/best_model.pth\"\nOUTPUT_DIR = \"/kaggle/working/test_predictions/\"\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n# These MUST match your training configuration exactly\nPATCH_SIZE = 128\nSTRIDE = 64\nBATCH_SIZE = 14\nSLICE_START = 20\nSLICE_END = 32\nNUM_WORKERS = 2\nPIN_MEMORY = True\n\nprint(f\"Using device: {DEVICE}\")\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# ============================================\n# IGNORE MASK GENERATION (Must match training)\n# ============================================\ndef generate_ignore_mask(ink_mask, distance_threshold=3, erosion_size=2):\n    \"\"\"Generate ignore mask for uncertain regions - EXACT COPY from training\"\"\"\n    ignore_mask = np.zeros_like(ink_mask, dtype=np.uint8)\n    \n    # Simple boundary detection\n    kernel = np.ones((3, 3), np.uint8)\n    eroded = cv2.erode(ink_mask.astype(np.uint8), kernel, iterations=1)\n    dilated = cv2.dilate(ink_mask.astype(np.uint8), kernel, iterations=1)\n    \n    # Boundaries are where eroded and dilated differ from original\n    boundaries = (dilated != eroded)\n    ignore_mask[boundaries] = 1\n    \n    # Remove small isolated ink dots\n    contours, _ = cv2.findContours(ink_mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for contour in contours:\n        if cv2.contourArea(contour) < 10:\n            cv2.drawContours(ignore_mask, [contour], -1, 1, -1)\n    \n    return ignore_mask\n\n# ============================================\n# MODEL DEFINITION (Must match training exactly)\n# ============================================\nclass MultiSlice2DUNet(nn.Module):\n    \"\"\"2D CNN that treats depth slices as input channels\"\"\"\n    def __init__(self, in_channels=12, out_channels=1):\n        super().__init__()\n        \n        self.enc1 = self._conv_block(in_channels, 32)\n        self.enc2 = self._conv_block(32, 64)\n        self.enc3 = self._conv_block(64, 128)\n        self.enc4 = self._conv_block(128, 256)\n        self.bottleneck = self._conv_block(256, 512)\n        self.dec4 = self._upconv_block(512 + 256, 256)\n        self.dec3 = self._upconv_block(256 + 128, 128)\n        self.dec2 = self._upconv_block(128 + 64, 64)\n        self.dec1 = self._upconv_block(64 + 32, 32)\n        self.final = nn.Conv2d(32, out_channels, kernel_size=1)\n        self.pool = nn.MaxPool2d(2)\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False)\n        \n    def _conv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def _upconv_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        e1 = self.enc1(x)\n        p1 = self.pool(e1)\n        e2 = self.enc2(p1)\n        p2 = self.pool(e2)\n        e3 = self.enc3(p2)\n        p3 = self.pool(e3)\n        e4 = self.enc4(p3)\n        p4 = self.pool(e4)\n        b = self.bottleneck(p4)\n        d4 = self.upsample(b)\n        d4 = torch.cat([d4, e4], dim=1)\n        d4 = self.dec4(d4)\n        d3 = self.upsample(d4)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        d2 = self.upsample(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        d1 = self.upsample(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        out = self.final(d1)\n        return out\n\n# ============================================\n# DATA LOADING (Must match training exactly)\n# ============================================\ndef load_volume_fast(fragment_path):\n    \"\"\"Fast volume loading - EXACT COPY from training\"\"\"\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    \n    volume = np.stack(slices, axis=-1)\n    volume = volume.astype(np.float32)\n    \n    # Fast normalization per slice\n    for i in range(volume.shape[-1]):\n        slice_data = volume[:, :, i]\n        mean, std = slice_data.mean(), slice_data.std()\n        volume[:, :, i] = (slice_data - mean) / (std + 1e-6)\n    \n    return volume\n\ndef extract_patches_with_positions(volume, mask, ignore_mask, stride=None):\n    \"\"\"Extract ALL patches with their positions for reconstruction\"\"\"\n    if stride is None:\n        stride = STRIDE\n    \n    H, W, _ = volume.shape\n    patches = []\n    positions = []\n    mask_patches = []\n    ignore_patches = []\n    \n    for y in range(0, H - PATCH_SIZE + 1, stride):\n        for x in range(0, W - PATCH_SIZE + 1, stride):\n            v_patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            m_patch = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            i_patch = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            \n            patches.append(v_patch)\n            positions.append((y, x))\n            mask_patches.append(m_patch)\n            ignore_patches.append(i_patch)\n    \n    return patches, positions, mask_patches, ignore_patches\n\nclass VesuviusDataset(Dataset):\n    \"\"\"Dataset class - EXACT COPY from training\"\"\"\n    def __init__(self, volumes, masks, ignores):\n        self.volumes = volumes\n        self.masks = masks\n        self.ignores = ignores\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        image = self.volumes[idx].copy()\n        mask = self.masks[idx].copy()\n        ignore = self.ignores[idx].copy()\n        \n        # Convert to tensors\n        image = torch.tensor(image).permute(2, 0, 1).float()\n        mask = torch.tensor(mask).float().unsqueeze(0)\n        ignore = torch.tensor(ignore).float().unsqueeze(0)\n        \n        return image, mask, ignore\n\ndef collate_fn(batch):\n    \"\"\"Custom collate function - EXACT COPY from training\"\"\"\n    images = torch.stack([item[0] for item in batch])\n    masks = torch.stack([item[1] for item in batch])\n    ignores = torch.stack([item[2] for item in batch])\n    return images, masks, ignores\n\n# ============================================\n# PATCH-BASED EVALUATION (Matches training)\n# ============================================\ndef evaluate_patch_based(model, test_loader):\n    \"\"\"Evaluate using patch-based method (same as training)\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"PATCH-BASED EVALUATION (Matches Training)\")\n    print(\"=\"*60)\n    \n    model.eval()\n    test_dice = 0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, ignores in tqdm(test_loader, desc=\"Evaluating patches\"):\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            ignores = ignores.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=True):\n                outputs = model(images)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            valid_mask = (1 - ignores)\n            \n            # Calculate dice\n            intersection = ((preds * masks) * valid_mask).sum()\n            union = ((preds * valid_mask).sum() + (masks * valid_mask).sum())\n            dice = (2 * intersection) / (union + 1e-6)\n            test_dice += dice.item()\n            \n            # Store for metrics\n            all_preds.append((preds * valid_mask).cpu().numpy())\n            all_targets.append((masks * valid_mask).cpu().numpy())\n    \n    avg_test_dice = test_dice / len(test_loader)\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_targets = np.concatenate([t.flatten() for t in all_targets])\n    \n    tn, fp, fn, tp = confusion_matrix(all_targets, all_preds, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    \n    return avg_test_dice, precision, recall, f1, tn, fp, fn, tp\n\n# ============================================\n# FULL VOLUME RECONSTRUCTION\n# ============================================\ndef reconstruct_full_volume(model, patches, positions, volume_shape):\n    \"\"\"Reconstruct full volume from patches\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"RECONSTRUCTING FULL VOLUME\")\n    print(\"=\"*60)\n    \n    H, W = volume_shape\n    pred_map = np.zeros((H, W), dtype=np.float32)\n    count_map = np.zeros((H, W), dtype=np.float32)\n    \n    # Process patches in batches\n    dataset = VesuviusDataset(patches, \n                             [np.zeros((PATCH_SIZE, PATCH_SIZE))]*len(patches), \n                             [np.zeros((PATCH_SIZE, PATCH_SIZE))]*len(patches))\n    loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=False, \n                       num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY, \n                       collate_fn=collate_fn)\n    \n    print(f\"Processing {len(patches)} patches...\")\n    \n    with torch.no_grad():\n        patch_idx = 0\n        for images, _, _ in tqdm(loader, desc=\"Reconstructing\"):\n            images = images.to(DEVICE)\n            \n            with torch.cuda.amp.autocast(enabled=True):\n                outputs = model(images)\n                preds = torch.sigmoid(outputs).cpu().numpy()[:, 0, :, :]\n            \n            # Place predictions back\n            for j in range(len(preds)):\n                y, x = positions[patch_idx]\n                pred_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += preds[j]\n                count_map[y:y+PATCH_SIZE, x:x+PATCH_SIZE] += 1\n                patch_idx += 1\n    \n    # Average overlapping predictions\n    count_map[count_map == 0] = 1\n    pred_map = pred_map / count_map\n    \n    return pred_map\n\n# ============================================\n# VISUALIZATION FUNCTIONS\n# ============================================\ndef create_comprehensive_visualization(volume, ground_truth, prediction, \n                                      prediction_binary, save_path):\n    \"\"\"Create comprehensive visualization\"\"\"\n    \n    fig = plt.figure(figsize=(20, 12))\n    \n    # 1. Input volume (middle slice)\n    ax1 = plt.subplot(2, 3, 1)\n    mid_slice = volume.shape[-1] // 2\n    ax1.imshow(volume[:, :, mid_slice], cmap='gray')\n    ax1.set_title('Input Volume (Middle Slice)', fontsize=14)\n    ax1.axis('off')\n    \n    # 2. Ground truth\n    ax2 = plt.subplot(2, 3, 2)\n    if ground_truth is not None:\n        ax2.imshow(ground_truth, cmap='Reds')\n        ax2.set_title(f'Ground Truth\\nInk pixels: {ground_truth.sum():,}', fontsize=14)\n    ax2.axis('off')\n    \n    # 3. Prediction probability\n    ax3 = plt.subplot(2, 3, 3)\n    im3 = ax3.imshow(prediction, cmap='hot', vmin=0, vmax=1)\n    ax3.set_title('Prediction Probability', fontsize=14)\n    ax3.axis('off')\n    plt.colorbar(im3, ax=ax3, fraction=0.046)\n    \n    # 4. Binary prediction\n    ax4 = plt.subplot(2, 3, 4)\n    ax4.imshow(prediction_binary, cmap='Greens')\n    ax4.set_title(f'Binary Prediction (Threshold=0.5)\\nInk pixels: {prediction_binary.sum():,}', \n                  fontsize=14)\n    ax4.axis('off')\n    \n    # 5. Overlay\n    ax5 = plt.subplot(2, 3, 5)\n    if ground_truth is not None:\n        overlay = np.zeros((*ground_truth.shape, 3))\n        overlay[..., 0] = ground_truth\n        overlay[..., 1] = prediction_binary\n        ax5.imshow(overlay)\n        ax5.set_title('Overlay (Red: GT, Green: Pred, Yellow: Overlap)', fontsize=14)\n    ax5.axis('off')\n    \n    # 6. Error map\n    ax6 = plt.subplot(2, 3, 6)\n    if ground_truth is not None:\n        diff = np.zeros((*ground_truth.shape, 3), dtype=np.uint8)\n        diff[(ground_truth == 1) & (prediction_binary == 1)] = [255, 255, 255]  # TP: White\n        diff[(ground_truth == 0) & (prediction_binary == 1)] = [0, 0, 255]      # FP: Blue\n        diff[(ground_truth == 1) & (prediction_binary == 0)] = [255, 0, 0]      # FN: Red\n        ax6.imshow(diff)\n        ax6.set_title('Error Map\\nWhite: TP, Blue: FP, Red: FN, Black: TN', fontsize=14)\n        \n        # Add legend\n        from matplotlib.patches import Patch\n        legend_elements = [\n            Patch(facecolor='white', label='True Positive'),\n            Patch(facecolor='blue', label='False Positive'),\n            Patch(facecolor='red', label='False Negative'),\n            Patch(facecolor='black', label='True Negative')\n        ]\n        ax6.legend(handles=legend_elements, loc='lower right', fontsize=10)\n    ax6.axis('off')\n    \n    plt.suptitle('Full Volume Prediction Analysis', fontsize=16, y=0.98)\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"✓ Visualization saved to: {save_path}\")\n\n# ============================================\n# MAIN FUNCTION\n# ============================================\ndef main():\n    print(\"\\n\" + \"=\"*60)\n    print(\"COMPREHENSIVE TEST EVALUATION\")\n    print(\"=\"*60)\n    \n    # 1. Load model\n    print(\"\\n1. Loading trained model...\")\n    model = MultiSlice2DUNet(in_channels=12, out_channels=1).to(DEVICE)\n    checkpoint = torch.load(MODEL_PATH, map_location=DEVICE)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    print(f\"✓ Model loaded (Epoch: {checkpoint.get('epoch', 'N/A')}, Val Dice: {checkpoint.get('val_dice', 'N/A'):.4f})\")\n    \n    # 2. Load test data\n    print(\"\\n2. Loading test data...\")\n    volume = load_volume_fast(fragment1_path)\n    mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    mask = (mask > 0).astype(np.uint8)\n    ignore_mask = generate_ignore_mask(mask)\n    \n    print(f\"   Volume shape: {volume.shape}\")\n    print(f\"   Mask shape: {mask.shape}\")\n    print(f\"   Ink pixels: {mask.sum():,}\")\n    \n    # 3. Extract patches for evaluation\n    print(\"\\n3. Extracting patches for evaluation...\")\n    patches, positions, mask_patches, ignore_patches = extract_patches_with_positions(\n        volume, mask, ignore_mask, stride=STRIDE\n    )\n    print(f\"   Total patches: {len(patches)}\")\n    \n    # 4. Patch-based evaluation (matches training)\n    print(\"\\n4. Running patch-based evaluation...\")\n    test_dataset = VesuviusDataset(patches, mask_patches, ignore_patches)\n    test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False,\n                            num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY,\n                            collate_fn=collate_fn)\n    \n    patch_dice, precision, recall, f1, tn, fp, fn, tp = evaluate_patch_based(model, test_loader)\n    \n    print(f\"\\n📊 PATCH-BASED RESULTS (Matches Training Evaluation):\")\n    print(f\"  Dice Score: {patch_dice:.4f}\")\n    print(f\"  F1 Score: {f1:.4f}\")\n    print(f\"  Precision: {precision:.4f}\")\n    print(f\"  Recall: {recall:.4f}\")\n    \n    # 5. Full volume reconstruction\n    print(\"\\n5. Reconstructing full volume...\")\n    full_prediction = reconstruct_full_volume(model, patches, positions, volume.shape[:2])\n    full_prediction_binary = (full_prediction > 0.5).astype(np.uint8)\n    \n    # 6. Full volume evaluation\n    print(\"\\n6. Evaluating on full volume...\")\n    valid_mask_full = (1 - ignore_mask)\n    \n    # Calculate metrics on full volume\n    tp_full = np.sum((mask == 1) & (full_prediction_binary == 1) & (valid_mask_full == 1))\n    tn_full = np.sum((mask == 0) & (full_prediction_binary == 0) & (valid_mask_full == 1))\n    fp_full = np.sum((mask == 0) & (full_prediction_binary == 1) & (valid_mask_full == 1))\n    fn_full = np.sum((mask == 1) & (full_prediction_binary == 0) & (valid_mask_full == 1))\n    \n    dice_full = (2 * tp_full) / (2 * tp_full + fp_full + fn_full) if (2 * tp_full + fp_full + fn_full) > 0 else 0\n    f1_full = 2 * tp_full / (2 * tp_full + fp_full + fn_full) if (2 * tp_full + fp_full + fn_full) > 0 else 0\n    \n    print(f\"\\n🌍 FULL VOLUME RESULTS:\")\n    print(f\"  Dice Score: {dice_full:.4f}\")\n    print(f\"  F1 Score: {f1_full:.4f}\")\n    print(f\"  TP: {tp_full:,}, TN: {tn_full:,}, FP: {fp_full:,}, FN: {fn_full:,}\")\n    \n    # 7. Save results\n    print(\"\\n7. Saving results...\")\n    \n    # Save binary prediction\n    cv2.imwrite(os.path.join(OUTPUT_DIR, \"full_binary_prediction.png\"), \n                (full_prediction_binary * 255).astype(np.uint8))\n    \n    # Save probability map\n    np.save(os.path.join(OUTPUT_DIR, \"full_probability_map.npy\"), full_prediction)\n    tifffile.imwrite(os.path.join(OUTPUT_DIR, \"full_probability_map.tiff\"), \n                     (full_prediction * 255).astype(np.uint8))\n    \n    # Create visualization\n    create_comprehensive_visualization(volume, mask, full_prediction, \n                                      full_prediction_binary,\n                                      os.path.join(OUTPUT_DIR, \"comprehensive_analysis.png\"))\n    \n    # Save statistics\n    with open(os.path.join(OUTPUT_DIR, \"evaluation_results.txt\"), 'w') as f:\n        f.write(\"COMPREHENSIVE EVALUATION RESULTS\\n\")\n        f.write(\"=\"*60 + \"\\n\\n\")\n        \n        f.write(\"PATCH-BASED EVALUATION (Matches Training):\\n\")\n        f.write(f\"  Dice Score: {patch_dice:.6f}\\n\")\n        f.write(f\"  F1 Score: {f1:.6f}\\n\")\n        f.write(f\"  Precision: {precision:.6f}\\n\")\n        f.write(f\"  Recall: {recall:.6f}\\n\")\n        f.write(f\"  TP: {tp:,}, TN: {tn:,}, FP: {fp:,}, FN: {fn:,}\\n\\n\")\n        \n        f.write(\"FULL VOLUME EVALUATION:\\n\")\n        f.write(f\"  Dice Score: {dice_full:.6f}\\n\")\n        f.write(f\"  F1 Score: {f1_full:.6f}\\n\")\n        f.write(f\"  TP: {tp_full:,}, TN: {tn_full:,}, FP: {fp_full:,}, FN: {fn_full:,}\\n\\n\")\n        \n        f.write(\"DIFFERENCE ANALYSIS:\\n\")\n        f.write(f\"  Patch Dice - Full Volume Dice: {patch_dice - dice_full:.6f}\\n\")\n        f.write(f\"  Reason: Patch-based evaluation samples regions differently\\n\")\n        f.write(f\"  Recommendation: Use full volume evaluation for final assessment\\n\")\n    \n    # 8. Summary\n    print(\"\\n\" + \"=\"*60)\n    print(\"EVALUATION COMPLETE - SUMMARY\")\n    print(\"=\"*60)\n    print(f\"\\n📊 Patch-Based Dice (matches your training): {patch_dice:.4f}\")\n    print(f\"🌍 Full Volume Dice: {dice_full:.4f}\")\n    print(f\"📉 Difference: {patch_dice - dice_full:.4f}\")\n    print(f\"\\n✅ All results saved to: {OUTPUT_DIR}\")\n    print(\"=\"*60)\n    \n    return patch_dice, dice_full\n\nif __name__ == \"__main__\":\n    try:\n        patch_dice, full_dice = main()\n    except Exception as e:\n        print(f\"\\n❌ Error: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"hh\n\nimport os\nimport cv2\nimport numpy as np\nimport tifffile as tiff\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom tqdm import tqdm\nclass CFG:\n    \n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    \n    PATCH_SIZE = 64\n    Z_DIM = 20\n    \n    BATCH_SIZE = 32\n    EPOCHS = 20\n    \n    LR = 1e-4\n#/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume/00.tif    \n    DATA_PATH = \"/kaggle/input/vesuvius-challenge-ink-detection/\"\n    \n    FRAGMENTS = ['1']\ndef load_volume(fragment_id):\n\n    path = \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\"\n    \n    slices = []\n    \n    for i in range(65):\n        img = tiff.imread(path + f\"{i:02}.tif\")\n        slices.append(img)\n\n    volume = np.stack(slices, axis=-1)\n    \n    return volume\n\ndef load_labels(fragment_id):\n\n    label = cv2.imread(\n        f\"{CFG.DATA_PATH}/train/{fragment_id}/inklabels.png\",\n        0\n    )\n    \n    mask = cv2.imread(\n        f\"{CFG.DATA_PATH}/train/{fragment_id}/mask.png\",\n        0\n    )\n    \n    return label, mask\n\ndef sample_patch(volume, label, x, y):\n\n    half = CFG.PATCH_SIZE // 2\n\n    patch = volume[\n        y-half:y+half,\n        x-half:x+half,\n        :CFG.Z_DIM\n    ]\n\n    mask = label[\n        y-half:y+half,\n        x-half:x+half\n    ]\n\n    return patch, mask\nclass InkDataset(Dataset):\n\n    def __init__(self, volume, label, mask):\n\n        self.volume = volume\n        self.label = label\n        self.mask = mask\n        \n        ys, xs = np.where(mask > 0)\n        self.coords = list(zip(xs, ys))\n\n    def __len__(self):\n        return len(self.coords)\n\n    def __getitem__(self, idx):\n\n        x, y = self.coords[idx]\n\n        patch, target = sample_patch(\n            self.volume,\n            self.label,\n            x,\n            y\n        )\n\n        patch = patch.astype(np.float32) / 65535\n        \n        patch = np.transpose(patch, (2,0,1))\n        \n        target = target.astype(np.float32) / 255\n        \n        return torch.tensor(patch), torch.tensor(target)\n\nclass DoubleConv(nn.Module):\n\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch,3,padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(out_ch,out_ch,3,padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self,x):\n        return self.conv(x)\n\nclass UNet(nn.Module):\n\n    def __init__(self, in_ch=20):\n\n        super().__init__()\n\n        self.down1 = DoubleConv(in_ch,64)\n        self.pool1 = nn.MaxPool2d(2)\n\n        self.down2 = DoubleConv(64,128)\n        self.pool2 = nn.MaxPool2d(2)\n\n        self.down3 = DoubleConv(128,256)\n        self.pool3 = nn.MaxPool2d(2)\n\n        self.middle = DoubleConv(256,512)\n\n        self.up3 = nn.ConvTranspose2d(512,256,2,2)\n        self.conv3 = DoubleConv(512,256)\n\n        self.up2 = nn.ConvTranspose2d(256,128,2,2)\n        self.conv2 = DoubleConv(256,128)\n\n        self.up1 = nn.ConvTranspose2d(128,64,2,2)\n        self.conv1 = DoubleConv(128,64)\n\n        self.out = nn.Conv2d(64,1,1)\n\n    def forward(self,x):\n\n        d1 = self.down1(x)\n        d2 = self.down2(self.pool1(d1))\n        d3 = self.down3(self.pool2(d2))\n\n        m = self.middle(self.pool3(d3))\n\n        u3 = self.up3(m)\n        u3 = torch.cat([u3,d3],dim=1)\n        u3 = self.conv3(u3)\n\n        u2 = self.up2(u3)\n        u2 = torch.cat([u2,d2],dim=1)\n        u2 = self.conv2(u2)\n\n        u1 = self.up1(u2)\n        u1 = torch.cat([u1,d1],dim=1)\n        u1 = self.conv1(u1)\n\n        return torch.sigmoid(self.out(u1))\n\ndef dice_loss(pred, target):\n\n    smooth = 1\n    \n    pred = pred.view(-1)\n    target = target.view(-1)\n\n    intersection = (pred * target).sum()\n\n    dice = (2.*intersection + smooth) / \\\n           (pred.sum() + target.sum() + smooth)\n\n    return 1 - dice\ndef train_model(model, loader):\n\n    optimizer = torch.optim.Adam(model.parameters(), lr=CFG.LR)\n\n    model.train()\n\n    for epoch in range(CFG.EPOCHS):\n\n        total_loss = 0\n\n        for x,y in tqdm(loader):\n\n            x = x.to(CFG.DEVICE)\n            y = y.to(CFG.DEVICE)\n\n            pred = model(x)\n\n            loss = dice_loss(pred,y)\n\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item()\n\n        print(\"Epoch\",epoch,\"Loss\",total_loss/len(loader))\n\nvolume = load_volume(\"fragment_1\")\nlabel, mask = load_labels(\"fragment_1\")\n\ndataset = InkDataset(volume,label,mask)\n\nloader = DataLoader(\n    dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=True\n)\n\nmodel = UNet().to(CFG.DEVICE)\n\ntrain_model(model,loader)\ndef predict_patch(model, patch):\n\n    model.eval()\n\n    patch = patch.astype(np.float32)/65535\n    patch = np.transpose(patch,(2,0,1))\n\n    x = torch.tensor(patch).unsqueeze(0).to(CFG.DEVICE)\n\n    with torch.no_grad():\n        pred = model(x)\n\n    return pred.squeeze().cpu().numpy()\n\nfinal = (pred_unet + pred_attention + pred_autoencoder) / 3\ndef postprocess(mask):\n\n    kernel = np.ones((3,3),np.uint8)\n\n    mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel)\n\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask)\n\n    cleaned = np.zeros_like(mask)\n\n    for i in range(1,num_labels):\n\n        if stats[i,cv2.CC_STAT_AREA] > 50:\n            cleaned[labels==i] = 1\n\n    return cleaned\n\ndef dice_score(pred, target):\n\n    pred = pred.flatten()\n    target = target.flatten()\n\n    intersection = (pred*target).sum()\n\n    return (2*intersection) / (pred.sum()+target.sum()+1e-6)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume/00.tif\n!pip install segmentation-models-pytorch==0.2.0\nimport os\nimport cv2\nimport numpy as np\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nimport albumentations as A\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport torch.nn.functional as F\nfrom scipy.ndimage import distance_transform_edt\nimport gc\nfrom collections import OrderedDict\nimport math\n#/kaggle/input/vesuvius-challenge-ink-detection/test\n# ============================================\n# CONFIGURATION\n# ============================================\nfragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\ntrain_paths = [\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n    '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n]\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {DEVICE}\")\n\n# Memory-efficient settings\nPATCH_SIZE = 128  # Reduced from 224 to avoid OOM\nSTRIDE = 64\nBATCH_SIZE = 4  # Reduced batch size\nACCUMULATION_STEPS = 2  # Gradient accumulation to simulate larger batch\nEPOCHS = 10\nLR = 3e-4\nWEIGHT_DECAY = 1e-4\nSLICE_START = 26\nSLICE_END = 38  # Using 12 slices as before\nOUTPUT_DIR = \"/kaggle/working/\"\nVALIDATION_SPLIT = 0.15\n\n# Advanced training settings\nUSE_AMP = True  # Mixed precision training\nPIN_MEMORY = True\nNUM_WORKERS = 2\nGRAD_CLIP = 1.0\nSCHEDULER_PATIENCE = 3\nEARLY_STOPPING_PATIENCE = 10\n\n# ============================================\n# IGNORE MASK GENERATION\n# ============================================\ndef generate_ignore_mask(ink_mask, distance_threshold=3, erosion_size=2):\n    \"\"\"\n    Generate ignore mask for uncertain regions:\n    - Edges of ink (distance transform)\n    - Small isolated dots\n    - Border regions\n    \"\"\"\n    ignore_mask = np.zeros_like(ink_mask, dtype=np.uint8)\n    \n    # Distance transform from ink boundaries\n    dist_to_ink = distance_transform_edt(1 - ink_mask)\n    dist_from_ink = distance_transform_edt(ink_mask)\n    \n    # Mark uncertain regions near boundaries\n    uncertain_boundary = (dist_to_ink <= distance_threshold) & (dist_from_ink <= distance_threshold)\n    ignore_mask[uncertain_boundary] = 1\n    \n    # Erode ink mask to remove thin edges\n    kernel = np.ones((erosion_size, erosion_size), np.uint8)\n    eroded_ink = cv2.erode(ink_mask.astype(np.uint8), kernel, iterations=1)\n    \n    # Areas removed by erosion become uncertain\n    ignore_mask[(ink_mask == 1) & (eroded_ink == 0)] = 1\n    \n    # Remove small isolated ink dots (noise)\n    contours, _ = cv2.findContours(ink_mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    for contour in contours:\n        if cv2.contourArea(contour) < 10:  # Small dots\n            x, y, w, h = cv2.boundingRect(contour)\n            ignore_mask[y:y+h, x:x+w] = 1\n    \n    return ignore_mask\n\n# ============================================\n# MINIUNETR ARCHITECTURE (Memory Efficient)\n# ============================================\nclass PatchEmbed(nn.Module):\n    \"\"\"Image to Patch Embedding\"\"\"\n    def __init__(self, img_size=128, patch_size=16, in_chans=12, embed_dim=384):\n        super().__init__()\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.n_patches = (img_size // patch_size) ** 2\n        \n        self.proj = nn.Conv3d(in_chans, embed_dim, \n                              kernel_size=(patch_size, patch_size, 1),  # 2D patches through depth\n                              stride=(patch_size, patch_size, 1))\n        \n    def forward(self, x):\n        x = self.proj(x)  # [B, embed_dim, H/patch, W/patch, D]\n        x = x.flatten(2).transpose(1, 2)  # [B, n_patches, embed_dim]\n        return x\n\nclass TransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=8, mlp_ratio=4., dropout=0.1):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(dim)\n        self.attn = nn.MultiheadAttention(dim, num_heads, dropout=dropout, batch_first=True)\n        self.norm2 = nn.LayerNorm(dim)\n        self.mlp = nn.Sequential(\n            nn.Linear(dim, int(dim * mlp_ratio)),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(int(dim * mlp_ratio), dim),\n            nn.Dropout(dropout)\n        )\n        \n    def forward(self, x):\n        x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]\n        x = x + self.mlp(self.norm2(x))\n        return x\n\nclass MiniUNETR(nn.Module):\n    \"\"\"\n    Memory-efficient MiniUNETR for Vesuvius Challenge\n    \"\"\"\n    def __init__(self, in_channels=12, out_channels=1, img_size=128, \n                 patch_size=16, embed_dim=192, depth=6, num_heads=6):\n        super().__init__()\n        \n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.embed_dim = embed_dim\n        \n        # Patch embedding\n        self.patch_embed = PatchEmbed(img_size, patch_size, in_channels, embed_dim)\n        n_patches = self.patch_embed.n_patches\n        \n        # Positional embedding\n        self.pos_embed = nn.Parameter(torch.zeros(1, n_patches, embed_dim))\n        self.pos_drop = nn.Dropout(0.1)\n        \n        # Transformer encoder\n        self.blocks = nn.ModuleList([\n            TransformerBlock(embed_dim, num_heads, dropout=0.1)\n            for _ in range(depth)\n        ])\n        self.norm = nn.LayerNorm(embed_dim)\n        \n        # Decoder (lightweight)\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(embed_dim, 128, kernel_size=4, stride=4),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(128, 64, kernel_size=3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 32, kernel_size=3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, out_channels, kernel_size=1)\n        )\n        \n        self._init_weights()\n        \n    def _init_weights(self):\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n        self.apply(self._init_module_weights)\n        \n    def _init_module_weights(self, m):\n        if isinstance(m, nn.Linear):\n            nn.init.trunc_normal_(m.weight, std=0.02)\n            if m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n            \n    def forward(self, x):\n        B, C, D, H, W = x.shape\n        \n        # Ensure input is in correct format (B, C, D, H, W)\n        if D != 1:\n            # Take depth average for transformer\n            x_depth = x.mean(dim=2, keepdim=True)  # [B, C, 1, H, W]\n        else:\n            x_depth = x\n            \n        # Patch embedding\n        x_patch = self.patch_embed(x_depth)  # [B, n_patches, embed_dim]\n        \n        # Add positional embedding\n        x_patch = self.pos_drop(x_patch + self.pos_embed)\n        \n        # Transformer blocks\n        for block in self.blocks:\n            x_patch = block(x_patch)\n        x_patch = self.norm(x_patch)\n        \n        # Reshape for decoder\n        grid_size = self.img_size // self.patch_size\n        x_patch = x_patch.transpose(1, 2).view(B, self.embed_dim, grid_size, grid_size)\n        \n        # Decode to segmentation map\n        out = self.decoder(x_patch)\n        \n        # Resize to original size\n        out = F.interpolate(out, size=(H, W), mode='bilinear', align_corners=False)\n        \n        return out\n\n# ============================================\n# DATA LOADING WITH IGNORE MASKS\n# ============================================\ndef load_volume(fragment_path):\n    volume_dir = os.path.join(fragment_path, \"surface_volume\")\n    slices = []\n    \n    for i in range(SLICE_START, SLICE_END):\n        img_path = os.path.join(volume_dir, f\"{i:02}.tif\")\n        if os.path.exists(img_path):\n            img = tifffile.imread(img_path)\n            slices.append(img)\n    \n    if not slices:\n        raise ValueError(f\"No slices found in {fragment_path}\")\n    \n    volume = np.stack(slices, axis=-1)\n    volume = volume.astype(np.float32)\n    \n    # Z-score normalization per slice\n    for i in range(volume.shape[-1]):\n        slice_data = volume[:, :, i]\n        mean, std = slice_data.mean(), slice_data.std()\n        volume[:, :, i] = (slice_data - mean) / (std + 1e-6)\n    \n    return volume\n\ndef extract_patches_with_ignore(volume, mask, ignore_mask, min_ink_ratio=0.01):\n    \"\"\"\n    Extract patches with both mask and ignore mask\n    \"\"\"\n    patches = []\n    mask_patches = []\n    ignore_patches = []\n    coords = []\n    \n    H, W, _ = volume.shape\n    \n    for y in range(0, H - PATCH_SIZE, STRIDE):\n        for x in range(0, W - PATCH_SIZE, STRIDE):\n            v_patch = volume[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            m_patch = mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            i_patch = ignore_mask[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n            \n            # Only keep patches with some ink signal (but respect ignore mask)\n            ink_ratio = (m_patch.sum() > 0).astype(float)\n            if ink_ratio > min_ink_ratio or np.random.random() < 0.1:  # Keep some background patches\n                patches.append(v_patch)\n                mask_patches.append(m_patch)\n                ignore_patches.append(i_patch)\n                coords.append((y, x))\n    \n    return patches, mask_patches, ignore_patches, coords\n\nclass VesuviusDatasetWithIgnore(Dataset):\n    def __init__(self, volumes, masks, ignore_masks, coords=None, transform=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.ignore_masks = ignore_masks\n        self.coords = coords\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        image = self.volumes[idx].copy()\n        mask = self.masks[idx].copy()\n        ignore = self.ignore_masks[idx].copy()\n        \n        if self.transform:\n            transformed = self.transform(image=image, mask=mask, ignore_mask=ignore)\n            image = transformed['image']\n            mask = transformed['mask']\n            ignore = transformed['ignore_mask']\n        \n        # Convert to tensors\n        image = torch.tensor(image).permute(2, 0, 1).float()  # [D, H, W]\n        mask = torch.tensor(mask).float()  # [H, W]\n        ignore = torch.tensor(ignore).float()  # [H, W]\n        \n        return image, mask, ignore\n\n# ============================================\n# CUSTOM LOSS WITH IGNORE MASK\n# ============================================\nclass DiceBCELossWithIgnore(nn.Module):\n    def __init__(self, dice_weight=0.5, bce_weight=0.5, smooth=1e-6):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.smooth = smooth\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n        \n    def forward(self, pred, target, ignore_mask):\n        \"\"\"\n        pred: [B, 1, H, W] - logits\n        target: [B, H, W] - binary mask\n        ignore_mask: [B, H, W] - 1 for ignore, 0 for keep\n        \"\"\"\n        # Ensure pred has correct shape\n        if pred.dim() == 4 and pred.size(1) == 1:\n            pred = pred.squeeze(1)  # [B, H, W]\n        \n        # Create valid mask (inverse of ignore mask)\n        valid_mask = (1 - ignore_mask).float()\n        \n        # Apply valid mask to target (set ignored regions to 0 for dice)\n        target_masked = target * valid_mask\n        \n        # BCE loss with ignore mask\n        bce_loss = self.bce(pred, target)\n        bce_loss = (bce_loss * valid_mask).sum() / (valid_mask.sum() + self.smooth)\n        \n        # Dice loss with ignore mask\n        pred_probs = torch.sigmoid(pred)\n        \n        # Flatten\n        pred_flat = pred_probs.view(-1)\n        target_flat = target_masked.view(-1)\n        valid_flat = valid_mask.view(-1)\n        \n        # Apply valid mask\n        pred_valid = pred_flat * valid_flat\n        target_valid = target_flat * valid_flat\n        \n        intersection = (pred_valid * target_valid).sum()\n        dice_score = (2. * intersection + self.smooth) / (pred_valid.sum() + target_valid.sum() + self.smooth)\n        dice_loss = 1 - dice_score\n        \n        # Combined loss\n        total_loss = self.dice_weight * dice_loss + self.bce_weight * bce_loss\n        \n        return total_loss\n\n# ============================================\n# DISTANCE-BASED VALIDATION SPLIT\n# ============================================\ndef create_distance_based_split(coords, val_ratio=0.15, min_distance=512):\n    \"\"\"\n    Create validation split ensuring patches are far apart\n    \"\"\"\n    coords = np.array(coords)\n    n_samples = len(coords)\n    \n    if n_samples == 0:\n        return [], []\n    \n    # Randomly select initial validation point\n    val_indices = []\n    train_indices = list(range(n_samples))\n    \n    # Select validation points ensuring minimum distance\n    while len(val_indices) < int(n_samples * val_ratio) and train_indices:\n        # Randomly select from remaining training indices\n        candidate_idx = np.random.choice(train_indices)\n        candidate_coord = coords[candidate_idx]\n        \n        # Check distance to existing validation points\n        if val_indices:\n            val_coords = coords[val_indices]\n            distances = np.sqrt(((val_coords - candidate_coord) ** 2).sum(axis=1))\n            if np.min(distances) >= min_distance:\n                val_indices.append(candidate_idx)\n                train_indices.remove(candidate_idx)\n        else:\n            val_indices.append(candidate_idx)\n            train_indices.remove(candidate_idx)\n    \n    # If we couldn't select enough validation points, fall back to random\n    if len(val_indices) < int(n_samples * val_ratio):\n        print(f\"Warning: Could not create distance-based split. Falling back to random.\")\n        indices = np.random.permutation(n_samples)\n        split = int(n_samples * val_ratio)\n        train_indices = indices[split:].tolist()\n        val_indices = indices[:split].tolist()\n    \n    return train_indices, val_indices\n\n# ============================================\n# CUSTOM COLLATE FUNCTION FOR MEMORY EFFICIENCY\n# ============================================\ndef collate_fn(batch):\n    \"\"\"Custom collate function to handle variable sized tensors\"\"\"\n    images = torch.stack([item[0] for item in batch])\n    masks = torch.stack([item[1] for item in batch])\n    ignores = torch.stack([item[2] for item in batch])\n    return images, masks, ignores\n\n# ============================================\n# MAIN PIPELINE\n# ============================================\ndef main():\n    print(\"=\"*60)\n    print(\"VESUVIUS INK DETECTION - ADVANCED PIPELINE\")\n    print(\"=\"*60)\n    \n    # Clear cache\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # ============================================\n    # LOAD AND PREPARE DATA\n    # ============================================\n    print(\"\\n1. Loading training data...\")\n    all_patches = []\n    all_masks = []\n    all_ignores = []\n    all_coords = []\n    \n    for path in train_paths:\n        print(f\"   Processing {os.path.basename(path)}...\")\n        \n        # Load volume\n        volume = load_volume(path)\n        \n        # Load ink mask\n        mask_path = os.path.join(path, \"inklabels.png\")\n        if not os.path.exists(mask_path):\n            print(f\"   Warning: Mask not found at {mask_path}\")\n            continue\n            \n        mask = cv2.imread(mask_path, 0)\n        mask = (mask > 0).astype(np.uint8)\n        \n        # Generate ignore mask\n        ignore_mask = generate_ignore_mask(mask)\n        \n        # Extract patches\n        patches, mask_patches, ignore_patches, coords = extract_patches_with_ignore(\n            volume, mask, ignore_mask\n        )\n        \n        all_patches.extend(patches)\n        all_masks.extend(mask_patches)\n        all_ignores.extend(ignore_patches)\n        all_coords.extend(coords)\n        \n        print(f\"   Extracted {len(patches)} patches\")\n    \n    print(f\"\\nTotal patches: {len(all_patches)}\")\n    \n    # ============================================\n    # CREATE TRAIN/VAL SPLIT\n    # ============================================\n    print(\"\\n2. Creating train/validation split...\")\n    \n    # Use distance-based split for better validation\n    train_indices, val_indices = create_distance_based_split(all_coords)\n    \n    print(f\"   Train samples: {len(train_indices)}\")\n    print(f\"   Validation samples: {len(val_indices)}\")\n    \n    # Split data\n    train_patches = [all_patches[i] for i in train_indices]\n    train_masks = [all_masks[i] for i in train_indices]\n    train_ignores = [all_ignores[i] for i in train_indices]\n    \n    val_patches = [all_patches[i] for i in val_indices]\n    val_masks = [all_masks[i] for i in val_indices]\n    val_ignores = [all_ignores[i] for i in val_indices]\n    \n    # ============================================\n    # DATA AUGMENTATION\n    # ============================================\n    print(\"\\n3. Setting up data augmentation...\")\n    \n    # Training augmentations (3D-aware)\n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.5),\n        A.GaussNoise(var_limit=(0, 0.01), p=0.3),\n        A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=0.2),\n        A.CoarseDropout(max_holes=8, max_height=16, max_width=16, fill_value=0, p=0.3),\n    ], additional_targets={'ignore_mask': 'mask'})\n    \n    val_transform = None\n    \n    # Create datasets\n    train_dataset = VesuviusDatasetWithIgnore(\n        train_patches, train_masks, train_ignores, \n        [all_coords[i] for i in train_indices], \n        transform=train_transform\n    )\n    \n    val_dataset = VesuviusDatasetWithIgnore(\n        val_patches, val_masks, val_ignores,\n        [all_coords[i] for i in val_indices],\n        transform=val_transform\n    )\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=True, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn,\n        drop_last=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    # ============================================\n    # MODEL INITIALIZATION\n    # ============================================\n    print(\"\\n4. Initializing MiniUNETR model...\")\n    \n    model = MiniUNETR(\n        in_channels=12,\n        out_channels=1,\n        img_size=PATCH_SIZE,\n        patch_size=16,\n        embed_dim=192,\n        depth=4,  # Reduced depth for memory efficiency\n        num_heads=6\n    ).to(DEVICE)\n    \n    # Count parameters\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"   Total parameters: {total_params:,}\")\n    print(f\"   Trainable parameters: {trainable_params:,}\")\n    \n    # ============================================\n    # LOSS, OPTIMIZER, SCHEDULER\n    # ============================================\n    criterion = DiceBCELossWithIgnore(dice_weight=0.5, bce_weight=0.5)\n    \n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    \n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='max', factor=0.5, patience=SCHEDULER_PATIENCE, verbose=True\n    )\n    \n    # Mixed precision scaler\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n    \n    # ============================================\n    # TRAINING LOOP\n    # ============================================\n    print(\"\\n5. Starting training...\")\n    print(\"=\"*60)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    train_losses = []\n    val_dice_scores = []\n    \n    for epoch in range(EPOCHS):\n        # Training phase\n        model.train()\n        train_loss = 0\n        train_steps = 0\n        optimizer.zero_grad()\n        \n        progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Train]')\n        for batch_idx, (images, masks, ignores) in enumerate(progress_bar):\n            images = images.to(DEVICE, non_blocking=True)\n            masks = masks.to(DEVICE, non_blocking=True)\n            ignores = ignores.to(DEVICE, non_blocking=True)\n            \n            # Add channel dimension to masks if needed\n            if masks.dim() == 3:\n                masks = masks.unsqueeze(1)\n            \n            # Forward pass with mixed precision\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n                loss = criterion(outputs, masks, ignores)\n                loss = loss / ACCUMULATION_STEPS  # Normalize loss\n            \n            # Backward pass\n            scaler.scale(loss).backward()\n            \n            # Gradient accumulation\n            if (batch_idx + 1) % ACCUMULATION_STEPS == 0:\n                # Gradient clipping\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                \n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            train_loss += loss.item() * ACCUMULATION_STEPS\n            train_steps += 1\n            \n            # Update progress bar\n            progress_bar.set_postfix({'loss': loss.item() * ACCUMULATION_STEPS})\n            \n            # Clear cache periodically\n            if batch_idx % 50 == 49:\n                torch.cuda.empty_cache()\n        \n        avg_train_loss = train_loss / train_steps\n        train_losses.append(avg_train_loss)\n        \n        # Validation phase\n        model.eval()\n        val_dice = 0\n        val_steps = 0\n        \n        with torch.no_grad():\n            for images, masks, ignores in tqdm(val_loader, desc=f'Epoch {epoch+1}/{EPOCHS} [Val]'):\n                images = images.to(DEVICE, non_blocking=True)\n                masks = masks.to(DEVICE, non_blocking=True)\n                ignores = ignores.to(DEVICE, non_blocking=True)\n                \n                if masks.dim() == 3:\n                    masks = masks.unsqueeze(1)\n                \n                with torch.cuda.amp.autocast(enabled=USE_AMP):\n                    outputs = model(images)\n                \n                # Calculate dice score (ignoring ignored regions)\n                preds = torch.sigmoid(outputs) > 0.5\n                \n                # Apply ignore mask\n                valid_mask = (1 - ignores).unsqueeze(1)\n                preds_valid = preds * valid_mask\n                masks_valid = masks * valid_mask\n                \n                intersection = (preds_valid * masks_valid).sum()\n                dice = (2 * intersection) / (preds_valid.sum() + masks_valid.sum() + 1e-6)\n                val_dice += dice.item()\n                val_steps += 1\n        \n        avg_val_dice = val_dice / val_steps\n        val_dice_scores.append(avg_val_dice)\n        \n        # Update learning rate\n        scheduler.step(avg_val_dice)\n        \n        # Print epoch results\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n        print(f\"  Train Loss: {avg_train_loss:.4f}\")\n        print(f\"  Val Dice: {avg_val_dice:.4f}\")\n        print(f\"  LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        # Save best model\n        if avg_val_dice > best_val_dice:\n            best_val_dice = avg_val_dice\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_dice': avg_val_dice,\n                'train_losses': train_losses,\n                'val_dice_scores': val_dice_scores\n            }, os.path.join(OUTPUT_DIR, 'best_model.pth'))\n            patience_counter = 0\n            print(f\"  ✓ New best model saved! Dice: {best_val_dice:.4f}\")\n        else:\n            patience_counter += 1\n            \n        # Early stopping\n        if patience_counter >= EARLY_STOPPING_PATIENCE:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n        \n        print(\"-\"*60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING COMPLETE\")\n    print(f\"Best Validation Dice: {best_val_dice:.4f}\")\n    print(\"=\"*60)\n    \n    # ============================================\n    # FINAL TEST ON FRAGMENT 1\n    # ============================================\n    print(\"\\n6. Testing on Fragment 1...\")\n    print(\"=\"*60)\n    \n    # Load best model\n    checkpoint = torch.load(os.path.join(OUTPUT_DIR, 'best_model.pth'))\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    # Load test data\n    test_volume = load_volume(fragment1_path)\n    test_mask = cv2.imread(os.path.join(fragment1_path, \"inklabels.png\"), 0)\n    test_mask = (test_mask > 0).astype(np.uint8)\n    test_ignore = generate_ignore_mask(test_mask)\n    \n    # Extract test patches\n    test_patches, test_masks, test_ignores, test_coords = extract_patches_with_ignore(\n        test_volume, test_mask, test_ignore\n    )\n    print(f\"Test patches: {len(test_patches)}\")\n    \n    # Create test dataset\n    test_dataset = VesuviusDatasetWithIgnore(\n        test_patches, test_masks, test_ignores, test_coords, transform=None\n    )\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=BATCH_SIZE, \n        shuffle=False, \n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        collate_fn=collate_fn\n    )\n    \n    # Evaluate\n    test_dice = 0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, ignores in tqdm(test_loader, desc=\"Testing\"):\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            ignores = ignores.to(DEVICE)\n            \n            if masks.dim() == 3:\n                masks = masks.unsqueeze(1)\n            \n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                outputs = model(images)\n            \n            preds = torch.sigmoid(outputs) > 0.5\n            \n            # Apply ignore mask\n            valid_mask = (1 - ignores).unsqueeze(1)\n            preds_valid = preds * valid_mask\n            masks_valid = masks * valid_mask\n            \n            intersection = (preds_valid * masks_valid).sum()\n            dice = (2 * intersection) / (preds_valid.sum() + masks_valid.sum() + 1e-6)\n            test_dice += dice.item()\n            \n            # Store for confusion matrix\n            all_preds.append(preds_valid.cpu().numpy())\n            all_targets.append(masks_valid.cpu().numpy())\n    \n    avg_test_dice = test_dice / len(test_loader)\n    \n    # Calculate confusion matrix\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_targets = np.concatenate([t.flatten() for t in all_targets])\n    \n    # Filter out ignored regions (where mask is 0 in all_targets due to valid_mask)\n    valid_pixels = all_targets != -1  # We need to adjust this based on how we stored\n    # For simplicity, we'll use the raw predictions and targets\n    \n    tn, fp, fn, tp = confusion_matrix(all_targets, all_preds, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n    \n    # ============================================\n    # FINAL RESULTS\n    # ============================================\n    print(\"\\n\" + \"=\"*60)\n    print(\"FINAL TEST RESULTS ON FRAGMENT 1\")\n    print(\"=\"*60)\n    print(f\"Test Dice Score: {avg_test_dice:.4f}\")\n    print(f\"Precision: {precision:.4f}\")\n    print(f\"Recall: {recall:.4f}\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(\"-\"*60)\n    print(\"Confusion Matrix:\")\n    print(f\"  True Positives: {tp}\")\n    print(f\"  True Negatives: {tn}\")\n    print(f\"  False Positives: {fp}\")\n    print(f\"  False Negatives: {fn}\")\n    print(\"=\"*60)\n    \n    # Save results\n    with open(os.path.join(OUTPUT_DIR, \"final_results.txt\"), \"w\") as f:\n        f.write(\"VESUVIUS INK DETECTION - FINAL RESULTS\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n        f.write(f\"Best Validation Dice: {best_val_dice:.4f}\\n\")\n        f.write(f\"Test Dice Score: {avg_test_dice:.4f}\\n\")\n        f.write(f\"Precision: {precision:.4f}\\n\")\n        f.write(f\"Recall: {recall:.4f}\\n\")\n        f.write(f\"F1 Score: {f1:.4f}\\n\")\n        f.write(\"-\"*60 + \"\\n\")\n        f.write(f\"True Positives: {tp}\\n\")\n        f.write(f\"True Negatives: {tn}\\n\")\n        f.write(f\"False Positives: {fp}\\n\")\n        f.write(f\"False Negatives: {fn}\\n\")\n        f.write(\"=\"*60 + \"\\n\")\n    \n    print(f\"\\nResults saved to: {os.path.join(OUTPUT_DIR, 'final_results.txt')}\")\n    \n    # Plot training curves\n    plt.figure(figsize=(12, 4))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses)\n    plt.title('Training Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(val_dice_scores)\n    plt.title('Validation Dice Score')\n    plt.xlabel('Epoch')\n    plt.ylabel('Dice')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'))\n    plt.show()\n    \n    return avg_test_dice\n\nif __name__ == \"__main__\":\n    test_dice = main()\n    print(f\"\\nFinal Test Dice Score: {test_dice:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Install libraries\n# =========================================================\n#!pip install segmentation-models-pytorch --quiet\n!pip install segmentation-models-pytorch==0.2.0\n\n# =========================================================\n# Imports\n# =========================================================\nimport os\nimport cv2\nimport torch\nimport numpy as np\nimport tifffile as tiff\nimport torch.nn as nn\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nimport segmentation_models_pytorch as smp\nfrom sklearn.metrics import precision_score, recall_score, fbeta_score\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# =========================================================\n# Dataset configuration\n# =========================================================\nBASE_PATH = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\n#/kaggle/input/vesuvius-challenge-ink-detection/test\nTRAIN_FRAGMENTS = [\"2\",\"3\"]\nTEST_FRAGMENT = [\"1\"]\n\nSLICE_START = 12\nSLICE_END = 30\n\nPATCH = 256\nSTRIDE = 128\n\n# =========================================================\n# Build patch index\n# =========================================================\ndef build_index(fragment_ids):\n\n    index = []\n\n    for fid in fragment_ids:\n\n        mask = cv2.imread(f\"{BASE_PATH}/{fid}/mask.png\",0)\n\n        H,W = mask.shape\n\n        for y in range(0,H-PATCH,STRIDE):\n            for x in range(0,W-PATCH,STRIDE):\n\n                mask_patch = mask[y:y+PATCH,x:x+PATCH]\n\n                if mask_patch.sum()==0:\n                    continue\n\n                index.append((fid,y,x))\n\n    return index\n\n\nprint(\"Building patch indices...\")\n\ntrain_index = build_index(TRAIN_FRAGMENTS)\ntest_index = build_index(TEST_FRAGMENT)\n\nprint(\"Train patches:\",len(train_index))\nprint(\"Test patches:\",len(test_index))\n\n# =========================================================\n# Dataset\n# =========================================================\nclass VesuviusDataset(Dataset):\n\n    def __init__(self,index):\n        self.index=index\n        self.cache={}\n\n    def load_volume(self,fid):\n\n        if fid in self.cache:\n            return self.cache[fid]\n\n        slices=[]\n\n        for i in range(SLICE_START,SLICE_END+1):\n\n            path=f\"{BASE_PATH}/{fid}/surface_volume/{i:02}.tif\"\n            img=tiff.imread(path)\n\n            slices.append(img)\n\n        volume=np.stack(slices)\n\n        self.cache[fid]=volume\n\n        return volume\n\n    def __len__(self):\n        return len(self.index)\n\n    def __getitem__(self,idx):\n\n        fid,y,x=self.index[idx]\n\n        volume=self.load_volume(fid)\n\n        ink=cv2.imread(f\"{BASE_PATH}/{fid}/inklabels.png\",0)\n\n        img_patch = volume[:,y:y+PATCH,x:x+PATCH]\n        label_patch = ink[y:y+PATCH,x:x+PATCH]\n\n        img_patch = img_patch.astype(np.float32)/65535.0\n        label_patch = (label_patch>0).astype(np.float32)\n\n        img=torch.tensor(img_patch).unsqueeze(0)\n        label=torch.tensor(label_patch).unsqueeze(0)\n\n        return img,label\n\n\ntrain_dataset=VesuviusDataset(train_index)\ntest_dataset=VesuviusDataset(test_index)\n\ntrain_loader=DataLoader(train_dataset,batch_size=2,shuffle=True,num_workers=2)\ntest_loader=DataLoader(test_dataset,batch_size=2,num_workers=2)\n\n# =========================================================\n# 3D UNet ResNet34\n# =========================================================\nclass UNet3D(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        self.unet = smp.Unet(\n            encoder_name=\"resnet34\",\n            encoder_weights=\"imagenet\",\n            in_channels=1,\n            classes=1\n        )\n\n    def forward(self,x):\n\n        B,C,D,H,W = x.shape\n\n        x = x.view(B*D,C,H,W)\n\n        x = self.unet(x)\n\n        x = x.view(B,1,D,H,W)\n\n        x = torch.mean(x,dim=2)\n\n        return x\n\n\nmodel = UNet3D().to(device)\n\n# =========================================================\n# Loss\n# =========================================================\nbce = nn.BCEWithLogitsLoss()\n\ndef dice_loss(pred,gt):\n\n    pred=torch.sigmoid(pred)\n\n    intersection=(pred*gt).sum()\n\n    dice=(2*intersection+1e-6)/(pred.sum()+gt.sum()+1e-6)\n\n    return 1-dice\n\ndef hybrid_loss(pred,gt):\n\n    return bce(pred,gt)+dice_loss(pred,gt)\n\noptimizer=torch.optim.Adam(model.parameters(),lr=1e-4)\n\n# =========================================================\n# Metrics\n# =========================================================\ndef dice_score(pred,gt):\n\n    intersection=(pred*gt).sum()\n\n    return (2*intersection+1e-6)/(pred.sum()+gt.sum()+1e-6)\n\ndef compute_metrics(preds,gts):\n\n    preds=preds.flatten()\n    gts=gts.flatten()\n\n    precision=precision_score(gts,preds,zero_division=0)\n    recall=recall_score(gts,preds,zero_division=0)\n    f05=fbeta_score(gts,preds,beta=0.5,zero_division=0)\n    dice=dice_score(preds,gts)\n\n    return precision,recall,f05,dice\n\n# =========================================================\n# Training\n# =========================================================\nEPOCHS=5\n\nfor epoch in range(EPOCHS):\n\n    model.train()\n\n    total_loss=0\n\n    for imgs,labels in tqdm(train_loader):\n\n        imgs=imgs.to(device)\n        labels=labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs=model(imgs)\n\n        loss=hybrid_loss(outputs,labels)\n\n        loss.backward()\n\n        optimizer.step()\n\n        total_loss+=loss.item()\n\n    print(\"Epoch\",epoch,\"Loss:\",total_loss/len(train_loader))\n\n# =========================================================\n# Testing\n# =========================================================\nmodel.eval()\n\nall_preds=[]\nall_gts=[]\n\nwith torch.no_grad():\n\n    for imgs,labels in tqdm(test_loader):\n\n        imgs=imgs.to(device)\n\n        outputs=torch.sigmoid(model(imgs))\n\n        preds=(outputs>0.5).cpu().numpy()\n        gts=labels.numpy()\n\n        all_preds.append(preds)\n        all_gts.append(gts)\n\nall_preds=np.concatenate(all_preds)\nall_gts=np.concatenate(all_gts)\n\nprecision,recall,f05,dice=compute_metrics(all_preds,all_gts)\n\nprint(\"Precision:\",precision)\nprint(\"Recall:\",recall)\nprint(\"F0.5:\",f05)\nprint(\"Dice:\",dice)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install segmentation-models-pytorch==0.2.0\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, random_split\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nfrom skimage.filters import threshold_otsu\nimport warnings\nimport random\nimport math\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport pandas as pd\nimport json\nfrom datetime import datetime\nimport pickle\nimport tifffile\nimport albumentations as A\n\nwarnings.filterwarnings('ignore')\n\n# Configuration\nclass Config:\n    # Data paths\n    fragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data - ONLY USED AT THE END\n    train_paths = [\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n    ]\n    \n    # Model parameters - USING 352 IMAGE SIZE FROM FIRST CODE\n    img_size = 352  # Changed from 224 to 352\n    slices = list(range(26, 38))  # Slice 26-37 (12 slices)\n    batch_size = 15  # Reduced due to larger image size\n    epochs = 5\n    lr = 3e-4  # Updated to 3e-4\n    \n    # 3D volumetric setup\n    num_input_slices = 5  # Updated to 5 slices\n    \n    # Memory optimization\n    max_train_samples = None  # Use all samples\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 5\n    \n    # Results saving\n    save_history = True\n    save_metrics = True\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Augmentation parameters\n    validation_split = 0.15  # Use 15% of training data for validation\n    patch_size = 352  # Using full image size from first code\n    stride = 176  # Half of patch size for overlap\n\nconfig = Config()\n\ndef visualize_preprocessing(base_path, slice_idx=30, save_path='preprocessing_comparison.png'):\n    \"\"\"Visualize raw vs preprocessed image for a specific slice\"\"\"\n    os.makedirs('preprocessing_visualizations', exist_ok=True)\n    \n    slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n    \n    if not os.path.exists(slice_path):\n        print(f\"Slice {slice_idx} not found at {slice_path}\")\n        return\n    \n    # Load raw image\n    raw_img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n    if raw_img is None:\n        print(f\"Could not load image from {slice_path}\")\n        return\n    \n    # Resize raw image for comparison\n    raw_resized = cv2.resize(raw_img, (config.img_size, config.img_size))\n    \n    # Apply preprocessing steps\n    # 1. Bilateral filter\n    img_filtered = cv2.bilateralFilter(raw_resized, 5, 50, 50)\n    \n    # 2. CLAHE\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    img_enhanced = clahe.apply(img_filtered)\n    \n    # 3. Normalized version (for display)\n    img_normalized = img_enhanced.astype(np.float32) / 255.0\n    \n    # Create comparison figure\n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    # Row 1: Original images\n    axes[0, 0].imshow(raw_resized, cmap='gray')\n    axes[0, 0].set_title(f'Raw Input\\n(Resized to {config.img_size}x{config.img_size})')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(img_filtered, cmap='gray')\n    axes[0, 1].set_title('After Bilateral Filter\\n(Noise reduction)')\n    axes[0, 1].axis('off')\n    \n    axes[0, 2].imshow(img_enhanced, cmap='gray')\n    axes[0, 2].set_title('After CLAHE\\n(Contrast enhancement)')\n    axes[0, 2].axis('off')\n    \n    # Row 2: Histograms and final normalized\n    axes[1, 0].hist(raw_resized.flatten(), bins=50, alpha=0.7, color='blue', label='Raw')\n    axes[1, 0].hist(img_enhanced.flatten(), bins=50, alpha=0.7, color='red', label='Processed')\n    axes[1, 0].set_title('Histogram Comparison')\n    axes[1, 0].set_xlabel('Pixel Intensity')\n    axes[1, 0].set_ylabel('Frequency')\n    axes[1, 0].legend()\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    axes[1, 1].imshow(img_normalized, cmap='gray')\n    axes[1, 1].set_title(f'Final Normalized\\n([0,1] range)')\n    axes[1, 1].axis('off')\n    \n    # Difference image\n    diff = np.abs(img_enhanced.astype(np.float32) - raw_resized.astype(np.float32))\n    diff_normalized = diff / diff.max() if diff.max() > 0 else diff\n    \n    im = axes[1, 2].imshow(diff_normalized, cmap='hot')\n    axes[1, 2].set_title('Enhancement Difference\\n(Brighter = More change)')\n    axes[1, 2].axis('off')\n    plt.colorbar(im, ax=axes[1, 2], fraction=0.046, pad=0.04)\n    \n    plt.suptitle(f'Preprocessing Pipeline Visualization - Slice {slice_idx}\\n'\n                 f'Bilateral Filter → CLAHE → Normalization', \n                 fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(os.path.join('preprocessing_visualizations', save_path), dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"Preprocessing visualization saved to preprocessing_visualizations/{save_path}\")\n    \n    # Print statistics\n    print(f\"\\nImage Statistics for Slice {slice_idx}:\")\n    print(f\"  Raw - Min: {raw_resized.min()}, Max: {raw_resized.max()}, Mean: {raw_resized.mean():.2f}, Std: {raw_resized.std():.2f}\")\n    print(f\"  Processed - Min: {img_enhanced.min()}, Max: {img_enhanced.max()}, Mean: {img_enhanced.mean():.2f}, Std: {img_enhanced.std():.2f}\")\n    print(f\"  Normalized - Min: {img_normalized.min():.3f}, Max: {img_normalized.max():.3f}, Mean: {img_normalized.mean():.3f}, Std: {img_normalized.std():.3f}\")\n\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data with enhanced preprocessing from first code\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                # RESIZE TO 352x352 (from first code)\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                \n                # APPLY ALL PREPROCESSING STEPS FROM FIRST CODE\n                # 1. Bilateral filter for noise reduction\n                img_filtered = cv2.bilateralFilter(img_resized, 5, 50, 50)\n                \n                # 2. CLAHE for contrast enhancement\n                clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n                img_enhanced = clahe.apply(img_filtered)\n                \n                # 3. Normalize to [0, 1]\n                img_normalized = img_enhanced.astype(np.float32) / 255.0\n                \n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=-1)  # Shape: (H, W, num_slices)\n        print(f\"Loaded volume from {base_path} with shape {volume.shape}\")\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            # Resize mask to match image size\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio for {base_path}: {ink_ratio:.4f} ({mask_binary.sum()} pixels)\")\n            return mask_binary\n    return None\n\ndef extract_patches_with_stride(volume, mask, patch_size=352, stride=176):\n    \"\"\"Extract patches from volume and mask with stride\"\"\"\n    patches = []\n    mask_patches = []\n    coordinates = []  # Store coordinates for reconstruction if needed\n    \n    H, W, _ = volume.shape\n    \n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            v_patch = volume[y:y+patch_size, x:x+patch_size, :]\n            m_patch = mask[y:y+patch_size, x:x+patch_size]\n            \n            patches.append(v_patch)\n            mask_patches.append(m_patch)\n            coordinates.append((y, x))\n    \n    print(f\"  Extracted {len(patches)} patches from {H}x{W} image with stride {stride}\")\n    return patches, mask_patches, coordinates\n\n# Dataset class with augmentations\nclass VesuviusDataset(Dataset):\n    def __init__(self, volumes, masks, transform=None, use_slices=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.transform = transform\n        self.use_slices = use_slices  # Number of slices to use (if None, use all)\n\n    def __len__(self):\n        return len(self.volumes)\n\n    def __getitem__(self, idx):\n        image = self.volumes[idx]  # Shape: (H, W, C)\n        mask = self.masks[idx]      # Shape: (H, W)\n        \n        # Select specific slices if configured\n        if self.use_slices is not None and self.use_slices < image.shape[-1]:\n            # Randomly select num_input_slices consecutive slices\n            total_slices = image.shape[-1]\n            start_idx = random.randint(0, total_slices - self.use_slices)\n            image = image[:, :, start_idx:start_idx + self.use_slices]\n\n        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image = augmented['image']\n            mask = augmented['mask']\n\n        # Convert to tensor: (C, H, W) for image, (1, H, W) for mask\n        image = torch.tensor(image).permute(2, 0, 1).float()\n        mask = torch.tensor(mask).unsqueeze(0).float()\n\n        return image, mask\n\n# BCE Loss only (from second code)\nclass BCELoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        \n    def forward(self, pred, target):\n        return self.bce(pred, target)\n\ndef calculate_all_metrics(pred, target, threshold=0.5):\n    \"\"\"Calculate all metrics including TP, TN, FP, FN\"\"\"\n    pred_binary = (pred > threshold).astype(np.float32)\n    pred_flat = pred_binary.flatten().astype(int)\n    target_flat = target.flatten().astype(int)\n    \n    tn, fp, fn, tp = confusion_matrix(target_flat, pred_flat, labels=[0, 1]).ravel()\n    \n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n    \n    beta_squared = 0.5 ** 2\n    f05 = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (beta_squared * precision + recall) > 0 else 0\n    \n    iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n    \n    return {\n        'precision': float(precision),\n        'recall': float(recall),\n        'dice': float(dice),\n        'f05': float(f05),\n        'iou': float(iou),\n        'tp': int(tp),\n        'tn': int(tn),\n        'fp': int(fp),\n        'fn': int(fn)\n    }\n\ndef dice_score(pred, target, smooth=1e-6):\n    pred = pred.flatten()\n    target = target.flatten()\n    intersection = (pred * target).sum()\n    return float((2. * intersection + smooth) / (pred.sum() + target.sum() + smooth))\n\ndef save_history(history, filename='training_history.json'):\n    \"\"\"Save training history to JSON file\"\"\"\n    os.makedirs('history', exist_ok=True)\n    \n    def convert_numpy_types(obj):\n        if isinstance(obj, np.integer):\n            return int(obj)\n        elif isinstance(obj, np.floating):\n            return float(obj)\n        elif isinstance(obj, np.ndarray):\n            return obj.tolist()\n        elif isinstance(obj, dict):\n            return {key: convert_numpy_types(value) for key, value in obj.items()}\n        elif isinstance(obj, list):\n            return [convert_numpy_types(item) for item in obj]\n        else:\n            return obj\n    \n    history_serializable = convert_numpy_types(history)\n    \n    with open(os.path.join('history', filename), 'w') as f:\n        json.dump(history_serializable, f, indent=2)\n    \n    print(f\"Training history saved to {filename}\")\n\ndef save_metrics_summary(train_metrics, val_metrics, test_metrics, filename='metrics_summary.csv'):\n    \"\"\"Save all metrics to CSV file\"\"\"\n    os.makedirs('metrics', exist_ok=True)\n    \n    summary_data = {\n        'Phase': ['Training', 'Validation', 'Test (Fragment 1)'],\n        'Dice_Score': [float(train_metrics['dice']), float(val_metrics['dice']), float(test_metrics['dice'])],\n        'F0.5_Score': [float(train_metrics['f05']), float(val_metrics['f05']), float(test_metrics['f05'])],\n        'IOU_Score': [float(train_metrics['iou']), float(val_metrics['iou']), float(test_metrics['iou'])],\n        'Precision': [float(train_metrics['precision']), float(val_metrics['precision']), float(test_metrics['precision'])],\n        'Recall': [float(train_metrics['recall']), float(val_metrics['recall']), float(test_metrics['recall'])],\n        'TP': [int(train_metrics['tp']), int(val_metrics['tp']), int(test_metrics['tp'])],\n        'TN': [int(train_metrics['tn']), int(val_metrics['tn']), int(test_metrics['tn'])],\n        'FP': [int(train_metrics['fp']), int(val_metrics['fp']), int(test_metrics['fp'])],\n        'FN': [int(train_metrics['fn']), int(val_metrics['fn']), int(test_metrics['fn'])]\n    }\n    \n    df = pd.DataFrame(summary_data)\n    df.to_csv(os.path.join('metrics', filename), index=False)\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"METRICS SUMMARY\")\n    print(\"=\"*80)\n    print(df.to_string())\n    print(f\"\\nMetrics summary saved to {filename}\")\n\ndef plot_training_history(history):\n    \"\"\"Plot training history\"\"\"\n    os.makedirs('training_plots', exist_ok=True)\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Loss plot\n    ax1.plot(epochs, history['train_loss'], 'b-', label='Training Loss', linewidth=2)\n    ax1.plot(epochs, history['val_loss'], 'r-', label='Validation Loss', linewidth=2)\n    ax1.set_title('Training and Validation Loss (BCE)')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Loss')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    # Dice score plot\n    ax2.plot(epochs, history['train_dice'], 'b-', label='Training Dice', linewidth=2)\n    ax2.plot(epochs, history['val_dice'], 'r-', label='Validation Dice', linewidth=2)\n    ax2.set_title('Dice Score')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Dice Score')\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    # F0.5 score plot\n    ax3.plot(epochs, history['train_f05'], 'b-', label='Training F0.5', linewidth=2)\n    ax3.plot(epochs, history['val_f05'], 'r-', label='Validation F0.5', linewidth=2)\n    ax3.set_title('F0.5 Score')\n    ax3.set_xlabel('Epoch')\n    ax3.set_ylabel('F0.5 Score')\n    ax3.legend()\n    ax3.grid(True, alpha=0.3)\n    \n    # IOU score plot\n    ax4.plot(epochs, history['train_iou'], 'b-', label='Training IOU', linewidth=2)\n    ax4.plot(epochs, history['val_iou'], 'r-', label='Validation IOU', linewidth=2)\n    ax4.set_title('IOU Score')\n    ax4.set_xlabel('Epoch')\n    ax4.set_ylabel('IOU Score')\n    ax4.legend()\n    ax4.grid(True, alpha=0.3)\n    \n    plt.suptitle(f'Training History - Image Size: {config.img_size}x{config.img_size}', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/training_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Precision and Recall plot\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n    \n    ax1.plot(epochs, history['train_precision'], 'g-', label='Training Precision', linewidth=2)\n    ax1.plot(epochs, history['val_precision'], 'orange', label='Validation Precision', linewidth=2)\n    ax1.set_title('Precision')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Precision')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    ax2.plot(epochs, history['train_recall'], 'g-', label='Training Recall', linewidth=2)\n    ax2.plot(epochs, history['val_recall'], 'orange', label='Validation Recall', linewidth=2)\n    ax2.set_title('Recall')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Recall')\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    plt.suptitle('Precision and Recall', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/precision_recall_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n\ndef save_test_visualization(test_volume, test_mask, test_predictions, test_metrics, fragment_name=\"fragment1\"):\n    \"\"\"Save comprehensive test visualization with metrics\"\"\"\n    os.makedirs('test_results', exist_ok=True)\n    \n    # Get middle slice for visualization\n    middle_slice_idx = test_volume.shape[2] // 2\n    middle_slice = test_volume[:, :, middle_slice_idx]\n    \n    # Create binary predictions\n    test_pred_binary = (test_predictions > test_metrics['threshold']).astype(np.float32)\n    \n    fig, axes = plt.subplots(2, 4, figsize=(20, 10))\n    \n    # Row 1\n    axes[0, 0].imshow(middle_slice, cmap='gray')\n    axes[0, 0].set_title('Input Image\\n(Preprocessed Middle Slice)')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(test_mask, cmap='gray')\n    axes[0, 1].set_title('Ground Truth')\n    axes[0, 1].axis('off')\n    \n    axes[0, 2].imshow(test_predictions, cmap='jet', vmin=0, vmax=1)\n    axes[0, 2].set_title(f'Test Predictions\\n(Probabilities)')\n    axes[0, 2].axis('off')\n    \n    axes[0, 3].imshow(test_pred_binary, cmap='jet')\n    axes[0, 3].set_title(f'Binary Predictions\\n(Threshold={test_metrics[\"threshold\"]:.2f})')\n    axes[0, 3].axis('off')\n    \n    # Row 2: Overlays\n    # Prediction overlay\n    pred_overlay = np.stack([\n        middle_slice * 0.5 + test_pred_binary * 0.5,\n        middle_slice,\n        middle_slice\n    ], axis=-1)\n    axes[1, 0].imshow(pred_overlay)\n    axes[1, 0].set_title('Input + Predictions\\n(Red = Predicted Ink)')\n    axes[1, 0].axis('off')\n    \n    # Error map (FP in red, FN in blue)\n    error_map = np.zeros((test_mask.shape[0], test_mask.shape[1], 3))\n    false_positives = (test_pred_binary == 1) & (test_mask == 0)\n    false_negatives = (test_pred_binary == 0) & (test_mask == 1)\n    \n    error_map[:, :, 0] = false_positives * 0.8  # Red = False Positives\n    error_map[:, :, 2] = false_negatives * 0.8  # Blue = False Negatives\n    \n    axes[1, 1].imshow(error_map)\n    axes[1, 1].set_title('Error Map\\n(Red=FP, Blue=FN)')\n    axes[1, 1].axis('off')\n    \n    # Confusion matrix visualization\n    confusion_img = np.zeros((100, 100, 3))\n    confusion_img[:50, :50, :] = 0.2  # TN area (gray)\n    confusion_img[:50, 50:, 0] = 0.8  # FP area (red)\n    confusion_img[50:, :50, 2] = 0.8  # FN area (blue)\n    confusion_img[50:, 50:, 1] = 0.8  # TP area (green)\n    \n    axes[1, 2].imshow(confusion_img)\n    axes[1, 2].set_title('Confusion Matrix\\n(Green=TP, Red=FP, Blue=FN, Gray=TN)')\n    axes[1, 2].axis('off')\n    \n    # Hide the last subplot\n    axes[1, 3].axis('off')\n    \n    # Add detailed metrics text\n    metrics_text = f\"TEST RESULTS - {fragment_name}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Dice Score: {test_metrics['dice']:.4f}\\n\"\n    metrics_text += f\"F0.5 Score: {test_metrics['f05']:.4f}\\n\"\n    metrics_text += f\"IOU Score: {test_metrics['iou']:.4f}\\n\"\n    metrics_text += f\"Precision: {test_metrics['precision']:.4f}\\n\"\n    metrics_text += f\"Recall: {test_metrics['recall']:.4f}\\n\"\n    metrics_text += f\"Optimal Threshold: {test_metrics['threshold']:.3f}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"True Positives: {test_metrics['tp']:,}\\n\"\n    metrics_text += f\"True Negatives: {test_metrics['tn']:,}\\n\"\n    metrics_text += f\"False Positives: {test_metrics['fp']:,}\\n\"\n    metrics_text += f\"False Negatives: {test_metrics['fn']:,}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Image Size: {config.img_size}x{config.img_size}\\n\"\n    metrics_text += f\"Input Slices: {config.num_input_slices}\\n\"\n    metrics_text += f\"Available Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} total)\"\n    \n    plt.figtext(0.02, 0.02, metrics_text, fontsize=8, fontfamily='monospace',\n                bbox=dict(boxstyle=\"round,pad=0.5\", facecolor=\"lightgray\", alpha=0.8))\n    \n    plt.suptitle(f'TEST RESULTS - {fragment_name} - Full Fragment Evaluation\\nImage Size: {config.img_size}x{config.img_size}', \n                 fontsize=12, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(f'test_results/{fragment_name}_test_results.png', \n                dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Save test predictions\n    np.save(f'test_results/{fragment_name}_test_predictions.npy', test_predictions)\n    np.save(f'test_results/{fragment_name}_test_predictions_binary.npy', test_pred_binary)\n    \n    print(f\"Test visualization saved to: test_results/{fragment_name}_test_results.png\")\n    print(f\"Test predictions saved to: test_results/{fragment_name}_test_predictions.npy\")\n\ndef main():\n    print(\"=\"*80)\n    print(\"3D UNet with EfficientNet-B0 - Training with Preprocessing (352x352)\")\n    print(\"=\"*80)\n    print(\"CONFIGURATION:\")\n    print(f\"  Image Size: {config.img_size}x{config.img_size} (from first code)\")\n    print(f\"  Preprocessing: Bilateral Filter + CLAHE + Normalization\")\n    print(f\"  Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\")\n    print(f\"  Input Slices: {config.num_input_slices} (random consecutive selection)\")\n    print(f\"  Batch Size: {config.batch_size}\")\n    print(f\"  Epochs: {config.epochs}\")\n    print(f\"  Learning Rate: {config.lr} (3e-4)\")\n    print(f\"  Loss Function: BCE Loss (from second code)\")\n    print(f\"  Model: UNet with EfficientNet-B0 encoder\")\n    print(f\"  Device: {config.device}\")\n    print(f\"  Validation Split: {config.validation_split*100:.0f}%\")\n    print(\"=\"*80)\n    \n    # Generate preprocessing visualization first\n    print(\"\\n\" + \"=\"*50)\n    print(\"GENERATING PREPROCESSING VISUALIZATION\")\n    print(\"=\"*50)\n    visualize_preprocessing(config.train_paths[0], slice_idx=30, save_path='fragment2_slice30_preprocessing.png')\n    \n    # Set random seeds\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        gpu_memory = torch.cuda.get_device_properties(0).total_memory / 1e9\n        print(f\"Available GPU memory: {gpu_memory:.1f} GB\")\n    \n    # ============================================\n    # LOAD TRAINING DATA (Fragments 2 & 3)\n    # ============================================\n    print(\"\\n\" + \"=\"*50)\n    print(\"LOADING TRAINING DATA FROM FRAGMENTS 2 & 3\")\n    print(\"=\"*50)\n    \n    all_train_volumes = []\n    all_train_masks = []\n    \n    for path in config.train_paths:\n        print(f\"\\nLoading from {path}...\")\n        \n        # Load volume with preprocessing from first code\n        volume = load_volume_data(path, config.slices)\n        if volume is None:\n            print(f\"Warning: Could not load volume from {path}\")\n            continue\n        \n        # Load mask\n        mask = load_mask_data(path)\n        if mask is None:\n            print(f\"Warning: Could not load mask from {path}\")\n            continue\n        \n        # Extract patches with stride\n        v_patches, m_patches, coords = extract_patches_with_stride(\n            volume, mask, \n            patch_size=config.patch_size, \n            stride=config.stride\n        )\n        \n        all_train_volumes.extend(v_patches)\n        all_train_masks.extend(m_patches)\n    \n    print(f\"\\nTotal training patches: {len(all_train_volumes)}\")\n    \n    if len(all_train_volumes) == 0:\n        raise ValueError(\"No training data loaded!\")\n    \n    # ============================================\n    # CREATE TRAIN/VALIDATION SPLIT\n    # ============================================\n    dataset_size = len(all_train_volumes)\n    val_size = max(1, int(dataset_size * config.validation_split))  # Ensure at least 1 validation sample\n    train_size = dataset_size - val_size\n    \n    print(f\"\\nSplitting data: Train samples: {train_size}, Validation samples: {val_size}\")\n    \n    # Define augmentations (from second code)\n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.5),\n        A.GridDistortion(p=0.3),\n        A.GaussianBlur(p=0.3),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),\n    ])\n    \n    val_transform = None  # No augmentation for validation\n    \n    # Create full dataset and split\n    full_dataset = VesuviusDataset(\n        all_train_volumes, \n        all_train_masks, \n        train_transform,\n        use_slices=config.num_input_slices\n    )\n    \n    # Use random_split with appropriate sizes\n    if val_size > 0:\n        train_dataset, val_dataset = random_split(\n            full_dataset, \n            [train_size, val_size],\n            generator=torch.Generator().manual_seed(42)\n        )\n        # Update transforms\n        train_dataset.dataset.transform = train_transform\n        val_dataset.dataset.transform = val_transform\n        \n        val_loader = DataLoader(\n            val_dataset, \n            batch_size=config.batch_size, \n            shuffle=False, \n            num_workers=2,\n            pin_memory=True if torch.cuda.is_available() else False\n        )\n        print(f\"  Validation loader created with {len(val_loader)} batches\")\n    else:\n        # If no validation samples, use all for training\n        train_dataset = full_dataset\n        val_loader = None\n        print(\"  Warning: No validation samples created. Using all data for training.\")\n    \n    # Create training loader\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=config.batch_size, \n        shuffle=True, \n        num_workers=2,\n        pin_memory=True if torch.cuda.is_available() else False\n    )\n    \n    print(f\"\\nData Loaders Created:\")\n    print(f\"  Training batches: {len(train_loader)}\")\n    if val_loader:\n        print(f\"  Validation batches: {len(val_loader)}\")\n    \n    # ============================================\n    # INITIALIZE MODEL\n    # ============================================\n    print(f\"\\nInitializing model...\")\n    \n    model = smp.Unet(\n        encoder_name=\"resnet34\",\n        encoder_weights=\"imagenet\",\n        in_channels=config.num_input_slices,  # Using 5 input slices\n        \n        classes=1\n    ).to(config.device)\n    \n    print(f\"Model initialized on {config.device}\")\n    print(f\"Input channels: {config.num_input_slices} (randomly selected from {len(config.slices)} total)\")\n    print(f\"Encoder: efficientnet-b0 with ImageNet weights\")\n    \n    # Initialize optimizer and loss (BCE only from second code)\n    optimizer = optim.Adam(model.parameters(), lr=config.lr, weight_decay=1e-4)\n    criterion = BCELoss()  # Using BCE loss only\n    \n    # Fix for ReduceLROnPlateau - removed 'verbose' parameter\n    scheduler = ReduceLROnPlateau(\n        optimizer, \n        mode='max', \n        factor=0.5, \n        patience=3\n        # verbose=True removed - this parameter doesn't exist in this version\n    )\n    \n    scaler = torch.cuda.amp.GradScaler()\n    \n    # Training history\n    history = {\n        'train_loss': [],\n        'train_dice': [],\n        'train_f05': [],\n        'train_iou': [],\n        'train_precision': [],\n        'train_recall': [],\n        'val_loss': [],\n        'val_dice': [],\n        'val_f05': [],\n        'val_iou': [],\n        'val_precision': [],\n        'val_recall': []\n    }\n    \n    best_val_dice = 0\n    best_model_state = None\n    patience = 5\n    patience_counter = 0\n    \n    # ============================================\n    # TRAINING LOOP\n    # ============================================\n    print(f\"\\nStarting training for {config.epochs} epochs...\")\n    print(\"=\"*80)\n    \n    for epoch in range(config.epochs):\n        # Training phase\n        model.train()\n        epoch_train_loss = 0\n        train_outputs = []\n        train_targets = []\n        \n        train_iterator = tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.epochs} - Training')\n        for images, masks in train_iterator:\n            images, masks = images.to(config.device), masks.to(config.device)\n            \n            optimizer.zero_grad()\n            \n            with torch.cuda.amp.autocast():\n                outputs = model(images)\n                loss = criterion(outputs, masks)\n            \n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            \n            epoch_train_loss += loss.item()\n            \n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_outputs.append(preds.cpu().numpy())\n                train_targets.append(masks.cpu().numpy())\n            \n            train_iterator.set_postfix(loss=loss.item())\n        \n        # Calculate training metrics\n        train_outputs = np.concatenate(train_outputs)\n        train_targets = np.concatenate(train_targets)\n        \n        best_train_threshold = 0.5\n        best_train_dice = 0\n        for thresh in np.arange(0.1, 0.9, 0.05):\n            preds = (train_outputs > thresh).astype(np.float32)\n            dice = dice_score(preds, train_targets)\n            if dice > best_train_dice:\n                best_train_dice = dice\n                best_train_threshold = thresh\n        \n        train_preds = (train_outputs > best_train_threshold).astype(np.float32)\n        train_metrics = calculate_all_metrics(train_preds, train_targets)\n        train_metrics['threshold'] = float(best_train_threshold)\n        \n        # Store training metrics\n        history['train_loss'].append(float(epoch_train_loss / len(train_loader)))\n        history['train_dice'].append(float(train_metrics['dice']))\n        history['train_f05'].append(float(train_metrics['f05']))\n        history['train_iou'].append(float(train_metrics['iou']))\n        history['train_precision'].append(float(train_metrics['precision']))\n        history['train_recall'].append(float(train_metrics['recall']))\n        \n        # Clear cache\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        \n        # Validation phase (if validation loader exists)\n        if val_loader:\n            model.eval()\n            val_loss = 0\n            val_outputs = []\n            val_targets = []\n            \n            with torch.no_grad():\n                val_iterator = tqdm(val_loader, desc=f'Epoch {epoch+1}/{config.epochs} - Validation')\n                for images, masks in val_iterator:\n                    images, masks = images.to(config.device), masks.to(config.device)\n                    \n                    outputs = model(images)\n                    loss = criterion(outputs, masks)\n                    val_loss += loss.item()\n                    \n                    preds = torch.sigmoid(outputs)\n                    val_outputs.append(preds.cpu().numpy())\n                    val_targets.append(masks.cpu().numpy())\n            \n            val_outputs = np.concatenate(val_outputs)\n            val_targets = np.concatenate(val_targets)\n            \n            best_val_threshold = 0.5\n            best_val_dice_current = 0\n            for thresh in np.arange(0.1, 0.9, 0.05):\n                preds = (val_outputs > thresh).astype(np.float32)\n                dice = dice_score(preds, val_targets)\n                if dice > best_val_dice_current:\n                    best_val_dice_current = dice\n                    best_val_threshold = thresh\n            \n            val_preds = (val_outputs > best_val_threshold).astype(np.float32)\n            val_metrics = calculate_all_metrics(val_preds, val_targets)\n            val_metrics['threshold'] = float(best_val_threshold)\n            \n            # Store validation metrics\n            history['val_loss'].append(float(val_loss / len(val_loader)))\n            history['val_dice'].append(float(val_metrics['dice']))\n            history['val_f05'].append(float(val_metrics['f05']))\n            history['val_iou'].append(float(val_metrics['iou']))\n            history['val_precision'].append(float(val_metrics['precision']))\n            history['val_recall'].append(float(val_metrics['recall']))\n            \n            # Update best model\n            if val_metrics['dice'] > best_val_dice:\n                best_val_dice = val_metrics['dice']\n                best_model_state = model.state_dict().copy()\n                patience_counter = 0\n                print(f\"  ✓ New best model! Validation Dice: {best_val_dice:.4f}\")\n            else:\n                patience_counter += 1\n            \n            # Print progress with validation\n            print(f\"\\nEpoch {epoch+1}/{config.epochs}:\")\n            print(f\"  Train - Loss: {history['train_loss'][-1]:.4f}, \"\n                  f\"Dice: {train_metrics['dice']:.4f}, F0.5: {train_metrics['f05']:.4f}\")\n            print(f\"  Val   - Loss: {history['val_loss'][-1]:.4f}, \"\n                  f\"Dice: {val_metrics['dice']:.4f}, F0.5: {val_metrics['f05']:.4f}, \"\n                  f\"Precision: {val_metrics['precision']:.4f}, Recall: {val_metrics['recall']:.4f}\")\n            \n            # Learning rate scheduling\n            scheduler.step(val_metrics['dice'])\n            \n            # Early stopping\n            if patience_counter >= patience:\n                print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n                break\n        else:\n            # Print progress without validation\n            print(f\"\\nEpoch {epoch+1}/{config.epochs}:\")\n            print(f\"  Train - Loss: {history['train_loss'][-1]:.4f}, \"\n                  f\"Dice: {train_metrics['dice']:.4f}, F0.5: {train_metrics['f05']:.4f}\")\n            # Update best model with training dice if no validation\n            if train_metrics['dice'] > best_val_dice:\n                best_val_dice = train_metrics['dice']\n                best_model_state = model.state_dict().copy()\n        \n        # Clear cache\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    # Load best model for testing\n    if best_model_state:\n        model.load_state_dict(best_model_state)\n        print(f\"\\nLoaded best model with best Dice: {best_val_dice:.4f}\")\n    \n    # ============================================\n    # FINAL TEST ON FRAGMENT 1 (ONLY AT THE END)\n    # ============================================\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL TESTING ON FRAGMENT 1 (COMPLETELY UNSEEN DURING TRAINING)\")\n    print(\"=\"*80)\n    \n    # Load Fragment 1 test data with preprocessing\n    print(\"\\nLoading Fragment 1 test data...\")\n    test_volume = load_volume_data(config.fragment1_path, config.slices)\n    test_mask = load_mask_data(config.fragment1_path)\n    \n    if test_volume is None or test_mask is None:\n        raise ValueError(f\"Failed to load Fragment 1 test data from {config.fragment1_path}\")\n    \n    print(f\"Test volume shape: {test_volume.shape}\")\n    print(f\"Test mask shape: {test_mask.shape}\")\n    \n    # Extract patches for evaluation\n    test_patches, test_mask_patches, test_coords = extract_patches_with_stride(\n        test_volume, test_mask,\n        patch_size=config.patch_size,\n        stride=config.stride\n    )\n    \n    print(f\"Extracted {len(test_patches)} test patches from Fragment 1\")\n    \n    # Create test dataset and loader\n    test_dataset = VesuviusDataset(\n        test_patches, \n        test_mask_patches, \n        transform=None,\n        use_slices=config.num_input_slices\n    )\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=config.batch_size, \n        shuffle=False, \n        num_workers=2\n    )\n    \n    # Evaluate on Fragment 1\n    model.eval()\n    test_outputs = []\n    test_targets = []\n    \n    with torch.no_grad():\n        test_iterator = tqdm(test_loader, desc=\"Testing on Fragment 1\")\n        for images, masks in test_iterator:\n            images = images.to(config.device)\n            \n            outputs = torch.sigmoid(model(images))\n            test_outputs.append(outputs.cpu().numpy())\n            test_targets.append(masks.cpu().numpy())\n    \n    test_outputs = np.concatenate(test_outputs)\n    test_targets = np.concatenate(test_targets)\n    \n    # Find best threshold for test\n    best_test_threshold = 0.5\n    best_test_dice = 0\n    for thresh in np.arange(0.1, 0.9, 0.05):\n        preds = (test_outputs > thresh).astype(np.float32)\n        dice = dice_score(preds, test_targets)\n        if dice > best_test_dice:\n            best_test_dice = dice\n            best_test_threshold = thresh\n    \n    # Calculate all test metrics\n    test_preds = (test_outputs > best_test_threshold).astype(np.float32)\n    test_metrics = calculate_all_metrics(test_preds, test_targets)\n    test_metrics['threshold'] = float(best_test_threshold)\n    \n    # Save test visualization\n    save_test_visualization(\n        test_volume, test_mask, \n        test_outputs.mean(axis=0)[0] if len(test_outputs.shape) > 2 else test_outputs[0], \n        test_metrics,\n        \"fragment1\"\n    )\n    \n    # Save training history\n    if config.save_history:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_history(history, filename=f'training_history_{timestamp}.json')\n        plot_training_history(history)\n    \n    # Get final training and validation metrics\n    final_train_metrics = {\n        'dice': float(history['train_dice'][-1]),\n        'f05': float(history['train_f05'][-1]),\n        'iou': float(history['train_iou'][-1]),\n        'precision': float(history['train_precision'][-1]),\n        'recall': float(history['train_recall'][-1]),\n        'threshold': float(best_train_threshold),\n        'tp': int(train_metrics['tp']),\n        'tn': int(train_metrics['tn']),\n        'fp': int(train_metrics['fp']),\n        'fn': int(train_metrics['fn'])\n    }\n    \n    if val_loader and len(history['val_dice']) > 0:\n        final_val_metrics = {\n            'dice': float(history['val_dice'][-1]),\n            'f05': float(history['val_f05'][-1]),\n            'iou': float(history['val_iou'][-1]),\n            'precision': float(history['val_precision'][-1]),\n            'recall': float(history['val_recall'][-1]),\n            'threshold': float(val_metrics['threshold']),\n            'tp': int(val_metrics['tp']),\n            'tn': int(val_metrics['tn']),\n            'fp': int(val_metrics['fp']),\n            'fn': int(val_metrics['fn'])\n        }\n    else:\n        final_val_metrics = {\n            'dice': 0.0,\n            'f05': 0.0,\n            'iou': 0.0,\n            'precision': 0.0,\n            'recall': 0.0,\n            'threshold': 0.5,\n            'tp': 0,\n            'tn': 0,\n            'fp': 0,\n            'fn': 0\n        }\n    \n    # Save metrics summary\n    if config.save_metrics:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_metrics_summary(\n            final_train_metrics,\n            final_val_metrics,\n            test_metrics,\n            filename=f'metrics_summary_{timestamp}.csv'\n        )\n    \n    # Save best model\n    if best_model_state:\n        torch.save(best_model_state, 'best_model.pth')\n        print(\"Best model saved to best_model.pth\")\n    \n    # ============================================\n    # FINAL SUMMARY\n    # ============================================\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL RESULTS SUMMARY\")\n    print(\"=\"*80)\n    \n    print(f\"\\nCONFIGURATION:\")\n    print(f\"  Image Size: {config.img_size}x{config.img_size} (from first code)\")\n    print(f\"  Preprocessing: Bilateral Filter + CLAHE + Normalization\")\n    print(f\"  Input Slices: {config.num_input_slices} (randomly selected from {len(config.slices)} total)\")\n    print(f\"  Loss Function: BCE Loss (from second code)\")\n    print(f\"  Learning Rate: {config.lr} (3e-4)\")\n    print(f\"  Model: UNet with EfficientNet-B0 encoder\")\n    print(f\"  Training Strategy: Fragments 2 & 3 only (Fragment 1 held out for final test)\")\n    \n    print(f\"\\nPERFORMANCE METRICS:\")\n    print(f\"  Best Validation Dice: {best_val_dice:.4f}\")\n    print(f\"  Training Dice: {final_train_metrics['dice']:.4f}\")\n    if val_loader:\n        print(f\"  Validation Dice: {final_val_metrics['dice']:.4f}\")\n    print(f\"  Test Dice (Fragment 1): {test_metrics['dice']:.4f}\")\n    print(f\"  Generalization Gap (Train-Test): {final_train_metrics['dice'] - test_metrics['dice']:.4f}\")\n    \n    print(f\"\\nTEST RESULTS DETAILS:\")\n    print(f\"  Dice Score: {test_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {test_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {test_metrics['iou']:.4f}\")\n    print(f\"  Precision: {test_metrics['precision']:.4f}\")\n    print(f\"  Recall: {test_metrics['recall']:.4f}\")\n    print(f\"  Optimal Threshold: {test_metrics['threshold']:.3f}\")\n    print(f\"  TP: {test_metrics['tp']:,} | TN: {test_metrics['tn']:,} | FP: {test_metrics['fp']:,} | FN: {test_metrics['fn']:,}\")\n    \n    print(f\"\\nDATA SAVED:\")\n    print(f\"  Preprocessing visualization: preprocessing_visualizations/\")\n    print(f\"  Test results: test_results/\")\n    print(f\"  Training plots: training_plots/\")\n    print(f\"  Best model: best_model.pth\")\n    \n    if config.save_history:\n        print(f\"  Training history: history/\")\n    if config.save_metrics:\n        print(f\"  Metrics summary: metrics/\")\n    \n    print(f\"\\nKEY FEATURES:\")\n    print(f\"  ✓ Image preprocessing: bilateral filter, CLAHE, normalization (from first code)\")\n    print(f\"  ✓ Image size: 352x352 (from first code)\")\n    print(f\"  ✓ Input slices: {config.num_input_slices} (random consecutive selection)\")\n    print(f\"  ✓ Learning rate: {config.lr} (3e-4)\")\n    print(f\"  ✓ Loss function: BCE only (from second code)\")\n    print(f\"  ✓ Model: UNet with EfficientNet-B0 encoder\")\n    print(f\"  ✓ Data augmentation: flips, brightness/contrast, distortion, blur, rotation\")\n    print(f\"  ✓ Patch extraction: {config.patch_size}x{config.patch_size} with stride {config.stride}\")\n    print(f\"  ✓ Early stopping with patience {patience}\")\n    print(f\"  ✓ Mixed precision training\")\n    print(f\"  ✓ Preprocessing visualization (raw vs processed comparison)\")\n    print(f\"  ✓ Complete test visualization and metrics saved\")\n    print(f\"  ✓ Fragment 1 used ONLY for final testing (no contamination)\")\n\n# Run the main function\nif __name__ == \"__main__\":\n    try:\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        \n        main()\n    except Exception as e:\n        print(f\"Error occurred: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#improved resnet model with paches \nimport os\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nfrom tqdm import tqdm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# INSTALL LIBRARIES\n# =========================================================\n\n#!pip install -q segmentation-models-pytorch albumentations\n!pip install segmentation-models-pytorch==0.2.0\n# =========================================================\n# IMPORTS\n# =========================================================\n\nimport os\nimport cv2\nimport random\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nimport albumentations as A\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device(\"cuda\")\n\n# =========================================================\n# PATHS\n# =========================================================\n\nfragment1_path = \"/kaggle/input/vesuvius-challenge-ink-detection/train/1\"\n\ntrain_paths = [\n    \"/kaggle/input/vesuvius-challenge-ink-detection/train/2\",\n    \"/kaggle/input/vesuvius-challenge-ink-detection/train/3\"\n]\n \n# =========================================================\n# CONFIG\n# =========================================================\n\nCFG = {\n    \"patch_size\":256,\n    \"stride\":128,\n    \"batch_size\":16,\n    \"epochs\":10,\n    \"lr\":3e-4,\n    \"slices\":list(range(28,38)),\n    \"num_workers\":2\n}\n\n# =========================================================\n# NORMALIZATION\n# =========================================================\n\ndef normalize_fragment(volume):\n\n    low, high = np.percentile(volume,(1,99))\n    volume = np.clip(volume,low,high)\n\n    volume = (volume-volume.mean())/(volume.std()+1e-6)\n\n    return volume\n\n# =========================================================\n# LOAD VOLUME\n# =========================================================\n\ndef load_volume(fragment_path):\n\n    slices=[]\n\n    for i in CFG[\"slices\"]:\n\n        img = cv2.imread(\n            os.path.join(fragment_path,\"surface_volume\",f\"{i:02}.tif\"),0\n        )\n\n        slices.append(img)\n\n    volume = np.stack(slices,axis=-1).astype(np.float32)\n\n    volume = normalize_fragment(volume)\n\n    return volume\n\n# =========================================================\n# LOAD LABELS (inklabels * mask)\n# =========================================================\n\ndef load_label(fragment_path):\n\n    ink = cv2.imread(os.path.join(fragment_path,\"inklabels.png\"),0)\n    mask = cv2.imread(os.path.join(fragment_path,\"mask.png\"),0)\n\n    ink = (ink>0).astype(np.float32)\n    mask = (mask>0).astype(np.float32)\n\n    label = ink * mask\n\n    return label\n\n# =========================================================\n# COORDINATE GENERATION\n# =========================================================\n\ndef generate_coords(label):\n\n    coords=[]\n    H,W = label.shape\n\n    for y in range(0,H-CFG[\"patch_size\"],CFG[\"stride\"]):\n        for x in range(0,W-CFG[\"patch_size\"],CFG[\"stride\"]):\n\n            patch = label[y:y+CFG[\"patch_size\"],x:x+CFG[\"patch_size\"]]\n\n            ink_ratio = patch.mean()\n\n            coords.append((y,x,ink_ratio))\n\n    return coords\n\n# =========================================================\n# TRAIN / VAL SPLIT\n# =========================================================\n\ndef split_coords(coords):\n\n    random.shuffle(coords)\n\n    split = int(len(coords)*0.8)\n\n    return coords[:split], coords[split:]\n\n# =========================================================\n# AUGMENTATIONS\n# =========================================================\n\ntrain_aug = A.Compose([\n\n    A.HorizontalFlip(p=0.5),\n\n    A.VerticalFlip(p=0.5),\n\n    A.RandomRotate90(p=0.5),\n\n    A.ShiftScaleRotate(\n        shift_limit=0.1,\n        scale_limit=0.1,\n        rotate_limit=20,\n        p=0.5\n    ),\n\n    A.ElasticTransform(p=0.3),\n\n])\n\n# =========================================================\n# DATASET\n# =========================================================\n\nclass InkDataset(Dataset):\n\n    def __init__(self, volumes, labels, coords, transform=None):\n\n        self.volumes = volumes\n        self.labels = labels\n        self.coords = coords\n        self.transform = transform\n\n    def __len__(self):\n        return 8000\n\n    def sample_coord(self, coords):\n\n        ink = [c for c in coords if c[2] > 0.01]\n        bg = [c for c in coords if c[2] <= 0.01]\n\n        if random.random() < 0.7:\n            return random.choice(ink)\n        else:\n            return random.choice(bg)\n\n    def __getitem__(self, idx):\n\n        frag = random.randint(0,len(self.volumes)-1)\n\n        volume = self.volumes[frag]\n        label = self.labels[frag]\n\n        y,x,_ = self.sample_coord(self.coords[frag])\n\n        patch = volume[y:y+CFG[\"patch_size\"],x:x+CFG[\"patch_size\"]]\n        patch_label = label[y:y+CFG[\"patch_size\"],x:x+CFG[\"patch_size\"]]\n\n        if self.transform:\n\n            aug = self.transform(image=patch,mask=patch_label)\n\n            patch = aug[\"image\"]\n            patch_label = aug[\"mask\"]\n\n        patch = torch.tensor(patch).permute(2,0,1).float()\n        patch_label = torch.tensor(patch_label).unsqueeze(0).float()\n\n        return patch, patch_label\n\n# =========================================================\n# LOAD TRAIN DATA\n# =========================================================\n\ntrain_volumes=[]\ntrain_labels=[]\ntrain_coords=[]\nval_coords=[]\n\nfor p in train_paths:\n\n    volume = load_volume(p)\n    label = load_label(p)\n\n    coords = generate_coords(label)\n\n    tr,val = split_coords(coords)\n\n    train_volumes.append(volume)\n    train_labels.append(label)\n\n    train_coords.append(tr)\n    val_coords.append(val)\n\n# =========================================================\n# DATALOADER\n# =========================================================\n\ntrain_dataset = InkDataset(train_volumes,train_labels,train_coords,train_aug)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG[\"batch_size\"],\n    shuffle=True,\n    num_workers=CFG[\"num_workers\"],\n    pin_memory=True\n)\n\n# =========================================================\n# MODEL\n# =========================================================\n\nmodel = smp.Unet(\n    encoder_name=\"efficientnet-b7\",\n    encoder_weights=\"imagenet\",\n    in_channels=len(CFG[\"slices\"]),\n    classes=1,\n    decoder_attention_type=\"scse\"\n).to(device)\n\n# =========================================================\n# LOSS\n# =========================================================\n\nbce = nn.BCEWithLogitsLoss()\n\ndef dice_loss(pred,target):\n\n    pred = torch.sigmoid(pred)\n\n    inter = (pred*target).sum()\n    union = pred.sum() + target.sum()\n\n    return 1 - (2*inter+1)/(union+1)\n\ndef loss_fn(pred,target):\n\n    return 0.5*bce(pred,target) + 0.5*dice_loss(pred,target)\n\n# =========================================================\n# METRICS\n# =========================================================\n\ndef metrics(pred,target):\n\n    pred = (pred > 0.4).float()\n\n    TP = (pred * target).sum()\n    FP = (pred * (1-target)).sum()\n    FN = ((1-pred) * target).sum()\n\n    dice = (2*TP) / (2*TP + FP + FN + 1e-6)\n    iou = TP / (TP + FP + FN + 1e-6)\n\n    precision = TP / (TP + FP + 1e-6)\n    recall = TP / (TP + FN + 1e-6)\n\n    return dice,iou,precision,recall\n\n# =========================================================\n# OPTIMIZER\n# =========================================================\n\noptimizer = torch.optim.AdamW(model.parameters(),lr=CFG[\"lr\"])\n\nscaler = torch.cuda.amp.GradScaler()\n\n# =========================================================\n# TRAINING LOOP\n# =========================================================\n\nfor epoch in range(CFG[\"epochs\"]):\n\n    model.train()\n\n    total_loss = 0\n\n    loop = tqdm(train_loader)\n\n    for imgs,labels in loop:\n\n        imgs = imgs.to(device)\n        labels = labels.to(device)\n\n        with torch.cuda.amp.autocast():\n\n            preds = model(imgs)\n\n            loss = loss_fn(preds,labels)\n            #loss = loss_fn(preds * mask, labels * mask)\n\n        optimizer.zero_grad()\n\n        scaler.scale(loss).backward()\n\n        scaler.step(optimizer)\n\n        scaler.update()\n\n        total_loss += loss.item()\n\n        loop.set_description(f\"Epoch {epoch+1}\")\n\n    print(\"Train Loss:\", total_loss/len(train_loader))\n\n# =========================================================\n# SLIDING WINDOW INFERENCE\n# =========================================================\n\ndef sliding_window(volume):\n\n    H,W,C = volume.shape\n\n    pred = np.zeros((H,W))\n    count = np.zeros((H,W))\n\n    for y in range(0,H-CFG[\"patch_size\"],CFG[\"stride\"]):\n        for x in range(0,W-CFG[\"patch_size\"],CFG[\"stride\"]):\n\n            patch = volume[y:y+CFG[\"patch_size\"],x:x+CFG[\"patch_size\"]]\n\n            patch = torch.tensor(patch).permute(2,0,1).unsqueeze(0).cuda()\n\n            with torch.no_grad():\n\n                p = torch.sigmoid(model(patch)).cpu().numpy()[0,0]\n\n            pred[y:y+CFG[\"patch_size\"],x:x+CFG[\"patch_size\"]] += p\n            count[y:y+CFG[\"patch_size\"],x:x+CFG[\"patch_size\"]] += 1\n\n    pred /= count\n\n    return pred\n\n# =========================================================\n# INFERENCE ON FRAGMENT 1\n# =========================================================\n\nvol1 = load_volume(fragment1_path)\n\npred_map = sliding_window(vol1)\n\nmask_pred = (pred_map > 0.4).astype(np.uint8)\n\ncv2.imwrite(\"fragment1_prediction.png\", mask_pred*255)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install segmentation-models-pytorch==0.2.0\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nfrom skimage.filters import threshold_otsu\nimport warnings\nimport random\nimport math\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport pandas as pd\nimport json\nfrom datetime import datetime\n\nwarnings.filterwarnings('ignore')\n#/kaggle/input/vesuvius-challenge-ink-detection/test\n# Configuration\nclass Config:\n    # Data paths\n    fragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\n    train_paths = [\n        '/kaggle/input/vesuvius-challenge/train-ink-detection/2',  # Training data\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n    ]\n    \n    # Model parameters\n    img_size = 352\n    slices = list(range(12, 31))  # Slice 12-30 (19 slices)\n    batch_size = 8  # Increased batch size since no validation\n    epochs = 50\n    lr = 3e-4\n    \n    # 3D volumetric setup\n    num_input_slices = 5\n    \n    # Memory optimization\n    max_train_samples = 1500  # Increased for training\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 5\n    \n    # Results saving\n    save_history = True\n    save_metrics = True\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Training strategy - NO VALIDATION, just train on fragments 2 & 3, test on fragment 1\n    use_validation = False  # Set to False to disable validation during training\n\nconfig = Config()\n\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data with enhanced preprocessing\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_filtered = cv2.bilateralFilter(img_resized, 5, 50, 50)\n                clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n                img_enhanced = clahe.apply(img_filtered)\n                img_normalized = img_enhanced.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        print(f\"Loaded volume from {base_path} with {volume.shape[0]} slices\")\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio for {base_path}: {ink_ratio:.4f} ({mask_binary.sum()} pixels)\")\n            return mask_binary\n    return None\n\n# Simple Training Dataset (no region masks)\nclass TrainingDataset(Dataset):\n    def __init__(self, volumes, masks, fragment_ids=None, augment=True, max_samples=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.augment = augment\n        self.fragment_ids = fragment_ids if fragment_ids else [0] * len(volumes)\n        self.max_samples = max_samples\n        self.samples = self._prepare_samples()\n        print(f\"Created training dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        \n        for vol_idx, volume in enumerate(self.volumes):\n            start_idx = config.num_input_slices // 2\n            end_idx = volume.shape[0] - config.num_input_slices // 2\n            \n            max_samples_per_volume = 500  # Increased for better coverage\n            for _ in range(max_samples_per_volume):\n                center_slice = np.random.randint(start_idx, end_idx)\n                y = np.random.randint(0, volume.shape[1])\n                x = np.random.randint(0, volume.shape[2])\n                \n                samples.append({\n                    'volume_idx': vol_idx,\n                    'center_slice': center_slice,\n                    'y': y,\n                    'x': x,\n                    'fragment_id': self.fragment_ids[vol_idx]\n                })\n        \n        if self.max_samples and len(samples) > self.max_samples:\n            samples = samples[:self.max_samples]\n        \n        return samples\n    \n    def _extract_patch(self, volume, center_slice, y, x):\n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        half_size = config.img_size // 2\n        y_start = max(0, y - half_size)\n        y_end = min(volume.shape[1], y + half_size)\n        x_start = max(0, x - half_size)\n        x_end = min(volume.shape[2], x + half_size)\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = volume[slice_idx]\n            patch = slice_data[y_start:y_end, x_start:x_end]\n            \n            if patch.shape != (config.img_size, config.img_size):\n                pad_y = config.img_size - patch.shape[0]\n                pad_x = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((0, pad_y), (0, pad_x)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        return volume_patch, y_start, y_end, x_start, x_end\n    \n    def _apply_augmentation(self, volume_slices, mask_patch):\n        volume_slices = volume_slices.copy()\n        mask_patch = mask_patch.copy()\n        \n        # Random horizontal flip\n        if random.random() > 0.5:\n            volume_slices = np.ascontiguousarray(np.flip(volume_slices, axis=2))\n            mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=1))\n        \n        # Random vertical flip\n        if random.random() > 0.5:\n            volume_slices = np.ascontiguousarray(np.flip(volume_slices, axis=1))\n            mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=0))\n        \n        # Random brightness/contrast adjustment\n        for i in range(volume_slices.shape[0]):\n            if random.random() > 0.5:\n                alpha = random.uniform(0.8, 1.2)\n                beta = random.uniform(-0.1, 0.1)\n                volume_slices[i] = np.clip(alpha * volume_slices[i] + beta, 0, 1)\n        \n        return volume_slices, mask_patch\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        vol_idx = sample['volume_idx']\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        fragment_id = sample['fragment_id']\n        \n        volume_patch, y_start, y_end, x_start, x_end = self._extract_patch(\n            self.volumes[vol_idx], center_slice, y, x\n        )\n        \n        mask = self.masks[vol_idx]\n        mask_patch = mask[y_start:y_end, x_start:x_end]\n        \n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y = config.img_size - mask_patch.shape[0]\n            pad_x = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((0, pad_y), (0, pad_x)), mode='constant')\n        \n        if self.augment:\n            volume_patch, mask_patch = self._apply_augmentation(volume_patch, mask_patch)\n        \n        volume_patch = np.ascontiguousarray(volume_patch)\n        mask_patch = np.ascontiguousarray(mask_patch)\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor\n\n# Test Dataset for Full Fragment 1\nclass TestDataset(Dataset):\n    def __init__(self, volume, mask, fragment_id=0, stride_factor=2):\n        self.volume = volume\n        self.mask = mask\n        self.fragment_id = fragment_id\n        self.stride_factor = stride_factor\n        self.samples = self._prepare_samples()\n        print(f\"Created test dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        start_idx = config.num_input_slices // 2\n        end_idx = self.volume.shape[0] - config.num_input_slices // 2\n        \n        stride = (config.img_size // 2) * self.stride_factor\n        for center_slice in range(start_idx, end_idx, 2):\n            for y in range(0, self.volume.shape[1], stride):\n                for x in range(0, self.volume.shape[2], stride):\n                    samples.append({\n                        'center_slice': center_slice,\n                        'y': y,\n                        'x': x\n                    })\n        return samples\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        \n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = self.volume[slice_idx]\n            patch = slice_data[y:y+config.img_size, x:x+config.img_size]\n            \n            if patch.shape != (config.img_size, config.img_size):\n                pad_y = config.img_size - patch.shape[0]\n                pad_x = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((0, pad_y), (0, pad_x)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        \n        mask_patch = self.mask[y:y+config.img_size, x:x+config.img_size]\n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y = config.img_size - mask_patch.shape[0]\n            pad_x = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((0, pad_y), (0, pad_x)), mode='constant')\n        \n        volume_patch = np.ascontiguousarray(volume_patch)\n        mask_patch = np.ascontiguousarray(mask_patch)\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor, y, x\n\n# 3D UNet Model\nclass VolumetricUNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = smp.Unet(\n            encoder_name=\"resnet34\",\n            encoder_weights=\"imagenet\",\n            in_channels=config.num_input_slices,\n            classes=1,\n            encoder_depth=4,\n            decoder_channels=[128, 64, 32, 16],\n            activation=None\n        )\n        \n    def forward(self, x):\n        return self.model(x.contiguous())\n\n# BCE Loss only\nclass BCEOnlyLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        \n    def forward(self, pred, target):\n        return self.bce_loss(pred, target)\n\ndef calculate_all_metrics(pred, target, threshold=0.5):\n    \"\"\"Calculate all metrics including TP, TN, FP, FN\"\"\"\n    pred_binary = (pred > threshold).astype(np.float32)\n    pred_flat = pred_binary.flatten().astype(int)\n    target_flat = target.flatten().astype(int)\n    \n    tn, fp, fn, tp = confusion_matrix(target_flat, pred_flat, labels=[0, 1]).ravel()\n    \n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n    \n    beta_squared = 0.5 ** 2\n    f05 = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (beta_squared * precision + recall) > 0 else 0\n    \n    iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n    \n    return {\n        'precision': float(precision),\n        'recall': float(recall),\n        'dice': float(dice),\n        'f05': float(f05),\n        'iou': float(iou),\n        'tp': int(tp),\n        'tn': int(tn),\n        'fp': int(fp),\n        'fn': int(fn)\n    }\n\ndef dice_score(pred, target, smooth=1e-6):\n    pred = pred.flatten()\n    target = target.flatten()\n    intersection = (pred * target).sum()\n    return float((2. * intersection + smooth) / (pred.sum() + target.sum() + smooth))\n\ndef save_test_visualization(volume, mask, predictions, metrics, fragment_name=\"fragment1\"):\n    \"\"\"Save comprehensive test visualization\"\"\"\n    os.makedirs('test_results', exist_ok=True)\n    \n    middle_slice_idx = volume.shape[0] // 2\n    middle_slice = volume[middle_slice_idx]\n    \n    pred_binary = (predictions > metrics['threshold']).astype(np.float32)\n    \n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n    \n    # Row 1\n    axes[0, 0].imshow(middle_slice, cmap='gray')\n    axes[0, 0].set_title(f'Input Image\\n(Middle Slice)')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(mask, cmap='gray')\n    axes[0, 1].set_title(f'Ground Truth\\nInk: {mask.sum()/mask.size:.2%}')\n    axes[0, 1].axis('off')\n    \n    axes[0, 2].imshow(predictions, cmap='jet', vmin=0, vmax=1)\n    axes[0, 2].set_title(f'Predictions\\n(Probabilities)')\n    axes[0, 2].axis('off')\n    \n    # Row 2\n    pred_overlay = np.stack([\n        middle_slice * 0.5 + pred_binary * 0.5,\n        middle_slice,\n        middle_slice\n    ], axis=-1)\n    axes[1, 0].imshow(pred_overlay)\n    axes[1, 0].set_title('Input + Binary Predictions\\n(Red = Predicted Ink)')\n    axes[1, 0].axis('off')\n    \n    error_map = np.zeros((mask.shape[0], mask.shape[1], 3))\n    false_positives = (pred_binary == 1) & (mask == 0)\n    false_negatives = (pred_binary == 0) & (mask == 1)\n    \n    error_map[:, :, 0] = false_positives * 0.8  # Red = False Positives\n    error_map[:, :, 2] = false_negatives * 0.8  # Blue = False Negatives\n    \n    axes[1, 1].imshow(error_map)\n    axes[1, 1].set_title('Error Map\\n(Red=FP, Blue=FN)')\n    axes[1, 1].axis('off')\n    \n    # Confusion matrix visualization\n    confusion_img = np.zeros((100, 100, 3))\n    confusion_img[:50, :50, :] = 0.2  # TN area (gray)\n    confusion_img[:50, 50:, 0] = 0.8  # FP area (red)\n    confusion_img[50:, :50, 2] = 0.8  # FN area (blue)\n    confusion_img[50:, 50:, 1] = 0.8  # TP area (green)\n    \n    axes[1, 2].imshow(confusion_img)\n    axes[1, 2].set_title('Confusion Matrix\\n(Green=TP, Red=FP, Blue=FN, Gray=TN)')\n    axes[1, 2].axis('off')\n    \n    # Add metrics text\n    metrics_text = f\"TEST RESULTS - Full Fragment 1\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Dice Score: {metrics['dice']:.4f}\\n\"\n    metrics_text += f\"F0.5 Score: {metrics['f05']:.4f}\\n\"\n    metrics_text += f\"IOU Score: {metrics['iou']:.4f}\\n\"\n    metrics_text += f\"Precision: {metrics['precision']:.4f}\\n\"\n    metrics_text += f\"Recall: {metrics['recall']:.4f}\\n\"\n    metrics_text += f\"Optimal Threshold: {metrics['threshold']:.3f}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"True Positives: {metrics['tp']:,}\\n\"\n    metrics_text += f\"True Negatives: {metrics['tn']:,}\\n\"\n    metrics_text += f\"False Positives: {metrics['fp']:,}\\n\"\n    metrics_text += f\"False Negatives: {metrics['fn']:,}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Training Data: Fragments 2 & 3 (Full)\\n\"\n    metrics_text += f\"Test Data: Fragment 1 (Full)\\n\"\n    metrics_text += f\"Image Size: {config.img_size}x{config.img_size}\\n\"\n    metrics_text += f\"Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\"\n    \n    plt.figtext(0.02, 0.02, metrics_text, fontsize=9, fontfamily='monospace',\n                bbox=dict(boxstyle=\"round,pad=0.5\", facecolor=\"lightgray\", alpha=0.8))\n    \n    plt.suptitle(f'TEST RESULTS - {fragment_name} - Full Fragment Evaluation\\n'\n                 f'Train on Fragments 2 & 3 | Test on Fragment 1\\n'\n                 f'Loss Function: BCE Only', \n                 fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(f'test_results/{fragment_name}_test_results.png', \n                dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Save predictions\n    np.save(f'test_results/{fragment_name}_predictions.npy', predictions)\n    np.save(f'test_results/{fragment_name}_predictions_binary.npy', pred_binary)\n    \n    print(f\"Test visualization saved to: test_results/{fragment_name}_test_results.png\")\n    print(f\"Test predictions saved to: test_results/{fragment_name}_predictions.npy\")\n\ndef save_history(history, filename='training_history.json'):\n    \"\"\"Save training history to JSON file\"\"\"\n    os.makedirs('history', exist_ok=True)\n    \n    def convert_numpy_types(obj):\n        if isinstance(obj, np.integer):\n            return int(obj)\n        elif isinstance(obj, np.floating):\n            return float(obj)\n        elif isinstance(obj, np.ndarray):\n            return obj.tolist()\n        elif isinstance(obj, dict):\n            return {key: convert_numpy_types(value) for key, value in obj.items()}\n        elif isinstance(obj, list):\n            return [convert_numpy_types(item) for item in obj]\n        else:\n            return obj\n    \n    history_serializable = convert_numpy_types(history)\n    \n    with open(os.path.join('history', filename), 'w') as f:\n        json.dump(history_serializable, f, indent=2)\n    \n    print(f\"Training history saved to {filename}\")\n\ndef save_metrics_summary(train_metrics, test_metrics, filename='metrics_summary.csv'):\n    \"\"\"Save metrics to CSV file\"\"\"\n    os.makedirs('metrics', exist_ok=True)\n    \n    summary_data = {\n        'Phase': ['Training (Fragments 2 & 3)', 'Test (Fragment 1)'],\n        'Loss_Function': ['BCE Only', 'BCE Only'],\n        'Dice_Score': [float(train_metrics['dice']), float(test_metrics['dice'])],\n        'F0.5_Score': [float(train_metrics['f05']), float(test_metrics['f05'])],\n        'IOU_Score': [float(train_metrics['iou']), float(test_metrics['iou'])],\n        'Precision': [float(train_metrics['precision']), float(test_metrics['precision'])],\n        'Recall': [float(train_metrics['recall']), float(test_metrics['recall'])],\n        'Optimal_Threshold': [float(train_metrics['threshold']), float(test_metrics['threshold'])],\n        'TP': [int(train_metrics['tp']), int(test_metrics['tp'])],\n        'TN': [int(train_metrics['tn']), int(test_metrics['tn'])],\n        'FP': [int(train_metrics['fp']), int(test_metrics['fp'])],\n        'FN': [int(train_metrics['fn']), int(test_metrics['fn'])]\n    }\n    \n    df = pd.DataFrame(summary_data)\n    df.to_csv(os.path.join('metrics', filename), index=False)\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"METRICS SUMMARY - Train on Fragments 2 & 3, Test on Fragment 1 (BCE Only Loss)\")\n    print(\"=\"*80)\n    print(df.to_string())\n    print(f\"\\nMetrics summary saved to {filename}\")\n\ndef plot_training_history(history):\n    \"\"\"Plot training history\"\"\"\n    os.makedirs('training_plots', exist_ok=True)\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Loss plot\n    axes[0, 0].plot(epochs, history['train_loss'], 'b-', linewidth=2)\n    axes[0, 0].set_title('Training Loss (BCE)')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Dice score plot\n    axes[0, 1].plot(epochs, history['train_dice'], 'g-', linewidth=2)\n    axes[0, 1].set_title('Training Dice Score')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Dice Score')\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # F0.5 score plot\n    axes[1, 0].plot(epochs, history['train_f05'], 'r-', linewidth=2)\n    axes[1, 0].set_title('Training F0.5 Score')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('F0.5 Score')\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # IOU score plot\n    axes[1, 1].plot(epochs, history['train_iou'], 'm-', linewidth=2)\n    axes[1, 1].set_title('Training IOU Score')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('IOU Score')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.suptitle('Training History - BCE Only Loss\\nTrain on Fragments 2 & 3', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/training_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Precision and Recall plot\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n    \n    ax1.plot(epochs, history['train_precision'], 'b-', label='Precision', linewidth=2)\n    ax1.plot(epochs, history['train_recall'], 'r-', label='Recall', linewidth=2)\n    ax1.set_title('Training Precision & Recall')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Score')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    ax2.plot(epochs, history['train_precision'], 'b-', label='Precision', linewidth=2)\n    ax2.plot(epochs, history['train_recall'], 'r-', label='Recall', linewidth=2)\n    ax2.set_title('Training Precision & Recall (Zoomed)')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Score')\n    ax2.set_ylim([0, 1])\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    plt.suptitle('Precision and Recall - BCE Only Loss\\nTraining on Fragments 2 & 3', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/precision_recall_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n\ndef main():\n    print(\"=\"*80)\n    print(\"3D UNet with ResNet34 - BCE Only Loss - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    print(\"TRAINING STRATEGY:\")\n    print(f\"  Training Data: Fragments 2 and 3 (100%)\")\n    print(f\"  Test Data: Fragment 1 (100%)\")\n    print(f\"  Validation: DISABLED (pure training on fragments 2 & 3)\")\n    print(f\"  Loss Function: BCE Only\")\n    print(\"=\"*80)\n    \n    # Set random seeds\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Load data\n    print(\"\\nLoading data...\")\n    \n    # Load fragment 1 (test data)\n    fragment1_volume = load_volume_data(config.fragment1_path, config.slices)\n    fragment1_mask = load_mask_data(config.fragment1_path)\n    \n    if fragment1_volume is None or fragment1_mask is None:\n        raise ValueError(f\"Failed to load fragment 1 data from {config.fragment1_path}\")\n    \n    # Load training fragments 2 and 3\n    train_volumes = []\n    train_masks = []\n    fragment_ids = []\n    \n    for i, path in enumerate(config.train_paths):\n        volume = load_volume_data(path, config.slices)\n        mask = load_mask_data(path)\n        \n        if volume is not None and mask is not None:\n            train_volumes.append(volume)\n            train_masks.append(mask)\n            fragment_ids.append(i + 1)\n            print(f\"Fragment {i+2} loaded - shape: {volume.shape}\")\n    \n    print(f\"\\nLoaded {len(train_volumes)} training fragments (2 and 3)\")\n    print(f\"Fragment 1 (test) shape: {fragment1_volume.shape}\")\n    \n    # Prepare datasets\n    print(f\"\\nPreparing datasets...\")\n    \n    # Training dataset (fragments 2 & 3 only)\n    train_dataset = TrainingDataset(\n        volumes=train_volumes,\n        masks=train_masks,\n        fragment_ids=fragment_ids,\n        augment=True,\n        max_samples=config.max_train_samples\n    )\n    \n    # Test dataset (full fragment 1)\n    test_dataset = TestDataset(\n        volume=fragment1_volume,\n        mask=fragment1_mask,\n        fragment_id=0,\n        stride_factor=2\n    )\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=1,\n        shuffle=False,\n        num_workers=0\n    )\n    \n    print(f\"\\nData Loaders Created:\")\n    print(f\"  Training samples: {len(train_dataset)} (Fragments 2 & 3)\")\n    print(f\"  Test samples: {len(test_dataset)} (Fragment 1)\")\n    \n    # Initialize model\n    print(f\"\\nInitializing model...\")\n    model = VolumetricUNet().to(config.device)\n    print(f\"Model initialized on {config.device}\")\n    \n    # Initialize optimizer and loss (BCE only)\n    optimizer = optim.Adam(model.parameters(), lr=config.lr, weight_decay=1e-5)\n    criterion = BCEOnlyLoss()  # Using BCE only\n    \n    # Training history\n    history = {\n        'train_loss': [],\n        'train_dice': [],\n        'train_f05': [],\n        'train_iou': [],\n        'train_precision': [],\n        'train_recall': []\n    }\n    \n    print(f\"\\nStarting training for {config.epochs} epochs...\")\n    print(\"=\"*80)\n    \n    for epoch in range(config.epochs):\n        # Training phase\n        model.train()\n        epoch_train_loss = 0\n        train_outputs = []\n        train_targets = []\n        \n        train_iterator = tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.epochs}')\n        for images, masks in train_iterator:\n            images, masks = images.to(config.device), masks.to(config.device)\n            images = images.contiguous()\n            masks = masks.contiguous()\n            \n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            loss.backward()\n            \n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            epoch_train_loss += loss.item()\n            \n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_outputs.append(preds.cpu().numpy())\n                train_targets.append(masks.cpu().numpy())\n            \n            train_iterator.set_postfix(loss=loss.item())\n            \n            if torch.cuda.is_available() and len(train_outputs) % 100 == 0:\n                torch.cuda.empty_cache()\n        \n        # Calculate training metrics\n        train_outputs = np.concatenate(train_outputs)\n        train_targets = np.concatenate(train_targets)\n        \n        best_train_threshold = 0.5\n        best_train_dice = 0\n        for thresh in np.arange(0.1, 0.9, 0.05):\n            preds = (train_outputs > thresh).astype(np.float32)\n            dice = dice_score(preds, train_targets)\n            if dice > best_train_dice:\n                best_train_dice = dice\n                best_train_threshold = thresh\n        \n        train_preds = (train_outputs > best_train_threshold).astype(np.float32)\n        train_metrics = calculate_all_metrics(train_preds, train_targets)\n        train_metrics['threshold'] = float(best_train_threshold)\n        \n        # Store training metrics\n        history['train_loss'].append(float(epoch_train_loss / len(train_loader)))\n        history['train_dice'].append(float(train_metrics['dice']))\n        history['train_f05'].append(float(train_metrics['f05']))\n        history['train_iou'].append(float(train_metrics['iou']))\n        history['train_precision'].append(float(train_metrics['precision']))\n        history['train_recall'].append(float(train_metrics['recall']))\n        \n        # Print progress\n        print(f\"\\nEpoch {epoch+1}/{config.epochs}:\")\n        print(f\"  Training - BCE Loss: {epoch_train_loss/len(train_loader):.4f}, \"\n              f\"Dice: {train_metrics['dice']:.4f}, F0.5: {train_metrics['f05']:.4f}\")\n        \n        # Clear cache\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    # Training complete - save model\n    torch.save(model.state_dict(), 'trained_model_fragments_2_3_bce_only.pth')\n    print(f\"\\nModel saved to: trained_model_fragments_2_3_bce_only.pth\")\n    \n    # Clear cache before testing\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Test phase on full fragment 1\n    print(\"\\n\" + \"=\"*80)\n    print(\"TESTING PHASE - Full Fragment 1 (BCE Only Model)\")\n    print(\"=\"*80)\n    \n    model.eval()\n    \n    pred_accumulator = np.zeros(fragment1_mask.shape, dtype=np.float32)\n    count_accumulator = np.zeros(fragment1_mask.shape, dtype=np.float32)\n    \n    with torch.no_grad():\n        test_iterator = tqdm(test_loader, desc=\"Testing on fragment 1\")\n        for images, masks, y_coords, x_coords in test_iterator:\n            images = images.to(config.device).contiguous()\n            outputs = torch.sigmoid(model(images))\n            \n            for i in range(outputs.shape[0]):\n                y = y_coords[i].item()\n                x = x_coords[i].item()\n                \n                patch = outputs[i, 0].cpu().numpy()\n                h, w = patch.shape\n                \n                pred_accumulator[y:y+h, x:x+w] += patch\n                count_accumulator[y:y+h, x:x+w] += 1\n            \n            if torch.cuda.is_available() and test_iterator.n % 10 == 0:\n                torch.cuda.empty_cache()\n    \n    # Average overlapping predictions\n    final_predictions = pred_accumulator / np.maximum(count_accumulator, 1)\n    \n    # Calculate test metrics on full fragment 1\n    best_test_threshold = 0.5\n    best_test_dice = 0\n    for thresh in np.arange(0.1, 0.9, 0.05):\n        preds = (final_predictions > thresh).astype(np.float32)\n        dice = dice_score(preds, fragment1_mask)\n        if dice > best_test_dice:\n            best_test_dice = dice\n            best_test_threshold = thresh\n    \n    # Calculate all test metrics\n    test_preds = (final_predictions > best_test_threshold).astype(np.float32)\n    test_metrics = calculate_all_metrics(test_preds, fragment1_mask)\n    test_metrics['threshold'] = float(best_test_threshold)\n    \n    # Save test visualization\n    print(\"\\nSaving test visualization and metrics...\")\n    save_test_visualization(\n        fragment1_volume, \n        fragment1_mask, \n        final_predictions, \n        test_metrics,\n        \"fragment1\"\n    )\n    \n    # Save training history\n    if config.save_history:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_history(history, filename=f'training_history_bce_only_{timestamp}.json')\n        plot_training_history(history)\n    \n    # Final training metrics\n    final_train_metrics = {\n        'dice': float(history['train_dice'][-1]),\n        'f05': float(history['train_f05'][-1]),\n        'iou': float(history['train_iou'][-1]),\n        'precision': float(history['train_precision'][-1]),\n        'recall': float(history['train_recall'][-1]),\n        'threshold': float(best_train_threshold),\n        'tp': int(train_metrics['tp']),\n        'tn': int(train_metrics['tn']),\n        'fp': int(train_metrics['fp']),\n        'fn': int(train_metrics['fn'])\n    }\n    \n    # Save metrics summary\n    if config.save_metrics:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_metrics_summary(\n            final_train_metrics,\n            test_metrics,\n            filename=f'metrics_summary_bce_only_{timestamp}.csv'\n        )\n    \n    # Final summary\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL RESULTS SUMMARY - BCE Only Loss\")\n    print(\"Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    \n    print(f\"\\nCONFIGURATION:\")\n    print(f\"  Image Size: {config.img_size}x{config.img_size}\")\n    print(f\"  Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\")\n    print(f\"  Loss Function: BCE Only\")\n    \n    print(f\"\\nTRAINING DETAILS:\")\n    print(f\"  Training Data: Fragments 2 & 3 (100%)\")\n    print(f\"  Test Data: Fragment 1 (100%)\")\n    print(f\"  Validation: DISABLED\")\n    \n    print(f\"\\nTRAINING PERFORMANCE (Final Epoch):\")\n    print(f\"  BCE Loss: {history['train_loss'][-1]:.4f}\")\n    print(f\"  Dice Score: {final_train_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {final_train_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {final_train_metrics['iou']:.4f}\")\n    print(f\"  Precision: {final_train_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_train_metrics['recall']:.4f}\")\n    \n    print(f\"\\nTEST RESULTS (Fragment 1):\")\n    print(f\"  Dice Score: {test_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {test_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {test_metrics['iou']:.4f}\")\n    print(f\"  Precision: {test_metrics['precision']:.4f}\")\n    print(f\"  Recall: {test_metrics['recall']:.4f}\")\n    print(f\"  Optimal Threshold: {test_metrics['threshold']:.3f}\")\n    print(f\"  TP: {test_metrics['tp']:,} | TN: {test_metrics['tn']:,} | FP: {test_metrics['fp']:,} | FN: {test_metrics['fn']:,}\")\n    \n    print(f\"\\nDATA SAVED:\")\n    print(f\"  Test results: test_results/\")\n    print(f\"  Training plots: training_plots/\")\n    \n    if config.save_history:\n        print(f\"  Training history: history/\")\n    if config.save_metrics:\n        print(f\"  Metrics summary: metrics/\")\n    \n    print(f\"\\nMODEL SAVED:\")\n    print(f\"  trained_model_fragments_2_3_bce_only.pth\")\n    \n    print(f\"\\nKEY FEATURES:\")\n    print(f\"  1. Trained ONLY on Fragments 2 & 3 (no validation)\")\n    print(f\"  2. Tested on FULL Fragment 1\")\n    print(f\"  3. Image Size: {config.img_size}x{config.img_size}\")\n    print(f\"  4. All slices 12-30 included\")\n    print(f\"  5. Pure cross-fragment generalization test\")\n    print(f\"  6. Loss Function: BCE Only\")\n\n# Run the main function\nif __name__ == \"__main__\":\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        gpu_memory = torch.cuda.get_device_properties(0).total_memory / 1e9\n        print(f\"Available GPU memory: {gpu_memory:.1f} GB\")\n    \n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install segmentation-models-pytorch==0.2.0\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nfrom skimage.filters import threshold_otsu\nimport warnings\nimport random\nimport math\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport pandas as pd\nimport json\nfrom datetime import datetime\n\nwarnings.filterwarnings('ignore')\n#/kaggle/input/vesuvius-challenge-ink-detection/train\n# Configuration\nclass Config:\n    # Data paths\n    fragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\n    train_paths = [\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n    ]\n    \n    # Model parameters\n    img_size = 352\n    slices = list(range(12, 31))  # Slice 12-30 (19 slices)\n    batch_size = 8  # Increased batch size since no validation\n    epochs = 20\n    lr = 3e-4\n    \n    # 3D volumetric setup\n    num_input_slices = 5\n    \n    # Memory optimization\n    max_train_samples = 1500  # Increased for training\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 5\n    \n    # Results saving\n    save_history = True\n    save_metrics = True\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Training strategy - NO VALIDATION, just train on fragments 2 & 3, test on fragment 1\n    use_validation = False  # Set to False to disable validation during training\n\nconfig = Config()\n\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data with enhanced preprocessing\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_filtered = cv2.bilateralFilter(img_resized, 5, 50, 50)\n                clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n                img_enhanced = clahe.apply(img_filtered)\n                img_normalized = img_enhanced.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        print(f\"Loaded volume from {base_path} with {volume.shape[0]} slices\")\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio for {base_path}: {ink_ratio:.4f} ({mask_binary.sum()} pixels)\")\n            return mask_binary\n    return None\n\n# Simple Training Dataset (no region masks)\nclass TrainingDataset(Dataset):\n    def __init__(self, volumes, masks, fragment_ids=None, augment=True, max_samples=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.augment = augment\n        self.fragment_ids = fragment_ids if fragment_ids else [0] * len(volumes)\n        self.max_samples = max_samples\n        self.samples = self._prepare_samples()\n        print(f\"Created training dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        \n        for vol_idx, volume in enumerate(self.volumes):\n            start_idx = config.num_input_slices // 2\n            end_idx = volume.shape[0] - config.num_input_slices // 2\n            \n            max_samples_per_volume = 500  # Increased for better coverage\n            for _ in range(max_samples_per_volume):\n                center_slice = np.random.randint(start_idx, end_idx)\n                y = np.random.randint(0, volume.shape[1])\n                x = np.random.randint(0, volume.shape[2])\n                \n                samples.append({\n                    'volume_idx': vol_idx,\n                    'center_slice': center_slice,\n                    'y': y,\n                    'x': x,\n                    'fragment_id': self.fragment_ids[vol_idx]\n                })\n        \n        if self.max_samples and len(samples) > self.max_samples:\n            samples = samples[:self.max_samples]\n        \n        return samples\n    \n    def _extract_patch(self, volume, center_slice, y, x):\n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        half_size = config.img_size // 2\n        y_start = max(0, y - half_size)\n        y_end = min(volume.shape[1], y + half_size)\n        x_start = max(0, x - half_size)\n        x_end = min(volume.shape[2], x + half_size)\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = volume[slice_idx]\n            patch = slice_data[y_start:y_end, x_start:x_end]\n            \n            if patch.shape != (config.img_size, config.img_size):\n                pad_y = config.img_size - patch.shape[0]\n                pad_x = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((0, pad_y), (0, pad_x)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        return volume_patch, y_start, y_end, x_start, x_end\n    \n    def _apply_augmentation(self, volume_slices, mask_patch):\n        volume_slices = volume_slices.copy()\n        mask_patch = mask_patch.copy()\n        \n        # Random horizontal flip\n        if random.random() > 0.5:\n            volume_slices = np.ascontiguousarray(np.flip(volume_slices, axis=2))\n            mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=1))\n        \n        # Random vertical flip\n        if random.random() > 0.5:\n            volume_slices = np.ascontiguousarray(np.flip(volume_slices, axis=1))\n            mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=0))\n        \n        # Random brightness/contrast adjustment\n        for i in range(volume_slices.shape[0]):\n            if random.random() > 0.5:\n                alpha = random.uniform(0.8, 1.2)\n                beta = random.uniform(-0.1, 0.1)\n                volume_slices[i] = np.clip(alpha * volume_slices[i] + beta, 0, 1)\n        \n        return volume_slices, mask_patch\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        vol_idx = sample['volume_idx']\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        fragment_id = sample['fragment_id']\n        \n        volume_patch, y_start, y_end, x_start, x_end = self._extract_patch(\n            self.volumes[vol_idx], center_slice, y, x\n        )\n        \n        mask = self.masks[vol_idx]\n        mask_patch = mask[y_start:y_end, x_start:x_end]\n        \n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y = config.img_size - mask_patch.shape[0]\n            pad_x = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((0, pad_y), (0, pad_x)), mode='constant')\n        \n        if self.augment:\n            volume_patch, mask_patch = self._apply_augmentation(volume_patch, mask_patch)\n        \n        volume_patch = np.ascontiguousarray(volume_patch)\n        mask_patch = np.ascontiguousarray(mask_patch)\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor\n\n# Test Dataset for Full Fragment 1\nclass TestDataset(Dataset):\n    def __init__(self, volume, mask, fragment_id=0, stride_factor=2):\n        self.volume = volume\n        self.mask = mask\n        self.fragment_id = fragment_id\n        self.stride_factor = stride_factor\n        self.samples = self._prepare_samples()\n        print(f\"Created test dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        start_idx = config.num_input_slices // 2\n        end_idx = self.volume.shape[0] - config.num_input_slices // 2\n        \n        stride = (config.img_size // 2) * self.stride_factor\n        for center_slice in range(start_idx, end_idx, 2):\n            for y in range(0, self.volume.shape[1], stride):\n                for x in range(0, self.volume.shape[2], stride):\n                    samples.append({\n                        'center_slice': center_slice,\n                        'y': y,\n                        'x': x\n                    })\n        return samples\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        \n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = self.volume[slice_idx]\n            patch = slice_data[y:y+config.img_size, x:x+config.img_size]\n            \n            if patch.shape != (config.img_size, config.img_size):\n                pad_y = config.img_size - patch.shape[0]\n                pad_x = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((0, pad_y), (0, pad_x)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        \n        mask_patch = self.mask[y:y+config.img_size, x:x+config.img_size]\n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y = config.img_size - mask_patch.shape[0]\n            pad_x = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((0, pad_y), (0, pad_x)), mode='constant')\n        \n        volume_patch = np.ascontiguousarray(volume_patch)\n        mask_patch = np.ascontiguousarray(mask_patch)\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor, y, x\n\n# 3D UNet Model\nclass VolumetricUNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = smp.Unet(\n            encoder_name=\"resnet34\",\n            encoder_weights=\"imagenet\",\n            in_channels=config.num_input_slices,\n            classes=1,\n            encoder_depth=4,\n            decoder_channels=[128, 64, 32, 16],\n            activation=None\n        )\n        \n    def forward(self, x):\n        return self.model(x.contiguous())\n\n# Loss function\nclass CustomCombinedLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        \n    def forward(self, pred, target):\n        pred_sigmoid = torch.sigmoid(pred)\n        intersection = (pred_sigmoid * target).sum()\n        dice = 1 - (2. * intersection + 1e-6) / (pred_sigmoid.sum() + target.sum() + 1e-6)\n        bce = self.bce_loss(pred, target)\n        return 0.7 * dice + 0.3 * bce\n\ndef calculate_all_metrics(pred, target, threshold=0.5):\n    \"\"\"Calculate all metrics including TP, TN, FP, FN\"\"\"\n    pred_binary = (pred > threshold).astype(np.float32)\n    pred_flat = pred_binary.flatten().astype(int)\n    target_flat = target.flatten().astype(int)\n    \n    tn, fp, fn, tp = confusion_matrix(target_flat, pred_flat, labels=[0, 1]).ravel()\n    \n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n    \n    beta_squared = 0.5 ** 2\n    f05 = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (beta_squared * precision + recall) > 0 else 0\n    \n    iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n    \n    return {\n        'precision': float(precision),\n        'recall': float(recall),\n        'dice': float(dice),\n        'f05': float(f05),\n        'iou': float(iou),\n        'tp': int(tp),\n        'tn': int(tn),\n        'fp': int(fp),\n        'fn': int(fn)\n    }\n\ndef dice_score(pred, target, smooth=1e-6):\n    pred = pred.flatten()\n    target = target.flatten()\n    intersection = (pred * target).sum()\n    return float((2. * intersection + smooth) / (pred.sum() + target.sum() + smooth))\n\ndef save_test_visualization(volume, mask, predictions, metrics, threshold_results, fragment_name=\"fragment1\"):\n    \"\"\"Save comprehensive test visualization\"\"\"\n    os.makedirs('test_results', exist_ok=True)\n    \n    middle_slice_idx = volume.shape[0] // 2\n    middle_slice = volume[middle_slice_idx]\n    \n    pred_binary = (predictions > metrics['threshold']).astype(np.float32)\n    \n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n    \n    # Row 1\n    axes[0, 0].imshow(middle_slice, cmap='gray')\n    axes[0, 0].set_title(f'Input Image\\n(Middle Slice)')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(mask, cmap='gray')\n    axes[0, 1].set_title(f'Ground Truth\\nInk: {mask.sum()/mask.size:.2%}')\n    axes[0, 1].axis('off')\n    \n    axes[0, 2].imshow(predictions, cmap='jet', vmin=0, vmax=1)\n    axes[0, 2].set_title(f'Predictions\\n(Probabilities)')\n    axes[0, 2].axis('off')\n    \n    # Row 2\n    pred_overlay = np.stack([\n        middle_slice * 0.5 + pred_binary * 0.5,\n        middle_slice,\n        middle_slice\n    ], axis=-1)\n    axes[1, 0].imshow(pred_overlay)\n    axes[1, 0].set_title('Input + Binary Predictions\\n(Red = Predicted Ink)')\n    axes[1, 0].axis('off')\n    \n    error_map = np.zeros((mask.shape[0], mask.shape[1], 3))\n    false_positives = (pred_binary == 1) & (mask == 0)\n    false_negatives = (pred_binary == 0) & (mask == 1)\n    \n    error_map[:, :, 0] = false_positives * 0.8  # Red = False Positives\n    error_map[:, :, 2] = false_negatives * 0.8  # Blue = False Negatives\n    \n    axes[1, 1].imshow(error_map)\n    axes[1, 1].set_title('Error Map\\n(Red=FP, Blue=FN)')\n    axes[1, 1].axis('off')\n    \n    # Threshold performance plot\n    thresh_vals = [r['threshold'] for r in threshold_results]\n    dice_vals = [r['dice'] for r in threshold_results]\n    \n    axes[1, 2].plot(thresh_vals, dice_vals, 'b-', linewidth=2, marker='o')\n    axes[1, 2].axvline(x=metrics['threshold'], color='r', linestyle='--', label=f\"Best: {metrics['threshold']:.2f}\")\n    axes[1, 2].set_title('Threshold Performance on Fragment 1')\n    axes[1, 2].set_xlabel('Threshold')\n    axes[1, 2].set_ylabel('Dice Score')\n    axes[1, 2].grid(True, alpha=0.3)\n    axes[1, 2].legend()\n    \n    # Add metrics text\n    metrics_text = f\"TEST RESULTS - Full Fragment 1\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Dice Score: {metrics['dice']:.4f}\\n\"\n    metrics_text += f\"F0.5 Score: {metrics['f05']:.4f}\\n\"\n    metrics_text += f\"IOU Score: {metrics['iou']:.4f}\\n\"\n    metrics_text += f\"Precision: {metrics['precision']:.4f}\\n\"\n    metrics_text += f\"Recall: {metrics['recall']:.4f}\\n\"\n    metrics_text += f\"Best Threshold: {metrics['threshold']:.3f}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"True Positives: {metrics['tp']:,}\\n\"\n    metrics_text += f\"True Negatives: {metrics['tn']:,}\\n\"\n    metrics_text += f\"False Positives: {metrics['fp']:,}\\n\"\n    metrics_text += f\"False Negatives: {metrics['fn']:,}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Training Data: Fragments 2 & 3 (Full)\\n\"\n    metrics_text += f\"Test Data: Fragment 1 (Full)\\n\"\n    metrics_text += f\"Image Size: {config.img_size}x{config.img_size}\\n\"\n    metrics_text += f\"Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\"\n    \n    plt.figtext(0.02, 0.02, metrics_text, fontsize=9, fontfamily='monospace',\n                bbox=dict(boxstyle=\"round,pad=0.5\", facecolor=\"lightgray\", alpha=0.8))\n    \n    plt.suptitle(f'TEST RESULTS - {fragment_name} - Full Fragment Evaluation\\n'\n                 f'Train on Fragments 2 & 3 | Test on Fragment 1', \n                 fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(f'test_results/{fragment_name}_test_results.png', \n                dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Save predictions\n    np.save(f'test_results/{fragment_name}_predictions.npy', predictions)\n    np.save(f'test_results/{fragment_name}_predictions_binary.npy', pred_binary)\n    \n    # Save threshold results\n    threshold_df = pd.DataFrame(threshold_results)\n    threshold_df.to_csv(f'test_results/{fragment_name}_threshold_results.csv', index=False)\n    \n    print(f\"Test visualization saved to: test_results/{fragment_name}_test_results.png\")\n    print(f\"Test predictions saved to: test_results/{fragment_name}_predictions.npy\")\n    print(f\"Threshold results saved to: test_results/{fragment_name}_threshold_results.csv\")\n\ndef save_history(history, filename='training_history.json'):\n    \"\"\"Save training history to JSON file\"\"\"\n    os.makedirs('history', exist_ok=True)\n    \n    def convert_numpy_types(obj):\n        if isinstance(obj, np.integer):\n            return int(obj)\n        elif isinstance(obj, np.floating):\n            return float(obj)\n        elif isinstance(obj, np.ndarray):\n            return obj.tolist()\n        elif isinstance(obj, dict):\n            return {key: convert_numpy_types(value) for key, value in obj.items()}\n        elif isinstance(obj, list):\n            return [convert_numpy_types(item) for item in obj]\n        else:\n            return obj\n    \n    history_serializable = convert_numpy_types(history)\n    \n    with open(os.path.join('history', filename), 'w') as f:\n        json.dump(history_serializable, f, indent=2)\n    \n    print(f\"Training history saved to {filename}\")\n\ndef save_metrics_summary(train_metrics, test_metrics, filename='metrics_summary.csv'):\n    \"\"\"Save metrics to CSV file\"\"\"\n    os.makedirs('metrics', exist_ok=True)\n    \n    summary_data = {\n        'Phase': ['Training (Fragments 2 & 3)', 'Test (Fragment 1)'],\n        'Dice_Score': [float(train_metrics['dice']), float(test_metrics['dice'])],\n        'F0.5_Score': [float(train_metrics['f05']), float(test_metrics['f05'])],\n        'IOU_Score': [float(train_metrics['iou']), float(test_metrics['iou'])],\n        'Precision': [float(train_metrics['precision']), float(test_metrics['precision'])],\n        'Recall': [float(train_metrics['recall']), float(test_metrics['recall'])],\n        'Optimal_Threshold': [float(train_metrics['threshold']), float(test_metrics['threshold'])],\n        'TP': [int(train_metrics['tp']), int(test_metrics['tp'])],\n        'TN': [int(train_metrics['tn']), int(test_metrics['tn'])],\n        'FP': [int(train_metrics['fp']), int(test_metrics['fp'])],\n        'FN': [int(train_metrics['fn']), int(test_metrics['fn'])]\n    }\n    \n    df = pd.DataFrame(summary_data)\n    df.to_csv(os.path.join('metrics', filename), index=False)\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"METRICS SUMMARY - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    print(df.to_string())\n    print(f\"\\nMetrics summary saved to {filename}\")\n\ndef plot_training_history(history):\n    \"\"\"Plot training history\"\"\"\n    os.makedirs('training_plots', exist_ok=True)\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Loss plot\n    axes[0, 0].plot(epochs, history['train_loss'], 'b-', linewidth=2)\n    axes[0, 0].set_title('Training Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Dice score plot\n    axes[0, 1].plot(epochs, history['train_dice'], 'g-', linewidth=2)\n    axes[0, 1].set_title('Training Dice Score')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Dice Score')\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # F0.5 score plot\n    axes[1, 0].plot(epochs, history['train_f05'], 'r-', linewidth=2)\n    axes[1, 0].set_title('Training F0.5 Score')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('F0.5 Score')\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # IOU score plot\n    axes[1, 1].plot(epochs, history['train_iou'], 'm-', linewidth=2)\n    axes[1, 1].set_title('Training IOU Score')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('IOU Score')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.suptitle('Training History\\nTrain on Fragments 2 & 3', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/training_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Precision and Recall plot\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n    \n    ax1.plot(epochs, history['train_precision'], 'b-', label='Precision', linewidth=2)\n    ax1.plot(epochs, history['train_recall'], 'r-', label='Recall', linewidth=2)\n    ax1.set_title('Training Precision & Recall')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Score')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    ax2.plot(epochs, history['train_precision'], 'b-', label='Precision', linewidth=2)\n    ax2.plot(epochs, history['train_recall'], 'r-', label='Recall', linewidth=2)\n    ax2.set_title('Training Precision & Recall (Zoomed)')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Score')\n    ax2.set_ylim([0, 1])\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    plt.suptitle('Precision and Recall - Training on Fragments 2 & 3', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/precision_recall_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n\ndef main():\n    print(\"=\"*80)\n    print(\"3D UNet with ResNet34 - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    print(\"TRAINING STRATEGY:\")\n    print(f\"  Training Data: Fragments 2 and 3 (100%)\")\n    print(f\"  Test Data: Fragment 1 (100%)\")\n    print(f\"  Validation: DISABLED (pure training on fragments 2 & 3)\")\n    print(\"=\"*80)\n    \n    # Set random seeds\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Load data\n    print(\"\\nLoading data...\")\n    \n    # Load fragment 1 (test data)\n    fragment1_volume = load_volume_data(config.fragment1_path, config.slices)\n    fragment1_mask = load_mask_data(config.fragment1_path)\n    \n    if fragment1_volume is None or fragment1_mask is None:\n        raise ValueError(f\"Failed to load fragment 1 data from {config.fragment1_path}\")\n    \n    # Load training fragments 2 and 3\n    train_volumes = []\n    train_masks = []\n    fragment_ids = []\n    \n    for i, path in enumerate(config.train_paths):\n        volume = load_volume_data(path, config.slices)\n        mask = load_mask_data(path)\n        \n        if volume is not None and mask is not None:\n            train_volumes.append(volume)\n            train_masks.append(mask)\n            fragment_ids.append(i + 1)\n            print(f\"Fragment {i+2} loaded - shape: {volume.shape}\")\n    \n    print(f\"\\nLoaded {len(train_volumes)} training fragments (2 and 3)\")\n    print(f\"Fragment 1 (test) shape: {fragment1_volume.shape}\")\n    \n    # Prepare datasets\n    print(f\"\\nPreparing datasets...\")\n    \n    # Training dataset (fragments 2 & 3 only)\n    train_dataset = TrainingDataset(\n        volumes=train_volumes,\n        masks=train_masks,\n        fragment_ids=fragment_ids,\n        augment=True,\n        max_samples=config.max_train_samples\n    )\n    \n    # Test dataset (full fragment 1)\n    test_dataset = TestDataset(\n        volume=fragment1_volume,\n        mask=fragment1_mask,\n        fragment_id=0,\n        stride_factor=2\n    )\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=1,\n        shuffle=False,\n        num_workers=0\n    )\n    \n    print(f\"\\nData Loaders Created:\")\n    print(f\"  Training samples: {len(train_dataset)} (Fragments 2 & 3)\")\n    print(f\"  Test samples: {len(test_dataset)} (Fragment 1)\")\n    \n    # Initialize model\n    print(f\"\\nInitializing model...\")\n    model = VolumetricUNet().to(config.device)\n    print(f\"Model initialized on {config.device}\")\n    \n    # Initialize optimizer and loss\n    optimizer = optim.Adam(model.parameters(), lr=config.lr, weight_decay=1e-5)\n    criterion = CustomCombinedLoss()\n    \n    # Training history\n    history = {\n        'train_loss': [],\n        'train_dice': [],\n        'train_f05': [],\n        'train_iou': [],\n        'train_precision': [],\n        'train_recall': []\n    }\n    \n    print(f\"\\nStarting training for {config.epochs} epochs...\")\n    print(\"=\"*80)\n    \n    for epoch in range(config.epochs):\n        # Training phase\n        model.train()\n        epoch_train_loss = 0\n        train_outputs = []\n        train_targets = []\n        \n        train_iterator = tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.epochs}')\n        for images, masks in train_iterator:\n            images, masks = images.to(config.device), masks.to(config.device)\n            images = images.contiguous()\n            masks = masks.contiguous()\n            \n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            loss.backward()\n            \n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            epoch_train_loss += loss.item()\n            \n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_outputs.append(preds.cpu().numpy())\n                train_targets.append(masks.cpu().numpy())\n            \n            train_iterator.set_postfix(loss=loss.item())\n            \n            if torch.cuda.is_available() and len(train_outputs) % 100 == 0:\n                torch.cuda.empty_cache()\n        \n        # Calculate training metrics\n        train_outputs = np.concatenate(train_outputs)\n        train_targets = np.concatenate(train_targets)\n        \n        best_train_threshold = 0.5\n        best_train_dice = 0\n        for thresh in np.arange(0.1, 0.9, 0.05):\n            preds = (train_outputs > thresh).astype(np.float32)\n            dice = dice_score(preds, train_targets)\n            if dice > best_train_dice:\n                best_train_dice = dice\n                best_train_threshold = thresh\n        \n        train_preds = (train_outputs > best_train_threshold).astype(np.float32)\n        train_metrics = calculate_all_metrics(train_preds, train_targets)\n        train_metrics['threshold'] = float(best_train_threshold)\n        \n        # Store training metrics\n        history['train_loss'].append(float(epoch_train_loss / len(train_loader)))\n        history['train_dice'].append(float(train_metrics['dice']))\n        history['train_f05'].append(float(train_metrics['f05']))\n        history['train_iou'].append(float(train_metrics['iou']))\n        history['train_precision'].append(float(train_metrics['precision']))\n        history['train_recall'].append(float(train_metrics['recall']))\n        \n        # Print progress\n        print(f\"\\nEpoch {epoch+1}/{config.epochs}:\")\n        print(f\"  Training - Loss: {epoch_train_loss/len(train_loader):.4f}, \"\n              f\"Dice: {train_metrics['dice']:.4f}, F0.5: {train_metrics['f05']:.4f}\")\n        \n        # Clear cache\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    # Training complete - save model\n    torch.save(model.state_dict(), 'trained_model_fragments_2_3.pth')\n    print(f\"\\nModel saved to: trained_model_fragments_2_3.pth\")\n    \n    # Clear cache before testing\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Test phase on full fragment 1\n    print(\"\\n\" + \"=\"*80)\n    print(\"TESTING PHASE - Full Fragment 1\")\n    print(\"=\"*80)\n    \n    model.eval()\n    \n    pred_accumulator = np.zeros(fragment1_mask.shape, dtype=np.float32)\n    count_accumulator = np.zeros(fragment1_mask.shape, dtype=np.float32)\n    \n    with torch.no_grad():\n        test_iterator = tqdm(test_loader, desc=\"Testing on fragment 1\")\n        for images, masks, y_coords, x_coords in test_iterator:\n            images = images.to(config.device).contiguous()\n            outputs = torch.sigmoid(model(images))\n            \n            for i in range(outputs.shape[0]):\n                y = y_coords[i].item()\n                x = x_coords[i].item()\n                \n                patch = outputs[i, 0].cpu().numpy()\n                h, w = patch.shape\n                \n                pred_accumulator[y:y+h, x:x+w] += patch\n                count_accumulator[y:y+h, x:x+w] += 1\n            \n            if torch.cuda.is_available() and test_iterator.n % 10 == 0:\n                torch.cuda.empty_cache()\n    \n    # Average overlapping predictions\n    final_predictions = pred_accumulator / np.maximum(count_accumulator, 1)\n    \n    # Calculate test metrics on full fragment 1 with MULTIPLE THRESHOLDS\n    print(\"\\nEvaluating thresholds on Fragment 1...\")\n    \n    # Try a wide range of thresholds\n    thresholds_to_try = [0.1, 0.15, 0.2, 0.25, 0.3, 0.35, 0.4, 0.45, 0.5, 0.55, 0.6, 0.65, 0.7, 0.75, 0.8]\n    \n    best_test_dice = 0\n    best_test_threshold = 0.5\n    threshold_results = []\n    \n    for thresh in thresholds_to_try:\n        preds = (final_predictions > thresh).astype(np.float32)\n        dice = dice_score(preds, fragment1_mask)\n        threshold_results.append({'threshold': thresh, 'dice': dice})\n        \n        if dice > best_test_dice:\n            best_test_dice = dice\n            best_test_threshold = thresh\n    \n    print(\"\\nThreshold performance on Fragment 1:\")\n    for result in threshold_results:\n        print(f\"  Threshold {result['threshold']:.2f}: Dice = {result['dice']:.4f}\")\n    \n    # Calculate all test metrics with the best threshold found on Fragment 1\n    test_preds = (final_predictions > best_test_threshold).astype(np.float32)\n    test_metrics = calculate_all_metrics(test_preds, fragment1_mask)\n    test_metrics['threshold'] = float(best_test_threshold)\n    \n    print(f\"\\nBest threshold for Fragment 1: {best_test_threshold:.2f} (Dice: {best_test_dice:.4f})\")\n    \n    # Save test visualization\n    print(\"\\nSaving test visualization and metrics...\")\n    save_test_visualization(\n        fragment1_volume, \n        fragment1_mask, \n        final_predictions, \n        test_metrics,\n        threshold_results,\n        \"fragment1\"\n    )\n    \n    # Save training history\n    if config.save_history:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_history(history, filename=f'training_history_{timestamp}.json')\n        plot_training_history(history)\n    \n    # Final training metrics\n    final_train_metrics = {\n        'dice': float(history['train_dice'][-1]),\n        'f05': float(history['train_f05'][-1]),\n        'iou': float(history['train_iou'][-1]),\n        'precision': float(history['train_precision'][-1]),\n        'recall': float(history['train_recall'][-1]),\n        'threshold': float(best_train_threshold),\n        'tp': int(train_metrics['tp']),\n        'tn': int(train_metrics['tn']),\n        'fp': int(train_metrics['fp']),\n        'fn': int(train_metrics['fn'])\n    }\n    \n    # Save metrics summary\n    if config.save_metrics:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_metrics_summary(\n            final_train_metrics,\n            test_metrics,\n            filename=f'metrics_summary_{timestamp}.csv'\n        )\n    \n    # Final summary\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL RESULTS SUMMARY - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    \n    print(f\"\\nCONFIGURATION:\")\n    print(f\"  Image Size: {config.img_size}x{config.img_size}\")\n    print(f\"  Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\")\n    \n    print(f\"\\nTRAINING DETAILS:\")\n    print(f\"  Training Data: Fragments 2 & 3 (100%)\")\n    print(f\"  Test Data: Fragment 1 (100%)\")\n    print(f\"  Validation: DISABLED\")\n    \n    print(f\"\\nTRAINING PERFORMANCE (Final Epoch):\")\n    print(f\"  Dice Score: {final_train_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {final_train_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {final_train_metrics['iou']:.4f}\")\n    print(f\"  Precision: {final_train_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_train_metrics['recall']:.4f}\")\n    \n    print(f\"\\nTEST RESULTS (Fragment 1) - WITH OPTIMAL THRESHOLD SEARCH:\")\n    print(f\"  Best Threshold: {test_metrics['threshold']:.2f}\")\n    print(f\"  Dice Score: {test_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {test_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {test_metrics['iou']:.4f}\")\n    print(f\"  Precision: {test_metrics['precision']:.4f}\")\n    print(f\"  Recall: {test_metrics['recall']:.4f}\")\n    print(f\"  TP: {test_metrics['tp']:,} | TN: {test_metrics['tn']:,} | FP: {test_metrics['fp']:,} | FN: {test_metrics['fn']:,}\")\n    \n    print(f\"\\nTHRESHOLD PERFORMANCE SUMMARY:\")\n    for result in threshold_results:\n        if result['threshold'] in [0.1, 0.2, 0.3, 0.4, 0.45, 0.5, 0.6, 0.7, 0.8]:\n            print(f\"  Threshold {result['threshold']:.2f}: Dice = {result['dice']:.4f}\")\n    \n    print(f\"\\nDATA SAVED:\")\n    print(f\"  Test results: test_results/\")\n    print(f\"  Training plots: training_plots/\")\n    print(f\"  Threshold results: test_results/fragment1_threshold_results.csv\")\n    \n    if config.save_history:\n        print(f\"  Training history: history/\")\n    if config.save_metrics:\n        print(f\"  Metrics summary: metrics/\")\n    \n    print(f\"\\nMODEL SAVED:\")\n    print(f\"  trained_model_fragments_2_3.pth\")\n    \n    print(f\"\\nKEY FEATURES:\")\n    print(f\"  1. Trained ONLY on Fragments 2 & 3 (no validation)\")\n    print(f\"  2. Tested on FULL Fragment 1\")\n    print(f\"  3. Optimal threshold SEARCHED on Fragment 1 (not using training threshold)\")\n    print(f\"  4. Image Size: {config.img_size}x{config.img_size}\")\n    print(f\"  5. All slices 12-30 included\")\n    print(f\"  6. Pure cross-fragment generalization test\")\n\n# Run the main function\nif __name__ == \"__main__\":\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        gpu_memory = torch.cuda.get_device_properties(0).total_memory / 1e9\n        print(f\"Available GPU memory: {gpu_memory:.1f} GB\")\n    \n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# #3D Unet&Resnet34 with test on full fragment 1\n# Install required packages\n!pip install segmentation-models-pytorch==0.2.0\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nfrom skimage.filters import threshold_otsu\nimport warnings\nimport random\nimport math\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport pandas as pd\nimport json\nfrom datetime import datetime\n\nwarnings.filterwarnings('ignore')\n#/kaggle/input/vesuvius-challenge-ink-detection/test\n# Configuration\nclass Config:\n    # Data paths\n    fragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'  # Test data\n    train_paths = [\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/2',  # Training data\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/3'   # Training data\n    ]\n    \n    # Model parameters\n    img_size = 352\n    slices = list(range(12, 31))  # Slice 12-30 (19 slices)\n    batch_size = 8  # Increased batch size since no validation\n    epochs = 20\n    lr = 3e-4\n    \n    # 3D volumetric setup\n    num_input_slices = 5\n    \n    # Memory optimization\n    max_train_samples = 2500  # Increased for training\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 5\n    \n    # Results saving\n    save_history = True\n    save_metrics = True\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Training strategy - NO VALIDATION, just train on fragments 2 & 3, test on fragment 1\n    use_validation = False  # Set to False to disable validation during training\n\nconfig = Config()\n\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data with enhanced preprocessing\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_filtered = cv2.bilateralFilter(img_resized, 5, 50, 50)\n                clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n                img_enhanced = clahe.apply(img_filtered)\n                img_normalized = img_enhanced.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        print(f\"Loaded volume from {base_path} with {volume.shape[0]} slices\")\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio for {base_path}: {ink_ratio:.4f} ({mask_binary.sum()} pixels)\")\n            return mask_binary\n    return None\n\n# Simple Training Dataset (no region masks)\nclass TrainingDataset(Dataset):\n    def __init__(self, volumes, masks, fragment_ids=None, augment=True, max_samples=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.augment = augment\n        self.fragment_ids = fragment_ids if fragment_ids else [0] * len(volumes)\n        self.max_samples = max_samples\n        self.samples = self._prepare_samples()\n        print(f\"Created training dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        \n        for vol_idx, volume in enumerate(self.volumes):\n            start_idx = config.num_input_slices // 2\n            end_idx = volume.shape[0] - config.num_input_slices // 2\n            \n            max_samples_per_volume = 500  # Increased for better coverage\n            for _ in range(max_samples_per_volume):\n                center_slice = np.random.randint(start_idx, end_idx)\n                y = np.random.randint(0, volume.shape[1])\n                x = np.random.randint(0, volume.shape[2])\n                \n                samples.append({\n                    'volume_idx': vol_idx,\n                    'center_slice': center_slice,\n                    'y': y,\n                    'x': x,\n                    'fragment_id': self.fragment_ids[vol_idx]\n                })\n        \n        if self.max_samples and len(samples) > self.max_samples:\n            samples = samples[:self.max_samples]\n        \n        return samples\n    \n    def _extract_patch(self, volume, center_slice, y, x):\n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        half_size = config.img_size // 2\n        y_start = max(0, y - half_size)\n        y_end = min(volume.shape[1], y + half_size)\n        x_start = max(0, x - half_size)\n        x_end = min(volume.shape[2], x + half_size)\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = volume[slice_idx]\n            patch = slice_data[y_start:y_end, x_start:x_end]\n            \n            if patch.shape != (config.img_size, config.img_size):\n                pad_y = config.img_size - patch.shape[0]\n                pad_x = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((0, pad_y), (0, pad_x)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        return volume_patch, y_start, y_end, x_start, x_end\n    \n    def _apply_augmentation(self, volume_slices, mask_patch):\n        volume_slices = volume_slices.copy()\n        mask_patch = mask_patch.copy()\n        \n        # Random horizontal flip\n        if random.random() > 0.5:\n            volume_slices = np.ascontiguousarray(np.flip(volume_slices, axis=2))\n            mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=1))\n        \n        # Random vertical flip\n        if random.random() > 0.5:\n            volume_slices = np.ascontiguousarray(np.flip(volume_slices, axis=1))\n            mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=0))\n        \n        # Random brightness/contrast adjustment\n        for i in range(volume_slices.shape[0]):\n            if random.random() > 0.5:\n                alpha = random.uniform(0.8, 1.2)\n                beta = random.uniform(-0.1, 0.1)\n                volume_slices[i] = np.clip(alpha * volume_slices[i] + beta, 0, 1)\n        \n        return volume_slices, mask_patch\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        vol_idx = sample['volume_idx']\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        fragment_id = sample['fragment_id']\n        \n        volume_patch, y_start, y_end, x_start, x_end = self._extract_patch(\n            self.volumes[vol_idx], center_slice, y, x\n        )\n        \n        mask = self.masks[vol_idx]\n        mask_patch = mask[y_start:y_end, x_start:x_end]\n        \n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y = config.img_size - mask_patch.shape[0]\n            pad_x = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((0, pad_y), (0, pad_x)), mode='constant')\n        \n        if self.augment:\n            volume_patch, mask_patch = self._apply_augmentation(volume_patch, mask_patch)\n        \n        volume_patch = np.ascontiguousarray(volume_patch)\n        mask_patch = np.ascontiguousarray(mask_patch)\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor\n\n# Test Dataset for Full Fragment 1\nclass TestDataset(Dataset):\n    def __init__(self, volume, mask, fragment_id=0, stride_factor=2):\n        self.volume = volume\n        self.mask = mask\n        self.fragment_id = fragment_id\n        self.stride_factor = stride_factor\n        self.samples = self._prepare_samples()\n        print(f\"Created test dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        start_idx = config.num_input_slices // 2\n        end_idx = self.volume.shape[0] - config.num_input_slices // 2\n        \n        stride = (config.img_size // 2) * self.stride_factor\n        for center_slice in range(start_idx, end_idx, 2):\n            for y in range(0, self.volume.shape[1], stride):\n                for x in range(0, self.volume.shape[2], stride):\n                    samples.append({\n                        'center_slice': center_slice,\n                        'y': y,\n                        'x': x\n                    })\n        return samples\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        \n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = self.volume[slice_idx]\n            patch = slice_data[y:y+config.img_size, x:x+config.img_size]\n            \n            if patch.shape != (config.img_size, config.img_size):\n                pad_y = config.img_size - patch.shape[0]\n                pad_x = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((0, pad_y), (0, pad_x)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        \n        mask_patch = self.mask[y:y+config.img_size, x:x+config.img_size]\n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y = config.img_size - mask_patch.shape[0]\n            pad_x = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((0, pad_y), (0, pad_x)), mode='constant')\n        \n        volume_patch = np.ascontiguousarray(volume_patch)\n        mask_patch = np.ascontiguousarray(mask_patch)\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor, y, x\n\n# 3D UNet Model\nclass VolumetricUNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = smp.Unet(\n            encoder_name=\"resnet34\",\n            encoder_weights=\"imagenet\",\n            in_channels=config.num_input_slices,\n            classes=1,\n            encoder_depth=4,\n            decoder_channels=[128, 64, 32, 16],\n            activation=None\n        )\n        \n    def forward(self, x):\n        return self.model(x.contiguous())\n\n# Loss function\nclass CustomCombinedLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        \n    def forward(self, pred, target):\n        pred_sigmoid = torch.sigmoid(pred)\n        intersection = (pred_sigmoid * target).sum()\n        dice = 1 - (2. * intersection + 1e-6) / (pred_sigmoid.sum() + target.sum() + 1e-6)\n        bce = self.bce_loss(pred, target)\n        return 0.7 * dice + 0.3 * bce\n\ndef calculate_all_metrics(pred, target, threshold=0.5):\n    \"\"\"Calculate all metrics including TP, TN, FP, FN\"\"\"\n    pred_binary = (pred > threshold).astype(np.float32)\n    pred_flat = pred_binary.flatten().astype(int)\n    target_flat = target.flatten().astype(int)\n    \n    tn, fp, fn, tp = confusion_matrix(target_flat, pred_flat, labels=[0, 1]).ravel()\n    \n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n    \n    beta_squared = 0.5 ** 2\n    f05 = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (beta_squared * precision + recall) > 0 else 0\n    \n    iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n    \n    return {\n        'precision': float(precision),\n        'recall': float(recall),\n        'dice': float(dice),\n        'f05': float(f05),\n        'iou': float(iou),\n        'tp': int(tp),\n        'tn': int(tn),\n        'fp': int(fp),\n        'fn': int(fn)\n    }\n\ndef dice_score(pred, target, smooth=1e-6):\n    pred = pred.flatten()\n    target = target.flatten()\n    intersection = (pred * target).sum()\n    return float((2. * intersection + smooth) / (pred.sum() + target.sum() + smooth))\n\ndef save_test_visualization(volume, mask, predictions, metrics, fragment_name=\"fragment1\"):\n    \"\"\"Save comprehensive test visualization\"\"\"\n    os.makedirs('test_results', exist_ok=True)\n    \n    middle_slice_idx = volume.shape[0] // 2\n    middle_slice = volume[middle_slice_idx]\n    \n    pred_binary = (predictions > metrics['threshold']).astype(np.float32)\n    \n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n    \n    # Row 1\n    axes[0, 0].imshow(middle_slice, cmap='gray')\n    axes[0, 0].set_title(f'Input Image\\n(Middle Slice)')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(mask, cmap='gray')\n    axes[0, 1].set_title(f'Ground Truth\\nInk: {mask.sum()/mask.size:.2%}')\n    axes[0, 1].axis('off')\n    \n    axes[0, 2].imshow(predictions, cmap='jet', vmin=0, vmax=1)\n    axes[0, 2].set_title(f'Predictions\\n(Probabilities)')\n    axes[0, 2].axis('off')\n    \n    # Row 2\n    pred_overlay = np.stack([\n        middle_slice * 0.5 + pred_binary * 0.5,\n        middle_slice,\n        middle_slice\n    ], axis=-1)\n    axes[1, 0].imshow(pred_overlay)\n    axes[1, 0].set_title('Input + Binary Predictions\\n(Red = Predicted Ink)')\n    axes[1, 0].axis('off')\n    \n    error_map = np.zeros((mask.shape[0], mask.shape[1], 3))\n    false_positives = (pred_binary == 1) & (mask == 0)\n    false_negatives = (pred_binary == 0) & (mask == 1)\n    \n    error_map[:, :, 0] = false_positives * 0.8  # Red = False Positives\n    error_map[:, :, 2] = false_negatives * 0.8  # Blue = False Negatives\n    \n    axes[1, 1].imshow(error_map)\n    axes[1, 1].set_title('Error Map\\n(Red=FP, Blue=FN)')\n    axes[1, 1].axis('off')\n    \n    # Confusion matrix visualization\n    confusion_img = np.zeros((100, 100, 3))\n    confusion_img[:50, :50, :] = 0.2  # TN area (gray)\n    confusion_img[:50, 50:, 0] = 0.8  # FP area (red)\n    confusion_img[50:, :50, 2] = 0.8  # FN area (blue)\n    confusion_img[50:, 50:, 1] = 0.8  # TP area (green)\n    \n    axes[1, 2].imshow(confusion_img)\n    axes[1, 2].set_title('Confusion Matrix\\n(Green=TP, Red=FP, Blue=FN, Gray=TN)')\n    axes[1, 2].axis('off')\n    \n    # Add metrics text\n    metrics_text = f\"TEST RESULTS - Full Fragment 1\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Dice Score: {metrics['dice']:.4f}\\n\"\n    metrics_text += f\"F0.5 Score: {metrics['f05']:.4f}\\n\"\n    metrics_text += f\"IOU Score: {metrics['iou']:.4f}\\n\"\n    metrics_text += f\"Precision: {metrics['precision']:.4f}\\n\"\n    metrics_text += f\"Recall: {metrics['recall']:.4f}\\n\"\n    metrics_text += f\"Optimal Threshold: {metrics['threshold']:.3f}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"True Positives: {metrics['tp']:,}\\n\"\n    metrics_text += f\"True Negatives: {metrics['tn']:,}\\n\"\n    metrics_text += f\"False Positives: {metrics['fp']:,}\\n\"\n    metrics_text += f\"False Negatives: {metrics['fn']:,}\\n\"\n    metrics_text += \"=\"*60 + \"\\n\"\n    metrics_text += f\"Training Data: Fragments 2 & 3 (Full)\\n\"\n    metrics_text += f\"Test Data: Fragment 1 (Full)\\n\"\n    metrics_text += f\"Image Size: {config.img_size}x{config.img_size}\\n\"\n    metrics_text += f\"Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\"\n    \n    plt.figtext(0.02, 0.02, metrics_text, fontsize=9, fontfamily='monospace',\n                bbox=dict(boxstyle=\"round,pad=0.5\", facecolor=\"lightgray\", alpha=0.8))\n    \n    plt.suptitle(f'TEST RESULTS - {fragment_name} - Full Fragment Evaluation\\n'\n                 f'Train on Fragments 2 & 3 | Test on Fragment 1', \n                 fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(f'test_results/{fragment_name}_test_results.png', \n                dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Save predictions\n    np.save(f'test_results/{fragment_name}_predictions.npy', predictions)\n    np.save(f'test_results/{fragment_name}_predictions_binary.npy', pred_binary)\n    \n    print(f\"Test visualization saved to: test_results/{fragment_name}_test_results.png\")\n    print(f\"Test predictions saved to: test_results/{fragment_name}_predictions.npy\")\n\ndef save_history(history, filename='training_history.json'):\n    \"\"\"Save training history to JSON file\"\"\"\n    os.makedirs('history', exist_ok=True)\n    \n    def convert_numpy_types(obj):\n        if isinstance(obj, np.integer):\n            return int(obj)\n        elif isinstance(obj, np.floating):\n            return float(obj)\n        elif isinstance(obj, np.ndarray):\n            return obj.tolist()\n        elif isinstance(obj, dict):\n            return {key: convert_numpy_types(value) for key, value in obj.items()}\n        elif isinstance(obj, list):\n            return [convert_numpy_types(item) for item in obj]\n        else:\n            return obj\n    \n    history_serializable = convert_numpy_types(history)\n    \n    with open(os.path.join('history', filename), 'w') as f:\n        json.dump(history_serializable, f, indent=2)\n    \n    print(f\"Training history saved to {filename}\")\n\ndef save_metrics_summary(train_metrics, test_metrics, filename='metrics_summary.csv'):\n    \"\"\"Save metrics to CSV file\"\"\"\n    os.makedirs('metrics', exist_ok=True)\n    \n    summary_data = {\n        'Phase': ['Training (Fragments 2 & 3)', 'Test (Fragment 1)'],\n        'Dice_Score': [float(train_metrics['dice']), float(test_metrics['dice'])],\n        'F0.5_Score': [float(train_metrics['f05']), float(test_metrics['f05'])],\n        'IOU_Score': [float(train_metrics['iou']), float(test_metrics['iou'])],\n        'Precision': [float(train_metrics['precision']), float(test_metrics['precision'])],\n        'Recall': [float(train_metrics['recall']), float(test_metrics['recall'])],\n        'Optimal_Threshold': [float(train_metrics['threshold']), float(test_metrics['threshold'])],\n        'TP': [int(train_metrics['tp']), int(test_metrics['tp'])],\n        'TN': [int(train_metrics['tn']), int(test_metrics['tn'])],\n        'FP': [int(train_metrics['fp']), int(test_metrics['fp'])],\n        'FN': [int(train_metrics['fn']), int(test_metrics['fn'])]\n    }\n    \n    df = pd.DataFrame(summary_data)\n    df.to_csv(os.path.join('metrics', filename), index=False)\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"METRICS SUMMARY - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    print(df.to_string())\n    print(f\"\\nMetrics summary saved to {filename}\")\n\ndef plot_training_history(history):\n    \"\"\"Plot training history\"\"\"\n    os.makedirs('training_plots', exist_ok=True)\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Loss plot\n    axes[0, 0].plot(epochs, history['train_loss'], 'b-', linewidth=2)\n    axes[0, 0].set_title('Training Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Dice score plot\n    axes[0, 1].plot(epochs, history['train_dice'], 'g-', linewidth=2)\n    axes[0, 1].set_title('Training Dice Score')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Dice Score')\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # F0.5 score plot\n    axes[1, 0].plot(epochs, history['train_f05'], 'r-', linewidth=2)\n    axes[1, 0].set_title('Training F0.5 Score')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('F0.5 Score')\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # IOU score plot\n    axes[1, 1].plot(epochs, history['train_iou'], 'm-', linewidth=2)\n    axes[1, 1].set_title('Training IOU Score')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('IOU Score')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.suptitle('Training History\\nTrain on Fragments 2 & 3', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/training_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Precision and Recall plot\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n    \n    ax1.plot(epochs, history['train_precision'], 'b-', label='Precision', linewidth=2)\n    ax1.plot(epochs, history['train_recall'], 'r-', label='Recall', linewidth=2)\n    ax1.set_title('Training Precision & Recall')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Score')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    ax2.plot(epochs, history['train_precision'], 'b-', label='Precision', linewidth=2)\n    ax2.plot(epochs, history['train_recall'], 'r-', label='Recall', linewidth=2)\n    ax2.set_title('Training Precision & Recall (Zoomed)')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Score')\n    ax2.set_ylim([0, 1])\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    plt.suptitle('Precision and Recall - Training on Fragments 2 & 3', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig('training_plots/precision_recall_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n\ndef main():\n    print(\"=\"*80)\n    print(\"3D UNet with ResNet34 - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    print(\"TRAINING STRATEGY:\")\n    print(f\"  Training Data: Fragments 2 and 3 (100%)\")\n    print(f\"  Test Data: Fragment 1 (100%)\")\n    print(f\"  Validation: DISABLED (pure training on fragments 2 & 3)\")\n    print(\"=\"*80)\n    \n    # Set random seeds\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Load data\n    print(\"\\nLoading data...\")\n    \n    # Load fragment 1 (test data)\n    fragment1_volume = load_volume_data(config.fragment1_path, config.slices)\n    fragment1_mask = load_mask_data(config.fragment1_path)\n    \n    if fragment1_volume is None or fragment1_mask is None:\n        raise ValueError(f\"Failed to load fragment 1 data from {config.fragment1_path}\")\n    \n    # Load training fragments 2 and 3\n    train_volumes = []\n    train_masks = []\n    fragment_ids = []\n    \n    for i, path in enumerate(config.train_paths):\n        volume = load_volume_data(path, config.slices)\n        mask = load_mask_data(path)\n        \n        if volume is not None and mask is not None:\n            train_volumes.append(volume)\n            train_masks.append(mask)\n            fragment_ids.append(i + 1)\n            print(f\"Fragment {i+2} loaded - shape: {volume.shape}\")\n    \n    print(f\"\\nLoaded {len(train_volumes)} training fragments (2 and 3)\")\n    print(f\"Fragment 1 (test) shape: {fragment1_volume.shape}\")\n    \n    # Prepare datasets\n    print(f\"\\nPreparing datasets...\")\n    \n    # Training dataset (fragments 2 & 3 only)\n    train_dataset = TrainingDataset(\n        volumes=train_volumes,\n        masks=train_masks,\n        fragment_ids=fragment_ids,\n        augment=True,\n        max_samples=config.max_train_samples\n    )\n    \n    # Test dataset (full fragment 1)\n    test_dataset = TestDataset(\n        volume=fragment1_volume,\n        mask=fragment1_mask,\n        fragment_id=0,\n        stride_factor=2\n    )\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=1,\n        shuffle=False,\n        num_workers=0\n    )\n    \n    print(f\"\\nData Loaders Created:\")\n    print(f\"  Training samples: {len(train_dataset)} (Fragments 2 & 3)\")\n    print(f\"  Test samples: {len(test_dataset)} (Fragment 1)\")\n    \n    # Initialize model\n    print(f\"\\nInitializing model...\")\n    model = VolumetricUNet().to(config.device)\n    print(f\"Model initialized on {config.device}\")\n    \n    # Initialize optimizer and loss\n    optimizer = optim.Adam(model.parameters(), lr=config.lr, weight_decay=1e-5)\n    criterion = CustomCombinedLoss()\n    \n    # Training history\n    history = {\n        'train_loss': [],\n        'train_dice': [],\n        'train_f05': [],\n        'train_iou': [],\n        'train_precision': [],\n        'train_recall': []\n    }\n    \n    print(f\"\\nStarting training for {config.epochs} epochs...\")\n    print(\"=\"*80)\n    \n    for epoch in range(config.epochs):\n        # Training phase\n        model.train()\n        epoch_train_loss = 0\n        train_outputs = []\n        train_targets = []\n        \n        train_iterator = tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.epochs}')\n        for images, masks in train_iterator:\n            images, masks = images.to(config.device), masks.to(config.device)\n            images = images.contiguous()\n            masks = masks.contiguous()\n            \n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            loss.backward()\n            \n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            epoch_train_loss += loss.item()\n            \n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_outputs.append(preds.cpu().numpy())\n                train_targets.append(masks.cpu().numpy())\n            \n            train_iterator.set_postfix(loss=loss.item())\n            \n            if torch.cuda.is_available() and len(train_outputs) % 100 == 0:\n                torch.cuda.empty_cache()\n        \n        # Calculate training metrics\n        train_outputs = np.concatenate(train_outputs)\n        train_targets = np.concatenate(train_targets)\n        \n        best_train_threshold = 0.5\n        best_train_dice = 0\n        for thresh in np.arange(0.1, 0.9, 0.05):\n            preds = (train_outputs > thresh).astype(np.float32)\n            dice = dice_score(preds, train_targets)\n            if dice > best_train_dice:\n                best_train_dice = dice\n                best_train_threshold = thresh\n        \n        train_preds = (train_outputs > best_train_threshold).astype(np.float32)\n        train_metrics = calculate_all_metrics(train_preds, train_targets)\n        train_metrics['threshold'] = float(best_train_threshold)\n        \n        # Store training metrics\n        history['train_loss'].append(float(epoch_train_loss / len(train_loader)))\n        history['train_dice'].append(float(train_metrics['dice']))\n        history['train_f05'].append(float(train_metrics['f05']))\n        history['train_iou'].append(float(train_metrics['iou']))\n        history['train_precision'].append(float(train_metrics['precision']))\n        history['train_recall'].append(float(train_metrics['recall']))\n        \n        # Print progress\n        print(f\"\\nEpoch {epoch+1}/{config.epochs}:\")\n        print(f\"  Training - Loss: {epoch_train_loss/len(train_loader):.4f}, \"\n              f\"Dice: {train_metrics['dice']:.4f}, F0.5: {train_metrics['f05']:.4f}\")\n        \n        # Clear cache\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    # Training complete - save model\n    torch.save(model.state_dict(), 'trained_model_fragments_2_3.pth')\n    print(f\"\\nModel saved to: trained_model_fragments_2_3.pth\")\n    \n    # Clear cache before testing\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Test phase on full fragment 1\n    print(\"\\n\" + \"=\"*80)\n    print(\"TESTING PHASE - Full Fragment 1\")\n    print(\"=\"*80)\n    \n    model.eval()\n    \n    pred_accumulator = np.zeros(fragment1_mask.shape, dtype=np.float32)\n    count_accumulator = np.zeros(fragment1_mask.shape, dtype=np.float32)\n    \n    with torch.no_grad():\n        test_iterator = tqdm(test_loader, desc=\"Testing on fragment 1\")\n        for images, masks, y_coords, x_coords in test_iterator:\n            images = images.to(config.device).contiguous()\n            outputs = torch.sigmoid(model(images))\n            \n            for i in range(outputs.shape[0]):\n                y = y_coords[i].item()\n                x = x_coords[i].item()\n                \n                patch = outputs[i, 0].cpu().numpy()\n                h, w = patch.shape\n                \n                pred_accumulator[y:y+h, x:x+w] += patch\n                count_accumulator[y:y+h, x:x+w] += 1\n            \n            if torch.cuda.is_available() and test_iterator.n % 10 == 0:\n                torch.cuda.empty_cache()\n    \n    # Average overlapping predictions\n    final_predictions = pred_accumulator / np.maximum(count_accumulator, 1)\n    \n    # Calculate test metrics on full fragment 1\n    best_test_threshold = 0.5\n    best_test_dice = 0\n    for thresh in np.arange(0.1, 0.9, 0.05):\n        preds = (final_predictions > thresh).astype(np.float32)\n        dice = dice_score(preds, fragment1_mask)\n        if dice > best_test_dice:\n            best_test_dice = dice\n            best_test_threshold = thresh\n    \n    # Calculate all test metrics\n    test_preds = (final_predictions > best_test_threshold).astype(np.float32)\n    test_metrics = calculate_all_metrics(test_preds, fragment1_mask)\n    test_metrics['threshold'] = float(best_test_threshold)\n    \n    # Save test visualization\n    print(\"\\nSaving test visualization and metrics...\")\n    save_test_visualization(\n        fragment1_volume, \n        fragment1_mask, \n        final_predictions, \n        test_metrics,\n        \"fragment1\"\n    )\n    \n    # Save training history\n    if config.save_history:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_history(history, filename=f'training_history_{timestamp}.json')\n        plot_training_history(history)\n    \n    # Final training metrics\n    final_train_metrics = {\n        'dice': float(history['train_dice'][-1]),\n        'f05': float(history['train_f05'][-1]),\n        'iou': float(history['train_iou'][-1]),\n        'precision': float(history['train_precision'][-1]),\n        'recall': float(history['train_recall'][-1]),\n        'threshold': float(best_train_threshold),\n        'tp': int(train_metrics['tp']),\n        'tn': int(train_metrics['tn']),\n        'fp': int(train_metrics['fp']),\n        'fn': int(train_metrics['fn'])\n    }\n    \n    # Save metrics summary\n    if config.save_metrics:\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        save_metrics_summary(\n            final_train_metrics,\n            test_metrics,\n            filename=f'metrics_summary_{timestamp}.csv'\n        )\n    \n    # Final summary\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL RESULTS SUMMARY - Train on Fragments 2 & 3, Test on Fragment 1\")\n    print(\"=\"*80)\n    \n    print(f\"\\nCONFIGURATION:\")\n    print(f\"  Image Size: {config.img_size}x{config.img_size}\")\n    print(f\"  Slices: {config.slices[0]}-{config.slices[-1]} ({len(config.slices)} slices)\")\n    \n    print(f\"\\nTRAINING DETAILS:\")\n    print(f\"  Training Data: Fragments 2 & 3 (100%)\")\n    print(f\"  Test Data: Fragment 1 (100%)\")\n    print(f\"  Validation: DISABLED\")\n    \n    print(f\"\\nTRAINING PERFORMANCE (Final Epoch):\")\n    print(f\"  Dice Score: {final_train_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {final_train_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {final_train_metrics['iou']:.4f}\")\n    print(f\"  Precision: {final_train_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_train_metrics['recall']:.4f}\")\n    \n    print(f\"\\nTEST RESULTS (Fragment 1):\")\n    print(f\"  Dice Score: {test_metrics['dice']:.4f}\")\n    print(f\"  F0.5 Score: {test_metrics['f05']:.4f}\")\n    print(f\"  IOU Score: {test_metrics['iou']:.4f}\")\n    print(f\"  Precision: {test_metrics['precision']:.4f}\")\n    print(f\"  Recall: {test_metrics['recall']:.4f}\")\n    print(f\"  Optimal Threshold: {test_metrics['threshold']:.3f}\")\n    print(f\"  TP: {test_metrics['tp']:,} | TN: {test_metrics['tn']:,} | FP: {test_metrics['fp']:,} | FN: {test_metrics['fn']:,}\")\n    \n    print(f\"\\nDATA SAVED:\")\n    print(f\"  Test results: test_results/\")\n    print(f\"  Training plots: training_plots/\")\n    \n    if config.save_history:\n        print(f\"  Training history: history/\")\n    if config.save_metrics:\n        print(f\"  Metrics summary: metrics/\")\n    \n    print(f\"\\nMODEL SAVED:\")\n    print(f\"  trained_model_fragments_2_3.pth\")\n    \n    print(f\"\\nKEY FEATURES:\")\n    print(f\"  1. Trained ONLY on Fragments 2 & 3 (no validation)\")\n    print(f\"  2. Tested on FULL Fragment 1\")\n    print(f\"  3. Image Size: {config.img_size}x{config.img_size}\")\n    print(f\"  4. All slices 12-30 included\")\n    print(f\"  5. Pure cross-fragment generalization test\")\n\n# Run the main function\nif __name__ == \"__main__\":\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        gpu_memory = torch.cuda.get_device_properties(0).total_memory / 1e9\n        print(f\"Available GPU memory: {gpu_memory:.1f} GB\")\n    \n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nfrom skimage.filters import threshold_otsu\nimport warnings\nimport random\nimport math\nfrom sklearn.model_selection import train_test_split, KFold\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport pandas as pd\nimport json\nfrom datetime import datetime\nimport pickle\n\nwarnings.filterwarnings('ignore')\n\n# Configuration\nclass Config:\n    # Data paths\n    fragment1_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\n    train_paths = [\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/2',\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\n    ]\n    \n    # Model parameters\n    img_size = 384\n    slices = list(range(12, 31))  # Slice 12-30 (19 slices)\n    batch_size = 4\n    epochs = 30\n    lr = 3e-4\n    \n    # 3D volumetric setup\n    num_input_slices = 5\n    \n    # Data split - True Spatial Disjoint Split\n    fragment1_val_split = 0.4  # 40% spatial region for validation\n    split_seed = 42\n    grid_size = 5  # For spatial grid split\n    \n    # Sample counts\n    train_samples_per_volume = 200\n    val_samples_total = 200\n    test_stride = 64\n    \n    # Memory optimization\n    use_amp = True\n    gradient_accumulation_steps = 2\n    \n    # Threshold settings - NOW USING OPTIMAL THRESHOLD EACH EPOCH\n    test_every_n_epochs = 1  # Test on full fragment every N epochs\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 1\n    \n    # Results saving\n    save_history = True\n    save_metrics = True\n    save_val_split = True\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nconfig = Config()\n\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data with enhanced preprocessing\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_filtered = cv2.bilateralFilter(img_resized, 5, 50, 50)\n                clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n                img_enhanced = clahe.apply(img_filtered)\n                img_normalized = img_enhanced.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        print(f\"Loaded volume from {base_path} with {volume.shape[0]} slices\")\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio for {base_path}: {ink_ratio:.4f} ({mask_binary.sum()} pixels)\")\n            return mask_binary\n    return None\n\ndef create_true_spatial_disjoint_split(img_size, val_ratio=0.4, grid_size=5, seed=42):\n    \"\"\"\n    Create a TRUE SPATIAL DISJOINT SPLIT using grid-based approach.\n    \"\"\"\n    np.random.seed(seed)\n    \n    if isinstance(img_size, int):\n        h, w = img_size, img_size\n    else:\n        h, w = img_size\n    \n    # Create grid\n    block_h = math.ceil(h / grid_size)\n    block_w = math.ceil(w / grid_size)\n    \n    # Initialize masks\n    val_mask = np.zeros((h, w), dtype=np.float32)\n    non_val_mask = np.zeros((h, w), dtype=np.float32)\n    \n    # Calculate number of validation blocks\n    total_blocks = grid_size * grid_size\n    val_blocks = int(total_blocks * val_ratio)\n    \n    # Randomly select blocks\n    block_indices = list(range(total_blocks))\n    np.random.shuffle(block_indices)\n    val_block_indices = block_indices[:val_blocks]\n    non_val_block_indices = block_indices[val_blocks:]\n    \n    def set_block_mask(mask, block_idx, value):\n        row = block_idx // grid_size\n        col = block_idx % grid_size\n        \n        y_start = row * block_h\n        y_end = min((row + 1) * block_h, h)\n        x_start = col * block_w\n        x_end = min((col + 1) * block_w, w)\n        \n        mask[y_start:y_end, x_start:x_end] = value\n    \n    # Assign blocks\n    for block_idx in val_block_indices:\n        set_block_mask(val_mask, block_idx, 1.0)\n    \n    for block_idx in non_val_block_indices:\n        set_block_mask(non_val_mask, block_idx, 1.0)\n    \n    # Ensure coverage\n    coverage = np.sum(val_mask + non_val_mask)\n    total_pixels = h * w\n    \n    if abs(coverage - total_pixels) >= 1e-6:\n        combined_mask = val_mask + non_val_mask\n        missing_mask = (combined_mask == 0).astype(np.float32)\n        if missing_mask.sum() > 0:\n            non_val_mask[missing_mask > 0] = 1.0\n    \n    return val_mask, non_val_mask\n\n# 3D Dataset for Fragments 2+3 ONLY\nclass TrainingDataset(Dataset):\n    def __init__(self, volumes, masks, augment=True, is_train=True, fragment_ids=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.augment = augment and is_train\n        self.fragment_ids = fragment_ids if fragment_ids else [0] * len(volumes)\n        self.samples = self._prepare_samples()\n        print(f\"Created training dataset with {len(self.samples)} 3D samples (Fragments 2+3 ONLY)\")\n    \n    def _prepare_samples(self):\n        samples = []\n        for vol_idx, volume in enumerate(self.volumes):\n            start_idx = config.num_input_slices // 2\n            end_idx = volume.shape[0] - config.num_input_slices // 2\n            \n            for _ in range(config.train_samples_per_volume):\n                center_slice = np.random.randint(start_idx, end_idx)\n                y = np.random.randint(0, volume.shape[1])\n                x = np.random.randint(0, volume.shape[2])\n                \n                samples.append({\n                    'volume_idx': vol_idx,\n                    'center_slice': center_slice,\n                    'y': y,\n                    'x': x,\n                    'fragment_id': self.fragment_ids[vol_idx]\n                })\n        \n        return samples\n    \n    def _extract_patch(self, volume, center_slice, y, x):\n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        half_size = config.img_size // 2\n        y_start = max(0, y - half_size)\n        y_end = min(volume.shape[1], y + half_size)\n        x_start = max(0, x - half_size)\n        x_end = min(volume.shape[2], x + half_size)\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = volume[slice_idx]\n            patch = slice_data[y_start:y_end, x_start:x_end]\n            \n            # Pad if necessary - FIXED: Ensure patch is exactly img_size x img_size\n            if patch.shape != (config.img_size, config.img_size):\n                pad_y_before = 0\n                pad_y_after = config.img_size - patch.shape[0]\n                pad_x_before = 0\n                pad_x_after = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((pad_y_before, pad_y_after), (pad_x_before, pad_x_after)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        return volume_patch, y_start, y_end, x_start, x_end\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        vol_idx = sample['volume_idx']\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        \n        volume_patch, y_start, y_end, x_start, x_end = self._extract_patch(\n            self.volumes[vol_idx], center_slice, y, x\n        )\n        \n        mask = self.masks[vol_idx]\n        mask_patch = mask[y_start:y_end, x_start:x_end]\n        \n        # Pad mask to match volume patch\n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y_before = 0\n            pad_y_after = config.img_size - mask_patch.shape[0]\n            pad_x_before = 0\n            pad_x_after = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((pad_y_before, pad_y_after), (pad_x_before, pad_x_after)), mode='constant')\n        \n        if self.augment:\n            if random.random() > 0.5:\n                volume_patch = np.ascontiguousarray(np.flip(volume_patch, axis=2))\n                mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=1))\n            \n            if random.random() > 0.5:\n                volume_patch = np.ascontiguousarray(np.flip(volume_patch, axis=1))\n                mask_patch = np.ascontiguousarray(np.flip(mask_patch, axis=0))\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor, sample['fragment_id']\n\n# Dataset for validation on 40% of Fragment 1\nclass ValidationDataset(Dataset):\n    def __init__(self, volumes, masks, region_masks, augment=False, is_train=False, fragment_ids=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.region_masks = region_masks\n        self.augment = augment and is_train\n        self.fragment_ids = fragment_ids if fragment_ids else [0] * len(volumes)\n        self.samples = self._prepare_samples()\n        print(f\"Created validation dataset with {len(self.samples)} 3D samples (40% of Fragment 1)\")\n    \n    def _prepare_samples(self):\n        samples = []\n        for vol_idx, (volume, region_mask) in enumerate(zip(self.volumes, self.region_masks)):\n            start_idx = config.num_input_slices // 2\n            end_idx = volume.shape[0] - config.num_input_slices // 2\n            \n            allowed_coords = np.argwhere(region_mask > 0.5)\n            if len(allowed_coords) == 0:\n                continue\n            \n            num_samples = min(config.val_samples_total, len(allowed_coords))\n            sampled_coords = allowed_coords[np.random.choice(len(allowed_coords), num_samples, replace=False)]\n            \n            slice_step = 3\n            slice_indices = list(range(start_idx, end_idx, slice_step))\n            if not slice_indices:\n                slice_indices = [start_idx]\n            \n            for y, x in sampled_coords:\n                for center_slice in slice_indices:\n                    samples.append({\n                        'volume_idx': vol_idx,\n                        'center_slice': center_slice,\n                        'y': y,\n                        'x': x,\n                        'fragment_id': self.fragment_ids[vol_idx]\n                    })\n        \n        return samples\n    \n    def _extract_patch(self, volume, center_slice, y, x):\n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        half_size = config.img_size // 2\n        y_start = max(0, y - half_size)\n        y_end = min(volume.shape[1], y + half_size)\n        x_start = max(0, x - half_size)\n        x_end = min(volume.shape[2], x + half_size)\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = volume[slice_idx]\n            patch = slice_data[y_start:y_end, x_start:x_end]\n            \n            # Pad if necessary\n            if patch.shape != (config.img_size, config.img_size):\n                pad_y_before = 0\n                pad_y_after = config.img_size - patch.shape[0]\n                pad_x_before = 0\n                pad_x_after = config.img_size - patch.shape[1]\n                patch = np.pad(patch, ((pad_y_before, pad_y_after), (pad_x_before, pad_x_after)), mode='constant')\n            \n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        return volume_patch, y_start, y_end, x_start, x_end\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        vol_idx = sample['volume_idx']\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        \n        volume_patch, y_start, y_end, x_start, x_end = self._extract_patch(\n            self.volumes[vol_idx], center_slice, y, x\n        )\n        \n        mask = self.masks[vol_idx]\n        mask_patch = mask[y_start:y_end, x_start:x_end]\n        \n        # Pad mask\n        if mask_patch.shape != (config.img_size, config.img_size):\n            pad_y_before = 0\n            pad_y_after = config.img_size - mask_patch.shape[0]\n            pad_x_before = 0\n            pad_x_after = config.img_size - mask_patch.shape[1]\n            mask_patch = np.pad(mask_patch, ((pad_y_before, pad_y_after), (pad_x_before, pad_x_after)), mode='constant')\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor, sample['fragment_id']\n\n# Fixed: Simple dataset for full inference on Fragment 1\nclass FullFragmentDataset(Dataset):\n    def __init__(self, volumes, masks, fragment_ids=None):\n        self.volumes = volumes\n        self.masks = masks\n        self.fragment_ids = fragment_ids if fragment_ids else [0] * len(volumes)\n        self.samples = self._prepare_samples()\n        print(f\"Created full fragment dataset with {len(self.samples)} 3D samples\")\n    \n    def _prepare_samples(self):\n        samples = []\n        for vol_idx, volume in enumerate(self.volumes):\n            start_idx = config.num_input_slices // 2\n            end_idx = volume.shape[0] - config.num_input_slices // 2\n            \n            stride = config.test_stride\n            for center_slice in range(start_idx, end_idx, 4):\n                # FIXED: Ensure we don't go out of bounds\n                for y in range(0, volume.shape[1] - config.img_size + 1, stride):\n                    for x in range(0, volume.shape[2] - config.img_size + 1, stride):\n                        samples.append({\n                            'volume_idx': vol_idx,\n                            'center_slice': center_slice,\n                            'y': y,\n                            'x': x,\n                            'fragment_id': self.fragment_ids[vol_idx]\n                        })\n        return samples\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        vol_idx = sample['volume_idx']\n        center_slice = sample['center_slice']\n        y, x = sample['y'], sample['x']\n        \n        start_slice = center_slice - config.num_input_slices // 2\n        end_slice = center_slice + config.num_input_slices // 2 + 1\n        \n        volume_patch = []\n        for slice_idx in range(start_slice, end_slice):\n            slice_data = self.volumes[vol_idx][slice_idx]\n            # FIXED: Extract exactly img_size x img_size patch\n            patch = slice_data[y:y+config.img_size, x:x+config.img_size]\n            volume_patch.append(patch)\n        \n        volume_patch = np.stack(volume_patch, axis=0)\n        \n        mask = self.masks[vol_idx]\n        mask_patch = mask[y:y+config.img_size, x:x+config.img_size]\n        \n        volume_tensor = torch.FloatTensor(volume_patch)\n        mask_tensor = torch.FloatTensor(mask_patch).unsqueeze(0)\n        \n        return volume_tensor, mask_tensor, sample['fragment_id'], y, x\n\n# 3D UNet Model\nclass VolumetricUNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = smp.Unet(\n            encoder_name=\"resnet34\",\n            encoder_weights=\"imagenet\",\n            in_channels=config.num_input_slices,\n            classes=1,\n            encoder_depth=5,\n            decoder_channels=[256, 128, 64, 32, 16],\n            activation=None\n        )\n        \n    def forward(self, x):\n        return self.model(x)\n\n# Advanced loss function\nclass Combined3DLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.dice_loss = smp.losses.DiceLoss(mode='binary', from_logits=True)\n        self.focal_loss = smp.losses.FocalLoss(mode='binary')\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        \n    def forward(self, pred, target):\n        dice = self.dice_loss(pred, target)\n        focal = self.focal_loss(pred, target)\n        bce = self.bce_loss(pred, target)\n        return 0.6 * dice + 0.3 * focal + 0.1 * bce\n\n# Metrics functions\ndef calculate_all_metrics(pred, target, threshold=0.5):\n    pred_binary = (pred > threshold).astype(np.float32)\n    \n    pred_flat = pred_binary.flatten().astype(int)\n    target_flat = target.flatten().astype(int)\n    \n    tn, fp, fn, tp = confusion_matrix(target_flat, pred_flat, labels=[0, 1]).ravel()\n    \n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    dice = (2 * tp) / (2 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0\n    \n    beta_squared = 0.5 ** 2\n    f05 = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (beta_squared * precision + recall) > 0 else 0\n    \n    iou = tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0\n    \n    return {\n        'precision': float(precision),\n        'recall': float(recall),\n        'dice': float(dice),\n        'f05': float(f05),\n        'iou': float(iou),\n        'tp': int(tp),\n        'tn': int(tn),\n        'fp': int(fp),\n        'fn': int(fn)\n    }\n\ndef dice_score(pred, target, smooth=1e-6):\n    pred = pred.flatten()\n    target = target.flatten()\n    intersection = (pred * target).sum()\n    return float((2. * intersection + smooth) / (pred.sum() + target.sum() + smooth))\n\ndef find_optimal_threshold(pred, target):\n    \"\"\"Find optimal threshold for binary classification\"\"\"\n    best_threshold = 0.5\n    best_dice = 0\n    for thresh in np.arange(0.1, 0.9, 0.05):\n        preds = (pred > thresh).astype(np.float32)\n        dice = dice_score(preds, target)\n        if dice > best_dice:\n            best_dice = dice\n            best_threshold = thresh\n    return best_threshold, best_dice\n\n# Fixed: Test model on full fragment\ndef test_model_on_full_fragment(model, test_loader, device, fragment_mask):\n    \"\"\"Test model on full fragment with proper patch accumulation\"\"\"\n    model.eval()\n    \n    # Initialize prediction accumulator\n    pred_accumulator = np.zeros(fragment_mask.shape, dtype=np.float32)\n    count_accumulator = np.zeros(fragment_mask.shape, dtype=np.float32)\n    \n    with torch.no_grad():\n        for images, masks, _, y_coords, x_coords in tqdm(test_loader, desc=\"Testing on full fragment\"):\n            images = images.to(device)\n            outputs = torch.sigmoid(model(images))\n            \n            for i in range(outputs.shape[0]):\n                y = int(y_coords[i].item())\n                x = int(x_coords[i].item())\n                \n                patch = outputs[i, 0].cpu().numpy()\n                h, w = patch.shape\n                \n                # FIXED: Ensure we don't go out of bounds\n                y_end = min(y + h, fragment_mask.shape[0])\n                x_end = min(x + w, fragment_mask.shape[1])\n                actual_h = y_end - y\n                actual_w = x_end - x\n                \n                if actual_h > 0 and actual_w > 0:\n                    pred_accumulator[y:y_end, x:x_end] += patch[:actual_h, :actual_w]\n                    count_accumulator[y:y_end, x:x_end] += 1\n    \n    # Average overlapping predictions\n    final_predictions = np.zeros_like(pred_accumulator)\n    valid_mask = count_accumulator > 0\n    final_predictions[valid_mask] = pred_accumulator[valid_mask] / count_accumulator[valid_mask]\n    \n    # Find optimal threshold for test data\n    best_threshold, best_dice = find_optimal_threshold(final_predictions, fragment_mask)\n    \n    # Calculate metrics with optimal threshold\n    test_preds = (final_predictions > best_threshold).astype(np.float32)\n    test_metrics = calculate_all_metrics(test_preds, fragment_mask)\n    test_metrics['threshold'] = float(best_threshold)\n    \n    return test_metrics, final_predictions\n\ndef evaluate_model_with_optimal_threshold(model, dataloader, device):\n    \"\"\"Evaluate model with optimal threshold search (for validation)\"\"\"\n    model.eval()\n    all_outputs = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, masks, _ in dataloader:\n            images = images.to(device)\n            outputs = torch.sigmoid(model(images))\n            \n            all_outputs.append(outputs.cpu().numpy())\n            all_targets.append(masks.cpu().numpy())\n    \n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n    \n    # Find optimal threshold\n    best_threshold, best_dice = find_optimal_threshold(all_outputs, all_targets)\n    \n    # Calculate all metrics with optimal threshold\n    preds = (all_outputs > best_threshold).astype(np.float32)\n    metrics = calculate_all_metrics(preds, all_targets)\n    metrics['threshold'] = float(best_threshold)\n    \n    return metrics, all_outputs, all_targets\n\ndef main():\n    print(\"=\"*80)\n    print(\"3D UNet with ResNet34 - TRAINING ON FRAGMENTS 2+3 ONLY\")\n    print(\"=\"*80)\n    print(\"TRAINING STRATEGY:\")\n    print(\"  Train on: Fragments 2 + 3 ONLY (full)\")\n    print(\"  Validate on: 40% of Fragment 1 (spatially disjoint, COMPLETELY UNSEEN)\")\n    print(\"  Test on: Full Fragment 1\")\n    print(\"=\"*80)\n    print(\"IMPORTANT: Using optimal threshold search for validation each epoch\")\n    print(\"Best model selected based on optimal validation Dice\")\n    print(\"=\"*80)\n    \n    # Set random seeds\n    torch.manual_seed(config.split_seed)\n    np.random.seed(config.split_seed)\n    random.seed(config.split_seed)\n    \n    # Load all data\n    print(\"\\nLoading data...\")\n    \n    # Load fragment 1\n    fragment1_volume = load_volume_data(config.fragment1_path, config.slices)\n    fragment1_mask = load_mask_data(config.fragment1_path)\n    \n    if fragment1_volume is None or fragment1_mask is None:\n        raise ValueError(f\"Failed to load fragment 1 data from {config.fragment1_path}\")\n    \n    # Load training fragments 2 and 3\n    train_volumes = []\n    train_masks = []\n    fragment_ids = []\n    \n    for i, path in enumerate(config.train_paths):\n        volume = load_volume_data(path, config.slices)\n        mask = load_mask_data(path)\n        \n        if volume is not None and mask is not None:\n            train_volumes.append(volume)\n            train_masks.append(mask)\n            fragment_ids.append(i + 1)\n    \n    print(f\"\\nLoaded {len(train_volumes)} training fragments (Fragments 2+3)\")\n    \n    # Create spatial split\n    print(f\"\\nCreating TRUE SPATIAL DISJOINT SPLIT for fragment 1...\")\n    val_mask, non_val_mask = create_true_spatial_disjoint_split(\n        config.img_size, \n        val_ratio=config.fragment1_val_split,\n        grid_size=config.grid_size,\n        seed=config.split_seed\n    )\n    \n    # Prepare datasets\n    print(f\"\\nPreparing datasets...\")\n    \n    train_dataset = TrainingDataset(\n        volumes=train_volumes,\n        masks=train_masks,\n        augment=True,\n        is_train=True,\n        fragment_ids=fragment_ids\n    )\n    \n    val_dataset = ValidationDataset(\n        volumes=[fragment1_volume],\n        masks=[fragment1_mask],\n        region_masks=[val_mask],\n        augment=False,\n        is_train=False,\n        fragment_ids=[0]\n    )\n    \n    test_dataset = FullFragmentDataset(\n        volumes=[fragment1_volume],\n        masks=[fragment1_mask],\n        fragment_ids=[0]\n    )\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=0\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        num_workers=0\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        num_workers=0\n    )\n    \n    print(f\"\\nData Loaders Created:\")\n    print(f\"  Training samples: {len(train_dataset)} (from Fragments 2+3 ONLY)\")\n    print(f\"  Validation samples: {len(val_dataset)} (from 40% UNSEEN region of Fragment 1)\")\n    print(f\"  Test samples: {len(test_dataset)} (full Fragment 1)\")\n    \n    # Initialize model\n    model = VolumetricUNet().to(config.device)\n    print(f\"\\nModel initialized on {config.device}\")\n    \n    # Initialize optimizer and loss\n    optimizer = optim.Adam(model.parameters(), lr=config.lr, weight_decay=1e-5)\n    scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=10, verbose=True)\n    criterion = Combined3DLoss()\n    \n    # Mixed precision training\n    if config.use_amp:\n        scaler = torch.cuda.amp.GradScaler()\n    \n    # Training history\n    history = {\n        'train_loss': [],\n        'train_dice': [],\n        'train_f05': [],\n        'train_iou': [],\n        'train_precision': [],\n        'train_recall': [],\n        'train_threshold': [],\n        'val_dice': [],\n        'val_f05': [],\n        'val_iou': [],\n        'val_precision': [],\n        'val_recall': [],\n        'val_threshold': [],\n        'test_dice': [],\n        'test_f05': [],\n        'test_iou': [],\n        'test_precision': [],\n        'test_recall': [],\n        'test_threshold': []\n    }\n    \n    best_val_dice = 0\n    best_model_state = None\n    best_epoch = 0\n    \n    print(f\"\\nStarting training for {config.epochs} epochs...\")\n    print(\"Using optimal threshold search for validation each epoch\")\n    print(\"Best model selected based on optimal validation Dice\")\n    print(\"=\"*80)\n    \n    for epoch in range(config.epochs):\n        # Training phase\n        model.train()\n        epoch_train_loss = 0\n        train_outputs = []\n        train_targets = []\n        \n        optimizer.zero_grad()\n        \n        for batch_idx, (images, masks, _) in enumerate(tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.epochs}')):\n            images, masks = images.to(config.device), masks.to(config.device)\n            \n            if config.use_amp:\n                with torch.cuda.amp.autocast():\n                    outputs = model(images)\n                    loss = criterion(outputs, masks)\n                \n                loss = loss / config.gradient_accumulation_steps\n                scaler.scale(loss).backward()\n                \n                if (batch_idx + 1) % config.gradient_accumulation_steps == 0:\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad()\n            else:\n                outputs = model(images)\n                loss = criterion(outputs, masks)\n                \n                loss = loss / config.gradient_accumulation_steps\n                loss.backward()\n                \n                if (batch_idx + 1) % config.gradient_accumulation_steps == 0:\n                    optimizer.step()\n                    optimizer.zero_grad()\n            \n            epoch_train_loss += loss.item() * config.gradient_accumulation_steps\n            \n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_outputs.append(preds.cpu().numpy())\n                train_targets.append(masks.cpu().numpy())\n        \n        # Step optimizer if remaining gradients\n        if config.use_amp and (batch_idx + 1) % config.gradient_accumulation_steps != 0:\n            scaler.step(optimizer)\n            scaler.update()\n        elif (batch_idx + 1) % config.gradient_accumulation_steps != 0:\n            optimizer.step()\n        \n        # Calculate training metrics with optimal threshold\n        train_outputs = np.concatenate(train_outputs)\n        train_targets = np.concatenate(train_targets)\n        \n        best_train_threshold, best_train_dice = find_optimal_threshold(train_outputs, train_targets)\n        train_preds = (train_outputs > best_train_threshold).astype(np.float32)\n        train_metrics = calculate_all_metrics(train_preds, train_targets)\n        train_metrics['threshold'] = float(best_train_threshold)\n        \n        # Store training metrics\n        history['train_loss'].append(float(epoch_train_loss / len(train_loader)))\n        history['train_dice'].append(float(train_metrics['dice']))\n        history['train_f05'].append(float(train_metrics['f05']))\n        history['train_iou'].append(float(train_metrics['iou']))\n        history['train_precision'].append(float(train_metrics['precision']))\n        history['train_recall'].append(float(train_metrics['recall']))\n        history['train_threshold'].append(float(best_train_threshold))\n        \n        # Validation phase WITH OPTIMAL THRESHOLD SEARCH (as requested)\n        val_metrics, _, _ = evaluate_model_with_optimal_threshold(\n            model, val_loader, config.device\n        )\n        \n        # Store validation metrics\n        history['val_dice'].append(float(val_metrics['dice']))\n        history['val_f05'].append(float(val_metrics['f05']))\n        history['val_iou'].append(float(val_metrics['iou']))\n        history['val_precision'].append(float(val_metrics['precision']))\n        history['val_recall'].append(float(val_metrics['recall']))\n        history['val_threshold'].append(float(val_metrics['threshold']))\n        \n        # Test phase (every N epochs)\n        test_metrics = None\n        if (epoch + 1) % config.test_every_n_epochs == 0 or epoch == config.epochs - 1:\n            print(f\"\\n  Running full fragment test (epoch {epoch+1})...\")\n            test_metrics, _ = test_model_on_full_fragment(\n                model, test_loader, config.device, fragment1_mask\n            )\n            \n            # Store test metrics\n            history['test_dice'].append(float(test_metrics['dice']))\n            history['test_f05'].append(float(test_metrics['f05']))\n            history['test_iou'].append(float(test_metrics['iou']))\n            history['test_precision'].append(float(test_metrics['precision']))\n            history['test_recall'].append(float(test_metrics['recall']))\n            history['test_threshold'].append(float(test_metrics['threshold']))\n        \n        # Update best model based on OPTIMAL validation Dice\n        if val_metrics['dice'] > best_val_dice:\n            best_val_dice = val_metrics['dice']\n            best_model_state = model.state_dict().copy()\n            best_epoch = epoch + 1\n            print(f\"  NEW BEST! Validation Dice: {best_val_dice:.4f} (epoch {best_epoch})\")\n        \n        # Print progress\n        print(f\"\\nEpoch {epoch+1}/{config.epochs}:\")\n        print(f\"  Training (Fragments 2+3):\")\n        print(f\"    Loss: {epoch_train_loss/len(train_loader):.4f}, Dice: {train_metrics['dice']:.4f}\")\n        print(f\"    F0.5: {train_metrics['f05']:.4f}, IOU: {train_metrics['iou']:.4f}\")\n        print(f\"    Precision: {train_metrics['precision']:.4f}, Recall: {train_metrics['recall']:.4f}\")\n        print(f\"    Optimal threshold: {best_train_threshold:.3f}\")\n        \n        print(f\"  Validation (40% UNSEEN Fragment 1, OPTIMAL threshold):\")\n        print(f\"    Dice: {val_metrics['dice']:.4f}, F0.5: {val_metrics['f05']:.4f}\")\n        print(f\"    IOU: {val_metrics['iou']:.4f}, Precision: {val_metrics['precision']:.4f}\")\n        print(f\"    Recall: {val_metrics['recall']:.4f}\")\n        print(f\"    Optimal threshold: {val_metrics['threshold']:.3f}\")\n        \n        if test_metrics:\n            print(f\"  Test (Full Fragment 1, OPTIMAL threshold):\")\n            print(f\"    Dice: {test_metrics['dice']:.4f}, F0.5: {test_metrics['f05']:.4f}\")\n            print(f\"    IOU: {test_metrics['iou']:.4f}, Precision: {test_metrics['precision']:.4f}\")\n            print(f\"    Recall: {test_metrics['recall']:.4f}\")\n            print(f\"    Optimal threshold: {test_metrics['threshold']:.3f}\")\n        \n        # Learning rate scheduling based on optimal validation Dice\n        scheduler.step(val_metrics['dice'])\n        \n        # Clear GPU memory\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    # Load best model for final evaluation\n    if best_model_state:\n        model.load_state_dict(best_model_state)\n        print(f\"\\nLoaded BEST model from epoch {best_epoch} with validation Dice: {best_val_dice:.4f}\")\n    \n    # FINAL EVALUATION\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL EVALUATION WITH BEST MODEL\")\n    print(\"=\"*80)\n    \n    # Final training evaluation\n    model.eval()\n    with torch.no_grad():\n        train_outputs_final = []\n        train_targets_final = []\n        for images, masks, _ in train_loader:\n            images = images.to(config.device)\n            outputs = torch.sigmoid(model(images))\n            train_outputs_final.append(outputs.cpu().numpy())\n            train_targets_final.append(masks.cpu().numpy())\n        \n        train_outputs_final = np.concatenate(train_outputs_final)\n        train_targets_final = np.concatenate(train_targets_final)\n        \n        best_train_threshold_final, _ = find_optimal_threshold(train_outputs_final, train_targets_final)\n        final_train_preds = (train_outputs_final > best_train_threshold_final).astype(np.float32)\n        final_train_metrics = calculate_all_metrics(final_train_preds, train_targets_final)\n        final_train_metrics['threshold'] = float(best_train_threshold_final)\n    \n    # Final validation evaluation\n    final_val_metrics, _, _ = evaluate_model_with_optimal_threshold(\n        model, val_loader, config.device\n    )\n    \n    # Final test evaluation\n    print(f\"\\nRunning final test on full Fragment 1 with best model...\")\n    final_test_metrics, final_predictions = test_model_on_full_fragment(\n        model, test_loader, config.device, fragment1_mask\n    )\n    \n    # Final summary\n    print(\"\\n\" + \"=\"*80)\n    print(\"FINAL RESULTS SUMMARY - TRAINING ON FRAGMENTS 2+3 ONLY\")\n    print(\"=\"*80)\n    \n    print(f\"\\nTRAINING STRATEGY:\")\n    print(f\"  ✓ Train: Fragments 2 + 3 ONLY (full)\")\n    print(f\"  ✓ Validate: 40% of Fragment 1 (spatially disjoint, COMPLETELY UNSEEN)\")\n    print(f\"  ✓ Test: Full Fragment 1\")\n    \n    print(f\"\\nBEST MODEL SELECTION:\")\n    print(f\"  Selected from epoch: {best_epoch}\")\n    print(f\"  Based on optimal validation Dice: {best_val_dice:.4f}\")\n    \n    print(f\"\\nFINAL TRAINING METRICS (Fragments 2+3 ONLY, optimal threshold):\")\n    print(f\"  Dice: {final_train_metrics['dice']:.4f} | F0.5: {final_train_metrics['f05']:.4f}\")\n    print(f\"  IOU: {final_train_metrics['iou']:.4f} | Precision: {final_train_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_train_metrics['recall']:.4f} | Threshold: {final_train_metrics['threshold']:.3f}\")\n    \n    print(f\"\\nFINAL VALIDATION METRICS (40% Fragment 1, optimal threshold):\")\n    print(f\"  Dice: {final_val_metrics['dice']:.4f} | F0.5: {final_val_metrics['f05']:.4f}\")\n    print(f\"  IOU: {final_val_metrics['iou']:.4f} | Precision: {final_val_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_val_metrics['recall']:.4f} | Threshold: {final_val_metrics['threshold']:.3f}\")\n    \n    print(f\"\\nFINAL TEST METRICS (Full Fragment 1, optimal threshold):\")\n    print(f\"  Dice: {final_test_metrics['dice']:.4f} | F0.5: {final_test_metrics['f05']:.4f}\")\n    print(f\"  IOU: {final_test_metrics['iou']:.4f} | Precision: {final_test_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_test_metrics['recall']:.4f} | Threshold: {final_test_metrics['threshold']:.3f}\")\n    \n    print(f\"\\nGENERALIZATION ANALYSIS:\")\n    print(f\"  Training Dice: {final_train_metrics['dice']:.4f} (Fragments 2+3)\")\n    print(f\"  Validation Dice: {final_val_metrics['dice']:.4f} (40% of Fragment 1, UNSEEN)\")\n    print(f\"  Test Dice: {final_test_metrics['dice']:.4f} (Full Fragment 1)\")\n    print(f\"  Generalization Gap (Train→Test): {final_train_metrics['dice'] - final_test_metrics['dice']:.4f}\")\n    \n    print(f\"\\nPERFORMANCE INSIGHTS:\")\n    print(f\"  1. Model trained on Fragments 2+3 generalizes to unseen Fragment 1\")\n    print(f\"  2. Optimal threshold search maximizes validation performance\")\n    print(f\"  3. Best weights selected based on optimal validation metrics\")\n    print(f\"  4. True test of generalization ability\")\n\n# Run the main function\nif __name__ == \"__main__\":\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        gpu_memory = torch.cuda.get_device_properties(0).total_memory / 1e9\n        print(f\"Available GPU memory: {gpu_memory:.1f} GB\")\n        print(f\"Using device: {torch.cuda.get_device_name(0)}\")\n    \n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport timm\nimport numpy as np\nimport cv2\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport os\nfrom tqdm import tqdm\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Configuration\nclass Config:\n    # Data - Updated to correct Vesuvius Challenge paths\n    train_dirs = ['/kaggle/input/vesuvius-challenge-ink-detection/train/1', '/kaggle/input/vesuvius-challenge-ink-detection/train/2']\n    test_dir = '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\n    slice_range = list(range(12, 31))\n    img_size = 480\n    num_slices = 19\n    \n    # Model\n    backbone = 'seresnet50'\n    encoder_weights = 'imagenet'\n    num_classes = 1\n    \n    # Training\n    batch_size = 2\n    epochs = 30\n    lr = 1e-4\n    weight_decay = 1e-5\n    num_folds = 3  # Reduced for faster testing\n    \n    # Post-processing\n    edge_margin = 20\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nconfig = Config()\n\n# Manual KFold implementation\nclass ManualKFold:\n    def __init__(self, n_splits=5, shuffle=True, random_state=None):\n        self.n_splits = n_splits\n        self.shuffle = shuffle\n        self.random_state = random_state\n        \n    def split(self, X):\n        n_samples = len(X)\n        indices = np.arange(n_samples)\n        \n        if self.shuffle:\n            if self.random_state is not None:\n                np.random.seed(self.random_state)\n            np.random.shuffle(indices)\n        \n        fold_sizes = np.full(self.n_splits, n_samples // self.n_splits, dtype=int)\n        fold_sizes[:n_samples % self.n_splits] += 1\n        \n        current = 0\n        for fold_size in fold_sizes:\n            start, stop = current, current + fold_size\n            val_indices = indices[start:stop]\n            train_indices = np.concatenate([indices[:start], indices[stop:]])\n            yield train_indices, val_indices\n            current = stop\n\n# Manual metrics implementation\ndef manual_precision_recall_f1(pred, target):\n    \"\"\"Manual implementation of precision, recall, and F1 score\"\"\"\n    pred_flat = pred.reshape(-1)\n    target_flat = target.reshape(-1)\n    \n    # True positives, false positives, false negatives\n    tp = np.sum((pred_flat == 1) & (target_flat == 1))\n    fp = np.sum((pred_flat == 1) & (target_flat == 0))\n    fn = np.sum((pred_flat == 0) & (target_flat == 1))\n    \n    # Calculate metrics\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n    \n    # F1 score\n    f1 = 2 * (precision * recall) / (precision + recall + 1e-8)\n    \n    return precision, recall, f1\n\n# Custom SegFormer with SEResNet backbone\nclass SegFormerWithSEResNet(nn.Module):\n    def __init__(self, num_classes=1, backbone='seresnet50', in_channels=4):\n        super().__init__()\n        \n        self.encoder = timm.create_model(backbone, features_only=True, pretrained=True)\n        \n        original_first_conv = self.encoder.conv1\n        self.encoder.conv1 = nn.Conv2d(\n            in_channels, \n            original_first_conv.out_channels,\n            kernel_size=original_first_conv.kernel_size,\n            stride=original_first_conv.stride,\n            padding=original_first_conv.padding,\n            bias=original_first_conv.bias is not None\n        )\n        \n        with torch.no_grad():\n            if original_first_conv.weight.shape[1] == 3:\n                new_weight = original_first_conv.weight.data\n                if in_channels > 3:\n                    extra_channels = in_channels - 3\n                    extra_weights = original_first_conv.weight.data.mean(dim=1, keepdim=True).repeat(1, extra_channels, 1, 1)\n                    extra_weights = extra_weights / extra_channels\n                    new_weight = torch.cat([new_weight, extra_weights], dim=1)\n                elif in_channels == 1:\n                    new_weight = original_first_conv.weight.data.mean(dim=1, keepdim=True)\n                self.encoder.conv1.weight.data = new_weight\n        \n        feature_info = self.encoder.feature_info\n        self.channels = [info['num_chs'] for info in feature_info]\n        self.strides = [info['reduction'] for info in feature_info]\n        \n        self.decoder = SegFormerDecoder(self.channels, 256, num_classes, output_size=config.img_size)\n        \n    def forward(self, x):\n        features = self.encoder(x)\n        return self.decoder(features)\n\nclass SegFormerDecoder(nn.Module):\n    def __init__(self, encoder_channels, decoder_dim, num_classes, output_size=480):\n        super().__init__()\n        self.output_size = output_size\n        \n        self.fusion_layers = nn.ModuleList()\n        for in_channels in encoder_channels:\n            self.fusion_layers.append(\n                nn.Sequential(\n                    nn.Conv2d(in_channels, decoder_dim, 1),\n                    nn.BatchNorm2d(decoder_dim),\n                    nn.ReLU(inplace=True)\n                )\n            )\n        \n        self.head = nn.Sequential(\n            nn.Conv2d(decoder_dim * len(encoder_channels), decoder_dim, 3, padding=1),\n            nn.BatchNorm2d(decoder_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout2d(0.1),\n            nn.Conv2d(decoder_dim, decoder_dim // 2, 3, padding=1),\n            nn.BatchNorm2d(decoder_dim // 2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(decoder_dim // 2, num_classes, 1)\n        )\n        \n    def forward(self, features):\n        target_size = features[-1].shape[2:]\n        \n        fused_features = []\n        for i, (feature, fusion_layer) in enumerate(zip(features, self.fusion_layers)):\n            if i < len(features) - 1:\n                feature = F.interpolate(feature, size=target_size, mode='bilinear', align_corners=False)\n            fused = fusion_layer(feature)\n            fused_features.append(fused)\n        \n        x = torch.cat(fused_features, dim=1)\n        x = self.head(x)\n        x = F.interpolate(x, size=self.output_size, mode='bilinear', align_corners=False)\n        \n        return x\n\ndef calculate_metrics(pred, target, threshold=0.5):\n    \"\"\"Calculate precision, recall, dice, and F0.5 score\"\"\"\n    pred_bin = (pred > threshold).float()\n    target_bin = (target > 0.5).float()\n    \n    pred_flat = pred_bin.view(-1).cpu().numpy()\n    target_flat = target_bin.view(-1).cpu().numpy()\n    \n    if np.sum(target_flat) == 0:\n        return {\n            'precision': 0.0,\n            'recall': 0.0,\n            'dice': 0.0,\n            'f0.5': 0.0\n        }\n    \n    # Use manual implementation\n    precision, recall, _ = manual_precision_recall_f1(pred_flat, target_flat)\n    \n    intersection = (pred_bin * target_bin).sum()\n    dice = (2. * intersection) / (pred_bin.sum() + target_bin.sum() + 1e-8)\n    \n    beta = 0.5\n    f_beta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall + 1e-8)\n    \n    return {\n        'precision': precision,\n        'recall': recall,\n        'dice': dice.item(),\n        'f0.5': f_beta\n    }\n\ndef find_optimal_threshold(predictions, masks, data_type=\"Validation\"):\n    \"\"\"Find optimal threshold for given predictions and masks\"\"\"\n    best_threshold = 0.5\n    best_f05 = 0.0\n    best_dice = 0.0\n    \n    thresholds = np.arange(0.1, 0.9, 0.05)\n    \n    for threshold in thresholds:\n        metrics = calculate_metrics(predictions, masks, threshold)\n        \n        if metrics['f0.5'] > best_f05:\n            best_f05 = metrics['f0.5']\n            best_threshold = threshold\n            best_dice = metrics['dice']\n    \n    print(f\"{data_type} - Optimal threshold: {best_threshold:.3f}\")\n    print(f\"  F0.5: {best_f05:.4f}, Dice: {best_dice:.4f}\")\n    \n    return best_threshold, best_f05, best_dice\n\n# Data preprocessing functions - COMPLETELY FIXED VERSION\ndef apply_filters(image):\n    \"\"\"Apply various filters for preprocessing - FIXED for 16-bit TIF images\"\"\"\n    # Handle different input types\n    if image.dtype == np.uint16:\n        # Convert 16-bit to 8-bit for OpenCV processing\n        image = cv2.normalize(image, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n    elif image.dtype == np.float32 or image.dtype == np.float64:\n        # Convert float to uint8\n        image = cv2.normalize(image, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n    \n    # Ensure we have uint8 for OpenCV operations\n    if image.dtype != np.uint8:\n        image = image.astype(np.uint8)\n    \n    # Apply Gaussian blur\n    image_blurred = cv2.GaussianBlur(image, (3, 3), 0)\n    \n    # Apply CLAHE for contrast enhancement\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    image_enhanced = clahe.apply(image_blurred)\n    \n    # Convert back to float32 and normalize\n    image_normalized = image_enhanced.astype(np.float32) / 255.0\n    image_normalized = (image_normalized - image_normalized.mean()) / (image_normalized.std() + 1e-8)\n    \n    return image_normalized\n\ndef load_and_preprocess_slices(data_dir, slice_range, target_size=(480, 480)):\n    \"\"\"Load and preprocess slices from a directory - FIXED for TIF loading\"\"\"\n    slices = []\n    print(f\"  Loading slices {slice_range[0]} to {slice_range[-1]}...\")\n    \n    for slice_idx in slice_range:\n        slice_path = os.path.join(data_dir, 'surface_volume', f'{slice_idx}.tif')\n        if os.path.exists(slice_path):\n            try:\n                # Load slice using PIL for better TIF handling\n                with Image.open(slice_path) as img:\n                    slice_img = np.array(img)\n                \n                if slice_img is None:\n                    print(f\"Warning: Could not load {slice_path}\")\n                    continue\n                \n                # Print image info for debugging\n                print(f\"    Slice {slice_idx}: shape={slice_img.shape}, dtype={slice_img.dtype}, range=({slice_img.min()}, {slice_img.max()})\")\n                \n                # Fix data type issues - convert big-endian to little-endian if needed\n                if slice_img.dtype.byteorder == '>' or slice_img.dtype.byteorder == '=':\n                    # Big-endian or native byte order - convert to little-endian\n                    slice_img = slice_img.byteswap().newbyteorder()\n                \n                # Ensure proper data type for OpenCV\n                if slice_img.dtype == np.uint16:\n                    # Convert 16-bit to 8-bit\n                    slice_img = cv2.normalize(slice_img, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n                elif slice_img.dtype != np.uint8:\n                    slice_img = slice_img.astype(np.uint8)\n                \n                # Resize\n                slice_img = cv2.resize(slice_img, target_size, interpolation=cv2.INTER_AREA)\n                \n                # Apply filters\n                slice_img = apply_filters(slice_img)\n                \n                slices.append(slice_img)\n                \n            except Exception as e:\n                print(f\"Error loading {slice_path}: {e}\")\n                continue\n        else:\n            print(f\"Warning: File not found {slice_path}\")\n    \n    if slices:\n        volume = np.stack(slices, axis=0)\n        print(f\"  Created volume with shape: {volume.shape}\")\n        return volume\n    else:\n        print(f\"  No slices loaded from {data_dir}\")\n        return None\n\ndef load_mask(mask_path, target_size=(480, 480)):\n    \"\"\"Load and preprocess mask\"\"\"\n    if not os.path.exists(mask_path):\n        raise ValueError(f\"Mask file not found: {mask_path}\")\n    \n    # Load mask using PIL for consistency\n    with Image.open(mask_path) as img:\n        mask = np.array(img)\n    \n    if mask is None:\n        raise ValueError(f\"Could not load mask from {mask_path}\")\n    \n    print(f\"  Mask: shape={mask.shape}, dtype={mask.dtype}, unique values={np.unique(mask)}\")\n    \n    # Fix data type if needed\n    if mask.dtype.byteorder == '>' or mask.dtype.byteorder == '=':\n        mask = mask.byteswap().newbyteorder()\n    \n    mask = cv2.resize(mask, target_size, interpolation=cv2.INTER_NEAREST)\n    return (mask > 0).astype(np.float32)\n\n# Dataset class\nclass VolumeDataset(Dataset):\n    def __init__(self, volumes, masks, is_train=True):\n        self.volumes = volumes\n        self.masks = masks\n        self.is_train = is_train\n        \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        volume = self.volumes[idx]\n        mask = self.masks[idx]\n        \n        num_slices = volume.shape[0]\n        if num_slices > 4:\n            selected_indices = [0, num_slices//3, 2*num_slices//3, -1]\n            volume = volume[selected_indices]\n        elif num_slices < 4:\n            padding = [volume[-1:]] * (4 - num_slices)\n            volume = np.concatenate([volume] + padding, axis=0)\n        \n        volume = torch.FloatTensor(volume)\n        mask = torch.FloatTensor(mask)\n        \n        if self.is_train and np.random.random() > 0.5:\n            if np.random.random() > 0.5:\n                volume = torch.flip(volume, [1])\n                mask = torch.flip(mask, [0])\n            if np.random.random() > 0.5:\n                volume = torch.flip(volume, [2])\n                mask = torch.flip(mask, [1])\n        \n        return volume, mask.unsqueeze(0)\n\n# Post-processing functions\ndef remove_edge_predictions(mask, margin=20):\n    h, w = mask.shape\n    edge_mask = np.zeros_like(mask)\n    edge_mask[margin:h-margin, margin:w-margin] = 1\n    return mask * edge_mask\n\ndef apply_post_processing(pred_mask, threshold=0.5, edge_margin=20):\n    binary_mask = (pred_mask > threshold).astype(np.uint8)\n    kernel = np.ones((3, 3), np.uint8)\n    binary_mask = cv2.morphologyEx(binary_mask, cv2.MORPH_OPEN, kernel)\n    binary_mask = remove_edge_predictions(binary_mask, edge_margin)\n    return binary_mask\n\n# Training function\ndef train_epoch(model, dataloader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    \n    for volumes, masks in tqdm(dataloader, desc='Training'):\n        volumes = volumes.to(device)\n        masks = masks.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(volumes)\n        loss = criterion(outputs, masks)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n    \n    return running_loss / len(dataloader)\n\n# Validation function (for training phase)\ndef validate_epoch_training(model, dataloader, criterion, device):\n    \"\"\"Validation during training phase - uses validation data\"\"\"\n    model.eval()\n    running_loss = 0.0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for volumes, masks in tqdm(dataloader, desc='Validation'):\n            volumes = volumes.to(device)\n            masks = masks.to(device)\n            \n            outputs = model(volumes)\n            loss = criterion(outputs, masks)\n            running_loss += loss.item()\n            \n            all_preds.append(outputs.sigmoid().cpu())\n            all_targets.append(masks.cpu())\n    \n    all_preds = torch.cat(all_preds)\n    all_targets = torch.cat(all_targets)\n    \n    # Find optimal threshold for validation data (training phase)\n    optimal_threshold, best_f05, best_dice = find_optimal_threshold(\n        all_preds, all_targets, \"Training Validation\"\n    )\n    \n    # Calculate metrics with optimal threshold\n    final_metrics = calculate_metrics(all_preds, all_targets, optimal_threshold)\n    final_metrics['optimal_threshold'] = optimal_threshold\n    \n    return running_loss / len(dataloader), final_metrics, all_preds, all_targets\n\n# Test evaluation function (separate phase)\ndef evaluate_test_data(model, test_volume, test_mask, fragment_name, device):\n    \"\"\"Evaluate on test data with test-specific threshold optimization\"\"\"\n    model.eval()\n    \n    # Prepare test volume\n    num_slices = test_volume.shape[0]\n    if num_slices > 4:\n        selected_indices = [0, num_slices//3, 2*num_slices//3, -1]\n        test_volume_processed = test_volume[selected_indices]\n    else:\n        test_volume_processed = test_volume\n    \n    test_volume_tensor = torch.FloatTensor(test_volume_processed).unsqueeze(0).to(device)\n    \n    with torch.no_grad():\n        test_output = model(test_volume_tensor)\n        test_pred = test_output.sigmoid().cpu()\n    \n    test_pred_np = test_pred.numpy()[0, 0]\n    test_mask_tensor = torch.FloatTensor(test_mask).unsqueeze(0).unsqueeze(0)\n    \n    # Find test-specific optimal threshold\n    print(f\"\\n=== TEST PHASE: Searching optimal threshold for {fragment_name} ===\")\n    optimal_threshold, best_f05, best_dice = find_optimal_threshold(\n        test_pred, test_mask_tensor, f\"Test {fragment_name}\"\n    )\n    \n    # Calculate metrics with test-optimized threshold\n    test_metrics = calculate_metrics(test_pred, test_mask_tensor, optimal_threshold)\n    test_metrics['optimal_threshold'] = optimal_threshold\n    \n    # Apply post-processing\n    final_prediction = apply_post_processing(test_pred_np, optimal_threshold, config.edge_margin)\n    \n    return test_metrics, test_pred_np, final_prediction\n\n# Visualization function\ndef visualize_comparison(volume, true_mask, raw_pred, processed_pred, metrics, epoch, fold, phase, save_path=None):\n    fig, axes = plt.subplots(2, 4, figsize=(24, 12))\n    \n    middle_slice = volume[len(volume) // 2]\n    axes[0, 0].imshow(middle_slice, cmap='gray')\n    axes[0, 0].set_title('Input Volume (Middle Slice)')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(true_mask, cmap='gray')\n    axes[0, 1].set_title('Ground Truth Mask')\n    axes[0, 1].axis('off')\n    \n    im_raw = axes[0, 2].imshow(raw_pred, cmap='jet', vmin=0, vmax=1)\n    axes[0, 2].set_title('Raw Prediction (Probability)')\n    axes[0, 2].axis('off')\n    plt.colorbar(im_raw, ax=axes[0, 2], fraction=0.046)\n    \n    axes[0, 3].imshow(processed_pred, cmap='jet')\n    axes[0, 3].set_title('Processed Prediction (Binary)')\n    axes[0, 3].axis('off')\n    \n    axes[1, 0].imshow(middle_slice, cmap='gray')\n    axes[1, 0].imshow(raw_pred, cmap='jet', alpha=0.5)\n    axes[1, 0].set_title('Raw Prediction Overlay')\n    axes[1, 0].axis('off')\n    \n    axes[1, 1].imshow(middle_slice, cmap='gray')\n    axes[1, 1].imshow(processed_pred, cmap='jet', alpha=0.5)\n    axes[1, 1].set_title('Processed Prediction Overlay')\n    axes[1, 1].axis('off')\n    \n    diff = np.abs(processed_pred - true_mask)\n    im_diff = axes[1, 2].imshow(diff, cmap='hot')\n    axes[1, 2].set_title('Difference Map\\n(Red = Error)')\n    axes[1, 2].axis('off')\n    plt.colorbar(im_diff, ax=axes[1, 2], fraction=0.046)\n    \n    axes[1, 3].axis('off')\n    metrics_text = f\"\"\"Fold {fold}, Epoch {epoch}\n{phase} Metrics:\n────────────────────\nPrecision: {metrics['precision']:.4f}\nRecall:    {metrics['recall']:.4f}\nDice:      {metrics['dice']:.4f}\nF0.5:      {metrics['f0.5']:.4f}\nOptimal Threshold: {metrics['optimal_threshold']:.3f}\n\nThreshold optimized for: {phase} data\"\"\"\n    \n    axes[1, 3].text(0.1, 0.9, metrics_text, transform=axes[1, 3].transAxes, \n                   fontsize=12, verticalalignment='top', fontfamily='monospace')\n    \n    plt.suptitle(f'Ink Detection - {phase} Results (Fold {fold}, Epoch {epoch})', fontsize=16, y=0.95)\n    plt.tight_layout()\n    \n    if save_path:\n        plt.savefig(save_path, dpi=300, bbox_inches='tight')\n        print(f\"Visualization saved to: {save_path}\")\n    \n    plt.show()\n\n# Main training loop\ndef main():\n    print(\"Loading and preprocessing data...\")\n    \n    # Load training data\n    train_volumes = []\n    train_masks = []\n    \n    for train_dir in config.train_dirs:\n        print(f\"Loading data from {train_dir}...\")\n        volume = load_and_preprocess_slices(train_dir, config.slice_range, \n                                          (config.img_size, config.img_size))\n        if volume is not None:\n            mask_path = os.path.join(train_dir, 'inklabels.png')\n            mask = load_mask(mask_path, (config.img_size, config.img_size))\n            \n            split_size = config.num_slices // 4\n            for i in range(4):\n                start_idx = i * split_size\n                end_idx = min((i + 1) * split_size, config.num_slices)\n                \n                if end_idx - start_idx >= 3:\n                    volume_split = volume[start_idx:end_idx]\n                    train_volumes.append(volume_split)\n                    train_masks.append(mask)\n            print(f\"  Added {4} splits from {train_dir}\")\n        else:\n            print(f\"  No volume data found in {train_dir}\")\n    \n    print(f\"Loaded {len(train_volumes)} training volumes\")\n    \n    if len(train_volumes) == 0:\n        raise ValueError(\"No training data loaded! Check your data paths.\")\n    \n    # Load test data\n    print(f\"Loading test data from {config.test_dir}...\")\n    test_volume = load_and_preprocess_slices(config.test_dir, config.slice_range,\n                                           (config.img_size, config.img_size))\n    test_mask = load_mask(os.path.join(config.test_dir, 'inklabels.png'),\n                         (config.img_size, config.img_size))\n    \n    if test_volume is None:\n        raise ValueError(f\"No test data found in {config.test_dir}\")\n    \n    print(f\"Test volume shape: {test_volume.shape}\")\n    print(f\"Test mask shape: {test_mask.shape}\")\n    \n    # Prepare for cross-validation\n    kfold = ManualKFold(n_splits=min(config.num_folds, len(train_volumes)), shuffle=True, random_state=42)\n    \n    # Store results\n    fold_results = []\n    best_models = []\n    training_metrics_history = []\n    testing_metrics_history = []\n    \n    print(f\"Starting {min(config.num_folds, len(train_volumes))}-fold cross-validation...\")\n    \n    for fold, (train_idx, val_idx) in enumerate(kfold.split(train_volumes)):\n        print(f\"\\n{'='*50}\")\n        print(f\"FOLD {fold + 1}/{min(config.num_folds, len(train_volumes))}\")\n        print(f\"{'='*50}\")\n        \n        # Create datasets\n        train_dataset = VolumeDataset(\n            [train_volumes[i] for i in train_idx],\n            [train_masks[i] for i in train_idx],\n            is_train=True\n        )\n        val_dataset = VolumeDataset(\n            [train_volumes[i] for i in val_idx],\n            [train_masks[i] for i in val_idx],\n            is_train=False\n        )\n        \n        train_loader = DataLoader(train_dataset, batch_size=config.batch_size, \n                                shuffle=True, num_workers=0, pin_memory=True)\n        val_loader = DataLoader(val_dataset, batch_size=config.batch_size,\n                              shuffle=False, num_workers=0, pin_memory=True)\n        \n        # Initialize model\n        model = SegFormerWithSEResNet(\n            num_classes=config.num_classes, \n            backbone=config.backbone,\n            in_channels=4\n        ).to(config.device)\n        \n        # Training setup\n        criterion = nn.BCEWithLogitsLoss()\n        optimizer = optim.AdamW(model.parameters(), lr=config.lr, weight_decay=config.weight_decay)\n        scheduler = CosineAnnealingLR(optimizer, T_max=config.epochs)\n        \n        best_val_f05 = 0.0\n        best_model_state = None\n        fold_training_metrics = []\n        \n        print(f\"\\n📊 PHASE 1: TRAINING (Fold {fold + 1})\")\n        print(\"Optimizing threshold for VALIDATION data during training...\")\n        \n        # TRAINING PHASE\n        for epoch in range(min(5, config.epochs)):  # Reduced epochs for testing\n            print(f\"\\nEpoch {epoch+1}/{min(5, config.epochs)}\")\n            \n            # Train\n            train_loss = train_epoch(model, train_loader, optimizer, criterion, config.device)\n            \n            # Validate on validation data (training phase)\n            val_loss, val_metrics, val_preds, val_targets = validate_epoch_training(\n                model, val_loader, criterion, config.device\n            )\n            \n            scheduler.step()\n            \n            print(f\"Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}\")\n            print(f\"VAL Precision: {val_metrics['precision']:.4f}, Recall: {val_metrics['recall']:.4f}\")\n            print(f\"VAL Dice: {val_metrics['dice']:.4f}, F0.5: {val_metrics['f0.5']:.4f}\")\n            print(f\"VAL Optimal Threshold: {val_metrics['optimal_threshold']:.4f}\")\n            \n            # Store training phase metrics\n            fold_training_metrics.append({\n                'epoch': epoch + 1,\n                'train_loss': train_loss,\n                'val_loss': val_loss,\n                'val_metrics': val_metrics.copy()\n            })\n            \n            # Save best model based on validation F0.5\n            if val_metrics['f0.5'] > best_val_f05:\n                best_val_f05 = val_metrics['f0.5']\n                best_model_state = model.state_dict().copy()\n        \n        print(f\"\\n✅ TRAINING PHASE COMPLETE (Fold {fold + 1})\")\n        print(f\"Best Validation F0.5: {best_val_f05:.4f}\")\n        \n        # TESTING PHASE\n        print(f\"\\n🎯 PHASE 2: TESTING (Fold {fold + 1})\")\n        print(\"Now optimizing threshold for TEST data...\")\n        \n        # Load best model from training phase\n        model.load_state_dict(best_model_state)\n        \n        # Evaluate on test data with test-specific threshold optimization\n        test_metrics, test_raw_pred, test_processed_pred = evaluate_test_data(\n            model, test_volume, test_mask, f\"Fold_{fold+1}\", config.device\n        )\n        \n        print(f\"\\n📊 TEST RESULTS (Fold {fold + 1}):\")\n        print(f\"TEST Precision: {test_metrics['precision']:.4f}\")\n        print(f\"TEST Recall:    {test_metrics['recall']:.4f}\")\n        print(f\"TEST Dice:      {test_metrics['dice']:.4f}\")\n        print(f\"TEST F0.5:      {test_metrics['f0.5']:.4f}\")\n        print(f\"TEST Optimal Threshold: {test_metrics['optimal_threshold']:.4f}\")\n        \n        # Store testing phase metrics\n        fold_testing_metrics = {\n            'fold': fold + 1,\n            'test_metrics': test_metrics.copy(),\n            'val_f0.5': best_val_f05\n        }\n        \n        # Save visualization for test results\n        os.makedirs('visualizations', exist_ok=True)\n        viz_path = f'visualizations/fold_{fold+1}_test_results.png'\n        visualize_comparison(\n            test_volume, test_mask, test_raw_pred, test_processed_pred,\n            test_metrics, \"Final\", fold + 1, \"TEST\", viz_path\n        )\n        \n        # Store results\n        best_models.append(best_model_state)\n        fold_results.append(test_metrics['f0.5'])\n        training_metrics_history.append(fold_training_metrics)\n        testing_metrics_history.append(fold_testing_metrics)\n        \n        # Clean up\n        del model, optimizer, scheduler\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    # FINAL ENSEMBLE EVALUATION\n    if best_models:\n        print(f\"\\n{'='*60}\")\n        print(\"🎯 FINAL PHASE: ENSEMBLE TESTING\")\n        print(f\"{'='*60}\")\n        print(\"Evaluating ensemble model with test-specific threshold optimization...\")\n        \n        # Prepare test volume\n        num_slices = test_volume.shape[0]\n        if num_slices > 4:\n            selected_indices = [0, num_slices//3, 2*num_slices//3, -1]\n            test_volume_processed = test_volume[selected_indices]\n        else:\n            test_volume_processed = test_volume\n        \n        test_volume_tensor = torch.FloatTensor(test_volume_processed).unsqueeze(0)\n        test_mask_tensor = torch.FloatTensor(test_mask).unsqueeze(0).unsqueeze(0)\n        \n        all_test_predictions = []\n        \n        for fold, model_state in enumerate(best_models):\n            print(f\"Generating ensemble prediction from fold {fold + 1}...\")\n            \n            model = SegFormerWithSEResNet(\n                num_classes=config.num_classes, \n                backbone=config.backbone,\n                in_channels=4\n            ).to(config.device)\n            model.load_state_dict(model_state)\n            model.eval()\n            \n            with torch.no_grad():\n                test_output = model(test_volume_tensor.to(config.device))\n                test_pred = test_output.sigmoid().cpu()\n                all_test_predictions.append(test_pred.numpy())\n        \n        # Ensemble predictions\n        ensemble_pred = np.mean(all_test_predictions, axis=0)[0, 0]\n        \n        # Find optimal threshold for ensemble (TEST-SPECIFIC)\n        print(f\"\\n=== ENSEMBLE TEST: Searching optimal threshold ===\")\n        ensemble_pred_tensor = torch.FloatTensor(ensemble_pred).unsqueeze(0).unsqueeze(0)\n        optimal_threshold, best_f05, best_dice = find_optimal_threshold(\n            ensemble_pred_tensor, test_mask_tensor, \"Ensemble Test\"\n        )\n        \n        # Apply post-processing\n        final_prediction = apply_post_processing(ensemble_pred, optimal_threshold, config.edge_margin)\n        \n        # Calculate final metrics\n        final_metrics = calculate_metrics(ensemble_pred_tensor, test_mask_tensor, optimal_threshold)\n        final_metrics['optimal_threshold'] = optimal_threshold\n        \n        print(f\"\\n🎉 FINAL ENSEMBLE TEST RESULTS:\")\n        print(f\"Precision: {final_metrics['precision']:.4f}\")\n        print(f\"Recall:    {final_metrics['recall']:.4f}\")\n        print(f\"Dice:      {final_metrics['dice']:.4f}\")\n        print(f\"F0.5:      {final_metrics['f0.5']:.4f}\")\n        print(f\"Optimal Threshold: {final_metrics['optimal_threshold']:.4f}\")\n        \n        # Save final visualization\n        viz_path = 'visualizations/final_ensemble_test_results.png'\n        visualize_comparison(\n            test_volume, test_mask, ensemble_pred, final_prediction,\n            final_metrics, \"Final\", \"Ensemble\", \"TEST\", viz_path\n        )\n    else:\n        print(\"No models available for ensemble evaluation!\")\n        final_metrics = {'f0.5': 0.0, 'dice': 0.0, 'precision': 0.0, 'recall': 0.0, 'optimal_threshold': 0.5}\n        ensemble_pred = np.zeros_like(test_mask)\n        final_prediction = np.zeros_like(test_mask)\n    \n    # Print comprehensive summary\n    print(f\"\\n{'='*60}\")\n    print(\"📊 COMPREHENSIVE RESULTS SUMMARY\")\n    print(f\"{'='*60}\")\n    \n    if training_metrics_history:\n        print(f\"\\n📈 TRAINING PHASE METRICS (Validation Data):\")\n        for fold, train_metrics in enumerate(training_metrics_history):\n            best_train_epoch = max(train_metrics, key=lambda x: x['val_metrics']['f0.5'])\n            print(f\"Fold {fold + 1}: Best Val F0.5 = {best_train_epoch['val_metrics']['f0.5']:.4f} \"\n                  f\"(Epoch {best_train_epoch['epoch']}, Threshold: {best_train_epoch['val_metrics']['optimal_threshold']:.3f})\")\n    \n    if testing_metrics_history:\n        print(f\"\\n🎯 TESTING PHASE METRICS (Test Data):\")\n        for fold, test_metrics in enumerate(testing_metrics_history):\n            print(f\"Fold {fold + 1}: Test F0.5 = {test_metrics['test_metrics']['f0.5']:.4f} \"\n                  f\"(Threshold: {test_metrics['test_metrics']['optimal_threshold']:.3f})\")\n    \n    print(f\"\\n📊 FINAL ENSEMBLE:\")\n    print(f\"Test F0.5: {final_metrics['f0.5']:.4f}\")\n    print(f\"Test Dice: {final_metrics['dice']:.4f}\")\n    print(f\"Optimal Threshold: {final_metrics['optimal_threshold']:.3f}\")\n    \n    return final_metrics, ensemble_pred, final_prediction, training_metrics_history, testing_metrics_history\n\n# Memory optimization\ndef optimize_memory():\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    gc.collect()\n\n# Run the main function\nif __name__ == \"__main__\":\n    torch.manual_seed(42)\n    np.random.seed(42)\n    \n    try:\n        final_metrics, ensemble_pred, final_prediction, training_history, testing_history = main()\n        \n        print(f\"\\n🎯 FINAL SUMMARY:\")\n        print(f\"Best Ensemble Test F0.5: {final_metrics['f0.5']:.4f}\")\n        print(f\"Best Ensemble Test Dice: {final_metrics['dice']:.4f}\")\n        print(f\"Final Optimal Threshold: {final_metrics['optimal_threshold']:.3f}\")\n        \n    except RuntimeError as e:\n        if \"out of memory\" in str(e):\n            print(\"OOM error detected. Trying to optimize...\")\n            optimize_memory()\n            config.batch_size = max(1, config.batch_size // 2)\n            print(f\"Reduced batch size to {config.batch_size}\")\n            final_metrics, ensemble_pred, final_prediction, training_history, testing_history = main()\n        else:\n            raise e","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install scikit-learn\n!pip install segmentation-models-pytorch==0.2.0\n!pip install torch torchvision\n!pip install opencv-python\n!pip install scikit-image\n!pip install tqdm\n!pip install matplotlib\n!pip install scipy\n!pip install networkx\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage, spatial, interpolate, signal\nfrom scipy.ndimage import binary_fill_holes, binary_closing, gaussian_filter, convolve\nfrom scipy.ndimage import label as ndi_label\nfrom scipy.spatial import KDTree, cKDTree\nfrom scipy.optimize import minimize\nfrom scipy.interpolate import splprep, splev, interp1d\nfrom skimage.morphology import remove_small_objects, skeletonize, thin, medial_axis\nfrom skimage.filters import threshold_otsu, gabor, frangi, hessian, sobel\nfrom skimage.feature import canny, peak_local_max, structure_tensor, hessian_matrix, hessian_matrix_eigvals\nfrom skimage.transform import warp_polar, rotate\nfrom skimage.segmentation import watershed\nfrom skimage.measure import regionprops, label, LineModelND, ransac\nfrom skimage.graph import route_through_array\nimport warnings\nimport random\nimport math\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import precision_score, recall_score, f1_score, precision_recall_curve\nfrom sklearn.cluster import DBSCAN\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport networkx as nx\nfrom collections import defaultdict, deque\nfrom itertools import combinations\nimport json\nimport time\n\nwarnings.filterwarnings('ignore')\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\nclass Config:\n    # Data paths\n    train_paths = [\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/2',\n        '/kaggle/input/vesuvius-challenge-ink-detection/train/3'\n    ]\n    test_path = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\n    \n    # Volume parameters\n    img_size = 384\n    slices = list(range(12, 31))\n    patch_size = 48  # Increased for better context\n    overlap = 12\n    \n    # Surface estimation\n    surface_iso_level = 0.3\n    surface_smooth_sigma = 1.0\n    surface_neighborhood_radius = 5\n    \n    # Flattening parameters\n    flatten_thickness = 4\n    flatten_resolution = 0.5\n    \n    # Ridge enhancement\n    ridge_scales = [1.0, 2.0, 3.0]  # More scales\n    ridge_alpha = 0.5\n    ridge_beta = 0.5\n    ridge_gamma = 15\n    \n    # Seed detection - stricter thresholds\n    seed_min_vesselness = 0.5  # Increased\n    seed_min_intensity = 0.15  # Lower - ink is darker\n    seed_min_distance = 5\n    \n    # Centerline tracing\n    trace_step_size = 0.5\n    trace_max_length = 50\n    trace_min_vesselness = 0.2  # Increased\n    trace_smoothness_weight = 0.1\n    \n    # Stroke scoring - stricter\n    contrast_weight = 0.3\n    tangent_weight = 0.2\n    aspect_weight = 0.2\n    consistency_weight = 0.3\n    min_score_threshold = 0.6  # Increased\n    \n    # Stroke to mask conversion\n    stroke_width = 2.0  # Width in pixels for stroke rendering\n    min_stroke_length = 8.0  # Minimum stroke length\n    \n    # Memory management\n    max_patches_per_batch = 30\n    use_gradient_checkpointing = False\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 1\n    min_score_threshold_vis = 0.5\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nconfig = Config()\n\n# ============================================================================\n# SURFACE ESTIMATION MODULE\n# ============================================================================\nclass SurfaceEstimator:\n    \"\"\"Estimate surface where ink resides using local morphology and PCA\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        \n    def estimate_surface_volume(self, volume):\n        \"\"\"Estimate surface position for each (x,y) coordinate\"\"\"\n        depth = volume.shape[0]\n        height = volume.shape[1]\n        width = volume.shape[2]\n        \n        # Use intensity-weighted depth for better surface estimation\n        surface_pos = np.zeros((height, width))\n        \n        for y in range(height):\n            for x in range(width):\n                # Find depth with maximum intensity gradient\n                intensity_profile = volume[:, y, x]\n                \n                # Compute gradient\n                grad = np.gradient(intensity_profile)\n                grad_mag = np.abs(grad)\n                \n                # Weight by intensity (ink is darker)\n                weight = 1.0 - intensity_profile\n                weighted_grad = grad_mag * weight\n                \n                # Find peak gradient\n                if len(weighted_grad) > 0:\n                    surface_pos[y, x] = np.argmax(weighted_grad)\n        \n        # Smooth surface\n        surface_pos = gaussian_filter(surface_pos.astype(float), \n                                     sigma=self.config.surface_smooth_sigma)\n        \n        return surface_pos\n\n# ============================================================================\n# SURFACE-CONSTRAINED FLATTENING\n# ============================================================================\nclass SurfaceFlattener:\n    \"\"\"Flatten thin shell around surface into 2D patches\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        \n    def flatten_patch(self, volume, surface_pos, center, size):\n        \"\"\"Flatten patch to 2D image\"\"\"\n        # Ensure size is odd for symmetric patch\n        if size % 2 == 0:\n            size = size - 1\n        \n        half_size = size // 2\n        x_center, y_center = int(center[0]), int(center[1])\n        \n        height, width = volume.shape[1], volume.shape[2]\n        depth = volume.shape[0]\n        \n        patch_2d = np.zeros((size, size))\n        \n        for i in range(size):\n            for j in range(size):\n                x_img = x_center + j - half_size\n                y_img = y_center + i - half_size\n                \n                if 0 <= x_img < width and 0 <= y_img < height:\n                    # Get surface position at this point\n                    z_surface = surface_pos[y_img, x_img]\n                    z_int = int(np.clip(z_surface, 0, depth - 1))\n                    \n                    # Sample around surface with more emphasis on surface\n                    values = []\n                    weights = []\n                    \n                    # Sample 5 slices centered at surface\n                    for k in range(-2, 3):\n                        z_sample = z_int + k\n                        if 0 <= z_sample < depth:\n                            values.append(volume[z_sample, y_img, x_img])\n                            # Gaussian weight centered at surface\n                            weight = np.exp(-(k**2) / 2.0)\n                            weights.append(weight)\n                    \n                    if values:\n                        values = np.array(values)\n                        weights = np.array(weights)\n                        patch_2d[i, j] = np.average(values, weights=weights)\n        \n        # Enhance contrast\n        patch_2d = self._enhance_contrast(patch_2d)\n        \n        return patch_2d\n    \n    def _enhance_contrast(self, image):\n        \"\"\"Enhance contrast of flattened patch\"\"\"\n        # Normalize\n        if image.max() > image.min():\n            image_norm = (image - image.min()) / (image.max() - image.min() + 1e-8)\n        else:\n            image_norm = image\n        \n        # Apply CLAHE-like enhancement\n        image_eq = np.power(image_norm, 0.7)  # Gamma correction\n        \n        return image_eq\n\n# ============================================================================\n# RIDGE/TUBULAR ENHANCEMENT\n# ============================================================================\nclass RidgeEnhancer:\n    \"\"\"Enhance ink-like strokes using multi-scale ridge filters\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        \n    def multi_scale_ridge_enhancement(self, image):\n        \"\"\"Combine ridge responses at multiple scales\"\"\"\n        if image.size == 0:\n            return np.zeros_like(image)\n        \n        responses = []\n        \n        for scale in self.config.ridge_scales:\n            # Hessian-based ridge detection\n            hessian_response = self._hessian_ridge_detection(image, scale)\n            responses.append(hessian_response)\n        \n        # Combine scales (maximum response)\n        if responses:\n            combined = np.max(np.stack(responses), axis=0)\n        else:\n            combined = np.zeros_like(image)\n        \n        # Post-process\n        combined = self._post_process_vesselness(combined)\n        \n        return combined\n    \n    def _hessian_ridge_detection(self, image, scale):\n        \"\"\"Hessian-based ridge detection\"\"\"\n        # Smooth image\n        smoothed = gaussian_filter(image, scale)\n        \n        # Compute second derivatives\n        gy, gx = np.gradient(smoothed)\n        gyy, gyx = np.gradient(gy)\n        gxy, gxx = np.gradient(gx)\n        \n        # Compute eigenvalues of Hessian\n        vesselness = np.zeros_like(image)\n        \n        for i in range(image.shape[0]):\n            for j in range(image.shape[1]):\n                H = np.array([[gxx[i, j], gxy[i, j]],\n                             [gyx[i, j], gyy[i, j]]])\n                \n                try:\n                    eigvals = np.linalg.eigvalsh(H)\n                    lambda1, lambda2 = sorted(np.abs(eigvals))\n                    \n                    # Dark ridges (ink) have negative eigenvalues\n                    if eigvals[0] < 0 and eigvals[1] < 0:\n                        # Ridge measure\n                        Rb = lambda1 / (lambda2 + 1e-8)\n                        S = np.sqrt(lambda1**2 + lambda2**2)\n                        \n                        # Frangi-like measure\n                        vesselness[i, j] = np.exp(-Rb**2 / (2 * self.config.ridge_alpha**2)) * \\\n                                          (1 - np.exp(-S**2 / (2 * self.config.ridge_beta**2)))\n                except:\n                    continue\n        \n        return vesselness\n    \n    def _post_process_vesselness(self, vesselness):\n        \"\"\"Post-process vesselness map\"\"\"\n        # Remove noise\n        vesselness = gaussian_filter(vesselness, 0.5)\n        \n        # Threshold\n        threshold = np.percentile(vesselness, 70)\n        vesselness[vesselness < threshold] = 0\n        \n        # Normalize\n        if vesselness.max() > vesselness.min():\n            vesselness = (vesselness - vesselness.min()) / (vesselness.max() - vesselness.min() + 1e-8)\n        \n        return vesselness\n    \n    def compute_ridge_orientation(self, image, vesselness):\n        \"\"\"Compute orientation of ridges from gradient\"\"\"\n        # Compute gradients\n        gy, gx = np.gradient(image)\n        \n        # Smooth gradients\n        gx_smooth = gaussian_filter(gx, 1.0)\n        gy_smooth = gaussian_filter(gy, 1.0)\n        \n        # Orientation perpendicular to gradient (along ridge)\n        orientation = np.arctan2(gy_smooth, gx_smooth) + np.pi/2\n        \n        # Normalize to [0, 2π]\n        orientation = np.mod(orientation, 2*np.pi)\n        \n        return orientation, np.zeros_like(image)\n\n# ============================================================================\n# SEED DETECTION\n# ============================================================================\nclass SeedDetector:\n    \"\"\"Detect high-confidence seeds for stroke tracing\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        \n    def detect_seeds(self, vesselness, patch_image):\n        \"\"\"Detect seed points for stroke tracing\"\"\"\n        if vesselness.size == 0:\n            return []\n        \n        # Threshold vesselness\n        vesselness_thresh = np.percentile(vesselness[vesselness > 0], 80)\n        vesselness_mask = vesselness > vesselness_thresh\n        \n        # Intensity threshold (ink is darker)\n        intensity_thresh = np.percentile(patch_image, 30)\n        intensity_mask = patch_image < intensity_thresh\n        \n        # Combined mask\n        combined_mask = vesselness_mask & intensity_mask\n        \n        # Find local maxima in vesselness within mask\n        seeds = []\n        height, width = vesselness.shape\n        \n        # Dilate mask slightly\n        dilated_mask = ndimage.binary_dilation(combined_mask, structure=np.ones((3,3)))\n        \n        for y in range(2, height-2):\n            for x in range(2, width-2):\n                if dilated_mask[y, x]:\n                    # Check if it's a local maximum in 5x5 neighborhood\n                    neighborhood = vesselness[y-2:y+3, x-2:x+3]\n                    if vesselness[y, x] >= neighborhood.max():\n                        seeds.append((x, y, vesselness[y, x]))\n        \n        # Sort by vesselness score\n        seeds.sort(key=lambda s: s[2], reverse=True)\n        \n        # Limit number of seeds and apply distance constraint\n        if seeds:\n            seeds = self._apply_distance_constraint(seeds)\n        \n        return seeds[:15]  # Limit to 15 best seeds\n    \n    def _apply_distance_constraint(self, seeds, min_distance=8):\n        \"\"\"Apply minimum distance constraint between seeds\"\"\"\n        if len(seeds) <= 1:\n            return seeds\n        \n        filtered_seeds = []\n        used_positions = []\n        \n        for x, y, score in seeds:\n            too_close = False\n            \n            for (x2, y2, _) in used_positions:\n                distance = np.sqrt((x - x2)**2 + (y - y2)**2)\n                if distance < min_distance:\n                    too_close = True\n                    break\n            \n            if not too_close:\n                filtered_seeds.append((x, y, score))\n                used_positions.append((x, y, score))\n        \n        return filtered_seeds\n\n# ============================================================================\n# CENTERLINE TRACING\n# ============================================================================\nclass CenterlineTracer:\n    \"\"\"Trace stroke centerlines from seeds\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        \n    def trace_from_seed(self, vesselness, orientation, seed, patch_image=None):\n        \"\"\"Trace centerline from seed point in both directions\"\"\"\n        x, y, _ = seed\n        height, width = vesselness.shape\n        \n        # Trace in both directions\n        forward_points = self._trace_direction(vesselness, orientation, (x, y), forward=True, image=patch_image)\n        backward_points = self._trace_direction(vesselness, orientation, (x, y), forward=False, image=patch_image)\n        \n        # Combine (backward needs to be reversed)\n        if backward_points:\n            backward_points = backward_points[::-1]\n        \n        # Combine points\n        if backward_points and forward_points:\n            points = backward_points[:-1] + forward_points  # Avoid duplicate middle point\n        elif backward_points:\n            points = backward_points\n        elif forward_points:\n            points = forward_points\n        else:\n            points = [(x, y)]\n        \n        # Filter points based on intensity\n        if patch_image is not None and len(points) > 2:\n            points = self._filter_by_intensity(points, patch_image)\n        \n        # Smooth polyline\n        if len(points) > 3:\n            points = self._smooth_polyline(points)\n        \n        return points\n    \n    def _trace_direction(self, vesselness, orientation, start, forward=True, image=None):\n        \"\"\"Trace in one direction\"\"\"\n        points = []\n        x, y = start\n        \n        height, width = vesselness.shape\n        visited = set()\n        \n        for step in range(self.config.trace_max_length):\n            # Mark as visited\n            visited.add((int(x), int(y)))\n            \n            # Get current orientation\n            x_int, y_int = int(round(x)), int(round(y))\n            if not (0 <= x_int < width and 0 <= y_int < height):\n                break\n            \n            orient = orientation[y_int, x_int]\n            \n            # Determine step direction\n            if forward:\n                dx = np.cos(orient) * self.config.trace_step_size\n                dy = np.sin(orient) * self.config.trace_step_size\n            else:\n                dx = -np.cos(orient) * self.config.trace_step_size\n                dy = -np.sin(orient) * self.config.trace_step_size\n            \n            # Take step\n            x_new = x + dx\n            y_new = y + dy\n            \n            # Check bounds\n            if not (0 <= x_new < width and 0 <= y_new < height):\n                break\n            \n            # Convert to integer for lookup\n            x_new_int, y_new_int = int(round(x_new)), int(round(y_new))\n            \n            # Check if visited\n            if (x_new_int, y_new_int) in visited:\n                break\n            \n            # Check vesselness threshold\n            if vesselness[y_new_int, x_new_int] < self.config.trace_min_vesselness:\n                break\n            \n            # Check intensity if image provided (ink should be dark)\n            if image is not None:\n                if image[y_new_int, x_new_int] > 0.5:  # Too bright for ink\n                    break\n            \n            # Add point\n            points.append((x_new, y_new))\n            x, y = x_new, y_new\n        \n        return points\n    \n    def _filter_by_intensity(self, points, image):\n        \"\"\"Filter points based on intensity (ink is dark)\"\"\"\n        filtered_points = []\n        \n        for x, y in points:\n            x_int, y_int = int(round(x)), int(round(y))\n            if 0 <= y_int < image.shape[0] and 0 <= x_int < image.shape[1]:\n                if image[y_int, x_int] < 0.6:  # Keep dark points\n                    filtered_points.append((x, y))\n        \n        return filtered_points if len(filtered_points) > 2 else points\n    \n    def _smooth_polyline(self, points, window_size=3):\n        \"\"\"Smooth polyline using moving average\"\"\"\n        if len(points) < window_size:\n            return points\n        \n        smoothed = []\n        for i in range(len(points)):\n            start = max(0, i - window_size // 2)\n            end = min(len(points), i + window_size // 2 + 1)\n            \n            window = points[start:end]\n            x_avg = np.mean([p[0] for p in window])\n            y_avg = np.mean([p[1] for p in window])\n            smoothed.append((x_avg, y_avg))\n        \n        return smoothed\n\n# ============================================================================\n# STROKE TO MASK CONVERSION\n# ============================================================================\nclass StrokeRenderer:\n    \"\"\"Convert strokes to binary mask\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        \n    def strokes_to_mask(self, strokes, volume_shape):\n        \"\"\"Convert list of strokes to binary mask\"\"\"\n        mask = np.zeros(volume_shape, dtype=np.uint8)\n        \n        for stroke in strokes:\n            if len(stroke) < 2:\n                continue\n                \n            # Render stroke as thick line\n            for i in range(len(stroke) - 1):\n                p1 = stroke[i]\n                p2 = stroke[i+1]\n                \n                # Draw line segment\n                self._draw_thick_line(mask, p1, p2)\n        \n        return mask\n    \n    def _draw_thick_line(self, mask, p1, p2):\n        \"\"\"Draw a thick line between two points\"\"\"\n        x1, y1, z1 = p1\n        x2, y2, z2 = p2\n        \n        # Number of interpolation points\n        num_points = max(2, int(np.sqrt((x2-x1)**2 + (y2-y1)**2 + (z2-z1)**2)))\n        \n        for t in np.linspace(0, 1, num_points):\n            # Interpolate\n            x = x1 + t * (x2 - x1)\n            y = y1 + t * (y2 - y1)\n            z = z1 + t * (z2 - z1)\n            \n            # Draw thick point\n            self._draw_thick_point(mask, (x, y, z))\n    \n    def _draw_thick_point(self, mask, center):\n        \"\"\"Draw a thick point (sphere)\"\"\"\n        x, y, z = center\n        radius = self.config.stroke_width / 2.0\n        \n        x_min = max(0, int(x - radius))\n        x_max = min(mask.shape[2], int(x + radius) + 1)\n        y_min = max(0, int(y - radius))\n        y_max = min(mask.shape[1], int(y + radius) + 1)\n        z_min = max(0, int(z - radius))\n        z_max = min(mask.shape[0], int(z + radius) + 1)\n        \n        for zi in range(z_min, z_max):\n            for yi in range(y_min, y_max):\n                for xi in range(x_min, x_max):\n                    dist = np.sqrt((xi - x)**2 + (yi - y)**2 + (zi - z)**2)\n                    if dist <= radius:\n                        mask[zi, yi, xi] = 1\n\n# ============================================================================\n# VISUALIZATION MODULE\n# ============================================================================\nclass Visualization:\n    \"\"\"Create comprehensive visualizations\"\"\"\n    \n    @staticmethod\n    def save_test_visualization(input_volume, prediction_mask, ground_truth_mask, \n                               strokes, scores, save_path, fragment_name=\"test\", metrics=None):\n        \"\"\"Save comprehensive test visualization\"\"\"\n        os.makedirs(os.path.dirname(save_path) if os.path.dirname(save_path) else '.', exist_ok=True)\n        \n        # Get middle slices\n        depth = input_volume.shape[0]\n        middle_slice = depth // 2\n        \n        input_slice = input_volume[middle_slice]\n        pred_slice = prediction_mask[middle_slice]\n        gt_slice = ground_truth_mask[middle_slice]\n        \n        # Create figure\n        fig = plt.figure(figsize=(20, 15))\n        \n        # 1. Input volume\n        ax1 = plt.subplot(3, 4, 1)\n        ax1.imshow(input_slice, cmap='gray')\n        ax1.set_title(f'Input Slice {middle_slice}')\n        ax1.axis('off')\n        \n        # 2. Predicted mask\n        ax2 = plt.subplot(3, 4, 2)\n        ax2.imshow(pred_slice, cmap='hot')\n        ax2.set_title('Predicted Ink')\n        ax2.axis('off')\n        \n        # 3. Ground truth\n        ax3 = plt.subplot(3, 4, 3)\n        ax3.imshow(gt_slice, cmap='hot')\n        ax3.set_title('Ground Truth')\n        ax3.axis('off')\n        \n        # 4. Overlay\n        ax4 = plt.subplot(3, 4, 4)\n        overlay = np.zeros((input_slice.shape[0], input_slice.shape[1], 3))\n        overlay[:, :, 0] = pred_slice * 0.8  # Red for prediction\n        overlay[:, :, 1] = gt_slice * 0.8    # Green for ground truth\n        overlay[:, :, 2] = input_slice * 0.5 # Blue for input\n        ax4.imshow(np.clip(overlay, 0, 1))\n        ax4.set_title('Overlay (Pred=Red, GT=Green)')\n        ax4.axis('off')\n        \n        # 5. 3D strokes visualization (projection)\n        ax5 = plt.subplot(3, 4, (5, 8))\n        if strokes:\n            # Color by score\n            cmap = plt.cm.viridis\n            if scores:\n                norm = plt.Normalize(min(scores), max(scores))\n            \n            for i, stroke in enumerate(strokes[:50]):  # Show first 50 strokes\n                if len(stroke) > 1:\n                    stroke_array = np.array(stroke)\n                    color = 'red'\n                    if scores and i < len(scores):\n                        color = cmap(norm(scores[i]))\n                    \n                    ax5.plot(stroke_array[:, 0], stroke_array[:, 1], \n                            color=color, linewidth=1.0, alpha=0.6)\n        \n        ax5.set_xlabel('X')\n        ax5.set_ylabel('Y')\n        ax5.set_title(f'Detected Strokes (n={len(strokes)})')\n        ax5.set_aspect('equal')\n        ax5.grid(True, alpha=0.3)\n        \n        # 6. Metrics text\n        ax6 = plt.subplot(3, 4, (9, 12))\n        ax6.axis('off')\n        \n        metrics_text = f\"FRAGMENT: {fragment_name}\\n\"\n        metrics_text += \"=\"*50 + \"\\n\"\n        \n        if metrics:\n            metrics_text += f\"Precision:   {metrics['precision']:.4f}\\n\"\n            metrics_text += f\"Recall:      {metrics['recall']:.4f}\\n\"\n            metrics_text += f\"F0.5 Score:  {metrics['f05']:.4f}\\n\"\n            metrics_text += f\"Dice Score:  {metrics['dice']:.4f}\\n\"\n            metrics_text += f\"True Positives:  {metrics['true_positives']}\\n\"\n            metrics_text += f\"False Positives: {metrics['false_positives']}\\n\"\n            metrics_text += f\"False Negatives: {metrics['false_negatives']}\\n\"\n        \n        metrics_text += f\"\\nStroke Statistics:\\n\"\n        metrics_text += f\"Total strokes: {len(strokes)}\\n\"\n        if strokes:\n            avg_length = np.mean([len(s) for s in strokes])\n            metrics_text += f\"Avg length: {avg_length:.1f} voxels\\n\"\n            if scores:\n                avg_score = np.mean(scores)\n                metrics_text += f\"Avg score:  {avg_score:.3f}\\n\"\n        \n        metrics_text += f\"\\nVolume Info:\\n\"\n        metrics_text += f\"Shape: {input_volume.shape}\\n\"\n        metrics_text += f\"Slice shown: {middle_slice}/{depth}\\n\"\n        \n        ax6.text(0.05, 0.95, metrics_text, fontsize=10, fontfamily='monospace',\n                verticalalignment='top', transform=ax6.transAxes,\n                bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"lightblue\", alpha=0.7))\n        \n        plt.suptitle(f'Stroke Detection Results - {fragment_name}', fontsize=16, fontweight='bold')\n        plt.tight_layout()\n        plt.savefig(save_path, dpi=120, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"Saved comprehensive visualization to {save_path}\")\n\n# ============================================================================\n# MAIN STROKE DETECTION PIPELINE\n# ============================================================================\nclass StrokeDetectionPipeline:\n    \"\"\"Main pipeline for stroke detection on 3D scroll surfaces\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        self.surface_estimator = SurfaceEstimator(config)\n        self.flattener = SurfaceFlattener(config)\n        self.ridge_enhancer = RidgeEnhancer(config)\n        self.seed_detector = SeedDetector(config)\n        self.tracer = CenterlineTracer(config)\n        self.renderer = StrokeRenderer(config)\n        self.visualizer = Visualization()\n        \n    def process_volume(self, volume, mask=None, fragment_name=\"unknown\"):\n        \"\"\"Main processing pipeline\"\"\"\n        print(f\"Processing volume of shape {volume.shape}\")\n        \n        # Step 1: Surface estimation\n        print(\"Step 1: Estimating surface...\")\n        try:\n            surface_pos = self.surface_estimator.estimate_surface_volume(volume)\n            print(f\"Surface estimation complete. Surface shape: {surface_pos.shape}\")\n        except Exception as e:\n            print(f\"Error in surface estimation: {e}\")\n            return [], [], np.zeros_like(volume, dtype=np.uint8)\n        \n        # Step 2: Extract patches\n        print(\"Step 2: Extracting patches...\")\n        patches = self._extract_patches(volume, surface_pos, mask)\n        \n        if not patches:\n            print(\"No patches extracted!\")\n            return [], [], np.zeros_like(volume, dtype=np.uint8)\n        \n        # Step 3-6: Process each patch\n        all_strokes = []\n        all_scores = []\n        \n        print(f\"Step 3-6: Processing {len(patches)} patches...\")\n        for i, (patch, patch_center) in enumerate(tqdm(patches, desc=\"Processing patches\")):\n            if i >= self.config.max_patches_per_batch:\n                print(f\"Reached max patches per batch ({self.config.max_patches_per_batch})\")\n                break\n            \n            try:\n                # Step 3: Ridge enhancement\n                vesselness = self.ridge_enhancer.multi_scale_ridge_enhancement(patch)\n                orientation, _ = self.ridge_enhancer.compute_ridge_orientation(patch, vesselness)\n                \n                # Step 4: Seed detection\n                seeds = self.seed_detector.detect_seeds(vesselness, patch)\n                \n                # Step 5: Centerline tracing\n                patch_strokes = []\n                for seed in seeds:\n                    points_2d = self.tracer.trace_from_seed(vesselness, orientation, seed, patch)\n                    \n                    if len(points_2d) > 3:  # Minimum 4 points for a valid stroke\n                        # Map to 3D using surface position\n                        points_3d = self._map_to_3d(points_2d, surface_pos, patch_center)\n                        \n                        if points_3d:\n                            # Score stroke (simple length-based score)\n                            score = min(len(points_3d) / 50.0, 1.0)\n                            if score >= self.config.min_score_threshold:\n                                patch_strokes.append(points_3d)\n                                all_scores.append(score)\n                \n                all_strokes.extend(patch_strokes)\n                        \n            except Exception as e:\n                print(f\"Error processing patch {i}: {e}\")\n                continue\n        \n        print(f\"Detected {len(all_strokes)} strokes\")\n        \n        # Filter short strokes\n        filtered_strokes = []\n        filtered_scores = []\n        for stroke, score in zip(all_strokes, all_scores):\n            if len(stroke) >= self.config.min_stroke_length:\n                filtered_strokes.append(stroke)\n                filtered_scores.append(score)\n        \n        print(f\"After length filtering: {len(filtered_strokes)} strokes\")\n        \n        # Convert strokes to mask\n        prediction_mask = self.renderer.strokes_to_mask(filtered_strokes, volume.shape)\n        \n        return filtered_strokes, filtered_scores, prediction_mask\n    \n    def _extract_patches(self, volume, surface_pos, mask=None):\n        \"\"\"Extract patches from volume focusing on potential ink regions\"\"\"\n        height, width = surface_pos.shape\n        patches = []\n        \n        step = max(1, self.config.patch_size - self.config.overlap)\n        \n        # Calculate grid\n        num_y = (height - self.config.patch_size) // step + 1\n        num_x = (width - self.config.patch_size) // step + 1\n        \n        print(f\"Grid: {num_y}x{num_x} possible patches\")\n        \n        # Focus on regions with ink (if mask provided)\n        if mask is not None:\n            # Find ink regions\n            ink_coords = np.argwhere(mask > 0)\n            if len(ink_coords) > 0:\n                # Sample patches centered on ink regions\n                np.random.shuffle(ink_coords)\n                for i, (y, x) in enumerate(ink_coords[:self.config.max_patches_per_batch * 2]):\n                    # Adjust to patch center\n                    y_center = max(self.config.patch_size//2, \n                                 min(height - self.config.patch_size//2 - 1, y))\n                    x_center = max(self.config.patch_size//2,\n                                 min(width - self.config.patch_size//2 - 1, x))\n                    \n                    y_start = y_center - self.config.patch_size//2\n                    x_start = x_center - self.config.patch_size//2\n                    \n                    patch = self.flattener.flatten_patch(volume, surface_pos, \n                                                        (x_start, y_start), \n                                                        self.config.patch_size)\n                    patches.append((patch, (x_start, y_start)))\n        \n        # If no mask or not enough patches, sample regularly\n        if len(patches) < self.config.max_patches_per_batch:\n            for y in range(0, height - self.config.patch_size + 1, step * 2):\n                for x in range(0, width - self.config.patch_size + 1, step * 2):\n                    if len(patches) >= self.config.max_patches_per_batch:\n                        break\n                    \n                    patch = self.flattener.flatten_patch(volume, surface_pos, \n                                                        (x, y), self.config.patch_size)\n                    patches.append((patch, (x, y)))\n        \n        print(f\"Extracted {len(patches)} valid patches\")\n        return patches[:self.config.max_patches_per_batch]\n    \n    def _map_to_3d(self, points_2d, surface_pos, patch_center):\n        \"\"\"Map 2D points to 3D using surface position\"\"\"\n        points_3d = []\n        \n        half_size = self.config.patch_size // 2\n        patch_x, patch_y = patch_center\n        \n        for x_2d, y_2d in points_2d:\n            # Convert to image coordinates\n            x_img = patch_x + (x_2d - half_size)\n            y_img = patch_y + (y_2d - half_size)\n            \n            x_int = int(round(x_img))\n            y_int = int(round(y_img))\n            \n            if (0 <= x_int < surface_pos.shape[1] and \n                0 <= y_int < surface_pos.shape[0]):\n                z = surface_pos[y_int, x_int]\n                points_3d.append((x_img, y_img, z))\n        \n        return points_3d\n    \n    def evaluate_strokes(self, detected_strokes, ground_truth_mask):\n        \"\"\"Evaluate stroke detection against ground truth\"\"\"\n        if not detected_strokes:\n            return {\n                'precision': 0.0,\n                'recall': 0.0,\n                'f05': 0.0,\n                'dice': 0.0,\n                'true_positives': 0,\n                'false_positives': 0,\n                'false_negatives': 0,\n                'num_strokes': 0,\n                'avg_stroke_length': 0.0\n            }\n        \n        # Convert strokes to mask\n        prediction_mask = self.renderer.strokes_to_mask(detected_strokes, ground_truth_mask.shape)\n        \n        # Ensure same dtype\n        prediction_mask = prediction_mask.astype(np.uint8)\n        ground_truth_mask = ground_truth_mask.astype(np.uint8)\n        \n        # Calculate metrics\n        true_positives = np.sum((prediction_mask == 1) & (ground_truth_mask == 1))\n        false_positives = np.sum((prediction_mask == 1) & (ground_truth_mask == 0))\n        false_negatives = np.sum((prediction_mask == 0) & (ground_truth_mask == 1))\n        \n        # Avoid division by zero\n        precision = true_positives / (true_positives + false_positives + 1e-8)\n        recall = true_positives / (true_positives + false_negatives + 1e-8)\n        \n        # F0.5 score (emphasizes precision)\n        if precision + recall > 0:\n            f05 = (1 + 0.5**2) * (precision * recall) / (0.5**2 * precision + recall + 1e-8)\n        else:\n            f05 = 0.0\n        \n        # Dice score\n        dice = 2 * true_positives / (2 * true_positives + false_positives + false_negatives + 1e-8)\n        \n        return {\n            'precision': float(precision),\n            'recall': float(recall),\n            'f05': float(f05),\n            'dice': float(dice),\n            'true_positives': int(true_positives),\n            'false_positives': int(false_positives),\n            'false_negatives': int(false_negatives),\n            'num_strokes': len(detected_strokes),\n            'avg_stroke_length': float(np.mean([len(s) for s in detected_strokes]) if detected_strokes else 0.0)\n        }\n    \n    def save_test_visualization(self, input_volume, strokes, scores, prediction_mask, \n                               ground_truth_mask, fragment_name, save_dir=\"test_visualizations\"):\n        \"\"\"Save comprehensive test visualization\"\"\"\n        os.makedirs(save_dir, exist_ok=True)\n        \n        # Evaluate to get metrics\n        metrics = self.evaluate_strokes(strokes, ground_truth_mask)\n        \n        # Create visualization\n        save_path = os.path.join(save_dir, f\"test_{fragment_name}.png\")\n        self.visualizer.save_test_visualization(\n            input_volume, prediction_mask, ground_truth_mask,\n            strokes, scores, save_path, fragment_name, metrics\n        )\n        \n        return metrics\n\n# ============================================================================\n# DATA LOADING FUNCTIONS\n# ============================================================================\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_normalized = img_resized.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        print(f\"Loaded volume with {volume.shape[0]} slices\")\n        return volume\n    print(f\"Failed to load volume from {base_path}\")\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.uint8)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio: {ink_ratio:.4f}\")\n            return mask_binary\n    print(f\"No mask found at {mask_path}\")\n    return None\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\ndef main():\n    \"\"\"Main execution function\"\"\"\n    print(\"Initializing Stroke Detection Pipeline...\")\n    print(f\"Using device: {config.device}\")\n    \n    # Load data\n    print(\"\\nLoading training data...\")\n    train_volumes = []\n    train_masks = []\n    fragment_names = []\n    \n    for path in config.train_paths:\n        volume = load_volume_data(path, config.slices)\n        mask = load_mask_data(path)\n        \n        if volume is not None and mask is not None:\n            train_volumes.append(volume)\n            train_masks.append(mask)\n            fragment_name = path.split('/')[-1]\n            fragment_names.append(fragment_name)\n            print(f\"Loaded fragment {fragment_name}\")\n    \n    if not train_volumes:\n        print(\"No training data loaded!\")\n        return\n    \n    # Initialize pipeline\n    pipeline = StrokeDetectionPipeline(config)\n    \n    # Process each volume\n    all_results = []\n    \n    for i, (volume, mask, fragment_name) in enumerate(zip(train_volumes, train_masks, fragment_names)):\n        print(f\"\\n{'='*60}\")\n        print(f\"Processing volume {i+1}/{len(train_volumes)} ({fragment_name})...\")\n        print(f\"Volume shape: {volume.shape}\")\n        print(f\"Mask shape: {mask.shape}\")\n        print(f\"{'='*60}\")\n        \n        try:\n            # Run pipeline\n            start_time = time.time()\n            strokes, scores, prediction_mask = pipeline.process_volume(volume, mask, fragment_name)\n            elapsed_time = time.time() - start_time\n            \n            print(f\"\\nProcessing completed in {elapsed_time:.1f} seconds\")\n            \n            # Create 3D mask from 2D mask\n            mask_3d = np.stack([mask] * volume.shape[0], axis=0)\n            \n            # Save visualization\n            if config.save_visualization:\n                print(\"\\nSaving test visualization...\")\n                metrics = pipeline.save_test_visualization(\n                    volume, strokes, scores, prediction_mask, \n                    mask_3d, fragment_name, \"test_visualizations\"\n                )\n                \n                print(f\"\\nEvaluation Metrics for {fragment_name}:\")\n                print(f\"  Precision:    {metrics['precision']:.4f}\")\n                print(f\"  Recall:       {metrics['recall']:.4f}\")\n                print(f\"  F0.5 Score:   {metrics['f05']:.4f}\")\n                print(f\"  Dice Score:   {metrics['dice']:.4f}\")\n                print(f\"  True Positives:  {metrics['true_positives']}\")\n                print(f\"  False Positives: {metrics['false_positives']}\")\n                print(f\"  False Negatives: {metrics['false_negatives']}\")\n                print(f\"  Detected strokes: {metrics['num_strokes']}\")\n                print(f\"  Average length:   {metrics['avg_stroke_length']:.1f} voxels\")\n                \n                all_results.append(metrics)\n            else:\n                print(\"Visualization saving disabled in config\")\n            \n            # Save strokes to file\n            if strokes:\n                pipeline.save_strokes_to_file(strokes, scores, f\"strokes_{fragment_name}.json\")\n            else:\n                print(\"No strokes to save\")\n                \n        except Exception as e:\n            print(f\"Error processing volume {i+1}: {e}\")\n            import traceback\n            traceback.print_exc()\n    \n    # Summary\n    if all_results:\n        print(\"\\n\" + \"=\"*60)\n        print(\"OVERALL RESULTS SUMMARY\")\n        print(\"=\"*60)\n        \n        for metric in ['precision', 'recall', 'f05', 'dice']:\n            values = [r[metric] for r in all_results]\n            print(f\"{metric.capitalize():12s}: {np.mean(values):.4f} ± {np.std(values):.4f}\")\n        \n        total_strokes = sum([r['num_strokes'] for r in all_results])\n        total_tp = sum([r['true_positives'] for r in all_results])\n        total_fp = sum([r['false_positives'] for r in all_results])\n        total_fn = sum([r['false_negatives'] for r in all_results])\n        \n        print(f\"\\nTotal strokes detected: {total_strokes}\")\n        print(f\"Total true positives:   {total_tp}\")\n        print(f\"Total false positives:  {total_fp}\")\n        print(f\"Total false negatives:  {total_fn}\")\n        \n        # Calculate overall metrics\n        overall_precision = total_tp / (total_tp + total_fp + 1e-8)\n        overall_recall = total_tp / (total_tp + total_fn + 1e-8)\n        overall_f05 = (1 + 0.5**2) * (overall_precision * overall_recall) / (0.5**2 * overall_precision + overall_recall + 1e-8)\n        overall_dice = 2 * total_tp / (2 * total_tp + total_fp + total_fn + 1e-8)\n        \n        print(f\"\\nOverall Metrics:\")\n        print(f\"  Precision:  {overall_precision:.4f}\")\n        print(f\"  Recall:     {overall_recall:.4f}\")\n        print(f\"  F0.5 Score: {overall_f05:.4f}\")\n        print(f\"  Dice Score: {overall_dice:.4f}\")\n        \n        # Memory usage\n        if torch.cuda.is_available():\n            allocated = torch.cuda.memory_allocated() / 1e9\n            cached = torch.cuda.memory_reserved() / 1e9\n            print(f\"\\nGPU Memory Usage:\")\n            print(f\"  Allocated: {allocated:.2f} GB\")\n            print(f\"  Cached:    {cached:.2f} GB\")\n    else:\n        print(\"\\nNo evaluation results available\")\n    \n    print(\"\\nVisualizations saved to 'test_visualizations/' directory\")\n\n# ============================================================================\n# ADDITIONAL UTILITY FUNCTIONS\n# ============================================================================\ndef save_strokes_to_file(strokes, scores, filename):\n    \"\"\"Save strokes to JSON file\"\"\"\n    import json\n    \n    # Convert to serializable format\n    strokes_list = []\n    for stroke in strokes:\n        strokes_list.append([list(point) for point in stroke])\n    \n    output = {\n        'strokes': strokes_list,\n        'scores': scores,\n        'num_strokes': len(strokes),\n        'config': {\n            'patch_size': config.patch_size,\n            'min_score_threshold': config.min_score_threshold,\n            'min_stroke_length': config.min_stroke_length,\n            'stroke_width': config.stroke_width\n        }\n    }\n    \n    with open(filename, 'w') as f:\n        json.dump(output, f, indent=2)\n    \n    print(f\"Saved {len(strokes)} strokes to {filename}\")\n\n# ============================================================================\n# RUN MAIN PIPELINE\n# ============================================================================\nif __name__ == \"__main__\":\n    # Set random seeds for reproducibility\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    # Clear GPU memory\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        print(f\"GPU Memory cleared\")\n    \n    # Run main pipeline\n    try:\n        main()\n    except KeyboardInterrupt:\n        print(\"\\nPipeline interrupted by user\")\n    except Exception as e:\n        print(f\"\\nFatal error in main pipeline: {e}\")\n        import traceback\n        traceback.print_exc()\n    \n    print(\"\\nStroke detection pipeline completed!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install scikit-learn\n!pip install segmentation-models-pytorch==0.3.0\n!pip install torch torchvision\n!pip install opencv-python\n!pip install scikit-image\n!pip install tqdm\n!pip install matplotlib\n!pip install scipy\n!pip install joblib\n!pip install albumentations\n!pip install torch-fft\n!pip install joblib\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nimport warnings\nimport random\nimport math\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import precision_score, recall_score, f1_score, accuracy_score\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport torch.fft\nfrom joblib import Parallel, delayed\nfrom multiprocessing import cpu_count\nfrom collections import defaultdict\n\nwarnings.filterwarnings('ignore')\n\n# Configuration\nclass Config:\n    # Data paths\n    train_fragment = '2'  # Train on fragment 2\n    val_fragment = '3'    # Validate on fragment 3\n    test_fragment = '1'   # Test on fragment 1\n    \n    # Model parameters\n    img_size = 384\n    slices = list(range(12, 31))  # 19 slices\n    batch_size = 64  # Larger batch for pixel classification\n    epochs = 100\n    lr = 3e-4\n    \n    # Split parameters\n    rows_per_part = 10  # Split each slice into parts of 10 rows\n    parts_per_slice = None  # Will be calculated\n    \n    # Pixel classification parameters\n    num_features = 3  # x, y, intensity\n    hidden_dim = 128  # Increased from 64\n    use_context = True  # Use neighboring pixels as context\n    \n    # Training parameters\n    pos_weight = 2.0  # For handling class imbalance\n    \n    # Parallel processing\n    num_workers = max(1, cpu_count() - 2)  # Use all but 2 cores\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 5\n    min_score_threshold = 0.25\n    \n    # Early stopping\n    early_stopping_patience = 15\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Paths\n    @property\n    def train_path(self):\n        return f'/kaggle/input/vesuvius-challenge-ink-detection/train/{self.train_fragment}'\n    \n    @property\n    def val_path(self):\n        return f'/kaggle/input/vesuvius-challenge-ink-detection/train/{self.val_fragment}'\n    \n    @property\n    def test_path(self):\n        return f'/kaggle/input/vesuvius-challenge-ink-detection/train/{self.test_fragment}'\n\nconfig = Config()\nconfig.parts_per_slice = math.ceil(config.img_size / config.rows_per_part)\n\n# Data loading functions\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_normalized = img_resized.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            return mask_binary\n    return None\n\n# Enhanced feature extraction with context\ndef process_slice_part_with_context(volume_slice, mask_slice, slice_idx, part_idx, rows_per_part, img_size):\n    \"\"\"Process a single part with context from neighboring pixels\"\"\"\n    height, width = volume_slice.shape\n    \n    # Calculate start and end rows\n    start_row = part_idx * rows_per_part\n    end_row = min((part_idx + 1) * rows_per_part, height)\n    \n    # Extract the part with padding for context\n    pad_size = 1  # Look at immediate neighbors\n    padded_start = max(0, start_row - pad_size)\n    padded_end = min(height, end_row + pad_size)\n    \n    volume_part = volume_slice[padded_start:padded_end, :]\n    mask_part = mask_slice[padded_start:padded_end, :]\n    \n    # Get pixel coordinates within the actual region (without padding)\n    actual_start = start_row - padded_start\n    actual_end = actual_start + (end_row - start_row)\n    \n    # Get ink and non-ink pixel coordinates in actual region\n    ink_mask = mask_part[actual_start:actual_end, :] > 0.5\n    non_ink_mask = mask_part[actual_start:actual_end, :] <= 0.5\n    \n    ink_coords = np.argwhere(ink_mask)\n    non_ink_coords = np.argwhere(non_ink_mask)\n    \n    # Balance the samples - MORE AGGRESSIVE SAMPLING\n    n_samples = min(8000, len(ink_coords) * 3)  # Increased target samples\n    \n    if len(ink_coords) > 0:\n        # Sample more ink pixels (class imbalance)\n        n_ink = min(n_samples // 2, len(ink_coords))\n        n_non_ink = min(n_samples - n_ink, len(non_ink_coords))\n        \n        if len(ink_coords) > n_ink:\n            ink_indices = np.random.choice(len(ink_coords), n_ink, replace=False)\n            ink_coords = ink_coords[ink_indices]\n        \n        if len(non_ink_coords) > n_non_ink:\n            non_ink_indices = np.random.choice(len(non_ink_coords), n_non_ink, replace=False)\n            non_ink_coords = non_ink_coords[non_ink_indices]\n    else:\n        n_non_ink = min(n_samples, len(non_ink_coords))\n        if len(non_ink_coords) > n_non_ink:\n            non_ink_indices = np.random.choice(len(non_ink_coords), n_non_ink, replace=False)\n            non_ink_coords = non_ink_coords[non_ink_indices]\n        n_ink = 0\n    \n    features_list = []\n    labels_list = []\n    \n    # Process ink pixels\n    for rel_y, x in ink_coords:\n        # Convert to coordinates in padded region\n        padded_y = actual_start + rel_y + pad_size\n        global_y = start_row + rel_y\n        global_x = x\n        \n        # Get intensity at pixel\n        intensity = volume_part[padded_y, x]\n        \n        # Get context (3x3 neighborhood)\n        context = []\n        for dy in [-1, 0, 1]:\n            for dx in [-1, 0, 1]:\n                ny = padded_y + dy\n                nx = x + dx\n                if 0 <= ny < volume_part.shape[0] and 0 <= nx < volume_part.shape[1]:\n                    context.append(volume_part[ny, nx])\n                else:\n                    context.append(0.0)  # Padding\n        \n        # Features: x, y, intensity + context\n        features = [global_x / img_size, global_y / img_size, intensity] + context\n        features_list.append(features)\n        labels_list.append(1)\n    \n    # Process non-ink pixels\n    for rel_y, x in non_ink_coords:\n        padded_y = actual_start + rel_y + pad_size\n        global_y = start_row + rel_y\n        global_x = x\n        \n        intensity = volume_part[padded_y, x]\n        \n        # Get context\n        context = []\n        for dy in [-1, 0, 1]:\n            for dx in [-1, 0, 1]:\n                ny = padded_y + dy\n                nx = x + dx\n                if 0 <= ny < volume_part.shape[0] and 0 <= nx < volume_part.shape[1]:\n                    context.append(volume_part[ny, nx])\n                else:\n                    context.append(0.0)\n        \n        features = [global_x / img_size, global_y / img_size, intensity] + context\n        features_list.append(features)\n        labels_list.append(0)\n    \n    return {\n        'features': np.array(features_list, dtype=np.float32),\n        'labels': np.array(labels_list, dtype=np.float32),\n        'slice_idx': slice_idx,\n        'part_idx': part_idx,\n        'start_row': start_row,\n        'end_row': end_row,\n        'original_shape': (height, width)\n    }\n\n# Main function to process all slices in parallel\ndef process_fragment_parallel(volume, mask, fragment_name):\n    \"\"\"Process all slices of a fragment in parallel\"\"\"\n    print(f\"Processing fragment {fragment_name} in parallel...\")\n    \n    all_features = []\n    all_labels = []\n    slice_info = []\n    \n    num_slices = volume.shape[0]\n    \n    # Prepare all tasks for parallel processing\n    tasks = []\n    for slice_idx in range(num_slices):\n        for part_idx in range(config.parts_per_slice):\n            tasks.append((volume[slice_idx], mask, slice_idx, part_idx))\n    \n    # Process in parallel\n    results = Parallel(n_jobs=config.num_workers, verbose=1)(\n        delayed(process_slice_part_with_context)(\n            vol_slice, mask, slice_idx, part_idx, \n            config.rows_per_part, config.img_size\n        )\n        for vol_slice, mask, slice_idx, part_idx in tasks\n    )\n    \n    # Collect results\n    total_pixels = 0\n    for result in results:\n        if len(result['features']) > 0:\n            all_features.append(result['features'])\n            all_labels.append(result['labels'])\n            slice_info.append({\n                'slice_idx': result['slice_idx'],\n                'part_idx': result['part_idx'],\n                'start_row': result['start_row'],\n                'end_row': result['end_row']\n            })\n            total_pixels += len(result['features'])\n    \n    if all_features:\n        all_features = np.vstack(all_features)\n        all_labels = np.concatenate(all_labels)\n        \n        print(f\"Fragment {fragment_name}: Processed {total_pixels:,} pixels\")\n        print(f\"  Ink pixels: {np.sum(all_labels):,} ({np.mean(all_labels)*100:.2f}%)\")\n        print(f\"  Non-ink pixels: {len(all_labels) - np.sum(all_labels):,}\")\n        \n        return all_features, all_labels, slice_info\n    else:\n        return None, None, None\n\n# Dataset for pixel classification\nclass PixelClassificationDataset(Dataset):\n    def __init__(self, features, labels, augment=True):\n        self.features = features\n        self.labels = labels\n        self.augment = augment\n        print(f\"Created dataset with {len(self.features):,} samples\")\n    \n    def __len__(self):\n        return len(self.features)\n    \n    def __getitem__(self, idx):\n        features = self.features[idx].copy()\n        label = self.labels[idx].copy()\n        \n        # Simple augmentation: add noise to features\n        if self.augment and random.random() > 0.5:\n            noise = np.random.normal(0, 0.01, features.shape).astype(np.float32)\n            features = np.clip(features + noise, 0, 1)\n        \n        return torch.FloatTensor(features), torch.FloatTensor([label])\n\n# Enhanced classifier with residual connections\nclass EnhancedPixelClassifier(nn.Module):\n    def __init__(self, input_dim, hidden_dim=128):\n        super().__init__()\n        \n        self.input_layer = nn.Linear(input_dim, hidden_dim)\n        self.bn1 = nn.BatchNorm1d(hidden_dim)\n        \n        # Residual block 1\n        self.res1_fc1 = nn.Linear(hidden_dim, hidden_dim)\n        self.res1_bn1 = nn.BatchNorm1d(hidden_dim)\n        self.res1_fc2 = nn.Linear(hidden_dim, hidden_dim)\n        self.res1_bn2 = nn.BatchNorm1d(hidden_dim)\n        \n        # Residual block 2\n        self.res2_fc1 = nn.Linear(hidden_dim, hidden_dim)\n        self.res2_bn1 = nn.BatchNorm1d(hidden_dim)\n        self.res2_fc2 = nn.Linear(hidden_dim, hidden_dim)\n        self.res2_bn2 = nn.BatchNorm1d(hidden_dim)\n        \n        self.output_layer = nn.Linear(hidden_dim, 1)\n        \n        self.relu = nn.ReLU()\n        self.dropout = nn.Dropout(0.3)\n    \n    def forward(self, x):\n        # Input layer\n        x = self.input_layer(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.dropout(x)\n        \n        # Residual block 1\n        identity = x\n        out = self.res1_fc1(x)\n        out = self.res1_bn1(out)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.res1_fc2(out)\n        out = self.res1_bn2(out)\n        x = self.relu(out + identity)\n        \n        # Residual block 2\n        identity = x\n        out = self.res2_fc1(x)\n        out = self.res2_bn1(out)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.res2_fc2(out)\n        out = self.res2_bn2(out)\n        x = self.relu(out + identity)\n        \n        # Output layer\n        x = self.output_layer(x)\n        \n        return x\n\n# Loss functions\nclass WeightedBCELoss(nn.Module):\n    def __init__(self, pos_weight=2.0):\n        super().__init__()\n        self.pos_weight = torch.tensor([pos_weight])\n    \n    def forward(self, predictions, targets):\n        # Move pos_weight to correct device\n        self.pos_weight = self.pos_weight.to(predictions.device)\n        \n        return F.binary_cross_entropy_with_logits(\n            predictions, targets, \n            pos_weight=self.pos_weight,\n            reduction='mean'\n        )\n\n# Enhanced metrics calculation\ndef calculate_comprehensive_metrics(predictions, targets, threshold=0.5):\n    \"\"\"Calculate comprehensive metrics including accuracy\"\"\"\n    pred_binary = (predictions > threshold).astype(np.float32)\n    \n    # Flatten arrays\n    pred_flat = pred_binary.flatten()\n    target_flat = targets.flatten()\n    \n    # Convert to int for sklearn metrics\n    pred_int = pred_flat.astype(int)\n    target_int = target_flat.astype(int)\n    \n    # Calculate all metrics\n    accuracy = accuracy_score(target_int, pred_int)\n    precision = precision_score(target_int, pred_int, zero_division=0)\n    recall = recall_score(target_int, pred_int, zero_division=0)\n    \n    # Dice score\n    intersection = (pred_flat * target_flat).sum()\n    union = pred_flat.sum() + target_flat.sum()\n    dice = (2. * intersection + 1e-6) / (union + 1e-6) if union > 0 else 0\n    \n    # F0.5 score\n    if precision + recall > 0:\n        f05 = (1 + 0.5**2) * (precision * recall) / ((0.5**2 * precision) + recall)\n    else:\n        f05 = 0\n    \n    return {\n        'accuracy': accuracy,\n        'precision': precision,\n        'recall': recall,\n        'dice': dice,\n        'f05': f05,\n        'threshold': threshold\n    }\n\n# Function to evaluate on pixel data (YOUR ORIGINAL METHOD - for 0.8088 Dice)\ndef evaluate_on_pixels(model, features, labels, batch_size=8192, threshold=0.5):\n    \"\"\"Evaluate model on pixel data - YOUR ORIGINAL METHOD\"\"\"\n    model.eval()\n    \n    # Predict in batches\n    num_batches = (len(features) + batch_size - 1) // batch_size\n    all_preds = []\n    \n    with torch.no_grad():\n        for batch_idx in tqdm(range(num_batches), desc=\"Pixel Evaluation\", leave=False):\n            start_idx = batch_idx * batch_size\n            end_idx = min((batch_idx + 1) * batch_size, len(features))\n            \n            batch_features = features[start_idx:end_idx]\n            \n            features_tensor = torch.FloatTensor(batch_features).to(config.device)\n            \n            outputs = torch.sigmoid(model(features_tensor))\n            preds = outputs.cpu().numpy().flatten()\n            \n            all_preds.extend(preds)\n    \n    all_preds = np.array(all_preds, dtype=np.float32)\n    \n    # Calculate comprehensive metrics\n    metrics = calculate_comprehensive_metrics(all_preds, labels, threshold)\n    \n    return metrics, all_preds\n\n# Function to reconstruct volume from predictions\ndef reconstruct_volume_from_predictions(predictions, coords, num_slices, img_size):\n    \"\"\"Reconstruct full 3D volume from pixel predictions\"\"\"\n    height = width = img_size\n    volume_masks = np.zeros((num_slices, height, width), dtype=np.float32)\n    count_maps = np.zeros((num_slices, height, width), dtype=np.int32)\n    \n    for pred, (slice_idx, y, x) in zip(predictions, coords):\n        volume_masks[slice_idx, y, x] += pred\n        count_maps[slice_idx, y, x] += 1\n    \n    # Average if any pixel was predicted multiple times\n    mask = count_maps > 0\n    volume_masks[mask] = volume_masks[mask] / count_maps[mask]\n    \n    return volume_masks\n\n# Optimized test processing for reconstruction\ndef process_fragment_for_reconstruction(model, volume, fragment_name, batch_size=8192):\n    \"\"\"Process fragment for full reconstruction\"\"\"\n    print(f\"Processing fragment {fragment_name} for reconstruction...\")\n    \n    model.eval()\n    all_features = []\n    all_coords = []\n    \n    num_slices = volume.shape[0]\n    height, width = volume.shape[1], volume.shape[2]\n    \n    # Pre-compute all features\n    for slice_idx in range(num_slices):\n        for y in range(height):\n            for x in range(width):\n                intensity = volume[slice_idx, y, x]\n                \n                # Get context (3x3 neighborhood)\n                context = []\n                for dy in [-1, 0, 1]:\n                    for dx in [-1, 0, 1]:\n                        ny = y + dy\n                        nx = x + dx\n                        if 0 <= ny < height and 0 <= nx < width:\n                            context.append(volume[slice_idx, ny, nx])\n                        else:\n                            context.append(0.0)\n                \n                features = [x / width, y / height, intensity] + context\n                all_features.append(features)\n                all_coords.append((slice_idx, y, x))\n    \n    # Convert to array\n    all_features = np.array(all_features, dtype=np.float32)\n    print(f\"  Total pixels: {len(all_features):,}\")\n    \n    # Predict in batches\n    num_batches = (len(all_features) + batch_size - 1) // batch_size\n    all_predictions = []\n    \n    for batch_idx in tqdm(range(num_batches), desc=f\"    Predicting\"):\n        start_idx = batch_idx * batch_size\n        end_idx = min((batch_idx + 1) * batch_size, len(all_features))\n        \n        batch_features = all_features[start_idx:end_idx]\n        features_tensor = torch.FloatTensor(batch_features).to(config.device)\n        \n        with torch.no_grad():\n            batch_preds = torch.sigmoid(model(features_tensor))\n            batch_preds_np = batch_preds.cpu().numpy().flatten()\n        \n        all_predictions.extend(batch_preds_np)\n    \n    all_predictions = np.array(all_predictions, dtype=np.float32)\n    \n    return all_predictions, all_coords\n\n# Visualization functions\ndef save_visualization(input_slice, pred_slice, true_slice, metrics, \n                       fragment_name, phase=\"test\", epoch=None):\n    \"\"\"Save visualization with comprehensive metrics\"\"\"\n    \n    os.makedirs(f'{phase}_visualizations', exist_ok=True)\n    \n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    # Input slice\n    axes[0, 0].imshow(input_slice, cmap='gray')\n    axes[0, 0].set_title('Input Slice')\n    axes[0, 0].axis('off')\n    \n    # Prediction probability\n    pred_prob = pred_slice\n    im1 = axes[0, 1].imshow(pred_prob, cmap='jet', vmin=0, vmax=1)\n    axes[0, 1].set_title('Prediction Probability')\n    axes[0, 1].axis('off')\n    plt.colorbar(im1, ax=axes[0, 1], fraction=0.046, pad=0.04)\n    \n    # Prediction binary\n    pred_binary = (pred_prob > 0.5).astype(np.float32)\n    axes[0, 2].imshow(pred_binary, cmap='jet')\n    axes[0, 2].set_title('Binary Prediction')\n    axes[0, 2].axis('off')\n    \n    # Ground truth\n    axes[1, 0].imshow(true_slice, cmap='jet')\n    axes[1, 0].set_title('Ground Truth')\n    axes[1, 0].axis('off')\n    \n    # Overlay\n    overlay = np.stack([pred_binary, true_slice, np.zeros_like(pred_binary)], axis=-1)\n    axes[1, 1].imshow(overlay)\n    axes[1, 1].set_title('Overlay (Pred=Red, GT=Green)')\n    axes[1, 1].axis('off')\n    \n    # Metrics text\n    axes[1, 2].axis('off')\n    title_str = f'{phase.upper()}'\n    if epoch is not None:\n        title_str += f' - Epoch {epoch}'\n    title_str += f' - Fragment {fragment_name}\\n\\n'\n    \n    metrics_text = title_str\n    metrics_text += f\"Dice: {metrics['dice']:.4f}\\n\"\n    metrics_text += f\"F0.5: {metrics['f05']:.4f}\\n\"\n    metrics_text += f\"Accuracy: {metrics['accuracy']:.4f}\\n\"\n    metrics_text += f\"Precision: {metrics['precision']:.4f}\\n\"\n    metrics_text += f\"Recall: {metrics['recall']:.4f}\"\n    \n    axes[1, 2].text(0.5, 0.5, metrics_text, ha='center', va='center', fontsize=12,\n                   bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"lightblue\", alpha=0.7))\n    \n    plt.suptitle(f'{phase.title()} - Dice: {metrics[\"dice\"]:.4f}, F0.5: {metrics[\"f05\"]:.4f}', \n                fontsize=14, y=1.02)\n    \n    plt.tight_layout()\n    \n    if epoch is not None:\n        filename = f'{phase}_visualizations/{phase}_epoch{epoch}_fragment{fragment_name}_dice{metrics[\"dice\"]:.4f}_f05{metrics[\"f05\"]:.4f}.png'\n    else:\n        filename = f'{phase}_visualizations/{phase}_fragment{fragment_name}_dice{metrics[\"dice\"]:.4f}_f05{metrics[\"f05\"]:.4f}.png'\n    \n    plt.savefig(filename, dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"✅ Saved {phase} visualization: {filename}\")\n    return True\n\ndef plot_training_history(history):\n    \"\"\"Plot training history\"\"\"\n    os.makedirs('training_plots', exist_ok=True)\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    fig, axes = plt.subplots(2, 3, figsize=(18, 10))\n    \n    # Loss\n    axes[0, 0].plot(epochs, history['train_loss'], 'b-', label='Train Loss', linewidth=2)\n    axes[0, 0].plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2)\n    axes[0, 0].set_title('Training and Validation Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Dice Score\n    axes[0, 1].plot(epochs, history['train_dice'], 'b-', label='Train Dice', linewidth=2)\n    axes[0, 1].plot(epochs, history['val_dice'], 'r-', label='Val Dice', linewidth=2)\n    axes[0, 1].set_title('Training and Validation Dice')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Dice Score')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # F0.5 Score\n    axes[0, 2].plot(epochs, history['train_f05'], 'b-', label='Train F0.5', linewidth=2)\n    axes[0, 2].plot(epochs, history['val_f05'], 'r-', label='Val F0.5', linewidth=2)\n    axes[0, 2].set_title('Training and Validation F0.5')\n    axes[0, 2].set_xlabel('Epoch')\n    axes[0, 2].set_ylabel('F0.5 Score')\n    axes[0, 2].legend()\n    axes[0, 2].grid(True, alpha=0.3)\n    \n    # Pixel Evaluation Dice (if available)\n    if 'test_pixel_dice' in history and history['test_pixel_dice']:\n        axes[1, 0].plot(range(1, len(history['test_pixel_dice']) + 1), \n                       history['test_pixel_dice'], 'g-', label='Pixel Eval Dice', linewidth=2)\n        axes[1, 0].set_title('Pixel Evaluation Dice')\n        axes[1, 0].set_xlabel('Epoch')\n        axes[1, 0].set_ylabel('Dice Score')\n        axes[1, 0].legend()\n        axes[1, 0].grid(True, alpha=0.3)\n    \n    # Reconstruction Evaluation Dice (if available)\n    if 'test_reconstruction_dice' in history and history['test_reconstruction_dice']:\n        recon_epochs = [i*5+1 for i in range(len(history['test_reconstruction_dice']))]\n        axes[1, 1].plot(recon_epochs, history['test_reconstruction_dice'], \n                       'm-', label='Reconstruction Dice', linewidth=2, marker='o')\n        axes[1, 1].set_title('Reconstruction Evaluation Dice')\n        axes[1, 1].set_xlabel('Epoch')\n        axes[1, 1].set_ylabel('Dice Score')\n        axes[1, 1].legend()\n        axes[1, 1].grid(True, alpha=0.3)\n    \n    # Final Comparison\n    final_metrics = ['Pixel Eval', 'Reconstruction']\n    final_dice = []\n    if history.get('final_pixel_dice'):\n        final_dice.append(history['final_pixel_dice'])\n    if history.get('final_reconstruction_dice'):\n        final_dice.append(history['final_reconstruction_dice'])\n    \n    if final_dice:\n        bars = axes[1, 2].bar(final_metrics, final_dice, color=['green', 'purple'])\n        axes[1, 2].set_title('Final Evaluation Comparison')\n        axes[1, 2].set_ylabel('Dice Score')\n        axes[1, 2].grid(True, alpha=0.3, axis='y')\n        \n        # Add value labels on bars\n        for bar, val in zip(bars, final_dice):\n            height = bar.get_height()\n            axes[1, 2].text(bar.get_x() + bar.get_width()/2., height + 0.01,\n                          f'{val:.4f}', ha='center', va='bottom')\n    \n    plt.tight_layout()\n    plt.savefig('training_plots/training_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    print(\"📊 Saved training history plot\")\n\n# ENHANCED Training function with BOTH evaluations\ndef train_pixel_classifier_enhanced(model, train_loader, val_loader, \n                                   train_features, train_labels,\n                                   val_features, val_labels,\n                                   test_features, test_labels,\n                                   test_volume, test_mask):\n    \"\"\"Train pixel classifier with BOTH evaluations\"\"\"\n    \n    optimizer = optim.AdamW(model.parameters(), lr=config.lr, weight_decay=1e-4)\n    scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=7, verbose=True)\n    \n    criterion = WeightedBCELoss(pos_weight=config.pos_weight)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    best_model_state = None\n    \n    # Track BOTH types of metrics\n    history = {\n        'train_loss': [],\n        'val_loss': [],\n        'train_dice': [],\n        'val_dice': [],\n        'train_f05': [],\n        'val_f05': [],\n        'test_pixel_dice': [],  # Your original method\n        'test_reconstruction_dice': []  # Reconstruction method\n    }\n    \n    for epoch in range(config.epochs):\n        print(f\"\\n{'='*60}\")\n        print(f\"Epoch {epoch+1}/{config.epochs}\")\n        print(f\"{'='*60}\")\n        \n        # Training phase\n        model.train()\n        train_loss = 0\n        train_preds = []\n        train_targets = []\n        \n        for features, labels in tqdm(train_loader, desc=f'Training'):\n            features, labels = features.to(config.device), labels.to(config.device)\n            \n            optimizer.zero_grad()\n            outputs = model(features)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            \n            # Store predictions for metrics\n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_preds.append(preds.cpu().numpy())\n                train_targets.append(labels.cpu().numpy())\n        \n        # Calculate training metrics\n        avg_train_loss = train_loss / len(train_loader)\n        train_preds = np.concatenate(train_preds).flatten()\n        train_targets = np.concatenate(train_targets).flatten()\n        train_metrics = calculate_comprehensive_metrics(train_preds, train_targets)\n        \n        # Update training history\n        history['train_loss'].append(avg_train_loss)\n        history['train_dice'].append(train_metrics['dice'])\n        history['train_f05'].append(train_metrics['f05'])\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0\n        val_preds = []\n        val_targets = []\n        \n        with torch.no_grad():\n            for features, labels in tqdm(val_loader, desc=f'Validation'):\n                features, labels = features.to(config.device), labels.to(config.device)\n                \n                outputs = model(features)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item()\n                \n                preds = torch.sigmoid(outputs)\n                val_preds.append(preds.cpu().numpy())\n                val_targets.append(labels.cpu().numpy())\n        \n        # Calculate validation metrics\n        avg_val_loss = val_loss / len(val_loader)\n        val_preds = np.concatenate(val_preds).flatten()\n        val_targets = np.concatenate(val_targets).flatten()\n        val_metrics = calculate_comprehensive_metrics(val_preds, val_targets)\n        \n        # Update validation history\n        history['val_loss'].append(avg_val_loss)\n        history['val_dice'].append(val_metrics['dice'])\n        history['val_f05'].append(val_metrics['f05'])\n        \n        # TEST: PIXEL EVALUATION (YOUR ORIGINAL METHOD - for 0.8088 Dice)\n        pixel_test_metrics = None\n        if test_features is not None and test_labels is not None and (epoch % 2 == 0 or epoch == 0):\n            print(f\"\\n🔍 Running PIXEL EVALUATION (Original Method)...\")\n            pixel_test_metrics, _ = evaluate_on_pixels(model, test_features, test_labels)\n            history['test_pixel_dice'].append(pixel_test_metrics['dice'])\n            \n            print(f\"  Pixel Eval - Dice: {pixel_test_metrics['dice']:.4f}, \"\n                  f\"F0.5: {pixel_test_metrics['f05']:.4f}\")\n        \n        # TEST: RECONSTRUCTION EVALUATION (every 5 epochs to save time)\n        reconstruction_test_metrics = None\n        if test_volume is not None and test_mask is not None and (epoch % 5 == 0 or epoch == 0):\n            print(f\"\\n🏗️  Running RECONSTRUCTION EVALUATION...\")\n            \n            # Process test fragment for reconstruction\n            test_predictions, test_coords = process_fragment_for_reconstruction(\n                model, test_volume, config.test_fragment\n            )\n            \n            # Reconstruct mask\n            test_reconstructed = reconstruct_volume_from_predictions(\n                test_predictions, test_coords, len(config.slices), config.img_size\n            )\n            \n            # Take middle slice for evaluation\n            middle_slice = len(config.slices) // 2\n            test_pred_slice = test_reconstructed[middle_slice]\n            test_pred_binary = (test_pred_slice > 0.5).astype(np.float32)\n            \n            # Calculate reconstruction metrics\n            reconstruction_test_metrics = calculate_comprehensive_metrics(test_pred_binary, test_mask)\n            history['test_reconstruction_dice'].append(reconstruction_test_metrics['dice'])\n            \n            print(f\"  Recon Eval - Dice: {reconstruction_test_metrics['dice']:.4f}, \"\n                  f\"F0.5: {reconstruction_test_metrics['f05']:.4f}\")\n            \n            # Save reconstruction visualization\n            if config.save_visualization:\n                input_slice = test_volume[middle_slice]\n                save_visualization(\n                    input_slice, test_pred_binary, test_mask,\n                    reconstruction_test_metrics, config.test_fragment, \n                    phase=\"reconstruction\", epoch=epoch+1\n                )\n        \n        # Early stopping based on validation Dice\n        if val_metrics['dice'] > best_val_dice:\n            best_val_dice = val_metrics['dice']\n            patience_counter = 0\n            best_model_state = model.state_dict().copy()\n            print(f\"\\n🎯 NEW BEST: Validation Dice: {best_val_dice:.4f}\")\n        else:\n            patience_counter += 1\n        \n        # Print epoch summary\n        print(f\"\\n📊 Epoch {epoch+1} Summary:\")\n        print(f\"  TRAIN: Loss={avg_train_loss:.4f}, Dice={train_metrics['dice']:.4f}, F0.5={train_metrics['f05']:.4f}\")\n        print(f\"  VAL:   Loss={avg_val_loss:.4f}, Dice={val_metrics['dice']:.4f}, F0.5={val_metrics['f05']:.4f}\")\n        \n        # Learning rate scheduling based on validation Dice\n        scheduler.step(val_metrics['dice'])\n        \n        # Early stopping\n        if patience_counter >= config.early_stopping_patience:\n            print(f\"\\n🛑 Early stopping triggered at epoch {epoch+1}\")\n            break\n    \n    # Load best model\n    if best_model_state:\n        model.load_state_dict(best_model_state)\n        print(f\"\\n✅ Loaded best model with Validation Dice: {best_val_dice:.4f}\")\n    \n    return model, history\n\n# MAIN FUNCTION\ndef main():\n    print(\"🚀 Starting Enhanced Pixel Classification...\")\n    print(f\"Training on fragment {config.train_fragment}\")\n    print(f\"Validating on fragment {config.val_fragment}\")\n    print(f\"Testing on fragment {config.test_fragment}\")\n    print(f\"Device: {config.device}\")\n    print(f\"Rows per part: {config.rows_per_part}\")\n    print(f\"Parallel workers: {config.num_workers}\")\n    \n    # Clear cache\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    # Load data\n    print(\"\\n📂 Loading data...\")\n    \n    # Load training fragment\n    train_volume = load_volume_data(config.train_path, config.slices)\n    train_mask = load_mask_data(config.train_path)\n    \n    if train_volume is None or train_mask is None:\n        print(f\"❌ Failed to load training fragment {config.train_fragment}\")\n        return None\n    \n    # Load validation fragment\n    val_volume = load_volume_data(config.val_path, config.slices)\n    val_mask = load_mask_data(config.val_path)\n    \n    if val_volume is None or val_mask is None:\n        print(f\"❌ Failed to load validation fragment {config.val_fragment}\")\n        return None\n    \n    # Load test fragment\n    test_volume = load_volume_data(config.test_path, config.slices)\n    test_mask = load_mask_data(config.test_path)\n    \n    if test_volume is None or test_mask is None:\n        print(f\"❌ Failed to load test fragment {config.test_fragment}\")\n        return None\n    \n    print(f\"\\n📊 Dataset Summary:\")\n    print(f\"  Train volume: {train_volume.shape}\")\n    print(f\"  Val volume: {val_volume.shape}\")\n    print(f\"  Test volume: {test_volume.shape}\")\n    \n    # Process fragments in parallel\n    print(\"\\n🔄 Processing fragments in parallel...\")\n    \n    # Process training fragment\n    train_features, train_labels, _ = process_fragment_parallel(\n        train_volume, train_mask, config.train_fragment\n    )\n    \n    if train_features is None:\n        print(\"❌ Failed to process training fragment\")\n        return None\n    \n    # Process validation fragment\n    val_features, val_labels, _ = process_fragment_parallel(\n        val_volume, val_mask, config.val_fragment\n    )\n    \n    if val_features is None:\n        print(\"❌ Failed to process validation fragment\")\n        return None\n    \n    # Process test fragment for pixel evaluation\n    test_features, test_labels, _ = process_fragment_parallel(\n        test_volume, test_mask, config.test_fragment\n    )\n    \n    if test_features is None:\n        print(\"❌ Failed to process test fragment\")\n        return None\n    \n    print(f\"\\n📈 Processed Data Summary:\")\n    print(f\"  Training pixels: {len(train_features):,}\")\n    print(f\"  Validation pixels: {len(val_features):,}\")\n    print(f\"  Test pixels: {len(test_features):,}\")\n    print(f\"  Input dimension: {train_features.shape[1]}\")\n    \n    # Create datasets\n    train_dataset = PixelClassificationDataset(train_features, train_labels, augment=True)\n    val_dataset = PixelClassificationDataset(val_features, val_labels, augment=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True if torch.cuda.is_available() else False\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        num_workers=2\n    )\n    \n    # Initialize model\n    input_dim = train_features.shape[1]\n    model = EnhancedPixelClassifier(input_dim, hidden_dim=config.hidden_dim).to(config.device)\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"\\n✅ Model initialized\")\n    print(f\"  Parameters: {total_params/1e6:.1f}M\")\n    print(f\"  Input dimension: {input_dim}\")\n    print(f\"  Hidden dimension: {config.hidden_dim}\")\n    \n    # Train model\n    print(\"\\n🎯 Starting training with BOTH evaluations...\")\n    model, history = train_pixel_classifier_enhanced(\n        model, train_loader, val_loader,\n        train_features, train_labels,\n        val_features, val_labels,\n        test_features, test_labels,\n        test_volume, test_mask\n    )\n    \n    # FINAL EVALUATIONS\n    print(f\"\\n{'='*60}\")\n    print(\"FINAL EVALUATION\")\n    print(f\"{'='*60}\")\n    \n    # 1. FINAL PIXEL EVALUATION (YOUR ORIGINAL METHOD)\n    print(\"\\n1️⃣ FINAL PIXEL EVALUATION (Original Method):\")\n    final_pixel_metrics, final_pixel_predictions = evaluate_on_pixels(\n        model, test_features, test_labels\n    )\n    \n    print(f\"  ✅ Dice: {final_pixel_metrics['dice']:.4f}\")\n    print(f\"  ✅ F0.5: {final_pixel_metrics['f05']:.4f}\")\n    print(f\"  ✅ Accuracy: {final_pixel_metrics['accuracy']:.4f}\")\n    print(f\"  ✅ Precision: {final_pixel_metrics['precision']:.4f}\")\n    print(f\"  ✅ Recall: {final_pixel_metrics['recall']:.4f}\")\n    \n    history['final_pixel_dice'] = final_pixel_metrics['dice']\n    \n    # 2. FINAL RECONSTRUCTION EVALUATION\n    print(\"\\n2️⃣ FINAL RECONSTRUCTION EVALUATION:\")\n    test_predictions, test_coords = process_fragment_for_reconstruction(\n        model, test_volume, config.test_fragment\n    )\n    \n    # Reconstruct mask\n    test_reconstructed = reconstruct_volume_from_predictions(\n        test_predictions, test_coords, len(config.slices), config.img_size\n    )\n    \n    # Take middle slice for evaluation\n    middle_slice = len(config.slices) // 2\n    test_pred_slice = test_reconstructed[middle_slice]\n    test_pred_binary = (test_pred_slice > 0.5).astype(np.float32)\n    \n    # Calculate final reconstruction metrics\n    final_recon_metrics = calculate_comprehensive_metrics(test_pred_binary, test_mask)\n    \n    print(f\"  ✅ Dice: {final_recon_metrics['dice']:.4f}\")\n    print(f\"  ✅ F0.5: {final_recon_metrics['f05']:.4f}\")\n    print(f\"  ✅ Accuracy: {final_recon_metrics['accuracy']:.4f}\")\n    print(f\"  ✅ Precision: {final_recon_metrics['precision']:.4f}\")\n    print(f\"  ✅ Recall: {final_recon_metrics['recall']:.4f}\")\n    \n    history['final_reconstruction_dice'] = final_recon_metrics['dice']\n    \n    # Save final visualizations\n    if config.save_visualization:\n        print(\"\\n📸 Saving final visualizations...\")\n        \n        # Pixel evaluation visualization\n        input_slice = test_volume[middle_slice]\n        \n        # For pixel evaluation, create a simple visualization\n        pixel_viz = np.zeros_like(test_mask)\n        # We can't visualize individual pixels easily, so we'll show a comparison\n        \n        save_visualization(\n            input_slice, test_pred_binary, test_mask,\n            final_recon_metrics, config.test_fragment, \n            phase=\"final_reconstruction\"\n        )\n    \n    # Plot training history\n    plot_training_history(history)\n    \n    print(f\"\\n{'='*60}\")\n    print(\"RESULTS SUMMARY\")\n    print(f\"{'='*60}\")\n    print(f\"Pixel Evaluation (Original Method):\")\n    print(f\"  Dice: {final_pixel_metrics['dice']:.4f} ← YOUR 0.8088 DICE METHOD\")\n    print(f\"  F0.5: {final_pixel_metrics['f05']:.4f}\")\n    print(f\"\\nReconstruction Evaluation:\")\n    print(f\"  Dice: {final_recon_metrics['dice']:.4f}\")\n    print(f\"  F0.5: {final_recon_metrics['f05']:.4f}\")\n    print(f\"\\n💡 Note: Pixel evaluation matches your training sampling method\")\n    print(f\"       while reconstruction evaluates full slice reconstruction\")\n    \n    return {\n        'pixel_metrics': final_pixel_metrics,\n        'reconstruction_metrics': final_recon_metrics,\n        'history': history,\n        'model': model\n    }\n\n# Run the main function\nif __name__ == \"__main__\":\n    # Set seeds for reproducibility\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        print(f\"✅ Using GPU: {torch.cuda.get_device_name(0)}\")\n        # Enable TF32 for faster computation on Ampere GPUs\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.allow_tf32 = True\n    \n    results = main()\n    \n    if results:\n        print(\"\\n🎉 Training completed successfully!\")\n        print(f\"📊 Final Pixel Dice: {results['pixel_metrics']['dice']:.4f}\")\n        print(f\"📊 Final Reconstruction Dice: {results['reconstruction_metrics']['dice']:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing\nfrom skimage.morphology import remove_small_objects\nimport warnings\nimport random\nimport math\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import precision_score, recall_score, f1_score, accuracy_score\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport torch.fft\nfrom joblib import Parallel, delayed\nfrom multiprocessing import cpu_count\nfrom collections import defaultdict\n\nwarnings.filterwarnings('ignore')\n\n# Configuration\nclass Config:\n    # Data paths\n    train_fragment = '2'  # Train on fragment 2\n    val_fragment = '3'    # Validate on fragment 3\n    test_fragment = '1'   # Test on fragment 1\n    \n    # Model parameters\n    img_size = 384\n    slices = list(range(12, 31))  # 19 slices\n    batch_size = 64  # Larger batch for pixel classification\n    epochs = 100\n    lr = 3e-4\n    \n    # Split parameters\n    rows_per_part = 10  # Split each slice into parts of 10 rows\n    parts_per_slice = None  # Will be calculated\n    \n    # Pixel classification parameters\n    num_features = 3  # x, y, intensity\n    hidden_dim = 64\n    use_context = True  # Use neighboring pixels as context\n    \n    # Training parameters\n    pos_weight = 2.0  # For handling class imbalance\n    sample_ink_pixels = 100000  # Max ink pixels to sample per fragment (for memory)\n    sample_non_ink_pixels = 200000  # Max non-ink pixels to sample per fragment\n    \n    # Parallel processing\n    num_workers = max(1, cpu_count() - 2)  # Use all but 2 cores\n    \n    # Visualization\n    save_visualization = True\n    vis_frequency = 5\n    min_score_threshold = 0.30  # Save when Dice or F0.5 > 0.30\n    \n    # Early stopping\n    early_stopping_patience = 15\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Paths\n    @property\n    def train_path(self):\n        return f'/kaggle/input/vesuvius-challenge-ink-detection/train/{self.train_fragment}'\n    \n    @property\n    def val_path(self):\n        return f'/kaggle/input/vesuvius-challenge-ink-detection/train/{self.val_fragment}'\n    \n    @property\n    def test_path(self):\n        return f'/kaggle/input/vesuvius-challenge-ink-detection/train/{self.test_fragment}'\n\nconfig = Config()\nconfig.parts_per_slice = math.ceil(config.img_size / config.rows_per_part)\n\n# Data loading functions\ndef load_volume_data(base_path, slices):\n    \"\"\"Load 3D volume data\"\"\"\n    volume = []\n    for slice_idx in slices:\n        slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n        if os.path.exists(slice_path):\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_resized = cv2.resize(img, (config.img_size, config.img_size))\n                img_normalized = img_resized.astype(np.float32) / 255.0\n                volume.append(img_normalized)\n    \n    if volume:\n        volume = np.stack(volume, axis=0)\n        return volume\n    return None\n\ndef load_mask_data(base_path):\n    \"\"\"Load and preprocess mask data\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            ink_ratio = mask_binary.sum() / mask_binary.size\n            print(f\"Ink ratio: {ink_ratio:.4f} ({mask_binary.sum()} pixels)\")\n            return mask_binary\n    return None\n\n# Function to process a single slice part\ndef process_slice_part_with_context(volume_slice, mask_slice, slice_idx, part_idx, rows_per_part, img_size):\n    \"\"\"Process a single part with context from neighboring pixels\"\"\"\n    height, width = volume_slice.shape\n    \n    # Calculate start and end rows\n    start_row = part_idx * rows_per_part\n    end_row = min((part_idx + 1) * rows_per_part, height)\n    \n    # Extract the part with padding for context\n    pad_size = 1  # Look at immediate neighbors\n    padded_start = max(0, start_row - pad_size)\n    padded_end = min(height, end_row + pad_size)\n    \n    volume_part = volume_slice[padded_start:padded_end, :]\n    mask_part = mask_slice[padded_start:padded_end, :]\n    \n    # Get pixel coordinates within the actual region (without padding)\n    actual_start = start_row - padded_start\n    actual_end = actual_start + (end_row - start_row)\n    \n    # Get ink and non-ink pixel coordinates in actual region\n    ink_mask = mask_part[actual_start:actual_end, :] > 0.5\n    non_ink_mask = mask_part[actual_start:actual_end, :] <= 0.5\n    \n    ink_coords = np.argwhere(ink_mask)\n    non_ink_coords = np.argwhere(non_ink_mask)\n    \n    # Balance the samples\n    n_samples = min(5000, len(ink_coords) * 2)  # Target samples per part\n    if len(ink_coords) > 0:\n        n_ink = min(n_samples // 2, len(ink_coords))\n        n_non_ink = min(n_samples - n_ink, len(non_ink_coords))\n        \n        if len(ink_coords) > n_ink:\n            ink_indices = np.random.choice(len(ink_coords), n_ink, replace=False)\n            ink_coords = ink_coords[ink_indices]\n        \n        if len(non_ink_coords) > n_non_ink:\n            non_ink_indices = np.random.choice(len(non_ink_coords), n_non_ink, replace=False)\n            non_ink_coords = non_ink_coords[non_ink_indices]\n    else:\n        n_non_ink = min(n_samples, len(non_ink_coords))\n        if len(non_ink_coords) > n_non_ink:\n            non_ink_indices = np.random.choice(len(non_ink_coords), n_non_ink, replace=False)\n            non_ink_coords = non_ink_coords[non_ink_indices]\n        n_ink = 0\n    \n    features_list = []\n    labels_list = []\n    \n    # Process ink pixels\n    for rel_y, x in ink_coords:\n        # Convert to coordinates in padded region\n        padded_y = actual_start + rel_y + pad_size\n        global_y = start_row + rel_y\n        global_x = x\n        \n        # Get intensity at pixel\n        intensity = volume_part[padded_y, x]\n        \n        # Get context (3x3 neighborhood)\n        context = []\n        for dy in [-1, 0, 1]:\n            for dx in [-1, 0, 1]:\n                ny = padded_y + dy\n                nx = x + dx\n                if 0 <= ny < volume_part.shape[0] and 0 <= nx < volume_part.shape[1]:\n                    context.append(volume_part[ny, nx])\n                else:\n                    context.append(0.0)  # Padding\n        \n        # Features: x, y, intensity + context\n        features = [global_x / img_size, global_y / img_size, intensity] + context\n        features_list.append(features)\n        labels_list.append(1)\n    \n    # Process non-ink pixels\n    for rel_y, x in non_ink_coords:\n        padded_y = actual_start + rel_y + pad_size\n        global_y = start_row + rel_y\n        global_x = x\n        \n        intensity = volume_part[padded_y, x]\n        \n        # Get context\n        context = []\n        for dy in [-1, 0, 1]:\n            for dx in [-1, 0, 1]:\n                ny = padded_y + dy\n                nx = x + dx\n                if 0 <= ny < volume_part.shape[0] and 0 <= nx < volume_part.shape[1]:\n                    context.append(volume_part[ny, nx])\n                else:\n                    context.append(0.0)\n        \n        features = [global_x / img_size, global_y / img_size, intensity] + context\n        features_list.append(features)\n        labels_list.append(0)\n    \n    return {\n        'features': np.array(features_list, dtype=np.float32),\n        'labels': np.array(labels_list, dtype=np.float32),\n        'slice_idx': slice_idx,\n        'part_idx': part_idx,\n        'start_row': start_row,\n        'end_row': end_row,\n        'original_shape': (height, width)\n    }\n\n# Main function to process all slices in parallel\ndef process_fragment_parallel(volume, mask, fragment_name, use_context=True):\n    \"\"\"Process all slices of a fragment in parallel\"\"\"\n    print(f\"Processing fragment {fragment_name} in parallel...\")\n    \n    all_features = []\n    all_labels = []\n    slice_info = []\n    \n    num_slices = volume.shape[0]\n    \n    # Prepare all tasks for parallel processing\n    tasks = []\n    for slice_idx in range(num_slices):\n        for part_idx in range(config.parts_per_slice):\n            tasks.append((volume[slice_idx], mask, slice_idx, part_idx))\n    \n    # Process in parallel\n    results = Parallel(n_jobs=config.num_workers, verbose=1)(\n        delayed(process_slice_part_with_context)(\n            vol_slice, mask, slice_idx, part_idx, \n            config.rows_per_part, config.img_size\n        )\n        for vol_slice, mask, slice_idx, part_idx in tasks\n    )\n    \n    # Collect results\n    total_pixels = 0\n    for result in results:\n        if len(result['features']) > 0:\n            all_features.append(result['features'])\n            all_labels.append(result['labels'])\n            slice_info.append({\n                'slice_idx': result['slice_idx'],\n                'part_idx': result['part_idx'],\n                'start_row': result['start_row'],\n                'end_row': result['end_row']\n            })\n            total_pixels += len(result['features'])\n    \n    if all_features:\n        all_features = np.vstack(all_features)\n        all_labels = np.concatenate(all_labels)\n        \n        print(f\"Fragment {fragment_name}: Processed {total_pixels:,} pixels\")\n        print(f\"  Ink pixels: {np.sum(all_labels):,} ({np.mean(all_labels)*100:.2f}%)\")\n        print(f\"  Non-ink pixels: {len(all_labels) - np.sum(all_labels):,}\")\n        \n        return all_features, all_labels, slice_info\n    else:\n        return None, None, None\n\n# Dataset for pixel classification\nclass PixelClassificationDataset(Dataset):\n    def __init__(self, features, labels, augment=True):\n        self.features = features\n        self.labels = labels\n        self.augment = augment\n        print(f\"Created dataset with {len(self.features):,} samples\")\n    \n    def __len__(self):\n        return len(self.features)\n    \n    def __getitem__(self, idx):\n        features = self.features[idx].copy()\n        label = self.labels[idx].copy()\n        \n        # Simple augmentation: add noise to features\n        if self.augment and random.random() > 0.5:\n            noise = np.random.normal(0, 0.01, features.shape).astype(np.float32)\n            features = np.clip(features + noise, 0, 1)\n        \n        return torch.FloatTensor(features), torch.FloatTensor([label])\n\n# Enhanced classifier with residual connections\nclass EnhancedPixelClassifier(nn.Module):\n    def __init__(self, input_dim, hidden_dim=128):\n        super().__init__()\n        \n        self.input_layer = nn.Linear(input_dim, hidden_dim)\n        self.bn1 = nn.BatchNorm1d(hidden_dim)\n        \n        # Residual block 1\n        self.res1_fc1 = nn.Linear(hidden_dim, hidden_dim)\n        self.res1_bn1 = nn.BatchNorm1d(hidden_dim)\n        self.res1_fc2 = nn.Linear(hidden_dim, hidden_dim)\n        self.res1_bn2 = nn.BatchNorm1d(hidden_dim)\n        \n        # Residual block 2\n        self.res2_fc1 = nn.Linear(hidden_dim, hidden_dim)\n        self.res2_bn1 = nn.BatchNorm1d(hidden_dim)\n        self.res2_fc2 = nn.Linear(hidden_dim, hidden_dim)\n        self.res2_bn2 = nn.BatchNorm1d(hidden_dim)\n        \n        self.output_layer = nn.Linear(hidden_dim, 1)\n        \n        self.relu = nn.ReLU()\n        self.dropout = nn.Dropout(0.3)\n    \n    def forward(self, x):\n        # Input layer\n        x = self.input_layer(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.dropout(x)\n        \n        # Residual block 1\n        identity = x\n        out = self.res1_fc1(x)\n        out = self.res1_bn1(out)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.res1_fc2(out)\n        out = self.res1_bn2(out)\n        x = self.relu(out + identity)\n        \n        # Residual block 2\n        identity = x\n        out = self.res2_fc1(x)\n        out = self.res2_bn1(out)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.res2_fc2(out)\n        out = self.res2_bn2(out)\n        x = self.relu(out + identity)\n        \n        # Output layer\n        x = self.output_layer(x)\n        \n        return x\n\n# Loss functions\nclass WeightedBCELoss(nn.Module):\n    def __init__(self, pos_weight=2.0):\n        super().__init__()\n        self.pos_weight = torch.tensor([pos_weight])\n    \n    def forward(self, predictions, targets):\n        # Move pos_weight to correct device\n        self.pos_weight = self.pos_weight.to(predictions.device)\n        \n        return F.binary_cross_entropy_with_logits(\n            predictions, targets, \n            pos_weight=self.pos_weight,\n            reduction='mean'\n        )\n\n# Metrics calculation with comprehensive metrics\ndef calculate_comprehensive_metrics(predictions, targets, threshold=0.5):\n    \"\"\"Calculate comprehensive metrics including accuracy\"\"\"\n    pred_binary = (predictions > threshold).astype(np.float32)\n    \n    # Flatten arrays\n    pred_flat = pred_binary.flatten()\n    target_flat = targets.flatten()\n    \n    # Convert to int for sklearn metrics\n    pred_int = pred_flat.astype(int)\n    target_int = target_flat.astype(int)\n    \n    # Calculate all metrics\n    accuracy = accuracy_score(target_int, pred_int)\n    precision = precision_score(target_int, pred_int, zero_division=0)\n    recall = recall_score(target_int, pred_int, zero_division=0)\n    \n    # Dice score\n    intersection = (pred_flat * target_flat).sum()\n    union = pred_flat.sum() + target_flat.sum()\n    dice = (2. * intersection + 1e-6) / (union + 1e-6) if union > 0 else 0\n    \n    # F0.5 score\n    if precision + recall > 0:\n        f05 = (1 + 0.5**2) * (precision * recall) / ((0.5**2 * precision) + recall)\n    else:\n        f05 = 0\n    \n    return {\n        'accuracy': accuracy,\n        'precision': precision,\n        'recall': recall,\n        'dice': dice,\n        'f05': f05,\n        'threshold': threshold\n    }\n\n# Function to evaluate model on pixel data with comprehensive metrics\ndef evaluate_pixel_model(model, features, labels, batch_size=4096, threshold=0.5):\n    \"\"\"Evaluate model on pixel data with comprehensive metrics\"\"\"\n    model.eval()\n    \n    # Predict in batches\n    num_batches = (len(features) + batch_size - 1) // batch_size\n    all_preds = []\n    \n    with torch.no_grad():\n        for batch_idx in range(num_batches):\n            start_idx = batch_idx * batch_size\n            end_idx = min((batch_idx + 1) * batch_size, len(features))\n            \n            batch_features = features[start_idx:end_idx]\n            \n            features_tensor = torch.FloatTensor(batch_features).to(config.device)\n            \n            outputs = torch.sigmoid(model(features_tensor))\n            preds = outputs.cpu().numpy().flatten()\n            \n            all_preds.extend(preds)\n    \n    all_preds = np.array(all_preds, dtype=np.float32)\n    \n    # Calculate comprehensive metrics\n    metrics = calculate_comprehensive_metrics(all_preds, labels, threshold)\n    \n    return metrics\n\n# Function to reconstruct volume from predictions\ndef reconstruct_volume_from_predictions(predictions, coords, num_slices, img_size):\n    \"\"\"Reconstruct full 3D volume from pixel predictions\"\"\"\n    height = width = img_size\n    volume_masks = np.zeros((num_slices, height, width), dtype=np.float32)\n    count_maps = np.zeros((num_slices, height, width), dtype=np.int32)\n    \n    for pred, (slice_idx, y, x) in zip(predictions, coords):\n        volume_masks[slice_idx, y, x] += pred\n        count_maps[slice_idx, y, x] += 1\n    \n    # Average if any pixel was predicted multiple times\n    mask = count_maps > 0\n    volume_masks[mask] = volume_masks[mask] / count_maps[mask]\n    \n    return volume_masks\n\n# Visualization functions with score checking\ndef save_validation_visualization(input_slice, pred_slice, true_slice, epoch, \n                                 metrics, fragment_name, phase=\"validation\"):\n    \"\"\"Save validation visualization when Dice or F0.5 > 0.30\"\"\"\n    \n    if metrics['dice'] >= config.min_score_threshold or metrics['f05'] >= config.min_score_threshold:\n        os.makedirs(f'{phase}_visualizations', exist_ok=True)\n        \n        fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n        \n        # Input slice\n        axes[0, 0].imshow(input_slice, cmap='gray')\n        axes[0, 0].set_title('Input Slice')\n        axes[0, 0].axis('off')\n        \n        # Prediction\n        axes[0, 1].imshow(pred_slice, cmap='jet')\n        axes[0, 1].set_title('Prediction')\n        axes[0, 1].axis('off')\n        \n        # Ground truth\n        axes[0, 2].imshow(true_slice, cmap='jet')\n        axes[0, 2].set_title('Ground Truth')\n        axes[0, 2].axis('off')\n        \n        # Overlay\n        overlay = np.stack([pred_slice, true_slice, np.zeros_like(pred_slice)], axis=-1)\n        axes[1, 0].imshow(overlay)\n        axes[1, 0].set_title('Overlay (Pred=Red, GT=Green)')\n        axes[1, 0].axis('off')\n        \n        # Difference\n        diff = np.abs(pred_slice - true_slice)\n        axes[1, 1].imshow(diff, cmap='hot')\n        axes[1, 1].set_title('Difference')\n        axes[1, 1].axis('off')\n        \n        # Metrics text\n        axes[1, 2].axis('off')\n        metrics_text = f'{phase.upper()} - Epoch {epoch} - Fragment {fragment_name}\\n\\n'\n        metrics_text += f\"Dice: {metrics['dice']:.4f}\\n\"\n        metrics_text += f\"F0.5: {metrics['f05']:.4f}\\n\"\n        metrics_text += f\"Accuracy: {metrics['accuracy']:.4f}\\n\"\n        metrics_text += f\"Precision: {metrics['precision']:.4f}\\n\"\n        metrics_text += f\"Recall: {metrics['recall']:.4f}\\n\"\n        metrics_text += f\"Threshold: {metrics['threshold']:.3f}\"\n        \n        axes[1, 2].text(0.5, 0.5, metrics_text, ha='center', va='center', fontsize=12,\n                       bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"lightblue\", alpha=0.7))\n        \n        plt.suptitle(f'{phase.title()} Visualization - Dice: {metrics[\"dice\"]:.4f}, F0.5: {metrics[\"f05\"]:.4f}', \n                    fontsize=14, y=1.02)\n        \n        plt.tight_layout()\n        filename = f'{phase}_visualizations/{phase}_epoch{epoch}_fragment{fragment_name}_dice{metrics[\"dice\"]:.4f}_f05{metrics[\"f05\"]:.4f}.png'\n        plt.savefig(filename, dpi=100, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"✅ Saved {phase} visualization: {filename}\")\n        return True\n    else:\n        print(f\"⏭️ Skipped {phase} visualization (Dice: {metrics['dice']:.4f}, F0.5: {metrics['f05']:.4f} < {config.min_score_threshold})\")\n        return False\n\n# Optimized test processing with batching\ndef process_fragment_for_evaluation(model, volume, fragment_name, batch_size=8192):\n    \"\"\"Process fragment for evaluation with batching\"\"\"\n    print(f\"  Processing fragment {fragment_name} for evaluation...\")\n    \n    model.eval()\n    all_features = []\n    all_coords = []\n    \n    num_slices = volume.shape[0]\n    height, width = volume.shape[1], volume.shape[2]\n    \n    # Pre-compute all features\n    for slice_idx in range(num_slices):\n        for y in range(height):\n            for x in range(width):\n                intensity = volume[slice_idx, y, x]\n                \n                # Get context (3x3 neighborhood)\n                context = []\n                for dy in [-1, 0, 1]:\n                    for dx in [-1, 0, 1]:\n                        ny = y + dy\n                        nx = x + dx\n                        if 0 <= ny < height and 0 <= nx < width:\n                            context.append(volume[slice_idx, ny, nx])\n                        else:\n                            context.append(0.0)\n                \n                features = [x / width, y / height, intensity] + context\n                all_features.append(features)\n                all_coords.append((slice_idx, y, x))\n    \n    # Convert to array\n    all_features = np.array(all_features, dtype=np.float32)\n    \n    # Predict in batches\n    num_batches = (len(all_features) + batch_size - 1) // batch_size\n    all_predictions = []\n    \n    for batch_idx in tqdm(range(num_batches), desc=f\"    Predicting\"):\n        start_idx = batch_idx * batch_size\n        end_idx = min((batch_idx + 1) * batch_size, len(all_features))\n        \n        batch_features = all_features[start_idx:end_idx]\n        features_tensor = torch.FloatTensor(batch_features).to(config.device)\n        \n        with torch.no_grad():\n            batch_preds = torch.sigmoid(model(features_tensor))\n            batch_preds_np = batch_preds.cpu().numpy().flatten()\n        \n        all_predictions.extend(batch_preds_np)\n    \n    all_predictions = np.array(all_predictions, dtype=np.float32)\n    \n    return all_predictions, all_coords\n\n# Enhanced training function with comprehensive metrics - FIXED VERSION\ndef train_pixel_classifier_with_metrics(model, train_loader, val_loader, \n                                       train_features, train_labels,\n                                       val_features, val_labels,\n                                       train_volume, train_mask,\n                                       val_volume, val_mask,\n                                       test_volume, test_mask, test_fragment):\n    \"\"\"Train pixel classifier with comprehensive metrics tracking\"\"\"\n    \n    optimizer = optim.AdamW(model.parameters(), lr=config.lr, weight_decay=1e-4)\n    scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=7, verbose=True)\n    \n    criterion = WeightedBCELoss(pos_weight=config.pos_weight)\n    \n    best_val_dice = 0\n    patience_counter = 0\n    best_model_state = None\n    \n    # Comprehensive history tracking\n    history = {\n        'train': {\n            'loss': [], 'accuracy': [], 'precision': [], 'recall': [], 'dice': [], 'f05': []\n        },\n        'val': {\n            'loss': [], 'accuracy': [], 'precision': [], 'recall': [], 'dice': [], 'f05': []\n        },\n        'test': {\n            'accuracy': [], 'precision': [], 'recall': [], 'dice': [], 'f05': []\n        }\n    }\n    \n    for epoch in range(config.epochs):\n        # Training phase\n        model.train()\n        train_loss = 0\n        train_preds = []\n        train_targets = []\n        \n        for features, labels in tqdm(train_loader, desc=f'Epoch {epoch+1}'):\n            features, labels = features.to(config.device), labels.to(config.device)\n            \n            optimizer.zero_grad()\n            outputs = model(features)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            \n            # Store predictions for metrics\n            with torch.no_grad():\n                preds = torch.sigmoid(outputs)\n                train_preds.append(preds.cpu().numpy())\n                train_targets.append(labels.cpu().numpy())\n        \n        # Calculate training metrics\n        avg_train_loss = train_loss / len(train_loader)\n        train_preds = np.concatenate(train_preds).flatten()\n        train_targets = np.concatenate(train_targets).flatten()\n        train_metrics = calculate_comprehensive_metrics(train_preds, train_targets)\n        \n        # Update training history\n        history['train']['loss'].append(avg_train_loss)\n        history['train']['accuracy'].append(train_metrics['accuracy'])\n        history['train']['precision'].append(train_metrics['precision'])\n        history['train']['recall'].append(train_metrics['recall'])\n        history['train']['dice'].append(train_metrics['dice'])\n        history['train']['f05'].append(train_metrics['f05'])\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0\n        val_preds = []\n        val_targets = []\n        \n        with torch.no_grad():\n            for features, labels in val_loader:\n                features, labels = features.to(config.device), labels.to(config.device)\n                \n                outputs = model(features)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item()\n                \n                preds = torch.sigmoid(outputs)\n                val_preds.append(preds.cpu().numpy())\n                val_targets.append(labels.cpu().numpy())\n        \n        # Calculate validation metrics\n        avg_val_loss = val_loss / len(val_loader)\n        val_preds = np.concatenate(val_preds).flatten()\n        val_targets = np.concatenate(val_targets).flatten()\n        val_metrics = calculate_comprehensive_metrics(val_preds, val_targets)\n        \n        # Update validation history\n        history['val']['loss'].append(avg_val_loss)\n        history['val']['accuracy'].append(val_metrics['accuracy'])\n        history['val']['precision'].append(val_metrics['precision'])\n        history['val']['recall'].append(val_metrics['recall'])\n        history['val']['dice'].append(val_metrics['dice'])\n        history['val']['f05'].append(val_metrics['f05'])\n        \n        # Test phase (on actual test fragment)\n        test_metrics = None\n        if test_volume is not None and test_mask is not None:\n            print(f\"\\n  Evaluating on test fragment {test_fragment}...\")\n            \n            # Process test fragment\n            test_predictions, test_coords = process_fragment_for_evaluation(\n                model, test_volume, test_fragment\n            )\n            \n            # Reconstruct mask\n            test_reconstructed = reconstruct_volume_from_predictions(\n                test_predictions, test_coords, len(config.slices), config.img_size\n            )\n            \n            # Take middle slice for evaluation\n            middle_slice = len(config.slices) // 2\n            test_pred_slice = test_reconstructed[middle_slice]\n            test_pred_binary = (test_pred_slice > 0.5).astype(np.float32)\n            \n            # Calculate test metrics\n            test_metrics = calculate_comprehensive_metrics(test_pred_binary, test_mask)\n            \n            # Update test history\n            history['test']['accuracy'].append(test_metrics['accuracy'])\n            history['test']['precision'].append(test_metrics['precision'])\n            history['test']['recall'].append(test_metrics['recall'])\n            history['test']['dice'].append(test_metrics['dice'])\n            history['test']['f05'].append(test_metrics['f05'])\n            \n            # Save test visualization if metrics > threshold\n            if (test_metrics['dice'] >= config.min_score_threshold or \n                test_metrics['f05'] >= config.min_score_threshold):\n                input_slice = test_volume[middle_slice]\n                save_validation_visualization(\n                    input_slice, test_pred_binary, test_mask,\n                    epoch + 1, test_metrics, test_fragment, phase=\"test\"\n                )\n        \n        # Save validation visualization if metrics > threshold\n        if (val_metrics['dice'] >= config.min_score_threshold or \n            val_metrics['f05'] >= config.min_score_threshold):\n            print(f\"\\n  Creating validation visualization...\")\n            \n            # Process validation fragment\n            val_predictions, val_coords = process_fragment_for_evaluation(\n                model, val_volume, config.val_fragment\n            )\n            \n            val_reconstructed = reconstruct_volume_from_predictions(\n                val_predictions, val_coords, len(config.slices), config.img_size\n            )\n            \n            middle_slice = len(config.slices) // 2\n            val_pred_slice = (val_reconstructed[middle_slice] > 0.5).astype(np.float32)\n            input_slice = val_volume[middle_slice]\n            \n            save_validation_visualization(\n                input_slice, val_pred_slice, val_mask,\n                epoch + 1, val_metrics, config.val_fragment, phase=\"validation\"\n            )\n        \n        # Early stopping based on validation Dice\n        if val_metrics['dice'] > best_val_dice:\n            best_val_dice = val_metrics['dice']\n            patience_counter = 0\n            best_model_state = model.state_dict().copy()\n            print(f\"\\n🎯 New best validation Dice: {best_val_dice:.4f} at epoch {epoch+1}\")\n        else:\n            patience_counter += 1\n        \n        # Print comprehensive metrics\n        print(f\"\\n📊 Epoch {epoch+1}/{config.epochs} Summary:\")\n        print(f\"  TRAIN:\")\n        print(f\"    Loss: {avg_train_loss:.4f} | Acc: {train_metrics['accuracy']:.4f}\")\n        print(f\"    Dice: {train_metrics['dice']:.4f} | F0.5: {train_metrics['f05']:.4f}\")\n        print(f\"    Prec: {train_metrics['precision']:.4f} | Rec: {train_metrics['recall']:.4f}\")\n        \n        print(f\"  VALIDATION:\")\n        print(f\"    Loss: {avg_val_loss:.4f} | Acc: {val_metrics['accuracy']:.4f}\")\n        print(f\"    Dice: {val_metrics['dice']:.4f} | F0.5: {val_metrics['f05']:.4f}\")\n        print(f\"    Prec: {val_metrics['precision']:.4f} | Rec: {val_metrics['recall']:.4f}\")\n        \n        if test_metrics:\n            print(f\"  TEST:\")\n            print(f\"    Acc: {test_metrics['accuracy']:.4f} | Dice: {test_metrics['dice']:.4f}\")\n            print(f\"    F0.5: {test_metrics['f05']:.4f} | Prec: {test_metrics['precision']:.4f}\")\n            print(f\"    Rec: {test_metrics['recall']:.4f}\")\n        \n        # Learning rate scheduling based on validation Dice\n        scheduler.step(val_metrics['dice'])\n        \n        # Early stopping\n        if patience_counter >= config.early_stopping_patience:\n            print(f\"\\n🛑 Early stopping at epoch {epoch+1}\")\n            break\n    \n    # Load best model\n    if best_model_state:\n        model.load_state_dict(best_model_state)\n        print(f\"\\n✅ Loaded best model (Validation Dice: {best_val_dice:.4f})\")\n    \n    return model, history, best_val_dice\n\n# Main function\ndef main():\n    print(\"🚀 Starting PIXEL CLASSIFICATION with COMPREHENSIVE METRICS...\")\n    print(f\"Train on fragment {config.train_fragment}\")\n    print(f\"Validate on fragment {config.val_fragment}\")\n    print(f\"Test on fragment {config.test_fragment}\")\n    print(f\"Save visualizations when Dice/F0.5 > {config.min_score_threshold}\")\n    \n    # Load data\n    print(\"\\n📂 Loading data...\")\n    \n    # Load training fragment\n    train_volume = load_volume_data(config.train_path, config.slices)\n    train_mask = load_mask_data(config.train_path)\n    \n    if train_volume is None or train_mask is None:\n        print(f\"❌ Failed to load training fragment {config.train_fragment}\")\n        return\n    \n    # Load validation fragment\n    val_volume = load_volume_data(config.val_path, config.slices)\n    val_mask = load_mask_data(config.val_path)\n    \n    if val_volume is None or val_mask is None:\n        print(f\"❌ Failed to load validation fragment {config.val_fragment}\")\n        return\n    \n    # Load test fragment\n    test_volume = load_volume_data(config.test_path, config.slices)\n    test_mask = load_mask_data(config.test_path)\n    \n    if test_volume is None or test_mask is None:\n        print(f\"❌ Failed to load test fragment {config.test_fragment}\")\n        return\n    \n    print(f\"\\n📊 Dataset Summary:\")\n    print(f\"  Training fragment {config.train_fragment}: {train_volume.shape}\")\n    print(f\"  Validation fragment {config.val_fragment}: {val_volume.shape}\")\n    print(f\"  Test fragment {config.test_fragment}: {test_volume.shape}\")\n    \n    # Process fragments in parallel\n    print(\"\\n🔄 Processing fragments in parallel...\")\n    \n    # Process training fragment\n    train_features, train_labels, _ = process_fragment_parallel(\n        train_volume, train_mask, config.train_fragment, \n        use_context=config.use_context\n    )\n    \n    if train_features is None:\n        print(\"❌ Failed to process training fragment\")\n        return\n    \n    # Process validation fragment\n    val_features, val_labels, _ = process_fragment_parallel(\n        val_volume, val_mask, config.val_fragment,\n        use_context=config.use_context\n    )\n    \n    if val_features is None:\n        print(\"❌ Failed to process validation fragment\")\n        return\n    \n    print(f\"\\n📈 Processed Data Summary:\")\n    print(f\"  Training: {len(train_features):,} pixels\")\n    print(f\"  Validation: {len(val_features):,} pixels\")\n    \n    # Create datasets\n    train_dataset = PixelClassificationDataset(train_features, train_labels, augment=True)\n    val_dataset = PixelClassificationDataset(val_features, val_labels, augment=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True if torch.cuda.is_available() else False\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        num_workers=2\n    )\n    \n    # Initialize model\n    input_dim = train_features.shape[1]\n    print(f\"\\n📐 Input dimension: {input_dim}\")\n    \n    model = EnhancedPixelClassifier(input_dim, hidden_dim=config.hidden_dim).to(config.device)\n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"✅ Model initialized (Parameters: {total_params/1e6:.1f}M)\")\n    \n    # Train with comprehensive metrics\n    print(\"\\n🎯 Training pixel classifier with comprehensive metrics...\")\n    model, history, best_val_dice = train_pixel_classifier_with_metrics(\n        model, train_loader, val_loader,\n        train_features, train_labels,\n        val_features, val_labels,\n        train_volume, train_mask,\n        val_volume, val_mask,\n        test_volume, test_mask, config.test_fragment\n    )\n    \n    # Final evaluation\n    print(\"\\n🔬 Final evaluation on test fragment...\")\n    \n    # Process test fragment\n    test_predictions, test_coords = process_fragment_for_evaluation(\n        model, test_volume, config.test_fragment\n    )\n    \n    test_reconstructed = reconstruct_volume_from_predictions(\n        test_predictions, test_coords, len(config.slices), config.img_size\n    )\n    \n    middle_slice = len(config.slices) // 2\n    test_pred_slice = (test_reconstructed[middle_slice] > 0.5).astype(np.float32)\n    \n    final_metrics = calculate_comprehensive_metrics(test_pred_slice, test_mask)\n    \n    print(f\"\\n🎯 FINAL TEST RESULTS:\")\n    print(f\"  Accuracy: {final_metrics['accuracy']:.4f}\")\n    print(f\"  Dice: {final_metrics['dice']:.4f}\")\n    print(f\"  F0.5: {final_metrics['f05']:.4f}\")\n    print(f\"  Precision: {final_metrics['precision']:.4f}\")\n    print(f\"  Recall: {final_metrics['recall']:.4f}\")\n    \n    # Save final visualization\n    if (final_metrics['dice'] >= config.min_score_threshold or \n        final_metrics['f05'] >= config.min_score_threshold):\n        input_slice = test_volume[middle_slice]\n        save_validation_visualization(\n            input_slice, test_pred_slice, test_mask,\n            'final', final_metrics, config.test_fragment, phase=\"test\"\n        )\n    \n    # Plot comprehensive training history\n    plot_comprehensive_training_history(history, final_metrics)\n    \n    return final_metrics\n\ndef plot_comprehensive_training_history(history, final_metrics):\n    \"\"\"Plot comprehensive training history with all metrics\"\"\"\n    os.makedirs('training_plots', exist_ok=True)\n    \n    epochs = range(1, len(history['train']['loss']) + 1)\n    \n    # Create 3x2 grid of plots\n    fig, axes = plt.subplots(3, 2, figsize=(16, 15))\n    \n    # Plot 1: Loss\n    axes[0, 0].plot(epochs, history['train']['loss'], 'b-', label='Train Loss', linewidth=2)\n    axes[0, 0].plot(epochs, history['val']['loss'], 'r-', label='Val Loss', linewidth=2)\n    axes[0, 0].set_title('Training and Validation Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Plot 2: Accuracy\n    axes[0, 1].plot(epochs, history['train']['accuracy'], 'b-', label='Train Acc', linewidth=2)\n    axes[0, 1].plot(epochs, history['val']['accuracy'], 'r-', label='Val Acc', linewidth=2)\n    if history['test']['accuracy']:\n        axes[0, 1].plot(epochs[:len(history['test']['accuracy'])], \n                       history['test']['accuracy'], 'g-', label='Test Acc', linewidth=2)\n    axes[0, 1].set_title('Accuracy')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Accuracy')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # Plot 3: Dice Score\n    axes[1, 0].plot(epochs, history['train']['dice'], 'b-', label='Train Dice', linewidth=2)\n    axes[1, 0].plot(epochs, history['val']['dice'], 'r-', label='Val Dice', linewidth=2)\n    if history['test']['dice']:\n        axes[1, 0].plot(epochs[:len(history['test']['dice'])], \n                       history['test']['dice'], 'g-', label='Test Dice', linewidth=2)\n    axes[1, 0].set_title('Dice Score')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('Dice')\n    axes[1, 0].legend()\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # Plot 4: F0.5 Score\n    axes[1, 1].plot(epochs, history['train']['f05'], 'b-', label='Train F0.5', linewidth=2)\n    axes[1, 1].plot(epochs, history['val']['f05'], 'r-', label='Val F0.5', linewidth=2)\n    if history['test']['f05']:\n        axes[1, 1].plot(epochs[:len(history['test']['f05'])], \n                       history['test']['f05'], 'g-', label='Test F0.5', linewidth=2)\n    axes[1, 1].set_title('F0.5 Score')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('F0.5')\n    axes[1, 1].legend()\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    # Plot 5: Precision\n    axes[2, 0].plot(epochs, history['train']['precision'], 'b-', label='Train Prec', linewidth=2)\n    axes[2, 0].plot(epochs, history['val']['precision'], 'r-', label='Val Prec', linewidth=2)\n    if history['test']['precision']:\n        axes[2, 0].plot(epochs[:len(history['test']['precision'])], \n                       history['test']['precision'], 'g-', label='Test Prec', linewidth=2)\n    axes[2, 0].set_title('Precision')\n    axes[2, 0].set_xlabel('Epoch')\n    axes[2, 0].set_ylabel('Precision')\n    axes[2, 0].legend()\n    axes[2, 0].grid(True, alpha=0.3)\n    \n    # Plot 6: Recall\n    axes[2, 1].plot(epochs, history['train']['recall'], 'b-', label='Train Rec', linewidth=2)\n    axes[2, 1].plot(epochs, history['val']['recall'], 'r-', label='Val Rec', linewidth=2)\n    if history['test']['recall']:\n        axes[2, 1].plot(epochs[:len(history['test']['recall'])], \n                       history['test']['recall'], 'g-', label='Test Rec', linewidth=2)\n    axes[2, 1].set_title('Recall')\n    axes[2, 1].set_xlabel('Epoch')\n    axes[2, 1].set_ylabel('Recall')\n    axes[2, 1].legend()\n    axes[2, 1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.savefig('training_plots/comprehensive_training_history.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    # Plot final test metrics comparison\n    fig, ax = plt.subplots(figsize=(12, 6))\n    \n    metrics_names = ['Accuracy', 'Dice', 'F0.5', 'Precision', 'Recall']\n    train_final = [\n        history['train']['accuracy'][-1] if history['train']['accuracy'] else 0,\n        history['train']['dice'][-1] if history['train']['dice'] else 0,\n        history['train']['f05'][-1] if history['train']['f05'] else 0,\n        history['train']['precision'][-1] if history['train']['precision'] else 0,\n        history['train']['recall'][-1] if history['train']['recall'] else 0\n    ]\n    val_final = [\n        history['val']['accuracy'][-1] if history['val']['accuracy'] else 0,\n        history['val']['dice'][-1] if history['val']['dice'] else 0,\n        history['val']['f05'][-1] if history['val']['f05'] else 0,\n        history['val']['precision'][-1] if history['val']['precision'] else 0,\n        history['val']['recall'][-1] if history['val']['recall'] else 0\n    ]\n    test_final = [\n        final_metrics['accuracy'],\n        final_metrics['dice'],\n        final_metrics['f05'],\n        final_metrics['precision'],\n        final_metrics['recall']\n    ]\n    \n    x = np.arange(len(metrics_names))\n    width = 0.25\n    \n    ax.bar(x - width, train_final, width, label='Train', alpha=0.8)\n    ax.bar(x, val_final, width, label='Validation', alpha=0.8)\n    ax.bar(x + width, test_final, width, label='Test', alpha=0.8)\n    \n    ax.set_xlabel('Metrics')\n    ax.set_ylabel('Score')\n    ax.set_title('Final Metrics Comparison')\n    ax.set_xticks(x)\n    ax.set_xticklabels(metrics_names)\n    ax.legend()\n    ax.grid(True, alpha=0.3, axis='y')\n    \n    # Add value labels\n    for i, (train, val, test) in enumerate(zip(train_final, val_final, test_final)):\n        ax.text(i - width, train + 0.01, f'{train:.3f}', ha='center', va='bottom', fontsize=9)\n        ax.text(i, val + 0.01, f'{val:.3f}', ha='center', va='bottom', fontsize=9)\n        ax.text(i + width, test + 0.01, f'{test:.3f}', ha='center', va='bottom', fontsize=9)\n    \n    plt.tight_layout()\n    plt.savefig('training_plots/final_metrics_comparison.png', dpi=100, bbox_inches='tight')\n    plt.close()\n    \n    print(\"📊 Saved comprehensive training history plots\")\n\n# Run the main function\nif __name__ == \"__main__\":\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        print(f\"✅ Using GPU: {torch.cuda.get_device_name(0)}\")\n    \n    # Run pixel classification with comprehensive metrics\n    final_metrics = main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install scikit-learn\n!pip install segmentation-models-pytorch==0.3.0\n!pip install torch torchvision\n!pip install opencv-python\n!pip install scikit-image\n!pip install tqdm\n!pip install matplotlib\n!pip install scipy\n!pip install joblib\n!pip install albumentations\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing, binary_opening\nfrom skimage.morphology import remove_small_objects\nimport warnings\nimport random\nimport math\nfrom sklearn.metrics import precision_score, recall_score, f1_score, precision_recall_curve, roc_curve, auc\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, CosineAnnealingLR, OneCycleLR\nimport segmentation_models_pytorch as smp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport json\nimport pickle\n\nwarnings.filterwarnings('ignore')\n\n# ============================================================================\n# SIMPLIFIED CONFIGURATION - FIXED LOSS ISSUES\n# ============================================================================\nclass Config:\n    # Data paths\n    train_fragments = ['3', '2']  # Train on fragments 1 and 2\n    val_fragment = '1'  # Validate on fragment 3\n    test_fragment = '1'  # Test on fragment 3\n    base_path = '/kaggle/input/vesuvius-challenge-ink-detection/train'\n    \n    # Model parameters\n    img_size = 256  # Reduced for faster training\n    slice_range = list(range(20, 30))  # 10 slices\n    batch_size = 16  # Larger batch size\n    epochs = 100  # Reduced epochs for testing\n    \n    # Learning rate\n    lr = 1e-3  # Standard learning rate\n    \n    # SIMPLIFIED LOSS CONFIGURATION\n    # Use only BCE with high pos_weight to handle imbalance\n    pos_weight = 10.0  # Reasonable pos_weight\n    \n    # Augmentation\n    use_augmentation = True\n    \n    # Post-processing\n    min_area_size = 30\n    \n    # Early stopping\n    patience = 20\n    \n    # Threshold optimization\n    threshold_search_range = np.linspace(0.1, 0.9, 9).tolist()  # Simple search\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nconfig = Config()\n\n# ============================================================================\n# SIMPLE 2D UNET\n# ============================================================================\nclass SimpleUNet(nn.Module):\n    \"\"\"Simple 2D UNet\"\"\"\n    def __init__(self, in_channels=1, out_channels=1):\n        super().__init__()\n        \n        self.unet = smp.Unet(\n            encoder_name=\"resnet34\",  # Lighter backbone\n            encoder_weights=\"imagenet\",\n            in_channels=in_channels,\n            classes=out_channels,\n            encoder_depth=4,  # Shallower\n            decoder_channels=[128, 64, 32, 16],\n            activation='sigmoid'\n        )\n    \n    def forward(self, x):\n        return self.unet(x)\n\n# ============================================================================\n# BASIC DATA AUGMENTATION\n# ============================================================================\ndef basic_augment(image, mask):\n    \"\"\"Basic augmentations\"\"\"\n    # Random flips\n    if random.random() > 0.5:\n        image = np.fliplr(image).copy()\n        mask = np.fliplr(mask).copy()\n    \n    if random.random() > 0.5:\n        image = np.flipud(image).copy()\n        mask = np.flipud(mask).copy()\n    \n    # Random rotation (90 degree multiples)\n    if random.random() > 0.5:\n        angle = random.choice([0, 90, 180, 270])\n        if angle == 90:\n            image = np.rot90(image).copy()\n            mask = np.rot90(mask).copy()\n        elif angle == 180:\n            image = np.rot90(image, 2).copy()\n            mask = np.rot90(mask, 2).copy()\n        elif angle == 270:\n            image = np.rot90(image, 3).copy()\n            mask = np.rot90(mask, 3).copy()\n    \n    # Simple brightness/contrast\n    if random.random() > 0.5:\n        alpha = random.uniform(0.8, 1.2)\n        beta = random.uniform(-0.1, 0.1)\n        image = np.clip(alpha * image + beta, 0, 1)\n    \n    return image, mask\n\n# ============================================================================\n# SIMPLE DATA LOADING\n# ============================================================================\ndef load_single_slice_simple(base_path, slice_idx):\n    \"\"\"Load a single slice with simple preprocessing\"\"\"\n    slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n    if os.path.exists(slice_path):\n        img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n        if img is not None:\n            img_resized = cv2.resize(img, (config.img_size, config.img_size))\n            \n            # Simple preprocessing\n            # 1. CLAHE for contrast\n            clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n            img_clahe = clahe.apply(img_resized)\n            \n            # 2. Normalize\n            img_normalized = img_clahe.astype(np.float32) / 255.0\n            \n            return img_normalized\n    return None\n\ndef load_mask_simple(base_path):\n    \"\"\"Load and enhance mask\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            \n            return mask_binary\n    return None\n\nclass SimpleDataset(Dataset):\n    \"\"\"Simple dataset\"\"\"\n    \n    def __init__(self, fragment_paths, slice_indices, augment=True):\n        self.samples = []\n        self.augment = augment\n        \n        for frag_path in fragment_paths:\n            # Load mask\n            mask = load_mask_simple(frag_path)\n            if mask is None:\n                continue\n            \n            # Load slices\n            for slice_idx in slice_indices:\n                slice_img = load_single_slice_simple(frag_path, slice_idx)\n                if slice_img is not None:\n                    self.samples.append({\n                        'image': slice_img,\n                        'mask': mask,\n                        'fragment': os.path.basename(frag_path),\n                        'slice_idx': slice_idx\n                    })\n        \n        print(f\"Created dataset with {len(self.samples)} samples\")\n        \n        # Calculate positive ratio\n        if self.samples:\n            all_masks = np.stack([s['mask'] for s in self.samples])\n            positive_ratio = all_masks.sum() / all_masks.size\n            print(f\"Positive pixel ratio: {positive_ratio:.4f}\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        image = sample['image'].copy()\n        mask = sample['mask'].copy()\n        \n        # Apply augmentation\n        if self.augment:\n            image, mask = basic_augment(image, mask)\n        \n        # Add channel dimension\n        image_tensor = torch.FloatTensor(image).unsqueeze(0)  # [1, H, W]\n        mask_tensor = torch.FloatTensor(mask).unsqueeze(0)    # [1, H, W]\n        \n        return image_tensor, mask_tensor\n\n# ============================================================================\n# SIMPLE LOSS FUNCTION - FIXED!\n# ============================================================================\nclass SimpleDiceLoss(nn.Module):\n    \"\"\"Simple Dice loss\"\"\"\n    def __init__(self, smooth=1e-5):\n        super().__init__()\n        self.smooth = smooth\n\n    def forward(self, predictions, targets):\n        predictions = torch.sigmoid(predictions)\n        \n        # Flatten predictions and targets\n        predictions_flat = predictions.view(-1)\n        targets_flat = targets.view(-1)\n        \n        intersection = (predictions_flat * targets_flat).sum()\n        union = predictions_flat.sum() + targets_flat.sum()\n        \n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        return 1 - dice\n\nclass SimpleCombinedLoss(nn.Module):\n    \"\"\"Simple combined loss with proper scaling\"\"\"\n    def __init__(self, dice_weight=0.5, bce_weight=0.5, pos_weight=10.0):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.pos_weight = pos_weight\n        \n        self.dice_loss = SimpleDiceLoss()\n    \n    def forward(self, pred, target):\n        # BCE with proper pos_weight\n        bce_loss = F.binary_cross_entropy_with_logits(\n            pred, \n            target,\n            pos_weight=torch.tensor([self.pos_weight], device=pred.device),\n            reduction='mean'\n        )\n        \n        # Dice loss\n        dice = self.dice_loss(pred, target)\n        \n        # PROPERLY SCALED COMBINED LOSS\n        total_loss = self.dice_weight * dice + self.bce_weight * bce_loss\n        \n        return total_loss, {\n            'dice': dice.item(),\n            'bce': bce_loss.item(),\n            'total': total_loss.item()\n        }\n\n# ============================================================================\n# THRESHOLD OPTIMIZATION\n# ============================================================================\ndef find_optimal_threshold_simple(pred_probs, targets, thresholds=None):\n    \"\"\"Simple threshold optimization\"\"\"\n    if thresholds is None:\n        thresholds = config.threshold_search_range\n    \n    best_threshold = 0.5\n    best_dice = 0\n    \n    for thresh in thresholds:\n        pred_binary = (pred_probs > thresh).astype(np.float32)\n        \n        # Calculate Dice\n        intersection = (pred_binary * targets).sum()\n        union = pred_binary.sum() + targets.sum()\n        dice = (2. * intersection) / (union + 1e-6)\n        \n        if dice > best_dice:\n            best_dice = dice\n            best_threshold = thresh\n    \n    return best_threshold, best_dice\n\n# ============================================================================\n# SIMPLE TRAINING FUNCTION - CORRECTED!\n# ============================================================================\ndef train_model_simple(train_loader, val_loader, model_name=\"model\"):\n    \"\"\"Simple training function with correct loss scaling\"\"\"\n    model = SimpleUNet(in_channels=1).to(config.device)\n    \n    # SIMPLE LOSS - properly scaled\n    criterion = SimpleCombinedLoss(\n        dice_weight=0.5,\n        bce_weight=0.5,\n        pos_weight=config.pos_weight\n    )\n    \n    # Simple optimizer\n    optimizer = optim.Adam(model.parameters(), lr=config.lr)\n    \n    # Simple scheduler\n    scheduler = CosineAnnealingLR(optimizer, T_max=config.epochs, eta_min=1e-5)\n    \n    best_val_dice = 0\n    best_model_state = None\n    train_history = []\n    val_history = []\n    \n    for epoch in range(1, config.epochs + 1):\n        # Training phase\n        model.train()\n        train_losses = []\n        all_train_preds = []\n        all_train_targets = []\n        \n        pbar = tqdm(train_loader, desc=f'Epoch {epoch}/{config.epochs}')\n        for images, masks in pbar:\n            images = images.to(config.device)\n            masks = masks.to(config.device)\n            \n            optimizer.zero_grad()\n            outputs = model(images)\n            loss, loss_components = criterion(outputs, masks)\n            loss.backward()\n            \n            # Gentle gradient clipping\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            \n            optimizer.step()\n            \n            train_losses.append(loss_components['total'])\n            \n            # Store predictions for metrics\n            with torch.no_grad():\n                preds = torch.sigmoid(outputs).cpu().numpy()\n                all_train_preds.append(preds)\n                all_train_targets.append(masks.cpu().numpy())\n            \n            pbar.set_postfix({\n                'Loss': f'{loss_components[\"total\"]:.4f}',\n                'DiceL': f'{loss_components[\"dice\"]:.4f}'\n            })\n        \n        # Calculate training metrics\n        if all_train_preds:\n            train_preds = np.concatenate(all_train_preds)\n            train_targets = np.concatenate(all_train_targets)\n            \n            # Find optimal threshold\n            train_optimal_threshold, train_dice_opt = find_optimal_threshold_simple(train_preds, train_targets)\n            \n            # Calculate default threshold dice\n            train_preds_binary_default = (train_preds > 0.5).astype(np.float32)\n            intersection_default = (train_preds_binary_default * train_targets).sum()\n            union_default = train_preds_binary_default.sum() + train_targets.sum()\n            train_dice_default = (2. * intersection_default) / (union_default + 1e-6)\n        else:\n            train_optimal_threshold = 0.5\n            train_dice_opt = 0\n            train_dice_default = 0\n        \n        avg_train_loss = np.mean(train_losses)\n        \n        # Validation phase\n        model.eval()\n        val_losses = []\n        all_val_preds = []\n        all_val_targets = []\n        \n        with torch.no_grad():\n            for images, masks in val_loader:\n                images = images.to(config.device)\n                masks = masks.to(config.device)\n                \n                outputs = model(images)\n                loss, loss_components = criterion(outputs, masks)\n                val_losses.append(loss_components['total'])\n                \n                preds = torch.sigmoid(outputs).cpu().numpy()\n                all_val_preds.append(preds)\n                all_val_targets.append(masks.cpu().numpy())\n        \n        # Calculate validation metrics\n        if all_val_preds:\n            val_preds = np.concatenate(all_val_preds)\n            val_targets = np.concatenate(all_val_targets)\n            \n            # Find optimal threshold\n            val_optimal_threshold, val_dice_opt = find_optimal_threshold_simple(val_preds, val_targets)\n            \n            # Calculate default threshold dice\n            val_preds_binary_default = (val_preds > 0.5).astype(np.float32)\n            intersection_default = (val_preds_binary_default * val_targets).sum()\n            union_default = val_preds_binary_default.sum() + val_targets.sum()\n            val_dice_default = (2. * intersection_default) / (union_default + 1e-6)\n        else:\n            val_optimal_threshold = 0.5\n            val_dice_opt = 0\n            val_dice_default = 0\n        \n        avg_val_loss = np.mean(val_losses)\n        \n        # Update learning rate\n        scheduler.step()\n        current_lr = optimizer.param_groups[0]['lr']\n        \n        # Print metrics\n        print(f\"\\nEpoch {epoch}/{config.epochs}:\")\n        print(f\"  LR: {current_lr:.2e}, Train Loss: {avg_train_loss:.4f}, Val Loss: {avg_val_loss:.4f}\")\n        print(f\"  Train Dice: {train_dice_opt:.4f} (thresh={train_optimal_threshold:.2f}), \"\n              f\"Default: {train_dice_default:.4f}\")\n        print(f\"  Val Dice:   {val_dice_opt:.4f} (thresh={val_optimal_threshold:.2f}), \"\n              f\"Default: {val_dice_default:.4f}\")\n        \n        # Store history\n        train_history.append({\n            'epoch': epoch,\n            'loss': float(avg_train_loss),\n            'dice_optimal': float(train_dice_opt),\n            'dice_default': float(train_dice_default),\n            'threshold': float(train_optimal_threshold)\n        })\n        \n        val_history.append({\n            'epoch': epoch,\n            'loss': float(avg_val_loss),\n            'dice_optimal': float(val_dice_opt),\n            'dice_default': float(val_dice_default),\n            'threshold': float(val_optimal_threshold)\n        })\n        \n        # Save best model\n        if val_dice_opt > best_val_dice:\n            best_val_dice = val_dice_opt\n            best_model_state = model.state_dict().copy()\n            best_epoch = epoch\n            best_threshold = val_optimal_threshold\n            print(f\"  🎯 New best validation Dice: {best_val_dice:.4f} at threshold {best_threshold:.2f}\")\n        \n        # Early stopping\n        if epoch > config.patience:\n            recent_dice = [v['dice_optimal'] for v in val_history[-config.patience:]]\n            if max(recent_dice) < best_val_dice * 0.98:\n                print(f\"  ⏹️  Early stopping at epoch {epoch}\")\n                break\n    \n    # Load best model\n    if best_model_state:\n        model.load_state_dict(best_model_state)\n        print(f\"\\nBest model from epoch {best_epoch}:\")\n        print(f\"  Dice: {best_val_dice:.4f} at threshold {best_threshold:.2f}\")\n    \n    # Save model\n    torch.save(model.state_dict(), f'{model_name}.pth')\n    \n    return model, best_val_dice, best_threshold, train_history, val_history\n\n# ============================================================================\n# MAIN PIPELINE\n# ============================================================================\ndef main_pipeline_simple():\n    \"\"\"Simple main pipeline\"\"\"\n    print(\"=\" * 70)\n    print(\"SIMPLE TRAINING PIPELINE - FIXED LOSS\")\n    print(\"=\" * 70)\n    print(f\"Train: {config.train_fragments}, Val/Test: {config.val_fragment}\")\n    print(\"=\" * 70)\n    \n    # Set seeds\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    print(f\"Device: {config.device}\")\n    print(f\"Image size: {config.img_size}\")\n    print(f\"Batch size: {config.batch_size}\")\n    print(f\"Epochs: {config.epochs}\")\n    print(f\"LR: {config.lr}, Pos weight: {config.pos_weight}\")\n    \n    # Create datasets\n    train_paths = [os.path.join(config.base_path, f) for f in config.train_fragments]\n    val_paths = [os.path.join(config.base_path, config.val_fragment)]\n    \n    train_dataset = SimpleDataset(train_paths, config.slice_range, augment=True)\n    val_dataset = SimpleDataset(val_paths, config.slice_range, augment=False)\n    \n    print(f\"\\nDataset sizes:\")\n    print(f\"  Train: {len(train_dataset)} samples\")\n    print(f\"  Val:   {len(val_dataset)} samples\")\n    \n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=config.batch_size, \n        shuffle=True,\n        num_workers=0, \n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, \n        batch_size=config.batch_size, \n        shuffle=False, \n        num_workers=0, \n        pin_memory=True\n    )\n    \n    # Train model\n    model_name = f\"simple_model_fragment_{config.val_fragment}\"\n    print(f\"\\n{'='*70}\")\n    print(f\"STARTING TRAINING\")\n    print(f\"{'='*70}\")\n    \n    model, val_dice, optimal_threshold, train_history, val_history = train_model_simple(\n        train_loader, val_loader, model_name\n    )\n    \n    # Store results\n    results = {\n        'train_fragments': config.train_fragments,\n        'val_fragment': config.val_fragment,\n        'val_dice': float(val_dice),\n        'optimal_threshold': float(optimal_threshold),\n        'model_path': f'{model_name}.pth'\n    }\n    \n    print(f\"\\nTraining completed!\")\n    print(f\"  Best validation Dice: {val_dice:.4f}\")\n    print(f\"  Optimal threshold: {optimal_threshold:.2f}\")\n    \n    return results, model, optimal_threshold, train_history, val_history\n\n# ============================================================================\n# TESTING FUNCTION\n# ============================================================================\ndef test_model_simple(model, fragment_path, optimal_threshold):\n    \"\"\"Test the model\"\"\"\n    model.eval()\n    \n    # Load all slices\n    all_slices = []\n    for slice_idx in config.slice_range:\n        slice_img = load_single_slice_simple(fragment_path, slice_idx)\n        if slice_img is not None:\n            all_slices.append(slice_img)\n    \n    if not all_slices:\n        return None, None, None\n    \n    # Predict each slice\n    all_predictions = []\n    with torch.no_grad():\n        for slice_img in all_slices:\n            slice_tensor = torch.FloatTensor(slice_img).unsqueeze(0).unsqueeze(0).to(config.device)\n            pred = model(slice_tensor)\n            pred_sigmoid = torch.sigmoid(pred)\n            all_predictions.append(pred_sigmoid[0, 0].cpu().numpy())\n    \n    # Average predictions\n    avg_prediction = np.mean(all_predictions, axis=0)\n    \n    # Load ground truth\n    test_mask = load_mask_simple(fragment_path)\n    \n    if test_mask is None:\n        return None, None, None\n    \n    # Calculate metrics with optimal threshold\n    pred_binary_opt = (avg_prediction > optimal_threshold).astype(np.float32)\n    intersection_opt = (pred_binary_opt * test_mask).sum()\n    union_opt = pred_binary_opt.sum() + test_mask.sum()\n    dice_opt = (2. * intersection_opt) / (union_opt + 1e-6)\n    \n    # Calculate metrics with default threshold\n    pred_binary_default = (avg_prediction > 0.5).astype(np.float32)\n    intersection_default = (pred_binary_default * test_mask).sum()\n    union_default = pred_binary_default.sum() + test_mask.sum()\n    dice_default = (2. * intersection_default) / (union_default + 1e-6)\n    \n    return dice_opt, dice_default, avg_prediction, test_mask\n\n# ============================================================================\n# RUN PIPELINE\n# ============================================================================\nif __name__ == \"__main__\":\n    print(\"Starting simple training pipeline with fixed loss...\")\n    \n    # Run training\n    results, model, optimal_threshold, train_history, val_history = main_pipeline_simple()\n    \n    # Run testing\n    print(f\"\\n{'='*70}\")\n    print(f\"TESTING ON FRAGMENT {config.val_fragment}\")\n    print(f\"{'='*70}\")\n    \n    test_path = os.path.join(config.base_path, config.val_fragment)\n    dice_opt, dice_default, predictions, ground_truth = test_model_simple(\n        model, test_path, optimal_threshold\n    )\n    \n    if dice_opt is not None:\n        print(f\"\\nTest Results:\")\n        print(f\"  Optimal threshold ({optimal_threshold:.2f}): Dice = {dice_opt:.4f}\")\n        print(f\"  Default threshold (0.50): Dice = {dice_default:.4f}\")\n        print(f\"  Improvement: +{(dice_opt - dice_default):.4f} \"\n              f\"(+{(dice_opt/dice_default - 1)*100:.1f}%)\")\n    \n    # Plot training curves\n    if train_history and val_history:\n        epochs = [h['epoch'] for h in train_history]\n        \n        fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n        \n        # Loss\n        axes[0].plot(epochs, [h['loss'] for h in train_history], 'b-', label='Train')\n        axes[0].plot(epochs, [h['loss'] for h in val_history], 'r-', label='Val')\n        axes[0].set_xlabel('Epoch')\n        axes[0].set_ylabel('Loss')\n        axes[0].set_title('Training Loss')\n        axes[0].legend()\n        axes[0].grid(True)\n        \n        # Dice (optimal)\n        axes[1].plot(epochs, [h['dice_optimal'] for h in train_history], 'b-', label='Train')\n        axes[1].plot(epochs, [h['dice_optimal'] for h in val_history], 'r-', label='Val')\n        axes[1].set_xlabel('Epoch')\n        axes[1].set_ylabel('Dice Score')\n        axes[1].set_title('Dice Score (Optimal Threshold)')\n        axes[1].legend()\n        axes[1].grid(True)\n        \n        # Dice (default vs optimal)\n        axes[2].plot(epochs, [h['dice_default'] for h in val_history], 'g-', label='Val (Default 0.5)')\n        axes[2].plot(epochs, [h['dice_optimal'] for h in val_history], 'r-', label='Val (Optimal)')\n        axes[2].set_xlabel('Epoch')\n        axes[2].set_ylabel('Dice Score')\n        axes[2].set_title('Default vs Optimal Threshold')\n        axes[2].legend()\n        axes[2].grid(True)\n        \n        plt.tight_layout()\n        plt.savefig('simple_training_results.png', dpi=150, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"\\nTraining curves saved to 'simple_training_results.png'\")\n    \n    # Save results\n    final_summary = {\n        'config': config.__dict__,\n        'results': results,\n        'train_history': train_history,\n        'val_history': val_history\n    }\n    \n    with open('simple_results.json', 'w') as f:\n        json.dump(final_summary, f, indent=2, default=str)\n    \n    print(f\"\\nResults saved to 'simple_results.json'\")\n    print(f\"Model saved as: {results['model_path']}\")\n    print(\"\\n\" + \"=\"*70)\n    print(\"PIPELINE COMPLETE!\")\n    print(\"=\"*70)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install scikit-learn\n!pip install segmentation-models-pytorch==0.3.0\n!pip install torch torchvision\n!pip install opencv-python\n!pip install scikit-image\n!pip install tqdm\n!pip install matplotlib\n!pip install scipy\n!pip install joblib\n!pip install albumentations\n\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom scipy import ndimage\nfrom scipy.ndimage import binary_fill_holes, binary_closing, binary_opening\nfrom skimage.morphology import remove_small_objects\nimport warnings\nimport random\nimport math\nfrom sklearn.metrics import precision_score, recall_score, f1_score\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nwarnings.filterwarnings('ignore')\n\n# ============================================================================\n# CONFIGURATION - SIMPLIFIED FOR BETTER LEARNING\n# ============================================================================\nclass Config:\n    # Data paths\n    train_fragments = ['1', '2', '3']\n    base_path = '/kaggle/input/vesuvius-challenge-ink-detection/train'\n    \n    # Model parameters - GOING BACK TO BASICS\n    img_size = 384\n    slice_range = list(range(20, 25))  # FOCUS ON KEY SLICES (5 slices instead of 19)\n    batch_size = 4\n    epochs = 200\n    lr = 3e-4  # Higher learning rate\n    \n    # Processing - SINGLE SLICE APPROACH (proven to work)\n    use_multi_slice = False  # Disable 2.5D for now\n    \n    # Loss weights - ADJUST FOR IMBALANCED DATA\n    dice_weight = 0.5\n    wbce_weight = 0.5\n    pos_weight = 10.0  # Weight for positive class (ink is rare)\n    \n    # Augmentation\n    use_augmentation = True\n    \n    # Post-processing\n    min_area_size = 50\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nconfig = Config()\n\n# ============================================================================\n# SIMPLE 2D UNET - BACK TO BASICS\n# ============================================================================\nclass SimpleUNet(nn.Module):\n    \"\"\"Simple 2D UNet that works\"\"\"\n    def __init__(self, in_channels=1, out_channels=1):\n        super().__init__()\n        \n        self.unet = smp.Unet(\n            encoder_name=\"resnet34\",\n            encoder_weights=\"imagenet\",\n            in_channels=in_channels,\n            classes=out_channels,\n            encoder_depth=4,\n            decoder_channels=[128, 64, 32, 16]\n        )\n    \n    def forward(self, x):\n        return self.unet(x)\n\n# ============================================================================\n# SIMPLE DATA AUGMENTATION\n# ============================================================================\ndef simple_augment(image, mask):\n    \"\"\"Simple augmentations that work\"\"\"\n    # Random flips\n    if random.random() > 0.5:\n        image = np.fliplr(image).copy()\n        mask = np.fliplr(mask).copy()\n    \n    if random.random() > 0.5:\n        image = np.flipud(image).copy()\n        mask = np.flipud(mask).copy()\n    \n    # Random rotation\n    if random.random() > 0.5:\n        k = random.randint(1, 3)\n        image = np.rot90(image, k).copy()\n        mask = np.rot90(mask, k).copy()\n    \n    # Brightness/contrast\n    if random.random() > 0.5:\n        alpha = random.uniform(0.9, 1.1)\n        beta = random.uniform(-0.05, 0.05)\n        image = np.clip(alpha * image + beta, 0, 1)\n    \n    # Add noise\n    if random.random() > 0.5:\n        noise = np.random.normal(0, 0.02, image.shape)\n        image = np.clip(image + noise, 0, 1)\n    \n    return image, mask\n\n# ============================================================================\n# DATA LOADING - SIMPLIFIED\n# ============================================================================\ndef load_single_slice(base_path, slice_idx):\n    \"\"\"Load a single slice with preprocessing\"\"\"\n    slice_path = os.path.join(base_path, 'surface_volume', f'{slice_idx:02d}.tif')\n    if os.path.exists(slice_path):\n        img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n        if img is not None:\n            img_resized = cv2.resize(img, (config.img_size, config.img_size))\n            \n            # Basic preprocessing\n            img_filtered = cv2.bilateralFilter(img_resized, 5, 50, 50)\n            clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n            img_enhanced = clahe.apply(img_filtered)\n            \n            # Normalize\n            img_normalized = img_enhanced.astype(np.float32) / 255.0\n            \n            return img_normalized\n    return None\n\ndef load_mask(base_path):\n    \"\"\"Load mask\"\"\"\n    mask_path = os.path.join(base_path, 'inklabels.png')\n    if os.path.exists(mask_path):\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is not None:\n            mask_resized = cv2.resize(mask, (config.img_size, config.img_size))\n            mask_binary = (mask_resized > 0).astype(np.float32)\n            return mask_binary\n    return None\n\nclass SingleSliceDataset(Dataset):\n    \"\"\"Dataset for single slice processing\"\"\"\n    \n    def __init__(self, fragment_paths, slice_indices, augment=True):\n        self.samples = []\n        self.augment = augment\n        \n        for frag_path in fragment_paths:\n            # Load mask once per fragment\n            mask = load_mask(frag_path)\n            if mask is None:\n                continue\n            \n            # Load slices\n            for slice_idx in slice_indices:\n                slice_img = load_single_slice(frag_path, slice_idx)\n                if slice_img is not None:\n                    self.samples.append({\n                        'image': slice_img,\n                        'mask': mask,\n                        'fragment': os.path.basename(frag_path),\n                        'slice_idx': slice_idx\n                    })\n        \n        print(f\"Created dataset with {len(self.samples)} samples\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        image = sample['image'].copy()\n        mask = sample['mask'].copy()\n        \n        # Apply augmentation\n        if self.augment:\n            image, mask = simple_augment(image, mask)\n        \n        # Add channel dimension\n        image_tensor = torch.FloatTensor(image).unsqueeze(0)  # [1, H, W]\n        mask_tensor = torch.FloatTensor(mask).unsqueeze(0)    # [1, H, W]\n        \n        return image_tensor, mask_tensor\n\n# ============================================================================\n# LOSS FUNCTIONS - WITH CLASS WEIGHTS\n# ============================================================================\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1e-6):\n        super(DiceLoss, self).__init__()\n        self.smooth = smooth\n\n    def forward(self, predictions, targets):\n        predictions = torch.sigmoid(predictions)\n        \n        predictions_flat = predictions.view(-1)\n        targets_flat = targets.view(-1)\n        \n        intersection = (predictions_flat * targets_flat).sum()\n        union = predictions_flat.sum() + targets_flat.sum()\n        \n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        return 1 - dice\n\nclass WeightedBCELoss(nn.Module):\n    def __init__(self, pos_weight=1.0):\n        super(WeightedBCELoss, self).__init__()\n        self.pos_weight = torch.tensor([pos_weight])\n        \n    def forward(self, predictions, targets):\n        if self.pos_weight.device != predictions.device:\n            self.pos_weight = self.pos_weight.to(predictions.device)\n        \n        return F.binary_cross_entropy_with_logits(\n            predictions, targets, pos_weight=self.pos_weight\n        )\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, dice_weight=0.5, wbce_weight=0.5, pos_weight=10.0):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.wbce_weight = wbce_weight\n        \n        self.dice_loss = DiceLoss()\n        self.wbce_loss = WeightedBCELoss(pos_weight=pos_weight)\n    \n    def forward(self, pred, target):\n        dice = self.dice_loss(pred, target)\n        wbce = self.wbce_loss(pred, target)\n        \n        total_loss = self.dice_weight * dice + self.wbce_weight * wbce\n        \n        return total_loss, {'dice': dice.item(), 'wbce': wbce.item()}\n\n# ============================================================================\n# METRICS\n# ============================================================================\ndef dice_score(pred, target, smooth=1e-6):\n    pred = pred.flatten()\n    target = target.flatten()\n    intersection = (pred * target).sum()\n    return (2. * intersection + smooth) / (pred.sum() + target.sum() + smooth)\n\ndef calculate_metrics(pred, target, threshold=0.5):\n    \"\"\"Calculate metrics with adaptive thresholding\"\"\"\n    # Try multiple thresholds\n    best_dice = 0\n    best_threshold = 0.5\n    \n    for thresh in [0.3, 0.4, 0.5, 0.6, 0.7]:\n        pred_binary = (pred > thresh).astype(np.float32)\n        dice = dice_score(pred_binary, target)\n        \n        if dice > best_dice:\n            best_dice = dice\n            best_threshold = thresh\n    \n    # Use best threshold\n    pred_binary = (pred > best_threshold).astype(np.float32)\n    \n    if pred_binary.sum() == 0 and target.sum() == 0:\n        return {'dice': 1.0, 'f05': 1.0, 'threshold': best_threshold}\n    \n    metrics = {'dice': best_dice, 'threshold': best_threshold}\n    \n    if pred_binary.sum() > 0 or target.sum() > 0:\n        pred_flat = pred_binary.flatten().astype(int)\n        target_flat = target.flatten().astype(int)\n        \n        # Calculate F0.5\n        precision = precision_score(target_flat, pred_flat, zero_division=0)\n        recall = recall_score(target_flat, pred_flat, zero_division=0)\n        \n        if precision + recall > 0:\n            metrics['f05'] = (1 + 0.5**2) * (precision * recall) / ((0.5**2 * precision) + recall)\n        else:\n            metrics['f05'] = 0.0\n        \n        metrics['precision'] = precision\n        metrics['recall'] = recall\n    \n    return metrics\n\n# ============================================================================\n# POST-PROCESSING - FIXED\n# ============================================================================\ndef post_process(prediction, min_area=50):\n    \"\"\"Apply post-processing to prediction\"\"\"\n    # Adaptive threshold\n    if prediction.max() > 0:\n        threshold = np.percentile(prediction[prediction > 0], 70)\n    else:\n        threshold = 0.5\n    \n    binary = (prediction > threshold).astype(np.uint8)\n    \n    # Morphological operations\n    kernel = np.ones((3, 3), np.uint8)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=1)\n    binary = cv2.morphologyEx(binary, cv2.MORPH_OPEN, kernel, iterations=1)\n    \n    # Remove small objects\n    if binary.sum() > 0:\n        binary = remove_small_objects(binary.astype(bool), min_size=min_area)\n        binary = binary.astype(np.uint8)\n    \n    # Fill holes\n    binary = binary_fill_holes(binary).astype(np.uint8)\n    \n    return binary.astype(np.float32)\n\n# ============================================================================\n# TRAINING FUNCTION\n# ============================================================================\ndef train_model(train_loader, val_loader, model_name=\"model\"):\n    \"\"\"Train a single model\"\"\"\n    model = SimpleUNet(in_channels=1).to(config.device)\n    \n    criterion = CombinedLoss(\n        dice_weight=config.dice_weight,\n        wbce_weight=config.wbce_weight,\n        pos_weight=config.pos_weight\n    )\n    \n    optimizer = optim.Adam(model.parameters(), lr=config.lr)\n    scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=10, verbose=True)\n    \n    best_val_dice = 0\n    best_model_state = None\n    \n    for epoch in range(config.epochs):\n        # Training\n        model.train()\n        train_loss = 0\n        train_dice_loss = 0\n        train_wbce_loss = 0\n        \n        pbar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.epochs}')\n        for images, masks in pbar:\n            images = images.to(config.device)\n            masks = masks.to(config.device)\n            \n            optimizer.zero_grad()\n            outputs = model(images)\n            loss, loss_components = criterion(outputs, masks)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            train_dice_loss += loss_components['dice']\n            train_wbce_loss += loss_components['wbce']\n            \n            pbar.set_postfix({\n                'Loss': f'{loss.item():.4f}',\n                'DiceL': f'{loss_components[\"dice\"]:.4f}'\n            })\n        \n        avg_train_loss = train_loss / len(train_loader)\n        avg_train_dice = train_dice_loss / len(train_loader)\n        avg_train_wbce = train_wbce_loss / len(train_loader)\n        \n        # Validation\n        model.eval()\n        val_predictions = []\n        val_targets = []\n        \n        with torch.no_grad():\n            for images, masks in val_loader:\n                images = images.to(config.device)\n                outputs = torch.sigmoid(model(images))\n                val_predictions.append(outputs.cpu().numpy())\n                val_targets.append(masks.cpu().numpy())\n        \n        if val_predictions:\n            val_predictions = np.concatenate(val_predictions)\n            val_targets = np.concatenate(val_targets)\n            val_metrics = calculate_metrics(val_predictions, val_targets)\n        else:\n            val_metrics = {'dice': 0, 'f05': 0}\n        \n        print(f\"Epoch {epoch+1}: Train Loss={avg_train_loss:.4f}, \"\n              f\"Val Dice={val_metrics['dice']:.4f}, Val F0.5={val_metrics.get('f05', 0):.4f}\")\n        \n        # Update scheduler\n        scheduler.step(val_metrics['dice'])\n        \n        # Save best model\n        if val_metrics['dice'] > best_val_dice:\n            best_val_dice = val_metrics['dice']\n            best_model_state = model.state_dict().copy()\n            print(f\"  New best validation Dice: {best_val_dice:.4f}\")\n        \n        # Early stopping\n        if epoch > 20 and val_metrics['dice'] < best_val_dice * 0.95:\n            patience_counter = getattr(model, 'patience_counter', 0) + 1\n            model.patience_counter = patience_counter\n            \n            if patience_counter >= 15:\n                print(f\"Early stopping at epoch {epoch+1}\")\n                break\n    \n    # Load best model\n    if best_model_state:\n        model.load_state_dict(best_model_state)\n    \n    # Save model\n    torch.save(model.state_dict(), f'{model_name}.pth')\n    \n    return model, best_val_dice\n\n# ============================================================================\n# PREDICTION FUNCTION\n# ============================================================================\ndef predict_fragment(model, fragment_path):\n    \"\"\"Predict on all slices of a fragment and average results\"\"\"\n    model.eval()\n    \n    # Load all slices\n    all_slices = []\n    for slice_idx in config.slice_range:\n        slice_img = load_single_slice(fragment_path, slice_idx)\n        if slice_img is not None:\n            all_slices.append(slice_img)\n    \n    if not all_slices:\n        return None\n    \n    # Predict each slice\n    all_predictions = []\n    for slice_img in all_slices:\n        slice_tensor = torch.FloatTensor(slice_img).unsqueeze(0).unsqueeze(0).to(config.device)\n        \n        with torch.no_grad():\n            pred = torch.sigmoid(model(slice_tensor))\n            all_predictions.append(pred[0, 0].cpu().numpy())\n    \n    # Average predictions\n    avg_prediction = np.mean(all_predictions, axis=0)\n    \n    return avg_prediction\n\n# ============================================================================\n# LEAVE-ONE-OUT CROSS VALIDATION\n# ============================================================================\ndef leave_one_out_cross_validation():\n    \"\"\"Perform leave-one-out cross validation\"\"\"\n    print(\"Starting leave-one-out cross validation...\")\n    \n    all_fragments = config.train_fragments\n    results = {}\n    all_models = []\n    \n    for test_fragment in all_fragments:\n        print(f\"\\n{'='*60}\")\n        print(f\"Testing on fragment {test_fragment}\")\n        print(f\"{'='*60}\")\n        \n        # Prepare training fragments\n        train_fragments = [f for f in all_fragments if f != test_fragment]\n        \n        # Create datasets\n        train_paths = [os.path.join(config.base_path, f) for f in train_fragments]\n        val_paths = [os.path.join(config.base_path, test_fragment)]\n        \n        train_dataset = SingleSliceDataset(train_paths, config.slice_range, augment=config.use_augmentation)\n        val_dataset = SingleSliceDataset(val_paths, config.slice_range, augment=False)\n        \n        print(f\"Training samples: {len(train_dataset)}\")\n        print(f\"Validation samples: {len(val_dataset)}\")\n        \n        # Create dataloaders\n        train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True)\n        val_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False)\n        \n        # Train model\n        model_name = f\"model_fragment_{test_fragment}\"\n        model, val_dice = train_model(train_loader, val_loader, model_name)\n        \n        # Store results\n        results[test_fragment] = val_dice\n        all_models.append(model)\n        \n        # Cleanup\n        del model\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    # Print results\n    print(f\"\\n{'='*60}\")\n    print(\"CROSS-VALIDATION RESULTS\")\n    print(f\"{'='*60}\")\n    \n    for fragment, dice in results.items():\n        print(f\"Fragment {fragment}: Dice = {dice:.4f}\")\n    \n    mean_dice = np.mean(list(results.values()))\n    std_dice = np.std(list(results.values()))\n    print(f\"\\nAverage Dice: {mean_dice:.4f} ± {std_dice:.4f}\")\n    \n    # Test ensemble\n    print(f\"\\n{'='*60}\")\n    print(\"ENSEMBLE TESTING\")\n    print(f\"{'='*60}\")\n    \n    ensemble_results = {}\n    \n    for test_fragment in all_fragments:\n        print(f\"\\nTesting ensemble on fragment {test_fragment}...\")\n        \n        # Load test mask\n        test_path = os.path.join(config.base_path, test_fragment)\n        test_mask = load_mask(test_path)\n        \n        if test_mask is None:\n            continue\n        \n        # Get predictions from all models\n        all_predictions = []\n        \n        for model in all_models:\n            model.eval()\n            pred = predict_fragment(model, test_path)\n            if pred is not None:\n                all_predictions.append(pred)\n        \n        if not all_predictions:\n            continue\n        \n        # Ensemble prediction (average)\n        ensemble_pred = np.mean(all_predictions, axis=0)\n        \n        # Apply post-processing\n        final_pred = post_process(ensemble_pred, min_area=config.min_area_size)\n        \n        # Calculate metrics\n        ensemble_metrics = calculate_metrics(final_pred, test_mask)\n        ensemble_results[test_fragment] = ensemble_metrics\n        \n        print(f\"Fragment {test_fragment}:\")\n        print(f\"  Dice: {ensemble_metrics['dice']:.4f}\")\n        print(f\"  F0.5: {ensemble_metrics.get('f05', 0):.4f}\")\n        print(f\"  Threshold: {ensemble_metrics.get('threshold', 0.5):.3f}\")\n        \n        # Save visualization\n        save_visualization(\n            test_path, \n            final_pred, \n            test_mask, \n            test_fragment,\n            ensemble_metrics\n        )\n    \n    # Print ensemble results\n    print(f\"\\n{'='*60}\")\n    print(\"FINAL ENSEMBLE RESULTS\")\n    print(f\"{'='*60}\")\n    \n    for fragment, metrics in ensemble_results.items():\n        print(f\"Fragment {fragment}: Dice={metrics['dice']:.4f}, F0.5={metrics.get('f05', 0):.4f}\")\n    \n    mean_ensemble_dice = np.mean([m['dice'] for m in ensemble_results.values()])\n    print(f\"\\nEnsemble Average Dice: {mean_ensemble_dice:.4f}\")\n    \n    if mean_ensemble_dice > 0.70:\n        print(f\"\\n🎉 SUCCESS: Achieved target Dice > 0.70! 🎉\")\n    elif mean_ensemble_dice > 0.50:\n        print(f\"\\n⚠️  Getting there! Try increasing epochs or adjusting parameters.\")\n    else:\n        print(f\"\\n❌ Need improvement. Consider changing the approach.\")\n    \n    return results, ensemble_results\n\n# ============================================================================\n# VISUALIZATION\n# ============================================================================\ndef save_visualization(fragment_path, prediction, target, fragment_name, metrics):\n    \"\"\"Save visualization of predictions\"\"\"\n    os.makedirs('visualizations', exist_ok=True)\n    \n    # Load a sample slice for display\n    sample_slice = load_single_slice(fragment_path, config.slice_range[0])\n    \n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    # Input slice\n    axes[0, 0].imshow(sample_slice, cmap='gray')\n    axes[0, 0].set_title('Input Slice')\n    axes[0, 0].axis('off')\n    \n    # Prediction probability (before thresholding)\n    pred_prob = prediction.copy()\n    im1 = axes[0, 1].imshow(pred_prob, cmap='jet', vmin=0, vmax=1)\n    axes[0, 1].set_title('Prediction (Probability)')\n    axes[0, 1].axis('off')\n    plt.colorbar(im1, ax=axes[0, 1], fraction=0.046, pad=0.04)\n    \n    # Prediction binary\n    pred_binary = (prediction > 0.5).astype(np.float32)\n    axes[0, 2].imshow(pred_binary, cmap='jet')\n    axes[0, 2].set_title('Prediction (Binary)')\n    axes[0, 2].axis('off')\n    \n    # Ground truth\n    axes[1, 0].imshow(target, cmap='jet')\n    axes[1, 0].set_title('Ground Truth')\n    axes[1, 0].axis('off')\n    \n    # Overlay\n    axes[1, 1].imshow(sample_slice, cmap='gray')\n    axes[1, 1].imshow(pred_binary, cmap='jet', alpha=0.5)\n    axes[1, 1].set_title('Overlay')\n    axes[1, 1].axis('off')\n    \n    # Difference\n    diff = np.abs(pred_binary - target)\n    axes[1, 2].imshow(diff, cmap='hot')\n    axes[1, 2].set_title('Difference')\n    axes[1, 2].axis('off')\n    \n    plt.suptitle(f\"Fragment {fragment_name} - Dice: {metrics['dice']:.4f} | \"\n                f\"F0.5: {metrics.get('f05', 0):.4f}\", fontsize=14)\n    \n    plt.tight_layout()\n    plt.savefig(f'visualizations/fragment_{fragment_name}.png', dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"Visualization saved for fragment {fragment_name}\")\n\n# ============================================================================\n# HYPERPARAMETER TUNING\n# ============================================================================\ndef tune_hyperparameters():\n    \"\"\"Quick hyperparameter tuning\"\"\"\n    print(\"Tuning hyperparameters...\")\n    \n    # Try different configurations\n    configs = [\n        {'lr': 1e-4, 'batch_size': 4, 'pos_weight': 5.0, 'slice_range': range(20, 25)},\n        {'lr': 3e-4, 'batch_size': 8, 'pos_weight': 10.0, 'slice_range': range(18, 27)},\n        {'lr': 5e-4, 'batch_size': 4, 'pos_weight': 15.0, 'slice_range': range(15, 30)},\n    ]\n    \n    best_score = 0\n    best_config = None\n    \n    for i, cfg in enumerate(configs):\n        print(f\"\\nTesting config {i+1}/{len(configs)}: {cfg}\")\n        \n        # Update config\n        config.lr = cfg['lr']\n        config.batch_size = cfg['batch_size']\n        config.pos_weight = cfg['pos_weight']\n        config.slice_range = list(cfg['slice_range'])\n        \n        # Quick test on fragment 1 only\n        test_fragment = '1'\n        train_fragments = ['2', '3']\n        \n        # Create datasets\n        train_paths = [os.path.join(config.base_path, f) for f in train_fragments]\n        val_paths = [os.path.join(config.base_path, test_fragment)]\n        \n        train_dataset = SingleSliceDataset(train_paths, config.slice_range, augment=True)\n        val_dataset = SingleSliceDataset(val_paths, config.slice_range, augment=False)\n        \n        train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True)\n        val_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False)\n        \n        # Train for fewer epochs\n        original_epochs = config.epochs\n        config.epochs = 20  # Quick test\n        \n        model, val_dice = train_model(train_loader, val_loader, f\"test_config_{i}\")\n        \n        print(f\"Config {i+1}: Val Dice = {val_dice:.4f}\")\n        \n        if val_dice > best_score:\n            best_score = val_dice\n            best_config = cfg.copy()\n        \n        # Reset epochs\n        config.epochs = original_epochs\n        \n        # Cleanup\n        del model\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n    print(f\"\\nBest config: {best_config} with Dice = {best_score:.4f}\")\n    \n    # Apply best config\n    if best_config:\n        config.lr = best_config['lr']\n        config.batch_size = best_config['batch_size']\n        config.pos_weight = best_config['pos_weight']\n        config.slice_range = list(best_config['slice_range'])\n    \n    return best_config\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\nif __name__ == \"__main__\":\n    # Set random seeds\n    torch.manual_seed(42)\n    np.random.seed(42)\n    random.seed(42)\n    \n    # Clear GPU memory\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    print(\"Starting training pipeline...\")\n    print(f\"Device: {config.device}\")\n    print(f\"Image size: {config.img_size}\")\n    print(f\"Slices: {config.slice_range}\")\n    \n    # Option 1: Tune hyperparameters first (recommended)\n    # best_config = tune_hyperparameters()\n    # print(f\"Using best config: {best_config}\")\n    \n    # Option 2: Run full cross-validation\n    results, ensemble_results = leave_one_out_cross_validation()\n    \n    print(\"\\nTraining complete!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nfrom scipy import ndimage\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# IMPROVED Configuration\nclass Config:\n    img_size = 384\n    slices = list(range(12, 31))  # Slices 12 to 30\n    batch_size = 8\n    learning_rate = 1e-3  # Higher learning rate for faster convergence\n    num_epochs = 50  # More epochs\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Loss parameters\n    focal_alpha = 0.8\n    focal_gamma = 2.0\n    dice_smooth = 1.0\n    loss_patience = 5  # Patience for loss improvement\n    \n    # Threshold search parameters\n    threshold_range = np.linspace(0.1, 0.9, 17)\n    \n    # Ensemble\n    num_models = 3\n\nconfig = Config()\n\n# Data paths\ntrain_paths = [\n    ('/kaggle/input/vesuvius-challenge-ink-detection/train/2', '/kaggle/input/vesuvius-challenge-ink-detection/train/2/inklabels.png'),\n    ('/kaggle/input/vesuvius-challenge-ink-detection/train/3', '/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png')\n]\ntest_path = ('/kaggle/input/vesuvius-challenge-ink-detection/train/1', '/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png')\n\n# IMPROVED Loss Functions with monitoring\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n        \n    def forward(self, inputs, targets):\n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n        \n        if self.reduction == 'mean':\n            return torch.mean(F_loss)\n        elif self.reduction == 'sum':\n            return torch.sum(F_loss)\n        else:\n            return F_loss\n\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n        \n    def forward(self, inputs, targets):\n        inputs = torch.sigmoid(inputs)\n        \n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()\n        dice = (2. * intersection + self.smooth) / (inputs.sum() + targets.sum() + self.smooth)\n        \n        return 1 - dice\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2.0, smooth=1.0, focal_weight=0.6, dice_weight=0.4):\n        super().__init__()\n        self.focal = FocalLoss(alpha=alpha, gamma=gamma)\n        self.dice = DiceLoss(smooth=smooth)\n        self.focal_weight = focal_weight\n        self.dice_weight = dice_weight\n        \n    def forward(self, inputs, targets):\n        focal_loss = self.focal(inputs, targets)\n        dice_loss = self.dice(inputs, targets)\n        \n        return self.focal_weight * focal_loss + self.dice_weight * dice_loss\n\n# IMPROVED U-Net with Dropout for better generalization\nclass ImprovedUNet(nn.Module):\n    def __init__(self, in_channels=19, base_channels=32):\n        super().__init__()\n        \n        # Encoder with dropout\n        self.enc1 = self._block(in_channels, base_channels, dropout=0.1)\n        self.enc2 = self._block(base_channels, base_channels * 2, dropout=0.2)\n        self.enc3 = self._block(base_channels * 2, base_channels * 4, dropout=0.3)\n        self.enc4 = self._block(base_channels * 4, base_channels * 8, dropout=0.4)\n        \n        # Bottleneck\n        self.bottleneck = self._block(base_channels * 8, base_channels * 16, dropout=0.5)\n        \n        # Decoder\n        self.up4 = nn.ConvTranspose2d(base_channels * 16, base_channels * 8, 2, stride=2)\n        self.dec4 = self._block(base_channels * 16, base_channels * 8, dropout=0.4)\n        \n        self.up3 = nn.ConvTranspose2d(base_channels * 8, base_channels * 4, 2, stride=2)\n        self.dec3 = self._block(base_channels * 8, base_channels * 4, dropout=0.3)\n        \n        self.up2 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, 2, stride=2)\n        self.dec2 = self._block(base_channels * 4, base_channels * 2, dropout=0.2)\n        \n        self.up1 = nn.ConvTranspose2d(base_channels * 2, base_channels, 2, stride=2)\n        self.dec1 = self._block(base_channels * 2, base_channels, dropout=0.1)\n        \n        # Final convolution\n        self.final = nn.Conv2d(base_channels, 1, 1)\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        \n    def _block(self, in_channels, out_channels, dropout=0.0):\n        layers = [\n            nn.Conv2d(in_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        ]\n        \n        if dropout > 0:\n            layers.append(nn.Dropout2d(dropout))\n            \n        return nn.Sequential(*layers)\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.enc1(x)\n        enc2 = self.enc2(self.pool(enc1))\n        enc3 = self.enc3(self.pool(enc2))\n        enc4 = self.enc4(self.pool(enc3))\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc4))\n        \n        # Decoder with skip connections\n        dec4 = self.up4(bottleneck)\n        dec4 = torch.cat([dec4, enc4], dim=1)\n        dec4 = self.dec4(dec4)\n        \n        dec3 = self.up3(dec4)\n        dec3 = torch.cat([dec3, enc3], dim=1)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.up2(dec3)\n        dec2 = torch.cat([dec2, enc2], dim=1)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.up1(dec2)\n        dec1 = torch.cat([dec1, enc1], dim=1)\n        dec1 = self.dec1(dec1)\n        \n        output = self.final(dec1)\n        return output\n\n# Dataset Class with MORE AUGMENTATION\nclass VesuviusDataset(Dataset):\n    def __init__(self, paths, is_training=True):\n        self.paths = paths\n        self.is_training = is_training\n        self.config = Config()\n        \n        self.samples = []\n        for volume_path, label_path in paths:\n            try:\n                volume = self.load_volume(volume_path)\n                label = self.load_label(label_path)\n                self.samples.append((volume, label))\n            except Exception as e:\n                print(f\"Error loading {volume_path}: {e}\")\n                continue\n        \n        if len(self.samples) == 0:\n            raise ValueError(\"No samples were successfully loaded!\")\n    \n    def load_volume(self, volume_path):\n        volume = []\n        for slice_idx in self.config.slices:\n            slice_path = os.path.join(volume_path, 'surface_volume', f\"{slice_idx}.tif\")\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Could not load image: {slice_path}\")\n            img = cv2.resize(img, (self.config.img_size, self.config.img_size))\n            volume.append(img)\n        return np.stack(volume, axis=0)\n    \n    def load_label(self, label_path):\n        label = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE)\n        if label is None:\n            raise ValueError(f\"Could not load label: {label_path}\")\n        label = cv2.resize(label, (self.config.img_size, self.config.img_size))\n        return (label > 0).astype(np.float32)\n    \n    def __len__(self):\n        return len(self.samples) * 32  # More augmentations\n    \n    def __getitem__(self, idx):\n        volume_idx = idx % len(self.samples)\n        volume, label = self.samples[volume_idx]\n        \n        if self.is_training:\n            h, w = volume.shape[1], volume.shape[2]\n            patch_size = 256\n            \n            # Random crop\n            top = np.random.randint(0, h - patch_size)\n            left = np.random.randint(0, w - patch_size)\n            \n            volume = volume[:, top:top+patch_size, left:left+patch_size]\n            label = label[top:top+patch_size, left:left+patch_size]\n            \n            # MORE AUGMENTATIONS\n            # Random flip\n            if np.random.random() > 0.5:\n                volume = np.flip(volume, axis=2).copy()\n                label = np.flip(label, axis=1).copy()\n            \n            if np.random.random() > 0.5:\n                volume = np.flip(volume, axis=1).copy()\n                label = np.flip(label, axis=0).copy()\n            \n            # Random rotation\n            if np.random.random() > 0.5:\n                k = np.random.randint(1, 4)\n                volume = np.rot90(volume, k, axes=(1, 2)).copy()\n                label = np.rot90(label, k).copy()\n            \n            # Random brightness adjustment\n            if np.random.random() > 0.5:\n                factor = 0.8 + np.random.random() * 0.4  # 0.8-1.2\n                volume = np.clip(volume * factor, 0, 255)\n            \n            # Random Gaussian noise\n            if np.random.random() > 0.5:\n                noise = np.random.randn(*volume.shape) * 10\n                volume = np.clip(volume + noise, 0, 255)\n        \n        # Normalize\n        volume = volume.astype(np.float32) / 255.0\n        \n        # Convert to tensors\n        volume_tensor = torch.FloatTensor(volume)\n        label_tensor = torch.FloatTensor(label).unsqueeze(0)\n        \n        return volume_tensor, label_tensor\n\n# Functions for threshold optimization\ndef compute_metrics_for_threshold(pred_probs, targets, threshold):\n    pred_binary = (pred_probs > threshold).float()\n    \n    pred_np = pred_binary.cpu().numpy().flatten()\n    target_np = targets.cpu().numpy().flatten()\n    \n    if np.sum(target_np) == 0 and np.sum(pred_np) == 0:\n        return {'dice': 1.0, 'f05': 1.0, 'precision': 1.0, 'recall': 1.0}\n    elif np.sum(target_np) == 0:\n        return {'dice': 0.0, 'f05': 0.0, 'precision': 0.0, 'recall': 0.0}\n    \n    # Dice score\n    intersection = (pred_binary * targets).sum().item()\n    union = pred_binary.sum().item() + targets.sum().item()\n    dice = (2. * intersection + 1e-6) / (union + 1e-6)\n    \n    # F0.5 score\n    try:\n        f05 = fbeta_score(target_np, pred_np, beta=0.5, zero_division=0)\n    except:\n        f05 = 0.0\n    \n    # Precision and recall\n    try:\n        precision = precision_score(target_np, pred_np, zero_division=0)\n        recall = recall_score(target_np, pred_np, zero_division=0)\n    except:\n        precision = 0.0\n        recall = 0.0\n    \n    return {'dice': dice, 'f05': f05, 'precision': precision, 'recall': recall}\n\ndef find_optimal_thresholds(pred_probs, targets, threshold_range):\n    thresholds = {}\n    scores = {}\n    \n    for metric in ['dice', 'f05']:\n        best_thresh = 0.5\n        best_score = -1.0\n        \n        for threshold in threshold_range:\n            metrics = compute_metrics_for_threshold(pred_probs, targets, threshold)\n            score = metrics[metric]\n            \n            if score > best_score:\n                best_score = score\n                best_thresh = threshold\n        \n        thresholds[metric] = best_thresh\n        scores[metric] = best_score\n    \n    return thresholds, scores\n\n# IMPROVED Training with BETTER OPTIMIZATION\ndef train_supervised():\n    print(\"Starting IMPROVED Supervised Training with Threshold Optimization...\")\n    print(f\"Target: Train Loss < 0.75\")\n    print(f\"Threshold search range: {config.threshold_range}\")\n    \n    try:\n        # Create datasets\n        train_dataset = VesuviusDataset([train_paths[0]], is_training=True)\n        val_dataset = VesuviusDataset([train_paths[1]], is_training=False)\n        \n        train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=2, pin_memory=True)\n        val_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n        \n        print(f\"Train samples: {len(train_dataset)}, Val samples: {len(val_dataset)}\")\n        \n        # Create ensemble of models\n        models = []\n        optimizers = []\n        schedulers = []\n        \n        for i in range(config.num_models):\n            model = ImprovedUNet(in_channels=len(config.slices)).to(config.device)\n            \n            # Better weight initialization\n            def init_weights(m):\n                if isinstance(m, nn.Conv2d):\n                    nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                    if m.bias is not None:\n                        nn.init.constant_(m.bias, 0)\n                elif isinstance(m, nn.BatchNorm2d):\n                    nn.init.constant_(m.weight, 1)\n                    nn.init.constant_(m.bias, 0)\n            \n            model.apply(init_weights)\n            \n            # AdamW with higher weight decay\n            optimizer = optim.AdamW(model.parameters(), lr=config.learning_rate, \n                                   weight_decay=1e-4, betas=(0.9, 0.999))\n            \n            # Cosine annealing scheduler\n            scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=config.num_epochs)\n            \n            models.append(model)\n            optimizers.append(optimizer)\n            schedulers.append(scheduler)\n        \n        # Use Combined Loss\n        criterion = CombinedLoss(\n            alpha=config.focal_alpha,\n            gamma=config.focal_gamma,\n            smooth=config.dice_smooth\n        )\n        \n        # Training history with loss monitoring\n        history = {\n            'train_loss': [[] for _ in range(config.num_models)],\n            'val_dice': [[] for _ in range(config.num_models)],\n            'val_f05': [[] for _ in range(config.num_models)],\n            'val_precision': [[] for _ in range(config.num_models)],\n            'val_recall': [[] for _ in range(config.num_models)],\n            'val_threshold_dice': [[] for _ in range(config.num_models)],\n            'val_threshold_f05': [[] for _ in range(config.num_models)],\n            'learning_rate': [[] for _ in range(config.num_models)]\n        }\n        \n        best_dice = 0.0\n        best_models = [None] * config.num_models\n        best_val_thresholds = [0.5] * config.num_models\n        \n        # Early stopping variables\n        no_improvement_count = 0\n        best_avg_loss = float('inf')\n        \n        for epoch in range(config.num_epochs):\n            print(f\"\\n{'='*60}\")\n            print(f\"Epoch {epoch+1}/{config.num_epochs}\")\n            print('='*60)\n            \n            # Training phase with LOSS MINIMIZATION as priority\n            for model_idx, (model, optimizer, scheduler) in enumerate(zip(models, optimizers, schedulers)):\n                model.train()\n                epoch_loss = 0\n                batch_count = 0\n                \n                pbar = tqdm(train_loader, desc=f\"Model {model_idx+1} Training\", leave=False)\n                for volume, labels in pbar:\n                    volume = volume.to(config.device)\n                    labels = labels.to(config.device)\n                    \n                    # Forward pass\n                    outputs = model(volume)\n                    loss = criterion(outputs, labels)\n                    \n                    # Backward pass with gradient accumulation for stability\n                    optimizer.zero_grad()\n                    loss.backward()\n                    \n                    # Gradient clipping for stability\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.5)\n                    \n                    optimizer.step()\n                    \n                    epoch_loss += loss.item()\n                    batch_count += 1\n                    \n                    # Update progress bar with loss target\n                    pbar.set_postfix({'Loss': f'{loss.item():.4f}', 'Target': '<0.75'})\n                \n                if batch_count > 0:\n                    avg_loss = epoch_loss / batch_count\n                    history['train_loss'][model_idx].append(avg_loss)\n                    \n                    # Store learning rate\n                    history['learning_rate'][model_idx].append(optimizer.param_groups[0]['lr'])\n                \n                # Step scheduler\n                scheduler.step()\n            \n            # Calculate average training loss\n            current_losses = [history['train_loss'][i][-1] for i in range(config.num_models)]\n            avg_train_loss = np.mean(current_losses)\n            \n            print(f\"\\nAverage Training Loss: {avg_train_loss:.4f} (Target: <0.75)\")\n            \n            # Validation phase\n            print(\"\\nValidation with Threshold Optimization:\")\n            for model_idx, model in enumerate(models):\n                model.eval()\n                \n                # Collect predictions for threshold search\n                all_preds = []\n                all_targets = []\n                \n                with torch.no_grad():\n                    for volume, labels in val_loader:\n                        volume = volume.to(config.device)\n                        labels = labels.to(config.device)\n                        \n                        outputs = model(volume)\n                        preds = torch.sigmoid(outputs)\n                        \n                        all_preds.append(preds.cpu())\n                        all_targets.append(labels.cpu())\n                \n                all_preds = torch.cat(all_preds, dim=0)\n                all_targets = torch.cat(all_targets, dim=0)\n                \n                # Find optimal thresholds\n                thresholds, scores = find_optimal_thresholds(\n                    all_preds, all_targets, config.threshold_range\n                )\n                \n                # Use F0.5 optimal threshold\n                optimal_threshold = thresholds['f05']\n                optimal_score = scores['f05']\n                \n                # Store thresholds\n                history['val_threshold_dice'][model_idx].append(thresholds['dice'])\n                history['val_threshold_f05'][model_idx].append(thresholds['f05'])\n                \n                # Compute final metrics\n                val_metrics = compute_metrics_for_threshold(\n                    all_preds, all_targets, optimal_threshold\n                )\n                \n                # Store metrics\n                for metric, value in val_metrics.items():\n                    history[f'val_{metric}'][model_idx].append(value)\n                \n                # Check if training loss is on target\n                train_loss = history['train_loss'][model_idx][-1]\n                loss_status = \"✓\" if train_loss < 0.75 else \"✗\"\n                \n                print(f\"\\nModel {model_idx+1} {loss_status}:\")\n                print(f\"  Train Loss: {train_loss:.4f} {'(TARGET ACHIEVED!)' if train_loss < 0.75 else ''}\")\n                print(f\"  Val Dice: {val_metrics['dice']:.4f} (t={thresholds['dice']:.3f})\")\n                print(f\"  Val F0.5: {val_metrics['f05']:.4f} (t={thresholds['f05']:.3f})\")\n                print(f\"  LR: {history['learning_rate'][model_idx][-1]:.6f}\")\n                \n                # Save best model based on F0.5\n                if val_metrics['f05'] > best_dice:\n                    best_dice = val_metrics['f05']\n                    best_models[model_idx] = model.state_dict().copy()\n                    best_val_thresholds[model_idx] = optimal_threshold\n                    torch.save({\n                        'model_state_dict': model.state_dict(),\n                        'threshold': optimal_threshold,\n                        'f05_score': val_metrics['f05'],\n                        'dice_score': val_metrics['dice'],\n                        'train_loss': train_loss,\n                        'epoch': epoch\n                    }, f'best_model_{model_idx}.pth')\n            \n            # Early stopping based on loss improvement\n            if avg_train_loss < best_avg_loss:\n                best_avg_loss = avg_train_loss\n                no_improvement_count = 0\n                print(f\"\\n✓ Loss improved to {avg_train_loss:.4f}\")\n            else:\n                no_improvement_count += 1\n                print(f\"\\n✗ Loss didn't improve for {no_improvement_count} epochs\")\n            \n            # Check stopping conditions\n            if avg_train_loss < 0.75:\n                print(f\"\\n🎯 TARGET ACHIEVED! Average loss {avg_train_loss:.4f} < 0.75\")\n                print(\"Continuing for better metrics...\")\n            \n            if no_improvement_count >= config.loss_patience and epoch > 10:\n                print(f\"\\n🛑 Early stopping: No improvement for {config.loss_patience} epochs\")\n                break\n            \n            if epoch >= 15 and avg_train_loss < 0.5:\n                print(f\"\\n✅ Good convergence achieved, stopping early\")\n                break\n        \n        # Load best models\n        for i in range(config.num_models):\n            if best_models[i] is not None:\n                models[i].load_state_dict(best_models[i])\n        \n        return models, history, best_val_thresholds\n    \n    except Exception as e:\n        print(f\"Supervised training failed: {e}\")\n        import traceback\n        traceback.print_exc()\n        return None, None, None\n\ndef test_ensemble_with_threshold_optimization(models, test_path, val_thresholds=None):\n    \"\"\"COMPLETE testing with threshold optimization\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"TESTING ENSEMBLE WITH THRESHOLD OPTIMIZATION\")\n    print(\"=\"*60)\n    \n    if models is None:\n        print(\"No models to test!\")\n        return {'metrics': {'dice': 0.0, 'f05': 0.0, 'precision': 0.0, 'recall': 0.0},\n                'thresholds': None}\n    \n    try:\n        test_dataset = VesuviusDataset([test_path], is_training=False)\n        test_loader = DataLoader(test_dataset, batch_size=config.batch_size, shuffle=False, \n                               num_workers=2, pin_memory=True)\n        \n        # Store predictions from all models\n        all_ensemble_preds = []\n        all_test_targets = []\n        \n        with torch.no_grad():\n            for batch_idx, (volume, labels) in enumerate(test_loader):\n                volume = volume.to(config.device)\n                labels = labels.to(config.device)\n                \n                # Ensemble predictions\n                ensemble_preds = []\n                for model in models:\n                    model.eval()\n                    output = model(volume)\n                    pred = torch.sigmoid(output)\n                    ensemble_preds.append(pred)\n                \n                # Average predictions\n                ensemble_pred = torch.mean(torch.stack(ensemble_preds), dim=0)\n                \n                all_ensemble_preds.append(ensemble_pred.cpu())\n                all_test_targets.append(labels.cpu())\n        \n        # Concatenate all batches\n        all_ensemble_preds = torch.cat(all_ensemble_preds, dim=0)\n        all_test_targets = torch.cat(all_test_targets, dim=0)\n        \n        print(\"\\n🔍 Searching for optimal TEST thresholds (0.1 to 0.9)...\")\n        \n        # Find optimal thresholds for TEST set\n        test_thresholds, test_scores = find_optimal_thresholds(\n            all_ensemble_preds, all_test_targets, config.threshold_range\n        )\n        \n        print(f\"\\n📊 TEST SET OPTIMAL THRESHOLDS:\")\n        print(f\"  For Dice: {test_thresholds['dice']:.3f} (Dice: {test_scores['dice']:.4f})\")\n        print(f\"  For F0.5: {test_thresholds['f05']:.3f} (F0.5: {test_scores['f05']:.4f})\")\n        \n        # Compare with validation thresholds\n        if val_thresholds is not None:\n            avg_val_threshold = np.mean(val_thresholds)\n            print(f\"\\n📈 VALIDATION vs TEST COMPARISON:\")\n            print(f\"  Average Validation Threshold: {avg_val_threshold:.3f}\")\n            print(f\"  Test F0.5 Threshold: {test_thresholds['f05']:.3f}\")\n            print(f\"  Difference: {abs(test_thresholds['f05'] - avg_val_threshold):.3f}\")\n            \n            # Test with validation thresholds too\n            print(f\"\\n🔬 Testing with validation thresholds:\")\n            for i, val_thresh in enumerate(val_thresholds):\n                val_metrics = compute_metrics_for_threshold(\n                    all_ensemble_preds, all_test_targets, val_thresh\n                )\n                print(f\"  Model {i+1} Val Threshold {val_thresh:.3f}: \"\n                      f\"Dice={val_metrics['dice']:.4f}, F0.5={val_metrics['f05']:.4f}\")\n        \n        # Use F0.5 optimal threshold for test set\n        optimal_test_threshold = test_thresholds['f05']\n        \n        # Compute final metrics\n        final_metrics = compute_metrics_for_threshold(\n            all_ensemble_preds, all_test_targets, optimal_test_threshold\n        )\n        \n        print(f\"\\n🎯 FINAL TEST RESULTS (Threshold: {optimal_test_threshold:.3f}):\")\n        print(f\"  Dice Score:  {final_metrics['dice']:.4f}\")\n        print(f\"  F0.5 Score:  {final_metrics['f05']:.4f}\")\n        print(f\"  Precision:   {final_metrics['precision']:.4f}\")\n        print(f\"  Recall:      {final_metrics['recall']:.4f}\")\n        \n        # Save visualization\n        with torch.no_grad():\n            volume, labels = next(iter(test_loader))\n            volume = volume.to(config.device)\n            labels = labels.to(config.device)\n            \n            ensemble_preds = []\n            for model in models:\n                model.eval()\n                output = model(volume)\n                pred = torch.sigmoid(output)\n                ensemble_preds.append(pred)\n            \n            ensemble_pred = torch.mean(torch.stack(ensemble_preds), dim=0)\n            \n            save_visualization(volume, ensemble_pred, labels, 0, 0, 'ensemble', \n                             'testing', optimal_test_threshold)\n        \n        results = {\n            'metrics': final_metrics,\n            'thresholds': {\n                'test_optimal_f05': test_thresholds['f05'],\n                'test_optimal_dice': test_thresholds['dice'],\n                'test_scores': test_scores,\n                'val_thresholds': val_thresholds if val_thresholds else None\n            }\n        }\n        \n        return results\n    \n    except Exception as e:\n        print(f\"Testing failed: {e}\")\n        return {'metrics': {'dice': 0.0, 'f05': 0.0, 'precision': 0.0, 'recall': 0.0},\n                'thresholds': None}\n\ndef save_visualization(volume, pred, label, batch_idx, epoch, model_idx, phase, threshold):\n    os.makedirs(f'visualizations/{phase}', exist_ok=True)\n    \n    try:\n        vol_sample = volume[0].cpu().numpy()\n        pred_sample = pred[0, 0].cpu().numpy()\n        label_sample = label[0, 0].cpu().numpy()\n        \n        pred_binary = (pred_sample > threshold).astype(np.float32)\n        vol_viz = np.mean(vol_sample, axis=0)\n        \n        fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n        \n        axes[0].imshow(vol_viz, cmap='gray')\n        axes[0].set_title('Input (Mean Slice)')\n        axes[0].axis('off')\n        \n        im1 = axes[1].imshow(pred_sample, cmap='jet', vmin=0, vmax=1)\n        axes[1].set_title('Prediction (Probability)')\n        axes[1].axis('off')\n        plt.colorbar(im1, ax=axes[1], fraction=0.046, pad=0.04)\n        \n        axes[2].imshow(pred_binary, cmap='jet', vmin=0, vmax=1)\n        axes[2].set_title(f'Thresholded (t={threshold:.2f})')\n        axes[2].axis('off')\n        \n        axes[3].imshow(label_sample, cmap='jet', vmin=0, vmax=1)\n        axes[3].set_title('Ground Truth')\n        axes[3].axis('off')\n        \n        plt.suptitle(f'{phase.capitalize()} - Epoch {epoch}, Model {model_idx}, Threshold={threshold:.3f}', \n                    fontsize=16, y=1.02)\n        plt.tight_layout()\n        plt.savefig(f'visualizations/{phase}/epoch_{epoch}_model_{model_idx}_batch_{batch_idx}.png', \n                    dpi=150, bbox_inches='tight')\n        plt.close()\n        print(f\"Saved visualization: visualizations/{phase}/epoch_{epoch}_model_{model_idx}_batch_{batch_idx}.png\")\n    except Exception as e:\n        print(f\"Failed to save visualization: {e}\")\n\n# MAIN EXECUTION - COMPLETE TRAINING AND TESTING\nif __name__ == \"__main__\":\n    print(f\"Using device: {config.device}\")\n    print(f\"Target: Training Loss < 0.75\")\n    \n    os.makedirs('visualizations/validation', exist_ok=True)\n    os.makedirs('visualizations/testing', exist_ok=True)\n    \n    try:\n        # 1. TRAIN MODELS\n        print(\"\\n\" + \"=\"*60)\n        print(\"PHASE 1: TRAINING\")\n        print(\"=\"*60)\n        models, history, best_val_thresholds = train_supervised()\n        \n        # 2. TEST ENSEMBLE\n        print(\"\\n\" + \"=\"*60)\n        print(\"PHASE 2: TESTING\")\n        print(\"=\"*60)\n        test_results = test_ensemble_with_threshold_optimization(models, test_path, best_val_thresholds)\n        \n        # 3. FINAL RESULTS\n        print(\"\\n\" + \"=\"*60)\n        print(\"FINAL RESULTS\")\n        print(\"=\"*60)\n        \n        if test_results['thresholds']:\n            print(f\"\\n📊 OPTIMAL THRESHOLDS:\")\n            print(f\"  Test Set (F0.5): {test_results['thresholds']['test_optimal_f05']:.3f}\")\n            print(f\"  Test Set (Dice): {test_results['thresholds']['test_optimal_dice']:.3f}\")\n            if test_results['thresholds']['val_thresholds']:\n                avg_val = np.mean(test_results['thresholds']['val_thresholds'])\n                print(f\"  Validation Avg: {avg_val:.3f}\")\n                print(f\"  Test-Val Diff: {abs(test_results['thresholds']['test_optimal_f05'] - avg_val):.3f}\")\n        \n        print(f\"\\n🎯 TEST METRICS:\")\n        print(f\"  Dice Score:  {test_results['metrics']['dice']:.4f}\")\n        print(f\"  F0.5 Score:  {test_results['metrics']['f05']:.4f}\")\n        print(f\"  Precision:   {test_results['metrics']['precision']:.4f}\")\n        print(f\"  Recall:      {test_results['metrics']['recall']:.4f}\")\n        print(\"=\"*60)\n        \n        # 4. ANALYSIS\n        if history is not None:\n            np.save('training_history.npy', history)\n            \n            # Plot comprehensive results\n            fig = plt.figure(figsize=(20, 12))\n            \n            # Plot 1: Training Loss with target\n            ax1 = plt.subplot(2, 3, 1)\n            colors = ['b', 'g', 'r', 'c', 'm']\n            for i in range(config.num_models):\n                ax1.plot(history['train_loss'][i], label=f'Model {i+1}', color=colors[i], linewidth=2)\n            ax1.axhline(y=0.75, color='r', linestyle='--', linewidth=2, label='Target (0.75)')\n            ax1.set_title('Training Loss (Target: <0.75)', fontsize=14, fontweight='bold')\n            ax1.set_xlabel('Epoch')\n            ax1.set_ylabel('Loss')\n            ax1.legend()\n            ax1.grid(True, alpha=0.3)\n            \n            # Plot 2: Validation Dice\n            ax2 = plt.subplot(2, 3, 2)\n            for i in range(config.num_models):\n                ax2.plot(history['val_dice'][i], label=f'Model {i+1}', color=colors[i], linewidth=2)\n            ax2.set_title('Validation Dice Score', fontsize=14, fontweight='bold')\n            ax2.set_xlabel('Epoch')\n            ax2.set_ylabel('Dice')\n            ax2.legend()\n            ax2.grid(True, alpha=0.3)\n            \n            # Plot 3: Validation F0.5\n            ax3 = plt.subplot(2, 3, 3)\n            for i in range(config.num_models):\n                ax3.plot(history['val_f05'][i], label=f'Model {i+1}', color=colors[i], linewidth=2)\n            ax3.set_title('Validation F0.5 Score', fontsize=14, fontweight='bold')\n            ax3.set_xlabel('Epoch')\n            ax3.set_ylabel('F0.5')\n            ax3.legend()\n            ax3.grid(True, alpha=0.3)\n            \n            # Plot 4: Learning Rate\n            ax4 = plt.subplot(2, 3, 4)\n            for i in range(config.num_models):\n                ax4.plot(history['learning_rate'][i], label=f'Model {i+1}', color=colors[i], linewidth=2)\n            ax4.set_title('Learning Rate Schedule', fontsize=14, fontweight='bold')\n            ax4.set_xlabel('Epoch')\n            ax4.set_ylabel('Learning Rate')\n            ax4.set_yscale('log')\n            ax4.legend()\n            ax4.grid(True, alpha=0.3)\n            \n            # Plot 5: Optimal Thresholds (F0.5)\n            ax5 = plt.subplot(2, 3, 5)\n            for i in range(config.num_models):\n                ax5.plot(history['val_threshold_f05'][i], label=f'Model {i+1}', color=colors[i], linewidth=2)\n            ax5.set_title('Optimal Thresholds for F0.5', fontsize=14, fontweight='bold')\n            ax5.set_xlabel('Epoch')\n            ax5.set_ylabel('Threshold')\n            ax5.legend()\n            ax5.grid(True, alpha=0.3)\n            \n            # Plot 6: Test Results Summary\n            ax6 = plt.subplot(2, 3, 6)\n            if test_results['metrics']:\n                metrics_names = ['Dice', 'F0.5', 'Precision', 'Recall']\n                metrics_values = [\n                    test_results['metrics']['dice'],\n                    test_results['metrics']['f05'],\n                    test_results['metrics']['precision'],\n                    test_results['metrics']['recall']\n                ]\n                bars = ax6.bar(metrics_names, metrics_values, color=['blue', 'green', 'orange', 'red'])\n                ax6.set_title('Final Test Metrics', fontsize=14, fontweight='bold')\n                ax6.set_ylabel('Score')\n                ax6.set_ylim(0, 1)\n                \n                # Add value labels on bars\n                for bar, val in zip(bars, metrics_values):\n                    height = bar.get_height()\n                    ax6.text(bar.get_x() + bar.get_width()/2., height + 0.02,\n                            f'{val:.3f}', ha='center', va='bottom', fontweight='bold')\n            \n            plt.suptitle('Complete Training and Testing Results', fontsize=16, fontweight='bold', y=1.02)\n            plt.tight_layout()\n            plt.savefig('complete_results.png', dpi=150, bbox_inches='tight')\n            plt.close()\n            \n            # Print training summary\n            print(f\"\\n📈 TRAINING SUMMARY:\")\n            for i in range(config.num_models):\n                if history['train_loss'][i]:\n                    final_loss = history['train_loss'][i][-1]\n                    best_loss = min(history['train_loss'][i])\n                    best_dice = max(history['val_dice'][i]) if history['val_dice'][i] else 0\n                    best_f05 = max(history['val_f05'][i]) if history['val_f05'][i] else 0\n                    \n                    target_status = \"✓ ACHIEVED\" if final_loss < 0.75 else \"✗ NOT ACHIEVED\"\n                    \n                    print(f\"\\n  Model {i+1}:\")\n                    print(f\"    Final Loss: {final_loss:.4f} {target_status}\")\n                    print(f\"    Best Loss:  {best_loss:.4f}\")\n                    print(f\"    Best Dice:  {best_dice:.4f}\")\n                    print(f\"    Best F0.5:  {best_f05:.4f}\")\n        \n        print(\"\\n\" + \"=\"*60)\n        print(\"✅ PROCESS COMPLETE!\")\n        print(\"=\"*60)\n        print(\"Check 'visualizations/' for prediction images\")\n        print(\"Check 'complete_results.png' for training history\")\n        print(\"Check 'training_history.npy' for all data\")\n        print(\"=\"*60)\n        \n    except Exception as e:\n        print(f\"Error occurred: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport numpy as np\nimport cv2\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nfrom scipy import ndimage\nfrom torchvision import transforms\nimport torchvision.transforms.functional as TF\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Configuration\nclass Config:\n    img_size = 384\n    slices = list(range(12, 31))  # Slices 12 to 30\n    batch_size = 8\n    learning_rate = 3e-4\n    num_epochs = 30\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Loss parameters\n    focal_alpha = 0.8\n    focal_gamma = 2.0\n    dice_smooth = 1.0\n    \n    # Threshold search parameters\n    threshold_range = np.linspace(0.1, 0.9, 17)  # 0.1, 0.15, 0.2, ..., 0.9\n    default_threshold = 0.5\n    \n    # Ensemble\n    num_models = 3\n\nconfig = Config()\n\n# Data paths\ntrain_paths = [\n    ('/kaggle/input/vesuvius-challenge-ink-detection/train/2', '/kaggle/input/vesuvius-challenge-ink-detection/train/2/inklabels.png'),\n    ('/kaggle/input/vesuvius-challenge-ink-detection/train/3', '/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png')\n]\ntest_path = ('/kaggle/input/vesuvius-challenge-ink-detection/train/1', '/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png')\n\n# Advanced Loss Functions\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n        \n    def forward(self, inputs, targets):\n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n        \n        if self.reduction == 'mean':\n            return torch.mean(F_loss)\n        elif self.reduction == 'sum':\n            return torch.sum(F_loss)\n        else:\n            return F_loss\n\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n        \n    def forward(self, inputs, targets):\n        inputs = torch.sigmoid(inputs)\n        \n        # Flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()\n        dice = (2. * intersection + self.smooth) / (inputs.sum() + targets.sum() + self.smooth)\n        \n        return 1 - dice\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2.0, smooth=1.0, focal_weight=0.6, dice_weight=0.4):\n        super().__init__()\n        self.focal = FocalLoss(alpha=alpha, gamma=gamma)\n        self.dice = DiceLoss(smooth=smooth)\n        self.focal_weight = focal_weight\n        self.dice_weight = dice_weight\n        \n    def forward(self, inputs, targets):\n        focal_loss = self.focal(inputs, targets)\n        dice_loss = self.dice(inputs, targets)\n        \n        return self.focal_weight * focal_loss + self.dice_weight * dice_loss\n\n# SIMPLE BUT EFFECTIVE U-Net\nclass SimpleUNet(nn.Module):\n    def __init__(self, in_channels=19, base_channels=32):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._block(in_channels, base_channels)\n        self.enc2 = self._block(base_channels, base_channels * 2)\n        self.enc3 = self._block(base_channels * 2, base_channels * 4)\n        self.enc4 = self._block(base_channels * 4, base_channels * 8)\n        \n        # Bottleneck\n        self.bottleneck = self._block(base_channels * 8, base_channels * 16)\n        \n        # Decoder\n        self.up4 = nn.ConvTranspose2d(base_channels * 16, base_channels * 8, 2, stride=2)\n        self.dec4 = self._block(base_channels * 16, base_channels * 8)\n        \n        self.up3 = nn.ConvTranspose2d(base_channels * 8, base_channels * 4, 2, stride=2)\n        self.dec3 = self._block(base_channels * 8, base_channels * 4)\n        \n        self.up2 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, 2, stride=2)\n        self.dec2 = self._block(base_channels * 4, base_channels * 2)\n        \n        self.up1 = nn.ConvTranspose2d(base_channels * 2, base_channels, 2, stride=2)\n        self.dec1 = self._block(base_channels * 2, base_channels)\n        \n        # Final convolution\n        self.final = nn.Conv2d(base_channels, 1, 1)\n        \n        # Pooling\n        self.pool = nn.MaxPool2d(2)\n        \n    def _block(self, in_channels, out_channels):\n        return nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.enc1(x)                    # 32\n        enc2 = self.enc2(self.pool(enc1))      # 64\n        enc3 = self.enc3(self.pool(enc2))      # 128\n        enc4 = self.enc4(self.pool(enc3))      # 256\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc4))  # 512\n        \n        # Decoder with skip connections\n        dec4 = self.up4(bottleneck)\n        dec4 = torch.cat([dec4, enc4], dim=1)\n        dec4 = self.dec4(dec4)\n        \n        dec3 = self.up3(dec4)\n        dec3 = torch.cat([dec3, enc3], dim=1)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.up2(dec3)\n        dec2 = torch.cat([dec2, enc2], dim=1)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.up1(dec2)\n        dec1 = torch.cat([dec1, enc1], dim=1)\n        dec1 = self.dec1(dec1)\n        \n        output = self.final(dec1)\n        return output\n\n# Dataset Class\nclass VesuviusDataset(Dataset):\n    def __init__(self, paths, is_training=True):\n        self.paths = paths\n        self.is_training = is_training\n        self.config = Config()\n        \n        self.samples = []\n        for volume_path, label_path in paths:\n            try:\n                volume = self.load_volume(volume_path)\n                label = self.load_label(label_path)\n                self.samples.append((volume, label))\n            except Exception as e:\n                print(f\"Error loading {volume_path}: {e}\")\n                continue\n        \n        if len(self.samples) == 0:\n            raise ValueError(\"No samples were successfully loaded!\")\n    \n    def load_volume(self, volume_path):\n        volume = []\n        for slice_idx in self.config.slices:\n            slice_path = os.path.join(volume_path, 'surface_volume', f\"{slice_idx}.tif\")\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Could not load image: {slice_path}\")\n            img = cv2.resize(img, (self.config.img_size, self.config.img_size))\n            volume.append(img)\n        return np.stack(volume, axis=0)\n    \n    def load_label(self, label_path):\n        label = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE)\n        if label is None:\n            raise ValueError(f\"Could not load label: {label_path}\")\n        label = cv2.resize(label, (self.config.img_size, self.config.img_size))\n        return (label > 0).astype(np.float32)\n    \n    def __len__(self):\n        return len(self.samples) * 16  # Augment by using multiple patches\n    \n    def __getitem__(self, idx):\n        volume_idx = idx % len(self.samples)\n        volume, label = self.samples[volume_idx]\n        \n        # Create patches for training\n        if self.is_training:\n            h, w = volume.shape[1], volume.shape[2]\n            patch_size = 256\n            \n            # Random crop\n            top = np.random.randint(0, h - patch_size)\n            left = np.random.randint(0, w - patch_size)\n            \n            volume = volume[:, top:top+patch_size, left:left+patch_size]\n            label = label[top:top+patch_size, left:left+patch_size]\n            \n            # Random augmentations\n            if np.random.random() > 0.5:\n                volume = np.flip(volume, axis=2).copy()\n                label = np.flip(label, axis=1).copy()\n            \n            if np.random.random() > 0.5:\n                volume = np.flip(volume, axis=1).copy()\n                label = np.flip(label, axis=0).copy()\n            \n            if np.random.random() > 0.5:\n                k = np.random.randint(1, 4)\n                volume = np.rot90(volume, k, axes=(1, 2)).copy()\n                label = np.rot90(label, k).copy()\n        \n        # Normalize\n        volume = volume.astype(np.float32) / 255.0\n        \n        # Convert to tensors\n        volume_tensor = torch.FloatTensor(volume)\n        label_tensor = torch.FloatTensor(label).unsqueeze(0)\n        \n        return volume_tensor, label_tensor\n\n# Functions for threshold optimization\ndef compute_metrics_for_threshold(pred_probs, targets, threshold):\n    \"\"\"Compute metrics for a specific threshold\"\"\"\n    # Apply threshold\n    pred_binary = (pred_probs > threshold).float()\n    \n    # Convert to numpy for sklearn metrics\n    pred_np = pred_binary.cpu().numpy().flatten()\n    target_np = targets.cpu().numpy().flatten()\n    \n    # Handle case where there are no positive samples\n    if np.sum(target_np) == 0 and np.sum(pred_np) == 0:\n        return {'dice': 1.0, 'f05': 1.0, 'precision': 1.0, 'recall': 1.0}\n    elif np.sum(target_np) == 0:\n        return {'dice': 0.0, 'f05': 0.0, 'precision': 0.0, 'recall': 0.0}\n    \n    # Dice score\n    intersection = (pred_binary * targets).sum().item()\n    union = pred_binary.sum().item() + targets.sum().item()\n    dice = (2. * intersection + 1e-6) / (union + 1e-6)\n    \n    # F0.5 score\n    try:\n        f05 = fbeta_score(target_np, pred_np, beta=0.5, zero_division=0)\n    except:\n        f05 = 0.0\n    \n    # Precision and recall\n    try:\n        precision = precision_score(target_np, pred_np, zero_division=0)\n        recall = recall_score(target_np, pred_np, zero_division=0)\n    except:\n        precision = 0.0\n        recall = 0.0\n    \n    return {'dice': dice, 'f05': f05, 'precision': precision, 'recall': recall}\n\ndef find_optimal_threshold(pred_probs, targets, threshold_range, metric='f05'):\n    \"\"\"Find optimal threshold for a specific metric\"\"\"\n    best_threshold = 0.5\n    best_score = -1.0\n    \n    for threshold in threshold_range:\n        metrics = compute_metrics_for_threshold(pred_probs, targets, threshold)\n        score = metrics[metric]\n        \n        if score > best_score:\n            best_score = score\n            best_threshold = threshold\n    \n    return best_threshold, best_score\n\ndef find_optimal_thresholds(pred_probs, targets, threshold_range):\n    \"\"\"Find optimal thresholds for multiple metrics\"\"\"\n    thresholds = {}\n    scores = {}\n    \n    for metric in ['dice', 'f05']:\n        best_thresh, best_score = find_optimal_threshold(\n            pred_probs, targets, threshold_range, metric\n        )\n        thresholds[metric] = best_thresh\n        scores[metric] = best_score\n    \n    return thresholds, scores\n\n# Supervised Training with THRESHOLD OPTIMIZATION\ndef train_supervised():\n    print(\"Starting Supervised Training with Threshold Optimization...\")\n    print(f\"Threshold search range: {config.threshold_range}\")\n    \n    try:\n        # Create datasets\n        train_dataset = VesuviusDataset([train_paths[0]], is_training=True)\n        val_dataset = VesuviusDataset([train_paths[1]], is_training=False)\n        \n        train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=2, pin_memory=True)\n        val_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n        \n        print(f\"Train samples: {len(train_dataset)}, Val samples: {len(val_dataset)}\")\n        \n        # Create ensemble of models\n        models = []\n        optimizers = []\n        \n        for i in range(config.num_models):\n            model = SimpleUNet(in_channels=len(config.slices)).to(config.device)\n            \n            # Initialize weights\n            def init_weights(m):\n                if isinstance(m, nn.Conv2d):\n                    nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                    if m.bias is not None:\n                        nn.init.constant_(m.bias, 0)\n            \n            model.apply(init_weights)\n            \n            optimizer = optim.AdamW(model.parameters(), lr=config.learning_rate, weight_decay=1e-5)\n            models.append(model)\n            optimizers.append(optimizer)\n        \n        # Use Combined Loss\n        criterion = CombinedLoss(\n            alpha=config.focal_alpha,\n            gamma=config.focal_gamma,\n            smooth=config.dice_smooth\n        )\n        \n        # Training history\n        history = {\n            'train_loss': [[] for _ in range(config.num_models)],\n            'val_dice': [[] for _ in range(config.num_models)],\n            'val_f05': [[] for _ in range(config.num_models)],\n            'val_precision': [[] for _ in range(config.num_models)],\n            'val_recall': [[] for _ in range(config.num_models)],\n            'val_threshold_dice': [[] for _ in range(config.num_models)],  # Optimal threshold for Dice\n            'val_threshold_f05': [[] for _ in range(config.num_models)],   # Optimal threshold for F0.5\n            'val_threshold_used': [[] for _ in range(config.num_models)]   # Which threshold was used\n        }\n        \n        best_dice = 0.0\n        best_models = [None] * config.num_models\n        best_val_thresholds = [0.5] * config.num_models  # Store best thresholds per model\n        \n        for epoch in range(config.num_epochs):\n            print(f\"\\n{'='*60}\")\n            print(f\"Epoch {epoch+1}/{config.num_epochs}\")\n            print('='*60)\n            \n            # Training phase\n            for model_idx, (model, optimizer) in enumerate(zip(models, optimizers)):\n                model.train()\n                epoch_loss = 0\n                batch_count = 0\n                \n                pbar = tqdm(train_loader, desc=f\"Model {model_idx+1} Training\", leave=False)\n                for volume, labels in pbar:\n                    volume = volume.to(config.device)\n                    labels = labels.to(config.device)\n                    \n                    # Forward pass\n                    outputs = model(volume)\n                    loss = criterion(outputs, labels)\n                    \n                    # Backward pass\n                    optimizer.zero_grad()\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                    optimizer.step()\n                    \n                    epoch_loss += loss.item()\n                    batch_count += 1\n                    \n                    # Update progress bar\n                    pbar.set_postfix({'Loss': f'{loss.item():.4f}'})\n                \n                if batch_count > 0:\n                    avg_loss = epoch_loss / batch_count\n                    history['train_loss'][model_idx].append(avg_loss)\n            \n            # Validation phase with threshold optimization\n            print(\"\\nValidation with Threshold Optimization:\")\n            for model_idx, model in enumerate(models):\n                model.eval()\n                \n                # Collect all predictions and targets for threshold search\n                all_preds = []\n                all_targets = []\n                \n                with torch.no_grad():\n                    for volume, labels in val_loader:\n                        volume = volume.to(config.device)\n                        labels = labels.to(config.device)\n                        \n                        outputs = model(volume)\n                        preds = torch.sigmoid(outputs)\n                        \n                        all_preds.append(preds.cpu())\n                        all_targets.append(labels.cpu())\n                \n                # Concatenate all batches\n                all_preds = torch.cat(all_preds, dim=0)\n                all_targets = torch.cat(all_targets, dim=0)\n                \n                # Find optimal thresholds for validation set\n                print(f\"\\nModel {model_idx+1} - Searching optimal thresholds...\")\n                thresholds, scores = find_optimal_thresholds(\n                    all_preds, all_targets, config.threshold_range\n                )\n                \n                # Choose best threshold based on F0.5 score (competition metric)\n                optimal_threshold = thresholds['f05']\n                optimal_score = scores['f05']\n                \n                # Store thresholds in history\n                history['val_threshold_dice'][model_idx].append(thresholds['dice'])\n                history['val_threshold_f05'][model_idx].append(thresholds['f05'])\n                history['val_threshold_used'][model_idx].append(optimal_threshold)\n                \n                # Compute metrics with optimal threshold\n                val_metrics = compute_metrics_for_threshold(\n                    all_preds, all_targets, optimal_threshold\n                )\n                \n                # Store metrics\n                for metric, value in val_metrics.items():\n                    history[f'val_{metric}'][model_idx].append(value)\n                \n                print(f\"Model {model_idx+1} Results:\")\n                print(f\"  Training Loss: {history['train_loss'][model_idx][-1]:.4f}\")\n                print(f\"  Optimal Threshold (Dice): {thresholds['dice']:.3f}, Dice: {scores['dice']:.4f}\")\n                print(f\"  Optimal Threshold (F0.5): {thresholds['f05']:.3f}, F0.5: {scores['f05']:.4f}\")\n                print(f\"  Using Threshold: {optimal_threshold:.3f}\")\n                print(f\"  Final Metrics - Dice: {val_metrics['dice']:.4f}, \"\n                      f\"F0.5: {val_metrics['f05']:.4f}, \"\n                      f\"Precision: {val_metrics['precision']:.4f}, \"\n                      f\"Recall: {val_metrics['recall']:.4f}\")\n                \n                # Save best model based on F0.5 score\n                if val_metrics['f05'] > best_dice:\n                    best_dice = val_metrics['f05']\n                    best_models[model_idx] = model.state_dict().copy()\n                    best_val_thresholds[model_idx] = optimal_threshold\n                    torch.save({\n                        'model_state_dict': model.state_dict(),\n                        'threshold': optimal_threshold,\n                        'f05_score': val_metrics['f05'],\n                        'dice_score': val_metrics['dice'],\n                        'epoch': epoch\n                    }, f'best_model_{model_idx}.pth')\n                \n                # Save visualizations for first batch with optimal threshold\n                if val_metrics['dice'] > 0.50 or val_metrics['f05'] > 0.50:\n                    # Get first batch for visualization\n                    with torch.no_grad():\n                        volume, labels = next(iter(val_loader))\n                        volume = volume.to(config.device)\n                        labels = labels.to(config.device)\n                        outputs = model(volume)\n                        preds = torch.sigmoid(outputs)\n                        save_visualization(volume, preds, labels, 0, epoch, model_idx, \n                                         'validation', optimal_threshold)\n            \n            # Early stopping check\n            current_loss = np.mean([history['train_loss'][i][-1] for i in range(config.num_models)])\n            if current_loss < 0.3 and epoch > 5:\n                print(f\"\\nGood average loss achieved: {current_loss:.4f}\")\n                break\n            \n            if current_loss > 5.0:\n                print(f\"\\nLoss exploding: {current_loss:.4f}\")\n                break\n        \n        # Load best models and their thresholds\n        for i in range(config.num_models):\n            if best_models[i] is not None:\n                models[i].load_state_dict(best_models[i])\n        \n        return models, history, best_val_thresholds\n    \n    except Exception as e:\n        print(f\"Supervised training failed: {e}\")\n        import traceback\n        traceback.print_exc()\n        return None, None, None\n\ndef test_ensemble_with_threshold_optimization(models, test_path, val_thresholds=None):\n    \"\"\"Test ensemble with separate threshold optimization for test set\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"Testing Ensemble with Threshold Optimization...\")\n    print(\"=\"*60)\n    \n    if models is None:\n        print(\"No models to test!\")\n        return {'dice': 0.0, 'f05': 0.0, 'precision': 0.0, 'recall': 0.0}\n    \n    try:\n        test_dataset = VesuviusDataset([test_path], is_training=False)\n        test_loader = DataLoader(test_dataset, batch_size=config.batch_size, shuffle=False, \n                               num_workers=2, pin_memory=True)\n        \n        # Store predictions from all models\n        all_ensemble_preds = []\n        all_test_targets = []\n        \n        with torch.no_grad():\n            for batch_idx, (volume, labels) in enumerate(test_loader):\n                volume = volume.to(config.device)\n                labels = labels.to(config.device)\n                \n                # Ensemble predictions\n                ensemble_preds = []\n                for model in models:\n                    model.eval()\n                    output = model(volume)\n                    pred = torch.sigmoid(output)\n                    ensemble_preds.append(pred)\n                \n                # Average predictions\n                ensemble_pred = torch.mean(torch.stack(ensemble_preds), dim=0)\n                \n                all_ensemble_preds.append(ensemble_pred.cpu())\n                all_test_targets.append(labels.cpu())\n        \n        # Concatenate all batches\n        all_ensemble_preds = torch.cat(all_ensemble_preds, dim=0)\n        all_test_targets = torch.cat(all_test_targets, dim=0)\n        \n        print(\"\\nSearching for optimal test thresholds...\")\n        \n        # Find optimal thresholds for TEST set (separate from validation)\n        test_thresholds, test_scores = find_optimal_thresholds(\n            all_ensemble_preds, all_test_targets, config.threshold_range\n        )\n        \n        print(f\"Test Set Optimal Thresholds:\")\n        print(f\"  For Dice: {test_thresholds['dice']:.3f} (Dice: {test_scores['dice']:.4f})\")\n        print(f\"  For F0.5: {test_thresholds['f05']:.3f} (F0.5: {test_scores['f05']:.4f})\")\n        \n        # Compare with validation thresholds if provided\n        if val_thresholds is not None:\n            avg_val_threshold = np.mean(val_thresholds)\n            print(f\"\\nAverage Validation Threshold: {avg_val_threshold:.3f}\")\n            print(f\"Difference from Test Threshold (F0.5): {abs(test_thresholds['f05'] - avg_val_threshold):.3f}\")\n        \n        # Use F0.5 optimal threshold for test set (since that's the competition metric)\n        optimal_test_threshold = test_thresholds['f05']\n        \n        # Compute final metrics with optimal test threshold\n        final_metrics = compute_metrics_for_threshold(\n            all_ensemble_preds, all_test_targets, optimal_test_threshold\n        )\n        \n        # Also compute metrics with validation thresholds for comparison\n        if val_thresholds is not None:\n            print(f\"\\nComparison with Validation Thresholds:\")\n            for i, val_thresh in enumerate(val_thresholds):\n                val_metrics = compute_metrics_for_threshold(\n                    all_ensemble_preds, all_test_targets, val_thresh\n                )\n                print(f\"  Model {i+1} Val Threshold {val_thresh:.3f}: \"\n                      f\"Dice={val_metrics['dice']:.4f}, F0.5={val_metrics['f05']:.4f}\")\n        \n        print(f\"\\nFinal Test Results with Optimal Test Threshold ({optimal_test_threshold:.3f}):\")\n        print(f\"Dice Score: {final_metrics['dice']:.4f}\")\n        print(f\"F0.5 Score: {final_metrics['f05']:.4f}\")\n        print(f\"Precision: {final_metrics['precision']:.4f}\")\n        print(f\"Recall: {final_metrics['recall']:.4f}\")\n        \n        # Save visualization with optimal test threshold\n        with torch.no_grad():\n            volume, labels = next(iter(test_loader))\n            volume = volume.to(config.device)\n            labels = labels.to(config.device)\n            \n            # Ensemble predictions for visualization\n            ensemble_preds = []\n            for model in models:\n                model.eval()\n                output = model(volume)\n                pred = torch.sigmoid(output)\n                ensemble_preds.append(pred)\n            \n            ensemble_pred = torch.mean(torch.stack(ensemble_preds), dim=0)\n            \n            save_visualization(volume, ensemble_pred, labels, 0, 0, 'ensemble', \n                             'testing', optimal_test_threshold)\n        \n        # Return results including thresholds\n        results = {\n            'metrics': final_metrics,\n            'thresholds': {\n                'test_optimal_f05': test_thresholds['f05'],\n                'test_optimal_dice': test_thresholds['dice'],\n                'test_scores': test_scores,\n                'val_thresholds': val_thresholds if val_thresholds else None\n            }\n        }\n        \n        return results\n    \n    except Exception as e:\n        print(f\"Testing failed: {e}\")\n        return {'metrics': {'dice': 0.0, 'f05': 0.0, 'precision': 0.0, 'recall': 0.0},\n                'thresholds': None}\n\ndef save_visualization(volume, pred, label, batch_idx, epoch, model_idx, phase, threshold):\n    os.makedirs(f'visualizations/{phase}', exist_ok=True)\n    \n    try:\n        # Get first sample in batch\n        vol_sample = volume[0].cpu().numpy()\n        pred_sample = pred[0, 0].cpu().numpy()\n        label_sample = label[0, 0].cpu().numpy()\n        \n        # Apply threshold to prediction\n        pred_binary = (pred_sample > threshold).astype(np.float32)\n        \n        # Take mean across slices for visualization\n        vol_viz = np.mean(vol_sample, axis=0)\n        \n        # Create figure\n        fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n        \n        # Row 1: Input and predictions\n        axes[0, 0].imshow(vol_viz, cmap='gray')\n        axes[0, 0].set_title('Input (Mean Slice)')\n        axes[0, 0].axis('off')\n        \n        im1 = axes[0, 1].imshow(pred_sample, cmap='jet', vmin=0, vmax=1)\n        axes[0, 1].set_title('Prediction (Probability)')\n        axes[0, 1].axis('off')\n        plt.colorbar(im1, ax=axes[0, 1], fraction=0.046, pad=0.04)\n        \n        axes[0, 2].imshow(pred_binary, cmap='jet', vmin=0, vmax=1)\n        axes[0, 2].set_title(f'Thresholded (t={threshold:.2f})')\n        axes[0, 2].axis('off')\n        \n        # Row 2: Ground truth and comparison\n        axes[1, 0].imshow(label_sample, cmap='jet', vmin=0, vmax=1)\n        axes[1, 0].set_title('Ground Truth')\n        axes[1, 0].axis('off')\n        \n        # Overlay prediction on input\n        overlay = vol_viz.copy()\n        overlay = (overlay - overlay.min()) / (overlay.max() - overlay.min())\n        overlay[pred_binary > 0.5] = 1.0  # Highlight predictions\n        axes[1, 1].imshow(overlay, cmap='gray')\n        axes[1, 1].set_title('Prediction Overlay')\n        axes[1, 1].axis('off')\n        \n        # Error map (FP in red, FN in blue)\n        error_map = np.zeros((*pred_binary.shape, 3))\n        error_map[pred_binary > 0.5] = [1, 0, 0]  # FP in red\n        error_map[label_sample > 0.5] = [0, 0, 1]  # FN in blue\n        overlap = (pred_binary > 0.5) & (label_sample > 0.5)\n        error_map[overlap] = [0, 1, 0]  # TP in green\n        axes[1, 2].imshow(error_map)\n        axes[1, 2].set_title('Error Map (Red=FP, Blue=FN, Green=TP)')\n        axes[1, 2].axis('off')\n        \n        plt.suptitle(f'{phase.capitalize()} - Epoch {epoch}, Model {model_idx}, Threshold={threshold:.3f}', \n                    fontsize=16, y=1.02)\n        plt.tight_layout()\n        plt.savefig(f'visualizations/{phase}/epoch_{epoch}_model_{model_idx}_batch_{batch_idx}.png', \n                    dpi=150, bbox_inches='tight')\n        plt.close()\n        print(f\"Saved visualization: visualizations/{phase}/epoch_{epoch}_model_{model_idx}_batch_{batch_idx}.png\")\n    except Exception as e:\n        print(f\"Failed to save visualization: {e}\")\n\n# Main execution\nif __name__ == \"__main__\":\n    print(f\"Using device: {config.device}\")\n    print(f\"Threshold search range ({len(config.threshold_range)} points): {config.threshold_range}\")\n    \n    os.makedirs('visualizations/validation', exist_ok=True)\n    os.makedirs('visualizations/testing', exist_ok=True)\n    \n    try:\n        # Start supervised training with threshold optimization\n        models, history, best_val_thresholds = train_supervised()\n        \n        # Test ensemble with separate threshold optimization\n        test_results = test_ensemble_with_threshold_optimization(models, test_path, best_val_thresholds)\n        \n        # Print final results\n        print(\"\\n\" + \"=\"*60)\n        print(\"FINAL RESULTS SUMMARY\")\n        print(\"=\"*60)\n        \n        if test_results['thresholds']:\n            print(f\"\\nOptimal Thresholds Found:\")\n            print(f\"  Test Set (F0.5): {test_results['thresholds']['test_optimal_f05']:.3f}\")\n            print(f\"  Test Set (Dice): {test_results['thresholds']['test_optimal_dice']:.3f}\")\n            if test_results['thresholds']['val_thresholds']:\n                print(f\"  Validation Set Avg: {np.mean(test_results['thresholds']['val_thresholds']):.3f}\")\n        \n        print(f\"\\nFinal Test Metrics:\")\n        print(f\"  Dice Score: {test_results['metrics']['dice']:.4f}\")\n        print(f\"  F0.5 Score: {test_results['metrics']['f05']:.4f}\")\n        print(f\"  Precision: {test_results['metrics']['precision']:.4f}\")\n        print(f\"  Recall: {test_results['metrics']['recall']:.4f}\")\n        print(\"=\"*60)\n        \n        if history is not None:\n            np.save('training_history.npy', history)\n            \n            # Plot comprehensive training curves\n            fig = plt.figure(figsize=(20, 12))\n            \n            # Plot 1: Training Loss\n            ax1 = plt.subplot(2, 3, 1)\n            for i in range(config.num_models):\n                ax1.plot(history['train_loss'][i], label=f'Model {i+1}', linewidth=2)\n            ax1.axhline(y=0.75, color='r', linestyle='--', linewidth=2, label='Target (0.75)')\n            ax1.set_title('Training Loss', fontsize=14, fontweight='bold')\n            ax1.set_xlabel('Epoch')\n            ax1.set_ylabel('Loss')\n            ax1.legend()\n            ax1.grid(True, alpha=0.3)\n            \n            # Plot 2: Validation Dice\n            ax2 = plt.subplot(2, 3, 2)\n            for i in range(config.num_models):\n                ax2.plot(history['val_dice'][i], label=f'Model {i+1}', linewidth=2)\n            ax2.set_title('Validation Dice Score', fontsize=14, fontweight='bold')\n            ax2.set_xlabel('Epoch')\n            ax2.set_ylabel('Dice')\n            ax2.legend()\n            ax2.grid(True, alpha=0.3)\n            \n            # Plot 3: Validation F0.5\n            ax3 = plt.subplot(2, 3, 3)\n            for i in range(config.num_models):\n                ax3.plot(history['val_f05'][i], label=f'Model {i+1}', linewidth=2)\n            ax3.set_title('Validation F0.5 Score', fontsize=14, fontweight='bold')\n            ax3.set_xlabel('Epoch')\n            ax3.set_ylabel('F0.5')\n            ax3.legend()\n            ax3.grid(True, alpha=0.3)\n            \n            # Plot 4: Optimal Thresholds (Dice)\n            ax4 = plt.subplot(2, 3, 4)\n            for i in range(config.num_models):\n                ax4.plot(history['val_threshold_dice'][i], label=f'Model {i+1}', linewidth=2)\n            ax4.set_title('Optimal Thresholds for Dice', fontsize=14, fontweight='bold')\n            ax4.set_xlabel('Epoch')\n            ax4.set_ylabel('Threshold')\n            ax4.legend()\n            ax4.grid(True, alpha=0.3)\n            \n            # Plot 5: Optimal Thresholds (F0.5)\n            ax5 = plt.subplot(2, 3, 5)\n            for i in range(config.num_models):\n                ax5.plot(history['val_threshold_f05'][i], label=f'Model {i+1}', linewidth=2)\n            ax5.set_title('Optimal Thresholds for F0.5', fontsize=14, fontweight='bold')\n            ax5.set_xlabel('Epoch')\n            ax5.set_ylabel('Threshold')\n            ax5.legend()\n            ax5.grid(True, alpha=0.3)\n            \n            # Plot 6: Used Thresholds\n            ax6 = plt.subplot(2, 3, 6)\n            for i in range(config.num_models):\n                ax6.plot(history['val_threshold_used'][i], label=f'Model {i+1}', linewidth=2)\n            ax6.set_title('Used Thresholds in Validation', fontsize=14, fontweight='bold')\n            ax6.set_xlabel('Epoch')\n            ax6.set_ylabel('Threshold')\n            ax6.legend()\n            ax6.grid(True, alpha=0.3)\n            \n            plt.suptitle('Training Progress with Threshold Optimization', fontsize=16, fontweight='bold', y=1.02)\n            plt.tight_layout()\n            plt.savefig('training_curves_with_thresholds.png', dpi=150, bbox_inches='tight')\n            plt.close()\n            \n            # Print best metrics\n            print(f\"\\nBest Validation Metrics per Model:\")\n            for i in range(config.num_models):\n                if history['val_f05'][i]:  # Check if list is not empty\n                    best_loss = min(history['train_loss'][i])\n                    best_dice = max(history['val_dice'][i])\n                    best_f05 = max(history['val_f05'][i])\n                    best_thresh_f05 = history['val_threshold_f05'][i][np.argmax(history['val_f05'][i])]\n                    best_thresh_dice = history['val_threshold_dice'][i][np.argmax(history['val_dice'][i])]\n                    print(f\"Model {i+1}: Loss={best_loss:.4f}, Dice={best_dice:.4f} (t={best_thresh_dice:.3f}), \"\n                          f\"F0.5={best_f05:.4f} (t={best_thresh_f05:.3f})\")\n        \n        print(\"\\n\" + \"=\"*60)\n        print(\"Training completed!\")\n        print(\"Check 'visualizations' folder for prediction visualizations\")\n        print(\"Check 'training_curves_with_thresholds.png' for comprehensive training history\")\n        print(\"=\"*60)\n        \n    except Exception as e:\n        print(f\"Error occurred: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#CONVNEXT WITH SPLIT TO 12 \nimport os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score, f1_score\nimport torch.cuda as cuda\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 224,\n    \"num_splits\": 12,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        # Apply transforms that don't change the aspect ratio\n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n        ]) if is_train else None\n\n    def _load_volume(self):\n        \"\"\"Load volume slices with robust error handling\"\"\"\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            if img.size == 0:\n                raise ValueError(f\"Empty image {p}\")\n            volume.append(img)\n        return np.clip(np.stack(volume), 0, 255).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        \"\"\"Load mask with robust error handling\"\"\"\n        if self.mask_path is None:\n            return None\n            \n        try:\n            mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"Failed to load mask {self.mask_path}\")\n            if mask.size == 0:\n                raise ValueError(f\"Empty mask {self.mask_path}\")\n            mask = mask.astype(np.float32) / 255.0\n            print(f\"Mask loaded - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")  # Debug\n            return mask\n        except Exception as e:\n            raise ValueError(f\"Failed to load mask {self.mask_path}: {str(e)}\")\n\n    def _generate_coords(self):\n        \"\"\"Generate coordinates for valid patches\"\"\"\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        \"\"\"Get item with robust error handling\"\"\"\n        try:\n            if self.is_train:\n                orig_idx = idx // CFG['num_splits']\n                split_idx = idx % CFG['num_splits']\n            else:\n                orig_idx = idx\n                split_idx = None\n\n            i, j = self.coords[orig_idx]\n            \n            # Get volume patch\n            patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Process mask\n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n                mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply split first to ensure consistent sizes\n            if split_idx is not None and mask is not None:\n                h, w = mask.shape[:2]\n                rows, cols = 3, 4  # 3x4 grid for 12 splits\n                split_h, split_w = h // rows, w // cols\n                \n                # Make sure all splits have the same size\n                split_h = (h // rows) \n                split_w = (w // cols)\n                \n                row = split_idx // cols\n                col = split_idx % cols\n                \n                # Ensure we don't go out of bounds\n                end_h = min((row+1)*split_h, h)\n                end_w = min((col+1)*split_w, w)\n                \n                patch = patch[row*split_h:end_h, col*split_w:end_w]\n                mask = mask[row*split_h:end_h, col*split_w:end_w]\n                \n                # Resize to consistent size if needed\n                if patch.shape[0] != split_h or patch.shape[1] != split_w:\n                    patch = cv2.resize(patch, (split_w, split_h))\n                    mask = cv2.resize(mask, (split_w, split_h), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply transforms after splitting (only flips, no rotations)\n            if self.is_train and mask is not None and self.transform is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                patch = transformed['image']\n                mask = transformed['mask']\n            \n            # Convert to tensor\n            image = torch.tensor(patch).permute(2, 0, 1)  # CHW\n            if mask is not None:\n                mask = torch.tensor(mask)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx}: {str(e)}\")\n            # Return consistent dummy data\n            dummy_img = torch.zeros((CFG['slice_end']-CFG['slice_start'], CFG['target_size']//3, CFG['target_size']//4))\n            dummy_mask = torch.zeros((CFG['target_size']//3, CFG['target_size']//4)) if self.mask is not None else None\n            return (dummy_img, dummy_mask) if dummy_mask is not None else dummy_img\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        # Update final upsample to match split size\n        split_h = CFG['target_size'] // 3\n        split_w = CFG['target_size'] // 4\n        self.final_upsample = nn.Upsample(size=(split_h, split_w))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = nn.functional.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(preds, size=masks.shape[-2:], mode='bilinear')\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef dice_score(y_true, y_pred):\n    \"\"\"Calculate Dice coefficient\"\"\"\n    y_true_f = y_true.flatten()\n    y_pred_f = y_pred.flatten()\n    intersection = np.sum(y_true_f * y_pred_f)\n    return (2. * intersection + 1) / (np.sum(y_true_f) + np.sum(y_pred_f) + 1)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    # Debug mask values\n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    # Binarize labels properly\n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5, 0.0, 0.0, 0.0\n    \n    # Find optimal threshold from 0.3 to 0.9\n    best_thresh, best_f05 = 0.3, 0\n    thresholds = np.linspace(0.3, 0.9, 13)\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n            \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    f1 = f1_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    fb = fbeta_score(all_labels.flatten(), final_binarized.flatten(), beta=2, zero_division=0)\n    dice_val = dice_score(all_labels, final_binarized)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh, dice_val, f1, fb\n\ndef reconstruct_full_prediction(model, dataset, device):\n    \"\"\"Reconstruct full prediction for visualization\"\"\"\n    model.eval()\n    z, h, w = dataset.volume.shape\n    full_pred = np.zeros((h, w), dtype=np.float32)\n    count = np.zeros((h, w), dtype=np.float32)\n    \n    with torch.no_grad():\n        for i, j in tqdm(dataset.coords, desc=\"Reconstructing full prediction\"):\n            # Get volume patch\n            patch = dataset.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Convert to tensor and predict\n            image_tensor = torch.tensor(patch).permute(2, 0, 1).unsqueeze(0).to(device)\n            \n            with autocast(enabled=CFG['amp']):\n                pred = torch.sigmoid(model(image_tensor)).cpu().numpy()[0, 0]\n            \n            # Resize prediction back to original patch size\n            pred_resized = cv2.resize(pred, (CFG['target_size'], CFG['target_size']))\n            \n            # Add to full prediction\n            full_pred[i:i+CFG['target_size'], j:j+CFG['target_size']] += pred_resized\n            count[i:i+CFG['target_size'], j:j+CFG['target_size']] += 1\n    \n    # Average overlapping regions\n    full_pred = np.divide(full_pred, count, out=np.zeros_like(full_pred), where=count != 0)\n    return full_pred\n\ndef visualize_comparison(full_input, full_label, full_pred, threshold=0.5, metrics=None):\n    \"\"\"Visualize comparison of full input, label and prediction\"\"\"\n    plt.figure(figsize=(20, 10))\n    \n    # Input (middle slice)\n    plt.subplot(2, 3, 1)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.title(\"Input Volume (Middle Slice)\")\n    plt.axis('off')\n    \n    # Ground truth\n    plt.subplot(2, 3, 2)\n    plt.imshow(full_label, cmap='gray')\n    plt.title(\"Ground Truth\")\n    plt.axis('off')\n    \n    # Prediction\n    plt.subplot(2, 3, 3)\n    plt.imshow(full_pred, cmap='gray')\n    plt.title(f\"Prediction (Threshold: {threshold:.2f})\")\n    plt.axis('off')\n    \n    # Thresholded prediction\n    plt.subplot(2, 3, 4)\n    plt.imshow(full_pred > threshold, cmap='gray')\n    plt.title(\"Thresholded Prediction\")\n    plt.axis('off')\n    \n    # Overlay\n    plt.subplot(2, 3, 5)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.imshow(full_pred > threshold, cmap='jet', alpha=0.3)\n    plt.title(\"Input with Prediction Overlay\")\n    plt.axis('off')\n    \n    # Metrics text\n    plt.subplot(2, 3, 6)\n    plt.axis('off')\n    if metrics:\n        metrics_text = f\"\"\"\n        Evaluation Metrics:\n        Dice Score: {metrics['dice']:.4f}\n        Precision: {metrics['precision']:.4f}\n        Recall: {metrics['recall']:.4f}\n        F0.5 Score: {metrics['f05']:.4f}\n        F1 Score: {metrics['f1']:.4f}\n        F2 Score: {metrics['f2']:.4f}\n        Optimal Threshold: {threshold:.2f}\n        \"\"\"\n        plt.text(0.1, 0.5, metrics_text, fontsize=14, va='center')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh, dice_val, f1, fb = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - Dice: {dice_val:.4f}, F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, F1: {f1:.4f}, F2: {fb:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nReconstructing full prediction for visualization...\")\n                    full_pred = reconstruct_full_prediction(model, test_dataset, device)\n                    \n                    # Get full input and label\n                    full_input = test_dataset.volume\n                    full_label = test_dataset.mask\n                    \n                    # Visualize comparison\n                    metrics = {\n                        'dice': dice_val,\n                        'precision': prec,\n                        'recall': rec,\n                        'f05': f05,\n                        'f1': f1,\n                        'f2': fb\n                    }\n                    visualize_comparison(full_input, full_label, full_pred, threshold=best_thresh, metrics=metrics)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score, f1_score\nimport torch.cuda as cuda\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 224,\n    \"num_splits\": 12,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        # Apply transforms that don't change the aspect ratio\n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n        ]) if is_train else None\n\n    def _load_volume(self):\n        \"\"\"Load volume slices with robust error handling\"\"\"\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            if img.size == 0:\n                raise ValueError(f\"Empty image {p}\")\n            volume.append(img)\n        return np.clip(np.stack(volume), 0, 255).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        \"\"\"Load mask with robust error handling\"\"\"\n        if self.mask_path is None:\n            return None\n            \n        try:\n            mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"Failed to load mask {self.mask_path}\")\n            if mask.size == 0:\n                raise ValueError(f\"Empty mask {self.mask_path}\")\n            mask = mask.astype(np.float32) / 255.0\n            print(f\"Mask loaded - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")  # Debug\n            return mask\n        except Exception as e:\n            raise ValueError(f\"Failed to load mask {self.mask_path}: {str(e)}\")\n\n    def _generate_coords(self):\n        \"\"\"Generate coordinates for valid patches\"\"\"\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        \"\"\"Get item with robust error handling\"\"\"\n        try:\n            if self.is_train:\n                orig_idx = idx // CFG['num_splits']\n                split_idx = idx % CFG['num_splits']\n            else:\n                orig_idx = idx\n                split_idx = None\n\n            i, j = self.coords[orig_idx]\n            \n            # Get volume patch\n            patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Process mask\n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n                mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply split first to ensure consistent sizes\n            if split_idx is not None and mask is not None:\n                h, w = mask.shape[:2]\n                rows, cols = 3, 4  # 3x4 grid for 12 splits\n                split_h, split_w = h // rows, w // cols\n                \n                # Make sure all splits have the same size\n                split_h = (h // rows) \n                split_w = (w // cols)\n                \n                row = split_idx // cols\n                col = split_idx % cols\n                \n                # Ensure we don't go out of bounds\n                end_h = min((row+1)*split_h, h)\n                end_w = min((col+1)*split_w, w)\n                \n                patch = patch[row*split_h:end_h, col*split_w:end_w]\n                mask = mask[row*split_h:end_h, col*split_w:end_w]\n                \n                # Resize to consistent size if needed\n                if patch.shape[0] != split_h or patch.shape[1] != split_w:\n                    patch = cv2.resize(patch, (split_w, split_h))\n                    mask = cv2.resize(mask, (split_w, split_h), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply transforms after splitting (only flips, no rotations)\n            if self.is_train and mask is not None and self.transform is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                patch = transformed['image']\n                mask = transformed['mask']\n            \n            # Convert to tensor\n            image = torch.tensor(patch).permute(2, 0, 1)  # CHW\n            if mask is not None:\n                mask = torch.tensor(mask)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx}: {str(e)}\")\n            # Return consistent dummy data\n            dummy_img = torch.zeros((CFG['slice_end']-CFG['slice_start'], CFG['target_size']//3, CFG['target_size']//4))\n            dummy_mask = torch.zeros((CFG['target_size']//3, CFG['target_size']//4)) if self.mask is not None else None\n            return (dummy_img, dummy_mask) if dummy_mask is not None else dummy_img\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        # Update final upsample to match split size\n        split_h = CFG['target_size'] // 3\n        split_w = CFG['target_size'] // 4\n        self.final_upsample = nn.Upsample(size=(split_h, split_w))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            # Convert to float32 for interpolation to avoid AMP issues\n            inputs = nn.functional.interpolate(inputs.float(), size=targets.shape[-2:], mode='bilinear').to(inputs.dtype)\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                # Convert to float32 for interpolation to avoid AMP issues\n                preds = nn.functional.interpolate(preds.float(), size=masks.shape[-2:], mode='bilinear').to(preds.dtype)\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef dice_score(y_true, y_pred):\n    \"\"\"Calculate Dice coefficient\"\"\"\n    y_true_f = y_true.flatten()\n    y_pred_f = y_pred.flatten()\n    intersection = np.sum(y_true_f * y_pred_f)\n    return (2. * intersection + 1) / (np.sum(y_true_f) + np.sum(y_pred_f) + 1)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                # Convert to float32 for interpolation to avoid AMP issues\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds).float(), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    # Debug mask values\n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    # Binarize labels properly\n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5, 0.0, 0.0, 0.0\n    \n    # Find optimal threshold from 0.3 to 0.9\n    best_thresh, best_f05 = 0.3, 0\n    thresholds = np.linspace(0.3, 0.9, 13)\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n            \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    f1 = f1_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    fb = fbeta_score(all_labels.flatten(), final_binarized.flatten(), beta=2, zero_division=0)\n    dice_val = dice_score(all_labels, final_binarized)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh, dice_val, f1, fb\n\ndef reconstruct_full_prediction(model, dataset, device):\n    \"\"\"Reconstruct full prediction for visualization\"\"\"\n    model.eval()\n    z, h, w = dataset.volume.shape\n    full_pred = np.zeros((h, w), dtype=np.float32)\n    count = np.zeros((h, w), dtype=np.float32)\n    \n    with torch.no_grad():\n        for i, j in tqdm(dataset.coords, desc=\"Reconstructing full prediction\"):\n            # Get volume patch\n            patch = dataset.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Convert to tensor and predict\n            image_tensor = torch.tensor(patch).permute(2, 0, 1).unsqueeze(0).to(device)\n            \n            with autocast(enabled=CFG['amp']):\n                pred = torch.sigmoid(model(image_tensor)).cpu().numpy()[0, 0]\n            \n            # Resize prediction back to original patch size\n            pred_resized = cv2.resize(pred, (CFG['target_size'], CFG['target_size']))\n            \n            # Add to full prediction\n            full_pred[i:i+CFG['target_size'], j:j+CFG['target_size']] += pred_resized\n            count[i:i+CFG['target_size'], j:j+CFG['target_size']] += 1\n    \n    # Average overlapping regions\n    full_pred = np.divide(full_pred, count, out=np.zeros_like(full_pred), where=count != 0)\n    return full_pred\n\ndef visualize_comparison(full_input, full_label, full_pred, threshold=0.5, metrics=None):\n    \"\"\"Visualize comparison of full input, label and prediction\"\"\"\n    plt.figure(figsize=(20, 10))\n    \n    # Input (middle slice)\n    plt.subplot(2, 3, 1)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.title(\"Input Volume (Middle Slice)\")\n    plt.axis('off')\n    \n    # Ground truth\n    plt.subplot(2, 3, 2)\n    plt.imshow(full_label, cmap='gray')\n    plt.title(\"Ground Truth\")\n    plt.axis('off')\n    \n    # Prediction\n    plt.subplot(2, 3, 3)\n    plt.imshow(full_pred, cmap='gray')\n    plt.title(f\"Prediction (Threshold: {threshold:.2f})\")\n    plt.axis('off')\n    \n    # Thresholded prediction\n    plt.subplot(2, 3, 4)\n    plt.imshow(full_pred > threshold, cmap='gray')\n    plt.title(\"Thresholded Prediction\")\n    plt.axis('off')\n    \n    # Overlay\n    plt.subplot(2, 3, 5)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.imshow(full_pred > threshold, cmap='jet', alpha=0.3)\n    plt.title(\"Input with Prediction Overlay\")\n    plt.axis('off')\n    \n    # Metrics text\n    plt.subplot(2, 3, 6)\n    plt.axis('off')\n    if metrics:\n        metrics_text = f\"\"\"\n        Evaluation Metrics:\n        Dice Score: {metrics['dice']:.4f}\n        Precision: {metrics['precision']:.4f}\n        Recall: {metrics['recall']:.4f}\n        F0.5 Score: {metrics['f05']:.4f}\n        F1 Score: {metrics['f1']:.4f}\n        F2 Score: {metrics['f2']:.4f}\n        Optimal Threshold: {threshold:.2f}\n        \"\"\"\n        plt.text(0.1, 0.5, metrics_text, fontsize=14, va='center')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh, dice_val, f1, fb = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - Dice: {dice_val:.4f}, F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, F1: {f1:.4f}, F2: {fb:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nReconstructing full prediction for visualization...\")\n                    full_pred = reconstruct_full_prediction(model, test_dataset, device)\n                    \n                    # Get full input and label\n                    full_input = test_dataset.volume\n                    full_label = test_dataset.mask\n                    \n                    # Visualize comparison\n                    metrics = {\n                        'dice': dice_val,\n                        'precision': prec,\n                        'recall': rec,\n                        'f05': f05,\n                        'f1': f1,\n                        'f2': fb\n                    }\n                    visualize_comparison(full_input, full_label, full_pred, threshold=best_thresh, metrics=metrics)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score, f1_score\nimport torch.cuda as cuda\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 224,\n    \"num_splits\": 12,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        # Apply transforms that don't change the aspect ratio\n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n        ]) if is_train else None\n\n    def _load_volume(self):\n        \"\"\"Load volume slices with robust error handling\"\"\"\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            if img.size == 0:\n                raise ValueError(f\"Empty image {p}\")\n            volume.append(img)\n        return np.clip(np.stack(volume), 0, 255).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        \"\"\"Load mask with robust error handling\"\"\"\n        if self.mask_path is None:\n            return None\n            \n        try:\n            mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"Failed to load mask {self.mask_path}\")\n            if mask.size == 0:\n                raise ValueError(f\"Empty mask {self.mask_path}\")\n            mask = mask.astype(np.float32) / 255.0\n            print(f\"Mask loaded - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")  # Debug\n            return mask\n        except Exception as e:\n            raise ValueError(f\"Failed to load mask {self.mask_path}: {str(e)}\")\n\n    def _generate_coords(self):\n        \"\"\"Generate coordinates for valid patches\"\"\"\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        \"\"\"Get item with robust error handling\"\"\"\n        try:\n            if self.is_train:\n                orig_idx = idx // CFG['num_splits']\n                split_idx = idx % CFG['num_splits']\n            else:\n                orig_idx = idx\n                split_idx = None\n\n            i, j = self.coords[orig_idx]\n            \n            # Get volume patch\n            patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Process mask\n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n                mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply split first to ensure consistent sizes\n            if split_idx is not None and mask is not None:\n                h, w = mask.shape[:2]\n                rows, cols = 3, 4  # 3x4 grid for 12 splits\n                split_h, split_w = h // rows, w // cols\n                \n                # Make sure all splits have the same size\n                split_h = (h // rows) \n                split_w = (w // cols)\n                \n                row = split_idx // cols\n                col = split_idx % cols\n                \n                # Ensure we don't go out of bounds\n                end_h = min((row+1)*split_h, h)\n                end_w = min((col+1)*split_w, w)\n                \n                patch = patch[row*split_h:end_h, col*split_w:end_w]\n                mask = mask[row*split_h:end_h, col*split_w:end_w]\n                \n                # Resize to consistent size if needed\n                if patch.shape[0] != split_h or patch.shape[1] != split_w:\n                    patch = cv2.resize(patch, (split_w, split_h))\n                    mask = cv2.resize(mask, (split_w, split_h), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply transforms after splitting (only flips, no rotations)\n            if self.is_train and mask is not None and self.transform is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                patch = transformed['image']\n                mask = transformed['mask']\n            \n            # Convert to tensor\n            image = torch.tensor(patch).permute(2, 0, 1)  # CHW\n            if mask is not None:\n                mask = torch.tensor(mask)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx}: {str(e)}\")\n            # Return consistent dummy data\n            dummy_img = torch.zeros((CFG['slice_end']-CFG['slice_start'], CFG['target_size']//3, CFG['target_size']//4))\n            dummy_mask = torch.zeros((CFG['target_size']//3, CFG['target_size']//4)) if self.mask is not None else None\n            return (dummy_img, dummy_mask) if dummy_mask is not None else dummy_img\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        # Update final upsample to match split size\n        split_h = CFG['target_size'] // 3\n        split_w = CFG['target_size'] // 4\n        self.final_upsample = nn.Upsample(size=(split_h, split_w))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            # Convert to float32 for interpolation to avoid AMP issues\n            inputs = nn.functional.interpolate(inputs.float(), size=targets.shape[-2:], mode='bilinear').to(inputs.dtype)\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                # Convert to float32 for interpolation to avoid AMP issues\n                preds = nn.functional.interpolate(preds.float(), size=masks.shape[-2:], mode='bilinear').to(preds.dtype)\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef dice_score(y_true, y_pred):\n    \"\"\"Calculate Dice coefficient\"\"\"\n    y_true_f = y_true.flatten()\n    y_pred_f = y_pred.flatten()\n    intersection = np.sum(y_true_f * y_pred_f)\n    return (2. * intersection + 1) / (np.sum(y_true_f) + np.sum(y_pred_f) + 1)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                # Convert to float32 for interpolation to avoid AMP issues\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds).float(), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    # Debug mask values\n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    # Binarize labels properly\n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5, 0.0, 0.0, 0.0\n    \n    # Find optimal threshold from 0.3 to 0.9\n    best_thresh, best_f05 = 0.3, 0\n    thresholds = np.linspace(0.3, 0.9, 13)\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n            \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    f1 = f1_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    fb = fbeta_score(all_labels.flatten(), final_binarized.flatten(), beta=2, zero_division=0)\n    dice_val = dice_score(all_labels, final_binarized)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh, dice_val, f1, fb\n\ndef reconstruct_full_prediction(model, dataset, device):\n    \"\"\"Reconstruct full prediction for visualization\"\"\"\n    model.eval()\n    z, h, w = dataset.volume.shape\n    full_pred = np.zeros((h, w), dtype=np.float32)\n    count = np.zeros((h, w), dtype=np.float32)\n    \n    with torch.no_grad():\n        for i, j in tqdm(dataset.coords, desc=\"Reconstructing full prediction\"):\n            # Get volume patch\n            patch = dataset.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Convert to tensor and predict\n            image_tensor = torch.tensor(patch).permute(2, 0, 1).unsqueeze(0).to(device)\n            \n            with autocast(enabled=CFG['amp']):\n                pred = torch.sigmoid(model(image_tensor)).cpu().numpy()[0, 0]\n            \n            # Ensure the prediction is in the correct format for OpenCV\n            pred = np.clip(pred, 0, 1).astype(np.float32)\n            \n            # Resize prediction back to original patch size\n            pred_resized = cv2.resize(pred, (CFG['target_size'], CFG['target_size']))\n            \n            # Add to full prediction\n            full_pred[i:i+CFG['target_size'], j:j+CFG['target_size']] += pred_resized\n            count[i:i+CFG['target_size'], j:j+CFG['target_size']] += 1\n    \n    # Average overlapping regions\n    full_pred = np.divide(full_pred, count, out=np.zeros_like(full_pred), where=count != 0)\n    return full_pred\n\ndef visualize_comparison(full_input, full_label, full_pred, threshold=0.5, metrics=None):\n    \"\"\"Visualize comparison of full input, label and prediction\"\"\"\n    plt.figure(figsize=(20, 10))\n    \n    # Input (middle slice)\n    plt.subplot(2, 3, 1)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.title(\"Input Volume (Middle Slice)\")\n    plt.axis('off')\n    \n    # Ground truth\n    plt.subplot(2, 3, 2)\n    plt.imshow(full_label, cmap='gray')\n    plt.title(\"Ground Truth\")\n    plt.axis('off')\n    \n    # Prediction\n    plt.subplot(2, 3, 3)\n    plt.imshow(full_pred, cmap='gray')\n    plt.title(f\"Prediction (Threshold: {threshold:.2f})\")\n    plt.axis('off')\n    \n    # Thresholded prediction\n    plt.subplot(2, 3, 4)\n    plt.imshow(full_pred > threshold, cmap='gray')\n    plt.title(\"Thresholded Prediction\")\n    plt.axis('off')\n    \n    # Overlay\n    plt.subplot(2, 3, 5)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.imshow(full_pred > threshold, cmap='jet', alpha=0.3)\n    plt.title(\"Input with Prediction Overlay\")\n    plt.axis('off')\n    \n    # Metrics text\n    plt.subplot(2, 3, 6)\n    plt.axis('off')\n    if metrics:\n        metrics_text = f\"\"\"\n        Evaluation Metrics:\n        Dice Score: {metrics['dice']:.4f}\n        Precision: {metrics['precision']:.4f}\n        Recall: {metrics['recall']:.4f}\n        F0.5 Score: {metrics['f05']:.4f}\n        F1 Score: {metrics['f1']:.4f}\n        F2 Score: {metrics['f2']:.4f}\n        Optimal Threshold: {threshold:.2f}\n        \"\"\"\n        plt.text(0.1, 0.5, metrics_text, fontsize=14, va='center')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh, dice_val, f1, fb = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - Dice: {dice_val:.4f}, F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, F1: {f1:.4f, F2: {fb:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nReconstructing full prediction for visualization...\")\n                    full_pred = reconstruct_full_prediction(model, test_dataset, device)\n                    \n                    # Get full input and label\n                    full_input = test_dataset.volume\n                    full_label = test_dataset.mask\n                    \n                    # Visualize comparison\n                    metrics = {\n                        'dice': dice_val,\n                        'precision': prec,\n                        'recall': rec,\n                        'f05': f05,\n                        'f1': f1,\n                        'f2': fb\n                    }\n                    visualize_comparison(full_input, full_label, full_pred, threshold=best_thresh, metrics=metrics)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score, f1_score\nimport torch.cuda as cuda\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 224,\n    \"num_splits\": 12,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        # Apply transforms that don't change the aspect ratio\n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n        ]) if is_train else None\n\n    def _load_volume(self):\n        \"\"\"Load volume slices with robust error handling\"\"\"\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            if img.size == 0:\n                raise ValueError(f\"Empty image {p}\")\n            volume.append(img)\n        return np.clip(np.stack(volume), 0, 255).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        \"\"\"Load mask with robust error handling\"\"\"\n        if self.mask_path is None:\n            return None\n            \n        try:\n            mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"Failed to load mask {self.mask_path}\")\n            if mask.size == 0:\n                raise ValueError(f\"Empty mask {self.mask_path}\")\n            mask = mask.astype(np.float32) / 255.0\n            print(f\"Mask loaded - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")  # Debug\n            return mask\n        except Exception as e:\n            raise ValueError(f\"Failed to load mask {self.mask_path}: {str(e)}\")\n\n    def _generate_coords(self):\n        \"\"\"Generate coordinates for valid patches\"\"\"\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        \"\"\"Get item with robust error handling\"\"\"\n        try:\n            if self.is_train:\n                orig_idx = idx // CFG['num_splits']\n                split_idx = idx % CFG['num_splits']\n            else:\n                orig_idx = idx\n                split_idx = None\n\n            i, j = self.coords[orig_idx]\n            \n            # Get volume patch\n            patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Process mask\n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n                mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply split first to ensure consistent sizes\n            if split_idx is not None and mask is not None:\n                h, w = mask.shape[:2]\n                rows, cols = 3, 4  # 3x4 grid for 12 splits\n                split_h, split_w = h // rows, w // cols\n                \n                # Make sure all splits have the same size\n                split_h = (h // rows) \n                split_w = (w // cols)\n                \n                row = split_idx // cols\n                col = split_idx % cols\n                \n                # Ensure we don't go out of bounds\n                end_h = min((row+1)*split_h, h)\n                end_w = min((col+1)*split_w, w)\n                \n                patch = patch[row*split_h:end_h, col*split_w:end_w]\n                mask = mask[row*split_h:end_h, col*split_w:end_w]\n                \n                # Resize to consistent size if needed\n                if patch.shape[0] != split_h or patch.shape[1] != split_w:\n                    patch = cv2.resize(patch, (split_w, split_h))\n                    mask = cv2.resize(mask, (split_w, split_h), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply transforms after splitting (only flips, no rotations)\n            if self.is_train and mask is not None and self.transform is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                patch = transformed['image']\n                mask = transformed['mask']\n            \n            # Convert to tensor\n            image = torch.tensor(patch).permute(2, 0, 1)  # CHW\n            if mask is not None:\n                mask = torch.tensor(mask)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx}: {str(e)}\")\n            # Return consistent dummy data\n            dummy_img = torch.zeros((CFG['slice_end']-CFG['slice_start'], CFG['target_size']//3, CFG['target_size']//4))\n            dummy_mask = torch.zeros((CFG['target_size']//3, CFG['target_size']//4)) if self.mask is not None else None\n            return (dummy_img, dummy_mask) if dummy_mask is not None else dummy_img\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        # Update final upsample to match split size\n        split_h = CFG['target_size'] // 3\n        split_w = CFG['target_size'] // 4\n        self.final_upsample = nn.Upsample(size=(split_h, split_w))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            # Convert to float32 for interpolation to avoid AMP issues\n            inputs = nn.functional.interpolate(inputs.float(), size=targets.shape[-2:], mode='bilinear').to(inputs.dtype)\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                # Convert to float32 for interpolation to avoid AMP issues\n                preds = nn.functional.interpolate(preds.float(), size=masks.shape[-2:], mode='bilinear').to(preds.dtype)\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef dice_score(y_true, y_pred):\n    \"\"\"Calculate Dice coefficient\"\"\"\n    y_true_f = y_true.flatten()\n    y_pred_f = y_pred.flatten()\n    intersection = np.sum(y_true_f * y_pred_f)\n    return (2. * intersection + 1) / (np.sum(y_true_f) + np.sum(y_pred_f) + 1)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                # Convert to float32 for interpolation to avoid AMP issues\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds).float(), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    # Debug mask values\n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    # Binarize labels properly\n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5, 0.0, 0.0, 0.0\n    \n    # Find optimal threshold from 0.3 to 0.9\n    best_thresh, best_f05 = 0.3, 0\n    thresholds = np.linspace(0.3, 0.9, 13)\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n            \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    f1 = f1_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    fb = fbeta_score(all_labels.flatten(), final_binarized.flatten(), beta=2, zero_division=0)\n    dice_val = dice_score(all_labels, final_binarized)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh, dice_val, f1, fb\n\ndef reconstruct_full_prediction(model, dataset, device):\n    \"\"\"Reconstruct full prediction for visualization\"\"\"\n    model.eval()\n    z, h, w = dataset.volume.shape\n    full_pred = np.zeros((h, w), dtype=np.float32)\n    count = np.zeros((h, w), dtype=np.float32)\n    \n    with torch.no_grad():\n        for i, j in tqdm(dataset.coords, desc=\"Reconstructing full prediction\"):\n            # Get volume patch\n            patch = dataset.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Convert to tensor and predict\n            image_tensor = torch.tensor(patch).permute(2, 0, 1).unsqueeze(0).to(device)\n            \n            with autocast(enabled=CFG['amp']):\n                pred = torch.sigmoid(model(image_tensor)).cpu().numpy()[0, 0]\n            \n            # Ensure the prediction is in the correct format for OpenCV\n            pred = np.clip(pred, 0, 1).astype(np.float32)\n            \n            # Resize prediction back to original patch size\n            pred_resized = cv2.resize(pred, (CFG['target_size'], CFG['target_size']))\n            \n            # Add to full prediction\n            full_pred[i:i+CFG['target_size'], j:j+CFG['target_size']] += pred_resized\n            count[i:i+CFG['target_size'], j:j+CFG['target_size']] += 1\n    \n    # Average overlapping regions\n    full_pred = np.divide(full_pred, count, out=np.zeros_like(full_pred), where=count != 0)\n    return full_pred\n\ndef visualize_comparison(full_input, full_label, full_pred, threshold=0.5, metrics=None):\n    \"\"\"Visualize comparison of full input, label and prediction\"\"\"\n    plt.figure(figsize=(20, 10))\n    \n    # Input (middle slice)\n    plt.subplot(2, 3, 1)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.title(\"Input Volume (Middle Slice)\")\n    plt.axis('off')\n    \n    # Ground truth\n    plt.subplot(2, 3, 2)\n    plt.imshow(full_label, cmap='gray')\n    plt.title(\"Ground Truth\")\n    plt.axis('off')\n    \n    # Prediction\n    plt.subplot(2, 3, 3)\n    plt.imshow(full_pred, cmap='gray')\n    plt.title(f\"Prediction (Threshold: {threshold:.2f})\")\n    plt.axis('off')\n    \n    # Thresholded prediction\n    plt.subplot(2, 3, 4)\n    plt.imshow(full_pred > threshold, cmap='gray')\n    plt.title(\"Thresholded Prediction\")\n    plt.axis('off')\n    \n    # Overlay\n    plt.subplot(2, 3, 5)\n    plt.imshow(full_input[len(full_input)//2], cmap='gray')\n    plt.imshow(full_pred > threshold, cmap='jet', alpha=0.3)\n    plt.title(\"Input with Prediction Overlay\")\n    plt.axis('off')\n    \n    # Metrics text\n    plt.subplot(2, 3, 6)\n    plt.axis('off')\n    if metrics:\n        metrics_text = f\"\"\"\n        Evaluation Metrics:\n        Dice Score: {metrics['dice']:.4f}\n        Precision: {metrics['precision']:.4f}\n        Recall: {metrics['recall']:.4f}\n        F0.5 Score: {metrics['f05']:.4f}\n        F1 Score: {metrics['f1']:.4f}\n        F2 Score: {metrics['f2']:.4f}\n        Optimal Threshold: {threshold:.2f}\n        \"\"\"\n        plt.text(0.1, 0.5, metrics_text, fontsize=14, va='center')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh, dice_val, f1, fb = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - Dice: {dice_val:.4f}, F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, F1: {f1:.4f}, F2: {fb:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nReconstructing full prediction for visualization...\")\n                    full_pred = reconstruct_full_prediction(model, test_dataset, device)\n                    \n                    # Get full input and label\n                    full_input = test_dataset.volume\n                    full_label = test_dataset.mask\n                    \n                    # Visualize comparison\n                    metrics = {\n                        'dice': dice_val,\n                        'precision': prec,\n                        'recall': rec,\n                        'f05': f05,\n                        'f1': f1,\n                        'f2': fb\n                    }\n                    visualize_comparison(full_input, full_label, full_pred, threshold=best_thresh, metrics=metrics)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install monai ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pre-trained Unet\nimport os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom skimage.morphology import binary_opening, disk\nfrom sklearn.model_selection import KFold\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom tqdm import tqdm\nfrom monai.networks.nets import UNet\nfrom monai.losses import DiceLoss, FocalLoss\n#from early_stopping import EarlyStopping\n\n# Configuration\nCFG = {\n    \"train_case\": (\"2\", \"/kaggle/input/vesuvius-challenge/train/2/surface_volume\", \"/kaggle/input/vesuvius-challenge/train/2/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge/train/1/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 8,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10,\n    \"seed\": 42,\n    \"num_workers\": 2,\n    \"device\": torch.device('cuda' if torch.cuda.is_available() else 'cpu'),\n    \"pretrained\": True,\n    \"gradient_accumulation_steps\": 2\n}\n\n# Set seed for reproducibility\ntorch.manual_seed(CFG[\"seed\"])\nnp.random.seed(CFG[\"seed\"])\n\n# Data Loading and Preprocessing\ndef load_volume(surface_volume_path):\n    images = []\n    for i in range(CFG[\"slice_start\"], CFG[\"slice_end\"]):\n        img_path = os.path.join(surface_volume_path, f\"{i:02}.tif\")\n        img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n        if (img.shape[0], img.shape[1]) != (CFG[\"target_size\"], CFG[\"target_size\"]):\n            img = cv2.resize(img, (CFG[\"target_size\"], CFG[\"target_size\"]), interpolation=cv2.INTER_AREA)\n        images.append(img)\n    return np.stack(images, axis=0)\n\ndef load_mask(fragment_id, data_dir):\n    mask_path = os.path.join(data_dir, f\"{fragment_id}/mask.png\")\n    mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    if (mask.shape[0], mask.shape[1]) != (CFG[\"target_size\"], CFG[\"target_size\"]):\n        mask = cv2.resize(mask, (CFG[\"target_size\"], CFG[\"target_size\"]), interpolation=cv2.INTER_NEAREST)\n    return (mask > 0).astype(np.uint8)\n\ndef load_inklabels(inklabels_path):\n    label = cv2.imread(inklabels_path, cv2.IMREAD_GRAYSCALE)\n    if (label.shape[0], label.shape[1]) != (CFG[\"target_size\"], CFG[\"target_size\"]):\n        label = cv2.resize(label, (CFG[\"target_size\"], CFG[\"target_size\"]), interpolation=cv2.INTER_NEAREST)\n    return (label > 0).astype(np.uint8)\n\n# Dataset Class\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume, mask, labels=None, is_train=True):\n        self.volume = volume\n        self.mask = mask\n        self.labels = labels\n        self.is_train = is_train\n        self.patch_size = CFG[\"target_size\"]\n        self.stride = CFG[\"stride\"]\n        \n        # Get valid coordinates\n        self.coords = []\n        for y in range(0, self.volume.shape[1] - self.patch_size + 1, self.stride):\n            for x in range(0, self.volume.shape[2] - self.patch_size + 1, self.stride):\n                if self.mask[y + self.patch_size//2, x + self.patch_size//2] > 0:\n                    if self.labels is None or self.labels[y:y+self.patch_size, x:x+self.patch_size].sum() >= CFG[\"min_ink_pixels\"]:\n                        self.coords.append((y, x))\n    \n    def __len__(self):\n        return len(self.coords)\n    \n    def __getitem__(self, idx):\n        y, x = self.coords[idx]\n        subvolume = self.volume[:, y:y+self.patch_size, x:x+self.patch_size]\n        \n        # Normalize\n        subvolume = subvolume.astype(np.float32) / 255.0\n        \n        # Add augmentations if training\n        if self.is_train:\n            if np.random.rand() > 0.5:\n                subvolume = np.flip(subvolume, axis=1)\n            if np.random.rand() > 0.5:\n                subvolume = np.flip(subvolume, axis=2)\n            if np.random.rand() > 0.5:\n                subvolume = np.rot90(subvolume, k=1, axes=(1,2))\n        \n        subvolume = torch.from_numpy(subvolume.copy()).float()\n        \n        if self.labels is not None:\n            label = self.labels[y:y+self.patch_size, x:x+self.patch_size]\n            label = torch.from_numpy(label).float()\n            return subvolume, label\n        return subvolume\n\n# Model Initialization with Pretrained UNet\ndef create_model():\n    model = UNet(\n        spatial_dims=3,\n        in_channels=1,\n        out_channels=1,\n        channels=(16, 32, 64, 128, 256),\n        strides=(2, 2, 2, 2),\n        num_res_units=2,\n        norm=\"BATCH\"\n    )\n    \n    if CFG[\"pretrained\"]:\n        try:\n            # Load pretrained weights (this is a placeholder - you'll need actual pretrained weights)\n            pretrained_path = \"/kaggle/input/pretrained-3d-unet/model.pth\"\n            if os.path.exists(pretrained_path):\n                model.load_state_dict(torch.load(pretrained_path))\n                print(\"Loaded pretrained weights!\")\n            else:\n                print(\"Pretrained weights not found, training from scratch\")\n        except:\n            print(\"Failed to load pretrained weights, training from scratch\")\n    \n    return model.to(CFG[\"device\"])\n\n# Combined Loss Function\nclass CombinedLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.dice_loss = DiceLoss(sigmoid=True)\n        self.focal_loss = FocalLoss(to_onehot_y=False)\n    \n    def forward(self, preds, targets):\n        dice = self.dice_loss(preds, targets.unsqueeze(1))\n        focal = self.focal_loss(preds, targets.unsqueeze(1))\n        return 0.5 * dice + 0.5 * focal\n\n# Training Function\ndef train_one_epoch(model, loader, optimizer, scheduler, scaler, device):\n    model.train()\n    running_loss = 0.0\n    criterion = CombinedLoss()\n    \n    for batch_idx, (data, target) in enumerate(tqdm(loader, desc=\"Training\")):\n        data, target = data.to(device), target.to(device)\n        \n        with autocast(enabled=CFG[\"amp\"]):\n            output = model(data.unsqueeze(1))  # Add channel dimension\n            loss = criterion(output, target)\n            \n            # Gradient accumulation\n            loss = loss / CFG[\"gradient_accumulation_steps\"]\n        \n        scaler.scale(loss).backward()\n        \n        if (batch_idx + 1) % CFG[\"gradient_accumulation_steps\"] == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n        \n        running_loss += loss.item() * CFG[\"gradient_accumulation_steps\"]\n    \n    if scheduler is not None:\n        scheduler.step()\n    \n    return running_loss / len(loader)\n\n# Validation Function\ndef validate(model, loader, device):\n    model.eval()\n    running_loss = 0.0\n    criterion = CombinedLoss()\n    \n    with torch.no_grad():\n        for data, target in tqdm(loader, desc=\"Validation\"):\n            data, target = data.to(device), target.to(device)\n            output = model(data.unsqueeze(1))\n            loss = criterion(output, target)\n            running_loss += loss.item()\n    \n    return running_loss / len(loader)\n\n# Prediction Function\ndef predict(model, loader, device):\n    model.eval()\n    preds = []\n    \n    with torch.no_grad():\n        for data in tqdm(loader, desc=\"Predicting\"):\n            data = data.to(device)\n            with autocast(enabled=CFG[\"amp\"]):\n                output = model(data.unsqueeze(1))\n                output = torch.sigmoid(output)\n            preds.append(output.cpu())\n    \n    return torch.cat(preds, dim=0)\n\n# Main Training Function\ndef train_and_validate():\n    # Load data\n    print(\"Loading training data...\")\n    train_volume = load_volume(CFG[\"train_case\"][1])\n    train_mask = load_mask(CFG[\"train_case\"][0], os.path.dirname(os.path.dirname(CFG[\"train_case\"][1])))\n    train_labels = load_inklabels(CFG[\"train_case\"][2])\n    \n    print(\"Loading test data...\")\n    test_volume = load_volume(CFG[\"test_case\"][1])\n    test_mask = load_mask(CFG[\"test_case\"][0], os.path.dirname(os.path.dirname(CFG[\"test_case\"][1])))\n    test_labels = load_inklabels(CFG[\"test_case\"][2])\n    \n    # Create datasets\n    train_dataset = VesuviusDataset(train_volume, train_mask, train_labels, is_train=True)\n    test_dataset = VesuviusDataset(test_volume, test_mask, test_labels, is_train=False)\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG[\"train_batch_size\"],\n        shuffle=True,\n        num_workers=CFG[\"num_workers\"],\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG[\"test_batch_size\"],\n        shuffle=False,\n        num_workers=CFG[\"num_workers\"],\n        pin_memory=True\n    )\n    \n    # Initialize model, optimizer, etc.\n    model = create_model()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG[\"lr\"])\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG[\"epochs\"])\n    scaler = GradScaler(enabled=CFG[\"amp\"])\n    early_stopping = EarlyStopping(patience=CFG[\"early_stop_patience\"], verbose=True)\n    \n    best_loss = float('inf')\n    for epoch in range(CFG[\"epochs\"]):\n        print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n        \n        train_loss = train_one_epoch(model, train_loader, optimizer, scheduler, scaler, CFG[\"device\"])\n        val_loss = validate(model, test_loader, CFG[\"device\"])\n        \n        print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n        \n        # Early stopping check\n        early_stopping(val_loss, model)\n        if early_stopping.early_stop:\n            print(\"Early stopping triggered\")\n            break\n        \n        if val_loss < best_loss:\n            best_loss = val_loss\n            torch.save(model.state_dict(), \"best_model.pth\")\n            print(\"Saved best model!\")\n    \n    # Load best model\n    model.load_state_dict(torch.load(\"best_model.pth\"))\n    \n    return model\n\n# Post-processing\ndef apply_post_processing(pred, threshold=0.5, kernel_size=3):\n    pred = (pred > threshold).astype(np.uint8)\n    kernel = disk(kernel_size)\n    pred = binary_opening(pred, kernel)\n    return pred\n\n# Main Execution\nif __name__ == \"__main__\":\n    print(\"Starting training...\")\n    model = train_and_validate()\n    \n    # Example of how to use the trained model for prediction\n    print(\"\\nRunning example prediction...\")\n    example_loader = DataLoader(\n        VesuviusDataset(load_volume(CFG[\"test_case\"][1]), \n        load_mask(CFG[\"test_case\"][0], os.path.dirname(os.path.dirname(CFG[\"test_case\"][1]))), \n        None, \n        is_train=False),\n        batch_size=CFG[\"test_batch_size\"],\n        shuffle=False,\n        num_workers=CFG[\"num_workers\"],\n        pin_memory=True\n    )\n    \n    predictions = predict(model, example_loader, CFG[\"device\"])\n    print(f\"Prediction shape: {predictions.shape}\")\n    \n    # Clean up\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#VNET\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc\n\n# Memory-optimized configuration\nCFG = {\n    \"data_path\": \"/kaggle/input/vesuvius-challenge-ink-detection\",\n    \"train_fold\": \"1\",\n    \"test_fold\": \"2\",\n    \"target_size\": 128,\n    \"num_splits\": 2,\n    \"slice_start\": 16,\n    \"slice_end\": 32,\n    \"train_batch_size\": 4,\n    \"test_batch_size\": 2,\n    \"stride\": 96,\n    \"epochs\": 10,\n    \"lr\": 2e-4,\n    \"amp\": True,\n    \"min_ink_threshold\": 0.1\n}\n\nclass InkDetectionDataset(Dataset):\n    def __init__(self, fold, is_train=True):\n        self.is_train = is_train\n        self.volume_dir = f\"{CFG['data_path']}/train/{fold}/surface_volume\"\n        self.mask_path = f\"{CFG['data_path']}/train/{fold}/inklabels.png\"\n        self.volume_paths = sorted(glob(f\"{self.volume_dir}/*.tif\"))\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.3),\n            A.VerticalFlip(p=0.3),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_mask(self):\n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        return (mask / 255.0).astype(np.float32) if mask is not None else None\n\n    def _generate_coords(self):\n        h, w = cv2.imread(self.volume_paths[0], cv2.IMREAD_GRAYSCALE).shape\n        coords = []\n        for i in range(0, h - CFG['target_size'] + 1, CFG['stride']):\n            for j in range(0, w - CFG['target_size'] + 1, CFG['stride']):\n                if self.mask is None or np.mean(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_threshold']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        orig_idx = idx // CFG['num_splits'] if self.is_train else idx\n        i, j = self.coords[orig_idx]\n        \n        # Load slices on-demand\n        patch = []\n        for p in self.volume_paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            patch.append(cv2.resize(img, (CFG['target_size'], CFG['target_size'])))\n        patch = np.stack(patch, axis=-1).astype(np.float32) / 255.0  # [H,W,C]\n        \n        mask = None\n        if self.mask is not None:\n            mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n        \n        if self.is_train and mask is not None:\n            transformed = self.transform(image=patch, mask=mask)\n            return transformed['image'], transformed['mask']\n        return torch.tensor(patch).permute(2, 0, 1).float(), torch.tensor(mask).float() if mask is not None else None\n\nclass LiteVNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = nn.Sequential(\n            nn.Conv3d(in_channels, 16, kernel_size=3, padding=1),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        self.enc2 = nn.Sequential(\n            nn.Conv3d(16, 32, kernel_size=3, padding=1),\n            nn.BatchNorm3d(32),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        # Decoder\n        self.dec1 = nn.Sequential(\n            nn.ConvTranspose3d(32, 16, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU()\n        )\n        \n        # Remove final sigmoid and output logits directly\n        self.dec2 = nn.ConvTranspose3d(16, 1, kernel_size=(1,2,2), stride=(1,2,2))\n    \n    def forward(self, x):\n        x = x.unsqueeze(2)  # Add depth dim [B,C,1,H,W]\n        x = self.enc1(x)\n        x = self.enc2(x)\n        x = self.dec1(x)\n        x = self.dec2(x)\n        return x.squeeze(1).squeeze(1)  # [B,H,W]\n\ndef train_and_validate():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Initialize model and optimizer\n    model = LiteVNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\n    scaler = GradScaler(enabled=CFG['amp'])\n    \n    # Data loaders\n    train_dataset = InkDetectionDataset(CFG['train_fold'], is_train=True)\n    test_dataset = InkDetectionDataset(CFG['test_fold'], is_train=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG['train_batch_size'],\n        shuffle=True,\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG['test_batch_size'],\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    print(f\"Training samples: {len(train_dataset)}\")\n    print(f\"Validation samples: {len(test_dataset)}\")\n    \n    best_score = 0\n    for epoch in range(CFG['epochs']):\n        # Training phase\n        model.train()\n        train_loss = 0\n        for x, y in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{CFG['epochs']}\"):\n            x, y = x.to(device), y.to(device)\n            \n            optimizer.zero_grad()\n            with autocast(enabled=CFG['amp']):\n                pred = model(x)\n                # Use BCEWithLogitsLoss instead of manual sigmoid + BCE\n                loss = F.binary_cross_entropy_with_logits(pred, y)\n            \n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            train_loss += loss.item()\n        \n        # Validation phase\n        model.eval()\n        preds, targets = [], []\n        with torch.no_grad():\n            for x, y in test_loader:\n                x = x.to(device)\n                # Apply sigmoid only during validation\n                pred = torch.sigmoid(model(x)).cpu()\n                preds.append(pred)\n                targets.append(y.cpu())\n        \n        preds = torch.cat(preds).numpy()\n        targets = torch.cat(targets).numpy()\n        f05 = fbeta_score(targets.flatten(), (preds > 0.5).flatten(), beta=0.5)\n        \n        print(f\"Epoch {epoch+1} | Train Loss: {train_loss/len(train_loader):.4f} | F0.5: {f05:.4f}\")\n        \n        # Save best model\n        if f05 > best_score:\n            best_score = f05\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"Saved new best model with F0.5: {best_score:.4f}\")\n            \n            # Visualize predictions\n            visualize_results(x[:4].cpu(), preds[:4], targets[:4])\n        \n        # Clean up\n        torch.cuda.empty_cache()\n        gc.collect()\n\ndef visualize_results(inputs, preds, targets, num_samples=4):\n    plt.figure(figsize=(15, 5*num_samples))\n    for i in range(min(num_samples, len(inputs))):\n        # Input middle slice\n        plt.subplot(num_samples, 3, i*3+1)\n        plt.imshow(inputs[i][len(inputs[i])//2], cmap='gray')\n        plt.title(f\"Input {i+1}\")\n        plt.axis('off')\n        \n        # Ground truth\n        plt.subplot(num_samples, 3, i*3+2)\n        plt.imshow(targets[i], cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        # Prediction\n        plt.subplot(num_samples, 3, i*3+3)\n        plt.imshow(preds[i] > 0.5, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\nif __name__ == '__main__':\n    train_and_validate()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"jh","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc\n\n# Memory-optimized configuration\nCFG = {\n    \"data_path\": \"/kaggle/input/vesuvius-challenge-ink-detection\",\n    \"train_fold\": \"1\",\n    \"test_fold\": \"2\",\n    \"target_size\": 128,\n    \"num_splits\": 2,\n    \"slice_start\": 16,\n    \"slice_end\": 32,\n    \"train_batch_size\": 4,\n    \"test_batch_size\": 2,\n    \"stride\": 96,\n    \"epochs\": 10,\n    \"lr\": 2e-4,\n    \"amp\": True,\n    \"min_ink_threshold\": 0.1\n}\n\nclass InkDetectionDataset(Dataset):\n    def __init__(self, fold, is_train=True):\n        self.is_train = is_train\n        self.volume_dir = f\"{CFG['data_path']}/train/{fold}/surface_volume\"\n        self.mask_path = f\"{CFG['data_path']}/train/{fold}/inklabels.png\"\n        self.volume_paths = sorted(glob(f\"{self.volume_dir}/*.tif\"))\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.3),\n            A.VerticalFlip(p=0.3),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_mask(self):\n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        return (mask / 255.0).astype(np.float32) if mask is not None else None\n\n    def _generate_coords(self):\n        h, w = cv2.imread(self.volume_paths[0], cv2.IMREAD_GRAYSCALE).shape\n        coords = []\n        for i in range(0, h - CFG['target_size'] + 1, CFG['stride']):\n            for j in range(0, w - CFG['target_size'] + 1, CFG['stride']):\n                if self.mask is None or np.mean(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']) > CFG['min_ink_threshold']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        orig_idx = idx // CFG['num_splits'] if self.is_train else idx\n        i, j = self.coords[orig_idx]\n        \n        # Load slices on-demand\n        patch = []\n        for p in self.volume_paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            patch.append(cv2.resize(img, (CFG['target_size'], CFG['target_size'])))\n        patch = np.stack(patch, axis=-1).astype(np.float32) / 255.0  # [H,W,C]\n        \n        mask = None\n        if self.mask is not None:\n            mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n        \n        if self.is_train and mask is not None:\n            transformed = self.transform(image=patch, mask=mask)\n            return transformed['image'], transformed['mask']\n        return torch.tensor(patch).permute(2, 0, 1).float(), torch.tensor(mask).float() if mask is not None else None\n\nclass LiteVNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = nn.Sequential(\n            nn.Conv3d(in_channels, 16, kernel_size=3, padding=1),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        self.enc2 = nn.Sequential(\n            nn.Conv3d(16, 32, kernel_size=3, padding=1),\n            nn.BatchNorm3d(32),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        \n        # Decoder\n        self.dec1 = nn.Sequential(\n            nn.ConvTranspose3d(32, 16, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU()\n        )\n        \n        self.dec2 = nn.Sequential(\n            nn.ConvTranspose3d(16, 1, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        x = x.unsqueeze(2)  # Add depth dim [B,C,1,H,W]\n        x = self.enc1(x)\n        x = self.enc2(x)\n        x = self.dec1(x)\n        x = self.dec2(x)\n        return x.squeeze(1).squeeze(1)  # [B,H,W]\n\ndef train_and_validate():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Initialize model and optimizer\n    model = LiteVNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\n    scaler = GradScaler(enabled=CFG['amp'])\n    \n    # Data loaders\n    train_dataset = InkDetectionDataset(CFG['train_fold'], is_train=True)\n    test_dataset = InkDetectionDataset(CFG['test_fold'], is_train=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG['train_batch_size'],\n        shuffle=True,\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG['test_batch_size'],\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    print(f\"Training samples: {len(train_dataset)}\")\n    print(f\"Validation samples: {len(test_dataset)}\")\n    \n    best_score = 0\n    for epoch in range(CFG['epochs']):\n        # Training phase\n        model.train()\n        train_loss = 0\n        for x, y in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{CFG['epochs']}\"):\n            x, y = x.to(device), y.to(device)\n            \n            optimizer.zero_grad()\n            with autocast(enabled=CFG['amp']):\n                pred = model(x)\n                loss = F.binary_cross_entropy(pred, y)\n            \n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            train_loss += loss.item()\n        \n        # Validation phase\n        model.eval()\n        preds, targets = [], []\n        with torch.no_grad():\n            for x, y in test_loader:\n                x = x.to(device)\n                pred = torch.sigmoid(model(x)).cpu()\n                preds.append(pred)\n                targets.append(y.cpu())\n        \n        preds = torch.cat(preds).numpy()\n        targets = torch.cat(targets).numpy()\n        f05 = fbeta_score(targets.flatten(), (preds > 0.5).flatten(), beta=0.5)\n        \n        print(f\"Epoch {epoch+1} | Train Loss: {train_loss/len(train_loader):.4f} | F0.5: {f05:.4f}\")\n        \n        # Save best model\n        if f05 > best_score:\n            best_score = f05\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"Saved new best model with F0.5: {best_score:.4f}\")\n            \n            # Visualize predictions\n            visualize_results(x[:4].cpu(), preds[:4], targets[:4])\n        \n        # Clean up\n        torch.cuda.empty_cache()\n        gc.collect()\n\ndef visualize_results(inputs, preds, targets, num_samples=4):\n    plt.figure(figsize=(15, 5*num_samples))\n    for i in range(min(num_samples, len(inputs))):\n        # Input middle slice\n        plt.subplot(num_samples, 3, i*3+1)\n        plt.imshow(inputs[i][len(inputs[i])//2], cmap='gray')\n        plt.title(f\"Input {i+1}\")\n        plt.axis('off')\n        \n        # Ground truth\n        plt.subplot(num_samples, 3, i*3+2)\n        plt.imshow(targets[i], cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        # Prediction\n        plt.subplot(num_samples, 3, i*3+3)\n        plt.imshow(preds[i] > 0.5, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\nif __name__ == '__main__':\n    train_and_validate()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc\n\n# Memory-optimized configuration\nCFG = {\n    \"data_path\": \"/kaggle/input/vesuvius-challenge-ink-detection\",\n    \"train_fold\": \"1\",\n    \"test_fold\": \"2\",\n    \"target_size\": 128,\n    \"num_splits\": 2,\n    \"slice_start\": 16,\n    \"slice_end\": 32,\n    \"train_batch_size\": 4,\n    \"test_batch_size\": 2,\n    \"stride\": 96,\n    \"epochs\": 10,\n    \"lr\": 2e-4,\n    \"amp\": True,\n    \"min_ink_threshold\": 0.1\n}\n\nclass InkDetectionDataset(Dataset):\n    def __init__(self, fold, is_train=True):\n        self.is_train = is_train\n        self.volume_dir = f\"{CFG['data_path']}/train/{fold}/surface_volume\"\n        self.mask_path = f\"{CFG['data_path']}/train/{fold}/inklabels.png\"\n        self.volume_paths = sorted(glob(f\"{self.volume_dir}/*.tif\"))\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.3),\n            A.VerticalFlip(p=0.3),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_mask(self):\n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        return (mask / 255.0).astype(np.float32) if mask is not None else None\n\n    def _generate_coords(self):\n        h, w = cv2.imread(self.volume_paths[0], cv2.IMREAD_GRAYSCALE).shape\n        coords = []\n        for i in range(0, h - CFG['target_size'] + 1, CFG['stride']):\n            for j in range(0, w - CFG['target_size'] + 1, CFG['stride']):\n                if self.mask is None or np.mean(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_threshold']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        orig_idx = idx // CFG['num_splits'] if self.is_train else idx\n        i, j = self.coords[orig_idx]\n        \n        # Load slices on-demand\n        patch = []\n        for p in self.volume_paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            patch.append(cv2.resize(img, (CFG['target_size'], CFG['target_size'])))\n        patch = np.stack(patch, axis=-1).astype(np.float32) / 255.0  # [H,W,C]\n        \n        mask = None\n        if self.mask is not None:\n            mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n        \n        if self.is_train and mask is not None:\n            transformed = self.transform(image=patch, mask=mask)\n            return transformed['image'], transformed['mask']\n        return torch.tensor(patch).permute(2, 0, 1).float(), torch.tensor(mask).float() if mask is not None else None\n\nclass LiteVNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = nn.Sequential(\n            nn.Conv3d(in_channels, 16, kernel_size=3, padding=1),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        self.enc2 = nn.Sequential(\n            nn.Conv3d(16, 32, kernel_size=3, padding=1),\n            nn.BatchNorm3d(32),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        # Decoder\n        self.dec1 = nn.Sequential(\n            nn.ConvTranspose3d(32, 16, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU()\n        )\n        \n        self.dec2 = nn.Sequential(\n            nn.ConvTranspose3d(16, 1, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        x = x.unsqueeze(2)  # Add depth dim [B,C,1,H,W]\n        x = self.enc1(x)\n        x = self.enc2(x)\n        x = self.dec1(x)\n        x = self.dec2(x)\n        return x.squeeze(1).squeeze(1)  # [B,H,W]\n\ndef train_and_validate():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Initialize model and optimizer\n    model = LiteVNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\n    scaler = GradScaler(enabled=CFG['amp'])\n    \n    # Data loaders\n    train_dataset = InkDetectionDataset(CFG['train_fold'], is_train=True)\n    test_dataset = InkDetectionDataset(CFG['test_fold'], is_train=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG['train_batch_size'],\n        shuffle=True,\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG['test_batch_size'],\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    print(f\"Training samples: {len(train_dataset)}\")\n    print(f\"Validation samples: {len(test_dataset)}\")\n    \n    best_score = 0\n    for epoch in range(CFG['epochs']):\n        # Training phase\n        model.train()\n        train_loss = 0\n        for x, y in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{CFG['epochs']}\"):\n            x, y = x.to(device), y.to(device)\n            \n            optimizer.zero_grad()\n            with autocast(enabled=CFG['amp']):\n                pred = model(x)\n                loss = F.binary_cross_entropy(pred, y)\n            \n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            train_loss += loss.item()\n        \n        # Validation phase\n        model.eval()\n        preds, targets = [], []\n        with torch.no_grad():\n            for x, y in test_loader:\n                x = x.to(device)\n                pred = torch.sigmoid(model(x)).cpu()\n                preds.append(pred)\n                targets.append(y.cpu())\n        \n        preds = torch.cat(preds).numpy()\n        targets = torch.cat(targets).numpy()\n        f05 = fbeta_score(targets.flatten(), (preds > 0.5).flatten(), beta=0.5)\n        \n        print(f\"Epoch {epoch+1} | Train Loss: {train_loss/len(train_loader):.4f} | F0.5: {f05:.4f}\")\n        \n        # Save best model\n        if f05 > best_score:\n            best_score = f05\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"Saved new best model with F0.5: {best_score:.4f}\")\n            \n            # Visualize predictions\n            visualize_results(x[:4].cpu(), preds[:4], targets[:4])\n        \n        # Clean up\n        torch.cuda.empty_cache()\n        gc.collect()\n\ndef visualize_results(inputs, preds, targets, num_samples=4):\n    plt.figure(figsize=(15, 5*num_samples))\n    for i in range(min(num_samples, len(inputs))):\n        # Input middle slice\n        plt.subplot(num_samples, 3, i*3+1)\n        plt.imshow(inputs[i][len(inputs[i])//2], cmap='gray')\n        plt.title(f\"Input {i+1}\")\n        plt.axis('off')\n        \n        # Ground truth\n        plt.subplot(num_samples, 3, i*3+2)\n        plt.imshow(targets[i], cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        # Prediction\n        plt.subplot(num_samples, 3, i*3+3)\n        plt.imshow(preds[i] > 0.5, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\nif __name__ == '__main__':\n    train_and_validate()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc\n\n# Memory-optimized configuration\nCFG = {\n    \"data_path\": \"/kaggle/input/vesuvius-challenge-ink-detection\",\n    \"train_fold\": \"1\",\n    \"test_fold\": \"2\",\n    \"target_size\": 128,\n    \"num_splits\": 2,          # Reduced from original 8\n    \"slice_start\": 16,\n    \"slice_end\": 32,          # Reduced from 40\n    \"train_batch_size\": 4,    # Reduced from 8\n    \"test_batch_size\": 2,     # Reduced from 4\n    \"stride\": 96,             # Increased from 64\n    \"epochs\": 10,             # Reduced from 20\n    \"lr\": 2e-4,               # Reduced learning rate\n    \"amp\": True,              # Mixed precision training\n    \"min_ink_threshold\": 0.1  # Lower ink threshold\n}\n\nclass InkDetectionDataset(Dataset):\n    def __init__(self, fold, is_train=True):\n        self.is_train = is_train\n        self.volume_dir = f\"{CFG['data_path']}/train/{fold}/surface_volume\"\n        self.mask_path = f\"{CFG['data_path']}/train/{fold}/inklabels.png\"\n        self.volume_paths = sorted(glob(f\"{self.volume_dir}/*.tif\"))\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.3),\n            A.VerticalFlip(p=0.3),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_mask(self):\n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        return (mask / 255.0).astype(np.float32) if mask is not None else None\n\n    def _generate_coords(self):\n        h, w = cv2.imread(self.volume_paths[0], cv2.IMREAD_GRAYSCALE).shape\n        coords = []\n        for i in range(0, h - CFG['target_size'] + 1, CFG['stride']):\n            for j in range(0, w - CFG['target_size'] + 1, CFG['stride']):\n                if self.mask is None or np.mean(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_threshold']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        orig_idx = idx // CFG['num_splits'] if self.is_train else idx\n        i, j = self.coords[orig_idx]\n        \n        # Load slices on-demand\n        patch = []\n        for p in self.volume_paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            patch.append(cv2.resize(img, (CFG['target_size'], CFG['target_size'])))\n        patch = np.stack(patch, axis=-1).astype(np.float32) / 255.0  # [H,W,C]\n        \n        mask = None\n        if self.mask is not None:\n            mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n        \n        if self.is_train and mask is not None:\n            transformed = self.transform(image=patch, mask=mask)\n            return transformed['image'], transformed['mask']\n        return torch.tensor(patch).permute(2, 0, 1).float(), torch.tensor(mask).float() if mask is not None else None\n\nclass LiteVNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = nn.Sequential(\n            nn.Conv3d(in_channels, 16, kernel_size=3, padding=1),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        self.enc2 = nn.Sequential(\n            nn.Conv3d(16, 32, kernel_size=3, padding=1),\n            nn.BatchNorm3d(32),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2))\n        )\n        \n        # Decoder\n        self.dec1 = nn.Sequential(\n            nn.ConvTranspose3d(32, 16, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU()\n        )\n        \n        self.dec2 = nn.Sequential(\n            nn.ConvTranspose3d(16, 1, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        x = x.unsqueeze(2)  # Add depth dim [B,C,1,H,W]\n        x = self.enc1(x)\n        x = self.enc2(x)\n        x = self.dec1(x)\n        x = self.dec2(x)\n        return x.squeeze(1).squeeze(1)  # [B,H,W]\n\ndef train_and_validate():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Initialize model and optimizer\n    model = LiteVNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\n    scaler = GradScaler(enabled=CFG['amp'])\n    \n    # Data loaders\n    train_dataset = InkDetectionDataset(CFG['train_fold'], is_train=True)\n    test_dataset = InkDetectionDataset(CFG['test_fold'], is_train=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG['train_batch_size'],\n        shuffle=True,\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG['test_batch_size'],\n        pin_memory=True,\n        num_workers=2\n    )\n    \n    print(f\"Training samples: {len(train_dataset)}\")\n    print(f\"Validation samples: {len(test_dataset)}\")\n    \n    best_score = 0\n    for epoch in range(CFG['epochs']):\n        # Training phase\n        model.train()\n        train_loss = 0\n        for x, y in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{CFG['epochs']}\"):\n            x, y = x.to(device), y.to(device)\n            \n            optimizer.zero_grad()\n            with autocast(enabled=CFG['amp']):\n                pred = model(x)\n                loss = F.binary_cross_entropy(pred, y)\n            \n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            train_loss += loss.item()\n        \n        # Validation phase\n        model.eval()\n        preds, targets = [], []\n        with torch.no_grad():\n            for x, y in test_loader:\n                x = x.to(device)\n                pred = torch.sigmoid(model(x)).cpu()\n                preds.append(pred)\n                targets.append(y.cpu())\n        \n        preds = torch.cat(preds).numpy()\n        targets = torch.cat(targets).numpy()\n        f05 = fbeta_score(targets.flatten(), (preds > 0.5).flatten(), beta=0.5)\n        \n        print(f\"Epoch {epoch+1} | Train Loss: {train_loss/len(train_loader):.4f} | F0.5: {f05:.4f}\")\n        \n        # Save best model\n        if f05 > best_score:\n            best_score = f05\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"Saved new best model with F0.5: {best_score:.4f}\")\n            \n            # Visualize predictions\n            visualize_results(x[:4].cpu(), preds[:4], targets[:4])\n        \n        # Clean up\n        torch.cuda.empty_cache()\n        gc.collect()\n\ndef visualize_results(inputs, preds, targets, num_samples=4):\n    plt.figure(figsize=(15, 5*num_samples))\n    for i in range(min(num_samples, len(inputs))):\n        # Input middle slice\n        plt.subplot(num_samples, 3, i*3+1)\n        plt.imshow(inputs[i][len(inputs[i])//2], cmap='gray')\n        plt.title(f\"Input {i+1}\")\n        plt.axis('off')\n        \n        # Ground truth\n        plt.subplot(num_samples, 3, i*3+2)\n        plt.imshow(targets[i], cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        # Prediction\n        plt.subplot(num_samples, 3, i*3+3)\n        plt.imshow(preds[i] > 0.5, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\nif __name__ == '__main__':\n    train_and_validate()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc\n\n# Configuration with memory optimizations\nCFG = {\n    \"train_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"test_case\": (\"2\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 4,  # Reduced from 8 to save memory\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 4,  # Reduced from 8\n    \"test_batch_size\": 2,   # Reduced from 4\n    \"stride\": 96,  # Increased from 64 to reduce samples\n    \"epochs\": 15,  # Reduced from 20\n    \"lr\": 2e-4,    # Reduced from 3e-4\n    \"amp\": True,\n    \"min_ink_pixels\": 20\n}\n\n# Memory-optimized dataset\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume_paths = sorted(glob(os.path.join(volume_dir, '*.tif')))\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_mask(self):\n        if self.mask_path is None:\n            return None\n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        return mask.astype(np.float32) / 255.0 if mask is not None else None\n\n    def _generate_coords(self):\n        sample_img = cv2.imread(self.volume_paths[0], cv2.IMREAD_GRAYSCALE)\n        h, w = sample_img.shape\n        coords = []\n        for i in range(0, h - CFG['target_size'] + 1, CFG['stride']):\n            for j in range(0, w - CFG['target_size'] + 1, CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        orig_idx = idx // CFG['num_splits'] if self.is_train else idx\n        i, j = self.coords[orig_idx]\n        \n        # Load slices on-the-fly\n        patch = []\n        for p in self.volume_paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            patch.append(cv2.resize(img, (CFG['target_size'], CFG['target_size'])))\n        patch = np.stack(patch).astype(np.float32) / 255.0\n        patch = np.transpose(patch, (1, 2, 0))  # [H, W, C]\n        \n        mask = None\n        if self.mask is not None:\n            mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n        \n        if self.is_train and mask is not None:\n            transformed = self.transform(image=patch, mask=mask)\n            return transformed['image'], transformed['mask']\n        return torch.tensor(patch).permute(2, 0, 1).float(), torch.tensor(mask).float() if mask is not None else None\n\n# Simplified VNet with memory optimizations\nclass VNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super().__init__()\n        self.encoder = nn.Sequential(\n            nn.Conv3d(in_channels, 16, kernel_size=3, padding=1),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2)),\n            \n            nn.Conv3d(16, 32, kernel_size=3, padding=1),\n            nn.BatchNorm3d(32),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((1,2,2)),\n        )\n        \n        self.decoder = nn.Sequential(\n            nn.ConvTranspose3d(32, 16, kernel_size=(1,2,2), stride=(1,2,2)),\n            nn.BatchNorm3d(16),\n            nn.LeakyReLU(),\n            \n            nn.Conv3d(16, 1, kernel_size=3, padding=1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        x = x.unsqueeze(2)  # [B,C,1,H,W]\n        x = self.encoder(x)\n        x = self.decoder(x)\n        return x.squeeze(2).squeeze(1)  # [B,H,W]\n\n# Training utilities\ndef train_model():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    model = VNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\n    scaler = GradScaler(enabled=CFG['amp'])\n    \n    train_set = VesuviusDataset(*CFG['train_case'][1:], is_train=True)\n    test_set = VesuviusDataset(*CFG['test_case'][1:], is_train=False)\n    \n    train_loader = DataLoader(train_set, batch_size=CFG['train_batch_size'], shuffle=True, pin_memory=True)\n    test_loader = DataLoader(test_set, batch_size=CFG['test_batch_size'], pin_memory=True)\n    \n    best_score = 0\n    for epoch in range(CFG['epochs']):\n        model.train()\n        for x, y in tqdm(train_loader, desc=f\"Epoch {epoch+1}\"):\n            x, y = x.to(device), y.to(device)\n            \n            optimizer.zero_grad()\n            with autocast(enabled=CFG['amp']):\n                pred = model(x)\n                loss = F.binary_cross_entropy(pred, y)\n            \n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        \n        # Validation\n        model.eval()\n        preds, targets = [], []\n        with torch.no_grad():\n            for x, y in test_loader:\n                x = x.to(device)\n                preds.append(torch.sigmoid(model(x)).cpu())\n                targets.append(y.cpu())\n        \n        preds = torch.cat(preds).numpy()\n        targets = torch.cat(targets).numpy()\n        score = fbeta_score(targets.flatten(), (preds > 0.5).flatten(), beta=0.5)\n        \n        if score > best_score:\n            best_score = score\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"New best score: {score:.4f}\")\n        \n        torch.cuda.empty_cache()\n        gc.collect()\n\nif __name__ == '__main__':\n    train_model()\n    ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#########************* 3D VNet\nimport os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"test_case\": (\"2\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 8,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_volume(self):\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            volume.append(img)\n        return np.stack(volume).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        if self.mask_path is None:\n            return None\n            \n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is None:\n            raise ValueError(f\"Failed to load mask {self.mask_path}\")\n        return mask.astype(np.float32) / 255.0\n\n    def _generate_coords(self):\n        z, h, w = self.volume.shape\n        coords = []\n        target_size = CFG['target_size']\n        \n        for i in range(0, h - target_size + 1, CFG['stride']):\n            for j in range(0, w - target_size + 1, CFG['stride']):\n                if i + target_size > h or j + target_size > w:\n                    continue\n                    \n                if self.mask is None or np.sum(self.mask[i:i+target_size, j:j+target_size]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            orig_idx = idx // CFG['num_splits']\n            split_idx = idx % CFG['num_splits']\n        else:\n            orig_idx = idx\n            split_idx = None\n\n        i, j = self.coords[orig_idx]\n        target_size = CFG['target_size']\n        \n        try:\n            # Extract patch\n            patch = self.volume[:, i:i+target_size, j:j+target_size]\n            patch = np.stack([cv2.resize(slice_img, (target_size, target_size)) \n                            for slice_img in patch])\n            patch = np.transpose(patch, (1, 2, 0))  # [H, W, C]\n            \n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+target_size, j:j+target_size]\n                mask = cv2.resize(mask, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            \n            if split_idx is not None and mask is not None:\n                # 8-way splitting logic\n                h, w = patch.shape[:2]\n                h_half, w_half = h//2, w//2\n                h_quarter, w_quarter = h//4, w//4\n                \n                if split_idx == 0:  # Top-left\n                    patch = patch[:h_half, :w_half]\n                    mask = mask[:h_half, :w_half]\n                elif split_idx == 1:  # Top-right\n                    patch = patch[:h_half, w_half:]\n                    mask = mask[:h_half, w_half:]\n                elif split_idx == 2:  # Bottom-left\n                    patch = patch[h_half:, :w_half]\n                    mask = mask[h_half:, :w_half]\n                elif split_idx == 3:  # Bottom-right\n                    patch = patch[h_half:, w_half:]\n                    mask = mask[h_half:, w_half:]\n                elif split_idx == 4:  # Top-center\n                    patch = patch[:h_half, w_quarter:w_quarter*3]\n                    mask = mask[:h_half, w_quarter:w_quarter*3]\n                elif split_idx == 5:  # Bottom-center\n                    patch = patch[h_half:, w_quarter:w_quarter*3]\n                    mask = mask[h_half:, w_quarter:w_quarter*3]\n                elif split_idx == 6:  # Middle-left\n                    patch = patch[h_quarter:h_quarter*3, :w_half]\n                    mask = mask[h_quarter:h_quarter*3, :w_half]\n                elif split_idx == 7:  # Middle-right\n                    patch = patch[h_quarter:h_quarter*3, w_half:]\n                    mask = mask[h_quarter:h_quarter*3, w_half:]\n                \n                # Resize back to target size\n                patch = cv2.resize(patch, (target_size, target_size))\n                mask = cv2.resize(mask, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            \n            if self.is_train and mask is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                image = transformed['image']  # [C, H, W]\n                mask = transformed['mask']    # [H, W]\n            else:\n                image = torch.tensor(patch).permute(2, 0, 1).float()  # [C, H, W]\n                if mask is not None:\n                    mask = torch.tensor(mask).float()  # [H, W]\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error processing patch at ({i}, {j}): {str(e)}\")\n            image = torch.zeros((CFG['slice_end']-CFG['slice_start'], target_size, target_size), dtype=torch.float32)\n            mask = torch.zeros((target_size, target_size), dtype=torch.float32) if self.mask is not None else None\n            return (image, mask) if mask is not None else image\n\nclass VNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super(VNet, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = 1\n        \n        # Encoder\n        self.encoder1 = self._make_encoder_block(in_channels, 16)\n        self.encoder2 = self._make_encoder_block(16, 32, stride=(1,2,2))\n        self.encoder3 = self._make_encoder_block(32, 64, stride=(1,2,2))\n        self.encoder4 = self._make_encoder_block(64, 128, stride=(1,2,2))\n        \n        # Decoder\n        self.decoder4 = self._make_decoder_block(128, 64, scale_factor=(1,2,2))\n        self.decoder3 = self._make_decoder_block(64, 32, scale_factor=(1,2,2))\n        self.decoder2 = self._make_decoder_block(32, 16, scale_factor=(1,2,2))\n        \n        # Final projection\n        self.final_proj = nn.Sequential(\n            nn.Conv3d(16, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n    def _make_encoder_block(self, in_channels, out_channels, stride=1):\n        return nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True),\n            nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True)\n        )\n    \n    def _make_decoder_block(self, in_channels, out_channels, scale_factor):\n        return nn.Sequential(\n            nn.Upsample(scale_factor=scale_factor, mode='trilinear', align_corners=True),\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Input: [B, C, H, W]\n        x = x.unsqueeze(2)  # Add depth dim: [B, C, 1, H, W]\n        \n        # Encoder\n        enc1 = self.encoder1(x)        # [B, 16, 1, H, W]\n        enc2 = self.encoder2(enc1)     # [B, 32, 1, H/2, W/2]\n        enc3 = self.encoder3(enc2)     # [B, 64, 1, H/4, W/4]\n        enc4 = self.encoder4(enc3)     # [B, 128, 1, H/8, W/8]\n        \n        # Decoder with skip connections\n        dec4 = self.decoder4(enc4) + enc3\n        dec3 = self.decoder3(dec4) + enc2\n        dec2 = self.decoder2(dec3) + enc1\n        \n        # Final output\n        out = self.final_proj(dec2)    # [B, 1, 1, H, W]\n        out = out.squeeze(2)           # [B, 1, H, W]\n        \n        return out\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        targets = targets.unsqueeze(1)  # [B, H, W] -> [B, 1, H, W]\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = F.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs = imgs.to(device)          # [B, C, H, W]\n        masks = masks.to(device)        # [B, H, W]\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)         # [B, 1, H, W]\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        \n    return total_loss / len(loader)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.cpu().numpy()  # [B, H, W]\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()  # [B, 1, H, W]\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            \n    all_preds = np.concatenate(all_preds)  # [N, 1, H, W]\n    all_labels = np.concatenate(all_labels)  # [N, H, W]\n    \n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    thresholds = [0.41, 0.42, 0.43, 0.44, 0.45, 0.46, 0.47, 0.48, 0.49, \n                  0.5, 0.51, 0.52, 0.53, 0.54, 0.55, 0.56, 0.57, 0.58, 0.3, 0.2]\n    \n    best_thresh, best_f05 = 0.3, 0\n    threshold_results = []\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n    \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh\n\ndef visualize_predictions(inputs, preds, labels, threshold=0.5, num_samples=5):\n    # Convert to numpy arrays if they're tensors\n    if torch.is_tensor(inputs):\n        inputs = inputs.cpu().numpy()\n    if torch.is_tensor(preds):\n        preds = preds.cpu().numpy()\n    if torch.is_tensor(labels):\n        labels = labels.cpu().numpy()\n    \n    # Ensure proper shapes\n    inputs = np.array(inputs)  # [N, C, H, W]\n    preds = np.array(preds)    # [N, 1, H, W]\n    labels = np.array(labels)  # [N, H, W]\n    \n    # Select random samples\n    indices = np.random.choice(len(inputs), min(num_samples, len(inputs)), replace=False)\n    \n    plt.figure(figsize=(20, 4*num_samples))\n    for i, idx in enumerate(indices):\n        # Get middle slice from input volume\n        input_img = inputs[idx][len(inputs[idx])//2]  # [H, W]\n        \n        # Get prediction and label\n        pred = preds[idx][0] > threshold  # [H, W]\n        label = labels[idx]               # [H, W]\n        \n        # Plotting\n        plt.subplot(num_samples, 3, i*3 + 1)\n        plt.imshow(input_img, cmap='gray')\n        plt.title(f\"Input (Sample {idx})\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 2)\n        plt.imshow(label, cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 3)\n        plt.imshow(pred, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        model = VNet().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nVisualizing best predictions...\")\n                    sample_inputs = []\n                    sample_masks = []\n                    for batch in test_loader:\n                        sample_inputs.extend(batch[0].cpu().numpy())\n                        sample_masks.extend(batch[1].cpu().numpy())\n                        if len(sample_inputs) >= 20:\n                            break\n                    visualize_predictions(\n                        sample_inputs, \n                        all_preds[:len(sample_inputs)], \n                        all_labels[:len(sample_inputs)], \n                        threshold=best_thresh\n                    )\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3Dvnet\nimport os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"test_case\": (\"2\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 8,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_volume(self):\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            volume.append(img)\n        return np.stack(volume).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        if self.mask_path is None:\n            return None\n            \n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is None:\n            raise ValueError(f\"Failed to load mask {self.mask_path}\")\n        return mask.astype(np.float32) / 255.0\n\n    def _generate_coords(self):\n        z, h, w = self.volume.shape\n        coords = []\n        target_size = CFG['target_size']\n        \n        for i in range(0, h - target_size + 1, CFG['stride']):\n            for j in range(0, w - target_size + 1, CFG['stride']):\n                if i + target_size > h or j + target_size > w:\n                    continue\n                    \n                if self.mask is None or np.sum(self.mask[i:i+target_size, j:j+target_size]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            orig_idx = idx // CFG['num_splits']\n            split_idx = idx % CFG['num_splits']\n        else:\n            orig_idx = idx\n            split_idx = None\n\n        i, j = self.coords[orig_idx]\n        target_size = CFG['target_size']\n        \n        try:\n            # Extract patch\n            patch = self.volume[:, i:i+target_size, j:j+target_size]\n            patch = np.stack([cv2.resize(slice_img, (target_size, target_size)) \n                            for slice_img in patch])\n            patch = np.transpose(patch, (1, 2, 0))  # [H, W, C]\n            \n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+target_size, j:j+target_size]\n                mask = cv2.resize(mask, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            \n            if split_idx is not None and mask is not None:\n                # 8-way splitting logic\n                h, w = patch.shape[:2]\n                h_half, w_half = h//2, w//2\n                h_quarter, w_quarter = h//4, w//4\n                \n                if split_idx == 0:  # Top-left\n                    patch = patch[:h_half, :w_half]\n                    mask = mask[:h_half, :w_half]\n                elif split_idx == 1:  # Top-right\n                    patch = patch[:h_half, w_half:]\n                    mask = mask[:h_half, w_half:]\n                elif split_idx == 2:  # Bottom-left\n                    patch = patch[h_half:, :w_half]\n                    mask = mask[h_half:, :w_half]\n                elif split_idx == 3:  # Bottom-right\n                    patch = patch[h_half:, w_half:]\n                    mask = mask[h_half:, w_half:]\n                elif split_idx == 4:  # Top-center\n                    patch = patch[:h_half, w_quarter:w_quarter*3]\n                    mask = mask[:h_half, w_quarter:w_quarter*3]\n                elif split_idx == 5:  # Bottom-center\n                    patch = patch[h_half:, w_quarter:w_quarter*3]\n                    mask = mask[h_half:, w_quarter:w_quarter*3]\n                elif split_idx == 6:  # Middle-left\n                    patch = patch[h_quarter:h_quarter*3, :w_half]\n                    mask = mask[h_quarter:h_quarter*3, :w_half]\n                elif split_idx == 7:  # Middle-right\n                    patch = patch[h_quarter:h_quarter*3, w_half:]\n                    mask = mask[h_quarter:h_quarter*3, w_half:]\n                \n                # Resize back to target size\n                patch = cv2.resize(patch, (target_size, target_size))\n                mask = cv2.resize(mask, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            \n            if self.is_train and mask is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                image = transformed['image']  # [C, H, W]\n                mask = transformed['mask']    # [H, W]\n            else:\n                image = torch.tensor(patch).permute(2, 0, 1).float()  # [C, H, W]\n                if mask is not None:\n                    mask = torch.tensor(mask).float()  # [H, W]\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error processing patch at ({i}, {j}): {str(e)}\")\n            image = torch.zeros((CFG['slice_end']-CFG['slice_start'], target_size, target_size), dtype=torch.float32)\n            mask = torch.zeros((target_size, target_size), dtype=torch.float32) if self.mask is not None else None\n            return (image, mask) if mask is not None else image\n\nclass VNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super(VNet, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = 1\n        \n        # Encoder\n        self.encoder1 = self._make_encoder_block(in_channels, 16)\n        self.encoder2 = self._make_encoder_block(16, 32, stride=(1,2,2))\n        self.encoder3 = self._make_encoder_block(32, 64, stride=(1,2,2))\n        self.encoder4 = self._make_encoder_block(64, 128, stride=(1,2,2))\n        \n        # Decoder\n        self.decoder4 = self._make_decoder_block(128, 64, scale_factor=(1,2,2))\n        self.decoder3 = self._make_decoder_block(64, 32, scale_factor=(1,2,2))\n        self.decoder2 = self._make_decoder_block(32, 16, scale_factor=(1,2,2))\n        \n        # Final projection\n        self.final_proj = nn.Sequential(\n            nn.Conv3d(16, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n    def _make_encoder_block(self, in_channels, out_channels, stride=1):\n        return nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True),\n            nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True)\n        )\n    \n    def _make_decoder_block(self, in_channels, out_channels, scale_factor):\n        return nn.Sequential(\n            nn.Upsample(scale_factor=scale_factor, mode='trilinear', align_corners=True),\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Input: [B, C, H, W]\n        x = x.unsqueeze(2)  # Add depth dim: [B, C, 1, H, W]\n        \n        # Encoder\n        enc1 = self.encoder1(x)        # [B, 16, 1, H, W]\n        enc2 = self.encoder2(enc1)     # [B, 32, 1, H/2, W/2]\n        enc3 = self.encoder3(enc2)     # [B, 64, 1, H/4, W/4]\n        enc4 = self.encoder4(enc3)     # [B, 128, 1, H/8, W/8]\n        \n        # Decoder with skip connections\n        dec4 = self.decoder4(enc4) + enc3\n        dec3 = self.decoder3(dec4) + enc2\n        dec2 = self.decoder2(dec3) + enc1\n        \n        # Final output\n        out = self.final_proj(dec2)    # [B, 1, 1, H, W]\n        out = out.squeeze(2)           # [B, 1, H, W]\n        \n        return out\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        targets = targets.unsqueeze(1)  # [B, H, W] -> [B, 1, H, W]\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = F.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs = imgs.to(device)          # [B, C, H, W]\n        masks = masks.to(device)         # [B, H, W]\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)          # [B, 1, H, W]\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        \n    return total_loss / len(loader)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.cpu().numpy()  # [B, H, W]\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()  # [B, 1, H, W]\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    thresholds = [0.3, 0.35, 0.4, 0.45, 0.5, 0.55, 0.6]\n    best_thresh, best_f05 = 0.3, 0\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n    \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh\n\ndef visualize_predictions(inputs, preds, labels, threshold=0.5, num_samples=5):\n    indices = np.random.choice(len(inputs), num_samples, replace=False)\n    \n    plt.figure(figsize=(20, 4*num_samples))\n    for i, idx in enumerate(indices):\n        input_img = inputs[idx][len(inputs[idx])//2]\n        pred = preds[idx][0] > threshold\n        label = labels[idx][0]\n        \n        plt.subplot(num_samples, 3, i*3 + 1)\n        plt.imshow(input_img, cmap='gray')\n        plt.title(f\"Input (Sample {idx})\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 2)\n        plt.imshow(label, cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 3)\n        plt.imshow(pred, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        model = VNet().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nVisualizing best predictions...\")\n                    sample_inputs = []\n                    for batch in test_loader:\n                        sample_inputs.extend(batch[0].cpu().numpy())\n                        if len(sample_inputs) >= 20:\n                            break\n                    visualize_predictions(sample_inputs, all_preds, all_labels, threshold=best_thresh)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"######*************************** 3D VNet with split into 8 splits\nimport os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"test_case\": (\"2\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 8,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_volume(self):\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            volume.append(img)\n        return np.stack(volume).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        if self.mask_path is None:\n            return None\n            \n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is None:\n            raise ValueError(f\"Failed to load mask {self.mask_path}\")\n        return mask.astype(np.float32) / 255.0\n\n    def _generate_coords(self):\n        z, h, w = self.volume.shape\n        coords = []\n        target_size = CFG['target_size']\n        \n        for i in range(0, h - target_size + 1, CFG['stride']):\n            for j in range(0, w - target_size + 1, CFG['stride']):\n                if i + target_size > h or j + target_size > w:\n                    continue\n                    \n                if self.mask is None or np.sum(self.mask[i:i+target_size, j:j+target_size]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            orig_idx = idx // CFG['num_splits']\n            split_idx = idx % CFG['num_splits']\n        else:\n            orig_idx = idx\n            split_idx = None\n\n        i, j = self.coords[orig_idx]\n        target_size = CFG['target_size']\n        \n        try:\n            # Extract patch\n            patch = self.volume[:, i:i+target_size, j:j+target_size]\n            patch = np.stack([cv2.resize(slice_img, (target_size, target_size)) \n                            for slice_img in patch])\n            patch = np.transpose(patch, (1, 2, 0))\n            \n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+target_size, j:j+target_size]\n                mask = cv2.resize(mask, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            \n            if split_idx is not None and mask is not None:\n                # 8-way splitting logic\n                h, w = patch.shape[:2]\n                h_half, w_half = h//2, w//2\n                h_quarter, w_quarter = h//4, w//4\n                \n                if split_idx == 0:  # Top-left\n                    patch = patch[:h_half, :w_half]\n                    mask = mask[:h_half, :w_half]\n                elif split_idx == 1:  # Top-right\n                    patch = patch[:h_half, w_half:]\n                    mask = mask[:h_half, w_half:]\n                elif split_idx == 2:  # Bottom-left\n                    patch = patch[h_half:, :w_half]\n                    mask = mask[h_half:, :w_half]\n                elif split_idx == 3:  # Bottom-right\n                    patch = patch[h_half:, w_half:]\n                    mask = mask[h_half:, w_half:]\n                elif split_idx == 4:  # Top-center\n                    patch = patch[:h_half, w_quarter:w_quarter*3]\n                    mask = mask[:h_half, w_quarter:w_quarter*3]\n                elif split_idx == 5:  # Bottom-center\n                    patch = patch[h_half:, w_quarter:w_quarter*3]\n                    mask = mask[h_half:, w_quarter:w_quarter*3]\n                elif split_idx == 6:  # Middle-left\n                    patch = patch[h_quarter:h_quarter*3, :w_half]\n                    mask = mask[h_quarter:h_quarter*3, :w_half]\n                elif split_idx == 7:  # Middle-right\n                    patch = patch[h_quarter:h_quarter*3, w_half:]\n                    mask = mask[h_quarter:h_quarter*3, w_half:]\n                \n                # Resize back to target size\n                patch = cv2.resize(patch, (target_size, target_size))\n                mask = cv2.resize(mask, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            \n            if self.is_train and mask is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                image = transformed['image']\n                mask = transformed['mask']\n            else:\n                image = torch.tensor(patch).permute(2, 0, 1)\n                if mask is not None:\n                    mask = torch.tensor(mask)\n            \n            # Ensure consistent dimensions\n            if image.shape[-2:] != (target_size, target_size):\n                image = F.interpolate(image.unsqueeze(0), size=(target_size, target_size), mode='bilinear').squeeze(0)\n            \n            if mask is not None and mask.shape[-2:] != (target_size, target_size):\n                mask = F.interpolate(mask.unsqueeze(0).unsqueeze(0), size=(target_size, target_size), mode='nearest').squeeze(0).squeeze(0)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error processing patch at ({i}, {j}): {str(e)}\")\n            # Return zero tensors with correct dimensions\n            image = torch.zeros((CFG['slice_end']-CFG['slice_start'], target_size, target_size), dtype=torch.float32)\n            mask = torch.zeros((target_size, target_size), dtype=torch.float32) if self.mask is not None else None\n            return (image, mask) if mask is not None else image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass VNet(nn.Module):\n    def __init__(self, in_channels=CFG['slice_end']-CFG['slice_start']):\n        super(VNet, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = 1\n        \n        # Encoder\n        self.encoder1 = self._make_encoder_block(in_channels, 16)\n        self.encoder2 = self._make_encoder_block(16, 32, stride=2)\n        self.encoder3 = self._make_encoder_block(32, 64, stride=2)\n        self.encoder4 = self._make_encoder_block(64, 128, stride=2)\n        \n        # Decoder\n        self.decoder4 = self._make_decoder_block(128, 64)\n        self.decoder3 = self._make_decoder_block(64, 32)\n        self.decoder2 = self._make_decoder_block(32, 16)\n        \n        # Final output (remove depth dimension)\n        self.final_conv = nn.Sequential(\n            nn.Conv3d(16, self.out_channels, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n        # Spatial upsampling if needed\n        self.upsample = nn.Upsample(size=(CFG['target_size'], CFG['target_size']), mode='bilinear')\n        \n    def _make_encoder_block(self, in_channels, out_channels, stride=1):\n        return nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True),\n            nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True)\n        )\n    \n    def _make_decoder_block(self, in_channels, out_channels):\n        return nn.Sequential(\n            nn.ConvTranspose3d(in_channels, out_channels, kernel_size=2, stride=2),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True),\n            nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.LeakyReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Input shape: [B, C, H, W]\n        # Add dummy depth dimension: [B, C, H, W, 1]\n        x = x.unsqueeze(-1)\n        \n        # Encoder\n        enc1 = self.encoder1(x)        # [B, 16, H, W, 1]\n        enc2 = self.encoder2(enc1)     # [B, 32, H/2, W/2, 1]\n        enc3 = self.encoder3(enc2)     # [B, 64, H/4, W/4, 1]\n        enc4 = self.encoder4(enc3)     # [B, 128, H/8, W/8, 1]\n        \n        # Decoder\n        dec4 = self.decoder4(enc4) + enc3  # [B, 64, H/4, W/4, 1]\n        dec3 = self.decoder3(dec4) + enc2  # [B, 32, H/2, W/2, 1]\n        dec2 = self.decoder2(dec3) + enc1  # [B, 16, H, W, 1]\n        \n        # Final output\n        out = self.final_conv(dec2)    # [B, 1, H, W, 1]\n        out = out.squeeze(-1)          # Remove depth: [B, 1, H, W]\n        \n        # Upsample if needed\n        if out.shape[-2:] != (CFG['target_size'], CFG['target_size']):\n            out = self.upsample(out)\n            \n        return out\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = nn.functional.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(preds, size=masks.shape[-2:], mode='bilinear')\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    thresholds = [0.3, 0.35, 0.4, 0.45, 0.5, 0.55, 0.6]\n    best_thresh, best_f05 = 0.3, 0\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n    \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh\n\ndef visualize_predictions(inputs, preds, labels, threshold=0.5, num_samples=5):\n    indices = np.random.choice(len(inputs), num_samples, replace=False)\n    \n    plt.figure(figsize=(20, 4*num_samples))\n    for i, idx in enumerate(indices):\n        input_img = inputs[idx][len(inputs[idx])//2]\n        pred = preds[idx][0] > threshold\n        label = labels[idx][0]\n        \n        plt.subplot(num_samples, 3, i*3 + 1)\n        plt.imshow(input_img, cmap='gray')\n        plt.title(f\"Input (Sample {idx})\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 2)\n        plt.imshow(label, cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 3)\n        plt.imshow(pred, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        model = VNet().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nVisualizing best predictions...\")\n                    sample_inputs = []\n                    for batch in test_loader:\n                        sample_inputs.extend(batch[0].cpu().numpy())\n                        if len(sample_inputs) >= 20:\n                            break\n                    visualize_predictions(sample_inputs, all_preds, all_labels, threshold=best_thresh)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    thresholds = [0.41, 0.42, 0.43, 0.44, 0.45, 0.46, 0.47, 0.48, 0.49, \n                  0.5, 0.51, 0.52, 0.53, 0.54, 0.55, 0.56, 0.57, 0.58, 0.59, 0.6]\n    \n    best_thresh, best_f05 = 0.3, 0\n    threshold_results = []\n    \n    print(\"\\nTesting thresholds:\")\n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            precision = precision_score(all_labels.flatten(), binarized.flatten(), zero_division=0)\n            recall = recall_score(all_labels.flatten(), binarized.flatten(), zero_division=0)\n            \n            threshold_results.append((t, f05, precision, recall))\n            print(f\"Threshold {t:.2f} - F0.5: {f05:.4f}, Precision: {precision:.4f}, Recall: {recall:.4f}\")\n            \n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n                print(f\"  ↳ New best!\")\n        except Exception as e:\n            print(f\"Error calculating metrics for threshold {t}: {str(e)}\")\n            continue\n    \n    # Sort and show all results for clarity\n    threshold_results.sort(key=lambda x: x[1], reverse=True)\n    print(\"\\nThreshold results ranked:\")\n    for t, f05, prec, rec in threshold_results[:10]:  # Show top 10\n        print(f\"Thresh {t:.2f}: F0.5={f05:.4f}, P={prec:.4f}, R={rec:.4f}\")\n    \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#************************* ConvNeXT model with threshold range\n#pre-trained model with data split into 8 parts\nimport os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 4,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_volume(self):\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            volume.append(img)\n        return np.stack(volume).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        if self.mask_path is None:\n            return None\n            \n        mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n        if mask is None:\n            raise ValueError(f\"Failed to load mask {self.mask_path}\")\n        return mask.astype(np.float32) / 255.0\n\n    def _generate_coords(self):\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            orig_idx = idx // CFG['num_splits']\n            split_idx = idx % CFG['num_splits']\n        else:\n            orig_idx = idx\n            split_idx = None\n\n        i, j = self.coords[orig_idx]\n        \n        patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n        patch = np.stack([cv2.resize(slice_img, (CFG['target_size'], CFG['target_size'])) \n                         for slice_img in patch])\n        patch = np.transpose(patch, (1, 2, 0))\n        \n        mask = None\n        if self.mask is not None:\n            mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n            mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n        \n        if split_idx is not None and mask is not None:\n            def safe_split(x):\n                h, w = x.shape[:2]\n                if split_idx == 0: return x[:h//2, :w//2]\n                elif split_idx == 1: return x[:h//2, w//2:]\n                elif split_idx == 2: return x[h//2:, :w//2]\n                elif split_idx == 3: return x[h//2:, w//2:]\n            \n            patch = safe_split(patch)\n            mask = safe_split(mask)\n        \n        if self.is_train and mask is not None:\n            transformed = self.transform(image=patch, mask=mask)\n            image = transformed['image']\n            mask = transformed['mask']\n        else:\n            image = torch.tensor(patch).permute(2, 0, 1)\n            if mask is not None:\n                mask = torch.tensor(mask)\n        \n        return (image, mask) if mask is not None else image\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Enable pretrained weights\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        self.final_upsample = nn.Upsample(size=(CFG['target_size'], CFG['target_size']))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = nn.functional.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(preds, size=masks.shape[-2:], mode='bilinear')\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    thresholds = [0.41, 0.42, 0.43, 0.44, 0.45, 0.46, 0.47, 0.48, 0.49, \n                  0.5, 0.51, 0.52, 0.53, 0.54, 0.55, 0.56, 0.57, 0.58, 0.59, 0.6]\n    \n    best_thresh, best_f05 = 0.3, 0\n    threshold_results = []\n    \n    print(\"\\nTesting thresholds:\")\n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            precision = precision_score(all_labels.flatten(), binarized.flatten(), zero_division=0)\n            recall = recall_score(all_labels.flatten(), binarized.flatten(), zero_division=0)\n            \n            threshold_results.append((t, f05, precision, recall))\n            print(f\"Threshold {t:.2f} - F0.5: {f05:.4f}, Precision: {precision:.4f}, Recall: {recall:.4f}\")\n            \n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n                print(f\"  ↳ New best!\")\n        except Exception as e:\n            print(f\"Error calculating metrics for threshold {t}: {str(e)}\")\n            continue\n    \n    # Sort and show all results for clarity\n    threshold_results.sort(key=lambda x: x[1], reverse=True)\n    print(\"\\nThreshold results ranked:\")\n    for t, f05, prec, rec in threshold_results[:10]:  # Show top 10\n        print(f\"Thresh {t:.2f}: F0.5={f05:.4f}, P={prec:.4f}, R={rec:.4f}\")\n    \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh\n\ndef visualize_predictions(inputs, preds, labels, threshold=0.5, num_samples=5):\n    indices = np.random.choice(len(inputs), num_samples, replace=False)\n    \n    plt.figure(figsize=(20, 4*num_samples))\n    for i, idx in enumerate(indices):\n        input_img = inputs[idx][len(inputs[idx])//2]\n        pred = preds[idx][0] > threshold\n        label = labels[idx][0]\n        \n        plt.subplot(num_samples, 3, i*3 + 1)\n        plt.imshow(input_img, cmap='gray')\n        plt.title(f\"Input (Sample {idx})\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 2)\n        plt.imshow(label, cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 3)\n        plt.imshow(pred, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model with pretrained weights\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nVisualizing best predictions...\")\n                    sample_inputs = []\n                    for batch in test_loader:\n                        sample_inputs.extend(batch[0].cpu().numpy())\n                        if len(sample_inputs) >= 20:\n                            break\n                    visualize_predictions(sample_inputs, all_preds, all_labels, threshold=best_thresh)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport tifffile as tiff\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# Path configuration\nbase_path = \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/\"\nsurface_volume_path = os.path.join(base_path, \"surface_volume\")\nink_labels_path = os.path.join(base_path, \"inklabels.png\")\noutput_dir = \"/kaggle/working/\"  # Directory to save outputs\n\n# Create output directory if it doesn't exist\nos.makedirs(output_dir, exist_ok=True)\n\n# Parameters\ntarget_size = (128, 128)\nnum_slices = 65  # 00.tif to 64.tif\n\ndef load_and_resize_slices(slice_dir, num_slices, target_size):\n    \"\"\"Load and resize all surface volume slices\"\"\"\n    slices = []\n    for i in range(num_slices):\n        slice_path = os.path.join(slice_dir, f\"{i:02d}.tif\")\n        img = tiff.imread(slice_path)\n        img_pil = Image.fromarray(img)\n        if img_pil.mode != 'L':\n            img_pil = img_pil.convert('L')  # Convert to grayscale if needed\n        img_resized = img_pil.resize(target_size, Image.Resampling.LANCZOS)\n        slices.append(np.array(img_resized))\n    return np.stack(slices)\n\ndef load_and_resize_labels(label_path, target_size):\n    \"\"\"Load and resize ink labels\"\"\"\n    labels = Image.open(label_path)\n    if labels.mode != '1':\n        labels = labels.convert('1')  # Convert to binary if needed\n    labels_resized = labels.resize(target_size, Image.Resampling.NEAREST)\n    return np.array(labels_resized)\n\n# Load and process the data\nprint(\"Loading and processing data...\")\nsurface_volume = load_and_resize_slices(surface_volume_path, num_slices, target_size)\nink_labels = load_and_resize_labels(ink_labels_path, target_size)\n\n# Binarize the labels (0 = non-inked, 1 = inked)\nink_labels_binary = (ink_labels > 0).astype(np.uint8)\n\n# Analyze each slice\nprint(\"Analyzing slices...\")\nslice_stats = []\npixel_data = []\n\nfor slice_idx in range(num_slices):\n    current_slice = surface_volume[slice_idx]\n    \n    # Get pixel values for inked and non-inked areas\n    inked_pixels = current_slice[ink_labels_binary == 1]\n    non_inked_pixels = current_slice[ink_labels_binary == 0]\n    \n    # Calculate statistics\n    mean_inked = np.mean(inked_pixels) if len(inked_pixels) > 0 else 0\n    mean_non_inked = np.mean(non_inked_pixels) if len(non_inked_pixels) > 0 else 0\n    diff = mean_inked - mean_non_inked\n    \n    # Store slice statistics\n    slice_stats.append({\n        'slice': slice_idx,\n        'mean_inked': mean_inked,\n        'mean_non_inked': mean_non_inked,\n        'difference': diff,\n        'num_inked_pixels': len(inked_pixels),\n        'num_non_inked_pixels': len(non_inked_pixels)\n    })\n    \n    # Store pixel-level data for CSV (sample every 10th pixel to reduce size)\n    for y in range(0, target_size[0], 10):\n        for x in range(0, target_size[1], 10):\n            pixel_data.append({\n                'slice': slice_idx,\n                'x': x,\n                'y': y,\n                'value': current_slice[y, x],\n                'is_inked': ink_labels_binary[y, x]\n            })\n\n# Convert to DataFrames\nstats_df = pd.DataFrame(slice_stats)\npixels_df = pd.DataFrame(pixel_data)\n\n# Save to CSV files\nprint(\"Saving results to CSV...\")\nstats_csv_path = os.path.join(output_dir, \"slice_statistics_fold3.csv\")\npixels_csv_path = os.path.join(output_dir, \"pixel_data_fold3.csv\")\n\nstats_df.to_csv(stats_csv_path, index=False)\npixels_df.to_csv(pixels_csv_path, index=False)\n\n# Print summary statistics\nprint(\"\\nSlice Statistics Summary:\")\nprint(f\"{'Slice':<6} {'Inked Mean':<12} {'Non-Inked Mean':<15} {'Difference':<12} {'Inked Pixels':<13} {'Non-Inked Pixels':<15}\")\nfor stat in slice_stats:\n    print(f\"{stat['slice']:<6} {stat['mean_inked']:<12.2f} {stat['mean_non_inked']:<15.2f} \"\n          f\"{stat['difference']:<12.2f} {stat['num_inked_pixels']:<13} {stat['num_non_inked_pixels']:<15}\")\n\n# Plot the differences\nplt.figure(figsize=(12, 6))\nplt.plot(stats_df['slice'], stats_df['difference'], marker='o')\nplt.title('Difference between Inked and Non-Inked Pixel Values per Slice')\nplt.xlabel('Slice Number')\nplt.ylabel('Mean Difference (Inked - Non-Inked)')\nplt.grid(True)\nplot_path = os.path.join(output_dir, \"differences_plot_fold3.png\")\nplt.savefig(plot_path)\nplt.show()\n\n# Find slices with the largest differences\nsorted_slices = stats_df.copy()\nsorted_slices['abs_difference'] = sorted_slices['difference'].abs()\nsorted_slices = sorted_slices.sort_values('abs_difference', ascending=False)\n\nprint(\"\\nTop 5 slices with largest absolute differences:\")\nprint(sorted_slices[['slice', 'difference']].head(5).to_string(index=False))\n\n# Save the sorted differences\nsorted_csv_path = os.path.join(output_dir, \"sorted_differences_fold3.csv\")\nsorted_slices.to_csv(sorted_csv_path, index=False)\n\nprint(\"\\nProcessing complete! Files saved to:\")\nprint(f\"- Slice statistics: {stats_csv_path}\")\nprint(f\"- Pixel data: {pixels_csv_path}\")\nprint(f\"- Sorted differences: {sorted_csv_path}\")\nprint(f\"- Differences plot: {plot_path}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport tifffile as tiff\nfrom PIL import Image, ImageOps\nimport matplotlib.pyplot as plt\n\n# Path configuration\nbase_path = \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/\"\nsurface_volume_path = os.path.join(base_path, \"surface_volume\")\nink_labels_path = os.path.join(base_path, \"inklabels.png\")\n\n# Parameters\ntarget_size = (128, 128)\nnum_slices = 65  # 00.tif to 64.tif\n\ndef load_and_resize_slices(slice_dir, num_slices, target_size):\n    \"\"\"Load and resize all surface volume slices\"\"\"\n    slices = []\n    for i in range(num_slices):\n        slice_path = os.path.join(slice_dir, f\"{i:02d}.tif\")\n        img = tiff.imread(slice_path)\n        \n        # Convert to PIL Image and ensure proper mode\n        img_pil = Image.fromarray(img)\n        if img_pil.mode != 'L':\n            img_pil = img_pil.convert('L')\n            \n        # Use updated resampling method\n        img_resized = img_pil.resize(target_size, Image.Resampling.LANCZOS)\n        slices.append(np.array(img_resized))\n    return np.stack(slices)\n\ndef load_and_resize_labels(label_path, target_size):\n    \"\"\"Load and resize ink labels\"\"\"\n    labels = Image.open(label_path)\n    # Ensure binary mode for labels\n    if labels.mode != '1':\n        labels = labels.convert('1')\n    labels_resized = labels.resize(target_size, Image.Resampling.NEAREST)\n    return np.array(labels_resized)\n\n# Load and process the data\nsurface_volume = load_and_resize_slices(surface_volume_path, num_slices, target_size)\nink_labels = load_and_resize_labels(ink_labels_path, target_size)\n\n# Binarize the labels (0 = non-inked, 1 = inked)\nink_labels_binary = (ink_labels > 0).astype(np.uint8)\n\n# Analyze each slice\nslice_stats = []\n\nfor slice_idx in range(num_slices):\n    current_slice = surface_volume[slice_idx]\n    \n    # Get pixel values for inked and non-inked areas\n    inked_pixels = current_slice[ink_labels_binary == 1]\n    non_inked_pixels = current_slice[ink_labels_binary == 0]\n    \n    # Calculate statistics\n    mean_inked = np.mean(inked_pixels) if len(inked_pixels) > 0 else 0\n    mean_non_inked = np.mean(non_inked_pixels) if len(non_inked_pixels) > 0 else 0\n    diff = mean_inked - mean_non_inked\n    \n    slice_stats.append({\n        'slice': slice_idx,\n        'mean_inked': mean_inked,\n        'mean_non_inked': mean_non_inked,\n        'difference': diff,\n        'num_inked_pixels': len(inked_pixels),\n        'num_non_inked_pixels': len(non_inked_pixels)\n    })\n\n# Print summary statistics\nprint(f\"{'Slice':<6} {'Inked Mean':<12} {'Non-Inked Mean':<15} {'Difference':<12} {'Inked Pixels':<13} {'Non-Inked Pixels':<15}\")\nfor stat in slice_stats:\n    print(f\"{stat['slice']:<6} {stat['mean_inked']:<12.2f} {stat['mean_non_inked']:<15.2f} \"\n          f\"{stat['difference']:<12.2f} {stat['num_inked_pixels']:<13} {stat['num_non_inked_pixels']:<15}\")\n\n# Plot the differences\ndifferences = [s['difference'] for s in slice_stats]\nplt.figure(figsize=(12, 6))\nplt.plot(differences, marker='o')\nplt.title('Difference between Inked and Non-Inked Pixel Values per Slice')\nplt.xlabel('Slice Number')\nplt.ylabel('Mean Difference (Inked - Non-Inked)')\nplt.grid(True)\nplt.show()\n\n# Find slices with the largest differences\nsorted_slices = sorted(slice_stats, key=lambda x: abs(x['difference']), reverse=True)\nprint(\"\\nTop 5 slices with largest absolute differences:\")\nfor s in sorted_slices[:5]:\n    print(f\"Slice {s['slice']}: Difference = {s['difference']:.2f}\")\n\n# Optional: Save the processed data\n# np.save('processed_surface_volume_128x128.npy', surface_volume)\n# np.save('processed_ink_labels_128x128.npy', ink_labels_binary)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport tifffile as tiff\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# Path configuration\nbase_path = \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/\"\nsurface_volume_path = os.path.join(base_path, \"surface_volume\")\nink_labels_path = os.path.join(base_path, \"inklabels.png\")\n\n# Parameters\ntarget_size = (128, 128)\nnum_slices = 65  # 00.tif to 64.tif\n\ndef load_and_resize_slices(slice_dir, num_slices, target_size):\n    \"\"\"Load and resize all surface volume slices\"\"\"\n    slices = []\n    for i in range(num_slices):\n        slice_path = os.path.join(slice_dir, f\"{i:02d}.tif\")\n        img = tiff.imread(slice_path)\n        img_pil = Image.fromarray(img)\n        img_resized = img_pil.resize(target_size, Image.LANCZOS)\n        slices.append(np.array(img_resized))\n    return np.stack(slices)\n\ndef load_and_resize_labels(label_path, target_size):\n    \"\"\"Load and resize ink labels\"\"\"\n    labels = Image.open(label_path)\n    labels_resized = labels.resize(target_size, Image.NEAREST)\n    return np.array(labels_resized)\n\n# Load and process the data\nsurface_volume = load_and_resize_slices(surface_volume_path, num_slices, target_size)\nink_labels = load_and_resize_labels(ink_labels_path, target_size)\n\n# Binarize the labels (0 = non-inked, 1 = inked)\nink_labels_binary = (ink_labels > 0).astype(np.uint8)\n\n# Analyze each slice\nslice_stats = []\n\nfor slice_idx in range(num_slices):\n    current_slice = surface_volume[slice_idx]\n    \n    # Get pixel values for inked and non-inked areas\n    inked_pixels = current_slice[ink_labels_binary == 1]\n    non_inked_pixels = current_slice[ink_labels_binary == 0]\n    \n    # Calculate statistics\n    mean_inked = np.mean(inked_pixels) if len(inked_pixels) > 0 else 0\n    mean_non_inked = np.mean(non_inked_pixels) if len(non_inked_pixels) > 0 else 0\n    diff = mean_inked - mean_non_inked\n    \n    slice_stats.append({\n        'slice': slice_idx,\n        'mean_inked': mean_inked,\n        'mean_non_inked': mean_non_inked,\n        'difference': diff,\n        'num_inked_pixels': len(inked_pixels),\n        'num_non_inked_pixels': len(non_inked_pixels)\n    })\n\n# Print summary statistics\nprint(f\"{'Slice':<6} {'Inked Mean':<12} {'Non-Inked Mean':<15} {'Difference':<12} {'Inked Pixels':<13} {'Non-Inked Pixels':<15}\")\nfor stat in slice_stats:\n    print(f\"{stat['slice']:<6} {stat['mean_inked']:<12.2f} {stat['mean_non_inked']:<15.2f} \"\n          f\"{stat['difference']:<12.2f} {stat['num_inked_pixels']:<13} {stat['num_non_inked_pixels']:<15}\")\n\n# Plot the differences\ndifferences = [s['difference'] for s in slice_stats]\nplt.figure(figsize=(12, 6))\nplt.plot(differences, marker='o')\nplt.title('Difference between Inked and Non-Inked Pixel Values per Slice')\nplt.xlabel('Slice Number')\nplt.ylabel('Mean Difference (Inked - Non-Inked)')\nplt.grid(True)\nplt.show()\n\n# Find slices with the largest differences\nsorted_slices = sorted(slice_stats, key=lambda x: abs(x['difference']), reverse=True)\nprint(\"\\nTop 5 slices with largest absolute differences:\")\nfor s in sorted_slices[:5]:\n    print(f\"Slice {s['slice']}: Difference = {s['difference']:.2f}\")\n\n# Optional: Save the processed data\n# np.save('processed_surface_volume_128x128.npy', surface_volume)\n# np.save('processed_ink_labels_128x128.npy', ink_labels_binary)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport tifffile as tiff\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.metrics import classification_report, accuracy_score\nimport time\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection/train/2/\"\nNUM_SLICES = 65\nSLICE_HEIGHT = 8184  # Update with actual dimensions if different\nSLICE_WIDTH = 8320   # Update with actual dimensions if different\n\ndef load_surface_volume(data_dir, num_slices):\n    \"\"\"Load all surface volume slices into a 3D numpy array\"\"\"\n    volume = np.zeros((num_slices, SLICE_HEIGHT, SLICE_WIDTH), dtype=np.float32)\n    \n    for i in range(num_slices):\n        slice_path = os.path.join(data_dir, \"surface_volume\", f\"{i:02d}.tif\")\n        if os.path.exists(slice_path):\n            volume[i] = tiff.imread(slice_path)\n        else:\n            print(f\"Warning: Missing slice {i:02d}.tif\")\n    \n    return volume\n\ndef load_ink_labels(data_dir):\n    \"\"\"Load the ink labels as a binary mask\"\"\"\n    label_path = os.path.join(data_dir, \"inklabels.png\")\n    labels = np.array(Image.open(label_path))\n    return labels\n\ndef preprocess_data(volume, labels):\n    \"\"\"Prepare data for machine learning\"\"\"\n    # Flatten the volume to 2D (pixels x slices)\n    X = volume.reshape(volume.shape[0], -1).T  # Shape: (height*width, num_slices)\n    y = labels.ravel()  # Flatten labels\n    \n    # Remove pixels where we have no information (zeros in all slices)\n    valid_pixels = np.any(X != 0, axis=1)\n    X = X[valid_pixels]\n    y = y[valid_pixels]\n    \n    return X, y, valid_pixels\n\ndef train_model(X, y):\n    \"\"\"Train a Random Forest classifier\"\"\"\n    # Split data into training and validation sets\n    X_train, X_val, y_train, y_val = train_test_split(\n        X, y, test_size=0.2, random_state=42, stratify=y\n    )\n    \n    # Initialize and train the model\n    model = RandomForestClassifier(\n        n_estimators=100,\n        max_depth=10,\n        random_state=42,\n        n_jobs=-1,\n        class_weight='balanced'\n    )\n    \n    print(\"Training model...\")\n    start_time = time.time()\n    model.fit(X_train, y_train)\n    training_time = time.time() - start_time\n    print(f\"Training completed in {training_time:.2f} seconds\")\n    \n    # Evaluate on validation set\n    y_pred = model.predict(X_val)\n    print(\"\\nValidation Results:\")\n    print(classification_report(y_val, y_pred))\n    print(f\"Accuracy: {accuracy_score(y_val, y_pred):.4f}\")\n    \n    return model\n\ndef analyze_differences(volume, labels, model, valid_pixels):\n    \"\"\"Analyze differences between inked and non-inked pixels\"\"\"\n    # Get all pixel values\n    X_full = volume.reshape(volume.shape[0], -1).T\n    \n    # Predict probabilities for all valid pixels\n    proba = model.predict_proba(X_full[valid_pixels])[:, 1]\n    \n    # Reshape predictions to original image dimensions\n    predictions = np.zeros(SLICE_HEIGHT * SLICE_WIDTH)\n    predictions[valid_pixels] = proba\n    predictions = predictions.reshape(SLICE_HEIGHT, SLICE_WIDTH)\n    \n    # Calculate mean values for inked vs non-inked across slices\n    inked_mask = labels > 0\n    non_inked_mask = labels == 0\n    \n    # Calculate mean intensity per slice for inked vs non-inked\n    inked_means = []\n    non_inked_means = []\n    differences = []\n    \n    for i in range(volume.shape[0]):\n        slice_data = volume[i]\n        inked_mean = np.mean(slice_data[inked_mask])\n        non_inked_mean = np.mean(slice_data[non_inked_mask])\n        difference = inked_mean - non_inked_mean\n        \n        inked_means.append(inked_mean)\n        non_inked_means.append(non_inked_mean)\n        differences.append(difference)\n    \n    # Plot the results\n    plt.figure(figsize=(15, 6))\n    plt.plot(inked_means, label='Inked Pixels Mean Intensity')\n    plt.plot(non_inked_means, label='Non-Inked Pixels Mean Intensity')\n    plt.plot(differences, label='Difference (Inked - Non-Inked)')\n    plt.xlabel('Slice Number')\n    plt.ylabel('Intensity')\n    plt.title('Mean Intensity Across Slices')\n    plt.legend()\n    plt.grid()\n    plt.show()\n    \n    # Find slices with largest differences\n    best_slices = np.argsort(differences)[-5:][::-1]\n    print(\"\\nTop 5 slices with largest intensity differences:\")\n    for i, slice_num in enumerate(best_slices):\n        print(f\"{i+1}. Slice {slice_num:02d}: Difference = {differences[slice_num]:.2f}\")\n    \n    return predictions, differences\n\ndef visualize_results(volume, labels, predictions, best_slice_idx):\n    \"\"\"Visualize original data and predictions\"\"\"\n    plt.figure(figsize=(20, 15))\n    \n    # Plot ground truth labels\n    plt.subplot(2, 2, 1)\n    plt.imshow(labels, cmap='gray')\n    plt.title('Ground Truth Ink Labels')\n    \n    # Plot predictions\n    plt.subplot(2, 2, 2)\n    plt.imshow(predictions, cmap='viridis')\n    plt.title('Model Predictions (Probability)')\n    \n    # Plot best slice\n    plt.subplot(2, 2, 3)\n    plt.imshow(volume[best_slice_idx], cmap='gray')\n    plt.title(f'Best Slice {best_slice_idx:02d}')\n    \n    # Plot histogram of intensities\n    plt.subplot(2, 2, 4)\n    plt.hist(volume[best_slice_idx][labels > 0].ravel(), bins=50, alpha=0.5, label='Inked')\n    plt.hist(volume[best_slice_idx][labels == 0].ravel(), bins=50, alpha=0.5, label='Non-Inked')\n    plt.yscale('log')\n    plt.title('Intensity Distribution')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    # Load data\n    print(\"Loading surface volume slices...\")\n    volume = load_surface_volume(DATA_DIR, NUM_SLICES)\n    print(f\"Loaded volume shape: {volume.shape}\")\n    \n    print(\"\\nLoading ink labels...\")\n    labels = load_ink_labels(DATA_DIR)\n    print(f\"Loaded labels shape: {labels.shape}\")\n    \n    # Preprocess data\n    print(\"\\nPreprocessing data...\")\n    X, y, valid_pixels = preprocess_data(volume, labels)\n    print(f\"Preprocessed data shape: {X.shape}\")\n    print(f\"Class distribution: {np.bincount(y)} (0: non-inked, 1: inked)\")\n    \n    # Train model\n    model = train_model(X, y)\n    \n    # Analyze differences\n    print(\"\\nAnalyzing differences between inked and non-inked pixels...\")\n    predictions, differences = analyze_differences(volume, labels, model, valid_pixels)\n    \n    # Visualize results\n    best_slice_idx = np.argmax(differences)\n    visualize_results(volume, labels, predictions, best_slice_idx)\n    \n    # Save predictions\n    output_dir = \"/kaggle/working/\"\n    os.makedirs(output_dir, exist_ok=True)\n    np.save(os.path.join(output_dir, \"ink_predictions.npy\"), predictions)\n    print(f\"\\nPredictions saved to {output_dir}\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Training: Uses OneCycleLR scheduler for better learning rate management\n\n#The thresholds are now tested in this exact order:\n#[0.3, 0.35, 0.39, 0.4, 0.41, 0.42, 0.43, 0.44, 0.45, 0.46, 0.47, 0.48, 0.49, 0.5, 0.51, 0.52, 0.53, 0.54, 0.55, 0.56, 0.57, 0.58, 0.59, 0.6]\n\nimport os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nimport torch.cuda as cuda\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 4,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_volume(self):\n        \"\"\"Load volume slices with robust error handling\"\"\"\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            if img.size == 0:\n                raise ValueError(f\"Empty image {p}\")\n            volume.append(img)\n        return np.clip(np.stack(volume), 0, 255).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        \"\"\"Load mask with robust error handling\"\"\"\n        if self.mask_path is None:\n            return None\n            \n        try:\n            mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"Failed to load mask {self.mask_path}\")\n            if mask.size == 0:\n                raise ValueError(f\"Empty mask {self.mask_path}\")\n            mask = mask.astype(np.float32) / 255.0\n            print(f\"Mask loaded - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n            return mask\n        except Exception as e:\n            raise ValueError(f\"Failed to load mask {self.mask_path}: {str(e)}\")\n\n    def _generate_coords(self):\n        \"\"\"Generate coordinates for valid patches\"\"\"\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        \"\"\"Get item with robust error handling\"\"\"\n        try:\n            if self.is_train:\n                orig_idx = idx // CFG['num_splits']\n                split_idx = idx % CFG['num_splits']\n            else:\n                orig_idx = idx\n                split_idx = None\n\n            i, j = self.coords[orig_idx]\n            \n            # Get volume patch\n            patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Process mask\n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n                mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply split if training\n            if split_idx is not None and mask is not None:\n                def safe_split(x):\n                    h, w = x.shape[:2]\n                    if split_idx == 0: return x[:h//2, :w//2]\n                    elif split_idx == 1: return x[:h//2, w//2:]\n                    elif split_idx == 2: return x[h//2:, :w//2]\n                    elif split_idx == 3: return x[h//2:, w//2:]\n                \n                patch = safe_split(patch)\n                mask = safe_split(mask)\n            \n            # Apply transforms\n            if self.is_train and mask is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                image = transformed['image']\n                mask = transformed['mask']\n            else:\n                image = torch.tensor(patch).permute(2, 0, 1)  # CHW\n                if mask is not None:\n                    mask = torch.tensor(mask)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx}: {str(e)}\")\n            dummy_img = torch.zeros((CFG['slice_end']-CFG['slice_start'], CFG['target_size'], CFG['target_size']))\n            dummy_mask = torch.zeros((CFG['target_size'], CFG['target_size'])) if self.mask is not None else None\n            return (dummy_img, dummy_mask) if dummy_mask is not None else dummy_img\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        self.final_upsample = nn.Upsample(size=(CFG['target_size'], CFG['target_size']))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = nn.functional.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(preds, size=masks.shape[-2:], mode='bilinear')\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    # Debug mask values\n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    # Binarize labels properly\n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    # Custom threshold range as requested\n    thresholds = [0.3, 0.35, 0.39, 0.4, 0.41, 0.42, 0.43, 0.44, 0.45, \n                  0.46, 0.47, 0.48, 0.49, 0.5, 0.51, 0.52, 0.53, 0.54, \n                  0.55, 0.56, 0.57, 0.58, 0.59, 0.6]\n    \n    best_thresh, best_f05 = 0.3, 0\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n            \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh\n\ndef visualize_predictions(inputs, preds, labels, threshold=0.5, num_samples=5):\n    indices = np.random.choice(len(inputs), num_samples, replace=False)\n    \n    plt.figure(figsize=(20, 4*num_samples))\n    for i, idx in enumerate(indices):\n        input_img = inputs[idx][len(inputs[idx])//2]\n        pred = preds[idx][0] > threshold\n        label = labels[idx][0]\n        \n        plt.subplot(num_samples, 3, i*3 + 1)\n        plt.imshow(input_img, cmap='gray')\n        plt.title(f\"Input (Sample {idx})\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 2)\n        plt.imshow(label, cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 3)\n        plt.imshow(pred, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nVisualizing best predictions...\")\n                    sample_inputs = []\n                    for batch in test_loader:\n                        sample_inputs.extend(batch[0].cpu().numpy())\n                        if len(sample_inputs) >= 20:\n                            break\n                    visualize_predictions(sample_inputs, all_preds, all_labels, threshold=best_thresh)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom glob import glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm import create_model\nimport torch.nn as nn\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom sklearn.metrics import fbeta_score, precision_score, recall_score\nimport torch.cuda as cuda\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import GradScaler, autocast\nimport warnings\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Configuration\nCFG = {\n    \"train_case\": (\"3\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/3/inklabels.png\"),\n    \"test_case\": (\"1\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\", \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"),\n    \"target_size\": 128,\n    \"num_splits\": 4,\n    \"slice_start\": 16,\n    \"slice_end\": 40,\n    \"train_batch_size\": 8,\n    \"test_batch_size\": 4,\n    \"stride\": 64,\n    \"epochs\": 20,\n    \"lr\": 3e-4,\n    \"early_stop_patience\": 5,\n    \"amp\": True,\n    \"min_ink_pixels\": 10\n}\n\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_dir, mask_path=None, is_train=True):\n        self.volume_dir = volume_dir\n        self.mask_path = mask_path\n        self.is_train = is_train\n        self.volume = self._load_volume()\n        self.mask = self._load_mask()\n        self.coords = self._generate_coords()\n        \n        if len(self.coords) == 0:\n            raise ValueError(f\"No valid patches found in {volume_dir}\")\n            \n        self.transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            ToTensorV2()\n        ]) if is_train else ToTensorV2()\n\n    def _load_volume(self):\n        \"\"\"Load volume slices with robust error handling\"\"\"\n        paths = sorted(glob(os.path.join(self.volume_dir, '*.tif')))\n        if not paths:\n            raise FileNotFoundError(f\"No TIFF files found in {self.volume_dir}\")\n            \n        volume = []\n        for p in paths[CFG['slice_start']:CFG['slice_end']]:\n            img = cv2.imread(p, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise ValueError(f\"Failed to load {p}\")\n            if img.size == 0:\n                raise ValueError(f\"Empty image {p}\")\n            volume.append(img)\n        return np.clip(np.stack(volume), 0, 255).astype(np.float32) / 255.0\n\n    def _load_mask(self):\n        \"\"\"Load mask with robust error handling\"\"\"\n        if self.mask_path is None:\n            return None\n            \n        try:\n            mask = cv2.imread(self.mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"Failed to load mask {self.mask_path}\")\n            if mask.size == 0:\n                raise ValueError(f\"Empty mask {self.mask_path}\")\n            mask = mask.astype(np.float32) / 255.0\n            print(f\"Mask loaded - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n            return mask\n        except Exception as e:\n            raise ValueError(f\"Failed to load mask {self.mask_path}: {str(e)}\")\n\n    def _generate_coords(self):\n        \"\"\"Generate coordinates for valid patches\"\"\"\n        z, h, w = self.volume.shape\n        coords = []\n        \n        for i in range(0, h - CFG['target_size'], CFG['stride']):\n            for j in range(0, w - CFG['target_size'], CFG['stride']):\n                if self.mask is None or np.sum(self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]) > CFG['min_ink_pixels']:\n                    coords.append((i, j))\n        return coords\n\n    def __len__(self):\n        return len(self.coords) * CFG['num_splits'] if self.is_train else len(self.coords)\n\n    def __getitem__(self, idx):\n        \"\"\"Get item with robust error handling\"\"\"\n        try:\n            if self.is_train:\n                orig_idx = idx // CFG['num_splits']\n                split_idx = idx % CFG['num_splits']\n            else:\n                orig_idx = idx\n                split_idx = None\n\n            i, j = self.coords[orig_idx]\n            \n            # Get volume patch\n            patch = self.volume[:, i:i+CFG['target_size'], j:j+CFG['target_size']]\n            \n            # Resize each slice\n            resized_slices = []\n            for s in range(patch.shape[0]):\n                resized = cv2.resize(patch[s], (CFG['target_size'], CFG['target_size']))\n                resized_slices.append(resized)\n                \n            patch = np.stack(resized_slices)\n            patch = np.transpose(patch, (1, 2, 0))  # HWC\n            \n            # Process mask\n            mask = None\n            if self.mask is not None:\n                mask = self.mask[i:i+CFG['target_size'], j:j+CFG['target_size']]\n                mask = cv2.resize(mask, (CFG['target_size'], CFG['target_size']), interpolation=cv2.INTER_NEAREST)\n            \n            # Apply split if training\n            if split_idx is not None and mask is not None:\n                def safe_split(x):\n                    h, w = x.shape[:2]\n                    if split_idx == 0: return x[:h//2, :w//2]\n                    elif split_idx == 1: return x[:h//2, w//2:]\n                    elif split_idx == 2: return x[h//2:, :w//2]\n                    elif split_idx == 3: return x[h//2:, w//2:]\n                \n                patch = safe_split(patch)\n                mask = safe_split(mask)\n            \n            # Apply transforms\n            if self.is_train and mask is not None:\n                transformed = self.transform(image=patch, mask=mask)\n                image = transformed['image']\n                mask = transformed['mask']\n            else:\n                image = torch.tensor(patch).permute(2, 0, 1)  # CHW\n                if mask is not None:\n                    mask = torch.tensor(mask)\n            \n            return (image, mask) if mask is not None else image\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx}: {str(e)}\")\n            dummy_img = torch.zeros((CFG['slice_end']-CFG['slice_start'], CFG['target_size'], CFG['target_size']))\n            dummy_mask = torch.zeros((CFG['target_size'], CFG['target_size'])) if self.mask is not None else None\n            return (dummy_img, dummy_mask) if dummy_mask is not None else dummy_img\n\nclass InkDetector(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = create_model('convnext_tiny', pretrained=True, features_only=True, \n                                  in_chans=CFG['slice_end']-CFG['slice_start'])\n        enc_chs = [f['num_chs'] for f in self.encoder.feature_info]\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(enc_chs[-1], 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(256, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(128, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Upsample(scale_factor=2),\n            \n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.Conv2d(32, 1, 1)\n        )\n        self.final_upsample = nn.Upsample(size=(CFG['target_size'], CFG['target_size']))\n\n    def forward(self, x):\n        feats = self.encoder(x)[-1]\n        x = self.decoder(feats)\n        return self.final_upsample(x)\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self, alpha=0.8):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = BCEWithLogitsFocalLoss(alpha=alpha, gamma=2)\n        \n    def forward(self, inputs, targets):\n        if inputs.shape[-2:] != targets.shape[-2:]:\n            inputs = nn.functional.interpolate(inputs, size=targets.shape[-2:], mode='bilinear')\n        return 0.4*self.dice(inputs, targets) + 0.6*self.focal(inputs, targets)\n\nclass DiceLoss(nn.Module):\n    def forward(self, inputs, targets, smooth=1):\n        inputs = torch.sigmoid(inputs).view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        return 1 - ((2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth))\n\nclass BCEWithLogitsFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.bce_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        BCE = self.bce_logits(inputs, targets)\n        pt = torch.exp(-BCE)\n        return (self.alpha * (1 - pt) ** self.gamma * BCE).mean()\n\ndef get_class_weights(loader):\n    pos = 0\n    total = 0\n    for _, masks in loader:\n        pos += masks.sum()\n        total += masks.numel()\n    return torch.tensor([(total - pos)/pos])\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, pos_weight, scaler):\n    model.train()\n    total_loss = 0\n    \n    for imgs, masks in tqdm(loader, desc=\"Training\"):\n        imgs, masks = imgs.to(device), masks.to(device).unsqueeze(1)\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG['amp']):\n            preds = model(imgs)\n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(preds, size=masks.shape[-2:], mode='bilinear')\n            loss = criterion(preds, masks) * pos_weight.to(device)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        torch.cuda.empty_cache()\n        \n    return total_loss / len(loader)\n\ndef validate_model(model, loader, device):\n    model.eval()\n    all_preds, all_labels = [], []\n    \n    with torch.no_grad():\n        for imgs, masks in tqdm(loader, desc=\"Validating\"):\n            imgs = imgs.to(device)\n            masks = masks.unsqueeze(1).cpu().numpy()\n            \n            with autocast(enabled=CFG['amp']):\n                preds = torch.sigmoid(model(imgs)).cpu().numpy()\n            \n            if preds.shape[-2:] != masks.shape[-2:]:\n                preds = nn.functional.interpolate(\n                    torch.tensor(preds), \n                    size=masks.shape[-2:], \n                    mode='bilinear'\n                ).numpy()\n            \n            all_preds.append(preds)\n            all_labels.append(masks)\n            torch.cuda.empty_cache()\n            \n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    \n    # Debug mask values\n    print(f\"\\nMask value range: [{np.min(all_labels)}, {np.max(all_labels)}]\")\n    print(f\"Unique mask values: {np.unique(all_labels)}\")\n    \n    # Binarize labels properly\n    all_labels = (all_labels > 0.5).astype(np.uint8)\n    positive_pixels = np.sum(all_labels)\n    \n    if positive_pixels == 0:\n        print(\"Warning: No positive pixels found in validation masks!\")\n        return 0.0, 0.0, 0.0, all_preds, all_labels, 0.5\n    \n    # Custom threshold range as requested\n    thresholds = [0.3, 0.35, 0.39, 0.4, 0.41, 0.42, 0.43, 0.44, 0.45, \n                  0.46, 0.47, 0.48, 0.49, 0.5, 0.51, 0.52, 0.53, 0.54, \n                  0.55, 0.56, 0.57, 0.58, 0.59, 0.6]\n    \n    best_thresh, best_f05 = 0.3, 0\n    \n    for t in thresholds:\n        binarized = (all_preds > t).astype(np.uint8)\n        try:\n            f05 = fbeta_score(all_labels.flatten(), binarized.flatten(), beta=0.5, zero_division=0)\n            if f05 > best_f05:\n                best_thresh, best_f05 = t, f05\n        except:\n            continue\n            \n    final_binarized = (all_preds > best_thresh).astype(np.uint8)\n    precision = precision_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    recall = recall_score(all_labels.flatten(), final_binarized.flatten(), zero_division=0)\n    \n    print(f\"\\nDebug Info - Pred range: [{np.min(all_preds):.3f}, {np.max(all_preds):.3f}]\")\n    print(f\"Positive pixels: {positive_pixels}/{all_labels.size} ({positive_pixels/all_labels.size*100:.4f}%)\")\n    \n    return best_f05, precision, recall, all_preds, all_labels, best_thresh\n\ndef visualize_predictions(inputs, preds, labels, threshold=0.5, num_samples=5):\n    indices = np.random.choice(len(inputs), num_samples, replace=False)\n    \n    plt.figure(figsize=(20, 4*num_samples))\n    for i, idx in enumerate(indices):\n        input_img = inputs[idx][len(inputs[idx])//2]\n        pred = preds[idx][0] > threshold\n        label = labels[idx][0]\n        \n        plt.subplot(num_samples, 3, i*3 + 1)\n        plt.imshow(input_img, cmap='gray')\n        plt.title(f\"Input (Sample {idx})\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 2)\n        plt.imshow(label, cmap='gray')\n        plt.title(\"Ground Truth\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 3, i*3 + 3)\n        plt.imshow(pred, cmap='gray')\n        plt.title(\"Prediction\")\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    try:\n        # Initialize model\n        model = InkDetector().to(device)\n        print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n        \n        # Load datasets\n        print(\"Loading training data...\")\n        train_case, train_volume, train_mask = CFG['train_case']\n        train_dataset = VesuviusDataset(train_volume, train_mask, is_train=True)\n        print(f\"Found {len(train_dataset)} training samples\")\n        \n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=CFG['train_batch_size'],\n            shuffle=True,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        print(\"Loading test data...\")\n        test_case, test_volume, test_mask = CFG['test_case']\n        test_dataset = VesuviusDataset(test_volume, test_mask, is_train=False)\n        print(f\"Found {len(test_dataset)} test samples\")\n        \n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG['test_batch_size'],\n            shuffle=False,\n            pin_memory=True,\n            num_workers=2\n        )\n        \n        # Verify we have data\n        if len(train_dataset) == 0 or len(test_dataset) == 0:\n            raise ValueError(\"No samples found in datasets\")\n        \n        # Get class weights\n        pos_weight = get_class_weights(train_loader)\n        print(f\"Positive weight: {pos_weight.item():.2f}\")\n        \n        # Training setup\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'], weight_decay=1e-4)\n        criterion = ImprovedLoss(alpha=0.8)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, \n            max_lr=CFG['lr'],\n            steps_per_epoch=len(train_loader),\n            epochs=CFG['epochs'],\n            pct_start=0.3\n        )\n        scaler = GradScaler(enabled=CFG['amp'])\n        \n        best_f05 = 0\n        no_improve = 0\n        \n        for epoch in range(CFG['epochs']):\n            print(f\"\\nEpoch {epoch + 1}/{CFG['epochs']}\")\n            \n            # Train\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, device, pos_weight, scaler)\n            scheduler.step()\n            print(f\"Train Loss: {loss:.4f}\")\n            print(f\"Current LR: {optimizer.param_groups[0]['lr']:.2e}\")\n            \n            # Validate\n            print(\"\\nEvaluating...\")\n            f05, prec, rec, all_preds, all_labels, best_thresh = validate_model(model, test_loader, device)\n            print(f\"Test Metrics - F0.5: {f05:.4f}, Precision: {prec:.4f}, Recall: {rec:.4f}, Threshold: {best_thresh:.2f}\")\n            \n            if f05 > best_f05 + 1e-4:\n                best_f05 = f05\n                no_improve = 0\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(\"Saved new best model\")\n                \n                if f05 > 0:\n                    print(\"\\nVisualizing best predictions...\")\n                    sample_inputs = []\n                    for batch in test_loader:\n                        sample_inputs.extend(batch[0].cpu().numpy())\n                        if len(sample_inputs) >= 20:\n                            break\n                    visualize_predictions(sample_inputs, all_preds, all_labels, threshold=best_thresh)\n            else:\n                no_improve += 1\n                if no_improve >= CFG['early_stop_patience']:\n                    print(f\"\\nEarly stopping after {no_improve} epochs without improvement\")\n                    break\n            \n            torch.cuda.empty_cache()\n        \n        print(\"\\nTraining complete!\")\n        print(f\"Best F0.5 Score: {best_f05:.4f}\")\n        \n    except Exception as e:\n        print(f\"Error: {str(e)}\")\n        print(\"\\nDebug Info:\")\n        print(f\"Train volume path: {CFG['train_case'][1]}\")\n        print(f\"Train mask path: {CFG['train_case'][2]}\")\n        print(f\"Test volume path: {CFG['test_case'][1]}\")\n        print(f\"Test mask path: {CFG['test_case'][2]}\")\n        \n        # Check if files exist\n        print(\"\\nFile checks:\")\n        for case in [CFG['train_case'], CFG['test_case']]:\n            vol_path = case[1]\n            mask_path = case[2]\n            print(f\"\\nVolume: {vol_path}\")\n            print(f\"Files exist: {os.path.exists(vol_path)}\")\n            if os.path.exists(vol_path):\n                print(f\"TIFF files found: {len(glob(os.path.join(vol_path, '*.tif')))}\")\n            \n            print(f\"\\nMask: {mask_path}\")\n            print(f\"Exists: {os.path.exists(mask_path)}\")\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                print(f\"Mask values - min: {np.min(mask)}, max: {np.max(mask)}, mean: {np.mean(mask)}\")\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Welcome","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import f1_score, jaccard_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport gc\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nPRETRAINED_PATH = \"/kaggle/input/3d-unet/3d_unet_3dunet_b3_final_swa.pth\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)  # From 00.tif to 64.tif\nBATCH_SIZE = 2\nEPOCHS = 10\nLR = 0.0001\nNUM_SPLITS = 8  # Split volume 1 into 8 parts\n\n# Corrected UNet3D class matching pretrained weights\nclass UNet3D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet3D, self).__init__()\n        \n        def SingleConv(in_channels, out_channels):\n            return nn.Sequential(\n                nn.GroupNorm(16, in_channels),\n                nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1)\n            )\n        \n        def DoubleConv(in_channels, out_channels):\n            return nn.Sequential(\n                SingleConv(in_channels, out_channels),\n                nn.ReLU(inplace=True),\n                SingleConv(out_channels, out_channels),\n                nn.ReLU(inplace=True)\n            )\n        \n        # Encoders\n        self.encoders = nn.ModuleList([\n            DoubleConv(in_channels, 64),\n            DoubleConv(64, 128),\n            DoubleConv(128, 256),\n            DoubleConv(256, 512)\n        ])\n        \n        self.pool = nn.MaxPool3d(2)\n        \n        # Decoders\n        self.decoders = nn.ModuleList([\n            DoubleConv(512 + 256, 256),\n            DoubleConv(256 + 128, 128),\n            DoubleConv(128 + 64, 64)\n        ])\n        \n        # Upsampling\n        self.upconvs = nn.ModuleList([\n            nn.ConvTranspose3d(512, 256, kernel_size=2, stride=2),\n            nn.ConvTranspose3d(256, 128, kernel_size=2, stride=2),\n            nn.ConvTranspose3d(128, 64, kernel_size=2, stride=2)\n        ])\n        \n        # Output\n        self.conv_out = nn.Conv3d(64, out_channels, kernel_size=1)\n    \n    def forward(self, x):\n        # Encoder path\n        encoder_features = []\n        for encoder in self.encoders:\n            x = encoder(x)\n            encoder_features.append(x)\n            x = self.pool(x)\n        \n        # Decoder path\n        for i, (decoder, upconv) in enumerate(zip(self.decoders, self.upconvs)):\n            x = upconv(x)\n            x = torch.cat([x, encoder_features[-(i+2)]], dim=1)\n            x = decoder(x)\n        \n        return torch.sigmoid(self.conv_out(x))\n\n# Dataset class that splits volume into parts\nclass SplitVolumeDataset(Dataset):\n    def __init__(self, volume_path, label_path, resize_shape=None, split_index=0, total_splits=8):\n        self.volume_path = volume_path\n        self.label_path = label_path\n        self.resize_shape = resize_shape\n        self.split_index = split_index\n        self.total_splits = total_splits\n        \n        # Load and process label\n        self.label = np.array(Image.open(label_path)) > 0\n        if resize_shape:\n            self.label = cv2.resize(self.label.astype(np.float32), resize_shape,\n                                 interpolation=cv2.INTER_NEAREST)\n            self.label = self.label > 0.5\n        \n        # Calculate split boundaries\n        self.slice_indices = self._get_split_indices()\n        \n    def _get_split_indices(self):\n        total_slices = SLICE_RANGE[1] - SLICE_RANGE[0] + 1\n        slices_per_split = total_slices // self.total_splits\n        start = SLICE_RANGE[0] + self.split_index * slices_per_split\n        end = start + slices_per_split\n        \n        # Handle remainder slices\n        if self.split_index == self.total_splits - 1:\n            end = SLICE_RANGE[1] + 1\n            \n        return list(range(start, end))\n    \n    def __len__(self):\n        return len(self.slice_indices)\n    \n    def __getitem__(self, idx):\n        slice_idx = self.slice_indices[idx]\n        slice_path = os.path.join(self.volume_path, f\"{slice_idx:02d}.tif\")\n        img = tiff.imread(slice_path)\n        \n        if self.resize_shape:\n            img = cv2.resize(img, self.resize_shape, interpolation=cv2.INTER_AREA)\n        \n        img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n        \n        return (\n            torch.FloatTensor(img).unsqueeze(0),  # (1, H, W)\n            torch.FloatTensor(self.label).unsqueeze(0)  # (1, H, W)\n        )\n\n# Function to create all splits\ndef create_all_splits(volume_path, label_path, resize_shape, num_splits):\n    return [SplitVolumeDataset(volume_path, label_path, resize_shape, i, num_splits) \n            for i in range(num_splits)]\n\n# Load pretrained model\ndef load_pretrained():\n    model = UNet3D().to(DEVICE)\n    try:\n        state_dict = torch.load(PRETRAINED_PATH, map_location=DEVICE)\n        \n        # Remove 'model.' prefix and handle GroupNorm naming\n        new_state_dict = {}\n        for k, v in state_dict.items():\n            k = k.replace('model.', '')\n            k = k.replace('basic_SingleConv1', '0.0')\n            k = k.replace('basic_SingleConv2', '0.3')\n            new_state_dict[k] = v\n        \n        model.load_state_dict(new_state_dict, strict=False)\n        print(\"Successfully loaded pretrained model (some layers may not match exactly)\")\n        return model\n    except Exception as e:\n        print(f\"Error loading pretrained model: {e}\")\n        return None\n\n# Training and evaluation function\ndef train_and_evaluate():\n    # Prepare data\n    train_vol = os.path.join(DATA_DIR, \"train/1/surface_volume\")\n    train_label = os.path.join(DATA_DIR, \"train/1/inklabels.png\")\n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create splits\n    train_splits = create_all_splits(train_vol, train_label, RESIZE_SHAPE, NUM_SPLITS)\n    test_dataset = SplitVolumeDataset(test_vol, test_label, RESIZE_SHAPE)\n    \n    # Initialize model\n    model = load_pretrained()\n    if model is None:\n        raise RuntimeError(\"Failed to load pretrained model\")\n    \n    criterion = nn.BCELoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    \n    # Training loop over splits\n    for split_idx, split_dataset in enumerate(train_splits):\n        print(f\"\\n=== Training on split {split_idx+1}/{NUM_SPLITS} ===\")\n        \n        train_loader = DataLoader(\n            split_dataset,\n            batch_size=BATCH_SIZE,\n            shuffle=True,\n            pin_memory=True\n        )\n        \n        for epoch in range(EPOCHS):\n            model.train()\n            epoch_loss = 0.0\n            \n            for slices, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n                slices = slices.to(DEVICE)  # (B, 1, H, W)\n                labels = labels.to(DEVICE)  # (B, 1, H, W)\n                \n                # Add dummy depth dimension\n                slices = slices.unsqueeze(2)  # (B, 1, 1, H, W)\n                \n                optimizer.zero_grad()\n                outputs = model(slices)  # (B, 1, 1, H, W)\n                loss = criterion(outputs.squeeze(2), labels)\n                loss.backward()\n                optimizer.step()\n                \n                epoch_loss += loss.item()\n            \n            print(f\"Split {split_idx+1} | Epoch {epoch+1} Loss: {epoch_loss/len(train_loader):.4f}\")\n            \n            # Clear memory\n            torch.cuda.empty_cache()\n            gc.collect()\n    \n    # Evaluation\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    model.eval()\n    \n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for slices, labels in test_loader:\n            slices = slices.to(DEVICE).unsqueeze(2)  # (1, 1, 1, H, W)\n            outputs = model(slices)  # (1, 1, 1, H, W)\n            \n            pred = (outputs.squeeze().cpu().numpy() > 0.5).astype(np.uint8)\n            label = labels.squeeze().cpu().numpy().astype(np.uint8)\n            \n            all_preds.append(pred)\n            all_labels.append(label)\n            \n            # Visualize first sample\n            if len(all_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                plt.subplot(1, 3, 1)\n                plt.imshow(slices[0, 0, 0].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 2)\n                plt.imshow(label, cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_labels = np.concatenate([l.flatten() for l in all_labels])\n    \n    metrics = {\n        'F1': f1_score(all_labels, all_preds),\n        'IoU': jaccard_score(all_labels, all_preds),\n        'Precision': precision_score(all_labels, all_preds),\n        'Recall': recall_score(all_labels, all_preds),\n        'Ink Coverage (GT)': f\"{all_labels.mean():.2%}\",\n        'Ink Coverage (Pred)': f\"{all_preds.mean():.2%}\"\n    }\n    \n    # Print metrics\n    print(\"\\n\" + \"=\"*50)\n    print(f\"{'Evaluation Metrics':^50}\")\n    print(\"=\"*50)\n    for name, value in metrics.items():\n        if isinstance(value, float):\n            print(f\"{name+':':<20}{value:.4f}\")\n        else:\n            print(f\"{name+':':<20}{value}\")\n    print(\"=\"*50)\n    \n    return model\n\n# Run the complete pipeline\nif __name__ == \"__main__\":\n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Run training and evaluation\n    model = train_and_evaluate()\n    \n    # Save fine-tuned model\n    torch.save(model.state_dict(), \"fine_tuned_split_model.pth\")\n    print(\"\\nFine-tuned model saved successfully!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import f1_score, jaccard_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport gc\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nPRETRAINED_PATH = \"/kaggle/input/3d-unet/3d_unet_3dunet_b3_final_swa.pth\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)  # From 00.tif to 64.tif\nBATCH_SIZE = 2\nEPOCHS = 10\nLR = 0.0001\nNUM_SPLITS = 8  # Split volume 1 into 8 parts\nNUM_GROUPS = 16  # For GroupNorm\n\n# Corrected UNet3D class matching pretrained weights\nclass UNet3D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet3D, self).__init__()\n        \n        def SingleConv(in_channels, out_channels):\n            # Ensure channels are divisible by num_groups\n            groups = min(NUM_GROUPS, in_channels)\n            return nn.Sequential(\n                nn.GroupNorm(groups, in_channels),\n                nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1)\n            )\n        \n        def DoubleConv(in_channels, out_channels):\n            return nn.Sequential(\n                SingleConv(in_channels, out_channels),\n                nn.ReLU(inplace=True),\n                SingleConv(out_channels, out_channels),\n                nn.ReLU(inplace=True)\n            )\n        \n        # Encoders with proper channel counts for GroupNorm\n        self.encoders = nn.ModuleList([\n            DoubleConv(in_channels, 64),\n            DoubleConv(64, 128),\n            DoubleConv(128, 256),\n            DoubleConv(256, 512)\n        ])\n        \n        self.pool = nn.MaxPool3d(2)\n        \n        # Decoders\n        self.decoders = nn.ModuleList([\n            DoubleConv(512 + 256, 256),\n            DoubleConv(256 + 128, 128),\n            DoubleConv(128 + 64, 64)\n        ])\n        \n        # Upsampling\n        self.upconvs = nn.ModuleList([\n            nn.ConvTranspose3d(512, 256, kernel_size=2, stride=2),\n            nn.ConvTranspose3d(256, 128, kernel_size=2, stride=2),\n            nn.ConvTranspose3d(128, 64, kernel_size=2, stride=2)\n        ])\n        \n        # Output\n        self.conv_out = nn.Conv3d(64, out_channels, kernel_size=1)\n    \n    def forward(self, x):\n        # Encoder path\n        encoder_features = []\n        for encoder in self.encoders:\n            x = encoder(x)\n            encoder_features.append(x)\n            x = self.pool(x)\n        \n        # Decoder path\n        for i, (decoder, upconv) in enumerate(zip(self.decoders, self.upconvs)):\n            x = upconv(x)\n            x = torch.cat([x, encoder_features[-(i+2)]], dim=1)\n            x = decoder(x)\n        \n        return torch.sigmoid(self.conv_out(x))\n\n# Dataset class that splits volume into parts\nclass SplitVolumeDataset(Dataset):\n    def __init__(self, volume_path, label_path, resize_shape=None, split_index=0, total_splits=8):\n        self.volume_path = volume_path\n        self.label_path = label_path\n        self.resize_shape = resize_shape\n        self.split_index = split_index\n        self.total_splits = total_splits\n        \n        # Load and process label\n        self.label = np.array(Image.open(label_path)) > 0\n        if resize_shape:\n            self.label = cv2.resize(self.label.astype(np.float32), resize_shape,\n                                 interpolation=cv2.INTER_NEAREST)\n            self.label = self.label > 0.5\n        \n        # Calculate split boundaries\n        self.slice_indices = self._get_split_indices()\n        \n    def _get_split_indices(self):\n        total_slices = SLICE_RANGE[1] - SLICE_RANGE[0] + 1\n        slices_per_split = total_slices // self.total_splits\n        start = SLICE_RANGE[0] + self.split_index * slices_per_split\n        end = start + slices_per_split\n        \n        # Handle remainder slices\n        if self.split_index == self.total_splits - 1:\n            end = SLICE_RANGE[1] + 1\n            \n        return list(range(start, end))\n    \n    def __len__(self):\n        return len(self.slice_indices)\n    \n    def __getitem__(self, idx):\n        slice_idx = self.slice_indices[idx]\n        slice_path = os.path.join(self.volume_path, f\"{slice_idx:02d}.tif\")\n        img = tiff.imread(slice_path)\n        \n        if self.resize_shape:\n            img = cv2.resize(img, self.resize_shape, interpolation=cv2.INTER_AREA)\n        \n        img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n        \n        return (\n            torch.FloatTensor(img).unsqueeze(0),  # (1, H, W)\n            torch.FloatTensor(self.label).unsqueeze(0)  # (1, H, W)\n        )\n\n# Function to create all splits\ndef create_all_splits(volume_path, label_path, resize_shape, num_splits):\n    return [SplitVolumeDataset(volume_path, label_path, resize_shape, i, num_splits) \n            for i in range(num_splits)]\n\n# Load pretrained model with proper state dict handling\ndef load_pretrained():\n    model = UNet3D().to(DEVICE)\n    try:\n        state_dict = torch.load(PRETRAINED_PATH, map_location=DEVICE)\n        \n        # Create new state dict with matching keys\n        new_state_dict = {}\n        for k, v in state_dict.items():\n            # Remove 'model.' prefix\n            k = k.replace('model.', '')\n            \n            # Handle encoder blocks\n            if 'encoders' in k:\n                parts = k.split('.')\n                layer_num = int(parts[1])\n                block_part = parts[3]\n                \n                if 'basic_SingleConv1' in k:\n                    new_key = f'encoders.{layer_num}.0.{0 if \"groupnorm\" in block_part else 1}'\n                    if 'weight' in block_part:\n                        new_key += '.weight'\n                    else:\n                        new_key += '.bias'\n                elif 'basic_SingleConv2' in k:\n                    new_key = f'encoders.{layer_num}.3.{0 if \"groupnorm\" in block_part else 1}'\n                    if 'weight' in block_part:\n                        new_key += '.weight'\n                    else:\n                        new_key += '.bias'\n                else:\n                    continue\n                    \n                new_state_dict[new_key] = v\n            \n            # Handle other layers\n            elif 'upconvs' in k:\n                parts = k.split('.')\n                layer_num = int(parts[1])\n                new_key = f'upconvs.{layer_num}.{parts[-1]}'\n                new_state_dict[new_key] = v\n            \n            elif 'conv_out' in k:\n                new_state_dict[k] = v\n        \n        # Load state dict\n        model.load_state_dict(new_state_dict, strict=False)\n        print(\"Successfully loaded pretrained model (some layers may not match exactly)\")\n        return model\n    except Exception as e:\n        print(f\"Error loading pretrained model: {e}\")\n        return None\n\n# Training and evaluation function\ndef train_and_evaluate():\n    # Prepare data\n    train_vol = os.path.join(DATA_DIR, \"train/1/surface_volume\")\n    train_label = os.path.join(DATA_DIR, \"train/1/inklabels.png\")\n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create splits\n    train_splits = create_all_splits(train_vol, train_label, RESIZE_SHAPE, NUM_SPLITS)\n    test_dataset = SplitVolumeDataset(test_vol, test_label, RESIZE_SHAPE)\n    \n    # Initialize model\n    model = load_pretrained()\n    if model is None:\n        raise RuntimeError(\"Failed to load pretrained model\")\n    \n    criterion = nn.BCELoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    \n    # Training loop over splits\n    for split_idx, split_dataset in enumerate(train_splits):\n        print(f\"\\n=== Training on split {split_idx+1}/{NUM_SPLITS} ===\")\n        \n        train_loader = DataLoader(\n            split_dataset,\n            batch_size=BATCH_SIZE,\n            shuffle=True,\n            pin_memory=True\n        )\n        \n        for epoch in range(EPOCHS):\n            model.train()\n            epoch_loss = 0.0\n            \n            for slices, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n                slices = slices.to(DEVICE).unsqueeze(2)  # (B, 1, 1, H, W)\n                labels = labels.to(DEVICE)  # (B, 1, H, W)\n                \n                optimizer.zero_grad()\n                outputs = model(slices)  # (B, 1, 1, H, W)\n                loss = criterion(outputs.squeeze(2), labels)\n                loss.backward()\n                optimizer.step()\n                \n                epoch_loss += loss.item()\n            \n            print(f\"Split {split_idx+1} | Epoch {epoch+1} Loss: {epoch_loss/len(train_loader):.4f}\")\n            \n            # Clear memory\n            torch.cuda.empty_cache()\n            gc.collect()\n    \n    # Evaluation\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    model.eval()\n    \n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for slices, labels in test_loader:\n            slices = slices.to(DEVICE).unsqueeze(2)  # (1, 1, 1, H, W)\n            outputs = model(slices)  # (1, 1, 1, H, W)\n            \n            pred = (outputs.squeeze().cpu().numpy() > 0.5).astype(np.uint8)\n            label = labels.squeeze().cpu().numpy().astype(np.uint8)\n            \n            all_preds.append(pred)\n            all_labels.append(label)\n            \n            # Visualize first sample\n            if len(all_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                plt.subplot(1, 3, 1)\n                plt.imshow(slices[0, 0, 0].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 2)\n                plt.imshow(label, cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_labels = np.concatenate([l.flatten() for l in all_labels])\n    \n    metrics = {\n        'F1': f1_score(all_labels, all_preds),\n        'IoU': jaccard_score(all_labels, all_preds),\n        'Precision': precision_score(all_labels, all_preds),\n        'Recall': recall_score(all_labels, all_preds),\n        'Ink Coverage (GT)': f\"{all_labels.mean():.2%}\",\n        'Ink Coverage (Pred)': f\"{all_preds.mean():.2%}\"\n    }\n    \n    # Print metrics\n    print(\"\\n\" + \"=\"*50)\n    print(f\"{'Evaluation Metrics':^50}\")\n    print(\"=\"*50)\n    for name, value in metrics.items():\n        if isinstance(value, float):\n            print(f\"{name+':':<20}{value:.4f}\")\n        else:\n            print(f\"{name+':':<20}{value}\")\n    print(\"=\"*50)\n    \n    return model\n\n# Run the complete pipeline\nif __name__ == \"__main__\":\n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Run training and evaluation\n    model = train_and_evaluate()\n    \n    # Save fine-tuned model\n    torch.save(model.state_dict(), \"fine_tuned_split_model.pth\")\n    print(\"\\nFine-tuned model saved successfully!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Pre_trained fine tuned model U-Net\nimport os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import f1_score, jaccard_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport gc\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nPRETRAINED_PATH = \"/kaggle/input/3d-unet/3d_unet_3dunet_b3_final_swa.pth\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)  # From 00.tif to 64.tif\nBATCH_SIZE = 2\nEPOCHS = 10\nLR = 0.0001\nNUM_SPLITS = 10  # Split volume 1 into 8 parts\n\n# Dataset class that splits volume into parts\nclass SplitVolumeDataset(Dataset):\n    def __init__(self, volume_path, label_path, resize_shape=None, split_index=0, total_splits=8):\n        self.volume_path = volume_path\n        self.label_path = label_path\n        self.resize_shape = resize_shape\n        self.split_index = split_index\n        self.total_splits = total_splits\n        \n        # Load and process label\n        self.label = np.array(Image.open(label_path)) > 0\n        if resize_shape:\n            self.label = cv2.resize(self.label.astype(np.float32), resize_shape,\n                                 interpolation=cv2.INTER_NEAREST)\n            self.label = self.label > 0.5\n        \n        # Calculate split boundaries\n        self.slice_indices = self._get_split_indices()\n        \n    def _get_split_indices(self):\n        total_slices = SLICE_RANGE[1] - SLICE_RANGE[0] + 1\n        slices_per_split = total_slices // self.total_splits\n        start = SLICE_RANGE[0] + self.split_index * slices_per_split\n        end = start + slices_per_split\n        \n        # Handle remainder slices\n        if self.split_index == self.total_splits - 1:\n            end = SLICE_RANGE[1] + 1\n            \n        return list(range(start, end))\n    \n    def __len__(self):\n        return len(self.slice_indices)\n    \n    def __getitem__(self, idx):\n        slice_idx = self.slice_indices[idx]\n        slice_path = os.path.join(self.volume_path, f\"{slice_idx:02d}.tif\")\n        img = tiff.imread(slice_path)\n        \n        if self.resize_shape:\n            img = cv2.resize(img, self.resize_shape, interpolation=cv2.INTER_AREA)\n        \n        img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n        \n        return (\n            torch.FloatTensor(img).unsqueeze(0),  # (1, H, W)\n            torch.FloatTensor(self.label).unsqueeze(0)  # (1, H, W)\n        )\n\n# Function to create all splits\ndef create_all_splits(volume_path, label_path, resize_shape, num_splits):\n    return [SplitVolumeDataset(volume_path, label_path, resize_shape, i, num_splits) \n            for i in range(num_splits)]\n\n# Simplified UNet3D that matches pretrained architecture\nimport torch\nimport torch.nn as nn\n\nclass UNet3D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet3D, self).__init__()\n        \n        # Basic building blocks\n        def SingleConv(in_c, out_c):\n            return nn.Sequential(\n                nn.Conv3d(in_c, out_c, kernel_size=3, padding=1),\n                nn.GroupNorm(8, out_c),\n                nn.ReLU(inplace=True)\n            )\n        \n        # Modified pooling to preserve dimensions\n        self.pool = nn.MaxPool3d(kernel_size=(1, 2, 2), stride=(1, 2, 2))  # Only pool spatial dimensions\n        \n        # Encoders with adjusted channel dimensions\n        self.encoder1 = nn.Sequential(\n            SingleConv(in_channels, 32),  # Reduced from 64\n            SingleConv(32, 32)\n        )\n        self.encoder2 = nn.Sequential(\n            SingleConv(32, 64),  # Reduced from 128\n            SingleConv(64, 64)\n        )\n        self.encoder3 = nn.Sequential(\n            SingleConv(64, 128),  # Reduced from 256\n            SingleConv(128, 128)\n        )\n        \n        # Bottleneck with reduced channels\n        self.bottleneck = nn.Sequential(\n            SingleConv(128, 256),  # Reduced from 512\n            SingleConv(256, 256)\n        )\n        \n        # Decoders with adjusted channels\n        self.upconv3 = nn.ConvTranspose3d(256, 128, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.decoder3 = nn.Sequential(\n            SingleConv(256, 128),\n            SingleConv(128, 128)\n        )\n        \n        self.upconv2 = nn.ConvTranspose3d(128, 64, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.decoder2 = nn.Sequential(\n            SingleConv(128, 64),\n            SingleConv(64, 64)\n        )\n        \n        self.upconv1 = nn.ConvTranspose3d(64, 32, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.decoder1 = nn.Sequential(\n            SingleConv(64, 32),\n            SingleConv(32, 32)\n        )\n        \n        # Output\n        self.conv_out = nn.Conv3d(32, out_channels, kernel_size=1)\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.encoder1(x)\n        enc2 = self.encoder2(self.pool(enc1))\n        enc3 = self.encoder3(self.pool(enc2))\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc3))\n        \n        # Decoder\n        dec3 = self.upconv3(bottleneck)\n        dec3 = torch.cat([dec3, enc3], dim=1)\n        dec3 = self.decoder3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = torch.cat([dec2, enc2], dim=1)\n        dec2 = self.decoder2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = torch.cat([dec1, enc1], dim=1)\n        dec1 = self.decoder1(dec1)\n        \n        return torch.sigmoid(self.conv_out(dec1))\n\n# Modified load_pretrained function\ndef load_pretrained():\n    model = UNet3D().to(DEVICE)\n    try:\n        state_dict = torch.load(PRETRAINED_PATH, map_location=DEVICE)\n        \n        # Create new state dict with matching keys\n        new_state_dict = {}\n        for k, v in state_dict.items():\n            k = k.replace('module.', '')\n            \n            # Handle channel dimension mismatches\n            if v.ndim == 5:  # Conv3d weights\n                in_c, out_c = v.shape[:2]\n                if in_c > 32 or out_c > 256:  # Skip weights that won't fit\n                    continue\n            \n            # Simple key mapping (won't match all layers due to architecture changes)\n            if 'encoders.0' in k:\n                k = k.replace('encoders.0', 'encoder1')\n            elif 'encoders.1' in k:\n                k = k.replace('encoders.1', 'encoder2')\n            elif 'encoders.2' in k:\n                k = k.replace('encoders.2', 'encoder3')\n            elif 'decoders.0' in k:\n                k = k.replace('decoders.0', 'decoder3')\n            elif 'decoders.1' in k:\n                k = k.replace('decoders.1', 'decoder2')\n            elif 'decoders.2' in k:\n                k = k.replace('decoders.2', 'decoder1')\n            elif 'upconvs.0' in k:\n                k = k.replace('upconvs.0', 'upconv3')\n            elif 'upconvs.1' in k:\n                k = k.replace('upconvs.1', 'upconv2')\n            elif 'upconvs.2' in k:\n                k = k.replace('upconvs.2', 'upconv1')\n            \n            new_state_dict[k] = v\n        \n        # Load state dict\n        model.load_state_dict(new_state_dict, strict=False)\n        print(\"Successfully loaded compatible pretrained weights\")\n        return model\n    except Exception as e:\n        print(f\"Error loading pretrained model: {e}\")\n        print(\"Initializing new model instead\")\n        return UNet3D().to(DEVICE)\n\n# Training and evaluation function\ndef train_and_evaluate():\n    # Prepare data\n    train_vol = os.path.join(DATA_DIR, \"train/1/surface_volume\")\n    train_label = os.path.join(DATA_DIR, \"train/1/inklabels.png\")\n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create splits\n    train_splits = create_all_splits(train_vol, train_label, RESIZE_SHAPE, NUM_SPLITS)\n    test_dataset = SplitVolumeDataset(test_vol, test_label, RESIZE_SHAPE)\n    \n    # Initialize model\n    model = load_pretrained()\n    if model is None:\n        # If pretrained fails, initialize fresh model\n        model = UNet3D().to(DEVICE)\n        print(\"Initialized new model (pretrained weights not used)\")\n    \n    criterion = nn.BCELoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    \n    # Training loop over splits\n    for split_idx, split_dataset in enumerate(train_splits):\n        print(f\"\\n=== Training on split {split_idx+1}/{NUM_SPLITS} ===\")\n        \n        train_loader = DataLoader(\n            split_dataset,\n            batch_size=BATCH_SIZE,\n            shuffle=True,\n            pin_memory=True\n        )\n        \n        for epoch in range(EPOCHS):\n            model.train()\n            epoch_loss = 0.0\n            \n            for slices, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n                slices = slices.to(DEVICE).unsqueeze(2)  # (B, 1, 1, H, W)\n                labels = labels.to(DEVICE)  # (B, 1, H, W)\n                \n                optimizer.zero_grad()\n                outputs = model(slices)  # (B, 1, 1, H, W)\n                loss = criterion(outputs.squeeze(2), labels)\n                loss.backward()\n                optimizer.step()\n                \n                epoch_loss += loss.item()\n            \n            print(f\"Split {split_idx+1} | Epoch {epoch+1} Loss: {epoch_loss/len(train_loader):.4f}\")\n            \n            # Clear memory\n            torch.cuda.empty_cache()\n            gc.collect()\n    \n    # Evaluation\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    model.eval()\n    \n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for slices, labels in test_loader:\n            slices = slices.to(DEVICE).unsqueeze(2)  # (1, 1, 1, H, W)\n            outputs = model(slices)  # (1, 1, 1, H, W)\n            \n            pred = (outputs.squeeze().cpu().numpy() > 0.5).astype(np.uint8)\n            label = labels.squeeze().cpu().numpy().astype(np.uint8)\n            \n            all_preds.append(pred)\n            all_labels.append(label)\n            \n            # Visualize first sample\n            if len(all_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                plt.subplot(1, 3, 1)\n                plt.imshow(slices[0, 0, 0].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 2)\n                plt.imshow(label, cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_labels = np.concatenate([l.flatten() for l in all_labels])\n    \n    metrics = {\n        'F1': f1_score(all_labels, all_preds),\n        'IoU': jaccard_score(all_labels, all_preds),\n        'Precision': precision_score(all_labels, all_preds),\n        'Recall': recall_score(all_labels, all_preds),\n        'Ink Coverage (GT)': f\"{all_labels.mean():.2%}\",\n        'Ink Coverage (Pred)': f\"{all_preds.mean():.2%}\"\n    }\n    \n    # Print metrics\n    print(\"\\n\" + \"=\"*50)\n    print(f\"{'Evaluation Metrics':^50}\")\n    print(\"=\"*50)\n    for name, value in metrics.items():\n        if isinstance(value, float):\n            print(f\"{name+':':<20}{value:.4f}\")\n        else:\n            print(f\"{name+':':<20}{value}\")\n    print(\"=\"*50)\n    \n    return model\n\n# Run the complete pipeline\nif __name__ == \"__main__\":\n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Run training and evaluation\n    model = train_and_evaluate()\n    \n    # Save fine-tuned model\n    torch.save(model.state_dict(), \"fine_tuned_split_model.pth\")\n    print(\"\\nFine-tuned model saved successfully!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score, jaccard_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport random\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom skimage import morphology\nimport gc\n \n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)\nBATCH_SIZE = 8\nEPOCHS = 100\nLR = 0.001\nVAL_SPLIT = 0.2\n\n# Enhanced 3D U-Net with Residual Connections\nclass InkDetectionModel(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(InkDetectionModel, self).__init__()\n        \n        def conv_block(in_c, out_c):\n            return nn.Sequential(\n                nn.Conv3d(in_c, out_c, kernel_size=3, padding=1),\n                nn.InstanceNorm3d(out_c),\n                nn.LeakyReLU(0.2),\n                nn.Conv3d(out_c, out_c, kernel_size=3, padding=1),\n                nn.InstanceNorm3d(out_c),\n                nn.LeakyReLU(0.2)\n            )\n        \n        # Encoders\n        self.enc1 = conv_block(in_channels, 64)\n        self.enc2 = conv_block(64, 128)\n        self.enc3 = conv_block(128, 256)\n        self.enc4 = conv_block(256, 512)\n        \n        self.pool = nn.MaxPool3d((1, 2, 2))  # Only pool spatial dimensions\n        \n        # Bottleneck with residual\n        self.bottleneck = nn.Sequential(\n            nn.Conv3d(512, 1024, kernel_size=3, padding=1),\n            nn.InstanceNorm3d(1024),\n            nn.LeakyReLU(0.2),\n            nn.Conv3d(1024, 1024, kernel_size=3, padding=1),\n            nn.InstanceNorm3d(1024)\n        )\n        self.res_conv = nn.Conv3d(512, 1024, kernel_size=1)\n        \n        # Decoders\n        self.upconv4 = nn.ConvTranspose3d(1024, 512, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec4 = conv_block(1024, 512)\n        \n        self.upconv3 = nn.ConvTranspose3d(512, 256, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec3 = conv_block(512, 256)\n        \n        self.upconv2 = nn.ConvTranspose3d(256, 128, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec2 = conv_block(256, 128)\n        \n        self.upconv1 = nn.ConvTranspose3d(128, 64, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec1 = conv_block(128, 64)\n        \n        # Output\n        self.conv_out = nn.Conv3d(64, out_channels, kernel_size=1)\n        \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        e4 = self.enc4(self.pool(e3))\n        \n        # Bottleneck with residual\n        b = self.bottleneck(self.pool(e4)) + self.res_conv(self.pool(e4))\n        \n        # Decoder\n        d4 = self.upconv4(b)\n        d4 = torch.cat([d4, e4], dim=1)\n        d4 = self.dec4(d4)\n        \n        d3 = self.upconv3(d4)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.upconv2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.upconv1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        return torch.sigmoid(self.conv_out(d1))\n\n# Augmented Dataset\nclass InkDataset(Dataset):\n    def __init__(self, volume_paths, label_paths, resize_shape=None, train=True):\n        self.volumes = []\n        self.labels = []\n        self.resize_shape = resize_shape\n        self.train = train\n        \n        for vol_path, lbl_path in zip(volume_paths, label_paths):\n            # Load volume slices\n            volume = []\n            for i in range(SLICE_RANGE[0], SLICE_RANGE[1] + 1):\n                img = tiff.imread(os.path.join(vol_path, f\"{i:02d}.tif\"))\n                if resize_shape:\n                    img = cv2.resize(img, resize_shape, interpolation=cv2.INTER_AREA)\n                img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n                volume.append(img)\n            self.volumes.append(np.stack(volume, axis=0))\n            \n            # Load label\n            label = np.array(Image.open(lbl_path)) > 0\n            if resize_shape:\n                label = cv2.resize(label.astype(np.float32), resize_shape, \n                                 interpolation=cv2.INTER_NEAREST)\n                label = label > 0.5\n            self.labels.append(label)\n    \n    def __len__(self):\n        return len(self.volumes)\n    \n    def __getitem__(self, idx):\n        img = self.volumes[idx]\n        label = self.labels[idx]\n        \n        if self.train:\n            # Random horizontal flip\n            if random.random() > 0.5:\n                img = np.flip(img, axis=2)\n                label = np.flip(label, axis=1)\n            \n            # Random vertical flip\n            if random.random() > 0.5:\n                img = np.flip(img, axis=1)\n                label = np.flip(label, axis=0)\n            \n            # Random gamma correction\n            gamma = random.uniform(0.7, 1.3)\n            img = img ** gamma\n        \n        return (\n            torch.FloatTensor(img).unsqueeze(0),  # (1, depth, H, W)\n            torch.FloatTensor(label).unsqueeze(0)  # (1, H, W)\n        )\n\n# Post-processing\ndef postprocess(prediction, threshold=0.5, min_size=32):\n    prediction = prediction.squeeze()\n    binary = (prediction > threshold).astype(bool)\n    \n    # Remove small objects\n    cleaned = morphology.remove_small_objects(binary, min_size=min_size)\n    \n    # Fill small holes\n    cleaned = morphology.remove_small_holes(cleaned, area_threshold=min_size)\n    \n    return cleaned.astype(np.float32)\n\n# Training function\ndef train_model():\n    # Prepare data paths\n    train_volumes = [\n        os.path.join(DATA_DIR, \"train/1/surface_volume\"),\n        os.path.join(DATA_DIR, \"train/2/surface_volume\")\n    ]\n    train_labels = [\n        os.path.join(DATA_DIR, \"train/1/inklabels.png\"),\n        os.path.join(DATA_DIR, \"train/2/inklabels.png\")\n    ]\n    \n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create datasets\n    full_dataset = InkDataset(train_volumes, train_labels, RESIZE_SHAPE, train=True)\n    train_size = int((1 - VAL_SPLIT) * len(full_dataset))\n    val_size = len(full_dataset) - train_size\n    train_dataset, val_dataset = torch.utils.data.random_split(full_dataset, [train_size, val_size])\n    \n    test_dataset = InkDataset([test_vol], [test_label], RESIZE_SHAPE, train=False)\n    \n    # Create dataloaders\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    # Initialize model\n    model = InkDetectionModel().to(DEVICE)\n    \n    # Loss and optimizer\n    pos_weight = torch.tensor([9.0]).to(DEVICE)  # Adjust based on your class imbalance\n    criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS*len(train_loader))\n    \n    # Training loop\n    best_f1 = 0\n    for epoch in range(EPOCHS):\n        model.train()\n        epoch_loss = 0\n        \n        for slices, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n            slices = slices.to(DEVICE)  # (B, 1, depth, H, W)\n            labels = labels.to(DEVICE)  # (B, 1, H, W)\n            \n            optimizer.zero_grad()\n            outputs = model(slices)  # (B, 1, depth, H, W)\n            \n            # Use middle slice for training\n            mid_slice = outputs.shape[2] // 2\n            loss = criterion(outputs[:, :, mid_slice], labels)\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            \n            epoch_loss += loss.item()\n        \n        # Validation\n        model.eval()\n        val_preds = []\n        val_labels = []\n        \n        with torch.no_grad():\n            for slices, labels in val_loader:\n                slices = slices.to(DEVICE)\n                outputs = model(slices)\n                \n                mid_slice = outputs.shape[2] // 2\n                pred = postprocess(outputs[:, :, mid_slice].cpu().numpy())\n                val_preds.append(pred)\n                val_labels.append(labels.squeeze().cpu().numpy())\n        \n        # Calculate metrics\n        val_preds = np.concatenate([p.flatten() for p in val_preds])\n        val_labels = np.concatenate([l.flatten() for l in val_labels])\n        \n        val_f1 = f1_score(val_labels, val_preds)\n        val_iou = jaccard_score(val_labels, val_preds)\n        \n        print(f\"\\nEpoch {epoch+1} | Loss: {epoch_loss/len(train_loader):.4f}\")\n        print(f\"Val F1: {val_f1:.4f} | Val IoU: {val_iou:.4f}\")\n        \n        # Save best model\n        if val_f1 > best_f1:\n            best_f1 = val_f1\n            torch.save(model.state_dict(), \"best_model.pth\")\n            print(f\"New best model saved with F1: {best_f1:.4f}\")\n        \n        # Early stopping condition\n        if best_f1 > 0.8:\n            print(\"Target F1 reached!\")\n            break\n    \n    # Final evaluation on test set\n    model.load_state_dict(torch.load(\"best_model.pth\"))\n    model.eval()\n    \n    test_preds = []\n    test_labels = []\n    \n    with torch.no_grad():\n        for slices, labels in test_loader:\n            slices = slices.to(DEVICE)\n            outputs = model(slices)\n            \n            mid_slice = outputs.shape[2] // 2\n            pred = postprocess(outputs[:, :, mid_slice].cpu().numpy())\n            test_preds.append(pred)\n            test_labels.append(labels.squeeze().cpu().numpy())\n            \n            # Visualize first sample\n            if len(test_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                plt.subplot(1, 3, 1)\n                plt.imshow(slices[0, 0, mid_slice].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 2)\n                plt.imshow(labels[0, 0].cpu(), cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate test metrics\n    test_preds = np.concatenate([p.flatten() for p in test_preds])\n    test_labels = np.concatenate([l.flatten() for l in test_labels])\n    \n    metrics = {\n        'F1': f1_score(test_labels, test_preds),\n        'IoU': jaccard_score(test_labels, test_preds),\n        'Precision': precision_score(test_labels, test_preds),\n        'Recall': recall_score(test_labels, test_preds),\n        'Ink Coverage (GT)': f\"{test_labels.mean():.2%}\",\n        'Ink Coverage (Pred)': f\"{test_preds.mean():.2%}\"\n    }\n    \n    print(\"\\n\" + \"=\"*50)\n    print(f\"{'Test Metrics':^50}\")\n    print(\"=\"*50)\n    for name, value in metrics.items():\n        if isinstance(value, float):\n            print(f\"{name+':':<20}{value:.4f}\")\n        else:\n            print(f\"{name+':':<20}{value}\")\n    print(\"=\"*50)\n    \n    return model\n\n# Run training\nif __name__ == \"__main__\":\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    model = train_model()\n    print(\"Training completed!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score, jaccard_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport random\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom skimage import morphology\nimport gc\nimport warnings\n\n# Suppress PIL decompression bomb warning\nwarnings.simplefilter('ignore', Image.DecompressionBombWarning)\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)  # Reduced size to handle memory\nSLICE_RANGE = (0, 64)  # From 00.tif to 64.tif\nBATCH_SIZE = 4  # Reduced batch size\nEPOCHS = 50\nLR = 0.001\nVAL_SPLIT = 0.2\n\n# Enhanced 3D U-Net with proper memory handling\nclass InkDetectionModel(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(InkDetectionModel, self).__init__()\n        \n        def conv_block(in_c, out_c):\n            return nn.Sequential(\n                nn.Conv3d(in_c, out_c, kernel_size=3, padding=1),\n                nn.InstanceNorm3d(out_c),\n                nn.LeakyReLU(0.2),\n                nn.Conv3d(out_c, out_c, kernel_size=3, padding=1),\n                nn.InstanceNorm3d(out_c),\n                nn.LeakyReLU(0.2)\n            )\n        \n        # Encoders\n        self.enc1 = conv_block(in_channels, 32)  # Reduced channels\n        self.enc2 = conv_block(32, 64)\n        self.enc3 = conv_block(64, 128)\n        \n        self.pool = nn.MaxPool3d((1, 2, 2))  # Only pool spatial dimensions\n        \n        # Bottleneck\n        self.bottleneck = conv_block(128, 256)\n        \n        # Decoders\n        self.upconv3 = nn.ConvTranspose3d(256, 128, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec3 = conv_block(256, 128)\n        \n        self.upconv2 = nn.ConvTranspose3d(128, 64, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec2 = conv_block(128, 64)\n        \n        self.upconv1 = nn.ConvTranspose3d(64, 32, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec1 = conv_block(64, 32)\n        \n        # Output\n        self.conv_out = nn.Conv3d(32, out_channels, kernel_size=1)\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        \n        # Bottleneck\n        b = self.bottleneck(self.pool(e3))\n        \n        # Decoder\n        d3 = self.upconv3(b)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.upconv2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.upconv1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        return torch.sigmoid(self.conv_out(d1))\n\n# Fixed Dataset class with proper array handling\nclass InkDataset(Dataset):\n    def __init__(self, volume_paths, label_paths, resize_shape=None, train=True):\n        self.volume_paths = volume_paths\n        self.label_paths = label_paths\n        self.resize_shape = resize_shape\n        self.train = train\n    \n    def __len__(self):\n        return len(self.volume_paths)\n    \n    def __getitem__(self, idx):\n        # Load volume slices with proper memory handling\n        volume = []\n        for i in range(SLICE_RANGE[0], SLICE_RANGE[1] + 1):\n            img = tiff.imread(os.path.join(self.volume_paths[idx], f\"{i:02d}.tif\"))\n            if self.resize_shape:\n                img = cv2.resize(img, self.resize_shape, interpolation=cv2.INTER_AREA)\n            img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n            volume.append(img)\n        img_volume = np.stack(volume, axis=0)\n        \n        # Load label with proper memory handling\n        label = np.array(Image.open(self.label_paths[idx]).convert('L')) > 0\n        if self.resize_shape:\n            label = cv2.resize(label.astype(np.float32), self.resize_shape, \n                             interpolation=cv2.INTER_NEAREST)\n            label = label > 0.5\n        \n        # Ensure arrays are contiguous\n        img_volume = np.ascontiguousarray(img_volume)\n        label = np.ascontiguousarray(label)\n        \n        # Data augmentation\n        if self.train:\n            # Random horizontal flip\n            if random.random() > 0.5:\n                img_volume = np.flip(img_volume, axis=2).copy()\n                label = np.flip(label, axis=1).copy()\n            \n            # Random vertical flip\n            if random.random() > 0.5:\n                img_volume = np.flip(img_volume, axis=1).copy()\n                label = np.flip(label, axis=0).copy()\n            \n            # Random gamma correction\n            gamma = random.uniform(0.8, 1.2)\n            img_volume = np.power(img_volume, gamma)\n        \n        return (\n            torch.FloatTensor(img_volume).unsqueeze(0),  # (1, depth, H, W)\n            torch.FloatTensor(label).unsqueeze(0)  # (1, H, W)\n        )\n\n# Post-processing\ndef postprocess(prediction, threshold=0.5, min_size=32):\n    prediction = prediction.squeeze()\n    binary = (prediction > threshold).astype(bool)\n    \n    # Remove small objects\n    cleaned = morphology.remove_small_objects(binary, min_size=min_size)\n    \n    # Fill small holes\n    cleaned = morphology.remove_small_holes(cleaned, area_threshold=min_size)\n    \n    return cleaned.astype(np.float32)\n\n# Training function\ndef train_model():\n    # Prepare data paths\n    train_volumes = [\n        os.path.join(DATA_DIR, \"train/1/surface_volume\"),\n        os.path.join(DATA_DIR, \"train/2/surface_volume\")\n    ]\n    train_labels = [\n        os.path.join(DATA_DIR, \"train/1/inklabels.png\"),\n        os.path.join(DATA_DIR, \"train/2/inklabels.png\")\n    ]\n    \n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create datasets\n    train_dataset = InkDataset(train_volumes[:1], train_labels[:1], RESIZE_SHAPE, train=True)  # Start with just volume 1\n    test_dataset = InkDataset([test_vol], [test_label], RESIZE_SHAPE, train=False)\n    \n    # Split train into train/val\n    train_size = int((1 - VAL_SPLIT) * len(train_dataset))\n    val_size = len(train_dataset) - train_size\n    train_dataset, val_dataset = torch.utils.data.random_split(train_dataset, [train_size, val_size])\n    \n    # Create dataloaders\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    # Initialize model\n    model = InkDetectionModel().to(DEVICE)\n    \n    # Loss and optimizer\n    pos_weight = torch.tensor([9.0]).to(DEVICE)  # Adjust based on your class imbalance\n    criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS*len(train_loader))\n    \n    # Training loop\n    best_f1 = 0\n    for epoch in range(EPOCHS):\n        model.train()\n        epoch_loss = 0\n        \n        for slices, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n            slices = slices.to(DEVICE)  # (B, 1, depth, H, W)\n            labels = labels.to(DEVICE)  # (B, 1, H, W)\n            \n            optimizer.zero_grad()\n            outputs = model(slices)  # (B, 1, depth, H, W)\n            \n            # Use middle slice for training\n            mid_slice = outputs.shape[2] // 2\n            loss = criterion(outputs[:, :, mid_slice], labels)\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            \n            epoch_loss += loss.item()\n        \n        # Validation\n        model.eval()\n        val_preds = []\n        val_labels = []\n        \n        with torch.no_grad():\n            for slices, labels in val_loader:\n                slices = slices.to(DEVICE)\n                outputs = model(slices)\n                \n                mid_slice = outputs.shape[2] // 2\n                pred = postprocess(outputs[:, :, mid_slice].cpu().numpy())\n                val_preds.append(pred)\n                val_labels.append(labels.squeeze().cpu().numpy())\n        \n        # Calculate metrics\n        val_preds = np.concatenate([p.flatten() for p in val_preds])\n        val_labels = np.concatenate([l.flatten() for l in val_labels])\n        \n        val_f1 = f1_score(val_labels, val_preds)\n        val_iou = jaccard_score(val_labels, val_preds)\n        \n        print(f\"\\nEpoch {epoch+1} | Loss: {epoch_loss/len(train_loader):.4f}\")\n        print(f\"Val F1: {val_f1:.4f} | Val IoU: {val_iou:.4f}\")\n        \n        # Save best model\n        if val_f1 > best_f1:\n            best_f1 = val_f1\n            torch.save(model.state_dict(), \"best_model.pth\")\n            print(f\"New best model saved with F1: {best_f1:.4f}\")\n        \n        # Early stopping condition\n        if best_f1 > 0.8:\n            print(\"Target F1 reached!\")\n            break\n    \n    # Final evaluation on test set\n    model.load_state_dict(torch.load(\"best_model.pth\"))\n    model.eval()\n    \n    test_preds = []\n    test_labels = []\n    \n    with torch.no_grad():\n        for slices, labels in test_loader:\n            slices = slices.to(DEVICE)\n            outputs = model(slices)\n            \n            mid_slice = outputs.shape[2] // 2\n            pred = postprocess(outputs[:, :, mid_slice].cpu().numpy())\n            test_preds.append(pred)\n            test_labels.append(labels.squeeze().cpu().numpy())\n            \n            # Visualize first sample\n            if len(test_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                plt.subplot(1, 3, 1)\n                plt.imshow(slices[0, 0, mid_slice].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 2)\n                plt.imshow(labels[0, 0].cpu(), cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate test metrics\n    test_preds = np.concatenate([p.flatten() for p in test_preds])\n    test_labels = np.concatenate([l.flatten() for l in test_labels])\n    \n    metrics = {\n        'F1': f1_score(test_labels, test_preds),\n        'IoU': jaccard_score(test_labels, test_preds),\n        'Precision': precision_score(test_labels, test_preds),\n        'Recall': recall_score(test_labels, test_preds),\n        'Ink Coverage (GT)': f\"{test_labels.mean():.2%}\",\n        'Ink Coverage (Pred)': f\"{test_preds.mean():.2%}\"\n    }\n    \n    print(\"\\n\" + \"=\"*50)\n    print(f\"{'Test Metrics':^50}\")\n    print(\"=\"*50)\n    for name, value in metrics.items():\n        if isinstance(value, float):\n            print(f\"{name+':':<20}{value:.4f}\")\n        else:\n            print(f\"{name+':':<20}{value}\")\n    print(\"=\"*50)\n    \n    return model\n\n# Run training\nif __name__ == \"__main__\":\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    model = train_model()\n    print(\"Training completed!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import f1_score, jaccard_score, precision_score, recall_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport gc\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nPRETRAINED_PATH = \"/kaggle/input/3d-unet/3d_unet_3dunet_b3_final_swa.pth\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)  # From 00.tif to 64.tif\nBATCH_SIZE = 2  # Reduced to prevent OOM errors\nEPOCHS = 10\nLR = 0.0001  # Small LR for fine-tuning\n\n# Enhanced Dataset Class with memory optimization\nclass Vesuvius3DDataset(Dataset):\n    def __init__(self, volume_path, label_path, resize_shape=None):\n        self.volume_path = volume_path\n        self.label_path = label_path\n        self.resize_shape = resize_shape\n        \n        # Load and process label first\n        self.label = self._process_label()\n        \n        # We'll load slices on demand to save memory\n        self.slice_paths = [os.path.join(volume_path, f\"{i:02d}.tif\") \n                          for i in range(SLICE_RANGE[0], SLICE_RANGE[1]+1)]\n    \n    def _process_label(self):\n        label = np.array(Image.open(self.label_path)) > 0\n        if self.resize_shape:\n            label = cv2.resize(label.astype(np.float32), self.resize_shape,\n                             interpolation=cv2.INTER_NEAREST)\n            label = label > 0.5\n        return label\n    \n    def __len__(self):\n        return 1  # Each dataset contains one complete volume\n    \n    def __getitem__(self, idx):\n        # Load and process all slices for this volume\n        volume = []\n        for slice_path in self.slice_paths:\n            img = tiff.imread(slice_path)\n            if self.resize_shape:\n                img = cv2.resize(img, self.resize_shape, interpolation=cv2.INTER_AREA)\n            img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n            volume.append(img)\n        \n        volume = np.stack(volume, axis=0)  # (65, H, W)\n        \n        return (\n            torch.FloatTensor(volume).unsqueeze(0),  # (1, 65, H, W) - add channel dim\n            torch.FloatTensor(self.label).unsqueeze(0)  # (1, H, W)\n        )\n\n# 3D U-Net Model compatible with pretrained weights\nclass UNet3D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet3D, self).__init__()\n        \n        # Encoder\n        self.enc1 = self._block(in_channels, 64)\n        self.enc2 = self._block(64, 128)\n        self.enc3 = self._block(128, 256)\n        self.pool = nn.MaxPool3d((1, 2, 2))  # Only pool spatial dimensions\n        \n        # Bottleneck\n        self.bottleneck = self._block(256, 512)\n        \n        # Decoder\n        self.upconv3 = nn.ConvTranspose3d(512, 256, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec3 = self._block(512, 256)\n        self.upconv2 = nn.ConvTranspose3d(256, 128, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec2 = self._block(256, 128)\n        self.upconv1 = nn.ConvTranspose3d(128, 64, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec1 = self._block(128, 64)\n        \n        # Output\n        self.conv_out = nn.Conv3d(64, out_channels, kernel_size=1)\n    \n    def _block(self, in_channels, features):\n        return nn.Sequential(\n            nn.Conv3d(in_channels, features, kernel_size=3, padding=1),\n            nn.BatchNorm3d(features),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features, features, kernel_size=3, padding=1),\n            nn.BatchNorm3d(features),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.enc1(x)\n        enc2 = self.enc2(self.pool(enc1))\n        enc3 = self.enc3(self.pool(enc2))\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc3))\n        \n        # Decoder\n        dec3 = self.upconv3(bottleneck)\n        dec3 = torch.cat((dec3, enc3), dim=1)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = torch.cat((dec2, enc2), dim=1)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = torch.cat((dec1, enc1), dim=1)\n        dec1 = self.dec1(dec1)\n        \n        return torch.sigmoid(self.conv_out(dec1))\n\n# Load pretrained model with error handling\ndef load_pretrained_model():\n    model = UNet3D().to(DEVICE)\n    try:\n        state_dict = torch.load(PRETRAINED_PATH, map_location=DEVICE)\n        \n        # Handle DataParallel wrapping if present\n        if all(k.startswith('module.') for k in state_dict.keys()):\n            state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}\n        \n        model.load_state_dict(state_dict)\n        print(\"Successfully loaded pretrained model\")\n        return model\n    except Exception as e:\n        print(f\"Error loading pretrained model: {e}\")\n        return None\n\n# Training and evaluation function\ndef fine_tune_and_evaluate():\n    # Prepare data paths\n    train_vol = os.path.join(DATA_DIR, \"train/1/surface_volume\")\n    train_label = os.path.join(DATA_DIR, \"train/1/inklabels.png\")\n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create datasets\n    train_dataset = Vesuvius3DDataset(train_vol, train_label, RESIZE_SHAPE)\n    test_dataset = Vesuvius3DDataset(test_vol, test_label, RESIZE_SHAPE)\n    \n    # Create dataloaders\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, pin_memory=True)\n    \n    # Initialize model\n    model = load_pretrained_model()\n    if model is None:\n        raise RuntimeError(\"Failed to load pretrained model\")\n    \n    criterion = nn.BCELoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=2)\n    \n    # Fine-tuning loop\n    print(\"\\nStarting fine-tuning...\")\n    for epoch in range(EPOCHS):\n        model.train()\n        epoch_loss = 0.0\n        \n        for volumes, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n            volumes = volumes.to(DEVICE)  # (B, 1, 65, H, W)\n            labels = labels.to(DEVICE)    # (B, 1, H, W)\n            \n            optimizer.zero_grad()\n            outputs = model(volumes)      # (B, 1, 65, H, W)\n            \n            # Use middle slice for training\n            mid_slice = outputs.shape[2] // 2\n            loss = criterion(outputs[:, :, mid_slice], labels)\n            loss.backward()\n            optimizer.step()\n            \n            epoch_loss += loss.item()\n        \n        avg_loss = epoch_loss / len(train_loader)\n        scheduler.step(avg_loss)\n        print(f\"Epoch {epoch+1} Complete - Loss: {avg_loss:.4f}, LR: {optimizer.param_groups[0]['lr']:.6f}\")\n    \n    # Evaluation\n    model.eval()\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for volumes, labels in test_loader:\n            volumes = volumes.to(DEVICE)\n            outputs = model(volumes)\n            \n            # Get middle slice prediction\n            mid_slice = outputs.shape[2] // 2\n            pred = (outputs[0, 0, mid_slice].cpu().numpy() > 0.5).astype(np.uint8)\n            label = labels[0, 0].cpu().numpy().astype(np.uint8)\n            \n            all_preds.append(pred)\n            all_labels.append(label)\n            \n            # Visualize first sample\n            if len(all_preds) == 1:\n                plt.figure(figsize=(18, 6))\n                \n                # Input middle slice\n                plt.subplot(1, 3, 1)\n                plt.imshow(volumes[0, 0, mid_slice].cpu(), cmap='gray')\n                plt.title(\"Input Slice (Middle)\", fontsize=12)\n                plt.axis('off')\n                \n                # Ground truth\n                plt.subplot(1, 3, 2)\n                plt.imshow(label, cmap='gray')\n                plt.title(\"Ground Truth\", fontsize=12)\n                plt.axis('off')\n                \n                # Prediction\n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\", fontsize=12)\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_labels = np.concatenate([l.flatten() for l in all_labels])\n    \n    metrics = {\n        'F1': f1_score(all_labels, all_preds),\n        'IoU': jaccard_score(all_labels, all_preds),\n        'Precision': precision_score(all_labels, all_preds),\n        'Recall': recall_score(all_labels, all_preds),\n        'Ink Coverage (GT)': f\"{all_labels.mean():.2%}\",\n        'Ink Coverage (Pred)': f\"{all_preds.mean():.2%}\"\n    }\n    \n    # Print metrics\n    print(\"\\n\" + \"=\"*50)\n    print(f\"{'Evaluation Metrics':^50}\")\n    print(\"=\"*50)\n    for name, value in metrics.items():\n        if isinstance(value, float):\n            print(f\"{name+':':<20}{value:.4f}\")\n        else:\n            print(f\"{name+':':<20}{value}\")\n    print(\"=\"*50)\n    \n    return model\n\n# Run the complete pipeline\nif __name__ == \"__main__\":\n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Run fine-tuning and evaluation\n    model = fine_tune_and_evaluate()\n    \n    # Save fine-tuned model\n    torch.save(model.state_dict(), \"fine_tuned_ink_detection_model.pth\")\n    print(\"\\nFine-tuned model saved successfully!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import f1_score, jaccard_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\nimport gc\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)  # 00.tif to 64.tif\nBATCH_SIZE = 10\nEPOCHS = 8\nLR = 0.001\nNUM_SPLITS = 4\n\n# Corrected Dataset Class\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_path, label_path, resize_shape=None):\n        self.volume_path = volume_path\n        self.label_path = label_path\n        self.resize_shape = resize_shape\n        \n        # Load all slices at once (but we'll process them individually)\n        self.slices = []\n        for i in range(SLICE_RANGE[0], SLICE_RANGE[1] + 1):\n            slice_path = os.path.join(volume_path, f\"{i:02d}.tif\")\n            img = tiff.imread(slice_path)\n            if resize_shape:\n                img = cv2.resize(img, resize_shape, interpolation=cv2.INTER_AREA)\n            img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n            self.slices.append(img)\n        \n        # Load and process label\n        self.label = np.array(Image.open(label_path)) > 0\n        if resize_shape:\n            self.label = cv2.resize(self.label.astype(np.float32), resize_shape, \n                                 interpolation=cv2.INTER_NEAREST)\n            self.label = self.label > 0.5\n    \n    def __len__(self):\n        return len(self.slices)\n    \n    def __getitem__(self, idx):\n        return (\n            torch.FloatTensor(self.slices[idx]).unsqueeze(0),  # (1, H, W)\n            torch.FloatTensor(self.label)  # (H, W)\n        )\n\n# Corrected LiteUNet3D\nclass LiteUNet3D(nn.Module):\n    def __init__(self):\n        super(LiteUNet3D, self).__init__()\n        \n        # Encoder\n        self.enc1 = self._block(1, 16)  # Input channels = 1\n        self.enc2 = self._block(16, 32)\n        self.enc3 = self._block(32, 64)\n        self.pool = nn.MaxPool2d(2)  # Using 2D pooling for memory efficiency\n        \n        # Bottleneck\n        self.bottleneck = self._block(64, 128)\n        \n        # Decoder\n        self.upconv3 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        self.dec3 = self._block(128, 64)\n        self.upconv2 = nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2)\n        self.dec2 = self._block(64, 32)\n        self.upconv1 = nn.ConvTranspose2d(32, 16, kernel_size=2, stride=2)\n        self.dec1 = self._block(32, 16)\n        \n        # Output\n        self.conv_out = nn.Conv2d(16, 1, kernel_size=1)\n    \n    def _block(self, in_channels, features):\n        return nn.Sequential(\n            nn.Conv2d(in_channels, features, kernel_size=3, padding=1),\n            nn.InstanceNorm2d(features),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(features, features, kernel_size=3, padding=1),\n            nn.InstanceNorm2d(features),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # x shape: (batch_size, 1, H, W)\n        \n        # Encoder\n        enc1 = self.enc1(x)  # (B, 16, H, W)\n        enc2 = self.enc2(self.pool(enc1))  # (B, 32, H/2, W/2)\n        enc3 = self.enc3(self.pool(enc2))  # (B, 64, H/4, W/4)\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc3))  # (B, 128, H/8, W/8)\n        \n        # Decoder\n        dec3 = self.upconv3(bottleneck)  # (B, 64, H/4, W/4)\n        dec3 = torch.cat((dec3, enc3), dim=1)  # (B, 128, H/4, W/4)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.upconv2(dec3)  # (B, 32, H/2, W/2)\n        dec2 = torch.cat((dec2, enc2), dim=1)  # (B, 64, H/2, W/2)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.upconv1(dec2)  # (B, 16, H, W)\n        dec1 = torch.cat((dec1, enc1), dim=1)  # (B, 32, H, W)\n        dec1 = self.dec1(dec1)\n        \n        return torch.sigmoid(self.conv_out(dec1))  # (B, 1, H, W)\n\n# Training function\ndef train_and_evaluate():\n    # Prepare data\n    train_vol = os.path.join(DATA_DIR, \"train/1/surface_volume\")\n    train_label = os.path.join(DATA_DIR, \"train/1/inklabels.png\")\n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create datasets\n    train_dataset = VesuviusDataset(train_vol, train_label, RESIZE_SHAPE)\n    test_dataset = VesuviusDataset(test_vol, test_label, RESIZE_SHAPE)\n    \n    # Split training data into 8 parts\n    kfold = KFold(n_splits=NUM_SPLITS, shuffle=True)\n    split_datasets = []\n    for _, fold_indices in kfold.split(range(len(train_dataset))):\n        split_datasets.append(Subset(train_dataset, fold_indices))\n    \n    # Initialize model\n    model = LiteUNet3D().to(DEVICE)\n    criterion = nn.BCELoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    \n    # Training loop\n    for split_idx, train_subset in enumerate(split_datasets):\n        print(f\"\\n=== Training on split {split_idx+1}/{NUM_SPLITS} ===\")\n        \n        train_loader = DataLoader(\n            train_subset,\n            batch_size=BATCH_SIZE,\n            shuffle=True,\n            pin_memory=True\n        )\n        \n        for epoch in range(EPOCHS):\n            model.train()\n            epoch_loss = 0.0\n            \n            for slices, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n                slices = slices.to(DEVICE)  # (B, 1, H, W)\n                labels = labels.to(DEVICE)  # (B, H, W)\n                \n                optimizer.zero_grad()\n                outputs = model(slices)  # (B, 1, H, W)\n                loss = criterion(outputs.squeeze(1), labels)  # Remove channel dim for BCELoss\n                loss.backward()\n                optimizer.step()\n                \n                epoch_loss += loss.item()\n            \n            print(f\"Split {split_idx+1} | Epoch {epoch+1} Loss: {epoch_loss/len(train_loader):.4f}\")\n            \n            # Clear memory\n            torch.cuda.empty_cache()\n            gc.collect()\n    \n    # Evaluation\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    model.eval()\n    \n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for slices, labels in test_loader:\n            slices = slices.to(DEVICE)\n            outputs = model(slices)\n            \n            pred = (outputs.squeeze().cpu().numpy() > 0.5).astype(np.uint8)\n            label = labels.squeeze().cpu().numpy().astype(np.uint8)\n            \n            all_preds.append(pred)\n            all_labels.append(label)\n            \n            # Visualize first sample\n            if len(all_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                plt.subplot(1, 3, 1)\n                plt.imshow(slices[0, 0].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 2)\n                plt.imshow(label, cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_labels = np.concatenate([l.flatten() for l in all_labels])\n    \n    f1 = f1_score(all_labels, all_preds)\n    iou = jaccard_score(all_labels, all_preds)\n    \n    print(\"\\n=== Final Test Metrics ===\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(f\"IoU (Jaccard): {iou:.4f}\")\n    \n    return model\n\n# Run training\ntrained_model = train_and_evaluate()\n\n# Save model\ntorch.save(trained_model.state_dict(), \"ink_detection_model_2d.pth\")\nprint(\"Model saved successfully\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import f1_score, jaccard_score\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\n\n# Configuration\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nPRETRAINED_PATH = \"/kaggle/input/3d-unet/3d_unet_3dunet_b3_final_swa.pth\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nRESIZE_SHAPE = (128, 128)\nSLICE_RANGE = (0, 64)  # From 00.tif to 64.tif\nBATCH_SIZE = 2\nEPOCHS = 10\nLR = 0.0001\n\n# Dataset Class\nclass VesuviusDataset(Dataset):\n    def __init__(self, volume_path, label_path, resize_shape=None):\n        self.volume_path = volume_path\n        self.label_path = label_path\n        self.resize_shape = resize_shape\n        \n        # Load all slices\n        self.volume = []\n        for i in range(SLICE_RANGE[0], SLICE_RANGE[1] + 1):\n            img = tiff.imread(os.path.join(volume_path, f\"{i:02d}.tif\"))\n            if resize_shape:\n                img = cv2.resize(img, resize_shape, interpolation=cv2.INTER_AREA)\n            self.volume.append(img)\n        self.volume = np.stack(self.volume, axis=0)  # (65, H, W)\n        \n        # Load and process label\n        self.label = np.array(Image.open(label_path)) > 0\n        if resize_shape:\n            self.label = cv2.resize(self.label.astype(np.float32), resize_shape,\n                                 interpolation=cv2.INTER_NEAREST)\n            self.label = self.label > 0.5\n    \n    def __len__(self):\n        return 1  # Return 1 for whole volume\n    \n    def __getitem__(self, idx):\n        # Normalize volume\n        volume = (self.volume - self.volume.min()) / (self.volume.max() - self.volume.min())\n        return (\n            torch.FloatTensor(volume).unsqueeze(0),  # (1, 65, H, W)\n            torch.FloatTensor(self.label).unsqueeze(0)  # (1, H, W)\n        )\n\n# 3D U-Net Model (must match pretrained architecture)\nclass UNet3D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet3D, self).__init__()\n        \n        # Encoder\n        self.enc1 = self._block(in_channels, 64)\n        self.enc2 = self._block(64, 128)\n        self.enc3 = self._block(128, 256)\n        self.pool = nn.MaxPool3d((1, 2, 2))  # Only pool spatial dimensions\n        \n        # Bottleneck\n        self.bottleneck = self._block(256, 512)\n        \n        # Decoder\n        self.upconv3 = nn.ConvTranspose3d(512, 256, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec3 = self._block(512, 256)\n        self.upconv2 = nn.ConvTranspose3d(256, 128, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec2 = self._block(256, 128)\n        self.upconv1 = nn.ConvTranspose3d(128, 64, kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec1 = self._block(128, 64)\n        \n        # Output\n        self.conv_out = nn.Conv3d(64, out_channels, kernel_size=1)\n    \n    def _block(self, in_channels, features):\n        return nn.Sequential(\n            nn.Conv3d(in_channels, features, kernel_size=3, padding=1),\n            nn.BatchNorm3d(features),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features, features, kernel_size=3, padding=1),\n            nn.BatchNorm3d(features),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.enc1(x)\n        enc2 = self.enc2(self.pool(enc1))\n        enc3 = self.enc3(self.pool(enc2))\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc3))\n        \n        # Decoder\n        dec3 = self.upconv3(bottleneck)\n        dec3 = torch.cat((dec3, enc3), dim=1)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = torch.cat((dec2, enc2), dim=1)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = torch.cat((dec1, enc1), dim=1)\n        dec1 = self.dec1(dec1)\n        \n        return torch.sigmoid(self.conv_out(dec1))\n\n# Load pretrained model\ndef load_pretrained():\n    model = UNet3D().to(DEVICE)\n    state_dict = torch.load(PRETRAINED_PATH, map_location=DEVICE)\n    \n    # Handle DataParallel if present\n    if all(k.startswith('module.') for k in state_dict.keys()):\n        state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}\n    \n    model.load_state_dict(state_dict)\n    return model\n\n# Fine-tuning function\ndef fine_tune_and_evaluate():\n    # Load pretrained model\n    model = load_pretrained()\n    \n    # Prepare data\n    train_vol = os.path.join(DATA_DIR, \"train/1/surface_volume\")\n    train_label = os.path.join(DATA_DIR, \"train/1/inklabels.png\")\n    test_vol = os.path.join(DATA_DIR, \"train/3/surface_volume\")\n    test_label = os.path.join(DATA_DIR, \"train/3/inklabels.png\")\n    \n    # Create datasets\n    train_dataset = VesuviusDataset(train_vol, train_label, RESIZE_SHAPE)\n    test_dataset = VesuviusDataset(test_vol, test_label, RESIZE_SHAPE)\n    \n    # Create dataloaders\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    # Loss and optimizer\n    criterion = nn.BCELoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    \n    # Fine-tuning loop\n    for epoch in range(EPOCHS):\n        model.train()\n        epoch_loss = 0.0\n        \n        for volumes, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n            volumes = volumes.to(DEVICE)  # (B, 1, 65, H, W)\n            labels = labels.to(DEVICE)  # (B, 1, H, W)\n            \n            optimizer.zero_grad()\n            outputs = model(volumes)  # (B, 1, 65, H, W)\n            \n            # Take middle slice for training\n            mid_slice = outputs.shape[2] // 2\n            loss = criterion(outputs[:, :, mid_slice], labels)\n            loss.backward()\n            optimizer.step()\n            \n            epoch_loss += loss.item()\n        \n        print(f\"Epoch {epoch+1} Loss: {epoch_loss/len(train_loader):.4f}\")\n    \n    # Evaluation\n    model.eval()\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for volumes, labels in test_loader:\n            volumes = volumes.to(DEVICE)\n            outputs = model(volumes)\n            \n            # Get middle slice prediction\n            mid_slice = outputs.shape[2] // 2\n            pred = (outputs[0, 0, mid_slice].cpu().numpy() > 0.5).astype(np.uint8)\n            label = labels[0, 0].cpu().numpy().astype(np.uint8)\n            \n            all_preds.append(pred)\n            all_labels.append(label)\n            \n            # Visualize first sample\n            if len(all_preds) == 1:\n                plt.figure(figsize=(15, 5))\n                \n                # Input middle slice\n                plt.subplot(1, 3, 1)\n                plt.imshow(volumes[0, 0, mid_slice].cpu(), cmap='gray')\n                plt.title(\"Input Slice\")\n                plt.axis('off')\n                \n                # Ground truth\n                plt.subplot(1, 3, 2)\n                plt.imshow(label, cmap='gray')\n                plt.title(\"Ground Truth\")\n                plt.axis('off')\n                \n                # Prediction\n                plt.subplot(1, 3, 3)\n                plt.imshow(pred, cmap='gray')\n                plt.title(\"Prediction\")\n                plt.axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n    \n    # Calculate metrics\n    all_preds = np.concatenate([p.flatten() for p in all_preds])\n    all_labels = np.concatenate([l.flatten() for l in all_labels])\n    \n    f1 = f1_score(all_labels, all_preds)\n    iou = jaccard_score(all_labels, all_preds)\n    \n    print(\"\\n=== Evaluation Metrics ===\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(f\"IoU (Jaccard): {iou:.4f}\")\n    \n    return model\n\n# Run fine-tuning and evaluation\ntrained_model = fine_tune_and_evaluate()\n\n# Save fine-tuned model\ntorch.save(trained_model.state_dict(), \"fine_tuned_ink_detection_model.pth\")\nprint(\"Fine-tuned model saved successfully\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. Install monai and einops\n# ✅ Install monai and einops together from local wheels\n!pip install --no-index --find-links=/kaggle/input/monai-packages monai > /dev/null\n!pip install --no-index --find-links=/kaggle/input/einops einops > /dev/null\n\n# 2. Imports\nimport os, sys, cv2\nimport numpy as np\nfrom glob import glob\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport torch\nimport torch.nn.functional as F\nfrom sklearn.metrics import fbeta_score\nsys.path.append(\"/kaggle/input/unet3d/\")\nfrom unetr import UNETR\n\n# 3. Load pretrained model\nmodel = UNETR(\n    in_channels=1,\n    out_channels=1,\n    img_size=(256, 256, 64),  # ✅ must be divisible by patch size\n    feature_size=16,\n    hidden_size=768,\n    num_heads=12,\n    pos_embed='perceptron',\n    norm_name='instance',\n    res_block=True,\n    dropout_rate=0.0\n)\nmodel.load_state_dict(torch.load(\"/kaggle/input/vesuvius-models/Unet_fold1_best.pth\", map_location=\"cpu\"))\nmodel.eval()\n\n# 4. Load and preprocess volume\nresize_shape = (256, 256)\nvolume_dir = \"/kaggle/input/vesuvius-challenge-ink-detection/train/1/surface_volume\"\nfiles = sorted(glob(os.path.join(volume_dir, \"*.tif\")))[:64]  # ✅ trim to 64 slices\nslices = [cv2.resize(cv2.imread(f, 0), resize_shape) for f in files]\nvolume = np.stack(slices).astype(np.float32) / 255.0\nvolume = volume[np.newaxis, np.newaxis, ...]  # (1, 1, D, H, W)\ninput_tensor = torch.tensor(volume, dtype=torch.float32)\n\n# 5. Predict\nwith torch.no_grad():\n    output = model(input_tensor)\n    output = torch.sigmoid(output)\n    pred_volume = (output.squeeze().numpy() > 0.5).astype(np.uint8)  # (D, H, W)\n\n# 6. Load ground truth and resize\ngt = np.array(Image.open(\"/kaggle/input/vesuvius-challenge-ink-detection/train/1/inklabels.png\"))\ngt = cv2.resize(gt, resize_shape, interpolation=cv2.INTER_NEAREST)\ngt_stack = np.broadcast_to((gt > 0).astype(np.uint8), pred_volume.shape)\n\n# 7. Metrics\nflat_pred = pred_volume.flatten()\nflat_gt = gt_stack.flatten()\nintersection = (flat_pred * flat_gt).sum()\nunion = flat_pred.sum() + flat_gt.sum()\ndice = 2 * intersection / (union + 1e-8)\nf05 = fbeta_score(flat_gt, flat_pred, beta=0.5)\n\nprint(f\"DICE score: {dice:.4f}\")\nprint(f\"F0.5 score: {f05:.4f}\")\n\n# 8. Save masks\nos.makedirs(\"predicted_ink_mask\", exist_ok=True)\nfor i in range(pred_volume.shape[0]):\n    img = (pred_volume[i] * 255).astype(np.uint8)\n    Image.fromarray(img).save(f\"predicted_ink_mask/slice_{i:02d}.png\")\n\n# 9. Visual comparison\ndef show_overlay(i):\n    fig, axs = plt.subplots(1, 3, figsize=(18, 6))\n    axs[0].imshow(volume[0, 0, i], cmap='gray')\n    axs[0].set_title(\"Original Volume Slice\")\n    axs[1].imshow(pred_volume[i], cmap='hot')\n    axs[1].set_title(\"Predicted Ink Mask\")\n    axs[2].imshow(gt_stack[i], cmap='gray')\n    axs[2].set_title(\"Ground Truth Ink Mask\")\n    plt.show()\n\nshow_overlay(32)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index --find-links=/kaggle/input/einops/einops-0.6.1-py3-none-any.whl > /dev/null\n!pip install --no-index --find-links=/kaggle/input/monai monai > /dev/null\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install /kaggle/input/einops/einops-0.6.1-py3-none-any.whl","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import einops","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport pandas as pd\n\n# Paths\ndata_id = \"1\"\nbase_path = f\"/kaggle/input/vesuvius-challenge-ink-detection/train/{data_id}\"\nmask_path = os.path.join(base_path, \"mask.png\")\ninklabels_path = os.path.join(base_path, \"inklabels.png\")\n\n# Resize dimensions\nresize_shape = (256, 256)\n\n# Load and resize mask.png\nmask_img = np.array(Image.open(mask_path).convert(\"L\"))\nmask_img = cv2.resize(mask_img, resize_shape, interpolation=cv2.INTER_AREA)\nmask_img = mask_img.astype(np.float32) / 255.0\n\n# Load and resize inklabels (binary mask)\ninklabels = np.array(Image.open(inklabels_path).convert(\"L\"))\ninklabels = cv2.resize(inklabels, resize_shape, interpolation=cv2.INTER_NEAREST)\ninklabels = (inklabels > 0).astype(np.uint8)\n\n# Multiply ink * mask\nink_result = inklabels * mask_img\n\n# Multiply ~ink * mask\nnonink_mask = 1 - inklabels\nnonink_result = nonink_mask * mask_img\n\n# Extract positions and intensities of inked and non-inked regions\nink_positions = np.argwhere(ink_result > 0)\nink_values = ink_result[ink_result > 0]\nnonink_positions = np.argwhere(nonink_result > 0)\nnonink_values = nonink_result[nonink_result > 0]\n\n# Create DataFrames for analysis\nink_df = pd.DataFrame(ink_positions, columns=[\"y\", \"x\"])\nink_df[\"intensity\"] = ink_values\n\nnonink_df = pd.DataFrame(nonink_positions, columns=[\"y\", \"x\"])\nnonink_df[\"intensity\"] = nonink_values\n\n# Export to CSV\nink_df.to_csv(\"inked_mask_pixels.csv\", index=False)\nnonink_df.to_csv(\"noninked_mask_pixels.csv\", index=False)\n\n# Preview some stats\nprint(\"Inked region pixel stats:\")\nprint(ink_df.describe())\nprint(\"\\nNon-inked region pixel stats:\")\nprint(nonink_df.describe())\n\n# Visualize\nfig, axs = plt.subplots(1, 3, figsize=(18, 6))\naxs[0].imshow(mask_img, cmap='gray')\naxs[0].set_title(\"Original Mask\")\naxs[1].imshow(ink_result, cmap='hot')\naxs[1].set_title(\"Inked Region * Mask\")\naxs[2].imshow(nonink_result, cmap='bone')\naxs[2].set_title(\"Non-Inked Region * Mask\")\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport cv2\nfrom glob import glob\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom PIL import Image\nimport pandas as pd\n\n# Paths\ndata_id = \"1\"\nbase_path = f\"/kaggle/input/vesuvius-challenge-ink-detection/train/{data_id}\"\nsurface_volume_path = os.path.join(base_path, \"surface_volume\")\ninklabels_path = os.path.join(base_path, \"inklabels.png\")\n\n# Resize dimensions\nresize_shape = (256, 256)\n\n# Load and resize surface_volume (stack of 65 .tif slices)\nsurface_files = sorted(glob(os.path.join(surface_volume_path, \"*.tif\")))\nsurface_volume = np.stack([\n    cv2.resize(cv2.imread(f, cv2.IMREAD_GRAYSCALE), resize_shape, interpolation=cv2.INTER_AREA)\n    for f in tqdm(surface_files)\n])\n\n# Normalize to 0.0 - 1.0 range\nsurface_volume = surface_volume.astype(np.float32) / 255.0\n\n# Load and resize inklabels (binary mask)\ninklabels = np.array(Image.open(inklabels_path).convert(\"L\"))\ninklabels = cv2.resize(inklabels, resize_shape, interpolation=cv2.INTER_NEAREST)\ninklabels = (inklabels > 0).astype(np.uint8)  # binary mask\n\n# Broadcast inklabels across 65 slices\nink_mask_3d = np.broadcast_to(inklabels, surface_volume.shape)\n\n# Multiply ink * surface_volume\nink_result = ink_mask_3d * surface_volume\n\n# Multiply ~ink * surface_volume\nnonink_mask_3d = 1 - ink_mask_3d\nnonink_result = nonink_mask_3d * surface_volume\n\n# Extract positions and intensities of inked and non-inked regions\nink_positions = np.argwhere(ink_result > 0)\nink_values = ink_result[ink_result > 0]\nnonink_positions = np.argwhere(nonink_result > 0)\nnonink_values = nonink_result[nonink_result > 0]\n\n# Create DataFrames for analysis\nink_df = pd.DataFrame(ink_positions, columns=[\"slice\", \"y\", \"x\"])\nink_df[\"intensity\"] = ink_values\n\nnonink_df = pd.DataFrame(nonink_positions, columns=[\"slice\", \"y\", \"x\"])\nnonink_df[\"intensity\"] = nonink_values\n\n# Export to CSV\nink_df.to_csv(\"inked_voxels.csv\", index=False)\nnonink_df.to_csv(\"noninked_voxels.csv\", index=False)\n\n# Preview some stats\nprint(\"Inked region voxel stats:\")\nprint(ink_df.describe())\nprint(\"\\nNon-inked region voxel stats:\")\nprint(nonink_df.describe())\n\n# Preview slice comparison\ndef show_comparison(slice_idx):\n    fig, axs = plt.subplots(1, 3, figsize=(18, 6))\n    axs[0].imshow(surface_volume[slice_idx], cmap='gray')\n    axs[0].set_title(\"Original Surface\")\n    axs[1].imshow(ink_result[slice_idx], cmap='hot')\n    axs[1].set_title(\"Inked Region * Surface\")\n    axs[2].imshow(nonink_result[slice_idx], cmap='bone')\n    axs[2].set_title(\"Non-Inked Region * Surface\")\n    plt.show()\n\n# Example output slice\nshow_comparison(32)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport cv2\nfrom glob import glob\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom PIL import Image\nimport pandas as pd\n\n# Paths\ndata_id = \"1\"\nbase_path = f\"/kaggle/input/vesuvius-challenge-ink-detection/train/{data_id}\"\nsurface_volume_path = os.path.join(base_path, \"surface_volume\")\ninklabels_path = os.path.join(base_path, \"inklabels.png\")\n\n# Resize dimensions\nresize_shape = (256, 256)\n\n# Load and resize surface_volume (stack of 65 .tif slices)\nsurface_files = sorted(glob(os.path.join(surface_volume_path, \"*.tif\")))\nsurface_volume = np.stack([\n    cv2.resize(cv2.imread(f, cv2.IMREAD_GRAYSCALE), resize_shape, interpolation=cv2.INTER_AREA)\n    for f in tqdm(surface_files)\n])\n\n# Normalize to 0.0 - 1.0 range\nsurface_volume = surface_volume.astype(np.float32) / 255.0\n\n# Load and resize inklabels (binary mask)\ninklabels = np.array(Image.open(inklabels_path).convert(\"L\"))\ninklabels = cv2.resize(inklabels, resize_shape, interpolation=cv2.INTER_NEAREST)\ninklabels = (inklabels > 0).astype(np.uint8)  # binary mask\n\n# Broadcast inklabels across 65 slices\nink_mask_3d = np.broadcast_to(inklabels, surface_volume.shape)\n\n# Multiply ink * surface_volume\nink_result = ink_mask_3d * surface_volume\n\n# Multiply ~ink * surface_volume\nnonink_mask_3d = 1 - ink_mask_3d\nnonink_result = nonink_mask_3d * surface_volume\n\n# Extract positions and intensities of inked and non-inked regions\nink_positions = np.argwhere(ink_result > 0)\nink_values = ink_result[ink_result > 0]\nnonink_positions = np.argwhere(nonink_result > 0)\nnonink_values = nonink_result[nonink_result > 0]\n\n# Create DataFrames for analysis\nink_df = pd.DataFrame(ink_positions, columns=[\"slice\", \"y\", \"x\"])\nink_df[\"intensity\"] = ink_values\n\nnonink_df = pd.DataFrame(nonink_positions, columns=[\"slice\", \"y\", \"x\"])\nnonink_df[\"intensity\"] = nonink_values\n\n# Preview some results\nprint(\"Inked region voxel stats:\")\nprint(ink_df.describe())\nprint(\"\\nNon-inked region voxel stats:\")\nprint(nonink_df.describe())\n\n# Preview slice comparison\ndef show_comparison(slice_idx):\n    fig, axs = plt.subplots(1, 3, figsize=(18, 6))\n    axs[0].imshow(surface_volume[slice_idx], cmap='gray')\n    axs[0].set_title(\"Original Surface\")\n    axs[1].imshow(ink_result[slice_idx], cmap='hot')\n    axs[1].set_title(\"Inked Region * Surface\")\n    axs[2].imshow(nonink_result[slice_idx], cmap='bone')\n    axs[2].set_title(\"Non-Inked Region * Surface\")\n    plt.show()\n\n# Example output slice\nshow_comparison(32)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport cv2\nfrom glob import glob\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom PIL import Image\n\n# Paths\ndata_id = \"1\"\nbase_path = f\"/kaggle/input/vesuvius-challenge-ink-detection/train/{data_id}\"\nsurface_volume_path = os.path.join(base_path, \"surface_volume\")\ninklabels_path = os.path.join(base_path, \"inklabels.png\")\n\n# Resize dimensions\nresize_shape = (256, 256)\n\n# Load and resize surface_volume (stack of 65 .tif slices)\nsurface_files = sorted(glob(os.path.join(surface_volume_path, \"*.tif\")))\nsurface_volume = np.stack([\n    cv2.resize(cv2.imread(f, cv2.IMREAD_GRAYSCALE), resize_shape, interpolation=cv2.INTER_AREA)\n    for f in tqdm(surface_files)\n])\n\n# Normalize to 0.0 - 1.0 range\nsurface_volume = surface_volume.astype(np.float32) / 255.0\n\n# Load and resize inklabels (binary mask)\ninklabels = np.array(Image.open(inklabels_path).convert(\"L\"))\ninklabels = cv2.resize(inklabels, resize_shape, interpolation=cv2.INTER_NEAREST)\ninklabels = (inklabels > 0).astype(np.uint8)  # binary mask\n\n# Broadcast inklabels across 65 slices\nink_mask_3d = np.broadcast_to(inklabels, surface_volume.shape)\n\n# Multiply ink * surface_volume\nink_result = ink_mask_3d * surface_volume\n\n# Multiply ~ink * surface_volume\nnonink_mask_3d = 1 - ink_mask_3d\nnonink_result = nonink_mask_3d * surface_volume\n\n# Preview some slices\ndef show_comparison(slice_idx):\n    fig, axs = plt.subplots(1, 3, figsize=(18, 6))\n    axs[0].imshow(surface_volume[slice_idx], cmap='gray')\n    axs[0].set_title(\"Original Surface\")\n    axs[1].imshow(ink_result[slice_idx], cmap='hot')\n    axs[1].set_title(\"Inked Region * Surface\")\n    axs[2].imshow(nonink_result[slice_idx], cmap='bone')\n    axs[2].set_title(\"Non-Inked Region * Surface\")\n    plt.show()\n\n# Example output slice\nshow_comparison(32)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nimport pickle\nfrom torch.utils.data import DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport warnings\nimport sys\nimport pandas as pd\nimport os\nimport gc\nimport sys\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nimport cv2\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport argparse\nimport importlib\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam, SGD, AdamW\n\nimport datetime\nimport wandb","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/einops/einops-0.6.1-py3-none-any.whl\n!pip install /kaggle/input/monai-packages/monai-1.1.0-202212191849-py3-none-any.whl[\"einops\"]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/input/pretrainedmodels/pretrainedmodels-0.7.4')\nsys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch-master')\nsys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master')\nsys.path.append('/kaggle/input/segmentation-models-pytorch/segmentation_models.pytorch-master')\nsys.path.append('/kaggle/input/unet3d/pytorch3dunet/pytorch3dunet')\nsys.path.append('/kaggle/input/unet3d/pytorch3dunet')\nsys.path.append('/kaggle/input/unet3d/')\n\nimport segmentation_models_pytorch as smp\nfrom unet3d.model import get_model\nfrom unetr import UNETR","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom torch.utils.data import DataLoader, Dataset\nimport cv2\nimport torch\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## config","metadata":{}},{"cell_type":"code","source":"import os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nclass CFG:\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n\n    # comp_dir_path = './'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'vesuvius-challenge-ink-detection'\n    # comp_dataset_path = f'{comp_dir_path}datasets/{comp_folder_name}/'\n    comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'\n    \n    exp_name = '3d_unet_subv2'\n\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = '3d_unet_segformer'\n    backbone = 'None'\n#     backbone = 'se_resnext50_32x4d'\n\n    in_chans = 16\n    # ============== training cfg =============\n    size = 1024\n    tile_size = 1024\n    stride = tile_size // 4\n\n    batch_size = 3 # 32\n    use_amp = True\n\n    scheduler = 'GradualWarmupSchedulerV2'\n    # scheduler = 'CosineAnnealingLR'\n    epochs = 15\n\n    warmup_factor = 10\n    lr = 1e-4 / warmup_factor\n\n    # ============== fold =============\n    valid_id = 2\n\n    objective_cv = 'binary'  # 'binary', 'multiclass', 'regression'\n    metric_direction = 'maximize'  # maximize, 'minimize'\n    # metrics = 'dice_coef'\n\n    # ============== fixed =============\n    pretrained = True\n    inf_weight = 'best'  # 'best'\n\n    min_lr = 1e-6\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n\n    print_freq = 50\n    num_workers = 2\n\n    seed = 42\n\n    # ============== augmentation =============\n    train_aug_list = [\n        # A.RandomResizedCrop(\n        #     size, size, scale=(0.85, 1.0)),\n        A.Resize(size, size),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.75),\n        A.ShiftScaleRotate(p=0.75),\n        A.OneOf([\n                A.GaussNoise(var_limit=[10, 50]),\n                A.GaussianBlur(),\n                A.MotionBlur(),\n                ], p=0.4),\n        A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n        A.CoarseDropout(max_holes=1, max_width=int(size * 0.3), max_height=int(size * 0.3), \n                        mask_fill_value=0, p=0.5),\n        # A.Cutout(max_h_size=int(size * 0.6),\n        #          max_w_size=int(size * 0.6), num_holes=1, p=1.0),\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\n\n    valid_aug_list = [\n        A.Resize(size, size),\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IS_DEBUG = False\nmode = 'train' if IS_DEBUG else 'test'\nTH = 0.5","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## helper","metadata":{}},{"cell_type":"code","source":"# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    # pixels = (pixels >= thr).astype(int)\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"def read_image(fragment_id):\n    images = []\n\n#     idxs = range(65)\n    mid = 65 // 2\n    start = mid - CFG.in_chans // 2\n    end = mid + CFG.in_chans // 2\n    idxs = range(start, end)\n\n    for i in tqdm(idxs):\n        \n        image = cv2.imread(CFG.comp_dataset_path + f\"{mode}/{fragment_id}/surface_volume/{i:02}.tif\", 0)\n\n        pad0 = (CFG.tile_size - image.shape[0] % CFG.tile_size)\n        pad1 = (CFG.tile_size - image.shape[1] % CFG.tile_size)\n\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0)\n\n        images.append(image)\n    images = np.stack(images, axis=2)\n    \n    return images","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(cfg.train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(cfg.valid_aug_list)\n\n    # print(aug)\n    return aug\n\nclass CustomDataset(Dataset):\n    def __init__(self, images, cfg, labels=None, transform=None):\n        self.images = np.array(images)\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        # return len(self.xyxys)\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image = np.load(self.images[idx])\n        data = self.transform(image=image)\n        image = data['image']\n        return image[None, :, :, :]\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_test_dataset(fragment_id):\n    test_images = read_image(fragment_id)\n    \n    x1_list = list(range(0, test_images.shape[1]-CFG.tile_size+1, CFG.stride))\n    y1_list = list(range(0, test_images.shape[0]-CFG.tile_size+1, CFG.stride))\n    \n    test_images_list = []\n    xyxys = []\n    for y1 in y1_list:\n        for x1 in x1_list:\n            y2 = y1 + CFG.tile_size\n            x2 = x1 + CFG.tile_size\n            if test_images[y1:y2, x1:x2].max() != 0:\n                if not os.path.exists(f\"{x1}_{y1}_{x2}_{y2}.npy\"):\n                    np.save(f\"{x1}_{y1}_{x2}_{y2}.npy\", test_images[y1:y2, x1:x2])\n                test_images_list.append(f\"{x1}_{y1}_{x2}_{y2}.npy\")\n                xyxys.append((x1, y1, x2, y2))\n    del test_images\n    gc.collect()\n    xyxys = np.stack(xyxys)\n            \n    test_dataset = CustomDataset(test_images_list, CFG, transform=get_transforms(data='valid', cfg=CFG))\n    \n    test_loader = DataLoader(test_dataset,\n                          batch_size=CFG.batch_size,\n                          shuffle=False,\n                          num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n    return test_loader, xyxys","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"from transformers import SegformerForSemanticSegmentation, SegformerModel, SegformerConfig","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_3d_segformer_b1_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 256,\n  \"depths\": [\n    2,\n    2,\n    2,\n    2\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 32,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"transformers_version\": \"4.12.0.dev0\",\n  \"num_labels\":1,\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_3d_segformer_b2_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    4,\n    6,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 32,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"transformers_version\": \"4.12.0.dev0\",\n  \"num_labels\":1\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_3d_segformer_b4_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    8,\n    27,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 32,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"transformers_version\": \"4.12.0.dev0\",\n  \"num_labels\":1\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_3d_segformer_b5_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    6,\n    40,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 32,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"transformers_version\": \"4.12.0.dev0\",\n  \"num_labels\":1\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_3d_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    4,\n    18,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"num_labels\":1,\n  \"num_channels\":32})\ncnn_3d_more_filters_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    4,\n    18,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"num_labels\":1,\n  \"num_channels\":64})\n\nunet_3d_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    4,\n    18,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 3,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"num_labels\":1,\n  \"num_channels\":16})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unet_3d_jumbo_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    6,\n    40,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 32,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"transformers_version\": \"4.12.0.dev0\",\n  \"num_labels\":1\n})\n\nunetr_multiclass_config = SegformerConfig(**{\n  \"architectures\": [\n    \"SegformerForImageClassification\"\n  ],\n  \"attention_probs_dropout_prob\": 0.0,\n  \"classifier_dropout_prob\": 0.1,\n  \"decoder_hidden_size\": 768,\n  \"depths\": [\n    3,\n    6,\n    40,\n    3\n  ],\n  \"downsampling_rates\": [\n    1,\n    4,\n    8,\n    16\n  ],\n  \"drop_path_rate\": 0.1,\n  \"hidden_act\": \"gelu\",\n  \"hidden_dropout_prob\": 0.0,\n  \"hidden_sizes\": [\n    64,\n    128,\n    320,\n    512\n  ],\n  \"image_size\": 224,\n  \"initializer_range\": 0.02,\n  \"layer_norm_eps\": 1e-06,\n  \"mlp_ratios\": [\n    4,\n    4,\n    4,\n    4\n  ],\n  \"model_type\": \"segformer\",\n  \"num_attention_heads\": [\n    1,\n    2,\n    5,\n    8\n  ],\n  \"num_channels\": 32,\n  \"num_encoder_blocks\": 4,\n  \"patch_sizes\": [\n    7,\n    3,\n    3,\n    3\n  ],\n  \"sr_ratios\": [\n    8,\n    4,\n    2,\n    1\n  ],\n  \"strides\": [\n    4,\n    2,\n    2,\n    2\n  ],\n  \"torch_dtype\": \"float32\",\n  \"transformers_version\": \"4.12.0.dev0\",\n  \"num_labels\":3\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from unetr import UNETR\nclass UNETR_Segformer(nn.Module):\n    def __init__(self, cfg, dropout = .2):\n        super().__init__()\n        self.cfg = cfg\n        self.dropout = nn.Dropout2d(dropout)\n        self.encoder = UNETR(\n            in_channels=1,\n            out_channels=32,\n            img_size=(16, self.cfg.size, self.cfg.size),\n            conv_block=True\n        )\n        self.encoder_2d = SegformerForSemanticSegmentation(self.cfg.segformer_config)\n        self.upscaler1 = nn.ConvTranspose2d(\n            1, 1, kernel_size=(4, 4), stride=2, padding=1)\n        self.upscaler2 = nn.ConvTranspose2d(\n            1, 1, kernel_size=(4, 4), stride=2, padding=1)\n\n    def forward(self, image):\n        output = self.encoder(image).max(axis=2)[0]\n        output = self.dropout(output)\n        output = self.encoder_2d(output).logits\n        output = self.upscaler1(output)\n        output = self.upscaler2(output)\n        return output\nclass UNETR_SegformerMC(nn.Module):\n    def __init__(self, cfg, dropout = .2):\n        super().__init__()\n        self.cfg = cfg\n        self.dropout = nn.Dropout2d(dropout)\n        self.encoder = UNETR(\n            in_channels=1,\n            out_channels=32,\n            img_size=(16, self.cfg.size, self.cfg.size),\n#             conv_block=True\n        )\n        self.encoder_2d = SegformerForSemanticSegmentation(self.cfg.segformer_config)\n        self.upscaler1 = nn.ConvTranspose2d(\n            3, 3, kernel_size=(4, 4), stride=2, padding=1)\n        self.upscaler2 = nn.ConvTranspose2d(\n            3, 3, kernel_size=(4, 4), stride=2, padding=1)\n\n    def forward(self, image):\n        output = self.encoder(image).max(axis=2)[0]\n        output = self.dropout(output)\n        output = self.encoder_2d(output).logits\n        output = self.upscaler1(output)\n        output = self.upscaler2(output)\n        return output[:, 2:, :, :]\n    \nclass cnn3d_segformer(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        \n        self.conv3d_1 = nn.Conv3d(1, 4, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n        self.conv3d_2 = nn.Conv3d(4, 8, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n        self.conv3d_3 = nn.Conv3d(8, 16, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n        self.conv3d_4 = nn.Conv3d(16, 32, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n\n        self.xy_encoder_2d = SegformerForSemanticSegmentation(self.cfg.segformer_config)\n        self.upscaler1 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n        self.upscaler2 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n        \n    def forward(self, image):\n        output = self.conv3d_1(image)\n        output = self.conv3d_2(output)\n        output = self.conv3d_3(output)\n        output = self.conv3d_4(output).max(axis = 2)[0]\n        output = self.xy_encoder_2d(output).logits\n        output = self.upscaler1(output)\n        output = self.upscaler2(output)\n        return output\n    \nclass cnn3d_segformer_more_filters(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        \n        self.conv3d_1 = nn.Conv3d(1, 4, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n        self.conv3d_2 = nn.Conv3d(4, 8, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n        self.conv3d_3 = nn.Conv3d(8, 16, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n        self.conv3d_4 = nn.Conv3d(16, 64, kernel_size=(3, 3, 3), stride=1, padding=(1, 1, 1))\n\n        self.xy_encoder_2d = SegformerForSemanticSegmentation(self.cfg.segformer_config)\n        self.upscaler1 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n        self.upscaler2 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n        \n    def forward(self, image):\n        output = self.conv3d_1(image)\n        output = self.conv3d_2(output)\n        output = self.conv3d_3(output)\n        output = self.conv3d_4(output).max(axis = 2)[0]\n        output = self.xy_encoder_2d(output).logits\n        output = self.upscaler1(output)\n        output = self.upscaler2(output)\n        return output\n    \nclass unet3d_segformer(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        \n        self.model = get_model({\"name\":\"UNet3D\", \"in_channels\":1, \"out_channels\":16, \"f_maps\":8, \"num_groups\":4, \"is_segmentation\":False})\n        self.encoder_2d = SegformerForSemanticSegmentation(self.cfg.segformer_config)\n        self.upscaler1 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n        self.upscaler2 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n    def forward(self, image):\n        output = self.model(image).max(axis = 2)[0]\n        output = self.encoder_2d(output).logits\n        output = self.upscaler1(output)\n        output = self.upscaler2(output)\n        return output\n    \nclass unet3d_segformer_jumbo(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        \n        self.model = get_model({\"name\":\"UNet3D\", \"in_channels\":1, \"out_channels\":32, \"f_maps\":8, \"num_groups\":4, \"is_segmentation\":False})\n        self.encoder_2d = SegformerForSemanticSegmentation(self.cfg.segformer_config)\n        self.upscaler1 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n        self.upscaler2 = nn.ConvTranspose2d(1, 1, kernel_size=(4, 4), stride = 2, padding=1)\n    def forward(self, image):\n        output = self.model(image).max(axis = 2)[0]\n        output = self.encoder_2d(output).logits\n        output = self.upscaler1(output)\n        output = self.upscaler2(output)\n        return output\n    \ndef build_model(cfg, model_arch = None):\n    print('model_name', cfg.model_name)\n    if model_arch == \"cnn3d\":\n        model = cnn3d_segformer(cfg)\n    if model_arch == \"cnn3d_more_filters\":\n        model = cnn3d_segformer_more_filters(cfg)\n    if model_arch == \"unet3d\":\n        model = unet3d_segformer(cfg)\n    if model_arch == \"unet3d_jumbo\":\n        model = unet3d_segformer_jumbo(cfg)\n    if model_arch == \"unetr\":\n        model = UNETR_Segformer(cfg)\n    if model_arch == \"unetr_mc\":\n        model = UNETR_SegformerMC(cfg)\n\n    return model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EnsembleModel:\n    def __init__(self, use_tta=False):\n        self.models = []\n        self.use_tta = use_tta\n    def tta_infer(self, model:nn.Module, x):\n        #x.shape=(batch,c,h,w)\n        shape=x.shape\n        x=[x,*[torch.rot90(x,k=i,dims=(-2,-1)) for i in range(1,4)]]\n        x=[model(single_x) for single_x in x]\n        x=torch.cat(x,dim=0)\n        x=x.reshape(4,shape[0],*shape[3:])\n        x=[torch.rot90(x[i],k=-i,dims=(-2,-1)) for i in range(4)]\n        x=torch.stack(x,dim=0)\n        return x.mean(0)\n                \n    def __call__(self, x):\n        if self.use_tta:\n            outputs = [self.tta_infer(model, x).to('cpu').numpy()\n                   for model in self.models]\n        else:\n            outputs = [model(x).mean(axis = 1).to('cpu').numpy()\n                       for model in self.models]\n        avg_preds = np.mean(outputs, axis=0)\n        return avg_preds\n\n    def add_model(self, model):\n        self.models.append(model)\n\ndef build_ensemble_model(model_path, model_arch):\n    model = EnsembleModel(use_tta = True)\n    _model = build_model(CFG, model_arch)\n    _model.to(device)\n    state = torch.load(model_path)\n    try:\n        _model.load_state_dict(state)\n    except:\n        _model = nn.DataParallel(_model)\n        _model.load_state_dict(state)\n    _model.eval()\n\n    model.add_model(_model)\n    \n    return model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if mode == 'test':\n    fragment_ids = sorted(os.listdir(CFG.comp_dataset_path + mode))\nelse:\n    fragment_ids = [3]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_tuples = [\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_1024_3dcnn_segformer_best.pth\", \"segformer_config\": cnn_3d_config, \"score\": .75},\n    {\"model_arch\": \"unet3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n     \"weight_path\": \"/kaggle/input/3d-unet/3d_unet_segformer_1024_3d_unet_segformer_final_all_train.pth\", \"segformer_config\": unet_3d_config, \"score\":.78},\n#     {\"model_arch\": \"unet3d\", \"tile_size\": 512, \"size\": 512, \"batch_size\": 4,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3d_unet_segformer_512_3d_unet_segformer_final.pth\", \"segformer_config\": unet_3d_config, \"score\":.77},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_1024_full_train_3dcnn_segformer_final.pth\", \"segformer_config\": cnn_3d_config, \"score\":.76},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_1024_swa_slow_3dcnn_segformer_final_swa.pth\", \"segformer_config\": cnn_3d_config, \"score\".74},\n#     {\"model_arch\": \"unet3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dunet_segformer_1024_swa_slow_3dunet_segformer_final_swa.pth\", \"segformer_config\": unet_3d_config, \"score\":.75},\n    {\"model_arch\": \"unet3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n     \"weight_path\": \"/kaggle/input/3d-unet/3dunet_segformer_1024_swa_slow_all_train_3dunet_segformer_final.pth\", \"segformer_config\": unet_3d_config, \"score\":.78},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_all_train_swa_3dcnn_segformer_10_final.pth\", \"segformer_config\": cnn_3d_config, \"score\":.76},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_all_train_swa_3dcnn_segformer_15_final.pth\", \"segformer_config\": cnn_3d_config, \"score\":.76},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_all_train_swa_3dcnn_segformer_20_final.pth\", \"segformer_config\": cnn_3d_config, \"score\":.77},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_all_train_swa_3dcnn_segformer_25_final.pth\", \"segformer_config\": cnn_3d_config},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 3,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_all_train_swa_3dcnn_segformer_final_swa.pth\", \"segformer_config\": cnn_3d_config, \"score\":.78},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 5,\n#      \"weight_path\": \"/kaggle/input/3d-unet/b1_3dcnn_segformer_b1_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b1_config, \"score\":.71},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 4,\n#      \"weight_path\": \"/kaggle/input/3d-unet/b2_3dcnn_segformer_b2_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b2_config, \"score\":.68},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 2,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_segformer_b4_3dcnn_segformer_b4_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b4_config, \"score\":.74},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_bigsegformer_3dcnn_bigsegformer_final.pth\", \"segformer_config\": cnn_3d_segformer_b5_config, \"score\":.76},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 2,\n#      \"weight_path\": \"/kaggle/input/3d-unet/b5_long_train_all_frags_3dcnn_segformer_b5_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b5_config, \"score\": .77},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/b5_long_train_all_frags_3dcnn_segformer_b5_10_final.pth\", \"segformer_config\": cnn_3d_segformer_b5_config, \"score\": .74},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/b5_long_train_all_frags_3dcnn_segformer_b5_30_final.pth\", \"segformer_config\": cnn_3d_segformer_b5_config, \"score\":.76},\n#     {\"model_arch\": \"cnn3d_more_filters\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/b3_more_fmaps_3dcnn_segformerb364_final_swa.pth\", \"segformer_config\": cnn_3d_more_filters_config, \"score\": .74},\n    {\"model_arch\": \"cnn3d_more_filters\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 2,\n     \"weight_path\": \"/kaggle/input/3d-unet/b3_more_fmaps_all_train_3dcnn_segformerb364_final_swa.pth\", \"segformer_config\": cnn_3d_more_filters_config, \"score\":.78},\n#     {\"model_arch\": \"unet3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3d_unet_3dunet_b3_final_swa.pth\", \"segformer_config\": unet_3d_config, \"score\":.73},\n#     {\"model_arch\": \"unet3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3d_unet_all_train_3dunet_b3_final_swa.pth\", \"segformer_config\": unet_3d_config, \"score\":.76},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 512, \"size\": 512, \"batch_size\": 4,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_512_b2_all_train_3dcnn_b2_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b2_config},\n    # ran at wrong resolution. Scored .73\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 2,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_768_b4_adam_3dcnn_b4_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b4_config, \"score\":.73},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 768, \"size\": 768, \"batch_size\": 2,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_768_b4_adam_3dcnn_b4_final_swa.pth\", \"segformer_config\": cnn_3d_segformer_b4_config, \"score\":.74},\n#     {\"model_arch\": \"cnn3d\", \"tile_size\": 768, \"size\": 768, \"batch_size\": 2,\n#      \"weight_path\": \"/kaggle/input/3d-unet/3dcnn_768_b4_adam_3dcnn_b4_final_swa_all_train.pth\", \"segformer_config\": cnn_3d_segformer_b4_config, \"score\":.75},\n    {\"model_arch\": \"unet3d_jumbo\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n     \"weight_path\": \"/kaggle/input/3d-unet/Jumbo_Unet_Jumbo_Unet_69_final_swa_all_train.pth\", \"segformer_config\": unet_3d_jumbo_config, \"score\":.79},\n#     {\"model_arch\": \"unet3d_jumbo\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/Jumbo_Unet_Jumbo_Unet_69_new_label_final_swa_all_train.pth\", \"segformer_config\": unet_3d_jumbo_config, \"score\":.77},\n#     {\"model_arch\": \"unet3d_jumbo\", \"tile_size\": 1024, \"size\": 1024, \"batch_size\": 1,\n#      \"weight_path\": \"/kaggle/input/3d-unet/Jumbo_Unet_Jumbo_Unet_5_final_swa_all_train.pth\", \"segformer_config\": unet_3d_jumbo_config, \"score\": .69},\n#     {\"model_arch\": \"unetr\", \"tile_size\": 512, \"size\": 512, \"batch_size\": 8,\n#      \"weight_path\": \"/kaggle/input/3d-unet/jumbo_unetr_unetr_1245_final_swa_all_train.pth\", \"segformer_config\": unet_3d_jumbo_config},\n#     {\"model_arch\": \"unetr_mc\", \"tile_size\": 512, \"size\": 512, \"batch_size\": 8,\n#      \"weight_path\": \"/kaggle/input/3d-unet/unetr_multiclass_512_b5_unet_final_4Ryan.pth\", \"segformer_config\": unetr_multiclass_config, \"score\":.77},\n    {\"model_arch\": \"unetr\", \"tile_size\": 512, \"size\": 512, \"batch_size\": 8,\n     \"weight_path\": \"/kaggle/input/3d-unet/jumbo_unetr_unetr_888_final_swa_all_train_long.pth\", \"segformer_config\": unet_3d_jumbo_config, \"score\":.82},\n    {\"model_arch\": \"unetr_mc\", \"tile_size\": 512, \"size\": 512, \"batch_size\": 8,\n     \"weight_path\": \"/kaggle/input/3d-unet/unetr_multiclass_NOVALIDATION_512_b5_unet_final_swa_all_train.pth\", \"segformer_config\": unetr_multiclass_config},\n\n]\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## main","metadata":{}},{"cell_type":"code","source":"def post_process(probability, threshold, min_size = 20000):\n    \"\"\"\n    Post processing of each predicted mask, components with lesser number of pixels\n    than `min_size` are ignored\n    \"\"\"\n    # don't remember where I saw it\n    mask = cv2.threshold(probability, threshold, 1, cv2.THRESH_BINARY)[1]\n    num_component, component = cv2.connectedComponents(mask.astype(np.uint8))\n    predictions = np.zeros_like(probability, np.float32)\n    num = 0\n    for c in range(1, num_component):\n        p = (component == c)\n        if p.sum() > min_size:\n            predictions[p] = 1\n    return predictions","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nimport time","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = []\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    for fragment_id in fragment_ids:\n        mask_preds = None\n        last_res = None\n        for model_config in model_tuples:\n            mask_pred = None\n            mask_count = None\n            if last_res != model_config[\"size\"]:\n                for file in glob.glob(\"*.npy\"):\n                    os.remove(file)\n                last_res = model_config[\"size\"]\n            CFG.tile_size = model_config[\"tile_size\"]\n            CFG.size = model_config[\"size\"]\n            CFG.batch_size = model_config[\"batch_size\"]\n            CFG.stride = CFG.tile_size // 4\n            CFG.valid_aug_list = [\n                A.Resize(CFG.size, CFG.size),\n                A.Normalize(\n                    mean= [0] * CFG.in_chans,\n                    std= [1] * CFG.in_chans\n                ),\n                ToTensorV2(transpose_mask=True),\n            ]\n            CFG.segformer_config = model_config[\"segformer_config\"]\n            model = build_ensemble_model(model_config[\"weight_path\"], model_config[\"model_arch\"])\n            test_loader, xyxys = make_test_dataset(fragment_id)\n\n            binary_mask = cv2.imread(CFG.comp_dataset_path + f\"{mode}/{fragment_id}/mask.png\", 0)\n            binary_mask = (binary_mask / 255).astype(int)\n\n            ori_h = binary_mask.shape[0]\n            ori_w = binary_mask.shape[1]\n            # mask = mask / 255\n\n            pad0 = (CFG.tile_size - binary_mask.shape[0] % CFG.tile_size)\n            pad1 = (CFG.tile_size - binary_mask.shape[1] % CFG.tile_size)\n\n            binary_mask = np.pad(binary_mask, [(0, pad0), (0, pad1)], constant_values=0)\n            if mask_pred is None:\n                mask_pred = np.zeros(binary_mask.shape)\n                mask_count = np.zeros(binary_mask.shape)\n\n            for step, (images) in tqdm(enumerate(test_loader), total=len(test_loader)):\n                images = images.to(device)\n                batch_size = images.size(0)\n                with autocast():            \n                    with torch.no_grad():\n                        y_preds = model(images)\n\n                start_idx = step*CFG.batch_size\n                end_idx = start_idx + batch_size\n                for i, (x1, y1, x2, y2) in enumerate(xyxys[start_idx:end_idx]):\n                    mask_pred[y1:y2, x1:x2] += y_preds[i]\n                    mask_count[y1:y2, x1:x2] += np.ones((CFG.tile_size, CFG.tile_size))\n            del test_loader\n            del model\n            gc.collect()\n            torch.cuda.empty_cache()\n            mask_pred = mask_pred[:ori_h, :ori_w]\n            mask_count = mask_count[:ori_h, :ori_w]\n            binary_mask = binary_mask[:ori_h, :ori_w]\n\n            print(f'mask_count_min: {mask_count.min()}')\n            mask_pred = mask_pred/mask_count\n            mask_pred = torch.sigmoid(torch.tensor(mask_pred)).numpy()\n            if mask_preds is None:\n                mask_preds = mask_pred/len(model_tuples)\n            else:\n                mask_preds += mask_pred/len(model_tuples)\n\n        mask_pred = (mask_preds >= TH).astype(int)\n        mask_pred *= binary_mask\n        mask_pred = post_process(mask_pred.astype(float), TH, 10000).astype(int)\n        plt.imshow(mask_pred)\n        inklabels_rle = rle(mask_pred)\n        results.append((fragment_id, inklabels_rle))\n        del mask_pred, mask_count\n        gc.collect()\n        torch.cuda.empty_cache()\n        for file in glob.glob(\"*.npy\"):\n            os.remove(file)\nelse:\n    pass\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## submission","metadata":{}},{"cell_type":"code","source":"sub = pd.DataFrame(results, columns=['Id', 'Predicted'])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_sub = pd.read_csv(CFG.comp_dataset_path + 'sample_submission.csv')\nsample_sub = pd.merge(sample_sub[['Id']], sub, on='Id', how='left')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_sub","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_sub.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}