[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:
Changyi Yang
2026-03-21 21:06:49 -07:00
committed by GitHub
parent c32e35a2a5
commit c1794e2944

View File

@@ -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,
)