[VLM] Adopt jit qk_norm kernel in VLM (#16171)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-01-01 10:10:36 +08:00
committed by GitHub
parent 12b89e51d8
commit 3a42c5e341
4 changed files with 26 additions and 7 deletions

View File

@@ -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)

View File

@@ -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:

View File

@@ -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,

View File

@@ -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