{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":730451,"sourceType":"modelInstanceVersion","modelInstanceId":556400,"modelId":568964}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":65.870912,"end_time":"2026-01-25T08:27:31.242412","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-01-25T08:26:25.371500","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"6eb5a43e","cell_type":"markdown","source":"# Vesuvius Challenge: Surface Detection with GPU Acceleration\n\nThis notebook is optimized for both local GPU training and Kaggle's GPU environment. It will automatically detect your hardware and configure accordingly.","metadata":{"papermill":{"duration":0.005958,"end_time":"2026-01-25T08:26:33.374263","exception":false,"start_time":"2026-01-25T08:26:33.368305","status":"completed"},"tags":[]}},{"id":"165d11e3","cell_type":"code","source":"import os\nimport sys\nimport warnings\nfrom pathlib import Path\nwarnings.filterwarnings('ignore')\n\n# ============================================================\n# CONFIGURATION: Set to True to skip training and load pre-trained checkpoint\n# ============================================================\nINFERENCE_ONLY = False  # Set to True for inference-only runs\n\n# If INFERENCE_ONLY=True, specify the Kaggle input path to your checkpoint\n# Example: '/kaggle/input/vesuvius-checkpoint/checkpoint_epoch_2.pt'\nCHECKPOINT_INPUT_PATH = '/kaggle/input/vesuvius-unet-checkpoint/pytorch/default/1/checkpoint_epoch_2.pt'\n\n# ============================================================\n# SUBMISSION CONFIG: Threshold and format\n# ============================================================\nPRED_THRESHOLD = 0.30  # Tunable threshold for binary conversion (0.3-0.7 range typically optimal)\n\n# ============================================================\n# POST-PROCESSING CONFIG (Hysteresis + Closing + Dust Removal)\n# ============================================================\nUSE_TOPO_POSTPROCESS = False\nPOST_T_LOW = 0.20\nPOST_T_HIGH = 0.50\nPOST_Z_RADIUS = 3\nPOST_XY_RADIUS = 2\nPOST_DUST_MIN_SIZE = 0\n\nif POST_T_LOW > POST_T_HIGH:\n    raise ValueError(\"POST_T_LOW must be <= POST_T_HIGH\")\n\n# Check if running on Kaggle\nIS_KAGGLE = 'KAGGLE_DATA_PROXY_URL' in os.environ\nprint(f\"Running on Kaggle: {IS_KAGGLE}\")\nprint(f\"Inference-only mode: {INFERENCE_ONLY}\")\nprint(f\"Prediction threshold: {PRED_THRESHOLD}\")\nprint(f\"Post-process enabled: {USE_TOPO_POSTPROCESS}\")\nif USE_TOPO_POSTPROCESS:\n    print(f\"  Hysteresis: T_low={POST_T_LOW}, T_high={POST_T_HIGH}\")\n    print(f\"  Closing: z_radius={POST_Z_RADIUS}, xy_radius={POST_XY_RADIUS}\")\n    print(f\"  Dust min size: {POST_DUST_MIN_SIZE}\")\n\n# Detect environment paths\nif IS_KAGGLE:\n    DATA_PATH = Path('/kaggle/input/vesuvius-challenge-surface-detection')\n    OUTPUT_PATH = Path('/kaggle/working')\nelse:\n    # Notebook lives in notebooks\n    DATA_PATH = Path('vesuvius-challenge-surface-detection')\n    OUTPUT_PATH = Path('output')\n\n# Create output directory if it doesn't exist\nOUTPUT_PATH.mkdir(parents=True, exist_ok=True)\nprint(f\"Data path: {DATA_PATH}\")\nprint(f\"Output path: {OUTPUT_PATH}\")","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:39.828023Z","iopub.execute_input":"2026-02-06T05:31:39.828457Z","iopub.status.idle":"2026-02-06T05:31:39.839764Z","shell.execute_reply.started":"2026-02-06T05:31:39.828427Z","shell.execute_reply":"2026-02-06T05:31:39.838776Z"},"papermill":{"duration":0.015071,"end_time":"2026-01-25T08:26:33.395388","exception":false,"start_time":"2026-01-25T08:26:33.380317","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"c022878f","cell_type":"markdown","source":"## Section 1: Environment & Dependency Setup (Local GPU + Kaggle GPU)\n\nInstall and configure required packages for both local and Kaggle environments.\n\n**IMPORTANT for Kaggle**: Before running this notebook, install `imagecodecs` via Kaggle Add-ons:\n1. Click **Add-ons** at the top of the notebook editor\n2. Click **Install Dependencies**\n3. Add: `pip install imagecodecs`\n4. Click **Save**\n\nThis enables faster TIFF I/O without requiring internet access during notebook execution.","metadata":{"papermill":{"duration":0.005564,"end_time":"2026-01-25T08:26:33.406983","exception":false,"start_time":"2026-01-25T08:26:33.401419","status":"completed"},"tags":[]}},{"id":"a4977f58","cell_type":"code","source":"# Core packages are already available on Kaggle\n# Only install if missing (mainly for local runs)\nimport subprocess\n\n# Critical packages for efficient TIFF handling and submission generation\npackages_to_check = [\n    'tifffile',\n    'imagecodecs',  # Essential for fast TIFF compression/decompression\n    'pillow',       # Fallback for TIFF reading\n]\n\nprint(\"Checking package availability...\")\nmissing_packages = []\n\nfor package in packages_to_check:\n    package_import_name = package.replace('-', '_')\n    if package == 'pillow':\n        package_import_name = 'PIL'\n    \n    try:\n        __import__(package_import_name)\n        print(f\"✓ {package} available\")\n    except ImportError:\n        print(f\"✗ {package} MISSING\")\n        missing_packages.append(package)\n\nif missing_packages:\n    print(f\"\\n⚠ WARNING: Missing packages: {', '.join(missing_packages)}\")\n    if IS_KAGGLE:\n        print(\"\\n\" + \"=\"*60)\n        print(\"CRITICAL: Install missing packages via Kaggle Add-ons:\")\n        print(\"1. Click 'Add-ons' at top of notebook\")\n        print(\"2. Click 'Install Dependencies'\")\n        print(\"3. Add these packages:\")\n        for pkg in missing_packages:\n            print(f\"   pip install {pkg}\")\n        print(\"4. Click 'Save' and restart notebook\")\n        print(\"=\"*60)\n    else:\n        print(\"\\nAttempting to install locally...\")\n        for package in missing_packages:\n            try:\n                subprocess.check_call([sys.executable, '-m', 'pip', 'install', package, '-q'])\n                print(f\"✓ {package} installed\")\n            except Exception as e:\n                print(f\"⚠ Could not install {package}: {e}\")\nelse:\n    print(\"\\n✓ All required packages available!\")\n    \n# Verify imagecodecs is working with tifffile\ntry:\n    import tifffile\n    import imagecodecs\n    print(f\"✓ tifffile can use imagecodecs backend (fast compression)\")\nexcept ImportError:\n    print(\"⚠ imagecodecs not available - TIFF operations will be slower\")\n\nprint(\"\\n✓ Package check complete!\")","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:39.841289Z","iopub.execute_input":"2026-02-06T05:31:39.841564Z","iopub.status.idle":"2026-02-06T05:31:39.866891Z","shell.execute_reply.started":"2026-02-06T05:31:39.841542Z","shell.execute_reply":"2026-02-06T05:31:39.865762Z"},"papermill":{"duration":0.245881,"end_time":"2026-01-25T08:26:33.658528","exception":false,"start_time":"2026-01-25T08:26:33.412647","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"92de20f3","cell_type":"code","source":"# Import core libraries\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nfrom pathlib import Path\n\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"Torchvision version: {transforms.__version__ if hasattr(transforms, '__version__') else 'N/A'}\")\n\n# Set deterministic behavior for reproducibility\ntorch.manual_seed(42)\nnp.random.seed(42)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(42)","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:39.869044Z","iopub.execute_input":"2026-02-06T05:31:39.869363Z","iopub.status.idle":"2026-02-06T05:31:39.897419Z","shell.execute_reply.started":"2026-02-06T05:31:39.869338Z","shell.execute_reply":"2026-02-06T05:31:39.896357Z"},"papermill":{"duration":16.680917,"end_time":"2026-01-25T08:26:50.345756","exception":false,"start_time":"2026-01-25T08:26:33.664839","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"d602173d","cell_type":"markdown","source":"## Section 2: GPU/TPU Detection and Device Selection\n\nAutomatically detect CUDA, MPS, or CPU and configure the device accordingly.","metadata":{"papermill":{"duration":0.006036,"end_time":"2026-01-25T08:26:50.358256","exception":false,"start_time":"2026-01-25T08:26:50.352220","status":"completed"},"tags":[]}},{"id":"7e270a33","cell_type":"code","source":"# Device selection and GPU detection\ndef get_device():\n    \"\"\"Detect and configure GPU/TPU device\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device('cuda')\n        print(f\"✓ GPU detected: {torch.cuda.get_device_name(0)}\")\n        print(f\"  GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n        print(f\"  CUDA Version: {torch.version.cuda}\")\n        print(f\"  Number of GPUs: {torch.cuda.device_count()}\")\n    elif torch.backends.mps.is_available():\n        device = torch.device('mps')\n        print(\"✓ MPS (Apple Silicon) detected\")\n    else:\n        device = torch.device('cpu')\n        print(\"⚠ No GPU detected. Using CPU (training will be slower)\")\n    \n    return device\n\ndevice = get_device()\nprint(f\"\\n✓ Using device: {device}\")","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:39.898629Z","iopub.execute_input":"2026-02-06T05:31:39.898940Z","iopub.status.idle":"2026-02-06T05:31:39.930625Z","shell.execute_reply.started":"2026-02-06T05:31:39.898913Z","shell.execute_reply":"2026-02-06T05:31:39.929278Z"},"papermill":{"duration":0.047017,"end_time":"2026-01-25T08:26:50.411393","exception":false,"start_time":"2026-01-25T08:26:50.364376","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"af015b37","cell_type":"markdown","source":"## Section 3: Dataset Pathing and Download Check\n\nVerify dataset paths and files exist. Resolve paths for local disk vs Kaggle environment.","metadata":{"papermill":{"duration":0.006061,"end_time":"2026-01-25T08:26:50.424529","exception":false,"start_time":"2026-01-25T08:26:50.418468","status":"completed"},"tags":[]}},{"id":"b2c8a50e","cell_type":"code","source":"# Dataset path verification\nfrom pathlib import Path\n\ndef verify_dataset():\n    \"\"\"Verify dataset exists and list available files (excluding deprecated).\"\"\"\n    data_path = Path(DATA_PATH)\n\n    if not data_path.exists():\n        print(f\"⚠ Dataset path not found: {DATA_PATH}\")\n        print(\"Please ensure the dataset is downloaded to the correct location.\")\n        return False\n\n    print(f\"✓ Dataset path verified: {DATA_PATH}\")\n\n    required_dirs = [\"train_images\", \"train_labels\", \"test_images\"]\n    missing_dirs = [d for d in required_dirs if not (data_path / d).exists()]\n    if missing_dirs:\n        print(f\"⚠ Missing required directories: {missing_dirs}\")\n        return False\n\n    # CSVs\n    train_csv = data_path / \"train.csv\"\n    test_csv = data_path / \"test.csv\"\n    if not train_csv.exists() or not test_csv.exists():\n        print(f\"⚠ Missing train.csv or test.csv in {data_path}\")\n        return False\n\n    # Show counts\n    for d in required_dirs:\n        count = len(list((data_path / d).glob('*.tif')))\n        print(f\"  - {d}: {count} files\")\n\n    # Warn about deprecated data (we do NOT use it)\n    deprecated_dirs = [\"deprecated_train_images\", \"deprecated_train_labels\"]\n    found_deprecated = [d for d in deprecated_dirs if (data_path / d).exists()]\n    if found_deprecated:\n        print(\"⚠ Deprecated directories found (ignored):\")\n        for d in found_deprecated:\n            print(f\"  - {d}\")\n\n    return True\n\n\ndataset_ok = verify_dataset()","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:39.933912Z","iopub.execute_input":"2026-02-06T05:31:39.934665Z","iopub.status.idle":"2026-02-06T05:31:39.977254Z","shell.execute_reply.started":"2026-02-06T05:31:39.934606Z","shell.execute_reply":"2026-02-06T05:31:39.975537Z"},"papermill":{"duration":1.751475,"end_time":"2026-01-25T08:26:52.182024","exception":false,"start_time":"2026-01-25T08:26:50.430549","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"7fbb2351","cell_type":"markdown","source":"## Section 4: Data Loading and Preprocessing\n\nImplement dataset class and preprocessing pipeline for surface detection.","metadata":{"execution":{"iopub.execute_input":"2026-01-19T06:50:39.891968Z","iopub.status.busy":"2026-01-19T06:50:39.891194Z","iopub.status.idle":"2026-01-19T06:50:39.896889Z","shell.execute_reply":"2026-01-19T06:50:39.895843Z","shell.execute_reply.started":"2026-01-19T06:50:39.891936Z"},"papermill":{"duration":0.006116,"end_time":"2026-01-25T08:26:52.195033","exception":false,"start_time":"2026-01-25T08:26:52.188917","status":"completed"},"tags":[]}},{"id":"30df35ab","cell_type":"code","source":"import random\nfrom typing import List, Tuple\n\nimport tifffile as tiff\nfrom PIL import Image, ImageSequence\n\n# Patch and sampling configuration\nPATCH_SIZE = 256\nZ_CONTEXT = 2  # number of slices on each side (total channels = 2*Z_CONTEXT + 1)\nTRAIN_SAMPLES_PER_ID = 64\nVAL_SAMPLES_PER_ID = 16\nVAL_SPLIT = 0.15\nMAX_TRAIN_IDS = 48  # limit for faster local runs; set None to use all ids\n\n\ndef read_volume(path: Path):\n    \"\"\"Read a 3D volume (.tif) with tifffile, fallback to PIL if needed.\"\"\"\n    try:\n        vol = tiff.imread(str(path))\n    except Exception as e:\n        print(f\"tifffile failed for {path}: {e}; falling back to PIL ImageSequence\")\n        im = Image.open(str(path))\n        frames = [np.array(frame) for frame in ImageSequence.Iterator(im)]\n        if len(frames) == 0:\n            raise RuntimeError(f\"No frames found in {path}\")\n        vol = np.stack(frames, axis=0)\n    if vol.ndim == 2:  # single slice\n        vol = vol[None, ...]\n    return vol\n\n\ndef normalize_patch(patch: np.ndarray) -> np.ndarray:\n    \"\"\"Per-patch normalization with percentile clipping.\"\"\"\n    patch = patch.astype(np.float32)\n    lo, hi = np.percentile(patch, 1), np.percentile(patch, 99)\n    patch = np.clip(patch, lo, hi)\n    mean = patch.mean()\n    std = patch.std() + 1e-6\n    patch = (patch - mean) / std\n    return patch\n\n\ndef make_coords(z_max: int, h: int, w: int, samples: int, patch: int, z_context: int, rng: np.random.Generator):\n    \"\"\"Sample valid (z, y, x) coordinates for patch extraction.\"\"\"\n    coords = []\n    z_choices = rng.integers(z_context, z_max - z_context, size=samples)\n    y_choices = rng.integers(0, max(1, h - patch), size=samples)\n    x_choices = rng.integers(0, max(1, w - patch), size=samples)\n    for z, y, x in zip(z_choices, y_choices, x_choices):\n        coords.append((int(z), int(y), int(x)))\n    return coords\n\n\nclass VesuviusSurfaceDataset(Dataset):\n    \"\"\"Patch-based 2.5D dataset for surface detection.\"\"\"\n\n    def __init__(\n        self,\n        ids: List[str],\n        images_dir: Path,\n        labels_dir: Path,\n        samples_per_id: int,\n        patch_size: int = PATCH_SIZE,\n        z_context: int = Z_CONTEXT,\n        augment: bool = False,\n        seed: int = 42,\n    ):\n        self.ids = ids\n        self.images_dir = images_dir\n        self.labels_dir = labels_dir\n        self.samples_per_id = samples_per_id\n        self.patch_size = patch_size\n        self.z_context = z_context\n        self.augment = augment\n        self.rng = np.random.default_rng(seed)\n\n        self.volumes = {}\n        self.labels = {}\n        self.coords = []  # list of (id, z, y, x)\n\n        for vid in self.ids:\n            img_path = self.images_dir / f\"{vid}.tif\"\n            lbl_path = self.labels_dir / f\"{vid}.tif\"\n            if not img_path.exists() or not lbl_path.exists():\n                continue\n            vol = read_volume(img_path)\n            lbl = read_volume(lbl_path)\n            if lbl.ndim == 3:\n                lbl = lbl[0]  # labels are single channel\n            z_max, h, w = vol.shape\n            if h < self.patch_size or w < self.patch_size or z_max <= (2 * self.z_context):\n                continue\n            self.volumes[vid] = vol\n            self.labels[vid] = lbl\n            coords = make_coords(z_max, h, w, self.samples_per_id, self.patch_size, self.z_context, self.rng)\n            self.coords.extend([(vid, z, y, x) for (z, y, x) in coords])\n\n        print(f\"Dataset built: {len(self.coords)} patches from {len(self.volumes)} volumes\")\n\n    def __len__(self):\n        return len(self.coords)\n\n    def _augment(self, img: np.ndarray, mask: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:\n        if random.random() < 0.5:\n            img = img[..., ::-1, :]\n            mask = mask[..., ::-1]\n        if random.random() < 0.5:\n            img = img[..., :, ::-1]\n            mask = mask[..., ::-1]\n        if random.random() < 0.2:\n            img = np.transpose(img, (0, 2, 1))\n            mask = np.transpose(mask, (1, 0))\n        return img, mask\n\n    def __getitem__(self, idx):\n        vid, z, y, x = self.coords[idx]\n        vol = self.volumes[vid]\n        lbl = self.labels[vid]\n\n        z0, z1 = z - self.z_context, z + self.z_context + 1\n        patch_img = vol[z0:z1, y : y + self.patch_size, x : x + self.patch_size]\n        patch_mask = lbl[y : y + self.patch_size, x : x + self.patch_size]\n\n        patch_img = normalize_patch(patch_img)\n        patch_mask = (patch_mask > 0).astype(np.float32)\n\n        if self.augment:\n            patch_img, patch_mask = self._augment(patch_img, patch_mask)\n\n        # Make contiguous copies to avoid negative strides from flipping/transposing\n        patch_img = np.ascontiguousarray(patch_img, dtype=np.float32)\n        patch_mask = np.ascontiguousarray(patch_mask, dtype=np.float32)\n\n        patch_img = torch.from_numpy(patch_img).float()\n        patch_mask = torch.from_numpy(patch_mask).float().unsqueeze(0)\n\n        return patch_img, patch_mask\n\n\nprint(\"✓ Patch-based dataset and config defined\")","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:39.978650Z","iopub.execute_input":"2026-02-06T05:31:39.978950Z","iopub.status.idle":"2026-02-06T05:31:40.005195Z","shell.execute_reply.started":"2026-02-06T05:31:39.978924Z","shell.execute_reply":"2026-02-06T05:31:40.004114Z"},"papermill":{"duration":0.027466,"end_time":"2026-01-25T08:26:52.228630","exception":false,"start_time":"2026-01-25T08:26:52.201164","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"4bad7847","cell_type":"markdown","source":"## Section 5: DataLoader Configuration for GPU Throughput\n\nCreate optimized DataLoaders with GPU acceleration in mind.","metadata":{"papermill":{"duration":0.006188,"end_time":"2026-01-25T08:26:52.241209","exception":false,"start_time":"2026-01-25T08:26:52.235021","status":"completed"},"tags":[]}},{"id":"f96fc3db","cell_type":"code","source":"# Build train/val dataloaders using real volumes\nif INFERENCE_ONLY:\n    print(\"⏭ Skipping dataset loading (INFERENCE_ONLY=True)\")\n    train_loader = None\n    val_loader = None\nelse:\n    train_df = pd.read_csv(DATA_PATH / 'train.csv')\n    all_ids = train_df['id'].astype(str).tolist()\n\n    # Shuffle and optionally limit ids for faster experimentation\n    rng = np.random.default_rng(42)\n    rng.shuffle(all_ids)\n    if MAX_TRAIN_IDS is not None:\n        all_ids = all_ids[:MAX_TRAIN_IDS]\n\n    split_idx = int(len(all_ids) * (1 - VAL_SPLIT))\n    train_ids = all_ids[:split_idx]\n    val_ids = all_ids[split_idx:]\n\n    train_dataset = VesuviusSurfaceDataset(\n        ids=train_ids,\n        images_dir=DATA_PATH / 'train_images',\n        labels_dir=DATA_PATH / 'train_labels',\n        samples_per_id=TRAIN_SAMPLES_PER_ID,\n        patch_size=PATCH_SIZE,\n        z_context=Z_CONTEXT,\n        augment=True,\n        seed=42,\n    )\n\n    val_dataset = VesuviusSurfaceDataset(\n        ids=val_ids,\n        images_dir=DATA_PATH / 'train_images',\n        labels_dir=DATA_PATH / 'train_labels',\n        samples_per_id=VAL_SAMPLES_PER_ID,\n        patch_size=PATCH_SIZE,\n        z_context=Z_CONTEXT,\n        augment=False,\n        seed=123,\n    )\n\n    # With imagecodecs installed, we can safely use multiple workers\n    num_workers = 4 if IS_KAGGLE else 2\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=4,\n        shuffle=True,\n        num_workers=num_workers,\n        pin_memory=(device.type == 'cuda'),\n    )\n\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=4,\n        shuffle=False,\n        num_workers=num_workers,\n        pin_memory=(device.type == 'cuda'),\n    )\n\n    print(f\"✓ Train patches: {len(train_dataset)} from {len(train_ids)} ids\")\n    print(f\"✓ Val patches:   {len(val_dataset)} from {len(val_ids)} ids\")\n    print(f\"✓ DataLoader workers: {num_workers}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-02-06T05:31:40.006630Z","iopub.execute_input":"2026-02-06T05:31:40.006924Z"},"papermill":{"duration":0.016317,"end_time":"2026-01-25T08:26:52.263707","exception":false,"start_time":"2026-01-25T08:26:52.247390","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"fe37af1a","cell_type":"markdown","source":"## Section 6: Model Definition\n\nDefine a neural network for surface detection and move it to the selected device.","metadata":{"papermill":{"duration":0.006322,"end_time":"2026-01-25T08:26:52.276159","exception":false,"start_time":"2026-01-25T08:26:52.269837","status":"completed"},"tags":[]}},{"id":"16999546","cell_type":"code","source":"# 2.5D U-Net model for surface detection\nclass ConvBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass UNet2p5D(nn.Module):\n    def __init__(self, in_ch: int, base: int = 32):\n        super().__init__()\n        self.enc1 = ConvBlock(in_ch, base)\n        self.enc2 = ConvBlock(base, base * 2)\n        self.enc3 = ConvBlock(base * 2, base * 4)\n        self.pool = nn.MaxPool2d(2)\n\n        self.bottleneck = ConvBlock(base * 4, base * 8)\n\n        self.up3 = nn.ConvTranspose2d(base * 8, base * 4, kernel_size=2, stride=2)\n        self.dec3 = ConvBlock(base * 8, base * 4)\n        self.up2 = nn.ConvTranspose2d(base * 4, base * 2, kernel_size=2, stride=2)\n        self.dec2 = ConvBlock(base * 4, base * 2)\n        self.up1 = nn.ConvTranspose2d(base * 2, base, kernel_size=2, stride=2)\n        self.dec1 = ConvBlock(base * 2, base)\n\n        self.out_conv = nn.Conv2d(base, 1, kernel_size=1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n\n        b = self.bottleneck(self.pool(e3))\n\n        d3 = self.up3(b)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n\n        d2 = self.up2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n\n        d1 = self.up1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n\n        return self.out_conv(d1)  # logits\n\n\nin_channels = 2 * Z_CONTEXT + 1\nmodel = UNet2p5D(in_ch=in_channels, base=32)\nmodel = model.to(device)\n\nprint(\"✓ 2.5D U-Net initialized\")\nprint(f\"  Input channels: {in_channels}\")\nprint(f\"  Parameters: {sum(p.numel() for p in model.parameters()):,}\")","metadata":{"papermill":{"duration":0.249038,"end_time":"2026-01-25T08:26:52.531568","exception":false,"start_time":"2026-01-25T08:26:52.282530","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"16274c7b","cell_type":"markdown","source":"## Section 7: Training Loop with Device Placement\n\nImplement training loop with proper GPU device placement and loss computation.","metadata":{"papermill":{"duration":0.00616,"end_time":"2026-01-25T08:26:52.544366","exception":false,"start_time":"2026-01-25T08:26:52.538206","status":"completed"},"tags":[]}},{"id":"d15b8d69","cell_type":"code","source":"# Losses and optimizer\nbce_loss = nn.BCEWithLogitsLoss()\n\n\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits)\n    num = 2.0 * (probs * targets).sum(dim=(1, 2, 3))\n    den = probs.sum(dim=(1, 2, 3)) + targets.sum(dim=(1, 2, 3)) + eps\n    dice = 1 - num / den\n    return dice.mean()\n\n\ndef combined_loss(logits, targets):\n    return 0.5 * bce_loss(logits, targets) + 0.5 * dice_loss(logits, targets)\n\noptimizer = optim.Adam(model.parameters(), lr=1e-3)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5)\nuse_amp = device.type == 'cuda'\nscaler = torch.cuda.amp.GradScaler(enabled=use_amp)\n\n\n# Training function\n\ndef train_epoch(model, train_loader, optimizer, scaler, device):\n    model.train()\n    total_loss = 0.0\n    total_dice = 0.0\n\n    progress_bar = tqdm(train_loader, desc=\"Training\")\n\n    for batch_idx, (inputs, masks) in enumerate(progress_bar):\n        inputs = inputs.to(device)\n        masks = masks.to(device)\n\n        optimizer.zero_grad()\n        with torch.cuda.amp.autocast(enabled=use_amp):\n            logits = model(inputs)\n            loss = combined_loss(logits, masks)\n            batch_dice = 1 - dice_loss(logits.detach(), masks).item()\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item()\n        total_dice += batch_dice\n\n        progress_bar.set_postfix({\n            'Loss': total_loss / (batch_idx + 1),\n            'Dice': total_dice / (batch_idx + 1)\n        })\n\n    return total_loss / len(train_loader), total_dice / len(train_loader)\n\n\nprint(\"✓ Training components ready (BCE + Dice, AMP enabled on CUDA)\")","metadata":{"papermill":{"duration":0.018432,"end_time":"2026-01-25T08:26:52.568744","exception":false,"start_time":"2026-01-25T08:26:52.550312","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"49813cda","cell_type":"markdown","source":"## Section 8: Evaluation and Metrics\n\nRun validation on the device and compute accuracy metrics.","metadata":{"papermill":{"duration":0.006438,"end_time":"2026-01-25T08:26:52.581642","exception":false,"start_time":"2026-01-25T08:26:52.575204","status":"completed"},"tags":[]}},{"id":"f0aceb72","cell_type":"code","source":"# Validation function\n\ndef validate(model, val_loader, device):\n    model.eval()\n    total_loss = 0.0\n    total_dice = 0.0\n\n    with torch.no_grad():\n        progress_bar = tqdm(val_loader, desc=\"Validating\")\n\n        for inputs, masks in progress_bar:\n            inputs = inputs.to(device)\n            masks = masks.to(device)\n\n            logits = model(inputs)\n            loss = combined_loss(logits, masks)\n            dice = 1 - dice_loss(logits, masks).item()\n\n            total_loss += loss.item()\n            total_dice += dice\n\n            progress_bar.set_postfix({\n                'Loss': total_loss / (progress_bar.n + 1),\n                'Dice': total_dice / (progress_bar.n + 1)\n            })\n\n    return total_loss / len(val_loader), total_dice / len(val_loader)\n\n\nprint(\"✓ Validation function ready (Dice tracked)\")","metadata":{"papermill":{"duration":0.014678,"end_time":"2026-01-25T08:26:52.602327","exception":false,"start_time":"2026-01-25T08:26:52.587649","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"d5d678d8","cell_type":"markdown","source":"## Section 9: Checkpointing and Kaggle-Compatible Output\n\nSave model weights and predictions with Kaggle-compatible output paths.","metadata":{"papermill":{"duration":0.006005,"end_time":"2026-01-25T08:26:52.614558","exception":false,"start_time":"2026-01-25T08:26:52.608553","status":"completed"},"tags":[]}},{"id":"17883df4","cell_type":"code","source":"# Checkpointing and model saving\ndef save_checkpoint(model, optimizer, epoch, loss, output_dir):\n    \"\"\"Save model checkpoint\"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    \n    checkpoint_path = os.path.join(output_dir, f'checkpoint_epoch_{epoch}.pt')\n    \n    torch.save({\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': loss,\n    }, checkpoint_path)\n    \n    print(f\"✓ Checkpoint saved: {checkpoint_path}\")\n    return checkpoint_path\n\ndef save_model_onnx(model, output_dir, input_shape=(1, 3, 256, 256)):\n    \"\"\"Save model in ONNX format for deployment\"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    \n    onnx_path = os.path.join(output_dir, 'model.onnx')\n    \n    # Create dummy input\n    dummy_input = torch.randn(input_shape).to(device)\n    \n    try:\n        torch.onnx.export(\n            model,\n            dummy_input,\n            onnx_path,\n            export_params=True,\n            opset_version=11,\n            do_constant_folding=True,\n            input_names=['input'],\n            output_names=['output'],\n            dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}},\n            verbose=False\n        )\n        print(f\"✓ Model exported to ONNX: {onnx_path}\")\n    except Exception as e:\n        print(f\"⚠ ONNX export failed: {e}\")\n\nprint(\"✓ Checkpointing functions defined\")","metadata":{"papermill":{"duration":0.015143,"end_time":"2026-01-25T08:26:52.635722","exception":false,"start_time":"2026-01-25T08:26:52.620579","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"137e33d1","cell_type":"markdown","source":"## Section 10: Full Training Pipeline\n\nExecute the complete training pipeline with GPU acceleration.","metadata":{"papermill":{"duration":0.005906,"end_time":"2026-01-25T08:26:52.647982","exception":false,"start_time":"2026-01-25T08:26:52.642076","status":"completed"},"tags":[]}},{"id":"710fe910","cell_type":"code","source":"# Full training pipeline\nif INFERENCE_ONLY:\n    print(\"=\"*60)\n    print(\"⏭ SKIPPING TRAINING (INFERENCE_ONLY=True)\")\n    print(\"=\"*60)\n    history = {'train_loss': [], 'val_loss': [], 'train_dice': [], 'val_dice': []}\nelse:\n    NUM_EPOCHS = 3\n    best_val_dice = 0.0\n    history = {'train_loss': [], 'val_loss': [], 'train_dice': [], 'val_dice': []}\n\n    print(\"Starting training...\")\n    print(f\"Device: {device}\")\n    print(f\"Epochs: {NUM_EPOCHS}\")\n    print(\"-\" * 50)\n\n    for epoch in range(NUM_EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS}\")\n\n        # Train\n        train_loss, train_dice = train_epoch(model, train_loader, optimizer, scaler, device)\n\n        # Validate\n        val_loss, val_dice = validate(model, val_loader, device)\n\n        # Record history\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        history['train_dice'].append(train_dice)\n        history['val_dice'].append(val_dice)\n\n        # Print results\n        print(f\"Train Loss: {train_loss:.4f}, Train Dice: {train_dice:.4f}\")\n        print(f\"Val Loss:   {val_loss:.4f}, Val Dice:   {val_dice:.4f}\")\n\n        # Save checkpoint\n        if val_dice > best_val_dice:\n            best_val_dice = val_dice\n            save_checkpoint(model, optimizer, epoch, val_loss, OUTPUT_PATH)\n\n        # Update learning rate\n        scheduler.step()\n\n    print(\"\\n\" + \"=\" * 50)\n    print(f\"Training completed! Best validation Dice: {best_val_dice:.4f}\")\n    print(\"=\" * 50)","metadata":{"papermill":{"duration":0.014799,"end_time":"2026-01-25T08:26:52.668824","exception":false,"start_time":"2026-01-25T08:26:52.654025","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"89d73d63","cell_type":"markdown","source":"## Section 11: Results Visualization\n\nPlot training history and visualize model performance.","metadata":{"papermill":{"duration":0.006308,"end_time":"2026-01-25T08:26:52.681602","exception":false,"start_time":"2026-01-25T08:26:52.675294","status":"completed"},"tags":[]}},{"id":"4f34973b","cell_type":"code","source":"# Visualize training history\nif not INFERENCE_ONLY and history['train_loss']:\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\n    # Loss plot\n    axes[0].plot(history['train_loss'], label='Train Loss', marker='o')\n    axes[0].plot(history['val_loss'], label='Val Loss', marker='o')\n    axes[0].set_xlabel('Epoch')\n    axes[0].set_ylabel('Loss')\n    axes[0].set_title('Training & Validation Loss')\n    axes[0].legend()\n    axes[0].grid(True)\n\n    # Dice plot\n    axes[1].plot(history['train_dice'], label='Train Dice', marker='o')\n    axes[1].plot(history['val_dice'], label='Val Dice', marker='o')\n    axes[1].set_xlabel('Epoch')\n    axes[1].set_ylabel('Dice')\n    axes[1].set_title('Training & Validation Dice')\n    axes[1].legend()\n    axes[1].grid(True)\n\n    plt.tight_layout()\n    plt.savefig(OUTPUT_PATH / 'training_history.png', dpi=100, bbox_inches='tight')\n    print(f\"✓ Training history saved to {OUTPUT_PATH}/training_history.png\")\n    plt.show()\nelse:\n    print(\"⏭ Skipping training visualization (no training data)\")","metadata":{"papermill":{"duration":0.015321,"end_time":"2026-01-25T08:26:52.703133","exception":false,"start_time":"2026-01-25T08:26:52.687812","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"a578cb47","cell_type":"markdown","source":"## Section 12: Final Model Export and Summary\n\nExport the trained model and create a summary report.","metadata":{"papermill":{"duration":0.006311,"end_time":"2026-01-25T08:26:52.715379","exception":false,"start_time":"2026-01-25T08:26:52.709068","status":"completed"},"tags":[]}},{"id":"3cbdb392","cell_type":"code","source":"# Export model in PyTorch format\nif not INFERENCE_ONLY:\n    model_path = OUTPUT_PATH / 'vesuvius_model.pt'\n    torch.save(model.state_dict(), model_path)\n    print(f\"✓ Model saved: {model_path}\")\n\n    # Export to ONNX for cross-platform compatibility\n    save_model_onnx(model, OUTPUT_PATH)\nelse:\n    print(\"⏭ Skipping model export (INFERENCE_ONLY=True)\")\n\n# Create training summary\nif not INFERENCE_ONLY:\n    summary = f\"\"\"\nVESUVIUS CHALLENGE - TRAINING SUMMARY\n=====================================\n\nTraining Configuration:\n- Device: {device}\n- GPU Available: {torch.cuda.is_available()}\n- Model: 2.5D U-Net\n- Input Channels: {2 * Z_CONTEXT + 1}\n- Patch Size: {PATCH_SIZE}\n- Total Parameters: {sum(p.numel() for p in model.parameters()):,}\n- Batch Size: 4\n- Learning Rate: 0.001\n- Optimizer: Adam\n- Loss: 0.5*BCE + 0.5*Dice\n- Num Epochs: {NUM_EPOCHS}\n\nFinal Results:\n- Best Validation Dice: {max(history['val_dice']) if history['val_dice'] else 0:.4f}\n- Final Train Loss: {history['train_loss'][-1]:.4f}\n- Final Val Loss: {history['val_loss'][-1]:.4f}\n- Final Train Dice: {history['train_dice'][-1]:.4f}\n- Final Val Dice: {history['val_dice'][-1]:.4f}\n\nOutput Files:\n- Model (PyTorch): {model_path}\n- Model (ONNX): {OUTPUT_PATH / 'model.onnx'}\n- Training History: {OUTPUT_PATH / 'training_history.png'}\n\nEnvironment:\n- Running on Kaggle: {IS_KAGGLE}\n- Data Path: {DATA_PATH}\n- Output Path: {OUTPUT_PATH}\n\"\"\"\n    summary_path = OUTPUT_PATH / 'training_summary.txt'\n    with open(summary_path, 'w') as f:\n        f.write(summary)\n\n    print(\"\\n\" + summary)\n    print(f\"✓ Summary saved: {summary_path}\")\nelse:\n    print(\"⏭ Skipping training summary (INFERENCE_ONLY=True)\")","metadata":{"papermill":{"duration":0.016348,"end_time":"2026-01-25T08:26:52.738295","exception":false,"start_time":"2026-01-25T08:26:52.721947","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"29345af4","cell_type":"markdown","source":"## Submission: Generate `submission.zip` with U-Net predictions\n\nThis section loads the trained U-Net model and uses it to generate predictions for all test volumes.\nFormat: Sliding-window inference with 2.5D patches, aggregated with overlap voting for robust predictions.","metadata":{"papermill":{"duration":0.006134,"end_time":"2026-01-25T08:26:52.750808","exception":false,"start_time":"2026-01-25T08:26:52.744674","status":"completed"},"tags":[]}},{"id":"6b42aced","cell_type":"code","source":"# 1. SETUP: Verify dependencies and paths\nimport os\nimport sys\nimport shutil\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\nimport tifffile as tiff\n\n# Verify imagecodecs is available for optimal performance (read-only)\ntry:\n    import imagecodecs\n    print(\"✓ imagecodecs available - using optimized TIFF I/O\")\n    HAS_IMAGECODECS = True\nexcept ImportError:\n    print(\"⚠ imagecodecs not available - using fallback (slower)\")\n    HAS_IMAGECODECS = False\n\n# Post-processing dependencies (required when USE_TOPO_POSTPROCESS=True)\nif USE_TOPO_POSTPROCESS:\n    try:\n        from scipy import ndimage as ndi\n        from skimage.morphology import remove_small_objects\n        print(\"✓ scipy/skimage available - post-processing enabled\")\n    except Exception as e:\n        raise RuntimeError(\n            \"Post-processing requires scipy and scikit-image. \"\n            \"Install them or set USE_TOPO_POSTPROCESS=False.\"\n        ) from e\n\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Setup\")\nprint(\"=\"*60)\n\n# Setup paths\nDATA_PATH = Path(DATA_PATH)\nOUTPUT_PATH = Path(OUTPUT_PATH)\nTEST_CSV = DATA_PATH / 'test.csv'\nTEST_DIR = DATA_PATH / 'test_images'\n\n# Validate paths exist\nif not TEST_CSV.exists():\n    raise FileNotFoundError(f\"test.csv not found: {TEST_CSV}\")\nif not TEST_DIR.exists():\n    raise FileNotFoundError(f\"test_images not found: {TEST_DIR}\")\n\nprint(f\"✓ Data paths verified\")\nprint(f\"  TEST_CSV: {TEST_CSV}\")\nprint(f\"  TEST_DIR: {TEST_DIR}\")\n\n# Write predictions to a temporary directory so only submission.zip remains\nif IS_KAGGLE:\n    PRED_DIR = Path('/kaggle/temp/vesuvius_preds')\nelse:\n    PRED_DIR = OUTPUT_PATH / '_pred_tmp'\nif PRED_DIR.exists():\n    shutil.rmtree(PRED_DIR)\nPRED_DIR.mkdir(parents=True, exist_ok=True)\nprint(f\"✓ Temporary prediction directory: {PRED_DIR}\")\n\n# Post-processing helpers\ndef build_anisotropic_struct(z_radius: int, xy_radius: int):\n    z, r = z_radius, xy_radius\n    if z == 0 and r == 0:\n        return None\n    if z == 0 and r > 0:\n        size = 2 * r + 1\n        struct = np.zeros((1, size, size), dtype=bool)\n        cy, cx = r, r\n        for dy in range(-r, r + 1):\n            for dx in range(-r, r + 1):\n                if dy * dy + dx * dx <= r * r:\n                    struct[0, cy + dy, cx + dx] = True\n        return struct\n    if z > 0 and r == 0:\n        struct = np.zeros((2 * z + 1, 1, 1), dtype=bool)\n        struct[:, 0, 0] = True\n        return struct\n    depth = 2 * z + 1\n    size = 2 * r + 1\n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz, cy, cx = z, r, r\n    for dz in range(-z, z + 1):\n        for dy in range(-r, r + 1):\n            for dx in range(-r, r + 1):\n                if dy * dy + dx * dx <= r * r:\n                    struct[cz + dz, cy + dy, cx + dx] = True\n    return struct\n\n\ndef topo_postprocess(\n    probs,\n    T_low=0.90,\n    T_high=0.90,\n    z_radius=1,\n    xy_radius=0,\n    dust_min_size=100,\n):\n    # Step 1: 3D Hysteresis\n    strong = probs >= T_high\n    weak = probs >= T_low\n\n    if not strong.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n\n    struct_hyst = ndi.generate_binary_structure(3, 3)\n    mask = ndi.binary_propagation(\n        strong, mask=weak, structure=struct_hyst\n    )\n\n    if not mask.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n\n    # Step 2: 3D Anisotropic Closing\n    if z_radius > 0 or xy_radius > 0:\n        struct_close = build_anisotropic_struct(z_radius, xy_radius)\n        if struct_close is not None:\n            mask = ndi.binary_closing(mask, structure=struct_close)\n\n    # Step 3: Dust Removal\n    if dust_min_size > 0:\n        mask = remove_small_objects(\n            mask.astype(bool), min_size=dust_min_size\n        )\n\n    return mask.astype(np.uint8)\n\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"9292b2a3","cell_type":"code","source":"# 2. MODEL LOADING: Load checkpoint and initialize U-Net\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Load Model\")\nprint(\"=\"*60)\n\n# Find or load the checkpoint\nif INFERENCE_ONLY and IS_KAGGLE:\n    # Load from Kaggle input\n    MODEL_PATH = Path(CHECKPOINT_INPUT_PATH)\n    if not MODEL_PATH.exists():\n        input_checkpoints = list(Path('/kaggle/input').rglob('checkpoint_*.pt'))\n        if input_checkpoints:\n            MODEL_PATH = input_checkpoints[0]\n            print(f\"⚠ Using auto-discovered checkpoint: {MODEL_PATH}\")\n        else:\n            raise FileNotFoundError(\"No checkpoint found. Set CHECKPOINT_INPUT_PATH or add checkpoint as Kaggle input.\")\n    print(f\"✓ Loading checkpoint from Kaggle input: {MODEL_PATH}\")\nelse:\n    # Load from training output\n    checkpoint_dir = OUTPUT_PATH\n    checkpoints = list(checkpoint_dir.glob('checkpoint_*.pt'))\n    if not checkpoints:\n        raise FileNotFoundError(f\"No checkpoints found in {checkpoint_dir}\")\n    MODEL_PATH = sorted(checkpoints)[-1]\n    print(f\"✓ Loading checkpoint from training: {MODEL_PATH}\")\n\n# Device setup\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"✓ Using device: {device}\")\n\n# Model architecture (must match training exactly)\nclass ConvBlock(torch.nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = torch.nn.Sequential(\n            torch.nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),\n            torch.nn.BatchNorm2d(out_ch),\n            torch.nn.ReLU(inplace=True),\n            torch.nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),\n            torch.nn.BatchNorm2d(out_ch),\n            torch.nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass UNet2p5D(torch.nn.Module):\n    def __init__(self, in_ch: int, base: int = 32):\n        super().__init__()\n        self.enc1 = ConvBlock(in_ch, base)\n        self.enc2 = ConvBlock(base, base * 2)\n        self.enc3 = ConvBlock(base * 2, base * 4)\n        self.pool = torch.nn.MaxPool2d(2)\n\n        self.bottleneck = ConvBlock(base * 4, base * 8)\n\n        self.up3 = torch.nn.ConvTranspose2d(base * 8, base * 4, kernel_size=2, stride=2)\n        self.dec3 = ConvBlock(base * 8, base * 4)\n        self.up2 = torch.nn.ConvTranspose2d(base * 4, base * 2, kernel_size=2, stride=2)\n        self.dec2 = ConvBlock(base * 4, base * 2)\n        self.up1 = torch.nn.ConvTranspose2d(base * 2, base, kernel_size=2, stride=2)\n        self.dec1 = ConvBlock(base * 2, base)\n\n        self.out_conv = torch.nn.Conv2d(base, 1, kernel_size=1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n\n        b = self.bottleneck(self.pool(e3))\n\n        d3 = self.up3(b)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n\n        d2 = self.up2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n\n        d1 = self.up1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n\n        return self.out_conv(d1)\n\n\n# Initialize and load model\nin_channels = 2 * Z_CONTEXT + 1\nmodel = UNet2p5D(in_ch=in_channels, base=32)\ncheckpoint = torch.load(str(MODEL_PATH), map_location=device, weights_only=False)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.to(device)\nmodel.eval()\nprint(f\"✓ Model loaded successfully\")\nprint(f\"  Input channels: {in_channels}\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"c70bd443","cell_type":"code","source":"# 3. DATA VALIDATION: Load test CSV and verify all test images exist\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Load & Validate Test Data\")\nprint(\"=\"*60)\n\n# Load test metadata\ntest_df = pd.read_csv(TEST_CSV)\nif 'id' not in test_df.columns:\n    raise ValueError(\"test.csv must contain an 'id' column\")\n\ntest_ids = test_df['id'].astype(str).tolist()\nprint(f\"✓ Loaded test.csv with {len(test_ids)} test IDs\")\n\n# Verify all test images exist\nmissing_images = [tid for tid in test_ids if not (TEST_DIR / f\"{tid}.tif\").exists()]\nif missing_images:\n    raise FileNotFoundError(f\"Missing test images for ids: {missing_images}\")\n\nprint(f\"✓ All {len(test_ids)} test images found\")\nprint(f\"  First 5 IDs: {test_ids[:5]}\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"6688dffc","cell_type":"code","source":"# 4. INFERENCE: Run sliding-window predictions on all test volumes\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Inference\")\nprint(\"=\"*60)\nif USE_TOPO_POSTPROCESS:\n    print(\"Post-processing enabled\")\n    print(f\"  Hysteresis: T_low={POST_T_LOW}, T_high={POST_T_HIGH}\")\n    print(f\"  Closing: z_radius={POST_Z_RADIUS}, xy_radius={POST_XY_RADIUS}\")\n    print(f\"  Dust min size: {POST_DUST_MIN_SIZE}\")\nelse:\n    print(f\"Prediction threshold: {PRED_THRESHOLD}\")\nprint(f\"Processing {len(test_ids)} test volumes...\\n\")\n\ncreated = []\n\n# Keep one sample for visualization\nviz_volume = None\nviz_mask = None\nviz_id = None\n\n# Inference with sliding window\nfor image_id in test_ids:\n    vol_path = TEST_DIR / f\"{image_id}.tif\"\n\n    print(f\"Processing {image_id}...\", end=\" \", flush=True)\n\n    # Read volume\n    try:\n        vol = tiff.imread(str(vol_path))\n    except Exception as e:\n        # Fallback to PIL if tifffile fails\n        try:\n            from PIL import Image, ImageSequence\n            im = Image.open(str(vol_path))\n            frames = [np.array(frame) for frame in ImageSequence.Iterator(im)]\n            vol = np.stack(frames, axis=0)\n            print(f\"[PIL fallback] \", end=\"\", flush=True)\n        except Exception as e2:\n            print(f\"✗ ERROR: tifffile failed: {e}, PIL failed: {e2}\")\n            continue\n\n    vol = np.asarray(vol, dtype=np.float32)\n    z_size, y_size, x_size = vol.shape\n\n    # Initialize output aggregation\n    pred_sum = np.zeros((z_size, y_size, x_size), dtype=np.float32)\n    pred_count = np.zeros((z_size, y_size, x_size), dtype=np.float32)\n\n    # Sliding window inference\n    with torch.no_grad():\n        for z in range(z_size):\n            z_start = max(0, z - Z_CONTEXT)\n            z_end = min(z_size, z + Z_CONTEXT + 1)\n\n            vol_slice = vol[z_start:z_end]\n\n            # Pad z-axis at boundaries\n            if z_start == 0:\n                vol_slice = np.pad(vol_slice, ((Z_CONTEXT - z, 0), (0, 0), (0, 0)), mode='edge')\n            if z_end == z_size:\n                vol_slice = np.pad(vol_slice, ((0, Z_CONTEXT - (z_size - z - 1)), (0, 0), (0, 0)), mode='edge')\n\n            # Process all patches in this z-slice\n            for y in range(0, y_size, PATCH_SIZE // 2):\n                for x in range(0, x_size, PATCH_SIZE // 2):\n                    y_end = min(y + PATCH_SIZE, y_size)\n                    x_end = min(x + PATCH_SIZE, x_size)\n\n                    patch_img = vol_slice[:, y:y_end, x:x_end]\n\n                    # Pad if at boundary\n                    if patch_img.shape[1] < PATCH_SIZE or patch_img.shape[2] < PATCH_SIZE:\n                        patch_img = np.pad(\n                            patch_img,\n                            ((0, 0), (0, PATCH_SIZE - patch_img.shape[1]), (0, PATCH_SIZE - patch_img.shape[2])),\n                            mode='edge'\n                        )\n\n                    # Normalize\n                    low = np.percentile(patch_img, 1.0)\n                    high = np.percentile(patch_img, 99.0)\n                    if high > low:\n                        patch_img = (patch_img - low) / (high - low)\n                    else:\n                        patch_img = (patch_img - low) / (np.abs(low) + 1e-8)\n                    patch_img = np.clip(patch_img, 0, 1)\n\n                    # Infer\n                    patch_tensor = torch.from_numpy(np.ascontiguousarray(patch_img, dtype=np.float32)).unsqueeze(0).to(device)\n                    logits = model(patch_tensor)\n                    pred = torch.sigmoid(logits).cpu().numpy()[0, 0]\n\n                    # Aggregate: only use the valid region from this patch\n                    valid_y = min(PATCH_SIZE, y_size - y)\n                    valid_x = min(PATCH_SIZE, x_size - x)\n                    pred_sum[z, y:y_end, x:x_end] += pred[:valid_y, :valid_x]\n                    pred_count[z, y:y_end, x:x_end] += 1\n\n    # Average probabilities\n    pred_vol = np.zeros((z_size, y_size, x_size), dtype=np.float32)\n    mask = pred_count > 0\n    pred_vol[mask] = pred_sum[mask] / pred_count[mask]\n\n    # Post-processing or simple thresholding\n    if USE_TOPO_POSTPROCESS:\n        mask_out = topo_postprocess(\n            pred_vol,\n            T_low=POST_T_LOW,\n            T_high=POST_T_HIGH,\n            z_radius=POST_Z_RADIUS,\n            xy_radius=POST_XY_RADIUS,\n            dust_min_size=POST_DUST_MIN_SIZE,\n        )\n    else:\n        mask_out = (pred_vol > PRED_THRESHOLD).astype(np.uint8)\n\n    # Cache one example for visualization (first successful volume)\n    if viz_volume is None:\n        viz_volume = vol\n        viz_mask = mask_out\n        viz_id = image_id\n\n    # Verify output shape matches input exactly\n    if mask_out.shape != (z_size, y_size, x_size):\n        print(f\"✗ Shape mismatch! Expected {(z_size, y_size, x_size)}, got {mask_out.shape}\")\n        continue\n\n    # Save without compression for maximum evaluator compatibility\n    out_path = PRED_DIR / f\"{image_id}.tif\"\n    try:\n        tiff.imwrite(str(out_path), mask_out.astype(np.uint8), compression=None)\n\n        # Validate saved file\n        saved = tiff.imread(str(out_path))\n\n        # Check 1: Exact shape match\n        if saved.shape != (z_size, y_size, x_size):\n            print(f\"✗ Saved shape {saved.shape} != input shape {(z_size, y_size, x_size)}\")\n            continue\n\n        # Check 2: Correct dtype\n        if saved.dtype != np.uint8:\n            print(f\"✗ Saved dtype {saved.dtype} != uint8\")\n            continue\n\n        # Check 3: Only 0s and 1s\n        unique_vals = np.unique(saved)\n        if not (set(unique_vals.tolist()) <= {0, 1}):\n            print(f\"✗ Invalid values found: {unique_vals}\")\n            continue\n\n        coverage = np.mean(saved)\n        print(f\"✓ (coverage: {coverage:.1%})\", flush=True)\n        created.append(out_path)\n\n    except Exception as e:\n        print(f\"✗ ERROR writing {image_id}: {e}\")\n        continue\n\nprint(f\"\\n{'='*60}\")\nprint(f\"Created {len(created)}/{len(test_ids)} predictions\")\nprint(f\"{'='*60}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"f34c06ab","cell_type":"code","source":"# 5. VALIDATION: Check predictions are complete and match test.csv\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Validate Predictions\")\nprint(\"=\"*60)\n\n# Ensure we have predictions for all test ids\nexpected = set(test_ids)\ncreated_ids = set(p.stem for p in created)\nmissing = sorted(expected - created_ids)\nextra = sorted(created_ids - expected)\n\nif missing:\n    raise RuntimeError(f\"❌ CRITICAL: Missing masks for ids: {missing}\")\nif extra:\n    raise RuntimeError(f\"❌ CRITICAL: Extra masks not in test.csv: {extra}\")\n\nprint(f\"✓ All {len(created)} predictions match test.csv\")\nprint(f\"✓ No missing or extra files\")\n\n# Spot check a few files\nprint(\"\\nSpot-checking output files:\")\nfor p in list(created)[:min(3, len(created))]:\n    arr = tiff.imread(str(p))\n    uniq = np.unique(arr)\n    print(f\"  {p.name}: shape={arr.shape}, dtype={arr.dtype}, values={sorted(uniq.tolist())}\")\n\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"3eabb917-0ec8-4dd6-ad94-eb74df911242","cell_type":"code","source":"# DEBUG: Check prediction value distribution\nprint(\"=\"*60)\nprint(\"DEBUG: Prediction Statistics\")\nprint(\"=\"*60)\nprint(f\"First prediction volume stats (BEFORE thresholding):\")\nprint(f\"  Min: {pred_vol.min():.6f}\")\nprint(f\"  Max: {pred_vol.max():.6f}\")\nprint(f\"  Mean: {pred_vol.mean():.6f}\")\nprint(f\"  Median: {np.median(pred_vol):.6f}\")\nprint(f\"  % > 0.1: {100 * (pred_vol > 0.1).mean():.2f}%\")\nprint(f\"  % > 0.3: {100 * (pred_vol > 0.3).mean():.2f}%\")\nprint(f\"  % > 0.5: {100 * (pred_vol > 0.5).mean():.2f}%\")\nprint(f\"  % > 0.7: {100 * (pred_vol > 0.7).mean():.2f}%\")\nprint(\"\")\nprint(f\"After thresholding at {PRED_THRESHOLD}:\")\nprint(f\"  % pixels = 1: {100 * mask_out.mean():.2f}%\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"847744a3-5c24-4a4b-8daf-1fa7e7eae86a","cell_type":"code","source":"\n# EXTENDED DEBUG: Check what was actually saved vs what we expect\nprint(\"=\"*60)\nprint(\"EXTENDED DEBUG: Saved Masks Analysis\")\nprint(\"=\"*60)\n\n# Check the submission.zip files directly\nimport zipfile\n\nzip_path = OUTPUT_PATH / 'submission.zip'\nif not zip_path.exists():\n    print(f\"⚠ submission.zip not found at {zip_path}\")\nelse:\n    print(f\"✓ Found submission.zip at {zip_path}\")\n    \n    with zipfile.ZipFile(str(zip_path), 'r') as zf:\n        names = zf.namelist()\n        print(f\"  Total files in zip: {len(names)}\")\n        \n        if names:\n            # Check first file\n            first_name = names[0]\n            print(f\"\\nAnalyzing first file: {first_name}\")\n            \n            with zf.open(first_name) as f:\n                mask_bytes = f.read()\n            \n            # Read it as TIFF\n            import io\n            mask_arr = tiff.imread(io.BytesIO(mask_bytes))\n            print(f\"  Shape: {mask_arr.shape}\")\n            print(f\"  Dtype: {mask_arr.dtype}\")\n            print(f\"  Unique values: {np.unique(mask_arr)}\")\n            print(f\"  Min: {mask_arr.min()}\")\n            print(f\"  Max: {mask_arr.max()}\")\n            print(f\"  Mean: {mask_arr.mean():.4f}\")\n            print(f\"  % pixels = 0: {100 * (mask_arr == 0).mean():.2f}%\")\n            print(f\"  % pixels = 1: {100 * (mask_arr == 1).mean():.2f}%\")\n            \n            # Show sample values\n            if mask_arr.ndim == 3:\n                z = mask_arr.shape[0] // 2\n                print(f\"\\n  Sample slice (z={z}) center (10x10):\")\n                h, w = mask_arr.shape[1], mask_arr.shape[2]\n                sample = mask_arr[z, h//2-5:h//2+5, w//2-5:w//2+5]\n                print(f\"    {sample}\")\n\n# Also check the temporary prediction directory\nprint(f\"\\n{'='*60}\")\nprint(\"EXTENDED DEBUG: Temporary Prediction Directory\")\nprint(f\"{'='*60}\")\n\nif PRED_DIR.exists():\n    pred_files = list(PRED_DIR.glob('*.tif'))\n    print(f\"✓ Found {len(pred_files)} prediction files in {PRED_DIR}\")\n    \n    if pred_files:\n        # Check first file\n        first_pred = pred_files[0]\n        print(f\"\\nAnalyzing first file: {first_pred.name}\")\n        \n        pred_mask = tiff.imread(str(first_pred))\n        print(f\"  Shape: {pred_mask.shape}\")\n        print(f\"  Dtype: {pred_mask.dtype}\")\n        print(f\"  Unique values: {np.unique(pred_mask)}\")\n        print(f\"  Min: {pred_mask.min()}\")\n        print(f\"  Max: {pred_mask.max()}\")\n        print(f\"  Mean: {pred_mask.mean():.4f}\")\n        print(f\"  % pixels = 0: {100 * (pred_mask == 0).mean():.2f}%\")\n        print(f\"  % pixels = 1: {100 * (pred_mask == 1).mean():.2f}%\")\nelse:\n    print(f\"⚠ Prediction directory not found: {PRED_DIR}\")\n\nprint(\"=\"*60)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"4410a15d-f73d-44f3-b356-04b0e15a6f94","cell_type":"code","source":"import zipfile\nimport tifffile as tiff\nimport numpy as np\nimport io\nfrom pathlib import Path\n\n# Check submission.zip\nzip_path = Path('/kaggle/working/submission.zip')\nprint(f\"Checking: {zip_path}\")\nprint(f\"Exists: {zip_path.exists()}\")\n\nif zip_path.exists():\n    with zipfile.ZipFile(str(zip_path), 'r') as zf:\n        names = zf.namelist()\n        if names:\n            first = names[0]\n            print(f\"\\nFirst file: {first}\")\n            with zf.open(first) as f:\n                arr = tiff.imread(io.BytesIO(f.read()))\n            print(f\"  Unique values: {np.unique(arr)}\")\n            print(f\"  % = 0: {100*(arr==0).mean():.1f}%\")\n            print(f\"  % = 1: {100*(arr==1).mean():.1f}%\")\n            print(f\"  Min={arr.min()}, Max={arr.max()}, Mean={arr.mean():.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"2ab1b764","cell_type":"code","source":"# 6. ZIP CREATION: Create submission.zip at root of output directory\nimport zipfile\n\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Create submission.zip\")\nprint(\"=\"*60)\n\nzip_path = OUTPUT_PATH / 'submission.zip'\nprint(f\"\\nCreating submission.zip at: {zip_path}\")\n\nwith zipfile.ZipFile(str(zip_path), 'w', zipfile.ZIP_DEFLATED) as zf:\n    for p in sorted(created):\n        # Add file at root level of zip with just filename\n        zf.write(str(p), arcname=p.name)\n\nprint(f\"✓ Created submission.zip with {len(created)} files\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"9b6016f4","cell_type":"code","source":"# 7. ZIP VALIDATION: Verify submission.zip structure and contents\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Validate submission.zip\")\nprint(\"=\"*60)\n\nwith zipfile.ZipFile(str(zip_path), 'r') as zf:\n    names = zf.namelist()\n    print(f\"\\nZip file analysis:\")\n    print(f\"  Total entries: {len(names)}\")\n    \n    # Check root level\n    bad_entries = [n for n in names if '/' in n or '\\\\' in n]\n    if bad_entries:\n        raise RuntimeError(f\"❌ CRITICAL: Files not at root level in zip: {bad_entries[:5]}\")\n    print(f\"  ✓ All {len(names)} files at root level (no subdirectories)\")\n    \n    # Check file extensions\n    extensions = set([Path(n).suffix for n in names])\n    if extensions != {'.tif'}:\n        raise RuntimeError(f\"❌ CRITICAL: Unexpected file extensions: {extensions}\")\n    print(f\"  ✓ All files are .tif (extension check passed)\")\n    \n    # Check matching test IDs\n    name_stems = [Path(n).stem for n in names]\n    if set(name_stems) != expected:\n        missing_in_zip = sorted(expected - set(name_stems))\n        extra_in_zip = sorted(set(name_stems) - expected)\n        raise RuntimeError(f\"❌ CRITICAL: ID mismatch. Missing: {missing_in_zip}, Extra: {extra_in_zip}\")\n    print(f\"  ✓ All {len(expected)} test IDs present\")\n\nprint(f\"\\n✓ submission.zip is valid and ready for Kaggle submission!\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"4dfb39fa","cell_type":"code","source":"# 8. CLEANUP: Remove temporary prediction directory\nprint(\"=\"*60)\nprint(\"SUBMISSION PIPELINE: Cleanup\")\nprint(\"=\"*60)\n\n# Cleanup temp predictions so only submission.zip remains in outputs\nif PRED_DIR.exists():\n    shutil.rmtree(PRED_DIR)\n    print(f\"✓ Cleaned up temporary directory: {PRED_DIR}\")\n\n# Final output summary\nprint(f\"\\n{'='*60}\")\nprint(f\"✅ SUBMISSION COMPLETE!\")\nprint(f\"{'='*60}\")\nprint(f\"Location: {zip_path}\")\nprint(f\"Files: {len(created)} test volumes\")\nprint(f\"Format: Binary masks (uint8, values {{0,1}}) in uncompressed TIFF\")\nprint(f\"Threshold used: {PRED_THRESHOLD}\")\nprint(f\"\\n📤 Ready to submit to Kaggle!\")\nprint(f\"{'='*60}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"1282e50f","cell_type":"code","source":"# 9. VISUALIZE: Sample slices from prediction vs input\nimport numpy as np\nimport matplotlib.pyplot as plt\n\ndef plot_sample(x, y, sample_idx=0, max_slices=16):\n    img = np.squeeze(x[sample_idx])  # (D, H, W)\n    mask = np.squeeze(y[sample_idx])  # (D, H, W)\n    D = img.shape[0]\n\n    # Decide which slices to plot\n    step = max(1, D // max_slices)\n    slices = range(0, D, step)\n\n    n_slices = len(slices)\n    fig, axes = plt.subplots(2, n_slices, figsize=(3 * n_slices, 6))\n\n    # Handle case with only 1 slice\n    if n_slices == 1:\n        axes = np.array([[axes[0]], [axes[1]]])\n\n    for i, s in enumerate(slices):\n        axes[0, i].imshow(img[s], cmap='gray')\n        axes[0, i].set_title(f\"Slice {s}\")\n        axes[0, i].axis('off')\n\n        axes[1, i].imshow(mask[s], cmap='gray')\n        axes[1, i].set_title(f\"Mask {s}\")\n        axes[1, i].axis('off')\n\n    plt.suptitle(f\"Sample {sample_idx}\")\n    plt.tight_layout()\n    plt.show()\n\nif viz_volume is None or viz_mask is None:\n    raise RuntimeError(\"No visualization sample found. Run inference cell first.\")\n\nprint(f\"Visualizing sample ID: {viz_id}\")\nplot_sample(\n    viz_volume[None],\n    viz_mask[None],\n    sample_idx=0,\n    max_slices=5\n )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"ffaf1b48-5786-408d-b34a-64194166c800","cell_type":"code","source":"\n# ROOT CAUSE ANALYSIS: Check model logits BEFORE sigmoid (with PIL fallback)\nprint(\"=\"*70)\nprint(\"ROOT CAUSE: Why is model predicting all 1s?\")\nprint(\"=\"*70)\n\ntest_id_check = test_ids[0] if 'test_ids' in dir() else '1407735'\nvol_path_check = TEST_DIR / f\"{test_id_check}.tif\"\n\nprint(f\"\\nTesting with volume: {test_id_check}\")\n\ntry:\n    # Try PIL ImageSequence (should work without imagecodecs for LZW)\n    from PIL import Image, ImageSequence\n    \n    im = Image.open(str(vol_path_check))\n    frames = [np.array(frame) for frame in ImageSequence.Iterator(im)]\n    vol_test = np.stack(frames, axis=0).astype(np.float32)\n    \n    print(f\"✓ Loaded volume with PIL: shape={vol_test.shape}\")\n    \n    # Extract one patch (middle z-slice, top-left corner)\n    z_test = vol_test.shape[0] // 2\n    z_start_test = max(0, z_test - Z_CONTEXT)\n    z_end_test = min(vol_test.shape[0], z_test + Z_CONTEXT + 1)\n    \n    patch_test = vol_test[z_start_test:z_end_test, 0:PATCH_SIZE, 0:PATCH_SIZE].copy()\n    \n    # Pad z-axis if at boundaries\n    if z_start_test == 0:\n        pad_top = Z_CONTEXT - z_test\n        patch_test = np.pad(patch_test, ((pad_top, 0), (0, 0), (0, 0)), mode='edge')\n    if z_end_test == vol_test.shape[0]:\n        pad_bottom = Z_CONTEXT - (vol_test.shape[0] - z_test - 1)\n        patch_test = np.pad(patch_test, ((0, pad_bottom), (0, 0), (0, 0)), mode='edge')\n    \n    # Normalize exactly as in inference\n    low = np.percentile(patch_test, 1.0)\n    high = np.percentile(patch_test, 99.0)\n    patch_norm = (patch_test - low) / (high - low) if high > low else (patch_test - low) / (np.abs(low) + 1e-8)\n    patch_norm = np.clip(patch_norm, 0, 1)\n    \n    print(f\"\\n✓ Patch shape after padding: {patch_norm.shape}\")\n    print(f\"  Input min={patch_norm.min():.4f}, max={patch_norm.max():.4f}, mean={patch_norm.mean():.4f}\")\n    \n    # Forward pass through model\n    patch_tensor = torch.from_numpy(np.ascontiguousarray(patch_norm, dtype=np.float32)).unsqueeze(0).to(device)\n    \n    with torch.no_grad():\n        logits_raw = model(patch_tensor)\n        sigmoid_raw = torch.sigmoid(logits_raw)\n    \n    logits_np = logits_raw.cpu().numpy()[0, 0]\n    sigmoid_np = sigmoid_raw.cpu().numpy()[0, 0]\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"MODEL LOGITS (raw output before sigmoid):\")\n    print(f\"{'='*70}\")\n    print(f\"  Min: {logits_np.min():.6f}\")\n    print(f\"  Max: {logits_np.max():.6f}\")\n    print(f\"  Mean: {logits_np.mean():.6f}\")\n    print(f\"  Median: {np.median(logits_np):.6f}\")\n    print(f\"  Std: {logits_np.std():.6f}\")\n    print(f\"  % > 0: {100 * (logits_np > 0).mean():.1f}%\")\n    print(f\"  % > 5: {100 * (logits_np > 5).mean():.1f}%\")\n    print(f\"  % > 10: {100 * (logits_np > 10).mean():.1f}%\")\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"MODEL SIGMOID (after sigmoid activation):\")\n    print(f\"{'='*70}\")\n    print(f\"  Min: {sigmoid_np.min():.6f}\")\n    print(f\"  Max: {sigmoid_np.max():.6f}\")\n    print(f\"  Mean: {sigmoid_np.mean():.6f}\")\n    print(f\"  Median: {np.median(sigmoid_np):.6f}\")\n    print(f\"  % > 0.5: {100 * (sigmoid_np > 0.5).mean():.1f}%\")\n    print(f\"  % > 0.9: {100 * (sigmoid_np > 0.9).mean():.1f}%\")\n    print(f\"  % > 0.99: {100 * (sigmoid_np > 0.99).mean():.1f}%\")\n    \n    # Show sample values\n    print(f\"\\nSample logits (center 5x5 region):\")\n    h, w = logits_np.shape\n    sample_logits = logits_np[h//2-2:h//2+3, w//2-2:w//2+3]\n    print(sample_logits.round(3))\n    \n    print(f\"\\nSample sigmoid (center 5x5 region):\")\n    sample_sigmoid = sigmoid_np[h//2-2:h//2+3, w//2-2:w//2+3]\n    print(sample_sigmoid.round(6))\n    \n    print(f\"\\n{'='*70}\")\n    print(\"DIAGNOSIS:\")\n    print(f\"{'='*70}\")\n    \n    if logits_np.mean() > 5:\n        print(\"❌ CRITICAL: Logits are VERY HIGH (mean > 5)\")\n        print(\"   → sigmoid(x) ≈ 1 for all x > 5\")\n        print(\"   → All predictions become 1s\")\n        print(\"\\nROOT CAUSE: Checkpoint is outputting excessive logits\")\n        print(\"   This typically means:\")\n        print(\"   1. Model was trained with wrong loss function (no normalization)\")\n        print(\"   2. Model weights are corrupted\")\n        print(\"   3. Model architecture mismatch\")\n        print(\"\\nSOLUTION: Need different/retrained checkpoint\")\n    elif logits_np.mean() > 1:\n        print(\"⚠ WARNING: Logits somewhat elevated (mean > 1)\")\n        print(f\"   Sigmoid mean: {sigmoid_np.mean():.4f}\")\n        if sigmoid_np.mean() > 0.95:\n            print(\"   → Most predictions near 1\")\n            print(\"   Solution: Lower threshold OR retrain model\")\n    elif sigmoid_np.mean() > 0.95:\n        print(\"❌ PROBLEM: Sigmoid outputs too high (mean > 0.95)\")\n        print(\"   Even thresholds at 0.3-0.5 will produce all 1s\")\n        print(\"   This is BAD - model lacks discrimination\")\n        print(\"   Solution: Need better model training\")\n    elif logits_np.std() < 0.1:\n        print(\"❌ PROBLEM: Logits have almost NO VARIANCE (std < 0.1)\")\n        print(\"   All outputs nearly identical\")\n        print(\"   Model is not learning anything useful\")\n        print(\"   Solution: Checkpoint is broken, need retraining\")\n    else:\n        print(\"✓ Model outputs look reasonable\")\n        print(f\"   Logits mean={logits_np.mean():.3f}, std={logits_np.std():.3f}\")\n        print(f\"   Sigmoid mean={sigmoid_np.mean():.4f}\")\n        print(\"   May be an issue elsewhere (aggregation, thresholding)\")\n\nexcept Exception as e:\n    print(f\"ERROR during analysis: {e}\")\n    import traceback\n    traceback.print_exc()\n\nprint(\"=\"*70)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}