[codex] Centralize Triton utility kernels (#26000)
Co-authored-by: BBuf <xiaoyu.zhang@radixark.net>
This commit is contained in:
@@ -5,11 +5,13 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import AttentionArch
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
|
from sglang.srt.layers.attention.triton_ops.metadata import (
|
||||||
|
normal_decode_set_metadata,
|
||||||
|
prepare_swa_spec_page_table_triton,
|
||||||
|
)
|
||||||
from sglang.srt.layers.radix_attention import AttentionType
|
from sglang.srt.layers.radix_attention import AttentionType
|
||||||
from sglang.srt.layers.utils.cp_utils import (
|
from sglang.srt.layers.utils.cp_utils import (
|
||||||
cp_allgather_and_save_kv_cache,
|
cp_allgather_and_save_kv_cache,
|
||||||
@@ -1256,7 +1258,6 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
cache_seqlens = metadata.cache_seqlens_int32
|
cache_seqlens = metadata.cache_seqlens_int32
|
||||||
cu_seqlens_k = metadata.cu_seqlens_k
|
|
||||||
max_seqlen_q = metadata.max_seq_len_q
|
max_seqlen_q = metadata.max_seq_len_q
|
||||||
q_reshaped = q.contiguous().view(
|
q_reshaped = q.contiguous().view(
|
||||||
-1, layer.tp_q_head_num, layer.head_dim
|
-1, layer.tp_q_head_num, layer.head_dim
|
||||||
@@ -1991,7 +1992,6 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata_expand = None
|
metadata_expand = None
|
||||||
|
|
||||||
if forward_mode.is_decode_or_idle():
|
if forward_mode.is_decode_or_idle():
|
||||||
|
|
||||||
if spec_info is not None:
|
if spec_info is not None:
|
||||||
# Draft Decode
|
# Draft Decode
|
||||||
if self.topk <= 1:
|
if self.topk <= 1:
|
||||||
@@ -2597,110 +2597,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.swa_spec_metadata = metadata_swa
|
metadata.swa_spec_metadata = metadata_swa
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _prepare_swa_spec_page_table_kernel(
|
|
||||||
dst_ptr,
|
|
||||||
src_a_ptr,
|
|
||||||
src_b_ptr,
|
|
||||||
seq_len_a_ptr,
|
|
||||||
seq_len_b_ptr,
|
|
||||||
dst_stride_m,
|
|
||||||
dst_stride_n,
|
|
||||||
a_stride_m,
|
|
||||||
a_stride_n,
|
|
||||||
b_stride_m,
|
|
||||||
b_stride_n,
|
|
||||||
LEN_A: tl.constexpr,
|
|
||||||
LEN_B: tl.constexpr,
|
|
||||||
REPEAT_STEP: tl.constexpr,
|
|
||||||
BLOCK_N: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid_m = tl.program_id(0)
|
|
||||||
pid_n = tl.program_id(1)
|
|
||||||
|
|
||||||
idx_a = pid_m // REPEAT_STEP
|
|
||||||
idx_b = pid_m
|
|
||||||
seq_len_a = tl.load(seq_len_a_ptr + idx_a)
|
|
||||||
seq_len_b = tl.load(seq_len_b_ptr + idx_b)
|
|
||||||
|
|
||||||
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
|
|
||||||
total_len = seq_len_a + seq_len_b
|
|
||||||
|
|
||||||
if pid_n * BLOCK_N >= total_len:
|
|
||||||
return
|
|
||||||
|
|
||||||
mask = offs_n < total_len
|
|
||||||
dst = dst_ptr + pid_m * dst_stride_m + offs_n * dst_stride_n
|
|
||||||
|
|
||||||
if (pid_n + 1) * BLOCK_N < seq_len_a:
|
|
||||||
a_ptr = src_a_ptr + idx_a * a_stride_m + offs_n * a_stride_n
|
|
||||||
a_mask = mask & (offs_n < LEN_A)
|
|
||||||
val = tl.load(a_ptr, mask=a_mask, other=0)
|
|
||||||
tl.store(dst, val, mask=mask)
|
|
||||||
elif pid_n * BLOCK_N >= seq_len_a:
|
|
||||||
offs_b = offs_n - seq_len_a
|
|
||||||
b_ptr = src_b_ptr + idx_b * b_stride_m + offs_b * b_stride_n
|
|
||||||
b_mask = mask & (offs_b < LEN_B)
|
|
||||||
val = tl.load(b_ptr, mask=b_mask, other=0)
|
|
||||||
tl.store(dst, val, mask=mask)
|
|
||||||
else:
|
|
||||||
# mixed part
|
|
||||||
a_offs = offs_n
|
|
||||||
a_mask = (a_offs < seq_len_a) & (a_offs < LEN_A)
|
|
||||||
a_ptr = src_a_ptr + idx_a * a_stride_m + a_offs * a_stride_n
|
|
||||||
a_val = tl.load(a_ptr, mask=a_mask, other=0)
|
|
||||||
|
|
||||||
b_offs = offs_n - seq_len_a
|
|
||||||
b_mask = (b_offs >= 0) & (b_offs < seq_len_b) & (b_offs < LEN_B)
|
|
||||||
b_ptr = src_b_ptr + idx_b * b_stride_m + b_offs * b_stride_n
|
|
||||||
b_val = tl.load(b_ptr, mask=b_mask, other=0)
|
|
||||||
|
|
||||||
result = tl.where(offs_n < seq_len_a, a_val, b_val)
|
|
||||||
tl.store(dst, result, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
def prepare_swa_spec_page_table_triton(
|
|
||||||
page_table_dst: torch.Tensor,
|
|
||||||
page_table_a: torch.Tensor,
|
|
||||||
page_table_b: torch.Tensor, # expand page table
|
|
||||||
seq_len_a: torch.Tensor,
|
|
||||||
seq_len_b: torch.Tensor, # expand seq lens
|
|
||||||
speculative_num_draft_tokens: int,
|
|
||||||
):
|
|
||||||
# concat page_table and expand page_table by kv seq length
|
|
||||||
bs = seq_len_a.numel()
|
|
||||||
bs_expand = seq_len_b.numel()
|
|
||||||
assert bs_expand == bs * speculative_num_draft_tokens
|
|
||||||
|
|
||||||
LEN_A = page_table_a.shape[1]
|
|
||||||
LEN_B = page_table_b.shape[1]
|
|
||||||
LEN_OUT = LEN_A + LEN_B
|
|
||||||
REPEAT_STEP = speculative_num_draft_tokens
|
|
||||||
BLOCK_N = 256
|
|
||||||
|
|
||||||
grid = (bs_expand, triton.cdiv(LEN_OUT, BLOCK_N))
|
|
||||||
_prepare_swa_spec_page_table_kernel[grid](
|
|
||||||
page_table_dst,
|
|
||||||
page_table_a,
|
|
||||||
page_table_b,
|
|
||||||
seq_len_a,
|
|
||||||
seq_len_b,
|
|
||||||
page_table_dst.stride(0),
|
|
||||||
page_table_dst.stride(1),
|
|
||||||
page_table_a.stride(0),
|
|
||||||
page_table_a.stride(1),
|
|
||||||
page_table_b.stride(0),
|
|
||||||
page_table_b.stride(1),
|
|
||||||
LEN_A=LEN_A,
|
|
||||||
LEN_B=LEN_B,
|
|
||||||
REPEAT_STEP=REPEAT_STEP,
|
|
||||||
BLOCK_N=BLOCK_N,
|
|
||||||
num_warps=4,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class FlashAttentionMultiStepBackend:
|
class FlashAttentionMultiStepBackend:
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model_runner: ModelRunner,
|
model_runner: ModelRunner,
|
||||||
@@ -2771,310 +2668,6 @@ class FlashAttentionMultiStepBackend:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _fused_metadata_kernel_general(
|
|
||||||
# Input tensors
|
|
||||||
seq_lens,
|
|
||||||
seq_lens_stride_0,
|
|
||||||
req_to_token,
|
|
||||||
req_to_token_stride_0,
|
|
||||||
req_to_token_stride_1,
|
|
||||||
req_pool_indices,
|
|
||||||
req_pool_indices_stride_0,
|
|
||||||
# Output buffers
|
|
||||||
cache_seqlens_int32,
|
|
||||||
cache_seqlens_int32_stride_0,
|
|
||||||
cu_seqlens_k,
|
|
||||||
cu_seqlens_k_stride_0,
|
|
||||||
page_table,
|
|
||||||
page_table_stride_0,
|
|
||||||
page_table_stride_1,
|
|
||||||
swa_page_table,
|
|
||||||
swa_page_table_stride_0,
|
|
||||||
swa_page_table_stride_1,
|
|
||||||
full_to_swa_mapping,
|
|
||||||
full_to_swa_mapping_stride_0,
|
|
||||||
# Scalar parameters
|
|
||||||
B,
|
|
||||||
max_seq_pages,
|
|
||||||
page_size: tl.constexpr,
|
|
||||||
seq_len_delta: tl.constexpr,
|
|
||||||
use_swa: tl.constexpr,
|
|
||||||
SHIFT: tl.constexpr,
|
|
||||||
BLOCK_COLS: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid_b = tl.program_id(0) # batch index
|
|
||||||
pid_c = tl.program_id(1) # column chunk index
|
|
||||||
|
|
||||||
# 1. Prefix sum (only one block does it)
|
|
||||||
if pid_b == 0 and pid_c == 0:
|
|
||||||
acc = 0
|
|
||||||
for idx in range(B):
|
|
||||||
seq = tl.load(seq_lens + idx * seq_lens_stride_0)
|
|
||||||
val = (seq + seq_len_delta).to(tl.int32)
|
|
||||||
tl.store(cache_seqlens_int32 + idx * cache_seqlens_int32_stride_0, val)
|
|
||||||
tl.store(cu_seqlens_k + idx * cu_seqlens_k_stride_0, acc)
|
|
||||||
acc += val
|
|
||||||
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
|
|
||||||
|
|
||||||
# 2. Gather for this batch and column chunk
|
|
||||||
if max_seq_pages == 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
i = pid_b
|
|
||||||
# Load row index for this batch (all threads in block have same i)
|
|
||||||
row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0)
|
|
||||||
row_offset = row_idx * req_to_token_stride_0
|
|
||||||
|
|
||||||
col_start = pid_c * BLOCK_COLS
|
|
||||||
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
|
|
||||||
mask = col_offsets < max_seq_pages
|
|
||||||
|
|
||||||
# Compute column indices in the source tensor (token offset)
|
|
||||||
if page_size == 1:
|
|
||||||
col_idx = col_offsets
|
|
||||||
else:
|
|
||||||
col_idx = col_offsets << SHIFT # faster than multiplication for power-of-two
|
|
||||||
|
|
||||||
# Load page indices from req_to_token
|
|
||||||
rt_offsets = row_offset + col_idx * req_to_token_stride_1
|
|
||||||
page_index = tl.load(
|
|
||||||
req_to_token + rt_offsets, mask=mask, other=0, cache_modifier=".cg"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Compute page_table
|
|
||||||
if page_size == 1:
|
|
||||||
page_table_val = page_index
|
|
||||||
else:
|
|
||||||
page_table_val = page_index >> SHIFT
|
|
||||||
|
|
||||||
# Store to page_table
|
|
||||||
pt_offsets = i * page_table_stride_0 + col_offsets * page_table_stride_1
|
|
||||||
tl.store(page_table + pt_offsets, page_table_val, mask=mask, cache_modifier=".cg")
|
|
||||||
|
|
||||||
if use_swa:
|
|
||||||
swa_slot = tl.load(
|
|
||||||
full_to_swa_mapping + page_index * full_to_swa_mapping_stride_0,
|
|
||||||
mask=mask,
|
|
||||||
other=0,
|
|
||||||
cache_modifier=".cg",
|
|
||||||
)
|
|
||||||
if page_size == 1:
|
|
||||||
swa_val = swa_slot
|
|
||||||
else:
|
|
||||||
swa_val = swa_slot >> SHIFT
|
|
||||||
swa_offsets = (
|
|
||||||
i * swa_page_table_stride_0 + col_offsets * swa_page_table_stride_1
|
|
||||||
)
|
|
||||||
tl.store(swa_page_table + swa_offsets, swa_val, mask=mask, cache_modifier=".cg")
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _fused_metadata_kernel_ps1_no_swa(
|
|
||||||
# Input tensors
|
|
||||||
seq_lens,
|
|
||||||
seq_lens_stride_0,
|
|
||||||
req_to_token,
|
|
||||||
req_to_token_stride_0,
|
|
||||||
req_to_token_stride_1,
|
|
||||||
req_pool_indices,
|
|
||||||
req_pool_indices_stride_0,
|
|
||||||
# Output buffers
|
|
||||||
cache_seqlens_int32,
|
|
||||||
cache_seqlens_int32_stride_0,
|
|
||||||
cu_seqlens_k,
|
|
||||||
cu_seqlens_k_stride_0,
|
|
||||||
page_table,
|
|
||||||
page_table_stride_0,
|
|
||||||
page_table_stride_1,
|
|
||||||
# Scalar parameters
|
|
||||||
B,
|
|
||||||
max_seq_pages,
|
|
||||||
seq_len_delta: tl.constexpr,
|
|
||||||
BLOCK_COLS: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid_b = tl.program_id(0) # batch index
|
|
||||||
pid_c = tl.program_id(1) # column chunk index
|
|
||||||
|
|
||||||
# 1. Prefix sum (only one block does it)
|
|
||||||
if pid_b == 0 and pid_c == 0:
|
|
||||||
acc = 0
|
|
||||||
for idx in range(B):
|
|
||||||
seq = tl.load(seq_lens + idx * seq_lens_stride_0)
|
|
||||||
val = (seq + seq_len_delta).to(tl.int32)
|
|
||||||
tl.store(cache_seqlens_int32 + idx * cache_seqlens_int32_stride_0, val)
|
|
||||||
tl.store(cu_seqlens_k + idx * cu_seqlens_k_stride_0, acc)
|
|
||||||
acc += val
|
|
||||||
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
|
|
||||||
|
|
||||||
# 2. Gather for this batch and column chunk
|
|
||||||
if max_seq_pages == 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
i = pid_b
|
|
||||||
# Load row index for this batch (all threads in block have same i)
|
|
||||||
row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0)
|
|
||||||
row_offset = row_idx * req_to_token_stride_0
|
|
||||||
|
|
||||||
col_start = pid_c * BLOCK_COLS
|
|
||||||
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
|
|
||||||
mask = col_offsets < max_seq_pages
|
|
||||||
|
|
||||||
# page_size = 1: col_idx = col_offsets
|
|
||||||
rt_offsets = row_offset + col_offsets * req_to_token_stride_1
|
|
||||||
page_index = tl.load(
|
|
||||||
req_to_token + rt_offsets, mask=mask, other=0, cache_modifier=".cg"
|
|
||||||
)
|
|
||||||
|
|
||||||
# page_table = page_index // 1 = page_index
|
|
||||||
pt_offsets = i * page_table_stride_0 + col_offsets * page_table_stride_1
|
|
||||||
tl.store(page_table + pt_offsets, page_index, mask=mask, cache_modifier=".cg")
|
|
||||||
|
|
||||||
|
|
||||||
# Fused Triton kernel implementation
|
|
||||||
def normal_decode_set_metadata(
|
|
||||||
cache_seqlens_int32: torch.Tensor,
|
|
||||||
cu_seqlens_k: torch.Tensor,
|
|
||||||
page_table: torch.Tensor,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
strided_indices: torch.Tensor,
|
|
||||||
max_seq_pages: torch.Tensor,
|
|
||||||
seq_lens: torch.Tensor,
|
|
||||||
seq_len_delta: int,
|
|
||||||
page_size: int,
|
|
||||||
swa_page_table: Optional[torch.Tensor] = None,
|
|
||||||
token_to_kv_pool: Optional[SWAKVPool] = None,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Fused Triton implementation that replaces 4-5 sequential CUDA kernels with 1-2 kernels:
|
|
||||||
1. cache_seqlens = seq_lens + seq_len_delta (int64→int32 cast)
|
|
||||||
2. cu_seqlens_k = cumsum(cache_seqlens) (prefix-sum)
|
|
||||||
3. page_indices = req_to_token[pool_idx, stride_idx] (2-D gather)
|
|
||||||
4. page_table = page_indices // page_size (floor-divide)
|
|
||||||
5. (optional) swa_page_table for sliding window attention
|
|
||||||
|
|
||||||
Achieves ~5.2x speedup on H200 hardware for typical decode workloads.
|
|
||||||
"""
|
|
||||||
assert (
|
|
||||||
page_size > 0 and (page_size & (page_size - 1)) == 0
|
|
||||||
), f"page_size must be a power of two, got {page_size}"
|
|
||||||
|
|
||||||
batch_size = cache_seqlens_int32.shape[0]
|
|
||||||
device = seq_lens.device
|
|
||||||
|
|
||||||
# Ensure contiguous memory layout for efficient Triton access
|
|
||||||
seq_lens = seq_lens.contiguous()
|
|
||||||
req_to_token = req_to_token.contiguous()
|
|
||||||
req_pool_indices = req_pool_indices.contiguous()
|
|
||||||
|
|
||||||
# Prepare tensor strides
|
|
||||||
seq_lens_stride_0 = seq_lens.stride(0)
|
|
||||||
req_to_token_stride_0 = req_to_token.stride(0)
|
|
||||||
req_to_token_stride_1 = req_to_token.stride(1)
|
|
||||||
req_pool_indices_stride_0 = req_pool_indices.stride(0)
|
|
||||||
cache_seqlens_int32_stride_0 = cache_seqlens_int32.stride(0)
|
|
||||||
cu_seqlens_k_stride_0 = cu_seqlens_k.stride(0)
|
|
||||||
page_table_stride_0 = page_table.stride(0)
|
|
||||||
page_table_stride_1 = page_table.stride(1)
|
|
||||||
|
|
||||||
# Check if we should use the specialized fast path for page_size=1, no SWA
|
|
||||||
use_swa = swa_page_table is not None and token_to_kv_pool is not None
|
|
||||||
|
|
||||||
if page_size == 1 and not use_swa:
|
|
||||||
# Specialized kernel for the common case (page_size=1, no SWA)
|
|
||||||
BLOCK_COLS = 256
|
|
||||||
if max_seq_pages == 0:
|
|
||||||
grid = (1, 1)
|
|
||||||
else:
|
|
||||||
num_blocks_j = triton.cdiv(max_seq_pages, BLOCK_COLS)
|
|
||||||
grid = (batch_size, num_blocks_j)
|
|
||||||
|
|
||||||
_fused_metadata_kernel_ps1_no_swa[grid](
|
|
||||||
seq_lens,
|
|
||||||
seq_lens_stride_0,
|
|
||||||
req_to_token,
|
|
||||||
req_to_token_stride_0,
|
|
||||||
req_to_token_stride_1,
|
|
||||||
req_pool_indices,
|
|
||||||
req_pool_indices_stride_0,
|
|
||||||
cache_seqlens_int32,
|
|
||||||
cache_seqlens_int32_stride_0,
|
|
||||||
cu_seqlens_k,
|
|
||||||
cu_seqlens_k_stride_0,
|
|
||||||
page_table,
|
|
||||||
page_table_stride_0,
|
|
||||||
page_table_stride_1,
|
|
||||||
batch_size,
|
|
||||||
max_seq_pages,
|
|
||||||
seq_len_delta,
|
|
||||||
BLOCK_COLS=BLOCK_COLS,
|
|
||||||
num_warps=8,
|
|
||||||
num_stages=3,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# General kernel for page_size > 1 or SWA cases
|
|
||||||
# SWA parameters
|
|
||||||
if use_swa:
|
|
||||||
assert isinstance(token_to_kv_pool, SWAKVPool)
|
|
||||||
swa_page_table = swa_page_table.contiguous()
|
|
||||||
swa_page_table_stride_0 = swa_page_table.stride(0)
|
|
||||||
swa_page_table_stride_1 = swa_page_table.stride(1)
|
|
||||||
# Extract the full_to_swa_index_mapping from token_to_kv_pool
|
|
||||||
full_to_swa_mapping = (
|
|
||||||
token_to_kv_pool.full_to_swa_index_mapping.contiguous()
|
|
||||||
)
|
|
||||||
full_to_swa_mapping_stride_0 = full_to_swa_mapping.stride(0)
|
|
||||||
else:
|
|
||||||
# Dummy tensors (not used)
|
|
||||||
swa_page_table = torch.empty(0, dtype=torch.int32, device=device)
|
|
||||||
swa_page_table_stride_0 = 0
|
|
||||||
swa_page_table_stride_1 = 0
|
|
||||||
full_to_swa_mapping = torch.empty(0, dtype=torch.int32, device=device)
|
|
||||||
full_to_swa_mapping_stride_0 = 0
|
|
||||||
|
|
||||||
# Kernel configuration
|
|
||||||
BLOCK_COLS = 128
|
|
||||||
shift = (page_size).bit_length() - 1 if page_size > 1 else 0
|
|
||||||
|
|
||||||
if max_seq_pages == 0:
|
|
||||||
grid = (1, 1)
|
|
||||||
else:
|
|
||||||
num_blocks_j = triton.cdiv(max_seq_pages, BLOCK_COLS)
|
|
||||||
grid = (batch_size, num_blocks_j)
|
|
||||||
|
|
||||||
_fused_metadata_kernel_general[grid](
|
|
||||||
seq_lens,
|
|
||||||
seq_lens_stride_0,
|
|
||||||
req_to_token,
|
|
||||||
req_to_token_stride_0,
|
|
||||||
req_to_token_stride_1,
|
|
||||||
req_pool_indices,
|
|
||||||
req_pool_indices_stride_0,
|
|
||||||
cache_seqlens_int32,
|
|
||||||
cache_seqlens_int32_stride_0,
|
|
||||||
cu_seqlens_k,
|
|
||||||
cu_seqlens_k_stride_0,
|
|
||||||
page_table,
|
|
||||||
page_table_stride_0,
|
|
||||||
page_table_stride_1,
|
|
||||||
swa_page_table,
|
|
||||||
swa_page_table_stride_0,
|
|
||||||
swa_page_table_stride_1,
|
|
||||||
full_to_swa_mapping,
|
|
||||||
full_to_swa_mapping_stride_0,
|
|
||||||
batch_size,
|
|
||||||
max_seq_pages,
|
|
||||||
page_size,
|
|
||||||
seq_len_delta,
|
|
||||||
use_swa,
|
|
||||||
shift,
|
|
||||||
BLOCK_COLS=BLOCK_COLS,
|
|
||||||
num_warps=4,
|
|
||||||
num_stages=3,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True, backend=get_compiler_backend())
|
@torch.compile(dynamic=True, backend=get_compiler_backend())
|
||||||
def draft_decode_set_expand_metadata(
|
def draft_decode_set_expand_metadata(
|
||||||
cache_seqlens_int32: torch.Tensor, # Modifies
|
cache_seqlens_int32: torch.Tensor, # Modifies
|
||||||
|
|||||||
@@ -5,11 +5,13 @@ from typing import TYPE_CHECKING, List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import AttentionArch
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
||||||
|
create_flashinfer_kv_indices_triton,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.layers.radix_attention import AttentionType
|
from sglang.srt.layers.radix_attention import AttentionType
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
@@ -1447,58 +1449,6 @@ class TritonMultiStepDraftBackend:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def get_num_kv_splits_triton(
|
|
||||||
num_kv_splits_ptr,
|
|
||||||
seq_lens_ptr,
|
|
||||||
num_seq,
|
|
||||||
num_group,
|
|
||||||
num_head,
|
|
||||||
num_kv_head,
|
|
||||||
max_kv_splits,
|
|
||||||
device_core_count,
|
|
||||||
MAX_NUM_SEQ: tl.constexpr,
|
|
||||||
):
|
|
||||||
# TODO: this method is tunable, we need more online serving data to tune it
|
|
||||||
offs_seq = tl.arange(0, MAX_NUM_SEQ)
|
|
||||||
mask_seq = offs_seq < num_seq
|
|
||||||
|
|
||||||
seq_lens = tl.load(seq_lens_ptr + offs_seq, mask=mask_seq, other=0)
|
|
||||||
max_seq_len = tl.max(seq_lens)
|
|
||||||
seq_lens = tl.load(seq_lens_ptr + offs_seq, mask=mask_seq, other=max_seq_len)
|
|
||||||
min_seq_len = tl.min(seq_lens)
|
|
||||||
if max_seq_len * 8 < min_seq_len * 10:
|
|
||||||
min_seq_len = max_seq_len
|
|
||||||
max_kv_splits_1 = tl.minimum(tl.cdiv(max_seq_len, min_seq_len), max_kv_splits)
|
|
||||||
kv_chunk_size_1 = tl.cdiv(max_seq_len, max_kv_splits_1)
|
|
||||||
|
|
||||||
# NOTE: this is a hack to let num_kv_split grows up with seqlen gradually
|
|
||||||
ext_seq_len = tl.cast(max_seq_len, tl.float32) / 64.0
|
|
||||||
ext_device_core_count = tl.cast(
|
|
||||||
device_core_count * tl.maximum(tl.log2(ext_seq_len), 1.0), tl.int32
|
|
||||||
)
|
|
||||||
block_h, num_kv_group = 16, num_head // num_kv_head
|
|
||||||
if num_kv_group == 1:
|
|
||||||
token_grid = num_seq * num_group * num_head
|
|
||||||
else:
|
|
||||||
# from triton_ops/decode_attention.py:_decode_grouped_att_m_fwd
|
|
||||||
block_h = tl.minimum(block_h, num_kv_group)
|
|
||||||
token_grid = num_seq * num_group * tl.cdiv(num_head, block_h)
|
|
||||||
max_kv_splits_2 = tl.minimum(
|
|
||||||
tl.cdiv(ext_device_core_count, token_grid), max_kv_splits
|
|
||||||
)
|
|
||||||
kv_chunk_size_2 = tl.cdiv(max_seq_len, max_kv_splits_2)
|
|
||||||
|
|
||||||
num_kv_splits = tl.maximum(
|
|
||||||
tl.cdiv(seq_lens, kv_chunk_size_1), tl.cdiv(seq_lens, kv_chunk_size_2)
|
|
||||||
)
|
|
||||||
|
|
||||||
offs_token = offs_seq * num_group
|
|
||||||
mask_token = offs_token < num_seq * num_group
|
|
||||||
for i in range(0, num_group):
|
|
||||||
tl.store(num_kv_splits_ptr + i + offs_token, num_kv_splits, mask=mask_token)
|
|
||||||
|
|
||||||
|
|
||||||
def update_sliding_window_buffer(
|
def update_sliding_window_buffer(
|
||||||
window_kv_indptr,
|
window_kv_indptr,
|
||||||
req_to_token,
|
req_to_token,
|
||||||
|
|||||||
@@ -0,0 +1,266 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def concat_and_cast_mha_k_kernel(
|
||||||
|
k_ptr,
|
||||||
|
k_nope_ptr,
|
||||||
|
k_rope_ptr,
|
||||||
|
head_cnt: tl.constexpr,
|
||||||
|
k_stride0: tl.constexpr,
|
||||||
|
k_stride1: tl.constexpr,
|
||||||
|
nope_stride0: tl.constexpr,
|
||||||
|
nope_stride1: tl.constexpr,
|
||||||
|
rope_stride0: tl.constexpr,
|
||||||
|
nope_dim: tl.constexpr,
|
||||||
|
rope_dim: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_loc = tl.program_id(0)
|
||||||
|
head_range = tl.arange(0, head_cnt)
|
||||||
|
|
||||||
|
k_head_ptr = k_ptr + pid_loc * k_stride0 + head_range[:, None] * k_stride1
|
||||||
|
|
||||||
|
nope_offs = tl.arange(0, nope_dim)
|
||||||
|
|
||||||
|
src_nope_ptr = (
|
||||||
|
k_nope_ptr
|
||||||
|
+ pid_loc * nope_stride0
|
||||||
|
+ head_range[:, None] * nope_stride1
|
||||||
|
+ nope_offs[None, :]
|
||||||
|
)
|
||||||
|
dst_nope_ptr = k_head_ptr + nope_offs[None, :]
|
||||||
|
|
||||||
|
src_nope = tl.load(src_nope_ptr)
|
||||||
|
tl.store(dst_nope_ptr, src_nope)
|
||||||
|
|
||||||
|
rope_offs = tl.arange(0, rope_dim)
|
||||||
|
src_rope_ptr = k_rope_ptr + pid_loc * rope_stride0 + rope_offs[None, :]
|
||||||
|
dst_rope_ptr = k_head_ptr + nope_dim + rope_offs[None, :]
|
||||||
|
src_rope = tl.load(src_rope_ptr)
|
||||||
|
tl.store(dst_rope_ptr, src_rope)
|
||||||
|
|
||||||
|
|
||||||
|
def concat_and_cast_mha_k_triton(
|
||||||
|
k: torch.Tensor,
|
||||||
|
k_nope: torch.Tensor,
|
||||||
|
k_rope: torch.Tensor,
|
||||||
|
):
|
||||||
|
# The source data type will be implicitly converted to the target data type.
|
||||||
|
assert (
|
||||||
|
len(k.shape) == 3 and len(k_nope.shape) == 3 and len(k_rope.shape) == 3
|
||||||
|
), f"shape should be 3d, but got {k.shape=}, {k_nope.shape=}, {k_rope.shape=}"
|
||||||
|
assert (
|
||||||
|
k.shape[0] == k_nope.shape[0] and k.shape[0] == k_rope.shape[0]
|
||||||
|
), f"invalid shape, got {k.shape=}, {k_nope.shape=}, {k_rope.shape=}"
|
||||||
|
assert (
|
||||||
|
k.shape[1] == k_nope.shape[1] and 1 == k_rope.shape[1]
|
||||||
|
), f"invalid shape, got {k.shape=}, {k_nope.shape=}, {k_rope.shape=}"
|
||||||
|
assert (
|
||||||
|
k.shape[-1] == k_nope.shape[-1] + k_rope.shape[-1]
|
||||||
|
), f"invalid shape, got {k.shape=}, {k_nope.shape=}, {k_rope.shape=}"
|
||||||
|
|
||||||
|
nope_dim = k_nope.shape[-1]
|
||||||
|
rope_dim = k_rope.shape[-1]
|
||||||
|
grid = (k.shape[0],)
|
||||||
|
|
||||||
|
concat_and_cast_mha_k_kernel[grid](
|
||||||
|
k,
|
||||||
|
k_nope,
|
||||||
|
k_rope,
|
||||||
|
k.shape[1],
|
||||||
|
k.stride(0),
|
||||||
|
k.stride(1),
|
||||||
|
k_nope.stride(0),
|
||||||
|
k_nope.stride(1),
|
||||||
|
k_rope.stride(0),
|
||||||
|
nope_dim,
|
||||||
|
rope_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def reshape_and_cache_flash(
|
||||||
|
key_ptr,
|
||||||
|
value_ptr,
|
||||||
|
key_cache_ptr,
|
||||||
|
value_cache_ptr,
|
||||||
|
slot_mapping_ptr,
|
||||||
|
swa_slot_mapping_ptr,
|
||||||
|
k_scale_ptr,
|
||||||
|
v_scale_ptr,
|
||||||
|
block_stride,
|
||||||
|
key_stride,
|
||||||
|
value_stride,
|
||||||
|
num_heads,
|
||||||
|
head_size,
|
||||||
|
block_size,
|
||||||
|
HEAD_BLOCK: tl.constexpr,
|
||||||
|
BLOCK_D: tl.constexpr,
|
||||||
|
HAS_SWA: tl.constexpr,
|
||||||
|
USE_SCALE: tl.constexpr,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Triton kernel for reshaping per-token K/V tensors into paged KV cache layout.
|
||||||
|
|
||||||
|
Source layout:
|
||||||
|
key/value: [num_tokens, num_heads, head_size]
|
||||||
|
|
||||||
|
Target cache layout:
|
||||||
|
cache: [num_blocks, block_size, num_heads, head_size]
|
||||||
|
|
||||||
|
Each Triton program instance handles:
|
||||||
|
- one token (program_id(0))
|
||||||
|
- one block of heads (program_id(1))
|
||||||
|
|
||||||
|
Features:
|
||||||
|
- optional SWA slot remapping
|
||||||
|
- optional FP8 scale dequantization before cache write
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key_ptr: Pointer to source key tensor.
|
||||||
|
value_ptr: Pointer to source value tensor.
|
||||||
|
key_cache_ptr: Pointer to destination key cache tensor.
|
||||||
|
value_cache_ptr: Pointer to destination value cache tensor.
|
||||||
|
slot_mapping_ptr: Maps token -> cache slot.
|
||||||
|
swa_slot_mapping_ptr: Optional second-stage slot remap for SWA mode.
|
||||||
|
k_scale_ptr: Optional key scaling factor pointer.
|
||||||
|
v_scale_ptr: Optional value scaling factor pointer.
|
||||||
|
block_stride: Stride between cache blocks.
|
||||||
|
key_stride: Stride between source key tokens.
|
||||||
|
value_stride: Stride between source value tokens.
|
||||||
|
num_heads: Number of attention heads.
|
||||||
|
head_size: Hidden dimension per head.
|
||||||
|
block_size: Number of slots per cache block.
|
||||||
|
HEAD_BLOCK: Number of heads processed per program.
|
||||||
|
BLOCK_D: Vectorized dimension size (power-of-2 padded).
|
||||||
|
HAS_SWA: Enable SWA remapping.
|
||||||
|
USE_SCALE: Enable scale division before storing.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# program ids
|
||||||
|
# pid0 = token
|
||||||
|
# pid1 = head block
|
||||||
|
# ----------------------------------
|
||||||
|
token_idx = tl.program_id(0)
|
||||||
|
head_block_idx = tl.program_id(1)
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# slot mapping
|
||||||
|
# ----------------------------------
|
||||||
|
slot_idx = tl.load(slot_mapping_ptr + token_idx)
|
||||||
|
|
||||||
|
if HAS_SWA:
|
||||||
|
slot_idx = tl.load(swa_slot_mapping_ptr + slot_idx)
|
||||||
|
|
||||||
|
if slot_idx < 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
block_idx = slot_idx // block_size
|
||||||
|
block_offset = slot_idx % block_size
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# head range
|
||||||
|
# ----------------------------------
|
||||||
|
head_idx = head_block_idx * HEAD_BLOCK + tl.arange(0, HEAD_BLOCK)
|
||||||
|
|
||||||
|
head_mask = head_idx < num_heads
|
||||||
|
|
||||||
|
dim_idx = tl.arange(0, BLOCK_D)
|
||||||
|
|
||||||
|
# shape = [HEAD_BLOCK, BLOCK_D]
|
||||||
|
offs = head_idx[:, None] * head_size + dim_idx[None, :]
|
||||||
|
|
||||||
|
mask = head_mask[:, None] & (dim_idx[None, :] < head_size)
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# source load
|
||||||
|
# ----------------------------------
|
||||||
|
src_key = token_idx * key_stride + offs
|
||||||
|
src_value = token_idx * value_stride + offs
|
||||||
|
|
||||||
|
k = tl.load(key_ptr + src_key, mask=mask)
|
||||||
|
v = tl.load(value_ptr + src_value, mask=mask)
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# optional scale
|
||||||
|
# ----------------------------------
|
||||||
|
if USE_SCALE:
|
||||||
|
k_scale = tl.load(k_scale_ptr)
|
||||||
|
v_scale = tl.load(v_scale_ptr)
|
||||||
|
|
||||||
|
k = k / k_scale
|
||||||
|
v = v / v_scale
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# target layout
|
||||||
|
# [block_idx, block_offset, head, dim]
|
||||||
|
# ----------------------------------
|
||||||
|
tgt = block_idx * block_stride + block_offset * num_heads * head_size + offs
|
||||||
|
|
||||||
|
tl.store(key_cache_ptr + tgt, k, mask=mask)
|
||||||
|
tl.store(value_cache_ptr + tgt, v, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def launch_reshape_and_cache_flash(
|
||||||
|
key,
|
||||||
|
value,
|
||||||
|
key_cache,
|
||||||
|
value_cache,
|
||||||
|
slot_mapping,
|
||||||
|
swa_slot_mapping=None,
|
||||||
|
k_scale=None,
|
||||||
|
v_scale=None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Launch wrapper for reshape_and_cache_flash Triton kernel.
|
||||||
|
|
||||||
|
This wrapper prepares launch configuration and dispatches the Triton kernel
|
||||||
|
that writes token-major K/V tensors into paged KV cache layout.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key: Source key tensor [num_tokens, num_heads, head_size]
|
||||||
|
value: Source value tensor [num_tokens, num_heads, head_size]
|
||||||
|
key_cache: Destination key cache [num_blocks, block_size, num_heads, head_size]
|
||||||
|
value_cache: Destination value cache [num_blocks, block_size, num_heads, head_size]
|
||||||
|
slot_mapping: Token-to-cache slot mapping
|
||||||
|
swa_slot_mapping: Optional SWA remapping table
|
||||||
|
k_scale: Optional key scaling factor
|
||||||
|
v_scale: Optional value scaling factor
|
||||||
|
"""
|
||||||
|
|
||||||
|
num_tokens = key.shape[0]
|
||||||
|
num_heads = key.shape[1]
|
||||||
|
head_size = key.shape[2]
|
||||||
|
|
||||||
|
HEAD_BLOCK = 4
|
||||||
|
|
||||||
|
BLOCK_D = triton.next_power_of_2(head_size)
|
||||||
|
|
||||||
|
grid = (
|
||||||
|
num_tokens,
|
||||||
|
triton.cdiv(num_heads, HEAD_BLOCK),
|
||||||
|
)
|
||||||
|
|
||||||
|
reshape_and_cache_flash[grid](
|
||||||
|
key,
|
||||||
|
value,
|
||||||
|
key_cache,
|
||||||
|
value_cache,
|
||||||
|
slot_mapping,
|
||||||
|
swa_slot_mapping,
|
||||||
|
k_scale if k_scale is not None else key,
|
||||||
|
v_scale if v_scale is not None else key,
|
||||||
|
key_cache.stride(0),
|
||||||
|
key.stride(0),
|
||||||
|
value.stride(0),
|
||||||
|
num_heads,
|
||||||
|
head_size,
|
||||||
|
key_cache.shape[1],
|
||||||
|
HEAD_BLOCK=HEAD_BLOCK,
|
||||||
|
BLOCK_D=BLOCK_D,
|
||||||
|
HAS_SWA=(swa_slot_mapping is not None),
|
||||||
|
USE_SCALE=(k_scale is not None),
|
||||||
|
)
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
_FLASHMLA_CREATE_KV_BLOCK_SIZE = 4096
|
||||||
|
FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON = tl.constexpr(_FLASHMLA_CREATE_KV_BLOCK_SIZE)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def create_flashinfer_kv_indices_triton(
|
||||||
|
req_to_token_ptr, # [max_batch, max_context_len]
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
page_kernel_lens_ptr,
|
||||||
|
kv_indptr,
|
||||||
|
kv_start_idx,
|
||||||
|
kv_indices_ptr,
|
||||||
|
req_to_token_ptr_stride: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 512
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
|
||||||
|
# find the req pool idx, this is for batch to token
|
||||||
|
req_pool_index = tl.load(req_pool_indices_ptr + pid)
|
||||||
|
kv_indices_offset = tl.load(kv_indptr + pid)
|
||||||
|
|
||||||
|
kv_start = 0
|
||||||
|
kv_end = 0
|
||||||
|
if kv_start_idx:
|
||||||
|
kv_start = tl.load(kv_start_idx + pid).to(tl.int32)
|
||||||
|
kv_end = kv_start
|
||||||
|
kv_end += tl.load(page_kernel_lens_ptr + pid).to(tl.int32)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
# index into req_to_token_ptr needs to be int64
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE
|
||||||
|
mask = offset < kv_end - kv_start
|
||||||
|
data = tl.load(
|
||||||
|
req_to_token_ptr
|
||||||
|
+ req_pool_index * req_to_token_ptr_stride
|
||||||
|
+ kv_start
|
||||||
|
+ offset,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
tl.store(kv_indices_ptr + kv_indices_offset + offset, data, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def get_num_page_per_block_flashmla(page_size: int = 64) -> int:
|
||||||
|
num_page_per_block = _FLASHMLA_CREATE_KV_BLOCK_SIZE // page_size
|
||||||
|
return num_page_per_block
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def create_flashmla_kv_indices_triton(
|
||||||
|
req_to_token_ptr, # [max_batch, max_context_len]
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
page_kernel_lens_ptr,
|
||||||
|
kv_start_idx,
|
||||||
|
kv_indices_ptr,
|
||||||
|
req_to_token_ptr_stride: tl.constexpr,
|
||||||
|
kv_indices_ptr_stride: tl.constexpr,
|
||||||
|
PAGED_SIZE: tl.constexpr = 64,
|
||||||
|
):
|
||||||
|
NUM_PAGE_PER_BLOCK: tl.constexpr = (
|
||||||
|
FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON // PAGED_SIZE
|
||||||
|
)
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
|
||||||
|
# find the req pool idx, this is for batch to token
|
||||||
|
req_pool_index = tl.load(req_pool_indices_ptr + pid)
|
||||||
|
|
||||||
|
kv_start = 0
|
||||||
|
kv_end = 0
|
||||||
|
if kv_start_idx:
|
||||||
|
kv_start = tl.load(kv_start_idx + pid).to(tl.int32)
|
||||||
|
kv_end = kv_start
|
||||||
|
|
||||||
|
kv_end += tl.load(page_kernel_lens_ptr + pid).to(tl.int32)
|
||||||
|
|
||||||
|
num_paged = tl.cdiv(kv_end - kv_start, PAGED_SIZE)
|
||||||
|
num_pages_loop = tl.cdiv(kv_end - kv_start, FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON)
|
||||||
|
|
||||||
|
for i in range(num_pages_loop):
|
||||||
|
# index into req_to_token_ptr needs to be int64
|
||||||
|
paged_offset = (
|
||||||
|
tl.arange(0, NUM_PAGE_PER_BLOCK).to(tl.int64) + i * NUM_PAGE_PER_BLOCK
|
||||||
|
) * PAGED_SIZE
|
||||||
|
paged_offset_out = tl.arange(0, NUM_PAGE_PER_BLOCK) + i * NUM_PAGE_PER_BLOCK
|
||||||
|
|
||||||
|
mask = paged_offset < num_paged * PAGED_SIZE
|
||||||
|
mask_out = paged_offset_out < num_paged
|
||||||
|
|
||||||
|
data = tl.load(
|
||||||
|
req_to_token_ptr
|
||||||
|
+ req_pool_index * req_to_token_ptr_stride
|
||||||
|
+ kv_start
|
||||||
|
+ paged_offset,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
kv_indices_ptr + pid * kv_indices_ptr_stride + paged_offset_out,
|
||||||
|
data // PAGED_SIZE,
|
||||||
|
mask=mask_out,
|
||||||
|
)
|
||||||
@@ -0,0 +1,467 @@
|
|||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def get_num_kv_splits_triton(
|
||||||
|
num_kv_splits_ptr,
|
||||||
|
seq_lens_ptr,
|
||||||
|
num_seq,
|
||||||
|
num_group,
|
||||||
|
num_head,
|
||||||
|
num_kv_head,
|
||||||
|
max_kv_splits,
|
||||||
|
device_core_count,
|
||||||
|
MAX_NUM_SEQ: tl.constexpr,
|
||||||
|
):
|
||||||
|
# TODO: this method is tunable, we need more online serving data to tune it
|
||||||
|
offs_seq = tl.arange(0, MAX_NUM_SEQ)
|
||||||
|
mask_seq = offs_seq < num_seq
|
||||||
|
|
||||||
|
seq_lens = tl.load(seq_lens_ptr + offs_seq, mask=mask_seq, other=0)
|
||||||
|
max_seq_len = tl.max(seq_lens)
|
||||||
|
seq_lens = tl.load(seq_lens_ptr + offs_seq, mask=mask_seq, other=max_seq_len)
|
||||||
|
min_seq_len = tl.min(seq_lens)
|
||||||
|
if max_seq_len * 8 < min_seq_len * 10:
|
||||||
|
min_seq_len = max_seq_len
|
||||||
|
max_kv_splits_1 = tl.minimum(tl.cdiv(max_seq_len, min_seq_len), max_kv_splits)
|
||||||
|
kv_chunk_size_1 = tl.cdiv(max_seq_len, max_kv_splits_1)
|
||||||
|
|
||||||
|
# NOTE: this is a hack to let num_kv_split grows up with seqlen gradually
|
||||||
|
ext_seq_len = tl.cast(max_seq_len, tl.float32) / 64.0
|
||||||
|
ext_device_core_count = tl.cast(
|
||||||
|
device_core_count * tl.maximum(tl.log2(ext_seq_len), 1.0), tl.int32
|
||||||
|
)
|
||||||
|
block_h, num_kv_group = 16, num_head // num_kv_head
|
||||||
|
if num_kv_group == 1:
|
||||||
|
token_grid = num_seq * num_group * num_head
|
||||||
|
else:
|
||||||
|
# from triton_ops/decode_attention.py:_decode_grouped_att_m_fwd
|
||||||
|
block_h = tl.minimum(block_h, num_kv_group)
|
||||||
|
token_grid = num_seq * num_group * tl.cdiv(num_head, block_h)
|
||||||
|
max_kv_splits_2 = tl.minimum(
|
||||||
|
tl.cdiv(ext_device_core_count, token_grid), max_kv_splits
|
||||||
|
)
|
||||||
|
kv_chunk_size_2 = tl.cdiv(max_seq_len, max_kv_splits_2)
|
||||||
|
|
||||||
|
num_kv_splits = tl.maximum(
|
||||||
|
tl.cdiv(seq_lens, kv_chunk_size_1), tl.cdiv(seq_lens, kv_chunk_size_2)
|
||||||
|
)
|
||||||
|
|
||||||
|
offs_token = offs_seq * num_group
|
||||||
|
mask_token = offs_token < num_seq * num_group
|
||||||
|
for i in range(0, num_group):
|
||||||
|
tl.store(num_kv_splits_ptr + i + offs_token, num_kv_splits, mask=mask_token)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _prepare_swa_spec_page_table_kernel(
|
||||||
|
dst_ptr,
|
||||||
|
src_a_ptr,
|
||||||
|
src_b_ptr,
|
||||||
|
seq_len_a_ptr,
|
||||||
|
seq_len_b_ptr,
|
||||||
|
dst_stride_m,
|
||||||
|
dst_stride_n,
|
||||||
|
a_stride_m,
|
||||||
|
a_stride_n,
|
||||||
|
b_stride_m,
|
||||||
|
b_stride_n,
|
||||||
|
LEN_A: tl.constexpr,
|
||||||
|
LEN_B: tl.constexpr,
|
||||||
|
REPEAT_STEP: tl.constexpr,
|
||||||
|
BLOCK_N: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_m = tl.program_id(0)
|
||||||
|
pid_n = tl.program_id(1)
|
||||||
|
|
||||||
|
idx_a = pid_m // REPEAT_STEP
|
||||||
|
idx_b = pid_m
|
||||||
|
seq_len_a = tl.load(seq_len_a_ptr + idx_a)
|
||||||
|
seq_len_b = tl.load(seq_len_b_ptr + idx_b)
|
||||||
|
|
||||||
|
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||||
|
total_len = seq_len_a + seq_len_b
|
||||||
|
|
||||||
|
if pid_n * BLOCK_N >= total_len:
|
||||||
|
return
|
||||||
|
|
||||||
|
mask = offs_n < total_len
|
||||||
|
dst = dst_ptr + pid_m * dst_stride_m + offs_n * dst_stride_n
|
||||||
|
|
||||||
|
if (pid_n + 1) * BLOCK_N < seq_len_a:
|
||||||
|
a_ptr = src_a_ptr + idx_a * a_stride_m + offs_n * a_stride_n
|
||||||
|
a_mask = mask & (offs_n < LEN_A)
|
||||||
|
val = tl.load(a_ptr, mask=a_mask, other=0)
|
||||||
|
tl.store(dst, val, mask=mask)
|
||||||
|
elif pid_n * BLOCK_N >= seq_len_a:
|
||||||
|
offs_b = offs_n - seq_len_a
|
||||||
|
b_ptr = src_b_ptr + idx_b * b_stride_m + offs_b * b_stride_n
|
||||||
|
b_mask = mask & (offs_b < LEN_B)
|
||||||
|
val = tl.load(b_ptr, mask=b_mask, other=0)
|
||||||
|
tl.store(dst, val, mask=mask)
|
||||||
|
else:
|
||||||
|
# mixed part
|
||||||
|
a_offs = offs_n
|
||||||
|
a_mask = (a_offs < seq_len_a) & (a_offs < LEN_A)
|
||||||
|
a_ptr = src_a_ptr + idx_a * a_stride_m + a_offs * a_stride_n
|
||||||
|
a_val = tl.load(a_ptr, mask=a_mask, other=0)
|
||||||
|
|
||||||
|
b_offs = offs_n - seq_len_a
|
||||||
|
b_mask = (b_offs >= 0) & (b_offs < seq_len_b) & (b_offs < LEN_B)
|
||||||
|
b_ptr = src_b_ptr + idx_b * b_stride_m + b_offs * b_stride_n
|
||||||
|
b_val = tl.load(b_ptr, mask=b_mask, other=0)
|
||||||
|
|
||||||
|
result = tl.where(offs_n < seq_len_a, a_val, b_val)
|
||||||
|
tl.store(dst, result, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_swa_spec_page_table_triton(
|
||||||
|
page_table_dst: torch.Tensor,
|
||||||
|
page_table_a: torch.Tensor,
|
||||||
|
page_table_b: torch.Tensor, # expand page table
|
||||||
|
seq_len_a: torch.Tensor,
|
||||||
|
seq_len_b: torch.Tensor, # expand seq lens
|
||||||
|
speculative_num_draft_tokens: int,
|
||||||
|
):
|
||||||
|
# concat page_table and expand page_table by kv seq length
|
||||||
|
bs = seq_len_a.numel()
|
||||||
|
bs_expand = seq_len_b.numel()
|
||||||
|
assert bs_expand == bs * speculative_num_draft_tokens
|
||||||
|
|
||||||
|
LEN_A = page_table_a.shape[1]
|
||||||
|
LEN_B = page_table_b.shape[1]
|
||||||
|
LEN_OUT = LEN_A + LEN_B
|
||||||
|
REPEAT_STEP = speculative_num_draft_tokens
|
||||||
|
BLOCK_N = 256
|
||||||
|
|
||||||
|
grid = (bs_expand, triton.cdiv(LEN_OUT, BLOCK_N))
|
||||||
|
_prepare_swa_spec_page_table_kernel[grid](
|
||||||
|
page_table_dst,
|
||||||
|
page_table_a,
|
||||||
|
page_table_b,
|
||||||
|
seq_len_a,
|
||||||
|
seq_len_b,
|
||||||
|
page_table_dst.stride(0),
|
||||||
|
page_table_dst.stride(1),
|
||||||
|
page_table_a.stride(0),
|
||||||
|
page_table_a.stride(1),
|
||||||
|
page_table_b.stride(0),
|
||||||
|
page_table_b.stride(1),
|
||||||
|
LEN_A=LEN_A,
|
||||||
|
LEN_B=LEN_B,
|
||||||
|
REPEAT_STEP=REPEAT_STEP,
|
||||||
|
BLOCK_N=BLOCK_N,
|
||||||
|
num_warps=4,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_metadata_kernel_general(
|
||||||
|
# Input tensors
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_stride_0,
|
||||||
|
req_to_token,
|
||||||
|
req_to_token_stride_0,
|
||||||
|
req_to_token_stride_1,
|
||||||
|
req_pool_indices,
|
||||||
|
req_pool_indices_stride_0,
|
||||||
|
# Output buffers
|
||||||
|
cache_seqlens_int32,
|
||||||
|
cache_seqlens_int32_stride_0,
|
||||||
|
cu_seqlens_k,
|
||||||
|
cu_seqlens_k_stride_0,
|
||||||
|
page_table,
|
||||||
|
page_table_stride_0,
|
||||||
|
page_table_stride_1,
|
||||||
|
swa_page_table,
|
||||||
|
swa_page_table_stride_0,
|
||||||
|
swa_page_table_stride_1,
|
||||||
|
full_to_swa_mapping,
|
||||||
|
full_to_swa_mapping_stride_0,
|
||||||
|
# Scalar parameters
|
||||||
|
B,
|
||||||
|
max_seq_pages,
|
||||||
|
page_size: tl.constexpr,
|
||||||
|
seq_len_delta: tl.constexpr,
|
||||||
|
use_swa: tl.constexpr,
|
||||||
|
SHIFT: tl.constexpr,
|
||||||
|
BLOCK_COLS: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_b = tl.program_id(0) # batch index
|
||||||
|
pid_c = tl.program_id(1) # column chunk index
|
||||||
|
|
||||||
|
# 1. Prefix sum (only one block does it)
|
||||||
|
if pid_b == 0 and pid_c == 0:
|
||||||
|
acc = 0
|
||||||
|
for idx in range(B):
|
||||||
|
seq = tl.load(seq_lens + idx * seq_lens_stride_0)
|
||||||
|
val = (seq + seq_len_delta).to(tl.int32)
|
||||||
|
tl.store(cache_seqlens_int32 + idx * cache_seqlens_int32_stride_0, val)
|
||||||
|
tl.store(cu_seqlens_k + idx * cu_seqlens_k_stride_0, acc)
|
||||||
|
acc += val
|
||||||
|
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
|
||||||
|
|
||||||
|
# 2. Gather for this batch and column chunk
|
||||||
|
if max_seq_pages == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
i = pid_b
|
||||||
|
# Load row index for this batch (all threads in block have same i)
|
||||||
|
row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0)
|
||||||
|
row_offset = row_idx * req_to_token_stride_0
|
||||||
|
|
||||||
|
col_start = pid_c * BLOCK_COLS
|
||||||
|
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
|
||||||
|
mask = col_offsets < max_seq_pages
|
||||||
|
|
||||||
|
# Compute column indices in the source tensor (token offset)
|
||||||
|
if page_size == 1:
|
||||||
|
col_idx = col_offsets
|
||||||
|
else:
|
||||||
|
col_idx = col_offsets << SHIFT # faster than multiplication for power-of-two
|
||||||
|
|
||||||
|
# Load page indices from req_to_token
|
||||||
|
rt_offsets = row_offset + col_idx * req_to_token_stride_1
|
||||||
|
page_index = tl.load(
|
||||||
|
req_to_token + rt_offsets, mask=mask, other=0, cache_modifier=".cg"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Compute page_table
|
||||||
|
if page_size == 1:
|
||||||
|
page_table_val = page_index
|
||||||
|
else:
|
||||||
|
page_table_val = page_index >> SHIFT
|
||||||
|
|
||||||
|
# Store to page_table
|
||||||
|
pt_offsets = i * page_table_stride_0 + col_offsets * page_table_stride_1
|
||||||
|
tl.store(page_table + pt_offsets, page_table_val, mask=mask, cache_modifier=".cg")
|
||||||
|
|
||||||
|
if use_swa:
|
||||||
|
swa_slot = tl.load(
|
||||||
|
full_to_swa_mapping + page_index * full_to_swa_mapping_stride_0,
|
||||||
|
mask=mask,
|
||||||
|
other=0,
|
||||||
|
cache_modifier=".cg",
|
||||||
|
)
|
||||||
|
if page_size == 1:
|
||||||
|
swa_val = swa_slot
|
||||||
|
else:
|
||||||
|
swa_val = swa_slot >> SHIFT
|
||||||
|
swa_offsets = (
|
||||||
|
i * swa_page_table_stride_0 + col_offsets * swa_page_table_stride_1
|
||||||
|
)
|
||||||
|
tl.store(swa_page_table + swa_offsets, swa_val, mask=mask, cache_modifier=".cg")
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_metadata_kernel_ps1_no_swa(
|
||||||
|
# Input tensors
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_stride_0,
|
||||||
|
req_to_token,
|
||||||
|
req_to_token_stride_0,
|
||||||
|
req_to_token_stride_1,
|
||||||
|
req_pool_indices,
|
||||||
|
req_pool_indices_stride_0,
|
||||||
|
# Output buffers
|
||||||
|
cache_seqlens_int32,
|
||||||
|
cache_seqlens_int32_stride_0,
|
||||||
|
cu_seqlens_k,
|
||||||
|
cu_seqlens_k_stride_0,
|
||||||
|
page_table,
|
||||||
|
page_table_stride_0,
|
||||||
|
page_table_stride_1,
|
||||||
|
# Scalar parameters
|
||||||
|
B,
|
||||||
|
max_seq_pages,
|
||||||
|
seq_len_delta: tl.constexpr,
|
||||||
|
BLOCK_COLS: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_b = tl.program_id(0) # batch index
|
||||||
|
pid_c = tl.program_id(1) # column chunk index
|
||||||
|
|
||||||
|
# 1. Prefix sum (only one block does it)
|
||||||
|
if pid_b == 0 and pid_c == 0:
|
||||||
|
acc = 0
|
||||||
|
for idx in range(B):
|
||||||
|
seq = tl.load(seq_lens + idx * seq_lens_stride_0)
|
||||||
|
val = (seq + seq_len_delta).to(tl.int32)
|
||||||
|
tl.store(cache_seqlens_int32 + idx * cache_seqlens_int32_stride_0, val)
|
||||||
|
tl.store(cu_seqlens_k + idx * cu_seqlens_k_stride_0, acc)
|
||||||
|
acc += val
|
||||||
|
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
|
||||||
|
|
||||||
|
# 2. Gather for this batch and column chunk
|
||||||
|
if max_seq_pages == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
i = pid_b
|
||||||
|
# Load row index for this batch (all threads in block have same i)
|
||||||
|
row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0)
|
||||||
|
row_offset = row_idx * req_to_token_stride_0
|
||||||
|
|
||||||
|
col_start = pid_c * BLOCK_COLS
|
||||||
|
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
|
||||||
|
mask = col_offsets < max_seq_pages
|
||||||
|
|
||||||
|
# page_size = 1: col_idx = col_offsets
|
||||||
|
rt_offsets = row_offset + col_offsets * req_to_token_stride_1
|
||||||
|
page_index = tl.load(
|
||||||
|
req_to_token + rt_offsets, mask=mask, other=0, cache_modifier=".cg"
|
||||||
|
)
|
||||||
|
|
||||||
|
# page_table = page_index // 1 = page_index
|
||||||
|
pt_offsets = i * page_table_stride_0 + col_offsets * page_table_stride_1
|
||||||
|
tl.store(page_table + pt_offsets, page_index, mask=mask, cache_modifier=".cg")
|
||||||
|
|
||||||
|
|
||||||
|
def normal_decode_set_metadata(
|
||||||
|
cache_seqlens_int32: torch.Tensor,
|
||||||
|
cu_seqlens_k: torch.Tensor,
|
||||||
|
page_table: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
strided_indices: torch.Tensor,
|
||||||
|
max_seq_pages: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
seq_len_delta: int,
|
||||||
|
page_size: int,
|
||||||
|
swa_page_table: Optional[torch.Tensor] = None,
|
||||||
|
token_to_kv_pool: Optional["SWAKVPool"] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Fused Triton implementation that replaces 4-5 sequential CUDA kernels with 1-2 kernels:
|
||||||
|
1. cache_seqlens = seq_lens + seq_len_delta (int64->int32 cast)
|
||||||
|
2. cu_seqlens_k = cumsum(cache_seqlens) (prefix-sum)
|
||||||
|
3. page_indices = req_to_token[pool_idx, stride_idx] (2-D gather)
|
||||||
|
4. page_table = page_indices // page_size (floor-divide)
|
||||||
|
5. (optional) swa_page_table for sliding window attention
|
||||||
|
|
||||||
|
Achieves ~5.2x speedup on H200 hardware for typical decode workloads.
|
||||||
|
"""
|
||||||
|
assert (
|
||||||
|
page_size > 0 and (page_size & (page_size - 1)) == 0
|
||||||
|
), f"page_size must be a power of two, got {page_size}"
|
||||||
|
|
||||||
|
batch_size = cache_seqlens_int32.shape[0]
|
||||||
|
device = seq_lens.device
|
||||||
|
|
||||||
|
# Ensure contiguous memory layout for efficient Triton access
|
||||||
|
seq_lens = seq_lens.contiguous()
|
||||||
|
req_to_token = req_to_token.contiguous()
|
||||||
|
req_pool_indices = req_pool_indices.contiguous()
|
||||||
|
|
||||||
|
# Prepare tensor strides
|
||||||
|
seq_lens_stride_0 = seq_lens.stride(0)
|
||||||
|
req_to_token_stride_0 = req_to_token.stride(0)
|
||||||
|
req_to_token_stride_1 = req_to_token.stride(1)
|
||||||
|
req_pool_indices_stride_0 = req_pool_indices.stride(0)
|
||||||
|
cache_seqlens_int32_stride_0 = cache_seqlens_int32.stride(0)
|
||||||
|
cu_seqlens_k_stride_0 = cu_seqlens_k.stride(0)
|
||||||
|
page_table_stride_0 = page_table.stride(0)
|
||||||
|
page_table_stride_1 = page_table.stride(1)
|
||||||
|
|
||||||
|
# Check if we should use the specialized fast path for page_size=1, no SWA
|
||||||
|
use_swa = swa_page_table is not None and token_to_kv_pool is not None
|
||||||
|
|
||||||
|
if page_size == 1 and not use_swa:
|
||||||
|
# Specialized kernel for the common case (page_size=1, no SWA)
|
||||||
|
BLOCK_COLS = 256
|
||||||
|
if max_seq_pages == 0:
|
||||||
|
grid = (1, 1)
|
||||||
|
else:
|
||||||
|
num_blocks_j = triton.cdiv(max_seq_pages, BLOCK_COLS)
|
||||||
|
grid = (batch_size, num_blocks_j)
|
||||||
|
|
||||||
|
_fused_metadata_kernel_ps1_no_swa[grid](
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_stride_0,
|
||||||
|
req_to_token,
|
||||||
|
req_to_token_stride_0,
|
||||||
|
req_to_token_stride_1,
|
||||||
|
req_pool_indices,
|
||||||
|
req_pool_indices_stride_0,
|
||||||
|
cache_seqlens_int32,
|
||||||
|
cache_seqlens_int32_stride_0,
|
||||||
|
cu_seqlens_k,
|
||||||
|
cu_seqlens_k_stride_0,
|
||||||
|
page_table,
|
||||||
|
page_table_stride_0,
|
||||||
|
page_table_stride_1,
|
||||||
|
batch_size,
|
||||||
|
max_seq_pages,
|
||||||
|
seq_len_delta,
|
||||||
|
BLOCK_COLS=BLOCK_COLS,
|
||||||
|
num_warps=8,
|
||||||
|
num_stages=3,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# General kernel for page_size > 1 or SWA cases
|
||||||
|
# SWA parameters
|
||||||
|
if use_swa:
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
|
|
||||||
|
assert isinstance(token_to_kv_pool, SWAKVPool)
|
||||||
|
swa_page_table = swa_page_table.contiguous()
|
||||||
|
swa_page_table_stride_0 = swa_page_table.stride(0)
|
||||||
|
swa_page_table_stride_1 = swa_page_table.stride(1)
|
||||||
|
# Extract the full_to_swa_index_mapping from token_to_kv_pool
|
||||||
|
full_to_swa_mapping = (
|
||||||
|
token_to_kv_pool.full_to_swa_index_mapping.contiguous()
|
||||||
|
)
|
||||||
|
full_to_swa_mapping_stride_0 = full_to_swa_mapping.stride(0)
|
||||||
|
else:
|
||||||
|
# Dummy tensors (not used)
|
||||||
|
swa_page_table = torch.empty(0, dtype=torch.int32, device=device)
|
||||||
|
swa_page_table_stride_0 = 0
|
||||||
|
swa_page_table_stride_1 = 0
|
||||||
|
full_to_swa_mapping = torch.empty(0, dtype=torch.int32, device=device)
|
||||||
|
full_to_swa_mapping_stride_0 = 0
|
||||||
|
|
||||||
|
# Kernel configuration
|
||||||
|
BLOCK_COLS = 128
|
||||||
|
shift = (page_size).bit_length() - 1 if page_size > 1 else 0
|
||||||
|
|
||||||
|
if max_seq_pages == 0:
|
||||||
|
grid = (1, 1)
|
||||||
|
else:
|
||||||
|
num_blocks_j = triton.cdiv(max_seq_pages, BLOCK_COLS)
|
||||||
|
grid = (batch_size, num_blocks_j)
|
||||||
|
|
||||||
|
_fused_metadata_kernel_general[grid](
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_stride_0,
|
||||||
|
req_to_token,
|
||||||
|
req_to_token_stride_0,
|
||||||
|
req_to_token_stride_1,
|
||||||
|
req_pool_indices,
|
||||||
|
req_pool_indices_stride_0,
|
||||||
|
cache_seqlens_int32,
|
||||||
|
cache_seqlens_int32_stride_0,
|
||||||
|
cu_seqlens_k,
|
||||||
|
cu_seqlens_k_stride_0,
|
||||||
|
page_table,
|
||||||
|
page_table_stride_0,
|
||||||
|
page_table_stride_1,
|
||||||
|
swa_page_table,
|
||||||
|
swa_page_table_stride_0,
|
||||||
|
swa_page_table_stride_1,
|
||||||
|
full_to_swa_mapping,
|
||||||
|
full_to_swa_mapping_stride_0,
|
||||||
|
batch_size,
|
||||||
|
max_seq_pages,
|
||||||
|
page_size,
|
||||||
|
seq_len_delta,
|
||||||
|
use_swa,
|
||||||
|
shift,
|
||||||
|
BLOCK_COLS=BLOCK_COLS,
|
||||||
|
num_warps=4,
|
||||||
|
num_stages=3,
|
||||||
|
)
|
||||||
@@ -0,0 +1,162 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def pad_sequence_with_mask_kernel(
|
||||||
|
input_ptr, # (total_tokens, hidden)
|
||||||
|
offsets_ptr, # (B,)
|
||||||
|
lengths_ptr, # (B,)
|
||||||
|
output_ptr, # (B, max_len, hidden)
|
||||||
|
mask_ptr, # (B, max_len)
|
||||||
|
max_len,
|
||||||
|
hidden_dim,
|
||||||
|
BLOCK_M: tl.constexpr, # seq block
|
||||||
|
BLOCK_D: tl.constexpr, # hidden block
|
||||||
|
):
|
||||||
|
b = tl.program_id(0) # batch index
|
||||||
|
m = tl.program_id(1) # seq block index
|
||||||
|
|
||||||
|
offset = tl.load(offsets_ptr + b)
|
||||||
|
length = tl.load(lengths_ptr + b)
|
||||||
|
|
||||||
|
seq_ids = m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||||
|
hid_ids = tl.arange(0, BLOCK_D)
|
||||||
|
|
||||||
|
seq_mask = seq_ids < max_len
|
||||||
|
valid_token = seq_ids < length
|
||||||
|
|
||||||
|
# input index
|
||||||
|
in_token = offset + seq_ids
|
||||||
|
in_ptr = input_ptr + in_token[:, None] * hidden_dim + hid_ids[None, :]
|
||||||
|
|
||||||
|
# output index
|
||||||
|
out_ptr = (
|
||||||
|
output_ptr
|
||||||
|
+ b * max_len * hidden_dim
|
||||||
|
+ seq_ids[:, None] * hidden_dim
|
||||||
|
+ hid_ids[None, :]
|
||||||
|
)
|
||||||
|
|
||||||
|
values = tl.load(
|
||||||
|
in_ptr,
|
||||||
|
mask=valid_token[:, None] & (hid_ids[None, :] < hidden_dim),
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
tl.store(
|
||||||
|
out_ptr,
|
||||||
|
values,
|
||||||
|
mask=seq_mask[:, None] & (hid_ids[None, :] < hidden_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
# attention mask
|
||||||
|
if tl.program_id(2) == 0:
|
||||||
|
mask_out_ptr = mask_ptr + b * max_len + seq_ids
|
||||||
|
tl.store(mask_out_ptr, valid_token, mask=seq_mask)
|
||||||
|
|
||||||
|
|
||||||
|
def pad_sequence_with_mask(
|
||||||
|
input_emb, # (total_tokens, hidden)
|
||||||
|
offsets, # (B,)
|
||||||
|
lengths, # (B,)
|
||||||
|
max_len,
|
||||||
|
):
|
||||||
|
B = offsets.shape[0]
|
||||||
|
hidden_dim = input_emb.shape[1]
|
||||||
|
|
||||||
|
output = torch.zeros(
|
||||||
|
(B, max_len, hidden_dim),
|
||||||
|
device=input_emb.device,
|
||||||
|
dtype=input_emb.dtype,
|
||||||
|
)
|
||||||
|
attn_mask = torch.empty(
|
||||||
|
(B * max_len),
|
||||||
|
device=input_emb.device,
|
||||||
|
dtype=torch.bool,
|
||||||
|
)
|
||||||
|
|
||||||
|
BLOCK_D = triton.next_power_of_2(hidden_dim)
|
||||||
|
BLOCK_M = triton.next_power_of_2(max_len)
|
||||||
|
|
||||||
|
grid = (
|
||||||
|
B,
|
||||||
|
triton.cdiv(max_len, BLOCK_M),
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
|
||||||
|
pad_sequence_with_mask_kernel[grid](
|
||||||
|
input_emb,
|
||||||
|
offsets,
|
||||||
|
lengths,
|
||||||
|
output,
|
||||||
|
attn_mask,
|
||||||
|
max_len,
|
||||||
|
hidden_dim,
|
||||||
|
BLOCK_M=BLOCK_M,
|
||||||
|
BLOCK_D=BLOCK_D,
|
||||||
|
)
|
||||||
|
|
||||||
|
return B, output, attn_mask
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def seqlens_expand_kernel(
|
||||||
|
extend_seq_lens_ptr, # [N]
|
||||||
|
seq_lens_ptr, # [N]
|
||||||
|
offsets_ptr, # [N+1]
|
||||||
|
output_ptr, # [sum(extend_seq_lens)]
|
||||||
|
N,
|
||||||
|
BLOCK: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
if pid >= N:
|
||||||
|
return
|
||||||
|
|
||||||
|
qo_len = tl.load(extend_seq_lens_ptr + pid)
|
||||||
|
kv_len = tl.load(seq_lens_ptr + pid)
|
||||||
|
|
||||||
|
start = kv_len - qo_len + 1
|
||||||
|
out_offset = tl.load(offsets_ptr + pid)
|
||||||
|
|
||||||
|
offs = tl.arange(0, BLOCK)
|
||||||
|
mask = offs < qo_len
|
||||||
|
|
||||||
|
values = start + offs
|
||||||
|
tl.store(output_ptr + out_offset + offs, values, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def seqlens_expand_triton(
|
||||||
|
extend_seq_lens: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
total_len: int,
|
||||||
|
max_q_len: int,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
extend_seq_lens: [N], int32, CUDA
|
||||||
|
seq_lens: [N], int32, CUDA
|
||||||
|
"""
|
||||||
|
assert extend_seq_lens.is_cuda
|
||||||
|
assert seq_lens.is_cuda
|
||||||
|
|
||||||
|
N = extend_seq_lens.numel()
|
||||||
|
|
||||||
|
offsets = torch.zeros(N + 1, device=extend_seq_lens.device, dtype=torch.int32)
|
||||||
|
offsets[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
||||||
|
output = torch.empty(total_len, device=extend_seq_lens.device, dtype=torch.int32)
|
||||||
|
|
||||||
|
BLOCK = triton.next_power_of_2(max_q_len)
|
||||||
|
grid = (N,)
|
||||||
|
|
||||||
|
seqlens_expand_kernel[grid](
|
||||||
|
extend_seq_lens,
|
||||||
|
seq_lens,
|
||||||
|
offsets,
|
||||||
|
output,
|
||||||
|
N,
|
||||||
|
BLOCK=BLOCK,
|
||||||
|
)
|
||||||
|
|
||||||
|
return output
|
||||||
@@ -0,0 +1,736 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _get_gptj_rotated_x(
|
||||||
|
x,
|
||||||
|
x_rotated_mask,
|
||||||
|
BLOCK_D: tl.constexpr,
|
||||||
|
BLOCK_D_HALF: tl.constexpr,
|
||||||
|
):
|
||||||
|
# GPT-J rotary layout:
|
||||||
|
# Pair adjacent dimensions and apply:
|
||||||
|
# [x0, x1, x2, x3] -> [-x1, x0, -x3, x2]
|
||||||
|
|
||||||
|
# Apply sign inversion on odd positions.
|
||||||
|
x_rotated = tl.where(x_rotated_mask, x, -x)
|
||||||
|
# Reshape into (D/2, 2) pairs.
|
||||||
|
x_rotated = tl.reshape(x_rotated, (BLOCK_D_HALF, 2))
|
||||||
|
# Swap each pair.
|
||||||
|
x_rotated = tl.flip(x_rotated, 1)
|
||||||
|
# Flatten back to original shape.
|
||||||
|
x_rotated = tl.reshape(x_rotated, (BLOCK_D,))
|
||||||
|
return x_rotated
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _get_neox_rotated_x(
|
||||||
|
x,
|
||||||
|
x_rotated_mask,
|
||||||
|
BLOCK_D: tl.constexpr,
|
||||||
|
BLOCK_D_HALF: tl.constexpr,
|
||||||
|
):
|
||||||
|
# GPT-NeoX rotary layout:
|
||||||
|
# Split head dimension into two halves:
|
||||||
|
# [x0, x1, x2, x3] -> [-x2, -x3, x0, x1]
|
||||||
|
|
||||||
|
# Keep first half positive, second half negative.
|
||||||
|
x_rotated = tl.where(x_rotated_mask, x, -x)
|
||||||
|
# Reshape into (2, D/2).
|
||||||
|
x_rotated = tl.reshape(x_rotated, (2, BLOCK_D_HALF))
|
||||||
|
# Reverse each half.
|
||||||
|
x_rotated = tl.flip(x_rotated, 1)
|
||||||
|
# Flatten and reverse full vector.
|
||||||
|
x_rotated = tl.reshape(x_rotated, (BLOCK_D,))
|
||||||
|
x_rotated = tl.flip(x_rotated, 0)
|
||||||
|
return x_rotated
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _unit_rope(
|
||||||
|
x_ptrs,
|
||||||
|
cos,
|
||||||
|
sin,
|
||||||
|
d_pe_offs,
|
||||||
|
IS_NEOX: tl.constexpr,
|
||||||
|
BLOCK_D_pe: tl.constexpr,
|
||||||
|
BLOCK_D_HALF_pe: tl.constexpr,
|
||||||
|
):
|
||||||
|
# Load one full attention head vector.
|
||||||
|
x_pe = tl.load(x_ptrs)
|
||||||
|
|
||||||
|
# Stage 1: Build rotated vector according to rotary layout.
|
||||||
|
if IS_NEOX:
|
||||||
|
x_rotated_mask = d_pe_offs < BLOCK_D_HALF_pe
|
||||||
|
x_pe_rotated = _get_neox_rotated_x(
|
||||||
|
x_pe, x_rotated_mask, BLOCK_D_pe, BLOCK_D_HALF_pe
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
x_rotated_mask = d_pe_offs % 2 == 0
|
||||||
|
x_pe_rotated = _get_gptj_rotated_x(
|
||||||
|
x_pe, x_rotated_mask, BLOCK_D_pe, BLOCK_D_HALF_pe
|
||||||
|
)
|
||||||
|
|
||||||
|
# Stage 2: Apply RoPE transform:
|
||||||
|
# x' = x*cos + rotate(x)*sin
|
||||||
|
x_pe = x_pe * cos + x_pe_rotated * sin
|
||||||
|
|
||||||
|
return x_pe
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _load_cos_sin(
|
||||||
|
cos_sin_ptr,
|
||||||
|
pos,
|
||||||
|
d_cos_offs,
|
||||||
|
stride_t,
|
||||||
|
stride_d,
|
||||||
|
freq_dim,
|
||||||
|
):
|
||||||
|
base = pos * stride_t
|
||||||
|
cos = tl.load(cos_sin_ptr + base + d_cos_offs * stride_d)
|
||||||
|
sin = tl.load(cos_sin_ptr + base + (d_cos_offs + freq_dim) * stride_d)
|
||||||
|
return cos, sin
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_qk_rope_reshape_and_cache_kernel(
|
||||||
|
q_ptr,
|
||||||
|
k_ptr,
|
||||||
|
v_ptr,
|
||||||
|
pos_ptr,
|
||||||
|
cos_sin_ptr,
|
||||||
|
offs_ptr,
|
||||||
|
key_cache_ptr,
|
||||||
|
value_cache_ptr,
|
||||||
|
slot_mapping_ptr,
|
||||||
|
swa_slot_mapping_ptr,
|
||||||
|
q_out_ptr,
|
||||||
|
k_out_ptr,
|
||||||
|
zeros_out_ptr,
|
||||||
|
T,
|
||||||
|
T_slot,
|
||||||
|
q_stride_t,
|
||||||
|
q_stride_h,
|
||||||
|
q_stride_d,
|
||||||
|
k_stride_t,
|
||||||
|
k_stride_h,
|
||||||
|
k_stride_d,
|
||||||
|
v_stride_t,
|
||||||
|
v_stride_h,
|
||||||
|
v_stride_d,
|
||||||
|
cos_sin_stride_t,
|
||||||
|
cos_sin_stride_d,
|
||||||
|
q_out_stride_t,
|
||||||
|
q_out_stride_h,
|
||||||
|
q_out_stride_d,
|
||||||
|
k_out_stride_t,
|
||||||
|
k_out_stride_h,
|
||||||
|
k_out_stride_d,
|
||||||
|
key_cache_stride_t,
|
||||||
|
key_cache_stride_h,
|
||||||
|
key_cache_stride_d,
|
||||||
|
key_cache_stride_b,
|
||||||
|
key_cache_stride_x,
|
||||||
|
value_cache_stride_t,
|
||||||
|
value_cache_stride_h,
|
||||||
|
value_cache_stride_d,
|
||||||
|
value_cache_stride_b,
|
||||||
|
value_cache_stride_slot_chunk,
|
||||||
|
value_cache_stride_x,
|
||||||
|
zeros_out_stride_t,
|
||||||
|
zeros_out_stride_h,
|
||||||
|
zeros_out_stride_d,
|
||||||
|
k_scale_ptr,
|
||||||
|
v_scale_ptr,
|
||||||
|
QH_PER_KH: tl.constexpr,
|
||||||
|
QH: tl.constexpr,
|
||||||
|
KH: tl.constexpr,
|
||||||
|
REUSE_FREQS_FRONT_PART: tl.constexpr,
|
||||||
|
IS_NEOX: tl.constexpr,
|
||||||
|
BLOCK_D_pe: tl.constexpr,
|
||||||
|
BLOCK_D_HALF_pe: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
X_SIZE: tl.constexpr,
|
||||||
|
FLASH_LAYOUT: tl.constexpr,
|
||||||
|
VALUE_SHUFFLE_LAYOUT: tl.constexpr = False,
|
||||||
|
HAVE_POS: tl.constexpr = False,
|
||||||
|
HAVE_K_SCALE: tl.constexpr = False,
|
||||||
|
HAVE_V_SCALE: tl.constexpr = False,
|
||||||
|
HAVE_ZEROS: tl.constexpr = False,
|
||||||
|
HAS_SWA: tl.constexpr = False,
|
||||||
|
):
|
||||||
|
# ============================================================
|
||||||
|
# Stage 0: Static stride assumptions for Triton compiler
|
||||||
|
#
|
||||||
|
# These assumptions help Triton optimize pointer arithmetic and
|
||||||
|
# simplify generated address calculations.
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
tl.assume(q_stride_t >= 0)
|
||||||
|
tl.assume(q_stride_h >= 0)
|
||||||
|
tl.assume(q_stride_d >= 0)
|
||||||
|
tl.assume(k_stride_t >= 0)
|
||||||
|
tl.assume(k_stride_h >= 0)
|
||||||
|
tl.assume(k_stride_d >= 0)
|
||||||
|
tl.assume(v_stride_t >= 0)
|
||||||
|
tl.assume(v_stride_h >= 0)
|
||||||
|
tl.assume(v_stride_d >= 0)
|
||||||
|
tl.assume(cos_sin_stride_t >= 0)
|
||||||
|
tl.assume(cos_sin_stride_d >= 0)
|
||||||
|
tl.assume(q_out_stride_t >= 0)
|
||||||
|
tl.assume(q_out_stride_h >= 0)
|
||||||
|
tl.assume(q_out_stride_d >= 0)
|
||||||
|
tl.assume(k_out_stride_t >= 0)
|
||||||
|
tl.assume(k_out_stride_h >= 0)
|
||||||
|
tl.assume(k_out_stride_d >= 0)
|
||||||
|
tl.assume(key_cache_stride_t >= 0)
|
||||||
|
tl.assume(key_cache_stride_h >= 0)
|
||||||
|
tl.assume(key_cache_stride_d >= 0)
|
||||||
|
tl.assume(key_cache_stride_b >= 0)
|
||||||
|
tl.assume(key_cache_stride_x >= 0)
|
||||||
|
tl.assume(value_cache_stride_t >= 0)
|
||||||
|
tl.assume(value_cache_stride_h >= 0)
|
||||||
|
tl.assume(value_cache_stride_d >= 0)
|
||||||
|
tl.assume(value_cache_stride_b >= 0)
|
||||||
|
tl.assume(value_cache_stride_slot_chunk >= 0)
|
||||||
|
tl.assume(value_cache_stride_x >= 0)
|
||||||
|
tl.assume(zeros_out_stride_t >= 0)
|
||||||
|
tl.assume(zeros_out_stride_h >= 0)
|
||||||
|
tl.assume(zeros_out_stride_d >= 0)
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# Stage 1: Program instance mapping
|
||||||
|
#
|
||||||
|
# Each program handles:
|
||||||
|
# - one (token, q_head) for Q path
|
||||||
|
# - selected KV ownership for cache write path
|
||||||
|
#
|
||||||
|
# pid layout:
|
||||||
|
# [0, T*QH) -> decode Q path
|
||||||
|
# [T*QH, extra KV) -> KV-only path
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
tl.assume(pid >= 0)
|
||||||
|
|
||||||
|
d_pe_offs = tl.arange(0, BLOCK_D_pe).to(tl.int64)
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# Stage 2: Main decode path (Q always active)
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
if pid < T * QH:
|
||||||
|
pid_t = pid // QH
|
||||||
|
pid_hq = pid % QH
|
||||||
|
|
||||||
|
# --------------------------------------------------------
|
||||||
|
# Stage 2.1: Compute rotary frequency offsets
|
||||||
|
#
|
||||||
|
# RoPE frequencies may be stored as:
|
||||||
|
# D/2 frequencies (shared front-half)
|
||||||
|
# D frequencies (full explicit)
|
||||||
|
# --------------------------------------------------------
|
||||||
|
|
||||||
|
if REUSE_FREQS_FRONT_PART:
|
||||||
|
if IS_NEOX:
|
||||||
|
d_cos_offs = d_pe_offs
|
||||||
|
d_cos_offs = tl.where(
|
||||||
|
(d_cos_offs >= BLOCK_D_HALF_pe) & (d_cos_offs < BLOCK_D_pe),
|
||||||
|
d_cos_offs - BLOCK_D_HALF_pe,
|
||||||
|
d_cos_offs,
|
||||||
|
).to(d_cos_offs.dtype)
|
||||||
|
# d_cos_mask = d_cos_offs < BLOCK_D_pe
|
||||||
|
else:
|
||||||
|
d_cos_offs = d_pe_offs // 2
|
||||||
|
# d_cos_mask = d_cos_offs < BLOCK_D_HALF_pe
|
||||||
|
else:
|
||||||
|
d_cos_offs = d_pe_offs
|
||||||
|
# d_cos_mask = d_cos_offs < BLOCK_D_pe
|
||||||
|
|
||||||
|
# --------------------------------------------------------
|
||||||
|
# Stage 2.2: Load token position and optional offset
|
||||||
|
#
|
||||||
|
# offs_ptr is used by chunked prefill / sliding-window decode.
|
||||||
|
# --------------------------------------------------------
|
||||||
|
pos = tl.load(pos_ptr + pid_t)
|
||||||
|
if HAVE_POS:
|
||||||
|
offset = tl.load(offs_ptr + pid_t)
|
||||||
|
pos = pos + offset
|
||||||
|
|
||||||
|
# --------------------------------------------------------
|
||||||
|
# Stage 2.3: Load cosine / sine table
|
||||||
|
# --------------------------------------------------------
|
||||||
|
# cos_offs = pos * cos_stride_t + d_cos_offs * cos_stride_d
|
||||||
|
# cos = tl.load(cos_ptr + cos_offs)
|
||||||
|
# sin = tl.load(sin_ptr + cos_offs)
|
||||||
|
|
||||||
|
freq_dim = BLOCK_D_HALF_pe if REUSE_FREQS_FRONT_PART else BLOCK_D_pe
|
||||||
|
|
||||||
|
cos, sin = _load_cos_sin(
|
||||||
|
cos_sin_ptr,
|
||||||
|
pos,
|
||||||
|
d_cos_offs,
|
||||||
|
cos_sin_stride_t,
|
||||||
|
cos_sin_stride_d,
|
||||||
|
freq_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
# --------------------------------------------------------
|
||||||
|
# Stage 2.4: Apply RoPE to Q
|
||||||
|
# --------------------------------------------------------
|
||||||
|
q_ptrs = (
|
||||||
|
q_ptr + pid_t * q_stride_t + pid_hq * q_stride_h + d_pe_offs * q_stride_d
|
||||||
|
)
|
||||||
|
q_pe = _unit_rope(
|
||||||
|
q_ptrs,
|
||||||
|
cos,
|
||||||
|
sin,
|
||||||
|
d_pe_offs,
|
||||||
|
IS_NEOX,
|
||||||
|
BLOCK_D_pe,
|
||||||
|
BLOCK_D_HALF_pe,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Store rotated Q output.
|
||||||
|
q_out_ptrs = (
|
||||||
|
q_out_ptr
|
||||||
|
+ pid_t * q_out_stride_t
|
||||||
|
+ pid_hq * q_out_stride_h
|
||||||
|
+ d_pe_offs * q_out_stride_d
|
||||||
|
)
|
||||||
|
tl.store(q_out_ptrs, q_pe.to(q_out_ptr.dtype.element_ty))
|
||||||
|
|
||||||
|
if HAVE_ZEROS:
|
||||||
|
z = tl.zeros((BLOCK_D_pe,), dtype=zeros_out_ptr.dtype.element_ty)
|
||||||
|
zeros_out_ptrs = (
|
||||||
|
zeros_out_ptr
|
||||||
|
+ pid_t * zeros_out_stride_t
|
||||||
|
+ pid_hq * zeros_out_stride_h
|
||||||
|
+ d_pe_offs * zeros_out_stride_d
|
||||||
|
)
|
||||||
|
tl.store(zeros_out_ptrs, z)
|
||||||
|
|
||||||
|
# ========================================================
|
||||||
|
# Stage 3: KV ownership path
|
||||||
|
#
|
||||||
|
# Only one Q group leader writes KV:
|
||||||
|
# pid_hq % QH_PER_KH == 0
|
||||||
|
#
|
||||||
|
# This prevents duplicated KV cache writes.
|
||||||
|
# ========================================================
|
||||||
|
|
||||||
|
if pid_hq % QH_PER_KH == 0:
|
||||||
|
# ----------------------------------------------------
|
||||||
|
# Stage 3.1: Resolve cache slot
|
||||||
|
# ----------------------------------------------------
|
||||||
|
pid_slot = tl.load(slot_mapping_ptr + pid_t).to(tl.int64)
|
||||||
|
if HAS_SWA:
|
||||||
|
pid_slot = tl.load(swa_slot_mapping_ptr + pid_slot)
|
||||||
|
|
||||||
|
# ------------------------------------------------
|
||||||
|
# Stage 3.2: Apply RoPE to K
|
||||||
|
# ------------------------------------------------
|
||||||
|
if pid_slot >= 0:
|
||||||
|
pid_t_slot = pid_slot // BLOCK_SIZE
|
||||||
|
pid_b = pid_slot % BLOCK_SIZE
|
||||||
|
pid_hk = pid_hq // QH_PER_KH
|
||||||
|
if HAVE_K_SCALE:
|
||||||
|
k_scale = tl.load(k_scale_ptr)
|
||||||
|
else:
|
||||||
|
k_scale = 1
|
||||||
|
k_ptrs = (
|
||||||
|
k_ptr
|
||||||
|
+ pid_t * k_stride_t
|
||||||
|
+ pid_hk * k_stride_h
|
||||||
|
+ d_pe_offs * k_stride_d
|
||||||
|
)
|
||||||
|
k_pe = _unit_rope(
|
||||||
|
k_ptrs,
|
||||||
|
cos,
|
||||||
|
sin,
|
||||||
|
d_pe_offs,
|
||||||
|
IS_NEOX,
|
||||||
|
BLOCK_D_pe,
|
||||||
|
BLOCK_D_HALF_pe,
|
||||||
|
)
|
||||||
|
|
||||||
|
k_out_ptrs = (
|
||||||
|
k_out_ptr
|
||||||
|
+ pid_t * k_out_stride_t
|
||||||
|
+ pid_hk * k_out_stride_h
|
||||||
|
+ d_pe_offs * k_out_stride_d
|
||||||
|
)
|
||||||
|
tl.store(k_out_ptrs, k_pe.to(k_out_ptr.dtype.element_ty))
|
||||||
|
|
||||||
|
# ------------------------------------------------
|
||||||
|
# Stage 3.3: Optional fp8 scaling before cache
|
||||||
|
# ------------------------------------------------
|
||||||
|
|
||||||
|
k_scale_rcprl = 1 / k_scale
|
||||||
|
k_pe = k_pe * k_scale_rcprl
|
||||||
|
|
||||||
|
# ------------------------------------------------
|
||||||
|
# Stage 3.4: Write K cache
|
||||||
|
#
|
||||||
|
# Two layouts supported:
|
||||||
|
# FLASH_LAYOUT
|
||||||
|
# paged KV layout
|
||||||
|
# ------------------------------------------------
|
||||||
|
|
||||||
|
if FLASH_LAYOUT:
|
||||||
|
k_out_ptrs = (
|
||||||
|
key_cache_ptr
|
||||||
|
+ pid_t_slot * key_cache_stride_t
|
||||||
|
+ pid_b * key_cache_stride_b
|
||||||
|
+ pid_hk * key_cache_stride_h
|
||||||
|
+ d_pe_offs * key_cache_stride_d
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
k_pe = tl.reshape(k_pe, (BLOCK_D_pe // X_SIZE, X_SIZE))
|
||||||
|
dx_offs = tl.arange(0, BLOCK_D_pe // X_SIZE).to(tl.int64)
|
||||||
|
x_offs = tl.arange(0, X_SIZE).to(tl.int64)
|
||||||
|
k_out_ptrs = (
|
||||||
|
key_cache_ptr
|
||||||
|
+ pid_t_slot * key_cache_stride_t
|
||||||
|
+ pid_hk * key_cache_stride_h
|
||||||
|
+ dx_offs[:, None] * key_cache_stride_d
|
||||||
|
+ pid_b * key_cache_stride_b
|
||||||
|
+ x_offs[None, :] * key_cache_stride_x
|
||||||
|
)
|
||||||
|
|
||||||
|
tl.store(k_out_ptrs, k_pe.to(key_cache_ptr.dtype.element_ty))
|
||||||
|
|
||||||
|
# ------------------------------------------------
|
||||||
|
# Stage 3.5: Write V cache
|
||||||
|
#
|
||||||
|
# Supports:
|
||||||
|
# normal layout
|
||||||
|
# shuffle layout
|
||||||
|
# ------------------------------------------------
|
||||||
|
|
||||||
|
v_ptrs = (
|
||||||
|
v_ptr
|
||||||
|
+ pid_t * v_stride_t
|
||||||
|
+ pid_hk * v_stride_h
|
||||||
|
+ d_pe_offs * v_stride_d
|
||||||
|
)
|
||||||
|
if HAVE_V_SCALE:
|
||||||
|
v_scale = tl.load(v_scale_ptr)
|
||||||
|
else:
|
||||||
|
v_scale = 1
|
||||||
|
v_scale_rcprl = 1 / v_scale
|
||||||
|
v = tl.load(v_ptrs) * v_scale_rcprl
|
||||||
|
if VALUE_SHUFFLE_LAYOUT:
|
||||||
|
slot_chunk = pid_b // X_SIZE
|
||||||
|
x_off = pid_b % X_SIZE
|
||||||
|
v_out_ptrs = (
|
||||||
|
value_cache_ptr
|
||||||
|
+ pid_t_slot * value_cache_stride_t
|
||||||
|
+ pid_hk * value_cache_stride_h
|
||||||
|
+ slot_chunk * value_cache_stride_slot_chunk
|
||||||
|
+ d_pe_offs.to(tl.int64) * value_cache_stride_d
|
||||||
|
+ x_off * value_cache_stride_x
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
v_out_ptrs = (
|
||||||
|
value_cache_ptr
|
||||||
|
+ pid_t_slot * value_cache_stride_t
|
||||||
|
+ pid_hk * value_cache_stride_h
|
||||||
|
+ d_pe_offs.to(tl.int64) * value_cache_stride_d
|
||||||
|
+ pid_b * value_cache_stride_b
|
||||||
|
)
|
||||||
|
tl.store(v_out_ptrs, v.to(value_cache_ptr.dtype.element_ty))
|
||||||
|
# ============================================================
|
||||||
|
# Stage 4: Extra KV-only path
|
||||||
|
#
|
||||||
|
# Handles tokens that only require cache update:
|
||||||
|
# T_slot > T
|
||||||
|
#
|
||||||
|
# No Q / no RoPE on Q branch.
|
||||||
|
# ============================================================
|
||||||
|
else:
|
||||||
|
pid = pid - T * QH + T * KH
|
||||||
|
if pid < T_slot * KH:
|
||||||
|
pid_t = pid // KH
|
||||||
|
pid_hk = pid % KH
|
||||||
|
pid_slot = tl.load(slot_mapping_ptr + pid_t).to(tl.int64)
|
||||||
|
if HAS_SWA:
|
||||||
|
pid_slot = tl.load(swa_slot_mapping_ptr + pid_slot)
|
||||||
|
|
||||||
|
if pid_slot >= 0:
|
||||||
|
pid_t_slot = pid_slot // BLOCK_SIZE
|
||||||
|
pid_b = pid_slot % BLOCK_SIZE
|
||||||
|
if HAVE_K_SCALE:
|
||||||
|
k_scale = tl.load(k_scale_ptr)
|
||||||
|
else:
|
||||||
|
k_scale = 1
|
||||||
|
k_ptrs = (
|
||||||
|
k_ptr
|
||||||
|
+ pid_t * k_stride_t
|
||||||
|
+ pid_hk * k_stride_h
|
||||||
|
+ d_pe_offs * k_stride_d
|
||||||
|
)
|
||||||
|
|
||||||
|
k_pe = tl.load(k_ptrs)
|
||||||
|
|
||||||
|
k_out_ptrs = (
|
||||||
|
k_out_ptr
|
||||||
|
+ pid_t * k_out_stride_t
|
||||||
|
+ pid_hk * k_out_stride_h
|
||||||
|
+ d_pe_offs * k_out_stride_d
|
||||||
|
)
|
||||||
|
tl.store(k_out_ptrs, k_pe.to(k_out_ptr.dtype.element_ty))
|
||||||
|
|
||||||
|
k_scale_rcprl = 1 / k_scale
|
||||||
|
k_pe = k_pe * k_scale_rcprl
|
||||||
|
|
||||||
|
if FLASH_LAYOUT:
|
||||||
|
k_out_ptrs = (
|
||||||
|
key_cache_ptr
|
||||||
|
+ pid_t_slot * key_cache_stride_t
|
||||||
|
+ d_pe_offs * key_cache_stride_d
|
||||||
|
+ pid_b * key_cache_stride_b
|
||||||
|
+ pid_hk * key_cache_stride_h
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
k_pe = tl.reshape(k_pe, (BLOCK_D_pe // X_SIZE, X_SIZE))
|
||||||
|
dx_offs = tl.arange(0, BLOCK_D_pe // X_SIZE).to(tl.int64)
|
||||||
|
x_offs = tl.arange(0, X_SIZE).to(tl.int64)
|
||||||
|
k_out_ptrs = (
|
||||||
|
key_cache_ptr
|
||||||
|
+ pid_t_slot * key_cache_stride_t
|
||||||
|
+ pid_hk * key_cache_stride_h
|
||||||
|
+ dx_offs[:, None] * key_cache_stride_d
|
||||||
|
+ pid_b * key_cache_stride_b
|
||||||
|
+ x_offs[None, :] * key_cache_stride_x
|
||||||
|
)
|
||||||
|
tl.store(k_out_ptrs, k_pe.to(key_cache_ptr.dtype.element_ty))
|
||||||
|
|
||||||
|
v_ptrs = (
|
||||||
|
v_ptr
|
||||||
|
+ pid_t * v_stride_t
|
||||||
|
+ pid_hk * v_stride_h
|
||||||
|
+ d_pe_offs * v_stride_d
|
||||||
|
)
|
||||||
|
if HAVE_V_SCALE:
|
||||||
|
v_scale = tl.load(v_scale_ptr)
|
||||||
|
else:
|
||||||
|
v_scale = 1
|
||||||
|
v_scale_rcprl = 1 / v_scale
|
||||||
|
v = tl.load(v_ptrs) * v_scale_rcprl
|
||||||
|
if VALUE_SHUFFLE_LAYOUT:
|
||||||
|
slot_chunk = pid_b // X_SIZE
|
||||||
|
x_off = pid_b % X_SIZE
|
||||||
|
v_out_ptrs = (
|
||||||
|
value_cache_ptr
|
||||||
|
+ pid_t_slot * value_cache_stride_t
|
||||||
|
+ pid_hk * value_cache_stride_h
|
||||||
|
+ slot_chunk * value_cache_stride_slot_chunk
|
||||||
|
+ d_pe_offs * value_cache_stride_d
|
||||||
|
+ x_off * value_cache_stride_x
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
v_out_ptrs = (
|
||||||
|
value_cache_ptr
|
||||||
|
+ pid_t_slot * value_cache_stride_t
|
||||||
|
+ pid_hk * value_cache_stride_h
|
||||||
|
+ d_pe_offs * value_cache_stride_d
|
||||||
|
+ pid_b * value_cache_stride_b
|
||||||
|
)
|
||||||
|
tl.store(v_out_ptrs, v.to(value_cache_ptr.dtype.element_ty))
|
||||||
|
|
||||||
|
|
||||||
|
def fused_qk_rope_reshape_and_cache(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
key_cache: torch.Tensor,
|
||||||
|
value_cache: torch.Tensor,
|
||||||
|
slot_mapping: torch.Tensor,
|
||||||
|
pos: torch.Tensor,
|
||||||
|
cos_sin: torch.Tensor,
|
||||||
|
k_scale: torch.Tensor,
|
||||||
|
v_scale: torch.Tensor,
|
||||||
|
is_neox: bool,
|
||||||
|
flash_layout: bool,
|
||||||
|
apply_scale: bool = True,
|
||||||
|
offs: torch.Tensor = None,
|
||||||
|
q_out: torch.Tensor = None,
|
||||||
|
k_out: torch.Tensor = None,
|
||||||
|
output_zeros: bool = True,
|
||||||
|
zeros_out: torch.Tensor = None,
|
||||||
|
swa_slot_mapping=None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Perform RoPE on q and k and along the last dimension and copy k and v in to key_cache and value_cache inplace
|
||||||
|
|
||||||
|
Key parameters:
|
||||||
|
- q: shape (T, QH, D).
|
||||||
|
- k: shape (T_slot, KH, D).
|
||||||
|
- v: shape (T_slot, KH, D).
|
||||||
|
- if flash_layout:
|
||||||
|
- key_cache: shape (T_cache, block_size, KH, D).
|
||||||
|
- value_cache: shape (T_cache, block_size, KH, D).
|
||||||
|
- else:
|
||||||
|
- key_cache: shape (T_cache, KH, D // x, block_size, x).
|
||||||
|
- value_cache: shape (T_cache, KH, D, block_size).
|
||||||
|
- slot_mapping: shape (T_slot, ).
|
||||||
|
|
||||||
|
T is the number of decode tokens, T_cahce * block_size is the max number of tokens of kv_cache
|
||||||
|
QH must be multiple of KH
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
- q_out: same shape as input q.
|
||||||
|
- k_out: same shape as input k.
|
||||||
|
- key_cache: same shape as input key_cache (inplace).
|
||||||
|
- value_cache: same shape as input value_cache (inplace).
|
||||||
|
- zeros_out: same shape as input q.
|
||||||
|
"""
|
||||||
|
|
||||||
|
t, qh, d = q.shape
|
||||||
|
tk, kh, dk = k.shape
|
||||||
|
tv, vh, dv = v.shape
|
||||||
|
if flash_layout:
|
||||||
|
t_cache, block_size, kh_cache, dk_cache = key_cache.shape
|
||||||
|
t_cache_v, block_size_v, vh_cache, dv_cache = value_cache.shape
|
||||||
|
value_shuffle_layout = False
|
||||||
|
else:
|
||||||
|
t_cache, kh_cache, dkx_cache, block_size, x_cache = key_cache.shape
|
||||||
|
if value_cache.ndim == 5:
|
||||||
|
# value_cache shuffle: (num_blocks, num_kv_heads, block_size // x, head_size, x)
|
||||||
|
t_cache_v, vh_cache, slot_chunk_v, dv_cache, x_v = value_cache.shape
|
||||||
|
value_shuffle_layout = True
|
||||||
|
block_size_v = slot_chunk_v * x_v
|
||||||
|
assert block_size_v == block_size and x_v == x_cache, (
|
||||||
|
f"value_cache shuffle (T,KH,block_size//x,D,x) must match key: "
|
||||||
|
f"{block_size_v=} {block_size=} {x_v=} {x_cache=}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
t_cache_v, vh_cache, dv_cache, block_size_v = value_cache.shape
|
||||||
|
value_shuffle_layout = False
|
||||||
|
(t_slot,) = slot_mapping.shape
|
||||||
|
|
||||||
|
assert (
|
||||||
|
t == tk == tv and t_slot <= tk
|
||||||
|
), f"Number of tokens should be identical for q, kand v. The number of tokens of slot_mapping should no more than that of q, k and v, {t=} {tk=} {tv=} {t_slot=}"
|
||||||
|
assert (
|
||||||
|
block_size == block_size_v
|
||||||
|
), f"block size should be identical for key_cache, and value_cache {block_size} {block_size_v}"
|
||||||
|
assert (
|
||||||
|
kh == vh == kh_cache == vh_cache
|
||||||
|
), "KV head should be identical for k, v, key_cache, and value_cache"
|
||||||
|
assert (
|
||||||
|
t_cache == t_cache_v
|
||||||
|
), "Number of tokens should be identical for key_cache, and value_cache"
|
||||||
|
if flash_layout:
|
||||||
|
assert (
|
||||||
|
d == dk == dv == dk_cache == dv_cache
|
||||||
|
), "D dimension should be identical for q, k, and v"
|
||||||
|
else:
|
||||||
|
assert (
|
||||||
|
d == dk == dv == dkx_cache * x_cache == dv_cache
|
||||||
|
), "D dimension should be identical for q, k, and v"
|
||||||
|
assert x_cache == triton.next_power_of_2(x_cache), "x_size should be power of 2"
|
||||||
|
|
||||||
|
assert d == triton.next_power_of_2(d), "D dimension should be power of 2"
|
||||||
|
assert block_size == triton.next_power_of_2(
|
||||||
|
block_size
|
||||||
|
), "block_size should be power of 2"
|
||||||
|
assert qh % kh == 0, "Q heads must be multiple of H heads"
|
||||||
|
d_freq = cos_sin.shape[-1] // 2
|
||||||
|
assert (d_freq == d // 2) or (
|
||||||
|
d_freq == d
|
||||||
|
), "cos/sin last dim should be the same or half of the qk last dim"
|
||||||
|
reuse_freqs_front_part = d_freq == d // 2
|
||||||
|
|
||||||
|
if q_out is None:
|
||||||
|
q_out = torch.empty((t, qh, d), dtype=q.dtype, device=q.device)
|
||||||
|
|
||||||
|
if k_out is None:
|
||||||
|
k_out = torch.empty((tk, kh, dk), dtype=k.dtype, device=q.device)
|
||||||
|
|
||||||
|
if zeros_out is not None:
|
||||||
|
tz, qhz, dz = zeros_out.shape
|
||||||
|
assert (
|
||||||
|
t == tz and qh == qhz and d == dz
|
||||||
|
), f"q and zeros shape mismatch {q.shape=} {zeros_out.shape=}"
|
||||||
|
output_zeros = True
|
||||||
|
elif output_zeros:
|
||||||
|
zeros_out = torch.empty((t, qh, d), dtype=q.dtype, device=q.device)
|
||||||
|
else:
|
||||||
|
zeros_out = None
|
||||||
|
|
||||||
|
n_pid = t * qh + (t_slot - t) * kh if t_slot >= t else t * qh
|
||||||
|
grid = (n_pid, 1, 1)
|
||||||
|
_fused_qk_rope_reshape_and_cache_kernel[grid](
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
pos,
|
||||||
|
cos_sin,
|
||||||
|
offs,
|
||||||
|
key_cache,
|
||||||
|
value_cache,
|
||||||
|
slot_mapping,
|
||||||
|
swa_slot_mapping,
|
||||||
|
q_out,
|
||||||
|
k_out,
|
||||||
|
zeros_out,
|
||||||
|
t,
|
||||||
|
t_slot,
|
||||||
|
*q.stride(),
|
||||||
|
*k.stride(),
|
||||||
|
*v.stride(),
|
||||||
|
cos_sin.stride(0),
|
||||||
|
cos_sin.stride(-1),
|
||||||
|
*q_out.stride(),
|
||||||
|
*k_out.stride(),
|
||||||
|
key_cache.stride(0) if not flash_layout else key_cache.stride(0),
|
||||||
|
key_cache.stride(1) if not flash_layout else key_cache.stride(2),
|
||||||
|
key_cache.stride(2) if not flash_layout else key_cache.stride(3),
|
||||||
|
key_cache.stride(3) if not flash_layout else key_cache.stride(1),
|
||||||
|
key_cache.stride(4) if not flash_layout else 0,
|
||||||
|
value_cache.stride(0) if not flash_layout else value_cache.stride(0),
|
||||||
|
value_cache.stride(1) if not flash_layout else value_cache.stride(2),
|
||||||
|
(
|
||||||
|
value_cache.stride(3)
|
||||||
|
if (not flash_layout and value_shuffle_layout)
|
||||||
|
else (value_cache.stride(2) if not flash_layout else value_cache.stride(3))
|
||||||
|
),
|
||||||
|
(
|
||||||
|
0
|
||||||
|
if (not flash_layout and value_shuffle_layout)
|
||||||
|
else (value_cache.stride(3) if not flash_layout else value_cache.stride(1))
|
||||||
|
),
|
||||||
|
value_cache.stride(2) if (not flash_layout and value_shuffle_layout) else 0,
|
||||||
|
value_cache.stride(4) if (not flash_layout and value_shuffle_layout) else 0,
|
||||||
|
zeros_out.stride(0) if zeros_out is not None else 0,
|
||||||
|
zeros_out.stride(1) if zeros_out is not None else 0,
|
||||||
|
zeros_out.stride(2) if zeros_out is not None else 0,
|
||||||
|
k_scale_ptr=k_scale,
|
||||||
|
v_scale_ptr=v_scale,
|
||||||
|
QH_PER_KH=qh // kh,
|
||||||
|
QH=qh,
|
||||||
|
KH=kh,
|
||||||
|
REUSE_FREQS_FRONT_PART=reuse_freqs_front_part,
|
||||||
|
IS_NEOX=is_neox,
|
||||||
|
BLOCK_D_pe=d,
|
||||||
|
BLOCK_D_HALF_pe=d // 2,
|
||||||
|
BLOCK_SIZE=block_size,
|
||||||
|
X_SIZE=x_cache if not flash_layout else 0,
|
||||||
|
FLASH_LAYOUT=flash_layout,
|
||||||
|
VALUE_SHUFFLE_LAYOUT=value_shuffle_layout,
|
||||||
|
HAVE_POS=(offs is not None),
|
||||||
|
HAVE_K_SCALE=(k_scale is not None and apply_scale),
|
||||||
|
HAVE_V_SCALE=(v_scale is not None and apply_scale),
|
||||||
|
HAVE_ZEROS=output_zeros,
|
||||||
|
HAS_SWA=(swa_slot_mapping is not None),
|
||||||
|
num_warps=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
if zeros_out is not None:
|
||||||
|
return q_out.view(-1, qh * d), k_out, key_cache, value_cache, zeros_out
|
||||||
|
return q_out.view(-1, qh * d), k_out, key_cache, value_cache
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -6,10 +6,12 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
||||||
|
create_flashinfer_kv_indices_triton,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
||||||
@@ -22,58 +24,6 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def get_num_kv_splits_triton(
|
|
||||||
num_kv_splits_ptr,
|
|
||||||
seq_lens_ptr,
|
|
||||||
num_seq,
|
|
||||||
num_group,
|
|
||||||
num_head,
|
|
||||||
num_kv_head,
|
|
||||||
max_kv_splits,
|
|
||||||
device_core_count,
|
|
||||||
MAX_NUM_SEQ: tl.constexpr,
|
|
||||||
):
|
|
||||||
# TODO: this method is tunable, we need more online serving data to tune it
|
|
||||||
offs_seq = tl.arange(0, MAX_NUM_SEQ)
|
|
||||||
mask_seq = offs_seq < num_seq
|
|
||||||
|
|
||||||
seq_lens = tl.load(seq_lens_ptr + offs_seq, mask=mask_seq, other=0)
|
|
||||||
max_seq_len = tl.max(seq_lens)
|
|
||||||
seq_lens = tl.load(seq_lens_ptr + offs_seq, mask=mask_seq, other=max_seq_len)
|
|
||||||
min_seq_len = tl.min(seq_lens)
|
|
||||||
if max_seq_len * 8 < min_seq_len * 10:
|
|
||||||
min_seq_len = max_seq_len
|
|
||||||
max_kv_splits_1 = tl.minimum(tl.cdiv(max_seq_len, min_seq_len), max_kv_splits)
|
|
||||||
kv_chunk_size_1 = tl.cdiv(max_seq_len, max_kv_splits_1)
|
|
||||||
|
|
||||||
# NOTE: this is a hack to let num_kv_split grows up with seqlen gradually
|
|
||||||
ext_seq_len = tl.cast(max_seq_len, tl.float32) / 64.0
|
|
||||||
ext_device_core_count = tl.cast(
|
|
||||||
device_core_count * tl.maximum(tl.log2(ext_seq_len), 1.0), tl.int32
|
|
||||||
)
|
|
||||||
block_h, num_kv_group = 16, num_head // num_kv_head
|
|
||||||
if num_kv_group == 1:
|
|
||||||
token_grid = num_seq * num_group * num_head
|
|
||||||
else:
|
|
||||||
# from triton_ops/decode_attention.py:_decode_grouped_att_m_fwd
|
|
||||||
block_h = tl.minimum(block_h, num_kv_group)
|
|
||||||
token_grid = num_seq * num_group * tl.cdiv(num_head, block_h)
|
|
||||||
max_kv_splits_2 = tl.minimum(
|
|
||||||
tl.cdiv(ext_device_core_count, token_grid), max_kv_splits
|
|
||||||
)
|
|
||||||
kv_chunk_size_2 = tl.cdiv(max_seq_len, max_kv_splits_2)
|
|
||||||
|
|
||||||
num_kv_splits = tl.maximum(
|
|
||||||
tl.cdiv(seq_lens, kv_chunk_size_1), tl.cdiv(seq_lens, kv_chunk_size_2)
|
|
||||||
)
|
|
||||||
|
|
||||||
offs_token = offs_seq * num_group
|
|
||||||
mask_token = offs_token < num_seq * num_group
|
|
||||||
for i in range(0, num_group):
|
|
||||||
tl.store(num_kv_splits_ptr + i + offs_token, num_kv_splits, mask=mask_token)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ForwardMetadata:
|
class ForwardMetadata:
|
||||||
attn_logits: torch.Tensor
|
attn_logits: torch.Tensor
|
||||||
|
|||||||
@@ -4,74 +4,13 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.srt.layers.triton_ops.softcap import softcap_out as fused_softcap
|
||||||
from sglang.srt.utils import is_hip
|
from sglang.srt.utils import is_hip
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
|
|
||||||
|
|
||||||
fused_softcap_autotune = triton.autotune(
|
|
||||||
configs=[
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 128}, num_warps=4),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 128}, num_warps=8),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 128}, num_warps=16),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 256}, num_warps=4),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 256}, num_warps=8),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 512}, num_warps=4),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 512}, num_warps=8),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 512}, num_warps=16),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=4),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=8),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=16),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=32),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 2048}, num_warps=32),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 4096}, num_warps=32),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 8192}, num_warps=32),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 16384}, num_warps=32),
|
|
||||||
triton.Config(kwargs={"BLOCK_SIZE": 32768}, num_warps=32),
|
|
||||||
],
|
|
||||||
key=["n_ele"],
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def fused_softcap_kernel(
|
|
||||||
output_ptr,
|
|
||||||
input_ptr,
|
|
||||||
n_ele,
|
|
||||||
softcap_const: tl.constexpr,
|
|
||||||
BLOCK_SIZE: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
block_start = pid * BLOCK_SIZE
|
|
||||||
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
|
||||||
mask = offsets < n_ele
|
|
||||||
x = tl.load(input_ptr + offsets, mask=mask)
|
|
||||||
fx = x.to(tl.float32)
|
|
||||||
fxs = fx / softcap_const
|
|
||||||
exped = tl.exp(2 * fxs)
|
|
||||||
top = exped - 1
|
|
||||||
bottom = exped + 1
|
|
||||||
output = top / bottom * softcap_const
|
|
||||||
tl.store(output_ptr + offsets, output, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
fused_softcap_kernel_autotuned = fused_softcap_autotune(fused_softcap_kernel)
|
|
||||||
|
|
||||||
|
|
||||||
def fused_softcap(x, softcap_const, autotune=False):
|
|
||||||
output = torch.empty_like(x, dtype=torch.float32)
|
|
||||||
n_elements = output.numel()
|
|
||||||
if autotune:
|
|
||||||
grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),)
|
|
||||||
fused_softcap_kernel_autotuned[grid](output, x, n_elements, softcap_const)
|
|
||||||
else:
|
|
||||||
fused_softcap_kernel[(triton.cdiv(n_elements, 128),)](
|
|
||||||
output, x, n_elements, softcap_const, BLOCK_SIZE=128, num_warps=8
|
|
||||||
)
|
|
||||||
return output
|
|
||||||
|
|
||||||
|
|
||||||
# cast to float + softcap
|
# cast to float + softcap
|
||||||
class Softcap:
|
class Softcap:
|
||||||
def __init__(self, softcap_const: float):
|
def __init__(self, softcap_const: float):
|
||||||
|
|||||||
@@ -18,10 +18,7 @@ import logging
|
|||||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from triton.language.extra import libdevice
|
|
||||||
|
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
@@ -41,6 +38,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
get_dp_dtype,
|
get_dp_dtype,
|
||||||
get_dp_hidden_size,
|
get_dp_hidden_size,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.triton_ops.softcap import softcap_inplace_logits as fused_softcap
|
||||||
from sglang.srt.layers.utils.logprob import (
|
from sglang.srt.layers.utils.logprob import (
|
||||||
InputLogprobsResult,
|
InputLogprobsResult,
|
||||||
get_token_ids_logprobs_chunk,
|
get_token_ids_logprobs_chunk,
|
||||||
@@ -208,7 +206,6 @@ class LogitsMetadata:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def compute_dp_attention_metadata(self):
|
def compute_dp_attention_metadata(self):
|
||||||
|
|
||||||
cumtokens = torch.cumsum(self.global_num_tokens_for_logprob_gpu, dim=0)
|
cumtokens = torch.cumsum(self.global_num_tokens_for_logprob_gpu, dim=0)
|
||||||
dp_rank = get_attention_dp_rank()
|
dp_rank = get_attention_dp_rank()
|
||||||
if dp_rank == 0:
|
if dp_rank == 0:
|
||||||
@@ -1073,55 +1070,3 @@ class LogitsProcessor(nn.Module):
|
|||||||
# They should be moved to GenerationBatchResult to keep this class clean.
|
# They should be moved to GenerationBatchResult to keep this class clean.
|
||||||
mm_input_embeds=logits_metadata.mm_input_embeds,
|
mm_input_embeds=logits_metadata.mm_input_embeds,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def fused_softcap_kernel(
|
|
||||||
full_logits_ptr,
|
|
||||||
softcapping_value,
|
|
||||||
ncols,
|
|
||||||
row_stride,
|
|
||||||
BLOCK_SIZE: tl.constexpr,
|
|
||||||
):
|
|
||||||
row = tl.program_id(1).to(tl.int64)
|
|
||||||
pid = tl.program_id(0).to(tl.int64)
|
|
||||||
block_start = pid * BLOCK_SIZE
|
|
||||||
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
|
||||||
mask = offsets < ncols
|
|
||||||
|
|
||||||
# Load values
|
|
||||||
row_ptr = full_logits_ptr + row * row_stride
|
|
||||||
x = tl.load(row_ptr + offsets, mask=mask)
|
|
||||||
|
|
||||||
# Perform operations in-place
|
|
||||||
x = x / softcapping_value
|
|
||||||
x = libdevice.tanh(x)
|
|
||||||
x = x * softcapping_value
|
|
||||||
|
|
||||||
# Store result
|
|
||||||
tl.store(row_ptr + offsets, x, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
def fused_softcap(full_logits, final_logit_softcapping):
|
|
||||||
if full_logits.is_contiguous():
|
|
||||||
nrows, ncols = 1, full_logits.numel()
|
|
||||||
row_stride = ncols
|
|
||||||
else:
|
|
||||||
assert full_logits.ndim == 2, "non-contiguous softcap requires 2D tensor"
|
|
||||||
assert (
|
|
||||||
full_logits.stride(1) == 1
|
|
||||||
), "non-contiguous softcap requires contiguous columns"
|
|
||||||
nrows, ncols = full_logits.shape
|
|
||||||
row_stride = full_logits.stride(0)
|
|
||||||
|
|
||||||
BLOCK_SIZE = 1024
|
|
||||||
grid = ((ncols + BLOCK_SIZE - 1) // BLOCK_SIZE, nrows)
|
|
||||||
|
|
||||||
fused_softcap_kernel[grid](
|
|
||||||
full_logits_ptr=full_logits,
|
|
||||||
softcapping_value=final_logit_softcapping,
|
|
||||||
ncols=ncols,
|
|
||||||
row_stride=row_stride,
|
|
||||||
BLOCK_SIZE=BLOCK_SIZE,
|
|
||||||
)
|
|
||||||
return full_logits
|
|
||||||
|
|||||||
@@ -0,0 +1,120 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
from triton.language.extra import libdevice
|
||||||
|
|
||||||
|
softcap_out_autotune = triton.autotune(
|
||||||
|
configs=[
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 128}, num_warps=4),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 128}, num_warps=8),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 128}, num_warps=16),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 256}, num_warps=4),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 256}, num_warps=8),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 512}, num_warps=4),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 512}, num_warps=8),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 512}, num_warps=16),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=4),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=8),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=16),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=32),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 2048}, num_warps=32),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 4096}, num_warps=32),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 8192}, num_warps=32),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 16384}, num_warps=32),
|
||||||
|
triton.Config(kwargs={"BLOCK_SIZE": 32768}, num_warps=32),
|
||||||
|
],
|
||||||
|
key=["n_ele"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def softcap_out_kernel(
|
||||||
|
output_ptr,
|
||||||
|
input_ptr,
|
||||||
|
n_ele,
|
||||||
|
softcap_const: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
block_start = pid * BLOCK_SIZE
|
||||||
|
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offsets < n_ele
|
||||||
|
x = tl.load(input_ptr + offsets, mask=mask)
|
||||||
|
fx = x.to(tl.float32)
|
||||||
|
fxs = fx / softcap_const
|
||||||
|
exped = tl.exp(2 * fxs)
|
||||||
|
top = exped - 1
|
||||||
|
bottom = exped + 1
|
||||||
|
output = top / bottom * softcap_const
|
||||||
|
tl.store(output_ptr + offsets, output, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
softcap_out_kernel_autotuned = softcap_out_autotune(softcap_out_kernel)
|
||||||
|
|
||||||
|
|
||||||
|
def softcap_out(x, softcap_const, autotune=False):
|
||||||
|
output = torch.empty_like(x, dtype=torch.float32)
|
||||||
|
n_elements = output.numel()
|
||||||
|
if autotune:
|
||||||
|
|
||||||
|
def grid(meta):
|
||||||
|
return (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),)
|
||||||
|
|
||||||
|
softcap_out_kernel_autotuned[grid](output, x, n_elements, softcap_const)
|
||||||
|
else:
|
||||||
|
softcap_out_kernel[(triton.cdiv(n_elements, 128),)](
|
||||||
|
output, x, n_elements, softcap_const, BLOCK_SIZE=128, num_warps=8
|
||||||
|
)
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def softcap_inplace_logits_kernel(
|
||||||
|
full_logits_ptr,
|
||||||
|
softcapping_value,
|
||||||
|
ncols,
|
||||||
|
row_stride,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
row = tl.program_id(1).to(tl.int64)
|
||||||
|
pid = tl.program_id(0).to(tl.int64)
|
||||||
|
block_start = pid * BLOCK_SIZE
|
||||||
|
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offsets < ncols
|
||||||
|
|
||||||
|
# Load values
|
||||||
|
row_ptr = full_logits_ptr + row * row_stride
|
||||||
|
x = tl.load(row_ptr + offsets, mask=mask)
|
||||||
|
|
||||||
|
# Perform operations in-place
|
||||||
|
x = x / softcapping_value
|
||||||
|
x = libdevice.tanh(x)
|
||||||
|
x = x * softcapping_value
|
||||||
|
|
||||||
|
# Store result
|
||||||
|
tl.store(row_ptr + offsets, x, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def softcap_inplace_logits(full_logits, final_logit_softcapping):
|
||||||
|
if full_logits.is_contiguous():
|
||||||
|
nrows, ncols = 1, full_logits.numel()
|
||||||
|
row_stride = ncols
|
||||||
|
else:
|
||||||
|
assert full_logits.ndim == 2, "non-contiguous softcap requires 2D tensor"
|
||||||
|
assert (
|
||||||
|
full_logits.stride(1) == 1
|
||||||
|
), "non-contiguous softcap requires contiguous columns"
|
||||||
|
nrows, ncols = full_logits.shape
|
||||||
|
row_stride = full_logits.stride(0)
|
||||||
|
|
||||||
|
BLOCK_SIZE = 1024
|
||||||
|
grid = ((ncols + BLOCK_SIZE - 1) // BLOCK_SIZE, nrows)
|
||||||
|
|
||||||
|
softcap_inplace_logits_kernel[grid](
|
||||||
|
full_logits_ptr=full_logits,
|
||||||
|
softcapping_value=final_logit_softcapping,
|
||||||
|
ncols=ncols,
|
||||||
|
row_stride=row_stride,
|
||||||
|
BLOCK_SIZE=BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
return full_logits
|
||||||
@@ -19,13 +19,16 @@ from __future__ import annotations
|
|||||||
Page-aligned memory pool.
|
Page-aligned memory pool.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
||||||
|
from sglang.srt.mem_cache.triton_ops.allocator import (
|
||||||
|
alloc_decode_kernel,
|
||||||
|
alloc_extend_kernel,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var, get_num_new_pages, next_power_of_2
|
from sglang.srt.utils import get_bool_env_var, get_num_new_pages, next_power_of_2
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -92,128 +95,6 @@ def alloc_extend_naive(
|
|||||||
).view(-1)
|
).view(-1)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def alloc_extend_kernel(
|
|
||||||
pre_lens_ptr,
|
|
||||||
seq_lens_ptr,
|
|
||||||
last_loc_ptr,
|
|
||||||
free_page_ptr,
|
|
||||||
out_indices,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
page_size: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
|
|
||||||
load_offset = tl.arange(0, bs_upper)
|
|
||||||
seq_lens = tl.load(seq_lens_ptr + load_offset, mask=load_offset <= pid)
|
|
||||||
pre_lens = tl.load(pre_lens_ptr + load_offset, mask=load_offset <= pid)
|
|
||||||
extend_lens = seq_lens - pre_lens
|
|
||||||
|
|
||||||
seq_len = tl.load(seq_lens_ptr + pid)
|
|
||||||
pre_len = tl.load(pre_lens_ptr + pid)
|
|
||||||
extend_len = seq_len - pre_len
|
|
||||||
|
|
||||||
sum_extend_lens = tl.sum(extend_lens)
|
|
||||||
output_start_loc = sum_extend_lens - extend_len
|
|
||||||
|
|
||||||
num_pages_after = (seq_lens + page_size - 1) // page_size
|
|
||||||
num_pages_before = (pre_lens + page_size - 1) // page_size
|
|
||||||
num_new_pages = num_pages_after - num_pages_before
|
|
||||||
|
|
||||||
num_page_start_loc_self = (seq_len + page_size - 1) // page_size - (
|
|
||||||
pre_len + page_size - 1
|
|
||||||
) // page_size
|
|
||||||
sum_num_new_pages = tl.sum(num_new_pages)
|
|
||||||
new_page_start_loc = sum_num_new_pages - num_page_start_loc_self
|
|
||||||
|
|
||||||
# Part 1: fill the old partial page
|
|
||||||
last_loc = tl.load(last_loc_ptr + pid)
|
|
||||||
num_part1 = (
|
|
||||||
min(seq_len, (pre_len + page_size - 1) // page_size * page_size) - pre_len
|
|
||||||
)
|
|
||||||
offset_one_page = tl.arange(0, page_size)
|
|
||||||
tl.store(
|
|
||||||
out_indices + output_start_loc + offset_one_page,
|
|
||||||
last_loc + 1 + offset_one_page,
|
|
||||||
mask=offset_one_page < num_part1,
|
|
||||||
)
|
|
||||||
if pre_len + num_part1 == seq_len:
|
|
||||||
return
|
|
||||||
|
|
||||||
# Part 2: fill the new full pages using a dynamic blocked loop.
|
|
||||||
# The loop bound is derived from num_part2 (runtime value), so Triton
|
|
||||||
# generates a real loop instead of unrolling — no constexpr dependency
|
|
||||||
# on extend size and only one kernel compilation.
|
|
||||||
num_part2 = (
|
|
||||||
seq_len // page_size * page_size
|
|
||||||
- (pre_len + page_size - 1) // page_size * page_size
|
|
||||||
)
|
|
||||||
BLOCK_EXTEND: tl.constexpr = 4096
|
|
||||||
num_blocks = (num_part2 + BLOCK_EXTEND - 1) // BLOCK_EXTEND
|
|
||||||
for block_id in range(num_blocks):
|
|
||||||
offset_in_block = tl.arange(0, BLOCK_EXTEND)
|
|
||||||
offset = block_id * BLOCK_EXTEND + offset_in_block
|
|
||||||
mask = offset < num_part2
|
|
||||||
page_start = tl.load(
|
|
||||||
free_page_ptr + new_page_start_loc + offset // page_size,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
out_indices + output_start_loc + num_part1 + offset,
|
|
||||||
page_start * page_size + offset % page_size,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
if pre_len + num_part1 + num_part2 == seq_len:
|
|
||||||
return
|
|
||||||
|
|
||||||
# Part 3: fill the new partial page
|
|
||||||
num_part3 = seq_len - seq_len // page_size * page_size
|
|
||||||
start_loc = tl.load(
|
|
||||||
free_page_ptr + new_page_start_loc + num_page_start_loc_self - 1
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
out_indices + output_start_loc + num_part1 + num_part2 + offset_one_page,
|
|
||||||
start_loc * page_size + offset_one_page,
|
|
||||||
mask=offset_one_page < num_part3,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def alloc_decode_kernel(
|
|
||||||
seq_lens_ptr,
|
|
||||||
last_loc_ptr,
|
|
||||||
free_page_ptr,
|
|
||||||
out_indices,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
page_size: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
|
|
||||||
load_offset = tl.arange(0, bs_upper)
|
|
||||||
seq_lens = tl.load(seq_lens_ptr + load_offset, mask=load_offset <= pid)
|
|
||||||
pre_lens = tl.where(load_offset <= pid, seq_lens - 1, seq_lens)
|
|
||||||
|
|
||||||
seq_len = tl.load(seq_lens_ptr + pid)
|
|
||||||
pre_len = seq_len - 1
|
|
||||||
|
|
||||||
num_pages_after = (seq_lens + page_size - 1) // page_size
|
|
||||||
num_pages_before = (pre_lens + page_size - 1) // page_size
|
|
||||||
num_new_pages = num_pages_after - num_pages_before
|
|
||||||
|
|
||||||
num_page_start_loc_self = (seq_len + page_size - 1) // page_size - (
|
|
||||||
pre_len + page_size - 1
|
|
||||||
) // page_size
|
|
||||||
sum_num_new_pages = tl.sum(num_new_pages)
|
|
||||||
new_page_start_loc = sum_num_new_pages - num_page_start_loc_self
|
|
||||||
|
|
||||||
if num_page_start_loc_self == 0:
|
|
||||||
last_loc = tl.load(last_loc_ptr + pid)
|
|
||||||
tl.store(out_indices + pid, last_loc + 1)
|
|
||||||
else:
|
|
||||||
page = tl.load(free_page_ptr + new_page_start_loc)
|
|
||||||
tl.store(out_indices + pid, page * page_size)
|
|
||||||
|
|
||||||
|
|
||||||
class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||||
"""
|
"""
|
||||||
An allocator managing the indices to kv cache data.
|
An allocator managing the indices to kv cache data.
|
||||||
|
|||||||
@@ -5,12 +5,21 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
|
from sglang.srt.mem_cache.triton_ops.common import (
|
||||||
|
_get_last_loc_safe_kernel as _get_last_loc_safe_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.common import (
|
||||||
|
get_last_loc_kernel as get_last_loc_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.common import (
|
||||||
|
get_last_loc_triton,
|
||||||
|
get_last_loc_triton_safe,
|
||||||
|
write_req_to_token_pool_triton,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import is_hip, support_triton
|
from sglang.srt.utils import is_hip, support_triton
|
||||||
from sglang.srt.utils.common import ceil_align
|
from sglang.srt.utils.common import ceil_align
|
||||||
@@ -50,57 +59,6 @@ def maybe_cache_unfinished_req(req: Req, tree_cache: BasePrefixCache, **kwargs):
|
|||||||
tree_cache.cache_unfinished_req(req, **kwargs)
|
tree_cache.cache_unfinished_req(req, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def write_req_to_token_pool_triton(
|
|
||||||
req_to_token_ptr, # [max_batch, max_context_len]
|
|
||||||
req_pool_indices,
|
|
||||||
prefix_tensors,
|
|
||||||
pre_lens,
|
|
||||||
seq_lens,
|
|
||||||
extend_lens,
|
|
||||||
out_cache_loc,
|
|
||||||
req_to_token_ptr_stride: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 512
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
|
|
||||||
req_pool_index = tl.load(req_pool_indices + pid)
|
|
||||||
pre_len = tl.load(pre_lens + pid)
|
|
||||||
seq_len = tl.load(seq_lens + pid)
|
|
||||||
prefix_tensor = tl.load(prefix_tensors + pid).to(tl.pointer_type(tl.int64))
|
|
||||||
|
|
||||||
# write prefix
|
|
||||||
num_loop = tl.cdiv(pre_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
mask = offset < pre_len
|
|
||||||
value = tl.load(prefix_tensor + offset, mask=mask)
|
|
||||||
tl.store(
|
|
||||||
req_to_token_ptr + req_pool_index * req_to_token_ptr_stride + offset,
|
|
||||||
value,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
# NOTE: This can be slow for large bs
|
|
||||||
cumsum_start = tl.cast(0, tl.int64)
|
|
||||||
for i in range(pid):
|
|
||||||
cumsum_start += tl.load(extend_lens + i)
|
|
||||||
|
|
||||||
num_loop = tl.cdiv(seq_len - pre_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
mask = offset < (seq_len - pre_len)
|
|
||||||
value = tl.load(out_cache_loc + cumsum_start + offset, mask=mask)
|
|
||||||
tl.store(
|
|
||||||
req_to_token_ptr
|
|
||||||
+ req_pool_index * req_to_token_ptr_stride
|
|
||||||
+ offset
|
|
||||||
+ pre_len,
|
|
||||||
value,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def write_cache_indices(
|
def write_cache_indices(
|
||||||
out_cache_loc: torch.Tensor,
|
out_cache_loc: torch.Tensor,
|
||||||
req_pool_indices_tensor: torch.Tensor,
|
req_pool_indices_tensor: torch.Tensor,
|
||||||
@@ -192,113 +150,6 @@ def get_last_loc_torch(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _get_last_loc_safe_kernel(
|
|
||||||
req_to_token,
|
|
||||||
req_pool_indices_tensor,
|
|
||||||
prefix_lens_tensor,
|
|
||||||
result_i32,
|
|
||||||
num_tokens,
|
|
||||||
req_to_token_stride,
|
|
||||||
BLOCK_SIZE: tl.constexpr,
|
|
||||||
PREFIX_DTYPE_IS_I64: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE
|
|
||||||
mask = offset < num_tokens
|
|
||||||
|
|
||||||
if PREFIX_DTYPE_IS_I64:
|
|
||||||
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
|
||||||
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
|
||||||
token_index = req_pool_indices * req_to_token_stride + (prefix_lens - 1)
|
|
||||||
else:
|
|
||||||
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
|
||||||
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
|
||||||
token_index = req_pool_indices.to(tl.int64) * req_to_token_stride + (
|
|
||||||
prefix_lens.to(tl.int64) - 1
|
|
||||||
)
|
|
||||||
|
|
||||||
token_mask = mask & (prefix_lens > 0)
|
|
||||||
tokens = tl.load(req_to_token + token_index, mask=token_mask, other=-1)
|
|
||||||
# Result stays int32 (req_to_token dtype); caller promotes after return.
|
|
||||||
tl.store(result_i32 + offset, tokens, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
def get_last_loc_triton_safe(
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
req_pool_indices_tensor: torch.Tensor,
|
|
||||||
prefix_lens_tensor: torch.Tensor,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Fused `last_loc` Triton kernel whose in-kernel result buffer is int32
|
|
||||||
(the dtype of req_to_token). The consumer-dtype promotion happens in
|
|
||||||
torch after the kernel returns, so Triton never issues a mixed-width
|
|
||||||
store — avoiding the HIP int32->int64 store bug hit by the legacy kernel.
|
|
||||||
"""
|
|
||||||
num_tokens = prefix_lens_tensor.shape[0]
|
|
||||||
BLOCK_SIZE = 256
|
|
||||||
result_i32 = torch.empty(
|
|
||||||
num_tokens, dtype=torch.int32, device=prefix_lens_tensor.device
|
|
||||||
)
|
|
||||||
grid = (triton.cdiv(num_tokens, BLOCK_SIZE),)
|
|
||||||
_get_last_loc_safe_kernel[grid](
|
|
||||||
req_to_token,
|
|
||||||
req_pool_indices_tensor,
|
|
||||||
prefix_lens_tensor,
|
|
||||||
result_i32,
|
|
||||||
num_tokens,
|
|
||||||
req_to_token.stride(0),
|
|
||||||
BLOCK_SIZE=BLOCK_SIZE,
|
|
||||||
PREFIX_DTYPE_IS_I64=(prefix_lens_tensor.dtype == torch.int64),
|
|
||||||
)
|
|
||||||
return result_i32.to(prefix_lens_tensor.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def get_last_loc_kernel(
|
|
||||||
req_to_token,
|
|
||||||
req_pool_indices_tensor,
|
|
||||||
prefix_lens_tensor,
|
|
||||||
result,
|
|
||||||
num_tokens,
|
|
||||||
req_to_token_stride,
|
|
||||||
BLOCK_SIZE: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE
|
|
||||||
mask = offset < num_tokens
|
|
||||||
|
|
||||||
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
|
||||||
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
|
||||||
|
|
||||||
token_mask = prefix_lens > 0
|
|
||||||
token_index = req_pool_indices * req_to_token_stride + (prefix_lens - 1)
|
|
||||||
tokens = tl.load(req_to_token + token_index, mask=token_mask, other=-1)
|
|
||||||
|
|
||||||
tl.store(result + offset, tokens, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
def get_last_loc_triton(
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
req_pool_indices_tensor: torch.Tensor,
|
|
||||||
prefix_lens_tensor: torch.Tensor,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
BLOCK_SIZE = 256
|
|
||||||
num_tokens = prefix_lens_tensor.shape[0]
|
|
||||||
result = torch.empty_like(prefix_lens_tensor)
|
|
||||||
grid = (triton.cdiv(num_tokens, BLOCK_SIZE),)
|
|
||||||
|
|
||||||
get_last_loc_kernel[grid](
|
|
||||||
req_to_token,
|
|
||||||
req_pool_indices_tensor,
|
|
||||||
prefix_lens_tensor,
|
|
||||||
result,
|
|
||||||
num_tokens,
|
|
||||||
req_to_token.stride(0),
|
|
||||||
BLOCK_SIZE,
|
|
||||||
)
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def alloc_token_slots(
|
def alloc_token_slots(
|
||||||
tree_cache: BasePrefixCache,
|
tree_cache: BasePrefixCache,
|
||||||
num_tokens: int,
|
num_tokens: int,
|
||||||
|
|||||||
@@ -11,11 +11,7 @@ distributed under the License is distributed on an "AS IS" BASIS,
|
|||||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
See the License for the specific language governing permissions and
|
See the License for the specific language governing permissions and
|
||||||
limitations under the License.
|
limitations under the License.
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
"""
|
|
||||||
Memory pool.
|
Memory pool.
|
||||||
|
|
||||||
SGLang has two levels of memory pool.
|
SGLang has two levels of memory pool.
|
||||||
@@ -24,6 +20,8 @@ TokenToKVPoolAllocator manages the indices to kv cache data.
|
|||||||
KVCache actually holds the physical kv cache.
|
KVCache actually holds the physical kv cache.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import abc
|
import abc
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import logging
|
import logging
|
||||||
@@ -33,8 +31,6 @@ from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.jit_kernel.kvcache import can_use_store_cache, store_cache
|
from sglang.jit_kernel.kvcache import can_use_store_cache, store_cache
|
||||||
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
||||||
@@ -48,6 +44,9 @@ from sglang.srt.layers.attention.dsa.quant_k_cache import (
|
|||||||
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
|
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.mem_cache.triton_ops.cache_move import (
|
||||||
|
copy_all_layer_kv_cache_tiled,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.utils import (
|
from sglang.srt.mem_cache.utils import (
|
||||||
get_mla_kv_buffer_triton,
|
get_mla_kv_buffer_triton,
|
||||||
maybe_init_custom_mem_pool,
|
maybe_init_custom_mem_pool,
|
||||||
@@ -610,11 +609,11 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
mamba_ping_pong_track_buffers.append(req.mamba_ping_pong_track_buffer)
|
mamba_ping_pong_track_buffers.append(req.mamba_ping_pong_track_buffer)
|
||||||
assert len(select_index) == len(
|
assert len(select_index) == len(
|
||||||
mamba_indices
|
mamba_indices
|
||||||
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
), "Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
assert len(select_index) == len(
|
assert len(select_index) == len(
|
||||||
mamba_ping_pong_track_buffers
|
mamba_ping_pong_track_buffers
|
||||||
), f"Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
), "Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
||||||
mamba_index_tensor = torch.stack(mamba_indices).to(dtype=torch.int32)
|
mamba_index_tensor = torch.stack(mamba_indices).to(dtype=torch.int32)
|
||||||
self.req_index_to_mamba_index_mapping[select_index] = mamba_index_tensor
|
self.req_index_to_mamba_index_mapping[select_index] = mamba_index_tensor
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
@@ -795,7 +794,6 @@ class KVCache(abc.ABC):
|
|||||||
|
|
||||||
|
|
||||||
class MHATokenToKVPool(KVCache):
|
class MHATokenToKVPool(KVCache):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
size: int,
|
size: int,
|
||||||
@@ -1257,7 +1255,6 @@ class NoOpMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
|
|
||||||
|
|
||||||
class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
||||||
|
|
||||||
def _create_buffers(self):
|
def _create_buffers(self):
|
||||||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||||
with (
|
with (
|
||||||
@@ -1434,7 +1431,6 @@ class HybridLinearKVPool(KVCache):
|
|||||||
assert not enable_kvcache_transpose
|
assert not enable_kvcache_transpose
|
||||||
self.use_mla = use_mla
|
self.use_mla = use_mla
|
||||||
if not use_mla:
|
if not use_mla:
|
||||||
|
|
||||||
TokenToKVPoolClass = MHATokenToKVPool
|
TokenToKVPoolClass = MHATokenToKVPool
|
||||||
|
|
||||||
if current_platform.is_out_of_tree():
|
if current_platform.is_out_of_tree():
|
||||||
@@ -1457,7 +1453,6 @@ class HybridLinearKVPool(KVCache):
|
|||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|
||||||
TokenToKVPoolClass = MLATokenToKVPool
|
TokenToKVPoolClass = MLATokenToKVPool
|
||||||
|
|
||||||
if current_platform.is_out_of_tree():
|
if current_platform.is_out_of_tree():
|
||||||
@@ -1543,7 +1538,6 @@ class HybridLinearKVPool(KVCache):
|
|||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _transfer_id_context(self, layer: RadixAttention):
|
def _transfer_id_context(self, layer: RadixAttention):
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _patch_layer_id(layer):
|
def _patch_layer_id(layer):
|
||||||
original_layer_id = layer.layer_id
|
original_layer_id = layer.layer_id
|
||||||
@@ -1863,7 +1857,6 @@ class MLATokenToKVPool(KVCache):
|
|||||||
|
|
||||||
|
|
||||||
class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
||||||
|
|
||||||
def _create_buffers(self):
|
def _create_buffers(self):
|
||||||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||||
with (
|
with (
|
||||||
@@ -2012,7 +2005,6 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
|||||||
end_layer: Optional[int] = None,
|
end_layer: Optional[int] = None,
|
||||||
index_buf_size: Optional[int] = None,
|
index_buf_size: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
|
||||||
override_dim = (
|
override_dim = (
|
||||||
kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None
|
kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None
|
||||||
)
|
)
|
||||||
@@ -2232,39 +2224,3 @@ def move_kv_cache_native(
|
|||||||
for k_cache, v_cache in zip(k_buffer, v_buffer):
|
for k_cache, v_cache in zip(k_buffer, v_buffer):
|
||||||
k_cache[tgt_loc_flat] = k_cache[src_loc_flat]
|
k_cache[tgt_loc_flat] = k_cache[src_loc_flat]
|
||||||
v_cache[tgt_loc_flat] = v_cache[src_loc_flat]
|
v_cache[tgt_loc_flat] = v_cache[src_loc_flat]
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def copy_all_layer_kv_cache_tiled(
|
|
||||||
data_ptrs,
|
|
||||||
strides,
|
|
||||||
tgt_loc_ptr,
|
|
||||||
src_loc_ptr,
|
|
||||||
num_locs,
|
|
||||||
num_locs_upper: tl.constexpr,
|
|
||||||
BYTES_PER_TILE: tl.constexpr,
|
|
||||||
):
|
|
||||||
"""2D tiled kernel. Safe for in-place copy."""
|
|
||||||
bid = tl.program_id(0)
|
|
||||||
tid = tl.program_id(1)
|
|
||||||
|
|
||||||
stride = tl.load(strides + bid)
|
|
||||||
base_ptr = tl.load(data_ptrs + bid)
|
|
||||||
base_ptr = tl.cast(base_ptr, tl.pointer_type(tl.uint8))
|
|
||||||
|
|
||||||
byte_off = tid * BYTES_PER_TILE + tl.arange(0, BYTES_PER_TILE)
|
|
||||||
mask_byte = byte_off < stride
|
|
||||||
tl.multiple_of(byte_off, 16)
|
|
||||||
|
|
||||||
loc_idx = tl.arange(0, num_locs_upper)
|
|
||||||
mask_loc = loc_idx < num_locs
|
|
||||||
|
|
||||||
src = tl.load(src_loc_ptr + loc_idx, mask=mask_loc, other=0)
|
|
||||||
tgt = tl.load(tgt_loc_ptr + loc_idx, mask=mask_loc, other=0)
|
|
||||||
|
|
||||||
src_ptr = base_ptr + src[:, None] * stride + byte_off[None, :]
|
|
||||||
tgt_ptr = base_ptr + tgt[:, None] * stride + byte_off[None, :]
|
|
||||||
|
|
||||||
mask = mask_loc[:, None] & mask_byte[None, :]
|
|
||||||
vals = tl.load(src_ptr, mask=mask)
|
|
||||||
tl.store(tgt_ptr, vals, mask=mask)
|
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Triton kernels for memory cache operations."""
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def alloc_extend_kernel(
|
||||||
|
pre_lens_ptr,
|
||||||
|
seq_lens_ptr,
|
||||||
|
last_loc_ptr,
|
||||||
|
free_page_ptr,
|
||||||
|
out_indices,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
page_size: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
load_offset = tl.arange(0, bs_upper)
|
||||||
|
seq_lens = tl.load(seq_lens_ptr + load_offset, mask=load_offset <= pid)
|
||||||
|
pre_lens = tl.load(pre_lens_ptr + load_offset, mask=load_offset <= pid)
|
||||||
|
extend_lens = seq_lens - pre_lens
|
||||||
|
|
||||||
|
seq_len = tl.load(seq_lens_ptr + pid)
|
||||||
|
pre_len = tl.load(pre_lens_ptr + pid)
|
||||||
|
extend_len = seq_len - pre_len
|
||||||
|
|
||||||
|
sum_extend_lens = tl.sum(extend_lens)
|
||||||
|
output_start_loc = sum_extend_lens - extend_len
|
||||||
|
|
||||||
|
num_pages_after = (seq_lens + page_size - 1) // page_size
|
||||||
|
num_pages_before = (pre_lens + page_size - 1) // page_size
|
||||||
|
num_new_pages = num_pages_after - num_pages_before
|
||||||
|
|
||||||
|
num_page_start_loc_self = (seq_len + page_size - 1) // page_size - (
|
||||||
|
pre_len + page_size - 1
|
||||||
|
) // page_size
|
||||||
|
sum_num_new_pages = tl.sum(num_new_pages)
|
||||||
|
new_page_start_loc = sum_num_new_pages - num_page_start_loc_self
|
||||||
|
|
||||||
|
# Part 1: fill the old partial page
|
||||||
|
last_loc = tl.load(last_loc_ptr + pid)
|
||||||
|
num_part1 = (
|
||||||
|
min(seq_len, (pre_len + page_size - 1) // page_size * page_size) - pre_len
|
||||||
|
)
|
||||||
|
offset_one_page = tl.arange(0, page_size)
|
||||||
|
tl.store(
|
||||||
|
out_indices + output_start_loc + offset_one_page,
|
||||||
|
last_loc + 1 + offset_one_page,
|
||||||
|
mask=offset_one_page < num_part1,
|
||||||
|
)
|
||||||
|
if pre_len + num_part1 == seq_len:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Part 2: fill the new full pages using a dynamic blocked loop.
|
||||||
|
# The loop bound is derived from num_part2 (runtime value), so Triton
|
||||||
|
# generates a real loop instead of unrolling -- no constexpr dependency
|
||||||
|
# on extend size and only one kernel compilation.
|
||||||
|
num_part2 = (
|
||||||
|
seq_len // page_size * page_size
|
||||||
|
- (pre_len + page_size - 1) // page_size * page_size
|
||||||
|
)
|
||||||
|
BLOCK_EXTEND: tl.constexpr = 4096
|
||||||
|
num_blocks = (num_part2 + BLOCK_EXTEND - 1) // BLOCK_EXTEND
|
||||||
|
for block_id in range(num_blocks):
|
||||||
|
offset_in_block = tl.arange(0, BLOCK_EXTEND)
|
||||||
|
offset = block_id * BLOCK_EXTEND + offset_in_block
|
||||||
|
mask = offset < num_part2
|
||||||
|
page_start = tl.load(
|
||||||
|
free_page_ptr + new_page_start_loc + offset // page_size,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
out_indices + output_start_loc + num_part1 + offset,
|
||||||
|
page_start * page_size + offset % page_size,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
if pre_len + num_part1 + num_part2 == seq_len:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Part 3: fill the new partial page
|
||||||
|
num_part3 = seq_len - seq_len // page_size * page_size
|
||||||
|
start_loc = tl.load(
|
||||||
|
free_page_ptr + new_page_start_loc + num_page_start_loc_self - 1
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
out_indices + output_start_loc + num_part1 + num_part2 + offset_one_page,
|
||||||
|
start_loc * page_size + offset_one_page,
|
||||||
|
mask=offset_one_page < num_part3,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def alloc_decode_kernel(
|
||||||
|
seq_lens_ptr,
|
||||||
|
last_loc_ptr,
|
||||||
|
free_page_ptr,
|
||||||
|
out_indices,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
page_size: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
load_offset = tl.arange(0, bs_upper)
|
||||||
|
seq_lens = tl.load(seq_lens_ptr + load_offset, mask=load_offset <= pid)
|
||||||
|
pre_lens = tl.where(load_offset <= pid, seq_lens - 1, seq_lens)
|
||||||
|
|
||||||
|
seq_len = tl.load(seq_lens_ptr + pid)
|
||||||
|
pre_len = seq_len - 1
|
||||||
|
|
||||||
|
num_pages_after = (seq_lens + page_size - 1) // page_size
|
||||||
|
num_pages_before = (pre_lens + page_size - 1) // page_size
|
||||||
|
num_new_pages = num_pages_after - num_pages_before
|
||||||
|
|
||||||
|
num_page_start_loc_self = (seq_len + page_size - 1) // page_size - (
|
||||||
|
pre_len + page_size - 1
|
||||||
|
) // page_size
|
||||||
|
sum_num_new_pages = tl.sum(num_new_pages)
|
||||||
|
new_page_start_loc = sum_num_new_pages - num_page_start_loc_self
|
||||||
|
|
||||||
|
if num_page_start_loc_self == 0:
|
||||||
|
last_loc = tl.load(last_loc_ptr + pid)
|
||||||
|
tl.store(out_indices + pid, last_loc + 1)
|
||||||
|
else:
|
||||||
|
page = tl.load(free_page_ptr + new_page_start_loc)
|
||||||
|
tl.store(out_indices + pid, page * page_size)
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def copy_all_layer_kv_cache_tiled(
|
||||||
|
data_ptrs,
|
||||||
|
strides,
|
||||||
|
tgt_loc_ptr,
|
||||||
|
src_loc_ptr,
|
||||||
|
num_locs,
|
||||||
|
num_locs_upper: tl.constexpr,
|
||||||
|
BYTES_PER_TILE: tl.constexpr,
|
||||||
|
):
|
||||||
|
"""2D tiled kernel. Safe for in-place copy."""
|
||||||
|
bid = tl.program_id(0)
|
||||||
|
tid = tl.program_id(1)
|
||||||
|
|
||||||
|
stride = tl.load(strides + bid)
|
||||||
|
base_ptr = tl.load(data_ptrs + bid)
|
||||||
|
base_ptr = tl.cast(base_ptr, tl.pointer_type(tl.uint8))
|
||||||
|
|
||||||
|
byte_off = tid * BYTES_PER_TILE + tl.arange(0, BYTES_PER_TILE)
|
||||||
|
mask_byte = byte_off < stride
|
||||||
|
tl.multiple_of(byte_off, 16)
|
||||||
|
|
||||||
|
loc_idx = tl.arange(0, num_locs_upper)
|
||||||
|
mask_loc = loc_idx < num_locs
|
||||||
|
|
||||||
|
src = tl.load(src_loc_ptr + loc_idx, mask=mask_loc, other=0)
|
||||||
|
tgt = tl.load(tgt_loc_ptr + loc_idx, mask=mask_loc, other=0)
|
||||||
|
|
||||||
|
src_ptr = base_ptr + src[:, None] * stride + byte_off[None, :]
|
||||||
|
tgt_ptr = base_ptr + tgt[:, None] * stride + byte_off[None, :]
|
||||||
|
|
||||||
|
mask = mask_loc[:, None] & mask_byte[None, :]
|
||||||
|
vals = tl.load(src_ptr, mask=mask)
|
||||||
|
tl.store(tgt_ptr, vals, mask=mask)
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def write_req_to_token_pool_triton(
|
||||||
|
req_to_token_ptr, # [max_batch, max_context_len]
|
||||||
|
req_pool_indices,
|
||||||
|
prefix_tensors,
|
||||||
|
pre_lens,
|
||||||
|
seq_lens,
|
||||||
|
extend_lens,
|
||||||
|
out_cache_loc,
|
||||||
|
req_to_token_ptr_stride: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 512
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
req_pool_index = tl.load(req_pool_indices + pid)
|
||||||
|
pre_len = tl.load(pre_lens + pid)
|
||||||
|
seq_len = tl.load(seq_lens + pid)
|
||||||
|
prefix_tensor = tl.load(prefix_tensors + pid).to(tl.pointer_type(tl.int64))
|
||||||
|
|
||||||
|
# write prefix
|
||||||
|
num_loop = tl.cdiv(pre_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
mask = offset < pre_len
|
||||||
|
value = tl.load(prefix_tensor + offset, mask=mask)
|
||||||
|
tl.store(
|
||||||
|
req_to_token_ptr + req_pool_index * req_to_token_ptr_stride + offset,
|
||||||
|
value,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
# NOTE: This can be slow for large bs
|
||||||
|
cumsum_start = tl.cast(0, tl.int64)
|
||||||
|
for i in range(pid):
|
||||||
|
cumsum_start += tl.load(extend_lens + i)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(seq_len - pre_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
mask = offset < (seq_len - pre_len)
|
||||||
|
value = tl.load(out_cache_loc + cumsum_start + offset, mask=mask)
|
||||||
|
tl.store(
|
||||||
|
req_to_token_ptr
|
||||||
|
+ req_pool_index * req_to_token_ptr_stride
|
||||||
|
+ offset
|
||||||
|
+ pre_len,
|
||||||
|
value,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _get_last_loc_safe_kernel(
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices_tensor,
|
||||||
|
prefix_lens_tensor,
|
||||||
|
result_i32,
|
||||||
|
num_tokens,
|
||||||
|
req_to_token_stride,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
PREFIX_DTYPE_IS_I64: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE
|
||||||
|
mask = offset < num_tokens
|
||||||
|
|
||||||
|
if PREFIX_DTYPE_IS_I64:
|
||||||
|
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
||||||
|
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
||||||
|
token_index = req_pool_indices * req_to_token_stride + (prefix_lens - 1)
|
||||||
|
else:
|
||||||
|
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
||||||
|
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
||||||
|
token_index = req_pool_indices.to(tl.int64) * req_to_token_stride + (
|
||||||
|
prefix_lens.to(tl.int64) - 1
|
||||||
|
)
|
||||||
|
|
||||||
|
token_mask = mask & (prefix_lens > 0)
|
||||||
|
tokens = tl.load(req_to_token + token_index, mask=token_mask, other=-1)
|
||||||
|
# Result stays int32 (req_to_token dtype); caller promotes after return.
|
||||||
|
tl.store(result_i32 + offset, tokens, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def get_last_loc_triton_safe(
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices_tensor: torch.Tensor,
|
||||||
|
prefix_lens_tensor: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Fused `last_loc` Triton kernel whose in-kernel result buffer is int32
|
||||||
|
(the dtype of req_to_token). The consumer-dtype promotion happens in
|
||||||
|
torch after the kernel returns, so Triton never issues a mixed-width
|
||||||
|
store -- avoiding the HIP int32->int64 store bug hit by the legacy kernel.
|
||||||
|
"""
|
||||||
|
num_tokens = prefix_lens_tensor.shape[0]
|
||||||
|
BLOCK_SIZE = 256
|
||||||
|
result_i32 = torch.empty(
|
||||||
|
num_tokens, dtype=torch.int32, device=prefix_lens_tensor.device
|
||||||
|
)
|
||||||
|
grid = (triton.cdiv(num_tokens, BLOCK_SIZE),)
|
||||||
|
_get_last_loc_safe_kernel[grid](
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices_tensor,
|
||||||
|
prefix_lens_tensor,
|
||||||
|
result_i32,
|
||||||
|
num_tokens,
|
||||||
|
req_to_token.stride(0),
|
||||||
|
BLOCK_SIZE=BLOCK_SIZE,
|
||||||
|
PREFIX_DTYPE_IS_I64=(prefix_lens_tensor.dtype == torch.int64),
|
||||||
|
)
|
||||||
|
return result_i32.to(prefix_lens_tensor.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def get_last_loc_kernel(
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices_tensor,
|
||||||
|
prefix_lens_tensor,
|
||||||
|
result,
|
||||||
|
num_tokens,
|
||||||
|
req_to_token_stride,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE
|
||||||
|
mask = offset < num_tokens
|
||||||
|
|
||||||
|
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
||||||
|
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
||||||
|
|
||||||
|
token_mask = prefix_lens > 0
|
||||||
|
token_index = req_pool_indices * req_to_token_stride + (prefix_lens - 1)
|
||||||
|
tokens = tl.load(req_to_token + token_index, mask=token_mask, other=-1)
|
||||||
|
|
||||||
|
tl.store(result + offset, tokens, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def get_last_loc_triton(
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices_tensor: torch.Tensor,
|
||||||
|
prefix_lens_tensor: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
BLOCK_SIZE = 256
|
||||||
|
num_tokens = prefix_lens_tensor.shape[0]
|
||||||
|
result = torch.empty_like(prefix_lens_tensor)
|
||||||
|
grid = (triton.cdiv(num_tokens, BLOCK_SIZE),)
|
||||||
|
|
||||||
|
get_last_loc_kernel[grid](
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices_tensor,
|
||||||
|
prefix_lens_tensor,
|
||||||
|
result,
|
||||||
|
num_tokens,
|
||||||
|
req_to_token.stride(0),
|
||||||
|
BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
return result
|
||||||
@@ -0,0 +1,377 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.jit_kernel.utils import is_arch_support_pdl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def set_mla_kv_buffer_kernel(
|
||||||
|
kv_buffer_ptr,
|
||||||
|
cache_k_nope_ptr,
|
||||||
|
cache_k_rope_ptr,
|
||||||
|
loc_ptr,
|
||||||
|
buffer_stride: tl.constexpr,
|
||||||
|
nope_stride: tl.constexpr,
|
||||||
|
rope_stride: tl.constexpr,
|
||||||
|
nope_dim: tl.constexpr,
|
||||||
|
rope_dim: tl.constexpr,
|
||||||
|
BLOCK: tl.constexpr,
|
||||||
|
USE_GDC: tl.constexpr = False,
|
||||||
|
):
|
||||||
|
pid_loc = tl.program_id(0)
|
||||||
|
pid_blk = tl.program_id(1)
|
||||||
|
|
||||||
|
base = pid_blk * BLOCK
|
||||||
|
offs = base + tl.arange(0, BLOCK)
|
||||||
|
total_dim = nope_dim + rope_dim
|
||||||
|
mask = offs < total_dim
|
||||||
|
|
||||||
|
if USE_GDC:
|
||||||
|
tl.extra.cuda.gdc_wait()
|
||||||
|
|
||||||
|
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
||||||
|
dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs
|
||||||
|
|
||||||
|
# Three-way branch to handle boundary correctly while preserving fast path
|
||||||
|
if base + BLOCK <= nope_dim:
|
||||||
|
# Fast path: entire block is in nope region
|
||||||
|
src = tl.load(
|
||||||
|
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
elif base >= nope_dim:
|
||||||
|
# Fast path: entire block is in rope region
|
||||||
|
offs_rope = offs - nope_dim
|
||||||
|
src = tl.load(
|
||||||
|
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Boundary case: block spans nope/rope boundary (e.g., FP8 with nope_dim=528)
|
||||||
|
# Handle each offset individually to avoid negative indexing
|
||||||
|
is_nope = offs < nope_dim
|
||||||
|
is_rope = (offs >= nope_dim) & (offs < (nope_dim + rope_dim))
|
||||||
|
|
||||||
|
src_nope = tl.load(
|
||||||
|
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
||||||
|
mask=mask & is_nope,
|
||||||
|
other=0,
|
||||||
|
)
|
||||||
|
src_rope = tl.load(
|
||||||
|
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
|
||||||
|
mask=mask & is_rope,
|
||||||
|
other=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
src = tl.where(is_nope, src_nope, src_rope)
|
||||||
|
|
||||||
|
tl.store(dst_ptr, src, mask=mask)
|
||||||
|
|
||||||
|
if USE_GDC:
|
||||||
|
tl.extra.cuda.gdc_launch_dependents()
|
||||||
|
|
||||||
|
|
||||||
|
# Above this loc count the TMA bulk-store path overtakes the single-CTA-per-loc
|
||||||
|
# Triton kernel. Below it, Triton with BLOCK = next_pow2(total_dim) (one CTA
|
||||||
|
# does the whole row in one tile, no boundary fan-out) is the winning fallback.
|
||||||
|
# Tuned on GB300 with DSv4 row widths.
|
||||||
|
_TMA_BULK_STORE_MIN_LOCS = 768
|
||||||
|
|
||||||
|
|
||||||
|
def set_mla_kv_buffer_triton(
|
||||||
|
kv_buffer: torch.Tensor,
|
||||||
|
loc: torch.Tensor,
|
||||||
|
cache_k_nope: torch.Tensor,
|
||||||
|
cache_k_rope: torch.Tensor,
|
||||||
|
):
|
||||||
|
"""Dispatch MLA paged-KV scatter writes to the fastest available path.
|
||||||
|
|
||||||
|
Two paths, chosen on ``n_loc``:
|
||||||
|
|
||||||
|
- ``n_loc >= 768`` (and SM90+ with TMA-compatible row widths): JIT CUDA
|
||||||
|
kernel where each warp loads one (nope, rope) row into shared memory and
|
||||||
|
issues a single ``cp.async.bulk.global.shared::cta`` store to scatter the
|
||||||
|
row at ``kv_buffer[loc[item]]``. Wins at large bs because it packs 4-8
|
||||||
|
items per CTA, drastically reducing the CTA count vs single-CTA-per-loc.
|
||||||
|
- Otherwise: Triton kernel with ``BLOCK = next_pow2(nope_dim + rope_dim)``,
|
||||||
|
i.e. one CTA per loc covering the entire row in one tile. Wins at small
|
||||||
|
bs because there's no per-loc CTA fan-out (5x fewer CTAs than the old
|
||||||
|
BLOCK=128 dispatch) and the row-spanning block makes the boundary branch
|
||||||
|
a one-shot per CTA. This is also the path for SM<90 and for shapes that
|
||||||
|
violate the TMA 16-byte alignment.
|
||||||
|
|
||||||
|
Speedup vs the legacy BLOCK=128 Triton kernel on GB300 (BF16, nope=512,
|
||||||
|
rope=64): ~1.05x at bs=8, ~1.5x at bs=128, 3.5x at bs=512, **11.7x at
|
||||||
|
bs=16384**.
|
||||||
|
|
||||||
|
Name retained for caller compatibility; the implementation is no longer
|
||||||
|
Triton-only.
|
||||||
|
"""
|
||||||
|
from sglang.jit_kernel.set_mla_kv_buffer import (
|
||||||
|
can_use_set_mla_kv_buffer,
|
||||||
|
)
|
||||||
|
from sglang.jit_kernel.set_mla_kv_buffer import (
|
||||||
|
set_mla_kv_buffer as jit_set_mla_kv_buffer,
|
||||||
|
)
|
||||||
|
|
||||||
|
n_loc = loc.numel()
|
||||||
|
nope_bytes = cache_k_nope.shape[-1] * cache_k_nope.element_size()
|
||||||
|
rope_bytes = cache_k_rope.shape[-1] * cache_k_rope.element_size()
|
||||||
|
if (
|
||||||
|
n_loc >= _TMA_BULK_STORE_MIN_LOCS
|
||||||
|
and is_arch_support_pdl()
|
||||||
|
and can_use_set_mla_kv_buffer(nope_bytes, rope_bytes)
|
||||||
|
):
|
||||||
|
jit_set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Fallback: Triton with BLOCK = next_pow2(total_dim). One CTA per loc; the
|
||||||
|
# whole row in one tile (the existing 3-way nope/rope/boundary branch in
|
||||||
|
# ``set_mla_kv_buffer_kernel`` handles the over-allocation past total_dim
|
||||||
|
# via the offs<total_dim mask). Beats BLOCK=128 by 60-2700 ns across the
|
||||||
|
# 2 <= bs <= 512 range on GB300.
|
||||||
|
nope_dim = cache_k_nope.shape[-1]
|
||||||
|
rope_dim = cache_k_rope.shape[-1]
|
||||||
|
total_dim = nope_dim + rope_dim
|
||||||
|
BLOCK = triton.next_power_of_2(total_dim)
|
||||||
|
grid = (n_loc, 1)
|
||||||
|
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
||||||
|
set_mla_kv_buffer_kernel[grid](
|
||||||
|
kv_buffer,
|
||||||
|
cache_k_nope,
|
||||||
|
cache_k_rope,
|
||||||
|
loc,
|
||||||
|
kv_buffer.stride(0),
|
||||||
|
cache_k_nope.stride(0),
|
||||||
|
cache_k_rope.stride(0),
|
||||||
|
nope_dim,
|
||||||
|
rope_dim,
|
||||||
|
BLOCK=BLOCK,
|
||||||
|
**pdl_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def set_mla_kv_buffer_fp8_quant_kernel(
|
||||||
|
kv_buffer_fp8_ptr,
|
||||||
|
cache_k_nope_ptr,
|
||||||
|
cache_k_rope_ptr,
|
||||||
|
loc_ptr,
|
||||||
|
buffer_stride: tl.constexpr,
|
||||||
|
nope_stride: tl.constexpr,
|
||||||
|
rope_stride: tl.constexpr,
|
||||||
|
nope_dim: tl.constexpr,
|
||||||
|
rope_dim: tl.constexpr,
|
||||||
|
BLOCK: tl.constexpr,
|
||||||
|
USE_GDC: tl.constexpr = False,
|
||||||
|
):
|
||||||
|
"""Fuse BF16/FP16->FP8 cast with paged KV write."""
|
||||||
|
pid_loc = tl.program_id(0)
|
||||||
|
pid_blk = tl.program_id(1)
|
||||||
|
|
||||||
|
base = pid_blk * BLOCK
|
||||||
|
offs = base + tl.arange(0, BLOCK)
|
||||||
|
total_dim = nope_dim + rope_dim
|
||||||
|
mask = offs < total_dim
|
||||||
|
|
||||||
|
if USE_GDC:
|
||||||
|
tl.extra.cuda.gdc_wait()
|
||||||
|
|
||||||
|
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
||||||
|
dst_ptr = kv_buffer_fp8_ptr + loc * buffer_stride + offs
|
||||||
|
|
||||||
|
if base + BLOCK <= nope_dim:
|
||||||
|
src = tl.load(
|
||||||
|
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
||||||
|
mask=mask,
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
elif base >= nope_dim:
|
||||||
|
offs_rope = offs - nope_dim
|
||||||
|
src = tl.load(
|
||||||
|
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
|
||||||
|
mask=mask,
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
is_nope = offs < nope_dim
|
||||||
|
src_nope = tl.load(
|
||||||
|
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
||||||
|
mask=mask & is_nope,
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
src_rope = tl.load(
|
||||||
|
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
|
||||||
|
mask=mask & ~is_nope,
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
src = tl.where(is_nope, src_nope, src_rope)
|
||||||
|
|
||||||
|
# Destination pointer is FP8-typed view; tl.store performs downcast.
|
||||||
|
tl.store(dst_ptr, src, mask=mask)
|
||||||
|
|
||||||
|
if USE_GDC:
|
||||||
|
tl.extra.cuda.gdc_launch_dependents()
|
||||||
|
|
||||||
|
|
||||||
|
def set_mla_kv_buffer_triton_fp8_quant(
|
||||||
|
kv_buffer: torch.Tensor,
|
||||||
|
loc: torch.Tensor,
|
||||||
|
cache_k_nope: torch.Tensor,
|
||||||
|
cache_k_rope: torch.Tensor,
|
||||||
|
fp8_dtype: torch.dtype,
|
||||||
|
):
|
||||||
|
"""Fuse BF16/FP16 MLA K quantization with paged KV write."""
|
||||||
|
kv_buffer_fp8 = kv_buffer.view(fp8_dtype)
|
||||||
|
|
||||||
|
nope_dim = cache_k_nope.shape[-1]
|
||||||
|
rope_dim = cache_k_rope.shape[-1]
|
||||||
|
total_dim = nope_dim + rope_dim
|
||||||
|
BLOCK = 128
|
||||||
|
n_loc = loc.numel()
|
||||||
|
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
|
||||||
|
|
||||||
|
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
||||||
|
|
||||||
|
set_mla_kv_buffer_fp8_quant_kernel[grid](
|
||||||
|
kv_buffer_fp8,
|
||||||
|
cache_k_nope,
|
||||||
|
cache_k_rope,
|
||||||
|
loc,
|
||||||
|
kv_buffer_fp8.stride(0),
|
||||||
|
cache_k_nope.stride(0),
|
||||||
|
cache_k_rope.stride(0),
|
||||||
|
nope_dim,
|
||||||
|
rope_dim,
|
||||||
|
BLOCK=BLOCK,
|
||||||
|
**pdl_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def set_mla_kv_scale_buffer_kernel(
|
||||||
|
kv_buffer_ptr,
|
||||||
|
cache_k_nope_ptr,
|
||||||
|
cache_k_rope_ptr,
|
||||||
|
loc_ptr,
|
||||||
|
buffer_stride: tl.constexpr,
|
||||||
|
nope_stride: tl.constexpr,
|
||||||
|
rope_stride: tl.constexpr,
|
||||||
|
nope_dim: tl.constexpr,
|
||||||
|
rope_dim: tl.constexpr,
|
||||||
|
BLOCK: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_loc = tl.program_id(0)
|
||||||
|
pid_blk = tl.program_id(1)
|
||||||
|
|
||||||
|
base = pid_blk * BLOCK
|
||||||
|
offs = base + tl.arange(0, BLOCK)
|
||||||
|
total_dim = nope_dim + rope_dim
|
||||||
|
mask = offs < total_dim # Make sure don't cross the boundary
|
||||||
|
|
||||||
|
loc = tl.load(loc_ptr + pid_loc)
|
||||||
|
dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs
|
||||||
|
|
||||||
|
# Check each offs should read 'nope' or 'rope'
|
||||||
|
is_nope = offs < nope_dim
|
||||||
|
src_nope = tl.load(
|
||||||
|
cache_k_nope_ptr + pid_loc * nope_stride + offs, mask=mask & is_nope, other=0.0
|
||||||
|
)
|
||||||
|
src_rope = tl.load(
|
||||||
|
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
|
||||||
|
mask=mask & ~is_nope,
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Combine nope + rope
|
||||||
|
src = src_nope + src_rope
|
||||||
|
tl.store(dst_ptr, src, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def set_mla_kv_scale_buffer_triton(
|
||||||
|
kv_buffer: torch.Tensor,
|
||||||
|
loc: torch.Tensor,
|
||||||
|
cache_k_nope: torch.Tensor,
|
||||||
|
cache_k_rope: torch.Tensor,
|
||||||
|
):
|
||||||
|
nope_dim = cache_k_nope.shape[-1]
|
||||||
|
rope_dim = cache_k_rope.shape[-1]
|
||||||
|
total_dim = nope_dim + rope_dim
|
||||||
|
BLOCK = 128 # Keep origin, works for smaller total_dim as well.
|
||||||
|
n_loc = loc.numel()
|
||||||
|
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
|
||||||
|
|
||||||
|
set_mla_kv_scale_buffer_kernel[grid](
|
||||||
|
kv_buffer,
|
||||||
|
cache_k_nope,
|
||||||
|
cache_k_rope,
|
||||||
|
loc,
|
||||||
|
kv_buffer.stride(0),
|
||||||
|
cache_k_nope.stride(0),
|
||||||
|
cache_k_rope.stride(0),
|
||||||
|
nope_dim,
|
||||||
|
rope_dim,
|
||||||
|
BLOCK=BLOCK,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def get_mla_kv_buffer_kernel(
|
||||||
|
kv_buffer_ptr,
|
||||||
|
cache_k_nope_ptr,
|
||||||
|
cache_k_rope_ptr,
|
||||||
|
loc_ptr,
|
||||||
|
buffer_stride: tl.constexpr,
|
||||||
|
nope_stride: tl.constexpr,
|
||||||
|
rope_stride: tl.constexpr,
|
||||||
|
nope_dim: tl.constexpr,
|
||||||
|
rope_dim: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid_loc = tl.program_id(0)
|
||||||
|
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
||||||
|
loc_src_ptr = kv_buffer_ptr + loc * buffer_stride
|
||||||
|
|
||||||
|
nope_offs = tl.arange(0, nope_dim)
|
||||||
|
nope_src_ptr = loc_src_ptr + nope_offs
|
||||||
|
nope_src = tl.load(nope_src_ptr)
|
||||||
|
|
||||||
|
tl.store(
|
||||||
|
cache_k_nope_ptr + pid_loc * nope_stride + nope_offs,
|
||||||
|
nope_src,
|
||||||
|
)
|
||||||
|
|
||||||
|
rope_offs = tl.arange(0, rope_dim)
|
||||||
|
rope_src_ptr = loc_src_ptr + nope_dim + rope_offs
|
||||||
|
rope_src = tl.load(rope_src_ptr)
|
||||||
|
tl.store(
|
||||||
|
cache_k_rope_ptr + pid_loc * rope_stride + rope_offs,
|
||||||
|
rope_src,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_mla_kv_buffer_triton(
|
||||||
|
kv_buffer: torch.Tensor,
|
||||||
|
loc: torch.Tensor,
|
||||||
|
cache_k_nope: torch.Tensor,
|
||||||
|
cache_k_rope: torch.Tensor,
|
||||||
|
):
|
||||||
|
# The source data type will be implicitly converted to the target data type.
|
||||||
|
nope_dim = cache_k_nope.shape[-1] # 512
|
||||||
|
rope_dim = cache_k_rope.shape[-1] # 64
|
||||||
|
n_loc = loc.numel()
|
||||||
|
grid = (n_loc,)
|
||||||
|
|
||||||
|
get_mla_kv_buffer_kernel[grid](
|
||||||
|
kv_buffer,
|
||||||
|
cache_k_nope,
|
||||||
|
cache_k_rope,
|
||||||
|
loc,
|
||||||
|
kv_buffer.stride(0),
|
||||||
|
cache_k_nope.stride(0),
|
||||||
|
cache_k_rope.stride(0),
|
||||||
|
nope_dim,
|
||||||
|
rope_dim,
|
||||||
|
)
|
||||||
@@ -16,11 +16,6 @@
|
|||||||
import hashlib
|
import hashlib
|
||||||
from typing import Any, Callable, List, Optional, Tuple
|
from typing import Any, Callable, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.jit_kernel.utils import is_arch_support_pdl
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.mem_cache.evict_policy import (
|
from sglang.srt.mem_cache.evict_policy import (
|
||||||
EvictionStrategy,
|
EvictionStrategy,
|
||||||
@@ -32,6 +27,30 @@ from sglang.srt.mem_cache.evict_policy import (
|
|||||||
PriorityStrategy,
|
PriorityStrategy,
|
||||||
SLRUStrategy,
|
SLRUStrategy,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
get_mla_kv_buffer_kernel as get_mla_kv_buffer_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
get_mla_kv_buffer_triton as get_mla_kv_buffer_triton,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
set_mla_kv_buffer_fp8_quant_kernel as set_mla_kv_buffer_fp8_quant_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
set_mla_kv_buffer_kernel as set_mla_kv_buffer_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
set_mla_kv_buffer_triton as set_mla_kv_buffer_triton,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
set_mla_kv_buffer_triton_fp8_quant as set_mla_kv_buffer_triton_fp8_quant,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
set_mla_kv_scale_buffer_kernel as set_mla_kv_scale_buffer_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.triton_ops.mla_buffer import (
|
||||||
|
set_mla_kv_scale_buffer_triton as set_mla_kv_scale_buffer_triton,
|
||||||
|
)
|
||||||
|
|
||||||
_EVICTION_POLICY_FACTORIES: dict[str, Callable[[], EvictionStrategy]] = {
|
_EVICTION_POLICY_FACTORIES: dict[str, Callable[[], EvictionStrategy]] = {
|
||||||
"lru": LRUStrategy,
|
"lru": LRUStrategy,
|
||||||
@@ -55,376 +74,6 @@ def get_eviction_strategy(eviction_policy: str) -> EvictionStrategy:
|
|||||||
) from None
|
) from None
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def set_mla_kv_buffer_kernel(
|
|
||||||
kv_buffer_ptr,
|
|
||||||
cache_k_nope_ptr,
|
|
||||||
cache_k_rope_ptr,
|
|
||||||
loc_ptr,
|
|
||||||
buffer_stride: tl.constexpr,
|
|
||||||
nope_stride: tl.constexpr,
|
|
||||||
rope_stride: tl.constexpr,
|
|
||||||
nope_dim: tl.constexpr,
|
|
||||||
rope_dim: tl.constexpr,
|
|
||||||
BLOCK: tl.constexpr,
|
|
||||||
USE_GDC: tl.constexpr = False,
|
|
||||||
):
|
|
||||||
pid_loc = tl.program_id(0)
|
|
||||||
pid_blk = tl.program_id(1)
|
|
||||||
|
|
||||||
base = pid_blk * BLOCK
|
|
||||||
offs = base + tl.arange(0, BLOCK)
|
|
||||||
total_dim = nope_dim + rope_dim
|
|
||||||
mask = offs < total_dim
|
|
||||||
|
|
||||||
if USE_GDC:
|
|
||||||
tl.extra.cuda.gdc_wait()
|
|
||||||
|
|
||||||
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
|
||||||
dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs
|
|
||||||
|
|
||||||
# Three-way branch to handle boundary correctly while preserving fast path
|
|
||||||
if base + BLOCK <= nope_dim:
|
|
||||||
# Fast path: entire block is in nope region
|
|
||||||
src = tl.load(
|
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
elif base >= nope_dim:
|
|
||||||
# Fast path: entire block is in rope region
|
|
||||||
offs_rope = offs - nope_dim
|
|
||||||
src = tl.load(
|
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# Boundary case: block spans nope/rope boundary (e.g., FP8 with nope_dim=528)
|
|
||||||
# Handle each offset individually to avoid negative indexing
|
|
||||||
is_nope = offs < nope_dim
|
|
||||||
is_rope = (offs >= nope_dim) & (offs < (nope_dim + rope_dim))
|
|
||||||
|
|
||||||
src_nope = tl.load(
|
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
|
||||||
mask=mask & is_nope,
|
|
||||||
other=0,
|
|
||||||
)
|
|
||||||
src_rope = tl.load(
|
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
|
|
||||||
mask=mask & is_rope,
|
|
||||||
other=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
src = tl.where(is_nope, src_nope, src_rope)
|
|
||||||
|
|
||||||
tl.store(dst_ptr, src, mask=mask)
|
|
||||||
|
|
||||||
if USE_GDC:
|
|
||||||
tl.extra.cuda.gdc_launch_dependents()
|
|
||||||
|
|
||||||
|
|
||||||
# Above this loc count the TMA bulk-store path overtakes the single-CTA-per-loc
|
|
||||||
# Triton kernel. Below it, Triton with BLOCK = next_pow2(total_dim) (one CTA
|
|
||||||
# does the whole row in one tile, no boundary fan-out) is the winning fallback.
|
|
||||||
# Tuned on GB300 with DSv4 row widths.
|
|
||||||
_TMA_BULK_STORE_MIN_LOCS = 768
|
|
||||||
|
|
||||||
|
|
||||||
def set_mla_kv_buffer_triton(
|
|
||||||
kv_buffer: torch.Tensor,
|
|
||||||
loc: torch.Tensor,
|
|
||||||
cache_k_nope: torch.Tensor,
|
|
||||||
cache_k_rope: torch.Tensor,
|
|
||||||
):
|
|
||||||
"""Dispatch MLA paged-KV scatter writes to the fastest available path.
|
|
||||||
|
|
||||||
Two paths, chosen on ``n_loc``:
|
|
||||||
|
|
||||||
- ``n_loc >= 768`` (and SM90+ with TMA-compatible row widths): JIT CUDA
|
|
||||||
kernel where each warp loads one (nope, rope) row into shared memory and
|
|
||||||
issues a single ``cp.async.bulk.global.shared::cta`` store to scatter the
|
|
||||||
row at ``kv_buffer[loc[item]]``. Wins at large bs because it packs 4-8
|
|
||||||
items per CTA, drastically reducing the CTA count vs single-CTA-per-loc.
|
|
||||||
- Otherwise: Triton kernel with ``BLOCK = next_pow2(nope_dim + rope_dim)``,
|
|
||||||
i.e. one CTA per loc covering the entire row in one tile. Wins at small
|
|
||||||
bs because there's no per-loc CTA fan-out (5× fewer CTAs than the old
|
|
||||||
BLOCK=128 dispatch) and the row-spanning block makes the boundary branch
|
|
||||||
a one-shot per CTA. This is also the path for SM<90 and for shapes that
|
|
||||||
violate the TMA 16-byte alignment.
|
|
||||||
|
|
||||||
Speedup vs the legacy BLOCK=128 Triton kernel on GB300 (BF16, nope=512,
|
|
||||||
rope=64): ~1.05× at bs=8, ~1.5× at bs=128, 3.5× at bs=512, **11.7× at
|
|
||||||
bs=16384**.
|
|
||||||
|
|
||||||
Name retained for caller compatibility; the implementation is no longer
|
|
||||||
Triton-only.
|
|
||||||
"""
|
|
||||||
from sglang.jit_kernel.set_mla_kv_buffer import (
|
|
||||||
can_use_set_mla_kv_buffer,
|
|
||||||
)
|
|
||||||
from sglang.jit_kernel.set_mla_kv_buffer import (
|
|
||||||
set_mla_kv_buffer as jit_set_mla_kv_buffer,
|
|
||||||
)
|
|
||||||
|
|
||||||
n_loc = loc.numel()
|
|
||||||
nope_bytes = cache_k_nope.shape[-1] * cache_k_nope.element_size()
|
|
||||||
rope_bytes = cache_k_rope.shape[-1] * cache_k_rope.element_size()
|
|
||||||
if (
|
|
||||||
n_loc >= _TMA_BULK_STORE_MIN_LOCS
|
|
||||||
and is_arch_support_pdl()
|
|
||||||
and can_use_set_mla_kv_buffer(nope_bytes, rope_bytes)
|
|
||||||
):
|
|
||||||
jit_set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
|
||||||
return
|
|
||||||
|
|
||||||
# Fallback: Triton with BLOCK = next_pow2(total_dim). One CTA per loc; the
|
|
||||||
# whole row in one tile (the existing 3-way nope/rope/boundary branch in
|
|
||||||
# ``set_mla_kv_buffer_kernel`` handles the over-allocation past total_dim
|
|
||||||
# via the offs<total_dim mask). Beats BLOCK=128 by 60-2700 ns across the
|
|
||||||
# 2 ≤ bs ≤ 512 range on GB300.
|
|
||||||
nope_dim = cache_k_nope.shape[-1]
|
|
||||||
rope_dim = cache_k_rope.shape[-1]
|
|
||||||
total_dim = nope_dim + rope_dim
|
|
||||||
BLOCK = triton.next_power_of_2(total_dim)
|
|
||||||
grid = (n_loc, 1)
|
|
||||||
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
|
||||||
set_mla_kv_buffer_kernel[grid](
|
|
||||||
kv_buffer,
|
|
||||||
cache_k_nope,
|
|
||||||
cache_k_rope,
|
|
||||||
loc,
|
|
||||||
kv_buffer.stride(0),
|
|
||||||
cache_k_nope.stride(0),
|
|
||||||
cache_k_rope.stride(0),
|
|
||||||
nope_dim,
|
|
||||||
rope_dim,
|
|
||||||
BLOCK=BLOCK,
|
|
||||||
**pdl_kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def set_mla_kv_buffer_fp8_quant_kernel(
|
|
||||||
kv_buffer_fp8_ptr,
|
|
||||||
cache_k_nope_ptr,
|
|
||||||
cache_k_rope_ptr,
|
|
||||||
loc_ptr,
|
|
||||||
buffer_stride: tl.constexpr,
|
|
||||||
nope_stride: tl.constexpr,
|
|
||||||
rope_stride: tl.constexpr,
|
|
||||||
nope_dim: tl.constexpr,
|
|
||||||
rope_dim: tl.constexpr,
|
|
||||||
BLOCK: tl.constexpr,
|
|
||||||
USE_GDC: tl.constexpr = False,
|
|
||||||
):
|
|
||||||
"""Fuse BF16/FP16->FP8 cast with paged KV write."""
|
|
||||||
pid_loc = tl.program_id(0)
|
|
||||||
pid_blk = tl.program_id(1)
|
|
||||||
|
|
||||||
base = pid_blk * BLOCK
|
|
||||||
offs = base + tl.arange(0, BLOCK)
|
|
||||||
total_dim = nope_dim + rope_dim
|
|
||||||
mask = offs < total_dim
|
|
||||||
|
|
||||||
if USE_GDC:
|
|
||||||
tl.extra.cuda.gdc_wait()
|
|
||||||
|
|
||||||
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
|
||||||
dst_ptr = kv_buffer_fp8_ptr + loc * buffer_stride + offs
|
|
||||||
|
|
||||||
if base + BLOCK <= nope_dim:
|
|
||||||
src = tl.load(
|
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
|
||||||
mask=mask,
|
|
||||||
other=0.0,
|
|
||||||
)
|
|
||||||
elif base >= nope_dim:
|
|
||||||
offs_rope = offs - nope_dim
|
|
||||||
src = tl.load(
|
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
|
|
||||||
mask=mask,
|
|
||||||
other=0.0,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
is_nope = offs < nope_dim
|
|
||||||
src_nope = tl.load(
|
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
|
||||||
mask=mask & is_nope,
|
|
||||||
other=0.0,
|
|
||||||
)
|
|
||||||
src_rope = tl.load(
|
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
|
|
||||||
mask=mask & ~is_nope,
|
|
||||||
other=0.0,
|
|
||||||
)
|
|
||||||
src = tl.where(is_nope, src_nope, src_rope)
|
|
||||||
|
|
||||||
# Destination pointer is FP8-typed view; tl.store performs downcast.
|
|
||||||
tl.store(dst_ptr, src, mask=mask)
|
|
||||||
|
|
||||||
if USE_GDC:
|
|
||||||
tl.extra.cuda.gdc_launch_dependents()
|
|
||||||
|
|
||||||
|
|
||||||
def set_mla_kv_buffer_triton_fp8_quant(
|
|
||||||
kv_buffer: torch.Tensor,
|
|
||||||
loc: torch.Tensor,
|
|
||||||
cache_k_nope: torch.Tensor,
|
|
||||||
cache_k_rope: torch.Tensor,
|
|
||||||
fp8_dtype: torch.dtype,
|
|
||||||
):
|
|
||||||
"""Fuse BF16/FP16 MLA K quantization with paged KV write."""
|
|
||||||
kv_buffer_fp8 = kv_buffer.view(fp8_dtype)
|
|
||||||
|
|
||||||
nope_dim = cache_k_nope.shape[-1]
|
|
||||||
rope_dim = cache_k_rope.shape[-1]
|
|
||||||
total_dim = nope_dim + rope_dim
|
|
||||||
BLOCK = 128
|
|
||||||
n_loc = loc.numel()
|
|
||||||
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
|
|
||||||
|
|
||||||
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
|
||||||
|
|
||||||
set_mla_kv_buffer_fp8_quant_kernel[grid](
|
|
||||||
kv_buffer_fp8,
|
|
||||||
cache_k_nope,
|
|
||||||
cache_k_rope,
|
|
||||||
loc,
|
|
||||||
kv_buffer_fp8.stride(0),
|
|
||||||
cache_k_nope.stride(0),
|
|
||||||
cache_k_rope.stride(0),
|
|
||||||
nope_dim,
|
|
||||||
rope_dim,
|
|
||||||
BLOCK=BLOCK,
|
|
||||||
**pdl_kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def set_mla_kv_scale_buffer_kernel(
|
|
||||||
kv_buffer_ptr,
|
|
||||||
cache_k_nope_ptr,
|
|
||||||
cache_k_rope_ptr,
|
|
||||||
loc_ptr,
|
|
||||||
buffer_stride: tl.constexpr,
|
|
||||||
nope_stride: tl.constexpr,
|
|
||||||
rope_stride: tl.constexpr,
|
|
||||||
nope_dim: tl.constexpr,
|
|
||||||
rope_dim: tl.constexpr,
|
|
||||||
BLOCK: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid_loc = tl.program_id(0)
|
|
||||||
pid_blk = tl.program_id(1)
|
|
||||||
|
|
||||||
base = pid_blk * BLOCK
|
|
||||||
offs = base + tl.arange(0, BLOCK)
|
|
||||||
total_dim = nope_dim + rope_dim
|
|
||||||
mask = offs < total_dim # Make sure don't cross the boundary
|
|
||||||
|
|
||||||
loc = tl.load(loc_ptr + pid_loc)
|
|
||||||
dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs
|
|
||||||
|
|
||||||
# Check each offs should read 'nope' or 'rope'
|
|
||||||
is_nope = offs < nope_dim
|
|
||||||
src_nope = tl.load(
|
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + offs, mask=mask & is_nope, other=0.0
|
|
||||||
)
|
|
||||||
src_rope = tl.load(
|
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
|
|
||||||
mask=mask & ~is_nope,
|
|
||||||
other=0.0,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Combine nope + rope
|
|
||||||
src = src_nope + src_rope
|
|
||||||
tl.store(dst_ptr, src, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
def set_mla_kv_scale_buffer_triton(
|
|
||||||
kv_buffer: torch.Tensor,
|
|
||||||
loc: torch.Tensor,
|
|
||||||
cache_k_nope: torch.Tensor,
|
|
||||||
cache_k_rope: torch.Tensor,
|
|
||||||
):
|
|
||||||
nope_dim = cache_k_nope.shape[-1]
|
|
||||||
rope_dim = cache_k_rope.shape[-1]
|
|
||||||
total_dim = nope_dim + rope_dim
|
|
||||||
BLOCK = 128 # Keep origin, works for smaller total_dim as well.
|
|
||||||
n_loc = loc.numel()
|
|
||||||
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
|
|
||||||
|
|
||||||
set_mla_kv_scale_buffer_kernel[grid](
|
|
||||||
kv_buffer,
|
|
||||||
cache_k_nope,
|
|
||||||
cache_k_rope,
|
|
||||||
loc,
|
|
||||||
kv_buffer.stride(0),
|
|
||||||
cache_k_nope.stride(0),
|
|
||||||
cache_k_rope.stride(0),
|
|
||||||
nope_dim,
|
|
||||||
rope_dim,
|
|
||||||
BLOCK=BLOCK,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def get_mla_kv_buffer_kernel(
|
|
||||||
kv_buffer_ptr,
|
|
||||||
cache_k_nope_ptr,
|
|
||||||
cache_k_rope_ptr,
|
|
||||||
loc_ptr,
|
|
||||||
buffer_stride: tl.constexpr,
|
|
||||||
nope_stride: tl.constexpr,
|
|
||||||
rope_stride: tl.constexpr,
|
|
||||||
nope_dim: tl.constexpr,
|
|
||||||
rope_dim: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid_loc = tl.program_id(0)
|
|
||||||
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
|
||||||
loc_src_ptr = kv_buffer_ptr + loc * buffer_stride
|
|
||||||
|
|
||||||
nope_offs = tl.arange(0, nope_dim)
|
|
||||||
nope_src_ptr = loc_src_ptr + nope_offs
|
|
||||||
nope_src = tl.load(nope_src_ptr)
|
|
||||||
|
|
||||||
tl.store(
|
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + nope_offs,
|
|
||||||
nope_src,
|
|
||||||
)
|
|
||||||
|
|
||||||
rope_offs = tl.arange(0, rope_dim)
|
|
||||||
rope_src_ptr = loc_src_ptr + nope_dim + rope_offs
|
|
||||||
rope_src = tl.load(rope_src_ptr)
|
|
||||||
tl.store(
|
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + rope_offs,
|
|
||||||
rope_src,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_mla_kv_buffer_triton(
|
|
||||||
kv_buffer: torch.Tensor,
|
|
||||||
loc: torch.Tensor,
|
|
||||||
cache_k_nope: torch.Tensor,
|
|
||||||
cache_k_rope: torch.Tensor,
|
|
||||||
):
|
|
||||||
# The source data type will be implicitly converted to the target data type.
|
|
||||||
nope_dim = cache_k_nope.shape[-1] # 512
|
|
||||||
rope_dim = cache_k_rope.shape[-1] # 64
|
|
||||||
n_loc = loc.numel()
|
|
||||||
grid = (n_loc,)
|
|
||||||
|
|
||||||
get_mla_kv_buffer_kernel[grid](
|
|
||||||
kv_buffer,
|
|
||||||
cache_k_nope,
|
|
||||||
cache_k_rope,
|
|
||||||
loc,
|
|
||||||
kv_buffer.stride(0),
|
|
||||||
cache_k_nope.stride(0),
|
|
||||||
cache_k_rope.stride(0),
|
|
||||||
nope_dim,
|
|
||||||
rope_dim,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def maybe_init_custom_mem_pool(
|
def maybe_init_custom_mem_pool(
|
||||||
device: str,
|
device: str,
|
||||||
) -> Tuple[bool, Optional[Any], Optional[str]]:
|
) -> Tuple[bool, Optional[Any], Optional[str]]:
|
||||||
|
|||||||
@@ -34,8 +34,6 @@ from functools import total_ordering
|
|||||||
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple, Union
|
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
get_moe_expert_parallel_world_size,
|
get_moe_expert_parallel_world_size,
|
||||||
@@ -57,6 +55,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
||||||
ForwardBatchDeepSeekMHAMixin,
|
ForwardBatchDeepSeekMHAMixin,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.triton_ops.position import compute_position_triton
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
is_cuda,
|
is_cuda,
|
||||||
@@ -1171,7 +1170,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size)
|
self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size)
|
||||||
|
|
||||||
def post_forward_mlp_sync_batch(self, logits_output: LogitsProcessorOutput):
|
def post_forward_mlp_sync_batch(self, logits_output: LogitsProcessorOutput):
|
||||||
|
|
||||||
self.forward_mode = getattr(self, "_original_forward_mode", self.forward_mode)
|
self.forward_mode = getattr(self, "_original_forward_mode", self.forward_mode)
|
||||||
self.batch_size = getattr(self, "_original_batch_size", self.batch_size)
|
self.batch_size = getattr(self, "_original_batch_size", self.batch_size)
|
||||||
bs = self.batch_size
|
bs = self.batch_size
|
||||||
@@ -1284,62 +1282,6 @@ def compute_position(
|
|||||||
return positions, extend_start_loc
|
return positions, extend_start_loc
|
||||||
|
|
||||||
|
|
||||||
def compute_position_triton(
|
|
||||||
extend_prefix_lens: torch.Tensor, extend_seq_lens: torch.Tensor, extend_seq_lens_sum
|
|
||||||
):
|
|
||||||
"""Compute positions. It is a fused version of `compute_position_torch`."""
|
|
||||||
batch_size = extend_seq_lens.shape[0]
|
|
||||||
has_prefix = extend_prefix_lens.shape[0] == batch_size
|
|
||||||
|
|
||||||
positions = torch.empty(
|
|
||||||
extend_seq_lens_sum, dtype=torch.int64, device=extend_seq_lens.device
|
|
||||||
)
|
|
||||||
extend_start_loc = torch.empty(
|
|
||||||
batch_size, dtype=torch.int32, device=extend_seq_lens.device
|
|
||||||
)
|
|
||||||
|
|
||||||
# Launch kernel
|
|
||||||
compute_position_kernel[(batch_size,)](
|
|
||||||
positions,
|
|
||||||
extend_start_loc,
|
|
||||||
extend_prefix_lens,
|
|
||||||
extend_seq_lens,
|
|
||||||
has_prefix,
|
|
||||||
)
|
|
||||||
|
|
||||||
return positions, extend_start_loc
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def compute_position_kernel(
|
|
||||||
positions,
|
|
||||||
extend_start_loc,
|
|
||||||
extend_prefix_lens,
|
|
||||||
extend_seq_lens,
|
|
||||||
has_prefix: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 512
|
|
||||||
pid = tl.program_id(0).to(tl.int64)
|
|
||||||
|
|
||||||
prefix_len = tl.load(extend_prefix_lens + pid) if has_prefix else 0
|
|
||||||
seq_len = tl.load(extend_seq_lens + pid)
|
|
||||||
|
|
||||||
# NOTE: This can be slow for large bs
|
|
||||||
cumsum_start = tl.cast(0, tl.int64)
|
|
||||||
for i in range(pid):
|
|
||||||
cumsum_start += tl.load(extend_seq_lens + i)
|
|
||||||
|
|
||||||
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
tl.store(
|
|
||||||
positions + cumsum_start + offset,
|
|
||||||
prefix_len + offset,
|
|
||||||
mask=offset < seq_len,
|
|
||||||
)
|
|
||||||
tl.store(extend_start_loc + pid, cumsum_start)
|
|
||||||
|
|
||||||
|
|
||||||
def compute_position_torch(
|
def compute_position_torch(
|
||||||
extend_prefix_lens: torch.Tensor, extend_seq_lens: torch.Tensor
|
extend_prefix_lens: torch.Tensor, extend_seq_lens: torch.Tensor
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
def compute_position_triton(
|
||||||
|
extend_prefix_lens: torch.Tensor, extend_seq_lens: torch.Tensor, extend_seq_lens_sum
|
||||||
|
):
|
||||||
|
"""Compute positions. It is a fused version of `compute_position_torch`."""
|
||||||
|
batch_size = extend_seq_lens.shape[0]
|
||||||
|
has_prefix = extend_prefix_lens.shape[0] == batch_size
|
||||||
|
|
||||||
|
positions = torch.empty(
|
||||||
|
extend_seq_lens_sum, dtype=torch.int64, device=extend_seq_lens.device
|
||||||
|
)
|
||||||
|
extend_start_loc = torch.empty(
|
||||||
|
batch_size, dtype=torch.int32, device=extend_seq_lens.device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Launch kernel
|
||||||
|
compute_position_kernel[(batch_size,)](
|
||||||
|
positions,
|
||||||
|
extend_start_loc,
|
||||||
|
extend_prefix_lens,
|
||||||
|
extend_seq_lens,
|
||||||
|
has_prefix,
|
||||||
|
)
|
||||||
|
|
||||||
|
return positions, extend_start_loc
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def compute_position_kernel(
|
||||||
|
positions,
|
||||||
|
extend_start_loc,
|
||||||
|
extend_prefix_lens,
|
||||||
|
extend_seq_lens,
|
||||||
|
has_prefix: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 512
|
||||||
|
pid = tl.program_id(0).to(tl.int64)
|
||||||
|
|
||||||
|
prefix_len = tl.load(extend_prefix_lens + pid) if has_prefix else 0
|
||||||
|
seq_len = tl.load(extend_seq_lens + pid)
|
||||||
|
|
||||||
|
# NOTE: This can be slow for large bs
|
||||||
|
cumsum_start = tl.cast(0, tl.int64)
|
||||||
|
for i in range(pid):
|
||||||
|
cumsum_start += tl.load(extend_seq_lens + i)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
tl.store(
|
||||||
|
positions + cumsum_start + offset,
|
||||||
|
prefix_len + offset,
|
||||||
|
mask=offset < seq_len,
|
||||||
|
)
|
||||||
|
tl.store(extend_start_loc + pid, cumsum_start)
|
||||||
@@ -5,8 +5,6 @@ from typing import TYPE_CHECKING, Any
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import get_tp_group
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
@@ -38,8 +36,23 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
SIMULATE_ACC_LEN,
|
SIMULATE_ACC_LEN,
|
||||||
generate_simulated_accept_index,
|
generate_simulated_accept_index,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_draft_cache_locs_page_size_1 as assign_draft_cache_locs_page_size_1,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_extend_cache_locs as assign_extend_cache_locs,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.eagle import (
|
||||||
|
fill_accepted_out_cache_loc as fill_accepted_out_cache_loc,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.eagle import (
|
||||||
|
fill_bonus_tokens as fill_bonus_tokens,
|
||||||
|
)
|
||||||
from sglang.srt.utils.async_probe import maybe_detect_nan, maybe_detect_oob
|
from sglang.srt.utils.async_probe import maybe_detect_nan, maybe_detect_oob
|
||||||
from sglang.srt.utils.common import is_cuda, is_hip, is_musa, is_npu, next_power_of_2
|
from sglang.srt.utils.common import is_cuda, is_hip, is_musa, is_npu
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -61,33 +74,6 @@ if is_cuda() or is_musa():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_draft_cache_locs_page_size_1(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
seq_lens,
|
|
||||||
out_cache_loc,
|
|
||||||
pool_len: tl.constexpr,
|
|
||||||
topk: tl.constexpr,
|
|
||||||
speculative_num_steps: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 128
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
|
|
||||||
copy_len = topk * speculative_num_steps
|
|
||||||
out_cache_ptr = out_cache_loc + pid * topk * speculative_num_steps
|
|
||||||
|
|
||||||
# Copy from req_to_token to out_cache_loc
|
|
||||||
kv_start = tl.load(seq_lens + pid)
|
|
||||||
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
|
||||||
num_loop = tl.cdiv(copy_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
copy_offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
mask = copy_offset < copy_len
|
|
||||||
data = tl.load(token_pool + kv_start + copy_offset, mask=mask)
|
|
||||||
tl.store(out_cache_ptr + copy_offset, data, mask=mask)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EagleDraftInputV2Mixin:
|
class EagleDraftInputV2Mixin:
|
||||||
def prepare_for_decode(self: EagleDraftInput, batch: ScheduleBatch):
|
def prepare_for_decode(self: EagleDraftInput, batch: ScheduleBatch):
|
||||||
@@ -499,118 +485,3 @@ class EagleVerifyInputV2Mixin:
|
|||||||
# tensor includes the trailing/bonus token via out-of-place +1 so the
|
# tensor includes the trailing/bonus token via out-of-place +1 so the
|
||||||
# name no longer flips semantics mid-function (naming doc C2).
|
# name no longer flips semantics mid-function (naming doc C2).
|
||||||
return predict, num_correct_drafts + 1, accept_index
|
return predict, num_correct_drafts + 1, accept_index
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def fill_bonus_tokens(
|
|
||||||
accept_tokens,
|
|
||||||
accept_lens,
|
|
||||||
bonus_tokens_ptr,
|
|
||||||
num_draft_tokens: tl.constexpr,
|
|
||||||
):
|
|
||||||
# NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel
|
|
||||||
# because this kernel reads accept_lens
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
# `accept_lens` includes the bonus token; the last accepted slot is at -1.
|
|
||||||
accept_len = tl.load(accept_lens + pid)
|
|
||||||
|
|
||||||
bonus_token_idx = num_draft_tokens * pid + accept_len - 1
|
|
||||||
bonus_token = tl.load(accept_tokens + bonus_token_idx)
|
|
||||||
tl.store(bonus_tokens_ptr + pid, bonus_token)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def fill_accepted_out_cache_loc(
|
|
||||||
accept_index,
|
|
||||||
out_cache_loc,
|
|
||||||
accepted_out_cache_loc,
|
|
||||||
size_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
offset = tl.arange(0, size_upper)
|
|
||||||
|
|
||||||
masks = (tl.load(accept_index + offset, offset < pid, other=-1) != -1).to(tl.int64)
|
|
||||||
dst = tl.sum(masks)
|
|
||||||
src = tl.load(accept_index + pid)
|
|
||||||
if src > -1:
|
|
||||||
value = tl.load(out_cache_loc + src)
|
|
||||||
tl.store(accepted_out_cache_loc + dst, value)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_extend_cache_locs(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
pool_len: tl.constexpr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 32
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
kv_start = tl.load(start_offset + pid)
|
|
||||||
kv_end = tl.load(end_offset + pid)
|
|
||||||
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
|
||||||
|
|
||||||
length_offset = tl.arange(0, bs_upper)
|
|
||||||
start = tl.load(start_offset + length_offset, mask=length_offset < pid, other=0)
|
|
||||||
end = tl.load(end_offset + length_offset, mask=length_offset < pid, other=0)
|
|
||||||
out_offset = tl.sum(end - start, axis=0)
|
|
||||||
|
|
||||||
out_cache_ptr = out_cache_loc + out_offset
|
|
||||||
|
|
||||||
load_offset = tl.arange(0, BLOCK_SIZE) + kv_start
|
|
||||||
save_offset = tl.arange(0, BLOCK_SIZE)
|
|
||||||
|
|
||||||
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
|
||||||
for _ in range(num_loop):
|
|
||||||
mask = load_offset < kv_end
|
|
||||||
data = tl.load(token_pool + load_offset, mask=mask)
|
|
||||||
tl.store(out_cache_ptr + save_offset, data, mask=mask)
|
|
||||||
load_offset += BLOCK_SIZE
|
|
||||||
save_offset += BLOCK_SIZE
|
|
||||||
|
|
||||||
|
|
||||||
def assign_extend_cache_locs_func(
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
start_offset: torch.Tensor,
|
|
||||||
end_offset: torch.Tensor,
|
|
||||||
batch_size: int,
|
|
||||||
draft_token_num: int,
|
|
||||||
device,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
if _is_cuda or _is_hip or _is_musa:
|
|
||||||
out_cache_loc = torch.empty(
|
|
||||||
(batch_size * draft_token_num,),
|
|
||||||
dtype=torch.int64,
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
assign_extend_cache_locs[(batch_size,)](
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
req_to_token.shape[1],
|
|
||||||
next_power_of_2(batch_size),
|
|
||||||
)
|
|
||||||
|
|
||||||
return out_cache_loc
|
|
||||||
|
|
||||||
elif _is_npu:
|
|
||||||
out_cache_loc = torch.empty(
|
|
||||||
(batch_size * draft_token_num,),
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
torch.ops.npu.cache_loc_update(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
)
|
|
||||||
|
|
||||||
return out_cache_loc
|
|
||||||
|
|||||||
@@ -12,339 +12,22 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
import torch
|
from sglang.srt.speculative.triton_ops.multi_layer_eagle import (
|
||||||
import triton
|
assign_hidden_states_pool_kernel,
|
||||||
import triton.language as tl
|
assign_hidden_states_pool_torch,
|
||||||
|
assign_hidden_states_pool_triton,
|
||||||
|
assign_new_state_kernel,
|
||||||
|
assign_new_state_triton,
|
||||||
|
rotate_input_ids_kernel,
|
||||||
|
rotate_input_ids_triton,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
@triton.jit
|
"assign_hidden_states_pool_kernel",
|
||||||
def rotate_input_ids_kernel(
|
"assign_hidden_states_pool_torch",
|
||||||
input_ids_ptr,
|
"assign_hidden_states_pool_triton",
|
||||||
extend_start_loc_ptr,
|
"assign_new_state_kernel",
|
||||||
extend_seq_lens_ptr,
|
"assign_new_state_triton",
|
||||||
topk_index_ptr,
|
"rotate_input_ids_kernel",
|
||||||
select_index_ptr,
|
"rotate_input_ids_triton",
|
||||||
BLOCK_SIZE: tl.constexpr,
|
]
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
|
|
||||||
start_loc = tl.load(extend_start_loc_ptr + pid)
|
|
||||||
seq_len = tl.load(extend_seq_lens_ptr + pid)
|
|
||||||
new_token = tl.load(topk_index_ptr + pid)
|
|
||||||
|
|
||||||
num_elements_to_shift = seq_len - 1
|
|
||||||
|
|
||||||
for off in range(0, num_elements_to_shift, BLOCK_SIZE):
|
|
||||||
offsets = off + tl.arange(0, BLOCK_SIZE)
|
|
||||||
mask = offsets < num_elements_to_shift
|
|
||||||
|
|
||||||
read_ptr = input_ids_ptr + start_loc + offsets + 1
|
|
||||||
val = tl.load(read_ptr, mask=mask)
|
|
||||||
tl.debug_barrier()
|
|
||||||
|
|
||||||
write_ptr = input_ids_ptr + start_loc + offsets
|
|
||||||
tl.store(write_ptr, val, mask=mask)
|
|
||||||
tl.debug_barrier()
|
|
||||||
|
|
||||||
if seq_len > 0:
|
|
||||||
if select_index_ptr is not None:
|
|
||||||
last_pos_ptr = input_ids_ptr + tl.load(select_index_ptr + pid)
|
|
||||||
else:
|
|
||||||
last_pos_ptr = input_ids_ptr + start_loc + seq_len - 1
|
|
||||||
tl.store(last_pos_ptr, new_token)
|
|
||||||
|
|
||||||
|
|
||||||
def rotate_input_ids_triton(
|
|
||||||
input_ids, extend_start_loc, extend_seq_lens, topk_index, select_index=None
|
|
||||||
):
|
|
||||||
batch_size = extend_seq_lens.shape[0]
|
|
||||||
BLOCK_SIZE = 4096 if select_index is not None else 8
|
|
||||||
grid = (batch_size,)
|
|
||||||
|
|
||||||
rotate_input_ids_kernel[grid](
|
|
||||||
input_ids,
|
|
||||||
extend_start_loc,
|
|
||||||
extend_seq_lens,
|
|
||||||
topk_index,
|
|
||||||
select_index,
|
|
||||||
BLOCK_SIZE=BLOCK_SIZE,
|
|
||||||
)
|
|
||||||
return input_ids
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_new_state_kernel(
|
|
||||||
# Source pointers
|
|
||||||
old_input_ids_ptr,
|
|
||||||
old_positions_ptr,
|
|
||||||
old_hidden_states_ptr,
|
|
||||||
old_out_cache_loc_ptr,
|
|
||||||
old_extend_seq_lens_ptr,
|
|
||||||
old_extend_start_loc_ptr,
|
|
||||||
# Destination pointers
|
|
||||||
input_ids_ptr,
|
|
||||||
positions_ptr,
|
|
||||||
hidden_states_ptr,
|
|
||||||
out_cache_loc_ptr,
|
|
||||||
extend_seq_lens_ptr,
|
|
||||||
extend_start_loc_ptr,
|
|
||||||
# Auxiliary data pointers
|
|
||||||
next_token_ids_ptr,
|
|
||||||
seq_lens_ptr,
|
|
||||||
padding_lens_ptr,
|
|
||||||
req_pool_indices_ptr,
|
|
||||||
req_to_token_ptr,
|
|
||||||
req_to_hidden_states_pool_ptr,
|
|
||||||
# Scalars and Strides
|
|
||||||
step,
|
|
||||||
stride_hidden_seq,
|
|
||||||
stride_hidden_dim, # hidden_states strides
|
|
||||||
stride_pool_req,
|
|
||||||
stride_pool_step,
|
|
||||||
stride_pool_dim, # pool strides
|
|
||||||
stride_req_token_0,
|
|
||||||
stride_req_token_1, # req_to_token strides
|
|
||||||
# Meta-parameters
|
|
||||||
HIDDEN_DIM: tl.constexpr,
|
|
||||||
BLOCK_SEQ: tl.constexpr,
|
|
||||||
BLOCK_HID: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
|
|
||||||
seq_len: tl.tensor = tl.load(seq_lens_ptr + pid)
|
|
||||||
old_extend_len = tl.load(old_extend_seq_lens_ptr + pid)
|
|
||||||
old_start = tl.load(old_extend_start_loc_ptr + pid)
|
|
||||||
new_extend_len = old_extend_len + 1
|
|
||||||
new_start = old_start + pid
|
|
||||||
|
|
||||||
tl.store(extend_seq_lens_ptr + pid, new_extend_len)
|
|
||||||
tl.store(extend_start_loc_ptr + pid, new_start)
|
|
||||||
|
|
||||||
offs_seq = tl.arange(0, BLOCK_SEQ)
|
|
||||||
mask_seq = offs_seq < old_extend_len
|
|
||||||
|
|
||||||
old_ids = tl.load(old_input_ids_ptr + old_start + offs_seq, mask=mask_seq)
|
|
||||||
tl.store(input_ids_ptr + new_start + offs_seq, old_ids, mask=mask_seq)
|
|
||||||
padding_len = tl.load(padding_lens_ptr + pid)
|
|
||||||
tl.store(
|
|
||||||
input_ids_ptr + new_start + old_extend_len - padding_len,
|
|
||||||
tl.load(next_token_ids_ptr + pid),
|
|
||||||
)
|
|
||||||
|
|
||||||
old_pos = tl.load(old_positions_ptr + old_start + offs_seq, mask=mask_seq)
|
|
||||||
tl.store(positions_ptr + new_start + 1 + offs_seq, old_pos, mask=mask_seq)
|
|
||||||
tl.store(
|
|
||||||
positions_ptr + new_start, max(tl.load(old_positions_ptr + old_start) - 1, 0)
|
|
||||||
)
|
|
||||||
|
|
||||||
old_cache = tl.load(old_out_cache_loc_ptr + old_start + offs_seq, mask=mask_seq)
|
|
||||||
tl.store(out_cache_loc_ptr + new_start + 1 + offs_seq, old_cache, mask=mask_seq)
|
|
||||||
|
|
||||||
req_idx = tl.load(req_pool_indices_ptr + pid)
|
|
||||||
token_idx_col = seq_len - old_extend_len - 1
|
|
||||||
if token_idx_col >= 0:
|
|
||||||
req_token_ptr_loc = (
|
|
||||||
req_to_token_ptr
|
|
||||||
+ (req_idx * stride_req_token_0)
|
|
||||||
+ (token_idx_col * stride_req_token_1)
|
|
||||||
)
|
|
||||||
last_cache_loc = tl.load(req_token_ptr_loc)
|
|
||||||
tl.store(out_cache_loc_ptr + new_start, last_cache_loc)
|
|
||||||
|
|
||||||
pool_vec_offset_base = ((req_idx + 1) * stride_pool_req) + (
|
|
||||||
-(step + 1) * stride_pool_step
|
|
||||||
)
|
|
||||||
|
|
||||||
for off_h in range(0, HIDDEN_DIM, BLOCK_HID):
|
|
||||||
offs_h = off_h + tl.arange(0, BLOCK_HID)
|
|
||||||
mask_h = offs_h < HIDDEN_DIM
|
|
||||||
|
|
||||||
for i in range(BLOCK_SEQ):
|
|
||||||
if i < old_extend_len:
|
|
||||||
old_h_ptr = (
|
|
||||||
old_hidden_states_ptr
|
|
||||||
+ (old_start + i) * stride_hidden_seq
|
|
||||||
+ (offs_h * stride_hidden_dim)
|
|
||||||
)
|
|
||||||
new_h_ptr = (
|
|
||||||
hidden_states_ptr
|
|
||||||
+ (new_start + 1 + i) * stride_hidden_seq
|
|
||||||
+ (offs_h * stride_hidden_dim)
|
|
||||||
)
|
|
||||||
|
|
||||||
chunk_old = tl.load(old_h_ptr, mask=mask_h)
|
|
||||||
tl.store(new_h_ptr, chunk_old, mask=mask_h)
|
|
||||||
|
|
||||||
pool_ptrs = (
|
|
||||||
req_to_hidden_states_pool_ptr
|
|
||||||
+ pool_vec_offset_base
|
|
||||||
+ (offs_h * stride_pool_dim)
|
|
||||||
)
|
|
||||||
pool_val = tl.load(pool_ptrs, mask=mask_h)
|
|
||||||
|
|
||||||
new_h_start_ptrs = (
|
|
||||||
hidden_states_ptr
|
|
||||||
+ (new_start * stride_hidden_seq)
|
|
||||||
+ (offs_h * stride_hidden_dim)
|
|
||||||
)
|
|
||||||
tl.store(new_h_start_ptrs, pool_val, mask=mask_h)
|
|
||||||
|
|
||||||
|
|
||||||
def assign_new_state_triton(
|
|
||||||
next_token_ids: torch.Tensor,
|
|
||||||
old_input_ids: torch.Tensor,
|
|
||||||
old_positions: torch.Tensor,
|
|
||||||
old_hidden_states: torch.Tensor,
|
|
||||||
old_out_cache_loc: torch.Tensor,
|
|
||||||
old_extend_seq_lens: torch.Tensor,
|
|
||||||
old_extend_start_loc: torch.Tensor,
|
|
||||||
input_ids: torch.Tensor,
|
|
||||||
positions: torch.Tensor,
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
out_cache_loc: torch.Tensor,
|
|
||||||
extend_seq_lens: torch.Tensor,
|
|
||||||
extend_start_loc: torch.Tensor,
|
|
||||||
seq_lens: torch.Tensor,
|
|
||||||
padding_lens: torch.Tensor,
|
|
||||||
num_seqs: int,
|
|
||||||
step: int,
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
req_to_hidden_states_pool: torch.Tensor,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Wrapper function to calculate offsets and launch the Triton kernel.
|
|
||||||
"""
|
|
||||||
hidden_dim = hidden_states.shape[1]
|
|
||||||
|
|
||||||
BLOCK_SEQ = 8
|
|
||||||
BLOCK_HID = 64
|
|
||||||
|
|
||||||
grid = (num_seqs,)
|
|
||||||
|
|
||||||
assign_new_state_kernel[grid](
|
|
||||||
# Pointers
|
|
||||||
old_input_ids,
|
|
||||||
old_positions,
|
|
||||||
old_hidden_states,
|
|
||||||
old_out_cache_loc,
|
|
||||||
old_extend_seq_lens,
|
|
||||||
old_extend_start_loc,
|
|
||||||
input_ids,
|
|
||||||
positions,
|
|
||||||
hidden_states,
|
|
||||||
out_cache_loc,
|
|
||||||
extend_seq_lens,
|
|
||||||
extend_start_loc,
|
|
||||||
next_token_ids,
|
|
||||||
seq_lens,
|
|
||||||
padding_lens,
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
req_to_hidden_states_pool,
|
|
||||||
# Constants/Strides
|
|
||||||
step,
|
|
||||||
old_hidden_states.stride(0),
|
|
||||||
old_hidden_states.stride(1),
|
|
||||||
req_to_hidden_states_pool.stride(0),
|
|
||||||
req_to_hidden_states_pool.stride(1),
|
|
||||||
req_to_hidden_states_pool.stride(2),
|
|
||||||
req_to_token.stride(0),
|
|
||||||
req_to_token.stride(1),
|
|
||||||
# Meta
|
|
||||||
HIDDEN_DIM=hidden_dim,
|
|
||||||
BLOCK_SEQ=BLOCK_SEQ,
|
|
||||||
BLOCK_HID=BLOCK_HID,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_hidden_states_pool_kernel(
|
|
||||||
hidden_states_ptr,
|
|
||||||
req_pool_indices_ptr,
|
|
||||||
req_to_hidden_states_pool_ptr,
|
|
||||||
extend_seq_lens_ptr,
|
|
||||||
extend_start_loc_ptr,
|
|
||||||
stride_hidden_seq,
|
|
||||||
stride_hidden_dim,
|
|
||||||
stride_pool_req,
|
|
||||||
stride_pool_step,
|
|
||||||
stride_pool_dim,
|
|
||||||
HIDDEN_DIM: tl.constexpr,
|
|
||||||
pool_size: tl.constexpr,
|
|
||||||
BLOCK_HID: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
|
|
||||||
extend_len = tl.load(extend_seq_lens_ptr + pid)
|
|
||||||
start_loc = tl.load(extend_start_loc_ptr + pid)
|
|
||||||
end_loc = start_loc + extend_len
|
|
||||||
|
|
||||||
req_idx = tl.load(req_pool_indices_ptr + pid)
|
|
||||||
pool_vec_offset_base = req_idx * stride_pool_req
|
|
||||||
|
|
||||||
for i in range(pool_size):
|
|
||||||
for off_h in range(0, HIDDEN_DIM, BLOCK_HID):
|
|
||||||
offs_h = off_h + tl.arange(0, BLOCK_HID)
|
|
||||||
mask_h = offs_h < HIDDEN_DIM
|
|
||||||
|
|
||||||
hid_ptr = (
|
|
||||||
hidden_states_ptr
|
|
||||||
+ (end_loc - pool_size + i) * stride_hidden_seq
|
|
||||||
+ offs_h * stride_hidden_dim
|
|
||||||
)
|
|
||||||
hid_val = tl.load(hid_ptr, mask=mask_h)
|
|
||||||
|
|
||||||
pool_ptr = (
|
|
||||||
req_to_hidden_states_pool_ptr
|
|
||||||
+ pool_vec_offset_base
|
|
||||||
+ i * stride_pool_step
|
|
||||||
+ offs_h * stride_pool_dim
|
|
||||||
)
|
|
||||||
tl.store(pool_ptr, hid_val, mask=mask_h)
|
|
||||||
|
|
||||||
|
|
||||||
def assign_hidden_states_pool_triton(
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_hidden_states_pool: torch.Tensor,
|
|
||||||
pool_size: int,
|
|
||||||
num_seqs: int,
|
|
||||||
extend_seq_lens: torch.Tensor,
|
|
||||||
extend_start_loc: torch.Tensor,
|
|
||||||
):
|
|
||||||
grid = (num_seqs,)
|
|
||||||
assign_hidden_states_pool_kernel[grid](
|
|
||||||
hidden_states,
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_hidden_states_pool,
|
|
||||||
extend_seq_lens,
|
|
||||||
extend_start_loc,
|
|
||||||
hidden_states.stride(0),
|
|
||||||
hidden_states.stride(1),
|
|
||||||
req_to_hidden_states_pool.stride(0),
|
|
||||||
req_to_hidden_states_pool.stride(1),
|
|
||||||
req_to_hidden_states_pool.stride(2),
|
|
||||||
HIDDEN_DIM=hidden_states.shape[1],
|
|
||||||
pool_size=pool_size,
|
|
||||||
BLOCK_HID=64,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def assign_hidden_states_pool_torch(
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_hidden_states_pool: torch.Tensor,
|
|
||||||
pool_size: int,
|
|
||||||
num_seqs: int,
|
|
||||||
extend_seq_lens: torch.Tensor,
|
|
||||||
extend_start_loc: torch.Tensor,
|
|
||||||
):
|
|
||||||
for req in range(num_seqs):
|
|
||||||
pool_idx = req_pool_indices[req]
|
|
||||||
extend_len = extend_seq_lens[req]
|
|
||||||
start_loc = extend_start_loc[req]
|
|
||||||
end_loc = start_loc + extend_len
|
|
||||||
req_to_hidden_states_pool[pool_idx, :pool_size, :].copy_(
|
|
||||||
hidden_states[end_loc - pool_size : end_loc, :]
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -7,8 +7,6 @@ from contextlib import contextmanager
|
|||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
from huggingface_hub import snapshot_download
|
from huggingface_hub import snapshot_download
|
||||||
|
|
||||||
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
||||||
@@ -20,7 +18,34 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
from sglang.srt.mem_cache.common import get_last_loc
|
from sglang.srt.mem_cache.common import get_last_loc
|
||||||
from sglang.srt.server_args import ServerArgs, get_global_server_args
|
from sglang.srt.server_args import ServerArgs, get_global_server_args
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, next_power_of_2
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
align_evict_mask_to_page_size as align_evict_mask_to_page_size,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_draft_cache_locs as assign_draft_cache_locs,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_req_to_token_pool as assign_req_to_token_pool,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
assign_req_to_token_pool_func as assign_req_to_token_pool_func,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
create_extend_after_decode_spec_info as create_extend_after_decode_spec_info,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
filter_finished_cache_loc_kernel as filter_finished_cache_loc_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
generate_draft_decode_kv_indices as generate_draft_decode_kv_indices,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
get_src_tgt_cache_loc as get_src_tgt_cache_loc,
|
||||||
|
)
|
||||||
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
|
get_target_cache_loc as get_target_cache_loc,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -111,403 +136,6 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool:
|
|||||||
return not server_args.enable_multi_layer_eagle
|
return not server_args.enable_multi_layer_eagle
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def create_extend_after_decode_spec_info(
|
|
||||||
accept_tokens,
|
|
||||||
seq_lens,
|
|
||||||
accept_lens,
|
|
||||||
positions,
|
|
||||||
bonus_tokens_ptr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
offsets = tl.arange(0, bs_upper)
|
|
||||||
seq_length = tl.load(seq_lens + pid)
|
|
||||||
# `accept_lens` includes the bonus token; load this req's value.
|
|
||||||
accept_len = tl.load(accept_lens + pid)
|
|
||||||
|
|
||||||
accept_len_cumsum = tl.sum(
|
|
||||||
tl.load(accept_lens + offsets, mask=offsets < pid, other=0)
|
|
||||||
)
|
|
||||||
positions_ptr = positions + accept_len_cumsum
|
|
||||||
mask = offsets < accept_len
|
|
||||||
tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask)
|
|
||||||
|
|
||||||
accept_len_cumsum += accept_len - 1
|
|
||||||
bonus_token = tl.load(accept_tokens + accept_len_cumsum)
|
|
||||||
tl.store(bonus_tokens_ptr + pid, bonus_token)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_req_to_token_pool(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
pool_len: tl.constexpr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 32
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
kv_start = tl.load(start_offset + pid)
|
|
||||||
kv_end = tl.load(end_offset + pid)
|
|
||||||
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
|
||||||
|
|
||||||
length_offset = tl.arange(0, bs_upper)
|
|
||||||
start = tl.load(start_offset + length_offset, mask=length_offset < pid, other=0)
|
|
||||||
end = tl.load(end_offset + length_offset, mask=length_offset < pid, other=0)
|
|
||||||
out_offset = tl.sum(end - start, axis=0)
|
|
||||||
|
|
||||||
out_cache_ptr = out_cache_loc + out_offset
|
|
||||||
|
|
||||||
save_offset = tl.arange(0, BLOCK_SIZE) + kv_start
|
|
||||||
load_offset = tl.arange(0, BLOCK_SIZE)
|
|
||||||
|
|
||||||
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
|
||||||
for _ in range(num_loop):
|
|
||||||
mask = save_offset < kv_end
|
|
||||||
data = tl.load(out_cache_ptr + load_offset, mask=mask)
|
|
||||||
tl.store(token_pool + save_offset, data, mask=mask)
|
|
||||||
save_offset += BLOCK_SIZE
|
|
||||||
load_offset += BLOCK_SIZE
|
|
||||||
|
|
||||||
|
|
||||||
def assign_req_to_token_pool_func(
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
start_offset: torch.Tensor,
|
|
||||||
end_offset: torch.Tensor,
|
|
||||||
out_cache_loc: torch.Tensor,
|
|
||||||
batch_size: int,
|
|
||||||
):
|
|
||||||
assign_req_to_token_pool[(batch_size,)](
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
req_to_token.shape[1],
|
|
||||||
next_power_of_2(batch_size),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_draft_cache_locs(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
seq_lens,
|
|
||||||
extend_lens,
|
|
||||||
num_new_pages_per_topk,
|
|
||||||
out_cache_loc,
|
|
||||||
source_cache_loc,
|
|
||||||
target_cache_loc,
|
|
||||||
last_page_lens_cumsum,
|
|
||||||
duplicate_cache_len: tl.constexpr,
|
|
||||||
pool_len: tl.constexpr,
|
|
||||||
topk: tl.constexpr,
|
|
||||||
speculative_num_steps: tl.constexpr,
|
|
||||||
page_size: tl.constexpr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
iter_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 128
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
|
|
||||||
if page_size == 1 or topk == 1:
|
|
||||||
copy_len = topk * speculative_num_steps
|
|
||||||
out_cache_ptr = out_cache_loc + pid * topk * speculative_num_steps
|
|
||||||
else:
|
|
||||||
bs_offset = tl.arange(0, bs_upper)
|
|
||||||
copy_len = tl.load(extend_lens + pid)
|
|
||||||
cum_copy_len = tl.sum(tl.load(extend_lens + bs_offset, mask=bs_offset < pid))
|
|
||||||
out_cache_ptr = out_cache_loc + cum_copy_len
|
|
||||||
|
|
||||||
# Part 1: Copy from out_cache_loc to req_to_token
|
|
||||||
kv_start = tl.load(seq_lens + pid)
|
|
||||||
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
|
||||||
num_loop = tl.cdiv(copy_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
copy_offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
mask = copy_offset < copy_len
|
|
||||||
data = tl.load(out_cache_ptr + copy_offset, mask=mask)
|
|
||||||
tl.store(token_pool + kv_start + copy_offset, data, mask=mask)
|
|
||||||
# XXX (MUSA): Triton issue: chained boolean operators (A or B or C) are not supported.
|
|
||||||
if (page_size != 1 and topk != 1) and duplicate_cache_len > 0:
|
|
||||||
# Part 2: Copy indices into source_cache_loc and target_cache_loc
|
|
||||||
# Expected output: src:[8,9,10,8,9,10...] tgt:[16,17,18,24,25,26...]
|
|
||||||
prefix_len = tl.load(seq_lens + pid)
|
|
||||||
last_page_len = prefix_len % page_size
|
|
||||||
offsets = tl.arange(0, page_size)
|
|
||||||
mask = offsets < last_page_len
|
|
||||||
num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid)
|
|
||||||
prefix_base = token_pool + prefix_len - last_page_len
|
|
||||||
src_indices = tl.load(prefix_base + offsets, mask=mask)
|
|
||||||
last_page_lens_cumsum_ = tl.load(last_page_lens_cumsum + pid)
|
|
||||||
# Skip the first one since no copy is needed
|
|
||||||
for topk_id in range(1, topk):
|
|
||||||
tl.store(
|
|
||||||
source_cache_loc
|
|
||||||
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
|
|
||||||
+ (topk_id - 1) * last_page_len
|
|
||||||
+ offsets,
|
|
||||||
src_indices,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
tgt_indices = tl.load(
|
|
||||||
prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
target_cache_loc
|
|
||||||
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
|
|
||||||
+ (topk_id - 1) * last_page_len
|
|
||||||
+ offsets,
|
|
||||||
tgt_indices,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
# Part 3: Copy and remove the used indices for duplication
|
|
||||||
# speculative_num_steps=5, page_size=4, num_new_pages_per_topk_=2, last_page_len=1
|
|
||||||
# - xxxxx .. | - xxxxx .. |
|
|
||||||
# topk=0 topk=1
|
|
||||||
# "-" means prefix tokens
|
|
||||||
# "x" means speculative draft tokens
|
|
||||||
# "." means padded tokens
|
|
||||||
# we only want to copy the "x" part.
|
|
||||||
iter_offset = tl.arange(0, iter_upper)
|
|
||||||
for topk_id in range(topk):
|
|
||||||
mask_upper = iter_offset < (speculative_num_steps + last_page_len)
|
|
||||||
mask_lower = iter_offset >= last_page_len
|
|
||||||
combined_mask = mask_upper & mask_lower
|
|
||||||
indices = tl.load(
|
|
||||||
prefix_base
|
|
||||||
+ topk_id * num_new_pages_per_topk_ * page_size
|
|
||||||
+ iter_offset,
|
|
||||||
mask=combined_mask,
|
|
||||||
other=0,
|
|
||||||
)
|
|
||||||
# Shift from previous batches
|
|
||||||
ptr_offset = pid * speculative_num_steps * topk
|
|
||||||
# Subtract last_page_len to fill the gap of duplicated last page tokens.
|
|
||||||
# For example, token pool is (1, 2, 3, 4 ,5) and last page is 1,
|
|
||||||
# we write 2, 3, 4 to the front of out_cache_loc.
|
|
||||||
tl.store(
|
|
||||||
out_cache_loc
|
|
||||||
+ ptr_offset
|
|
||||||
+ topk_id * speculative_num_steps
|
|
||||||
- last_page_len
|
|
||||||
+ iter_offset,
|
|
||||||
indices,
|
|
||||||
mask=combined_mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def generate_draft_decode_kv_indices(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
paged_kernel_lens,
|
|
||||||
kv_indices,
|
|
||||||
kv_indptr,
|
|
||||||
positions,
|
|
||||||
pool_len: tl.constexpr,
|
|
||||||
kv_indices_stride: tl.constexpr,
|
|
||||||
kv_indptr_stride: tl.constexpr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
iter_upper: tl.constexpr,
|
|
||||||
num_tokens_upper: tl.constexpr,
|
|
||||||
page_size: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 128
|
|
||||||
iters = tl.program_id(axis=0)
|
|
||||||
bid = tl.program_id(axis=1)
|
|
||||||
topk_id = tl.program_id(axis=2)
|
|
||||||
|
|
||||||
num_steps = tl.num_programs(axis=0)
|
|
||||||
num_seqs = tl.num_programs(axis=1)
|
|
||||||
topk = tl.num_programs(axis=2)
|
|
||||||
|
|
||||||
kv_indices += kv_indices_stride * iters
|
|
||||||
kv_indptr += kv_indptr_stride * iters
|
|
||||||
iters += 1
|
|
||||||
|
|
||||||
load_offset = tl.arange(0, bs_upper)
|
|
||||||
seq_lens = tl.load(paged_kernel_lens + load_offset, mask=load_offset < bid, other=0)
|
|
||||||
seq_len = tl.load(paged_kernel_lens + bid)
|
|
||||||
cum_seq_len = tl.sum(seq_lens)
|
|
||||||
|
|
||||||
# Update kv_indices
|
|
||||||
kv_offset = cum_seq_len * topk + bid * iters * topk + topk_id * (seq_len + iters)
|
|
||||||
kv_ptr = kv_indices + kv_offset
|
|
||||||
token_pool_ptr = req_to_token + tl.load(req_pool_indices + bid) * pool_len
|
|
||||||
|
|
||||||
kv_offset = tl.arange(0, BLOCK_SIZE)
|
|
||||||
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
|
||||||
for _ in range(num_loop):
|
|
||||||
mask = kv_offset < seq_len
|
|
||||||
data = tl.load(token_pool_ptr + kv_offset, mask=mask)
|
|
||||||
tl.store(kv_ptr + kv_offset, data, mask=mask)
|
|
||||||
kv_offset += BLOCK_SIZE
|
|
||||||
|
|
||||||
extend_offset = tl.arange(0, iter_upper)
|
|
||||||
if page_size == 1 or topk == 1:
|
|
||||||
extend_data = tl.load(
|
|
||||||
token_pool_ptr + seq_len + topk_id * num_steps + tl.arange(0, iter_upper),
|
|
||||||
mask=extend_offset < iters,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
prefix_len = seq_len
|
|
||||||
last_page_len = prefix_len % page_size
|
|
||||||
num_new_pages_per_topk = (
|
|
||||||
last_page_len + num_steps + page_size - 1
|
|
||||||
) // page_size
|
|
||||||
prefix_base = seq_len // page_size * page_size
|
|
||||||
start = (
|
|
||||||
prefix_base + topk_id * num_new_pages_per_topk * page_size + last_page_len
|
|
||||||
)
|
|
||||||
extend_data = tl.load(
|
|
||||||
token_pool_ptr + start + extend_offset,
|
|
||||||
mask=extend_offset < iters,
|
|
||||||
)
|
|
||||||
|
|
||||||
tl.store(kv_ptr + seq_len + extend_offset, extend_data, mask=extend_offset < iters)
|
|
||||||
|
|
||||||
# Update kv_indptr
|
|
||||||
bs_offset = tl.arange(0, num_tokens_upper)
|
|
||||||
|
|
||||||
zid = bid * topk + topk_id
|
|
||||||
if zid == 0:
|
|
||||||
zid = num_seqs * topk
|
|
||||||
positions = tl.load(positions + bs_offset, mask=bs_offset < zid, other=0)
|
|
||||||
base = tl.sum(positions)
|
|
||||||
tl.store(kv_indptr + zid, base + zid * iters)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def align_evict_mask_to_page_size(
|
|
||||||
seq_lens,
|
|
||||||
evict_mask,
|
|
||||||
page_size: tl.constexpr,
|
|
||||||
num_draft_tokens: tl.constexpr,
|
|
||||||
BLOCK_SIZE: tl.constexpr,
|
|
||||||
):
|
|
||||||
t_range = tl.arange(0, BLOCK_SIZE)
|
|
||||||
|
|
||||||
bid = tl.program_id(axis=0)
|
|
||||||
seq_len = tl.load(seq_lens + bid)
|
|
||||||
io_mask = t_range < num_draft_tokens
|
|
||||||
mask_row = tl.load(
|
|
||||||
evict_mask + bid * num_draft_tokens + t_range, mask=io_mask, other=0
|
|
||||||
)
|
|
||||||
|
|
||||||
num_trues = tl.sum(mask_row)
|
|
||||||
num_false = num_draft_tokens - num_trues
|
|
||||||
|
|
||||||
start = (seq_len + num_false - 1) // page_size * page_size - seq_len
|
|
||||||
for i in range(max(start, 0), min(start + page_size, num_draft_tokens)):
|
|
||||||
tl.store(evict_mask + bid * num_draft_tokens + i, False)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def get_target_cache_loc(
|
|
||||||
tgt_cache_loc,
|
|
||||||
to_free_slots,
|
|
||||||
num_correct_drafts,
|
|
||||||
to_free_num_slots,
|
|
||||||
out_cache_loc,
|
|
||||||
num_verify_tokens: tl.constexpr,
|
|
||||||
num_verify_tokens_upper: tl.constexpr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
bid = tl.program_id(axis=0)
|
|
||||||
offset = tl.arange(0, num_verify_tokens_upper)
|
|
||||||
bs_offset = tl.arange(0, bs_upper)
|
|
||||||
|
|
||||||
# write the first part to tgt_cache_loc
|
|
||||||
accept_len_all = tl.load(num_correct_drafts + bs_offset, mask=bs_offset < bid)
|
|
||||||
tgt_cache_loc_start = tl.sum(accept_len_all) + bid
|
|
||||||
copy_len = tl.load(num_correct_drafts + bid) + 1
|
|
||||||
out_cache_loc_row = tl.load(
|
|
||||||
out_cache_loc + bid * num_verify_tokens + offset, mask=offset < copy_len
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
tgt_cache_loc + tgt_cache_loc_start + offset,
|
|
||||||
out_cache_loc_row,
|
|
||||||
mask=offset < copy_len,
|
|
||||||
)
|
|
||||||
|
|
||||||
# write the second part to to_free_num_pages
|
|
||||||
to_free_num_slots_all = tl.load(to_free_num_slots + bs_offset, mask=bs_offset < bid)
|
|
||||||
to_free_num_slots_cur = tl.load(to_free_num_slots + bid)
|
|
||||||
out_cache_loc_start = num_verify_tokens - to_free_num_slots_cur
|
|
||||||
to_free_slots_start = tl.sum(to_free_num_slots_all)
|
|
||||||
|
|
||||||
copy_len = to_free_num_slots_cur
|
|
||||||
out_cache_loc_row = tl.load(
|
|
||||||
out_cache_loc + bid * num_verify_tokens + out_cache_loc_start + offset,
|
|
||||||
mask=offset < copy_len,
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
to_free_slots + to_free_slots_start + offset,
|
|
||||||
out_cache_loc_row,
|
|
||||||
mask=offset < copy_len,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True, disable=_is_npu)
|
|
||||||
def get_src_tgt_cache_loc(
|
|
||||||
seq_lens: torch.Tensor,
|
|
||||||
out_cache_loc: torch.Tensor,
|
|
||||||
accept_index: torch.Tensor,
|
|
||||||
num_correct_drafts: torch.Tensor,
|
|
||||||
draft_token_num: int,
|
|
||||||
page_size: int,
|
|
||||||
):
|
|
||||||
src_cache_loc = out_cache_loc[accept_index]
|
|
||||||
# zeros_like, not empty_like: any uncovered tail stays at slot 0 (padding)
|
|
||||||
# instead of caching-allocator garbage.
|
|
||||||
tgt_cache_loc = torch.zeros_like(src_cache_loc)
|
|
||||||
extended_len = seq_lens + draft_token_num
|
|
||||||
keep_len = torch.minimum(
|
|
||||||
(seq_lens + num_correct_drafts + 1 + page_size - 1) // page_size * page_size,
|
|
||||||
extended_len,
|
|
||||||
)
|
|
||||||
to_free_num_slots = extended_len - keep_len
|
|
||||||
return src_cache_loc, tgt_cache_loc, to_free_num_slots
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def filter_finished_cache_loc_kernel(
|
|
||||||
out_cache_loc,
|
|
||||||
tgt_cache_loc,
|
|
||||||
num_correct_drafts,
|
|
||||||
num_accept_tokens_filter,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
num_verify_tokens_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
bid = tl.program_id(0)
|
|
||||||
bs_offset = tl.arange(0, bs_upper)
|
|
||||||
|
|
||||||
num_correct_drafts_all = tl.load(
|
|
||||||
num_correct_drafts + bs_offset, mask=bs_offset < bid
|
|
||||||
)
|
|
||||||
old_start = tl.sum(num_correct_drafts_all) + bid
|
|
||||||
|
|
||||||
num_accept_tokens_filter_all = tl.load(
|
|
||||||
num_accept_tokens_filter + bs_offset, mask=bs_offset < bid
|
|
||||||
)
|
|
||||||
new_start = tl.sum(num_accept_tokens_filter_all)
|
|
||||||
|
|
||||||
copy_len = tl.load(num_accept_tokens_filter + bid)
|
|
||||||
copy_offset = tl.arange(0, num_verify_tokens_upper)
|
|
||||||
value = tl.load(
|
|
||||||
tgt_cache_loc + old_start + copy_offset, mask=copy_offset < copy_len
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
out_cache_loc + new_start + copy_offset, value, mask=copy_offset < copy_len
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True, disable=_is_npu)
|
@torch.compile(dynamic=True, disable=_is_npu)
|
||||||
def create_num_accept_tokens_filter(
|
def create_num_accept_tokens_filter(
|
||||||
num_correct_drafts: torch.Tensor,
|
num_correct_drafts: torch.Tensor,
|
||||||
|
|||||||
@@ -0,0 +1,515 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, next_power_of_2
|
||||||
|
|
||||||
|
_is_cuda = is_cuda()
|
||||||
|
_is_hip = is_hip()
|
||||||
|
_is_npu = is_npu()
|
||||||
|
_is_musa = is_musa()
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def create_extend_after_decode_spec_info(
|
||||||
|
accept_tokens,
|
||||||
|
seq_lens,
|
||||||
|
accept_lens,
|
||||||
|
positions,
|
||||||
|
bonus_tokens_ptr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
offsets = tl.arange(0, bs_upper)
|
||||||
|
seq_length = tl.load(seq_lens + pid)
|
||||||
|
# `accept_lens` includes the bonus token; load this req's value.
|
||||||
|
accept_len = tl.load(accept_lens + pid)
|
||||||
|
|
||||||
|
accept_len_cumsum = tl.sum(
|
||||||
|
tl.load(accept_lens + offsets, mask=offsets < pid, other=0)
|
||||||
|
)
|
||||||
|
positions_ptr = positions + accept_len_cumsum
|
||||||
|
mask = offsets < accept_len
|
||||||
|
tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask)
|
||||||
|
|
||||||
|
accept_len_cumsum += accept_len - 1
|
||||||
|
bonus_token = tl.load(accept_tokens + accept_len_cumsum)
|
||||||
|
tl.store(bonus_tokens_ptr + pid, bonus_token)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_req_to_token_pool(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
pool_len: tl.constexpr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 32
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
kv_start = tl.load(start_offset + pid)
|
||||||
|
kv_end = tl.load(end_offset + pid)
|
||||||
|
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
||||||
|
|
||||||
|
length_offset = tl.arange(0, bs_upper)
|
||||||
|
start = tl.load(start_offset + length_offset, mask=length_offset < pid, other=0)
|
||||||
|
end = tl.load(end_offset + length_offset, mask=length_offset < pid, other=0)
|
||||||
|
out_offset = tl.sum(end - start, axis=0)
|
||||||
|
|
||||||
|
out_cache_ptr = out_cache_loc + out_offset
|
||||||
|
|
||||||
|
save_offset = tl.arange(0, BLOCK_SIZE) + kv_start
|
||||||
|
load_offset = tl.arange(0, BLOCK_SIZE)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
||||||
|
for _ in range(num_loop):
|
||||||
|
mask = save_offset < kv_end
|
||||||
|
data = tl.load(out_cache_ptr + load_offset, mask=mask)
|
||||||
|
tl.store(token_pool + save_offset, data, mask=mask)
|
||||||
|
save_offset += BLOCK_SIZE
|
||||||
|
load_offset += BLOCK_SIZE
|
||||||
|
|
||||||
|
|
||||||
|
def assign_req_to_token_pool_func(
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
start_offset: torch.Tensor,
|
||||||
|
end_offset: torch.Tensor,
|
||||||
|
out_cache_loc: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
):
|
||||||
|
assign_req_to_token_pool[(batch_size,)](
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
req_to_token.shape[1],
|
||||||
|
next_power_of_2(batch_size),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_draft_cache_locs(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
seq_lens,
|
||||||
|
extend_lens,
|
||||||
|
num_new_pages_per_topk,
|
||||||
|
out_cache_loc,
|
||||||
|
source_cache_loc,
|
||||||
|
target_cache_loc,
|
||||||
|
last_page_lens_cumsum,
|
||||||
|
duplicate_cache_len: tl.constexpr,
|
||||||
|
pool_len: tl.constexpr,
|
||||||
|
topk: tl.constexpr,
|
||||||
|
speculative_num_steps: tl.constexpr,
|
||||||
|
page_size: tl.constexpr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
iter_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 128
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
|
||||||
|
if page_size == 1 or topk == 1:
|
||||||
|
copy_len = topk * speculative_num_steps
|
||||||
|
out_cache_ptr = out_cache_loc + pid * topk * speculative_num_steps
|
||||||
|
else:
|
||||||
|
bs_offset = tl.arange(0, bs_upper)
|
||||||
|
copy_len = tl.load(extend_lens + pid)
|
||||||
|
cum_copy_len = tl.sum(tl.load(extend_lens + bs_offset, mask=bs_offset < pid))
|
||||||
|
out_cache_ptr = out_cache_loc + cum_copy_len
|
||||||
|
|
||||||
|
# Part 1: Copy from out_cache_loc to req_to_token
|
||||||
|
kv_start = tl.load(seq_lens + pid)
|
||||||
|
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
||||||
|
num_loop = tl.cdiv(copy_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
copy_offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
mask = copy_offset < copy_len
|
||||||
|
data = tl.load(out_cache_ptr + copy_offset, mask=mask)
|
||||||
|
tl.store(token_pool + kv_start + copy_offset, data, mask=mask)
|
||||||
|
# XXX (MUSA): Triton issue: chained boolean operators (A or B or C) are not supported.
|
||||||
|
if (page_size != 1 and topk != 1) and duplicate_cache_len > 0:
|
||||||
|
# Part 2: Copy indices into source_cache_loc and target_cache_loc
|
||||||
|
# Expected output: src:[8,9,10,8,9,10...] tgt:[16,17,18,24,25,26...]
|
||||||
|
prefix_len = tl.load(seq_lens + pid)
|
||||||
|
last_page_len = prefix_len % page_size
|
||||||
|
offsets = tl.arange(0, page_size)
|
||||||
|
mask = offsets < last_page_len
|
||||||
|
num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid)
|
||||||
|
prefix_base = token_pool + prefix_len - last_page_len
|
||||||
|
src_indices = tl.load(prefix_base + offsets, mask=mask)
|
||||||
|
last_page_lens_cumsum_ = tl.load(last_page_lens_cumsum + pid)
|
||||||
|
# Skip the first one since no copy is needed
|
||||||
|
for topk_id in range(1, topk):
|
||||||
|
tl.store(
|
||||||
|
source_cache_loc
|
||||||
|
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
|
||||||
|
+ (topk_id - 1) * last_page_len
|
||||||
|
+ offsets,
|
||||||
|
src_indices,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
tgt_indices = tl.load(
|
||||||
|
prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
target_cache_loc
|
||||||
|
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
|
||||||
|
+ (topk_id - 1) * last_page_len
|
||||||
|
+ offsets,
|
||||||
|
tgt_indices,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
# Part 3: Copy and remove the used indices for duplication
|
||||||
|
# speculative_num_steps=5, page_size=4, num_new_pages_per_topk_=2, last_page_len=1
|
||||||
|
# - xxxxx .. | - xxxxx .. |
|
||||||
|
# topk=0 topk=1
|
||||||
|
# "-" means prefix tokens
|
||||||
|
# "x" means speculative draft tokens
|
||||||
|
# "." means padded tokens
|
||||||
|
# we only want to copy the "x" part.
|
||||||
|
iter_offset = tl.arange(0, iter_upper)
|
||||||
|
for topk_id in range(topk):
|
||||||
|
mask_upper = iter_offset < (speculative_num_steps + last_page_len)
|
||||||
|
mask_lower = iter_offset >= last_page_len
|
||||||
|
combined_mask = mask_upper & mask_lower
|
||||||
|
indices = tl.load(
|
||||||
|
prefix_base
|
||||||
|
+ topk_id * num_new_pages_per_topk_ * page_size
|
||||||
|
+ iter_offset,
|
||||||
|
mask=combined_mask,
|
||||||
|
other=0,
|
||||||
|
)
|
||||||
|
# Shift from previous batches
|
||||||
|
ptr_offset = pid * speculative_num_steps * topk
|
||||||
|
# Subtract last_page_len to fill the gap of duplicated last page tokens.
|
||||||
|
# For example, token pool is (1, 2, 3, 4 ,5) and last page is 1,
|
||||||
|
# we write 2, 3, 4 to the front of out_cache_loc.
|
||||||
|
tl.store(
|
||||||
|
out_cache_loc
|
||||||
|
+ ptr_offset
|
||||||
|
+ topk_id * speculative_num_steps
|
||||||
|
- last_page_len
|
||||||
|
+ iter_offset,
|
||||||
|
indices,
|
||||||
|
mask=combined_mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_draft_cache_locs_page_size_1(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
seq_lens,
|
||||||
|
out_cache_loc,
|
||||||
|
pool_len: tl.constexpr,
|
||||||
|
topk: tl.constexpr,
|
||||||
|
speculative_num_steps: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 128
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
|
||||||
|
copy_len = topk * speculative_num_steps
|
||||||
|
out_cache_ptr = out_cache_loc + pid * topk * speculative_num_steps
|
||||||
|
|
||||||
|
# Copy from req_to_token to out_cache_loc
|
||||||
|
kv_start = tl.load(seq_lens + pid)
|
||||||
|
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
||||||
|
num_loop = tl.cdiv(copy_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
copy_offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
mask = copy_offset < copy_len
|
||||||
|
data = tl.load(token_pool + kv_start + copy_offset, mask=mask)
|
||||||
|
tl.store(out_cache_ptr + copy_offset, data, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def generate_draft_decode_kv_indices(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
paged_kernel_lens,
|
||||||
|
kv_indices,
|
||||||
|
kv_indptr,
|
||||||
|
positions,
|
||||||
|
pool_len: tl.constexpr,
|
||||||
|
kv_indices_stride: tl.constexpr,
|
||||||
|
kv_indptr_stride: tl.constexpr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
iter_upper: tl.constexpr,
|
||||||
|
num_tokens_upper: tl.constexpr,
|
||||||
|
page_size: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 128
|
||||||
|
iters = tl.program_id(axis=0)
|
||||||
|
bid = tl.program_id(axis=1)
|
||||||
|
topk_id = tl.program_id(axis=2)
|
||||||
|
|
||||||
|
num_steps = tl.num_programs(axis=0)
|
||||||
|
num_seqs = tl.num_programs(axis=1)
|
||||||
|
topk = tl.num_programs(axis=2)
|
||||||
|
|
||||||
|
kv_indices += kv_indices_stride * iters
|
||||||
|
kv_indptr += kv_indptr_stride * iters
|
||||||
|
iters += 1
|
||||||
|
|
||||||
|
load_offset = tl.arange(0, bs_upper)
|
||||||
|
seq_lens = tl.load(paged_kernel_lens + load_offset, mask=load_offset < bid, other=0)
|
||||||
|
seq_len = tl.load(paged_kernel_lens + bid)
|
||||||
|
cum_seq_len = tl.sum(seq_lens)
|
||||||
|
|
||||||
|
# Update kv_indices
|
||||||
|
kv_offset = cum_seq_len * topk + bid * iters * topk + topk_id * (seq_len + iters)
|
||||||
|
kv_ptr = kv_indices + kv_offset
|
||||||
|
token_pool_ptr = req_to_token + tl.load(req_pool_indices + bid) * pool_len
|
||||||
|
|
||||||
|
kv_offset = tl.arange(0, BLOCK_SIZE)
|
||||||
|
num_loop = tl.cdiv(seq_len, BLOCK_SIZE)
|
||||||
|
for _ in range(num_loop):
|
||||||
|
mask = kv_offset < seq_len
|
||||||
|
data = tl.load(token_pool_ptr + kv_offset, mask=mask)
|
||||||
|
tl.store(kv_ptr + kv_offset, data, mask=mask)
|
||||||
|
kv_offset += BLOCK_SIZE
|
||||||
|
|
||||||
|
extend_offset = tl.arange(0, iter_upper)
|
||||||
|
if page_size == 1 or topk == 1:
|
||||||
|
extend_data = tl.load(
|
||||||
|
token_pool_ptr + seq_len + topk_id * num_steps + tl.arange(0, iter_upper),
|
||||||
|
mask=extend_offset < iters,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
prefix_len = seq_len
|
||||||
|
last_page_len = prefix_len % page_size
|
||||||
|
num_new_pages_per_topk = (
|
||||||
|
last_page_len + num_steps + page_size - 1
|
||||||
|
) // page_size
|
||||||
|
prefix_base = seq_len // page_size * page_size
|
||||||
|
start = (
|
||||||
|
prefix_base + topk_id * num_new_pages_per_topk * page_size + last_page_len
|
||||||
|
)
|
||||||
|
extend_data = tl.load(
|
||||||
|
token_pool_ptr + start + extend_offset,
|
||||||
|
mask=extend_offset < iters,
|
||||||
|
)
|
||||||
|
|
||||||
|
tl.store(kv_ptr + seq_len + extend_offset, extend_data, mask=extend_offset < iters)
|
||||||
|
|
||||||
|
# Update kv_indptr
|
||||||
|
bs_offset = tl.arange(0, num_tokens_upper)
|
||||||
|
|
||||||
|
zid = bid * topk + topk_id
|
||||||
|
if zid == 0:
|
||||||
|
zid = num_seqs * topk
|
||||||
|
positions = tl.load(positions + bs_offset, mask=bs_offset < zid, other=0)
|
||||||
|
base = tl.sum(positions)
|
||||||
|
tl.store(kv_indptr + zid, base + zid * iters)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def align_evict_mask_to_page_size(
|
||||||
|
seq_lens,
|
||||||
|
evict_mask,
|
||||||
|
page_size: tl.constexpr,
|
||||||
|
num_draft_tokens: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
t_range = tl.arange(0, BLOCK_SIZE)
|
||||||
|
|
||||||
|
bid = tl.program_id(axis=0)
|
||||||
|
seq_len = tl.load(seq_lens + bid)
|
||||||
|
io_mask = t_range < num_draft_tokens
|
||||||
|
mask_row = tl.load(
|
||||||
|
evict_mask + bid * num_draft_tokens + t_range, mask=io_mask, other=0
|
||||||
|
)
|
||||||
|
|
||||||
|
num_trues = tl.sum(mask_row)
|
||||||
|
num_false = num_draft_tokens - num_trues
|
||||||
|
|
||||||
|
start = (seq_len + num_false - 1) // page_size * page_size - seq_len
|
||||||
|
for i in range(max(start, 0), min(start + page_size, num_draft_tokens)):
|
||||||
|
tl.store(evict_mask + bid * num_draft_tokens + i, False)
|
||||||
|
|
||||||
|
|
||||||
|
@torch.compile(dynamic=True, disable=_is_npu)
|
||||||
|
def get_src_tgt_cache_loc(
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
out_cache_loc: torch.Tensor,
|
||||||
|
accept_index: torch.Tensor,
|
||||||
|
num_correct_drafts: torch.Tensor,
|
||||||
|
draft_token_num: int,
|
||||||
|
page_size: int,
|
||||||
|
):
|
||||||
|
src_cache_loc = out_cache_loc[accept_index]
|
||||||
|
# zeros_like, not empty_like: any uncovered tail stays at slot 0 (padding)
|
||||||
|
# instead of caching-allocator garbage.
|
||||||
|
tgt_cache_loc = torch.zeros_like(src_cache_loc)
|
||||||
|
extended_len = seq_lens + draft_token_num
|
||||||
|
keep_len = torch.minimum(
|
||||||
|
(seq_lens + num_correct_drafts + 1 + page_size - 1) // page_size * page_size,
|
||||||
|
extended_len,
|
||||||
|
)
|
||||||
|
to_free_num_slots = extended_len - keep_len
|
||||||
|
return src_cache_loc, tgt_cache_loc, to_free_num_slots
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def get_target_cache_loc(
|
||||||
|
tgt_cache_loc,
|
||||||
|
to_free_slots,
|
||||||
|
num_correct_drafts,
|
||||||
|
to_free_num_slots,
|
||||||
|
out_cache_loc,
|
||||||
|
num_verify_tokens: tl.constexpr,
|
||||||
|
num_verify_tokens_upper: tl.constexpr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
bid = tl.program_id(axis=0)
|
||||||
|
offset = tl.arange(0, num_verify_tokens_upper)
|
||||||
|
bs_offset = tl.arange(0, bs_upper)
|
||||||
|
|
||||||
|
# write the first part to tgt_cache_loc
|
||||||
|
accept_len_all = tl.load(num_correct_drafts + bs_offset, mask=bs_offset < bid)
|
||||||
|
tgt_cache_loc_start = tl.sum(accept_len_all) + bid
|
||||||
|
copy_len = tl.load(num_correct_drafts + bid) + 1
|
||||||
|
out_cache_loc_row = tl.load(
|
||||||
|
out_cache_loc + bid * num_verify_tokens + offset, mask=offset < copy_len
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
tgt_cache_loc + tgt_cache_loc_start + offset,
|
||||||
|
out_cache_loc_row,
|
||||||
|
mask=offset < copy_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
# write the second part to to_free_num_pages
|
||||||
|
to_free_num_slots_all = tl.load(to_free_num_slots + bs_offset, mask=bs_offset < bid)
|
||||||
|
to_free_num_slots_cur = tl.load(to_free_num_slots + bid)
|
||||||
|
out_cache_loc_start = num_verify_tokens - to_free_num_slots_cur
|
||||||
|
to_free_slots_start = tl.sum(to_free_num_slots_all)
|
||||||
|
|
||||||
|
copy_len = to_free_num_slots_cur
|
||||||
|
out_cache_loc_row = tl.load(
|
||||||
|
out_cache_loc + bid * num_verify_tokens + out_cache_loc_start + offset,
|
||||||
|
mask=offset < copy_len,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
to_free_slots + to_free_slots_start + offset,
|
||||||
|
out_cache_loc_row,
|
||||||
|
mask=offset < copy_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def filter_finished_cache_loc_kernel(
|
||||||
|
out_cache_loc,
|
||||||
|
tgt_cache_loc,
|
||||||
|
num_correct_drafts,
|
||||||
|
num_accept_tokens_filter,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
num_verify_tokens_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
bid = tl.program_id(0)
|
||||||
|
bs_offset = tl.arange(0, bs_upper)
|
||||||
|
|
||||||
|
num_correct_drafts_all = tl.load(
|
||||||
|
num_correct_drafts + bs_offset, mask=bs_offset < bid
|
||||||
|
)
|
||||||
|
old_start = tl.sum(num_correct_drafts_all) + bid
|
||||||
|
|
||||||
|
num_accept_tokens_filter_all = tl.load(
|
||||||
|
num_accept_tokens_filter + bs_offset, mask=bs_offset < bid
|
||||||
|
)
|
||||||
|
new_start = tl.sum(num_accept_tokens_filter_all)
|
||||||
|
|
||||||
|
copy_len = tl.load(num_accept_tokens_filter + bid)
|
||||||
|
copy_offset = tl.arange(0, num_verify_tokens_upper)
|
||||||
|
value = tl.load(
|
||||||
|
tgt_cache_loc + old_start + copy_offset, mask=copy_offset < copy_len
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
out_cache_loc + new_start + copy_offset, value, mask=copy_offset < copy_len
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_extend_cache_locs(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
pool_len: tl.constexpr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 32
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
kv_start = tl.load(start_offset + pid)
|
||||||
|
kv_end = tl.load(end_offset + pid)
|
||||||
|
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
||||||
|
|
||||||
|
length_offset = tl.arange(0, bs_upper)
|
||||||
|
start = tl.load(start_offset + length_offset, mask=length_offset < pid, other=0)
|
||||||
|
end = tl.load(end_offset + length_offset, mask=length_offset < pid, other=0)
|
||||||
|
out_offset = tl.sum(end - start, axis=0)
|
||||||
|
|
||||||
|
out_cache_ptr = out_cache_loc + out_offset
|
||||||
|
|
||||||
|
load_offset = tl.arange(0, BLOCK_SIZE) + kv_start
|
||||||
|
save_offset = tl.arange(0, BLOCK_SIZE)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
||||||
|
for _ in range(num_loop):
|
||||||
|
mask = load_offset < kv_end
|
||||||
|
data = tl.load(token_pool + load_offset, mask=mask)
|
||||||
|
tl.store(out_cache_ptr + save_offset, data, mask=mask)
|
||||||
|
load_offset += BLOCK_SIZE
|
||||||
|
save_offset += BLOCK_SIZE
|
||||||
|
|
||||||
|
|
||||||
|
def assign_extend_cache_locs_func(
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
start_offset: torch.Tensor,
|
||||||
|
end_offset: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
draft_token_num: int,
|
||||||
|
device,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if _is_cuda or _is_hip or _is_musa:
|
||||||
|
out_cache_loc = torch.empty(
|
||||||
|
(batch_size * draft_token_num,),
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
assign_extend_cache_locs[(batch_size,)](
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
req_to_token.shape[1],
|
||||||
|
next_power_of_2(batch_size),
|
||||||
|
)
|
||||||
|
|
||||||
|
return out_cache_loc
|
||||||
|
|
||||||
|
elif _is_npu:
|
||||||
|
out_cache_loc = torch.empty(
|
||||||
|
(batch_size * draft_token_num,),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
torch.ops.npu.cache_loc_update(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
)
|
||||||
|
|
||||||
|
return out_cache_loc
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def fill_bonus_tokens(
|
||||||
|
accept_tokens,
|
||||||
|
accept_lens,
|
||||||
|
bonus_tokens_ptr,
|
||||||
|
num_draft_tokens: tl.constexpr,
|
||||||
|
):
|
||||||
|
# NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel
|
||||||
|
# because this kernel reads accept_lens
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
# `accept_lens` includes the bonus token; the last accepted slot is at -1.
|
||||||
|
accept_len = tl.load(accept_lens + pid)
|
||||||
|
|
||||||
|
bonus_token_idx = num_draft_tokens * pid + accept_len - 1
|
||||||
|
bonus_token = tl.load(accept_tokens + bonus_token_idx)
|
||||||
|
tl.store(bonus_tokens_ptr + pid, bonus_token)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def fill_accepted_out_cache_loc(
|
||||||
|
accept_index,
|
||||||
|
out_cache_loc,
|
||||||
|
accepted_out_cache_loc,
|
||||||
|
size_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
offset = tl.arange(0, size_upper)
|
||||||
|
|
||||||
|
masks = (tl.load(accept_index + offset, offset < pid, other=-1) != -1).to(tl.int64)
|
||||||
|
dst = tl.sum(masks)
|
||||||
|
src = tl.load(accept_index + pid)
|
||||||
|
if src > -1:
|
||||||
|
value = tl.load(out_cache_loc + src)
|
||||||
|
tl.store(accepted_out_cache_loc + dst, value)
|
||||||
@@ -0,0 +1,350 @@
|
|||||||
|
# Copyright 2023-2024 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def rotate_input_ids_kernel(
|
||||||
|
input_ids_ptr,
|
||||||
|
extend_start_loc_ptr,
|
||||||
|
extend_seq_lens_ptr,
|
||||||
|
topk_index_ptr,
|
||||||
|
select_index_ptr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
start_loc = tl.load(extend_start_loc_ptr + pid)
|
||||||
|
seq_len = tl.load(extend_seq_lens_ptr + pid)
|
||||||
|
new_token = tl.load(topk_index_ptr + pid)
|
||||||
|
|
||||||
|
num_elements_to_shift = seq_len - 1
|
||||||
|
|
||||||
|
for off in range(0, num_elements_to_shift, BLOCK_SIZE):
|
||||||
|
offsets = off + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offsets < num_elements_to_shift
|
||||||
|
|
||||||
|
read_ptr = input_ids_ptr + start_loc + offsets + 1
|
||||||
|
val = tl.load(read_ptr, mask=mask)
|
||||||
|
tl.debug_barrier()
|
||||||
|
|
||||||
|
write_ptr = input_ids_ptr + start_loc + offsets
|
||||||
|
tl.store(write_ptr, val, mask=mask)
|
||||||
|
tl.debug_barrier()
|
||||||
|
|
||||||
|
if seq_len > 0:
|
||||||
|
if select_index_ptr is not None:
|
||||||
|
last_pos_ptr = input_ids_ptr + tl.load(select_index_ptr + pid)
|
||||||
|
else:
|
||||||
|
last_pos_ptr = input_ids_ptr + start_loc + seq_len - 1
|
||||||
|
tl.store(last_pos_ptr, new_token)
|
||||||
|
|
||||||
|
|
||||||
|
def rotate_input_ids_triton(
|
||||||
|
input_ids, extend_start_loc, extend_seq_lens, topk_index, select_index=None
|
||||||
|
):
|
||||||
|
batch_size = extend_seq_lens.shape[0]
|
||||||
|
BLOCK_SIZE = 4096 if select_index is not None else 8
|
||||||
|
grid = (batch_size,)
|
||||||
|
|
||||||
|
rotate_input_ids_kernel[grid](
|
||||||
|
input_ids,
|
||||||
|
extend_start_loc,
|
||||||
|
extend_seq_lens,
|
||||||
|
topk_index,
|
||||||
|
select_index,
|
||||||
|
BLOCK_SIZE=BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
return input_ids
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_new_state_kernel(
|
||||||
|
# Source pointers
|
||||||
|
old_input_ids_ptr,
|
||||||
|
old_positions_ptr,
|
||||||
|
old_hidden_states_ptr,
|
||||||
|
old_out_cache_loc_ptr,
|
||||||
|
old_extend_seq_lens_ptr,
|
||||||
|
old_extend_start_loc_ptr,
|
||||||
|
# Destination pointers
|
||||||
|
input_ids_ptr,
|
||||||
|
positions_ptr,
|
||||||
|
hidden_states_ptr,
|
||||||
|
out_cache_loc_ptr,
|
||||||
|
extend_seq_lens_ptr,
|
||||||
|
extend_start_loc_ptr,
|
||||||
|
# Auxiliary data pointers
|
||||||
|
next_token_ids_ptr,
|
||||||
|
seq_lens_ptr,
|
||||||
|
padding_lens_ptr,
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
req_to_token_ptr,
|
||||||
|
req_to_hidden_states_pool_ptr,
|
||||||
|
# Scalars and Strides
|
||||||
|
step,
|
||||||
|
stride_hidden_seq,
|
||||||
|
stride_hidden_dim, # hidden_states strides
|
||||||
|
stride_pool_req,
|
||||||
|
stride_pool_step,
|
||||||
|
stride_pool_dim, # pool strides
|
||||||
|
stride_req_token_0,
|
||||||
|
stride_req_token_1, # req_to_token strides
|
||||||
|
# Meta-parameters
|
||||||
|
HIDDEN_DIM: tl.constexpr,
|
||||||
|
BLOCK_SEQ: tl.constexpr,
|
||||||
|
BLOCK_HID: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
seq_len: tl.tensor = tl.load(seq_lens_ptr + pid)
|
||||||
|
old_extend_len = tl.load(old_extend_seq_lens_ptr + pid)
|
||||||
|
old_start = tl.load(old_extend_start_loc_ptr + pid)
|
||||||
|
new_extend_len = old_extend_len + 1
|
||||||
|
new_start = old_start + pid
|
||||||
|
|
||||||
|
tl.store(extend_seq_lens_ptr + pid, new_extend_len)
|
||||||
|
tl.store(extend_start_loc_ptr + pid, new_start)
|
||||||
|
|
||||||
|
offs_seq = tl.arange(0, BLOCK_SEQ)
|
||||||
|
mask_seq = offs_seq < old_extend_len
|
||||||
|
|
||||||
|
old_ids = tl.load(old_input_ids_ptr + old_start + offs_seq, mask=mask_seq)
|
||||||
|
tl.store(input_ids_ptr + new_start + offs_seq, old_ids, mask=mask_seq)
|
||||||
|
padding_len = tl.load(padding_lens_ptr + pid)
|
||||||
|
tl.store(
|
||||||
|
input_ids_ptr + new_start + old_extend_len - padding_len,
|
||||||
|
tl.load(next_token_ids_ptr + pid),
|
||||||
|
)
|
||||||
|
|
||||||
|
old_pos = tl.load(old_positions_ptr + old_start + offs_seq, mask=mask_seq)
|
||||||
|
tl.store(positions_ptr + new_start + 1 + offs_seq, old_pos, mask=mask_seq)
|
||||||
|
tl.store(
|
||||||
|
positions_ptr + new_start, max(tl.load(old_positions_ptr + old_start) - 1, 0)
|
||||||
|
)
|
||||||
|
|
||||||
|
old_cache = tl.load(old_out_cache_loc_ptr + old_start + offs_seq, mask=mask_seq)
|
||||||
|
tl.store(out_cache_loc_ptr + new_start + 1 + offs_seq, old_cache, mask=mask_seq)
|
||||||
|
|
||||||
|
req_idx = tl.load(req_pool_indices_ptr + pid)
|
||||||
|
token_idx_col = seq_len - old_extend_len - 1
|
||||||
|
if token_idx_col >= 0:
|
||||||
|
req_token_ptr_loc = (
|
||||||
|
req_to_token_ptr
|
||||||
|
+ (req_idx * stride_req_token_0)
|
||||||
|
+ (token_idx_col * stride_req_token_1)
|
||||||
|
)
|
||||||
|
last_cache_loc = tl.load(req_token_ptr_loc)
|
||||||
|
tl.store(out_cache_loc_ptr + new_start, last_cache_loc)
|
||||||
|
|
||||||
|
pool_vec_offset_base = ((req_idx + 1) * stride_pool_req) + (
|
||||||
|
-(step + 1) * stride_pool_step
|
||||||
|
)
|
||||||
|
|
||||||
|
for off_h in range(0, HIDDEN_DIM, BLOCK_HID):
|
||||||
|
offs_h = off_h + tl.arange(0, BLOCK_HID)
|
||||||
|
mask_h = offs_h < HIDDEN_DIM
|
||||||
|
|
||||||
|
for i in range(BLOCK_SEQ):
|
||||||
|
if i < old_extend_len:
|
||||||
|
old_h_ptr = (
|
||||||
|
old_hidden_states_ptr
|
||||||
|
+ (old_start + i) * stride_hidden_seq
|
||||||
|
+ (offs_h * stride_hidden_dim)
|
||||||
|
)
|
||||||
|
new_h_ptr = (
|
||||||
|
hidden_states_ptr
|
||||||
|
+ (new_start + 1 + i) * stride_hidden_seq
|
||||||
|
+ (offs_h * stride_hidden_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
chunk_old = tl.load(old_h_ptr, mask=mask_h)
|
||||||
|
tl.store(new_h_ptr, chunk_old, mask=mask_h)
|
||||||
|
|
||||||
|
pool_ptrs = (
|
||||||
|
req_to_hidden_states_pool_ptr
|
||||||
|
+ pool_vec_offset_base
|
||||||
|
+ (offs_h * stride_pool_dim)
|
||||||
|
)
|
||||||
|
pool_val = tl.load(pool_ptrs, mask=mask_h)
|
||||||
|
|
||||||
|
new_h_start_ptrs = (
|
||||||
|
hidden_states_ptr
|
||||||
|
+ (new_start * stride_hidden_seq)
|
||||||
|
+ (offs_h * stride_hidden_dim)
|
||||||
|
)
|
||||||
|
tl.store(new_h_start_ptrs, pool_val, mask=mask_h)
|
||||||
|
|
||||||
|
|
||||||
|
def assign_new_state_triton(
|
||||||
|
next_token_ids: torch.Tensor,
|
||||||
|
old_input_ids: torch.Tensor,
|
||||||
|
old_positions: torch.Tensor,
|
||||||
|
old_hidden_states: torch.Tensor,
|
||||||
|
old_out_cache_loc: torch.Tensor,
|
||||||
|
old_extend_seq_lens: torch.Tensor,
|
||||||
|
old_extend_start_loc: torch.Tensor,
|
||||||
|
input_ids: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
out_cache_loc: torch.Tensor,
|
||||||
|
extend_seq_lens: torch.Tensor,
|
||||||
|
extend_start_loc: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
padding_lens: torch.Tensor,
|
||||||
|
num_seqs: int,
|
||||||
|
step: int,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_to_hidden_states_pool: torch.Tensor,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Wrapper function to calculate offsets and launch the Triton kernel.
|
||||||
|
"""
|
||||||
|
hidden_dim = hidden_states.shape[1]
|
||||||
|
|
||||||
|
BLOCK_SEQ = 8
|
||||||
|
BLOCK_HID = 64
|
||||||
|
|
||||||
|
grid = (num_seqs,)
|
||||||
|
|
||||||
|
assign_new_state_kernel[grid](
|
||||||
|
# Pointers
|
||||||
|
old_input_ids,
|
||||||
|
old_positions,
|
||||||
|
old_hidden_states,
|
||||||
|
old_out_cache_loc,
|
||||||
|
old_extend_seq_lens,
|
||||||
|
old_extend_start_loc,
|
||||||
|
input_ids,
|
||||||
|
positions,
|
||||||
|
hidden_states,
|
||||||
|
out_cache_loc,
|
||||||
|
extend_seq_lens,
|
||||||
|
extend_start_loc,
|
||||||
|
next_token_ids,
|
||||||
|
seq_lens,
|
||||||
|
padding_lens,
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
req_to_hidden_states_pool,
|
||||||
|
# Constants/Strides
|
||||||
|
step,
|
||||||
|
old_hidden_states.stride(0),
|
||||||
|
old_hidden_states.stride(1),
|
||||||
|
req_to_hidden_states_pool.stride(0),
|
||||||
|
req_to_hidden_states_pool.stride(1),
|
||||||
|
req_to_hidden_states_pool.stride(2),
|
||||||
|
req_to_token.stride(0),
|
||||||
|
req_to_token.stride(1),
|
||||||
|
# Meta
|
||||||
|
HIDDEN_DIM=hidden_dim,
|
||||||
|
BLOCK_SEQ=BLOCK_SEQ,
|
||||||
|
BLOCK_HID=BLOCK_HID,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_hidden_states_pool_kernel(
|
||||||
|
hidden_states_ptr,
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
req_to_hidden_states_pool_ptr,
|
||||||
|
extend_seq_lens_ptr,
|
||||||
|
extend_start_loc_ptr,
|
||||||
|
stride_hidden_seq,
|
||||||
|
stride_hidden_dim,
|
||||||
|
stride_pool_req,
|
||||||
|
stride_pool_step,
|
||||||
|
stride_pool_dim,
|
||||||
|
HIDDEN_DIM: tl.constexpr,
|
||||||
|
pool_size: tl.constexpr,
|
||||||
|
BLOCK_HID: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
extend_len = tl.load(extend_seq_lens_ptr + pid)
|
||||||
|
start_loc = tl.load(extend_start_loc_ptr + pid)
|
||||||
|
end_loc = start_loc + extend_len
|
||||||
|
|
||||||
|
req_idx = tl.load(req_pool_indices_ptr + pid)
|
||||||
|
pool_vec_offset_base = req_idx * stride_pool_req
|
||||||
|
|
||||||
|
for i in range(pool_size):
|
||||||
|
for off_h in range(0, HIDDEN_DIM, BLOCK_HID):
|
||||||
|
offs_h = off_h + tl.arange(0, BLOCK_HID)
|
||||||
|
mask_h = offs_h < HIDDEN_DIM
|
||||||
|
|
||||||
|
hid_ptr = (
|
||||||
|
hidden_states_ptr
|
||||||
|
+ (end_loc - pool_size + i) * stride_hidden_seq
|
||||||
|
+ offs_h * stride_hidden_dim
|
||||||
|
)
|
||||||
|
hid_val = tl.load(hid_ptr, mask=mask_h)
|
||||||
|
|
||||||
|
pool_ptr = (
|
||||||
|
req_to_hidden_states_pool_ptr
|
||||||
|
+ pool_vec_offset_base
|
||||||
|
+ i * stride_pool_step
|
||||||
|
+ offs_h * stride_pool_dim
|
||||||
|
)
|
||||||
|
tl.store(pool_ptr, hid_val, mask=mask_h)
|
||||||
|
|
||||||
|
|
||||||
|
def assign_hidden_states_pool_triton(
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_hidden_states_pool: torch.Tensor,
|
||||||
|
pool_size: int,
|
||||||
|
num_seqs: int,
|
||||||
|
extend_seq_lens: torch.Tensor,
|
||||||
|
extend_start_loc: torch.Tensor,
|
||||||
|
):
|
||||||
|
grid = (num_seqs,)
|
||||||
|
assign_hidden_states_pool_kernel[grid](
|
||||||
|
hidden_states,
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_hidden_states_pool,
|
||||||
|
extend_seq_lens,
|
||||||
|
extend_start_loc,
|
||||||
|
hidden_states.stride(0),
|
||||||
|
hidden_states.stride(1),
|
||||||
|
req_to_hidden_states_pool.stride(0),
|
||||||
|
req_to_hidden_states_pool.stride(1),
|
||||||
|
req_to_hidden_states_pool.stride(2),
|
||||||
|
HIDDEN_DIM=hidden_states.shape[1],
|
||||||
|
pool_size=pool_size,
|
||||||
|
BLOCK_HID=64,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def assign_hidden_states_pool_torch(
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_hidden_states_pool: torch.Tensor,
|
||||||
|
pool_size: int,
|
||||||
|
num_seqs: int,
|
||||||
|
extend_seq_lens: torch.Tensor,
|
||||||
|
extend_start_loc: torch.Tensor,
|
||||||
|
):
|
||||||
|
for req in range(num_seqs):
|
||||||
|
pool_idx = req_pool_indices[req]
|
||||||
|
extend_len = extend_seq_lens[req]
|
||||||
|
start_loc = extend_start_loc[req]
|
||||||
|
end_loc = start_loc + extend_len
|
||||||
|
req_to_hidden_states_pool[pool_idx, :pool_size, :].copy_(
|
||||||
|
hidden_states[end_loc - pool_size : end_loc, :]
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user