[Feature] JIT Fused QK norm + qk norm clean up (#15835)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"])
|
||||
|
||||
Reference in New Issue
Block a user