[Feature] JIT Fused QK norm + qk norm clean up (#15835)

This commit is contained in:
DarkSharpness
2025-12-28 11:53:50 +08:00
committed by GitHub
parent 474a4699c5
commit 8e43980ebb
15 changed files with 827 additions and 127 deletions
+9 -23
View File
@@ -75,6 +75,7 @@ from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import (
apply_qk_norm,
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
@@ -507,28 +508,6 @@ class BailingMoEAttention(nn.Module):
self.alt_stream = alt_stream
def _apply_qk_norm(
self, q: torch.Tensor, k: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# overlap qk norm
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.query_layernorm(q_by_head)
with torch.cuda.stream(self.alt_stream):
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.key_layernorm(k_by_head)
current_stream.wait_stream(self.alt_stream)
else:
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.query_layernorm(q_by_head)
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.key_layernorm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
def forward(
self,
positions: torch.Tensor,
@@ -540,7 +519,14 @@ class BailingMoEAttention(nn.Module):
qkv, _ = self.query_key_value(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
if self.use_qk_norm:
q, k = self._apply_qk_norm(q, k)
q, k = apply_qk_norm(
q=q,
k=k,
q_norm=self.query_layernorm,
k_norm=self.key_layernorm,
head_dim=self.head_dim,
alt_stream=self.alt_stream,
)
q, k = self.rotary_emb(
positions,
q,
+9 -23
View File
@@ -75,6 +75,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
@@ -250,28 +251,6 @@ class Glm4MoeAttention(nn.Module):
self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
self.alt_stream = alt_stream
def _apply_qk_norm(
self, q: torch.Tensor, k: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# overlap qk norm
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.q_norm(q_by_head)
with torch.cuda.stream(self.alt_stream):
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.k_norm(k_by_head)
current_stream.wait_stream(self.alt_stream)
else:
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.q_norm(q_by_head)
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.k_norm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
def op_prepare(self, state):
state.attn_intermediate_state = self.forward_prepare(
positions=state.positions,
@@ -295,7 +274,14 @@ class Glm4MoeAttention(nn.Module):
qkv, _ = self.qkv_proj(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
if self.use_qk_norm:
q, k = self._apply_qk_norm(q, k)
q, k = apply_qk_norm(
q=q,
k=k,
q_norm=self.q_norm,
k_norm=self.k_norm,
head_dim=self.head_dim,
alt_stream=self.alt_stream,
)
q, k = self.rotary_emb(positions, q, k)
inner_state = q, k, v, forward_batch
return None, forward_batch, inner_state
+9 -23
View File
@@ -71,6 +71,7 @@ from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import (
apply_qk_norm,
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
@@ -492,28 +493,6 @@ class LLaDA2MoeAttention(nn.Module):
self.alt_stream = alt_stream
def _apply_qk_norm(
self, q: torch.Tensor, k: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# overlap qk norm
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.query_layernorm(q_by_head)
with torch.cuda.stream(self.alt_stream):
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.key_layernorm(k_by_head)
current_stream.wait_stream(self.alt_stream)
else:
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.query_layernorm(q_by_head)
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.key_layernorm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
def forward(
self,
positions: torch.Tensor,
@@ -525,7 +504,14 @@ class LLaDA2MoeAttention(nn.Module):
qkv, _ = self.query_key_value(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
if self.use_qk_norm:
q, k = self._apply_qk_norm(q, k)
q, k = apply_qk_norm(
q=q,
k=k,
q_norm=self.query_layernorm,
k_norm=self.key_layernorm,
head_dim=self.head_dim,
alt_stream=self.alt_stream,
)
q, k = self.rotary_emb(
positions,
q,
+9 -24
View File
@@ -21,7 +21,6 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import get_rope
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import (
default_weight_loader,
@@ -29,6 +28,7 @@ from sglang.srt.model_loader.weight_utils import (
)
from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP
from sglang.srt.models.qwen2 import Qwen2Model
from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_npu
@@ -138,32 +138,17 @@ class Qwen3Attention(nn.Module):
)
self.alt_stream = alt_stream
def _apply_qk_norm(
self, q: torch.Tensor, k: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# overlap qk norm
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.q_norm(q_by_head)
with torch.cuda.stream(self.alt_stream):
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.k_norm(k_by_head)
current_stream.wait_stream(self.alt_stream)
else:
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.q_norm(q_by_head)
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.k_norm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
def forward_prepare_native(self, positions, hidden_states):
qkv, _ = self.qkv_proj(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q, k = self._apply_qk_norm(q, k)
q, k = apply_qk_norm(
q=q,
k=k,
q_norm=self.q_norm,
k_norm=self.k_norm,
head_dim=self.head_dim,
alt_stream=self.alt_stream,
)
q, k = self.rotary_emb(positions, q, k)
return q, k, v
+9 -27
View File
@@ -57,12 +57,12 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import MRotaryEmbedding, get_rope
from sglang.srt.layers.utils import get_layer_id
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP as Qwen3MoeMLP
from sglang.srt.models.qwen2_moe import Qwen2MoeModel
from sglang.srt.models.utils import (
apply_qk_norm,
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
@@ -498,31 +498,6 @@ class Qwen3MoeAttention(nn.Module):
self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
self.alt_stream = alt_stream
def _apply_qk_norm(
self, q: torch.Tensor, k: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# overlap qk norm
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.q_norm(q_by_head)
with torch.cuda.stream(self.alt_stream):
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.k_norm(k_by_head)
current_stream.wait_stream(self.alt_stream)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
else:
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.q_norm(q_by_head)
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.k_norm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
def op_prepare(self, state):
state.attn_intermediate_state = self.forward_prepare(
positions=state.positions,
@@ -604,7 +579,14 @@ class Qwen3MoeAttention(nn.Module):
else:
# Fallback to non-fused QK Norm & RoPE implementation
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q, k = self._apply_qk_norm(q, k)
q, k = apply_qk_norm(
q=q,
k=k,
q_norm=self.q_norm,
k_norm=self.k_norm,
head_dim=self.head_dim,
alt_stream=self.alt_stream,
)
q, k = self.rotary_emb(
positions,
q,
+80 -5
View File
@@ -11,25 +11,28 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
from __future__ import annotations
from collections.abc import Iterable, Mapping
from dataclasses import dataclass, field
from functools import lru_cache
from typing import Any, Optional
from typing import TYPE_CHECKING, Any, Optional, Tuple
import numpy as np
import torch
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
from sglang.jit_kernel.utils import register_jit_op
from sglang.srt.environ import envs
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.utils import is_cuda
if TYPE_CHECKING:
from sglang.srt.layers.layernorm import RMSNorm
_is_cuda = is_cuda()
if _is_cuda:
from sgl_kernel import FusedSetKVBufferArg
WeightsMapping = Mapping[str, Optional[str]]
"""If a key maps to a value of `None`, the corresponding weight is ignored."""
@@ -113,6 +116,8 @@ def create_fused_set_kv_buffer_arg(
layer: RadixAttention,
forward_batch: ForwardBatch,
):
from sgl_kernel import FusedSetKVBufferArg
layer_id = layer.layer_id
token_to_kv_pool = forward_batch.token_to_kv_pool
@@ -191,3 +196,73 @@ class RotaryPosMixin:
wpos_ids = wpos_ids.flatten()
return torch.from_numpy(np.stack([hpos_ids, wpos_ids], axis=-1))
def apply_qk_norm(
q: torch.Tensor,
k: torch.Tensor,
q_norm: RMSNorm,
k_norm: RMSNorm,
head_dim: int,
alt_stream: Optional[torch.cuda.Stream] = None,
allow_inplace: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply QK normalization for query and key tensors.
If eligible, we will use JIT fused inplace QK normalization for better performance.
Args:
q: Query tensor of shape [batch_size, ...]
k: Key tensor of shape [batch_size, ...]
q_norm: RMSNorm layer for query normalization
k_norm: RMSNorm layer for key normalization
head_dim: Dimension of each attention head
alt_stream: Optional alternative CUDA stream for overlapping computation
allow_inplace: Whether to allow inplace normalization. (True for better performance)
Returns:
Tuple of normalized query and key tensors
"""
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
batch_size = q.size(0)
q_eps = q_norm.variance_epsilon
k_eps = k_norm.variance_epsilon
if (
_is_cuda # TODO(dark): have not tested on ROCm or other backends
and allow_inplace # TODO(dark): this can be relaxed if needed
and (q_eps == k_eps) # TODO(dark): this can also be relaxed
and not envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get()
and can_use_fused_inplace_qknorm(head_dim)
):
fused_inplace_qknorm(
q=q.view(batch_size, -1, head_dim),
k=k.view(batch_size, -1, head_dim),
q_weight=q_norm.weight,
k_weight=k_norm.weight,
head_dim=head_dim,
eps=q_eps,
)
return q, k
if alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, head_dim)
q_by_head = q_norm(q_by_head)
with torch.cuda.stream(alt_stream):
k_by_head = k.reshape(-1, head_dim)
k_by_head = k_norm(k_by_head)
current_stream.wait_stream(alt_stream)
else:
q_by_head = q.reshape(-1, head_dim)
q_by_head = q_norm(q_by_head)
k_by_head = k.reshape(-1, head_dim)
k_by_head = k_norm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
# Register the inplace op
fused_inplace_qknorm = register_jit_op(fused_inplace_qknorm, out_args=["q", "k"])