diff --git a/HunYuanDiT/conf.py b/HunYuanDiT/conf.py index d057a39..4322753 100644 --- a/HunYuanDiT/conf.py +++ b/HunYuanDiT/conf.py @@ -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 diff --git a/HunYuanDiT/models/attn_layers.py b/HunYuanDiT/models/attn_layers.py index 4308af9..b767d83 100644 --- a/HunYuanDiT/models/attn_layers.py +++ b/HunYuanDiT/models/attn_layers.py @@ -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) diff --git a/HunYuanDiT/models/models.py b/HunYuanDiT/models/models.py index e481206..7dae413 100644 --- a/HunYuanDiT/models/models.py +++ b/HunYuanDiT/models/models.py @@ -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]