[Triton] Use dynamic loop bound in alloc_extend_kernel (#19898)
This commit is contained in:
@@ -240,7 +240,6 @@ def alloc_extend_kernel(
|
|||||||
out_indices,
|
out_indices,
|
||||||
bs_upper: tl.constexpr,
|
bs_upper: tl.constexpr,
|
||||||
page_size: tl.constexpr,
|
page_size: tl.constexpr,
|
||||||
max_num_extend_tokens: tl.constexpr,
|
|
||||||
):
|
):
|
||||||
pid = tl.program_id(0)
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
@@ -280,15 +279,16 @@ def alloc_extend_kernel(
|
|||||||
if pre_len + num_part1 == seq_len:
|
if pre_len + num_part1 == seq_len:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Part 2: fill the new full pages. Use blocked loop (tl.arange(0, BLOCK_EXTEND))
|
# Part 2: fill the new full pages using a dynamic blocked loop.
|
||||||
# instead of tl.arange(0, max_num_extend_tokens): Triton does not handle very
|
# The loop bound is derived from num_part2 (runtime value), so Triton
|
||||||
# large arange well; block size 4096 avoids the bottleneck.
|
# generates a real loop instead of unrolling — no constexpr dependency
|
||||||
|
# on extend size and only one kernel compilation.
|
||||||
num_part2 = (
|
num_part2 = (
|
||||||
seq_len // page_size * page_size
|
seq_len // page_size * page_size
|
||||||
- (pre_len + page_size - 1) // page_size * page_size
|
- (pre_len + page_size - 1) // page_size * page_size
|
||||||
)
|
)
|
||||||
BLOCK_EXTEND: tl.constexpr = 4096
|
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):
|
for block_id in range(num_blocks):
|
||||||
offset_in_block = tl.arange(0, BLOCK_EXTEND)
|
offset_in_block = tl.arange(0, BLOCK_EXTEND)
|
||||||
offset = block_id * BLOCK_EXTEND + offset_in_block
|
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)
|
super().__init__(size, page_size, dtype, device, kvcache, need_sort)
|
||||||
self.num_pages = size // page_size
|
self.num_pages = size // page_size
|
||||||
self.debug_mode = get_bool_env_var("SGLANG_DEBUG_MEMORY_POOL")
|
self.debug_mode = get_bool_env_var("SGLANG_DEBUG_MEMORY_POOL")
|
||||||
self.seen_max_num_extend_tokens_next_power_of_2 = 1
|
|
||||||
self.clear()
|
self.clear()
|
||||||
|
|
||||||
def alloc(self, need_size: int):
|
def alloc(self, need_size: int):
|
||||||
@@ -415,11 +414,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
(last_loc + 1) % self.page_size == prefix_lens % self.page_size
|
(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)
|
bs = len(prefix_lens)
|
||||||
if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len(
|
if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len(
|
||||||
self.free_pages
|
self.free_pages
|
||||||
@@ -430,27 +424,15 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
(extend_num_tokens,), dtype=torch.int64, device=self.device
|
(extend_num_tokens,), dtype=torch.int64, device=self.device
|
||||||
)
|
)
|
||||||
|
|
||||||
if extend_num_tokens < tl.core.TRITON_MAX_TENSOR_NUMEL:
|
alloc_extend_kernel[(bs,)](
|
||||||
alloc_extend_kernel[(bs,)](
|
prefix_lens,
|
||||||
prefix_lens,
|
seq_lens,
|
||||||
seq_lens,
|
last_loc,
|
||||||
last_loc,
|
self.free_pages,
|
||||||
self.free_pages,
|
out_indices,
|
||||||
out_indices,
|
next_power_of_2(bs),
|
||||||
next_power_of_2(bs),
|
self.page_size,
|
||||||
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.debug_mode:
|
if self.debug_mode:
|
||||||
assert len(torch.unique(out_indices)) == len(out_indices)
|
assert len(torch.unique(out_indices)) == len(out_indices)
|
||||||
|
|||||||
Reference in New Issue
Block a user