{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":290917305,"sourceType":"kernelVersion"},{"sourceId":655294,"sourceType":"modelInstanceVersion","modelInstanceId":495238,"modelId":510647},{"sourceId":660383,"sourceType":"modelInstanceVersion","modelInstanceId":499479,"modelId":510647},{"sourceId":665924,"sourceType":"modelInstanceVersion","modelInstanceId":504051,"modelId":510647},{"sourceId":672178,"sourceType":"modelInstanceVersion","modelInstanceId":495238,"modelId":510647},{"sourceId":673516,"sourceType":"modelInstanceVersion","modelInstanceId":499479,"modelId":510647},{"sourceId":674747,"sourceType":"modelInstanceVersion","modelInstanceId":503784,"modelId":510647},{"sourceId":681152,"sourceType":"modelInstanceVersion","modelInstanceId":516822,"modelId":510647},{"sourceId":732880,"sourceType":"modelInstanceVersion","modelInstanceId":516822,"modelId":510647}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setup and Installation","metadata":{}},{"cell_type":"code","source":"from IPython.display import clear_output\n\nvar=\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"\n!pip install \\\n  \"$var\"/keras_nightly-*.whl \\\n  \"$var\"/tifffile-*.whl \\\n  \"$var\"/imagecodecs-*.whl \\\n  \"$var\"/medicai-*.whl \\\n  --no-index \\\n  --find-links \"$var\"\n\nclear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:03:49.373920Z","iopub.execute_input":"2026-01-31T17:03:49.374113Z","iopub.status.idle":"2026-01-31T17:03:53.043674Z","shell.execute_reply.started":"2026-01-31T17:03:49.374097Z","shell.execute_reply":"2026-01-31T17:03:53.042831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n\nimport keras\nfrom keras import ops\nfrom medicai.transforms import (\n    Compose,\n    ScaleIntensityRange,\n    NormalizeIntensity\n)\nfrom medicai.models import SegFormer, TransUNet\nfrom medicai.utils.inference import SlidingWindowInference\n\nimport numpy as np\nimport pandas as pd\nimport zipfile\nimport tifffile\nimport scipy.ndimage as ndi\nfrom scipy.ndimage import zoom\nfrom skimage.morphology import remove_small_objects, binary_opening, binary_closing\nfrom skimage.measure import label, regionprops\nfrom matplotlib import pyplot as plt\nfrom typing import List, Tuple, Optional\nimport gc\n\nprint(f\"Keras Backend: {keras.config.backend()}, Version: {keras.version()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:03:53.044705Z","iopub.execute_input":"2026-01-31T17:03:53.045065Z","iopub.status.idle":"2026-01-31T17:03:57.590196Z","shell.execute_reply.started":"2026-01-31T17:03:53.045041Z","shell.execute_reply":"2026-01-31T17:03:57.589491Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"# Paths\nroot_dir = \"/kaggle/input/vesuvius-challenge-surface-detection\"\ntest_dir = f\"{root_dir}/test_images\"\noutput_dir = \"/kaggle/working/submission_masks\"\nzip_path = \"/kaggle/working/submission.zip\"\nkaggle_model_path = \"/kaggle/input/vsd-model/keras/\"\nos.makedirs(output_dir, exist_ok=True)\n\n# Model configurations\nMODEL_CONFIGS = [\n    {\n        'name': 'TransUNet-160-3cls',\n        'input_shape': (160, 160, 160),\n        'num_classes': 3,\n        'window_size': (160, 160, 160),\n        'overlap': 0.25,\n        'weight_path': f\"{kaggle_model_path}/transunet/2/transunet.seresnext50.160px.weights.h5\"\n    },\n    {\n        'name': 'TransUNet-192-3cls',\n        'input_shape': (192, 192, 192),\n        'num_classes': 3,\n        'window_size': (192, 192, 192),\n        'overlap': 0.25,\n        'weight_path': f\"{kaggle_model_path}/transunet/2/transunet.seresnext50.192px.weights.h5\"\n    }\n]\n\n# Ensemble weights (can be tuned based on validation performance)\nENSEMBLE_WEIGHTS = [0.4, 0.6]  # Favor larger model slightly\n\n# Post-processing parameters\nPOSTPROCESS_CONFIG = {\n    'T_low': 0.30,           # Lower threshold for hysteresis\n    'T_high': 0.80,          # Higher threshold for hysteresis\n    'z_radius': 3,           # Anisotropic closing z-axis\n    'xy_radius': 2,          # Anisotropic closing xy-plane\n    'dust_min_size': 150,    # Minimum object size to keep\n    'min_area': 200,         # Minimum area for connected components\n    'use_adaptive': True     # Use adaptive thresholding\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:03:57.592427Z","iopub.execute_input":"2026-01-31T17:03:57.592904Z","iopub.status.idle":"2026-01-31T17:03:57.598476Z","shell.execute_reply.started":"2026-01-31T17:03:57.592883Z","shell.execute_reply":"2026-01-31T17:03:57.597744Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df = pd.read_csv(f\"{root_dir}/test.csv\")\nprint(f\"Test samples: {len(test_df)}\")\ntest_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:03:57.599197Z","iopub.execute_input":"2026-01-31T17:03:57.599375Z","iopub.status.idle":"2026-01-31T17:03:57.627514Z","shell.execute_reply.started":"2026-01-31T17:03:57.599361Z","shell.execute_reply":"2026-01-31T17:03:57.626771Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Loading and Preprocessing","metadata":{}},{"cell_type":"code","source":"def load_volume(tif_path: str) -> np.ndarray:\n    \"\"\"Load 3D volume with memory efficiency.\"\"\"\n    volume = tifffile.imread(tif_path)\n    \n    # Convert to float32\n    volume = volume.astype(np.float32)\n    \n    # Add batch and channel dimensions: (D, H, W) -> (1, D, H, W, 1)\n    volume = volume[None, ..., None]\n    \n    return volume\n\n\ndef advanced_normalization(volume: np.ndarray, clip_percentile: float = 99.5) -> np.ndarray:\n    \"\"\"Advanced normalization with outlier clipping.\"\"\"\n    # Clip extreme values\n    nonzero = volume[volume > 0]\n    if len(nonzero) > 0:\n        upper = np.percentile(nonzero, clip_percentile)\n        volume = np.clip(volume, 0, upper)\n    \n    # Normalize using medicai\n    data = {\"image\": volume}\n    pipeline = Compose([\n        NormalizeIntensity(\n            keys=[\"image\"],\n            nonzero=True,\n            channel_wise=False\n        ),\n    ])\n    result = pipeline(data)\n    \n    # Ensure we return a numpy array, not a tensor\n    normalized = result[\"image\"]\n    if hasattr(normalized, 'numpy'):\n        normalized = normalized.numpy()\n    elif not isinstance(normalized, np.ndarray):\n        normalized = np.array(normalized)\n    \n    return normalized\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:03:57.628225Z","iopub.execute_input":"2026-01-31T17:03:57.628446Z","iopub.status.idle":"2026-01-31T17:03:57.634377Z","shell.execute_reply.started":"2026-01-31T17:03:57.628430Z","shell.execute_reply":"2026-01-31T17:03:57.633834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Loading and Ensemble Setup","metadata":{}},{"cell_type":"code","source":"def create_model(config: dict, apply_softmax: bool = False):\n    \"\"\"Create and load a single model.\"\"\"\n    model = TransUNet(\n        input_shape=(*config['input_shape'], 1),\n        encoder_name='seresnext50',\n        classifier_activation='softmax' if apply_softmax else None,\n        num_classes=config['num_classes'],\n    )\n    \n    if os.path.exists(config['weight_path']):\n        model.load_weights(config['weight_path'])\n        print(f\"Loaded weights for {config['name']}\")\n    else:\n        print(f\"Warning: Weights not found for {config['name']}\")\n    \n    return model\n\n\ndef setup_ensemble_models(configs: List[dict], apply_softmax: bool = False):\n    \"\"\"Setup all ensemble models with their inference configurations.\"\"\"\n    ensemble_models = []\n    \n    for config in configs:\n        print(f\"\\nSetting up {config['name']}...\")\n        \n        # Create model\n        model = create_model(config, apply_softmax=apply_softmax)\n        \n        # Create sliding window inference\n        swi = SlidingWindowInference(\n            model=model,\n            num_classes=config['num_classes'],\n            roi_size=config['window_size'],\n            sw_batch_size=1,\n            overlap=config['overlap'],\n            mode=\"gaussian\",\n        )\n        \n        ensemble_models.append((model, swi, config))\n    \n    return ensemble_models\n\n\n# Initialize ensemble\nprint(\"Initializing ensemble models...\")\nensemble_models = setup_ensemble_models(MODEL_CONFIGS, apply_softmax=False)\nprint(f\"\\nEnsemble ready with {len(ensemble_models)} models\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:03:57.635082Z","iopub.execute_input":"2026-01-31T17:03:57.635277Z","iopub.status.idle":"2026-01-31T17:04:16.416981Z","shell.execute_reply.started":"2026-01-31T17:03:57.635261Z","shell.execute_reply":"2026-01-31T17:04:16.416272Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Advanced Inference Functions","metadata":{}},{"cell_type":"code","source":"def to_numpy(x):\n    \"\"\"Convert tensor to numpy array if needed.\"\"\"\n    if hasattr(x, 'numpy'):\n        return x.numpy()\n    elif not isinstance(x, np.ndarray):\n        return np.array(x)\n    return x\n\n\ndef predict_with_extended_tta(inputs: np.ndarray, swi, apply_softmax: bool = False) -> np.ndarray:\n    \"\"\"Extended test-time augmentation with more transformations.\"\"\"\n    logits_list = []\n    \n    # Original\n    logits = swi(inputs)\n    logits = to_numpy(logits)\n    if apply_softmax:\n        logits = to_numpy(ops.softmax(logits, axis=-1))\n    logits_list.append(logits)\n    \n    # Horizontal flip\n    img_flip = np.flip(inputs, axis=2)\n    logits = swi(img_flip)\n    logits = to_numpy(logits)\n    if apply_softmax:\n        logits = to_numpy(ops.softmax(logits, axis=-1))\n    logits = np.flip(logits, axis=2)\n    logits_list.append(logits)\n    \n    # Vertical flip\n    img_flip = np.flip(inputs, axis=3)\n    logits = swi(img_flip)\n    logits = to_numpy(logits)\n    if apply_softmax:\n        logits = to_numpy(ops.softmax(logits, axis=-1))\n    logits = np.flip(logits, axis=3)\n    logits_list.append(logits)\n    \n    # Axial rotations (90°, 180°, 270°)\n    for k in [1, 2, 3]:\n        img_rot = np.rot90(inputs, k=k, axes=(2, 3))\n        logits = swi(img_rot)\n        logits = to_numpy(logits)\n        if apply_softmax:\n            logits = to_numpy(ops.softmax(logits, axis=-1))\n        logits = np.rot90(logits, k=-k, axes=(2, 3))\n        logits_list.append(logits)\n    \n    # Transpose\n    img_t = np.transpose(inputs, (0, 1, 3, 2, 4))\n    logits = swi(img_t)\n    logits = to_numpy(logits)\n    if apply_softmax:\n        logits = to_numpy(ops.softmax(logits, axis=-1))\n    logits = np.transpose(logits, (0, 1, 3, 2, 4))\n    logits_list.append(logits)\n    \n    # Average all augmentations\n    mean_logits = np.mean(logits_list, axis=0)\n    \n    return mean_logits\n\n\ndef uncertainty_aware_ensemble(\n    inputs: np.ndarray,\n    ensemble_models: List,\n    ensemble_weights: Optional[List[float]] = None,\n    use_extended_tta: bool = True\n) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"Ensemble with uncertainty estimation.\"\"\"\n    \n    if ensemble_weights is None:\n        ensemble_weights = [1.0 / len(ensemble_models)] * len(ensemble_models)\n    \n    all_predictions = []\n    ensemble_logits = []\n    \n    for i, (model, swi, config) in enumerate(ensemble_models):\n        print(f\"  Model {i+1}/{len(ensemble_models)}: {config['name']}\")\n        \n        if use_extended_tta:\n            logits = predict_with_extended_tta(inputs, swi, apply_softmax=False)\n        else:\n            logits = swi(inputs)\n            logits = to_numpy(logits)\n        \n        # Apply weight\n        weighted_logits = logits * ensemble_weights[i]\n        ensemble_logits.append(weighted_logits)\n        \n        # Store probability for uncertainty calculation\n        probs = to_numpy(ops.softmax(logits, axis=-1))\n        all_predictions.append(probs)\n    \n    # Weighted average of logits\n    mean_logits = np.sum(ensemble_logits, axis=0)\n    final_probs = to_numpy(ops.softmax(mean_logits, axis=-1))\n    \n    # Calculate prediction variance as uncertainty measure\n    all_preds_stack = np.stack(all_predictions, axis=0)\n    uncertainty = np.var(all_preds_stack, axis=0)\n    \n    return final_probs, uncertainty\n\n\ndef adaptive_threshold_selection(probs: np.ndarray, uncertainty: np.ndarray) -> Tuple[float, float]:\n    \"\"\"Adaptively select thresholds based on prediction distribution.\"\"\"\n    # Get foreground probability (class 1)\n    fg_probs = probs[..., 1]\n    \n    # Calculate percentiles\n    percentiles = np.percentile(fg_probs, [10, 50, 90])\n    \n    # Adjust thresholds based on distribution\n    if percentiles[1] < 0.3:\n        # Low confidence overall -> lower thresholds\n        T_low = max(0.20, percentiles[0])\n        T_high = max(0.60, percentiles[2])\n    elif percentiles[1] > 0.7:\n        # High confidence -> higher thresholds\n        T_low = min(0.40, percentiles[0])\n        T_high = min(0.85, percentiles[2])\n    else:\n        # Default moderate thresholds\n        T_low = 0.30\n        T_high = 0.80\n    \n    return T_low, T_high\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:04:16.417836Z","iopub.execute_input":"2026-01-31T17:04:16.418078Z","iopub.status.idle":"2026-01-31T17:04:16.431733Z","shell.execute_reply.started":"2026-01-31T17:04:16.418058Z","shell.execute_reply":"2026-01-31T17:04:16.430973Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Advanced Post-Processing","metadata":{}},{"cell_type":"code","source":"def build_anisotropic_struct(z_radius: int, xy_radius: int) -> Optional[np.ndarray]:\n    \"\"\"Build anisotropic structuring element.\"\"\"\n    z, r = z_radius, xy_radius\n    \n    if z == 0 and r == 0:\n        return None\n    \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    \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    \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    \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    \n    return struct\n\n\ndef advanced_postprocess(\n    probs: np.ndarray,\n    uncertainty: Optional[np.ndarray] = None,\n    T_low: float = 0.30,\n    T_high: float = 0.80,\n    z_radius: int = 3,\n    xy_radius: int = 2,\n    dust_min_size: int = 150,\n    min_area: int = 200,\n    use_adaptive: bool = True\n) -> np.ndarray:\n    \"\"\"Enhanced post-processing with connected component analysis.\"\"\"\n    \n    # Get foreground probability (handle both 3D and 4D inputs)\n    if probs.ndim == 4:\n        # Shape is (1, D, H, W, num_classes) or (D, H, W, num_classes)\n        if probs.shape[0] == 1:\n            probs = probs[0]  # Remove batch dimension\n        fg_prob = probs[..., 1]  # Get class 1 (foreground)\n    else:\n        # Already 3D probability map\n        fg_prob = probs\n    \n    # Adaptive threshold selection\n    if use_adaptive and uncertainty is not None:\n        # Make sure probs has correct shape for adaptive selection\n        if probs.ndim == 3:\n            # Need to reconstruct class probabilities for adaptive selection\n            # Assume binary: [1-fg_prob, fg_prob]\n            probs_for_adapt = np.stack([1 - fg_prob, fg_prob], axis=-1)\n        else:\n            probs_for_adapt = probs\n        \n        T_low_adapted, T_high_adapted = adaptive_threshold_selection(probs_for_adapt, uncertainty)\n        print(f\"  Adaptive thresholds: T_low={T_low_adapted:.3f}, T_high={T_high_adapted:.3f}\")\n    else:\n        T_low_adapted, T_high_adapted = T_low, T_high\n    \n    # Step 1: 3D Hysteresis thresholding\n    strong = fg_prob >= T_high_adapted\n    weak = fg_prob >= T_low_adapted\n    \n    if not strong.any():\n        return np.zeros_like(fg_prob, dtype=np.uint8)\n    \n    struct_hyst = ndi.generate_binary_structure(3, 3)\n    mask = ndi.binary_propagation(strong, mask=weak, structure=struct_hyst)\n    \n    if not mask.any():\n        return np.zeros_like(fg_prob, dtype=np.uint8)\n    \n    # Step 2: Morphological opening (remove small noise)\n    struct_open = ndi.generate_binary_structure(3, 1)\n    mask = binary_opening(mask, footprint=struct_open)\n    \n    # Step 3: Anisotropic closing (fill gaps)\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 = binary_closing(mask, footprint=struct_close)\n    \n    # Step 4: Connected component analysis\n    labeled = label(mask)\n    regions = regionprops(labeled)\n    \n    # Filter components by size\n    filtered_mask = np.zeros_like(mask, dtype=bool)\n    for region in regions:\n        if region.area >= min_area:\n            # Keep component\n            coords = region.coords\n            filtered_mask[coords[:, 0], coords[:, 1], coords[:, 2]] = True\n    \n    # Step 5: Final dust removal\n    if dust_min_size > 0:\n        filtered_mask = remove_small_objects(filtered_mask, min_size=dust_min_size)\n    \n    return filtered_mask.astype(np.uint8)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:04:16.432567Z","iopub.execute_input":"2026-01-31T17:04:16.432762Z","iopub.status.idle":"2026-01-31T17:04:16.450698Z","shell.execute_reply.started":"2026-01-31T17:04:16.432745Z","shell.execute_reply":"2026-01-31T17:04:16.450154Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main Inference Pipeline","metadata":{}},{"cell_type":"code","source":"def to_numpy(x):\n    \"\"\"Convert tensor to numpy array if needed.\"\"\"\n    if hasattr(x, 'numpy'):\n        return x.numpy()\n    elif not isinstance(x, np.ndarray):\n        return np.array(x)\n    return x\n\n\ndef inference_pipeline(\n    volume: np.ndarray,\n    ensemble_models: List,\n    ensemble_weights: Optional[List[float]] = None,\n    use_extended_tta: bool = True,\n    **postprocess_kwargs\n) -> np.ndarray:\n    \"\"\"Complete inference pipeline.\"\"\"\n    \n    # Ensure input is numpy\n    volume = to_numpy(volume)\n    \n    # Prediction with uncertainty\n    probs, uncertainty = uncertainty_aware_ensemble(\n        volume,\n        ensemble_models,\n        ensemble_weights=ensemble_weights,\n        use_extended_tta=use_extended_tta\n    )\n    \n    # Ensure outputs are numpy\n    probs = to_numpy(probs)\n    uncertainty = to_numpy(uncertainty)\n    \n    # Post-processing\n    final_mask = advanced_postprocess(\n        probs,\n        uncertainty=uncertainty,\n        **postprocess_kwargs\n    )\n    \n    return final_mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:04:16.451308Z","iopub.execute_input":"2026-01-31T17:04:16.452311Z","iopub.status.idle":"2026-01-31T17:04:16.467716Z","shell.execute_reply.started":"2026-01-31T17:04:16.452292Z","shell.execute_reply":"2026-01-31T17:04:16.467185Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generate Submission","metadata":{}},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STARTING INFERENCE\")\nprint(\"=\"*60)\n\nwith zipfile.ZipFile(zip_path, \"w\", compression=zipfile.ZIP_DEFLATED) as z:\n    for idx, image_id in enumerate(test_df[\"id\"], 1):\n        print(f\"\\n[{idx}/{len(test_df)}] Processing: {image_id}\")\n        print(\"-\" * 60)\n        \n        # Load volume\n        tif_path = f\"{test_dir}/{image_id}.tif\"\n        volume = load_volume(tif_path)\n        print(f\"  Volume shape: {volume.shape}\")\n        \n        # Normalize\n        volume = advanced_normalization(volume, clip_percentile=99.5)\n        print(f\"  Normalized (min={volume.min():.3f}, max={volume.max():.3f})\")\n        \n        # Run inference\n        print(f\"  Running ensemble inference...\")\n        output = inference_pipeline(\n            volume,\n            ensemble_models,\n            ensemble_weights=ENSEMBLE_WEIGHTS,\n            use_extended_tta=True,\n            **POSTPROCESS_CONFIG\n        )\n        \n        # Save result\n        out_path = f\"{output_dir}/{image_id}.tif\"\n        tifffile.imwrite(out_path, output.astype(np.uint8))\n        \n        # Add to zip\n        z.write(out_path, arcname=f\"{image_id}.tif\")\n        os.remove(out_path)\n        \n        # Memory cleanup\n        del volume, output\n        gc.collect()\n        \n        print(f\"  ✓ Saved to submission.zip\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(f\"SUBMISSION COMPLETE: {zip_path}\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T17:04:16.468523Z","iopub.execute_input":"2026-01-31T17:04:16.468770Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization","metadata":{}},{"cell_type":"code","source":"def plot_sample_with_uncertainty(\n    volume: np.ndarray,\n    mask: np.ndarray,\n    uncertainty: Optional[np.ndarray] = None,\n    sample_idx: int = 0,\n    max_slices: int = 8\n):\n    \"\"\"Visualize volume, mask, and uncertainty.\"\"\"\n    img = np.squeeze(volume)\n    msk = np.squeeze(mask)\n    D = img.shape[0]\n    \n    step = max(1, D // max_slices)\n    slices = range(0, D, step)\n    n_slices = len(slices)\n    \n    rows = 3 if uncertainty is not None else 2\n    fig, axes = plt.subplots(rows, n_slices, figsize=(3*n_slices, 3*rows))\n    \n    for i, s in enumerate(slices):\n        # Input\n        axes[0, i].imshow(img[s], cmap='gray')\n        axes[0, i].set_title(f\"Input {s}\")\n        axes[0, i].axis('off')\n        \n        # Mask\n        axes[1, i].imshow(msk[s], cmap='gray')\n        axes[1, i].set_title(f\"Mask {s}\")\n        axes[1, i].axis('off')\n        \n        # Uncertainty\n        if uncertainty is not None:\n            unc = np.squeeze(uncertainty)\n            axes[2, i].imshow(unc[s], cmap='hot')\n            axes[2, i].set_title(f\"Uncertainty {s}\")\n            axes[2, i].axis('off')\n    \n    plt.suptitle(f\"Sample Visualization\", fontsize=14)\n    plt.tight_layout()\n    plt.show()\n\n\n# Example: Visualize last processed volume (if available)\n# plot_sample_with_uncertainty(volume, output, uncertainty, max_slices=6)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}