mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-24 10:56:37 +08:00
style changes
This commit is contained in:
parent
6316bb4cda
commit
f2a2d6b7b0
@ -71,7 +71,8 @@ class Dino2Block(torch.nn.Module):
|
|||||||
self.layer_scale2 = LayerScale(dim, dtype, device, operations)
|
self.layer_scale2 = LayerScale(dim, dtype, device, operations)
|
||||||
if use_swiglu_ffn:
|
if use_swiglu_ffn:
|
||||||
self.mlp = SwiGLUFFN(dim, dtype, device, operations)
|
self.mlp = SwiGLUFFN(dim, dtype, device, operations)
|
||||||
else: self.mlp = Dinov2MLP(dim, dtype, device, operations)
|
else:
|
||||||
|
self.mlp = Dinov2MLP(dim, dtype, device, operations)
|
||||||
self.norm1 = operations.LayerNorm(dim, eps=layer_norm_eps, dtype=dtype, device=device)
|
self.norm1 = operations.LayerNorm(dim, eps=layer_norm_eps, dtype=dtype, device=device)
|
||||||
self.norm2 = operations.LayerNorm(dim, eps=layer_norm_eps, dtype=dtype, device=device)
|
self.norm2 = operations.LayerNorm(dim, eps=layer_norm_eps, dtype=dtype, device=device)
|
||||||
|
|
||||||
@ -95,8 +96,8 @@ class Dino2Encoder(torch.nn.Module):
|
|||||||
intermediate_output = len(self.layer) + intermediate_output
|
intermediate_output = len(self.layer) + intermediate_output
|
||||||
|
|
||||||
intermediate = None
|
intermediate = None
|
||||||
for i, l in enumerate(self.layer):
|
for i, layer in enumerate(self.layer):
|
||||||
x = l(x, optimized_attention)
|
x = layer(x, optimized_attention)
|
||||||
if i == intermediate_output:
|
if i == intermediate_output:
|
||||||
intermediate = x.clone()
|
intermediate = x.clone()
|
||||||
return x, intermediate
|
return x, intermediate
|
||||||
|
|||||||
@ -46,7 +46,8 @@ def fps(src: torch.Tensor, batch: torch.Tensor, sampling_ratio: float, start_ran
|
|||||||
# select a random start point
|
# select a random start point
|
||||||
if start_random:
|
if start_random:
|
||||||
farthest = torch.randint(0, num_points, (1,), device = src.device)
|
farthest = torch.randint(0, num_points, (1,), device = src.device)
|
||||||
else: farthest = torch.tensor([0], device = src.device, dtype = torch.long)
|
else:
|
||||||
|
farthest = torch.tensor([0], device = src.device, dtype = torch.long)
|
||||||
|
|
||||||
for i in range(num_samples):
|
for i in range(num_samples):
|
||||||
selected[i] = farthest
|
selected[i] = farthest
|
||||||
@ -147,7 +148,8 @@ class PointCrossAttention(nn.Module):
|
|||||||
sharpedge_input_pc = torch.zeros(B, 0, D, dtype = random_input_pc.dtype).to(point_cloud.device)
|
sharpedge_input_pc = torch.zeros(B, 0, D, dtype = random_input_pc.dtype).to(point_cloud.device)
|
||||||
sharpedge_query_pc = torch.zeros(B, 0, D, dtype= random_query_pc.dtype).to(point_cloud.device)
|
sharpedge_query_pc = torch.zeros(B, 0, D, dtype= random_query_pc.dtype).to(point_cloud.device)
|
||||||
|
|
||||||
else: sharpedge_query_pc, sharpedge_input_pc, sharpedge_idx_pc, sharpedge_idx_query = \
|
else:
|
||||||
|
sharpedge_query_pc, sharpedge_input_pc, sharpedge_idx_pc, sharpedge_idx_query = \
|
||||||
self.subsample(pc = sharpedge_pc, num_query = num_sharpedge_query, input_pc_size = input_sharpedge_pc_size)
|
self.subsample(pc = sharpedge_pc, num_query = num_sharpedge_query, input_pc_size = input_sharpedge_pc_size)
|
||||||
|
|
||||||
# concat the random and sharpedges
|
# concat the random and sharpedges
|
||||||
@ -252,33 +254,6 @@ class PointCrossAttention(nn.Module):
|
|||||||
flattent_input_features.shape[-1])
|
flattent_input_features.shape[-1])
|
||||||
|
|
||||||
return input_surface_features, query_features
|
return input_surface_features, query_features
|
||||||
|
|
||||||
def forward(self, pc, feats):
|
|
||||||
"""
|
|
||||||
|
|
||||||
Args:
|
|
||||||
pc (torch.FloatTensor): [B, N, 3]
|
|
||||||
feats (torch.FloatTensor or None): [B, N, C]
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
|
|
||||||
"""
|
|
||||||
|
|
||||||
query, data = self.sample_points_and_latents(pc, feats)
|
|
||||||
|
|
||||||
query = self.input_proj(query)
|
|
||||||
query = query
|
|
||||||
data = self.input_proj(data)
|
|
||||||
data = data
|
|
||||||
|
|
||||||
latents = self.cross_attn(query, data)
|
|
||||||
if self.self_attn is not None:
|
|
||||||
latents = self.self_attn(latents)
|
|
||||||
|
|
||||||
if self.ln_post is not None:
|
|
||||||
latents = self.ln_post(latents)
|
|
||||||
|
|
||||||
return latents
|
|
||||||
|
|
||||||
def normalize_mesh(mesh, scale = 0.9999):
|
def normalize_mesh(mesh, scale = 0.9999):
|
||||||
"""Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""
|
"""Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""
|
||||||
@ -428,27 +403,6 @@ class SharpEdgeSurfaceLoader:
|
|||||||
return combined
|
return combined
|
||||||
|
|
||||||
return mesh
|
return mesh
|
||||||
|
|
||||||
class FourierEmbedder(nn.Module):
|
|
||||||
def __init__(self, num_freq: int = 8, input_dim: int = 3, include_pi: bool = False):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
frequencies = 2.0 ** torch.arange(
|
|
||||||
num_freq,
|
|
||||||
dtype = torch.float32
|
|
||||||
)
|
|
||||||
|
|
||||||
if include_pi:
|
|
||||||
frequencies *= torch.pi
|
|
||||||
|
|
||||||
self.register_buffer("frequencies", frequencies, persistent = False)
|
|
||||||
|
|
||||||
self.out_dim = input_dim * (num_freq * 2 + 1)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
||||||
|
|
||||||
embed = (x[..., None].contiguous() * self.frequencies).view(*x.shape[:-1], -1)
|
|
||||||
return torch.cat((x, embed.sin(), embed.cos()), dim = -1)
|
|
||||||
|
|
||||||
class DiagonalGaussianDistribution:
|
class DiagonalGaussianDistribution:
|
||||||
def __init__(self, params: torch.Tensor, feature_dim: int = -1):
|
def __init__(self, params: torch.Tensor, feature_dim: int = -1):
|
||||||
@ -489,7 +443,7 @@ class VanillaVolumeDecoder():
|
|||||||
grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1]
|
grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1]
|
||||||
|
|
||||||
batch_logits = []
|
batch_logits = []
|
||||||
for start in tqdm(range(0, xyz.shape[0], num_chunks), desc=f"Volume Decoding",
|
for start in tqdm(range(0, xyz.shape[0], num_chunks), desc="Volume Decoding",
|
||||||
disable=not enable_pbar):
|
disable=not enable_pbar):
|
||||||
|
|
||||||
chunk_queries = xyz[start: start + num_chunks, :]
|
chunk_queries = xyz[start: start + num_chunks, :]
|
||||||
@ -908,7 +862,7 @@ class CrossAttentionDecoder(nn.Module):
|
|||||||
self.query_proj = ops.Linear(self.fourier_embedder.out_dim, width)
|
self.query_proj = ops.Linear(self.fourier_embedder.out_dim, width)
|
||||||
if self.downsample_ratio != 1:
|
if self.downsample_ratio != 1:
|
||||||
self.latents_proj = ops.Linear(width * downsample_ratio, width)
|
self.latents_proj = ops.Linear(width * downsample_ratio, width)
|
||||||
if self.enable_ln_post == False:
|
if not self.enable_ln_post:
|
||||||
qk_norm = False
|
qk_norm = False
|
||||||
self.cross_attn_decoder = ResidualCrossAttentionBlock(
|
self.cross_attn_decoder = ResidualCrossAttentionBlock(
|
||||||
width=width,
|
width=width,
|
||||||
|
|||||||
@ -213,9 +213,6 @@ class Timesteps(nn.Module):
|
|||||||
|
|
||||||
x = timesteps.float().unsqueeze(1) * self.inv_freq.to(timesteps.device).unsqueeze(0)
|
x = timesteps.float().unsqueeze(1) * self.inv_freq.to(timesteps.device).unsqueeze(0)
|
||||||
|
|
||||||
# scale factor
|
|
||||||
if self.scale != 1.0:
|
|
||||||
emb = emb * self.scale
|
|
||||||
|
|
||||||
# fused CUDA kernels for sin and cos
|
# fused CUDA kernels for sin and cos
|
||||||
sin_emb = x.sin()
|
sin_emb = x.sin()
|
||||||
@ -223,6 +220,10 @@ class Timesteps(nn.Module):
|
|||||||
|
|
||||||
emb = torch.cat([sin_emb, cos_emb], dim = 1)
|
emb = torch.cat([sin_emb, cos_emb], dim = 1)
|
||||||
|
|
||||||
|
# scale factor
|
||||||
|
if self.scale != 1.0:
|
||||||
|
emb = emb * self.scale
|
||||||
|
|
||||||
# If we padded inv_freq for odd, emb is already wide enough; otherwise:
|
# If we padded inv_freq for odd, emb is already wide enough; otherwise:
|
||||||
if emb.shape[1] > self.num_channels:
|
if emb.shape[1] > self.num_channels:
|
||||||
emb = emb[:, :self.num_channels]
|
emb = emb[:, :self.num_channels]
|
||||||
@ -371,7 +372,8 @@ class Attention(nn.Module):
|
|||||||
|
|
||||||
if use_fp16:
|
if use_fp16:
|
||||||
eps = 1.0 / 65504
|
eps = 1.0 / 65504
|
||||||
else: eps = 1e-6
|
else:
|
||||||
|
eps = 1e-6
|
||||||
|
|
||||||
self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity()
|
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.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps = eps) if qk_norm else nn.Identity()
|
||||||
@ -434,7 +436,8 @@ class HunYuanDiTBlock(nn.Module):
|
|||||||
# eps can't be 1e-6 in fp16 mode because of numerical stability issues
|
# eps can't be 1e-6 in fp16 mode because of numerical stability issues
|
||||||
if use_fp16:
|
if use_fp16:
|
||||||
eps = 1.0 / 65504
|
eps = 1.0 / 65504
|
||||||
else: eps = 1e-6
|
else:
|
||||||
|
eps = 1e-6
|
||||||
|
|
||||||
self.norm1 = norm_layer(hidden_size, elementwise_affine = True, eps = eps)
|
self.norm1 = norm_layer(hidden_size, elementwise_affine = True, eps = eps)
|
||||||
|
|
||||||
@ -509,7 +512,8 @@ class FinalLayer(nn.Module):
|
|||||||
|
|
||||||
if use_fp16:
|
if use_fp16:
|
||||||
eps = 1.0 / 65504
|
eps = 1.0 / 65504
|
||||||
else: eps = 1e-6
|
else:
|
||||||
|
eps = 1e-6
|
||||||
|
|
||||||
self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine = True, eps = eps)
|
self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine = True, eps = eps)
|
||||||
self.linear = nn.Linear(final_hidden_size, out_channels, bias = True)
|
self.linear = nn.Linear(final_hidden_size, out_channels, bias = True)
|
||||||
@ -640,69 +644,4 @@ class HunYuanDiTPlain(nn.Module):
|
|||||||
output = output.movedim(-2, -1) * (-1.0)
|
output = output.movedim(-2, -1) * (-1.0)
|
||||||
|
|
||||||
cond_emb, uncond_emb = output.chunk(2, dim = 0)
|
cond_emb, uncond_emb = output.chunk(2, dim = 0)
|
||||||
return torch.cat([uncond_emb, cond_emb])
|
return torch.cat([uncond_emb, cond_emb])
|
||||||
|
|
||||||
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 = "hunyuan3dv2_1.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")
|
|
||||||
3
nodes.py
3
nodes.py
@ -1014,7 +1014,8 @@ class CLIPVisionEncode:
|
|||||||
|
|
||||||
if crop == "recenter":
|
if crop == "recenter":
|
||||||
crop_image = True
|
crop_image = True
|
||||||
else: border_ratio = None
|
else:
|
||||||
|
border_ratio = None
|
||||||
|
|
||||||
output = clip_vision.encode_image(image, crop=crop_image, border_ratio = border_ratio)
|
output = clip_vision.encode_image(image, crop=crop_image, border_ratio = border_ratio)
|
||||||
return (output,)
|
return (output,)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user