{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"modelInstanceVersion","sourceId":748260,"databundleVersionId":15661884,"modelInstanceId":555553,"modelId":568111},{"sourceType":"modelInstanceVersion","sourceId":764664,"databundleVersionId":15819439,"modelInstanceId":555553,"modelId":568111}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import scipy\nimport numpy as np\nfrom scipy.ndimage import binary_dilation, binary_fill_holes, gaussian_filter, binary_erosion, distance_transform_edt, binary_closing, generate_binary_structure\nimport scipy.ndimage as ndi\nfrom scipy.interpolate import griddata\nfrom skimage.measure import label\nfrom skimage.morphology import ball\nfrom scipy.ndimage import median_filter\nfrom collections import deque\nfrom scipy.spatial import cKDTree\nfrom numba import jit, prange\nfrom concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor\nimport warnings\nfrom skimage.measure import euler_number\n\n# ============================================================================\n# NUMBA-OPTIMIZED RASTERIZATION (10-50x faster)\n# ============================================================================\n\n@jit(nopython=True, fastmath=True)\ndef rasterize_triangle_numba(p1, p2, p3, volume):\n    \"\"\"Numba-optimized triangle rasterization. Expected speedup: 10-50x.\"\"\"\n    min_z = max(0, int(np.floor(min(p1[0], p2[0], p3[0]))))\n    max_z = min(volume.shape[0] - 1, int(np.ceil(max(p1[0], p2[0], p3[0]))))\n    min_y = max(0, int(np.floor(min(p1[1], p2[1], p3[1]))))\n    max_y = min(volume.shape[1] - 1, int(np.ceil(max(p1[1], p2[1], p3[1]))))\n    min_x = max(0, int(np.floor(min(p1[2], p2[2], p3[2]))))\n    max_x = min(volume.shape[2] - 1, int(np.ceil(max(p1[2], p2[2], p3[2]))))\n\n    v0 = p2 - p1\n    v1 = p3 - p1\n    d00 = np.dot(v0, v0)\n    d01 = np.dot(v0, v1)\n    d11 = np.dot(v1, v1)\n    denom = d00 * d11 - d01 * d01\n\n    if abs(denom) < 1e-10:\n        return\n\n    inv_denom = 1.0 / denom\n\n    for z in range(min_z, max_z + 1):\n        for y in range(min_y, max_y + 1):\n            for x in range(min_x, max_x + 1):\n                v2_0 = z - p1[0]\n                v2_1 = y - p1[1]\n                v2_2 = x - p1[2]\n                d20 = v2_0 * v0[0] + v2_1 * v0[1] + v2_2 * v0[2]\n                d21 = v2_0 * v1[0] + v2_1 * v1[1] + v2_2 * v1[2]\n                v = (d11 * d20 - d01 * d21) * inv_denom\n                w = (d00 * d21 - d01 * d20) * inv_denom\n                u = 1.0 - v - w\n                if u >= -0.01 and v >= -0.01 and w >= -0.01:\n                    volume[z, y, x] = True\n\n\n@jit(nopython=True, fastmath=True, parallel=True)\ndef rasterize_surface_numba(grid_points, volume, samples_per_edge=5):\n    \"\"\"\n    Numba-optimized surface rasterization with parallel processing.\n    Expected speedup: 20-100x faster than pure Python.\n    \"\"\"\n    grid_resolution = grid_points.shape[0]\n\n    for i in prange(grid_resolution - 1):\n        for j in range(grid_resolution - 1):\n            p1 = grid_points[i, j]\n            p2 = grid_points[i+1, j]\n            p3 = grid_points[i, j+1]\n            p4 = grid_points[i+1, j+1]\n\n            if (np.isnan(p1[0]) or np.isnan(p2[0]) or\n                np.isnan(p3[0]) or np.isnan(p4[0])):\n                continue\n\n            for u_idx in range(samples_per_edge):\n                u = u_idx / (samples_per_edge - 1) if samples_per_edge > 1 else 0.5\n                for v_idx in range(samples_per_edge):\n                    v = v_idx / (samples_per_edge - 1) if samples_per_edge > 1 else 0.5\n\n                    point_0 = ((1-u)*(1-v)*p1[0] + u*(1-v)*p2[0] +\n                              (1-u)*v*p3[0] + u*v*p4[0])\n                    point_1 = ((1-u)*(1-v)*p1[1] + u*(1-v)*p2[1] +\n                              (1-u)*v*p3[1] + u*v*p4[1])\n                    point_2 = ((1-u)*(1-v)*p1[2] + u*(1-v)*p2[2] +\n                              (1-u)*v*p3[2] + u*v*p4[2])\n\n                    iz = int(np.round(point_0))\n                    iy = int(np.round(point_1))\n                    ix = int(np.round(point_2))\n\n                    if (0 <= iz < volume.shape[0] and\n                        0 <= iy < volume.shape[1] and\n                        0 <= ix < volume.shape[2]):\n                        volume[iz, iy, ix] = True\n\n# ============================================================================\n# ALGORITHMIC OPTIMIZATIONS\n# ============================================================================\n\ndef adaptive_grid_resolution(component, base_resolution=100, max_resolution=150):\n    \"\"\"Dynamically adjust grid resolution based on component size.\"\"\"\n    num_voxels = np.sum(component)\n    if num_voxels < 500:\n        return min(30, base_resolution)\n    elif num_voxels < 2000:\n        return min(50, base_resolution)\n    elif num_voxels < 5000:\n        return min(70, base_resolution)\n    elif num_voxels < 15000:\n        return base_resolution\n    else:\n        return min(max_resolution, base_resolution + 20)\n\n\ndef should_skip_smoothing(component, coverage_threshold=0.8):\n    \"\"\"Determine if a component needs smoothing based on planarity.\"\"\"\n    coords = np.column_stack(np.nonzero(component))\n    coords_mean = coords.mean(axis=0)\n    U, S, Vt = np.linalg.svd(coords - coords_mean, full_matrices=False)\n    if S[2] / S[0] < 0.05:\n        return True\n    return False\n\n\ndef zero_volume_faces(volume, thickness=5):\n    \"\"\"Optimized face zeroing using slicing.\"\"\"\n    result = volume.copy()\n    result[:thickness, :, :] = False\n    result[-thickness:, :, :] = False\n    result[:, :thickness, :] = False\n    result[:, -thickness:, :] = False\n    result[:, :, :thickness] = False\n    result[:, :, -thickness:] = False\n    return result\n\n\n# ============================================================================\n# OPTIMIZED MAIN FITTING FUNCTION\n# ============================================================================\n\ndef fit_curved_sheet_to_component_optimized(\n    component,\n    grid_resolution=100,\n    thickness=3,\n    smoothing=1.0,\n    use_median_filter=True,\n    max_distance=10,\n    use_numba=True,\n    adaptive_resolution=True,\n    samples_per_edge=8\n):\n    \"\"\"\n    OPTIMIZED version of fit_curved_sheet_to_component.\n    Key optimizations:\n    1. Numba JIT compilation for rasterization (10-50x speedup)\n    2. Adaptive grid resolution (2-4x speedup for small components)\n    3. Skip smoothing when not needed\n    \"\"\"\n    coords = np.column_stack(np.nonzero(component))\n    if len(coords) < 10:\n        return component.copy()\n\n    if adaptive_resolution:\n        grid_resolution = adaptive_grid_resolution(component, grid_resolution)\n        print(f\"    Using adaptive grid resolution: {grid_resolution}\")\n\n    coords_mean = coords.mean(axis=0)\n    U, S, Vt = np.linalg.svd(coords - coords_mean, full_matrices=False)\n    tangent1, tangent2 = Vt[0], Vt[1]\n    normal_guess = Vt[2]\n\n    uv_coords = (coords - coords_mean) @ np.column_stack([tangent1, tangent2])\n    w_coords = (coords - coords_mean) @ normal_guess\n\n    if len(coords) > 5000:\n        indices = np.random.choice(len(coords), 5000, replace=False)\n        uv_coords_sample = uv_coords[indices]\n        w_coords_sample = w_coords[indices]\n    else:\n        uv_coords_sample = uv_coords\n        w_coords_sample = w_coords\n\n    u_min, u_max = uv_coords[:,0].min(), uv_coords[:,0].max()\n    v_min, v_max = uv_coords[:,1].min(), uv_coords[:,1].max()\n    u_padding = (u_max - u_min) * 0.05\n    v_padding = (v_max - v_min) * 0.05\n\n    grid_u, grid_v = np.meshgrid(\n        np.linspace(u_min - u_padding, u_max + u_padding, num=grid_resolution),\n        np.linspace(v_min - v_padding, v_max + v_padding, num=grid_resolution),\n        indexing='ij'\n    )\n\n    try:\n        w_grid = griddata(uv_coords_sample, w_coords_sample, (grid_u, grid_v), method='linear')\n    except:\n        w_grid = griddata(uv_coords_sample, w_coords_sample, (grid_u, grid_v), method='nearest')\n\n    if np.any(np.isnan(w_grid)):\n        mask = np.isnan(w_grid)\n        w_grid_nearest = griddata(uv_coords_sample, w_coords_sample, (grid_u, grid_v), method='nearest')\n        w_grid[mask] = w_grid_nearest[mask]\n\n    if use_median_filter:\n        w_grid = median_filter(w_grid, size=3)\n\n    skip_smooth = should_skip_smoothing(component)\n    \n    if smoothing > 0 and not skip_smooth:\n        w_grid = gaussian_filter(w_grid, sigma=smoothing)\n    elif skip_smooth:\n        print(f\"    Skipping smoothing (component already planar)\")\n\n    # grid_padding = 0.02\n    # if grid_padding is not None:\n    #     u_data_min, u_data_max = uv_coords[:,0].min(), uv_coords[:,0].max()\n    #     v_data_min, v_data_max = uv_coords[:,1].min(), uv_coords[:,1].max()\n    #     u_range = u_data_max - u_data_min\n    #     v_range = v_data_max - v_data_min\n    #     u_pad = u_range * grid_padding\n    #     v_pad = v_range * grid_padding\n    #     grid_mask = ((grid_u >= u_data_min - u_pad) & (grid_u <= u_data_max + u_pad) &\n    #                  (grid_v >= v_data_min - v_pad) & (grid_v <= v_data_max + v_pad))\n    #     w_grid[~grid_mask] = np.nan\n\n    #Grid trimming with KDTree (already optimized in original)\n    tree = cKDTree(uv_coords)\n    threshold = (u_max - u_min + v_max - v_min) / (2 * grid_resolution) * 2\n    \n    grid_uv_flat = np.column_stack([grid_u.ravel(), grid_v.ravel()])\n    distances, _ = tree.query(grid_uv_flat, k=1)\n    distances = distances.reshape(grid_resolution, grid_resolution)\n    \n    original_data_mask = distances <= threshold\n    \n    # Flood-fill from edges\n    grid_mask = np.ones_like(w_grid, dtype=bool)\n    visited = np.zeros_like(w_grid, dtype=bool)\n    queue = deque()\n    \n    for i in range(grid_resolution):\n        queue.append((i, 0))\n        queue.append((i, grid_resolution - 1))\n        visited[i, 0] = True\n        visited[i, grid_resolution - 1] = True\n    \n    for j in range(1, grid_resolution - 1):\n        queue.append((0, j))\n        queue.append((grid_resolution - 1, j))\n        visited[0, j] = True\n        visited[grid_resolution - 1, j] = True\n    \n    while queue:\n        i, j = queue.popleft()\n        \n        if not original_data_mask[i, j]:\n            grid_mask[i, j] = False\n            \n            for di, dj in [(-1, 0), (1, 0), (0, -1), (0, 1)]:\n                ni, nj = i + di, j + dj\n                if (0 <= ni < grid_resolution and \n                    0 <= nj < grid_resolution and \n                    not visited[ni, nj]):\n                    visited[ni, nj] = True\n                    queue.append((ni, nj))\n    \n    w_grid[~grid_mask] = np.nan\n    \n    grid_points = (coords_mean +\n                   grid_u[...,None] * tangent1 +\n                   grid_v[...,None] * tangent2 +\n                   w_grid[...,None] * normal_guess)\n\n    Z, Y, X = component.shape\n    sheet_volume = np.zeros_like(component, dtype=bool)\n\n    if use_numba:\n        rasterize_surface_numba(grid_points, sheet_volume, samples_per_edge=samples_per_edge)\n    else:\n        rasterize_surface_dense_sampling_original(grid_points, sheet_volume, samples_per_quad=samples_per_edge)\n\n    sheet_volume = zero_volume_faces(sheet_volume, thickness=5)\n\n    if thickness > 0:\n        iterations = max(1, thickness // 2)\n        struct_elem = np.array([\n            [[0,1,0], [1,1,1], [0,1,0]],\n            [[1,1,1], [1,1,1], [1,1,1]],\n            [[0,1,0], [1,1,1], [0,1,0]]\n        ], dtype=bool)\n        sheet_volume = binary_dilation(sheet_volume, structure=struct_elem, iterations=iterations)\n\n    for z in range(Z):\n        if np.any(sheet_volume[z]):\n            sheet_volume[z] = binary_fill_holes(sheet_volume[z])\n\n    # struct = ndi.generate_binary_structure(3, 3)\n    # sheet_volume = ndi.binary_closing(sheet_volume, structure=struct, iterations=1)\n\n    return sheet_volume\n\n\ndef rasterize_surface_dense_sampling_original(grid_points, volume, samples_per_quad=5):\n    \"\"\"Original Python implementation for fallback.\"\"\"\n    grid_resolution = grid_points.shape[0]\n    for i in range(grid_resolution - 1):\n        for j in range(grid_resolution - 1):\n            p1 = grid_points[i, j]\n            p2 = grid_points[i+1, j]\n            p3 = grid_points[i, j+1]\n            p4 = grid_points[i+1, j+1]\n            if (np.isnan(p1).any() or np.isnan(p2).any() or\n                np.isnan(p3).any() or np.isnan(p4).any()):\n                continue\n            for u in np.linspace(0, 1, samples_per_quad):\n                for v in np.linspace(0, 1, samples_per_quad):\n                    point = ((1-u)*(1-v)*p1 + u*(1-v)*p2 + (1-u)*v*p3 + u*v*p4)\n                    point = point.round().astype(int)\n                    if (0 <= point[0] < volume.shape[0] and\n                        0 <= point[1] < volume.shape[1] and\n                        0 <= point[2] < volume.shape[2]):\n                        volume[point[0], point[1], point[2]] = True\n\n# ============================================================================\n# HELPER FUNCTIONS\n# ============================================================================\n\ndef calculate_dice_score(mask1, mask2):\n    \"\"\"Calculate Dice coefficient between two binary masks.\"\"\"\n    intersection = np.sum(mask1 & mask2)\n    sum_masks = np.sum(mask1) + np.sum(mask2)\n    if sum_masks == 0:\n        return 0.0\n    return 2.0 * intersection / sum_masks\n\n\ndef calculate_coverage_score(original, fitted):\n    \"\"\"Calculate how well the fitted sheet covers the original positive pixels.\"\"\"\n    original_pixels = np.sum(original)\n    if original_pixels == 0:\n        return 0.0\n    return np.sum(original & fitted) / original_pixels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-11T20:25:58.685022Z","iopub.execute_input":"2026-03-11T20:25:58.685387Z","iopub.status.idle":"2026-03-11T20:25:58.724213Z","shell.execute_reply.started":"2026-03-11T20:25:58.685356Z","shell.execute_reply":"2026-03-11T20:25:58.723553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom scipy.ndimage import binary_dilation, binary_fill_holes, gaussian_filter, binary_erosion, distance_transform_edt, binary_closing, generate_binary_structure\nimport scipy.ndimage as ndi\nfrom scipy.interpolate import griddata\nfrom skimage.measure import label\nfrom skimage.morphology import ball\nfrom scipy.ndimage import median_filter\nfrom collections import deque\nfrom scipy.spatial import cKDTree\nfrom numba import jit, prange\nfrom concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor\nimport warnings\nfrom skimage.measure import euler_number\n\n# [Keep all the NUMBA functions and other helpers from the original code]\n# ... (rasterize_triangle_numba, rasterize_surface_numba, etc.) ...\n\n# ============================================================================\n# PATCH-WISE (CUBE-BASED) TOPOLOGY CHECKING AND INTERPOLATION\n# ============================================================================\n\ndef divide_component_into_cubes(component, cube_size=32, overlap=8):\n    \"\"\"\n    Divide a component into smaller cubes with optional overlap.\n    \n    Parameters:\n    -----------\n    component : ndarray (bool)\n        Binary mask of the component\n    cube_size : int\n        Size of each cube dimension\n    overlap : int\n        Overlap between adjacent cubes (helps with boundary artifacts)\n    \n    Returns:\n    --------\n    list of dict\n        Each dict contains:\n        - 'coords': (z_start, z_end, y_start, y_end, x_start, x_end)\n        - 'center': (z_center, y_center, x_center)\n        - 'mask': boolean cube region from component\n        - 'at_volume_boundary': dict indicating which faces touch volume boundaries\n    \"\"\"\n    shape = component.shape\n    cubes = []\n    \n    # Calculate stride (cube_size - overlap ensures overlap between cubes)\n    stride = cube_size - overlap\n    \n    for z in range(0, shape[0], stride):\n        z_end = min(z + cube_size, shape[0])\n        if z_end - z < cube_size // 2:  # Skip tiny remainder cubes\n            continue\n            \n        for y in range(0, shape[1], stride):\n            y_end = min(y + cube_size, shape[1])\n            if y_end - y < cube_size // 2:\n                continue\n                \n            for x in range(0, shape[2], stride):\n                x_end = min(x + cube_size, shape[2])\n                if x_end - x < cube_size // 2:\n                    continue\n                \n                # Extract cube region\n                cube_mask = component[z:z_end, y:y_end, x:x_end]\n                \n                # Only include cubes that contain component voxels\n                if np.sum(cube_mask) > 10:  # Minimum voxel threshold\n                    # Detect which faces are at volume boundaries\n                    at_volume_boundary = {\n                        'z_min': (z == 0),\n                        'z_max': (z_end == shape[0]),\n                        'y_min': (y == 0),\n                        'y_max': (y_end == shape[1]),\n                        'x_min': (x == 0),\n                        'x_max': (x_end == shape[2])\n                    }\n                    \n                    cubes.append({\n                        'coords': (z, z_end, y, y_end, x, x_end),\n                        'center': (\n                            (z + z_end) // 2,\n                            (y + y_end) // 2,\n                            (x + x_end) // 2\n                        ),\n                        'mask': cube_mask,\n                        'shape': cube_mask.shape,\n                        'at_volume_boundary': at_volume_boundary\n                    })\n    \n    return cubes\n\n\ndef check_cube_topology(cube_mask, min_voxels=20):\n    \"\"\"\n    Check if a cube region has topological issues (holes).\n    \n    Parameters:\n    -----------\n    cube_mask : ndarray (bool)\n        Binary mask of cube region\n    min_voxels : int\n        Minimum voxels required to perform check\n    \n    Returns:\n    --------\n    tuple (has_hole, beta1, chi)\n        has_hole: True if beta1 > 0\n        beta1: First Betti number (number of holes)\n        chi: Euler characteristic\n    \"\"\"\n    if np.sum(cube_mask) < min_voxels:\n        return False, 0, 1\n    \n    try:\n        chi = euler_number(cube_mask.astype(int), connectivity=1)\n        beta1 = 1 - chi\n        has_hole = beta1 > 0\n        return has_hole, beta1, chi\n    except:\n        return False, 0, 1\n\n\ndef find_overlapping_cube_groups(cubes_with_holes):\n    \"\"\"\n    Group cubes that overlap or are adjacent into clusters.\n    Uses connected components on a cube adjacency graph.\n    \n    Parameters:\n    -----------\n    cubes_with_holes : list of dict\n        List of cube dictionaries that have holes\n    \n    Returns:\n    --------\n    list of list\n        Each inner list contains indices of cubes in a connected group\n    \"\"\"\n    n = len(cubes_with_holes)\n    if n == 0:\n        return []\n    if n == 1:\n        return [[0]]\n    \n    # Build adjacency matrix\n    adjacent = np.zeros((n, n), dtype=bool)\n    \n    for i in range(n):\n        for j in range(i + 1, n):\n            if cubes_overlap_or_adjacent(\n                cubes_with_holes[i]['coords'],\n                cubes_with_holes[j]['coords'],\n                adjacency_distance=2  # Consider adjacent if within 2 voxels\n            ):\n                adjacent[i, j] = True\n                adjacent[j, i] = True\n    \n    # Find connected components in adjacency graph\n    visited = np.zeros(n, dtype=bool)\n    groups = []\n    \n    for i in range(n):\n        if not visited[i]:\n            # BFS to find connected component\n            group = []\n            queue = [i]\n            visited[i] = True\n            \n            while queue:\n                idx = queue.pop(0)\n                group.append(idx)\n                \n                for j in range(n):\n                    if adjacent[idx, j] and not visited[j]:\n                        visited[j] = True\n                        queue.append(j)\n            \n            groups.append(group)\n    \n    return groups\n\n\ndef cubes_overlap_or_adjacent(coords1, coords2, adjacency_distance=2):\n    \"\"\"\n    Check if two cubes overlap or are adjacent (within distance threshold).\n    \n    Parameters:\n    -----------\n    coords1, coords2 : tuple\n        (z_start, z_end, y_start, y_end, x_start, x_end)\n    adjacency_distance : int\n        Maximum gap distance to consider cubes as adjacent\n    \n    Returns:\n    --------\n    bool\n        True if cubes overlap or are adjacent\n    \"\"\"\n    z1_s, z1_e, y1_s, y1_e, x1_s, x1_e = coords1\n    z2_s, z2_e, y2_s, y2_e, x2_s, x2_e = coords2\n    \n    # Check for overlap or adjacency in each dimension\n    z_overlap = not (z1_e + adjacency_distance < z2_s or z2_e + adjacency_distance < z1_s)\n    y_overlap = not (y1_e + adjacency_distance < y2_s or y2_e + adjacency_distance < y1_s)\n    x_overlap = not (x1_e + adjacency_distance < x2_s or x2_e + adjacency_distance < x2_s)\n    \n    return z_overlap and y_overlap and x_overlap\n\n\ndef merge_cube_regions(cubes, cube_indices, component_shape):\n    \"\"\"\n    Merge multiple cubes into a single bounding region.\n    Also merges the volume boundary information.\n    \n    Returns:\n    --------\n    tuple\n        (coords, at_volume_boundary) where:\n        - coords: (z_start, z_end, y_start, y_end, x_start, x_end) of merged region\n        - at_volume_boundary: dict of boundary flags for merged region\n    \"\"\"\n    z_min = component_shape[0]\n    z_max = 0\n    y_min = component_shape[1]\n    y_max = 0\n    x_min = component_shape[2]\n    x_max = 0\n    \n    # Merge boundaries - if ANY cube touches a volume boundary, the merged region does\n    at_volume_boundary = {\n        'z_min': False, 'z_max': False,\n        'y_min': False, 'y_max': False,\n        'x_min': False, 'x_max': False\n    }\n    \n    for idx in cube_indices:\n        z_s, z_e, y_s, y_e, x_s, x_e = cubes[idx]['coords']\n        z_min = min(z_min, z_s)\n        z_max = max(z_max, z_e)\n        y_min = min(y_min, y_s)\n        y_max = max(y_max, y_e)\n        x_min = min(x_min, x_s)\n        x_max = max(x_max, x_e)\n        \n        # Merge boundary flags\n        for key in at_volume_boundary:\n            at_volume_boundary[key] |= cubes[idx]['at_volume_boundary'][key]\n    \n    coords = (z_min, z_max, y_min, y_max, x_min, x_max)\n    return coords, at_volume_boundary","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-03-11T20:25:58.725371Z","iopub.execute_input":"2026-03-11T20:25:58.725611Z","iopub.status.idle":"2026-03-11T20:25:58.746852Z","shell.execute_reply.started":"2026-03-11T20:25:58.725589Z","shell.execute_reply":"2026-03-11T20:25:58.746246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def interpolate_cube_region_with_quality_check(\n    component,\n    cube_coords,\n    at_volume_boundary,\n    alternative_volumes=None,\n    border_thickness=4,\n    grid_resolution=80,\n    thickness=3,\n    smoothing=2.0,\n    max_distance=10,\n    samples_per_edge=8,\n    min_dice=0.65,\n    min_coverage=0.65,\n    alt_min_dice=0.60,\n    alt_min_coverage=0.60,\n    min_component_size=100\n):\n    \"\"\"\n    Interpolate a cube region with quality checking and alternative volumes.\n    \n    Strategy:\n    1. Extract and interpolate cube region\n    2. Check dice/coverage scores\n    3. If scores too low, try alternative volumes (pre-segmented at different thresholds)\n    4. Only replace interior voxels, but skip border preservation at volume edges\n    \n    Parameters:\n    -----------\n    component : ndarray (bool)\n        Full component mask\n    cube_coords : tuple\n        (z_start, z_end, y_start, y_end, x_start, x_end)\n    at_volume_boundary : dict\n        Flags indicating which cube faces touch volume boundaries\n    alternative_volumes : list of ndarray (bool), optional\n        Pre-computed segmentations at different thresholds to try if quality is low\n    border_thickness : int\n        Thickness of border to preserve (except at volume boundaries)\n    min_dice : float\n        Minimum Dice score required for main interpolation\n    min_coverage : float\n        Minimum coverage score required for main interpolation\n    alt_min_dice : float\n        Minimum Dice score required for alternative volumes\n    alt_min_coverage : float\n        Minimum coverage score required for alternative volumes\n    min_component_size : int\n        Minimum voxels for components to keep (default: 100)\n    \n    Returns:\n    --------\n    tuple (result_component, success, metrics)\n        result_component: Modified component\n        success: True if interpolation met quality thresholds\n        metrics: dict with dice, coverage, etc.\n    \"\"\"\n    z_s, z_e, y_s, y_e, x_s, x_e = cube_coords\n    \n    # Extract cube region\n    cube_mask = component[z_s:z_e, y_s:y_e, x_s:x_e].copy()\n    original_cube = cube_mask.copy()\n    \n    if np.sum(cube_mask) < 20:\n        return component, False, {'dice': 0.0, 'coverage': 0.0, 'reason': 'too_small'}\n    \n    print(f\"    Interpolating cube region: \"\n          f\"z[{z_s}:{z_e}], y[{y_s}:{y_e}], x[{x_s}:{x_e}]\")\n    \n    # ========================================================================\n    # ATTEMPT 1: Try original mask\n    # ========================================================================\n    print(f\"      Attempt 1: Original mask\")\n    \n    try:\n        interpolated_cube = fit_curved_sheet_to_component_optimized(\n            cube_mask,\n            grid_resolution=grid_resolution,\n            thickness=thickness,\n            smoothing=smoothing,\n            max_distance=max_distance,\n            use_numba=True,\n            adaptive_resolution=True,\n            samples_per_edge=samples_per_edge\n        )\n    except Exception as e:\n        print(f\"      Failed to interpolate: {e}\")\n        interpolated_cube = None\n    \n    if interpolated_cube is not None:\n        dice = calculate_dice_score(original_cube, interpolated_cube)\n        coverage = calculate_coverage_score(original_cube, interpolated_cube)\n        print(f\"      Dice={dice:.3f}, Coverage={coverage:.3f}\")\n        \n        if dice >= min_dice and coverage >= min_coverage:\n            print(f\"      ✓ Quality threshold met!\")\n            best_result = interpolated_cube\n            best_metrics = {'dice': dice, 'coverage': coverage, 'attempt': 0, 'num_parts': 1}\n            # Skip to replacement section\n            success = True\n            interpolated_cube = best_result\n            final_dice = best_metrics['dice']\n            final_coverage = best_metrics['coverage']\n            \n            # Jump to interior mask creation\n            cube_shape = cube_mask.shape\n            interior_mask = np.ones(cube_shape, dtype=bool)\n            \n            # ... (rest of interior mask code below)\n            # For brevity, let me continue from here\n        else:\n            print(f\"      ✗ Quality too low, trying alternative volumes...\")\n            best_result = interpolated_cube\n            best_metrics = {'dice': dice, 'coverage': coverage, 'attempt': 0, 'num_parts': 1}\n    else:\n        best_result = None\n        best_metrics = {'dice': 0.0, 'coverage': 0.0}\n    \n    # ========================================================================\n    # ATTEMPT 2+: Try alternative volumes (if quality was low)\n    # ========================================================================\n    if best_metrics['dice'] < min_dice or best_metrics['coverage'] < min_coverage:\n        if alternative_volumes is None or len(alternative_volumes) == 0:\n            print(f\"      No alternative volumes available\")\n            if best_result is None:\n                return component, False, {'dice': 0.0, 'coverage': 0.0, 'reason': 'no_alternatives'}\n        else:\n            print(f\"      Trying {len(alternative_volumes)} alternative volume(s)...\")\n            \n            for alt_idx, alt_volume in enumerate(alternative_volumes):\n                print(f\"      Attempt {alt_idx + 2}: Alternative volume {alt_idx + 1}\")\n                \n                # Extract cube region from alternative volume\n                alt_cube = alt_volume[z_s:z_e, y_s:y_e, x_s:x_e]\n                \n                # Find overlap with original cube mask\n                alt_mask = alt_cube & original_cube\n                \n                if not np.any(alt_mask):\n                    print(f\"        No overlap with original mask\")\n                    continue\n                \n                # Label components in alternative within this region\n                alt_labeled = label(alt_mask)\n                num_alt_comps = alt_labeled.max()\n                \n                if num_alt_comps == 0:\n                    print(f\"        No components found\")\n                    continue\n                \n                print(f\"        Found {num_alt_comps} component(s)\")\n                \n                # Filter out small components\n                valid_components = []\n                for comp_id in range(1, num_alt_comps + 1):\n                    comp_mask = (alt_labeled == comp_id)\n                    comp_size = np.sum(comp_mask)\n                    \n                    if comp_size < min_component_size:\n                        print(f\"          Component {comp_id}: Too small ({comp_size} vx)\")\n                        continue\n                    \n                    valid_components.append((comp_id, comp_mask, comp_size))\n                \n                if len(valid_components) == 0:\n                    print(f\"        No valid components (all <{min_component_size} voxels)\")\n                    continue\n                \n                print(f\"        Processing {len(valid_components)} valid component(s)\")\n                \n                # Try interpolating each valid component\n                interpolated_parts = []\n                for comp_id, comp_mask, comp_size in valid_components:\n                    print(f\"          Component {comp_id}: {comp_size} voxels\", end=\"\")\n                    \n                    try:\n                        comp_fitted = fit_curved_sheet_to_component_optimized(\n                            comp_mask,\n                            grid_resolution=grid_resolution,\n                            thickness=thickness,\n                            smoothing=smoothing,\n                            max_distance=max_distance,\n                            use_numba=True,\n                            adaptive_resolution=True,\n                            samples_per_edge=samples_per_edge\n                        )\n                        \n                        comp_dice = calculate_dice_score(comp_mask, comp_fitted)\n                        comp_coverage = calculate_coverage_score(comp_mask, comp_fitted)\n                        print(f\" → Dice={comp_dice:.3f}, Cov={comp_coverage:.3f}\", end=\"\")\n                        \n                        if comp_dice >= alt_min_dice and comp_coverage >= alt_min_coverage:\n                            interpolated_parts.append(comp_fitted)\n                            print(f\" ✓\")\n                        else:\n                            print(f\" ✗\")\n                    \n                    except Exception as e:\n                        print(f\" ✗ ({e})\")\n                \n                if len(interpolated_parts) == 0:\n                    print(f\"        No parts met quality threshold\")\n                    continue\n                \n                # Combine all interpolated parts\n                combined_mask = np.zeros_like(cube_mask, dtype=bool)\n                for part in interpolated_parts:\n                    combined_mask |= part\n                \n                # Check overall quality\n                dice = calculate_dice_score(original_cube, combined_mask)\n                coverage = calculate_coverage_score(original_cube, combined_mask)\n                print(f\"        Combined: Dice={dice:.3f}, Coverage={coverage:.3f}, \"\n                      f\"{len(interpolated_parts)} part(s)\")\n                \n                if dice > best_metrics['dice']:\n                    best_result = combined_mask\n                    best_metrics = {\n                        'dice': dice,\n                        'coverage': coverage,\n                        'attempt': alt_idx + 2,\n                        'num_parts': len(interpolated_parts),\n                        'alt_idx': alt_idx\n                    }\n                    print(f\"        ✓ New best result!\")\n                \n                # If we meet the threshold, we can stop\n                if dice >= alt_min_dice and coverage >= alt_min_coverage:\n                    print(f\"        ✓ Alternative threshold met, stopping search\")\n                    break\n    \n    # ========================================================================\n    # Use best result found\n    # ========================================================================\n    if best_result is None:\n        return component, False, {'dice': 0.0, 'coverage': 0.0, 'reason': 'no_valid_result'}\n    \n    interpolated_cube = best_result\n    final_dice = best_metrics['dice']\n    final_coverage = best_metrics['coverage']\n    \n    print(f\"      Final: Dice={final_dice:.3f}, Coverage={final_coverage:.3f} \"\n          f\"(attempt {best_metrics.get('attempt', 0)}, {best_metrics.get('num_parts', 1)} part(s))\")\n    \n    # ========================================================================\n    # Create interior mask (exclude border_thickness from non-boundary faces)\n    # ========================================================================\n    cube_shape = cube_mask.shape\n    interior_mask = np.ones(cube_shape, dtype=bool)\n    \n    # Z dimension borders\n    if cube_shape[0] > 2 * border_thickness:\n        if not at_volume_boundary['z_min']:\n            interior_mask[:border_thickness, :, :] = False\n        if not at_volume_boundary['z_max']:\n            interior_mask[-border_thickness:, :, :] = False\n    else:\n        bt = max(1, cube_shape[0] // 4)\n        if not at_volume_boundary['z_min']:\n            interior_mask[:bt, :, :] = False\n        if not at_volume_boundary['z_max']:\n            interior_mask[-bt:, :, :] = False\n    \n    # Y dimension borders\n    if cube_shape[1] > 2 * border_thickness:\n        if not at_volume_boundary['y_min']:\n            interior_mask[:, :border_thickness, :] = False\n        if not at_volume_boundary['y_max']:\n            interior_mask[:, -border_thickness:, :] = False\n    else:\n        bt = max(1, cube_shape[1] // 4)\n        if not at_volume_boundary['y_min']:\n            interior_mask[:, :bt, :] = False\n        if not at_volume_boundary['y_max']:\n            interior_mask[:, -bt:, :] = False\n    \n    # X dimension borders\n    if cube_shape[2] > 2 * border_thickness:\n        if not at_volume_boundary['x_min']:\n            interior_mask[:, :, :border_thickness] = False\n        if not at_volume_boundary['x_max']:\n            interior_mask[:, :, -border_thickness:] = False\n    else:\n        bt = max(1, cube_shape[2] // 4)\n        if not at_volume_boundary['x_min']:\n            interior_mask[:, :, :bt] = False\n        if not at_volume_boundary['x_max']:\n            interior_mask[:, :, -bt:] = False\n    \n    # Replace ONLY interior voxels in the original component\n    result = component.copy()\n    \n    # Clear interior region\n    result[z_s:z_e, y_s:y_e, x_s:x_e][interior_mask] = False\n    \n    # Place interpolated interior\n    result[z_s:z_e, y_s:y_e, x_s:x_e][interior_mask] = interpolated_cube[interior_mask]\n    \n    interior_voxels = np.sum(interior_mask)\n    replaced_voxels = np.sum(interpolated_cube[interior_mask])\n    print(f\"      Replaced {replaced_voxels}/{interior_voxels} interior voxels\")\n    \n    # Report boundary preservation\n    boundaries_preserved = []\n    if not at_volume_boundary['z_min']:\n        boundaries_preserved.append('z_min')\n    if not at_volume_boundary['z_max']:\n        boundaries_preserved.append('z_max')\n    if not at_volume_boundary['y_min']:\n        boundaries_preserved.append('y_min')\n    if not at_volume_boundary['y_max']:\n        boundaries_preserved.append('y_max')\n    if not at_volume_boundary['x_min']:\n        boundaries_preserved.append('x_min')\n    if not at_volume_boundary['x_max']:\n        boundaries_preserved.append('x_max')\n    \n    if boundaries_preserved:\n        print(f\"      Preserved borders: {', '.join(boundaries_preserved)}\")\n    else:\n        print(f\"      No borders preserved (all at volume boundary)\")\n    \n    success = (final_dice >= min_dice and final_coverage >= min_coverage) or \\\n              (final_dice >= alt_min_dice and final_coverage >= alt_min_coverage)\n    \n    metrics = {\n        'dice': final_dice,\n        'coverage': final_coverage,\n        'attempt': best_metrics.get('attempt', 0),\n        'num_parts': best_metrics.get('num_parts', 1),\n        'replaced_voxels': replaced_voxels,\n        'interior_voxels': interior_voxels\n    }\n    \n    return result, success, metrics\n\n\ndef process_component_patchwise(\n    component,\n    component_id,\n    alternative_volumes=None,\n    cube_size=32,\n    overlap=8,\n    border_thickness=4,\n    grid_resolution=80,\n    thickness=3,\n    smoothing=2.0,\n    max_distance=10,\n    samples_per_edge=8,\n    min_cube_voxels=20,\n    min_dice=0.65,\n    min_coverage=0.65,\n    alt_min_dice=0.60,\n    alt_min_coverage=0.60,\n    min_component_size=100\n):\n    \"\"\"\n    Main function: Process a component using patch-wise topology checking.\n    \n    Workflow:\n    1. Divide component into cubes\n    2. Check each cube for holes (Euler number)\n    3. Group overlapping/adjacent cubes with holes\n    4. Interpolate grouped regions (or single cubes) with quality checking\n    5. If quality low, try alternative volumes (pre-segmented at different thresholds)\n    6. Filter out small components (<100 voxels) from alternatives\n    7. Only replace interior, preserve borders for connectivity (except at volume edges)\n    \n    Parameters:\n    -----------\n    component : ndarray (bool)\n        Binary mask of component to process\n    component_id : int\n        ID for logging\n    alternative_volumes : list of ndarray (bool), optional\n        Pre-computed segmentations at different thresholds\n    cube_size : int\n        Size of cube patches\n    overlap : int\n        Overlap between adjacent cubes\n    border_thickness : int\n        Border region to preserve (voxels from each face, except at volume boundary)\n    min_dice : float\n        Minimum Dice score required for main interpolation\n    min_coverage : float\n        Minimum coverage score required for main interpolation\n    alt_min_dice : float\n        Minimum Dice score for alternative volumes\n    alt_min_coverage : float\n        Minimum coverage score for alternative volumes\n    min_component_size : int\n        Minimum voxels for components to keep (default: 100)\n    \n    Returns:\n    --------\n    tuple (result, success_rate, metrics)\n        result: Modified component with holes interpolated\n        success_rate: Fraction of problematic cubes successfully interpolated\n        metrics: Summary statistics\n    \"\"\"\n    print(f\"\\n{'='*70}\")\n    print(f\"Processing component {component_id} (patch-wise)\")\n    print(f\"{'='*70}\")\n    \n    # Step 1: Divide into cubes\n    print(f\"Step 1: Dividing into cubes (size={cube_size}, overlap={overlap})...\")\n    cubes = divide_component_into_cubes(component, cube_size=cube_size, overlap=overlap)\n    print(f\"  Generated {len(cubes)} cubes\")\n    \n    if len(cubes) == 0:\n        print(\"  No cubes generated (component too small?)\")\n        return component, 0.0, {'total_cubes': 0}\n    \n    # Step 2: Check topology of each cube\n    print(f\"\\nStep 2: Checking topology of each cube...\")\n    cubes_with_holes = []\n    cubes_with_holes_indices = []\n    \n    for i, cube_info in enumerate(cubes):\n        has_hole, beta1, chi = check_cube_topology(cube_info['mask'], min_voxels=min_cube_voxels)\n        \n        if has_hole:\n            print(f\"  Cube {i} at {cube_info['center']}: β1={beta1} (χ={chi}) ⚠ HAS HOLE\")\n            cubes_with_holes.append(cube_info)\n            cubes_with_holes_indices.append(i)\n        else:\n            print(f\"  Cube {i} at {cube_info['center']}: β1={beta1} (χ={chi}) ✓ OK\")\n    \n    if len(cubes_with_holes) == 0:\n        print(\"\\n✓ No cubes have holes - component is topologically correct!\")\n        return component, 1.0, {\n            'total_cubes': len(cubes),\n            'problematic_cubes': 0,\n            'successful_interpolations': 0\n        }\n    \n    print(f\"\\nFound {len(cubes_with_holes)} cube(s) with holes\")\n    \n    # Step 3: Group overlapping/adjacent cubes\n    print(f\"\\nStep 3: Grouping overlapping/adjacent cubes...\")\n    groups = find_overlapping_cube_groups(cubes_with_holes)\n    print(f\"  Found {len(groups)} group(s)\")\n    \n    for g_idx, group in enumerate(groups):\n        print(f\"  Group {g_idx + 1}: {len(group)} cube(s)\")\n    \n    # Step 4: Interpolate each group with quality checking\n    print(f\"\\nStep 4: Interpolating cube groups with quality checking...\")\n    result = component.copy()\n    successful_interpolations = 0\n    failed_interpolations = 0\n    all_metrics = []\n    \n    for g_idx, group_indices in enumerate(groups):\n        print(f\"\\n  Group {g_idx + 1}/{len(groups)}: {len(group_indices)} cube(s)\")\n        \n        if len(group_indices) == 1:\n            # Single cube - interpolate just that cube\n            cube_idx = group_indices[0]\n            cube_coords = cubes_with_holes[cube_idx]['coords']\n            at_boundary = cubes_with_holes[cube_idx]['at_volume_boundary']\n            \n            print(f\"    Single cube - interpolating independently\")\n            result, success, metrics = interpolate_cube_region_with_quality_check(\n                result,\n                cube_coords,\n                at_boundary,\n                alternative_volumes=alternative_volumes,\n                border_thickness=border_thickness,\n                grid_resolution=grid_resolution,\n                thickness=thickness,\n                smoothing=smoothing,\n                max_distance=max_distance,\n                samples_per_edge=samples_per_edge,\n                min_dice=min_dice,\n                min_coverage=min_coverage,\n                alt_min_dice=alt_min_dice,\n                alt_min_coverage=alt_min_coverage,\n                min_component_size=min_component_size\n            )\n            \n            all_metrics.append(metrics)\n            if success:\n                successful_interpolations += 1\n                print(f\"    ✓ Interpolation successful\")\n            else:\n                failed_interpolations += 1\n                print(f\"    ✗ Interpolation failed or low quality\")\n        else:\n            # Multiple cubes - merge into single region\n            merged_coords, merged_boundary = merge_cube_regions(\n                cubes_with_holes,\n                group_indices,\n                component.shape\n            )\n            \n            print(f\"    Multiple cubes - merging into region: \"\n                  f\"z[{merged_coords[0]}:{merged_coords[1]}], \"\n                  f\"y[{merged_coords[2]}:{merged_coords[3]}], \"\n                  f\"x[{merged_coords[4]}:{merged_coords[5]}]\")\n            \n            result, success, metrics = interpolate_cube_region_with_quality_check(\n                result,\n                merged_coords,\n                merged_boundary,\n                alternative_volumes=alternative_volumes,\n                border_thickness=border_thickness,\n                grid_resolution=grid_resolution,\n                thickness=thickness,\n                smoothing=smoothing,\n                max_distance=max_distance,\n                samples_per_edge=samples_per_edge,\n                min_dice=min_dice,\n                min_coverage=min_coverage,\n                alt_min_dice=alt_min_dice,\n                alt_min_coverage=alt_min_coverage,\n                min_component_size=min_component_size\n            )\n            \n            all_metrics.append(metrics)\n            if success:\n                successful_interpolations += 1\n                print(f\"    ✓ Interpolation successful\")\n            else:\n                failed_interpolations += 1\n                print(f\"    ✗ Interpolation failed or low quality\")\n    \n    success_rate = successful_interpolations / len(groups) if len(groups) > 0 else 0.0\n    \n    summary_metrics = {\n        'total_cubes': len(cubes),\n        'problematic_cubes': len(cubes_with_holes),\n        'groups': len(groups),\n        'successful_interpolations': successful_interpolations,\n        'failed_interpolations': failed_interpolations,\n        'success_rate': success_rate,\n        'avg_dice': np.mean([m['dice'] for m in all_metrics]) if all_metrics else 0.0,\n        'avg_coverage': np.mean([m['coverage'] for m in all_metrics]) if all_metrics else 0.0\n    }\n    \n    print(f\"\\n✓ Component {component_id} processing complete\")\n    print(f\"  Success rate: {success_rate:.1%} ({successful_interpolations}/{len(groups)} groups)\")\n    print(f\"  Avg Dice: {summary_metrics['avg_dice']:.3f}\")\n    print(f\"  Avg Coverage: {summary_metrics['avg_coverage']:.3f}\")\n    \n    return result, success_rate, summary_metrics\n\n\n# ============================================================================\n# INTEGRATION WITH EXISTING PIPELINE\n# ============================================================================\n\ndef process_multiple_components_patchwise(\n    volume,\n    alternative_volumes=None,\n    cube_size=32,\n    overlap=8,\n    border_thickness=4,\n    grid_resolution=80,\n    thickness=3,\n    smoothing=2.0,\n    max_distance=10,\n    samples_per_edge=8,\n    min_dice=0.65,\n    min_coverage=0.65,\n    alt_min_dice=0.60,\n    alt_min_coverage=0.60,\n    min_component_size=100,\n    use_parallel=True,\n    n_jobs=-1\n):\n    \"\"\"\n    Process all components in a volume using patch-wise approach with quality checking.\n    \n    Parameters:\n    -----------\n    volume : ndarray (bool)\n        Binary volume to process\n    alternative_volumes : list of ndarray (bool), optional\n        Pre-computed segmentations at different thresholds (e.g., different sigma values)\n        These are used when main interpolation quality is too low\n    cube_size : int\n        Size of cube patches\n    overlap : int\n        Overlap between adjacent cubes\n    min_dice : float\n        Minimum Dice score for main interpolation\n    min_coverage : float\n        Minimum coverage score for main interpolation\n    alt_min_dice : float\n        Minimum Dice score for alternative volumes\n    alt_min_coverage : float\n        Minimum coverage score for alternative volumes\n    min_component_size : int\n        Minimum voxels for components in alternative volumes (filters out noise)\n    \"\"\"\n    print(f\"\\n{'='*70}\")\n    print(f\"PATCH-WISE COMPONENT PROCESSING WITH QUALITY CHECKING\")\n    print(f\"{'='*70}\")\n    \n    labeled_volume = label(volume)\n    num_components = labeled_volume.max()\n    \n    print(f\"Found {num_components} components\")\n    print(f\"Cube parameters: size={cube_size}, overlap={overlap}, border={border_thickness}\")\n    print(f\"Main thresholds: Dice≥{min_dice:.2f}, Coverage≥{min_coverage:.2f}\")\n    print(f\"Alternative thresholds: Dice≥{alt_min_dice:.2f}, Coverage≥{alt_min_coverage:.2f}\")\n    print(f\"Min component size: {min_component_size} voxels\")\n    if alternative_volumes is not None:\n        print(f\"Using {len(alternative_volumes)} alternative volume(s)\")\n    \n    result_labeled = np.zeros_like(volume, dtype=np.int32)\n    all_success_rates = []\n    all_metrics = []\n    \n    for i in range(1, num_components + 1):\n        component_mask = (labeled_volume == i)\n        \n        # Process component patch-wise\n        processed_component, success_rate, metrics = process_component_patchwise(\n            component_mask,\n            component_id=i,\n            alternative_volumes=alternative_volumes,\n            cube_size=cube_size,\n            overlap=overlap,\n            border_thickness=border_thickness,\n            grid_resolution=grid_resolution,\n            thickness=thickness,\n            smoothing=smoothing,\n            max_distance=max_distance,\n            samples_per_edge=samples_per_edge,\n            min_dice=min_dice,\n            min_coverage=min_coverage,\n            alt_min_dice=alt_min_dice,\n            alt_min_coverage=alt_min_coverage,\n            min_component_size=min_component_size\n        )\n        \n        # Assign to result\n        result_labeled[processed_component] = i\n        all_success_rates.append(success_rate)\n        all_metrics.append(metrics)\n    \n    result_binary = result_labeled > 0\n    \n    # Overall statistics\n    avg_success_rate = np.mean(all_success_rates) if all_success_rates else 0.0\n    total_problematic_cubes = sum(m.get('problematic_cubes', 0) for m in all_metrics)\n    total_successful = sum(m.get('successful_interpolations', 0) for m in all_metrics)\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"PROCESSING COMPLETE\")\n    print(f\"{'='*70}\")\n    print(f\"Components processed: {num_components}\")\n    print(f\"Final components: {label(result_binary).max()}\")\n    print(f\"Total problematic cubes: {total_problematic_cubes}\")\n    print(f\"Successfully interpolated: {total_successful}\")\n    print(f\"Average success rate: {avg_success_rate:.1%}\")\n    \n    return result_binary, result_labeled, all_metrics","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-03-11T20:25:58.793424Z","iopub.execute_input":"2026-03-11T20:25:58.794177Z","iopub.status.idle":"2026-03-11T20:25:58.837116Z","shell.execute_reply.started":"2026-03-11T20:25:58.794148Z","shell.execute_reply":"2026-03-11T20:25:58.836500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# OPTIMIZED DUAL-GPU ENSEMBLE INFERENCE\n# Model 1 on GPU:0, Model 2 on GPU:1, Parallel Processing\n# ============================================================\n\n# ============================================================\n# 1. INFERENCE CONFIGURATION\n# ============================================================\n\nCONFIG = {\n    # --- MODEL 1 (GPU:0) ---\n    \"model1_folder\": \"/kaggle/input/models/nguyncdngs/nnunet-resenc/pytorch/default/3/nnUNet_ResEnc\",\n    \"model1_checkpoint\": \"checkpoint_final.pth\",\n    \"model1_weight\": 0.6,\n    \"model1_device\": \"cuda:0\",  # First T4\n    \n    # --- MODEL 2 (GPU:1) ---\n    \"model2_folder\": \"/kaggle/input/models/nguyncdngs/nnunet-resenc/pytorch/default/2/nnUNet_ResEnc\",\n    \"model2_checkpoint\": \"checkpoint_final.pth\",\n    \"model2_weight\": 0.4,\n    \"model2_device\": \"cuda:1\",  # Second T4\n    \n    # Đường dẫn data\n    \"input_dir\": \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\",\n    \"output_zip\": \"submission.zip\",\n    \n    # Cấu hình Predictor\n    \"use_folds\": (\"all\",),          \n    \"tile_step_size\": 0.5,      \n    \"use_gaussian\": True,        \n    \"use_mirroring\": True,\n    \n    # Spacing\n    \"spacing\": [1.0, 1.0, 1.0],\n\n    # --- POST-PROCESSING CONFIG ---\n    \"enable_postproc\": True,\n    \"T_low\": 0.2,\n    \"T_high\": 0.83,\n    \"z_radius\": 1,\n    \"xy_radius\": 0,\n    \"min_object_size\": 2000,\n}\n\n# ============================================================\n# 2. IMPORTS\n# ============================================================\n\nimport os\nimport glob\nimport zipfile\nimport numpy as np\nimport torch\nimport tifffile\nfrom tqdm import tqdm\nfrom skimage import morphology\nimport scipy.ndimage as ndi\nfrom concurrent.futures import ThreadPoolExecutor\nimport threading\n\nfrom nnunetv2.inference.predict_from_raw_data import nnUNetPredictor\nfrom nnunetv2.imageio.tif_reader_writer import Tiff3DIO\n\n# ============================================================\n# 3. POST-PROCESSING FUNCTIONS\n# ============================================================\n\ndef build_anisotropic_struct(z_radius: int, xy_radius: int):\n    \"\"\"Tạo kernel 3D hình elip/trụ để đóng lỗ hổng.\"\"\"\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    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\ndef topo_postprocess(\n                    probs,          # (D, H, W)\n                    T_low=0.6,\n                    T_high=0.9,\n                    z_radius=1,\n                    xy_radius=1,\n                    dust_min_size=500,\n                ):\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(strong, mask=weak, structure=struct_hyst)\n\n    if not mask.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n\n    # --- Step 2: 3D Anisotropic Closing ---\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 = morphology.remove_small_objects(mask.astype(bool), min_size=dust_min_size)\n\n    return mask.astype(np.uint8)\n\n# ============================================================\n# 4. PARALLEL PREDICTION CLASS\n# ============================================================\n\nclass ParallelPredictor:\n    \"\"\"Wrapper to run predictions in parallel on separate GPUs.\"\"\"\n    \n    def __init__(self, predictor, device_name):\n        self.predictor = predictor\n        self.device_name = device_name\n        self.lock = threading.Lock()\n    \n    def predict(self, image, properties):\n        \"\"\"Thread-safe prediction on assigned GPU.\"\"\"\n        with self.lock:\n            ret = self.predictor.predict_single_npy_array(\n                image, \n                properties, \n                segmentation_previous_stage=None, \n                output_file_truncated=None, \n                save_or_return_probabilities=True\n            )\n        return ret\n\n# ============================================================\n# 5. MAIN INFERENCE ENGINE\n# ============================================================\n\ndef run_inference():\n    print(\"--> [INIT] Setting up Dual-GPU Parallel Ensemble Inference...\")\n    \n    # 1. Check GPU availability\n    if not torch.cuda.is_available():\n        raise RuntimeError(\"CUDA not available!\")\n    \n    gpu_count = torch.cuda.device_count()\n    print(f\"--> [GPU CHECK] Found {gpu_count} GPU(s)\")\n    \n    if gpu_count < 2:\n        print(f\"--> [WARNING] Only {gpu_count} GPU found. Falling back to single GPU mode.\")\n        device1 = torch.device('cuda:0')\n        device2 = torch.device('cuda:0')\n    else:\n        device1 = torch.device(CONFIG[\"model1_device\"])\n        device2 = torch.device(CONFIG[\"model2_device\"])\n        print(f\"--> [GPU ASSIGNMENT] Model 1 -> {device1}, Model 2 -> {device2}\")\n\n    # 2. Initialize Model 1 on GPU:0\n    print(f\"--> [MODEL 1] Loading on {device1}...\")\n    predictor1 = nnUNetPredictor(\n        tile_step_size=CONFIG[\"tile_step_size\"],\n        use_gaussian=CONFIG[\"use_gaussian\"],\n        use_mirroring=CONFIG[\"use_mirroring\"],\n        perform_everything_on_device=True,\n        device=device1,\n        verbose=False,\n        verbose_preprocessing=False,\n        allow_tqdm=False\n    )\n    \n    predictor1.initialize_from_trained_model_folder(\n        CONFIG[\"model1_folder\"],\n        use_folds=CONFIG[\"use_folds\"],\n        checkpoint_name=CONFIG[\"model1_checkpoint\"]\n    )\n    print(f\"   -> Model 1 loaded on {device1}\")\n\n    # 3. Initialize Model 2 on GPU:1\n    print(f\"--> [MODEL 2] Loading on {device2}...\")\n    predictor2 = nnUNetPredictor(\n        tile_step_size=CONFIG[\"tile_step_size\"],\n        use_gaussian=CONFIG[\"use_gaussian\"],\n        use_mirroring=CONFIG[\"use_mirroring\"],\n        perform_everything_on_device=True,\n        device=device2,\n        verbose=False,\n        verbose_preprocessing=False,\n        allow_tqdm=False\n    )\n    \n    predictor2.initialize_from_trained_model_folder(\n        CONFIG[\"model2_folder\"],\n        use_folds=CONFIG[\"use_folds\"],\n        checkpoint_name=CONFIG[\"model2_checkpoint\"]\n    )\n    print(f\"   -> Model 2 loaded on {device2}\")\n    \n    # 4. Wrap predictors for parallel execution\n    parallel_pred1 = ParallelPredictor(predictor1, device1)\n    parallel_pred2 = ParallelPredictor(predictor2, device2)\n    \n    # 5. Setup Image Reader\n    reader = Tiff3DIO()\n    \n    # 6. Prepare Input Files\n    test_files = sorted(glob.glob(os.path.join(CONFIG[\"input_dir\"], \"*.tif\")))\n    if not test_files:\n        print(\"--> [WARNING] No .tif files found.\")\n        return \n    \n    print(f\"--> [DATA] Found {len(test_files)} files to process.\")\n    print(f\"--> [ENSEMBLE] Weights: Model1={CONFIG['model1_weight']}, Model2={CONFIG['model2_weight']}\")\n\n    # 7. Processing Loop with Parallel Execution\n    with zipfile.ZipFile(CONFIG[\"output_zip\"], 'w', zipfile.ZIP_DEFLATED) as zf:\n        \n        for file_path in tqdm(test_files, desc=\"Parallel Dual-GPU Inference\"):\n            filename = os.path.basename(file_path)\n            \n            # Read image\n            image, _ = reader.read_images([file_path])\n            \n            properties = {\n                'spacing': CONFIG['spacing']\n            }\n\n            # --- PARALLEL PREDICTION ON BOTH GPUS ---\n            with ThreadPoolExecutor(max_workers=2) as executor:\n                # Submit both predictions simultaneously\n                future1 = executor.submit(parallel_pred1.predict, image, properties)\n                future2 = executor.submit(parallel_pred2.predict, image, properties)\n                \n                # Wait for both to complete\n                ret1 = future1.result()\n                ret2 = future2.result()\n\n            # --- Extract Probabilities ---\n            if isinstance(ret1, tuple):\n                _, probabilities1 = ret1\n            else:\n                probabilities1 = None\n\n            if isinstance(ret2, tuple):\n                _, probabilities2 = ret2\n            else:\n                probabilities2 = None\n\n            # --- ENSEMBLE PROBABILITIES ---\n            if probabilities1 is not None and probabilities2 is not None:\n                # Weighted ensemble\n                total_weight = CONFIG[\"model1_weight\"] + CONFIG[\"model2_weight\"]\n                w1 = CONFIG[\"model1_weight\"] / total_weight\n                w2 = CONFIG[\"model2_weight\"] / total_weight\n                \n                ensembled_probs = w1 * probabilities1 + w2 * probabilities2\n                prob_map = ensembled_probs[1]\n                # Apply multiple threshold levels\n                pred = topo_postprocess(prob_map, T_low=0.2,\n                                            T_high=0.83,\n                                            z_radius=1,\n                                            xy_radius=0,\n                                            dust_min_size=100)\n                pred = zero_volume_faces(pred.astype(bool), thickness=3)\n                pred = morphology.remove_small_objects(\n                    pred.astype(bool),\n                    min_size=1000,\n                    connectivity=3\n                )\n\n                pred2 = topo_postprocess(prob_map, T_low=0.5,\n                                            T_high=0.83,\n                                            z_radius=1,\n                                            xy_radius=0,\n                                            dust_min_size=100)\n                pred2 = zero_volume_faces(pred2.astype(bool), thickness=3)\n                pred2 = morphology.remove_small_objects(\n                    pred2.astype(bool),\n                    min_size=1000,\n                    connectivity=3\n                )\n                \n                pred3 = topo_postprocess(prob_map, T_low=0.6,\n                                            T_high=0.83,\n                                            z_radius=1,\n                                            xy_radius=0,\n                                            dust_min_size=100)\n                pred3 = zero_volume_faces(pred3.astype(bool), thickness=3)\n                pred3 = morphology.remove_small_objects(\n                    pred3.astype(bool),\n                    min_size=1000,\n                    connectivity=3\n                )\n                \n                pred5 = topo_postprocess(prob_map, T_low=0.7,\n                                            T_high=0.83,\n                                            z_radius=1,\n                                            xy_radius=0,\n                                            dust_min_size=100)\n                pred5 = zero_volume_faces(pred5.astype(bool), thickness=3)\n                pred5 = morphology.remove_small_objects(\n                    pred5.astype(bool),\n                    min_size=1000,\n                    connectivity=3\n                )\n                pred1_uint8 = pred2.astype(np.uint8)\n                pred2_uint8 = pred2.astype(np.uint8)\n                pred3_uint8 = pred3.astype(np.uint8)\n                pred5_uint8 = pred5.astype(np.uint8)\n                # Fill holes\n                pred = binary_fill_holes(pred.astype(bool))\n                \n                # Crop borders\n                pred = pred[3:-3, 3:-3, 3:-3]\n                pred2 = pred2[3:-3, 3:-3, 3:-3]\n                pred3 = pred3[3:-3, 3:-3, 3:-3]\n                pred5 = pred5[3:-3, 3:-3, 3:-3]\n                \n                # Advanced component processing\n                pred, masks, metrics = process_multiple_components_patchwise(\n                    volume=pred,\n                    alternative_volumes=[pred2, pred3, pred5],  # Pass pre-computed alternatives\n                    cube_size=80,\n                    overlap=40,\n                    border_thickness=8,\n                    grid_resolution=80,\n                    thickness=3,\n                    smoothing=1.0,\n                    samples_per_edge=8,\n                    min_dice=0.75,           # Main thresholds\n                    min_coverage=0.75,\n                    alt_min_dice=0.45,       # More lenient for alternatives\n                    alt_min_coverage=0.80,\n                    min_component_size=100,  # Filter small noise\n                    use_parallel=True,\n                    n_jobs=-1\n                )\n                \n                # Final cleanup\n                pred = morphology.remove_small_objects(\n                    pred, \n                    min_size=1000, \n                    connectivity=3\n                )\n                \n                # Restore border padding\n                border = 3\n                pred = np.pad(\n                    pred,\n                    pad_width=(\n                        (border, border),\n                        (border, border),\n                        (border, border)\n                    ),\n                    mode=\"constant\",\n                    constant_values=0\n                )\n                \n                pred_uint8 = pred.astype(np.uint8)\n                \n            else:\n                # Fallback without post-processing\n                if segmentation_mask.ndim == 4:\n                    segmentation_mask = segmentation_mask[0]\n                pred_uint8 = segmentation_mask.astype(np.uint8)\n\n            # --- Save to Zip ---\n            temp_output_path = filename \n            tifffile.imwrite(\"base.tif\", pred1_uint8)\n            tifffile.imwrite(\"pred2.tif\", pred2_uint8)\n            tifffile.imwrite(\"pred3.tif\", pred3_uint8)\n            tifffile.imwrite(\"pred5.tif\", pred5_uint8)\n            tifffile.imwrite(temp_output_path, pred_uint8)\n            zf.write(temp_output_path, arcname=filename)\n            \n            if os.path.exists(temp_output_path):\n                os.remove(temp_output_path)\n\n    print(f\"\\n--> [DONE] Parallel Dual-GPU Inference Complete!\")\n    print(f\"    Output: {CONFIG['output_zip']}\")\n\nif __name__ == \"__main__\":\n    run_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-11T20:25:58.838511Z","iopub.execute_input":"2026-03-11T20:25:58.839238Z","iopub.status.idle":"2026-03-11T20:26:06.080086Z","shell.execute_reply.started":"2026-03-11T20:25:58.839205Z","shell.execute_reply":"2026-03-11T20:26:06.079120Z"}},"outputs":[],"execution_count":null}]}