From f49e4b541106fd466f19c3de5deada34049203d5 Mon Sep 17 00:00:00 2001 From: Yousef Rafat <81116377+yousef-rafat@users.noreply.github.com> Date: Thu, 10 Jul 2025 19:58:59 +0300 Subject: [PATCH] integerated hunyuan3dv2_1 --- comfy/latent_formats.py | 5 + comfy/ldm/hunyuan3d/model_/dinov2.py | 401 ------------ comfy/ldm/hunyuan3d/vae.py | 587 ++++++++++++++++++ comfy/ldm/hunyuan3d/vae/fps.py | 81 --- comfy/ldm/hunyuan3d/vae/point_attention.py | 452 -------------- comfy/ldm/hunyuan3d/vae/postprocess.py | 106 ---- comfy/ldm/hunyuan3d/vae/preprocess.py | 165 ----- comfy/ldm/hunyuan3d/vae/transformer.py | 156 ----- comfy/ldm/hunyuan3d/vae/vae.py | 185 ------ .../model_ => hunyuan3dv2_1}/conditioner.py | 31 +- .../model_ => hunyuan3dv2_1}/hunyuandit.py | 5 +- .../image_processor.py | 0 .../model_ => hunyuan3dv2_1}/moe.py | 0 .../model_ => hunyuan3dv2_1}/pipeline.py | 0 .../model_ => hunyuan3dv2_1}/scheduler.py | 0 .../model_ => hunyuan3dv2_1}/vae.py | 12 +- comfy/model_base.py | 42 ++ comfy/model_detection.py | 14 + comfy/sd.py | 24 + comfy/supported_models.py | 13 +- comfy_extras/nodes_hunyuan3d.py | 81 ++- 21 files changed, 780 insertions(+), 1580 deletions(-) delete mode 100644 comfy/ldm/hunyuan3d/model_/dinov2.py create mode 100644 comfy/ldm/hunyuan3d/vae.py delete mode 100644 comfy/ldm/hunyuan3d/vae/fps.py delete mode 100644 comfy/ldm/hunyuan3d/vae/point_attention.py delete mode 100644 comfy/ldm/hunyuan3d/vae/postprocess.py delete mode 100644 comfy/ldm/hunyuan3d/vae/preprocess.py delete mode 100644 comfy/ldm/hunyuan3d/vae/transformer.py delete mode 100644 comfy/ldm/hunyuan3d/vae/vae.py rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/conditioner.py (83%) rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/hunyuandit.py (99%) rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/image_processor.py (100%) rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/moe.py (100%) rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/pipeline.py (100%) rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/scheduler.py (100%) rename comfy/ldm/{hunyuan3d/model_ => hunyuan3dv2_1}/vae.py (99%) diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index 82d9f9bb8..32ae1fb72 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -462,6 +462,11 @@ class Hunyuan3Dv2(LatentFormat): latent_dimensions = 1 scale_factor = 0.9990943042622529 +class Hunyuan3Dv2_1(LatentFormat): + scale_factor = 1.0039506158752403 + latent_channels = 64 + latent_dimensions = 1 + class Hunyuan3Dv2mini(LatentFormat): latent_channels = 64 latent_dimensions = 1 diff --git a/comfy/ldm/hunyuan3d/model_/dinov2.py b/comfy/ldm/hunyuan3d/model_/dinov2.py deleted file mode 100644 index 79f813a1a..000000000 --- a/comfy/ldm/hunyuan3d/model_/dinov2.py +++ /dev/null @@ -1,401 +0,0 @@ -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 - device: str = "cuda" - dtype = torch.float16 - -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.device = config.device - self.dtype = config.dtype - - 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 - sequence_output = self.layernorm(sequence_output) - - return sequence_output \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae.py b/comfy/ldm/hunyuan3d/vae.py new file mode 100644 index 000000000..6595d96d6 --- /dev/null +++ b/comfy/ldm/hunyuan3d/vae.py @@ -0,0 +1,587 @@ +# Original: https://github.com/Tencent/Hunyuan3D-2/blob/main/hy3dgen/shapegen/models/autoencoders/model.py +# Since the header on their VAE source file was a bit confusing we asked for permission to use this code from tencent under the GPL license used in ComfyUI. + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +from typing import Union, Tuple, List, Callable, Optional + +import numpy as np +from einops import repeat, rearrange +from tqdm import tqdm +import logging + +import comfy.ops +ops = comfy.ops.disable_weight_init + +def generate_dense_grid_points( + bbox_min: np.ndarray, + bbox_max: np.ndarray, + octree_resolution: int, + indexing: str = "ij", +): + length = bbox_max - bbox_min + num_cells = octree_resolution + + x = np.linspace(bbox_min[0], bbox_max[0], int(num_cells) + 1, dtype=np.float32) + y = np.linspace(bbox_min[1], bbox_max[1], int(num_cells) + 1, dtype=np.float32) + z = np.linspace(bbox_min[2], bbox_max[2], int(num_cells) + 1, dtype=np.float32) + [xs, ys, zs] = np.meshgrid(x, y, z, indexing=indexing) + xyz = np.stack((xs, ys, zs), axis=-1) + grid_size = [int(num_cells) + 1, int(num_cells) + 1, int(num_cells) + 1] + + return xyz, grid_size, length + + +class VanillaVolumeDecoder: + @torch.no_grad() + def __call__( + self, + latents: torch.FloatTensor, + geo_decoder: Callable, + bounds: Union[Tuple[float], List[float], float] = 1.01, + num_chunks: int = 10000, + octree_resolution: int = None, + enable_pbar: bool = True, + **kwargs, + ): + device = latents.device + dtype = latents.dtype + batch_size = latents.shape[0] + + # 1. generate query points + if isinstance(bounds, float): + bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds] + + bbox_min, bbox_max = np.array(bounds[0:3]), np.array(bounds[3:6]) + xyz_samples, grid_size, length = generate_dense_grid_points( + bbox_min=bbox_min, + bbox_max=bbox_max, + octree_resolution=octree_resolution, + indexing="ij" + ) + xyz_samples = torch.from_numpy(xyz_samples).to(device, dtype=dtype).contiguous().reshape(-1, 3) + + # 2. latents to 3d volume + batch_logits = [] + for start in tqdm(range(0, xyz_samples.shape[0], num_chunks), desc="Volume Decoding", + disable=not enable_pbar): + chunk_queries = xyz_samples[start: start + num_chunks, :] + chunk_queries = repeat(chunk_queries, "p c -> b p c", b=batch_size) + logits = geo_decoder(queries=chunk_queries, latents=latents) + batch_logits.append(logits) + + grid_logits = torch.cat(batch_logits, dim=1) + grid_logits = grid_logits.view((batch_size, *grid_size)).float() + + return grid_logits + + +class FourierEmbedder(nn.Module): + """The sin/cosine positional embedding. Given an input tensor `x` of shape [n_batch, ..., c_dim], it converts + each feature dimension of `x[..., i]` into: + [ + sin(x[..., i]), + sin(f_1*x[..., i]), + sin(f_2*x[..., i]), + ... + sin(f_N * x[..., i]), + cos(x[..., i]), + cos(f_1*x[..., i]), + cos(f_2*x[..., i]), + ... + cos(f_N * x[..., i]), + x[..., i] # only present if include_input is True. + ], here f_i is the frequency. + + Denote the space is [0 / num_freqs, 1 / num_freqs, 2 / num_freqs, 3 / num_freqs, ..., (num_freqs - 1) / num_freqs]. + If logspace is True, then the frequency f_i is [2^(0 / num_freqs), ..., 2^(i / num_freqs), ...]; + Otherwise, the frequencies are linearly spaced between [1.0, 2^(num_freqs - 1)]. + + Args: + num_freqs (int): the number of frequencies, default is 6; + logspace (bool): If logspace is True, then the frequency f_i is [..., 2^(i / num_freqs), ...], + otherwise, the frequencies are linearly spaced between [1.0, 2^(num_freqs - 1)]; + input_dim (int): the input dimension, default is 3; + include_input (bool): include the input tensor or not, default is True. + + Attributes: + frequencies (torch.Tensor): If logspace is True, then the frequency f_i is [..., 2^(i / num_freqs), ...], + otherwise, the frequencies are linearly spaced between [1.0, 2^(num_freqs - 1); + + out_dim (int): the embedding size, if include_input is True, it is input_dim * (num_freqs * 2 + 1), + otherwise, it is input_dim * num_freqs * 2. + + """ + + def __init__(self, + num_freqs: int = 6, + logspace: bool = True, + input_dim: int = 3, + include_input: bool = True, + include_pi: bool = True) -> None: + + """The initialization""" + + super().__init__() + + if logspace: + frequencies = 2.0 ** torch.arange( + num_freqs, + dtype=torch.float32 + ) + else: + frequencies = torch.linspace( + 1.0, + 2.0 ** (num_freqs - 1), + num_freqs, + dtype=torch.float32 + ) + + if include_pi: + frequencies *= torch.pi + + self.register_buffer("frequencies", frequencies, persistent=False) + self.include_input = include_input + self.num_freqs = num_freqs + + self.out_dim = self.get_dims(input_dim) + + def get_dims(self, input_dim): + temp = 1 if self.include_input or self.num_freqs == 0 else 0 + out_dim = input_dim * (self.num_freqs * 2 + temp) + + return out_dim + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """ Forward process. + + Args: + x: tensor of shape [..., dim] + + Returns: + embedding: an embedding of `x` of shape [..., dim * (num_freqs * 2 + temp)] + where temp is 1 if include_input is True and 0 otherwise. + """ + + if self.num_freqs > 0: + embed = (x[..., None].contiguous() * self.frequencies.to(device=x.device, dtype=x.dtype)).view(*x.shape[:-1], -1) + if self.include_input: + return torch.cat((x, embed.sin(), embed.cos()), dim=-1) + else: + return torch.cat((embed.sin(), embed.cos()), dim=-1) + else: + return x + + +class CrossAttentionProcessor: + def __call__(self, attn, q, k, v): + out = F.scaled_dot_product_attention(q, k, v) + return out + + +class DropPath(nn.Module): + """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks). + """ + + def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + self.scale_by_keep = scale_by_keep + + def forward(self, x): + """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks). + + This is the same as the DropConnect impl I created for EfficientNet, etc networks, however, + the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper... + See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for + changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use + 'survival rate' as the argument. + + """ + if self.drop_prob == 0. or not self.training: + return x + keep_prob = 1 - self.drop_prob + shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets + random_tensor = x.new_empty(shape).bernoulli_(keep_prob) + if keep_prob > 0.0 and self.scale_by_keep: + random_tensor.div_(keep_prob) + return x * random_tensor + + def extra_repr(self): + return f'drop_prob={round(self.drop_prob, 3):0.3f}' + + +class MLP(nn.Module): + def __init__( + self, *, + width: int, + expand_ratio: int = 4, + output_width: int = None, + drop_path_rate: float = 0.0 + ): + super().__init__() + self.width = width + self.c_fc = ops.Linear(width, width * expand_ratio) + self.c_proj = ops.Linear(width * expand_ratio, output_width if output_width is not None else width) + self.gelu = nn.GELU() + self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity() + + def forward(self, x): + return self.drop_path(self.c_proj(self.gelu(self.c_fc(x)))) + + +class QKVMultiheadCrossAttention(nn.Module): + def __init__( + self, + *, + heads: int, + width=None, + qk_norm=False, + norm_layer=ops.LayerNorm + ): + super().__init__() + self.heads = heads + self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() + self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() + + self.attn_processor = CrossAttentionProcessor() + + def forward(self, q, kv): + _, n_ctx, _ = q.shape + bs, n_data, width = kv.shape + attn_ch = width // self.heads // 2 + q = q.view(bs, n_ctx, self.heads, -1) + kv = kv.view(bs, n_data, self.heads, -1) + k, v = torch.split(kv, attn_ch, dim=-1) + + q = self.q_norm(q) + k = self.k_norm(k) + q, k, v = map(lambda t: rearrange(t, 'b n h d -> b h n d', h=self.heads), (q, k, v)) + out = self.attn_processor(self, q, k, v) + out = out.transpose(1, 2).reshape(bs, n_ctx, -1) + return out + + +class MultiheadCrossAttention(nn.Module): + def __init__( + self, + *, + width: int, + heads: int, + qkv_bias: bool = True, + data_width: Optional[int] = None, + norm_layer=ops.LayerNorm, + qk_norm: bool = False, + kv_cache: bool = False, + ): + super().__init__() + self.width = width + self.heads = heads + self.data_width = width if data_width is None else data_width + self.c_q = ops.Linear(width, width, bias=qkv_bias) + self.c_kv = ops.Linear(self.data_width, width * 2, bias=qkv_bias) + self.c_proj = ops.Linear(width, width) + self.attention = QKVMultiheadCrossAttention( + heads=heads, + width=width, + norm_layer=norm_layer, + qk_norm=qk_norm + ) + self.kv_cache = kv_cache + self.data = None + + def forward(self, x, data): + x = self.c_q(x) + if self.kv_cache: + if self.data is None: + self.data = self.c_kv(data) + logging.info('Save kv cache,this should be called only once for one mesh') + data = self.data + else: + data = self.c_kv(data) + x = self.attention(x, data) + x = self.c_proj(x) + return x + + +class ResidualCrossAttentionBlock(nn.Module): + def __init__( + self, + *, + width: int, + heads: int, + mlp_expand_ratio: int = 4, + data_width: Optional[int] = None, + qkv_bias: bool = True, + norm_layer=ops.LayerNorm, + qk_norm: bool = False + ): + super().__init__() + + if data_width is None: + data_width = width + + self.attn = MultiheadCrossAttention( + width=width, + heads=heads, + data_width=data_width, + qkv_bias=qkv_bias, + norm_layer=norm_layer, + qk_norm=qk_norm + ) + self.ln_1 = norm_layer(width, elementwise_affine=True, eps=1e-6) + self.ln_2 = norm_layer(data_width, elementwise_affine=True, eps=1e-6) + self.ln_3 = norm_layer(width, elementwise_affine=True, eps=1e-6) + self.mlp = MLP(width=width, expand_ratio=mlp_expand_ratio) + + def forward(self, x: torch.Tensor, data: torch.Tensor): + x = x + self.attn(self.ln_1(x), self.ln_2(data)) + x = x + self.mlp(self.ln_3(x)) + return x + + +class QKVMultiheadAttention(nn.Module): + def __init__( + self, + *, + heads: int, + width=None, + qk_norm=False, + norm_layer=ops.LayerNorm + ): + super().__init__() + self.heads = heads + self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() + self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() + + def forward(self, qkv): + bs, n_ctx, width = qkv.shape + attn_ch = width // self.heads // 3 + qkv = qkv.view(bs, n_ctx, self.heads, -1) + q, k, v = torch.split(qkv, attn_ch, dim=-1) + + q = self.q_norm(q) + k = self.k_norm(k) + + q, k, v = map(lambda t: rearrange(t, 'b n h d -> b h n d', h=self.heads), (q, k, v)) + out = F.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(bs, n_ctx, -1) + return out + + +class MultiheadAttention(nn.Module): + def __init__( + self, + *, + width: int, + heads: int, + qkv_bias: bool, + norm_layer=ops.LayerNorm, + qk_norm: bool = False, + drop_path_rate: float = 0.0 + ): + super().__init__() + self.width = width + self.heads = heads + self.c_qkv = ops.Linear(width, width * 3, bias=qkv_bias) + self.c_proj = ops.Linear(width, width) + self.attention = QKVMultiheadAttention( + heads=heads, + width=width, + norm_layer=norm_layer, + qk_norm=qk_norm + ) + self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity() + + def forward(self, x): + x = self.c_qkv(x) + x = self.attention(x) + x = self.drop_path(self.c_proj(x)) + return x + + +class ResidualAttentionBlock(nn.Module): + def __init__( + self, + *, + width: int, + heads: int, + qkv_bias: bool = True, + norm_layer=ops.LayerNorm, + qk_norm: bool = False, + drop_path_rate: float = 0.0, + ): + super().__init__() + self.attn = MultiheadAttention( + width=width, + heads=heads, + qkv_bias=qkv_bias, + norm_layer=norm_layer, + qk_norm=qk_norm, + drop_path_rate=drop_path_rate + ) + self.ln_1 = norm_layer(width, elementwise_affine=True, eps=1e-6) + self.mlp = MLP(width=width, drop_path_rate=drop_path_rate) + self.ln_2 = norm_layer(width, elementwise_affine=True, eps=1e-6) + + def forward(self, x: torch.Tensor): + x = x + self.attn(self.ln_1(x)) + x = x + self.mlp(self.ln_2(x)) + return x + + +class Transformer(nn.Module): + def __init__( + self, + *, + width: int, + layers: int, + heads: int, + qkv_bias: bool = True, + norm_layer=ops.LayerNorm, + qk_norm: bool = False, + drop_path_rate: float = 0.0 + ): + super().__init__() + self.width = width + self.layers = layers + self.resblocks = nn.ModuleList( + [ + ResidualAttentionBlock( + width=width, + heads=heads, + qkv_bias=qkv_bias, + norm_layer=norm_layer, + qk_norm=qk_norm, + drop_path_rate=drop_path_rate + ) + for _ in range(layers) + ] + ) + + def forward(self, x: torch.Tensor): + for block in self.resblocks: + x = block(x) + return x + + +class CrossAttentionDecoder(nn.Module): + + def __init__( + self, + *, + out_channels: int, + fourier_embedder: FourierEmbedder, + width: int, + heads: int, + mlp_expand_ratio: int = 4, + downsample_ratio: int = 1, + enable_ln_post: bool = True, + qkv_bias: bool = True, + qk_norm: bool = False, + label_type: str = "binary" + ): + super().__init__() + + self.enable_ln_post = enable_ln_post + self.fourier_embedder = fourier_embedder + self.downsample_ratio = downsample_ratio + self.query_proj = ops.Linear(self.fourier_embedder.out_dim, width) + if self.downsample_ratio != 1: + self.latents_proj = ops.Linear(width * downsample_ratio, width) + if self.enable_ln_post == False: + qk_norm = False + self.cross_attn_decoder = ResidualCrossAttentionBlock( + width=width, + mlp_expand_ratio=mlp_expand_ratio, + heads=heads, + qkv_bias=qkv_bias, + qk_norm=qk_norm + ) + + if self.enable_ln_post: + self.ln_post = ops.LayerNorm(width) + self.output_proj = ops.Linear(width, out_channels) + self.label_type = label_type + self.count = 0 + + def forward(self, queries=None, query_embeddings=None, latents=None): + if query_embeddings is None: + query_embeddings = self.query_proj(self.fourier_embedder(queries).to(latents.dtype)) + self.count += query_embeddings.shape[1] + if self.downsample_ratio != 1: + latents = self.latents_proj(latents) + x = self.cross_attn_decoder(query_embeddings, latents) + if self.enable_ln_post: + x = self.ln_post(x) + occ = self.output_proj(x) + return occ + + +class ShapeVAE(nn.Module): + def __init__( + self, + *, + embed_dim: int, + width: int, + heads: int, + num_decoder_layers: int, + geo_decoder_downsample_ratio: int = 1, + geo_decoder_mlp_expand_ratio: int = 4, + geo_decoder_ln_post: bool = True, + num_freqs: int = 8, + include_pi: bool = True, + qkv_bias: bool = True, + qk_norm: bool = False, + label_type: str = "binary", + drop_path_rate: float = 0.0, + scale_factor: float = 1.0, + ): + super().__init__() + self.geo_decoder_ln_post = geo_decoder_ln_post + + self.fourier_embedder = FourierEmbedder(num_freqs=num_freqs, include_pi=include_pi) + + self.post_kl = ops.Linear(embed_dim, width) + + self.transformer = Transformer( + width=width, + layers=num_decoder_layers, + heads=heads, + qkv_bias=qkv_bias, + qk_norm=qk_norm, + drop_path_rate=drop_path_rate + ) + + self.geo_decoder = CrossAttentionDecoder( + fourier_embedder=self.fourier_embedder, + out_channels=1, + mlp_expand_ratio=geo_decoder_mlp_expand_ratio, + downsample_ratio=geo_decoder_downsample_ratio, + enable_ln_post=self.geo_decoder_ln_post, + width=width // geo_decoder_downsample_ratio, + heads=heads // geo_decoder_downsample_ratio, + qkv_bias=qkv_bias, + qk_norm=qk_norm, + label_type=label_type, + ) + + self.volume_decoder = VanillaVolumeDecoder() + self.scale_factor = scale_factor + + def decode(self, latents, **kwargs): + latents = self.post_kl(latents.movedim(-2, -1)) + latents = self.transformer(latents) + + bounds = kwargs.get("bounds", 1.01) + num_chunks = kwargs.get("num_chunks", 8000) + octree_resolution = kwargs.get("octree_resolution", 256) + enable_pbar = kwargs.get("enable_pbar", True) + + grid_logits = self.volume_decoder(latents, self.geo_decoder, bounds=bounds, num_chunks=num_chunks, octree_resolution=octree_resolution, enable_pbar=enable_pbar) + return grid_logits.movedim(-2, -1) + + def encode(self, x): + return None \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae/fps.py b/comfy/ldm/hunyuan3d/vae/fps.py deleted file mode 100644 index 53b0793a7..000000000 --- a/comfy/ldm/hunyuan3d/vae/fps.py +++ /dev/null @@ -1,81 +0,0 @@ -# replaced torch.ops.torch_cluster.fps with a manual implementation -# to avoid having torch_cluster downloaded as dependency -# also the dependency takes a long time to install - -import torch -from torch import Tensor -import math - -def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True): - - # manually create the pointer vector - assert src.size(0) == batch.numel() - - batch_size = int(batch.max()) + 1 - deg = src.new_zeros(batch_size, dtype = torch.long) - - deg.scatter_add_(0, batch, torch.ones_like(batch)) - - ptr_vec = deg.new_zeros(batch_size + 1) - torch.cumsum(deg, 0, out=ptr_vec[1:]) - - #return fps_sampling(src, ptr_vec, ratio) - sampled_indicies = [] - - for b in range(batch_size): - # start and the end of each batch - start, end = ptr_vec[b].item(), ptr_vec[b + 1].item() - # points from the point cloud - points = src[start:end] - - num_points = points.size(0) - num_samples = max(1, math.ceil(num_points * sampling_ratio)) - - selected = torch.zeros(num_samples, device = src.device, dtype = torch.long) - distances = torch.full((num_points,), float("inf"), device = src.device) - - # select a random start point - if start_random: - farthest = torch.randint(0, num_points, (1,), device = src.device) - else: farthest = torch.tensor([0], device = src.device, dtype = torch.long) - - for i in range(num_samples): - selected[i] = farthest - centroid = points[farthest].squeeze(0) - dist = torch.norm(points - centroid, dim = 1) # compute euclidean distance - distances = torch.minimum(distances, dist) - farthest = torch.argmax(distances) - - sampled_indicies.append(torch.arange(start, end)[selected]) - - return torch.cat(sampled_indicies, dim = 0) - - -def test_fps(): - - torch.manual_seed(2025) - - # 2 batches with different numbers of points - points = torch.tensor([ - [0.0, 0.0, 0.0], # batch 0 - [1.0, 0.0, 0.0], # batch 0 - [2.0, 0.0, 0.0], # batch 0 - [0.0, 1.0, 0.0], # batch 1 - [0.0, 2.0, 0.0], # batch 1 - [0.0, 3.0, 0.0] # batch 1 - ], dtype=torch.float) - - batch = torch.tensor([0, 0, 0, 1, 1, 1]) # batch IDs - - ratio = 0.5 # sample 50% of points per batch - # jit compilation for speedups - #optimized_fps = torch.compile(fps) - - outputs = fps(points, batch, ratio, start_random = True) - #outputs2 = torch.ops.torch_cluster.fps(points, batch, ratio, True) # shouldn't work - - print(outputs) - -if __name__ == "__main__": - test_fps() - \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae/point_attention.py b/comfy/ldm/hunyuan3d/vae/point_attention.py deleted file mode 100644 index 523a61308..000000000 --- a/comfy/ldm/hunyuan3d/vae/point_attention.py +++ /dev/null @@ -1,452 +0,0 @@ -from transformer import MLP -from transformer import Transformer -import torch.nn.functional as F -import torch.nn as nn -from fps import fps -import torch - -class QKVMultiheadCrossAttention(nn.Module): - def __init__( - self, - heads: int, - n_data = None, - width=None, - qk_norm=False, - norm_layer=nn.LayerNorm - ): - super().__init__() - self.heads = heads - self.n_data = n_data - self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() - self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() - - def forward(self, q, kv): - - _, n_ctx, _ = q.shape - bs, n_data, width = kv.shape - - attn_ch = width // self.heads // 2 - q = q.view(bs, n_ctx, self.heads, -1) - - kv = kv.view(bs, n_data, self.heads, -1) - k, v = torch.split(kv, attn_ch, dim=-1) - - q = self.q_norm(q) - k = self.k_norm(k) - - q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)] - out = F.scaled_dot_product_attention(q, k, v) - - out = out.transpose(1, 2).reshape(bs, n_ctx, -1) - - return out - - -class MultiheadCrossAttention(nn.Module): - def __init__( - self, - width: int, - heads: int, - qkv_bias: bool = False, - n_data = None, - norm_layer = nn.LayerNorm, - qk_norm: bool = False, - kv_cache: bool = False, - ): - super().__init__() - - self.c_q = nn.Linear(width, width, bias=qkv_bias) - self.c_kv = nn.Linear(width, width * 2, bias=qkv_bias) - self.c_proj = nn.Linear(width, width) - - self.attention = QKVMultiheadCrossAttention( - heads = heads, - n_data = n_data, - width = width, - norm_layer = norm_layer, - qk_norm = qk_norm - ) - - self.kv_cache = kv_cache - self.data = None - - def forward(self, x, data): - x = self.c_q(x) - - if self.kv_cache: - if self.data is None: - self.data = self.c_kv(data) - - data = self.data - else: - data = self.c_kv(data) - - x = self.attention(x, data) - x = self.c_proj(x) - - return x - - -class ResidualCrossAttentionBlock(nn.Module): - def __init__( - self, - width: int, - heads: int, - n_data: int = None, - mlp_expand_ratio: int = 4, - qkv_bias: bool = False, - norm_layer=nn.LayerNorm, - qk_norm: bool = False - ): - super().__init__() - - self.attn = MultiheadCrossAttention( - n_data=n_data, - width = width, - heads=heads, - qkv_bias=qkv_bias, - norm_layer=norm_layer, - qk_norm=qk_norm - ) - - self.ln_1 = norm_layer(width, elementwise_affine = True, eps = 1e-6) - self.ln_2 = norm_layer(width, elementwise_affine = True, eps = 1e-6) - self.ln_3 = norm_layer(width, elementwise_affine = True, eps = 1e-6) - - self.mlp = MLP(width=width, ratio = mlp_expand_ratio) - - def forward(self, x: torch.Tensor, data: torch.Tensor): - x = x + self.attn(self.ln_1(x), self.ln_2(data)) - x = x + self.mlp(self.ln_3(x)) - return x - -class CrossAttentionDecoder(nn.Module): - def __init__( - self, - num_latents: int, - out_channels: int, - fourier_embedder, - width: int, - heads: int, - mlp_expand_ratio: int = 4, - downsample_ratio: int = 1, - enable_ln_post: bool = True, - qkv_bias: bool = False, - qk_norm: bool = False): - - super().__init__() - - self.enable_ln_post = enable_ln_post - self.fourier_embedder = fourier_embedder - self.downsample_ratio = downsample_ratio - - self.query_proj = nn.Linear(self.fourier_embedder.out_dim, width) - - if self.downsample_ratio != 1: - self.latents_proj = nn.Linear(width * downsample_ratio, width) - - if self.enable_ln_post == False: - qk_norm = False - - self.cross_attn_decoder = ResidualCrossAttentionBlock( - n_data=num_latents, - width=width, - mlp_expand_ratio=mlp_expand_ratio, - heads=heads, - qkv_bias=qkv_bias, - qk_norm=qk_norm - ) - - if self.enable_ln_post: - self.ln_post = nn.LayerNorm(width) - - self.output_proj = nn.Linear(width, out_channels) - self.count = 0 - - def forward(self, queries = None, query_embeddings = None, latents = None): - - if query_embeddings is None: - query_embeddings = self.query_proj(self.fourier_embedder(queries).to(latents.dtype)) - - self.count += query_embeddings.shape[1] - - if self.downsample_ratio != 1: - latents = self.latents_proj(latents) - - x = self.cross_attn_decoder(query_embeddings, latents) - - if self.enable_ln_post: - x = self.ln_post(x) - - out = self.output_proj(x) - - return out - - -class PointCrossAttention(nn.Module): - def __init__(self, - num_latents: int, - downsample_ratio: float, - pc_size: int, - pc_sharpedge_size: int, - point_feats: int, - width: int, - heads: int, - layers: int, - fourier_embedder, - normal_pe: bool = False, - qkv_bias: bool = False, - use_ln_post: bool = True, - qk_norm: bool = True): - - super().__init__() - - self.fourier_embedder = fourier_embedder - - self.pc_size = pc_size - self.normal_pe = normal_pe - self.downsample_ratio = downsample_ratio - self.pc_sharpedge_size = pc_sharpedge_size - self.num_latents = num_latents - self.point_feats = point_feats - - self.input_proj = nn.Linear(self.fourier_embedder.out_dim + point_feats, width) - - self.cross_attn = ResidualCrossAttentionBlock( - width = width, - heads = heads, - qkv_bias = qkv_bias, - qk_norm = qk_norm - ) - - self.self_attn = None - if layers > 0: - self.self_attn = Transformer( - n_ctx = num_latents, - width = width, - heads = heads, - qkv_bias = qkv_bias, - qk_norm = qk_norm, - depth = layers - ) - - if use_ln_post: - self.ln_post = nn.LayerNorm(width) - else: - self.ln_post = None - - def sample_points_and_latents(self, point_cloud: torch.Tensor, features: torch.Tensor): - - """ - Subsample points randomly from the point cloud (input_pc) - Further sample the subsampled points to get query_pc - take the fourier embeddings for both input and query pc - - Mental Note: FPS-sampled points (query_pc) act as latent tokens that attend to and learn from the broader context in input_pc. - Goal: get a smaller represenation (query_pc) to represent the entire scence structure by learning from a broader subset (input_pc). - More computationally efficient. - - Features are additional information for each point in the cloud - """ - - B, _, D = point_cloud.shape - - num_latents = int(self.num_latents) - - num_random_query = self.pc_size / (self.pc_size + self.pc_sharpedge_size) * num_latents - num_sharpedge_query = num_latents - num_random_query - - # Split random and sharpedge surface points - random_pc, sharpedge_pc = torch.split(point_cloud, [self.pc_size, self.pc_sharpedge_size], dim=1) - - # assert statements - assert random_pc.shape[1] <= self.pc_size, "Random surface points size must be less than or equal to pc_size" - assert sharpedge_pc.shape[1] <= self.pc_sharpedge_size, "Sharpedge surface points size must be less than or equal to pc_sharpedge_size" - - input_random_pc_size = int(num_random_query * self.downsample_ratio) - random_query_pc, random_input_pc, random_idx_pc, random_idx_query = \ - self.subsample(pc = random_pc, num_query = num_random_query, input_pc_size = input_random_pc_size) - - input_sharpedge_pc_size = int(num_sharpedge_query * self.downsample_ratio) - - if input_sharpedge_pc_size == 0: - sharpedge_input_pc = torch.zeros(B, 0, D, dtype = random_input_pc.dtype).to(point_cloud.device) - sharpedge_query_pc = torch.zeros(B, 0, D, dtype= random_query_pc.dtype).to(point_cloud.device) - - else: sharpedge_query_pc, sharpedge_input_pc, sharpedge_idx_pc, sharpedge_idx_query = \ - self.subsample(pc = sharpedge_pc, num_query = num_sharpedge_query, input_pc_size = input_sharpedge_pc_size) - - # concat the random and sharpedges - query_pc = torch.cat([random_query_pc, sharpedge_query_pc], dim = 1) - input_pc = torch.cat([random_input_pc, sharpedge_input_pc], dim = 1) - - query = self.fourier_embedder(query_pc) - data = self.fourier_embedder(input_pc) - - if self.point_feats > 0: - random_surface_features, sharpedge_surface_features = torch.split(features, [self.pc_size, self.pc_sharpedge_size], dim = 1) - - input_random_surface_features, query_random_features = \ - self.handle_features(features = random_surface_features, idx_pc = random_idx_pc, batch_size = B, - input_pc_size = input_random_pc_size, idx_query = random_idx_query) - - if input_sharpedge_pc_size == 0: - input_sharpedge_surface_features = torch.zeros(B, 0, self.point_feats, - dtype = input_random_surface_features.dtype, device = point_cloud.device) - - query_sharpedge_features = torch.zeros(B, 0, self.point_feats, - dtype = query_random_features.dtype, device = point_cloud.device) - else: - - input_sharpedge_surface_features, query_sharpedge_features = \ - self.handle_features(idx_pc = sharpedge_idx_pc, features = sharpedge_surface_features, - batch_size = B, idx_query = sharpedge_idx_query, input_pc_size = input_sharpedge_pc_size) - - query_features = torch.cat([query_random_features, query_sharpedge_features], dim = 1) - input_features = torch.cat([input_random_surface_features, input_sharpedge_surface_features], dim = 1) - - if self.normal_pe: - # apply the fourier embeddings on the first 3 dims (xyz) - input_features_pe = self.fourier_embedder(input_features[..., :3]) - query_features_pe = self.fourier_embedder(query_features[..., :3]) - # replace the first 3 dims with the new PE ones - input_features = torch.cat([input_features_pe, input_features[..., :3]], dim = -1) - query_features = torch.cat([query_features_pe, query_features[..., :3]], dim = -1) - - # concat at the channels dim - query = torch.cat([query, query_features], dim = -1) - data = torch.cat([data, input_features], dim = -1) - - # don't return pc_info to avoid unnecessary memory usuage - return query.view(B, -1, query.shape[-1]), data.view(B, -1, data.shape[-1]) - - def forward(self, point_cloud: torch.Tensor, features: torch.Tensor): - - query, data = self.sample_points_and_latents(point_cloud = point_cloud, features = features) - - # apply projections - query = self.input_proj(query) - data = self.input_proj(data) - - # apply cross attention between query and data - latents = self.cross_attn(query, data) - - if self.self_attn is not None: - latents = self.self_attn(latents) - - if self.ln_post is not None: - latents = self.ln_post(latents) - - return latents - - - def subsample(self, pc, num_query, input_pc_size: int): - - """ - num_query: number of points to keep after FPS - input_pc_size: number of points to select before FPS - """ - - B, _, D = pc.shape - query_ratio = num_query / input_pc_size - - # random subsampling of points inside the point cloud - idx_pc = torch.randperm(pc.shape[1], device = pc.device)[:input_pc_size] - input_pc = pc[:, idx_pc, :] - - # flatten to allow applying fps across the whole batch - flattent_input_pc = input_pc.view(B * input_pc_size, D) - - # construct a batch_down tensor to tell fps - # which points belong to which batch - N_down = int(flattent_input_pc.shape[0] / B) - batch_down = torch.arange(B).to(pc.device) - batch_down = torch.repeat_interleave(batch_down, N_down) - - idx_query = fps(flattent_input_pc, batch_down, sampling_ratio = query_ratio) - query_pc = flattent_input_pc[idx_query].view(B, -1, D) - - return query_pc, input_pc, idx_pc, idx_query - - def handle_features(self, features, idx_pc, input_pc_size, batch_size: int, idx_query): - - B = batch_size - - input_surface_features = features[:, idx_pc, :] - flattent_input_features = input_surface_features.view(B * input_pc_size, -1) - query_features = flattent_input_features[idx_query].view(B, -1, - flattent_input_features.shape[-1]) - - return input_surface_features, query_features - - def forward(self, pc, feats): - """ - - Args: - pc (torch.FloatTensor): [B, N, 3] - feats (torch.FloatTensor or None): [B, N, C] - - Returns: - - """ - - query, data = self.sample_points_and_latents(pc, feats) - - query = self.input_proj(query) - query = query - data = self.input_proj(data) - data = data - - latents = self.cross_attn(query, data) - if self.self_attn is not None: - latents = self.self_attn(latents) - - if self.ln_post is not None: - latents = self.ln_post(latents) - - return latents -def test_point_cross_attention(): - - from vae import FourierEmbedder - - torch.manual_seed(2025) - B = 2 # batch size - D = 3 # point dimension (x, y, z) - F = 16 # feature dimension - - pc_random = 96 # number of random surface points - pc_sharpedge = 32 # number of sharpedge points - total_points = pc_random + pc_sharpedge # = 128 - - L = 32 # num_latents (final tokens) - downsample_ratio = 2.0 # modest oversampling - width = 128 - heads = 4 - layers = 2 - - - point_cloud = torch.randn(B, total_points, D) - features = torch.randn(B, total_points, F) - - embedder = FourierEmbedder() - - model = PointCrossAttention( - num_latents=L, - downsample_ratio=downsample_ratio, - pc_size=pc_random, - pc_sharpedge_size=pc_sharpedge, - point_feats=F, - width=width, - heads=heads, - layers=layers, - fourier_embedder=embedder, - use_ln_post=True, - qkv_bias=True, - qk_norm=False - ) - - - output = model(point_cloud, features) - print(output[0]) -if __name__ == '__main__': - test_point_cross_attention() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae/postprocess.py b/comfy/ldm/hunyuan3d/vae/postprocess.py deleted file mode 100644 index c9f1909ea..000000000 --- a/comfy/ldm/hunyuan3d/vae/postprocess.py +++ /dev/null @@ -1,106 +0,0 @@ -import torch -from skimage import measure -from dataclasses import dataclass -import numpy as np - -@dataclass -class Latent2MeshOutput(): - # mesh for vertices and faces - mesh_v: None - mesh_f: None - -class SufraceExtractor(): - def compute_box_stat(self, bounds, octree_resolution: int): - - # if float, turn it into a cube - if isinstance(bounds, float): - bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds] - - bbox_min, bbox_max = np.array(bounds[0:3]), np.array(bounds[3:6]) - bbox_size = bbox_max - bbox_min - grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1] - return grid_size, bbox_min, bbox_size - - def run(self, grid_logit, *, bounds, octree_res, **kwargs): - # grid_logit from volume decoder - # use marching cube algo to turn an sdf to a mesh - vertices, faces, _, _ = measure.marching_cubes(grid_logit.cpu().numpy(), - 0.0, - method = "lewiner") - - grid_size, bbox_min, bbox_size = self.compute_box_stat(bounds = bounds, octree_resolution = octree_res) - vertices = vertices / grid_size * bbox_size + bbox_min - - return vertices, faces - - def __call__(self, grid_logits, **kwds): - - outputs = [] - # loop over the batches - for i in range(grid_logits.shape[0]): - try: - # process each batch - vertices, faces = self.run(grid_logits[i], **kwds) - vertices = vertices.astype(np.float32) - faces = np.ascontiguousarray(faces) - outputs.append(Latent2MeshOutput(mesh_v = vertices, mesh_f = faces)) - - except Exception: - import traceback - traceback.print_exc() - outputs.append(None) - - return outputs - -################################################ -# Volume Decoder -################################################ - -class VanillaVolumeDecoder(): - @torch.no_grad() - def __call__(self, latents: torch.Tensor, geo_decoder: callable, octree_res: int, bounds = 1.01, - num_chunks: int = 10_000): - - if isinstance(bounds, float): - bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds] - - bbox_min, bbox_max = torch.tensor(bounds[:3]), torch.tensor(bounds[3:]) - - x = torch.linspace(bbox_min[0], bbox_max[0], int(octree_res) + 1, dtype = torch.float32) - y = torch.linspace(bbox_min[1], bbox_max[1], int(octree_res) + 1, dtype = torch.float32) - z = torch.linspace(bbox_min[2], bbox_max[2], int(octree_res) + 1, dtype = torch.float32) - - [xs, ys, zs] = torch.meshgrid(x, y, z, indexing = "ij") - xyz = torch.stack((xs, ys, zs), axis=-1).to(latents.device, dtype = latents.dtype).contiguous().reshape(-1, 3) - grid_size = [int(octree_res) + 1, int(octree_res) + 1, int(octree_res) + 1] - - batch_logits = [] - for start in range(0, xyz.shape[0], num_chunks): - chunk_queries = xyz[start: start + num_chunks, :] - chunk_queries = chunk_queries.unsqueeze(0).repeat(latents.shape[0], 1, 1) - logits = geo_decoder(queries = chunk_queries, latents = latents) - batch_logits.append(logits) - - grid_logits = torch.cat(batch_logits, dim = 1) - grid_logits = grid_logits.view((latents.shape[0], *grid_size)).float() - - return grid_logits - -def export_to_trimesh(mesh_output): - import trimesh - - 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 - \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae/preprocess.py b/comfy/ldm/hunyuan3d/vae/preprocess.py deleted file mode 100644 index 1a946b27a..000000000 --- a/comfy/ldm/hunyuan3d/vae/preprocess.py +++ /dev/null @@ -1,165 +0,0 @@ -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]""" - - bbox = mesh.bounds - center = (bbox[1] + bbox[0]) / 2 - - max_extent = (bbox[1] - bbox[0]).max() - mesh.apply_translation(-center) - mesh.apply_scale((2 * scale) / max_extent) - - return mesh - -def sample_pointcloud(mesh, num = 200000): - """ Uniformly sample points from the surface of the mesh """ - - points, face_idx = mesh.sample(num, return_index = True) - normals = mesh.face_normals[face_idx] - return torch.from_numpy(points.astype(np.float32)), torch.from_numpy(normals.astype(np.float32)) - -def detect_sharp_edges(mesh, threshold=0.985): - """Return edge indices (a, b) that lie on sharp boundaries of the mesh.""" - - V, F = mesh.vertices, mesh.faces - VN, FN = mesh.vertex_normals, mesh.face_normals - - sharp_mask = np.ones(V.shape[0]) - for i in range(3): - indices = F[:, i] - alignment = np.einsum('ij,ij->i', VN[indices], FN) - dot_stack = np.stack((sharp_mask[indices], alignment), axis=-1) - sharp_mask[indices] = np.min(dot_stack, axis=-1) - - edge_a = np.concatenate([F[:, 0], F[:, 1], F[:, 2]]) - edge_b = np.concatenate([F[:, 1], F[:, 2], F[:, 0]]) - sharp_edges = (sharp_mask[edge_a] < threshold) & (sharp_mask[edge_b] < threshold) - - return edge_a[sharp_edges], edge_b[sharp_edges] - - -def sharp_sample_pointcloud(mesh, num = 16384): - """ Sample points preferentially from sharp edges in the mesh. """ - - edge_a, edge_b = detect_sharp_edges(mesh) - V, VN = mesh.vertices, mesh.vertex_normals - - va, vb = V[edge_a], V[edge_b] - na, nb = VN[edge_a], VN[edge_b] - - edge_lengths = np.linalg.norm(vb - va, axis=-1) - weights = edge_lengths / edge_lengths.sum() - - indices = np.searchsorted(np.cumsum(weights), np.random.rand(num)) - t = np.random.rand(num, 1) - - samples = t * va[indices] + (1 - t) * vb[indices] - normals = t * na[indices] + (1 - t) * nb[indices] - - return samples.astype(np.float32), normals.astype(np.float32) - -def load_surface_sharpedge(mesh, num_points=4096, num_sharp_points=4096, sharpedge_flag = True, device = "cuda"): - """Load a surface with optional sharp-edge annotations from a trimesh mesh.""" - - try: - mesh_full = trimesh.util.concatenate(mesh.dump()) - except Exception: - mesh_full = trimesh.util.concatenate(mesh) - - mesh_full = normalize_mesh(mesh_full) - - faces = mesh_full.faces - vertices = mesh_full.vertices - origin_face_count = faces.shape[0] - - mesh_surface = trimesh.Trimesh(vertices=vertices, faces=faces[:origin_face_count]) - mesh_fill = trimesh.Trimesh(vertices=vertices, faces=faces[origin_face_count:]) - - area_surface = mesh_surface.area - area_fill = mesh_fill.area - total_area = area_surface + area_fill - - sample_num = 499712 // 2 - fill_ratio = area_fill / total_area if total_area > 0 else 0 - - num_fill = int(sample_num * fill_ratio) - num_surface = sample_num - num_fill - - surf_pts, surf_normals = sample_pointcloud(mesh_surface, num_surface) - fill_pts, fill_normals = (torch.zeros(0, 3), torch.zeros(0, 3)) if num_fill == 0 else sample_pointcloud(mesh_fill, num_fill) - - sharp_pts, sharp_normals = sharp_sample_pointcloud(mesh_surface, sample_num) - - def assemble_tensor(points, normals, label=None): - - data = torch.cat([points, normals], dim=1).half().to(device) - - if label is not None: - label_tensor = torch.full((data.shape[0], 1), float(label), dtype=torch.float16).to(device) - data = torch.cat([data, label_tensor], dim=1) - - return data - - surface = assemble_tensor(torch.cat([surf_pts.to(device), fill_pts.to(device)], dim=0), - torch.cat([surf_normals.to(device), fill_normals.to(device)], dim=0), - label = 0 if sharpedge_flag else None) - - sharp_surface = assemble_tensor(torch.from_numpy(sharp_pts), torch.from_numpy(sharp_normals), - label = 1 if sharpedge_flag else None) - - rng = np.random.default_rng() - - surface = surface[rng.choice(surface.shape[0], num_points, replace = False)] - sharp_surface = sharp_surface[rng.choice(sharp_surface.shape[0], num_sharp_points, replace = False)] - - full = torch.cat([surface, sharp_surface], dim = 0).unsqueeze(0) - - return full - -class SharpEdgeSurfaceLoader: - """ Load mesh surface and sharp edge samples. """ - - def __init__(self, num_uniform_points = 8192, num_sharp_points = 8192): - - self.num_uniform_points = num_uniform_points - self.num_sharp_points = num_sharp_points - self.total_points = num_uniform_points + num_sharp_points - - def __call__(self, mesh_input, device = "cuda"): - mesh = self._load_mesh(mesh_input) - return load_surface_sharpedge(mesh, self.num_uniform_points, self.num_sharp_points, device = device) - - @staticmethod - def _load_mesh(mesh_input): - - if isinstance(mesh_input, str): - mesh = trimesh.load(mesh_input, force="mesh", merge_primitives = True) - else: - mesh = mesh_input - - if isinstance(mesh, trimesh.Scene): - combined = None - for obj in mesh.geometry.values(): - combined = obj if combined is None else combined + obj - return combined - - return mesh - -def test_preprocess(): - torch.set_default_device("cpu") - torch.manual_seed(2025) - np.random.seed(2025) - loader = SharpEdgeSurfaceLoader( - num_sharp_points = 0, - num_uniform_points = 81920, - ) - - mesh_demo = 'rock.glb' - surface = loader(mesh_demo, device = "cpu").to(dtype = torch.float16) - print(surface) - -if __name__ == "__main__": - test_preprocess() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/vae/transformer.py b/comfy/ldm/hunyuan3d/vae/transformer.py deleted file mode 100644 index b37e415e5..000000000 --- a/comfy/ldm/hunyuan3d/vae/transformer.py +++ /dev/null @@ -1,156 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F - -class DropPath(nn.Module): - """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks). - """ - - def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True): - super(DropPath, self).__init__() - self.drop_prob = drop_prob - self.scale_by_keep = scale_by_keep - - def forward(self, x): - - keep_prob = 1 - self.drop_prob - shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets - - random_tensor = x.new_empty(shape).bernoulli_(keep_prob) - - if keep_prob > 0.0 and self.scale_by_keep: - random_tensor.div_(keep_prob) - - return x * random_tensor - -class MLP(nn.Module): - def __init__(self, width: int, ratio: int = 4, drop_path_rate: float = 0): - super().__init__() - self.gelu = nn.GELU() - self.c_fc = nn.Linear(width, width * ratio) - self.c_proj = nn.Linear(width * ratio, width) - self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity() - - def forward(self, x): - return self.drop_path(self.c_proj(self.gelu(self.c_fc(x)))) - - -class QKVMultiheadAttention(nn.Module): - def __init__( - self, - heads: int, - n_ctx: int, - width=None, - qk_norm=False, - norm_layer=nn.LayerNorm - ): - super().__init__() - self.heads = heads - self.n_ctx = n_ctx - self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() - self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() - - def forward(self, qkv): - bs, n_ctx, width = qkv.shape - attn_ch = width // self.heads // 3 - qkv = qkv.view(bs, n_ctx, self.heads, -1) - q, k, v = torch.split(qkv, attn_ch, dim=-1) - - q = self.q_norm(q) - k = self.k_norm(k) - - q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)] - out = F.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(bs, n_ctx, -1) - return out - - -class MultiheadAttention(nn.Module): - def __init__( - self, - n_ctx: int, - width: int, - heads: int, - qkv_bias: bool, - norm_layer = nn.LayerNorm, - qk_norm: bool = False, - drop_path_rate: float = 0.0 - ): - super().__init__() - - self.c_qkv = nn.Linear(width, width * 3, bias=qkv_bias) - self.c_proj = nn.Linear(width, width) - - self.attention = QKVMultiheadAttention( - heads = heads, - n_ctx = n_ctx, - width = width, - norm_layer = norm_layer, - qk_norm = qk_norm - ) - self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity() - - def forward(self, x): - x = self.c_qkv(x) - x = self.attention(x) - x = self.drop_path(self.c_proj(x)) - return x - - -class ResAttnBlock(nn.Module): - def __init__( - self, - *, - n_ctx: int, - width: int, - heads: int, - qkv_bias: bool = True, - norm_layer=nn.LayerNorm, - qk_norm: bool = False, - drop_path_rate: float = 0.0, - ): - super().__init__() - self.attn = MultiheadAttention( - n_ctx=n_ctx, - width=width, - heads=heads, - qkv_bias=qkv_bias, - norm_layer=norm_layer, - qk_norm=qk_norm, - drop_path_rate=drop_path_rate - ) - self.ln_1 = norm_layer(width, elementwise_affine=True, eps=1e-6) - self.mlp = MLP(width=width, drop_path_rate=drop_path_rate) - self.ln_2 = norm_layer(width, elementwise_affine=True, eps=1e-6) - - def forward(self, x: torch.Tensor): - x = x + self.attn(self.ln_1(x)) - x = x + self.mlp(self.ln_2(x)) - return x - -class Transformer(nn.Module): - def __init__(self, n_ctx: int, heads: int, width: int, depth: int, - qkv_bias: bool = True, qk_norm: bool = False, drop_path_rate: float = 0.0): - super().__init__() - - self.resblocks = nn.ModuleList([ - ResAttnBlock(n_ctx = n_ctx, - heads = heads, - width = width, - qkv_bias = qkv_bias, - qk_norm = qk_norm, - drop_path_rate = drop_path_rate) - for _ in range(depth) - ]) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - - for resnet in self.resblocks: - x = resnet(x) - - return x - -if __name__ == "__main__": - torch.manual_seed(2025) - model = Transformer(512, 8, 224, 3) - outputs = model(x = torch.randn(1, 512, 224)) - print(outputs) diff --git a/comfy/ldm/hunyuan3d/vae/vae.py b/comfy/ldm/hunyuan3d/vae/vae.py deleted file mode 100644 index 9358a0ab2..000000000 --- a/comfy/ldm/hunyuan3d/vae/vae.py +++ /dev/null @@ -1,185 +0,0 @@ -import torch -import torch.nn as nn -from transformer import Transformer -from postprocess import VanillaVolumeDecoder, SufraceExtractor -from point_attention import PointCrossAttention, CrossAttentionDecoder - -class FourierEmbedder(nn.Module): - def __init__(self, num_freq: int = 8, input_dim: int = 3, include_pi: bool = False): - super().__init__() - - frequencies = 2.0 ** torch.arange( - num_freq, - dtype = torch.float32 - ) - - if include_pi: - frequencies *= torch.pi - - self.register_buffer("frequencies", frequencies, persistent = False) - - self.out_dim = input_dim * (num_freq * 2 + 1) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - - embed = (x[..., None].contiguous() * self.frequencies).view(*x.shape[:-1], -1) - return torch.cat((x, embed.sin(), embed.cos()), dim = -1) - -class DiagonalGaussianDistribution: - def __init__(self, params: torch.Tensor, feature_dim: int = -1): - - # divide quant channels (8) into mean and log variance - self.mean, self.logvar = torch.chunk(params, 2, dim = feature_dim) - - self.logvar = torch.clamp(self.logvar, -30.0, 20.0) - self.std = torch.exp(0.5 * self.logvar) - - def sample(self): - - eps = torch.randn_like(self.std) - z = self.mean + eps * self.std - - return z - -class VAE(nn.Module): - def __init__(self, - *, - num_latents: int = 4096, - embed_dim: int = 64, - width: int = 1024, - heads: int = 16, - num_decoder_layers: int = 16, - num_encoder_layers: int = 8, - pc_size: int = 81920, - pc_sharpedge_size: int = 0, - point_feats: int = 4, - downsample_ratio: int = 20, - geo_decoder_downsample_ratio: int = 1, - geo_decoder_mlp_expand_ratio: int = 4, - geo_decoder_ln_post: bool = True, - num_frequencies: int = 8, - qkv_bias: bool = False, - qk_norm: bool = True, - drop_path_rate: float = 0.0, - include_pi: bool = False, - scale_factor: float = 1.0039506158752403 - ): - - super().__init__() - - self.latent_shape = (num_latents, embed_dim) - self.scale_factor = scale_factor - - self.fourier_embedder = FourierEmbedder(num_freq = num_frequencies, include_pi = include_pi) - - self.encoder = PointCrossAttention(layers = num_encoder_layers, - num_latents = num_latents, - downsample_ratio = downsample_ratio, - heads = heads, - pc_size = pc_size, - width = width, - point_feats = point_feats, - fourier_embedder = self.fourier_embedder, - pc_sharpedge_size = pc_sharpedge_size) - - self.transformer = Transformer( - n_ctx=num_latents, - width=width, - depth=num_decoder_layers, - heads=heads, - qkv_bias=qkv_bias, - qk_norm=qk_norm, - drop_path_rate=drop_path_rate - ) - - self.geo_decoder = CrossAttentionDecoder( - fourier_embedder = self.fourier_embedder, - out_channels = 1, - num_latents = num_latents, - mlp_expand_ratio = geo_decoder_mlp_expand_ratio, - downsample_ratio = geo_decoder_downsample_ratio, - enable_ln_post = geo_decoder_ln_post, - width=width // geo_decoder_downsample_ratio, - heads=heads // geo_decoder_downsample_ratio, - qkv_bias = qkv_bias, - qk_norm= qk_norm - ) - - self.pre_kl = nn.Linear(width, embed_dim * 2) - self.post_kl = nn.Linear(embed_dim, width) - - self.volume_decoder = VanillaVolumeDecoder() - self.surface_extractor = SufraceExtractor() - - - def forward(self): - pass - - def encode(self, surface): - - pc, feats = surface[:, :, :3], surface[:, :, 3:] - latents = self.encoder(pc, feats) - - moments = self.pre_kl(latents) - posterior = DiagonalGaussianDistribution(moments, feature_dim = -1) - - latents = posterior.sample() - - return latents - - def decode(self, latents, to_mesh: bool = True, **kwargs): - - latents = self.post_kl(latents) - latents = self.transformer(latents) - - if not to_mesh: - return latents - - grid_logits = self.volume_decoder(latents = latents, geo_decoder = self.geo_decoder, **kwargs) - mesh = self.surface_extractor(grid_logits, **kwargs) - - return mesh - -def load_vae(vae): - - DEBUG = False - - checkpoint = "model.fp16.ckpt" - missing, unexpected = vae.load_state_dict(torch.load(checkpoint), strict = not DEBUG) - - if DEBUG: - print(f"Missing {len(missing)}: ", missing) - print(f"\nUnexpected {len(unexpected)}: ", unexpected) - - return vae - -def test_vae(): - - torch.manual_seed(2025) - vae = VAE() - vae = load_vae(vae) - - from preprocess import SharpEdgeSurfaceLoader - from postprocess import export_to_trimesh - - loader = SharpEdgeSurfaceLoader( - num_sharp_points = 0, - num_uniform_points = 81920, - ) - - mesh_demo = 'Duck.glb' - surface = loader(mesh_demo).to(dtype = torch.float16) - - latents = vae.encode(surface) - - mesh = vae.decode(latents, - num_chunks = 20000, - octree_res = 256, - to_mesh = True) - - mesh = export_to_trimesh(mesh)[0] - - mesh.export("duck_recreated.glb") - -if __name__ == "__main__": - test_vae() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3d/model_/conditioner.py b/comfy/ldm/hunyuan3dv2_1/conditioner.py similarity index 83% rename from comfy/ldm/hunyuan3d/model_/conditioner.py rename to comfy/ldm/hunyuan3dv2_1/conditioner.py index 93908ecb6..58de406e3 100644 --- a/comfy/ldm/hunyuan3d/model_/conditioner.py +++ b/comfy/ldm/hunyuan3dv2_1/conditioner.py @@ -1,8 +1,8 @@ import torch import torch.nn as nn import torch.nn.functional as F -from dinov2 import DinoConfig, Dinov2Model - +from image_encoders.dino2 import Dinov2Model +from dataclasses import dataclass, asdict # avoid using torchvision by recreating image processing functions def resize(img: torch.Tensor, size: int) -> torch.Tensor: @@ -61,6 +61,27 @@ def compose(transforms): return img return apply +# configuration for Dino Large +@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 + device: str = "cuda" + dtype = torch.float16 + class ImageEncoder(nn.Module): def __init__( self, @@ -71,7 +92,11 @@ class ImageEncoder(nn.Module): ): super().__init__() - self.model = Dinov2Model(config) + + import comfy.ops + ops = comfy.ops.disable_weight_init + + self.model = Dinov2Model(asdict(config), config.dtype, config.device, operations = ops) mean = [0.485, 0.456, 0.406] std = [0.229, 0.224, 0.225] diff --git a/comfy/ldm/hunyuan3d/model_/hunyuandit.py b/comfy/ldm/hunyuan3dv2_1/hunyuandit.py similarity index 99% rename from comfy/ldm/hunyuan3d/model_/hunyuandit.py rename to comfy/ldm/hunyuan3dv2_1/hunyuandit.py index d45d5c275..693590082 100644 --- a/comfy/ldm/hunyuan3d/model_/hunyuandit.py +++ b/comfy/ldm/hunyuan3dv2_1/hunyuandit.py @@ -355,7 +355,8 @@ class HunYuanDiTPlain(nn.Module): norm_type = 'layer', num_experts: int = 8, moe_top_k: int = 2, - use_fp16: bool = False + use_fp16: bool = False, + **kwargs ): super().__init__() @@ -404,7 +405,7 @@ class HunYuanDiTPlain(nn.Module): main_condition = contexts['main'] - time_embedded = self.t_embedder(t, condition=kwargs.get('guidance_cond')) + 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) diff --git a/comfy/ldm/hunyuan3d/model_/image_processor.py b/comfy/ldm/hunyuan3dv2_1/image_processor.py similarity index 100% rename from comfy/ldm/hunyuan3d/model_/image_processor.py rename to comfy/ldm/hunyuan3dv2_1/image_processor.py diff --git a/comfy/ldm/hunyuan3d/model_/moe.py b/comfy/ldm/hunyuan3dv2_1/moe.py similarity index 100% rename from comfy/ldm/hunyuan3d/model_/moe.py rename to comfy/ldm/hunyuan3dv2_1/moe.py diff --git a/comfy/ldm/hunyuan3d/model_/pipeline.py b/comfy/ldm/hunyuan3dv2_1/pipeline.py similarity index 100% rename from comfy/ldm/hunyuan3d/model_/pipeline.py rename to comfy/ldm/hunyuan3dv2_1/pipeline.py diff --git a/comfy/ldm/hunyuan3d/model_/scheduler.py b/comfy/ldm/hunyuan3dv2_1/scheduler.py similarity index 100% rename from comfy/ldm/hunyuan3d/model_/scheduler.py rename to comfy/ldm/hunyuan3dv2_1/scheduler.py diff --git a/comfy/ldm/hunyuan3d/model_/vae.py b/comfy/ldm/hunyuan3dv2_1/vae.py similarity index 99% rename from comfy/ldm/hunyuan3d/model_/vae.py rename to comfy/ldm/hunyuan3dv2_1/vae.py index 0ee313c9c..35986e917 100644 --- a/comfy/ldm/hunyuan3d/model_/vae.py +++ b/comfy/ldm/hunyuan3dv2_1/vae.py @@ -9,6 +9,8 @@ import trimesh import numpy as np from skimage import measure from dataclasses import dataclass +import torch.nn as nn +import torch.nn.functional as F def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True): @@ -53,11 +55,7 @@ def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = sampled_indicies.append(torch.arange(start, end)[selected]) return torch.cat(sampled_indicies, dim = 0) - -import torch -import torch.nn as nn -import torch.nn.functional as F - + class DropPath(nn.Module): """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks). """ @@ -624,11 +622,11 @@ class SufraceExtractor(): grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1] return grid_size, bbox_min, bbox_size - def run(self, grid_logit, *, bounds, octree_res, **kwargs): + def run(self, grid_logit, *, bounds, octree_res, level: float = 0.0, **kwargs): # grid_logit from volume decoder # use marching cube algo to turn an sdf to a mesh vertices, faces, _, _ = measure.marching_cubes(grid_logit.cpu().numpy(), - 0.0, + level, method = "lewiner") grid_size, bbox_min, bbox_size = self.compute_box_stat(bounds = bounds, octree_resolution = octree_res) diff --git a/comfy/model_base.py b/comfy/model_base.py index 4392355ea..6aa034a06 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -16,6 +16,8 @@ along with this program. If not, see . """ +import comfy.ldm.hunyuan3dv2_1 +import comfy.ldm.hunyuan3dv2_1.hunyuandit import torch import logging from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel, Timestep @@ -1196,6 +1198,46 @@ class Hunyuan3Dv2(BaseModel): if guidance is not None: out['guidance'] = comfy.conds.CONDRegular(torch.FloatTensor([guidance])) return out + +class Hunyuan3Dv2_1(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.hunyuan3dv2_1.hunyuandit.HunYuanDiTPlain) + + 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 extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + + guidance = kwargs.get("guidance", 5.0) + if guidance is not None: + + guidance_scale = torch.tensor([guidance], dtype = torch.float32, device = self.device) + + guidance_embed = self.get_guidance_scale_embedding(guidance_scale, + self.model.hidden_size, + dtype = next(self.model.parameters()).dtype) + + out['guidance_cond'] = comfy.conds.CONDRegular(guidance_embed) + + return out + class HiDream(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 18232ade3..44404f595 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -388,6 +388,20 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["guidance_embed"] = "{}guidance_in.in_layer.weight".format(key_prefix) in state_dict_keys return dit_config + if f"{key_prefix}t_embedder.mlp.2.weight" in state_dict_keys: # Hunyuan 3D 2.1 + + dit_config = {} + dit_config["image_model"] = "hunyuan3d2_1" + dit_config["in_channels"] = state_dict[f"{key_prefix}x_embedder.weight"].shape[1] + dit_config["context_dim"] = 1024 + dit_config["hidden_size"] = state_dict[f"{key_prefix}x_embedder.weight"].shape[0] + dit_config["mlp_ratio"] = 4.0 + dit_config["num_heads"] = 16 + dit_config["depth"] = count_blocks(state_dict_keys, f"{key_prefix}blocks.{{}}") + dit_config["qkv_bias"] = False + dit_config["guidance_cond_proj_dim"] = f"{key_prefix}t_embedder.cond_proj.weight" in state_dict_keys + return dit_config + if '{}caption_projection.0.linear.weight'.format(key_prefix) in state_dict_keys: # HiDream dit_config = {} dit_config["image_model"] = "hidream" diff --git a/comfy/sd.py b/comfy/sd.py index 5b95cf75a..33cadfcea 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -15,6 +15,7 @@ import comfy.ldm.lightricks.vae.causal_video_autoencoder import comfy.ldm.cosmos.vae import comfy.ldm.wan.vae import comfy.ldm.hunyuan3d.vae +import comfy.ldm.hunyuan3dv2_1.vae import comfy.ldm.ace.vae.music_dcae_pipeline import yaml import math @@ -441,6 +442,29 @@ class VAE: ddconfig = {"embed_dim": 64, "num_freqs": 8, "include_pi": False, "heads": 16, "width": 1024, "num_decoder_layers": 16, "qkv_bias": False, "qk_norm": True, "geo_decoder_mlp_expand_ratio": mlp_expand, "geo_decoder_downsample_ratio": downsample_ratio, "geo_decoder_ln_post": ln_post} self.first_stage_model = comfy.ldm.hunyuan3d.vae.ShapeVAE(**ddconfig) self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32] + + # Hunyuan 3d v2 2.1 + elif 'geo_decoder.cross_attn_decoder.mlp.c_proj.weight' in sd: + + self.latent_dim = 1 + + def estimate_memory(shape, dtype, num_layers = 16, kv_cache_multiplier = 2): + batch, num_tokens, hidden_dim = shape + dtype_size = model_management.dtype_size(dtype) + + total_mem = batch * num_tokens * hidden_dim * dtype_size * (1 + kv_cache_multiplier * num_layers) + return total_mem + + # better memory estimations + self.memory_used_encode = lambda shape, dtype, num_layers = 8, kv_cache_multiplier = 0:\ + estimate_memory(shape, dtype, num_layers, kv_cache_multiplier) + + self.memory_used_decode = lambda shape, dtype, num_layers = 16, kv_cache_multiplier = 2: \ + estimate_memory(shape, dtype, num_layers, kv_cache_multiplier) + + self.first_stage_model = comfy.ldm.hunyuan3dv2_1.vae.ShapeVAE() + self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32] + elif "vocoder.backbone.channel_layers.0.0.bias" in sd: #Ace Step Audio self.first_stage_model = comfy.ldm.ace.vae.music_dcae_pipeline.MusicDCAE(source_sample_rate=44100) self.memory_used_encode = lambda shape, dtype: (shape[2] * 330) * model_management.dtype_size(dtype) diff --git a/comfy/supported_models.py b/comfy/supported_models.py index 2669ca01e..b6b2bb237 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -1088,6 +1088,17 @@ class Hunyuan3Dv2(supported_models_base.BASE): def clip_target(self, state_dict={}): return None + +class Hunyuan3Dv2_1(Hunyuan3Dv2): + unet_config = { + "image_model": "hunyuan3d2_1", + } + + latent_format = latent_formats.Hunyuan3Dv2_1 + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.Hunyuan3Dv2_1(self, device = device) + return out class Hunyuan3Dv2mini(Hunyuan3Dv2): unet_config = { @@ -1217,6 +1228,6 @@ class Omnigen2(supported_models_base.BASE): return supported_models_base.ClipTarget(comfy.text_encoders.omnigen2.LuminaTokenizer, comfy.text_encoders.omnigen2.te(**hunyuan_detect)) -models = [LotusD, Stable_Zero123, SD15_instructpix2pix, SD15, SD20, SD21UnclipL, SD21UnclipH, SDXL_instructpix2pix, SDXLRefiner, SDXL, SSD1B, KOALA_700M, KOALA_1B, Segmind_Vega, SD_X4Upscaler, Stable_Cascade_C, Stable_Cascade_B, SV3D_u, SV3D_p, SD3, StableAudio, AuraFlow, PixArtAlpha, PixArtSigma, HunyuanDiT, HunyuanDiT1, FluxInpaint, Flux, FluxSchnell, GenmoMochi, LTXV, HunyuanVideoSkyreelsI2V, HunyuanVideoI2V, HunyuanVideo, CosmosT2V, CosmosI2V, CosmosT2IPredict2, CosmosI2VPredict2, Lumina2, WAN21_T2V, WAN21_I2V, WAN21_FunControl2V, WAN21_Vace, WAN21_Camera, Hunyuan3Dv2mini, Hunyuan3Dv2, HiDream, Chroma, ACEStep, Omnigen2] +models = [LotusD, Stable_Zero123, SD15_instructpix2pix, SD15, SD20, SD21UnclipL, SD21UnclipH, SDXL_instructpix2pix, SDXLRefiner, SDXL, SSD1B, KOALA_700M, KOALA_1B, Segmind_Vega, SD_X4Upscaler, Stable_Cascade_C, Stable_Cascade_B, SV3D_u, SV3D_p, SD3, StableAudio, AuraFlow, PixArtAlpha, PixArtSigma, HunyuanDiT, HunyuanDiT1, FluxInpaint, Flux, FluxSchnell, GenmoMochi, LTXV, HunyuanVideoSkyreelsI2V, HunyuanVideoI2V, HunyuanVideo, CosmosT2V, CosmosI2V, CosmosT2IPredict2, CosmosI2VPredict2, Lumina2, WAN21_T2V, WAN21_I2V, WAN21_FunControl2V, WAN21_Vace, WAN21_Camera, Hunyuan3Dv2mini, Hunyuan3Dv2, HiDream, Chroma, ACEStep, Omnigen2, Hunyuan3Dv2_1] models += [SVD_img2vid] diff --git a/comfy_extras/nodes_hunyuan3d.py b/comfy_extras/nodes_hunyuan3d.py index df0d1622c..1034f7fbc 100644 --- a/comfy_extras/nodes_hunyuan3d.py +++ b/comfy_extras/nodes_hunyuan3d.py @@ -8,22 +8,39 @@ import folder_paths import comfy.model_management from comfy.cli_args import args - class EmptyLatentHunyuan3Dv2: @classmethod def INPUT_TYPES(s): - return {"required": {"resolution": ("INT", {"default": 3072, "min": 1, "max": 8192}), - "batch_size": ("INT", {"default": 1, "min": 1, "max": 4096, "tooltip": "The number of latent images in the batch."}), - }} + return { + "required": { + "resolution": ("INT", {"default": 3072, "min": 1, "max": 8192}), + "batch_size": ("INT", { + "default": 1, + "min": 1, + "max": 4096, + "tooltip": "The number of latent images in the batch." + }), + "version": (["2.0", "2.1"], { + "default": "2.1", + "tooltip": "Choose latent layout version. 2.0: (B, C, N), 2.1: (B, N, C)" + }) + } + } + RETURN_TYPES = ("LATENT",) FUNCTION = "generate" - CATEGORY = "latent/3d" - def generate(self, resolution, batch_size): - latent = torch.zeros([batch_size, 64, resolution], device=comfy.model_management.intermediate_device()) - return ({"samples": latent, "type": "hunyuan3dv2"}, ) + def generate(self, resolution, batch_size, version): + embed_dim = 64 + if version == "2.0": + latent = torch.zeros([batch_size, embed_dim, resolution], + device = comfy.model_management.intermediate_device()) + else: # version = "2.1" + latent = torch.zeros([batch_size, resolution, embed_dim], + device = comfy.model_management.intermediate_device()) + return ({"samples": latent, "type": "hunyuan3dv2"}, ) class Hunyuan3Dv2Conditioning: @classmethod @@ -80,26 +97,48 @@ class Hunyuan3Dv2ConditioningMultiView: class VOXEL: def __init__(self, data): self.data = data - - class VAEDecodeHunyuan3D: @classmethod - def INPUT_TYPES(s): - return {"required": {"samples": ("LATENT", ), - "vae": ("VAE", ), - "num_chunks": ("INT", {"default": 8000, "min": 1000, "max": 500000}), - "octree_resolution": ("INT", {"default": 256, "min": 16, "max": 512}), - }} - RETURN_TYPES = ("VOXEL",) - FUNCTION = "decode" + def INPUT_TYPES(cls): + return { + "required": { + "samples": ("LATENT",), + "vae": ("VAE",), + "version": (["2.0", "2.1"], { + "default": "2.1", + "tooltip": "2.0 returns voxel grid; 2.1 returns implicit SDF function." + }), + "num_chunks": ("INT", { + "default": 8000, "min": 1000, "max": 500000, + "visible_if": {"version": "2.0"} + }), + "octree_resolution": ("INT", { + "default": 256, "min": 16, "max": 512, + "visible_if": {"version": "2.0"} + }), + } + } + RETURN_TYPES = ("VOXEL", "SDF_FUNCTION") + RETURN_NAMES = ("voxel", "sdf") + + FUNCTION = "decode" CATEGORY = "latent/3d" - def decode(self, vae, samples, num_chunks, octree_resolution): - voxels = VOXEL(vae.decode(samples["samples"], vae_options={"num_chunks": num_chunks, "octree_resolution": octree_resolution})) - return (voxels, ) + def decode(self, vae, samples, version, num_chunks, octree_resolution): + if version == "2.0": + voxel = vae.decode(samples["samples"], vae_options={ + "num_chunks": num_chunks, + "octree_resolution": octree_resolution + }) + return (VOXEL(voxel), None) + mesh = vae.decode(samples["samples"], to_mesh = True, + num_chunks = num_chunks, + octree_resolution = octree_resolution) + return (None, mesh) + def voxel_to_mesh(voxels, threshold=0.5, device=None): if device is None: device = torch.device("cpu")