fix the cross-attention num_head bug.
This commit is contained in:
+2
-3
@@ -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, {}]], )
|
||||
|
||||
|
||||
+2
-2
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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, {}]], )
|
||||
|
||||
|
||||
Reference in New Issue
Block a user