3d: make Voxel and Mesh types strict

This commit is contained in:
bigcat88 2025-07-26 11:52:29 +03:00
parent b007125398
commit ee54158e27
No known key found for this signature in database
GPG Key ID: 1F0BF0EC3CF22721
2 changed files with 33 additions and 31 deletions

View File

@ -619,13 +619,23 @@ class LossMap(ComfyTypeIO):
loss: list[torch.Tensor] loss: list[torch.Tensor]
Type = LossMapDict Type = LossMapDict
@comfytype(io_type="VOXEL") @comfytype(io_type="VOXEL")
class Voxel(ComfyTypeIO): class Voxel(ComfyTypeIO):
Type = Any # TODO: VOXEL class is defined in comfy_extras/nodes_hunyuan3d.py; should be moved to somewhere else before referenced directly in v3 class VoxelDict(TypedDict):
data: torch.Tensor
Type = VoxelDict
@comfytype(io_type="MESH") @comfytype(io_type="MESH")
class Mesh(ComfyTypeIO): class Mesh(ComfyTypeIO):
Type = Any # TODO: MESH class is defined in comfy_extras/nodes_hunyuan3d.py; should be moved to somewhere else before referenced directly in v3 class MeshDict(TypedDict):
vertices: list[torch.Tensor]
faces: list[torch.Tensor]
Type = MeshDict
@comfytype(io_type="HOOKS") @comfytype(io_type="HOOKS")
class Hooks(ComfyTypeIO): class Hooks(ComfyTypeIO):

View File

