[kernel slimming] Clean many useless sgl-kernel deprecated kernels (#20277)
This commit is contained in:
@@ -1,4 +1,3 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
@@ -332,121 +331,6 @@ if torch.version.hip is not None:
|
||||
return out
|
||||
|
||||
|
||||
@dataclass
|
||||
class FusedSetKVBufferArg:
|
||||
"""
|
||||
value : Optional[torch.Tensor]
|
||||
Value tensor, shape: ``(nnz, num_v_heads * head_size)``.
|
||||
k_buffer : Optional[torch.Tensor]
|
||||
Buffer for keys, shape: ``(nnz, num_k_heads * head_size)``.
|
||||
v_buffer : Optional[torch.Tensor]
|
||||
Buffer for values, shape: ``(nnz, num_v_heads * head_size)``.
|
||||
k_scale : Optional[float]
|
||||
Scale factor for keys.
|
||||
v_scale : Optional[float]
|
||||
Scale factor for values.
|
||||
cache_loc : Optional[torch.Tensor]
|
||||
Cache location tensor, used for indexing kv cache.
|
||||
"""
|
||||
|
||||
value: torch.Tensor
|
||||
k_buffer: torch.Tensor
|
||||
v_buffer: torch.Tensor
|
||||
k_scale: Optional[float]
|
||||
v_scale: Optional[float]
|
||||
cache_loc: torch.Tensor
|
||||
|
||||
|
||||
def _view_3d(x, head_size):
|
||||
return x.view(x.shape[0], -1, head_size)
|
||||
|
||||
|
||||
def apply_rope_with_cos_sin_cache_inplace(
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
head_size: int,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
is_neox: bool = True,
|
||||
fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None,
|
||||
enable_pdl: Optional[bool] = None,
|
||||
) -> None:
|
||||
r"""
|
||||
Apply rotary embedding to keys and queries with precomputed cos/sin values.
|
||||
This is designed to be compatible with the SGL/vLLM implementation.
|
||||
The result is inplace applied to the input tensors.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
positions : torch.Tensor
|
||||
Position indices, shape: ``(nnz)``.
|
||||
query : torch.Tensor
|
||||
Query tensor, shape: ``(nnz, num_q_heads * head_size)``.
|
||||
key : torch.Tensor
|
||||
Key tensor, shape: ``(nnz, num_k_heads * head_size)``.
|
||||
cos_sin_cache : torch.Tensor
|
||||
Cosine and Sine cache tensor, shape: ``(max_seq_len, rotary_dim)``.
|
||||
Cosine is the first half and Sine is the second half on rotary_dim.
|
||||
is_neox : bool
|
||||
Whether to use Neox style RoPE, default: ``True``.
|
||||
|
||||
* If ``True``, the last dimension of the query/key tensor is not interleaved, i.e.,
|
||||
we rotate the first half dimensions ``([..., :head_dim//2])`` and the second half
|
||||
dimensions ``([..., head_dim//2:])``.
|
||||
|
||||
* If ``False``, the last dimension of the query/key tensor is interleaved, i.e.,
|
||||
we rotate the even dimensions ``([..., ::2])`` and odd dimensions ``([..., 1::2])``.
|
||||
fused_set_kv_buffer_arg : FusedSetKVBufferArg
|
||||
Fuse the set-kv-buffer operation into this kernel
|
||||
|
||||
Note
|
||||
----
|
||||
The rotary dimension is determined by the cosine cache and sine cache.
|
||||
"""
|
||||
if cos_sin_cache.dtype != torch.float32:
|
||||
raise ValueError("cos_sin_cache should be float32")
|
||||
|
||||
if enable_pdl is None:
|
||||
# the non-fused branch does not yet support PDL, but after we switch to our impl for that branch it will
|
||||
enable_pdl = is_arch_support_pdl() and (fused_set_kv_buffer_arg is not None)
|
||||
|
||||
if (a := fused_set_kv_buffer_arg) is not None:
|
||||
assert a.k_scale is None, "k_scale is not yet supported"
|
||||
assert a.v_scale is None, "v_scale is not yet supported"
|
||||
assert a.cache_loc.dtype == torch.int64, f"{a.cache_loc.dtype=}"
|
||||
|
||||
torch.ops.sgl_kernel.apply_rope_pos_ids_cos_sin_cache.default(
|
||||
_view_3d(query, head_size),
|
||||
_view_3d(key, head_size),
|
||||
_view_3d(query, head_size),
|
||||
_view_3d(key, head_size),
|
||||
cos_sin_cache,
|
||||
positions.long(),
|
||||
(not is_neox),
|
||||
enable_pdl,
|
||||
(
|
||||
_view_3d(fused_set_kv_buffer_arg.value, head_size)
|
||||
if fused_set_kv_buffer_arg is not None
|
||||
else None
|
||||
),
|
||||
(
|
||||
_view_3d(fused_set_kv_buffer_arg.k_buffer, head_size)
|
||||
if fused_set_kv_buffer_arg is not None
|
||||
else None
|
||||
),
|
||||
(
|
||||
_view_3d(fused_set_kv_buffer_arg.v_buffer, head_size)
|
||||
if fused_set_kv_buffer_arg is not None
|
||||
else None
|
||||
),
|
||||
(
|
||||
fused_set_kv_buffer_arg.cache_loc
|
||||
if fused_set_kv_buffer_arg is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def rotary_embedding(
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
|
||||
Reference in New Issue
Block a user