[diffusion] Sana: pack self-attn q/k/v and cross-attn k/v into single GEMMs (#28393)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Xiaoyu Zhang
2026-06-19 13:25:52 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent 7e20e25848
commit 7da92e1112
2 changed files with 24 additions and 10 deletions
@@ -41,6 +41,13 @@ class SanaArchConfig(DiTArchConfig):
param_names_mapping: dict = field(
default_factory=lambda: {
# self linear-attn: merge q/k/v into to_qkv (concat order q, k, v)
r"^(transformer_blocks\.\d+\.attn1)\.to_q\.(.*)$": (r"\1.to_qkv.\2", 0, 3),
r"^(transformer_blocks\.\d+\.attn1)\.to_k\.(.*)$": (r"\1.to_qkv.\2", 1, 3),
r"^(transformer_blocks\.\d+\.attn1)\.to_v\.(.*)$": (r"\1.to_qkv.\2", 2, 3),
# cross-attn: merge k/v into to_kv (q stays separate)
r"^(transformer_blocks\.\d+\.attn2)\.to_k\.(.*)$": (r"\1.to_kv.\2", 0, 2),
r"^(transformer_blocks\.\d+\.attn2)\.to_v\.(.*)$": (r"\1.to_kv.\2", 1, 2),
r"^transformer\.(.*)$": r"\1",
}
)
@@ -7,6 +7,7 @@ from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbed
from sglang.multimodal_gen.configs.models.dits.sana import SanaConfig
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
from sglang.multimodal_gen.runtime.layers.linear import MergedColumnParallelLinear
from sglang.multimodal_gen.runtime.layers.visual_embedding import Timesteps
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
@@ -97,9 +98,11 @@ class SanaLinearAttention(nn.Module):
self.num_heads = num_heads
self.head_dim = head_dim
self.to_q = nn.Linear(query_dim, inner_dim, bias=bias)
self.to_k = nn.Linear(query_dim, inner_dim, bias=bias)
self.to_v = nn.Linear(query_dim, inner_dim, bias=bias)
self.inner_dim = inner_dim
# Self-attention q/k/v share the same input -> one packed GEMM.
self.to_qkv = MergedColumnParallelLinear(
query_dim, [inner_dim, inner_dim, inner_dim], bias=bias, gather_output=True
)
self.to_out = nn.ModuleList(
[nn.Linear(inner_dim, query_dim, bias=True), nn.Identity()]
)
@@ -107,9 +110,10 @@ class SanaLinearAttention(nn.Module):
def forward(self, hidden_states):
B, S, _ = hidden_states.shape
query = self.to_q(hidden_states)
key = self.to_k(hidden_states)
value = self.to_v(hidden_states)
qkv, _ = self.to_qkv(hidden_states)
query, key, value = qkv.split(
[self.inner_dim, self.inner_dim, self.inner_dim], dim=-1
)
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)
@@ -136,9 +140,12 @@ class SanaCrossAttention(nn.Module):
self.num_heads = num_heads
self.head_dim = head_dim
self.inner_dim = inner_dim
self.to_q = nn.Linear(query_dim, inner_dim, bias=bias)
self.to_k = nn.Linear(cross_attention_dim, inner_dim, bias=bias)
self.to_v = nn.Linear(cross_attention_dim, inner_dim, bias=bias)
# k/v share the (step-invariant) encoder input -> one packed GEMM.
self.to_kv = MergedColumnParallelLinear(
cross_attention_dim, [inner_dim, inner_dim], bias=bias, gather_output=True
)
self.to_out = nn.ModuleList(
[nn.Linear(inner_dim, query_dim, bias=True), nn.Identity()]
)
@@ -150,8 +157,8 @@ class SanaCrossAttention(nn.Module):
T = encoder_hidden_states.shape[1]
query = self.to_q(hidden_states)
key = self.to_k(encoder_hidden_states)
value = self.to_v(encoder_hidden_states)
kv, _ = self.to_kv(encoder_hidden_states)
key, value = kv.split([self.inner_dim, self.inner_dim], dim=-1)
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)