[diffusion] fix: fix Sana corrupted output by removing spurious QK norm layers (#20656)
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -89,7 +89,7 @@ class GLUMBConv(nn.Module):
|
||||
class SanaLinearAttention(nn.Module):
|
||||
"""Linear attention with O(N*D^2) complexity instead of O(N^2*D)."""
|
||||
|
||||
def __init__(self, query_dim, num_heads, head_dim, qk_norm_dim, bias=False):
|
||||
def __init__(self, query_dim, num_heads, head_dim, bias=False):
|
||||
super().__init__()
|
||||
inner_dim = num_heads * head_dim
|
||||
self.num_heads = num_heads
|
||||
@@ -101,8 +101,6 @@ class SanaLinearAttention(nn.Module):
|
||||
self.to_out = nn.ModuleList(
|
||||
[nn.Linear(inner_dim, query_dim, bias=True), nn.Identity()]
|
||||
)
|
||||
self.norm_q = RMSNorm(qk_norm_dim)
|
||||
self.norm_k = RMSNorm(qk_norm_dim)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
B, S, _ = hidden_states.shape
|
||||
@@ -111,9 +109,6 @@ class SanaLinearAttention(nn.Module):
|
||||
key = self.to_k(hidden_states)
|
||||
value = self.to_v(hidden_states)
|
||||
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
key = key.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
value = value.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
@@ -146,9 +141,6 @@ class SanaCrossAttention(nn.Module):
|
||||
[nn.Linear(inner_dim, query_dim, bias=True), nn.Identity()]
|
||||
)
|
||||
|
||||
self.norm_q = RMSNorm(inner_dim)
|
||||
self.norm_k = RMSNorm(inner_dim)
|
||||
|
||||
def forward(
|
||||
self, hidden_states, encoder_hidden_states, encoder_attention_mask=None
|
||||
):
|
||||
@@ -159,9 +151,6 @@ class SanaCrossAttention(nn.Module):
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
key = key.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
value = value.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
@@ -201,7 +190,6 @@ class SanaTransformerBlock(nn.Module):
|
||||
query_dim=dim,
|
||||
num_heads=num_attention_heads,
|
||||
head_dim=attention_head_dim,
|
||||
qk_norm_dim=num_attention_heads * attention_head_dim,
|
||||
bias=attention_bias,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user