@ -16,17 +16,6 @@ from comfy.ldm.modules.diffusionmodules.mmdit import (
from comfy_api.latest import io from comfy_api.latest import io
class VOXEL:
def __init__(self, data):
self.data = data
class MESH:
def __init__(self, vertices, faces):
self.vertices = vertices
self.faces = faces
def voxel_to_mesh(voxels, threshold=0.5, device=None): def voxel_to_mesh(voxels, threshold=0.5, device=None):
if device is None: if device is None:
device = torch.device("cpu") device = torch.device("cpu")
@ -485,11 +474,11 @@ class Hunyuan3Dv2Conditioning(io.ComfyNode):
node_id="Hunyuan3Dv2Conditioning_V3", node_id="Hunyuan3Dv2Conditioning_V3",
category="conditioning/video_models", category="conditioning/video_models",
inputs=[ inputs=[
io.ClipVisionOutput.Input("clip_vision_output") io.ClipVisionOutput.Input("clip_vision_output"),
], ],
outputs=[ outputs=[
io.Conditioning.Output(display_name="positive"), io.Conditioning.Output(display_name="positive"),
io.Conditioning.Output(display_name="negative") io.Conditioning.Output(display_name="negative"),
] ]
) )
@ -511,11 +500,11 @@ class Hunyuan3Dv2ConditioningMultiView(io.ComfyNode):
io.ClipVisionOutput.Input("front", optional=True), io.ClipVisionOutput.Input("front", optional=True),
io.ClipVisionOutput.Input("left", optional=True), io.ClipVisionOutput.Input("left", optional=True),
io.ClipVisionOutput.Input("back", optional=True), io.ClipVisionOutput.Input("back", optional=True),
io.ClipVisionOutput.Input("right", optional=True) io.ClipVisionOutput.Input("right", optional=True),
], ],
outputs=[ outputs=[
io.Conditioning.Output(display_name="positive"), io.Conditioning.Output(display_name="positive"),
io.Conditioning.Output(display_name="negative") io.Conditioning.Output(display_name="negative"),
] ]
) )
@ -552,7 +541,7 @@ class SaveGLB(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, mesh, filename_prefix): def execute(cls, mesh: io.Mesh.MeshDict, filename_prefix):
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, folder_paths.get_output_directory()) full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, folder_paths.get_output_directory())
results = [] results = []
@ -564,9 +553,9 @@ class SaveGLB(io.ComfyNode):
for x in cls.hidden.extra_pnginfo: for x in cls.hidden.extra_pnginfo:
metadata[x] = json.dumps(cls.hidden.extra_pnginfo[x]) metadata[x] = json.dumps(cls.hidden.extra_pnginfo[x])
for i in range(mesh.vertices.shape[0]): for i in range(mesh["vertices"].shape[0]):
f = f"{filename}_{counter:05}_.glb" f = f"{filename}_{counter:05}_.glb"
save_glb(mesh.vertices[i], mesh.faces[i], os.path.join(full_output_folder, f), metadata) save_glb(mesh["vertices"][i], mesh["faces"][i], os.path.join(full_output_folder, f), metadata)
results.append({ results.append({
"filename": f, "filename": f,
"subfolder": subfolder, "subfolder": subfolder,
@ -590,14 +579,17 @@ class VAEDecodeHunyuan3D(io.ComfyNode):
io.Int.Input("octree_resolution", default=256, min=16, max=512) io.Int.Input("octree_resolution", default=256, min=16, max=512)
], ],
outputs=[ outputs=[
io.Voxel.Output() io.Voxel.Output(),
] ]
) )
@classmethod @classmethod
def execute(cls, vae, samples, num_chunks, octree_resolution): def execute(cls, vae, samples, num_chunks, octree_resolution):
voxels = VOXEL(vae.decode(samples["samples"], vae_options={"num_chunks": num_chunks, "octree_resolution": octree_resolution})) return io.NodeOutput(
return io.NodeOutput(voxels) io.Voxel.VoxelDict(
data=vae.decode(samples["samples"], vae_options={"num_chunks": num_chunks, "octree_resolution": octree_resolution})
)
)
class VoxelToMesh(io.ComfyNode): class VoxelToMesh(io.ComfyNode):
@ -612,12 +604,12 @@ class VoxelToMesh(io.ComfyNode):
io.Float.Input("threshold", default=0.6, min=-1.0, max=1.0, step=0.01) io.Float.Input("threshold", default=0.6, min=-1.0, max=1.0, step=0.01)
], ],
outputs=[ outputs=[
io.Mesh.Output() io.Mesh.Output(),
] ]
) )
@classmethod @classmethod
def execute(cls, voxel, algorithm, threshold): def execute(cls, voxel: io.Voxel.VoxelDict, algorithm, threshold):
vertices = [] vertices = []
faces = [] faces = []
@ -626,12 +618,12 @@ class VoxelToMesh(io.ComfyNode):
elif algorithm == "surface net": elif algorithm == "surface net":
mesh_function = voxel_to_mesh_surfnet mesh_function = voxel_to_mesh_surfnet
for x in voxel.data: for x in voxel["data"]:
v, f = mesh_function(x, threshold=threshold, device=None) v, f = mesh_function(x, threshold=threshold, device=None)
vertices.append(v) vertices.append(v)
faces.append(f) faces.append(f)
return io.NodeOutput(MESH(torch.stack(vertices), torch.stack(faces))) return io.NodeOutput(io.Mesh.MeshDict(vertices=torch.stack(vertices), faces=torch.stack(faces)))
class VoxelToMeshBasic(io.ComfyNode): class VoxelToMeshBasic(io.ComfyNode):
@ -645,20 +637,20 @@ class VoxelToMeshBasic(io.ComfyNode):
io.Float.Input("threshold", default=0.6, min=-1.0, max=1.0, step=0.01) io.Float.Input("threshold", default=0.6, min=-1.0, max=1.0, step=0.01)
], ],
outputs=[ outputs=[
io.Mesh.Output() io.Mesh.Output(),
] ]
) )
@classmethod @classmethod
def execute(cls, voxel, threshold): def execute(cls, voxel: io.Voxel.VoxelDict, threshold):
vertices = [] vertices = []
faces = [] faces = []
for x in voxel.data: for x in voxel["data"]:
v, f = voxel_to_mesh(x, threshold=threshold, device=None) v, f = voxel_to_mesh(x, threshold=threshold, device=None)
vertices.append(v) vertices.append(v)
faces.append(f) faces.append(f)
return io.NodeOutput(MESH(torch.stack(vertices), torch.stack(faces))) return io.NodeOutput(io.Mesh.MeshDict(vertices=torch.stack(vertices), faces=torch.stack(faces)))
NODES_LIST: list[type[io.ComfyNode]] = [ NODES_LIST: list[type[io.ComfyNode]] = [