{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports & Environment Setup","metadata":{}},{"cell_type":"code","source":"#libraries\nimport os, time, random, shutil\nfrom pathlib import Path\nfrom io import BytesIO\nimport math\nimport zipfile\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\nfrom PIL import Image, ImageSequence\nimport matplotlib.pyplot as plt\n%matplotlib inline\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\nfrom skimage import measure\nimport tifffile\n\n# optional interactive widgets\ntry:\n    from ipywidgets import interact, IntSlider, Dropdown\nexcept Exception:\n    interact = None\n\nprint(\"torch:\", torch.__version__)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:18:18.575907Z","iopub.execute_input":"2025-11-16T15:18:18.576441Z","iopub.status.idle":"2025-11-16T15:18:22.787435Z","shell.execute_reply.started":"2025-11-16T15:18:18.576416Z","shell.execute_reply":"2025-11-16T15:18:22.786633Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Path Setup & Training Configuration","metadata":{}},{"cell_type":"code","source":"#parameters\nROOT = \"/kaggle/input/vesuvius-challenge-surface-detection\"   # <-- change to your root\nTRAIN_IMG_DIR = os.path.join(ROOT, \"train_images\")\nTRAIN_LABEL_DIR = os.path.join(ROOT, \"train_labels\")\nTEST_IMG_DIR  = os.path.join(ROOT, \"test_images\")\nTRAIN_CSV = os.path.join(ROOT, \"train.csv\")\nTEST_CSV  = os.path.join(ROOT, \"test.csv\")\n\nOUT_DIR = \"./predictions\"\nCKPT_DIR = \"./checkpoints\"\nCACHE_DIR = \"./cache_npy\"\nos.makedirs(OUT_DIR, exist_ok=True)\nos.makedirs(CKPT_DIR, exist_ok=True)\nos.makedirs(CACHE_DIR, exist_ok=True)\n\n\n#Training/dev control\nN_TRAIN_FILES = 500            #<= use only first N files (set None to use all)\nUSE_RANDOM_SAMPLE = False      #sample random N instead of first N\nPATCH_SIZE = (32, 128, 128)    #(D,H,W) for real run; use smaller for fast dev\nBATCH_SIZE = 1\nNUM_EPOCHS = 10\nLR = 1e-4\nNUM_WORKERS = 4\nSTEPS_PER_EPOCH = None         #set small int for dev (e.g. 400)\nVAL_N_SAMPLES = 50\n\n#Cache safety limits)\nMAX_CACHE_FILES = 50          #create at most this many .npy files\nMIN_FREE_BYTES = 1.5 * 1024**3  #keep ~1.5 GB free\n\n#Dev options (quick iteration)\nDEV_PATCH = (16, 96, 96)\nDEV_STEPS = 400","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:18:22.788563Z","iopub.execute_input":"2025-11-16T15:18:22.788882Z","iopub.status.idle":"2025-11-16T15:18:22.795246Z","shell.execute_reply.started":"2025-11-16T15:18:22.788864Z","shell.execute_reply":"2025-11-16T15:18:22.794494Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Storage & Cache Management Helpers","metadata":{}},{"cell_type":"code","source":"#helpers\ndef load_3d_tiff(path):\n    with Image.open(path) as img:\n        pages = [np.array(p) for p in ImageSequence.Iterator(img)]\n    return np.stack(pages, axis=0)\n\ndef write_tiff_safe(path, arr):\n    tifffile.imwrite(path, arr.astype(arr.dtype), compression=None)\n\ndef normalize_volume(vol):\n    vol = vol.astype(np.float32)\n    mu, sd = vol.mean(), vol.std()\n    if sd == 0: sd = 1.0\n    return (vol - mu) / sd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:18:22.795892Z","iopub.execute_input":"2025-11-16T15:18:22.796088Z","iopub.status.idle":"2025-11-16T15:18:22.811017Z","shell.execute_reply.started":"2025-11-16T15:18:22.796072Z","shell.execute_reply":"2025-11-16T15:18:22.810327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#caching\n_cache_meta_file = os.path.join(CACHE_DIR, \".cache_count\")\ntry:\n    _cached_count = int(open(_cache_meta_file).read().strip()) if os.path.exists(_cache_meta_file) else 0\nexcept Exception:\n    _cached_count = 0\n\ndef _inc_cache_count():\n    global _cached_count\n    _cached_count += 1\n    try:\n        with open(_cache_meta_file + \".tmp\", \"w\") as f:\n            f.write(str(_cached_count))\n        os.replace(_cache_meta_file + \".tmp\", _cache_meta_file)\n    except Exception:\n        pass\n\ndef _free_bytes(path=\".\"):\n    try:\n        return shutil.disk_usage(path).free\n    except Exception:\n        return 0\n\ndef safe_save_npy_atomic(arr, out_path):\n    tmp = out_path + \".tmp\"\n    try:\n        np.save(tmp, arr, allow_pickle=False)\n        os.replace(tmp, out_path)\n        return True\n    except Exception:\n        try:\n            if os.path.exists(tmp):\n                os.remove(tmp)\n        except Exception:\n            pass\n        return False\n\ndef cache_if_allowed(tif_path, cache_dir=CACHE_DIR):\n    \"\"\"\n    Create a normalized .npy cache if we haven't exceeded limits and there is free space.\n    Returns cache_path or None.\n    \"\"\"\n    global _cached_count\n    os.makedirs(cache_dir, exist_ok=True)\n    out_path = os.path.join(cache_dir, Path(tif_path).stem + \".npy\")\n    if os.path.exists(out_path):\n        return out_path\n    if _cached_count >= MAX_CACHE_FILES:\n        return None\n    if _free_bytes(cache_dir) < MIN_FREE_BYTES:\n        return None\n    # read and normalize then save atomically\n    try:\n        vol = load_3d_tiff(tif_path).astype(np.float32)\n        vol = normalize_volume(vol)\n        ok = safe_save_npy_atomic(vol, out_path)\n        if ok:\n            _inc_cache_count()\n            return out_path\n    except Exception:\n        return None\n    return None\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:18:22.812418Z","iopub.execute_input":"2025-11-16T15:18:22.812607Z","iopub.status.idle":"2025-11-16T15:18:22.825497Z","shell.execute_reply.started":"2025-11-16T15:18:22.812593Z","shell.execute_reply":"2025-11-16T15:18:22.824806Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Metadata Load + Basic Data Inspection","metadata":{}},{"cell_type":"code","source":"#gather shapes & label stats\nmeta = pd.read_csv(TRAIN_CSV)\nprint(\"Total train volumes:\", len(meta))\ndisplay(meta.head())\n\n\nsample_ids = meta['id'].tolist()[:min(200, len(meta))]\nshapes = {}\nvoxel_counts = {}\nclass_voxel_counts = []\nfor vid in tqdm(sample_ids, desc=\"Gather shapes (sample)\"):\n    imgp = os.path.join(TRAIN_IMG_DIR, f\"{vid}.tif\")\n    lblp = os.path.join(TRAIN_LABEL_DIR, f\"{vid}.tif\")\n    if not os.path.exists(imgp): continue\n    vol = load_3d_tiff(imgp)\n    shapes[vid] = vol.shape\n    voxel_counts[vid] = vol.size\n    if os.path.exists(lblp):\n        lbl = load_3d_tiff(lblp)\n        classes, counts = np.unique(lbl, return_counts=True)\n        d = dict(zip(classes, counts))\n        class_voxel_counts.append({\"id\": vid, **{f\"class_{c}\": d.get(c, 0) for c in [0,1,2]}})\n    else:\n        class_voxel_counts.append({\"id\": vid, \"class_0\": 0, \"class_1\": 0, \"class_2\": 0})\n\nshapes_df = pd.DataFrame.from_dict(shapes, orient='index', columns=['D','H','W'])\nshapes_df.index.name='id'\nprint(\"\\nShapes summary (sample):\")\ndisplay(shapes_df.describe())\n\nif len(class_voxel_counts) > 0:\n    class_df = pd.DataFrame(class_voxel_counts).set_index('id')\n    print(\"\\nLabel counts summary (sample):\")\n    display(class_df.describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:18:22.826082Z","iopub.execute_input":"2025-11-16T15:18:22.826280Z","iopub.status.idle":"2025-11-16T15:28:32.688413Z","shell.execute_reply.started":"2025-11-16T15:18:22.826257Z","shell.execute_reply":"2025-11-16T15:28:32.687758Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EDA: Histograms, Viewer & Slice Visualizations","metadata":{}},{"cell_type":"code","source":"#EDA and Visualization\nimport math\nimport random\nfrom matplotlib import cm\n\n#pick a few sample ids for histograms / grid\nsample_hist_ids = random.sample(list(shapes.keys()), min(6, max(1,len(shapes))))\nplt.figure(figsize=(12,8))\nfor i, sid in enumerate(sample_hist_ids):\n    vol = load_3d_tiff(os.path.join(TRAIN_IMG_DIR, f\"{sid}.tif\"))\n    plt.subplot(3,2,i+1)\n    plt.hist(vol.ravel(), bins=200)\n    plt.title(f\"Intensity histogram - id {sid} (min {vol.min():.1f}, max {vol.max():.1f})\")\nplt.tight_layout()\nplt.show()\n\n#interactive viewer\nexample_id = None\nfor vid in meta['id'].tolist():\n    if os.path.exists(os.path.join(TRAIN_LABEL_DIR, f\"{vid}.tif\")):\n        example_id = vid; break\nif example_id is None:\n    example_id = meta['id'].iloc[0]\nvol = load_3d_tiff(os.path.join(TRAIN_IMG_DIR, f\"{example_id}.tif\")).astype(np.float32)\nlbl = load_3d_tiff(os.path.join(TRAIN_LABEL_DIR, f\"{example_id}.tif\")).astype(np.uint8) if os.path.exists(os.path.join(TRAIN_LABEL_DIR, f\"{example_id}.tif\")) else None\nprint(\"Example id:\", example_id, \"shape:\", vol.shape)\n\ndef to_rgb(gray):\n    g = gray - gray.min(); g = g / (g.max()+1e-8); g = (g*255).astype(np.uint8)\n    return np.stack([g]*3, axis=-1)\n\ndef apply_colormap(gray, cmap='magma'):\n    normed = (gray - gray.min())/(gray.max()-gray.min()+1e-8)\n    return (cm.get_cmap(cmap)(normed)[...,:3]*255).astype(np.uint8)\n\ndef show_slice(idx=0, mode='gray'):\n    sl = vol[idx]\n    if mode=='gray':\n        rgb = to_rgb(sl)\n    else:\n        rgb = apply_colormap(sl, cmap=mode)\n    plt.figure(figsize=(6,6)); plt.imshow(rgb)\n    if lbl is not None:\n        mask = lbl[idx]==1\n        mask_rgb = np.zeros_like(rgb); mask_rgb[mask] = [255,0,0]\n        plt.imshow(mask_rgb, alpha=0.35)\n    plt.title(f\"id {example_id} slice {idx}\"); plt.axis('off'); plt.show()\n\nif interact is not None:\n    interact(show_slice, idx=IntSlider(min=0, max=vol.shape[0]-1, value=vol.shape[0]//2),\n             mode=Dropdown(options=['gray','magma','viridis','plasma'], value='gray'))\nelse:\n    print(\"ipywidgets not available; show middle slice:\")\n    show_slice(vol.shape[0]//2, mode='gray')\n\n#central-slice montage for sample volumes\nsample_grid_ids = list(shapes.keys())[:min(6, len(shapes))]\nn = len(sample_grid_ids)\nfig, axs = plt.subplots(2, math.ceil(n/2), figsize=(4*math.ceil(n/2),8))\naxs = axs.flatten()\nfor i, sid in enumerate(sample_grid_ids):\n    v = load_3d_tiff(os.path.join(TRAIN_IMG_DIR, f\"{sid}.tif\"))\n    mid = v.shape[0]//2\n    axs[i].imshow(v[mid], cmap='gray')\n    axs[i].set_title(f\"id {sid} (slice {mid})\")\n    axs[i].axis('off')\nplt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:32.689211Z","iopub.execute_input":"2025-11-16T15:28:32.689532Z","iopub.status.idle":"2025-11-16T15:28:55.235951Z","shell.execute_reply.started":"2025-11-16T15:28:32.689512Z","shell.execute_reply":"2025-11-16T15:28:55.234949Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Class (Patch Sampling + Foreground Sampling)","metadata":{}},{"cell_type":"code","source":"\nclass VesuviusDatasetFast(Dataset):\n    def __init__(self, img_paths, mask_paths=None, patch_size=PATCH_SIZE, mode='train', cache_dir=CACHE_DIR, sample_foreground_prob=0.75):\n        self.img_paths = [str(p) for p in img_paths]\n        self.mask_paths = [str(p) for p in mask_paths] if mask_paths is not None else None\n        self.patch_size = tuple(patch_size)\n        self.mode = mode\n        self.cache_dir = cache_dir\n        self.sample_foreground_prob = sample_foreground_prob if mode==\"train\" else 0.0\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def _read(self, path):\n        # Prefer cache if exists or can be created (limited)\n        if self.cache_dir:\n            cached = os.path.join(self.cache_dir, Path(path).stem + \".npy\")\n            if os.path.exists(cached):\n                try:\n                    return np.load(cached)\n                except Exception:\n                    try:\n                        os.remove(cached)\n                    except Exception:\n                        pass\n            # try to create limited cache\n            created = cache_if_allowed(path, cache_dir=self.cache_dir)\n            if created:\n                try:\n                    return np.load(created)\n                except Exception:\n                    pass\n        # fallback read from tif every time\n        vol = load_3d_tiff(path).astype(np.float32)\n        return normalize_volume(vol)\n\n    def __getitem__(self, idx):\n        if self.mode == 'train':\n            # sample random volume index for better mixing\n            idx = random.choice(range(len(self.img_paths)))\n        img_path = self.img_paths[idx]\n        vol = self._read(img_path)  # Z,H,W\n        if self.mask_paths is None:\n            mask = np.zeros_like(vol, dtype=np.uint8)\n        else:\n            lblp = self.mask_paths[idx]\n            if os.path.exists(lblp):\n                mask = load_3d_tiff(lblp).astype(np.uint8)\n            else:\n                mask = np.zeros_like(vol, dtype=np.uint8)\n\n        Z,H,W = vol.shape\n        pd,ph,pw = self.patch_size\n\n        want_fg = (self.mode=='train') and (random.random() < self.sample_foreground_prob) and np.any(mask==1)\n        if want_fg:\n            fg = np.argwhere(mask==1)\n            i = random.randrange(len(fg))\n            cz,cy,cx = fg[i]\n            z0 = max(0, min(Z-pd, cz - random.randint(0,pd-1)))\n            y0 = max(0, min(H-ph, cy - random.randint(0,ph-1)))\n            x0 = max(0, min(W-pw, cx - random.randint(0,pw-1)))\n        else:\n            z0 = random.randint(0, max(0, Z-pd))\n            y0 = random.randint(0, max(0, H-ph))\n            x0 = random.randint(0, max(0, W-pw))\n\n        vpatch = vol[z0:z0+pd, y0:y0+ph, x0:x0+pw]\n        mpatch = mask[z0:z0+pd, y0:y0+ph, x0:x0+pw].astype(np.uint8)\n\n        v_t = torch.from_numpy(vpatch[None]).float()   # 1 x D x H x W\n        m_t = torch.from_numpy(mpatch[None]).long()    # 1 x D x H x W (for CE)\n        return v_t, m_t\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.237045Z","iopub.execute_input":"2025-11-16T15:28:55.237631Z","iopub.status.idle":"2025-11-16T15:28:55.249943Z","shell.execute_reply.started":"2025-11-16T15:28:55.237599Z","shell.execute_reply":"2025-11-16T15:28:55.249394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"USE_PYVOI = False\ntry:\n    import pyvoi\n    USE_PYVOI = True\nexcept Exception:\n    USE_PYVOI = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.250623Z","iopub.execute_input":"2025-11-16T15:28:55.250792Z","iopub.status.idle":"2025-11-16T15:28:55.265073Z","shell.execute_reply.started":"2025-11-16T15:28:55.250777Z","shell.execute_reply":"2025-11-16T15:28:55.264578Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Metric Functions (Dice, Surface Dice, TopoScore, VOI)","metadata":{}},{"cell_type":"code","source":"#surface Dice + placeholders for TopoScore and VOI\nimport numpy as np\nfrom scipy import ndimage as ndi\nfrom skimage import morphology\nfrom typing import Tuple, Dict\n\n\n#Surface Dice\ndef surface_points(binary: np.ndarray) -> np.ndarray:\n    se = morphology.ball(1)  # spherical structuring element radius=1\n    eroded = ndi.binary_erosion(binary.astype(bool), structure=se)\n    surf = binary.astype(bool) & (~eroded)\n    return surf\n\ndef surface_dice(pred: np.ndarray, gt: np.ndarray, tol_vox: int = 2) -> float:\n    \"\"\"\n    Compute surface dice (symmetric) between pred and gt binary masks.\n    tol_vox: tolerance in voxels.\n    \"\"\"\n    p = (pred > 0).astype(np.uint8)\n    g = (gt > 0).astype(np.uint8)\n\n    sp = surface_points(p)\n    sg = surface_points(g)\n\n    if not sp.any() and not sg.any():\n        return 1.0\n    if not sp.any() or not sg.any():\n        return 0.0\n\n    dt_g = ndi.distance_transform_edt(~sg)\n    dt_p = ndi.distance_transform_edt(~sp)\n\n    d_p_to_g = dt_g[sp]\n    d_g_to_p = dt_p[sg]\n\n    a = float((d_p_to_g <= tol_vox).sum()) / float(sp.sum())\n    b = float((d_g_to_p <= tol_vox).sum()) / float(sg.sum())\n\n    if (a + b) == 0:\n        return 0.0\n    surf_dice = 2.0 * a * b / (a + b)\n    return float(surf_dice)\n\n\n\n#VOI (Variation of Information)\ndef voi_numpy(pred: np.ndarray, gt: np.ndarray) -> float:\n    \"\"\"Simple VI (Meila) implementation for binary masks (raw VI, lower is better).\"\"\"\n    a = (pred > 0).ravel().astype(np.int64)\n    b = (gt  > 0).ravel().astype(np.int64)\n    if a.size == 0:\n        return 0.0\n    vals_a = a\n    vals_b = b\n    # build 2x2 joint\n    joint = np.zeros((2,2), dtype=np.float64)\n    for i in (0,1):\n        for j in (0,1):\n            joint[i,j] = float(np.logical_and(vals_a==i, vals_b==j).sum())\n    total = joint.sum()\n    if total == 0:\n        return 0.0\n    pij = joint / total\n    pa = pij.sum(axis=1)\n    pb = pij.sum(axis=0)\n    def H(p):\n        p = p[p>0]\n        return -np.sum(p * np.log(p))\n    Ha = H(pa)\n    Hb = H(pb)\n    MI = 0.0\n    for i in range(pij.shape[0]):\n        for j in range(pij.shape[1]):\n            if pij[i,j] > 0:\n                MI += pij[i,j] * (math.log(pij[i,j]) - math.log(pa[i]) - math.log(pb[j]))\n    VI = Ha + Hb - 2.0 * MI\n    return float(VI)\n\ndef voi_wrapper(pred: np.ndarray, gt: np.ndarray) -> float:\n    \"\"\"Use pyvoi if present, otherwise fallback to voi_numpy.\"\"\"\n    if USE_PYVOI:\n        #pyvoi expects flattened arrays of ints\n        try:\n            vi = pyvoi.VI(pred.ravel().astype(np.int64), gt.ravel().astype(np.int64), torch=False)\n            if isinstance(vi, (tuple, list)):\n                return float(vi[0])\n            return float(vi)\n        except Exception:\n            return voi_numpy(pred, gt)\n    else:\n        return voi_numpy(pred, gt)\n\n#TopoScore (approximation)\ndef toposcore_approx(pred: np.ndarray, gt: np.ndarray) -> float:\n    \"\"\"\n    Heuristic TopoScore approximation:\n      - beta0 similarity based on 3D connected components (counts)\n      - beta1 similarity approximated by mean 2D hole counts in central slices\n    Returns score in [0,1] (1 = best).\n    Replace with official TopoScore implementation for exact competition scoring.\n    \"\"\"\n    p = (pred > 0).astype(np.uint8)\n    g = (gt  > 0).astype(np.uint8)\n\n    lab_p = measure.label(p, connectivity=1)\n    lab_g = measure.label(g, connectivity=1)\n    np_p = int(lab_p.max())\n    np_g = int(lab_g.max())\n\n    denom = max(1, (np_p + np_g))\n    beta0_score = 1.0 - abs(np_p - np_g) / denom\n\n    #approximate holes by looking at central slices\n    def approx_holes_slice(mask3d):\n        D = mask3d.shape[0]\n        if D <= 5:\n            sel = range(0, D)\n        else:\n            mid = D // 2\n            sel = range(max(0, mid-2), min(D, mid+3))\n        holes = []\n        for z in sel:\n            sl = mask3d[z].astype(np.uint8)\n            lab = measure.label(sl==0, connectivity=1)\n            #border-touching labels are not holes\n            border_mask = np.zeros_like(sl, dtype=bool)\n            border_mask[0,:] = border_mask[-1,:] = border_mask[:,0] = border_mask[:,-1] = True\n            border_labels = np.unique(lab[border_mask])\n            all_labels = np.unique(lab)\n            hole_labels = [L for L in all_labels if L != 0 and L not in border_labels]\n            holes.append(len(hole_labels))\n        return float(np.mean(holes)) if holes else 0.0\n\n    hp = approx_holes_slice(p)\n    hg = approx_holes_slice(g)\n    denom_h = max(1.0, (hp + hg))\n    beta1_score = 1.0 - abs(hp - hg) / denom_h\n\n    topo = 0.6 * beta0_score + 0.4 * beta1_score\n    topo = min(1.0, max(0.0, topo))\n    return float(topo)\n\n#Evaluator wrapper: compute all three metrics for a single volume\ndef compute_metrics_for_volume(pred: np.ndarray, gt: np.ndarray, tol_vox:int=2, min_voxels:int=250) -> Dict[str, float]:\n    \"\"\"\n    pred: predicted label volume (0/1/..)\n    gt: ground-truth label volume\n    returns dict: dice, surface_dice, toposcore, voi\n    \"\"\"\n    #ensure same shape\n    if pred.shape != gt.shape:\n        raise ValueError(f\"pred and gt shapes differ: {pred.shape} vs {gt.shape}\")\n\n    #remove small objects from pred as postprocessing (same you used in pipeline)\n    pclean = remove_small_objects(pred, min_voxels=min_voxels)\n    gclean = (gt == 1).astype(np.uint8)\n\n    #Simple Dice\n    pfg = (pclean == 1)\n    gfg = (gclean == 1)\n    inter = np.logical_and(pfg, gfg).sum()\n    denom = pfg.sum() + gfg.sum()\n    dice = float((2.*inter)/denom) if denom>0 else 1.0\n\n    s_dice = surface_dice(pclean, gclean, tol_vox=tol_vox)\n    topo = toposcore_approx(pclean, gclean)\n    voi = voi_wrapper(pclean, gclean)\n\n    return {'dice': dice, 'surface_dice': s_dice, 'toposcore': topo, 'voi': voi}\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.265737Z","iopub.execute_input":"2025-11-16T15:28:55.265902Z","iopub.status.idle":"2025-11-16T15:28:55.550753Z","shell.execute_reply.started":"2025-11-16T15:28:55.265888Z","shell.execute_reply":"2025-11-16T15:28:55.550192Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Safe DataLoader Creation","metadata":{}},{"cell_type":"code","source":"#safe dataloader creator (tries fork then spawn then fallback to num_workers=0)\nimport multiprocessing as mp\ndef make_dataloader_safe(dataset, batch_size=1, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True, persistent_workers=False):\n    try:\n        ctx = mp.get_context('fork')\n        return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle, num_workers=num_workers,\n                          pin_memory=pin_memory, persistent_workers=persistent_workers, multiprocessing_context=ctx)\n    except Exception as e_fork:\n        try:\n            ctx = mp.get_context('spawn')\n            return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle, num_workers=num_workers,\n                              pin_memory=pin_memory, persistent_workers=persistent_workers, multiprocessing_context=ctx)\n        except Exception as e_spawn:\n            print(\"Warning: falling back to num_workers=0 for stability.\")\n            print(\"fork error:\", e_fork)\n            print(\"spawn error:\", e_spawn)\n            return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle, num_workers=0, pin_memory=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.552822Z","iopub.execute_input":"2025-11-16T15:28:55.553171Z","iopub.status.idle":"2025-11-16T15:28:55.558210Z","shell.execute_reply.started":"2025-11-16T15:28:55.553154Z","shell.execute_reply":"2025-11-16T15:28:55.557517Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# UNet3DLight Model","metadata":{}},{"cell_type":"code","source":"class Conv3dBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, 3, padding=1, bias=False),\n            nn.InstanceNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, 3, padding=1, bias=False),\n            nn.InstanceNorm3d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n    def forward(self, x): return self.net(x)\n\nclass UpConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.up = nn.ConvTranspose3d(in_ch, out_ch, 2, 2)\n    def forward(self, x): return self.up(x)\n\nclass UNet3DLight(nn.Module):\n    def __init__(self, in_ch=1, out_ch=2, features=[16,32,64,128]):\n        super().__init__()\n        self.features = features\n\n        # Encoder\n        self.encs = nn.ModuleList()\n        self.pools = nn.ModuleList()\n        prev_ch = in_ch\n        for f in features:\n            self.encs.append(Conv3dBlock(prev_ch, f))\n            self.pools.append(nn.MaxPool3d(2))\n            prev_ch = f\n\n        # Bottleneck\n        bottleneck_ch = prev_ch * 2\n        self.bottleneck = Conv3dBlock(prev_ch, bottleneck_ch)\n\n        # Decoder / up path\n        self.upconvs = nn.ModuleList()\n        self.decs = nn.ModuleList()\n        curr_ch = bottleneck_ch\n        # iterate features reversed: for each skip feature f_skip, upconv maps curr_ch -> f_skip\n        for f_skip in reversed(features):\n            self.upconvs.append(nn.ConvTranspose3d(curr_ch, f_skip, kernel_size=2, stride=2))\n            # after upconv we will concat [up_out (f_skip) , skip (f_skip)] => 2*f_skip input to decoder -> produce f_skip\n            self.decs.append(Conv3dBlock(f_skip*2, f_skip))\n            curr_ch = f_skip  # next loop, curr_ch equals this decoder's output channels\n\n        # final conv\n        self.final = nn.Conv3d(features[0], out_ch, kernel_size=1)\n\n    def forward(self, x):\n        skips = []\n        out = x\n        # encoder\n        for enc, pool in zip(self.encs, self.pools):\n            out = enc(out)\n            skips.append(out)\n            out = pool(out)\n\n        # bottleneck\n        out = self.bottleneck(out)\n\n        # decoder (upsample + attention/concat not used here)\n        for upconv, dec, skip in zip(self.upconvs, self.decs, reversed(skips)):\n            out = upconv(out)\n            # ensure spatial alignment\n            if out.shape[2:] != skip.shape[2:]:\n                out = F.interpolate(out, size=skip.shape[2:], mode='trilinear', align_corners=False)\n            # concat\n            out = torch.cat([out, skip], dim=1)\n            out = dec(out)\n\n        out = self.final(out)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.558803Z","iopub.execute_input":"2025-11-16T15:28:55.558972Z","iopub.status.idle":"2025-11-16T15:28:55.579201Z","shell.execute_reply.started":"2025-11-16T15:28:55.558958Z","shell.execute_reply":"2025-11-16T15:28:55.578497Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss Function + Train Loop","metadata":{}},{"cell_type":"code","source":"scaler = GradScaler()\nfrom tqdm.auto import tqdm as tqdmlib\ndef criterion_from_logits(logits, labels):\n    # supports both binary single-channel logits or 2-channel logits\n    if logits.ndim==5 and logits.shape[1]==1:\n        loss = F.binary_cross_entropy_with_logits(logits[:,0], labels.float())\n    else:\n        loss = F.cross_entropy(logits, labels.squeeze(1).long(), ignore_index=2)\n    return loss\n\nfrom tqdm.auto import tqdm as tqdmlib\n\ndef train_one_epoch_verbose(model, loader, optimizer, steps_limit=None, log_every=50):\n    \"\"\"\n    Training loop with per-batch tqdm and mixed precision where available.\n    Uses torch.amp.autocast('cuda') when DEVICE is CUDA; no autocast on CPU.\n    \"\"\"\n    model.train()\n    running = 0.0\n    steps = 0\n    pbar = tqdmlib(total=(steps_limit if steps_limit else len(loader)), desc=\"Train\")\n    it = iter(loader)\n    use_cuda = (DEVICE.type == 'cuda')\n\n    while True:\n        try:\n            vols, lbls = next(it)\n        except StopIteration:\n            break\n\n        vols = vols.to(DEVICE, non_blocking=True)\n        lbls = lbls.to(DEVICE, non_blocking=True).squeeze(1)  # shape: B x D x H x W\n\n        optimizer.zero_grad()\n\n        if use_cuda:\n            # use AMP on CUDA\n            with torch.amp.autocast(\"cuda\"):\n                logits = model(vols)\n                loss = criterion_from_logits(logits, lbls)\n        else:\n            # CPU: run normal forward (autocast on CPU not necessary)\n            logits = model(vols)\n            loss = criterion_from_logits(logits, lbls)\n\n        # backward + step with GradScaler\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running += float(loss.item())\n        steps += 1\n\n        if steps % log_every == 0:\n            pbar.set_postfix({'loss': f\"{running/steps:.4f}\"})\n        pbar.update(1)\n\n        if steps_limit and steps >= steps_limit:\n            break\n\n    pbar.close()\n    return running / max(1, steps)\n\"\"\"\ndef validate_quick(model, csv_file=TRAIN_CSV, image_dir=TRAIN_IMG_DIR, n_samples=VAL_N_SAMPLES):\n    model.eval()\n    meta = pd.read_csv(csv_file)\n    ids = meta['id'].tolist()[:n_samples]\n    dices=[]\n    with torch.no_grad():\n        for vid in tqdm(ids, desc=\"Quick val\"):\n            imgp = os.path.join(image_dir, f\"{vid}.tif\")\n            lblp = os.path.join(TRAIN_LABEL_DIR, f\"{vid}.tif\")\n            if not os.path.exists(imgp) or not os.path.exists(lblp):\n                continue\n            vol = normalize_volume(load_3d_tiff(imgp)).astype(np.float32)\n            gt = load_3d_tiff(lblp).astype(np.uint8)\n            pred = sliding_window_inference(model, vol, patch_size=PATCH_SIZE, overlap=0.25)\n            pred = remove_small_objects(pred, min_voxels=250)\n            pfg = (pred==1); gfg = (gt==1)\n            inter = np.logical_and(pfg,gfg).sum()\n            denom = pfg.sum() + gfg.sum()\n            dice = (2.*inter)/denom if denom>0 else 1.0\n            dices.append(dice)\n    return float(np.mean(dices)) if dices else 0.0\n\n\"\"\"\ndef validate_quick(model,\n                   csv_file: str = TRAIN_CSV,\n                   image_dir: str = TRAIN_IMG_DIR,\n                   label_dir: str = TRAIN_LABEL_DIR,\n                   n_samples: int = 50,\n                   tol_vox: int = 2,\n                   patch_size = None,\n                   overlap: float = 0.25,\n                   min_voxels: int = 250,\n                   device = None):\n    \"\"\"\n    Validate on first n_samples volumes in csv_file, compute & print:\n      - Dice (volume)\n      - Surface Dice (tol_vox voxels)\n      - TopoScore (approximation)\n      - VOI (Variation of Information)\n    Returns average-metrics dict.\n    \"\"\"\n    model.eval()\n    if device is None:\n        device = DEVICE if 'DEVICE' in globals() else torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    if patch_size is None:\n        patch_size = PATCH_SIZE if 'PATCH_SIZE' in globals() else (32,128,128)\n\n    meta = pd.read_csv(csv_file)\n    ids = meta['id'].tolist()[:n_samples]\n\n    collected: List[Dict[str, float]] = []\n    print(f\"Quick val on {len(ids)} volumes (tol_vox={tol_vox}, min_vox={min_voxels})\")\n\n    with torch.no_grad():\n        for vid in tqdm(ids, desc=\"Validate volumes\"):\n            imgp = os.path.join(image_dir, f\"{vid}.tif\")\n            lblp = os.path.join(label_dir, f\"{vid}.tif\")\n            if not os.path.exists(imgp) or not os.path.exists(lblp):\n                print(f\"Skipping missing img/label for id {vid}\")\n                continue\n\n            vol = normalize_volume(load_3d_tiff(imgp)).astype(np.float32)\n            gt = load_3d_tiff(lblp).astype(np.uint8)\n\n            # inference (sliding window)\n            pred = sliding_window_inference(model, vol, patch_size=patch_size, overlap=overlap, device=device)\n\n            # postprocess small objects removal (already inside compute_metrics_for_volume too)\n            pred = remove_small_objects(pred, min_voxels=min_voxels)\n\n            metrics = compute_metrics_for_volume(pred, gt, tol_vox=tol_vox, min_voxels=min_voxels)\n            collected.append(metrics)\n\n            # print per-volume metrics\n            print(f\"id {vid} -> Dice: {metrics['dice']:.4f} | SurfaceDice: {metrics['surface_dice']:.4f} | Topo: {metrics['toposcore']:.4f} | VOI: {metrics['voi']:.4f}\")\n\n    if not collected:\n        avg = {'dice': 0.0, 'surface_dice': 0.0, 'toposcore': 0.0, 'voi': 0.0}\n    else:\n        avg = {k: float(np.mean([m[k] for m in collected])) for k in collected[0].keys()}\n    print(\"=== Quick-Validation Averages ===\")\n    print(f\"Dice: {avg['dice']:.4f} | SurfaceDice: {avg['surface_dice']:.4f} | Topo: {avg['toposcore']:.4f} | VOI: {avg['voi']:.4f}\")\n    return avg\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.580138Z","iopub.execute_input":"2025-11-16T15:28:55.580318Z","iopub.status.idle":"2025-11-16T15:28:55.599674Z","shell.execute_reply.started":"2025-11-16T15:28:55.580304Z","shell.execute_reply":"2025-11-16T15:28:55.598919Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Sliding Window Inference","metadata":{}},{"cell_type":"code","source":"def gaussian_weights_patch(patch_size):\n    grids = [np.linspace(-1, 1, n) for n in patch_size]\n    # build separable 1D Gaussian for each axis with sigma ~ 0.5\n    sig = 0.5\n    g = None\n    for axis in grids:\n        gw = np.exp(-(axis**2) / (2 * (sig**2)))\n        if g is None:\n            g = gw\n        else:\n            # outer product to expand\n            g = np.multiply.outer(g, gw)\n    # after loop, g has shape patch_size\n    g = np.array(g, dtype=np.float32)\n    # normalize to max 1.0\n    g = g / (g.max() + 1e-12)\n    return g\n\ndef sliding_window_inference(model, volume, patch_size=PATCH_SIZE, overlap=0.25, device=DEVICE):\n    if device is None:\n        device = DEVICE if 'DEVICE' in globals() else torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n    model.eval()\n\n    pd, ph, pw = patch_size\n    D, H, W = volume.shape\n    sd = max(1, int(pd * (1 - overlap)))\n    sh = max(1, int(ph * (1 - overlap)))\n    sw = max(1, int(pw * (1 - overlap)))\n\n    # build grid indices (include final patch ending at boundary)\n    grid_z = list(range(0, max(1, D - pd + 1), sd))\n    grid_y = list(range(0, max(1, H - ph + 1), sh))\n    grid_x = list(range(0, max(1, W - pw + 1), sw))\n    if grid_z[-1] != D - pd:\n        grid_z.append(D - pd)\n    if grid_y[-1] != H - ph:\n        grid_y.append(H - ph)\n    if grid_x[-1] != W - pw:\n        grid_x.append(W - pw)\n\n    # allocate score and weight maps\n    # score_map shape: (C, D, H, W) where C depends on model output (1 for binary logits or >1 for multi-class)\n    # to determine C we will check a single forward pass; but to avoid repeated forward, we will allocate with C=2 as default and adapt\n    # safer approach: accumulate per-patch and derive channel count from model output at runtime for first patch\n    weight_map = np.zeros((D, H, W), dtype=np.float32)\n    score_map = None\n    gw = gaussian_weights_patch((pd, ph, pw))  # shape (pd,ph,pw)\n\n    with torch.no_grad():\n        first = True\n        for z in grid_z:\n            for y in grid_y:\n                for x in grid_x:\n                    patch = volume[z:z+pd, y:y+ph, x:x+pw]\n                    # ensure correct shape (pd,ph,pw)\n                    if patch.shape != (pd,ph,pw):\n                        # pad if necessary (shouldn't happen with grid calc) — fallback\n                        padz = pd - patch.shape[0]\n                        pady = ph - patch.shape[1]\n                        padx = pw - patch.shape[2]\n                        patch = np.pad(patch,\n                                       ((0,padz),(0,pady),(0,padx)),\n                                       mode='constant', constant_values=0)\n                    inp = torch.from_numpy(patch[None, None]).float().to(device)  # 1 x 1 x pd x ph x pw\n                    logits = model(inp)  # expected shapes: (1,1,pd,ph,pw) or (1,C,pd,ph,pw)\n                    if first:\n                        # initialize score_map with appropriate num channels\n                        if logits.ndim == 5 and logits.shape[1] == 1:\n                            C = 2  # binary: we'll store probs for class 0 and 1\n                        else:\n                            C = int(logits.shape[1])\n                        score_map = np.zeros((C, D, H, W), dtype=np.float32)\n                        first = False\n\n                    # convert logits -> probs numpy\n                    if logits.ndim == 5 and logits.shape[1] == 1:\n                        # binary (single-channel logits)\n                        probs = torch.sigmoid(logits)[:, 0].cpu().numpy()[0]  # shape (pd,ph,pw)\n                        # add to score_map: channel 1 = probs, channel 0 = 1-probs\n                        score_map[1, z:z+pd, y:y+ph, x:x+pw] += probs * gw\n                        score_map[0, z:z+pd, y:y+ph, x:x+pw] += (1.0 - probs) * gw\n                    else:\n                        # multi-class logits\n                        probs = F.softmax(logits, dim=1).cpu().numpy()[0]  # shape (C, pd, ph, pw)\n                        # ensure ordering: probs[c] has shape (pd,ph,pw)\n                        for c in range(probs.shape[0]):\n                            score_map[c, z:z+pd, y:y+ph, x:x+pw] += probs[c] * gw\n\n                    # add window weights to weight_map (gw shape (pd,ph,pw) -> broadcast to slice (pd,ph,pw))\n                    weight_map[z:z+pd, y:y+ph, x:x+pw] += gw\n\n    # If for some reason score_map never created (empty volume?), fallback to zeros\n    if score_map is None:\n        C = 2\n        score_map = np.zeros((C, D, H, W), dtype=np.float32)\n        weight_map = np.ones((D,H,W), dtype=np.float32)\n\n    # normalize by weight map\n    weight_map[weight_map == 0] = 1.0\n    score_map = score_map / weight_map[np.newaxis, ...]  # broadcast weight_map across channel axis\n\n    # final argmax across channels -> predicted labels 0..C-1\n    pred = np.argmax(score_map, axis=0).astype(np.uint8)\n    return pred\n\ndef remove_small_objects(mask, min_voxels=300):\n    lab = measure.label(mask==1, connectivity=1)\n    out = np.zeros_like(mask, dtype=np.uint8)\n    for prop in measure.regionprops(lab):\n        if prop.area >= min_voxels:\n            out[lab==prop.label] = 1\n    return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.600501Z","iopub.execute_input":"2025-11-16T15:28:55.600726Z","iopub.status.idle":"2025-11-16T15:28:55.616617Z","shell.execute_reply.started":"2025-11-16T15:28:55.600700Z","shell.execute_reply":"2025-11-16T15:28:55.615901Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Quick EDA Utility Function","metadata":{}},{"cell_type":"code","source":"def quick_eda(sample_n=200):\n    meta = pd.read_csv(TRAIN_CSV)\n    print(\"Number of train volumes:\", len(meta))\n    sample_ids = meta['id'].tolist()[:min(sample_n, len(meta))]\n    shapes = {}\n    class_voxel_counts = []\n    for vid in tqdm(sample_ids, desc=\"Gather shapes\"):\n        imgp = os.path.join(TRAIN_IMG_DIR, f\"{vid}.tif\")\n        lblp = os.path.join(TRAIN_LABEL_DIR, f\"{vid}.tif\")\n        if not os.path.exists(imgp): continue\n        vol = load_3d_tiff(imgp)\n        shapes[vid] = vol.shape\n        if os.path.exists(lblp):\n            lbl = load_3d_tiff(lblp)\n            cls, cnt = np.unique(lbl, return_counts=True)\n            d = dict(zip(cls,cnt))\n            class_voxel_counts.append({\"id\":vid, **{f\"class_{c}\": d.get(c,0) for c in [0,1,2]}})\n    shapes_df = pd.DataFrame.from_dict(shapes, orient='index', columns=['D','H','W'])\n    shapes_df.index.name='id'\n    display(shapes_df.describe())\n    if len(class_voxel_counts)>0:\n        display(pd.DataFrame(class_voxel_counts).set_index('id').describe())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.617241Z","iopub.execute_input":"2025-11-16T15:28:55.617453Z","iopub.status.idle":"2025-11-16T15:28:55.630498Z","shell.execute_reply.started":"2025-11-16T15:28:55.617430Z","shell.execute_reply":"2025-11-16T15:28:55.629797Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train/Val Split + Build Dataloaders","metadata":{}},{"cell_type":"code","source":"meta = pd.read_csv(TRAIN_CSV)\nall_ids = meta['id'].tolist()\nif N_TRAIN_FILES is None or N_TRAIN_FILES >= len(all_ids):\n    sel_ids = all_ids\nelse:\n    if USE_RANDOM_SAMPLE:\n        random.seed(42)\n        sel_ids = random.sample(all_ids, N_TRAIN_FILES)\n    else:\n        sel_ids = all_ids[:N_TRAIN_FILES]\nprint(f\"Using {len(sel_ids)} volumes for training/validation (subset).\")\n\nimg_paths = [os.path.join(TRAIN_IMG_DIR, f\"{i}.tif\") for i in sel_ids]\nlbl_paths = [os.path.join(TRAIN_LABEL_DIR, f\"{i}.tif\") for i in sel_ids]\n\nsplit = int(0.9 * len(img_paths))\ntrain_imgs = img_paths[:split]\ntrain_lbls = lbl_paths[:split]\nval_imgs   = img_paths[split:]\nval_lbls   = lbl_paths[split:]\nprint(\"Train volumes:\", len(train_imgs), \"Val volumes:\", len(val_imgs))\n\n#Use dev patch & limited steps for quick iteration if desired\n_dev_mode = False\nif _dev_mode:\n    PATCH_SIZE = DEV_PATCH\n    STEPS_PER_EPOCH = DEV_STEPS\n\n#create datasets/dataloaders\ntrain_ds = VesuviusDatasetFast(train_imgs, train_lbls, patch_size=PATCH_SIZE, mode='train', cache_dir=CACHE_DIR)\nval_ds   = VesuviusDatasetFast(val_imgs,   val_lbls,   patch_size=PATCH_SIZE, mode='eval',  cache_dir=CACHE_DIR)\n\ntrain_loader = make_dataloader_safe(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=min(NUM_WORKERS, 4), persistent_workers=True)\nval_loader   = make_dataloader_safe(val_ds, batch_size=1, shuffle=False, num_workers=min(2, NUM_WORKERS), persistent_workers=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.631230Z","iopub.execute_input":"2025-11-16T15:28:55.631490Z","iopub.status.idle":"2025-11-16T15:28:55.651828Z","shell.execute_reply.started":"2025-11-16T15:28:55.631469Z","shell.execute_reply":"2025-11-16T15:28:55.651125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#build model, optimizer\nmodel = UNet3DLight(in_ch=1, out_ch=2).to(DEVICE)\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:55.652548Z","iopub.execute_input":"2025-11-16T15:28:55.653231Z","iopub.status.idle":"2025-11-16T15:28:58.585954Z","shell.execute_reply.started":"2025-11-16T15:28:55.653209Z","shell.execute_reply":"2025-11-16T15:28:58.585212Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"cell_type":"code","source":"best_val = -1.0   # global best combined score (update when validation improves)\nckpt_path = os.path.join(CKPT_DIR, \"best_model.pt\")\nprint(\"Checkpoint path:\", ckpt_path)\n\nfor epoch in range(1, NUM_EPOCHS+1):\n    t0 = time.time()\n    print(f\"\\n=== Epoch {epoch}/{NUM_EPOCHS} ===\")\n    train_loss = train_one_epoch_verbose(model, train_loader, optimizer, steps_limit=STEPS_PER_EPOCH)\n    \n    val_score = validate_quick(\n        model,\n        csv_file=TRAIN_CSV,\n        image_dir=TRAIN_IMG_DIR,\n        label_dir=TRAIN_LABEL_DIR,\n        n_samples=min(VAL_N_SAMPLES, len(val_imgs))   # or just 20–50\n    )\n# val_score might be: (a) a dict of averaged metrics (preferred), or (b) a single float (old behavior).\ncombined_score = None\n# if validate_quick returned a dict with a 'combined' entry (our newer function)\nif isinstance(val_score, dict):\n    # try keys in order of preference\n    for k in (\"combined\", \"combined_score\", \"surface_dice\", \"surface\", \"dice\"):\n        if k in val_score and val_score[k] is not None:\n            try:\n                combined_score = float(val_score[k])\n                break\n            except Exception:\n                continue\n# if validate_quick returned a single float (older code)\nelif isinstance(val_score, (int, float, np.floating)):\n    combined_score = float(val_score)\n\n# fallback: if we couldn't parse combined_score, compute a simple fallback using the printed avg dict if available\nif combined_score is None:\n    try:\n        # try to parse as float if val_score is string-like\n        combined_score = float(val_score)\n    except Exception:\n        combined_score = -1.0\n\n# now checkpointing\nif combined_score > best_val + 1e-8:\n    best_val = combined_score\n    try:\n        sd = model.state_dict()\n        # if trained with DataParallel, state_dict keys may be prefixed with 'module.'; saving as-is is fine.\n        torch.save(sd, ckpt_path)\n        print(f\"Saved new best checkpoint (score={combined_score:.6f}) -> {ckpt_path}\")\n    except Exception as e:\n        print(\"Failed to save checkpoint:\", e)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T15:28:58.586803Z","iopub.execute_input":"2025-11-16T15:28:58.587256Z","iopub.status.idle":"2025-11-16T21:39:05.474309Z","shell.execute_reply.started":"2025-11-16T15:28:58.587231Z","shell.execute_reply":"2025-11-16T21:39:05.473397Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Best Model & Generate Submission","metadata":{}},{"cell_type":"code","source":"#load best model and infer test set into zip\nckpt = os.path.join(CKPT_DIR, \"best_model.pt\")\nif os.path.exists(ckpt):\n    model = UNet3DLight(in_ch=1, out_ch=2).to(DEVICE)\n    model.load_state_dict(torch.load(ckpt, map_location=DEVICE))\n    print(\"Loaded checkpoint\", ckpt)\nelse:\n    print(\"No checkpoint found, using current model in memory\")\nzip_path2 = \"/kaggle/working/submission.zip\"\ntest_meta = pd.read_csv(TEST_CSV)\nzip_path2 = \"./submission.zip\"\nwith zipfile.ZipFile(zip_path2, 'w', compression=zipfile.ZIP_DEFLATED) as zf:\n    for vid in tqdm(test_meta['id'].tolist(), desc=\"Predict test\"):\n        imgp = os.path.join(TEST_IMG_DIR, f\"{vid}.tif\")\n        if not os.path.exists(imgp):\n            print(\"Missing:\", imgp); continue\n        vol = normalize_volume(load_3d_tiff(imgp)).astype(np.float32)\n        pred = sliding_window_inference(model, vol, patch_size=PATCH_SIZE, overlap=0.25)\n        pred = remove_small_objects(pred, min_voxels=250)\n        # write in-memory tif into zip\n        slices = [Image.fromarray(pred[z].astype(np.uint8)) for z in range(pred.shape[0])]\n        buffer = BytesIO()\n        slices[0].save(buffer, format='TIFF', save_all=True, append_images=slices[1:])\n        zf.writestr(f\"{vid}.tif\", buffer.getvalue())\nprint(\"Wrote submission:\", zip_path2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T21:40:09.664843Z","iopub.execute_input":"2025-11-16T21:40:09.665464Z","iopub.status.idle":"2025-11-16T21:40:22.026365Z","shell.execute_reply.started":"2025-11-16T21:40:09.665443Z","shell.execute_reply":"2025-11-16T21:40:22.025549Z"}},"outputs":[],"execution_count":null}]}