[Perf] Vectorize alloc_extend_naive to remove the per-request Python loop (#37938)
This commit is contained in:
@@ -54,52 +54,68 @@ def alloc_extend_naive(
|
|||||||
extend_lens = seq_lens - prefix_lens
|
extend_lens = seq_lens - prefix_lens
|
||||||
end_pos = torch.cumsum(extend_lens, 0)
|
end_pos = torch.cumsum(extend_lens, 0)
|
||||||
start_pos = end_pos - extend_lens
|
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]
|
|
||||||
)
|
|
||||||
if num1:
|
|
||||||
out_indices[start_pos[i] : start_pos[i] + num1] = (
|
|
||||||
last_loc[i] + 1 + pos_in_page[:num1].view(-1)
|
|
||||||
)
|
|
||||||
|
|
||||||
if prefix_lens[i] + num1 == seq_lens[i]:
|
extend_num_tokens = out_indices.shape[0]
|
||||||
continue
|
if extend_num_tokens == 0:
|
||||||
|
return
|
||||||
|
|
||||||
num2 = (
|
j = torch.arange(extend_num_tokens, device=device, dtype=torch.int64)
|
||||||
seq_lens[i] // page_size - (prefix_lens[i] + page_size - 1) // page_size
|
owner = torch.searchsorted(end_pos, j, right=True)
|
||||||
) * page_size
|
local = j - start_pos[owner]
|
||||||
if num2:
|
last_loc_g = last_loc[owner]
|
||||||
pages = (
|
|
||||||
free_pages[start_new_pages[i] : end_new_pages[i] - need_page[i]]
|
ceil_prefix = (prefix_lens + page_size - 1) // page_size * 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"
|
||||||
)
|
)
|
||||||
out_indices[start_pos[i] + num1 : start_pos[i] + num1 + num2] = (
|
out_indices.copy_(last_loc_g + 1 + local)
|
||||||
pages.view(-1, 1) + pos_in_page.view(1, -1)
|
return
|
||||||
).view(-1)
|
|
||||||
|
|
||||||
if prefix_lens[i] + num1 + num2 == seq_lens[i]:
|
num1 = torch.clamp(seq_lens, max=ceil_prefix) - prefix_lens
|
||||||
continue
|
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)
|
||||||
|
|
||||||
num3 = seq_lens[i] - seq_lens[i] // page_size * page_size
|
full_pages = num2 // page_size
|
||||||
if num3:
|
need_extra_page = (num3 > 0).to(torch.int64)
|
||||||
out_indices[end_pos[i] - num3 : end_pos[i]] = (
|
pages_per_req = full_pages + need_extra_page
|
||||||
free_pages[end_new_pages[i] - 1] * page_size + pos_in_page[:num3]
|
end_new_pages = torch.cumsum(pages_per_req, 0)
|
||||||
).view(-1)
|
start_new_pages = end_new_pages - pages_per_req
|
||||||
|
|
||||||
|
num1_g = num1[owner]
|
||||||
|
num2_g = num2[owner]
|
||||||
|
start_new_pages_g = start_new_pages[owner]
|
||||||
|
end_new_pages_g = end_new_pages[owner]
|
||||||
|
|
||||||
|
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):
|
class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||||
|
|||||||
Reference in New Issue
Block a user