From a44bb397a94f66c08b50145fe40819799d1d9ee1 Mon Sep 17 00:00:00 2001 From: Alex Nails Date: Fri, 4 Sep 2026 17:26:25 -0700 Subject: [PATCH] [Perf] Vectorize alloc_extend_naive to remove the per-request Python loop (#37938) --- .../sglang/srt/mem_cache/allocator/paged.py | 98 +++++++++++-------- 1 file changed, 57 insertions(+), 41 deletions(-) diff --git a/python/sglang/srt/mem_cache/allocator/paged.py b/python/sglang/srt/mem_cache/allocator/paged.py index fcb2b3ee2..ff2d7b52b 100755 --- a/python/sglang/srt/mem_cache/allocator/paged.py +++ b/python/sglang/srt/mem_cache/allocator/paged.py @@ -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):