From 7da92e11124d6acda336c01c5fd592797d66f5bd Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Fri, 19 Jun 2026 13:25:52 +0800 Subject: [PATCH] [diffusion] Sana: pack self-attn q/k/v and cross-attn k/v into single GEMMs (#28393) Co-authored-by: Claude Opus 4.8 --- .../configs/models/dits/sana.py | 7 +++++ .../runtime/models/dits/sana.py | 27 ++++++++++++------- 2 files changed, 24 insertions(+), 10 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/models/dits/sana.py b/python/sglang/multimodal_gen/configs/models/dits/sana.py index fa2a3f086..9e6b6b9ad 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/sana.py +++ b/python/sglang/multimodal_gen/configs/models/dits/sana.py @@ -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", } ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/sana.py b/python/sglang/multimodal_gen/runtime/models/dits/sana.py index 0a375f8c7..760ab97df 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/sana.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/sana.py @@ -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)