{"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":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"__author__=\"Kushvinth\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:36:38.877786Z","iopub.status.idle":"2025-11-14T07:36:38.878066Z","shell.execute_reply.started":"2025-11-14T07:36:38.877908Z","shell.execute_reply":"2025-11-14T07:36:38.877917Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Basic Module Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport random\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom PIL import Image\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\nimport albumentations as A\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:20.710704Z","iopub.execute_input":"2025-11-14T07:34:20.711387Z","iopub.status.idle":"2025-11-14T07:34:20.715761Z","shell.execute_reply.started":"2025-11-14T07:34:20.711363Z","shell.execute_reply":"2025-11-14T07:34:20.715043Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n# SAFE TIFF LOADER (Works in Kaggle)\n\n","metadata":{}},{"cell_type":"code","source":"\ndef load_3d_tiff(path):\n    img = Image.open(path)\n    slices = []\n    try:\n        for i in range(10000):\n            img.seek(i)\n            slices.append(np.array(img))\n    except EOFError:\n        pass\n    return np.stack(slices, axis=0)  # (Z,Y,X)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:21.732850Z","iopub.execute_input":"2025-11-14T07:34:21.733128Z","iopub.status.idle":"2025-11-14T07:34:21.737512Z","shell.execute_reply.started":"2025-11-14T07:34:21.733110Z","shell.execute_reply":"2025-11-14T07:34:21.736915Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"\nclass VesuviusDataset(Dataset):\n    def __init__(self, img_paths, mask_paths=None, patch_size=(64,128,128), transforms=None, mode='train'):\n        self.img_paths = img_paths\n        self.mask_paths = mask_paths\n        self.patch_size = patch_size\n        self.transforms = transforms\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def __getitem__(self, idx):\n        img = load_3d_tiff(self.img_paths[idx]).astype(np.float32)\n\n        if self.mask_paths is not None:\n            mask = load_3d_tiff(self.mask_paths[idx]).astype(np.uint8)\n        else:\n            mask = np.zeros_like(img, dtype=np.uint8)\n\n        Z,Y,X = img.shape\n        pz,py,px = self.patch_size\n\n        z0 = random.randint(0, max(0,Z-pz))\n        y0 = random.randint(0, max(0,Y-py))\n        x0 = random.randint(0, max(0,X-px))\n\n        img = img[z0:z0+pz, y0:y0+py, x0:x0+px]\n        mask = mask[z0:z0+pz, y0:y0+py, x0:x0+px]\n\n        if self.transforms:\n            xs, ys = [], []\n            for i in range(img.shape[0]):\n                aug = self.transforms(image=img[i].astype(np.uint8), mask=mask[i].astype(np.uint8))\n                xs.append(aug['image'])\n                ys.append(aug['mask'])\n            img = np.stack(xs)\n            mask = np.stack(ys)\n\n        img = img.astype(np.float32) / 255.0\n        mask = (mask == 1).astype(np.float32)\n\n        return (\n            torch.tensor(img[None], dtype=torch.float32),\n            torch.tensor(mask[None], dtype=torch.float32)\n        )\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:25.628631Z","iopub.execute_input":"2025-11-14T07:34:25.629328Z","iopub.status.idle":"2025-11-14T07:34:35.654769Z","shell.execute_reply.started":"2025-11-14T07:34:25.629302Z","shell.execute_reply":"2025-11-14T07:34:35.653785Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# UNET 3D MODEL (Compact version)","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),\n            nn.InstanceNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, 3, padding=1),\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 UNet3D(nn.Module):\n    def __init__(self, in_ch=1, out_ch=1, features=[16,32,64,128]):\n        super().__init__()\n        self.encs, self.pools = nn.ModuleList(), nn.ModuleList()\n\n        for f in features:\n            self.encs.append(Conv3dBlock(in_ch, f))\n            in_ch = f\n            self.pools.append(nn.MaxPool3d(2))\n\n        self.bottleneck = Conv3dBlock(features[-1], features[-1]*2)\n\n        self.upconvs, self.decs = nn.ModuleList(), nn.ModuleList()\n        for f in reversed(features):\n            self.upconvs.append(UpConv(features[-1]*2 if f==features[-1] else prev_f, f))\n            self.decs.append(Conv3dBlock(f*2, f))\n            prev_f = f\n\n        self.final = nn.Conv3d(features[0], out_ch, kernel_size=1)\n\n    def forward(self, x):\n        skips = []\n        for enc, pool in zip(self.encs, self.pools):\n            x = enc(x)\n            skips.append(x)\n            x = pool(x)\n\n        x = self.bottleneck(x)\n\n        for up, dec, skip in zip(self.upconvs, self.decs, reversed(skips)):\n            x = up(x)\n            if x.shape != skip.shape:\n                x = F.interpolate(x, size=skip.shape[2:], mode='trilinear', align_corners=False)\n            x = torch.cat([x, skip], dim=1)\n            x = dec(x)\n\n        return self.final(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:35.656347Z","iopub.execute_input":"2025-11-14T07:34:35.656637Z","iopub.status.idle":"2025-11-14T07:34:35.681228Z","shell.execute_reply.started":"2025-11-14T07:34:35.656612Z","shell.execute_reply":"2025-11-14T07:34:35.680481Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        num = 2 * (probs * targets).sum() + 1e-6\n        den = probs.sum() + targets.sum() + 1e-6\n        return 1 - num / den\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2.0, alpha=0.25):\n        super().__init__()\n        self.gamma, self.alpha = gamma, alpha\n    def forward(self, logits, targets):\n        p = torch.sigmoid(logits)\n        pt = p*targets + (1-p)*(1-targets)\n        w = self.alpha*targets + (1-self.alpha)*(1-targets)\n        return (-w*(1-pt)**self.gamma * torch.log(pt+1e-8)).mean()\n\nclass CombinedLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.dice = DiceLoss()\n        self.focal = FocalLoss()\n    def forward(self, logits, targets):\n        return self.dice(logits, targets) + 0.75*self.focal(logits, targets)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Sliding Window Inference","metadata":{}},{"cell_type":"code","source":"def sliding_window(volume, model, patch=(64,128,128), stride=(32,64,64)):\n    model.eval()\n    Z,Y,X = volume.shape\n    pz,py,px = patch\n    sz,sy,sx = stride\n\n    out = np.zeros((Z,Y,X), np.float32)\n    cnt = np.zeros_like(out)\n\n    with torch.no_grad():\n        for z in range(0, Z-pz+1, sz):\n            for y in range(0, Y-py+1, sy):\n                for x in range(0, X-px+1, sx):\n                    patch_data = volume[z:z+pz, y:y+py, x:x+px]\n                    patch_data = torch.tensor(patch_data[None,None]/255.0, dtype=torch.float32).cuda()\n                    pred = torch.sigmoid(model(patch_data))[0,0].cpu().numpy()\n                    out[z:z+pz, y:y+py, x:x+px] += pred\n                    cnt[z:z+pz, y:y+py, x:x+px] += 1\n\n    return out / np.maximum(cnt,1)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:41.058433Z","iopub.execute_input":"2025-11-14T07:34:41.058694Z","iopub.status.idle":"2025-11-14T07:34:41.065033Z","shell.execute_reply.started":"2025-11-14T07:34:41.058676Z","shell.execute_reply":"2025-11-14T07:34:41.064278Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup: Paths / Splits","metadata":{}},{"cell_type":"code","source":"ROOT = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\n\ntrain_imgs = sorted((ROOT/\"train_images\").glob(\"*.tif\"))\ntrain_masks = sorted((ROOT/\"train_labels\").glob(\"*.tif\"))\n\nsplit = int(0.9 * len(train_imgs))\ntrain_list = train_imgs[:split]\nval_list = train_imgs[split:]\ntrain_mask_list = train_masks[:split]\nval_mask_list = train_masks[split:]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:43.833682Z","iopub.execute_input":"2025-11-14T07:34:43.833990Z","iopub.status.idle":"2025-11-14T07:34:43.848307Z","shell.execute_reply.started":"2025-11-14T07:34:43.833970Z","shell.execute_reply":"2025-11-14T07:34:43.847558Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DataLoaders","metadata":{}},{"cell_type":"code","source":"transforms = A.Compose([\n    A.RandomBrightnessContrast(p=0.5),\n    A.HorizontalFlip(p=0.5),\n])\n\ntrain_ds = VesuviusDataset(train_list, train_mask_list, transforms=transforms)\nval_ds = VesuviusDataset(val_list, val_mask_list, transforms=None)\n\ntrain_dl = DataLoader(train_ds, batch_size=2, shuffle=True, num_workers=2)\nval_dl = DataLoader(val_ds, batch_size=1, shuffle=False, num_workers=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:34:45.134161Z","iopub.execute_input":"2025-11-14T07:34:45.134877Z","iopub.status.idle":"2025-11-14T07:34:45.142347Z","shell.execute_reply.started":"2025-11-14T07:34:45.134846Z","shell.execute_reply":"2025-11-14T07:34:45.141839Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"from torch.amp import autocast        # NEW API\nfrom torch.cuda.amp import GradScaler # Still valid\n\nmodel = UNet3D().cuda()\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\n\n# Correct initialization\nscaler = GradScaler()\ncriterion = CombinedLoss()\n\nEPOCHS = 3\n\ndef validate(model, loader):\n    model.eval()\n    dices = []\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.cuda(), y.cuda()\n            pred = torch.sigmoid(model(x))\n            pred = (pred > 0.5).float()\n            dice = (2*(pred*y).sum() + 1e-6) / (pred.sum()+y.sum()+1e-6)\n            dices.append(dice.item())\n    return np.mean(dices)\n\nbest_dice = 0\n\nfor epoch in range(EPOCHS):\n    model.train()\n    total_loss = 0.0\n\n    for x, y in tqdm(train_dl):\n        x, y = x.cuda(), y.cuda()\n        optimizer.zero_grad()\n\n        # NEW CORRECT AUTOTCAST\n        with autocast(\"cuda\"):\n            out = model(x)\n            loss = criterion(out, y)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item()\n\n    val_dice = validate(model, val_dl)\n    print(f\"Epoch {epoch} | TrainLoss={total_loss/len(train_dl):.4f} | ValDice={val_dice:.4f}\")\n\n    if val_dice > best_dice:\n        best_dice = val_dice\n        torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n        print(\"Saved best model!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T07:36:45.398024Z","iopub.execute_input":"2025-11-14T07:36:45.398816Z","iopub.status.idle":"2025-11-14T08:19:30.069360Z","shell.execute_reply.started":"2025-11-14T07:36:45.398775Z","shell.execute_reply":"2025-11-14T08:19:30.068363Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference on Test Set","metadata":{}},{"cell_type":"code","source":"def rle_encode(mask):\n    pixels = mask.flatten(order=\"F\")\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    return \" \".join(str(x) for x in runs)\n\ntest_df = pd.read_csv(ROOT/\"test.csv\")\ntest_images = ROOT/\"test_images\"\n\nmodel.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\"))\nmodel.eval()\n\nimport os\nimport zipfile\nimport tifffile as tiff\n\nos.makedirs(\"/kaggle/working/pred_masks\", exist_ok=True)\ntest_df = pd.read_csv(ROOT/\"test.csv\")\ntest_imgs = ROOT/\"test_images\"\n\nmodel.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\"))\nmodel.eval()\n\nfor vid in tqdm(test_df.id.values):\n\n    # Load test volume (safe loader)\n    vol = load_3d_tiff(test_imgs/f\"{vid}.tif\").astype(np.float32)\n    vol_norm = (vol - vol.mean()) / (vol.std() + 1e-8)\n\n    # Sliding window prediction\n    prob = sliding_window(vol_norm, model)\n\n    # Binarize → uint8 mask (0/1)\n    mask = (prob > 0.5).astype(\"uint8\")\n\n    # Save mask with SAME SHAPE + TYPE\n    out_path = f\"/kaggle/working/pred_masks/{vid}.tif\"\n    tiff.imwrite(out_path, mask, dtype=\"uint8\")\n\nprint(\"All TIFF masks written successfully.\")\nzip_path = \"/kaggle/working/submission.zip\"\n\nwith zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as z:\n    for vid in test_df.id.values:\n        file_path = f\"/kaggle/working/pred_masks/{vid}.tif\"\n        z.write(file_path, arcname=f\"{vid}.tif\")\n\nprint(\"Submission created:\", zip_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-14T08:19:47.127396Z","iopub.execute_input":"2025-11-14T08:19:47.127618Z","iopub.status.idle":"2025-11-14T08:20:03.240679Z","shell.execute_reply.started":"2025-11-14T08:19:47.127600Z","shell.execute_reply":"2025-11-14T08:20:03.239855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}