vmtk-centerline-api / endograft_generation.py
mosta02's picture
Upload 6 files
a1b81e3 verified
Raw
History Blame Contribute Delete
14.1 kB
import vtk
import numpy as np
import tempfile
def arc_length(cl):
s = np.zeros(len(cl))
for i in range(1, len(cl)):
s[i] = s[i-1] + np.linalg.norm(cl[i] - cl[i-1])
return s
def compute_parallel_frames(centerline):
frames = []
t0 = centerline[1] - centerline[0]
t0 /= np.linalg.norm(t0)
ref = np.array([0, 0, 1])
if abs(np.dot(t0, ref)) > 0.9:
ref = np.array([1, 0, 0])
n0 = np.cross(t0, ref)
n0 /= np.linalg.norm(n0)
b0 = np.cross(t0, n0)
frames.append((t0, n0, b0))
for i in range(1, len(centerline)):
ti = centerline[i] - centerline[i-1]
ti /= np.linalg.norm(ti)
t_prev, n_prev, b_prev = frames[-1]
v = np.cross(t_prev, ti)
if np.linalg.norm(v) < 1e-6:
frames.append((ti, n_prev, b_prev))
continue
v /= np.linalg.norm(v)
angle = np.arccos(np.clip(np.dot(t_prev, ti), -1, 1))
def rot(vec):
return (vec * np.cos(angle) + np.cross(v, vec) * np.sin(angle) + v * np.dot(v, vec) * (1 - np.cos(angle)))
frames.append((ti, rot(n_prev), rot(b_prev)))
return frames
def project_point_to_surface(p, surface_poly):
locator = vtk.vtkCellLocator()
locator.SetDataSet(surface_poly)
locator.BuildLocator()
closest = [0.0, 0.0, 0.0]
cellId = vtk.mutable(0)
subId = vtk.mutable(0)
dist2 = vtk.mutable(0.0)
locator.FindClosestPoint(p, closest, cellId, subId, dist2)
return np.array(closest)
def sample_centerline_by_s(centerline, s_cl, frames, s):
idx = np.searchsorted(s_cl, s)
if idx <= 0:
return centerline[0], frames[0]
if idx >= len(centerline):
return centerline[-1], frames[-1]
s0, s1 = s_cl[idx-1], s_cl[idx]
P0, P1 = centerline[idx-1], centerline[idx]
t = (s - s0) / (s1 - s0 + 1e-12)
P = (1 - t) * P0 + t * P1
t0, N0, B0 = frames[idx-1]
t1, N1, B1 = frames[idx]
t_vec = (1 - t) * t0 + t * t1
N_vec = (1 - t) * N0 + t * N1
B_vec = (1 - t) * B0 + t * B1
t_vec /= np.linalg.norm(t_vec)
N_vec /= np.linalg.norm(N_vec)
B_vec /= np.linalg.norm(B_vec)
return P, (t_vec, N_vec, B_vec)
# Fabric
def create_component_fabric(centerline, frames, start_s, length, diameter, start_amp_pp=2.0, waves=8):
# shift centerline sampling
shifted_centerline = centerline.copy()
s_cl = arc_length(centerline)
pts = vtk.vtkPoints()
polys = vtk.vtkCellArray()
rings = []
A = start_amp_pp / 2.0
n_sections = 180
n_theta = 64
for i in range(n_sections):
s_base = start_s + length * i / (n_sections - 1)
P_base, (t, N, B) = sample_centerline_by_s(centerline, s_cl, frames, s_base)
ring = []
for j in range(n_theta):
theta = 2 * np.pi * j / n_theta
axial_offset = 0.0
if i == 0 and start_amp_pp > 0:
axial_offset = A * np.sin(waves * theta) - A
P = P_base + axial_offset * t
r = diameter / 2.0
pos = P + r * (np.cos(theta) * N + np.sin(theta) * B)
ring.append(pts.InsertNextPoint(pos))
rings.append(ring)
for i in range(len(rings)-1):
for j in range(n_theta):
q = vtk.vtkQuad()
q.GetPointIds().SetId(0, rings[i][j])
q.GetPointIds().SetId(1, rings[i][(j+1)%n_theta])
q.GetPointIds().SetId(2, rings[i+1][(j+1)%n_theta])
q.GetPointIds().SetId(3, rings[i+1][j])
polys.InsertNextCell(q)
poly = vtk.vtkPolyData()
poly.SetPoints(pts)
poly.SetPolys(polys)
return poly
# Struts
def create_proximal_bare_stent(centerline, frames, fabric_start_s, main_diam_mm, length, main_fabric, peak_to_peak=15.0,
peak_spacing=7.0, component_type="proximal", wire_radius=0.35, waves=8, samples=240
):
s_cl = arc_length(centerline)
append = vtk.vtkAppendPolyData()
# --- build fabric normals & locator ONCE ---
normal_gen = vtk.vtkPolyDataNormals()
normal_gen.SetInputData(main_fabric)
normal_gen.ComputePointNormalsOn()
normal_gen.SplittingOff()
normal_gen.ConsistencyOn()
normal_gen.Update()
fabric_normals = normal_gen.GetOutput().GetPointData().GetNormals()
locator = vtk.vtkPointLocator()
locator.SetDataSet(main_fabric)
locator.BuildLocator()
# -----------------------------------------
A = peak_to_peak # peak-peak amplitude = 15
r = main_diam_mm / 2.0
gap = peak_spacing
step = A + gap
n_crowns = int(np.ceil(length / step))
s_peaks = [fabric_start_s + 1.0 + i * step for i in range(n_crowns)]
if ((component_type == "distal" or component_type == "extension") and (len(s_peaks) > 4)):
s_peaks = s_peaks[3:]
# =========================================================
# ===== EXTRA CLIPPED HALF sin =================
# =========================================================
A_extra = peak_to_peak / 2.0 # 7.5 mm height
extra_top = -1e-10 # peak at -7.5
if component_type == "proximal":
P0 = centerline[0]
t0, N0, B0 = frames[0]
P = P0 + extra_top * t0 # extrapolate backward
t, N, B = t0, N0, B0
elif component_type == "distal":
P, (t, N, B) = sample_centerline_by_s(centerline, s_cl, frames, fabric_start_s + length - extra_top)
else:
P, (t, N, B) = sample_centerline_by_s(centerline, s_cl, frames, extra_top)
pts = vtk.vtkPoints()
lines = vtk.vtkCellArray()
for i in range(samples + 1):
theta = 2*np.pi*i/samples
# keep ONLY negative or positive half of sine
if component_type == "proximal":
axial_offset = A_extra * min(0.0, np.sin(waves * theta))
raw_pos = (P
+ (r + wire_radius) * (np.cos(theta)*N + np.sin(theta)*B)
+ axial_offset * t)
elif (component_type == "distal" and length != 0):
axial_offset = A_extra * max(0.0, np.sin(waves * theta))
raw_pos = (P
+ (r + wire_radius) * (np.cos(theta)*N + np.sin(theta)*B)
+ axial_offset * t)
else:
continue
radial = (np.cos(theta)*N + np.sin(theta)*B)
radial /= np.linalg.norm(radial)
pos = P + (r + wire_radius) * radial + axial_offset * t
pts.InsertNextPoint(pos)
if i > 0:
l = vtk.vtkLine()
l.GetPointIds().SetId(0, i-1)
l.GetPointIds().SetId(1, i)
lines.InsertNextCell(l)
poly = vtk.vtkPolyData()
poly.SetPoints(pts)
poly.SetLines(lines)
tube = vtk.vtkTubeFilter()
tube.SetInputData(poly)
tube.SetRadius(wire_radius)
tube.SetNumberOfSides(18)
tube.CappingOff()
tube.Update()
append.AddInputData(tube.GetOutput())
# =========================================================
# ================= NORMAL FULL struts ====================
# =========================================================
for s0 in s_peaks:
P, (t, N, B) = sample_centerline_by_s(centerline, s_cl, frames, s0)
pts = vtk.vtkPoints()
lines = vtk.vtkCellArray()
for i in range(samples + 1):
theta = 2*np.pi*i/samples
axial_offset = A * (1 + np.sin(waves * theta)) / 2.0
raw_pos = (P
+ (r + wire_radius) * (np.cos(theta)*N + np.sin(theta)*B)
+ axial_offset * t)
surf_p = project_point_to_surface(raw_pos, main_fabric)
pid = locator.FindClosestPoint(surf_p)
n = np.array(fabric_normals.GetTuple(pid))
n /= np.linalg.norm(n)
pos = surf_p + wire_radius * n
pts.InsertNextPoint(pos)
if i > 0:
l = vtk.vtkLine()
l.GetPointIds().SetId(0, i-1)
l.GetPointIds().SetId(1, i)
lines.InsertNextCell(l)
poly = vtk.vtkPolyData()
poly.SetPoints(pts)
poly.SetLines(lines)
tube = vtk.vtkTubeFilter()
tube.SetInputData(poly)
tube.SetRadius(wire_radius)
tube.SetNumberOfSides(18)
tube.CappingOff()
tube.Update()
append.AddInputData(tube.GetOutput())
append.Update()
return append.GetOutput()
def create_component_stents(centerline, frames, fabric_start_s, length, diameter, fabric, peak_to_peak=15.0, peak_spacing=7.0, component_type="proximal"):
stents = create_proximal_bare_stent(centerline, frames, fabric_start_s, diameter, length, fabric, peak_to_peak, peak_spacing, component_type)
return stents
def merge_components(proximal, extension, distal):
merged_components = vtk.vtkAppendPolyData()
merged_components.AddInputData(proximal)
merged_components.AddInputData(extension)
merged_components.AddInputData(distal)
merged_components.Update()
clean = vtk.vtkCleanPolyData()
clean.SetInputData(merged_components.GetOutput())
clean.Update()
return clean
def create_endograft(centerline, diameter, prox_length,extension_length, distal_length, prox_gap, extension_gap, distal_gap):
frames = compute_parallel_frames(centerline)
peak_to_peak = 15.0
# -------------------------
# PROXIMAL COMPONENT
# -------------------------
prox_start = 0.0
prox_fabric = create_component_fabric(centerline, frames, prox_start, prox_length, diameter, start_amp_pp=2.0)
prox_stents = create_component_stents(centerline, frames, prox_start, prox_length, diameter, prox_fabric, peak_spacing=prox_gap, component_type="proximal")
# -------------------------
# EXTENSION COMPONENT
# -------------------------
if extension_length > 97:
prox_extension_overlap = 3*peak_to_peak + 2*extension_gap
else:
prox_extension_overlap = 2*peak_to_peak + extension_gap
extension_start = prox_length - prox_extension_overlap
extension_fabric = create_component_fabric(centerline, frames, extension_start, extension_length, diameter, start_amp_pp=0.0)
extension_stents = create_component_stents(centerline, frames, extension_start, extension_length, diameter, extension_fabric, peak_spacing=extension_gap, component_type="extension")
# -------------------------
# DISTAL COMPONENT
# -------------------------
distal_overlap = 3*peak_to_peak + 2*distal_gap
distal_start = prox_length + extension_length - prox_extension_overlap - distal_overlap
distal_fabric = create_component_fabric(centerline, frames, distal_start, distal_length, diameter, start_amp_pp=0.0)
distal_stents = create_component_stents(centerline, frames, distal_start, distal_length, diameter, distal_fabric, peak_spacing=distal_gap, component_type="distal")
# -------------------------
# MERGE COMPONENTS
# -------------------------
fabric_part = merge_components(prox_fabric, extension_fabric, distal_fabric)
stent_part = merge_components(prox_stents, extension_stents, distal_stents)
endograft = vtk.vtkAppendPolyData()
endograft.AddInputData(fabric_part.GetOutput())
endograft.AddInputData(stent_part.GetOutput())
endograft.Update()
clean = vtk.vtkCleanPolyData()
clean.SetInputData(endograft.GetOutput())
clean.Update()
return fabric_part.GetOutput(), stent_part.GetOutput(), clean.GetOutput()
def normalize_centerline_input(centerline):
if len(centerline) > 0 and isinstance(centerline[0], dict):
centerline = [[p["x"], p["y"], p["z"]] for p in centerline]
arr = np.asarray(centerline, dtype=np.float64)
if arr.ndim != 2 or arr.shape[1] != 3:
raise ValueError("centerline must have shape [N, 3]")
if arr.shape[0] < 2:
raise ValueError("centerline must contain at least 2 points")
return arr
# Returns generated VTK polydata parts as standalone objects
def build_endograft_parts(centerline, diameter=22.0, prox_length=105.0, extension_length=0.0, distal_length=0.0, prox_gap=7.0, extension_gap=0.0, distal_gap=0.0):
centerline_np = normalize_centerline_input(centerline)
endograft_poly, struts_poly, combined_poly = create_endograft(centerline_np, diameter, prox_length, extension_length, distal_length, prox_gap, extension_gap, distal_gap)
return {"endograft": endograft_poly, "struts": struts_poly, "combined": combined_poly,}
def polydata_to_stl_bytes(polydata):
tmp = tempfile.NamedTemporaryFile(suffix=".stl", delete=False)
tmp_path = tmp.name
tmp.close()
try:
writer = vtk.vtkSTLWriter()
writer.SetFileName(tmp_path)
writer.SetInputData(polydata)
writer.Write()
with open(tmp_path, "rb") as f:
return f.read()
finally:
try:
import os
os.unlink(tmp_path)
except OSError:
pass
def build_endograft_stl_payload(centerline, **params):
parts = build_endograft_parts(centerline, **params)
return {"endograft": polydata_to_stl_bytes(parts["endograft"]), "struts": polydata_to_stl_bytes(parts["struts"]), "combined": polydata_to_stl_bytes(parts["combined"])}
def generate_stl_files(centerline, output_prefix="endograft", **params):
parts = build_endograft_parts(centerline, **params)
outputs = {"endograft": f"{output_prefix}-endograft.stl", "struts": f"{output_prefix}-struts.stl", "combined": f"{output_prefix}-combined.stl"}
for key, path in outputs.items():
writer = vtk.vtkSTLWriter()
writer.SetFileName(path)
writer.SetInputData(parts[key])
writer.Write()
return outputs