Skip to content

Commit 4e024cb

Browse files
Lower trellis workflow remesh memory usage. (#16034)
1 parent 03468f4 commit 4e024cb

1 file changed

Lines changed: 21 additions & 6 deletions

File tree

comfy_extras/mesh3d/postprocess/remesh.py

Lines changed: 21 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -356,6 +356,8 @@ def _dual_contour(voxel_coords: torch.Tensor, corner_udf: torch.Tensor,
356356
torch.zeros_like(found)).reshape(Nv, 8)
357357
edge_valid = cv_per_voxel[:, EDGES[:, 0]] & cv_per_voxel[:, EDGES[:, 1]]
358358
crosses = crosses & edge_valid
359+
del cv_per_voxel, edge_valid
360+
del keys_per_voxel, idx_per_voxel, idx_clamped, found
359361
# Zero-crossing interp factor per edge
360362
t = a_sd / (a_sd - b_sd + 1e-20)
361363
t = t.clamp(0.0, 1.0).unsqueeze(-1)
@@ -364,6 +366,7 @@ def _dual_contour(voxel_coords: torch.Tensor, corner_udf: torch.Tensor,
364366
a_pos = corner_world[:, EDGES[:, 0]] # (Nv, 12, 3)
365367
b_pos = corner_world[:, EDGES[:, 1]]
366368
crossing_pts = torch.lerp(a_pos, b_pos, t) # (Nv, 12, 3)
369+
del corner_pos_per_voxel, a_pos, b_pos, t
367370

368371
# Default dual vert: centroid of crossings (also QEF/no-crossing fallback)
369372
crosses_f = crosses.float().unsqueeze(-1)
@@ -412,6 +415,12 @@ def _dual_contour(voxel_coords: torch.Tensor, corner_udf: torch.Tensor,
412415
qef_solution = torch.where(in_box.unsqueeze(-1), qef_solution, centroid_verts)
413416

414417
dual_verts = torch.where(has_cross, qef_solution, centre_world)
418+
del query_pts, qef_tri_idx, valid_q, normals_at_q, full_normals, n_per_edge
419+
del A, n_dot_p, b, qef_solution, lo, hi, in_box
420+
del flat_pts, flat_mask
421+
422+
del crossing_pts, corner_world, crosses_f, crossing_sum, n_cross
423+
del centroid_verts, centre_world, has_cross
415424

416425
# Topology: each crossing grid edge is shared by 4 voxels -> quad -> 2 tris.
417426
# NEIGHBOUR_OFFS lays out the 4 sharing voxels per axis; y-axis order is
@@ -993,6 +1002,7 @@ def tick():
9931002
unique_corner_keys, corner_inv = torch.unique(corner_keys, return_inverse=True)
9941003
unique_corners = torch.zeros((unique_corner_keys.shape[0], 3), dtype=torch.long, device=device)
9951004
unique_corners[corner_inv] = corners
1005+
del corners, corner_keys, corner_inv
9961006

9971007
if sign_mode == "sdf":
9981008
use_sdf = True
@@ -1003,8 +1013,6 @@ def tick():
10031013

10041014
# Step 3: distance field at every unique corner.
10051015
tri_verts_g = vertices[faces.long()]
1006-
centroids = tri_verts_g.mean(dim=1)
1007-
tri_radii = (tri_verts_g - centroids.unsqueeze(1)).norm(dim=-1).max(dim=-1).values
10081016
# face normals: needed for the SDF sign AND for QEF placement (QEF is sign-agnostic,
10091017
# so it works in UDF mode too — (n·(x-p))² is unchanged by normal orientation)
10101018
if use_sdf or qef:
@@ -1027,12 +1035,11 @@ def tick():
10271035
else:
10281036
# UDF mode: iso at UDF=eps; double surface on closed meshes, weld after
10291037
sdf = udf - eps
1038+
del unique_corners, corner_world, udf, corner_closest, corner_tri
1039+
if use_sdf:
1040+
del sign, n_for_corner, offset, sign_dot
10301041
tick() # SDF done
10311042

1032-
# Short-range hash reused by project_back / colors sampling (max_dist up to 4*cell)
1033-
short_hash_cell_t = torch.tensor(2.0 * cell_size, dtype=vertices.dtype, device=device)
1034-
short_hash = _build_tri_spatial_hash(centroids, tri_radii, short_hash_cell_t)
1035-
10361043
# Step 4 + 5: dual contouring + topology. QEF works in both modes (sign-agnostic);
10371044
# in UDF it pulls the ±eps crossing back onto the triangle planes → sharper edges.
10381045
if qef:
@@ -1060,12 +1067,20 @@ def _qef_query(pts):
10601067
tri_face_normals=tri_face_normals, qef_query=_qef_query,
10611068
# corner_valid filter only matters in SDF mode
10621069
corner_valid=corner_valid if use_sdf else None)
1070+
del voxel_coords, sdf, unique_corner_keys, corner_valid
1071+
del tri_face_normals, _qef_query
1072+
if use_sdf or qef:
1073+
del tri_face_normals_all
10631074
tick() # DC done
10641075

10651076
# Step 6: project_back and / or color sampling share one closest-point query
10661077
need_query = (project_back > 0 or colors is not None) and dual_verts.numel() > 0
10671078
out_colors = None
10681079
if need_query:
1080+
centroids = tri_verts_g.mean(dim=1)
1081+
tri_radii = (tri_verts_g - centroids.unsqueeze(1)).norm(dim=-1).max(dim=-1).values
1082+
short_hash_cell_t = torch.tensor(2.0 * cell_size, dtype=vertices.dtype, device=device)
1083+
short_hash = _build_tri_spatial_hash(centroids, tri_radii, short_hash_cell_t)
10691084
result = _udf_query(
10701085
dual_verts, tri_verts_g, short_hash, short_hash_cell_t,
10711086
max_dist=4.0 * cell_size,

0 commit comments

Comments
 (0)