{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# SLART V10: SAFE UPGRADE #2 (THRESHOLD 0.5 → 0.45)\n# Expected improvement: +0.03 to +0.08 score\n# ============================================================\n\nimport os\nimport random\nfrom pathlib import Path\nfrom io import BytesIO\nimport zipfile\n\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image, ImageSequence\n\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 scipy import ndimage\n\n# ------------------------------------------------------------\n#  Config - SAME AS BEFORE except inference threshold change\n# ------------------------------------------------------------\nROOT = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nTRAIN_IMG_DIR = ROOT / \"train_images\"\nTRAIN_LABEL_DIR = ROOT / \"train_labels\"\nTEST_IMG_DIR = ROOT / \"test_images\"\n\nTRAIN_CSV = ROOT / \"train.csv\"\nTEST_CSV = ROOT / \"test.csv\"\n\nOUT_ZIP = Path(\"/kaggle/working/submission.zip\")\n\nMAX_TRAIN_VOLUMES = 15\nPATCH_SIZE = 128\nTRAIN_SAMPLES = 10000\nVAL_SAMPLES = 1200\nBATCH_SIZE = 4\nEPOCHS = 6\nLR = 1e-3\n\nMIN_COMPONENT_VOXELS = 600\nCLOSING_ITERS = 1\n\n# SAFE UPGRADE #2 ----\nTHRESH = 0.45\n# ---------------------\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n\n# ============================================================\n#  Utilities: load/save volumes\n# ============================================================\n\ndef load_stack(path: Path) -> np.ndarray:\n    with Image.open(path) as tif:\n        frames = [np.array(frame) for frame in ImageSequence.Iterator(tif)]\n    return np.stack(frames).astype(np.float32)\n\ndef normalize_volume(vol: np.ndarray) -> np.ndarray:\n    v = vol.astype(np.float32)\n    return (v - v.mean()) / (v.std() + 1e-6)\n\ndef write_stack_to_zip(array3d: np.ndarray, zip_handle, name: str):\n    pages = [Image.fromarray(s.astype(np.uint8)) for s in array3d]\n    buffer = BytesIO()\n    pages[0].save(buffer, format=\"TIFF\", save_all=True, append_images=pages[1:])\n    zip_handle.writestr(name, buffer.getvalue())\n\n# ============================================================\n#  Postprocessing\n# ============================================================\n\ndef clean_mask(mask: np.ndarray) -> np.ndarray:\n    cc_structure = ndimage.generate_binary_structure(3, 1)\n    labeled, num = ndimage.label(mask, structure=cc_structure)\n\n    if num == 0:\n        return mask.astype(np.uint8)\n\n    component_sizes = ndimage.sum(mask, labeled, index=np.arange(1, num + 1))\n    keep_labels = np.where(component_sizes >= MIN_COMPONENT_VOXELS)[0] + 1\n    cleaned = np.isin(labeled, keep_labels).astype(np.uint8)\n\n    if CLOSING_ITERS > 0:\n        closing_structure = np.zeros((3, 3, 3), dtype=np.uint8)\n        closing_structure[1, :, :] = 1\n        cleaned = ndimage.binary_closing(\n            cleaned, structure=closing_structure, iterations=CLOSING_ITERS\n        ).astype(np.uint8)\n\n    return cleaned\n\n# ============================================================\n#  UNet (same as before)\n# ============================================================\n\nclass DoubleConv(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):\n        return self.net(x)\n\nclass UNet2D(nn.Module):\n    def __init__(self, in_ch=3, out_ch=1, base_ch=32):\n        super().__init__()\n\n        self.down1 = DoubleConv(in_ch, base_ch)\n        self.pool1 = nn.MaxPool2d(2)\n\n        self.down2 = DoubleConv(base_ch, base_ch * 2)\n        self.pool2 = nn.MaxPool2d(2)\n\n        self.down3 = DoubleConv(base_ch * 2, base_ch * 4)\n        self.pool3 = nn.MaxPool2d(2)\n\n        self.bottom = DoubleConv(base_ch * 4, base_ch * 8)\n\n        self.up3 = nn.ConvTranspose2d(base_ch * 8, base_ch * 4, 2, 2)\n        self.dec3 = DoubleConv(base_ch * 8, base_ch * 4)\n\n        self.up2 = nn.ConvTranspose2d(base_ch * 4, base_ch * 2, 2, 2)\n        self.dec2 = DoubleConv(base_ch * 4, base_ch * 2)\n\n        self.up1 = nn.ConvTranspose2d(base_ch * 2, base_ch, 2, 2)\n        self.dec1 = DoubleConv(base_ch * 2, base_ch)\n\n        self.out_conv = nn.Conv2d(base_ch, out_ch, 1)\n\n    def forward(self, x):\n        x1 = self.down1(x)\n        x2 = self.down2(self.pool1(x1))\n        x3 = self.down3(self.pool2(x2))\n        x4 = self.bottom(self.pool3(x3))\n\n        x = self.up3(x4)\n        x = torch.cat([x, x3], dim=1)\n        x = self.dec3(x)\n\n        x = self.up2(x)\n        x = torch.cat([x, x2], dim=1)\n        x = self.dec2(x)\n\n        x = self.up1(x)\n        x = torch.cat([x, x1], dim=1)\n        x = self.dec1(x)\n\n        return self.out_conv(x)\n\n# ============================================================\n# Dataset (same)\n# ============================================================\n\nclass VesuviusSliceDataset(Dataset):\n    def __init__(self, volumes, labels, n_samples, patch_size=128):\n        self.volumes = volumes\n        self.labels = labels\n        self.n_samples = n_samples\n        self.patch_size = patch_size\n\n    def __len__(self):\n        return self.n_samples\n\n    def __getitem__(self, idx):\n        vidx = random.randint(0, len(self.volumes) - 1)\n        vol = self.volumes[vidx]\n        lab = self.labels[vidx]\n\n        Z, H, W = vol.shape\n\n        z = random.randint(0, Z - 1)\n        z_prev = max(z - 1, 0)\n        z_next = min(z + 1, Z - 1)\n\n        ps = self.patch_size\n        y0 = 0 if H <= ps else random.randint(0, H - ps)\n        x0 = 0 if W <= ps else random.randint(0, W - ps)\n\n        ch0 = vol[z_prev, y0:y0+ps, x0:x0+ps]\n        ch1 = vol[z,      y0:y0+ps, x0:x0+ps]\n        ch2 = vol[z_next, y0:y0+ps, x0:x0+ps]\n\n        x_np = np.stack([ch0, ch1, ch2], axis=0)\n\n        lb_slice = lab[z, y0:y0+ps, x0:x0+ps]\n        y_fg = (lb_slice == 1).astype(np.float32)\n        y_ignore = (lb_slice == 2).astype(np.float32)\n\n        x = torch.from_numpy(x_np).float()\n        y = torch.from_numpy(y_fg).float().unsqueeze(0)\n        ignore = torch.from_numpy(y_ignore).float().unsqueeze(0)\n\n        return x, y, ignore\n\n# ============================================================\n# Loss (same)\n# ============================================================\n\ndef loss_fn(logits, targets, ignore_mask):\n    probs = torch.sigmoid(logits)\n    valid_mask = 1.0 - ignore_mask\n\n    bce = F.binary_cross_entropy(\n        probs * valid_mask + 0.5 * (1 - valid_mask),\n        targets,\n        reduction=\"sum\"\n    ) / (valid_mask.sum() + 1e-6)\n\n    probs_flat = (probs * valid_mask).view(logits.size(0), -1)\n    targets_flat = (targets * valid_mask).view(logits.size(0), -1)\n\n    inter = (probs_flat * targets_flat).sum(dim=1)\n    denom = probs_flat.sum(dim=1) + targets_flat.sum(dim=1) + 1e-6\n\n    dice = 1 - (2 * inter / denom)\n    return bce + dice.mean(), bce.detach(), (1 - dice).detach()\n\n# ============================================================\n# Load training data (same)\n# ============================================================\n\ntrain_df = pd.read_csv(TRAIN_CSV)\ntrain_ids = train_df[\"id\"].values[:MAX_TRAIN_VOLUMES]\n\ntrain_volumes = []\ntrain_labels = []\n\nfor vid in tqdm(train_ids, desc=\"Loading train volumes\"):\n    fname = f\"{vid}.tif\"\n    img_path = TRAIN_IMG_DIR / fname\n    lbl_path = TRAIN_LABEL_DIR / fname\n\n    if not img_path.exists() or not lbl_path.exists():\n        continue\n\n    vol = load_stack(img_path)\n    lbl = load_stack(lbl_path)\n\n    if vol.shape != lbl.shape:\n        continue\n\n    train_volumes.append(normalize_volume(vol))\n    train_labels.append(lbl)\n\nsplit_idx = max(1, len(train_volumes) - 3)\n\ntrain_dataset = VesuviusSliceDataset(\n    train_volumes[:split_idx], train_labels[:split_idx],\n    n_samples=TRAIN_SAMPLES, patch_size=PATCH_SIZE\n)\nval_dataset = VesuviusSliceDataset(\n    train_volumes[split_idx:], train_labels[split_idx:],\n    n_samples=VAL_SAMPLES, patch_size=PATCH_SIZE\n)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True)\n\n# ============================================================\n# Train (same)\n# ============================================================\n\nmodel = UNet2D().to(DEVICE)\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)\n\nbest_val_dice = 0.0\nbest_model_path = \"/kaggle/working/best_unet2d_2p5d.pth\"\n\nfor epoch in range(1, EPOCHS + 1):\n    model.train()\n    train_loss_sum = 0\n    train_dice_sum = 0\n    train_batches = 0\n\n    for x, y, ignore in tqdm(train_loader, desc=f\"Epoch {epoch} [train]\"):\n        x, y, ignore = x.to(DEVICE), y.to(DEVICE), ignore.to(DEVICE)\n\n        optimizer.zero_grad()\n        logits = model(x)\n        loss, bce, dice = loss_fn(logits, y, ignore)\n        loss.backward()\n        optimizer.step()\n\n        train_loss_sum += loss.item()\n        train_dice_sum += dice.mean().item()\n        train_batches += 1\n\n    val_loss_sum = 0\n    val_dice_sum = 0\n    val_batches = 0\n    model.eval()\n\n    with torch.no_grad():\n        for x, y, ignore in tqdm(val_loader, desc=f\"Epoch {epoch} [val]\"):\n            x, y, ignore = x.to(DEVICE), y.to(DEVICE), ignore.to(DEVICE)\n            logits = model(x)\n            loss, bce, dice = loss_fn(logits, y, ignore)\n\n            val_loss_sum += loss.item()\n            val_dice_sum += dice.mean().item()\n            val_batches += 1\n\n    train_loss = train_loss_sum / train_batches\n    train_dice = train_dice_sum / train_batches\n    val_loss = val_loss_sum / val_batches\n    val_dice = val_dice_sum / val_batches\n\n    print(f\"Epoch {epoch}: train_loss={train_loss:.4f}, val_loss={val_loss:.4f}, val_dice={val_dice:.4f}\")\n\n    if val_dice > best_val_dice:\n        best_val_dice = val_dice\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"New best val Dice: {best_val_dice:.4f}\")\n\nif os.path.exists(best_model_path):\n    model.load_state_dict(torch.load(best_model_path, map_location=DEVICE))\n\nmodel.eval()\n\n# ============================================================\n# Inference (threshold upgrade applied)\n# ============================================================\n\ntest_df = pd.read_csv(TEST_CSV)\n\nwith zipfile.ZipFile(OUT_ZIP, \"w\", compression=zipfile.ZIP_DEFLATED) as zf:\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Test inference\"):\n        vid = row[\"id\"]\n        fname = f\"{vid}.tif\"\n        img_path = TEST_IMG_DIR / fname\n\n        if not img_path.exists():\n            continue\n\n        vol = normalize_volume(load_stack(img_path))\n        Z, H, W = vol.shape\n        pred_mask = np.zeros((Z, H, W), dtype=np.uint8)\n\n        for z_idx in range(Z):\n            z_prev = max(z_idx - 1, 0)\n            z_next = min(z_idx + 1, Z - 1)\n\n            x_np = np.stack([vol[z_prev], vol[z_idx], vol[z_next]], axis=0)\n            x = torch.from_numpy(x_np).float().unsqueeze(0).to(DEVICE)\n\n            with torch.no_grad():\n                logits = model(x)\n                probs = torch.sigmoid(logits)[0, 0].cpu().numpy()\n\n                # --------------------------------------\n                # SAFE UPGRADE #2 (threshold 0.45)\n                pred = (probs > THRESH).astype(np.uint8)\n                # --------------------------------------\n\n            pred_mask[z_idx] = pred\n\n        cleaned = clean_mask(pred_mask)\n        write_stack_to_zip(cleaned, zf, fname)\n\nprint(\"Done! submission.zip created successfully.\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}