From 81f58d87cdc69a24ad425ac44da7bc81ee4682d7 Mon Sep 17 00:00:00 2001 From: junsong Date: Wed, 4 Dec 2024 03:25:35 -0800 Subject: [PATCH] fix the cross-attention num_head bug. --- Gemma/nodes.py | 5 ++--- Sana/conf.py | 4 ++-- Sana/loader.py | 4 ++-- Sana/models/sana.py | 3 ++- Sana/models/sana_multi_scale.py | 17 ++++++++--------- Sana/nodes.py | 6 +----- 6 files changed, 17 insertions(+), 22 deletions(-) diff --git a/Gemma/nodes.py b/Gemma/nodes.py index 55a8a03..55e0883 100644 --- a/Gemma/nodes.py +++ b/Gemma/nodes.py @@ -109,11 +109,10 @@ class GemmaTextEncode: return_tensors="pt" ).to(text_encoder.device) - cond = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None] + cond = text_encoder(tokens.input_ids, tokens.attention_mask)[0] emb_masks = tokens.attention_mask - # 利用emb_masks将有效的cond选出来,其他置零 - # cond = cond * emb_masks.unsqueeze(-1) + cond = cond * emb_masks.unsqueeze(-1) return ([[cond, {}]], ) diff --git a/Sana/conf.py b/Sana/conf.py index 7f5a046..3719c0f 100644 --- a/Sana/conf.py +++ b/Sana/conf.py @@ -14,7 +14,7 @@ sana_conf = { "depth": 28, "hidden_size": 1152, "patch_size": 1, - "num_heads": 36, + "num_heads": 16, "linear_head_dim": 32, "model_max_length": 300, "y_norm": True, @@ -36,7 +36,7 @@ sana_conf = { "depth": 20, "hidden_size": 2240, "patch_size": 1, - "num_heads": 70, + "num_heads": 20, "linear_head_dim": 32, "model_max_length": 300, "y_norm": True, diff --git a/Sana/loader.py b/Sana/loader.py index 34ba85f..806ca64 100644 --- a/Sana/loader.py +++ b/Sana/loader.py @@ -48,7 +48,7 @@ class EXM_Sana_Model(comfy.model_base.BaseModel): return out -def load_sana(model_path, model_conf, dtype): +def load_sana(model_path, model_conf): state_dict = comfy.utils.load_torch_file(model_path) state_dict = state_dict.get("model", state_dict) @@ -62,7 +62,7 @@ def load_sana(model_path, model_conf, dtype): state_dict = convert_state_dict(state_dict) # Diffusers parameters = comfy.utils.calculate_parameters(state_dict) - unet_dtype = dtype + unet_dtype = comfy.model_management.unet_dtype() load_device = comfy.model_management.get_torch_device() offload_device = comfy.model_management.unet_offload_device() diff --git a/Sana/models/sana.py b/Sana/models/sana.py index 0dd6551..9da2d71 100644 --- a/Sana/models/sana.py +++ b/Sana/models/sana.py @@ -159,7 +159,7 @@ class Sana(nn.Module): caption_channels=2304, pe_interpolation=1.0, config=None, - model_max_length=120, + model_max_length=300, qk_norm=False, y_norm=False, norm_eps=1e-5, @@ -182,6 +182,7 @@ class Sana(nn.Module): self.depth = depth self.use_pe = use_pe self.y_norm = y_norm + self.model_max_length = model_max_length self.fp32_attention = kwargs.get("use_fp32_attention", False) kernel_size = patch_embed_kernel or patch_size diff --git a/Sana/models/sana_multi_scale.py b/Sana/models/sana_multi_scale.py index 7cc3745..2d5452c 100644 --- a/Sana/models/sana_multi_scale.py +++ b/Sana/models/sana_multi_scale.py @@ -296,20 +296,19 @@ class SanaMS(Sana): t = self.t_embedder(timestep) # (N, D) + y_lens = ((y != 0).sum(dim=3) > 0).sum(dim=2).squeeze().tolist() + y_lens = [y_lens[1]] * bs + + mask = torch.zeros((len(y_lens), self.model_max_length), dtype=torch.int).to(x.device) + for i, count in enumerate(y_lens): + mask[i, :count] = 1 + t0 = self.t_block(t) y = self.y_embedder(y, self.training, mask=mask) # (N, D) if self.y_norm: y = self.attention_y_norm(y) - if mask is not None: - if mask.shape[0] != y.shape[0]: - mask = mask.repeat(y.shape[0] // mask.shape[0], 1) - mask = mask.squeeze(1).squeeze(1) - y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1]) - y_lens = mask.sum(dim=1).tolist() - else: - y_lens = [y.shape[2]] * y.shape[0] - y = y.squeeze(1).view(1, -1, x.shape[-1]) + y = y.squeeze(1).masked_select(mask.unsqueeze(-1).bool()).view(1, -1, y.shape[-1]) for block in self.blocks: x = auto_grad_checkpoint( diff --git a/Sana/nodes.py b/Sana/nodes.py index f400962..57b5928 100644 --- a/Sana/nodes.py +++ b/Sana/nodes.py @@ -23,7 +23,6 @@ class SanaCheckpointLoader: "required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), "model": (list(sana_conf.keys()),), - "dtype": (dtypes,), } } RETURN_TYPES = ("MODEL",) @@ -32,13 +31,12 @@ class SanaCheckpointLoader: CATEGORY = "ExtraModels/Sana" TITLE = "Sana Checkpoint Loader" - def load_checkpoint(self, ckpt_name, model, dtype): + def load_checkpoint(self, ckpt_name, model): ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) model_conf = sana_conf[model] model = load_sana( model_path = ckpt_path, model_conf = model_conf, - dtype = string_to_dtype(dtype, "text_encoder") ) return (model,) @@ -132,8 +130,6 @@ class SanaTextEncode: emb_masks = tokens.attention_mask[:, select_idx] # 利用emb_masks将有效的embs选出来,其他置零 embs = embs * emb_masks.unsqueeze(-1) - # import IPython - # IPython.embed() return ([[embs, {}]], )