fix the cross-attention num_head bug.
This commit is contained in:
+2
-3
@@ -109,11 +109,10 @@ class GemmaTextEncode:
|
|||||||
return_tensors="pt"
|
return_tensors="pt"
|
||||||
).to(text_encoder.device)
|
).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 = tokens.attention_mask
|
||||||
|
|
||||||
# 利用emb_masks将有效的cond选出来,其他置零
|
cond = cond * emb_masks.unsqueeze(-1)
|
||||||
# cond = cond * emb_masks.unsqueeze(-1)
|
|
||||||
|
|
||||||
return ([[cond, {}]], )
|
return ([[cond, {}]], )
|
||||||
|
|
||||||
|
|||||||
+2
-2
@@ -14,7 +14,7 @@ sana_conf = {
|
|||||||
"depth": 28,
|
"depth": 28,
|
||||||
"hidden_size": 1152,
|
"hidden_size": 1152,
|
||||||
"patch_size": 1,
|
"patch_size": 1,
|
||||||
"num_heads": 36,
|
"num_heads": 16,
|
||||||
"linear_head_dim": 32,
|
"linear_head_dim": 32,
|
||||||
"model_max_length": 300,
|
"model_max_length": 300,
|
||||||
"y_norm": True,
|
"y_norm": True,
|
||||||
@@ -36,7 +36,7 @@ sana_conf = {
|
|||||||
"depth": 20,
|
"depth": 20,
|
||||||
"hidden_size": 2240,
|
"hidden_size": 2240,
|
||||||
"patch_size": 1,
|
"patch_size": 1,
|
||||||
"num_heads": 70,
|
"num_heads": 20,
|
||||||
"linear_head_dim": 32,
|
"linear_head_dim": 32,
|
||||||
"model_max_length": 300,
|
"model_max_length": 300,
|
||||||
"y_norm": True,
|
"y_norm": True,
|
||||||
|
|||||||
+2
-2
@@ -48,7 +48,7 @@ class EXM_Sana_Model(comfy.model_base.BaseModel):
|
|||||||
return out
|
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 = comfy.utils.load_torch_file(model_path)
|
||||||
state_dict = state_dict.get("model", state_dict)
|
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
|
state_dict = convert_state_dict(state_dict) # Diffusers
|
||||||
|
|
||||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
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()
|
load_device = comfy.model_management.get_torch_device()
|
||||||
offload_device = comfy.model_management.unet_offload_device()
|
offload_device = comfy.model_management.unet_offload_device()
|
||||||
|
|
||||||
|
|||||||
+2
-1
@@ -159,7 +159,7 @@ class Sana(nn.Module):
|
|||||||
caption_channels=2304,
|
caption_channels=2304,
|
||||||
pe_interpolation=1.0,
|
pe_interpolation=1.0,
|
||||||
config=None,
|
config=None,
|
||||||
model_max_length=120,
|
model_max_length=300,
|
||||||
qk_norm=False,
|
qk_norm=False,
|
||||||
y_norm=False,
|
y_norm=False,
|
||||||
norm_eps=1e-5,
|
norm_eps=1e-5,
|
||||||
@@ -182,6 +182,7 @@ class Sana(nn.Module):
|
|||||||
self.depth = depth
|
self.depth = depth
|
||||||
self.use_pe = use_pe
|
self.use_pe = use_pe
|
||||||
self.y_norm = y_norm
|
self.y_norm = y_norm
|
||||||
|
self.model_max_length = model_max_length
|
||||||
self.fp32_attention = kwargs.get("use_fp32_attention", False)
|
self.fp32_attention = kwargs.get("use_fp32_attention", False)
|
||||||
|
|
||||||
kernel_size = patch_embed_kernel or patch_size
|
kernel_size = patch_embed_kernel or patch_size
|
||||||
|
|||||||
@@ -296,20 +296,19 @@ class SanaMS(Sana):
|
|||||||
|
|
||||||
t = self.t_embedder(timestep) # (N, D)
|
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)
|
t0 = self.t_block(t)
|
||||||
y = self.y_embedder(y, self.training, mask=mask) # (N, D)
|
y = self.y_embedder(y, self.training, mask=mask) # (N, D)
|
||||||
if self.y_norm:
|
if self.y_norm:
|
||||||
y = self.attention_y_norm(y)
|
y = self.attention_y_norm(y)
|
||||||
|
|
||||||
if mask is not None:
|
y = y.squeeze(1).masked_select(mask.unsqueeze(-1).bool()).view(1, -1, y.shape[-1])
|
||||||
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])
|
|
||||||
|
|
||||||
for block in self.blocks:
|
for block in self.blocks:
|
||||||
x = auto_grad_checkpoint(
|
x = auto_grad_checkpoint(
|
||||||
|
|||||||
+1
-5
@@ -23,7 +23,6 @@ class SanaCheckpointLoader:
|
|||||||
"required": {
|
"required": {
|
||||||
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
|
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
|
||||||
"model": (list(sana_conf.keys()),),
|
"model": (list(sana_conf.keys()),),
|
||||||
"dtype": (dtypes,),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("MODEL",)
|
RETURN_TYPES = ("MODEL",)
|
||||||
@@ -32,13 +31,12 @@ class SanaCheckpointLoader:
|
|||||||
CATEGORY = "ExtraModels/Sana"
|
CATEGORY = "ExtraModels/Sana"
|
||||||
TITLE = "Sana Checkpoint Loader"
|
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)
|
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||||
model_conf = sana_conf[model]
|
model_conf = sana_conf[model]
|
||||||
model = load_sana(
|
model = load_sana(
|
||||||
model_path = ckpt_path,
|
model_path = ckpt_path,
|
||||||
model_conf = model_conf,
|
model_conf = model_conf,
|
||||||
dtype = string_to_dtype(dtype, "text_encoder")
|
|
||||||
)
|
)
|
||||||
return (model,)
|
return (model,)
|
||||||
|
|
||||||
@@ -132,8 +130,6 @@ class SanaTextEncode:
|
|||||||
emb_masks = tokens.attention_mask[:, select_idx]
|
emb_masks = tokens.attention_mask[:, select_idx]
|
||||||
# 利用emb_masks将有效的embs选出来,其他置零
|
# 利用emb_masks将有效的embs选出来,其他置零
|
||||||
embs = embs * emb_masks.unsqueeze(-1)
|
embs = embs * emb_masks.unsqueeze(-1)
|
||||||
# import IPython
|
|
||||||
# IPython.embed()
|
|
||||||
|
|
||||||
return ([[embs, {}]], )
|
return ([[embs, {}]], )
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user