{"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"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom pathlib import Path\nimport random\n\nimport numpy as np\nimport pandas as pd\nimport tifffile as tiff\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:06.842255Z","iopub.execute_input":"2025-12-19T14:40:06.842548Z","iopub.status.idle":"2025-12-19T14:40:08.720576Z","shell.execute_reply.started":"2025-12-19T14:40:06.84252Z","shell.execute_reply":"2025-12-19T14:40:08.719708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_ROOT = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\n\nTRAIN_IMG_DIR = DATA_ROOT / \"train_images\"\nTRAIN_LBL_DIR = DATA_ROOT / \"train_labels\"\nTRAIN_CSV = DATA_ROOT / \"train.csv\"\n\n# Patch configuration (we'll tune later)\nPATCH_SIZE = 64          # 64³ patches (safe start)\nMIN_FG_RATIO = 0.01      # at least 1% papyrus in patch\nMAX_TRIES = 20           # avoid infinite loops\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:08.721796Z","iopub.execute_input":"2025-12-19T14:40:08.722169Z","iopub.status.idle":"2025-12-19T14:40:08.726413Z","shell.execute_reply.started":"2025-12-19T14:40:08.722145Z","shell.execute_reply":"2025-12-19T14:40:08.725601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV)\n\n# Choose validation scroll (smallest one for speed)\nval_scroll_id = train_df[\"scroll_id\"].value_counts().idxmin()\n\ntrain_ids = train_df[train_df[\"scroll_id\"] != val_scroll_id][\"id\"].astype(str).tolist()\nval_ids   = train_df[train_df[\"scroll_id\"] == val_scroll_id][\"id\"].astype(str).tolist()\n\nprint(f\"Validation scroll_id: {val_scroll_id}\")\nprint(f\"Train volumes: {len(train_ids)}\")\nprint(f\"Val volumes: {len(val_ids)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:08.727289Z","iopub.execute_input":"2025-12-19T14:40:08.727586Z","iopub.status.idle":"2025-12-19T14:40:08.748653Z","shell.execute_reply.started":"2025-12-19T14:40:08.727566Z","shell.execute_reply":"2025-12-19T14:40:08.747939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sample_patch(volume, label, patch_size):\n    \"\"\"\n    Randomly sample a 3D patch from volume + label.\n    Returns (img_patch, lbl_patch)\n    \"\"\"\n    D, H, W = volume.shape\n    ps = patch_size\n\n    z = random.randint(0, D - ps)\n    y = random.randint(0, H - ps)\n    x = random.randint(0, W - ps)\n\n    img_patch = volume[z:z+ps, y:y+ps, x:x+ps]\n    lbl_patch = label[z:z+ps, y:y+ps, x:x+ps]\n\n    return img_patch, lbl_patch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:08.75029Z","iopub.execute_input":"2025-12-19T14:40:08.750861Z","iopub.status.idle":"2025-12-19T14:40:08.755662Z","shell.execute_reply.started":"2025-12-19T14:40:08.750837Z","shell.execute_reply":"2025-12-19T14:40:08.754956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_balanced_patch(volume, label, patch_size,\n                       fg_prob=0.6,\n                       min_fg_ratio=0.01,\n                       max_tries=30):\n    \"\"\"\n    fg_prob: probability of sampling a foreground-heavy patch\n    \"\"\"\n\n    D, H, W = volume.shape\n    ps = patch_size\n\n    for _ in range(max_tries):\n        z = random.randint(0, D - ps)\n        y = random.randint(0, H - ps)\n        x = random.randint(0, W - ps)\n\n        img_p = volume[z:z+ps, y:y+ps, x:x+ps]\n        lbl_p = label[z:z+ps, y:y+ps, x:x+ps]\n\n        valid_mask = (lbl_p != 2)\n        if valid_mask.sum() == 0:\n            continue\n\n        fg_ratio = (lbl_p == 1).sum() / valid_mask.sum()\n\n        # Decide whether we WANT foreground or background\n        want_fg = random.random() < fg_prob\n\n        if want_fg and fg_ratio >= min_fg_ratio:\n            return img_p, lbl_p\n\n        if not want_fg and fg_ratio < min_fg_ratio:\n            return img_p, lbl_p\n\n    # fallback (rare)\n    return img_p, lbl_p\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:08.757181Z","iopub.execute_input":"2025-12-19T14:40:08.757455Z","iopub.status.idle":"2025-12-19T14:40:08.770644Z","shell.execute_reply.started":"2025-12-19T14:40:08.757433Z","shell.execute_reply":"2025-12-19T14:40:08.76997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PapyrusPatchDataset(Dataset):\n    def __init__(self, ids, img_dir, lbl_dir, patch_size):\n        self.ids = ids\n        self.img_dir = img_dir\n        self.lbl_dir = lbl_dir\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        vid = self.ids[idx]\n\n        img = tiff.imread(self.img_dir / f\"{vid}.tif\")\n        lbl = tiff.imread(self.lbl_dir / f\"{vid}.tif\")\n\n        img_p, lbl_p = get_balanced_patch(img, lbl, self.patch_size,  fg_prob=0.6) #important tuning point here\n\n        # normalize image\n        img_p = img_p.astype(np.float32) / 255.0\n\n        # create ignore mask\n        valid_mask = (lbl_p != 2)\n\n        # convert to torch\n        img_p = torch.from_numpy(img_p).unsqueeze(0)   # [1, D, H, W]\n        lbl_p = torch.from_numpy(lbl_p).long()\n        valid_mask = torch.from_numpy(valid_mask)\n\n        return img_p, lbl_p, valid_mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:08.771516Z","iopub.execute_input":"2025-12-19T14:40:08.771853Z","iopub.status.idle":"2025-12-19T14:40:08.788737Z","shell.execute_reply.started":"2025-12-19T14:40:08.771823Z","shell.execute_reply":"2025-12-19T14:40:08.788162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q imagecodecs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:08.789597Z","iopub.execute_input":"2025-12-19T14:40:08.789903Z","iopub.status.idle":"2025-12-19T14:40:11.833023Z","shell.execute_reply.started":"2025-12-19T14:40:08.789882Z","shell.execute_reply":"2025-12-19T14:40:11.83224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = PapyrusPatchDataset(\n    train_ids,\n    TRAIN_IMG_DIR,\n    TRAIN_LBL_DIR,\n    PATCH_SIZE\n)\n\nimg_p, lbl_p, valid_mask = train_dataset[0]\n\nprint(\"Image patch:\", img_p.shape, img_p.dtype)\nprint(\"Label patch:\", lbl_p.shape, lbl_p.dtype)\nprint(\"Valid mask:\", valid_mask.shape, valid_mask.dtype)\n\nprint(\"Foreground voxels:\", (lbl_p == 1).sum().item())\nprint(\"Ignored voxels:\", (lbl_p == 2).sum().item())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:11.834352Z","iopub.execute_input":"2025-12-19T14:40:11.834996Z","iopub.status.idle":"2025-12-19T14:40:12.745208Z","shell.execute_reply.started":"2025-12-19T14:40:11.834963Z","shell.execute_reply":"2025-12-19T14:40:12.744573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:12.746175Z","iopub.execute_input":"2025-12-19T14:40:12.746708Z","iopub.status.idle":"2025-12-19T14:40:12.751003Z","shell.execute_reply.started":"2025-12-19T14:40:12.746684Z","shell.execute_reply":"2025-12-19T14:40:12.750258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:12.75361Z","iopub.execute_input":"2025-12-19T14:40:12.753937Z","iopub.status.idle":"2025-12-19T14:40:12.767341Z","shell.execute_reply.started":"2025-12-19T14:40:12.753891Z","shell.execute_reply":"2025-12-19T14:40:12.766657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SmallUNet3D(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.enc1 = ConvBlock(1, 32)\n        self.pool1 = nn.MaxPool3d(2)\n\n        self.enc2 = ConvBlock(32, 64)\n        self.pool2 = nn.MaxPool3d(2)\n\n        self.bottleneck = ConvBlock(64, 128)\n\n        self.up2 = nn.ConvTranspose3d(128, 64, 2, stride=2)\n        self.dec2 = ConvBlock(128, 64)\n\n        self.up1 = nn.ConvTranspose3d(64, 32, 2, stride=2)\n        self.dec1 = ConvBlock(64, 32)\n\n        self.out = nn.Conv3d(32, 1, 1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n\n        b = self.bottleneck(self.pool2(e2))\n\n        d2 = self.dec2(torch.cat([self.up2(b), e2], dim=1))\n        d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))\n\n        return self.out(d1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:12.76808Z","iopub.execute_input":"2025-12-19T14:40:12.768301Z","iopub.status.idle":"2025-12-19T14:40:12.781377Z","shell.execute_reply.started":"2025-12-19T14:40:12.768282Z","shell.execute_reply":"2025-12-19T14:40:12.780668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = SmallUNet3D()\nwith torch.no_grad():\n    out = model(img_p.unsqueeze(0))  # add batch dim\n\nprint(\"Model output shape:\", out.shape)\n\n\n#DO NOT RE RUN THIS CELL U IDIOT SANDWICH","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:12.782148Z","iopub.execute_input":"2025-12-19T14:40:12.782398Z","iopub.status.idle":"2025-12-19T14:40:13.845222Z","shell.execute_reply.started":"2025-12-19T14:40:12.782377Z","shell.execute_reply":"2025-12-19T14:40:13.844482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MaskedBCELoss(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce = torch.nn.BCEWithLogitsLoss(reduction=\"none\")\n\n    def forward(self, logits, targets, valid_mask):\n        \"\"\"\n        logits: [B, 1, D, H, W]\n        targets: [B, D, H, W]  (0 or 1)\n        valid_mask: [B, D, H, W] (bool)\n        \"\"\"\n        targets = targets.float()\n        loss = self.bce(logits.squeeze(1), targets)\n        loss = loss * valid_mask.float()\n        return loss.sum() / (valid_mask.sum() + 1e-6)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:13.846212Z","iopub.execute_input":"2025-12-19T14:40:13.846567Z","iopub.status.idle":"2025-12-19T14:40:13.851983Z","shell.execute_reply.started":"2025-12-19T14:40:13.846512Z","shell.execute_reply":"2025-12-19T14:40:13.851257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 1  # 3D is heavy, start safe\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True,\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:13.852884Z","iopub.execute_input":"2025-12-19T14:40:13.853186Z","iopub.status.idle":"2025-12-19T14:40:13.868249Z","shell.execute_reply.started":"2025-12-19T14:40:13.853157Z","shell.execute_reply":"2025-12-19T14:40:13.867722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = SmallUNet3D().to(device)\ncriterion = MaskedBCELoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n\nprint(\"Using device:\", device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:13.869219Z","iopub.execute_input":"2025-12-19T14:40:13.869447Z","iopub.status.idle":"2025-12-19T14:40:16.878771Z","shell.execute_reply.started":"2025-12-19T14:40:13.869428Z","shell.execute_reply":"2025-12-19T14:40:16.878157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_STEPS = 10\nmodel.train()\n\nfor step, (img_p, lbl_p, valid_mask) in enumerate(train_loader):\n    if step >= NUM_STEPS:\n        break\n\n    img_p = img_p.to(device)\n    lbl_p = lbl_p.to(device)\n    valid_mask = valid_mask.to(device)\n\n    optimizer.zero_grad()\n    logits = model(img_p)\n    loss = criterion(logits, lbl_p, valid_mask)\n    loss.backward()\n    optimizer.step()\n\n    print(f\"Step {step:02d} | Loss: {loss.item():.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:16.87968Z","iopub.execute_input":"2025-12-19T14:40:16.880089Z","iopub.status.idle":"2025-12-19T14:40:23.331553Z","shell.execute_reply.started":"2025-12-19T14:40:16.880064Z","shell.execute_reply":"2025-12-19T14:40:23.330833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MaskedDiceLoss(torch.nn.Module):\n    def __init__(self, eps=1e-6):\n        super().__init__()\n        self.eps = eps\n\n    def forward(self, logits, targets, valid_mask):\n        \"\"\"\n        logits: [B, 1, D, H, W]\n        targets: [B, D, H, W]\n        valid_mask: [B, D, H, W]\n        \"\"\"\n        probs = torch.sigmoid(logits).squeeze(1)\n\n        targets = targets.float()\n        valid_mask = valid_mask.float()\n\n        probs = probs * valid_mask\n        targets = targets * valid_mask\n\n        intersection = (probs * targets).sum()\n        union = probs.sum() + targets.sum()\n\n        dice = (2 * intersection + self.eps) / (union + self.eps)\n        return 1 - dice\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:23.332605Z","iopub.execute_input":"2025-12-19T14:40:23.332834Z","iopub.status.idle":"2025-12-19T14:40:23.338747Z","shell.execute_reply.started":"2025-12-19T14:40:23.332809Z","shell.execute_reply":"2025-12-19T14:40:23.33798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CombinedLoss(torch.nn.Module):\n    def __init__(self, bce_weight=0.5):\n        super().__init__()\n        self.bce = MaskedBCELoss()\n        self.dice = MaskedDiceLoss()\n        self.bce_weight = bce_weight\n\n    def forward(self, logits, targets, valid_mask):\n        bce_loss = self.bce(logits, targets, valid_mask)\n        dice_loss = self.dice(logits, targets, valid_mask)\n        return self.bce_weight * bce_loss + (1 - self.bce_weight) * dice_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:23.339641Z","iopub.execute_input":"2025-12-19T14:40:23.34002Z","iopub.status.idle":"2025-12-19T14:40:23.35763Z","shell.execute_reply.started":"2025-12-19T14:40:23.33999Z","shell.execute_reply":"2025-12-19T14:40:23.356873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = CombinedLoss(bce_weight=0.5)\nNUM_STEPS = 10\nmodel.train()\n\nfor step, (img_p, lbl_p, valid_mask) in enumerate(train_loader):\n    if step >= NUM_STEPS:\n        break\n\n    img_p = img_p.to(device)\n    lbl_p = lbl_p.to(device)\n    valid_mask = valid_mask.to(device)\n\n    optimizer.zero_grad()\n    logits = model(img_p)\n    loss = criterion(logits, lbl_p, valid_mask)\n    loss.backward()\n    optimizer.step()\n\n    print(f\"Step {step:02d} | Loss: {loss.item():.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:23.358457Z","iopub.execute_input":"2025-12-19T14:40:23.358953Z","iopub.status.idle":"2025-12-19T14:40:28.521468Z","shell.execute_reply.started":"2025-12-19T14:40:23.358897Z","shell.execute_reply":"2025-12-19T14:40:28.520801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nwith torch.no_grad():\n    img_p, lbl_p, valid_mask = train_dataset[0]\n\n    img_p = img_p.to(device)\n    lbl_p = lbl_p.to(device)\n\n    logits = model(img_p.unsqueeze(0))\n    probs = torch.sigmoid(logits)[0, 0].cpu().numpy()\n    pred = (probs > 0.5).astype(np.uint8)\n\nimg_np = img_p[0].cpu().numpy()\nlbl_np = lbl_p.cpu().numpy()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:28.522447Z","iopub.execute_input":"2025-12-19T14:40:28.522719Z","iopub.status.idle":"2025-12-19T14:40:28.867945Z","shell.execute_reply.started":"2025-12-19T14:40:28.522694Z","shell.execute_reply":"2025-12-19T14:40:28.867099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nz = img_np.shape[0] // 2\n\nfig, axs = plt.subplots(1, 3, figsize=(15, 5))\n\naxs[0].imshow(img_np[z], cmap=\"gray\")\naxs[0].set_title(\"CT Image\")\n\naxs[1].imshow(lbl_np[z], cmap=\"viridis\")\naxs[1].set_title(\"Ground Truth\")\n\naxs[2].imshow(pred[z], cmap=\"viridis\")\naxs[2].set_title(\"Model Prediction\")\n\nfor ax in axs:\n    ax.axis(\"off\")\n\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:28.870246Z","iopub.execute_input":"2025-12-19T14:40:28.870489Z","iopub.status.idle":"2025-12-19T14:40:29.085568Z","shell.execute_reply.started":"2025-12-19T14:40:28.870468Z","shell.execute_reply":"2025-12-19T14:40:29.084769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_id = val_ids[0]\nprint(\"Validation volume ID:\", val_id)\n\nvol = tiff.imread(TRAIN_IMG_DIR / f\"{val_id}.tif\")\ngt  = tiff.imread(TRAIN_LBL_DIR / f\"{val_id}.tif\")\n\nprint(\"Volume shape:\", vol.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:29.086597Z","iopub.execute_input":"2025-12-19T14:40:29.087142Z","iopub.status.idle":"2025-12-19T14:40:29.67382Z","shell.execute_reply.started":"2025-12-19T14:40:29.087117Z","shell.execute_reply":"2025-12-19T14:40:29.673199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sliding_window_inference(\n    volume,\n    model,\n    patch_size=64,\n    stride=32,   # 50% overlap\n    device=\"cuda\"\n):\n    model.eval()\n\n    D, H, W = volume.shape\n\n    prob_map = np.zeros((D, H, W), dtype=np.float32)\n    count_map = np.zeros((D, H, W), dtype=np.float32)\n\n    with torch.no_grad():\n        for z in range(0, D - patch_size + 1, stride):\n            for y in range(0, H - patch_size + 1, stride):\n                for x in range(0, W - patch_size + 1, stride):\n\n                    patch = volume[z:z+patch_size, y:y+patch_size, x:x+patch_size]\n                    patch = patch.astype(np.float32) / 255.0\n\n                    patch = torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).to(device)\n\n                    logits = model(patch)\n                    probs = torch.sigmoid(logits)[0, 0].cpu().numpy()\n\n                    prob_map[z:z+patch_size, y:y+patch_size, x:x+patch_size] += probs\n                    count_map[z:z+patch_size, y:y+patch_size, x:x+patch_size] += 1.0\n\n    prob_map /= np.maximum(count_map, 1e-6)\n    return prob_map\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:29.674671Z","iopub.execute_input":"2025-12-19T14:40:29.675Z","iopub.status.idle":"2025-12-19T14:40:29.681946Z","shell.execute_reply.started":"2025-12-19T14:40:29.674969Z","shell.execute_reply":"2025-12-19T14:40:29.68125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\nprob_map = sliding_window_inference(\n    vol,\n    model,\n    patch_size=PATCH_SIZE,\n    stride=PATCH_SIZE // 2,\n    device=device\n)\n\nprint(\"Inference done.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:29.682779Z","iopub.execute_input":"2025-12-19T14:40:29.68303Z","iopub.status.idle":"2025-12-19T14:40:56.81375Z","shell.execute_reply.started":"2025-12-19T14:40:29.68301Z","shell.execute_reply":"2025-12-19T14:40:56.813098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"THRESHOLD = 0.4  # stronger than 0.5\n\npred_vol = (prob_map > THRESHOLD).astype(np.uint8)\n\nprint(\"Predicted foreground voxels (before masking):\", pred_vol.sum())\n\n# Ignore unlabeled regions (label == 2)\nvalid_region = (gt != 2)\n\npred_vol = pred_vol * valid_region\n\nprint(\"Foreground voxels after GT mask:\", pred_vol.sum())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:56.814723Z","iopub.execute_input":"2025-12-19T14:40:56.815048Z","iopub.status.idle":"2025-12-19T14:40:56.901058Z","shell.execute_reply.started":"2025-12-19T14:40:56.815014Z","shell.execute_reply":"2025-12-19T14:40:56.900313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import scipy.ndimage as ndi\n\nlabeled, num_components = ndi.label(pred_vol)\n\nsizes = ndi.sum(pred_vol, labeled, index=range(1, num_components + 1))\n\nprint(\"Number of connected components before filtering:\", num_components)\n\nMIN_COMPONENT_SIZE = 1500  # safe starting value\n\nclean_pred = np.zeros_like(pred_vol)\n\nfor i, size in enumerate(sizes, start=1):\n    if size >= MIN_COMPONENT_SIZE:\n        clean_pred[labeled == i] = 1\n\npred_vol = clean_pred\n\nprint(\"Number of connected components after filtering:\",\n      ndi.label(pred_vol)[1])\nprint(\"Foreground voxels after CC filtering:\", pred_vol.sum())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:56.903186Z","iopub.execute_input":"2025-12-19T14:40:56.903467Z","iopub.status.idle":"2025-12-19T14:40:58.087853Z","shell.execute_reply.started":"2025-12-19T14:40:56.903447Z","shell.execute_reply":"2025-12-19T14:40:58.087261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cz, cy, cx = vol.shape[0]//2, vol.shape[1]//2, vol.shape[2]//2\n\nfig, axs = plt.subplots(3, 3, figsize=(15, 15))\n\n# CT\naxs[0,0].imshow(vol[cz], cmap=\"gray\")\naxs[0,1].imshow(vol[:, cy, :], cmap=\"gray\")\naxs[0,2].imshow(vol[:, :, cx], cmap=\"gray\")\n\n# Ground Truth\naxs[1,0].imshow(gt[cz], cmap=\"viridis\")\naxs[1,1].imshow(gt[:, cy, :], cmap=\"viridis\")\naxs[1,2].imshow(gt[:, :, cx], cmap=\"viridis\")\n\n# Cleaned Prediction\naxs[2,0].imshow(pred_vol[cz], cmap=\"viridis\")\naxs[2,1].imshow(pred_vol[:, cy, :], cmap=\"viridis\")\naxs[2,2].imshow(pred_vol[:, :, cx], cmap=\"viridis\")\n\ntitles = [\"XY\", \"XZ\", \"YZ\"]\nfor i in range(3):\n    axs[0,i].set_title(f\"CT {titles[i]}\")\n    axs[1,i].set_title(f\"GT {titles[i]}\")\n    axs[2,i].set_title(f\"PRED (clean) {titles[i]}\")\n\nfor ax in axs.flatten():\n    ax.axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:58.088966Z","iopub.execute_input":"2025-12-19T14:40:58.089465Z","iopub.status.idle":"2025-12-19T14:40:59.084974Z","shell.execute_reply.started":"2025-12-19T14:40:58.089438Z","shell.execute_reply":"2025-12-19T14:40:59.084089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\n\nMETRIC_DIR = \"/kaggle/working/metrics\"\nos.makedirs(METRIC_DIR, exist_ok=True)\n\nsys.path.append(METRIC_DIR)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.085953Z","iopub.execute_input":"2025-12-19T14:40:59.086204Z","iopub.status.idle":"2025-12-19T14:40:59.090684Z","shell.execute_reply.started":"2025-12-19T14:40:59.086182Z","shell.execute_reply":"2025-12-19T14:40:59.089859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/metrics/surface_dice.py\nimport numpy as np\nfrom scipy.ndimage import binary_erosion, distance_transform_edt\n\ndef surface_dice(pred, gt, spacing=(1,1,1), tolerance=2.0, ignore_mask=None):\n    pred = pred.astype(bool)\n    gt = gt.astype(bool)\n\n    if ignore_mask is not None:\n        pred = pred & ~ignore_mask\n        gt = gt & ~ignore_mask\n\n    if pred.sum() == 0 and gt.sum() == 0:\n        return 1.0\n    if pred.sum() == 0 or gt.sum() == 0:\n        return 0.0\n\n    pred_surface = pred ^ binary_erosion(pred)\n    gt_surface = gt ^ binary_erosion(gt)\n\n    dt_gt = distance_transform_edt(~gt_surface, sampling=spacing)\n    dt_pred = distance_transform_edt(~pred_surface, sampling=spacing)\n\n    pred_to_gt = (dt_gt[pred_surface] <= tolerance).sum()\n    gt_to_pred = (dt_pred[gt_surface] <= tolerance).sum()\n\n    return (pred_to_gt + gt_to_pred) / (pred_surface.sum() + gt_surface.sum())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.091523Z","iopub.execute_input":"2025-12-19T14:40:59.091817Z","iopub.status.idle":"2025-12-19T14:40:59.105661Z","shell.execute_reply.started":"2025-12-19T14:40:59.091781Z","shell.execute_reply":"2025-12-19T14:40:59.10497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/metrics/voi.py\nimport numpy as np\nfrom scipy.ndimage import label\n\ndef voi_score(pred, gt, ignore_mask=None):\n    if ignore_mask is not None:\n        pred = pred & ~ignore_mask\n        gt = gt & ~ignore_mask\n\n    pred_lab, _ = label(pred)\n    gt_lab, _ = label(gt)\n\n    pred_ids = np.unique(pred_lab)\n    gt_ids = np.unique(gt_lab)\n\n    pred_ids = pred_ids[pred_ids != 0]\n    gt_ids = gt_ids[gt_ids != 0]\n\n    if len(pred_ids) == 0 and len(gt_ids) == 0:\n        return 1.0\n    if len(pred_ids) == 0 or len(gt_ids) == 0:\n        return 0.0\n\n    intersection = {}\n    for p in pred_ids:\n        for g in gt_ids:\n            inter = np.logical_and(pred_lab == p, gt_lab == g).sum()\n            if inter > 0:\n                intersection[(p, g)] = inter\n\n    total_pred = pred.sum()\n    total_gt = gt.sum()\n\n    voi_split = sum(\n        inter * np.log(inter / (pred_lab == p).sum())\n        for (p, g), inter in intersection.items()\n    ) / total_gt\n\n    voi_merge = sum(\n        inter * np.log(inter / (gt_lab == g).sum())\n        for (p, g), inter in intersection.items()\n    ) / total_pred\n\n    voi_total = voi_split + voi_merge\n    alpha = 0.3\n\n    return 1.0 / (1.0 + alpha * voi_total)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.106647Z","iopub.execute_input":"2025-12-19T14:40:59.107123Z","iopub.status.idle":"2025-12-19T14:40:59.123372Z","shell.execute_reply.started":"2025-12-19T14:40:59.107102Z","shell.execute_reply":"2025-12-19T14:40:59.122854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/metrics/toposcore.py\nimport numpy as np\nfrom scipy.ndimage import label\n\ndef toposcore(pred, gt, ignore_mask=None):\n    if ignore_mask is not None:\n        pred = pred & ~ignore_mask\n        gt = gt & ~ignore_mask\n\n    pred_cc, num_p = label(pred)\n    gt_cc, num_g = label(gt)\n\n    if num_p == 0 and num_g == 0:\n        return 1.0\n    if num_p == 0 or num_g == 0:\n        return 0.0\n\n    return min(num_g, num_p) / max(num_g, num_p)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.124252Z","iopub.execute_input":"2025-12-19T14:40:59.124547Z","iopub.status.idle":"2025-12-19T14:40:59.138768Z","shell.execute_reply.started":"2025-12-19T14:40:59.124527Z","shell.execute_reply":"2025-12-19T14:40:59.13817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from surface_dice import surface_dice\nfrom voi import voi_score\nfrom toposcore import toposcore\n\nprint(\"Metrics imported successfully\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.139521Z","iopub.execute_input":"2025-12-19T14:40:59.139739Z","iopub.status.idle":"2025-12-19T14:40:59.154633Z","shell.execute_reply.started":"2025-12-19T14:40:59.13972Z","shell.execute_reply":"2025-12-19T14:40:59.154031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prediction (binary, already cleaned)\npred_bin = pred_vol.astype(np.uint8)\n\n# Ground truth: papyrus only\ngt_bin = (gt == 1).astype(np.uint8)\n\n# Ignore mask: unlabeled regions\nignore_mask = (gt == 2)\n\nprint(\"pred_bin foreground:\", pred_bin.sum())\nprint(\"gt_bin foreground:\", gt_bin.sum())\nprint(\"ignored voxels:\", ignore_mask.sum())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.155327Z","iopub.execute_input":"2025-12-19T14:40:59.155564Z","iopub.status.idle":"2025-12-19T14:40:59.255775Z","shell.execute_reply.started":"2025-12-19T14:40:59.15554Z","shell.execute_reply":"2025-12-19T14:40:59.255043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sd = surface_dice(pred_bin, gt_bin, ignore_mask=ignore_mask)\nvoi = voi_score(pred_bin, gt_bin, ignore_mask=ignore_mask)\ntopo = toposcore(pred_bin, gt_bin, ignore_mask=ignore_mask)\n\nfinal_score = 0.30 * topo + 0.35 * sd + 0.35 * voi\n\nprint(\"SurfaceDice:\", sd)\nprint(\"VOI:\", voi)\nprint(\"TopoScore:\", topo)\nprint(\"Final score:\", final_score)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:40:59.256721Z","iopub.execute_input":"2025-12-19T14:40:59.256994Z","iopub.status.idle":"2025-12-19T14:41:15.839395Z","shell.execute_reply.started":"2025-12-19T14:40:59.256972Z","shell.execute_reply":"2025-12-19T14:41:15.838684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fg_ratios = []\n\nfor i in range(50):\n    _, lbl_p, valid_mask = train_dataset[i]\n    fg = (lbl_p == 1).sum().item()\n    valid = valid_mask.sum().item()\n    fg_ratios.append(fg / (valid + 1e-6))\n\nprint(\"FG ratio stats:\")\nprint(\"  min :\", min(fg_ratios))\nprint(\"  mean:\", sum(fg_ratios) / len(fg_ratios))\nprint(\"  max :\", max(fg_ratios))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:41:15.840317Z","iopub.execute_input":"2025-12-19T14:41:15.840569Z","iopub.status.idle":"2025-12-19T14:41:49.846183Z","shell.execute_reply.started":"2025-12-19T14:41:15.840547Z","shell.execute_reply":"2025-12-19T14:41:49.845464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_STEPS_PER_EPOCH = 20  # IMPORTANT: cap work per epoch\nNUM_EPOCHS = 3\n\nfor epoch in range(NUM_EPOCHS):\n    model.train()\n    losses = []\n\n    for step, (img_p, lbl_p, valid_mask) in enumerate(train_loader):\n        if step >= MAX_STEPS_PER_EPOCH:\n            break\n\n        img_p = img_p.to(device)\n        lbl_p = lbl_p.to(device)\n        valid_mask = valid_mask.to(device)\n\n        optimizer.zero_grad()\n        logits = model(img_p)\n        loss = criterion(logits, lbl_p, valid_mask)\n        loss.backward()\n        optimizer.step()\n\n        losses.append(loss.item())\n\n        # progress feedback so it never \"looks stuck\"\n        if step % 25 == 0:\n            print(f\"Epoch {epoch} | Step {step}/{MAX_STEPS_PER_EPOCH} | Loss {loss.item():.4f}\")\n\n    print(f\"Epoch {epoch} DONE | Mean loss: {sum(losses)/len(losses):.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:41:49.847152Z","iopub.execute_input":"2025-12-19T14:41:49.84745Z","iopub.status.idle":"2025-12-19T14:42:18.837509Z","shell.execute_reply.started":"2025-12-19T14:41:49.84742Z","shell.execute_reply":"2025-12-19T14:42:18.836783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================\n# FULL VOLUME VALIDATION PIPELINE (RELAXED)\n# ==============================\n\nimport numpy as np\nimport torch\nimport scipy.ndimage as ndi\nimport matplotlib.pyplot as plt\n\n# -------- CONFIG (UPDATED KNOBS) --------\nPATCH_SIZE = 64\nSTRIDE = PATCH_SIZE // 2\nTHRESHOLD = 0.42             # ↓ relaxed from 0.7\nMIN_COMPONENT_SIZE = 1500     # ↓ relaxed from 5000\n\n# pick one validation volume\nval_id = val_ids[0]\nprint(\"Validation volume:\", val_id)\n\n# load volume + GT\nvol = tiff.imread(TRAIN_IMG_DIR / f\"{val_id}.tif\")\ngt  = tiff.imread(TRAIN_LBL_DIR / f\"{val_id}.tif\")\n\nD, H, W = vol.shape\nprint(\"Volume shape:\", vol.shape)\n\n# -------- SLIDING WINDOW INFERENCE --------\nmodel.eval()\ndevice = next(model.parameters()).device\n\nprob_map = np.zeros((D, H, W), dtype=np.float32)\ncount_map = np.zeros((D, H, W), dtype=np.float32)\n\nwith torch.no_grad():\n    for z in range(0, D - PATCH_SIZE + 1, STRIDE):\n        for y in range(0, H - PATCH_SIZE + 1, STRIDE):\n            for x in range(0, W - PATCH_SIZE + 1, STRIDE):\n\n                patch = vol[z:z+PATCH_SIZE, y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n                patch = patch.astype(np.float32) / 255.0\n\n                patch = torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).to(device)\n\n                logits = model(patch)\n                probs = torch.sigmoid(logits)[0, 0].cpu().numpy()\n\n                prob_map[z:z+PATCH_SIZE, y:y+PATCH_SIZE, x:x+PATCH_SIZE] += probs\n                count_map[z:z+PATCH_SIZE, y:y+PATCH_SIZE, x:x+PATCH_SIZE] += 1\n\nprob_map /= np.maximum(count_map, 1e-6)\n\nprint(\"Inference done.\")\n\n# -------- THRESHOLD --------\npred_vol = (prob_map > THRESHOLD).astype(np.uint8)\nprint(\"Foreground after threshold:\", pred_vol.sum())\n\n# -------- GT MASK (validation only) --------\nignore_mask = (gt == 2)\npred_vol = pred_vol * (~ignore_mask)\n\nprint(\"Foreground after GT mask:\", pred_vol.sum())\n\n# -------- CONNECTED COMPONENT FILTERING --------\nlabeled, num = ndi.label(pred_vol)\nsizes = ndi.sum(pred_vol, labeled, range(1, num + 1))\n\nclean_pred = np.zeros_like(pred_vol)\n\nfor i, size in enumerate(sizes, start=1):\n    if size >= MIN_COMPONENT_SIZE:\n        clean_pred[labeled == i] = 1\n\npred_vol = clean_pred\n\nprint(\"Connected components after filtering:\", ndi.label(pred_vol)[1])\nprint(\"Foreground after CC filtering:\", pred_vol.sum())\n\n# -------- PREP METRICS --------\npred_bin = pred_vol.astype(np.uint8)\ngt_bin = (gt == 1).astype(np.uint8)\n\nprint(\"pred_bin foreground:\", pred_bin.sum())\nprint(\"gt_bin foreground:\", gt_bin.sum())\nprint(\"ignored voxels:\", ignore_mask.sum())\n\n# -------- METRICS --------\nsd = surface_dice(pred_bin, gt_bin, ignore_mask=ignore_mask)\nvoi = voi_score(pred_bin, gt_bin, ignore_mask=ignore_mask)\ntopo = toposcore(pred_bin, gt_bin, ignore_mask=ignore_mask)\n\nfinal_score = 0.30 * topo + 0.35 * sd + 0.35 * voi\n\nprint(\"\\n===== METRICS (1 VALIDATION VOLUME) =====\")\nprint(\"SurfaceDice :\", sd)\nprint(\"VOI         :\", voi)\nprint(\"TopoScore   :\", topo)\nprint(\"Final score :\", final_score)\n\n# -------- QUICK VISUAL CHECK --------\ncz, cy, cx = D//2, H//2, W//2\n\nfig, axs = plt.subplots(2, 3, figsize=(15, 10))\n\naxs[0,0].imshow(gt[cz], cmap=\"viridis\"); axs[0,0].set_title(\"GT XY\")\naxs[0,1].imshow(gt[:, cy, :], cmap=\"viridis\"); axs[0,1].set_title(\"GT XZ\")\naxs[0,2].imshow(gt[:, :, cx], cmap=\"viridis\"); axs[0,2].set_title(\"GT YZ\")\n\naxs[1,0].imshow(pred_vol[cz], cmap=\"viridis\"); axs[1,0].set_title(\"PRED XY\")\naxs[1,1].imshow(pred_vol[:, cy, :], cmap=\"viridis\"); axs[1,1].set_title(\"PRED XZ\")\naxs[1,2].imshow(pred_vol[:, :, cx], cmap=\"viridis\"); axs[1,2].set_title(\"PRED YZ\")\n\nfor ax in axs.flatten():\n    ax.axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T14:42:18.838743Z","iopub.execute_input":"2025-12-19T14:42:18.839052Z","iopub.status.idle":"2025-12-19T14:43:06.140199Z","shell.execute_reply.started":"2025-12-19T14:42:18.839023Z","shell.execute_reply":"2025-12-19T14:43:06.139462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ntrain_df = pd.read_csv(\"/kaggle/input/vesuvius-challenge-surface-detection/train.csv\")\n\nprint(\"Scroll distribution:\")\nprint(train_df.scroll_id.value_counts())\n\n# Choose validation scroll(s)\nVAL_SCROLLS = [train_df.scroll_id.unique()[0]]  # change if you want\n\ntrain_ids = train_df[~train_df.scroll_id.isin(VAL_SCROLLS)].id.tolist()\nval_ids   = train_df[ train_df.scroll_id.isin(VAL_SCROLLS)].id.tolist()\n\nprint(f\"\\nValidation scrolls: {VAL_SCROLLS}\")\nprint(f\"Train volumes: {len(train_ids)}\")\nprint(f\"Val volumes: {len(val_ids)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T15:21:16.446277Z","iopub.execute_input":"2025-12-19T15:21:16.446618Z","iopub.status.idle":"2025-12-19T15:21:16.460382Z","shell.execute_reply.started":"2025-12-19T15:21:16.446591Z","shell.execute_reply":"2025-12-19T15:21:16.459669Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\n\ndef infer_full_volume(\n    model,\n    volume,\n    patch_size=64,\n    stride=32,\n    threshold=0.5\n):\n    \"\"\"\n    volume: torch.Tensor [1, 1, D, H, W]\n    returns: numpy array [D, H, W] (binary)\n    \"\"\"\n    model.eval()\n    _, _, D, H, W = volume.shape\n\n    prob_map = torch.zeros((D, H, W), device=volume.device)\n    count_map = torch.zeros((D, H, W), device=volume.device)\n\n    with torch.no_grad():\n        for z in range(0, D - patch_size + 1, stride):\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[:, :, z:z+patch_size,\n                                      y:y+patch_size,\n                                      x:x+patch_size]\n\n                    pred = model(patch).sigmoid()[0, 0]\n\n                    prob_map[z:z+patch_size,\n                             y:y+patch_size,\n                             x:x+patch_size] += pred\n                    count_map[z:z+patch_size,\n                              y:y+patch_size,\n                              x:x+patch_size] += 1\n\n    prob_map /= torch.clamp(count_map, min=1)\n    return (prob_map > threshold).cpu().numpy()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T15:21:23.617084Z","iopub.execute_input":"2025-12-19T15:21:23.617703Z","iopub.status.idle":"2025-12-19T15:21:23.624602Z","shell.execute_reply.started":"2025-12-19T15:21:23.617676Z","shell.execute_reply":"2025-12-19T15:21:23.623735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tifffile as tiff\nfrom tqdm import tqdm\n\npredictions = {}\n\nfor vid in tqdm(val_ids):\n    img = tiff.imread(\n        f\"/kaggle/input/vesuvius-challenge-surface-detection/train_images/{vid}.tif\"\n    )\n\n    img_t = torch.from_numpy(img).float().unsqueeze(0).unsqueeze(0).cuda()\n\n    pred_bin = infer_full_volume(\n        model,\n        img_t,\n        patch_size=64,\n        stride=32,\n        threshold=0.5\n    )\n\n    predictions[vid] = pred_bin\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T15:21:30.957448Z","iopub.execute_input":"2025-12-19T15:21:30.957947Z","iopub.status.idle":"2025-12-19T16:04:40.731096Z","shell.execute_reply.started":"2025-12-19T15:21:30.957921Z","shell.execute_reply":"2025-12-19T16:04:40.730442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_single_volume(\n    model,\n    vol,\n    gt,\n    patch_size=PATCH_SIZE,\n    stride=STRIDE,\n    threshold=THRESHOLD,\n    min_cc_size=MIN_COMPONENT_SIZE,\n):\n    \"\"\"\n    Returns: surface_dice, voi, topo, final_score, pred_vol\n    \"\"\"\n\n    D, H, W = vol.shape\n    device = next(model.parameters()).device\n\n    # ---------- sliding window ----------\n    prob_map = np.zeros((D, H, W), dtype=np.float32)\n    count_map = np.zeros((D, H, W), dtype=np.float32)\n\n    model.eval()\n    with torch.no_grad():\n        for z in range(0, D - patch_size + 1, stride):\n            for y in range(0, H - patch_size + 1, stride):\n                for x in range(0, W - patch_size + 1, stride):\n\n                    patch = vol[z:z+patch_size, y:y+patch_size, x:x+patch_size]\n                    patch = patch.astype(np.float32) / 255.0\n                    patch = torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).to(device)\n\n                    probs = torch.sigmoid(model(patch))[0, 0].cpu().numpy()\n\n                    prob_map[z:z+patch_size,\n                             y:y+patch_size,\n                             x:x+patch_size] += probs\n                    count_map[z:z+patch_size,\n                              y:y+patch_size,\n                              x:x+patch_size] += 1\n\n    prob_map /= np.maximum(count_map, 1e-6)\n\n    # ---------- threshold ----------\n    pred_vol = (prob_map > threshold).astype(np.uint8)\n\n    # ---------- ignore mask ----------\n    ignore_mask = (gt == 2)\n    pred_vol = pred_vol * (~ignore_mask)\n\n    # ---------- connected component filtering ----------\n    labeled, num = ndi.label(pred_vol)\n    sizes = ndi.sum(pred_vol, labeled, range(1, num + 1))\n\n    clean_pred = np.zeros_like(pred_vol)\n    for i, size in enumerate(sizes, start=1):\n        if size >= min_cc_size:\n            clean_pred[labeled == i] = 1\n\n    pred_vol = clean_pred\n\n    # ---------- prep metrics ----------\n    pred_bin = pred_vol.astype(np.uint8)\n    gt_bin = (gt == 1).astype(np.uint8)\n\n    # ---------- metrics ----------\n    sd = surface_dice(pred_bin, gt_bin, ignore_mask=ignore_mask)\n    voi = voi_score(pred_bin, gt_bin, ignore_mask=ignore_mask)\n    topo = toposcore(pred_bin, gt_bin, ignore_mask=ignore_mask)\n\n    final = 0.30 * topo + 0.35 * sd + 0.35 * voi\n\n    return sd, voi, topo, final, pred_vol\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T16:14:00.554591Z","iopub.execute_input":"2025-12-19T16:14:00.554952Z","iopub.status.idle":"2025-12-19T16:14:00.564837Z","shell.execute_reply.started":"2025-12-19T16:14:00.554896Z","shell.execute_reply":"2025-12-19T16:14:00.56422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\nresults = {}\npred_cache = {}\n\nfor vid in tqdm(val_ids):\n    vol = tiff.imread(TRAIN_IMG_DIR / f\"{vid}.tif\")\n    gt  = tiff.imread(TRAIN_LBL_DIR / f\"{vid}.tif\")\n\n    sd, voi, topo, final, pred_vol = evaluate_single_volume(\n        model, vol, gt\n    )\n\n    results[vid] = {\n        \"surface_dice\": sd,\n        \"voi\": voi,\n        \"topo\": topo,\n        \"final\": final\n    }\n\n    pred_cache[vid] = pred_vol\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T16:14:13.747707Z","iopub.execute_input":"2025-12-19T16:14:13.748047Z","iopub.status.idle":"2025-12-19T17:29:18.372506Z","shell.execute_reply.started":"2025-12-19T16:14:13.74802Z","shell.execute_reply":"2025-12-19T17:29:18.371837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.DataFrame.from_dict(results, orient=\"index\")\ndf.index.name = \"volume_id\"\n\ndisplay(df)\ndisplay(df.describe())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T17:32:14.752126Z","iopub.execute_input":"2025-12-19T17:32:14.752702Z","iopub.status.idle":"2025-12-19T17:32:14.791453Z","shell.execute_reply.started":"2025-12-19T17:32:14.752674Z","shell.execute_reply":"2025-12-19T17:32:14.790743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sorted = df.sort_values(\"final\")\n\nprint(\"=== WORST 5 VOLUMES ===\")\ndisplay(df_sorted.head(5))\n\nprint(\"\\n=== BEST 5 VOLUMES ===\")\ndisplay(df_sorted.tail(5))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T17:32:17.151757Z","iopub.execute_input":"2025-12-19T17:32:17.152084Z","iopub.status.idle":"2025-12-19T17:32:17.169785Z","shell.execute_reply.started":"2025-12-19T17:32:17.152058Z","shell.execute_reply":"2025-12-19T17:32:17.169247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"worst_ids = df_sorted.head(2).index.tolist()\n\nfor vid in worst_ids:\n    vol = tiff.imread(TRAIN_IMG_DIR / f\"{vid}.tif\")\n    gt  = tiff.imread(TRAIN_LBL_DIR / f\"{vid}.tif\")\n    pred = pred_cache[vid]\n\n    D, H, W = vol.shape\n    cz, cy, cx = D//2, H//2, W//2\n\n    fig, axs = plt.subplots(2, 3, figsize=(15, 8))\n\n    axs[0,0].imshow(gt[cz], cmap=\"viridis\"); axs[0,0].set_title(\"GT XY\")\n    axs[0,1].imshow(gt[:, cy, :], cmap=\"viridis\"); axs[0,1].set_title(\"GT XZ\")\n    axs[0,2].imshow(gt[:, :, cx], cmap=\"viridis\"); axs[0,2].set_title(\"GT YZ\")\n\n    axs[1,0].imshow(pred[cz], cmap=\"viridis\"); axs[1,0].set_title(\"PRED XY\")\n    axs[1,1].imshow(pred[:, cy, :], cmap=\"viridis\"); axs[1,1].set_title(\"PRED XZ\")\n    axs[1,2].imshow(pred[:, :, cx], cmap=\"viridis\"); axs[1,2].set_title(\"PRED YZ\")\n\n    for ax in axs.flatten():\n        ax.axis(\"off\")\n\n    plt.suptitle(f\"Failure case: {vid}\")\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T17:32:19.89936Z","iopub.execute_input":"2025-12-19T17:32:19.899751Z","iopub.status.idle":"2025-12-19T17:32:21.337776Z","shell.execute_reply.started":"2025-12-19T17:32:19.899724Z","shell.execute_reply":"2025-12-19T17:32:21.337144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"        📊 VALIDATION SUMMARY (GLOBAL)\")\nprint(\"=\"*60)\n\nmean_vals = df.mean()\nstd_vals  = df.std()\n\nprint(f\"\\nMean Final Score      : {mean_vals['final']:.4f}\")\nprint(f\"Std  Final Score      : {std_vals['final']:.4f}\")\n\nprint(\"\\nMean Metric Breakdown:\")\nprint(f\"  SurfaceDice         : {mean_vals['surface_dice']:.4f}\")\nprint(f\"  VOI_score           : {mean_vals['voi']:.4f}\")\nprint(f\"  TopoScore           : {mean_vals['topo']:.4f}\")\n\n# identify weakest metric\nweakest_metric = mean_vals[['surface_dice', 'voi', 'topo']].idxmin()\n\nprint(\"\\n⚠️ Weakest Metric:\")\nprint(f\"  → {weakest_metric}\")\n\n# worst & best volumes\nworst_id = df['final'].idxmin()\nbest_id  = df['final'].idxmax()\n\nprint(\"\\n📉 Worst Validation Volume:\")\nprint(df.loc[worst_id])\nprint(f\"  Volume ID: {worst_id}\")\n\nprint(\"\\n📈 Best Validation Volume:\")\nprint(df.loc[best_id])\nprint(f\"  Volume ID: {best_id}\")\n\n# stability check\nif std_vals['final'] > 0.1:\n    stability = \"❌ Unstable (high variance)\"\nelif std_vals['final'] > 0.05:\n    stability = \"⚠️ Moderate variance\"\nelse:\n    stability = \"✅ Stable\"\n\nprint(\"\\n📐 Model Stability:\")\nprint(f\"  {stability}\")\n\nprint(\"\\n💡 Suggested Focus:\")\nif weakest_metric == \"topo\":\n    print(\"  → Fix topology: remove holes, bridges, broken surfaces\")\nelif weakest_metric == \"voi\":\n    print(\"  → Fix splits/merges: connectivity & component handling\")\nelse:\n    print(\"  → Fix boundary accuracy: surface thinning / alignment\")\n\nprint(\"\\n\" + \"=\"*60)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T17:33:59.207669Z","iopub.execute_input":"2025-12-19T17:33:59.208268Z","iopub.status.idle":"2025-12-19T17:33:59.222548Z","shell.execute_reply.started":"2025-12-19T17:33:59.208241Z","shell.execute_reply":"2025-12-19T17:33:59.221851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}