[Perf] Vectorize alloc_extend_naive to remove the per-request Python loop (#37938)

This commit is contained in:
Alex Nails
2026-09-04 17:26:25 -07:00
committed by GitHub
parent 613d87becd
commit a44bb397a9
+57 -41
View File
@@ -54,52 +54,68 @@ def alloc_extend_naive(
extend_lens = seq_lens - prefix_lens
end_pos = torch.cumsum(extend_lens, 0)
start_pos = end_pos - extend_lens
num_new_pages = (seq_lens + page_size - 1) // page_size - (
prefix_lens + page_size - 1
) // page_size
num_full_new_pages = (seq_lens) // page_size - (
prefix_lens + page_size - 1
) // page_size
need_page = num_new_pages - num_full_new_pages
end_new_pages = torch.cumsum(num_new_pages, 0)
start_new_pages = end_new_pages - num_new_pages
pos_in_page = torch.arange(page_size, device=device, dtype=torch.int32)
for i in range(len(prefix_lens)):
num1 = (
min(
seq_lens[i],
(prefix_lens[i] + page_size - 1) // page_size * page_size,
)
- prefix_lens[i]
extend_num_tokens = out_indices.shape[0]
if extend_num_tokens == 0:
return
j = torch.arange(extend_num_tokens, device=device, dtype=torch.int64)
owner = torch.searchsorted(end_pos, j, right=True)
local = j - start_pos[owner]
last_loc_g = last_loc[owner]
ceil_prefix = (prefix_lens + page_size - 1) // page_size * page_size
floor_seq = seq_lens // page_size * page_size
if free_pages.numel() == 0:
# Only valid when no request needs a new page; nothing below indexes
# the empty pool, so a short pool would silently return garbage.
ceil_seq = (seq_lens + page_size - 1) // page_size * page_size
assert torch.all(ceil_seq == ceil_prefix), (
"alloc_extend_naive: free_pages is empty but the batch requires "
"new pages; caller must ensure pool >= demand"
)
if num1:
out_indices[start_pos[i] : start_pos[i] + num1] = (
last_loc[i] + 1 + pos_in_page[:num1].view(-1)
)
out_indices.copy_(last_loc_g + 1 + local)
return
if prefix_lens[i] + num1 == seq_lens[i]:
continue
num1 = torch.clamp(seq_lens, max=ceil_prefix) - prefix_lens
done_after_1 = (prefix_lens + num1) == seq_lens
num2 = torch.where(done_after_1, torch.zeros_like(num1), floor_seq - ceil_prefix)
num3 = torch.where(done_after_1, torch.zeros_like(num1), seq_lens - floor_seq)
num2 = (
seq_lens[i] // page_size - (prefix_lens[i] + page_size - 1) // page_size
) * page_size
if num2:
pages = (
free_pages[start_new_pages[i] : end_new_pages[i] - need_page[i]]
* page_size
)
out_indices[start_pos[i] + num1 : start_pos[i] + num1 + num2] = (
pages.view(-1, 1) + pos_in_page.view(1, -1)
).view(-1)
full_pages = num2 // page_size
need_extra_page = (num3 > 0).to(torch.int64)
pages_per_req = full_pages + need_extra_page
end_new_pages = torch.cumsum(pages_per_req, 0)
start_new_pages = end_new_pages - pages_per_req
if prefix_lens[i] + num1 + num2 == seq_lens[i]:
continue
num1_g = num1[owner]
num2_g = num2[owner]
start_new_pages_g = start_new_pages[owner]
end_new_pages_g = end_new_pages[owner]
num3 = seq_lens[i] - seq_lens[i] // page_size * page_size
if num3:
out_indices[end_pos[i] - num3 : end_pos[i]] = (
free_pages[end_new_pages[i] - 1] * page_size + pos_in_page[:num3]
).view(-1)
is_phase1 = local < num1_g
is_phase2 = (~is_phase1) & (local < num1_g + num2_g)
val_phase1 = last_loc_g + 1 + local
rel2 = torch.clamp(local - num1_g, min=0)
# torch.where below evaluates both branches per slot, so dead-lane page
# indices must stay in range even where their phase is never selected.
page_idx2 = torch.clamp(
start_new_pages_g + rel2 // page_size, min=0, max=free_pages.numel() - 1
)
pos_in_page2 = rel2 % page_size
val_phase2 = free_pages[page_idx2] * page_size + pos_in_page2
rel3 = torch.clamp(local - num1_g - num2_g, min=0)
page_idx3 = torch.clamp(end_new_pages_g - 1, min=0, max=free_pages.numel() - 1)
val_phase3 = free_pages[page_idx3] * page_size + rel3
out = torch.where(
is_phase1, val_phase1, torch.where(is_phase2, val_phase2, val_phase3)
)
out_indices.copy_(out)
class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):