From 53240270075c2f0651a9909980c174da342e4a5f Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Thu, 22 Jan 2026 22:32:15 +0800 Subject: [PATCH] [Diffusion] Make the apply_qknorm function easier to use (#17537) --- .../runtime/layers/layernorm.py | 7 --- .../runtime/models/dits/flux.py | 53 ++++++------------- .../runtime/models/dits/flux_2.py | 53 ++++++------------- .../runtime/models/dits/zimage.py | 25 +++------ 4 files changed, 40 insertions(+), 98 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py index 86622ad17..6b42cafd1 100644 --- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py +++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py @@ -471,13 +471,6 @@ def apply_qk_norm( ) return q, k - # Fallback for AMD/ROCm: apply RMSNorm separately to q and k - import warnings - - warnings.warn( - "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) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index e6891cd6f..d30c083c0 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -27,7 +27,6 @@ from diffusers.models.normalization import ( ) from torch.nn import LayerNorm as LayerNorm -from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.layers.attention import USPAttention @@ -49,7 +48,6 @@ from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiT from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name -_is_cuda = current_platform.is_cuda() def _get_qkv_projections( @@ -167,47 +165,28 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin): query = query.unflatten(-1, (self.heads, -1)) key = key.unflatten(-1, (self.heads, -1)) value = value.unflatten(-1, (self.heads, -1)) - if ( - _is_cuda - and (self.norm_q.variance_epsilon == self.norm_k.variance_epsilon) - and can_use_fused_inplace_qknorm(self.head_dim, query.dtype) - ): - query, key = apply_qk_norm( - q=query, - k=key, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - else: - query = self.norm_q(query) - key = self.norm_k(key) + query, key = apply_qk_norm( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + allow_inplace=True, + ) if self.added_kv_proj_dim is not None: encoder_query = encoder_query.unflatten(-1, (self.heads, -1)) encoder_key = encoder_key.unflatten(-1, (self.heads, -1)) encoder_value = encoder_value.unflatten(-1, (self.heads, -1)) - if ( - _is_cuda - and ( - self.norm_added_q.variance_epsilon - == self.norm_added_k.variance_epsilon - ) - and can_use_fused_inplace_qknorm(self.head_dim, encoder_query.dtype) - ): - encoder_query, encoder_key = apply_qk_norm( - q=encoder_query, - k=encoder_key, - q_norm=self.norm_added_q, - k_norm=self.norm_added_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - else: - encoder_query = self.norm_added_q(encoder_query) - encoder_key = self.norm_added_k(encoder_key) + encoder_query, encoder_key = apply_qk_norm( + q=encoder_query, + k=encoder_key, + q_norm=self.norm_added_q, + k_norm=self.norm_added_k, + head_dim=self.head_dim, + allow_inplace=True, + ) bsz, seq_len, _, _ = query.shape query = torch.cat([encoder_query, query], dim=1) 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 907c39888..a5be6184d 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -20,7 +20,6 @@ from diffusers.models.attention import AttentionModuleMixin from diffusers.models.embeddings import TimestepEmbedding, Timesteps from diffusers.models.normalization import AdaLayerNormContinuous -from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm 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 @@ -35,7 +34,6 @@ from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiT from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name -_is_cuda = current_platform.is_cuda() def _get_qkv_projections( @@ -198,47 +196,28 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): key = key.unflatten(-1, (self.heads, -1)) value = value.unflatten(-1, (self.heads, -1)) - if ( - _is_cuda - and (self.norm_q.variance_epsilon == self.norm_k.variance_epsilon) - and can_use_fused_inplace_qknorm(self.head_dim, query.dtype) - ): - query, key = apply_qk_norm( - q=query, - k=key, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - else: - query = self.norm_q(query) - key = self.norm_k(key) + query, key = apply_qk_norm( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + allow_inplace=True, + ) if self.added_kv_proj_dim is not None: encoder_query = encoder_query.unflatten(-1, (self.heads, -1)) encoder_key = encoder_key.unflatten(-1, (self.heads, -1)) encoder_value = encoder_value.unflatten(-1, (self.heads, -1)) - if ( - _is_cuda - and ( - self.norm_added_q.variance_epsilon - == self.norm_added_k.variance_epsilon - ) - and can_use_fused_inplace_qknorm(self.head_dim, encoder_query.dtype) - ): - encoder_query, encoder_key = apply_qk_norm( - q=encoder_query, - k=encoder_key, - q_norm=self.norm_added_q, - k_norm=self.norm_added_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - else: - encoder_query = self.norm_added_q(encoder_query) - encoder_key = self.norm_added_k(encoder_key) + encoder_query, encoder_key = apply_qk_norm( + q=encoder_query, + k=encoder_key, + q_norm=self.norm_added_q, + k_norm=self.norm_added_k, + head_dim=self.head_dim, + allow_inplace=True, + ) query = torch.cat([encoder_query, query], dim=1) key = torch.cat([encoder_key, key], dim=1) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index 8bdac4807..c3d71c8a1 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -4,7 +4,6 @@ from typing import Any, List, Optional, Tuple import torch import torch.nn as nn -from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig from sglang.multimodal_gen.runtime.distributed import get_tp_world_size from sglang.multimodal_gen.runtime.layers.activation import SiluAndMul @@ -171,22 +170,14 @@ class ZImageAttention(nn.Module): v = v.view(*v.shape[:-1], self.local_num_kv_heads, self.head_dim) if self.qk_norm: - if ( - _is_cuda - and (self.norm_q.variance_epsilon == self.norm_k.variance_epsilon) - and can_use_fused_inplace_qknorm(self.head_dim, q.dtype) - ): - q, k = apply_qk_norm( - q=q, - k=k, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - else: - q = self.norm_q(q) - k = self.norm_k(k) + q, k = apply_qk_norm( + q=q, + k=k, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + allow_inplace=True, + ) if freqs_cis is not None: cos, sin = freqs_cis