testing and small fixes

This commit is contained in:
Yousef Rafat 2025-07-08 15:33:47 +03:00
parent 174655006c
commit 5b246946f6
3 changed files with 20 additions and 46 deletions

View File

@ -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))
return export_to_trimesh(mesh)

View File

@ -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

View File

@ -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]"""