{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"__author__=\"Kushvinth\"","metadata":{"execution":{"iopub.status.busy":"2025-11-14T15:29:01.970769Z","iopub.execute_input":"2025-11-14T15:29:01.971025Z","iopub.status.idle":"2025-11-14T15:29:01.979836Z","shell.execute_reply.started":"2025-11-14T15:29:01.971001Z","shell.execute_reply":"2025-11-14T15:29:01.979126Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:01.980588Z","iopub.execute_input":"2025-11-14T15:29:01.980901Z","iopub.status.idle":"2025-11-14T15:29:41.124205Z","shell.execute_reply.started":"2025-11-14T15:29:01.980878Z","shell.execute_reply":"2025-11-14T15:29:41.123487Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.125721Z","iopub.execute_input":"2025-11-14T15:29:41.126309Z","iopub.status.idle":"2025-11-14T15:29:41.130464Z","shell.execute_reply.started":"2025-11-14T15:29:41.126288Z","shell.execute_reply":"2025-11-14T15:29:41.129597Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.131125Z","iopub.execute_input":"2025-11-14T15:29:41.131340Z","iopub.status.idle":"2025-11-14T15:29:41.145597Z","shell.execute_reply.started":"2025-11-14T15:29:41.131300Z","shell.execute_reply":"2025-11-14T15:29:41.144962Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# UNET 3D MODEL (Compact version)","metadata":{}},{"cell_type":"code","source":"class ResidualBlock3D(nn.Module):\n    \"A small residual block with GroupNorm and optional dropout.\"\n    def __init__(self, in_ch, out_ch, dropout=0.0):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_ch, out_ch, 3, padding=1, bias=False)\n        self.gn1 = nn.GroupNorm(num_groups=max(1, out_ch//8), num_channels=out_ch)\n        self.act1 = nn.LeakyReLU(0.01, inplace=True)\n        self.conv2 = nn.Conv3d(out_ch, out_ch, 3, padding=1, bias=False)\n        self.gn2 = nn.GroupNorm(num_groups=max(1, out_ch//8), num_channels=out_ch)\n        self.act2 = nn.LeakyReLU(0.01, inplace=True)\n        self.dropout = nn.Dropout3d(dropout) if dropout>0 else nn.Identity()\n        if in_ch != out_ch:\n            self.res_conv = nn.Conv3d(in_ch, out_ch, 1, bias=False)\n        else:\n            self.res_conv = nn.Identity()\n\n    def forward(self, x):\n        res = self.res_conv(x)\n        x = self.conv1(x)\n        x = self.gn1(x)\n        x = self.act1(x)\n        x = self.dropout(x)\n        x = self.conv2(x)\n        x = self.gn2(x)\n        x = x + res\n        x = self.act2(x)\n        return x\n\nclass UpSampleConv(nn.Module):\n    \"Upsample by trilinear interpolation then 1x1 conv to reduce channels (safer than transposed conv).\"\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.conv = nn.Conv3d(in_ch, out_ch, 1)\n    def forward(self, x, target_shape=None):\n        if target_shape is not None:\n            x = F.interpolate(x, size=target_shape, mode='trilinear', align_corners=False)\n        else:\n            x = F.interpolate(x, scale_factor=2, mode='trilinear', align_corners=False)\n        return self.conv(x)\n\nclass ResUNet3D(nn.Module):\n    \"Residual UNet3D: encoder with ResidualBlock3D, bottleneck, and decoder with upsampling + residual blocks.\"\n    def __init__(self, in_ch=1, out_ch=1, features=[16,32,64,128], dropout=0.0):\n        super().__init__()\n        self.encs = nn.ModuleList()\n        self.pools = nn.ModuleList()\n        ch = in_ch\n        for f in features:\n            self.encs.append(ResidualBlock3D(ch, f, dropout=dropout))\n            self.pools.append(nn.MaxPool3d(2))\n            ch = f\n\n        self.bottleneck = ResidualBlock3D(features[-1], features[-1]*2, dropout=dropout)\n\n        self.upconvs = nn.ModuleList()\n        self.decs = nn.ModuleList()\n        prev_ch = features[-1]*2\n        for f in reversed(features):\n            self.upconvs.append(UpSampleConv(prev_ch, f))\n            self.decs.append(ResidualBlock3D(f*2, f, dropout=dropout))\n            prev_ch = 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            # upsample to the skip's spatial size to avoid shape mismatches\n            x = up(x, target_shape=skip.shape[2:])\n            x = torch.cat([x, skip], dim=1)\n            x = dec(x)\n\n        return self.final(x)\n","metadata":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.146434Z","iopub.execute_input":"2025-11-14T15:29:41.146690Z","iopub.status.idle":"2025-11-14T15:29:41.161874Z","shell.execute_reply.started":"2025-11-14T15:29:41.146669Z","shell.execute_reply":"2025-11-14T15:29:41.161364Z"},"trusted":true},"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,"execution":{"iopub.status.busy":"2025-11-14T15:29:41.162694Z","iopub.execute_input":"2025-11-14T15:29:41.163203Z","iopub.status.idle":"2025-11-14T15:29:41.179228Z","shell.execute_reply.started":"2025-11-14T15:29:41.163175Z","shell.execute_reply":"2025-11-14T15:29:41.178613Z"}},"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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.179782Z","iopub.execute_input":"2025-11-14T15:29:41.179954Z","iopub.status.idle":"2025-11-14T15:29:41.193972Z","shell.execute_reply.started":"2025-11-14T15:29:41.179939Z","shell.execute_reply":"2025-11-14T15:29:41.193333Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.194596Z","iopub.execute_input":"2025-11-14T15:29:41.194793Z","iopub.status.idle":"2025-11-14T15:29:41.250315Z","shell.execute_reply.started":"2025-11-14T15:29:41.194770Z","shell.execute_reply":"2025-11-14T15:29:41.249832Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.252078Z","iopub.execute_input":"2025-11-14T15:29:41.252269Z","iopub.status.idle":"2025-11-14T15:29:41.259045Z","shell.execute_reply.started":"2025-11-14T15:29:41.252255Z","shell.execute_reply":"2025-11-14T15:29:41.258368Z"},"trusted":true},"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 = ResUNet3D().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":{"execution":{"iopub.status.busy":"2025-11-14T15:29:41.259788Z","iopub.execute_input":"2025-11-14T15:29:41.260106Z","execution_failed":"2025-11-14T16:49:38.080Z"},"trusted":true},"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":{"execution":{"execution_failed":"2025-11-14T16:49:38.080Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}