[2/n jit_kernel restruct] unify rotary embedding entrypoints under rope.py (#20247)
This commit is contained in:
@@ -70,7 +70,7 @@ def sglang_pos_enc_rope(
|
|||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
is_neox: bool,
|
is_neox: bool,
|
||||||
) -> None:
|
) -> 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]
|
head_size = q.shape[-1]
|
||||||
rotary_embedding_with_key(
|
rotary_embedding_with_key(
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -17,6 +17,15 @@ if TYPE_CHECKING:
|
|||||||
from tvm_ffi.module import Module
|
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
|
@cache_once
|
||||||
def _jit_fused_rope_module(is_neox: bool, rope_dim: int, dtype: torch.dtype) -> Module:
|
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)
|
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
|
@dataclass
|
||||||
class FusedSetKVBufferArg:
|
class FusedSetKVBufferArg:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.jit_kernel.pos_enc import rotary_embedding
|
from sglang.jit_kernel.rope import rotary_embedding
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
|
|||||||
@@ -71,10 +71,10 @@ class RotaryEmbedding(MultiPlatformOp):
|
|||||||
and not (_is_npu)
|
and not (_is_npu)
|
||||||
and not (_is_musa)
|
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.
|
# TODO: Test on different devices and remove this conditional.
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sglang.jit_kernel.pos_enc import rotary_embedding
|
from sglang.jit_kernel.rope import rotary_embedding
|
||||||
elif _is_hip:
|
elif _is_hip:
|
||||||
from sgl_kernel import rotary_embedding
|
from sgl_kernel import rotary_embedding
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user