diff --git a/python/sglang/jit_kernel/benchmark/bench_ngram_update_token_table.py b/python/sglang/jit_kernel/benchmark/bench_ngram_update_token_table.py new file mode 100644 index 000000000..118459df3 --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/bench_ngram_update_token_table.py @@ -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) diff --git a/python/sglang/jit_kernel/csrc/ngram_embedding.cuh b/python/sglang/jit_kernel/csrc/ngram_embedding.cuh index 0683e0154..4eac7a946 100644 --- a/python/sglang/jit_kernel/csrc/ngram_embedding.cuh +++ b/python/sglang/jit_kernel/csrc/ngram_embedding.cuh @@ -15,21 +15,22 @@ namespace device::ngram_embedding { constexpr int kDecodeBlockSize = 256; constexpr int kMaxComputeNGramIdsDecodeBlocks = 65535; +constexpr int kMaxUpdateTokenTableDecodeBlocks = 1024; __global__ void ComputeNGramIdsKernel( int batch_size, int ne_n, int ne_k, - int* ne_weights, // [ne_n-1,ne_k,ne_n] - int* ne_mods, // [ne_n-1,ne_k] - int* exclusive_ne_embeder_size_sums, // [(ne_n-1)*ne_k] - int* tokens, // [token_num] - int* exclusive_req_len_sums, // [batch_size+1] - int* ne_token_table, // [max_running_reqs, max_context_len] - int max_context_len, // max_context_len - long* row_indices, // [batch_size] - int* column_starts, // [batch_size] - int* n_gram_ids // [ne_n-1,ne_k,token_num] + int* ne_weights, // [ne_n-1,ne_k,ne_n] + int* ne_mods, // [ne_n-1,ne_k] + int* exclusive_ne_embeder_size_sums, // [(ne_n-1)*ne_k] + int* tokens, // [token_num] + int* exclusive_req_len_sums, // [batch_size+1] + int* ne_token_table, // [max_running_reqs, max_context_len] + int max_context_len, // max_context_len + const int64_t* __restrict__ row_indices, // [batch_size] + int* column_starts, // [batch_size] + int* n_gram_ids // [ne_n-1,ne_k,token_num] ) { // 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) { uint64_t n_gram_id = 0; // 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 - 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(max_context_len); // 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++) { if (current_token_table_index - j < req_token_table_index) { // Out of this request's range, stop computing n_gram_id @@ -134,14 +135,14 @@ __global__ void ComputeNGramIdsDecodeKernel( __global__ void UpdateTokenTableKernel( int batch_size, - int* tokens, // [token_num] - int* ne_token_table, // [max_running_reqs, max_context_len] - int max_context_len, // max_context_len - long* row_indices, // [batch_size] - int* column_starts, // [batch_size] - int* req_lens, // [batch_size] - int ignore_token_num, // number of tokens to ignore - int* ignore_tokens // [ignore_token_num] + int* tokens, // [token_num] + int* ne_token_table, // [max_running_reqs, max_context_len] + int max_context_len, // max_context_len + const int64_t* __restrict__ row_indices, // [batch_size] + int* column_starts, // [batch_size] + int* req_lens, // [batch_size] + int ignore_token_num, // number of tokens to ignore + int* ignore_tokens // [ignore_token_num] ) { // Each block processes one request. const int req_id = blockIdx.x % batch_size; @@ -154,11 +155,11 @@ __global__ void UpdateTokenTableKernel( // stride loop for (int i = start + threadIdx.x; i < end; i += blockDim.x) { // 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 - 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(max_context_len); // 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]; for (int j = 0; j < ignore_token_num; j++) { 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(max_context_len) + column_starts[req_id]; + ne_token_table[token_table_offset] = tokens[req_id]; + } +} + } // namespace device::ngram_embedding namespace { @@ -256,7 +272,7 @@ struct NgramEmbeddingKernel { static_cast(exclusive_req_len_sums.data_ptr()), static_cast(ne_token_table.data_ptr()), max_context_len, - static_cast(row_indices.data_ptr()), + static_cast(row_indices.data_ptr()), static_cast(column_starts.data_ptr()), static_cast(n_gram_ids.data_ptr())); } @@ -412,12 +428,64 @@ struct NgramEmbeddingKernel { static_cast(tokens.data_ptr()), static_cast(ne_token_table.data_ptr()), max_context_len, - static_cast(row_indices.data_ptr()), + static_cast(row_indices.data_ptr()), static_cast(column_starts.data_ptr()), static_cast(req_lens.data_ptr()), ignore_token_num, 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() + .with_device(device_) + .verify(tokens); + + TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len] + .with_dtype() + .with_device() + .verify(ne_token_table); + + TensorMatcher({batch_size}) // [batch_size] + .with_dtype() + .with_device() + .verify(row_indices); + + TensorMatcher({batch_size}) // [batch_size] + .with_dtype() + .with_device() + .verify(column_starts); + + const int bs = static_cast(batch_size.unwrap()); + if (bs <= 0) { + return; + } + + const int max_context_len = static_cast(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(host::div_ceil(bs, kBlockSize))); + + LaunchKernel(grid_size, kBlockSize, stream)( + device::ngram_embedding::UpdateTokenTableDecodeKernel, + bs, + static_cast(tokens.data_ptr()), + static_cast(ne_token_table.data_ptr()), + max_context_len, + static_cast(row_indices.data_ptr()), + static_cast(column_starts.data_ptr())); + } }; } // namespace diff --git a/python/sglang/jit_kernel/ngram_embedding.py b/python/sglang/jit_kernel/ngram_embedding.py index 1d9613f98..4679b3828 100644 --- a/python/sglang/jit_kernel/ngram_embedding.py +++ b/python/sglang/jit_kernel/ngram_embedding.py @@ -22,6 +22,10 @@ def _jit_ngram_embedding_module() -> Module: "&NgramEmbeddingKernel::compute_n_gram_ids_decode", ), ("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, 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, + ) diff --git a/python/sglang/jit_kernel/tests/test_ngram_embedding.py b/python/sglang/jit_kernel/tests/test_ngram_embedding.py index fe0b0162f..0e0f34ca7 100644 --- a/python/sglang/jit_kernel/tests/test_ngram_embedding.py +++ b/python/sglang/jit_kernel/tests/test_ngram_embedding.py @@ -6,6 +6,8 @@ import torch from sglang.jit_kernel.ngram_embedding import ( compute_n_gram_ids, compute_n_gram_ids_decode, + update_token_table, + update_token_table_decode, ) 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) +@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__": sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 043bf713e..5c9dc6429 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -34,7 +34,7 @@ import torch import torch.distributed as dist 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 ( BailingHybridConfig, FalconH1Config, @@ -2840,13 +2840,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): forward_batch.seq_lens ) 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, tokens=next_token_ids.to(torch.int32), row_indices=forward_batch.req_pool_indices, 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):