{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":14014260,"sourceType":"datasetVersion","datasetId":8927911},{"sourceId":14147695,"sourceType":"datasetVersion","datasetId":9016590},{"sourceId":277417413,"sourceType":"kernelVersion"},{"sourceId":284402333,"sourceType":"kernelVersion"},{"sourceId":284426507,"sourceType":"kernelVersion"},{"sourceId":284681383,"sourceType":"kernelVersion"},{"sourceId":285110872,"sourceType":"kernelVersion"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nimport os\n\n\n# 3. Add the 'networks' folder to Python's path\n# The structure is CSA-Net -> CSANet -> networks\nsys.path.append('/kaggle/working/CSA-Net/CSANet')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:39.938663Z","iopub.execute_input":"2025-12-14T03:41:39.938983Z","iopub.status.idle":"2025-12-14T03:41:39.949337Z","shell.execute_reply.started":"2025-12-14T03:41:39.938952Z","shell.execute_reply":"2025-12-14T03:41:39.947903Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CSA-Net","metadata":{}},{"cell_type":"code","source":"!cp \"/kaggle/input/csanetimage/Screenshot 2025-12-14 090652.png\" ./my_image.png","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:48:26.498181Z","iopub.execute_input":"2025-12-14T03:48:26.498706Z","iopub.status.idle":"2025-12-14T03:48:26.827987Z","shell.execute_reply.started":"2025-12-14T03:48:26.498663Z","shell.execute_reply":"2025-12-14T03:48:26.826353Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"![img](my_image.png)\n","metadata":{}},{"cell_type":"markdown","source":"The reason I chose CSA-Net was because of the cross attention of the middle layer, with it's neighbor layers. This helps the model learn the continuity, and the self attention helps it learn the connections within.\n","metadata":{}},{"cell_type":"markdown","source":"I currently am not as motivated to work with this architecture, because of the training time it takes to learn things, But making this public so that someone else with the patience might make it reach it's potential.","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import sys\nimport os\n\n# --- 1. DEFINE THE EXACT ROOT PATH ---\n# We use the path you found, pointing to the folder *containing* the 'networks' module.\n# The 'd' in the path often means it's a dynamic path specific to your session, but we use it.\nCORRECT_IMPORT_DIR = \"/kaggle/input/d/choudharymanas/csa-net/CSA-Net-main/CSANet\"\n\n# 2. ADD THE CORRECT DIRECTORY TO PYTHON'S PATH\nif os.path.exists(os.path.join(CORRECT_IMPORT_DIR, \"networks\", \"vit_seg_modeling.py\")):\n    sys.path.append(CORRECT_IMPORT_DIR)\n    print(\"System path successfully updated to find CSANet source.\")\nelse:\n    print(f\"ERROR: Cannot find CSANet source at {CORRECT_IMPORT_DIR}.\")\n    print(\"Please check your uploaded dataset path for the 'CSA-Net-main/CSANet' folder.\")\n\n# 3. CORRECTLY IMPORT THE MODEL\n# Now that 'CSANet' is on the path, Python can find 'networks' inside it.\nfrom networks.vit_seg_modeling import VisionTransformer as CSANet\n\n# 4. Final verification: Check if it's on the path\nif CORRECT_IMPORT_DIR in sys.path:\n    print(\"Import configuration complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:39.951610Z","iopub.execute_input":"2025-12-14T03:41:39.951952Z","iopub.status.idle":"2025-12-14T03:41:45.579607Z","shell.execute_reply.started":"2025-12-14T03:41:39.951927Z","shell.execute_reply":"2025-12-14T03:41:45.578269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:45.580724Z","iopub.execute_input":"2025-12-14T03:41:45.581172Z","iopub.status.idle":"2025-12-14T03:41:45.586187Z","shell.execute_reply.started":"2025-12-14T03:41:45.581137Z","shell.execute_reply":"2025-12-14T03:41:45.584921Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch \nfrom torch.amp import autocast\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport tqdm as tqdm\nimport os\nimport random\nfrom ipywidgets import interact, IntSlider\nfrom torch.utils.data import Dataset, DataLoader\nimport glob\nfrom sklearn.model_selection import train_test_split\nimport ml_collections\nimport torch.optim as optim\nfrom torch.cuda.amp import autocast, GradScaler\nfrom tqdm.notebook import tqdm\nimport scipy.ndimage as ndi\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:45.587296Z","iopub.execute_input":"2025-12-14T03:41:45.587563Z","iopub.status.idle":"2025-12-14T03:41:46.854452Z","shell.execute_reply.started":"2025-12-14T03:41:45.587538Z","shell.execute_reply":"2025-12-14T03:41:46.853408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nimport os\nimport gc\nimport numpy as np\nimport torch\nimport tifffile\nfrom tqdm.notebook import tqdm\nfrom PIL import Image\n\n# --- CONFIGURATION ---\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nPATCH_SIZE = 224\n\n# STRIDE CONTROLS COMPUTE VS QUALITY\n# 112 = 50% overlap (High Quality, Slower). \n# 200 = Minimal overlap (Fastest). \n# If you get \"Notebook Timeout\", change this to 150 or 200.\nSTRIDE = 112  \nTHRESH = 0.7\nBATCH_SIZE = 16  # Increased slightly for better GPU utilization\n\n# --- ROBUST IMPORT ---\n# Ensures we fail fast if the path is wrong\nCORRECT_IMPORT_DIR = \"/kaggle/input/d/choudharymanas/csa-net/CSA-Net-main/CSANet\"\n\nif os.path.exists(os.path.join(CORRECT_IMPORT_DIR, \"networks\", \"vit_seg_modeling.py\")):\n    sys.path.append(CORRECT_IMPORT_DIR)\n    print(\"System path updated.\")\nelse:\n    raise FileNotFoundError(f\"CSANet source not found at {CORRECT_IMPORT_DIR}. Check dataset mount.\")\n\nfrom networks.vit_seg_modeling import VisionTransformer as CSANet\nprint(f\"Model imported successfully on {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:46.855396Z","iopub.execute_input":"2025-12-14T03:41:46.855845Z","iopub.status.idle":"2025-12-14T03:41:47.102341Z","shell.execute_reply.started":"2025-12-14T03:41:46.855821Z","shell.execute_reply":"2025-12-14T03:41:47.101210Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from scipy.ndimage import convolve, distance_transform_edt, gaussian_filter\nfrom scipy.spatial.distance import cdist\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.103322Z","iopub.execute_input":"2025-12-14T03:41:47.103632Z","iopub.status.idle":"2025-12-14T03:41:47.110358Z","shell.execute_reply.started":"2025-12-14T03:41:47.103605Z","shell.execute_reply":"2025-12-14T03:41:47.109304Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config/Hyperparameters","metadata":{}},{"cell_type":"code","source":"best_cfg = {\n    \"T_low\": 0.90,\n    \"T_high\": 0.98,\n    \"z_radius\": 2,\n    \"xy_radius\": 0,\n    \"dust_min_size\": 100,\n    \"enable_smoothing\": False,\n    \"smoothing_sigma\": 1.0,\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.113798Z","iopub.execute_input":"2025-12-14T03:41:47.114193Z","iopub.status.idle":"2025-12-14T03:41:47.135360Z","shell.execute_reply.started":"2025-12-14T03:41:47.114163Z","shell.execute_reply":"2025-12-14T03:41:47.133945Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Post Processing Functions","metadata":{}},{"cell_type":"code","source":"def smooth_surface_voxel(\n    mask: np.ndarray,\n    sigma: float = 1.0,\n) -> np.ndarray:\n    \"\"\"\n    Surface regularization in voxel space via signed distance smoothing.\n    mask: bool (Z,H,W)\n    sigma: Gaussian sigma in voxels; higher = smoother but more shrink/expand.\n    Returns: bool (Z,H,W)\n    \"\"\"\n    assert mask.ndim == 3\n    mask = mask.astype(bool)\n\n    # Inside / outside distance transforms\n    inside_dt = distance_transform_edt(mask)\n    outside_dt = distance_transform_edt(~mask)\n\n    # Signed distance: positive inside, negative outside\n    sdf = inside_dt - outside_dt\n\n    # Smooth the signed distance field\n    sdf_smooth = gaussian_filter(sdf, sigma=sigma)\n\n    # New mask = points where smoothed sdf is >= 0\n    new_mask = sdf_smooth >= 0.0\n    return new_mask.astype(bool)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.136572Z","iopub.execute_input":"2025-12-14T03:41:47.136931Z","iopub.status.idle":"2025-12-14T03:41:47.159606Z","shell.execute_reply.started":"2025-12-14T03:41:47.136899Z","shell.execute_reply":"2025-12-14T03:41:47.158590Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Hysteresis threshold\nIn simple words, We set a high threshold and a low threshold, The voxels with prob higher than the high threshold are selected, while the ones lower than the low threshold are rejected. And for the ones between these values, points connected (26 connectiviy in such cases) to high threshold voxels are selected. ","metadata":{}},{"cell_type":"code","source":"def hysteresis_threshold_3d(\n    probs: np.ndarray,\n    T_low: float = 0.35,\n    T_high: float = 0.75,\n) -> np.ndarray:\n    \"\"\"\n    3D probabilistic hysteresis thresholding.\n    probs: (Z,H,W) float16/float32 in [0,1].\n    Returns: bool mask (Z,H,W).\n    \"\"\"\n    assert probs.ndim == 3\n    assert 0.0 <= T_low <= T_high <= 1.0\n    if(T_low==T_high):\n        (probs>=T_low).astype(bool)\n    strong = probs >= T_high\n    weak   = probs >= T_low\n\n    # 26-connectivity\n    structure = np.ones((3, 3, 3), dtype=bool)\n\n    # Geodesic reconstruction: propagate strong inside weak\n    propagated = ndi.binary_propagation(strong, structure=structure, mask=weak)\n\n    return propagated.astype(bool)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.161179Z","iopub.execute_input":"2025-12-14T03:41:47.161566Z","iopub.status.idle":"2025-12-14T03:41:47.187897Z","shell.execute_reply.started":"2025-12-14T03:41:47.161541Z","shell.execute_reply":"2025-12-14T03:41:47.186903Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Remove small components\nAs the name suggests we remove super small components, they are just noise.","metadata":{}},{"cell_type":"code","source":"\ndef remove_small_components(\n    mask: np.ndarray,\n    min_size: int = 50,\n) -> np.ndarray:\n    \"\"\"\n    Remove small connected components (< min_size voxels) in 3D.\n    mask: bool (Z,H,W)\n    Uses 26-connectivity.\n    \"\"\"\n    assert mask.ndim == 3\n\n    structure = np.ones((3, 3, 3), dtype=bool)  # 26-connectivity\n    labeled, num = ndi.label(mask, structure=structure)\n    if num == 0:\n        return mask\n\n    counts = np.bincount(labeled.ravel())\n    # 0 = background, keep it as is\n    keep = counts >= min_size\n    keep[0] = False  # background always false\n\n    cleaned = keep[labeled]\n    return cleaned.astype(bool)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.189457Z","iopub.execute_input":"2025-12-14T03:41:47.189796Z","iopub.status.idle":"2025-12-14T03:41:47.209620Z","shell.execute_reply.started":"2025-12-14T03:41:47.189767Z","shell.execute_reply":"2025-12-14T03:41:47.208314Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Anisotropic closing along z\nWe make connections along z axis, I have dunno why since the start assumed that the roll will be along z. Will check this hypothesis sometime soon.","metadata":{}},{"cell_type":"code","source":"\ndef anisotropic_closing_z(\n    mask: np.ndarray,\n    z_radius: int = 1,\n    xy_radius: int = 0,\n    iterations: int = 1,\n) -> np.ndarray:\n    \"\"\"\n    mask: bool (Z,H,W)\n    Anisotropic closing, elongated along Z.\n    z_radius: how far to connect along Z (>=1)\n    xy_radius: usually 0 (or 1 if you want tiny XY linking).\n    \"\"\"\n    assert mask.ndim == 3\n\n    z = z_radius\n    r = xy_radius\n\n    # struct shape (2*z+1, 2*r+1, 2*r+1)\n    structure = np.zeros((2*z + 1, 2*r + 1, 2*r + 1), dtype=bool)\n    # central column along Z\n    structure[:, r, r] = True\n\n    closed = mask\n    for _ in range(iterations):\n        closed = ndi.binary_closing(closed, structure=structure)\n\n    return closed.astype(bool)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.210646Z","iopub.execute_input":"2025-12-14T03:41:47.211021Z","iopub.status.idle":"2025-12-14T03:41:47.237958Z","shell.execute_reply.started":"2025-12-14T03:41:47.210996Z","shell.execute_reply":"2025-12-14T03:41:47.236579Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Combining all these together","metadata":{}},{"cell_type":"code","source":"def topo_postprocess_core(\n    prob_volume: np.ndarray,\n    # hysteresis\n    T_low: float = 0.4,\n    T_high: float = 0.8,\n    # anisotropic closing\n    z_radius: int = 1,\n    xy_radius: int = 0,\n    dust_min_size: int = 100,\n    # surface smoothing\n    enable_smoothing: bool = True,\n    smoothing_sigma: float = 1.0,\n) -> np.ndarray:\n    \"\"\"\n    Topology-aware post-processing without expensive 3D skeletonization.\n\n    Stages:\n      1. Hysteresis thresholding\n      2. Anisotropic Z-closing\n      3. Dust removal\n      4. Optional voxel-space surface smoothing\n    \"\"\"\n\n    # 1) Hysteresis\n    mask = hysteresis_threshold_3d(\n        prob_volume.astype(np.float32),\n        T_low=T_low,\n        T_high=T_high,\n    )  # bool (Z,H,W)\n\n    # 2) Anisotropic closing along Z\n    mask = anisotropic_closing_z(\n        mask,\n        z_radius=z_radius,\n        xy_radius=xy_radius,\n        iterations=1,\n    )\n\n    # 3) Dust removal\n    mask = remove_small_components(\n        mask,\n        min_size=dust_min_size,\n    )\n\n    # 4) Surface smoothing (voxel-space)\n    if enable_smoothing:\n        mask = smooth_surface_voxel(\n            mask,\n            sigma=smoothing_sigma,\n        )\n\n    return mask.astype(np.uint8)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.239508Z","iopub.execute_input":"2025-12-14T03:41:47.239824Z","iopub.status.idle":"2025-12-14T03:41:47.261601Z","shell.execute_reply.started":"2025-12-14T03:41:47.239799Z","shell.execute_reply":"2025-12-14T03:41:47.260458Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Rolling inference","metadata":{}},{"cell_type":"code","source":"def make_positions(size, patch_size, stride):\n    \"\"\"Return list of starting positions that fully cover [0, size).\"\"\"\n    if size <= patch_size:\n        return [0]\n\n    positions = list(range(0, size - patch_size + 1, stride))\n\n    # Ensure we always hit the last possible start (size - patch_size)\n    last_start = size - patch_size\n    if positions[-1] != last_start:\n        positions.append(last_start)\n\n    return positions\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.262761Z","iopub.execute_input":"2025-12-14T03:41:47.263071Z","iopub.status.idle":"2025-12-14T03:41:47.285062Z","shell.execute_reply.started":"2025-12-14T03:41:47.263037Z","shell.execute_reply":"2025-12-14T03:41:47.284101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\n@torch.no_grad()\ndef sliding_window_inference_2d_batched(\n    model,\n    volume_3ch,\n    device,\n    patch_size=224,\n    stride=112,\n    batch_size=8,   # tune based on VRAM; try 8 or 4 on P100\n):\n    \"\"\"\n    volume_3ch: (3, H, W) float32 tensor.\n    Returns: prob_map (1, H, W) on device.\n    \"\"\"\n    model.eval()\n    volume_3ch = volume_3ch.to(device)\n    C, H, W = volume_3ch.shape\n    assert C == 3\n\n    prob_map   = torch.zeros(1, H, W, device=device)\n    weight_map = torch.zeros(1, H, W, device=device)\n\n    ys = make_positions(H, patch_size, stride)\n    xs = make_positions(W, patch_size, stride)\n\n    coords = [(y, x) for y in ys for x in xs]\n    n_patches = len(coords)\n\n    for start in range(0, n_patches, batch_size):\n        end = min(start + batch_size, n_patches)\n        batch_coords = coords[start:end]\n        bsz = len(batch_coords)\n\n        # Build batch (B, 3, ps, ps)\n        patches = torch.empty(\n            bsz, 3, patch_size, patch_size,\n            dtype=volume_3ch.dtype,\n            device=device,\n        )\n\n        for i, (y0, x0) in enumerate(batch_coords):\n            patches[i] = volume_3ch[:, y0:y0+patch_size, x0:x0+patch_size]\n\n        # Single forward for B patches\n        logits = model(\n            patches[:, 0:1, ...],\n            patches[:, 1:2, ...],\n            patches[:, 2:3, ...],\n        )  # (B, 1, ps, ps)\n        probs = torch.sigmoid(logits)  # (B, 1, ps, ps)\n\n        # Scatter back\n        for i, (y0, x0) in enumerate(batch_coords):\n            prob_map[:, y0:y0+patch_size, x0:x0+patch_size] += probs[i]\n            weight_map[:, y0:y0+patch_size, x0:x0+patch_size] += 1.0\n\n        del patches, logits, probs\n        torch.cuda.empty_cache()\n\n    weight_map = torch.clamp(weight_map, min=1.0)\n    prob_map = prob_map / weight_map\n    return prob_map  # (1, H, W)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.286246Z","iopub.execute_input":"2025-12-14T03:41:47.286556Z","iopub.status.idle":"2025-12-14T03:41:47.306854Z","shell.execute_reply.started":"2025-12-14T03:41:47.286525Z","shell.execute_reply":"2025-12-14T03:41:47.305595Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading Models","metadata":{}},{"cell_type":"markdown","source":"CSA-Net","metadata":{}},{"cell_type":"code","source":"def get_b16_config():\n    \"\"\"Returns the ViT-B/16 configuration.\"\"\"\n    config = ml_collections.ConfigDict()\n    config.patches = ml_collections.ConfigDict({'size': (16, 16)})\n    config.hidden_size = 768\n    config.transformer = ml_collections.ConfigDict()\n    config.transformer.mlp_dim = 3072\n    config.transformer.num_heads = 12\n    config.transformer.num_layers = 12\n    config.transformer.attention_dropout_rate = 0.0\n    config.transformer.dropout_rate = 0.1\n\n    config.classifier = 'seg'\n    config.representation_size = None\n    config.resnet_pretrained_path = None\n    config.pretrained_path = '../model/vit_checkpoint/imagenet21k/ViT-B_16.npz'\n    config.patch_size = 16\n\n    config.decoder_channels = (256, 128, 64, 16)\n    config.n_classes = 2\n    config.activation = 'softmax'\n    return config\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.308246Z","iopub.execute_input":"2025-12-14T03:41:47.308595Z","iopub.status.idle":"2025-12-14T03:41:47.355924Z","shell.execute_reply.started":"2025-12-14T03:41:47.308565Z","shell.execute_reply":"2025-12-14T03:41:47.354834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_r50_b16_config():\n    \"\"\"Returns the Resnet50 + ViT-B/16 configuration.\"\"\"\n    config = get_b16_config()\n    config.patches.grid = (16, 16)\n    config.resnet = ml_collections.ConfigDict()\n    config.resnet.num_layers = (3, 4, 9)\n    config.resnet.width_factor = 1\n\n    config.classifier = 'seg'\n    config.pretrained_path = '../model/vit_checkpoint/imagenet21k/R50+ViT-B_16.npz'\n    config.decoder_channels = (256, 128, 64, 16)\n    config.skip_channels = [512, 256, 64, 16]\n    config.n_classes = 2\n    config.n_skip = 3\n    config.activation = 'softmax'\n\n    return config\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.357105Z","iopub.execute_input":"2025-12-14T03:41:47.357449Z","iopub.status.idle":"2025-12-14T03:41:47.379081Z","shell.execute_reply.started":"2025-12-14T03:41:47.357405Z","shell.execute_reply":"2025-12-14T03:41:47.377790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = get_r50_b16_config()\nconfig.patches.grid = (PATCH_SIZE // 16, PATCH_SIZE // 16)\nconfig.n_classes    = 1\nconfig.activation   = None\n\n# 1) Create the model\nmodel = CSANet(config, img_size=PATCH_SIZE, num_classes=config.n_classes)\n\n# 2) Load checkpoint\nWEIGHTS_PATH = \"/kaggle/input/extended-crossattention/model_epoch_4_fullDice_0.7142.pth\"  \nckpt = torch.load(WEIGHTS_PATH, map_location=\"cpu\")\n\n\n# 3) Figure out how weights are stored\nif \"model\" in ckpt:\n    state_dict = ckpt[\"model\"]\nelif \"state_dict\" in ckpt:\n    state_dict = ckpt[\"state_dict\"]\nelif \"model_state_dict\" in ckpt:\n    state_dict = ckpt[\"model_state_dict\"]\nelse:\n    state_dict = ckpt  # raw state_dict\n\n# 4) If it was trained with DataParallel, strip \"module.\" prefix\nfrom collections import OrderedDict\nnew_sd = OrderedDict()\nfor k, v in state_dict.items():\n    new_k = k.replace(\"module.\", \"\", 1) if k.startswith(\"module.\") else k\n    new_sd[new_k] = v\n\n# 5) Finally load\nmissing, unexpected = model.load_state_dict(new_sd, strict=False)\nprint(\"Missing keys:\", missing)\nprint(\"Unexpected keys:\", unexpected)\n\nmodel = model.to(DEVICE)\nmodel.eval()\nprint(\"Model loaded.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:47.380295Z","iopub.execute_input":"2025-12-14T03:41:47.380568Z","iopub.status.idle":"2025-12-14T03:41:56.170296Z","shell.execute_reply.started":"2025-12-14T03:41:47.380548Z","shell.execute_reply":"2025-12-14T03:41:56.168907Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset and Dataloaders","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef infer_full_volume(\n    name,\n    model,\n    vol_np,\n    device,\n    patch_size=224,\n    stride=112,\n    batch_size=16,\n    # topo params (used by tuning_hyperparameters)\n    T_low=0.4,\n    T_high=0.8,\n    z_radius=1,\n    xy_radius=0,\n    dust_min_size=100,\n):\n    \"\"\"\n    Full-volume inference + topology-aware post-processing.\n    vol_np: np.ndarray (Z,H,W), typically uint8/uint16\n    Returns: pred_bin_volume: uint8 (Z,H,W)\n    \"\"\"\n    Z, H, W = vol_np.shape\n\n    # Store probabilities as float16 to save RAM\n    prob_volume = np.zeros((Z, H, W), dtype=np.float16)\n\n    # Adjust if your volume is 16-bit\n    norm_factor = 255.0\n\n    print(f\"  -> Input Shape: {vol_np.shape} | Stride: {stride} | Batch: {batch_size}\")\n    n_prev = 5\n    n_next = 5\n    for z in tqdm(range(Z), desc=\"Inference Z-slices\", leave=False):\n        prev_start = max(0, z - n_prev)\n        prev_end   = z \n        \n        # We want (z, z + n_next + 1) -> +1 because slice end is exclusive\n        next_start = min(Z, z + 1)\n        next_end   = min(Z, z + n_next + 1)\n    \n        # 2. Vectorized Fetching (No loops)\n        \n        # --- Current Slice ---\n        slice_curr = vol_np[z].astype(np.float32) / norm_factor\n    \n        # --- Previous Context ---\n        if prev_start < prev_end:\n            # Fetch block, take mean along Z axis (axis 0)\n            slice_prev = np.mean(vol_np[prev_start:prev_end], axis=0)\n            slice_prev = slice_prev.astype(np.float32) / norm_factor\n        else:\n            # Handle z=0 case (Zero padding or copy current)\n            slice_prev = slice_curr.copy()\n    \n        # --- Next Context ---\n        if next_start < next_end:\n            # Fetch block, take mean along Z axis\n            slice_next = np.mean(vol_np[next_start:next_end], axis=0)\n            slice_next = slice_next.astype(np.float32) / norm_factor\n        else:\n            # Handle z=Z-1 case\n            slice_next = slice_curr.copy()\n        vol_3ch = torch.from_numpy(\n            np.stack([slice_prev, slice_curr, slice_next], axis=0)\n        )  # (3,H,W) float32 CPU\n\n        # Sliding window → prob map (1,H,W) CPU float32\n        prob_map_tensor = sliding_window_inference_2d_batched(\n            model, vol_3ch, device, patch_size, stride, batch_size\n        )\n\n        prob_slice = prob_map_tensor.cpu().numpy()\n        if prob_slice.shape[0] == 1:\n            prob_slice = prob_slice[0]  # (H,W)\n\n        prob_volume[z] = prob_slice.astype(np.float16)\n        \n        # cleanup per slice\n        del vol_3ch, slice_prev, slice_curr, slice_next, prob_map_tensor, prob_slice\n    #     # === Debug: probability stats ===\n    # pv = prob_volume.astype(np.float32)\n    # print(\n    #     \"Prob stats:\",\n    #     \"min\", float(pv.min()),\n    #     \"max\", float(pv.max()),\n    #     \"mean\", float(pv.mean()),\n    #     \"p90\", float(np.percentile(pv, 90)),\n    #     \"p99\", float(np.percentile(pv, 99)),\n    # )\n    \n    # np.save(name,prob_volume)\n    # === Topology-aware post-processing ===\n    pred_bin_volume = topo_postprocess_core(\n        prob_volume,\n        **best_cfg,\n    ).astype(np.uint8)  # (Z,H,W)\n\n    del prob_volume\n    return pred_bin_volume\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.171861Z","iopub.execute_input":"2025-12-14T03:41:56.172407Z","iopub.status.idle":"2025-12-14T03:41:56.183555Z","shell.execute_reply.started":"2025-12-14T03:41:56.172379Z","shell.execute_reply":"2025-12-14T03:41:56.182291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pad_if_needed(vol_np, patch_size=224):\n    z, h, w = vol_np.shape\n    pad_h = max(0, patch_size - h)\n    pad_w = max(0, patch_size - w)\n    if pad_h > 0 or pad_w > 0:\n        vol_np = np.pad(vol_np, ((0, 0), (0, pad_h), (0, pad_w)), mode='constant', constant_values=0)\n    return vol_np, h, w\n\ndef load_tif_volume(path):\n    \"\"\"\n    Memory optimized loader. Pre-allocates array to avoid \n    memory spike during stacking.\n    \"\"\"\n    img = Image.open(path)\n    W, H = img.size\n    Z = getattr(img, 'n_frames', 1)\n    \n    print(f\"Loading {path} ({Z}x{H}x{W})...\")\n    \n    # Pre-allocate buffer (Standard RAM)\n    vol = np.zeros((Z, H, W), dtype=np.uint8)\n    \n    for i in range(Z):\n        try:\n            img.seek(i)\n            # Write directly to buffer\n            vol[i] = np.array(img, dtype=np.uint8)\n        except EOFError:\n            break\n            \n    return vol","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.184606Z","iopub.execute_input":"2025-12-14T03:41:56.184928Z","iopub.status.idle":"2025-12-14T03:41:56.208933Z","shell.execute_reply.started":"2025-12-14T03:41:56.184894Z","shell.execute_reply":"2025-12-14T03:41:56.207968Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" **Fine Tuning**","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport scipy.ndimage as ndi\n\ndef dice_score(pred: np.ndarray, gt: np.ndarray, eps: float = 1e-6) -> float:\n    \"\"\"\n    pred, gt: uint8 or bool, same shape (Z,H,W) or (H,W)\n    \"\"\"\n    pred = pred.astype(bool)\n    gt   = gt.astype(bool)\n\n    inter = np.logical_and(pred, gt).sum()\n    s_pred = pred.sum()\n    s_gt   = gt.sum()\n    return (2.0 * inter + eps) / (s_pred + s_gt + eps)\n\n\ndef voi_binary(pred: np.ndarray, gt: np.ndarray, eps: float = 1e-12) -> float:\n    \"\"\"\n    Variation of Information for binary segmentation (0/1 labels).\n    Any non-zero value in pred/gt is treated as 1.\n    Lower is better.\n    \"\"\"\n    # binarize first: anything >0 is foreground\n    pred = (pred > 0).astype(np.int32).ravel()\n    gt   = (gt > 0).astype(np.int32).ravel()\n\n    cm = np.zeros((2, 2), dtype=np.float64)\n    for p, g in zip(pred, gt):\n        cm[g, p] += 1.0\n\n    N = cm.sum()\n    if N == 0:\n        return 0.0\n\n    p_joint = cm / N\n    p_gt    = p_joint.sum(axis=1, keepdims=True)   # (2,1)\n    p_pred  = p_joint.sum(axis=0, keepdims=True)   # (1,2)\n\n    H_gt   = -np.sum(p_gt   * np.log2(p_gt   + eps))\n    H_pred = -np.sum(p_pred * np.log2(p_pred + eps))\n\n    denom = p_gt @ p_pred  # (2,2)\n    mask  = p_joint > 0\n    I = np.sum(p_joint[mask] * np.log2(p_joint[mask] / (denom[mask] + eps)))\n\n    voi = H_gt + H_pred - 2.0 * I\n    return float(voi)\n\n\n\ndef topology_proxy_b0(pred: np.ndarray, gt: np.ndarray) -> float:\n    \"\"\"\n    Simple topology proxy: difference in # of connected components (Betti-0-ish).\n    26-connectivity in 3D or 8-connectivity in 2D.\n    Lower is better (0 = perfect match).\n    \"\"\"\n    pred = pred.astype(bool)\n    gt   = gt.astype(bool)\n\n    if pred.ndim == 3:\n        struct = np.ones((3, 3, 3), dtype=bool)      # 26-connectivity\n    else:\n        struct = np.ones((3, 3), dtype=bool)         # 8-connectivity\n\n    _, n_pred = ndi.label(pred, structure=struct)\n    _, n_gt   = ndi.label(gt,   structure=struct)\n\n    return float(abs(n_pred - n_gt))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.209962Z","iopub.execute_input":"2025-12-14T03:41:56.210284Z","iopub.status.idle":"2025-12-14T03:41:56.233382Z","shell.execute_reply.started":"2025-12-14T03:41:56.210256Z","shell.execute_reply":"2025-12-14T03:41:56.232383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tifffile\n\ndef load_gt_mask_tif(path: str) -> np.ndarray:\n    \"\"\"\n    Simple GT loader; adapt if yours is npy or something else.\n    Returns uint8 or bool, shape (Z,H,W) or (H,W).\n    \"\"\"\n    gt = np.load(path)\n    # If GT is 2D but prediction is 3D, you can adapt later\n    return gt\n\n\ndef match_shape(pred_bin: np.ndarray, gt: np.ndarray) -> tuple[np.ndarray, np.ndarray]:\n    \"\"\"\n    Make shapes compatible by simple cropping along H,W. \n    (Assumes Z already matches or gt is 2D.)\n    \"\"\"\n    if pred_bin.ndim == 3 and gt.ndim == 2:\n        # project pred along Z (ink anywhere along depth)\n        pred_2d = (pred_bin > 0).any(axis=0).astype(np.uint8)\n        pred_bin = pred_2d\n\n    # Now shapes should be (H,W) or both (Z,H,W)\n    # Crop to min size along each spatial dimension\n    if pred_bin.ndim == 3:\n        Zp, Hp, Wp = pred_bin.shape\n        Zg, Hg, Wg = gt.shape\n        Z = min(Zp, Zg)\n        H = min(Hp, Hg)\n        W = min(Wp, Wg)\n        return pred_bin[:Z, :H, :W], gt[:Z, :H, :W]\n    else:\n        Hp, Wp = pred_bin.shape\n        Hg, Wg = gt.shape\n        H = min(Hp, Hg)\n        W = min(Wp, Wg)\n        return pred_bin[:H, :W], gt[:H, :W]\n\n\ndef evaluate_single_fragment(\n    fname,\n    model,\n    device,\n    vol_path,\n    gt_path,\n    patch_size,\n    stride,\n    batch_size,\n    topo_params,\n):\n    vol_np = np.load(vol_path)           # <--- here is the main fix\n    vol_np, orig_h, orig_w = pad_if_needed(vol_np, patch_size)\n\n    pred_bin = infer_full_volume(\n        fname,\n        model,\n        vol_np,\n        device,\n        patch_size=patch_size,\n        stride=stride,\n        batch_size=batch_size,\n        T_low=topo_params[\"T_low\"],\n        T_high=topo_params[\"T_high\"],\n        z_radius=topo_params[\"z_radius\"],\n        xy_radius=topo_params[\"xy_radius\"],\n        dust_min_size=topo_params[\"dust_min_size\"],\n    )\n\n    if orig_h < pred_bin.shape[1] or orig_w < pred_bin.shape[2]:\n        pred_bin = pred_bin[:, :orig_h, :orig_w]\n\n    # Load GT directly\n    gt = np.load(gt_path)               # <--- here also\n\n    # Shapes will match, so just for safety:\n    Z = min(pred_bin.shape[0], gt.shape[0])\n    H = min(pred_bin.shape[1], gt.shape[1])\n    W = min(pred_bin.shape[2], gt.shape[2])\n    pred_bin = pred_bin[:Z, :H, :W]\n    gt       = gt[:Z, :H, :W]\n\n    dsc    = dice_score(pred_bin, gt)\n    voi    = voi_binary(pred_bin, gt)\n    b0_err = topology_proxy_b0(pred_bin, gt)\n\n    return {\n        \"dice\": dsc,\n        \"voi\": voi,\n        \"b0_err\": b0_err,\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.234463Z","iopub.execute_input":"2025-12-14T03:41:56.234776Z","iopub.status.idle":"2025-12-14T03:41:56.259269Z","shell.execute_reply.started":"2025-12-14T03:41:56.234744Z","shell.execute_reply":"2025-12-14T03:41:56.257654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_single_config_on_val(\n    model,\n    device,\n    val_items: list,\n    patch_size: int,\n    stride: int,\n    batch_size: int,\n    topo_params: dict,\n) -> dict:\n    \"\"\"\n    val_items: list of dicts like:\n      {\n        \"id\": \"fragment01\",\n        \"vol_path\": \".../fragment01.tif\",\n        \"gt_path\":  \".../fragment01_gt.tif\",\n      }\n    Returns avg metrics over val_items.\n    \"\"\"\n    agg = {\n        \"dice\": 0.0,\n        \"voi\": 0.0,\n        \"b0_err\": 0.0,\n    }\n    n = len(val_items)\n\n    print(f\"\\nTesting config: {topo_params}\")\n    for item in val_items:\n        print(f\"  -> Fragment {item['id']}\")\n        m = evaluate_single_fragment(\n            model,\n            device,\n            vol_path=item[\"vol_path\"],\n            gt_path=item[\"gt_path\"],\n            patch_size=patch_size,\n            stride=stride,\n            batch_size=batch_size,\n            topo_params=topo_params,\n        )\n        for k in agg:\n            agg[k] += m[k]\n\n    for k in agg:\n        agg[k] /= max(1, n)\n\n    print(f\"  Avg Dice   : {agg['dice']:.4f}\")\n    print(f\"  Avg VOI    : {agg['voi']:.4f} (lower better)\")\n    print(f\"  Avg b0_err : {agg['b0_err']:.2f} (lower better)\")\n    return agg\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.263537Z","iopub.execute_input":"2025-12-14T03:41:56.263934Z","iopub.status.idle":"2025-12-14T03:41:56.284816Z","shell.execute_reply.started":"2025-12-14T03:41:56.263901Z","shell.execute_reply":"2025-12-14T03:41:56.283636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tuning_hyperparameters(\n    model,\n    device,\n    val_items: list,\n    patch_size: int,\n    stride: int,\n    batch_size: int,\n):\n    \"\"\"\n    Lightweight hyperparam sweep.\n    - Keep val_items small (e.g. 2-3 fragments).\n    - Each config runs full 3D inference over those fragments.\n    \"\"\"\n\n    # --- Very small search space on purpose ---\n    configs = [\n    # Slightly loose core, moderate band\n    {\"T_low\": 0.40, \"T_high\": 0.80, \"z_radius\": 1, \"xy_radius\": 0, \"dust_min_size\": 100},\n    {\"T_low\": 0.40, \"T_high\": 0.80, \"z_radius\": 2, \"xy_radius\": 0, \"dust_min_size\": 100},\n\n    # Stricter core, tighter band\n    {\"T_low\": 0.50, \"T_high\": 0.90, \"z_radius\": 1, \"xy_radius\": 0, \"dust_min_size\": 100},\n    {\"T_low\": 0.50, \"T_high\": 0.90, \"z_radius\": 2, \"xy_radius\": 0, \"dust_min_size\": 100},\n]\n\n    # You can tweak or reduce this further if runtime is too high.\n\n    results = []\n\n    for cfg in configs:\n        metrics = test_single_config_on_val(\n            model,\n            device,\n            val_items=val_items,\n            patch_size=patch_size,\n            stride=stride,\n            batch_size=batch_size,\n            topo_params=cfg,\n        )\n        results.append({\"config\": cfg, \"metrics\": metrics})\n\n    # Print summary sorted by Dice (desc)\n    print(\"\\n==== Hyperparameter Tuning Summary (sorted by Dice) ====\")\n    results_sorted = sorted(results, key=lambda x: x[\"metrics\"][\"dice\"], reverse=True)\n    for r in results_sorted:\n        cfg = r[\"config\"]\n        m   = r[\"metrics\"]\n        print(\n            f\"Config {cfg} -> \"\n            f\"Dice={m['dice']:.4f}, VOI={m['voi']:.4f}, b0_err={m['b0_err']:.2f}\"\n        )\n\n    return results_sorted\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.286284Z","iopub.execute_input":"2025-12-14T03:41:56.286667Z","iopub.status.idle":"2025-12-14T03:41:56.307124Z","shell.execute_reply.started":"2025-12-14T03:41:56.286642Z","shell.execute_reply":"2025-12-14T03:41:56.305806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Example small val set (KEEP IT SMALL)\n# val_items = [\n#     {\n#         \"id\": \"fragment01\",\n#         \"vol_path\": \"/kaggle/input/vesuvius-npy/train_images/1044587645.npy\",\n#         \"gt_path\":  \"/kaggle/input/vesuvius-npy/train_labels/1044587645.npy\",\n#     },\n#     {\n#         \"id\": \"fragment02\",\n#         \"vol_path\": \"/kaggle/input/vesuvius-npy/train_labels/1341999317.npy\",\n#         \"gt_path\":  \"/kaggle/input/vesuvius-npy/train_labels/1341999317.npy\",\n#     },\n# ]\n\n# results = tuning_hyperparameters(\n#     model,\n#     DEVICE,\n#     val_items=val_items,\n#     patch_size=PATCH_SIZE,\n#     stride=STRIDE,\n#     batch_size=BATCH_SIZE,\n# )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.308346Z","iopub.execute_input":"2025-12-14T03:41:56.308660Z","iopub.status.idle":"2025-12-14T03:41:56.330678Z","shell.execute_reply.started":"2025-12-14T03:41:56.308637Z","shell.execute_reply":"2025-12-14T03:41:56.329600Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# best_cfg = results[0][\"config\"]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.331654Z","iopub.execute_input":"2025-12-14T03:41:56.332026Z","iopub.status.idle":"2025-12-14T03:41:56.355515Z","shell.execute_reply.started":"2025-12-14T03:41:56.331988Z","shell.execute_reply":"2025-12-14T03:41:56.354556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(best_cfg)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.356703Z","iopub.execute_input":"2025-12-14T03:41:56.357059Z","iopub.status.idle":"2025-12-14T03:41:56.374387Z","shell.execute_reply.started":"2025-12-14T03:41:56.357034Z","shell.execute_reply":"2025-12-14T03:41:56.372966Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"\nTEST_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\"\nOUT_DIR  = \"/kaggle/working/predictions\"\nos.makedirs(OUT_DIR, exist_ok=True)\n\ntest_files = sorted([f for f in os.listdir(TEST_DIR) if f.lower().endswith(\".tif\")])\nprint(\"Found test tifs:\", test_files)\n# test_files = [\"1004283650.tif\",\"1006462223.tif\"] ## just for testing hyperparameters and all\nfor fname in test_files:\n    image_id = os.path.splitext(fname)[0]\n    in_path = os.path.join(TEST_DIR, fname)\n    \n    # Load\n    vol_np = load_tif_volume(in_path)\n    vol_np, orig_h, orig_w = pad_if_needed(vol_np, PATCH_SIZE)\n    \n    # Infer\n    pred_bin = infer_full_volume(\n        fname,\n        model, vol_np, DEVICE,\n        patch_size=PATCH_SIZE,\n        stride=STRIDE, # 112 (dense) or 200 (fast)\n        batch_size=BATCH_SIZE\n    )\n    \n    # Save & Cleanup\n    out_path = os.path.join(OUT_DIR, f\"{image_id}.tif\")\n    \n    # Crop back if we padded\n    if orig_h < pred_bin.shape[1] or orig_w < pred_bin.shape[2]:\n        pred_bin = pred_bin[:, :orig_h, :orig_w]\n        \n    tifffile.imwrite(out_path, pred_bin, compression=None)\n    \n    print(f\"Saved {out_path}\")\n    \n    # IMPORTANT: Free memory before next volume\n    del vol_np, pred_bin\n    gc.collect()\n    torch.cuda.empty_cache()\n    \n    \n    \n\nprint(\"All predictions saved in:\", OUT_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:41:56.375450Z","iopub.execute_input":"2025-12-14T03:41:56.376228Z","iopub.status.idle":"2025-12-14T03:42:00.444159Z","shell.execute_reply.started":"2025-12-14T03:41:56.376106Z","shell.execute_reply":"2025-12-14T03:42:00.442636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"All predictions saved in:\", OUT_DIR)\n\nimport zipfile\nimport os\n\n\nwith zipfile.ZipFile(\"submission.zip\", 'w', zipfile.ZIP_DEFLATED) as z:\n    for tif_name in os.listdir(OUT_DIR):\n        if tif_name.endswith(\".tif\"):\n            z.write(f\"{OUT_DIR}/{tif_name}\", tif_name)\nprint(\"✔ submission.zip created successfully!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.445052Z","iopub.status.idle":"2025-12-14T03:42:00.445420Z","shell.execute_reply.started":"2025-12-14T03:42:00.445237Z","shell.execute_reply":"2025-12-14T03:42:00.445252Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Comparision","metadata":{}},{"cell_type":"code","source":"# from PIL import Image, ImageSequence\n\n# def load_tif(path):\n#     img = Image.open(path)\n#     slices = []\n#     for page in ImageSequence.Iterator(img):\n#         slices.append(np.array(page))\n#     volume = np.stack(slices, axis=-1)  # H, W, Z\n#     return volume\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.447009Z","iopub.status.idle":"2025-12-14T03:42:00.447313Z","shell.execute_reply.started":"2025-12-14T03:42:00.447184Z","shell.execute_reply":"2025-12-14T03:42:00.447197Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import vesuvius_2025_metric_demo.vesuvius_2025_metric_demo as metric","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.448578Z","iopub.status.idle":"2025-12-14T03:42:00.448862Z","shell.execute_reply.started":"2025-12-14T03:42:00.448734Z","shell.execute_reply":"2025-12-14T03:42:00.448746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# PRED_DIR = \"/kaggle/working/predictions\"\n# SOLN_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection/train_labels\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.450487Z","iopub.status.idle":"2025-12-14T03:42:00.450748Z","shell.execute_reply.started":"2025-12-14T03:42:00.450630Z","shell.execute_reply":"2025-12-14T03:42:00.450641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def get_paths_csv(PRED_DIR, SOLN_DIR):\n#     metric.generate_standard_submission(SOLN_DIR)\n#     solution = pd.read_csv(\"/kaggle/working/submission.csv\")\n\n#     # 2) Generate prediction-style submission\n#     metric.generate_standard_submission(PRED_DIR)\n#     preds = pd.read_csv(\"/kaggle/working/submission.csv\")\n\n#     # 3) Keep only common ids\n#     common = set(solution[\"id\"]).intersection(preds[\"id\"])\n\n#     solution = solution[solution[\"id\"].isin(common)].copy()\n#     preds    = preds[preds[\"id\"].isin(common)].copy()\n\n#     # 4) Sort both by id and reset index so they are perfectly aligned\n#     solution = solution.sort_values(\"id\").reset_index(drop=True)\n#     preds    = preds.sort_values(\"id\").reset_index(drop=True)\n#     return solution,preds\n# soln, preds = get_paths_csv(PRED_DIR,SOLN_DIR)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.451697Z","iopub.status.idle":"2025-12-14T03:42:00.452009Z","shell.execute_reply.started":"2025-12-14T03:42:00.451844Z","shell.execute_reply":"2025-12-14T03:42:00.451856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# from ipywidgets import interact, IntSlider\n\n# def explore_side_by_side(gt, pred):\n#     \"\"\"\n#     Creates an interactive slider to scroll through the Z-axis.\n#     \"\"\"\n#     # 1. Determine the maximum depth based on your array shape (H, W, Depth)\n#     max_z = gt.shape[2] - 1\n    \n#     # 2. Define the callback function that runs every time the slider moves\n#     def plot_slice(z):\n#         fig, axs = plt.subplots(1, 2, figsize=(12, 4))\n        \n#         # Ground Truth\n#         axs[0].imshow(gt[:,:,z], cmap='gray')\n#         axs[0].set_title(f\"GT (z={z})\")\n        \n#         # Prediction\n#         axs[1].imshow(pred[:,:,z], cmap='gray')\n#         axs[1].set_title(f\"Pred (z={z})\")\n        \n#         for ax in axs: ax.axis('off')\n#         plt.show()\n\n#     # 3. Create the widget\n#     interact(plot_slice, \n#              z=IntSlider(min=0, max=max_z, step=1, value=max_z//2, description='Slice Z:'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.453236Z","iopub.status.idle":"2025-12-14T03:42:00.453617Z","shell.execute_reply.started":"2025-12-14T03:42:00.453431Z","shell.execute_reply":"2025-12-14T03:42:00.453447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def compare(PRED_DIR,SOLN_DIR):\n#     soln, preds = get_paths_csv(PRED_DIR,SOLN_DIR)\n#     for row in range(len(preds)):\n#         explore_side_by_side(load_tif(soln.loc[row][\"tif_paths\"]),load_tif(preds.loc[row][\"tif_paths\"]))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.454729Z","iopub.status.idle":"2025-12-14T03:42:00.455121Z","shell.execute_reply.started":"2025-12-14T03:42:00.454930Z","shell.execute_reply":"2025-12-14T03:42:00.454948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# compare(OUT_DIR,SOLN_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.456943Z","iopub.status.idle":"2025-12-14T03:42:00.457243Z","shell.execute_reply.started":"2025-12-14T03:42:00.457104Z","shell.execute_reply":"2025-12-14T03:42:00.457124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import numpy as np\n# import matplotlib.pyplot as plt\n# import torch\n# import cv2\n\n# def to_numpy(x):\n#     if isinstance(x, torch.Tensor):\n#         return x.detach().cpu().numpy()\n#     return np.asarray(x)\n\n# def show_test_overlay(input_stack, prediction, prob_threshold=0.5, title=\"\"):\n#     \"\"\"\n#     input_stack: (3, H, W)  - [prev, center, next] normalized like in your dataset\n#     prediction:  (H, W) or (1, H, W) or (H, W, 1), prob or binary\n#     \"\"\"\n#     inp = to_numpy(input_stack)\n#     pred = to_numpy(prediction)\n\n#     # center slice\n#     center_img = inp[99]  # (H, W)\n\n#     H, W = center_img.shape\n\n#     # normalize center slice for viz\n#     img_viz = (center_img - center_img.min()) / (center_img.max() - center_img.min() + 1e-8)\n#     img_gray = np.clip(img_viz, 0, 1)\n\n#     # squeeze prediction\n#     if pred.ndim == 3 and pred.shape[0] == 1:\n#         pred = pred[0]\n#     if pred.ndim == 3 and pred.shape[-1] == 1:\n#         pred = pred[..., 0]\n\n#     # resize prediction if needed\n#     if pred.shape != (H, W):\n#         pred = cv2.resize(pred, (W, H), interpolation=cv2.INTER_LINEAR)\n\n#     # if it's prob, keep as is; if it's 0/1 or logits, handle\n#     if pred.dtype != np.uint8 and pred.max() <= 1.0 and pred.min() >= 0.0:\n#         pred_prob = pred.astype(np.float32)\n#         pred_bin = (pred_prob >= prob_threshold).astype(np.uint8)\n#     else:\n#         # already binary-ish\n#         pred_bin = (pred != 0).astype(np.uint8)\n#         pred_prob = pred_bin.astype(np.float32)\n\n#     plt.figure(figsize=(6,6))\n#     plt.imshow(img_gray, cmap=\"gray\")\n#     # overlay prediction as red\n#     plt.imshow(pred_prob, cmap=\"Reds\", alpha=0.4)\n#     plt.title(title if title else \"Test image + prediction overlay\")\n#     plt.axis(\"off\")\n#     plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.458247Z","iopub.status.idle":"2025-12-14T03:42:00.458546Z","shell.execute_reply.started":"2025-12-14T03:42:00.458418Z","shell.execute_reply":"2025-12-14T03:42:00.458435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# v1 = load_tif_volume(\"/kaggle/working/predictions/1407735.tif\")\n# av1 = load_tif_volume(\"/kaggle/input/vesuvius-challenge-surface-detection/test_images/1407735.tif\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.459539Z","iopub.status.idle":"2025-12-14T03:42:00.459790Z","shell.execute_reply.started":"2025-12-14T03:42:00.459662Z","shell.execute_reply":"2025-12-14T03:42:00.459672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# show_test_overlay(av1,v1[99])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.461232Z","iopub.status.idle":"2025-12-14T03:42:00.461621Z","shell.execute_reply.started":"2025-12-14T03:42:00.461432Z","shell.execute_reply":"2025-12-14T03:42:00.461449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"raw","source":"","metadata":{}},{"cell_type":"code","source":"# v1 = load_tif_volume(\"/kaggle/input/vesuvius-challenge-surface-detection/train_labels/1004283650.tif\")\n# v2 = load_tif_volume(\"/kaggle/input/vesuvius-challenge-surface-detection/train_images/1006462223.tif\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.463586Z","iopub.status.idle":"2025-12-14T03:42:00.464054Z","shell.execute_reply.started":"2025-12-14T03:42:00.463783Z","shell.execute_reply":"2025-12-14T03:42:00.463800Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# v1.dtype","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.465076Z","iopub.status.idle":"2025-12-14T03:42:00.465463Z","shell.execute_reply.started":"2025-12-14T03:42:00.465333Z","shell.execute_reply":"2025-12-14T03:42:00.465346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# av1 = load_tif_volume(\"/kaggle/input/temper/1004283650.tif\")\n# av2 = load_tif_volume(\"/kaggle/input/temper/1006462223.tif\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.466487Z","iopub.status.idle":"2025-12-14T03:42:00.466746Z","shell.execute_reply.started":"2025-12-14T03:42:00.466627Z","shell.execute_reply":"2025-12-14T03:42:00.466638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import json\n# import os\n\n# def run_ipynb(path):\n#     \"\"\"\n#     Reads a raw .ipynb file and executes the code cells.\n#     \"\"\"\n#     if not os.path.exists(path):\n#         raise FileNotFoundError(f\"Could not find file at: {path}\")\n        \n#     print(f\"Executing notebook: {path}\")\n    \n#     with open(path, 'r', encoding='utf-8') as f:\n#         nb = json.load(f)\n        \n#     # Iterate over cells\n#     for cell in nb['cells']:\n#         if cell['cell_type'] == 'code':\n#             # Get the source code\n#             source = ''.join(cell['source'])\n            \n#             # Skip magic commands (like %time, !pip) which fail in exec()\n#             # or execute them if possible, but safer to skip lines starting with %, !\n#             cleaned_source = []\n#             for line in source.split('\\n'):\n#                 if not line.strip().startswith(('%', '!')):\n#                     cleaned_source.append(line)\n            \n#             cleaned_source = '\\n'.join(cleaned_source)\n            \n#             try:\n#                 # Execute in the current global scope\n#                 exec(cleaned_source, globals())\n#             except Exception as e:\n#                 print(f\"Error in cell execution: {e}\")\n    \n\n# # --- USE IT HERE ---\n# # Update this path to match exactly what you see in your error log or file tree\n\n# # Now you can use the function from that notebook\n# # e.g., metric = compute_metric(...)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.468303Z","iopub.status.idle":"2025-12-14T03:42:00.468554Z","shell.execute_reply.started":"2025-12-14T03:42:00.468439Z","shell.execute_reply":"2025-12-14T03:42:00.468449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(metric.score_single_tif(\"/kaggle/input/vesuvius-challenge-surface-detection/train_labels/1006462223.tif\",\"/kaggle/working/predictions/1006462223.tif\",2.0))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.469764Z","iopub.status.idle":"2025-12-14T03:42:00.470095Z","shell.execute_reply.started":"2025-12-14T03:42:00.469968Z","shell.execute_reply":"2025-12-14T03:42:00.469981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plt.imshow(av1[5,:,:],cmap = \"gray\" )\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.471565Z","iopub.status.idle":"2025-12-14T03:42:00.471934Z","shell.execute_reply.started":"2025-12-14T03:42:00.471733Z","shell.execute_reply":"2025-12-14T03:42:00.471744Z"}},"outputs":[],"execution_count":null},{"cell_type":"raw","source":"","metadata":{}},{"cell_type":"code","source":"# plt.imshow(v1[5,:,:],cmap = \"gray\" )\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T03:42:00.473500Z","iopub.status.idle":"2025-12-14T03:42:00.473780Z","shell.execute_reply.started":"2025-12-14T03:42:00.473654Z","shell.execute_reply":"2025-12-14T03:42:00.473666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}