{"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 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@jit(nopython=True, fastmath=True)\ndef check_triangle_in_bounds(p1, p2, p3, shape):\n    \"\"\"Check if triangle intersects volume bounds.\"\"\"\n    min_z = min(p1[0], p2[0], p3[0])\n    max_z = max(p1[0], p2[0], p3[0])\n    min_y = min(p1[1], p2[1], p3[1])\n    max_y = max(p1[1], p2[1], p3[1])\n    min_x = min(p1[2], p2[2], p3[2])\n    max_x = max(p1[2], p2[2], p3[2])\n\n    if max_z < 0 or min_z >= shape[0]: return False\n    if max_y < 0 or min_y >= shape[1]: return False\n    if max_x < 0 or min_x >= shape[2]: return False\n    return True\n\n\n# ============================================================================\n# VECTORIZED OVERLAP DETECTION (5-10x faster)\n# ============================================================================\n\ndef detect_overlaps_vectorized(fitted_sheets, num_components):\n    \"\"\"\n    Vectorized overlap detection using scipy operations.\n    Expected speedup: 5-10x faster than sequential method.\n    \"\"\"\n    shape = list(fitted_sheets.values())[0].shape\n    count_map = np.zeros(shape, dtype=np.int32)\n    for i in range(1, num_components + 1):\n        count_map += fitted_sheets[i].astype(np.int32)\n\n    potential_overlap = count_map > 1\n    if not np.any(potential_overlap):\n        return np.zeros(shape, dtype=bool)\n\n    labeled_result = np.zeros(shape, dtype=np.int32)\n    for i in range(1, num_components + 1):\n        labeled_result[fitted_sheets[i]] = i\n\n    from scipy.ndimage import generic_filter\n\n    def has_different_neighbor(values):\n        center = values[13]\n        if center == 0:\n            return 0\n        for val in values:\n            if val > 0 and val != center:\n                return 1\n        return 0\n\n    overlap_mask = np.zeros(shape, dtype=bool)\n    coords = np.column_stack(np.nonzero(potential_overlap))\n    if len(coords) == 0:\n        return overlap_mask\n\n    min_coords = np.maximum(coords.min(axis=0) - 1, 0)\n    max_coords = np.minimum(coords.max(axis=0) + 2, shape)\n    slices = tuple(slice(min_coords[i], max_coords[i]) for i in range(3))\n    roi_labeled = labeled_result[slices]\n    roi_potential = potential_overlap[slices]\n\n    roi_overlap = generic_filter(\n        roi_labeled,\n        has_different_neighbor,\n        size=3,\n        mode='constant',\n        cval=0\n    ).astype(bool)\n\n    roi_overlap = roi_overlap & roi_potential\n    overlap_mask[slices] = roi_overlap\n    return overlap_mask\n\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# ============================================================================\n# PARALLEL COMPONENT PROCESSING\n# ============================================================================\n\ndef process_component_wrapper(args):\n    \"\"\"Wrapper for parallel processing of components.\"\"\"\n    component_id, component_mask, grid_resolution, thickness, smoothing, max_distance, samples_per_edge = args\n    try:\n        fitted = fit_curved_sheet_to_component_optimized(\n            component_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        return component_id, fitted\n    except Exception as e:\n        print(f\"Error processing component {component_id}: {e}\")\n        return component_id, component_mask\n\n\ndef _evaluate_component_worker(args):\n    \"\"\"\n    Worker for parallel per-component evaluation (erode + quality check + alternatives).\n    Returns a dict describing what to write into result_labeled.\n\n    Keys in returned dict:\n        'id'              : component id\n        'status'          : 'correct' | 'fitted' | 'alternative' | 'removed' | 'lost'\n        'main_mask'       : boolean mask to assign label i (or None)\n        'extra_components': list of boolean masks for new sub-components (from alternatives)\n        'dice'            : float\n        'coverage'        : float\n    \"\"\"\n    (i, is_correct, component_mask, fitted_after_overlap,\n     grid_resolution, thickness, smoothing, max_distance, samples_per_edge,\n     alt_min_dice, alt_min_coverage, min_dice, min_coverage,\n     alternative_volumes, erosion_iterations, struct_elem) = args\n\n    original_component = component_mask\n\n    # ── Correct components: keep as-is ───────────────────────────────────────\n    if is_correct:\n        return {\n            'id': i, 'status': 'correct',\n            'main_mask': fitted_after_overlap,\n            'extra_components': [], 'dice': 1.0, 'coverage': 1.0,\n        }\n\n    # ── Lost during overlap removal ───────────────────────────────────────────\n    if not np.any(fitted_after_overlap):\n        print(f\"  Component {i}: Lost during overlap removal\")\n        # Fall through to alternatives below with dice/coverage = 0\n        dice, coverage = 0.0, 0.0\n        eroded = None\n    else:\n        # ── Erode ─────────────────────────────────────────────────────────────\n        if erosion_iterations > 0:\n            eroded = binary_erosion(fitted_after_overlap, structure=struct_elem,\n                                    iterations=erosion_iterations)\n        else:\n            eroded = fitted_after_overlap\n        eroded = binary_fill_holes(eroded)\n\n        dice = calculate_dice_score(original_component, eroded)\n        coverage = calculate_coverage_score(original_component, eroded)\n\n    if eroded is not None and dice >= min_dice and coverage >= min_coverage:\n        print(f\"  Component {i}: Dice={dice:.3f}, Coverage={coverage:.3f} ✓ (fitted)\")\n        return {\n            'id': i, 'status': 'fitted',\n            'main_mask': eroded,\n            'extra_components': [], 'dice': dice, 'coverage': coverage,\n        }\n\n    # ── Try alternative volumes ───────────────────────────────────────────────\n    if alternative_volumes is not None and len(alternative_volumes) > 0:\n        print(f\"  Component {i}: Dice={dice:.3f}, Coverage={coverage:.3f} — trying alternatives...\")\n\n        all_good_results = []\n        remaining_region = original_component.copy()\n\n        for alt_idx, alt_volume in enumerate(alternative_volumes):\n            if not np.any(remaining_region):\n                break\n\n            print(f\"    Alternative {alt_idx+1}/{len(alternative_volumes)} \"\n                  f\"(remaining: {np.sum(remaining_region)} vx)...\")\n\n            alt_mask = alt_volume & remaining_region\n            if not np.any(alt_mask):\n                print(f\"      No voxels in alternative within remaining region\")\n                continue\n\n            alt_labeled = label(alt_mask)\n            num_alt_comps = alt_labeled.max()\n            print(f\"      Found {num_alt_comps} component(s)\")\n\n            solved_in_this_alt = np.zeros_like(alt_volume, dtype=bool)\n            unsolved_in_this_alt = np.zeros_like(alt_volume, dtype=bool)\n\n            for comp_idx in range(1, num_alt_comps + 1):\n                alt_comp = (alt_labeled == comp_idx)\n\n                if np.sum(alt_comp) < 100:\n                    print(f\"        Component {comp_idx}: Too small ({np.sum(alt_comp)} vx)\")\n                    continue\n\n                # ── 3-faces check REMOVED ────────────────────────────────────\n                # (no face-touching filter applied here)\n\n                try:\n                    alt_fitted = fit_curved_sheet_to_component_optimized(\n                        alt_comp,\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                    alt_dice = calculate_dice_score(alt_comp, alt_fitted)\n                    alt_coverage = calculate_coverage_score(alt_comp, alt_fitted)\n                    print(f\"        Component {comp_idx}: Dice={alt_dice:.3f}, Cov={alt_coverage:.3f}\")\n\n                    if alt_dice >= alt_min_dice and alt_coverage >= alt_min_coverage:\n                        all_good_results.append({\n                            'fitted': alt_fitted,\n                            'dice': alt_dice, 'coverage': alt_coverage,\n                            'alt_idx': alt_idx, 'comp_idx': comp_idx,\n                            'source_comp': alt_comp,\n                        })\n                        solved_in_this_alt |= alt_comp\n                        print(f\"        Component {comp_idx}: ✓ Accepted\")\n                    else:\n                        unsolved_in_this_alt |= alt_comp\n                        print(f\"        Component {comp_idx}: ✗ Not good enough\")\n\n                except Exception as e:\n                    print(f\"        Component {comp_idx}: Failed to fit ({e})\")\n\n            if np.any(solved_in_this_alt):\n                remaining_region = unsolved_in_this_alt\n                print(f\"      Solved {np.sum(solved_in_this_alt)} vx; \"\n                      f\"remaining: {np.sum(unsolved_in_this_alt)} vx\")\n            else:\n                print(f\"      No components solved in this alternative\")\n\n        if len(all_good_results) > 0:\n            print(f\"    Total: {len(all_good_results)} good alternative(s)\")\n\n            # Overlap detection between alternatives\n            if len(all_good_results) > 1:\n                alt_fitted_sheets = {idx+1: r['fitted'] for idx, r in enumerate(all_good_results)}\n                alt_overlap = detect_overlaps_vectorized(alt_fitted_sheets, len(all_good_results))\n                alt_labeled_result = np.zeros_like(alt_volume, dtype=np.int32)\n                for idx in range(1, len(all_good_results) + 1):\n                    mask = alt_fitted_sheets[idx] & ~alt_overlap\n                    alt_labeled_result[mask] = idx\n\n                combined_alternatives = np.zeros_like(alt_volume, dtype=bool)\n                for idx in range(1, len(all_good_results) + 1):\n                    alt_comp_mask = (alt_labeled_result == idx)\n                    if not np.any(alt_comp_mask):\n                        continue\n                    if erosion_iterations > 0:\n                        ea = binary_erosion(alt_comp_mask, structure=struct_elem,\n                                            iterations=erosion_iterations)\n                    else:\n                        ea = alt_comp_mask\n                    ea = binary_fill_holes(ea)\n                    combined_alternatives |= ea\n            else:\n                combined_alternatives = all_good_results[0]['fitted']\n                if erosion_iterations > 0:\n                    combined_alternatives = binary_erosion(\n                        combined_alternatives, structure=struct_elem, iterations=erosion_iterations)\n                    combined_alternatives = binary_fill_holes(combined_alternatives)\n\n            # No final interpolation pass — beta1 re-interpolation loop will\n            # catch any remaining topology issues in a subsequent pass.\n            return {\n                'id': i, 'status': 'alternative',\n                'main_mask': None,\n                'extra_components': [combined_alternatives],\n                'dice': dice, 'coverage': coverage,\n            }\n\n    # ── No valid fit, remove component ───────────────────────────────────────\n    print(f\"  Component {i}: No valid fit found — REMOVING\")\n    return {\n        'id': i, 'status': 'removed',\n        'main_mask': None,\n        'extra_components': [], 'dice': dice, 'coverage': coverage,\n    }\n\n\n# ============================================================================\n# ITERATIVE BETA1 RE-INTERPOLATION\n# ============================================================================\n\ndef _reinterpolate_bad_components(\n    result_labeled,\n    grid_resolution, thickness, smoothing, max_distance, samples_per_edge,\n    overlap_buffer, min_dice, min_coverage, alt_min_dice, alt_min_coverage,\n    alternative_volumes, use_parallel, n_jobs,\n    max_iterations=3,\n    debug_output_dir=\"debug_reinterp\",\n):\n    \"\"\"\n    Check beta1 (= 1 - Euler number) for every component in result_labeled.\n    Components with beta1 > 0 are re-processed through the FULL threshold\n    pipeline — exactly as in the main loop:\n        1. Fit sheet\n        2. Overlap removal\n        3. Erode + evaluate vs min_dice / min_coverage\n        4. If that fails → try alternative volumes vs alt_min_dice / alt_min_coverage\n\n    Repeats up to max_iterations times, stopping early once all β1 ≤ 0.\n    Modifies result_labeled in-place and returns it.\n    \"\"\"\n    print(\"\\n\" + \"=\" * 70)\n    print(f\"ITERATIVE BETA1 RE-INTERPOLATION (max {max_iterations} passes)\")\n    print(\"=\" * 70)\n\n    erosion_iterations = overlap_buffer // 2\n    struct_elem = ball(1) if erosion_iterations > 0 else None\n\n    for iteration in range(max_iterations):\n        print(f\"\\n--- Pass {iteration + 1}/{max_iterations} ---\")\n\n        # ── Identify bad components ───────────────────────────────────────────\n        current_binary = result_labeled > 0\n        check_labeled = label(current_binary)\n        num_check = check_labeled.max()\n\n        bad_ids = []\n        for cid in range(1, num_check + 1):\n            comp = (check_labeled == cid)\n            chi = euler_number(comp.astype(int), connectivity=1)\n            beta1 = 1 - chi\n            if beta1 > 0:\n                bad_ids.append(cid)\n\n        if not bad_ids:\n            print(f\"  All {num_check} components have β1≤0 — done early!\")\n            break\n\n        print(f\"  {len(bad_ids)}/{num_check} components have β1>0\")\n\n        # ── Step A: Fit sheets in parallel ────────────────────────────────────\n        fit_args = [\n            (cid,\n             (check_labeled == cid),\n             grid_resolution,\n             thickness + overlap_buffer,\n             smoothing,\n             max_distance,\n             samples_per_edge)\n            for cid in bad_ids\n        ]\n\n        if use_parallel and len(fit_args) > 1:\n            max_workers = n_jobs if n_jobs > 0 else None\n            with ThreadPoolExecutor(max_workers=max_workers) as executor:\n                fit_results = list(executor.map(process_component_wrapper, fit_args))\n        else:\n            fit_results = [process_component_wrapper(a) for a in fit_args]\n\n        fitted_sheets = {cid: fitted for cid, fitted in fit_results}\n\n        # ── Step B: Overlap detection among re-fitted sheets ──────────────────\n        if len(fitted_sheets) > 1:\n            # Build a 1-indexed dict for detect_overlaps_vectorized\n            id_to_idx = {cid: idx + 1 for idx, cid in enumerate(bad_ids)}\n            idx_sheets = {id_to_idx[cid]: fitted_sheets[cid] for cid in bad_ids}\n            overlap_mask = detect_overlaps_vectorized(idx_sheets, len(bad_ids))\n        else:\n            shape = list(fitted_sheets.values())[0].shape\n            overlap_mask = np.zeros(shape, dtype=bool)\n\n        fitted_after_overlap = {\n            cid: fitted_sheets[cid] & ~overlap_mask for cid in bad_ids\n        }\n\n        # ── Step C: Evaluate each bad component through full threshold pipeline\n        eval_args = [\n            (cid,\n             False,                          # is_correct = False\n             (check_labeled == cid),         # component_mask (original)\n             fitted_after_overlap[cid],      # fitted after overlap removal\n             grid_resolution, thickness + overlap_buffer, smoothing, max_distance, samples_per_edge,\n             alt_min_dice, alt_min_coverage, min_dice, min_coverage,\n             alternative_volumes, erosion_iterations, struct_elem)\n            for cid in bad_ids\n        ]\n\n        if use_parallel and len(eval_args) > 1:\n            max_workers = n_jobs if n_jobs > 0 else None\n            with ThreadPoolExecutor(max_workers=max_workers) as executor:\n                eval_results = list(executor.map(_evaluate_component_worker, eval_args))\n        else:\n            eval_results = [_evaluate_component_worker(a) for a in eval_args]\n\n        # ── Step D: Update result_labeled ─────────────────────────────────────\n        next_label = result_labeled.max() + 1\n\n        for res in eval_results:\n            cid = res['id']\n            old_mask = (check_labeled == cid)\n\n            # Find the dominant original label this geometric region carried\n            orig_labels = result_labeled[old_mask]\n            dominant_label = (\n                int(np.bincount(orig_labels[orig_labels > 0]).argmax())\n                if np.any(orig_labels > 0) else 0\n            )\n            if dominant_label == 0:\n                dominant_label = next_label\n                next_label += 1\n\n            # Clear old voxels for this component\n            result_labeled[old_mask] = 0\n\n            if res['status'] in ('fitted', 'correct') and res['main_mask'] is not None:\n                result_labeled[res['main_mask'] & (result_labeled == 0)] = dominant_label\n\n            elif res['status'] == 'alternative':\n                for extra_mask in res['extra_components']:\n                    result_labeled[extra_mask & (result_labeled == 0)] = next_label\n                    next_label += 1\n\n            else:  # 'removed' — voxels already cleared, component is gone\n                print(f\"    Component {cid}: removed after re-interpolation\")\n\n        # ── Quick per-pass summary ─────────────────────────────────────────────\n        statuses = [r['status'] for r in eval_results]\n        print(f\"  Pass {iteration + 1} results: \"\n              f\"fitted={statuses.count('fitted')}, \"\n              f\"alternative={statuses.count('alternative')}, \"\n              f\"removed={statuses.count('removed')}\")\n\n    return result_labeled\n\n\n# ============================================================================\n# MAIN PARALLEL PROCESSING FUNCTION\n# ============================================================================\n\ndef process_multiple_components_parallel(\n    volume,\n    alternative_volumes=None,\n    grid_resolution=80,\n    thickness=3,\n    smoothing=2.0,\n    overlap_buffer=2,\n    min_coverage=0.60,\n    min_dice=0.6,\n    alt_min_coverage=None,\n    alt_min_dice=None,\n    max_distance=10,\n    use_parallel=True,\n    n_jobs=-1,\n    samples_per_edge=8,\n    max_reinterp_iterations=1,\n    debug_output_dir=\"debug_reinterp\",\n):\n    \"\"\"\n    Optimized parallel processing with:\n      • Euler-based pre-filtering (skip topologically correct components)\n      • Numba JIT rasterization\n      • Parallel fitting AND parallel per-component evaluation (incl. alternatives)\n      • Iterative beta1 re-interpolation loop (up to max_reinterp_iterations)\n      • 3-faces check REMOVED\n    \"\"\"\n    labeled_volume = label(volume)\n    num_components = labeled_volume.max()\n\n    if alt_min_coverage is None:\n        alt_min_coverage = min_coverage\n    if alt_min_dice is None:\n        alt_min_dice = min_dice\n\n    print(f\"Processing {num_components} components...\")\n    print(f\"Optimizations: Numba=True, Parallel={use_parallel}, AdaptiveRes=True\")\n    if alternative_volumes is not None:\n        print(f\"Using {len(alternative_volumes)} alternative volumes for fallback\")\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\n    # ── Extract component masks ───────────────────────────────────────────────\n    component_masks = {i: (labeled_volume == i) for i in range(1, num_components + 1)}\n\n    # ========================================================================\n    # STEP 1: Euler-based topology analysis\n    # ========================================================================\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 1: Euler-based topology analysis\")\n    print(\"=\" * 70)\n\n    correct_components = []\n    needs_interpolation = []\n\n    for i in range(1, num_components + 1):\n        chi = euler_number(component_masks[i].astype(int), connectivity=1)\n        beta1 = 1 - chi\n        if beta1 <= 0:\n            correct_components.append(i)\n            print(f\"  Component {i}: β1={beta1} (χ={chi}) ✓ CORRECT — keeping as-is\")\n        else:\n            needs_interpolation.append(i)\n            print(f\"  Component {i}: β1={beta1} (χ={chi}) ⚠  NEEDS INTERPOLATION\")\n\n    print(f\"\\nSummary: {len(correct_components)} correct, \"\n          f\"{len(needs_interpolation)} need interpolation\")\n\n    # ========================================================================\n    # STEP 2: Fit sheets to components that need it (parallel)\n    # ========================================================================\n    fitted_sheets = {}\n\n    # Correct components use original masks unchanged\n    for i in correct_components:\n        fitted_sheets[i] = component_masks[i]\n\n    if len(needs_interpolation) > 0:\n        print(\"\\n\" + \"=\" * 70)\n        print(f\"STEP 2: Fitting sheets for {len(needs_interpolation)} component(s)\")\n        print(\"=\" * 70)\n\n        fit_args = [\n            (i, component_masks[i], grid_resolution,\n             thickness + overlap_buffer, smoothing, max_distance, samples_per_edge)\n            for i in needs_interpolation\n        ]\n\n        if use_parallel and len(needs_interpolation) > 1:\n            print(\"  Running in parallel...\")\n            max_workers = n_jobs if n_jobs > 0 else None\n            with ThreadPoolExecutor(max_workers=max_workers) as executor:\n                fit_results = list(executor.map(process_component_wrapper, fit_args))\n        else:\n            fit_results = [process_component_wrapper(a) for a in fit_args]\n\n        for cid, fitted in fit_results:\n            fitted_sheets[cid] = fitted\n\n    # ========================================================================\n    # STEP 3: Overlap detection (vectorized, across ALL components)\n    # ========================================================================\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 3: Detecting overlaps (vectorized)\")\n    print(\"=\" * 70)\n\n    overlap_mask = detect_overlaps_vectorized(fitted_sheets, num_components)\n    print(f\"Removed {np.sum(overlap_mask)} overlapping voxels\")\n\n    # Build per-component overlap-free slices for evaluation\n    fitted_after_overlap = {}\n    for i in range(1, num_components + 1):\n        fitted_after_overlap[i] = fitted_sheets[i] & ~overlap_mask\n\n    # ========================================================================\n    # STEP 4: Evaluate components — erode, check quality, try alternatives\n    #         Run ALL components in parallel (correct ones short-circuit instantly)\n    # ========================================================================\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 4: Evaluating & rescuing components (parallel)\")\n    print(\"=\" * 70)\n\n    erosion_iterations = overlap_buffer // 2\n    struct_elem = ball(1) if erosion_iterations > 0 else None\n\n    eval_args = [\n        (i,\n         i in correct_components,\n         component_masks[i],\n         fitted_after_overlap[i],\n         grid_resolution, thickness + overlap_buffer, smoothing, max_distance, samples_per_edge,\n         alt_min_dice, alt_min_coverage, min_dice, min_coverage,\n         alternative_volumes, erosion_iterations, struct_elem)\n        for i in range(1, num_components + 1)\n    ]\n\n    if use_parallel and num_components > 1:\n        print(f\"  Running evaluation for {num_components} components in parallel...\")\n        max_workers = n_jobs if n_jobs > 0 else None\n        with ThreadPoolExecutor(max_workers=max_workers) as executor:\n            eval_results = list(executor.map(_evaluate_component_worker, eval_args))\n    else:\n        eval_results = [_evaluate_component_worker(a) for a in eval_args]\n\n    # ========================================================================\n    # STEP 5: Assemble result_labeled from evaluation results\n    # ========================================================================\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 5: Assembling final result\")\n    print(\"=\" * 70)\n\n    result_labeled = np.zeros_like(volume, dtype=np.int32)\n    dice_scores = {}\n    coverage_scores = {}\n\n    next_label = num_components + 1\n    kept_correct = 0\n    kept_fitted = 0\n    kept_alternative = 0\n    removed = 0\n\n    for res in eval_results:\n        i = res['id']\n        dice_scores[i] = res['dice']\n        coverage_scores[i] = res['coverage']\n\n        if res['status'] == 'correct':\n            if res['main_mask'] is not None:\n                result_labeled[res['main_mask']] = i\n            kept_correct += 1\n\n        elif res['status'] == 'fitted':\n            if res['main_mask'] is not None:\n                result_labeled[res['main_mask'] & (result_labeled == 0)] = i\n            kept_fitted += 1\n\n        elif res['status'] == 'alternative':\n            for extra_mask in res['extra_components']:\n                result_labeled[extra_mask & (result_labeled == 0)] = next_label\n                next_label += 1\n            kept_alternative += 1\n\n        else:  # 'removed' or 'lost'\n            removed += 1\n            print(f\"  Component {i}: {res['status'].upper()}\")\n\n    print(f\"\\n  Correct (β1=0):   {kept_correct}\")\n    print(f\"  Fitted:           {kept_fitted}\")\n    print(f\"  Via alternatives: {kept_alternative}\")\n    print(f\"  Removed:          {removed}\")\n    print(f\"  Total kept:       {kept_correct + kept_fitted + kept_alternative}/{num_components}\")\n\n    # ========================================================================\n    # STEP 6: Iterative beta1 re-interpolation (parallel, up to N passes)\n    # ========================================================================\n    result_labeled = _reinterpolate_bad_components(\n        result_labeled,\n        grid_resolution=grid_resolution,\n        thickness=thickness,\n        smoothing=smoothing,\n        max_distance=max_distance,\n        samples_per_edge=samples_per_edge,\n        overlap_buffer=overlap_buffer,\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        alternative_volumes=alternative_volumes,\n        use_parallel=use_parallel,\n        n_jobs=n_jobs,\n        max_iterations=max_reinterp_iterations,\n        debug_output_dir=debug_output_dir,\n    )\n\n    # ========================================================================\n    # Final summary\n    # ========================================================================\n    result_binary = result_labeled > 0\n    final_labeled = label(result_binary)\n    final_num = final_labeled.max()\n\n    valid_dice = [v for v in dice_scores.values() if isinstance(v, (int, float))]\n    avg_dice = float(np.mean(valid_dice)) if valid_dice else 0.0\n\n    print(f\"\\n{'=' * 70}\")\n    print(f\"FINAL SUMMARY\")\n    print(f\"{'=' * 70}\")\n    print(f\"Final component count : {final_num}\")\n    print(f\"Average Dice score    : {avg_dice:.3f}\")\n\n    return result_binary, result_labeled, dice_scores, coverage_scores\n\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-02-25T23:43:34.403322Z","iopub.execute_input":"2026-02-25T23:43:34.403945Z","iopub.status.idle":"2026-02-25T23:43:36.86554Z","shell.execute_reply.started":"2026-02-25T23:43:34.403913Z","shell.execute_reply":"2026-02-25T23:43:36.864943Z"}},"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    if z_radius > 0 or xy_radius > 0:\n        struct_close = build_anisotropic_struct(z_radius, xy_radius)\n        if struct_close is not None:\n            mask = ndi.binary_closing(mask, structure=struct_close)\n\n    # Step 3: Dust Removal\n    if dust_min_size > 0:\n        mask = 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                pred3 = 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                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                \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                pred3 = pred3[3:-3, 3:-3, 3:-3]\n                pred5 = pred5[3:-3, 3:-3, 3:-3]\n                \n                # Advanced component processing\n                pred = process_multiple_components_parallel(\n                    pred, \n                    grid_resolution=100, \n                    thickness=3, \n                    smoothing=1.0,\n                    overlap_buffer=0, \n                    min_coverage=0.65, \n                    min_dice=0.7,\n                    max_distance=10, \n                    alternative_volumes=[pred3, pred5], \n                    n_jobs=-1,\n                    alt_min_coverage=0.75,\n                    alt_min_dice=0.45,\n                    samples_per_edge=8\n                )\n                \n                # Final cleanup\n                pred = morphology.remove_small_objects(\n                    pred[0], \n                    min_size=2000, \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(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-02-25T23:43:36.866701Z","iopub.execute_input":"2026-02-25T23:43:36.867071Z","iopub.status.idle":"2026-02-25T23:43:43.715389Z","shell.execute_reply.started":"2026-02-25T23:43:36.867047Z","shell.execute_reply":"2026-02-25T23:43:43.714459Z"}},"outputs":[],"execution_count":null}]}