{"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":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# VESUVIUS CHALLENGE – PATCH-BASED BASELINE\n# ZIP VISIBILITY FIXED\n# ============================================================\n\n!pip install -q imagecodecs\n\nimport os, zipfile, random, sys\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport tifffile as tiff\nfrom tqdm import tqdm\nfrom scipy.ndimage import label, binary_fill_holes\n\n# -----------------------\n# CONFIG\n# -----------------------\nDATA_ROOT = \"/kaggle/input/vesuvius-challenge-surface-detection\"\nTRAIN_IMG = f\"{DATA_ROOT}/train_images\"\nTRAIN_LBL = f\"{DATA_ROOT}/train_labels\"\nTEST_IMG  = f\"{DATA_ROOT}/test_images\"\n\nWORK_DIR = \"/kaggle/working\"\nPRED_DIR = f\"{WORK_DIR}/preds\"\nZIP_PATH = f\"{WORK_DIR}/submission.zip\"\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nPATCH = 64\nPATCHES_PER_VOLUME = 4\nBATCH_SIZE = 2\nEPOCHS = 2\nLR = 1e-4\n\nTHRESH = 0.5\nMIN_COMPONENT = 800\n\nos.makedirs(PRED_DIR, exist_ok=True)\n\n# -----------------------\n# PATCH DATASET\n# -----------------------\nclass PatchDataset(Dataset):\n    def __init__(self, img_dir, lbl_dir):\n        self.img_dir = img_dir\n        self.lbl_dir = lbl_dir\n        self.ids = sorted(os.listdir(img_dir))\n\n    def __len__(self):\n        return len(self.ids) * PATCHES_PER_VOLUME\n\n    def __getitem__(self, idx):\n        vid = self.ids[idx // PATCHES_PER_VOLUME]\n        vol = tiff.imread(os.path.join(self.img_dir, vid))\n        lbl = tiff.imread(os.path.join(self.lbl_dir, vid))\n\n        z, y, x = vol.shape\n        zz = random.randint(0, z - PATCH)\n        yy = random.randint(0, y - PATCH)\n        xx = random.randint(0, x - PATCH)\n\n        vol = torch.from_numpy(vol[zz:zz+PATCH, yy:yy+PATCH, xx:xx+PATCH]).float().unsqueeze(0)\n        lbl = torch.from_numpy(lbl[zz:zz+PATCH, yy:yy+PATCH, xx:xx+PATCH]).float().unsqueeze(0)\n\n        return vol, lbl\n\n# -----------------------\n# MODEL\n# -----------------------\nclass Block(nn.Module):\n    def __init__(self, a, b):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv3d(a, b, 3, padding=1),\n            nn.InstanceNorm3d(b),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(b, b, 3, padding=1),\n            nn.InstanceNorm3d(b),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.net(x)\n\nclass UNet3D(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.e1 = Block(1, 16)\n        self.e2 = Block(16, 32)\n        self.d1 = Block(48, 16)\n        self.out = nn.Conv3d(16, 1, 1)\n\n    def forward(self, x):\n        x1 = self.e1(x)\n        x2 = self.e2(F.max_pool3d(x1, 2))\n        x = F.interpolate(x2, scale_factor=2, mode=\"nearest\")\n        x = self.d1(torch.cat([x, x1], dim=1))\n        return torch.sigmoid(self.out(x))\n\n# -----------------------\n# LOSS\n# -----------------------\ndef dice_loss(p, y, eps=1e-6):\n    num = 2 * (p * y).sum()\n    den = p.sum() + y.sum() + eps\n    return 1 - num / den\n\n# -----------------------\n# TRAIN\n# -----------------------\ntrain_ds = PatchDataset(TRAIN_IMG, TRAIN_LBL)\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nmodel = UNet3D().to(DEVICE)\nopt = torch.optim.AdamW(model.parameters(), lr=LR)\n\nmodel.train()\nfor epoch in range(EPOCHS):\n    bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}\")\n    for x, y in bar:\n        x = x.to(DEVICE, non_blocking=True)\n        y = y.to(DEVICE, non_blocking=True)\n\n        pred = model(x)\n        loss = dice_loss(pred, y)\n\n        opt.zero_grad()\n        loss.backward()\n        opt.step()\n\n        bar.set_postfix(loss=loss.detach().item())\n\n# -----------------------\n# POST-PROCESS\n# -----------------------\ndef clean_mask(prob):\n    m = prob > THRESH\n    m = binary_fill_holes(m)\n    lbl, n = label(m)\n    out = np.zeros_like(m)\n    for i in range(1, n + 1):\n        if (lbl == i).sum() >= MIN_COMPONENT:\n            out[lbl == i] = 1\n    return out.astype(np.uint8)\n\n# -----------------------\n# INFERENCE\n# -----------------------\nmodel.eval()\n\nwith torch.no_grad():\n    for vid in tqdm(os.listdir(TEST_IMG), desc=\"Inference\"):\n        vol = tiff.imread(os.path.join(TEST_IMG, vid))\n        z, y, x = vol.shape\n        prob = np.zeros(vol.shape, np.float32)\n        count = np.zeros(vol.shape, np.float32)\n\n        for zz in range(0, z-PATCH+1, PATCH):\n            for yy in range(0, y-PATCH+1, PATCH):\n                for xx in range(0, x-PATCH+1, PATCH):\n                    patch = vol[zz:zz+PATCH, yy:yy+PATCH, xx:xx+PATCH]\n                    patch = torch.from_numpy(patch).float().unsqueeze(0).unsqueeze(0).to(DEVICE)\n                    p = model(patch)[0,0].cpu().numpy()\n                    prob[zz:zz+PATCH, yy:yy+PATCH, xx:xx+PATCH] += p\n                    count[zz:zz+PATCH, yy:yy+PATCH, xx:xx+PATCH] += 1\n\n        prob /= np.maximum(count, 1)\n        mask = clean_mask(prob)\n        tiff.imwrite(os.path.join(PRED_DIR, vid), mask)\n\n# -----------------------\n# ZIP SUBMISSION (EXPLICIT PATH)\n# -----------------------\nwith zipfile.ZipFile(ZIP_PATH, \"w\", zipfile.ZIP_DEFLATED) as z:\n    for f in os.listdir(PRED_DIR):\n        z.write(os.path.join(PRED_DIR, f), arcname=f)\n\n# -----------------------\n# VERIFY OUTPUT\n# -----------------------\nprint(\"📁 /kaggle/working contents:\")\nprint(os.listdir(WORK_DIR))\n\nassert os.path.exists(ZIP_PATH), \"❌ submission.zip was NOT created\"\nprint(f\"✅ submission.zip created at: {ZIP_PATH}\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-03T21:36:58.890850Z","iopub.execute_input":"2026-02-03T21:36:58.891565Z","iopub.status.idle":"2026-02-03T22:06:17.181112Z","shell.execute_reply.started":"2026-02-03T21:36:58.891529Z","shell.execute_reply":"2026-02-03T22:06:17.180227Z"}},"outputs":[],"execution_count":null}]}