vmtk-centerline-api / centerline_extraction.py
mosta02's picture
Upload 6 files
a1b81e3 verified
Raw
History Blame Contribute Delete
37.8 kB
import math
from vmtk import vtkvmtk
from vtk.util import numpy_support
import numpy as np
import SimpleITK as sitk
import vtk
# Coordinate conversion helpers (Left-Posterior-Superior ↔ Right-Anterior-Superior) by flipping X & Y
# LPS: +X is patient's left, +Y is Posterior (back), +Z is Superior (head), RPS: +X is Patient's Right, +Y is Anterior (front)
def lps_ras_conversion(polydata):
transform = vtk.vtkTransform()
transform.Scale(-1.0, -1.0, 1.0)
tf = vtk.vtkTransformPolyDataFilter()
tf.SetInputData(polydata)
tf.SetTransform(transform)
tf.Update()
return tf.GetOutput()
def lps_endpoints_to_ras(endpoints):
return [(-x, -y, z, is_start) for x, y, z, is_start in endpoints]
# Reads the NRRD segmentation using SimpleITK instead of vmtkImageReader as it correctly preserves full 3x3 direction cosine matrix
def load_nrrd_as_vtk_image(path):
sitk_img = sitk.ReadImage(path)
arr = sitk.GetArrayFromImage(sitk_img)
# Auto-binarizes multi-label segmentations (any non-zero voxel → 1)
arr_max = arr.max()
if arr_max > 1 or (arr_max == 1 and arr.min() < 0):
arr = (arr > 0).astype(np.uint8)
elif arr.dtype != np.uint8:
arr = arr.astype(np.uint8)
spacing = sitk_img.GetSpacing() # (sx, sy, sz) in mm, indicates distance between the center of two adjacent voxels
origin = sitk_img.GetOrigin() # (ox, oy, oz) in mm
direction = sitk_img.GetDirection() # 3×3 direction cosine matrix
# Build vtkImageData matching the SimpleITK geometry (VTK functions only accept vtkImageData)
vtk_img = vtk.vtkImageData()
vtk_img.SetDimensions(arr.shape[::-1]) # SimpleITK is (z,y,x), VTK needs (x,y,z)
vtk_img.SetSpacing(spacing)
vtk_img.SetOrigin(origin)
# Flattens the 3D numpy array (shape [Z, Y, X]) into a 1D array in row-major (C) order, then converts it to vtkDataArray using deep copy to allow garbage collection
vtk_arr = numpy_support.numpy_to_vtk(num_array=arr.ravel(order="C"), deep=True)
vtk_arr.SetName("Image") # Names the scalar array so VTK knows it's the primary scalar dataset
vtk_img.GetPointData().SetScalars(vtk_arr)
return vtk_img, origin, direction
# Main centerline network computation pipeline, converts VTK image (representing segmentation file) into complete centerline network
def compute_centerlines_network(
image, origin, direction, label_value=1, surface_smoothing=0.3, surface_decimation=0.0,
keep_all_components=True, start_point=None, curve_sampling_distance=1.0
):
# Build closed surface mesh from segmentation, generating both LPS surface (for skeletonization) and RAS surface (for centerline extraction)
surface_polydata_lps = create_closed_surface_polydata(image, origin, direction, label_value, surface_smoothing, surface_decimation, keep_all_components, to_ras=False)
surface_polydata_ras = lps_ras_conversion(surface_polydata_lps)
# Skeletonize the surface mesh to get the network topolgy (vessel endpoints and radius at each point)
print("\nExtracting skeleton centerline network...")
network = skeletonize(
image, origin, direction, surface_polydata_lps, label_value, surface_smoothing, surface_decimation, keep_all_components,
target_number_of_points=5000, subdivide=False, start_point=start_point, compute_geometry=False,
)
# Detect the aorta root, iliac endpoints, branch endpoints from the skeleton network
print("\nDetecting aorta and iliac endpoints from network...")
aorta_iliac_endpoints, branch_endpoints = detect_centerline_endpoints(network)
if not aorta_iliac_endpoints:
print("Failed to detect aorta-iliac endpoints from network")
return None, []
# Convert detected endpoints from LPS to RAS (centerline extraction works in RAS)
aorta_iliac_endpoints_ras = lps_endpoints_to_ras(aorta_iliac_endpoints)
# Decimate the surface by reducing surface mesh points before voronoi centerline extraction as it scales badly with large surfaces
#surface_for_centerline = decimate_surface_for_voronoi(surface_polydata_ras, target_points=50000)
# Extract main aorta→iliacs centerline via Voronoi method
try:
print("Extracting main centerline...")
centerline_polydata = extract_centerline_from_surface(surface_polydata_ras, aorta_iliac_endpoints_ras, curve_sampling_distance, smooth_iterations=20)
print(f"Main centerline: {centerline_polydata.GetNumberOfCells()} cells, {centerline_polydata.GetNumberOfPoints()} points")
except Exception as e:
print(f"Main centerline extraction failed: {type(e).__name__}: {e}")
return None, []
# Extract branching arteries centerlines
branch_centerlines = []
branch_bifurcations = []
if centerline_polydata and branch_endpoints:
try:
branch_centerlines, branch_bifurcations = process_all_branches(branch_endpoints, centerline_polydata, surface_polydata_ras, curve_sampling_distance)
except Exception as e:
print(f"Branch extraction failed: {type(e).__name__}: {e}")
# Merge main + branch centerlines
if branch_centerlines:
try:
merged = merge_main_and_branch_centerlines(centerline_polydata, branch_centerlines)
except Exception as e:
print(f"Merging centerlines failed: {e}")
merged = centerline_polydata
else:
merged = centerline_polydata
# Convert final result from RAS back to LPS (to match expected coordinate system)
complete_centerline = lps_ras_conversion(merged)
branch_junctions_lps = [(-c[0], -c[1], c[2]) for c in branch_bifurcations]
return complete_centerline, branch_junctions_lps
# Converts a binary labelmap into a closed 3D surface mesh following the same algorithm as 3D Slicer's vtkBinaryLabelmapToClosedSurfaceConversionRule
def create_closed_surface_polydata(image, origin, direction, label_value, smoothing, decimation, keep_all_components, to_ras=True, conversion_method="flying_edges"):
smoothing = max(0.0, min(1.0, float(smoothing)))
decimation = max(0.0, min(1.0, float(decimation)))
label_value = int(label_value)
# Cast segmentation VTK image to unsigned short (discrete surface filters require integer types)
cast = vtk.vtkImageCast()
cast.SetInputData(image)
cast.SetOutputScalarTypeToUnsignedShort()
cast.Update()
image_data = cast.GetOutput()
# Pad by 1 voxel if border voxels are non-background (prevents open edges)
if is_labelmap_padding_necessary(image_data):
padder = vtk.vtkImageConstantPad()
padder.SetInputData(image_data)
extent = list(image_data.GetExtent())
padder.SetOutputWholeExtent(extent[0] - 1, extent[1] + 1, extent[2] - 1, extent[3] + 1, extent[4] - 1, extent[5] + 1)
padder.Update()
image_data = padder.GetOutput()
# Work in IJK space (identity geometry) for cleaner isosurface extraction which generate triangle vertices at positions inside voxel grid
# If the image has non-uniform spacing or a non-zero origin the triangles get stretched/shifted mid-computation producing degenerate results
image_identity = vtk.vtkImageData()
image_identity.ShallowCopy(image_data)
image_identity.SetOrigin(0.0, 0.0, 0.0)
image_identity.SetSpacing(1.0, 1.0, 1.0)
# Surface extraction
if conversion_method == "surface_nets":
try:
# generates a smoother, watertight mesh by placing vertices at the center of edges crossing the boundary
surface_filter = vtk.vtkSurfaceNets3D()
surface_filter.SetInputData(image_identity)
surface_filter.SmoothingOff()
surface_filter.SetValue(0, label_value)
surface_filter.Update()
except AttributeError:
raise RuntimeError("vtkSurfaceNets3D not available in this VTK build")
else:
try:
# Parallelized, cache-optimized reimplementation of Marching Cubes for labeled/discrete data
surface_filter = vtk.vtkDiscreteFlyingEdges3D()
except AttributeError:
# Each cube defined by 8 neighboring voxels is classified by which corners are "inside" vs "outside", then a lookup table selects triangle patterns
surface_filter = vtk.vtkDiscreteMarchingCubes()
surface_filter.SetInputData(image_identity)
surface_filter.ComputeGradientsOff() # Don't compute gradient vectors — not needed, saves time
surface_filter.ComputeNormalsOff() # Skip per-vertex normals here — we compute consistent ones later
surface_filter.SetValue(0, label_value) # Extract isosurface at label_value (usually 1)
surface_filter.Update()
surface = surface_filter.GetOutput()
# Optional decimation for faster computation
if decimation > 0.0:
decimator = vtk.vtkDecimatePro()
decimator.SetInputData(surface)
decimator.SetFeatureAngle(60) # Edges sharper than 60° are "features" and won't be collapsed
decimator.SplittingOff() # Don't split the mesh at feature edges (preserves topology)
decimator.PreserveTopologyOn() # Avoids creating holes during decimation
decimator.SetMaximumError(1) # Max displacement of any vertex during decimation = 1 unit
decimator.SetTargetReduction(decimation) # Fraction of triangles to remove (0.0–1.0)
decimator.Update()
surface = decimator.GetOutput()
# Windowed Sinc Smoothing with 3D Slicer's passband & iteration mapping
if smoothing > 0.0 and conversion_method != "surface_nets":
pass_band = pow(10.0, -4.0 * smoothing) # Lower passband = more smoothing
iterations = int(20 + smoothing * 40) # More smoothing = more iterations
smoother = vtk.vtkWindowedSincPolyDataFilter()
smoother.SetInputData(surface)
smoother.SetNumberOfIterations(iterations)
smoother.SetPassBand(pass_band)
smoother.BoundarySmoothingOff() # Keep open boundary edges fixed
smoother.FeatureEdgeSmoothingOff() # Keep sharp feature edges in place
smoother.NonManifoldSmoothingOn() # Allow smoothing at non-manifold intersections
smoother.NormalizeCoordinatesOn() # Scale coords to unit cube before smoothing for numerical stability
smoother.Update()
surface = smoother.GetOutput()
# Transform from IJK back to physical space via 4×4 affine matrix (combines direction cosine, spacing, origin into one transform)
d = (1, 0, 0, 0, 1, 0, 0, 0, 1)
if direction and len(direction) == 9:
d = direction
spacing = (1.0, 1.0, 1.0)
if hasattr(image, "GetSpacing"):
spacing = image.GetSpacing()
ox, oy, oz = origin
mat = vtk.vtkMatrix4x4()
mat.Identity()
mat.SetElement(0, 0, d[0] * spacing[0]); mat.SetElement(0, 1, d[1] * spacing[1]); mat.SetElement(0, 2, d[2] * spacing[2]); mat.SetElement(0, 3, ox)
mat.SetElement(1, 0, d[3] * spacing[0]); mat.SetElement(1, 1, d[4] * spacing[1]); mat.SetElement(1, 2, d[5] * spacing[2]); mat.SetElement(1, 3, oy)
mat.SetElement(2, 0, d[6] * spacing[0]); mat.SetElement(2, 1, d[7] * spacing[1]); mat.SetElement(2, 2, d[8] * spacing[2]); mat.SetElement(2, 3, oz)
transform = vtk.vtkTransform()
transform.SetMatrix(mat)
tf = vtk.vtkTransformPolyDataFilter()
tf.SetInputData(surface)
tf.SetTransform(transform)
tf.Update()
transformed = tf.GetOutput()
# Compute consistent normals (all pointing outward or inward) that matter for voronoi methods that depend on normals directions
if smoothing > 0.0 or conversion_method == "flying_edges":
normals_filter = vtk.vtkPolyDataNormals()
normals_filter.SetInputData(transformed)
normals_filter.ConsistencyOn()
normals_filter.SplittingOff()
normals_filter.Update()
transformed = normals_filter.GetOutput()
# Keep only largest connected component (removes noisy islands)
if not keep_all_components:
conn = vtk.vtkPolyDataConnectivityFilter()
conn.SetInputData(transformed)
conn.SetExtractionModeToLargestRegion()
conn.Update()
transformed = conn.GetOutput()
# Convert LPS → RAS if requested
if to_ras:
transformed = lps_ras_conversion(transformed)
return transformed
# Helper that checks if any non-background voxels touch the image border (would create open edges)
def is_labelmap_padding_necessary(image_data):
extent = image_data.GetExtent() # (xmin, xmax, ymin, ymax, zmin, zmax) in voxels
if extent[0] > extent[1] or extent[2] > extent[3] or extent[4] > extent[5]: # Empty image guard
return False
dims = image_data.GetDimensions()
scalars = image_data.GetPointData().GetScalars() # Flat array of voxel values
if scalars is None:
return False
arr = numpy_support.vtk_to_numpy(scalars).reshape((dims[2], dims[1], dims[0]))
# returns True if any voxel on that face is non-zero (the segmented vessel touches the image boundary)
return (
arr[0, :, :].any() or arr[-1, :, :].any() or
arr[:, 0, :].any() or arr[:, -1, :].any() or
arr[:, :, 0].any() or arr[:, :, -1].any()
)
# Generates a centerline skeleton from a segmentaion using the same approach as 3D Slicer's ExtractCenterline module
# The output has Radius, Topology, and Marks arrays on each point, Endpoints (degree-1) are vessel tips while degree-3+ points are bifurcations
def skeletonize(
image, origin, direction, surface_polydata, label_value, smoothing, decimation,
keep_all_components, target_number_of_points, subdivide, start_point, compute_geometry,
):
# Build closed surface from segmentation if not provided
if surface_polydata is None:
if image is None:
raise ValueError("image is required when surface_polydata is not provided")
if origin is None:
origin = (0.0, 0.0, 0.0)
if hasattr(image, "GetOrigin"):
origin = image.GetOrigin()
if direction is None:
direction = (1, 0, 0, 0, 1, 0, 0, 0, 1)
surface_polydata = create_closed_surface_polydata(image, origin, direction, label_value, smoothing, decimation, keep_all_components, to_ras=True)
# Decimate, smooth, clean, triangulate the surface
processed_surface = preprocess_surface_for_network_extraction(surface_polydata, target_number_of_points, subdivide)
# Choose where to open the surface
if start_point is not None:
start_position = start_point
else:
bounds = processed_surface.GetBounds()
start_position = (bounds[0], bounds[2], bounds[4]) # x_min, y_min, z_min (bottom left corner)
# Open a small hole as skeletonization algorithm requires opening points inside the vessel
open_surface_at_point(processed_surface, start_position)
# Run VMTK's network extraction which traces the medial axis using maximal inscribed sphere approach
network_extraction = vtkvmtk.vtkvmtkPolyDataNetworkExtraction()
network_extraction.SetInputData(processed_surface)
network_extraction.SetAdvancementRatio(1.05) # Advancing sphere can shrink by no more than ~5% between successive step
network_extraction.SetRadiusArrayName("Radius")
network_extraction.SetTopologyArrayName("Topology") # Integer array where each point stores connectivity degree (0: interior, 1: endpoint, ..)
network_extraction.SetMarksArrayName("Marks") # Used internally by the network extraction algorithm to track which points have been "visited"
network_extraction.Update()
network_output = network_extraction.GetOutput()
if not compute_geometry:
return network_output
# Compute geometric properties: Length, Curvature, Torsion, Tortuosity, Frenet frame
centerline_geometry = vtkvmtk.vtkvmtkCenterlineGeometry()
centerline_geometry.SetInputData(network_output)
centerline_geometry.SetLengthArrayName("Length")
centerline_geometry.SetCurvatureArrayName("Curvature")
centerline_geometry.SetTorsionArrayName("Torsion")
centerline_geometry.SetTortuosityArrayName("Tortuosity")
centerline_geometry.SetFrenetTangentArrayName("FrenetTangent")
centerline_geometry.SetFrenetNormalArrayName("FrenetNormal")
centerline_geometry.SetFrenetBinormalArrayName("FrenetBinormal")
centerline_geometry.Update()
return centerline_geometry.GetOutput()
# Prepares the surface mesh for VMTK's skeleton (network) extraction
def preprocess_surface_for_network_extraction(surface_polydata, target_number_of_points, subdivide):
num_points = surface_polydata.GetNumberOfPoints()
if num_points == 0:
raise ValueError("Input surface is empty")
# Optionally decimate surface by reducing point count to target via quadric decimation, decreasing computation time
decimated = surface_polydata
reduction_factor = (num_points - target_number_of_points) / num_points
if reduction_factor > 0.0:
decim = vtk.vtkQuadricDecimation()
decim.SetInputData(surface_polydata)
decim.SetTargetReduction(min(max(reduction_factor, 0.0), 0.99))
decim.AttributeErrorMetricOn()
decim.VolumePreservationOn()
decim.Update()
decimated = decim.GetOutput()
# Smooth to improve mesh quality for network extraction
smoother = vtk.vtkWindowedSincPolyDataFilter()
smoother.SetInputData(decimated)
smoother.SetNumberOfIterations(10)
smoother.SetPassBand(0.1)
smoother.NonManifoldSmoothingOn()
smoother.NormalizeCoordinatesOn()
smoother.Update()
# Cleaning the mesh which merges coincident points (removing degenerate cells, converting any polygons to triangle)
surface_cleaner = vtk.vtkCleanPolyData()
surface_cleaner.SetInputData(smoother.GetOutput())
surface_cleaner.Update()
# Triangulate as VMTK requires pure triangle meshes
surface_triangulator = vtk.vtkTriangleFilter()
surface_triangulator.SetInputData(surface_cleaner.GetOutput())
surface_triangulator.PassLinesOff()
surface_triangulator.PassVertsOff()
surface_triangulator.Update()
# Optional linear subdivision for extra detail (splits each triangle into 4 smaller triangles)
subdivided = surface_triangulator.GetOutput()
if subdivide:
subdiv = vtk.vtkLinearSubdivisionFilter()
subdiv.SetInputData(surface_triangulator.GetOutput())
subdiv.SetNumberOfSubdivisions(1)
subdiv.Update()
if subdiv.GetOutput().GetNumberOfPoints() > 0:
subdivided = subdiv.GetOutput()
# Orient normals consistently (ensures all face normals point outward from the vessel surface)
normals_filter = vtk.vtkPolyDataNormals()
normals_filter.SetInputData(subdivided)
normals_filter.SetAutoOrientNormals(1)
normals_filter.SetFlipNormals(0)
normals_filter.SetConsistency(1)
normals_filter.SplittingOff()
normals_filter.Update()
return normals_filter.GetOutput()
# Deletes one cell at the specified point to create a small opening in the surface mesh which VMTK's network extraction requires to enter the vessel
def open_surface_at_point(polydata, hole_position, hole_point_index=None):
if hole_point_index is None:
locator = vtk.vtkPointLocator()
locator.SetDataSet(polydata)
locator.BuildLocator()
hole_point_index = locator.FindClosestPoint(hole_position)
if hole_point_index < 0:
raise ValueError("open_surface_at_point: invalid hole point")
polydata.BuildLinks()
cell_ids = vtk.vtkIdList()
polydata.GetPointCells(hole_point_index, cell_ids)
if cell_ids.GetNumberOfIds() > 0:
polydata.DeleteCell(cell_ids.GetId(0)) # Only remove the first cell
polydata.RemoveDeletedCells()
# Analyzes the network skeleton to detect: Aortic root (endpoint with the largest inscribed sphere radius), Iliac endpoints, Branch endpoints
def detect_centerline_endpoints(network_polydata):
# Clean and prepare the network (remove duplicate points or zero length cells)
cleaner = vtk.vtkCleanPolyData()
cleaner.SetInputData(network_polydata)
cleaner.Update()
network = cleaner.GetOutput()
network.BuildCells()
network.BuildLinks(0)
points = network.GetPoints()
if points is None or network.GetNumberOfCells() == 0:
return [], []
radius_array = network.GetPointData().GetArray("Radius")
if radius_array is None:
radius_array = network.GetPointData().GetArray("MaximumInscribedSphereRadius")
# Find all endpoints (points connected to exactly 1 cell)
endpoint_ids = vtk.vtkIdList()
tmp_point_cells = vtk.vtkIdList()
for cell_id in range(network.GetNumberOfCells()):
cell = network.GetCell(cell_id)
num_pts = cell.GetNumberOfPoints()
if num_pts < 2:
continue
for pid in (cell.GetPointId(0), cell.GetPointId(num_pts - 1)):
network.GetPointCells(pid, tmp_point_cells)
if tmp_point_cells.GetNumberOfIds() == 1:
endpoint_ids.InsertUniqueId(pid)
if endpoint_ids.GetNumberOfIds() < 3:
print(f"Warning: only {endpoint_ids.GetNumberOfIds()} endpoints found (need >= 3)")
return [], []
# Collect information about each endpoint position and radius
endpoints_info = []
for i in range(endpoint_ids.GetNumberOfIds()):
pid = endpoint_ids.GetId(i)
pos = points.GetPoint(pid)
radius = 0.0
if radius_array:
radius = radius_array.GetValue(pid)
endpoints_info.append({'id': pid, 'pos': pos, 'x': pos[0], 'y': pos[1], 'z': pos[2], 'radius': radius})
# Aortic root = endpoint with the largest inscribed sphere radius
aortic_root = max(endpoints_info, key=lambda ep: ep['radius'])
# Iliac candidates = endpoints below the aortic root
iliac_candidates = [ep for ep in endpoints_info if ep['z'] < aortic_root['z']]
# Fallback if iliac candidates weren't found, we take the two furthest endpoints from aortic root as iliac endpoints
if len(iliac_candidates) < 2:
iliac_candidates = sorted(endpoints_info, key=lambda ep: vtk.vtkMath.Distance2BetweenPoints(aortic_root['pos'], ep['pos']), reverse=True)[:2]
# Keep only the most inferior endpoints first (iliac tips are the lowest), use up to 8 points so we still capture external/internal iliacs on both sides
iliac_candidates = sorted(iliac_candidates, key=lambda endpoint: endpoint['z'])[: min(8, len(iliac_candidates))]
# Split left and right candidates by sorted X into two balanced sets
iliac_candidates.sort(key=lambda endpoint: endpoint['x'])
n_iliac = len(iliac_candidates)
half = n_iliac // 2
if n_iliac <= 1:
left_iliacs = iliac_candidates[:]
right_iliacs = []
# Even count: exact half/half split
elif n_iliac % 2 == 0:
left_iliacs = iliac_candidates[:half]
right_iliacs = iliac_candidates[half:]
# Odd count: assign middle point to the side whose boundary X is closer.
else:
middle = iliac_candidates[half]
left_boundary = iliac_candidates[half - 1]
right_boundary = iliac_candidates[half + 1]
d_left = abs(middle['x'] - left_boundary['x'])
d_right = abs(right_boundary['x'] - middle['x'])
if d_left <= d_right:
left_iliacs = iliac_candidates[: half + 1]
right_iliacs = iliac_candidates[half + 1 :]
else:
left_iliacs = iliac_candidates[:half]
right_iliacs = iliac_candidates[half:]
# Fallback if one side is empty: split by mean X pivot
if not left_iliacs or not right_iliacs:
x_mean = sum(endpoint['x'] for endpoint in iliac_candidates) / max(1, len(iliac_candidates))
left_iliacs = [endpoint for endpoint in iliac_candidates if endpoint['x'] <= x_mean]
right_iliacs = [endpoint for endpoint in iliac_candidates if endpoint['x'] > x_mean]
def sort_by_z_and_radius(iliacs):
return sorted(iliacs, key=lambda endpoint: (endpoint['z'], -endpoint['radius']))
left_sorted = sort_by_z_and_radius(left_iliacs)
right_sorted = sort_by_z_and_radius(right_iliacs)
left_external = left_sorted[0] if len(left_sorted) > 0 else None
left_internal = left_sorted[1] if len(left_sorted) > 1 else None
right_external = right_sorted[0] if len(right_sorted) > 0 else None
right_internal = right_sorted[1] if len(right_sorted) > 1 else None
# Build aorta + iliac endpoint list where aortic root is the source (is_start=True) while iliac endpoints are the targets (is_start=False)
aorta_iliac_endpoints = [(aortic_root['x'], aortic_root['y'], aortic_root['z'], True)]
iliac_ids = set()
for endpoint in (right_external, right_internal, left_external, left_internal):
if endpoint:
aorta_iliac_endpoints.append((endpoint['x'], endpoint['y'], endpoint['z'], False))
iliac_ids.add(endpoint['id'])
# Other endpoints not detected as aortic root or iliac endpoints are branching arteries endpoints
branch_endpoints = []
for endpoint in endpoints_info:
if endpoint['id'] != aortic_root['id'] and endpoint['id'] not in iliac_ids:
bif_pt = find_junction_point(network, endpoint['id']) # Trace endpoint to its skeleton bifurcation
branch_endpoints.append({'endpoint': (endpoint['x'], endpoint['y'], endpoint['z']), 'bifurcation': bif_pt, 'radius': endpoint['radius']})
# Diagnostic text output
print(f"Aorta (source): {aortic_root['x']:.2f}, {aortic_root['y']:.2f}, {aortic_root['z']:.2f} (R={aortic_root['radius']:.2f})")
if right_external:
print(f"Right ext iliac: {right_external['x']:.2f}, {right_external['y']:.2f}, {right_external['z']:.2f}")
if left_external:
print(f"Left ext iliac: {left_external['x']:.2f}, {left_external['y']:.2f}, {left_external['z']:.2f}")
if branch_endpoints:
print(f"{len(branch_endpoints)} additional branch endpoint(s)")
return aorta_iliac_endpoints, branch_endpoints
# Traces back from an endpoint along the centerline skeleton to find its bifurcation point (point where artery branches from main aorta)
def find_junction_point(network_polydata, endpoint_id):
network_polydata.BuildLinks(0)
points = network_polydata.GetPoints()
current_id = endpoint_id
visited = set([current_id])
while True:
cell_ids = vtk.vtkIdList()
network_polydata.GetPointCells(current_id, cell_ids)
num_cells = cell_ids.GetNumberOfIds()
# Degree >= 3 means we reached a bifurcation
if num_cells >= 3:
return points.GetPoint(current_id)
# Find the next unvisited neighbour along the skeleton
next_id = None
for c in range(num_cells):
cell = network_polydata.GetCell(cell_ids.GetId(c))
for j in range(cell.GetNumberOfPoints()):
pid = cell.GetPointId(j)
if pid != current_id and pid not in visited:
next_id = pid
break
if next_id is not None:
break
# Reached dead-end, no bifurcation found
if next_id is None:
return None
visited.add(next_id)
current_id = next_id
# Decimates the surface mesh to speed up Voronoi centerline extraction. Voronoi computation scales roughly O(n^2) with surface point count
def decimate_surface_for_voronoi(surface_polydata, target_points, max_reduction=0.7):
num_points = surface_polydata.GetNumberOfPoints()
if num_points <= target_points:
return surface_polydata
reduction = (num_points - target_points) / num_points
reduction = min(reduction, max_reduction)
decim = vtk.vtkQuadricDecimation()
decim.SetInputData(surface_polydata)
decim.SetTargetReduction(reduction)
decim.AttributeErrorMetricOn()
decim.VolumePreservationOn()
decim.Update()
# Clean after decimation to remove degenerate cells and ensure triangle mesh
cleaner = vtk.vtkCleanPolyData()
cleaner.SetInputData(decim.GetOutput())
cleaner.Update()
tri = vtk.vtkTriangleFilter()
tri.SetInputData(cleaner.GetOutput())
tri.PassLinesOff()
tri.PassVertsOff()
tri.Update()
result = tri.GetOutput()
return result
# Voronoi-based centerline extraction from aorta to iliacs to compute accurate centerline paths, transforms closed surface to centerline network
def extract_centerline_from_surface(surface_polydata, endpoints, curve_sampling_distance, smooth_iterations=40, smooth_relaxation=0.1):
if not endpoints or len(endpoints) < 2:
raise ValueError("At least two endpoints are needed for centerline extraction")
# Cap open mesh ends (required for Voronoi computation inside the volume)
capper = vtkvmtk.vtkvmtkCapPolyData()
capper.SetInputData(surface_polydata)
# Cap is exactly a flat polygon in the plane of the opening. The Voronoi algorithm doesn't care whether caps are flat or domed
capper.SetDisplacement(0.0)
capper.SetInPlaneDisplacement(0.0)
capper.Update()
tube_polydata = capper.GetOutput()
# Classify seeds: is_start=True → source, is_start=False → target
source_ids = vtk.vtkIdList()
target_ids = vtk.vtkIdList()
locator = vtk.vtkPointLocator()
locator.SetDataSet(tube_polydata)
locator.BuildLocator()
has_explicit_start = any(p[3] for p in endpoints) # returns True if at least one endpoint has is_start=True
for idx, (x, y, z, is_start) in enumerate(endpoints):
is_target = not is_start
if not has_explicit_start and idx == 0:
is_target = False
pid = locator.FindClosestPoint((x, y, z))
if is_target:
target_ids.InsertNextId(pid)
else:
source_ids.InsertNextId(pid)
# Run VMTK's Voronoi-based centerline filter which finds minimum-cost paths between source and target seeds (const is inversely proportional to inscribed sphere radius)
cl_filter = vtkvmtk.vtkvmtkPolyDataCenterlines()
cl_filter.SetInputData(tube_polydata)
cl_filter.SetSourceSeedIds(source_ids)
cl_filter.SetTargetSeedIds(target_ids)
cl_filter.SetRadiusArrayName("Radius")
cl_filter.SetCostFunction("1/R")
cl_filter.SetFlipNormals(False) # Controls whether the Voronoi diagram is built from the inward normals (flipped) or outward normals
cl_filter.SetAppendEndPointsToCenterlines(0) # Path starts/ends at the nearest Voronoi node, which may be a fraction of a millimeter away from exact seed point
cl_filter.SetSimplifyVoronoi(False) # Voronoi smoothing disabled (slows up path finding but maintains narrow vessel sections)
cl_filter.SetCenterlineResampling(1) # Enable resampling every 1mm, without this the raw Voronoi path is returned which contains thousands of points per segment
cl_filter.SetResamplingStepLength(curve_sampling_distance) # Turns on resampler itself
cl_filter.Update()
if not cl_filter.GetOutput():
raise ValueError("Centerline extraction produced no output")
centerline_polydata = vtk.vtkPolyData()
centerline_polydata.DeepCopy(cl_filter.GetOutput())
"""voronoi_polydata = None
if cl_filter.GetVoronoiDiagram():
voronoi_polydata = vtk.vtkPolyData()
voronoi_polydata.DeepCopy(cl_filter.GetVoronoiDiagram())"""
# Optional Laplacian smoothing to remove zigzag artifacts
if smooth_iterations > 0:
centerline_polydata = smooth_centerline_polydata(centerline_polydata, smooth_iterations, smooth_relaxation)
return centerline_polydata
# Applies low pass smoothing filter to a centerline polydata while keeping endpoints fixed, it removes high-frequency zigzag noise while preserving vessel curvature
def smooth_centerline_polydata(centerline_polydata, iterations, relaxation_factor):
smoother = vtk.vtkWindowedSincPolyDataFilter()
smoother.SetInputData(centerline_polydata)
smoother.SetNumberOfIterations(iterations)
smoother.SetPassBand(relaxation_factor)
smoother.BoundarySmoothingOff() # Keep both endpoints exactly fixed
smoother.FeatureEdgeSmoothingOff()
smoother.NonManifoldSmoothingOn() # Required for polyline topology (not just surfaces)
smoother.NormalizeCoordinatesOn() # Numerical stability on long centerlines
smoother.Update()
smoothed = vtk.vtkPolyData()
smoothed.DeepCopy(smoother.GetOutput())
return smoothed
# Process all branches, generating their centerlines, computing their new bifurcation points on actual voronoi centerline (not skeleton)
def process_all_branches(branch_information, main_centerline_polydata, surface_polydata, curve_sampling_distance):
if not branch_information:
return [], []
# Convert branch info from LPS to RAS
branch_info_ras = []
for info in branch_information:
x, y, z = info['endpoint']
ep_ras = (-x, -y, z)
bif_ras = None
if info['bifurcation'] is not None:
bx, by, bz = info['bifurcation']
bif_ras = (-bx, -by, bz)
branch_info_ras.append({'endpoint': ep_ras, 'bifurcation': bif_ras})
branch_centerlines = []
branch_bifurcations = []
for _, info in enumerate(branch_info_ras, 1):
centerline, contact_point = extract_branch_centerline(info, main_centerline_polydata, surface_polydata, curve_sampling_distance)
if centerline:
branch_centerlines.append(centerline)
if contact_point:
branch_bifurcations.append(contact_point)
print(f"\nExtracted {len(branch_centerlines)}/{len(branch_info_ras)} branch centerlines")
return branch_centerlines, branch_bifurcations
# A separate Voronoi centerline is extracted between the branch tip and its contact point on the aorta
def extract_branch_centerline(branch_info, main_centerline_polydata, surface_polydata, curve_sampling_distance):
branch_endpoint = branch_info['endpoint']
bifurcation_point = branch_info['bifurcation']
# Find the contact point on the main centerline
contact_point = find_branch_contact_point(branch_endpoint, bifurcation_point, main_centerline_polydata)
# Branch tip = source, contact point = target
branch_endpoints = [(branch_endpoint[0], branch_endpoint[1], branch_endpoint[2], True), (contact_point[0], contact_point[1], contact_point[2], False),]
# Clip full segmentation surface to a local region covering only the artery for faster processing during voronoi centerline extraction of that branch
clipped_surface = clip_surface_around_points(surface_polydata, branch_endpoint, contact_point, radius_multiplier=1.25)
#clipped_surface = decimate_surface_for_voronoi(clipped_surface, target_points=15000)
# Extract the branch centerline
try:
branch_centerline = extract_centerline_from_surface(clipped_surface, branch_endpoints, curve_sampling_distance, smooth_iterations=20)
return branch_centerline, contact_point
except Exception as e:
print(f"\nWarning: branch centerline extraction failed: {e}")
return None, contact_point
# Finds where a branch meets the main centerline using the skeleton bifurcation point
def find_branch_contact_point(branch_endpoint, bifurcation_point, main_centerline_polydata):
locator = vtk.vtkPointLocator()
locator.SetDataSet(main_centerline_polydata)
locator.BuildLocator()
search_point = None
if bifurcation_point is None:
# Fallback: closest point on main centerline to the branch tip
search_point = branch_endpoint
else:
# Project the skeleton bifurcation onto the main centerline
search_point = bifurcation_point
closest_id = locator.FindClosestPoint(search_point)
return main_centerline_polydata.GetPoints().GetPoint(closest_id)
# Clips the surface mesh to a spherical region between two points (Used to reduce input mesh size from full vessel tree to local branch for faster extraction)
def clip_surface_around_points(surface_polydata, point1, point2, radius_multiplier):
center = [(point1[0] + point2[0]) / 2.0, (point1[1] + point2[1]) / 2.0, (point1[2] + point2[2]) / 2.0]
distance = math.sqrt(vtk.vtkMath.Distance2BetweenPoints(point1, point2))
radius = max(distance * radius_multiplier, 20.0) # 20 mm minimum
sphere = vtk.vtkSphere()
sphere.SetCenter(center)
sphere.SetRadius(radius)
clipper = vtk.vtkClipPolyData()
clipper.SetInputData(surface_polydata)
clipper.SetClipFunction(sphere)
clipper.InsideOutOn() # Keep points inside the sphere
clipper.Update()
cleaner = vtk.vtkCleanPolyData()
cleaner.SetInputData(clipper.GetOutput())
cleaner.Update()
return cleaner.GetOutput()
# Merges the main aorta→iliacs centerline with all branch centerlines into one polydata.
def merge_main_and_branch_centerlines(main_centerline_polydata, branch_centerlines):
if not branch_centerlines:
return main_centerline_polydata
append_filter = vtk.vtkAppendPolyData()
append_filter.AddInputData(main_centerline_polydata)
for cl in branch_centerlines:
append_filter.AddInputData(cl)
append_filter.Update()
# Applies localized smoothing near junction points so that branch-to-trunk connections look smooth
cleaner = vtk.vtkCleanPolyData()
cleaner.SetInputData(append_filter.GetOutput())
cleaner.SetToleranceIsAbsolute(True)
cleaner.SetAbsoluteTolerance(1.0) # 1mm tolerance to merge junction points from separately-extracted branch centerlines
cleaner.Update()
merged = cleaner.GetOutput()
return merged