From 33c92732f46891261c9d2b4c1415925a82d44395 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 4 Mar 2026 16:15:58 -0800 Subject: [PATCH] [Triton] Use dynamic loop bound in `alloc_extend_kernel` (#19898) --- python/sglang/srt/mem_cache/allocator.py | 46 ++++++++---------------- 1 file changed, 14 insertions(+), 32 deletions(-) diff --git a/python/sglang/srt/mem_cache/allocator.py b/python/sglang/srt/mem_cache/allocator.py index d0f646a18..c0ce15e20 100644 --- a/python/sglang/srt/mem_cache/allocator.py +++ b/python/sglang/srt/mem_cache/allocator.py @@ -240,7 +240,6 @@ def alloc_extend_kernel( out_indices, bs_upper: tl.constexpr, page_size: tl.constexpr, - max_num_extend_tokens: tl.constexpr, ): pid = tl.program_id(0) @@ -280,15 +279,16 @@ def alloc_extend_kernel( if pre_len + num_part1 == seq_len: return - # Part 2: fill the new full pages. Use blocked loop (tl.arange(0, BLOCK_EXTEND)) - # instead of tl.arange(0, max_num_extend_tokens): Triton does not handle very - # large arange well; block size 4096 avoids the bottleneck. + # 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 = (max_num_extend_tokens + BLOCK_EXTEND - 1) // BLOCK_EXTEND + 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 @@ -375,7 +375,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): super().__init__(size, page_size, dtype, device, kvcache, need_sort) self.num_pages = size // page_size self.debug_mode = get_bool_env_var("SGLANG_DEBUG_MEMORY_POOL") - self.seen_max_num_extend_tokens_next_power_of_2 = 1 self.clear() def alloc(self, need_size: int): @@ -415,11 +414,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): (last_loc + 1) % self.page_size == prefix_lens % self.page_size ) - self.seen_max_num_extend_tokens_next_power_of_2 = max( - self.seen_max_num_extend_tokens_next_power_of_2, - min(tl.core.TRITON_MAX_TENSOR_NUMEL, next_power_of_2(extend_num_tokens)), - ) - bs = len(prefix_lens) if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len( self.free_pages @@ -430,27 +424,15 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): (extend_num_tokens,), dtype=torch.int64, device=self.device ) - if extend_num_tokens < tl.core.TRITON_MAX_TENSOR_NUMEL: - alloc_extend_kernel[(bs,)]( - prefix_lens, - seq_lens, - last_loc, - self.free_pages, - out_indices, - next_power_of_2(bs), - self.page_size, - self.seen_max_num_extend_tokens_next_power_of_2, - ) - else: - alloc_extend_naive( - prefix_lens, - seq_lens, - last_loc, - self.free_pages, - out_indices, - self.page_size, - self.device, - ) + alloc_extend_kernel[(bs,)]( + prefix_lens, + seq_lens, + last_loc, + self.free_pages, + out_indices, + next_power_of_2(bs), + self.page_size, + ) if self.debug_mode: assert len(torch.unique(out_indices)) == len(out_indices)