From 3a42c5e34190f9a254477a74eaa9f024beb5477c Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Thu, 1 Jan 2026 10:10:36 +0800 Subject: [PATCH] [VLM] Adopt jit qk_norm kernel in VLM (#16171) Co-authored-by: luoyuan.luo --- .../jit_kernel/benchmark/bench_qknorm.py | 9 +++++---- python/sglang/jit_kernel/norm.py | 2 +- python/sglang/srt/layers/attention/vision.py | 20 ++++++++++++++++++- python/sglang/srt/models/utils.py | 2 +- 4 files changed, 26 insertions(+), 7 deletions(-) diff --git a/python/sglang/jit_kernel/benchmark/bench_qknorm.py b/python/sglang/jit_kernel/benchmark/bench_qknorm.py index 046fa02c0..a13613138 100644 --- a/python/sglang/jit_kernel/benchmark/bench_qknorm.py +++ b/python/sglang/jit_kernel/benchmark/bench_qknorm.py @@ -5,6 +5,10 @@ from typing import Tuple import torch import triton import triton.testing +from sgl_kernel import rmsnorm + +from sglang.jit_kernel.norm import fused_inplace_qknorm +from sglang.srt.utils import get_current_device_stream_fast IS_CI = ( os.getenv("CI", "false").lower() == "true" @@ -20,13 +24,12 @@ def sglang_aot_qknorm( q_weight: torch.Tensor, k_weight: torch.Tensor, ) -> None: - from sgl_kernel import rmsnorm head_dim = q.shape[-1] q = q.view(-1, head_dim) k = k.view(-1, head_dim) - current_stream = torch.cuda.current_stream() + current_stream = get_current_device_stream_fast() alt_stream.wait_stream(current_stream) rmsnorm(q, q_weight, out=q) with torch.cuda.stream(alt_stream): @@ -40,7 +43,6 @@ def sglang_jit_qknorm( q_weight: torch.Tensor, k_weight: torch.Tensor, ) -> None: - from sglang.jit_kernel.norm import fused_inplace_qknorm fused_inplace_qknorm(q, k, q_weight, k_weight) @@ -51,7 +53,6 @@ def flashinfer_qknorm( q_weight: torch.Tensor, k_weight: torch.Tensor, ) -> None: - from flashinfer.norm import rmsnorm rmsnorm(q, q_weight, out=q) rmsnorm(k, k_weight, out=k) diff --git a/python/sglang/jit_kernel/norm.py b/python/sglang/jit_kernel/norm.py index b5ed6df47..50f2315f8 100644 --- a/python/sglang/jit_kernel/norm.py +++ b/python/sglang/jit_kernel/norm.py @@ -30,7 +30,7 @@ def _jit_norm_module(head_dims: int) -> Module: @cache_once def can_use_fused_inplace_qknorm(head_dim: int) -> bool: logger = logging.getLogger(__name__) - if head_dim not in [64, 128, 256]: + if head_dim not in [64, 128, 256, 512, 1024]: logger.warning(f"Unsupported head_dim={head_dim} for JIT QK-Norm kernel") return False try: diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 7ad15e7ec..19584fb1f 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -11,8 +11,10 @@ import torch.nn as nn import torch.nn.functional as F from einops import rearrange +from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size +from sglang.srt.models.utils import apply_qk_norm from sglang.srt.utils import ( get_bool_env_var, get_device_capability, @@ -799,7 +801,23 @@ class VisionAttention(nn.Module): # internvl if self.qk_normalization: - q, k = self._apply_qk_norm(q, k) + # jit kernel + if can_use_jit_qk_norm(self.head_size): + + # q: [tokens, head, head_size] -> [tokens, embed_dim] + head_dim_for_norm = head * self.head_size + + q, k = apply_qk_norm( + q=q, + k=k, + q_norm=self.q_norm, + k_norm=self.k_norm, + head_dim=head_dim_for_norm, + alt_stream=self.aux_stream, + ) + + else: + q, k = self._apply_qk_norm(q, k) output = self.qkv_backend.forward( q=q, diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py index d27fd44ee..19bb72442 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -25,6 +25,7 @@ from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_q 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.cuda_graph_runner import get_is_capture_mode from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.utils import is_cuda @@ -223,7 +224,6 @@ def apply_qk_norm( 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