{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## 1) Setup: imports, device, reproducibility, and metadata\n\nThis notebook continues the Vesuvius **surface detection** pipeline. In this first step we:\n\n- Import core libraries for data handling (**NumPy/Pandas**), progress tracking (**tqdm**), and deep learning (**PyTorch**).\n- Detect whether a GPU is available and set the computation **DEVICE** accordingly.\n- Fix random seeds (Python / NumPy / PyTorch) to make results as reproducible as possible.\n- Load the official `train.csv` and `test.csv` metadata from the Kaggle dataset directory and print:\n  - dataset sizes\n  - class distribution proxy via `scroll_id` counts (how many samples per scroll)\n\n**Expected output**\n- A few lines confirming whether CUDA is available + GPU name (if any)\n- Shapes of `train_csv` and `test_csv`\n- A `scroll_id` frequency table (useful for understanding data balance across scrolls)\n","metadata":{}},{"cell_type":"code","source":"import os, random, zipfile\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom PIL import Image\nfrom scipy import ndimage as ndi\n\nUSE_CUDA = torch.cuda.is_available()\nDEVICE = \"cuda\" if USE_CUDA else \"cpu\"\n\nprint(\"cuda available:\", USE_CUDA)\nif USE_CUDA:\n    print(\"GPU count:\", torch.cuda.device_count())\n    print(\"GPU 0:\", torch.cuda.get_device_name(0))\nprint(\"DEVICE:\", DEVICE)\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything(42)\n\nDATA_ROOT = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\ntrain_csv = pd.read_csv(DATA_ROOT / \"train.csv\")\ntest_csv  = pd.read_csv(DATA_ROOT / \"test.csv\")\n\nprint(\"train:\", train_csv.shape, \"test:\", test_csv.shape)\nprint(\"scroll_id counts:\\n\", train_csv.scroll_id.value_counts())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:18.288156Z","iopub.execute_input":"2025-12-20T08:45:18.288771Z","iopub.status.idle":"2025-12-20T08:45:18.606464Z","shell.execute_reply.started":"2025-12-20T08:45:18.288734Z","shell.execute_reply":"2025-12-20T08:45:18.605642Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2) Reading 3D volumes (multi-page TIFF) and label sanity checks\n\nThe raw data is stored as **multi-page TIFF files** where each page corresponds to a depth slice. We load each example into a 3D NumPy array with shape:\n\n- **(Z, Y, X)** = (number of slices, height, width)\n\n### Why PIL?\nSome TIFFs are LZW-compressed and may fail with certain readers. \n\n### What we check\n1. **Sanity-load one training example** (`train_images/{id}.tif` and `train_labels/{id}.tif`)\n   - Confirm shapes and dtypes\n   - Inspect intensity range of the image volume\n   - Confirm label values present in the label volume\n\n2. **Verify whether label value `2` exists**\n   - We sample multiple training IDs and look for any label voxels equal to `2`.\n   - This is important because `2` typically represents **unlabeled/ignore regions** and should be excluded from training and evaluation.\n\n**Expected output**\n- Example image volume: `img.shape`, `dtype`, `min/max`\n- Example label volume: `lbl.shape`, `dtype`, `unique values`\n- Either a message that an example with `label=2` was found, or that it wasn’t found in the sampled subset.\n","metadata":{}},{"cell_type":"code","source":"def read_tiff_volume_pil(path: Path) -> np.ndarray:\n    \"\"\"Read multi-page TIFF -> (Z, Y, X) using PIL (works with LZW).\"\"\"\n    im = Image.open(str(path))\n    frames = []\n    i = 0\n    while True:\n        frames.append(np.array(im))\n        i += 1\n        try:\n            im.seek(i)\n        except EOFError:\n            break\n    return np.stack(frames, axis=0)\n\ndef read_volume(path: Path) -> np.ndarray:\n    return read_tiff_volume_pil(path)\n\n# sanity load\nex_id = int(train_csv[\"id\"].iloc[0])\nimg = read_volume(DATA_ROOT / \"train_images\" / f\"{ex_id}.tif\")\nlbl = read_volume(DATA_ROOT / \"train_labels\" / f\"{ex_id}.tif\")\n\nprint(\"Example:\", ex_id)\nprint(\" img:\", img.shape, img.dtype, \"min/max:\", img.min(), img.max())\nprint(\" lbl:\", lbl.shape, lbl.dtype, \"unique:\", np.unique(lbl))\n\n# prove label=2 exists in the dataset (it may not exist in this particular example)\nhas2 = None\nfor _id in train_csv[\"id\"].sample(80, random_state=0).astype(int).tolist():\n    _lbl = read_volume(DATA_ROOT / \"train_labels\" / f\"{_id}.tif\")\n    if np.any(_lbl == 2):\n        has2 = _id\n        print(\"Found a label with value 2:\", has2, \"| unique:\", np.unique(_lbl))\n        break\n\nif has2 is None:\n    print(\"Did not find label=2 in sample(80). It may still exist—sample more if needed.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:18.607606Z","iopub.execute_input":"2025-12-20T08:45:18.607951Z","iopub.status.idle":"2025-12-20T08:45:19.672206Z","shell.execute_reply.started":"2025-12-20T08:45:18.607917Z","shell.execute_reply":"2025-12-20T08:45:19.670586Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3) Experiment configuration (CFG) + train/validation split by scroll\n\nWe define a single configuration dictionary (`CFG`) that controls the full pipeline:\n\n- **Reproducibility:** `seed`\n- **2.5D context:** `in_slices=9` (each sample uses 9 neighboring slices as channels)\n- **Performance settings:** mixed precision (`amp`), `num_workers`, and a small `cache_size` for faster repeated access\n\n### GPU vs CPU defaults\nTo keep the notebook runnable everywhere, we choose different defaults depending on whether CUDA is available:\n\n**On GPU (Kaggle T4/P100-style):**\n- moderate patch size (`patch=160`)\n- larger batch (`batch=8`)\n- multi-epoch training (`epochs=6`)\n- sampling budget per epoch (`steps_per_epoch=600`)\n- positive sampling fraction (`pos_frac=0.65`) to handle class imbalance\n\n**On CPU (pipeline sanity-check mode):**\n- smaller patch size / batch\n- very few steps/epochs to verify everything runs end-to-end\n\n### Validation strategy: hold-out scrolls (anti-leakage split)\nWe split **by `scroll_id`**, not randomly by patches:\n- `val_scrolls = {26002, 26010}`\n- train = all other scrolls\n\nThis prevents leakage from highly correlated spatial regions and tests generalisation to unseen scrolls (same idea as in Assignment 2).\n","metadata":{}},{"cell_type":"code","source":"CFG = dict(\n    seed=42,\n    in_slices=9,\n    amp=True,          # will be used ONLY if GPU is available\n    num_workers=2,\n    cache_size=2,\n)\n\n# GPU vs CPU safe defaults\nif USE_CUDA:\n    # Good starting point for Kaggle T4/P100\n    CFG.update(dict(\n        patch=160,\n        batch=8,\n        lr=3e-4,\n        epochs=6,\n        steps_per_epoch=600,\n        pos_frac=0.65,\n    ))\nelse:\n    # CPU fallback (just to test the pipeline)\n    CFG.update(dict(\n        patch=96,\n        batch=2,\n        lr=3e-4,\n        epochs=1,\n        steps_per_epoch=50,\n        pos_frac=0\n        \n    CFG = dict(\n    seed=42,\n    in_slices=9,\n    amp=True,          # will be used ONLY if GPU is available\n    num_workers=2,\n    cache_size=2,\n)\n\n# GPU vs CPU safe defaults\nif USE_CUDA:\n    # Good starting point for Kaggle T4/P100\n    CFG.update(dict(\n        patch=160,\n        batch=8,\n        lr=3e-4,\n        epochs=6,\n        steps_per_epoch=600,\n        pos_frac=0.65,\n    ))\nelse:\n    # CPU fallback (just to test the pipeline)\n    CFG.update(dict(\n        patch=96,\n        batch=2,\n        lr=3e-4,\n        epochs=1,\n        steps_per_epoch=50,\n        pos_frac=0.65,\n        amp=False,\n        num_workers=0,\n        cache_size=1,\n    ))\n\nprint(\"CFG:\", CFG)\n\n\n# Hold out 2 scrolls for validation (same idea as A2; avoids leakage)\nval_scrolls = {26002, 26010}\ntrain_df = train_csv[~train_csv.scroll_id.isin(val_scrolls)].reset_index(drop=True)\nval_df   = train_csv[ train_csv.scroll_id.isin(val_scrolls)].reset_index(drop=True)\n\nprint(\"Train rows:\", len(train_df), \"Val rows:\", len(val_df))\nprint(\"Val scrolls:\", sorted(val_scrolls))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.672803Z","iopub.status.idle":"2025-12-20T08:45:19.673107Z","shell.execute_reply.started":"2025-12-20T08:45:19.672971Z","shell.execute_reply":"2025-12-20T08:45:19.672989Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4) Throughput tuning (Kaggle T4)\n\nBefore training/inference, we apply a few practical tweaks to improve data throughput on a Kaggle T4 GPU:\n\n- **`num_workers`**: parallelises data loading.\n- **`cache_size`**: keeps a small number of recently used volumes/tiles in memory to reduce repeated TIFF reads.\n- **`patch`**: keep at `160` as a good speed/quality trade-off on T4.\n- **`batch`**: `8` \n\nThe final print statement confirms the key runtime parameters that most strongly affect speed and GPU memory usage.\n","metadata":{}},{"cell_type":"code","source":"# Speed/throughput tweaks for T4\nCFG[\"num_workers\"] = 2      # try 4 if stable\nCFG[\"cache_size\"]  = 3\nCFG[\"patch\"]       = 160    # keep 160 on T4\nCFG[\"batch\"]       = 8      # if OOM, drop to 6 or 4\n\nprint(\"Updated CFG:\", {k: CFG[k] for k in [\"num_workers\",\"cache_size\",\"patch\",\"batch\"]})\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.674344Z","iopub.status.idle":"2025-12-20T08:45:19.674646Z","shell.execute_reply.started":"2025-12-20T08:45:19.674504Z","shell.execute_reply":"2025-12-20T08:45:19.674532Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5) Preprocessing + 2.5D patch datasets (train & validation)\n\nThis section defines the **core data pipeline** used for training and evaluation.\n\n### Normalisation\nWe standardise each 3D volume (per-example) to stabilise optimisation:\n- `normalize_volume(vol) = (vol - mean) / (std + 1e-6)`\n\n### Safe patch extraction (with padding)\n`crop_pad_2d(...)` extracts a `(patch, patch)` crop and **pads** when the crop hits image borders, so all samples have consistent shape.\n\n---\n\n## 5.1 Training dataset: `RandomPatch2p5D`\n\nWe train on randomly sampled patches from random volumes:\n\n- Input **x**: a **2.5D stack** of `in_slices` neighboring slices  \n  Shape: **(C, H, W)** where `C=in_slices`, `H=W=patch`\n- Target **y**: a 2D label patch from the center slice  \n  Shape: **(H, W)**\n\n### Handling unlabeled regions (label = 2)\n- Pixels with value `2` are treated as **ignore / unlabeled**.\n- Sampling avoids them via `valid = (sl != 2)`\n- When padding label patches, we pad with `2` so the loss can ignore padded regions.\n\n### Class imbalance control (`pos_frac`)\nBecause the surface class is rare, we oversample positive patches:\n- With probability `pos_frac`, we choose a center point from `label==1`\n- Otherwise we sample from `label==0` (background), always excluding `label==2`\n\n### Speed: small LRU cache (`cache_size`)\nTo avoid repeated TIFF I/O, we keep a small **LRU cache** of `(img,lbl)` volumes.\n\n---\n\n## 5.2 Validation dataset: `FixedValPatches2p5D`\n\nValidation should be **repeatable**, so we pre-sample a fixed set of patch locations:\n\n- We build `n_patches` fixed items `(vid, z, y0, x0)` once using a fixed RNG seed.\n- Each epoch evaluates on the **exact same patches**, making metric comparisons stable.\n\n---\n\n## 5.3 DataLoaders (performance settings)\n\nWe create `train_dl` and `val_dl` with:\n- `pin_memory=True` on GPU for faster host→device transfer\n- `persistent_workers` + `prefetch_factor` to improve throughput when `num_workers>0`\n\n**Expected output**\n- Number of batches per epoch for training and validation.\n","metadata":{}},{"cell_type":"code","source":"def normalize_volume(vol: np.ndarray) -> np.ndarray:\n    v = vol.astype(np.float32)\n    return (v - v.mean()) / (v.std() + 1e-6)\n\ndef crop_pad_2d(arr2d, y0, x0, ps, pad_value=0):\n    \"\"\"Crop arr2d[y0:y0+ps, x0:x0+ps] and pad if needed to (ps,ps).\"\"\"\n    H, W = arr2d.shape\n    y1 = min(H, y0 + ps)\n    x1 = min(W, x0 + ps)\n    crop = arr2d[y0:y1, x0:x1]\n\n    pad_h = ps - crop.shape[0]\n    pad_w = ps - crop.shape[1]\n    if pad_h > 0 or pad_w > 0:\n        crop = np.pad(\n            crop,\n            ((0, max(0, pad_h)), (0, max(0, pad_w))),\n            mode=\"constant\",\n            constant_values=pad_value\n        )\n    return crop\n\nclass RandomPatch2p5D(Dataset):\n    def __init__(self, df, steps, in_slices=9, patch=192, pos_frac=0.65, cache_size=2, seed=42):\n        self.ids = df[\"id\"].astype(int).values\n        self.steps = int(steps)\n        self.in_slices = int(in_slices)\n        self.half = self.in_slices // 2\n        self.patch = int(patch)\n        self.pos_frac = float(pos_frac)\n        self.cache_size = int(cache_size)\n        self.seed = int(seed)\n\n        self.cache = {}\n        self.cache_order = []\n\n    def _load(self, vid: int):\n        vid = int(vid)\n        if vid in self.cache:\n            # update LRU order\n            if vid in self.cache_order:\n                self.cache_order.remove(vid)\n            self.cache_order.append(vid)\n            return self.cache[vid]\n\n        img = read_volume(DATA_ROOT / \"train_images\" / f\"{vid}.tif\")\n        lbl = read_volume(DATA_ROOT / \"train_labels\" / f\"{vid}.tif\")\n        img = normalize_volume(img)\n\n        if len(self.cache_order) >= self.cache_size:\n            old = self.cache_order.pop(0)\n            self.cache.pop(old, None)\n\n        self.cache[vid] = (img, lbl)\n        self.cache_order.append(vid)\n        return img, lbl\n\n    def __len__(self):\n        return self.steps\n\n    def __getitem__(self, idx):\n        # deterministic-ish randomness per index\n        rng = np.random.RandomState(self.seed + idx)\n\n        vid = int(self.ids[rng.randint(0, len(self.ids))])\n        img, lbl = self._load(vid)  # (Z,Y,X)\n\n        Z, Y, X = img.shape\n        ps = self.patch\n\n        want_pos = (rng.rand() < self.pos_frac)\n        z = int(rng.randint(self.half, Z - self.half))\n\n        sl = lbl[z]\n        valid = (sl != 2)  # IGNORE unlabeled\n\n        if want_pos:\n            cand = np.argwhere(sl == 1)\n            if len(cand) == 0:\n                cand = np.argwhere(valid)\n        else:\n            cand = np.argwhere((sl == 0) & valid)\n            if len(cand) == 0:\n                y, x = Y // 2, X // 2\n            else:\n                y, x = cand[rng.randint(0, len(cand))]\n\n        y, x = cand[rng.randint(0, len(cand))]\n\n        y0 = int(np.clip(y - ps//2, 0, Y - ps))\n        x0 = int(np.clip(x - ps//2, 0, X - ps))\n        \n        # make z-range safe even if Z is small\n        z0 = max(0, z - self.half)\n        z1 = min(Z, z + self.half + 1)\n        stack = img[z0:z1]  # (<=C, Y, X)\n\n        # pad in Z if needed to exactly in_slices\n        need = self.in_slices - stack.shape[0]\n        if need > 0:\n            # pad by repeating edge slices\n            top = need // 2\n            bot = need - top\n            stack = np.pad(stack, ((top, bot), (0,0), (0,0)), mode=\"edge\")\n\n        # crop+pad each slice to (ps,ps)\n        stack = np.stack([crop_pad_2d(stack[i], y0, x0, ps, pad_value=0) for i in range(self.in_slices)], axis=0)\n\n        # label crop+pad: pad with 2 (unlabeled) so loss ignores it\n        target2d = crop_pad_2d(lbl[z], y0, x0, ps, pad_value=2)\n\n        x_t = torch.from_numpy(stack).float()\n        y_t = torch.from_numpy(target2d).long()\n        return x_t, y_t\n\nclass FixedValPatches2p5D(Dataset):\n    def __init__(self, df, n_patches=400, in_slices=9, patch=160, seed=123, cache_size=2):\n        self.ids = df[\"id\"].astype(int).values\n        self.in_slices = int(in_slices)\n        self.half = self.in_slices // 2\n        self.patch = int(patch)\n        self.seed = int(seed)\n\n        self.cache_size = int(cache_size)\n        self.cache = {}\n        self.cache_order = []\n\n        rng = np.random.RandomState(seed)\n        self.items = []  # list of (vid,z,y0,x0)\n\n        # pre-sample fixed positions from labeled regions\n        for _ in tqdm(range(n_patches), desc=\"build fixed val\", leave=False):\n            vid = int(self.ids[rng.randint(0, len(self.ids))])\n            img, lbl = self._load(vid)\n\n            Z, Y, X = img.shape\n            z = int(rng.randint(self.half, Z - self.half))\n            sl = lbl[z]\n            valid = np.argwhere(sl != 2)\n            if len(valid) == 0:\n                y, x = Y//2, X//2\n            else:\n                y, x = valid[rng.randint(0, len(valid))]\n\n            y0 = int(np.clip(y - self.patch//2, 0, Y - self.patch))\n            x0 = int(np.clip(x - self.patch//2, 0, X - self.patch))\n\n            self.items.append((vid, z, y0, x0))\n\n    def _load(self, vid):\n        vid = int(vid)\n        if vid in self.cache:\n            if vid in self.cache_order:\n                self.cache_order.remove(vid)\n            self.cache_order.append(vid)\n            return self.cache[vid]\n\n        img = read_volume(DATA_ROOT / \"train_images\" / f\"{vid}.tif\")\n        lbl = read_volume(DATA_ROOT / \"train_labels\" / f\"{vid}.tif\")\n        img = normalize_volume(img)\n\n        if len(self.cache_order) >= self.cache_size:\n            old = self.cache_order.pop(0)\n            self.cache.pop(old, None)\n\n        self.cache[vid] = (img, lbl)\n        self.cache_order.append(vid)\n        return img, lbl\n\n    def __len__(self):\n        return len(self.items)\n\n    def __getitem__(self, idx):\n        vid, z, y0, x0 = self.items[idx]\n        img, lbl = self._load(int(vid))\n\n        ps = self.patch\n        stack  = img[z-self.half:z+self.half+1, y0:y0+ps, x0:x0+ps]\n        target = lbl[z, y0:y0+ps, x0:x0+ps]\n\n        return torch.from_numpy(stack).float(), torch.from_numpy(target).long()\n\nPIN = True if USE_CUDA else False\n\ntrain_ds = RandomPatch2p5D(\n    train_df,\n    steps=CFG[\"steps_per_epoch\"],\n    in_slices=CFG[\"in_slices\"],\n    patch=CFG[\"patch\"],\n    pos_frac=CFG[\"pos_frac\"],\n    cache_size=CFG[\"cache_size\"],\n    seed=CFG[\"seed\"]\n)\n\nval_ds = FixedValPatches2p5D(\n    val_df,\n    n_patches=400,\n    in_slices=CFG[\"in_slices\"],\n    patch=CFG[\"patch\"],\n    seed=123,\n    cache_size=CFG[\"cache_size\"]\n)\n\ndl_kwargs = dict(\n    num_workers=CFG[\"num_workers\"],\n    pin_memory=PIN,\n    persistent_workers=(CFG[\"num_workers\"] > 0),\n)\nif CFG[\"num_workers\"] > 0:\n    dl_kwargs[\"prefetch_factor\"] = 2\n\ntrain_dl = DataLoader(train_ds, batch_size=CFG[\"batch\"], shuffle=True,  **dl_kwargs)\nval_dl   = DataLoader(val_ds,   batch_size=CFG[\"batch\"], shuffle=False, **dl_kwargs)\n\nprint(\"Train batches:\", len(train_dl), \"| Val batches:\", len(val_dl))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.675839Z","iopub.status.idle":"2025-12-20T08:45:19.676146Z","shell.execute_reply.started":"2025-12-20T08:45:19.67601Z","shell.execute_reply":"2025-12-20T08:45:19.676034Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6) Model: 2.5D U-Net for surface segmentation\n\nWe use a **U-Net-style encoder–decoder** architecture to predict a binary surface mask for the center slice of each sampled patch.\n\n### Input / output\n- **Input (`x`)**: a 2.5D stack treated as channels  \n  Shape: **(B, C, H, W)** where `C = in_slices`\n- **Output (`logits`)**: a single-channel segmentation logit map  \n  Shape: **(B, 1, H, W)**  \n  (Sigmoid will be applied later to convert logits → probabilities.)\n\n### Building blocks\n- **`ConvBlock`**: two `3×3` convolutions with BatchNorm + ReLU  \n  This is the basic feature extractor used at every level.\n\n### Encoder (downsampling path)\n- Four stages (`enc1` … `enc4`) with `MaxPool2d(2)` between them.\n- Channel width doubles each level: `base → 2×base → 4×base → 8×base`.\n\n### Bottleneck\n- A deeper block (`mid`) operating at the lowest resolution with `16×base` channels.\n\n### Decoder (upsampling path)\n- Upsampling is done with **transpose convolutions** (`ConvTranspose2d`).\n- Skip connections concatenate encoder features with decoder features at the same resolution:\n  - `cat([d4, e4])`, `cat([d3, e3])`, etc.\n- This preserves fine spatial detail while still using deep context.\n\n**Why this makes sense here**\nEven though the task uses 3D CT information, we keep the model simple by using **2.5D context** (`in_slices` slices) while still predicting a **2D mask**, which is a good trade-off between performance and compute.\n","metadata":{}},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.net(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_ch, base=32):\n        super().__init__()\n        self.enc1 = ConvBlock(in_ch, base)\n        self.enc2 = ConvBlock(base, base*2)\n        self.enc3 = ConvBlock(base*2, base*4)\n        self.enc4 = ConvBlock(base*4, base*8)\n        self.pool = nn.MaxPool2d(2)\n\n        self.mid  = ConvBlock(base*8, base*16)\n\n        self.up4  = nn.ConvTranspose2d(base*16, base*8, 2, stride=2)\n        self.dec4 = ConvBlock(base*16, base*8)\n\n        self.up3  = nn.ConvTranspose2d(base*8, base*4, 2, stride=2)\n        self.dec3 = ConvBlock(base*8, base*4)\n\n        self.up2  = nn.ConvTranspose2d(base*4, base*2, 2, stride=2)\n        self.dec2 = ConvBlock(base*4, base*2)\n\n        self.up1  = nn.ConvTranspose2d(base*2, base, 2, stride=2)\n        self.dec1 = ConvBlock(base*2, base)\n\n        self.out = nn.Conv2d(base, 1, 1)\n\n    def forward(self, x):\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        m  = self.mid(self.pool(e4))\n\n        d4 = self.up4(m); d4 = self.dec4(torch.cat([d4, e4], 1))\n        d3 = self.up3(d4); d3 = self.dec3(torch.cat([d3, e3], 1))\n        d2 = self.up2(d3); d2 = self.dec2(torch.cat([d2, e2], 1))\n        d1 = self.up1(d2); d1 = self.dec1(torch.cat([d1, e1], 1))\n        return self.out(d1)  # logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.677171Z","iopub.status.idle":"2025-12-20T08:45:19.677459Z","shell.execute_reply.started":"2025-12-20T08:45:19.677302Z","shell.execute_reply":"2025-12-20T08:45:19.677317Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7) Training objective and metric (ignore label = 2)\n\nThe label volumes may contain value **`2`**, which represents **unlabeled / ignore** regions.\nWe ensure both the loss and the evaluation metric exclude these pixels using a boolean mask:\n\n- `mask = (y != 2)`\n\n### Loss: BCE + Dice (from logits)\nWe optimise a **balanced combination** of:\n1. **Binary Cross-Entropy with logits (BCE)**  \n   - Stable for pixel-wise classification\n   - Uses the raw logits directly (no manual sigmoid needed)\n\n2. **Soft Dice loss**\n   - Encourages overlap between predicted probabilities and the target mask\n   - Helps with class imbalance and focuses on segmentation quality\n\nThe final loss is:\n- `loss = 0.5 * BCE + 0.5 * DiceLoss`\n\n### Metric: Dice score at a fixed threshold\nFor reporting, we compute **Dice/F1** after thresholding probabilities:\n\n- `pred = sigmoid(logits) > thr`\n- default `thr = 0.40` (chosen to trade off false positives and false negatives)\n\nDice score is computed as:\n- `Dice = 2TP / (2TP + FP + FN)`\n\n**Key point**\nBoth loss and metric operate only on `mask` pixels, so ignore regions (`y==2`) do not affect training or evaluation.\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\n\ndef dice_loss_from_logits(logits, y01, mask):\n    p = torch.sigmoid(logits).squeeze(1)\n    p = p[mask]\n    t = y01[mask].float()\n    num = 2*(p*t).sum() + 1e-6\n    den = (p + t).sum() + 1e-6\n    return 1 - num/den\n\ndef loss_fn(logits, y):\n    mask = (y != 2)\n    y01 = (y == 1).long()\n    bce = F.binary_cross_entropy_with_logits(logits.squeeze(1)[mask], y01[mask].float())\n    dsc = dice_loss_from_logits(logits, y01, mask)\n    return 0.5*bce + 0.5*dsc\n\n@torch.no_grad()\ndef dice_score(logits, y, thr=0.40):\n    mask = (y != 2)\n    y01 = (y == 1).long()\n    pred = (torch.sigmoid(logits) > thr).long().squeeze(1)\n    p = pred[mask]\n    t = y01[mask]\n    tp = (p*t).sum().item()\n    fp = (p*(1-t)).sum().item()\n    fn = ((1-p)*t).sum().item()\n    return (2*tp) / (2*tp + fp + fn + 1e-6)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.679001Z","iopub.status.idle":"2025-12-20T08:45:19.679253Z","shell.execute_reply.started":"2025-12-20T08:45:19.679138Z","shell.execute_reply":"2025-12-20T08:45:19.679153Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8) Training loop (single run) with checkpointing + AMP\n\nThis cell defines `train_one(seed, save_path)`, which trains **one U-Net model** end-to-end and saves the best checkpoint based on validation Dice.\n\n### What happens inside `train_one`\n1. **Reproducibility**\n   - We reseed Python/NumPy/PyTorch for each run so different seeds produce controlled diversity.\n\n2. **Model + optimiser**\n   - Model: `UNet(in_ch=in_slices, base=32)`\n   - Optimiser: `AdamW` with learning rate `CFG[\"lr\"]`\n\n3. **Mixed precision (AMP)**\n   - Uses `torch.amp.autocast(...)` + `GradScaler` when running on GPU and `CFG[\"amp\"]=True`\n   - On CPU, AMP is automatically disabled and standard FP32 training is used.\n\n4. **Training phase**\n   - Iterate over `train_dl`\n   - Forward → compute `loss_fn` (BCE + Dice, ignoring label=2) → backward → optimiser step\n   - Track average training loss per epoch\n\n5. **Validation phase**\n   - Evaluate on `val_dl` with `model.eval()` and `torch.no_grad()`\n   - Report average **validation Dice** using the fixed threshold `thr=0.4`\n\n6. **Best-checkpoint saving**\n   - If the current `val_dice` improves, we save `model.state_dict()` to `save_path`\n\n**Expected output**\nPer epoch:\n- `train_loss` (average over train batches)\n- `val_dice` (average over val batches)\n\nAt the end:\n- the best validation Dice and confirmation of the saved checkpoint path.\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom tqdm import tqdm\n\ndef seed_everything(seed=42):\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\ndef train_one(seed, save_path):\n    seed_everything(seed)\n\n    model = UNet(in_ch=CFG[\"in_slices\"], base=32).to(DEVICE)\n    opt = torch.optim.AdamW(model.parameters(), lr=CFG[\"lr\"])\n\n    # New AMP API (only active on GPU)\n    scaler = torch.amp.GradScaler(\"cuda\", enabled=(CFG[\"amp\"] and USE_CUDA))\n\n    best = 0.0\n    for ep in range(CFG[\"epochs\"]):\n        model.train()\n        tr_loss = 0.0\n\n        for x, y in tqdm(train_dl, desc=f\"train ep{ep+1}\", leave=False):\n            x, y = x.to(DEVICE, non_blocking=True), y.to(DEVICE, non_blocking=True)\n\n            opt.zero_grad(set_to_none=True)\n            with torch.amp.autocast(\"cuda\", enabled=(CFG[\"amp\"] and USE_CUDA)):\n                logits = model(x)\n                loss = loss_fn(logits, y)\n\n            if CFG[\"amp\"] and USE_CUDA:\n                scaler.scale(loss).backward()\n                scaler.step(opt)\n                scaler.update()\n            else:\n                loss.backward()\n                opt.step()\n\n            tr_loss += loss.item()\n\n        model.eval()\n        va = 0.0\n        with torch.no_grad():\n            for x, y in tqdm(val_dl, desc=f\"val ep{ep+1}\", leave=False):\n                x, y = x.to(DEVICE, non_blocking=True), y.to(DEVICE, non_blocking=True)\n                with torch.amp.autocast(\"cuda\", enabled=(CFG[\"amp\"] and USE_CUDA)):\n                    logits = model(x)\n                va += dice_score(logits, y, thr=0.4)\n\n        va /= len(val_dl)\n        print(f\"seed={seed} ep={ep+1}/{CFG['epochs']} | train_loss={tr_loss/len(train_dl):.4f} | val_dice={va:.4f}\")\n\n        if va > best:\n            best = va\n            torch.save(model.state_dict(), save_path)\n\n    print(\"BEST val_dice:\", best, \"saved:\", save_path)\n    return best\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.68031Z","iopub.status.idle":"2025-12-20T08:45:19.680644Z","shell.execute_reply.started":"2025-12-20T08:45:19.680468Z","shell.execute_reply":"2025-12-20T08:45:19.680487Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9) Train Model A (seed = 42)\n\nWe train the first model (**Model A**) using a fixed seed (`42`) and save the **best-performing checkpoint** (by validation Dice) to:\n\n- `modelA_seed42.pth`\n\nThis model will later be used:\n- for standalone evaluation, and\n- as a component in the final **ensemble** (Model A + Model B).\n","metadata":{}},{"cell_type":"code","source":"bestA = train_one(seed=42, save_path=\"modelA_seed42.pth\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.681913Z","iopub.status.idle":"2025-12-20T08:45:19.682189Z","shell.execute_reply.started":"2025-12-20T08:45:19.682062Z","shell.execute_reply":"2025-12-20T08:45:19.682079Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10) Train Model B (seed = 777)\n\nWe train a second model (**Model B**) with a different random seed (`777`) and save its best checkpoint to:\n\n- `modelB_seed777.pth`\n\nEven with the same architecture and hyperparameters, changing the seed typically leads to **slightly different decision boundaries**. This diversity is useful because averaging predictions from Model A and Model B often improves stability and overall Dice compared to either model alone.\n","metadata":{}},{"cell_type":"code","source":"bestB = train_one(seed=777, save_path=\"modelB_seed777.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T08:45:19.686792Z","iopub.status.idle":"2025-12-20T08:45:19.687154Z","shell.execute_reply.started":"2025-12-20T08:45:19.686963Z","shell.execute_reply":"2025-12-20T08:45:19.686985Z"}},"outputs":[],"execution_count":null}]}