diff --git a/comfy/ldm/hunyuan3d/model_/pipeline.py b/comfy/ldm/hunyuan3d/model_/pipeline.py index 65c50847c..ea144ecce 100644 --- a/comfy/ldm/hunyuan3d/model_/pipeline.py +++ b/comfy/ldm/hunyuan3d/model_/pipeline.py @@ -5,7 +5,6 @@ from PIL import Image from typing import List, Union from torch.utils._pytree import tree_map from torch.utils.data._utils.collate import default_collate -from vae import VAE def export_to_trimesh(mesh_output): if isinstance(mesh_output, list): @@ -170,7 +169,7 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module): noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond) # compute the previous noisy sample x_t -> x_t-1 - latents = self.scheduler.reverse_flow(noise_pred, latents) + latents = self.scheduler.step(noise_pred, latents) if callback is not None and i % callback_steps == 0: step_idx = i // getattr(self.scheduler, "order", 1) @@ -179,19 +178,4 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module): latents = 1. / self.vae.scale_factor * latents mesh = self.vae.decode(latents, bounds = bounds, octree_res = octree_res, num_chunks = num_chunks) - return export_to_trimesh(mesh) - -if __name__ == '__main__': - from scheduler import EulerScheduler - from conditioner import SingleImageEncoder - from image_processor import ImageProcessorV2 - from dinov2 import DinoConfig - from hunyuandit import HunYuanDiTPlain - - model = HunYuanDiTPlain(depth = 2) - - pipeline = Hunyuan3DDiTFlowMatchingPipeline(vae = VAE(), scheduler = EulerScheduler(), model = model, - conditioner = SingleImageEncoder(DinoConfig()), image_processor = ImageProcessorV2(), - device = "cpu", dtype = torch.bfloat16) - img = r"C:\Users\yrafa\Work\Hunyuan 3D\cat.jpg" - print(pipeline(img)) \ No newline at end of file + return export_to_trimesh(mesh) \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model_/scheduler.py b/comfy/ldm/hunyuan3d/model_/scheduler.py index a060f9623..200b9c56c 100644 --- a/comfy/ldm/hunyuan3d/model_/scheduler.py +++ b/comfy/ldm/hunyuan3d/model_/scheduler.py @@ -2,7 +2,7 @@ import torch class EulerScheduler(torch.nn.Module): def __init__(self, num_training_timesteps: int = 1_000, shift: float = 1, - num_inference_timesteps: int = 50, inference: bool = True): + num_inference_timesteps: int = 50, inference: bool = True, device: str = "cuda"): super(EulerScheduler, self).__init__() # compute timestep values so we can index into them later @@ -22,16 +22,13 @@ class EulerScheduler(torch.nn.Module): if inference: - sigmas = torch.linspace(0, 1, num_inference_timesteps) - sigmas = sigmas * shift / (1 + (shift - 1) * sigmas) - sigmas = sigmas.to(torch.float32) + sigmas = torch.linspace(0, 1, num_inference_timesteps, dtype = torch.float32, device = device) + timesteps = sigmas * self.num_training_timesteps - timesteps = sigmas * num_training_timesteps + self.timesteps = timesteps.to(device = device) + self.sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)]) - self.sigmas = torch.cat([sigmas, torch.ones(1, device = sigmas.device)]) - self.timesteps = timesteps.to(device = sigmas.device) - - self.step_index = 0 + self._step_index = 0 def sigma_to_timestep(self, sigma): return sigma * self.num_training_timesteps @@ -71,20 +68,18 @@ class EulerScheduler(torch.nn.Module): return noised_image @torch.no_grad() - def reverse_flow(self, current_sample: torch.Tensor, model_output: torch.FloatTensor): - - # upcast to avoid precision errors - current_sample = current_sample.to(torch.float32) + def step(self, model_output: torch.FloatTensor, sample: torch.FloatTensor,): + + sample = sample.to(torch.float32) - # get the current and next sigma and the change between them - current_sigma = self.sigmas[self.step_index] - next_sigma = self.sigmas[self.step_index + 1] - dt = next_sigma - current_sigma + sigma = self.sigmas[self._step_index] + sigma_next = self.sigmas[self._step_index + 1] + + prev_sample = sample + (sigma_next - sigma) * model_output - prev_sample = current_sample + dt * model_output prev_sample = prev_sample.to(model_output.dtype) - self.step_index += 1 + self._step_index += 1 return prev_sample diff --git a/comfy/ldm/hunyuan3d/model_/vae.py b/comfy/ldm/hunyuan3d/model_/vae.py index 7687b74c9..0ee313c9c 100644 --- a/comfy/ldm/hunyuan3d/model_/vae.py +++ b/comfy/ldm/hunyuan3d/model_/vae.py @@ -5,6 +5,10 @@ import torch from torch import Tensor import math +import trimesh +import numpy as np +from skimage import measure +from dataclasses import dataclass def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True): @@ -602,11 +606,6 @@ class PointCrossAttention(nn.Module): return latents -import torch -from skimage import measure -from dataclasses import dataclass -import numpy as np - @dataclass class Latent2MeshOutput(): # mesh for vertices and faces @@ -708,10 +707,6 @@ def export_to_trimesh(mesh_output): mesh_output = trimesh.Trimesh(mesh_output.mesh_v, mesh_output.mesh_f) return mesh_output -import trimesh -import torch -import numpy as np - def normalize_mesh(mesh, scale = 0.9999): """Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""