Fix triton alloc extend kernel (#19780)
This commit is contained in:
@@ -280,22 +280,28 @@ 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
|
# 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.
|
||||||
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
|
||||||
offset_many_page = tl.arange(0, max_num_extend_tokens)
|
num_blocks = (max_num_extend_tokens + BLOCK_EXTEND - 1) // BLOCK_EXTEND
|
||||||
page_start = tl.load(
|
for block_id in range(num_blocks):
|
||||||
free_page_ptr + new_page_start_loc + offset_many_page // page_size,
|
offset_in_block = tl.arange(0, BLOCK_EXTEND)
|
||||||
mask=offset_many_page < num_part2,
|
offset = block_id * BLOCK_EXTEND + offset_in_block
|
||||||
)
|
mask = offset < num_part2
|
||||||
tl.store(
|
page_start = tl.load(
|
||||||
out_indices + output_start_loc + num_part1 + offset_many_page,
|
free_page_ptr + new_page_start_loc + offset // page_size,
|
||||||
page_start * page_size + offset_many_page % page_size,
|
mask=mask,
|
||||||
mask=offset_many_page < num_part2,
|
)
|
||||||
)
|
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:
|
if pre_len + num_part1 + num_part2 == seq_len:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user