Optimize ngram decode token table update (#24756)
Co-authored-by: Codex <codex@example.com> Co-authored-by: BBuf <xiaoyu.zhang@radixark.net>
This commit is contained in:
co-authored by
Codex
BBuf
parent
9da88e32e0
commit
e513c13e2e
@@ -0,0 +1,76 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.testing
|
||||||
|
|
||||||
|
from sglang.jit_kernel.benchmark.utils import (
|
||||||
|
DEFAULT_DEVICE,
|
||||||
|
get_benchmark_range,
|
||||||
|
run_benchmark_no_cudagraph,
|
||||||
|
)
|
||||||
|
from sglang.jit_kernel.ngram_embedding import (
|
||||||
|
update_token_table,
|
||||||
|
update_token_table_decode,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=15, suite="base-b-kernel-benchmark-1-gpu-large")
|
||||||
|
|
||||||
|
MAX_CONTEXT_LEN = 4096
|
||||||
|
BATCH_SIZE_LIST = get_benchmark_range(
|
||||||
|
full_range=[1, 2, 8, 32, 128, 512, 1024, 2048, 4096],
|
||||||
|
ci_range=[32, 1024],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.testing.perf_report(
|
||||||
|
triton.testing.Benchmark(
|
||||||
|
x_names=["batch_size"],
|
||||||
|
x_vals=BATCH_SIZE_LIST,
|
||||||
|
line_arg="provider",
|
||||||
|
line_vals=["general", "decode"],
|
||||||
|
line_names=["general update_token_table", "decode fast path"],
|
||||||
|
styles=[("blue", "-"), ("orange", "-")],
|
||||||
|
ylabel="us",
|
||||||
|
plot_name="ngram-update-token-table",
|
||||||
|
args={},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
def benchmark(batch_size: int, provider: str):
|
||||||
|
max_running_reqs = batch_size + 8
|
||||||
|
tokens = torch.arange(batch_size, dtype=torch.int32, device=DEFAULT_DEVICE)
|
||||||
|
token_table = torch.empty(
|
||||||
|
(max_running_reqs, MAX_CONTEXT_LEN), dtype=torch.int32, device=DEFAULT_DEVICE
|
||||||
|
)
|
||||||
|
row_indices = torch.arange(batch_size, dtype=torch.int64, device=DEFAULT_DEVICE)
|
||||||
|
column_starts = torch.randint(
|
||||||
|
0, MAX_CONTEXT_LEN, (batch_size,), dtype=torch.int32, device=DEFAULT_DEVICE
|
||||||
|
)
|
||||||
|
req_lens = torch.ones(batch_size, dtype=torch.int32, device=DEFAULT_DEVICE)
|
||||||
|
|
||||||
|
if provider == "general":
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
update_token_table(
|
||||||
|
tokens,
|
||||||
|
token_table,
|
||||||
|
row_indices,
|
||||||
|
column_starts,
|
||||||
|
req_lens,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
update_token_table_decode(
|
||||||
|
tokens,
|
||||||
|
token_table,
|
||||||
|
row_indices,
|
||||||
|
column_starts,
|
||||||
|
)
|
||||||
|
|
||||||
|
return run_benchmark_no_cudagraph(fn)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
benchmark.run(print_data=True)
|
||||||
@@ -15,21 +15,22 @@ namespace device::ngram_embedding {
|
|||||||
|
|
||||||
constexpr int kDecodeBlockSize = 256;
|
constexpr int kDecodeBlockSize = 256;
|
||||||
constexpr int kMaxComputeNGramIdsDecodeBlocks = 65535;
|
constexpr int kMaxComputeNGramIdsDecodeBlocks = 65535;
|
||||||
|
constexpr int kMaxUpdateTokenTableDecodeBlocks = 1024;
|
||||||
|
|
||||||
__global__ void ComputeNGramIdsKernel(
|
__global__ void ComputeNGramIdsKernel(
|
||||||
int batch_size,
|
int batch_size,
|
||||||
int ne_n,
|
int ne_n,
|
||||||
int ne_k,
|
int ne_k,
|
||||||
int* ne_weights, // [ne_n-1,ne_k,ne_n]
|
int* ne_weights, // [ne_n-1,ne_k,ne_n]
|
||||||
int* ne_mods, // [ne_n-1,ne_k]
|
int* ne_mods, // [ne_n-1,ne_k]
|
||||||
int* exclusive_ne_embeder_size_sums, // [(ne_n-1)*ne_k]
|
int* exclusive_ne_embeder_size_sums, // [(ne_n-1)*ne_k]
|
||||||
int* tokens, // [token_num]
|
int* tokens, // [token_num]
|
||||||
int* exclusive_req_len_sums, // [batch_size+1]
|
int* exclusive_req_len_sums, // [batch_size+1]
|
||||||
int* ne_token_table, // [max_running_reqs, max_context_len]
|
int* ne_token_table, // [max_running_reqs, max_context_len]
|
||||||
int max_context_len, // max_context_len
|
int max_context_len, // max_context_len
|
||||||
long* row_indices, // [batch_size]
|
const int64_t* __restrict__ row_indices, // [batch_size]
|
||||||
int* column_starts, // [batch_size]
|
int* column_starts, // [batch_size]
|
||||||
int* n_gram_ids // [ne_n-1,ne_k,token_num]
|
int* n_gram_ids // [ne_n-1,ne_k,token_num]
|
||||||
) {
|
) {
|
||||||
// Determine which n, k, and request this block handles.
|
// Determine which n, k, and request this block handles.
|
||||||
/**
|
/**
|
||||||
@@ -62,11 +63,11 @@ __global__ void ComputeNGramIdsKernel(
|
|||||||
for (int i = exclusive_req_len_sums[req_id] + threadIdx.x; i < exclusive_req_len_sums[req_id + 1]; i += blockDim.x) {
|
for (int i = exclusive_req_len_sums[req_id] + threadIdx.x; i < exclusive_req_len_sums[req_id + 1]; i += blockDim.x) {
|
||||||
uint64_t n_gram_id = 0;
|
uint64_t n_gram_id = 0;
|
||||||
// Token offset within the current request
|
// Token offset within the current request
|
||||||
int current_token_offset = i - exclusive_req_len_sums[req_id];
|
const int64_t current_token_offset = i - exclusive_req_len_sums[req_id];
|
||||||
// Start index of this request in the token table; tokens before this belong to other requests
|
// Start index of this request in the token table; tokens before this belong to other requests
|
||||||
int req_token_table_index = row_indices[req_id] * max_context_len;
|
const int64_t req_token_table_index = row_indices[req_id] * static_cast<int64_t>(max_context_len);
|
||||||
// Position of the current token in the token table
|
// Position of the current token in the token table
|
||||||
int current_token_table_index = req_token_table_index + column_starts[req_id] + current_token_offset;
|
const int64_t current_token_table_index = req_token_table_index + column_starts[req_id] + current_token_offset;
|
||||||
for (int j = 0; j < n + 2; j++) {
|
for (int j = 0; j < n + 2; j++) {
|
||||||
if (current_token_table_index - j < req_token_table_index) {
|
if (current_token_table_index - j < req_token_table_index) {
|
||||||
// Out of this request's range, stop computing n_gram_id
|
// Out of this request's range, stop computing n_gram_id
|
||||||
@@ -134,14 +135,14 @@ __global__ void ComputeNGramIdsDecodeKernel(
|
|||||||
|
|
||||||
__global__ void UpdateTokenTableKernel(
|
__global__ void UpdateTokenTableKernel(
|
||||||
int batch_size,
|
int batch_size,
|
||||||
int* tokens, // [token_num]
|
int* tokens, // [token_num]
|
||||||
int* ne_token_table, // [max_running_reqs, max_context_len]
|
int* ne_token_table, // [max_running_reqs, max_context_len]
|
||||||
int max_context_len, // max_context_len
|
int max_context_len, // max_context_len
|
||||||
long* row_indices, // [batch_size]
|
const int64_t* __restrict__ row_indices, // [batch_size]
|
||||||
int* column_starts, // [batch_size]
|
int* column_starts, // [batch_size]
|
||||||
int* req_lens, // [batch_size]
|
int* req_lens, // [batch_size]
|
||||||
int ignore_token_num, // number of tokens to ignore
|
int ignore_token_num, // number of tokens to ignore
|
||||||
int* ignore_tokens // [ignore_token_num]
|
int* ignore_tokens // [ignore_token_num]
|
||||||
) {
|
) {
|
||||||
// Each block processes one request.
|
// Each block processes one request.
|
||||||
const int req_id = blockIdx.x % batch_size;
|
const int req_id = blockIdx.x % batch_size;
|
||||||
@@ -154,11 +155,11 @@ __global__ void UpdateTokenTableKernel(
|
|||||||
// stride loop
|
// stride loop
|
||||||
for (int i = start + threadIdx.x; i < end; i += blockDim.x) {
|
for (int i = start + threadIdx.x; i < end; i += blockDim.x) {
|
||||||
// Token offset within the current request
|
// Token offset within the current request
|
||||||
int current_token_offset = i - start;
|
const int64_t current_token_offset = i - start;
|
||||||
// Start index of this request in the token table
|
// Start index of this request in the token table
|
||||||
int req_token_table_index = row_indices[req_id] * max_context_len;
|
const int64_t req_token_table_index = row_indices[req_id] * static_cast<int64_t>(max_context_len);
|
||||||
// Position of the current token in the token table
|
// Position of the current token in the token table
|
||||||
int current_token_table_index = req_token_table_index + column_starts[req_id] + current_token_offset;
|
const int64_t current_token_table_index = req_token_table_index + column_starts[req_id] + current_token_offset;
|
||||||
ne_token_table[current_token_table_index] = tokens[i];
|
ne_token_table[current_token_table_index] = tokens[i];
|
||||||
for (int j = 0; j < ignore_token_num; j++) {
|
for (int j = 0; j < ignore_token_num; j++) {
|
||||||
if (ignore_tokens[j] == tokens[i]) {
|
if (ignore_tokens[j] == tokens[i]) {
|
||||||
@@ -169,6 +170,21 @@ __global__ void UpdateTokenTableKernel(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
__global__ void UpdateTokenTableDecodeKernel(
|
||||||
|
int batch_size,
|
||||||
|
const int* __restrict__ tokens, // [batch_size]
|
||||||
|
int* __restrict__ ne_token_table, // [max_running_reqs, max_context_len]
|
||||||
|
int max_context_len, // max_context_len
|
||||||
|
const int64_t* __restrict__ row_indices, // [batch_size]
|
||||||
|
const int* __restrict__ column_starts // [batch_size]
|
||||||
|
) {
|
||||||
|
for (int req_id = blockIdx.x * blockDim.x + threadIdx.x; req_id < batch_size; req_id += blockDim.x * gridDim.x) {
|
||||||
|
const int64_t token_table_offset =
|
||||||
|
row_indices[req_id] * static_cast<int64_t>(max_context_len) + column_starts[req_id];
|
||||||
|
ne_token_table[token_table_offset] = tokens[req_id];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
} // namespace device::ngram_embedding
|
} // namespace device::ngram_embedding
|
||||||
|
|
||||||
namespace {
|
namespace {
|
||||||
@@ -256,7 +272,7 @@ struct NgramEmbeddingKernel {
|
|||||||
static_cast<int*>(exclusive_req_len_sums.data_ptr()),
|
static_cast<int*>(exclusive_req_len_sums.data_ptr()),
|
||||||
static_cast<int*>(ne_token_table.data_ptr()),
|
static_cast<int*>(ne_token_table.data_ptr()),
|
||||||
max_context_len,
|
max_context_len,
|
||||||
static_cast<long*>(row_indices.data_ptr()),
|
static_cast<const int64_t*>(row_indices.data_ptr()),
|
||||||
static_cast<int*>(column_starts.data_ptr()),
|
static_cast<int*>(column_starts.data_ptr()),
|
||||||
static_cast<int*>(n_gram_ids.data_ptr()));
|
static_cast<int*>(n_gram_ids.data_ptr()));
|
||||||
}
|
}
|
||||||
@@ -412,12 +428,64 @@ struct NgramEmbeddingKernel {
|
|||||||
static_cast<int*>(tokens.data_ptr()),
|
static_cast<int*>(tokens.data_ptr()),
|
||||||
static_cast<int*>(ne_token_table.data_ptr()),
|
static_cast<int*>(ne_token_table.data_ptr()),
|
||||||
max_context_len,
|
max_context_len,
|
||||||
static_cast<long*>(row_indices.data_ptr()),
|
static_cast<const int64_t*>(row_indices.data_ptr()),
|
||||||
static_cast<int*>(column_starts.data_ptr()),
|
static_cast<int*>(column_starts.data_ptr()),
|
||||||
static_cast<int*>(req_lens.data_ptr()),
|
static_cast<int*>(req_lens.data_ptr()),
|
||||||
ignore_token_num,
|
ignore_token_num,
|
||||||
ignore_tokens_typed_ptr);
|
ignore_tokens_typed_ptr);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
static void update_token_table_decode(
|
||||||
|
const tvm::ffi::TensorView tokens,
|
||||||
|
const tvm::ffi::TensorView ne_token_table,
|
||||||
|
const tvm::ffi::TensorView row_indices,
|
||||||
|
const tvm::ffi::TensorView column_starts) {
|
||||||
|
using namespace host;
|
||||||
|
|
||||||
|
auto batch_size = SymbolicSize{"batch_size"};
|
||||||
|
auto device_ = SymbolicDevice{};
|
||||||
|
|
||||||
|
TensorMatcher({batch_size}) // [batch_size]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>(device_)
|
||||||
|
.verify(tokens);
|
||||||
|
|
||||||
|
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(ne_token_table);
|
||||||
|
|
||||||
|
TensorMatcher({batch_size}) // [batch_size]
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(row_indices);
|
||||||
|
|
||||||
|
TensorMatcher({batch_size}) // [batch_size]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(column_starts);
|
||||||
|
|
||||||
|
const int bs = static_cast<int>(batch_size.unwrap());
|
||||||
|
if (bs <= 0) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const int max_context_len = static_cast<int>(ne_token_table.size(1));
|
||||||
|
const auto stream = LaunchKernel::resolve_device(device_.unwrap());
|
||||||
|
|
||||||
|
constexpr int kBlockSize = device::ngram_embedding::kDecodeBlockSize;
|
||||||
|
const int grid_size = std::min(
|
||||||
|
device::ngram_embedding::kMaxUpdateTokenTableDecodeBlocks, static_cast<int>(host::div_ceil(bs, kBlockSize)));
|
||||||
|
|
||||||
|
LaunchKernel(grid_size, kBlockSize, stream)(
|
||||||
|
device::ngram_embedding::UpdateTokenTableDecodeKernel,
|
||||||
|
bs,
|
||||||
|
static_cast<const int*>(tokens.data_ptr()),
|
||||||
|
static_cast<int*>(ne_token_table.data_ptr()),
|
||||||
|
max_context_len,
|
||||||
|
static_cast<const int64_t*>(row_indices.data_ptr()),
|
||||||
|
static_cast<const int*>(column_starts.data_ptr()));
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
} // namespace
|
} // namespace
|
||||||
|
|||||||
@@ -22,6 +22,10 @@ def _jit_ngram_embedding_module() -> Module:
|
|||||||
"&NgramEmbeddingKernel::compute_n_gram_ids_decode",
|
"&NgramEmbeddingKernel::compute_n_gram_ids_decode",
|
||||||
),
|
),
|
||||||
("update_token_table", "&NgramEmbeddingKernel::update_token_table"),
|
("update_token_table", "&NgramEmbeddingKernel::update_token_table"),
|
||||||
|
(
|
||||||
|
"update_token_table_decode",
|
||||||
|
"&NgramEmbeddingKernel::update_token_table_decode",
|
||||||
|
),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -133,3 +137,24 @@ def update_token_table(
|
|||||||
req_lens,
|
req_lens,
|
||||||
ignore_tokens,
|
ignore_tokens,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@debug_kernel_api
|
||||||
|
def update_token_table_decode(
|
||||||
|
tokens: torch.Tensor,
|
||||||
|
ne_token_table: torch.Tensor,
|
||||||
|
row_indices: torch.Tensor,
|
||||||
|
column_starts: torch.Tensor,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Update one decoded token per request in the ngram embedding token table.
|
||||||
|
|
||||||
|
This is the decode-only fast path for req_lens == 1 and no ignored tokens.
|
||||||
|
"""
|
||||||
|
module = _jit_ngram_embedding_module()
|
||||||
|
module.update_token_table_decode(
|
||||||
|
tokens,
|
||||||
|
ne_token_table,
|
||||||
|
row_indices,
|
||||||
|
column_starts,
|
||||||
|
)
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import torch
|
|||||||
from sglang.jit_kernel.ngram_embedding import (
|
from sglang.jit_kernel.ngram_embedding import (
|
||||||
compute_n_gram_ids,
|
compute_n_gram_ids,
|
||||||
compute_n_gram_ids_decode,
|
compute_n_gram_ids_decode,
|
||||||
|
update_token_table,
|
||||||
|
update_token_table_decode,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
@@ -93,5 +95,41 @@ def test_compute_n_gram_ids_decode_matches_general(batch_size: int) -> None:
|
|||||||
torch.testing.assert_close(n_gram_ids_decode, n_gram_ids_general, atol=0, rtol=0)
|
torch.testing.assert_close(n_gram_ids_decode, n_gram_ids_general, atol=0, rtol=0)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("batch_size", [1, 2, 17, 128, 1024])
|
||||||
|
def test_update_token_table_decode_matches_general(batch_size: int) -> None:
|
||||||
|
max_context_len = 4096
|
||||||
|
max_running_reqs = batch_size + 8
|
||||||
|
tokens = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 100
|
||||||
|
row_indices = torch.randperm(max_running_reqs, device="cuda")[:batch_size].to(
|
||||||
|
torch.int64
|
||||||
|
)
|
||||||
|
column_starts = torch.randint(
|
||||||
|
0, max_context_len, (batch_size,), dtype=torch.int32, device="cuda"
|
||||||
|
)
|
||||||
|
req_lens = torch.ones(batch_size, dtype=torch.int32, device="cuda")
|
||||||
|
|
||||||
|
token_table_general = torch.full(
|
||||||
|
(max_running_reqs, max_context_len), -1, dtype=torch.int32, device="cuda"
|
||||||
|
)
|
||||||
|
token_table_decode = token_table_general.clone()
|
||||||
|
|
||||||
|
update_token_table(
|
||||||
|
tokens=tokens,
|
||||||
|
ne_token_table=token_table_general,
|
||||||
|
row_indices=row_indices,
|
||||||
|
column_starts=column_starts,
|
||||||
|
req_lens=req_lens,
|
||||||
|
ignore_tokens=None,
|
||||||
|
)
|
||||||
|
update_token_table_decode(
|
||||||
|
tokens=tokens,
|
||||||
|
ne_token_table=token_table_decode,
|
||||||
|
row_indices=row_indices,
|
||||||
|
column_starts=column_starts,
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.testing.assert_close(token_table_decode, token_table_general, atol=0, rtol=0)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ import torch
|
|||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from sglang.jit_kernel.ngram_embedding import update_token_table
|
from sglang.jit_kernel.ngram_embedding import update_token_table_decode
|
||||||
from sglang.srt.configs import (
|
from sglang.srt.configs import (
|
||||||
BailingHybridConfig,
|
BailingHybridConfig,
|
||||||
FalconH1Config,
|
FalconH1Config,
|
||||||
@@ -2840,13 +2840,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
forward_batch.seq_lens
|
forward_batch.seq_lens
|
||||||
)
|
)
|
||||||
ngram_embedding_info.out_req_lens[: forward_batch.batch_size] = 1
|
ngram_embedding_info.out_req_lens[: forward_batch.batch_size] = 1
|
||||||
update_token_table(
|
update_token_table_decode(
|
||||||
ne_token_table=ngram_embedding_info.token_table,
|
ne_token_table=ngram_embedding_info.token_table,
|
||||||
tokens=next_token_ids.to(torch.int32),
|
tokens=next_token_ids.to(torch.int32),
|
||||||
row_indices=forward_batch.req_pool_indices,
|
row_indices=forward_batch.req_pool_indices,
|
||||||
column_starts=ngram_embedding_info.out_column_starts,
|
column_starts=ngram_embedding_info.out_column_starts,
|
||||||
req_lens=torch.ones_like(ngram_embedding_info.out_column_starts),
|
|
||||||
ignore_tokens=None,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_device_graphs(self):
|
def init_device_graphs(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user