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 -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(