HunYuanDiT: Add 1.2 support

This commit is contained in:
City
2024-07-04 18:01:53 +02:00
parent 62b84531f4
commit 40d2a76310
3 changed files with 54 additions and 29 deletions
+26 -9
View File
@@ -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
+4 -7
View File
@@ -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
View File
@@ -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]