From dff570e364be2a606972d7308d435044924a1b04 Mon Sep 17 00:00:00 2001 From: Yousef Rafat <81116377+yousef-rafat@users.noreply.github.com> Date: Sat, 12 Jul 2025 00:00:25 +0300 Subject: [PATCH] file management and removed code redundancy --- comfy/ldm/hunyuan3d/vae.py | 110 +++----- comfy/ldm/hunyuan3dv2_1/conditioner.py | 189 ------------- comfy/ldm/hunyuan3dv2_1/hunyuandit.py | 182 +++++++++++- comfy/ldm/hunyuan3dv2_1/moe.py | 203 -------------- comfy/ldm/hunyuan3dv2_1/pipeline.py | 192 ------------- comfy/ldm/hunyuan3dv2_1/scheduler.py | 100 ------- comfy/ldm/hunyuan3dv2_1/vae.py | 368 +------------------------ 7 files changed, 220 insertions(+), 1124 deletions(-) delete mode 100644 comfy/ldm/hunyuan3dv2_1/conditioner.py delete mode 100644 comfy/ldm/hunyuan3dv2_1/moe.py delete mode 100644 comfy/ldm/hunyuan3dv2_1/pipeline.py delete mode 100644 comfy/ldm/hunyuan3dv2_1/scheduler.py diff --git a/comfy/ldm/hunyuan3d/vae.py b/comfy/ldm/hunyuan3d/vae.py index 6595d96d6..29d4b89d9 100644 --- a/comfy/ldm/hunyuan3d/vae.py +++ b/comfy/ldm/hunyuan3d/vae.py @@ -6,79 +6,47 @@ import torch.nn as nn import torch.nn.functional as F -from typing import Union, Tuple, List, Callable, Optional +from typing import 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 +################################################ +# Volume Decoder +################################################ - 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: +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 + def __call__(self, latents: torch.Tensor, geo_decoder: callable, octree_resolution: int, bounds = 1.01, + num_chunks: int = 10_000): + 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) + bbox_min, bbox_max = torch.tensor(bounds[:3]), torch.tensor(bounds[3:]) + + x = torch.linspace(bbox_min[0], bbox_max[0], int(octree_resolution) + 1, dtype = torch.float32) + y = torch.linspace(bbox_min[1], bbox_max[1], int(octree_resolution) + 1, dtype = torch.float32) + z = torch.linspace(bbox_min[2], bbox_max[2], int(octree_resolution) + 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_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1] - # 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) + 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((batch_size, *grid_size)).float() + grid_logits = torch.cat(batch_logits, dim = 1) + grid_logits = grid_logits.view((latents.shape[0], *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: @@ -175,13 +143,6 @@ class FourierEmbedder(nn.Module): 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). """ @@ -232,38 +193,41 @@ class MLP(nn.Module): 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, + n_data = None, width=None, qk_norm=False, norm_layer=ops.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() - 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 + 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__( @@ -306,7 +270,6 @@ class MultiheadCrossAttention(nn.Module): x = self.c_proj(x) return x - class ResidualCrossAttentionBlock(nn.Module): def __init__( self, @@ -366,7 +329,7 @@ class QKVMultiheadAttention(nn.Module): 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)) + 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 @@ -383,8 +346,7 @@ class MultiheadAttention(nn.Module): 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( diff --git a/comfy/ldm/hunyuan3dv2_1/conditioner.py b/comfy/ldm/hunyuan3dv2_1/conditioner.py deleted file mode 100644 index 58de406e3..000000000 --- a/comfy/ldm/hunyuan3dv2_1/conditioner.py +++ /dev/null @@ -1,189 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F -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: - - batched = img.ndim == 4 - - if not batched: - img = img.unsqueeze(0) - - _, _, 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 = F.interpolate(img, size = (new_h, new_w), mode = 'bilinear', align_corners = False, antialias = True ) - - if not batched: - img = img.squeeze(0) - - return img - - -def center_crop(img: torch.Tensor, size: int) -> torch.Tensor: - - batched = img.ndim == 4 - if not batched: - img = img.unsqueeze(0) - - _, _, h, w = img.shape - top = (h - size) // 2 - left = (w - size) // 2 - - - cropped = img[..., top:top + size, left:left + size] - - if not batched: - cropped = cropped.squeeze(0) - - return cropped - -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 - -# 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, - config: DinoConfig, - use_cls_token = True, - image_size = 518, - **kwargs, - ): - super().__init__() - - - 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] - - 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 image.ndim == 3: - image = image.unsqueeze(0) - - if value_range is not None: - low, high = value_range - image = (image - low) / (high - low) - - inputs = self.transform(image) - inputs = inputs.to(self.model.device, dtype=self.model.dtype) - last_hidden_state = self.model(inputs) - - 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, config): - super().__init__() - self.main_image_encoder = ImageEncoder(config) - - 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 test_image_encoder(): - - torch.manual_seed(2025) - config = DinoConfig() - image_encoder = SingleImageEncoder(config) - - image = torch.rand(3, 224, 224) - - outputs = image_encoder(image) - - print(outputs) - -if __name__ == "__main__": - #test_image_encoder() - conditioner = SingleImageEncoder(DinoConfig()) - torch.manual_seed(2025) - image = torch.rand(1, 3, 224, 224) - outputs = conditioner(image) - print(outputs["main"].size()) \ No newline at end of file diff --git a/comfy/ldm/hunyuan3dv2_1/hunyuandit.py b/comfy/ldm/hunyuan3dv2_1/hunyuandit.py index 693590082..2ea14f47c 100644 --- a/comfy/ldm/hunyuan3dv2_1/hunyuandit.py +++ b/comfy/ldm/hunyuan3dv2_1/hunyuandit.py @@ -2,9 +2,189 @@ 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 GELU(nn.Module): + + def __init__(self, dim_in: int, dim_out: int): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out) + + def gelu(self, gate: torch.Tensor) -> torch.Tensor: + + if gate.device.type == "mps": + return F.gelu(gate.to(dtype = torch.float32)).to(dtype = gate.dtype) + + return F.gelu(gate) + + def forward(self, hidden_states): + + hidden_states = self.proj(hidden_states) + hidden_states = self.gelu(hidden_states) + + return hidden_states + +class FeedForward(nn.Module): + + def __init__(self, dim: int, dim_out = None, mult: int = 4, + dropout: float = 0.0, inner_dim = None): + + super().__init__() + if inner_dim is None: + inner_dim = int(dim * mult) + + dim_out = dim_out if dim_out is not None else dim + + act_fn = GELU(dim, inner_dim) + + self.net = nn.ModuleList([]) + self.net.append(act_fn) + + self.net.append(nn.Dropout(dropout)) + self.net.append(nn.Linear(inner_dim, dim_out)) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + for module in self.net: + hidden_states = module(hidden_states) + return hidden_states + +class AddAuxLoss(torch.autograd.Function): + + @staticmethod + def forward(ctx, x, loss): + # do nothing in forward (no computation) + ctx.requires_aux_loss = loss.requires_grad + ctx.dtype = loss.dtype + + return x + + @staticmethod + def backward(ctx, grad_output): + # add the aux loss gradients + grad_loss = None + # put the aux grad the same as the main grad loss + # aux grad contributes equally + if ctx.requires_aux_loss: + grad_loss = torch.ones(1, dtype = ctx.dtype, device = grad_output.device) + + return grad_output, grad_loss + +class MoEGate(nn.Module): + + def __init__(self, embed_dim, num_experts=16, num_experts_per_tok=2, aux_loss_alpha=0.01): + + super().__init__() + self.top_k = num_experts_per_tok + self.n_routed_experts = num_experts + + self.alpha = aux_loss_alpha + + self.gating_dim = embed_dim + self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim))) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + + # flatten hidden states + hidden_states = hidden_states.view(-1, hidden_states.size(-1)) + + # get logits and pass it to softmax + logits = F.linear(hidden_states, self.weight, bias = None) + scores = logits.softmax(dim = -1) + + topk_weight, topk_idx = torch.topk(scores, k = self.top_k, dim = -1, sorted = False) + + if self.training and self.alpha > 0.0: + scores_for_aux = scores + + # used bincount instead of one hot encoding + counts = torch.bincount(topk_idx.view(-1), minlength = self.n_routed_experts).float() + ce = counts / topk_idx.numel() # normalized expert usage + + # mean expert score + Pi = scores_for_aux.mean(0) + + # expert balance loss + aux_loss = (Pi * ce * self.n_routed_experts).sum() * self.alpha + else: + aux_loss = None + + return topk_idx, topk_weight, aux_loss + +class MoEBlock(nn.Module): + def __init__(self, dim, num_experts: int = 6, moe_top_k: int = 2, dropout: float = 0.0, ff_inner_dim: int = None): + super().__init__() + + self.moe_top_k = moe_top_k + self.num_experts = num_experts + + self.experts = nn.ModuleList([ + FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim) + for _ in range(num_experts) + ]) + + self.gate = MoEGate(dim, num_experts = num_experts, num_experts_per_tok = moe_top_k) + self.shared_experts = FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim) + + def forward(self, hidden_states) -> torch.Tensor: + + identity = hidden_states + orig_shape = hidden_states.shape + topk_idx, topk_weight, aux_loss = self.gate(hidden_states) + + hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) + flat_topk_idx = topk_idx.view(-1) + + if self.training: + + hidden_states = hidden_states.repeat_interleave(self.moe_top_k, dim = 0) + y = torch.empty_like(hidden_states, dtype = hidden_states.dtype) + + for i, expert in enumerate(self.experts): + tmp = expert(hidden_states[flat_topk_idx == i]) + y[flat_topk_idx == i] = tmp.to(hidden_states.dtype) + + y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim = 1) + y = y.view(*orig_shape) + + y = AddAuxLoss.apply(y, aux_loss) + else: + y = self.moe_infer(hidden_states, flat_expert_indices = flat_topk_idx,flat_expert_weights = topk_weight.view(-1, 1)).view(*orig_shape) + + y = y + self.shared_experts(identity) + + return y + + @torch.no_grad() + def moe_infer(self, x, flat_expert_indices, flat_expert_weights): + + expert_cache = torch.zeros_like(x) + idxs = flat_expert_indices.argsort() + + # no need for .numpy().cpu() here + tokens_per_expert = flat_expert_indices.bincount().cumsum(0) + token_idxs = idxs // self.moe_top_k + + for i, end_idx in enumerate(tokens_per_expert): + + start_idx = 0 if i == 0 else tokens_per_expert[i-1] + + if start_idx == end_idx: + continue + + expert = self.experts[i] + exp_token_idx = token_idxs[start_idx:end_idx] + + expert_tokens = x[exp_token_idx] + expert_out = expert(expert_tokens) + + expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]]) + + # use index_add_ with a 1-D index tensor directly avoids building a large [N, D] index map and extra memcopy required by scatter_reduce_ + # + avoid dtype conversion + expert_cache.index_add_(0, exp_token_idx, expert_out) + + return expert_cache + class Timesteps(nn.Module): def __init__(self, num_channels: int, downscale_freq_shift: float = 0.0, scale: float = 1.0, max_period: int = 10000): diff --git a/comfy/ldm/hunyuan3dv2_1/moe.py b/comfy/ldm/hunyuan3dv2_1/moe.py deleted file mode 100644 index d09499714..000000000 --- a/comfy/ldm/hunyuan3dv2_1/moe.py +++ /dev/null @@ -1,203 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F - -class GELU(nn.Module): - - def __init__(self, dim_in: int, dim_out: int): - super().__init__() - self.proj = nn.Linear(dim_in, dim_out) - - def gelu(self, gate: torch.Tensor) -> torch.Tensor: - - if gate.device.type == "mps": - return F.gelu(gate.to(dtype = torch.float32)).to(dtype = gate.dtype) - - return F.gelu(gate) - - def forward(self, hidden_states): - - hidden_states = self.proj(hidden_states) - hidden_states = self.gelu(hidden_states) - - return hidden_states - -class FeedForward(nn.Module): - - def __init__(self, dim: int, dim_out = None, mult: int = 4, - dropout: float = 0.0, inner_dim = None): - - super().__init__() - if inner_dim is None: - inner_dim = int(dim * mult) - - dim_out = dim_out if dim_out is not None else dim - - act_fn = GELU(dim, inner_dim) - - self.net = nn.ModuleList([]) - self.net.append(act_fn) - - self.net.append(nn.Dropout(dropout)) - self.net.append(nn.Linear(inner_dim, dim_out)) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - for module in self.net: - hidden_states = module(hidden_states) - return hidden_states - -class AddAuxLoss(torch.autograd.Function): - - @staticmethod - def forward(ctx, x, loss): - # do nothing in forward (no computation) - ctx.requires_aux_loss = loss.requires_grad - ctx.dtype = loss.dtype - - return x - - @staticmethod - def backward(ctx, grad_output): - # add the aux loss gradients - grad_loss = None - # put the aux grad the same as the main grad loss - # aux grad contributes equally - if ctx.requires_aux_loss: - grad_loss = torch.ones(1, dtype = ctx.dtype, device = grad_output.device) - - return grad_output, grad_loss - -class MoEGate(nn.Module): - - def __init__(self, embed_dim, num_experts=16, num_experts_per_tok=2, aux_loss_alpha=0.01): - - super().__init__() - self.top_k = num_experts_per_tok - self.n_routed_experts = num_experts - - self.alpha = aux_loss_alpha - - self.gating_dim = embed_dim - self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim))) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - - # flatten hidden states - hidden_states = hidden_states.view(-1, hidden_states.size(-1)) - - # get logits and pass it to softmax - logits = F.linear(hidden_states, self.weight, bias = None) - scores = logits.softmax(dim = -1) - - topk_weight, topk_idx = torch.topk(scores, k = self.top_k, dim = -1, sorted = False) - - if self.training and self.alpha > 0.0: - scores_for_aux = scores - - # used bincount instead of one hot encoding - counts = torch.bincount(topk_idx.view(-1), minlength = self.n_routed_experts).float() - ce = counts / topk_idx.numel() # normalized expert usage - - # mean expert score - Pi = scores_for_aux.mean(0) - - # expert balance loss - aux_loss = (Pi * ce * self.n_routed_experts).sum() * self.alpha - else: - aux_loss = None - - return topk_idx, topk_weight, aux_loss - -class MoEBlock(nn.Module): - def __init__(self, dim, num_experts: int = 6, moe_top_k: int = 2, dropout: float = 0.0, ff_inner_dim: int = None): - super().__init__() - - self.moe_top_k = moe_top_k - self.num_experts = num_experts - - self.experts = nn.ModuleList([ - FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim) - for _ in range(num_experts) - ]) - - self.gate = MoEGate(dim, num_experts = num_experts, num_experts_per_tok = moe_top_k) - self.shared_experts = FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim) - - def forward(self, hidden_states) -> torch.Tensor: - - identity = hidden_states - orig_shape = hidden_states.shape - topk_idx, topk_weight, aux_loss = self.gate(hidden_states) - - hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) - flat_topk_idx = topk_idx.view(-1) - - if self.training: - - hidden_states = hidden_states.repeat_interleave(self.moe_top_k, dim = 0) - y = torch.empty_like(hidden_states, dtype = hidden_states.dtype) - - for i, expert in enumerate(self.experts): - tmp = expert(hidden_states[flat_topk_idx == i]) - y[flat_topk_idx == i] = tmp.to(hidden_states.dtype) - - y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim = 1) - y = y.view(*orig_shape) - - y = AddAuxLoss.apply(y, aux_loss) - else: - y = self.moe_infer(hidden_states, flat_expert_indices = flat_topk_idx,flat_expert_weights = topk_weight.view(-1, 1)).view(*orig_shape) - - y = y + self.shared_experts(identity) - - return y - - @torch.no_grad() - def moe_infer(self, x, flat_expert_indices, flat_expert_weights): - - expert_cache = torch.zeros_like(x) - idxs = flat_expert_indices.argsort() - - # no need for .numpy().cpu() here - tokens_per_expert = flat_expert_indices.bincount().cumsum(0) - token_idxs = idxs // self.moe_top_k - - for i, end_idx in enumerate(tokens_per_expert): - - start_idx = 0 if i == 0 else tokens_per_expert[i-1] - - if start_idx == end_idx: - continue - - expert = self.experts[i] - exp_token_idx = token_idxs[start_idx:end_idx] - - expert_tokens = x[exp_token_idx] - expert_out = expert(expert_tokens) - - expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]]) - - # use index_add_ with a 1-D index tensor directly avoids building a large [N, D] index map and extra memcopy required by scatter_reduce_ - # + avoid dtype conversion - expert_cache.index_add_(0, exp_token_idx, expert_out) - - return expert_cache - -def test_moe(): - - torch.manual_seed(2025) - import time - - start = time.time() - - moe_gate = MoEGate(512) - print(moe_gate(torch.rand(1, 71, 512))) - - #moe_block = MoEBlock(512) - #moe_block(torch.rand(1, 77, 512)) - - timing = time.time() - start - print(timing) - -if __name__ == "__main__": - test_moe() \ No newline at end of file diff --git a/comfy/ldm/hunyuan3dv2_1/pipeline.py b/comfy/ldm/hunyuan3dv2_1/pipeline.py deleted file mode 100644 index c1f5953f7..000000000 --- a/comfy/ldm/hunyuan3dv2_1/pipeline.py +++ /dev/null @@ -1,192 +0,0 @@ -import torch -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 - -import sys -import os - -def find_project_root(target_folder_name="ComfyUI"): - """ Walks directory until it finds ComfyUI base directroy """ - current = os.path.abspath(os.path.dirname(__file__)) - while True: - if os.path.basename(current) == target_folder_name: - return current - parent = os.path.dirname(current) - if parent == current: - raise RuntimeError(f"Could not find folder named '{target_folder_name}' in parent paths.") - current = parent - -comfyui_root = find_project_root() -sys.path.append(comfyui_root) - -from comfy_extras.nodes_hunyuan3d import save_glb - -class Hunyuan3DDiTFlowMatchingPipeline(nn.Module): - def __init__(self, model, vae, conditioner, image_processor, scheduler, device, dtype): - super().__init__() - - self.vae = vae - self.model = model - self.conditioner = conditioner - self.image_processor = image_processor - self.scheduler = scheduler - self.device = device - self.dtype = dtype - - def compile(self): - self.vae = torch.compile(self.vae) - self.model = torch.compile(self.model) - self.conditioner = torch.compile(self.conditioner) - - def load_ckpt(self, checkpoint_path: str): - - checkpoint = torch.load(checkpoint_path, weights_only = True) - self.model.load_state_dict(checkpoint["model"]) - self.vae.load_state_dict(checkpoint["vae"]) - self.conditioner.load_state_dict(checkpoint["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 to(self, device=None, dtype=None): - if dtype is not None: - self.dtype = dtype - self.vae.to(dtype=dtype) - self.model.to(dtype=dtype) - self.conditioner.to(dtype=dtype) - if device is not None: - self.device = torch.device(device) - self.vae.to(device) - self.model.to(device) - self.conditioner.to(device) - - def prepare_images(self, images): - - if isinstance(images, (str, Image.Image)): - return self.image_processor(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, - save_file = None, - **kwargs, - ): - - 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_images(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_training_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, 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) - - try: - if save_file is not None: - for i, output in enumerate(mesh): - output_file = f"{save_file}_{i}" if len(mesh) > 1 else save_file - save_glb(output.mesh_v, output.mesh_f, output_file, numpy_ready = True) - except Exception as e: - print(e) - - return mesh \ No newline at end of file diff --git a/comfy/ldm/hunyuan3dv2_1/scheduler.py b/comfy/ldm/hunyuan3dv2_1/scheduler.py deleted file mode 100644 index 200b9c56c..000000000 --- a/comfy/ldm/hunyuan3dv2_1/scheduler.py +++ /dev/null @@ -1,100 +0,0 @@ -import torch - -class EulerScheduler(torch.nn.Module): - def __init__(self, num_training_timesteps: int = 1_000, shift: float = 1, - num_inference_timesteps: int = 50, inference: bool = True, device: str = "cuda"): - 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, dtype = torch.float32, device = device) - timesteps = sigmas * self.num_training_timesteps - - self.timesteps = timesteps.to(device = device) - self.sigmas = torch.cat([sigmas, torch.ones(1, 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 step(self, model_output: torch.FloatTensor, sample: torch.FloatTensor,): - - sample = sample.to(torch.float32) - - sigma = self.sigmas[self._step_index] - sigma_next = self.sigmas[self._step_index + 1] - - prev_sample = sample + (sigma_next - sigma) * 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/hunyuan3dv2_1/vae.py b/comfy/ldm/hunyuan3dv2_1/vae.py index 219bc430e..d09b5614a 100644 --- a/comfy/ldm/hunyuan3dv2_1/vae.py +++ b/comfy/ldm/hunyuan3dv2_1/vae.py @@ -1,7 +1,3 @@ -# 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 @@ -9,7 +5,9 @@ import numpy as np from skimage import measure from dataclasses import dataclass import torch.nn as nn -import torch.nn.functional as F +from hunyuan3d.vae import ( + CrossAttentionDecoder, Transformer, ResidualCrossAttentionBlock, FourierEmbedder, VanillaVolumeDecoder +) def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True): @@ -54,332 +52,6 @@ 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) - -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 - -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, @@ -652,40 +324,6 @@ class SufraceExtractor(): 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 normalize_mesh(mesh, scale = 0.9999): """Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""