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