fix the cross-attention num_head bug.

This commit is contained in:
junsong
2024-12-04 03:25:35 -08:00
parent 9ec31c864f
commit 81f58d87cd
6 changed files with 17 additions and 22 deletions
+2 -2
View File
@@ -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,
+2 -2
View File
@@ -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()
+2 -1
View File
@@ -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
+8 -9
View File
@@ -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(
+1 -5
View File
@@ -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, {}]], )