{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14351778,"sourceType":"datasetVersion","datasetId":9163946}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\n\nDATASET_NAME = \"imagecodecs\"\n\n!{sys.executable} -m pip install --no-index --find-links /kaggle/input/{DATASET_NAME} imagecodecs -q","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:01:47.002985Z","iopub.execute_input":"2026-01-07T12:01:47.003253Z","iopub.status.idle":"2026-01-07T12:01:51.999417Z","shell.execute_reply.started":"2026-01-07T12:01:47.003220Z","shell.execute_reply":"2026-01-07T12:01:51.998467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nfrom pathlib import Path\nimport zipfile\nfrom io import BytesIO\nimport tifffile as tiff\nfrom tqdm.auto import tqdm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# --- 设置随机种子 ---\ndef set_seed(seed=42):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed()\n\n# --- 路径配置 ---\nCOMP_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nTRAIN_IMG_DIR = COMP_DIR / \"train_images\"\n\nif TRAIN_IMG_DIR.exists():\n    all_tifs = list(TRAIN_IMG_DIR.glob(\"*.tif\"))\n    valid_ids = [f.stem for f in all_tifs]\n    print(f\"Found {len(valid_ids)} training volumes. Example IDs: {valid_ids[:5]}\")\n    \n    # --- 关键修复：确保用于验证的 chunk 确实有 label ---\n    def has_label(cid):\n        return (COMP_DIR / f\"train_labels/{cid}.tif\").exists()\n    \n    # 过滤出有 label 的 ID\n    valid_ids_with_label = [cid for cid in valid_ids if has_label(cid)]\n    print(f\"Found {len(valid_ids_with_label)} volumes with labels.\")\n    \n    if len(valid_ids_with_label) < 3:\n        raise ValueError(\"Need at least 3 labeled volumes for train/val split!\")\n        \n    train_ids = valid_ids_with_label[:2]\n    valid_ids = valid_ids_with_label[2:3]  # 只取一个做验证\nelse:\n    print(\"Warning: Train directory not found.\")\n    train_ids = []\n    valid_ids = []\n\nCONFIG = {\n    \"train_chunks\": train_ids, \n    \"valid_chunks\": valid_ids,\n    \"batch_size\": 4,\n    \"lr\": 1e-3,\n    \"epochs\": 10,\n    \"device\": \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    \"slice_depth\": 5,\n    \"comp_dir\": COMP_DIR,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:01:52.001366Z","iopub.execute_input":"2026-01-07T12:01:52.001733Z","iopub.status.idle":"2026-01-07T12:02:31.338140Z","shell.execute_reply.started":"2026-01-07T12:01:52.001701Z","shell.execute_reply":"2026-01-07T12:02:31.337349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 2. MODEL: Standard U-Net (Lighter & More Stable) ---\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.conv = 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.conv(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_channels=5, out_channels=1):\n        super().__init__()\n        self.in_channels = in_channels\n        \n        self.down1 = DoubleConv(in_channels, 64)\n        self.down2 = DoubleConv(64, 128)\n        self.down3 = DoubleConv(128, 256)\n        self.down4 = DoubleConv(256, 512)\n        \n        self.pool = nn.MaxPool2d(2)\n        self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        \n        self.up3 = DoubleConv(512+256, 256)\n        self.up2 = DoubleConv(256+128, 128)\n        self.up1 = DoubleConv(128+64, 64)\n        \n        self.final = nn.Conv2d(64, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        # Contracting path\n        x1 = self.down1(x)      # 64\n        x2 = self.down2(self.pool(x1))  # 128\n        x3 = self.down3(self.pool(x2))  # 256\n        x4 = self.down4(self.pool(x3))  # 512\n        \n        # Expanding path\n        u3 = self.up(x4)\n        u3 = torch.cat([u3, x3], dim=1)\n        u3 = self.up3(u3)\n        \n        u2 = self.up(u3)\n        u2 = torch.cat([u2, x2], dim=1)\n        u2 = self.up2(u2)\n        \n        u1 = self.up(u2)\n        u1 = torch.cat([u1, x1], dim=1)\n        u1 = self.up1(u1)\n        \n        return self.final(u1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:02:31.339128Z","iopub.execute_input":"2026-01-07T12:02:31.339550Z","iopub.status.idle":"2026-01-07T12:02:31.349877Z","shell.execute_reply.started":"2026-01-07T12:02:31.339526Z","shell.execute_reply":"2026-01-07T12:02:31.349235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 3. DATASET ---\nclass VesuviusSliceDataset(Dataset):\n    def __init__(self, chunk_ids, mode='train'):\n        self.mode = mode\n        self.images = []\n        self.labels = []\n        \n        if self.mode == 'train':\n            self.transform = A.Compose([\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.RandomRotate90(p=0.5),\n                A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.3),\n                ToTensorV2()\n            ])\n        else:\n            self.transform = ToTensorV2()\n\n        if not chunk_ids:\n            return\n\n        print(f\"Loading {mode} chunks: {chunk_ids}...\")\n        total_slices = 0\n        for cid in chunk_ids:\n            img_path = CONFIG[\"comp_dir\"] / f\"train_images/{cid}.tif\"\n            lbl_path = CONFIG[\"comp_dir\"] / f\"train_labels/{cid}.tif\"\n            \n            if not lbl_path.exists():\n                print(f\"Skipping {cid}: label missing\")\n                continue\n                \n            vol = tiff.imread(img_path).astype(np.float32) / 65535.0\n            # --- 关键修复行 ---\n            mask = tiff.imread(lbl_path).astype(np.uint8)\n            mask = (mask > 0).astype(np.uint8) # ← 将 255 → 1\n            \n            # ↓ stride=1 for val to avoid missing positives\n            stride = 1 if mode == 'val' else 2\n            \n            count = 0\n            # --- 修正范围：确保能取到最后一个完整的 slice ---\n            max_start_z = vol.shape[0] - CONFIG[\"slice_depth\"]\n            for z in range(0, max_start_z + 1, stride):\n                mid_slice = z + CONFIG[\"slice_depth\"] // 2\n                if np.any(mask[mid_slice] == 1):\n                    self.images.append(vol[z : z + CONFIG[\"slice_depth\"]])\n                    self.labels.append(mask[mid_slice])\n                    count += 1\n            total_slices += count\n            print(f\"  Chunk {cid}: found {count} positive slices.\")\n        print(f\" Loaded {len(self.images)} slices ({total_slices} from {len(chunk_ids)} volumes)\")\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        img = np.transpose(self.images[idx], (1, 2, 0))\n        lbl = self.labels[idx]\n        augmented = self.transform(image=img, mask=lbl)\n        return augmented['image'], augmented['mask'].long()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:02:31.350720Z","iopub.execute_input":"2026-01-07T12:02:31.350958Z","iopub.status.idle":"2026-01-07T12:02:31.380931Z","shell.execute_reply.started":"2026-01-07T12:02:31.350925Z","shell.execute_reply":"2026-01-07T12:02:31.380277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 4. LOSS & TRAINING ---\nimport torch.nn.functional as F\n\nclass DiceBCELoss(nn.Module):\n    def __init__(self, weight_bce=0.3, smooth=1e-6):\n        super().__init__()\n        self.weight_bce = weight_bce\n        self.smooth = smooth\n\n    def forward(self, pred, target):\n        bce = F.binary_cross_entropy_with_logits(pred, target, reduction='mean')\n        pred_s = torch.sigmoid(pred)\n        intersection = (pred_s * target).sum()\n        dice_loss = 1 - (2. * intersection + self.smooth) / (\n            pred_s.sum() + target.sum() + self.smooth\n        )\n        return self.weight_bce * bce + (1 - self.weight_bce) * dice_loss\n\ndef train():\n    model = UNet(in_channels=CONFIG[\"slice_depth\"]).to(CONFIG[\"device\"])\n    \n    if not CONFIG[\"train_chunks\"]:\n        print(\"No training data!\")\n        return model\n\n    # Datasets\n    train_ds = VesuviusSliceDataset(CONFIG[\"train_chunks\"], mode='train')\n    val_ds = VesuviusSliceDataset(CONFIG[\"valid_chunks\"], mode='val')\n    \n    if len(train_ds) == 0:\n        print(\"Empty train dataset!\")\n        return model\n        \n    train_dl = DataLoader(train_ds, batch_size=CONFIG[\"batch_size\"], shuffle=True, num_workers=0)\n    val_dl = DataLoader(val_ds, batch_size=CONFIG[\"batch_size\"], shuffle=False, num_workers=0) if len(val_ds) > 0 else None\n\n    optimizer = optim.AdamW(model.parameters(), lr=CONFIG[\"lr\"])\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CONFIG[\"epochs\"])\n    criterion = DiceBCELoss(weight_bce=0.3)\n\n    best_dice = 0.0\n    print(\"Starting Training...\")\n    for epoch in range(CONFIG[\"epochs\"]):\n        model.train()\n        train_loss = 0.0\n        pbar = tqdm(train_dl, desc=f\"Epoch {epoch+1}/{CONFIG['epochs']}\")\n        for imgs, masks in pbar:\n            imgs = imgs.to(CONFIG[\"device\"])\n            masks = (masks == 1).float().to(CONFIG[\"device\"])\n            \n            optimizer.zero_grad()\n            preds = model(imgs).squeeze(1)\n            loss = criterion(preds, masks)\n            loss.backward()\n            \n            # --- 新增：梯度范数检查 ---\n            total_norm = 0\n            for p in model.parameters():\n                if p.grad is not None:\n                    param_norm = p.grad.data.norm(2)\n                    total_norm += param_norm.item() ** 2\n            total_norm = total_norm ** (1. / 2)\n            # 如果 total_norm 接近 0，说明梯度消失；如果极大，说明爆炸。\n            # 正常值应在 0.1 - 10 之间。\n            \n            optimizer.step()\n            train_loss += loss.item()\n            pbar.set_postfix(loss=loss.item(), grad_norm=f\"{total_norm:.2f}\")\n\n        # Validation\n        if val_dl is not None:\n            model.eval()\n            val_loss = 0.0\n            soft_dice = 0.0\n            with torch.no_grad():\n                for val_imgs, val_masks in val_dl:\n                    val_imgs = val_imgs.to(CONFIG[\"device\"])\n                    val_masks = (val_masks == 1).float().to(CONFIG[\"device\"])\n                    val_preds = model(val_imgs).squeeze(1)\n                    val_loss += criterion(val_preds, val_masks).item()\n                    \n                    probs = torch.sigmoid(val_preds)\n                    inter = (probs * val_masks).sum()\n                    dice = (2 * inter + 1e-6) / (probs.sum() + val_masks.sum() + 1e-6)\n                    soft_dice += dice.item()\n            \n            avg_val_loss = val_loss / len(val_dl)\n            avg_soft_dice = soft_dice / len(val_dl)\n            print(f\"\\nVal Loss: {avg_val_loss:.4f} | Soft Dice: {avg_soft_dice:.4f}\\n\")\n            \n            if avg_soft_dice > best_dice:\n                best_dice = avg_soft_dice\n                torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n                print(f\" Saved best model (Soft Dice: {best_dice:.4f})\")\n        \n        scheduler.step()\n    \n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:02:31.381873Z","iopub.execute_input":"2026-01-07T12:02:31.382121Z","iopub.status.idle":"2026-01-07T12:02:31.403736Z","shell.execute_reply.started":"2026-01-07T12:02:31.382082Z","shell.execute_reply":"2026-01-07T12:02:31.403162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 5. INFERENCE (CORRECTED FOR SURFACE DETECTION) ---\n\ndef inference(model, threshold=0.5):\n    \"\"\"\n    Performs inference for the Vesuvius Surface Detection challenge.\n    Generates a submission.zip containing .tif files, NOT a CSV with RLE.\n    \"\"\"\n    # --- Use the correct path and load IDs directly from test_images folder ---\n    TEST_IMG_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection/test_images\")\n    chunks = sorted([f.stem for f in TEST_IMG_DIR.glob(\"*.tif\")])\n    print(f\"Running Inference on {len(chunks)} volumes: {chunks}\")\n    \n    model.eval()\n    submission_path = Path(\"/kaggle/working/submission.zip\")\n    \n    with zipfile.ZipFile(submission_path, \"w\", compression=zipfile.ZIP_DEFLATED) as zf:\n        for cid in chunks:\n            img_path = TEST_IMG_DIR / f\"{cid}.tif\"\n            if not img_path.exists(): \n                print(f\"Warning: File missing: {img_path}. Creating dummy prediction.\")\n                _create_dummy_tiff(zf, cid)\n                continue\n            \n            print(f\"Processing {cid}...\")\n            try:\n                vol = tiff.imread(img_path).astype(np.float32) / 65535.0\n                D, H, W = vol.shape\n                pred_vol = np.zeros((D, H, W), dtype=np.uint8)\n                \n                half_depth = CONFIG[\"slice_depth\"] // 2\n                \n                # Handle each slice in the volume\n                for z in range(D):\n                    # Determine the slice window with padding for boundaries\n                    z_start = max(0, z - half_depth)\n                    z_end = min(D, z + half_depth + 1)\n                    input_slice = vol[z_start:z_end]\n                    \n                    # Pad if necessary to match expected input depth\n                    if input_slice.shape[0] < CONFIG[\"slice_depth\"]:\n                        pad_before = half_depth - (z - z_start)\n                        pad_after = half_depth - (z_end - 1 - z)\n                        input_slice = np.pad(input_slice, ((pad_before, pad_after), (0, 0), (0, 0)), mode='edge')\n                    \n                    tensor = torch.from_numpy(input_slice).unsqueeze(0).float().to(CONFIG[\"device\"])\n                    with torch.no_grad():\n                        out = model(tensor)\n                        out = torch.sigmoid(out)\n                        # Resize prediction back to original image dimensions\n                        out_resized = torch.nn.functional.interpolate(\n                            out, size=(H, W), mode='bilinear', align_corners=False\n                        )\n                        mask = (out_resized.squeeze().cpu().numpy() > threshold).astype(np.uint8)\n                    \n                    pred_vol[z] = mask\n\n                # Save the 3D prediction volume as a multi-page TIFF inside the zip\n                _save_volume_to_zip(zf, cid, pred_vol)\n                \n            except Exception as e:\n                print(f\"Error processing {cid}: {e}. Creating dummy prediction.\")\n                _create_dummy_tiff(zf, cid)\n\n    print(\"Done! Output saved to submission.zip\")\n\ndef _save_volume_to_zip(zip_file, chunk_id, volume_3d):\n    \"\"\"Helper function to save a 3D numpy array as a multi-page TIFF in a zip file.\"\"\"\n    img_buffer = BytesIO()\n    pages = [Image.fromarray((s * 255).astype(np.uint8), mode='L') for s in volume_3d]\n    pages[0].save(\n        img_buffer, \n        format=\"TIFF\", \n        save_all=True, \n        append_images=pages[1:],\n        compression=\"tiff_deflate\"\n    )\n    zip_file.writestr(f\"{chunk_id}.tif\", img_buffer.getvalue())\n\ndef _create_dummy_tiff(zip_file, chunk_id, shape=(10, 256, 256)):\n    \"\"\"Helper function to create a dummy all-black TIFF for error handling.\"\"\"\n    dummy_vol = np.zeros(shape, dtype=np.uint8)\n    _save_volume_to_zip(zip_file, chunk_id, dummy_vol)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:02:31.404504Z","iopub.execute_input":"2026-01-07T12:02:31.404684Z","iopub.status.idle":"2026-01-07T12:02:31.422625Z","shell.execute_reply.started":"2026-01-07T12:02:31.404665Z","shell.execute_reply":"2026-01-07T12:02:31.421972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 6. EXECUTE WITH THRESHOLD SEARCH ---\ndef find_best_threshold(model, val_dataset, thresholds=np.arange(0.05, 0.51, 0.05)):\n    \"\"\"\n    在验证集上搜索使 Hard Dice 最大的阈值。\n    \"\"\"\n    model.eval()\n    device = next(model.parameters()).device\n    \n    all_probs = []\n    all_masks = []\n    \n    # 收集所有验证集的预测概率和真实标签\n    for i in tqdm(range(len(val_dataset)), desc=\"Collecting Val Predictions\"):\n        img, mask = val_dataset[i]\n        img = img.unsqueeze(0).to(device)\n        with torch.no_grad():\n            pred = model(img).squeeze(0)\n            prob = torch.sigmoid(pred).cpu().numpy()\n        all_probs.append(prob)\n        all_masks.append(mask.numpy())\n    \n    best_thresh = 0.15\n    best_dice = 0.0\n    \n    for thresh in thresholds:\n        total_dice = 0.0\n        num_samples = 0\n        for prob, mask in zip(all_probs, all_masks):\n            pred_binary = (prob > thresh).astype(np.uint8)\n            mask_binary = (mask == 1).astype(np.uint8)\n            \n            intersection = np.sum(pred_binary * mask_binary)\n            dice = (2. * intersection) / (np.sum(pred_binary) + np.sum(mask_binary) + 1e-6)\n            total_dice += dice\n            num_samples += 1\n        \n        avg_dice = total_dice / num_samples\n        print(f\"Threshold: {thresh:.2f} -> Hard Dice: {avg_dice:.4f}\")\n        \n        if avg_dice > best_dice:\n            best_dice = avg_dice\n            best_thresh = thresh\n    \n    print(f\"\\n Best Threshold: {best_thresh:.2f} (Hard Dice: {best_dice:.4f})\")\n    return best_thresh\n\n# --- Main Execution Flow ---\nif __name__ == \"__main__\":\n    # 1. 训练模型\n    trained_model = train()\n    \n    # 2. 加载最佳模型\n    if Path(\"/kaggle/working/best_model.pth\").exists():\n        trained_model.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\", map_location=CONFIG[\"device\"]))\n        print(\"Loaded best model for inference.\")\n    \n    # 3. 在验证集上搜索最佳阈值\n    print(\"\\n Starting Threshold Search on Validation Set...\")\n    val_ds_for_thresh = VesuviusSliceDataset(CONFIG[\"valid_chunks\"], mode='val')\n    BEST_THRESHOLD = find_best_threshold(trained_model, val_ds_for_thresh)\n    \n    # 4. 使用最佳阈值进行最终推理\n    print(f\"\\n Running final inference with best threshold: {BEST_THRESHOLD:.2f}\")\n    inference(trained_model, threshold=BEST_THRESHOLD)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T12:02:31.424116Z","iopub.execute_input":"2026-01-07T12:02:31.424324Z","iopub.status.idle":"2026-01-07T12:06:06.433649Z","shell.execute_reply.started":"2026-01-07T12:02:31.424305Z","shell.execute_reply":"2026-01-07T12:06:06.432919Z"}},"outputs":[],"execution_count":null}]}