From 65155195075c5feab4f14dd91b5b8c84ca118163 Mon Sep 17 00:00:00 2001 From: Yousef Rafat <81116377+yousef-rafat@users.noreply.github.com> Date: Sat, 5 Jul 2025 20:57:24 +0300 Subject: [PATCH] Added Pipeline, Conditioner, Diffusion, Scheduler, and Image Processor --- comfy/ldm/hunyuan3d/model/conditioner.py | 145 ++++++ comfy/ldm/hunyuan3d/model/dinov2.py | 397 +++++++++++++++ comfy/ldm/hunyuan3d/model/hunyuandit.py | 490 +++++++++++++++++++ comfy/ldm/hunyuan3d/model/image_processor.py | 192 ++++++++ comfy/ldm/hunyuan3d/model/moe.py | 10 +- comfy/ldm/hunyuan3d/model/pipeline.py | 156 ++++++ comfy/ldm/hunyuan3d/model/scheduler.py | 105 ++++ comfy/ldm/hunyuan3d/vae/vae.py | 6 +- 8 files changed, 1495 insertions(+), 6 deletions(-) create mode 100644 comfy/ldm/hunyuan3d/model/conditioner.py create mode 100644 comfy/ldm/hunyuan3d/model/dinov2.py create mode 100644 comfy/ldm/hunyuan3d/model/hunyuandit.py create mode 100644 comfy/ldm/hunyuan3d/model/image_processor.py create mode 100644 comfy/ldm/hunyuan3d/model/pipeline.py create mode 100644 comfy/ldm/hunyuan3d/model/scheduler.py diff --git a/comfy/ldm/hunyuan3d/model/conditioner.py b/comfy/ldm/hunyuan3d/model/conditioner.py new file mode 100644 index 000000000..d56ffd34b --- /dev/null +++ b/comfy/ldm/hunyuan3d/model/conditioner.py @@ -0,0 +1,145 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +from dinov2 import Dinov2Model, DinoConfig + +# avoid using torchvision by recreating image processing functions + +def resize(img: torch.Tensor, size: int) -> torch.Tensor: + + _, h, w = img.shape + + # mantain aspect ratio + if h < w: + new_h = size + new_w = int(w * size / h) + else: + new_w = size + new_h = int(h * size / w) + + img = img.unsqueeze(0) + img = F.interpolate(img, size = (new_h, new_w), mode = 'bilinear', align_corners = False, antialias = True ) + return img.squeeze(0) + + +def center_crop(img: torch.Tensor, size: int) -> torch.Tensor: + + _, h, w = img.shape + top = (h - size) // 2 + left = (w - size) // 2 + + return img[:, top:top + size, left:left + size] + +def normalize(img: torch.Tensor, mean: list, std: list) -> torch.Tensor: + + mean = torch.tensor(mean, device = img.device).view(-1, 1, 1) + std = torch.tensor(std, device = img.device).view(-1, 1, 1) + return (img - mean) / std + +def compose(transforms): + def apply(img): + for t in transforms: + img = t(img) + return img + return apply + +class ImageEncoder(nn.Module): + def __init__( + self, + config: DinoConfig, + use_cls_token=True, + image_size=224, + **kwargs, + ): + super().__init__() + + self.model = Dinov2Model(config) + + mean = [0.485, 0.456, 0.406] + std = [0.229, 0.224, 0.225] + + self.model.eval() + self.model.requires_grad_(False) + + self.use_cls_token = use_cls_token + self.size = image_size // 14 + + self.num_patches = (image_size // 14) ** 2 + + if self.use_cls_token: + self.num_patches += 1 + + self.transform = compose([ + lambda x: resize(x, image_size), + lambda x: center_crop(x, image_size), + lambda x: normalize(x, mean, std), + ]) + + def forward(self, image, value_range=(-1, 1), **kwargs): + + if value_range is not None: + low, high = value_range + image = (image - low) / (high - low) + + image = image.to(self.model.device, dtype=self.model.dtype) + inputs = self.transform(image) + outputs = self.model(inputs) + + last_hidden_state = outputs.last_hidden_state + + if not self.use_cls_token: + last_hidden_state = last_hidden_state[:, 1:, :] + + return last_hidden_state + + def unconditional_embedding(self, batch_size, **kwargs): + + device = next(self.model.parameters()).device + dtype = next(self.model.parameters()).dtype + + zero = torch.zeros( + batch_size, + self.num_patches, + self.model.config.hidden_size, + device=device, + dtype=dtype, + ) + + return zero + +class SingleImageEncoder(nn.Module): + def __init__(self): + super().__init__() + self.main_image_encoder = ImageEncoder() + + def forward(self, image, **kwargs): + outputs = { + 'main': self.main_image_encoder(image, **kwargs), + } + return outputs + + def unconditional_embedding(self, batch_size, **kwargs): + outputs = { + 'main': self.main_image_encoder.unconditional_embedding(batch_size, **kwargs), + } + return outputs + +def load_dino2(dino2: Dinov2Model): + + checkpoint = "" + dino2.load_state_dict(torch.load(checkpoint)) + return dino2 + +def test_image_encoder(): + + torch.manual_seed(2025) + image_encoder = SingleImageEncoder() + + image = torch.rand(1, 3, 224, 224) + + outputs = image_encoder(image) + + print(outputs) + +if __name__ == "__main__": + test_image_encoder() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model/dinov2.py b/comfy/ldm/hunyuan3d/model/dinov2.py new file mode 100644 index 000000000..7ff44d466 --- /dev/null +++ b/comfy/ldm/hunyuan3d/model/dinov2.py @@ -0,0 +1,397 @@ +from dataclasses import dataclass +from typing import Optional +import collections.abc +import torch.nn as nn +import torch + +@dataclass +class DinoConfig(): + + hidden_size: int = 1024 + use_mask_token: bool = True + patch_size: int = 14 + image_size: int = 518 + num_channels: int = 3 + num_attention_heads: int = 16 + attention_probs_dropout_prob: float = 0.0 + hidden_dropout_prob: float = 0.0 + mlp_ratio: int = 4 + num_hidden_layers: int = 24 + layer_norm_eps: float = 1e-6 + qkv_bias: bool = True + layerscale_value: float = 1.0 + drop_path_rate: float = 0.0 + +class Dinov2Embeddings(nn.Module): + """ + Construct the CLS token, mask token, position and patch embeddings. + """ + + def __init__(self, config) -> None: + super().__init__() + + self.cls_token = nn.Parameter(torch.randn(1, 1, config.hidden_size)) + + if config.use_mask_token: + self.mask_token = nn.Parameter(torch.zeros(1, config.hidden_size)) + + self.patch_embeddings = Dinov2PatchEmbeddings(config) + num_patches = self.patch_embeddings.num_patches + + self.position_embeddings = nn.Parameter(torch.randn(1, num_patches + 1, config.hidden_size)) + self.dropout = nn.Dropout(config.hidden_dropout_prob) + + self.patch_size = config.patch_size + self.use_mask_token = config.use_mask_token + + def interpolate_pos_encoding(self, embeddings: torch.Tensor, height: int, width: int) -> torch.Tensor: + + num_patches = embeddings.shape[1] - 1 + num_positions = self.position_embeddings.shape[1] - 1 + + # always interpolate when tracing to ensure the exported model works for dynamic input shapes + if not torch.jit.is_tracing() and num_patches == num_positions and height == width: + return self.position_embeddings + + class_pos_embed = self.position_embeddings[:, :1] + patch_pos_embed = self.position_embeddings[:, 1:] + + dim = embeddings.shape[-1] + + new_height = height // self.patch_size + new_width = width // self.patch_size + + sqrt_num_positions = int(num_positions**0.5) + patch_pos_embed = patch_pos_embed.reshape(1, sqrt_num_positions, sqrt_num_positions, dim) + + patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2) + target_dtype = patch_pos_embed.dtype + + patch_pos_embed = nn.functional.interpolate( + patch_pos_embed.to(torch.float32), + size=(new_height, new_width), + mode="bicubic", + align_corners=False, + ).to(dtype=target_dtype) + + patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim) + + return torch.cat((class_pos_embed, patch_pos_embed), dim=1) + + def forward(self, pixel_values: torch.Tensor, bool_masked_pos: torch.Tensor = None) -> torch.Tensor: + + batch_size, _, height, width = pixel_values.shape + target_dtype = self.patch_embeddings.projection.weight.dtype + embeddings = self.patch_embeddings(pixel_values.to(dtype=target_dtype)) + + if bool_masked_pos is not None and self.use_mask_token: + embeddings = torch.where( + bool_masked_pos.unsqueeze(-1), self.mask_token.to(embeddings.dtype).unsqueeze(0), embeddings + ) + + # add the [CLS] token to the embedded patch tokens + cls_tokens = self.cls_token.expand(batch_size, -1, -1) + embeddings = torch.cat((cls_tokens, embeddings), dim=1) + + # add positional encoding to each token + embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width) + + embeddings = self.dropout(embeddings) + + return embeddings + + +class Dinov2PatchEmbeddings(nn.Module): + """ + This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial + `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a + Transformer. + """ + + def __init__(self, config): + super().__init__() + + image_size, patch_size = config.image_size, config.patch_size + num_channels, hidden_size = config.num_channels, config.hidden_size + + image_size = image_size if isinstance(image_size, collections.abc.Iterable) else (image_size, image_size) + patch_size = patch_size if isinstance(patch_size, collections.abc.Iterable) else (patch_size, patch_size) + num_patches = (image_size[1] // patch_size[1]) * (image_size[0] // patch_size[0]) + + self.image_size = image_size + self.patch_size = patch_size + + self.num_channels = num_channels + self.num_patches = num_patches + + self.projection = nn.Conv2d(num_channels, hidden_size, kernel_size=patch_size, stride=patch_size) + + def forward(self, pixel_values: torch.Tensor) -> torch.Tensor: + num_channels = pixel_values.shape[1] + if pixel_values.shape[1] != self.num_channels: + raise ValueError( + "Make sure that the channel dimension of the pixel values match with the one set in the configuration." + f" Expected {self.num_channels} but got {num_channels}." + ) + return self.projection(pixel_values).flatten(2).transpose(1, 2) + +def eager_attention_forward( + module: nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + scaling: float, + dropout: float = 0.0, + **kwargs, +): + # Take the dot product between "query" and "key" to get the raw attention scores. + attn_weights = torch.matmul(query, key.transpose(-1, -2)) * scaling + + # Normalize the attention scores to probabilities. + attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) + + # This is actually dropping out entire tokens to attend to, which might + # seem a bit unusual, but is taken from the original Transformer paper. + attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) + + attn_output = torch.matmul(attn_weights, value) + attn_output = attn_output.transpose(1, 2).contiguous() + + return attn_output, attn_weights + + + +# Copied from transformers.models.vit.modeling_vit.ViTSelfAttention with ViT->Dinov2 +class Dinov2SelfAttention(nn.Module): + def __init__(self, config) -> None: + super().__init__() + + self.config = config + self.num_attention_heads = config.num_attention_heads + + self.attention_head_size = int(config.hidden_size / config.num_attention_heads) + self.all_head_size = self.num_attention_heads * self.attention_head_size + + self.dropout_prob = config.attention_probs_dropout_prob + self.scaling = self.attention_head_size**-0.5 + + self.query = nn.Linear(config.hidden_size, self.all_head_size, bias=config.qkv_bias) + self.key = nn.Linear(config.hidden_size, self.all_head_size, bias=config.qkv_bias) + self.value = nn.Linear(config.hidden_size, self.all_head_size, bias=config.qkv_bias) + + def transpose_for_scores(self, x: torch.Tensor) -> torch.Tensor: + new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size) + x = x.view(new_x_shape) + return x.permute(0, 2, 1, 3) + + def forward( + self, hidden_states, head_mask: torch.Tensor = None + ): + key_layer = self.transpose_for_scores(self.key(hidden_states)) + value_layer = self.transpose_for_scores(self.value(hidden_states)) + query_layer = self.transpose_for_scores(self.query(hidden_states)) + + context_layer, _ = eager_attention_forward( + self, + query = query_layer, + key = key_layer, + value = value_layer, + scaling = self.scaling, + dropout = 0.0 if not self.training else self.dropout_prob, + ) + + new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,) + context_layer = context_layer.reshape(new_context_layer_shape) + + outputs = (context_layer,) + + return outputs + +class Dinov2SelfOutput(nn.Module): + """ + The residual connection is defined in Dinov2Layer instead of here (as is the case with other models), due to the + layernorm applied before each block. + """ + + def __init__(self, config: DinoConfig) -> None: + super().__init__() + self.dense = nn.Linear(config.hidden_size, config.hidden_size) + self.dropout = nn.Dropout(config.hidden_dropout_prob) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = self.dense(hidden_states) + hidden_states = self.dropout(hidden_states) + + return hidden_states + + +# Copied from transformers.models.vit.modeling_vit.ViTAttention with ViT->Dinov2 +class Dinov2Attention(nn.Module): + def __init__(self, config: DinoConfig) -> None: + super().__init__() + self.attention = Dinov2SelfAttention(config) + self.output = Dinov2SelfOutput(config) + self.pruned_heads = set() + + def forward( + self, + hidden_states: torch.Tensor, + head_mask: torch.Tensor = None, + ): + self_outputs = self.attention(hidden_states, head_mask) + + attention_output = self.output(self_outputs[0]) + + outputs = (attention_output,) + return outputs + + +class Dinov2LayerScale(nn.Module): + def __init__(self, config) -> None: + super().__init__() + self.lambda1 = nn.Parameter(config.layerscale_value * torch.ones(config.hidden_size)) + + def forward(self, hidden_state: torch.Tensor) -> torch.Tensor: + return hidden_state * self.lambda1 + +# Copied from transformers.models.beit.modeling_beit.drop_path +def drop_path(input: torch.Tensor, drop_prob: float = 0.0, training: bool = False) -> torch.Tensor: + + if drop_prob == 0.0 or not training: + return input + + keep_prob = 1 - drop_prob + shape = (input.shape[0],) + (1,) * (input.ndim - 1) # work with diff dim tensors, not just 2D ConvNets + + random_tensor = keep_prob + torch.rand(shape, dtype=input.dtype, device=input.device) + random_tensor.floor_() # binarize + + output = input.div(keep_prob) * random_tensor + + return output + + +# Copied from transformers.models.beit.modeling_beit.BeitDropPath +class Dinov2DropPath(nn.Module): + """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).""" + + def __init__(self, drop_prob: float = None) -> None: + super().__init__() + self.drop_prob = drop_prob + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return drop_path(hidden_states, self.drop_prob, self.training) + + +class Dinov2MLP(nn.Module): + def __init__(self, config): + super().__init__() + + in_features = out_features = config.hidden_size + + hidden_features = int(config.hidden_size * config.mlp_ratio) + self.fc1 = nn.Linear(in_features, hidden_features, bias=True) + + self.activation = nn.GELU() + self.fc2 = nn.Linear(hidden_features, out_features, bias=True) + + def forward(self, hidden_state: torch.Tensor) -> torch.Tensor: + + hidden_state = self.fc1(hidden_state) + hidden_state = self.activation(hidden_state) + hidden_state = self.fc2(hidden_state) + + return hidden_state + +class Dinov2Layer(nn.Module): + """This corresponds to the Block class in the original implementation.""" + + def __init__(self, config: DinoConfig) -> None: + super().__init__() + + self.norm1 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + self.attention = Dinov2Attention(config) + self.layer_scale1 = Dinov2LayerScale(config) + self.drop_path = Dinov2DropPath(config.drop_path_rate) if config.drop_path_rate > 0.0 else nn.Identity() + + self.norm2 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + + self.mlp = Dinov2MLP(config) + self.layer_scale2 = Dinov2LayerScale(config) + + def forward( + self, + hidden_states: torch.Tensor, + head_mask: torch.Tensor = None, + ): + self_attention_outputs = self.attention( + self.norm1(hidden_states), # in Dinov2, layernorm is applied before self-attention + head_mask, + ) + attention_output = self_attention_outputs[0] + + attention_output = self.layer_scale1(attention_output) + + # first residual connection + hidden_states = self.drop_path(attention_output) + hidden_states + + # in Dinov2, layernorm is also applied after self-attention + layer_output = self.norm2(hidden_states) + layer_output = self.mlp(layer_output) + layer_output = self.layer_scale2(layer_output) + + # second residual connection + layer_output = self.drop_path(layer_output) + hidden_states + + outputs = (layer_output,) + + return outputs + +class Dinov2Encoder(nn.Module): + def __init__(self, config: DinoConfig) -> None: + super().__init__() + self.layer = nn.ModuleList([Dinov2Layer(config) for _ in range(config.num_hidden_layers)]) + + def forward(self, hidden_states: torch.Tensor, head_mask: Optional[torch.Tensor] = None): + + for i, layer_module in enumerate(self.layer): + + layer_head_mask = head_mask[i] if head_mask is not None else None + + layer_outputs = layer_module(hidden_states, layer_head_mask) + + hidden_states = layer_outputs[0] + + return hidden_states + + +class Dinov2Model(nn.Module): + def __init__(self, config: DinoConfig): + super().__init__() + self.config = config + + self.embeddings = Dinov2Embeddings(config) + self.encoder = Dinov2Encoder(config) + + self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + + def get_input_embeddings(self) -> Dinov2PatchEmbeddings: + return self.embeddings.patch_embeddings + + def forward( + self, + pixel_values: torch.Tensor, + bool_masked_pos: Optional[torch.Tensor] = None, + head_mask: Optional[torch.Tensor] = None, + ): + + embedding_output = self.embeddings(pixel_values, bool_masked_pos=bool_masked_pos) + + encoder_outputs = self.encoder( + embedding_output, + head_mask = head_mask, + ) + sequence_output = encoder_outputs[0] + sequence_output = self.layernorm(sequence_output) + + return sequence_output \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model/hunyuandit.py b/comfy/ldm/hunyuan3d/model/hunyuandit.py new file mode 100644 index 000000000..d0894e46a --- /dev/null +++ b/comfy/ldm/hunyuan3d/model/hunyuandit.py @@ -0,0 +1,490 @@ +import math +import torch +import torch.nn as nn +import torch.nn.functional as F +from moe import MoEBlock +from torch.nn.attention import SDPBackend + +class Timesteps(nn.Module): + def __init__(self, num_channels: int, downscale_freq_shift: float = 0.0, + scale: float = 1.0, max_period: int = 10000): + super().__init__() + + self.num_channels = num_channels + half_dim = num_channels // 2 + + # precompute the “inv_freq” vector once + exponent = -math.log(max_period) * torch.arange( + half_dim, dtype=torch.float32 + ) / (half_dim - downscale_freq_shift) + + inv_freq = torch.exp(exponent) + + # pad + if num_channels % 2 == 1: + # we’ll pad a zero at the end of the cos-half + inv_freq = torch.cat([inv_freq, inv_freq.new_zeros(1)]) + + # register to buffer so it moves with the device + self.register_buffer("inv_freq", inv_freq, persistent = False) + self.scale = scale + + def forward(self, timesteps: torch.Tensor): + + x = timesteps.float().unsqueeze(1) * self.inv_freq.unsqueeze(0) + + # scale factor + if self.scale != 1.0: + emb = emb * self.scale + + # fused CUDA kernels for sin and cos + sin_emb = x.sin() + cos_emb = x.cos() + + emb = torch.cat([sin_emb, cos_emb], dim = 1) + + # If we padded inv_freq for odd, emb is already wide enough; otherwise: + if emb.shape[1] > self.num_channels: + emb = emb[:, :self.num_channels] + + return emb + +class TimestepEmbedder(nn.Module): + def __init__(self, hidden_size, frequency_embedding_size = 256, cond_proj_dim = None): + super().__init__() + + self.mlp = nn.Sequential( + nn.Linear(hidden_size, frequency_embedding_size, bias=True), + nn.GELU(), + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + + if cond_proj_dim is not None: + self.cond_proj = nn.Linear(cond_proj_dim, frequency_embedding_size, bias=False) + + self.time_embed = Timesteps(hidden_size) + + def forward(self, timesteps, condition): + + timestep_embed = self.time_embed(timesteps).type(self.mlp[0].weight.dtype) + + if condition is not None: + cond_embed = self.cond_proj(condition) + timestep_embed = timestep_embed + cond_embed + + time_conditioned = self.mlp(timestep_embed) + + # for broadcasting with image tokens + return time_conditioned.unsqueeze(1) + +class MLP(nn.Module): + def __init__(self, *, width: int): + super().__init__() + self.width = width + self.fc1 = nn.Linear(width, width * 4) + self.fc2 = nn.Linear(width * 4, width) + self.gelu = nn.GELU() + + def forward(self, x): + return self.fc2(self.gelu(self.fc1(x))) + +class CrossAttention(nn.Module): + def __init__( + self, + qdim, + kdim, + num_heads, + qkv_bias=True, + qk_norm=False, + norm_layer=nn.LayerNorm, + use_fp16: bool = False, + **kwargs, + ): + super().__init__() + self.qdim = qdim + self.kdim = kdim + + self.num_heads = num_heads + self.head_dim = self.qdim // num_heads + + self.scale = self.head_dim ** -0.5 + + self.to_q = nn.Linear(qdim, qdim, bias=qkv_bias) + self.to_k = nn.Linear(kdim, qdim, bias=qkv_bias) + self.to_v = nn.Linear(kdim, qdim, bias=qkv_bias) + + if use_fp16: + eps = 1.0 / 65504 + else: + eps = 1e-6 + + self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity() + self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity() + self.out_proj = nn.Linear(qdim, qdim, bias=True) + + self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity() + self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity() + self.out_proj = nn.Linear(qdim, qdim, bias=True) + + def forward(self, x, y): + + b, s1, _ = x.shape + _, s2, _ = y.shape + + q = self.to_q(x) + k = self.to_k(y) + v = self.to_v(y) + + kv = torch.cat((k, v), dim=-1) + split_size = kv.shape[-1] // self.num_heads // 2 + + kv = kv.view(1, -1, self.num_heads, split_size * 2) + k, v = torch.split(kv, split_size, dim=-1) + + q = q.view(b, s1, self.num_heads, self.head_dim) + k = k.view(b, s2, self.num_heads, self.head_dim) + v = v.view(b, s2, self.num_heads, self.head_dim) + + q = self.q_norm(q) + k = self.k_norm(k) + + # replaced with torch.nn.attention (avoid FutureWarning from backends.cuda.sdp_kerenl) + with torch.nn.attention.sdpa_kernel( + backends=[ + SDPBackend.FLASH_ATTENTION, + SDPBackend.MATH, + SDPBackend.EFFICIENT_ATTENTION, + ] + ): + q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)] + context = F.scaled_dot_product_attention( + q, k, v + ).transpose(1, 2).reshape(b, s1, -1) + + out = self.out_proj(context) + + return out + +class Attention(nn.Module): + + def __init__( + self, + dim, + num_heads, + qkv_bias = True, + qk_norm = False, + norm_layer = nn.LayerNorm, + use_fp16: bool = False + ): + super().__init__() + self.dim = dim + self.num_heads = num_heads + self.head_dim = self.dim // num_heads + self.scale = self.head_dim ** -0.5 + + self.to_q = nn.Linear(dim, dim, bias = qkv_bias) + self.to_k = nn.Linear(dim, dim, bias = qkv_bias) + self.to_v = nn.Linear(dim, dim, bias = qkv_bias) + + if use_fp16: + eps = 1.0 / 65504 + else: eps = 1e-6 + + self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity() + self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity() + self.out_proj = nn.Linear(dim, dim) + + def forward(self, x): + B, N, _ = x.shape + + query = self.to_q(x) + key = self.to_k(x) + value = self.to_v(x) + + qkv_combined = torch.cat((query, key, value), dim=-1) + split_size = qkv_combined.shape[-1] // self.num_heads // 3 + + qkv = qkv_combined.view(1, -1, self.num_heads, split_size * 3) + query, key, value = torch.split(qkv, split_size, dim=-1) + + query = query.reshape(B, N, self.num_heads, self.head_dim).transpose(1, 2) + key = key.reshape(B, N, self.num_heads, self.head_dim).transpose(1, 2) + value = value.reshape(B, N, self.num_heads, self.head_dim).transpose(1, 2) + + query = self.q_norm(query) + key = self.k_norm(key) + + # replaced with torch.nn.attention (avoid FutureWarning from backends.cuda.sdp_kerenl) + with torch.nn.attention.sdpa_kernel( + backends=[ + SDPBackend.FLASH_ATTENTION, + SDPBackend.MATH, + SDPBackend.EFFICIENT_ATTENTION, + ] + ): + x = F.scaled_dot_product_attention(query, key, value) + x = x.transpose(1, 2).reshape(B, N, -1) + + x = self.out_proj(x) + return x + +class HunYuanDiTBlock(nn.Module): + def __init__( + self, + hidden_size, + c_emb_size, + num_heads, + text_states_dim=1024, + qk_norm=False, + norm_layer=nn.LayerNorm, + qk_norm_layer=nn.RMSNorm, + qkv_bias=True, + skip_connection=True, + timested_modulate=False, + use_moe: bool = False, + num_experts: int = 8, + moe_top_k: int = 2, + use_fp16: bool = False + ): + super().__init__() + + # eps can't be 1e-6 in fp16 mode because of numerical stability issues + if use_fp16: + eps = 1.0 / 65504 + else: eps = 1e-6 + + self.norm1 = norm_layer(hidden_size, elementwise_affine = True, eps = eps) + + self.attn1 = Attention(hidden_size, num_heads=num_heads, qkv_bias=qkv_bias, qk_norm=qk_norm, + norm_layer=qk_norm_layer, use_fp16 = use_fp16) + + self.norm2 = norm_layer(hidden_size, elementwise_affine = True, eps = eps) + + self.timested_modulate = timested_modulate + if self.timested_modulate: + self.default_modulation = nn.Sequential( + nn.SiLU(), + nn.Linear(c_emb_size, hidden_size, bias=True) + ) + + self.attn2 = CrossAttention(hidden_size, text_states_dim, num_heads=num_heads, qkv_bias=qkv_bias, + qk_norm=qk_norm, norm_layer=qk_norm_layer, use_fp16 = use_fp16) + + self.norm3 = norm_layer(hidden_size, elementwise_affine = True, eps = eps) + + if skip_connection: + self.skip_norm = norm_layer(hidden_size, elementwise_affine = True, eps = eps) + self.skip_linear = nn.Linear(2 * hidden_size, hidden_size) + else: + self.skip_linear = None + + self.use_moe = use_moe + + if self.use_moe: + self.moe = MoEBlock( + hidden_size, + num_experts = num_experts, + moe_top_k = moe_top_k, + dropout = 0.0, + ff_inner_dim = int(hidden_size * 4.0), + ) + else: + self.mlp = MLP(width=hidden_size) + + def forward(self, hidden_states, conditioning=None, text_states=None, skip_tensor=None): + + if self.skip_linear is not None: + combined = torch.cat([skip_tensor, hidden_states], dim=-1) + hidden_states = self.skip_linear(combined) + hidden_states = self.skip_norm(hidden_states) + + # self attention + if self.timested_modulate: + modulation_shift = self.default_modulation(conditioning).unsqueeze(dim=1) + hidden_states = hidden_states + modulation_shift + + self_attn_out = self.attn1(self.norm1(hidden_states)) + hidden_states = hidden_states + self_attn_out + + # cross attention + hidden_states = hidden_states + self.attn2(self.norm2(hidden_states), text_states) + + # MLP Layer + mlp_input = self.norm3(hidden_states) + + if self.use_moe: + hidden_states = hidden_states + self.moe(mlp_input) + else: + hidden_states = hidden_states + self.mlp(mlp_input) + + return hidden_states + +class FinalLayer(nn.Module): + + def __init__(self, final_hidden_size, out_channels, use_fp16: bool = False): + super().__init__() + + if use_fp16: + eps = 1.0 / 65504 + else: eps = 1e-6 + + self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine = True, eps = eps) + self.linear = nn.Linear(final_hidden_size, out_channels, bias = True) + + def forward(self, x): + x = self.norm_final(x) + x = x[:, 1:] + x = self.linear(x) + return x + +class HunYuanDiTPlain(nn.Module): + + # init with the defaults values from https://huggingface.co/tencent/Hunyuan3D-2.1/blob/main/hunyuan3d-dit-v2-1/config.yaml + def __init__( + self, + in_channels: int = 64, + hidden_size: int = 2048, + context_dim: int = 1024, + depth: int = 21, + num_heads: int = 16, + qk_norm: bool = True, + qkv_bias: bool = False, + num_moe_layers: int = 6, + guidance_cond_proj_dim = None, + norm_type = 'layer', + num_experts: int = 8, + moe_top_k: int = 2, + use_fp16: bool = False + ): + + super().__init__() + + self.depth = depth + + self.in_channels = in_channels + self.out_channels = in_channels + + self.num_heads = num_heads + self.hidden_size = hidden_size + + norm = nn.LayerNorm if norm_type == 'layer' else nn.RMSNorm + qk_norm = nn.RMSNorm + + self.context_dim = context_dim + self.guidance_cond_proj_dim = guidance_cond_proj_dim + + self.x_embedder = nn.Linear(in_channels, hidden_size, bias = True) + self.t_embedder = TimestepEmbedder(hidden_size, hidden_size * 4, cond_proj_dim=guidance_cond_proj_dim) + + + # HUnYuanDiT Blocks + self.blocks = nn.ModuleList([ + HunYuanDiTBlock(hidden_size=hidden_size, + c_emb_size=hidden_size, + num_heads=num_heads, + text_states_dim=context_dim, + qk_norm=qk_norm, + norm_layer = norm, + qk_norm_layer = qk_norm, + skip_connection=layer > depth // 2, + qkv_bias=qkv_bias, + use_moe=True if depth - layer <= num_moe_layers else False, + num_experts=num_experts, + moe_top_k=moe_top_k, + use_fp16 = use_fp16) + for layer in range(depth) + ]) + + self.depth = depth + + self.final_layer = FinalLayer(hidden_size, self.out_channels, use_fp16 = use_fp16) + + def forward(self, x, t, contexts, **kwargs): + + main_condition = contexts['main'] + + time_embedded = self.t_embedder(t, condition=kwargs.get('guidance_cond')) + x_embedded = self.x_embedder(x) + + combined = torch.cat([time_embedded, x_embedded], dim=1) + + skip_stack = [] + for idx, block in enumerate(self.blocks): + if idx <= self.depth // 2: + skip_input = None + else: + skip_input = skip_stack.pop() + + combined = block(combined, time_embedded, main_condition, skip_tensor = skip_input) + + if idx < self.depth // 2: + skip_stack.append(combined) + + output = self.final_layer(combined) + return output + +def get_diffusion_checkpoint(): + import requests + + url = "https://huggingface.co/tencent/Hunyuan3D-2.1/resolve/main/hunyuan3d-dit-v2-1/model.fp16.ckpt" + output_path = "model.fp16.ckpt" + + response = requests.get(url, stream=True) + response.raise_for_status() + + with open(output_path, "wb") as f: + for chunk in response.iter_content(chunk_size=8192): + if chunk: + f.write(chunk) + + print(f"Downloaded to: {output_path}") + +def load_dit(dit: HunYuanDiTPlain): + + DEBUG = False + checkpoint = torch.load("model.fp16.ckpt") + missing, unexpected = dit.load_state_dict(checkpoint["model"], strict = not DEBUG) + + if DEBUG: + print(f"Missing {len(missing)}", missing) + print(f"Unexpected {len(unexpected)}", unexpected) + + return dit + +if __name__ == "__main__": + + torch.manual_seed(2025) + torch.set_default_device("cpu") + torch.set_default_dtype(torch.bfloat16) + + import time + + timings = {} + + start = time.time() + + model = HunYuanDiTPlain(depth = 10) + + timings["model_initialization"] = time.time() - start + + batch_size = 2 + seq_len = 1370 + in_channels = 64 + context_dim = 1024 + + # Random inputs + x = torch.randn(batch_size, seq_len, in_channels) + t = torch.randint(0, seq_len, (batch_size,)) + contexts = { + 'main': torch.randn(batch_size, seq_len, context_dim) + } + + # Forward pass + start = time.time() + output = model(x, t, contexts) + timings["forward_timing"] = time.time() - start + + print("\n=== Timing Summary ===") + for key, value in timings.items(): + print(f"{key}: {value:.3f} seconds") \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model/image_processor.py b/comfy/ldm/hunyuan3d/model/image_processor.py new file mode 100644 index 000000000..f98d895b8 --- /dev/null +++ b/comfy/ldm/hunyuan3d/model/image_processor.py @@ -0,0 +1,192 @@ +import torch +import torch.nn as nn +import numpy as np +from PIL import Image +import torch.nn.functional as F + +def to_tensor(image_pt): + + image_pt = image_pt / 255 * 2 - 1 + if image_pt.dim() == 4: + image_pt = image_pt.permute(0, 3, 1, 2) + else: + image_pt = image_pt.permute(2, 1, 0) + + return image_pt + +def resize_bilinear(img: torch.Tensor, size: int) -> torch.Tensor: + # pytorch implementation of cv2.INTER_LINEAR + + batched = (img.ndim == 4) + if img.ndim == 3: + img = img.unsqueeze(0) + + B, _, H, W = img.shape + H_out = W_out = size + + xs = torch.linspace(0, H_out - 1, H_out, device = img.device) + ys = torch.linspace(0, W_out - 1, W_out, device = img.device) + + xs = (xs + 0.5) * (H / H_out) - 0.5 + ys = (ys + 0.5) * (W / W_out) - 0.5 + + # normalize + xs = 2 * xs / (H - 1) - 1 + ys = 2 * ys / (W - 1) - 1 + + # meshgrid in “ij” order: first rows (xs), then cols (ys) + grid_i, grid_j = torch.meshgrid(xs, ys, indexing='ij') + + # stack into (x,y) where x=columns, y=rows + grid = torch.stack((grid_j, grid_i), dim=-1) + grid = grid.unsqueeze(0).expand(B, -1, -1, -1) + + out = F.grid_sample(img, grid, mode = 'bilinear', + padding_mode = 'zeros', align_corners = True) + + return out if batched else out.squeeze(0) + +def resize_bicubic(img: torch.Tensor, size: int) -> torch.Tensor: + # pytorch implementation of INTER_CUBIC + + was_batched = img.ndim == 4 + if img.ndim == 3: + img = img.unsqueeze(0) + + out = F.interpolate( + img.permute(0, 3, 2, 1), + size = (size, size), + mode = "bicubic", + align_corners = True + ) + + return out if was_batched else out.squeeze(0) + +def resize_area(img: torch.Tensor, size: tuple) -> torch.Tensor: + # pytorch implementation of INTER_AREA + + was_batched = img.ndim == 4 + if img.ndim == 3: + img = img.unsqueeze(0) + + image = F.interpolate(img.permute(0,3,1,2).float(), (size[1], size[0]), mode = "area") + + if was_batched: + image = image.permute(0, 2, 3, 1) # return to channel last + else: + image = image.squeeze(0).permute(1, 2, 0) + + return image if was_batched else image.squeeze(0) + +class ImageProcessorV2(nn.Module): + def __init__(self, size: int = 512, border_ratio: float = None): + self.size = size + self.border_ratio = border_ratio + + def load_image(self, pic, border_ratio: float = 0.15) -> torch.Tensor: + + if isinstance(pic, str): + img = Image.open(pic) + img = np.array(img) + + elif isinstance(pic, Image.Image): + img = np.array(pic) + + if img.ndim == 2: # grayscale + img = img[:, :, None] + + img = torch.from_numpy(img) + img, mask = self.recenter(img, border_ratio = border_ratio) + + img = resize_bicubic(img, size = self.size) + mask = resize_bilinear(mask.float(), size = self.size) + mask = mask[..., torch.newaxis] + + img = to_tensor(img) + mask = to_tensor(mask) + + return img, mask + + @staticmethod + def recenter(image, border_ratio: float = 0.2): + + if image.shape[-1] == 4: + mask = image[..., 3] + else: + mask = torch.ones_like(image[..., 0:1]) * 255 + image = torch.concatenate([image, mask], axis=-1) + mask = mask[..., 0] + + H, W, C = image.shape + + size = max(H, W) + result = torch.zeros((size, size, C), dtype = torch.uint8) + + # as_tuple to match numpy behaviour + x_coords, y_coords = torch.nonzero(mask, as_tuple=True) + + y_min, y_max = y_coords.min(), y_coords.max() + x_min, x_max = x_coords.min(), x_coords.max() + + h = x_max - x_min + w = y_max - y_min + + if h == 0 or w == 0: + raise ValueError('input image is empty') + + desired_size = int(size * (1 - border_ratio)) + scale = desired_size / max(h, w) + + h2 = int(h * scale) + w2 = int(w * scale) + + x2_min = (size - h2) // 2 + x2_max = x2_min + h2 + + y2_min = (size - w2) // 2 + y2_max = y2_min + w2 + + # note: opencv takes columns first (opposite to pytorch and numpy that take the row first) + result[x2_min:x2_max, y2_min:y2_max] = resize_area(image[x_min:x_max, y_min:y_max], (w2, h2)) + + bg = torch.ones((result.shape[0], result.shape[1], 3), dtype = torch.uint8) * 255 + + mask = result[..., 3:].to(torch.float32) / 255 + result = result[..., :3] * mask + bg * (1 - mask) + + mask = mask * 255 + result = result.clip(0, 255).to(torch.uint8) + mask = mask.clip(0, 255).to(torch.uint8) + + return result, mask + + def __call__(self, image, border_ratio = 0.15, **kwargs): + + if self.border_ratio is not None: + border_ratio = self.border_ratio + + image, mask = self.load_image(image, border_ratio = border_ratio) + + outputs = { + 'image': image, + 'mask': mask + } + + return outputs + +def test_image_processor(): + import matplotlib.pyplot as plt + + image_processor = ImageProcessorV2(size = 224) + import time + start = time.time() + outputs = image_processor(image = r"C:\Users\yrafa\Work\Hunyuan 3D\cat.jpg") + print(time.time() - start) + image = outputs["image"] + print(image.shape) + plt.imshow(image) + plt.axis("off") + plt.show() + +if __name__ == "__main__": + test_image_processor() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model/moe.py b/comfy/ldm/hunyuan3d/model/moe.py index f6ea7db3c..d09499714 100644 --- a/comfy/ldm/hunyuan3d/model/moe.py +++ b/comfy/ldm/hunyuan3d/model/moe.py @@ -51,7 +51,6 @@ class AddAuxLoss(torch.autograd.Function): @staticmethod def forward(ctx, x, loss): # do nothing in forward (no computation) - assert loss.numel() == 1 ctx.requires_aux_loss = loss.requires_grad ctx.dtype = loss.dtype @@ -77,7 +76,6 @@ class MoEGate(nn.Module): self.n_routed_experts = num_experts self.alpha = aux_loss_alpha - self.seq_aux = False self.gating_dim = embed_dim self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim))) @@ -152,6 +150,8 @@ class MoEBlock(nn.Module): y = y + self.shared_experts(identity) + return y + @torch.no_grad() def moe_infer(self, x, flat_expert_indices, flat_expert_weights): @@ -191,10 +191,10 @@ def test_moe(): start = time.time() moe_gate = MoEGate(512) - moe_gate(torch.rand(1, 71, 512)) + print(moe_gate(torch.rand(1, 71, 512))) - moe_block = MoEBlock(512) - moe_block(torch.rand(1, 77, 512)) + #moe_block = MoEBlock(512) + #moe_block(torch.rand(1, 77, 512)) timing = time.time() - start print(timing) diff --git a/comfy/ldm/hunyuan3d/model/pipeline.py b/comfy/ldm/hunyuan3d/model/pipeline.py new file mode 100644 index 000000000..93b9aa45b --- /dev/null +++ b/comfy/ldm/hunyuan3d/model/pipeline.py @@ -0,0 +1,156 @@ +import torch +import trimesh +import torch.nn as nn +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 + +def export_to_trimesh(mesh_output): + if isinstance(mesh_output, list): + outputs = [] + for mesh in mesh_output: + if mesh is None: + outputs.append(None) + else: + mesh.mesh_f = mesh.mesh_f[:, ::-1] + mesh_output = trimesh.Trimesh(mesh.mesh_v, mesh.mesh_f) + outputs.append(mesh_output) + return outputs + else: + mesh_output.mesh_f = mesh_output.mesh_f[:, ::-1] + mesh_output = trimesh.Trimesh(mesh_output.mesh_v, mesh_output.mesh_f) + return mesh_output + +class Hunyuan3DDiTFlowMatchingPipeline(nn.Module): + def __init__(self, model, vae, conditioner, image_processor, scheduler): + + self.vae = vae + self.model = model + self.conditioner = conditioner + self.image_processor = image_processor + self.scheduler = scheduler + + def compile(self): + self.vae = torch.compile(self.vae) + self.model = torch.compile(self.model) + self.conditioner = torch.compile(self.conditioner) + + def encode_cond(self, image, additional_cond_inputs, do_classifier_free_guidance): + + bsz = image.shape[0] + cond = self.conditioner(image=image, **additional_cond_inputs) + + if do_classifier_free_guidance: + un_cond = self.conditioner.unconditional_embedding(bsz, **additional_cond_inputs) + + # avoid python recursion by using tree_map + _fn = lambda x, y: torch.cat([x, y], dim=0).to(self.dtype) + + cond = tree_map(_fn, cond, un_cond) + + return cond + + def prepare_images(self, images): + + outputs = [] + for image in images: + output = self.image_processor(image) + outputs.append(output) + + return default_collate(outputs) + + def get_guidance_scale_embedding(self, w, embedding_dim=512, dtype=torch.float32): + + assert len(w.shape) == 1 + w = w * 1000.0 + + half_dim = embedding_dim // 2 + emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1) + + emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb) + emb = w.to(dtype)[:, None] * emb[None, :] + + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if embedding_dim % 2 == 1: # zero pad + emb = torch.nn.functional.pad(emb, (0, 1)) + + assert emb.shape == (w.shape[0], embedding_dim) + + return emb + + def prepare_latents(self, batch_size, dtype, device): + + shape = (batch_size, *self.vae.latent_shape) + latents = torch.randn(shape, dtype = dtype, device = device) + + return latents + + @torch.inference_mode() + def __call__( + self, + image: Union[str, List[str], Image.Image, dict, List[dict], torch.Tensor] = None, + guidance_scale: float = 5.0, + bounds = 1.01, + octree_res = 384, + num_chunks = 8000, + **kwargs, + ) -> List[List[trimesh.Trimesh]]: + + callback = kwargs.pop("callback", None) + callback_steps = kwargs.pop("callback_steps", None) + + device = self.device + dtype = self.dtype + do_classifier_free_guidance = guidance_scale >= 0 and not ( + hasattr(self.model, 'guidance_embed') and + self.model.guidance_embed is True + ) + + cond_inputs = self.prepare_image(image) + image = cond_inputs.pop('image') + + cond = self.encode_cond( + image = image, + additional_cond_inputs = cond_inputs, + do_classifier_free_guidance = do_classifier_free_guidance, + ) + + guidance = None + batch_size = image.shape[0] + + latents = self.prepare_latents(batch_size, dtype, device) + + if hasattr(self.model, 'guidance_embed') and \ + self.model.guidance_embed is True: + guidance = torch.tensor([guidance_scale] * batch_size, device=device, dtype=dtype) + + timesteps = self.scheduler.timesteps + + for i, t in enumerate(timesteps): + # expand the latents if we are doing classifier free guidance + if do_classifier_free_guidance: + latent_model_input = torch.cat([latents] * 2) + else: + latent_model_input = latents + + timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype) + timestep = timestep / self.scheduler.num_train_timesteps + noise_pred = self.model(latent_model_input, timestep, cond, guidance=guidance) + + if do_classifier_free_guidance: + noise_pred_cond, noise_pred_uncond = noise_pred.chunk(2) + 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.step(noise_pred, t, latents) + + if callback is not None and i % callback_steps == 0: + step_idx = i // getattr(self.scheduler, "order", 1) + callback(step_idx, t, latents) + + 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) + \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model/scheduler.py b/comfy/ldm/hunyuan3d/model/scheduler.py new file mode 100644 index 000000000..0d1f777c6 --- /dev/null +++ b/comfy/ldm/hunyuan3d/model/scheduler.py @@ -0,0 +1,105 @@ +import torch + +class EulerScheduler(torch.nn.Module): + def __init__(self, num_training_timesteps: int = 1_000, shift: float = 1, + num_inference_timesteps: int = 100, inference: bool = False): + super(EulerScheduler, self).__init__() + + # compute timestep values so we can index into them later + timesteps = torch.linspace(1, num_training_timesteps, num_training_timesteps).to(torch.float32) + + # normalize between 0 and 1 + sigmas = timesteps / num_training_timesteps + + # staticaly shift (fixed image size assumed) + self.sigmas = sigmas * shift / (1 + (shift - 1) * sigmas) + + # get timesteps after shifting + self.timesteps = self.sigmas * num_training_timesteps + + self.num_training_timesteps = num_training_timesteps + self.num_inference_timesteps = num_inference_timesteps + + if inference: + + sigmas = torch.linspace(0, 1, num_inference_timesteps) + sigmas = sigmas * shift / (1 + (shift - 1) * sigmas) + sigmas = sigmas.to(torch.float32) + + timesteps = sigmas * num_training_timesteps + + self.sigmas = torch.cat([sigmas, torch.ones(1, device = sigmas.device)]) + self.timesteps = timesteps.to(device = sigmas.device) + + self.step_index = 0 + + def sigma_to_timestep(self, sigma): + return sigma * self.num_training_timesteps + + def index_for_timestep(self, timestep, schedule_timesteps = None): + + indices = (schedule_timesteps == timestep).nonzero() + + return indices[0].item() + + + def add_noise(self, image: torch.FloatTensor, timestep: float): + + noise = torch.randn_like(image) + + if image.device.type == "mps" and torch.is_floating_point(timestep): + # mps does not support float64 + schedule_timesteps = self.timesteps.to(image.device, dtype=torch.float32) + timestep = timestep.to(image.device, dtype=torch.float32) + else: + schedule_timesteps = self.timesteps.to(image.device) + timestep = timestep.to(image.device) + + # supports a list and a float + if not isinstance(timestep, torch.Tensor) or timestep.ndim == 0: + step_indices = [self.index_for_timestep(timestep, schedule_timesteps)] + else: + step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep] + + sigma = self.sigmas[step_indices].flatten().to(dtype = image.dtype, device = image.device) + + while len(sigma.shape) < len(image.shape): + sigma = sigma.unsqueeze(-1) + + noised_image = (1.0 - sigma) * image + noise * sigma + + 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) + + # 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 + + prev_sample = current_sample + dt * model_output + prev_sample = prev_sample.to(model_output.dtype) + + self.step_index += 1 + + return prev_sample + +def test_scheduler(): + + scheduler = EulerScheduler() + + torch.manual_seed(2025) + image = torch.rand(1, 224, 224, dtype = torch.float32) + latent = torch.rand(1, 224, 224, dtype = torch.float32) + + output = scheduler.add_noise(image, timestep = torch.tensor([scheduler.timesteps[15]])) + output = scheduler.reverse_flow(model_output = image, current_sample = latent) + + print(output) + +if __name__ == "__main__": + test_scheduler() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae/vae.py b/comfy/ldm/hunyuan3d/vae/vae.py index d2bb795d7..2c4f47cfe 100644 --- a/comfy/ldm/hunyuan3d/vae/vae.py +++ b/comfy/ldm/hunyuan3d/vae/vae.py @@ -61,11 +61,15 @@ class VAE(nn.Module): qkv_bias: bool = False, qk_norm: bool = True, drop_path_rate: float = 0.0, - include_pi: bool = False + include_pi: bool = False, + scale_factor: float = 1.0039506158752403 ): super().__init__() + self.latent_shape = (num_latents, embed_dim) + self.scale = scale_factor + self.fourier_embedder = FourierEmbedder(num_freq = num_frequencies, include_pi = include_pi) self.encoder = PointCrossAttention(layers = num_encoder_layers,