HunYuanDiT: Add 1.2 support
This commit is contained in:
+26
-9
@@ -1,13 +1,6 @@
|
||||
"""
|
||||
List of all HYDiT model types / settings
|
||||
"""
|
||||
sampling_settings = {
|
||||
"beta_schedule" : "linear",
|
||||
"linear_start" : 0.00085,
|
||||
"linear_end" : 0.03,
|
||||
"timesteps" : 1000,
|
||||
}
|
||||
|
||||
from argparse import Namespace
|
||||
hydit_args = Namespace(**{ # normally from argparse
|
||||
"infer_mode": "torch",
|
||||
@@ -30,8 +23,32 @@ hydit_conf = {
|
||||
"input_size": (1024//8, 1024//8),
|
||||
"args": hydit_args,
|
||||
},
|
||||
"sampling_settings" : sampling_settings,
|
||||
"sampling_settings" : {
|
||||
"beta_schedule" : "linear",
|
||||
"linear_start" : 0.00085,
|
||||
"linear_end" : 0.03,
|
||||
"timesteps" : 1000,
|
||||
},
|
||||
},
|
||||
"G/2-1.2": {
|
||||
"unet_config": {
|
||||
"depth" : 40,
|
||||
"num_heads" : 16,
|
||||
"patch_size" : 2,
|
||||
"hidden_size" : 1408,
|
||||
"mlp_ratio" : 4.3637,
|
||||
"input_size": (1024//8, 1024//8),
|
||||
"cond_style": False,
|
||||
"cond_res" : False,
|
||||
"args": hydit_args,
|
||||
},
|
||||
"sampling_settings" : {
|
||||
"beta_schedule" : "linear",
|
||||
"linear_start" : 0.00085,
|
||||
"linear_end" : 0.018,
|
||||
"timesteps" : 1000,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
# these are the same as regular DiT, I think
|
||||
@@ -39,6 +56,6 @@ from ..DiT.conf import dit_conf
|
||||
for name in ["XL/2", "L/2", "B/2"]:
|
||||
hydit_conf[name] = {
|
||||
"unet_config": dit_conf[name]["unet_config"].copy(),
|
||||
"sampling_settings": sampling_settings,
|
||||
"sampling_settings": hydit_conf["G/2"]["sampling_settings"],
|
||||
}
|
||||
hydit_conf[name]["unet_config"]["args"] = hydit_args
|
||||
|
||||
@@ -362,13 +362,10 @@ class Attention(nn.Module):
|
||||
f'qq: {qq.shape}, q: {q.shape}, kk: {kk.shape}, k: {k.shape}'
|
||||
q, k = qq, kk
|
||||
|
||||
q = q * self.scale
|
||||
attn = q @ k.transpose(-2, -1) # [b, h, s, d] @ [b, h, d, s]
|
||||
attn = attn.softmax(dim=-1) # [b, h, s, s]
|
||||
attn = self.attn_drop(attn)
|
||||
x = attn @ v # [b, h, s, d]
|
||||
|
||||
x = x.transpose(1, 2).reshape(B, N, C) # [b, s, h, d]
|
||||
# just use SDP here for now
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
).permute(0, 2, 1, 3).contiguous().reshape(B, N, C)
|
||||
x = self.out_proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
|
||||
+24
-13
@@ -169,6 +169,8 @@ class HunYuanDiT(nn.Module):
|
||||
num_heads=16,
|
||||
mlp_ratio=4.0,
|
||||
log_fn=print,
|
||||
cond_style=True,
|
||||
cond_res=True,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -187,6 +189,8 @@ class HunYuanDiT(nn.Module):
|
||||
self.text_len = args.text_len
|
||||
self.text_len_t5 = args.text_len_t5
|
||||
self.norm = args.norm
|
||||
self.cond_res = cond_res
|
||||
self.cond_style = cond_style
|
||||
|
||||
use_flash_attn = args.infer_mode == 'fa'
|
||||
if use_flash_attn:
|
||||
@@ -205,11 +209,15 @@ class HunYuanDiT(nn.Module):
|
||||
# Attention pooling
|
||||
self.pooler = AttentionPool(self.text_len_t5, self.text_states_dim_t5, num_heads=8, output_dim=1024)
|
||||
|
||||
# Here we use a default learned embedder layer for future extension.
|
||||
self.style_embedder = nn.Embedding(1, hidden_size)
|
||||
|
||||
# Image size and crop size conditions
|
||||
self.extra_in_dim = 256 * 6 + hidden_size
|
||||
self.extra_in_dim = 0
|
||||
if self.cond_res:
|
||||
# Image size and crop size conditions
|
||||
self.extra_in_dim += 256 * 6
|
||||
if self.cond_style:
|
||||
# Here we use a default learned embedder layer for future extension.
|
||||
self.style_embedder = nn.Embedding(1, hidden_size)
|
||||
self.extra_in_dim += hidden_size
|
||||
|
||||
# Text embedding for `add`
|
||||
self.last_size = input_size
|
||||
@@ -310,16 +318,19 @@ class HunYuanDiT(nn.Module):
|
||||
# Build text tokens with pooling
|
||||
extra_vec = self.pooler(encoder_hidden_states_t5)
|
||||
|
||||
# Build image meta size tokens
|
||||
image_meta_size = timestep_embedding(image_meta_size.view(-1), 256) # [B * 6, 256]
|
||||
# if self.args.use_fp16:
|
||||
# image_meta_size = image_meta_size.half()
|
||||
image_meta_size = image_meta_size.view(-1, 6 * 256)
|
||||
extra_vec = torch.cat([extra_vec, image_meta_size], dim=1) # [B, D + 6 * 256]
|
||||
if self.cond_res:
|
||||
# Build image meta size tokens
|
||||
image_meta_size = timestep_embedding(image_meta_size.view(-1), 256) # [B * 6, 256]
|
||||
# if self.args.use_fp16:
|
||||
# image_meta_size = image_meta_size.half()
|
||||
|
||||
image_meta_size = image_meta_size.view(-1, 6 * 256)
|
||||
extra_vec = torch.cat([extra_vec, image_meta_size], dim=1) # [B, D + 6 * 256]
|
||||
|
||||
# Build style tokens
|
||||
style_embedding = self.style_embedder(style)
|
||||
extra_vec = torch.cat([extra_vec, style_embedding], dim=1)
|
||||
if self.cond_style:
|
||||
# Build style tokens
|
||||
style_embedding = self.style_embedder(style)
|
||||
extra_vec = torch.cat([extra_vec, style_embedding], dim=1)
|
||||
|
||||
# Concatenate all extra vectors
|
||||
c = t + self.extra_embedder(extra_vec.to(self.dtype)) # [B, D]
|
||||
|
||||
Reference in New Issue
Block a user