From 84c09913eb1458278f62f9dc393007141d3b67c3 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Wed, 4 Feb 2026 02:09:55 -0800 Subject: [PATCH] Moving _alloc_extend_naive out of npu allocator (#18200) --- .../srt/hardware_backend/npu/allocator_npu.py | 61 ++----------------- python/sglang/srt/mem_cache/allocator.py | 60 ++++++++++++++++++ 2 files changed, 65 insertions(+), 56 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/allocator_npu.py b/python/sglang/srt/hardware_backend/npu/allocator_npu.py index e863de62b..01842218e 100644 --- a/python/sglang/srt/hardware_backend/npu/allocator_npu.py +++ b/python/sglang/srt/hardware_backend/npu/allocator_npu.py @@ -2,67 +2,16 @@ from typing import TYPE_CHECKING import torch -from sglang.srt.mem_cache.allocator import PagedTokenToKVPoolAllocator +from sglang.srt.mem_cache.allocator import ( + PagedTokenToKVPoolAllocator, + alloc_extend_naive, +) from sglang.srt.utils import get_num_new_pages, next_power_of_2 if TYPE_CHECKING: from sglang.srt.mem_cache.memory_pool import KVCache -def _alloc_extend_naive( - prefix_lens, - seq_lens, - last_loc, - free_pages, - out_indices, - page_size, - device, -): - 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] - ) - if num1: - out_indices[start_pos[i] : start_pos[i] + num1] = ( - last_loc[i] + 1 + pos_in_page[:num1].view(-1) - ) - - 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) - - 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) - - class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): def __init__( self, @@ -128,7 +77,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): dtype=torch.int32, device=self.device, ) - _alloc_extend_naive( + alloc_extend_naive( prefix_lens, seq_lens, last_loc, diff --git a/python/sglang/srt/mem_cache/allocator.py b/python/sglang/srt/mem_cache/allocator.py index eaf29628b..9b6be180e 100644 --- a/python/sglang/srt/mem_cache/allocator.py +++ b/python/sglang/srt/mem_cache/allocator.py @@ -171,6 +171,66 @@ class TokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): return self._kvcache.load_cpu_copy(kv_cache_cpu, indices) +def alloc_extend_naive( + prefix_lens, + seq_lens, + last_loc, + free_pages, + out_indices, + page_size, + device, +): + 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] + ) + 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]: + continue + + 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) + + if prefix_lens[i] + num1 + num2 == seq_lens[i]: + continue + + 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) + + @triton.jit def alloc_extend_kernel( pre_lens_ptr,