[2/n jit_kernel restruct] unify rotary embedding entrypoints under rope.py (#20247)

This commit is contained in:
Xiaoyu Zhang
2026-03-10 17:49:57 +08:00
committed by GitHub
parent 6407891b4f
commit 51d9d34977
5 changed files with 63 additions and 90 deletions

View File

@@ -70,7 +70,7 @@ def sglang_pos_enc_rope(
positions: torch.Tensor,
is_neox: bool,
) -> None:
from sglang.jit_kernel.pos_enc import rotary_embedding_with_key
from sglang.jit_kernel.rope import rotary_embedding_with_key
head_size = q.shape[-1]
rotary_embedding_with_key(

View File

@@ -1,86 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.jit_kernel.utils import cache_once, load_jit
from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING:
from tvm_ffi.module import Module
@cache_once
def _jit_rotary_embedding_module() -> Module:
return load_jit(
"rotary_embedding",
cuda_files=["elementwise/pos_enc.cuh"],
cuda_wrappers=[("rotary_embedding", "RotaryEmbeddingKernel::run")],
)
@register_custom_op(
op_name="rotary_embedding_with_key",
mutates_args=["query", "key"],
)
def rotary_embedding_with_key(
positions: torch.Tensor, # [batch_size, seq_len] or [num_tokens]
query: torch.Tensor, # [batch_size, seq_len, num_heads * head_size] or
# [num_tokens, num_heads * head_size] or
# [batch_size, seq_len, num_heads, head_size] or
# [num_tokens, num_heads, head_size]
key: torch.Tensor, # [batch_size, seq_len, num_kv_heads * head_size] or
# [num_tokens, num_kv_heads * head_size] or
# [batch_size, seq_len, num_heads, head_size] or
# [num_tokens, num_heads, head_size]
head_size: int,
cos_sin_cache: torch.Tensor, # [max_position, rot_dim]
is_neox: bool = True,
) -> None:
"""
Apply rotary embedding to query and key tensors.
Args:
positions: Position indices of shape [num_tokens] or [batch_size, seq_len]
query: Query tensor of shape [num_tokens, num_heads, head_size] or [num_tokens, num_heads * head_size]
key: Key tensor of shape [num_tokens, num_kv_heads, head_size] or [num_tokens, num_kv_heads * head_size]
cos_sin_cache: Cosine and sine cache of shape [max_position, rot_dim]
is_neox: Whether to use GPT-NeoX style rotary embedding (True) or GPT-J style (False)
"""
module = _jit_rotary_embedding_module()
module.rotary_embedding(positions, query, key, head_size, cos_sin_cache, is_neox)
@register_custom_op(
op_name="rotary_embedding_without_key",
mutates_args=["query"],
)
def rotary_embedding_without_key(
positions: torch.Tensor,
query: torch.Tensor,
head_size: int,
cos_sin_cache: torch.Tensor,
is_neox: bool = True,
) -> None:
module = _jit_rotary_embedding_module()
module.rotary_embedding(positions, query, None, head_size, cos_sin_cache, is_neox)
def rotary_embedding(
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
head_size: int,
cos_sin_cache: torch.Tensor,
is_neox: bool = True,
):
if key is None:
rotary_embedding_without_key(
positions, query, head_size, cos_sin_cache, is_neox
)
else:
rotary_embedding_with_key(
positions, query, key, head_size, cos_sin_cache, is_neox
)
return query, key

View File

@@ -17,6 +17,15 @@ if TYPE_CHECKING:
from tvm_ffi.module import Module
@cache_once
def _jit_rotary_embedding_module() -> Module:
return load_jit(
"rotary_embedding",
cuda_files=["elementwise/pos_enc.cuh"],
cuda_wrappers=[("rotary_embedding", "RotaryEmbeddingKernel::run")],
)
@cache_once
def _jit_fused_rope_module(is_neox: bool, rope_dim: int, dtype: torch.dtype) -> Module:
args = make_cpp_args(is_neox, rope_dim, is_arch_support_pdl(), dtype)
@@ -31,6 +40,56 @@ def _jit_fused_rope_module(is_neox: bool, rope_dim: int, dtype: torch.dtype) ->
)
@register_custom_op(
op_name="rotary_embedding_with_key",
mutates_args=["query", "key"],
)
def rotary_embedding_with_key(
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
head_size: int,
cos_sin_cache: torch.Tensor,
is_neox: bool = True,
) -> None:
module = _jit_rotary_embedding_module()
module.rotary_embedding(positions, query, key, head_size, cos_sin_cache, is_neox)
@register_custom_op(
op_name="rotary_embedding_without_key",
mutates_args=["query"],
)
def rotary_embedding_without_key(
positions: torch.Tensor,
query: torch.Tensor,
head_size: int,
cos_sin_cache: torch.Tensor,
is_neox: bool = True,
) -> None:
module = _jit_rotary_embedding_module()
module.rotary_embedding(positions, query, None, head_size, cos_sin_cache, is_neox)
def rotary_embedding(
positions: torch.Tensor,
query: torch.Tensor,
key: Optional[torch.Tensor],
head_size: int,
cos_sin_cache: torch.Tensor,
is_neox: bool = True,
):
if key is None:
rotary_embedding_without_key(
positions, query, head_size, cos_sin_cache, is_neox
)
else:
rotary_embedding_with_key(
positions, query, key, head_size, cos_sin_cache, is_neox
)
return query, key
@dataclass
class FusedSetKVBufferArg:
"""

View File

@@ -6,7 +6,7 @@ import torch
import triton
import triton.language as tl
from sglang.jit_kernel.pos_enc import rotary_embedding
from sglang.jit_kernel.rope import rotary_embedding
@triton.jit

View File

@@ -71,10 +71,10 @@ class RotaryEmbedding(MultiPlatformOp):
and not (_is_npu)
and not (_is_musa)
):
# rotary_embedding from sglang.jit_kernel.pos_enc and vllm._custom_ops has the same implementation.
# rotary_embedding from sglang.jit_kernel.rope and vllm._custom_ops has the same implementation.
# TODO: Test on different devices and remove this conditional.
if _is_cuda:
from sglang.jit_kernel.pos_enc import rotary_embedding
from sglang.jit_kernel.rope import rotary_embedding
elif _is_hip:
from sgl_kernel import rotary_embedding
else: