[VLM] Adopt jit qk_norm kernel in VLM (#16171)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -5,6 +5,10 @@ from typing import Tuple
|
|||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
import triton.testing
|
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 = (
|
IS_CI = (
|
||||||
os.getenv("CI", "false").lower() == "true"
|
os.getenv("CI", "false").lower() == "true"
|
||||||
@@ -20,13 +24,12 @@ def sglang_aot_qknorm(
|
|||||||
q_weight: torch.Tensor,
|
q_weight: torch.Tensor,
|
||||||
k_weight: torch.Tensor,
|
k_weight: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
from sgl_kernel import rmsnorm
|
|
||||||
|
|
||||||
head_dim = q.shape[-1]
|
head_dim = q.shape[-1]
|
||||||
q = q.view(-1, head_dim)
|
q = q.view(-1, head_dim)
|
||||||
k = k.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)
|
alt_stream.wait_stream(current_stream)
|
||||||
rmsnorm(q, q_weight, out=q)
|
rmsnorm(q, q_weight, out=q)
|
||||||
with torch.cuda.stream(alt_stream):
|
with torch.cuda.stream(alt_stream):
|
||||||
@@ -40,7 +43,6 @@ def sglang_jit_qknorm(
|
|||||||
q_weight: torch.Tensor,
|
q_weight: torch.Tensor,
|
||||||
k_weight: torch.Tensor,
|
k_weight: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
from sglang.jit_kernel.norm import fused_inplace_qknorm
|
|
||||||
|
|
||||||
fused_inplace_qknorm(q, k, q_weight, k_weight)
|
fused_inplace_qknorm(q, k, q_weight, k_weight)
|
||||||
|
|
||||||
@@ -51,7 +53,6 @@ def flashinfer_qknorm(
|
|||||||
q_weight: torch.Tensor,
|
q_weight: torch.Tensor,
|
||||||
k_weight: torch.Tensor,
|
k_weight: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
from flashinfer.norm import rmsnorm
|
|
||||||
|
|
||||||
rmsnorm(q, q_weight, out=q)
|
rmsnorm(q, q_weight, out=q)
|
||||||
rmsnorm(k, k_weight, out=k)
|
rmsnorm(k, k_weight, out=k)
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ def _jit_norm_module(head_dims: int) -> Module:
|
|||||||
@cache_once
|
@cache_once
|
||||||
def can_use_fused_inplace_qknorm(head_dim: int) -> bool:
|
def can_use_fused_inplace_qknorm(head_dim: int) -> bool:
|
||||||
logger = logging.getLogger(__name__)
|
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")
|
logger.warning(f"Unsupported head_dim={head_dim} for JIT QK-Norm kernel")
|
||||||
return False
|
return False
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -11,8 +11,10 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from einops import rearrange
|
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.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size
|
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 (
|
from sglang.srt.utils import (
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
get_device_capability,
|
get_device_capability,
|
||||||
@@ -799,6 +801,22 @@ class VisionAttention(nn.Module):
|
|||||||
|
|
||||||
# internvl
|
# internvl
|
||||||
if self.qk_normalization:
|
if self.qk_normalization:
|
||||||
|
# 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)
|
q, k = self._apply_qk_norm(q, k)
|
||||||
|
|
||||||
output = self.qkv_backend.forward(
|
output = self.qkv_backend.forward(
|
||||||
|
|||||||
@@ -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.jit_kernel.utils import register_jit_op
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
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.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.utils import is_cuda
|
from sglang.srt.utils import is_cuda
|
||||||
|
|
||||||
@@ -223,7 +224,6 @@ def apply_qk_norm(
|
|||||||
Returns:
|
Returns:
|
||||||
Tuple of normalized query and key tensors
|
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)
|
batch_size = q.size(0)
|
||||||
q_eps = q_norm.variance_epsilon
|
q_eps = q_norm.variance_epsilon
|
||||||
|
|||||||
Reference in New Issue
Block a user