From a3d9a218826a1b70e66e8c7908f22d0c3e488044 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Mon, 19 Jan 2026 15:24:11 +0800 Subject: [PATCH] Revert "[Perf] fuse q, k norm for Flux2Attention (#17241)" (#17332) --- .../multimodal_gen/runtime/layers/layernorm.py | 14 +++----------- .../multimodal_gen/runtime/models/dits/flux_2.py | 8 +++++--- 2 files changed, 8 insertions(+), 14 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py index 4bd9368ee..a9f5811e2 100644 --- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py +++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py @@ -443,15 +443,6 @@ def apply_qk_norm( """ batch_size = q.size(0) - q_shape = q.shape - k_shape = k.shape - - # Ensure contiguous for view operation - if not q.is_contiguous(): - q = q.contiguous() - if not k.is_contiguous(): - k = k.contiguous() - q_eps = q_norm.variance_epsilon k_eps = k_norm.variance_epsilon # Only try fused path on CUDA and when it won't introduce implicit copies. @@ -469,8 +460,7 @@ def apply_qk_norm( head_dim=head_dim, eps=q_eps, ) - # Inplace kernel modifies q, k - return with original shape - return q.view(q_shape), k.view(k_shape) + return q, k # Fallback for AMD/ROCm: apply RMSNorm separately to q and k import warnings @@ -479,6 +469,8 @@ def apply_qk_norm( "Fused QK-norm not available, using RMSNorm fallback", stacklevel=2, ) + q_shape = q.shape + k_shape = k.shape q_out = q_norm(q.view(-1, head_dim)).view(q_shape) k_out = k_norm(k.view(-1, head_dim)).view(k_shape) return q_out, k_out diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index 585c30cb5..a157b88d3 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -22,7 +22,7 @@ from diffusers.models.normalization import AdaLayerNormContinuous from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.layers.attention import USPAttention -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( NDRotaryEmbedding, @@ -196,7 +196,8 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): key = key.unflatten(-1, (self.heads, -1)) value = value.unflatten(-1, (self.heads, -1)) - query, key = apply_qk_norm(query, key, self.norm_q, self.norm_k, self.head_dim) + query = self.norm_q(query) + key = self.norm_k(key) if self.added_kv_proj_dim is not None: encoder_query = encoder_query.unflatten(-1, (self.heads, -1)) @@ -339,7 +340,8 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin): key = key.unflatten(-1, (self.heads, -1)) value = value.unflatten(-1, (self.heads, -1)) - query, key = apply_qk_norm(query, key, self.norm_q, self.norm_k, self.head_dim) + query = self.norm_q(query) + key = self.norm_k(key) if freqs_cis is not None: cos, sin = freqs_cis