diff --git a/docs/advanced_features/nsa_prefill_cp_phase3_current_reuse_plan.md b/docs/advanced_features/nsa_prefill_cp_phase3_current_reuse_plan.md index 239edab05..04fbc7624 100644 --- a/docs/advanced_features/nsa_prefill_cp_phase3_current_reuse_plan.md +++ b/docs/advanced_features/nsa_prefill_cp_phase3_current_reuse_plan.md @@ -2,7 +2,7 @@ 本文档定义 **Phase 3**:在 Phase 2 已经实现 CP shared/sharded persistent KV 的基础上,减少 shared KV compatibility path 引入的 **可避免重复 KV/index materialize**。 -Phase 3 的边界是:**复用当前 forward/chunk 中已经 CP all-gather + rerange 过的 current KV/index**,避免它们再次从 sharded persistent KV/index pool 里 materialize。后续拆分为:Phase 4 做 page-aligned CP token split,Phase 5 做 compute-owner shared KV layout;更深层的 history KV shard-aware topk、selected KV exchange、distributed sparse attention 继续后移。 +Phase 3 的边界是:**复用当前 forward/chunk 中已经 CP all-gather + rerange 过的 current KV/index**,避免它们再次从 sharded persistent KV/index pool 里 materialize。更深层的 history KV shard-aware topk、selected KV exchange、distributed sparse attention 进入 Phase 4。 --- @@ -91,7 +91,7 @@ history 仍走 Phase 2 compatibility materialize ### 1.3 非目标 -以下不属于 Phase 3,进入 Phase 4/5 或后续: +以下不属于 Phase 3,进入 Phase 4 或后续: ```text - shard-aware NSA topk @@ -373,10 +373,13 @@ metadata.page_table_1 ```text SGLANG_CP_SHARED_KV_CURRENT_REUSE=0/1 SGLANG_CP_SHARED_KV_NVTX=0/1 +SGLANG_DEBUG_SORT_NVTX=0/1 +SGLANG_DEBUG_MOE_SORT_NVTX=0/1 ``` 建议初期默认关闭,验证稳定后再默认开启。 `SGLANG_CP_SHARED_KV_NVTX=1` 只打开 shared-KV materialize 相关 NVTX range, +`SGLANG_DEBUG_SORT_NVTX=1` 只标记 CP shared KV current remap 与 KV allocator free-page merge 中可能触发 ATen sort kernel 的 `torch.sort` callsite。MoE EP preprocess sort 标记被拆到 `SGLANG_DEBUG_MOE_SORT_NVTX=1`,因为 MoE sort 调用频率很高,常规 profile 不应默认开启,避免 Nsight/CPU 侧被动态 NVTX range 拖慢。两个开关默认关闭,并在模块导入时缓存,避免热路径反复读取环境变量。 用于 Nsight Systems 定位 Phase2/Phase3 materialize、local copy、all-reduce 开销。 统计项: @@ -486,7 +489,7 @@ history index materialize + current compact index -> topk result maintains logical loc semantics ``` -如果 topk loc semantic 不容易保证,允许保留 Phase 2 fallback,并把 mixed indexer 推迟到 Phase 5 之后的 shard-aware runtime 任务。 +如果 topk loc semantic 不容易保证,允许保留 Phase 2 fallback,并把 mixed indexer 推迟到 Phase 4 前置任务。 --- @@ -615,23 +618,17 @@ MLA current-only + MLA mixed + NSA indexer current-only --- -## 9. Phase 4 / Phase 5 边界 +## 9. Phase 4 边界 -Phase 4 先解决 CP split 与 KV page 不对齐的问题: +Phase 4 继续解决 history/shared KV runtime 成本: ```text -- page-aligned in-seq-split -- 保证一个 current KV page 只由一个 CP compute rank 计算 -- 不改变 shared KV persistent layout +- history selected-page materialize +- shard-aware NSA topk +- local topk + global merge +- selected KV owner exchange +- distributed sparse attention +- owner-routed current write,消除写入前 current chunk all-gather 后再过滤 owner 的浪费 ``` -Phase 5 再解决 persistent KV layout 与 compute owner 不一致的问题: - -```text -- compute-owner logical page allocation -- CP-local out_cache_loc -- local MLA KV / NSA index K direct persistent write -- 保持 materialize compatibility path 与 Mooncake PD transfer 正确 -``` - -history selected-page materialize、shard-aware NSA topk、local topk + global merge、selected KV owner exchange、distributed sparse attention 继续后移。Phase 3 不隐藏这些问题。Phase 3 只保证:**已经 gather 过的 current chunk 不再被 shared KV compatibility path 重复 gather/materialize。** +Phase 3 不隐藏 Phase 4 问题。Phase 3 只保证:**已经 gather 过的 current chunk 不再被 shared KV compatibility path 重复 gather/materialize。** diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 1098f2406..5f6ac05e5 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -43,7 +43,6 @@ from sglang.srt.disaggregation.utils import ( TransferBackend, get_kv_class, is_mla_backend, - kv_to_page_indices, poll_and_all_reduce, prepare_abort, ) @@ -78,6 +77,29 @@ if TYPE_CHECKING: CLIP_MAX_NEW_TOKEN = envs.SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION.get() +def _kv_locs_to_page_indices_cpu( + kv_locs: torch.Tensor, + page_size: int, + *, + num_tokens: Optional[int] = None, +): + """Return compact int32 page indices for a token-location tensor. + + PD bootstrap only needs page-level destination indices. Avoid materializing + a full token-level CPU copy before converting to pages. + """ + + if num_tokens is not None: + kv_locs = kv_locs[:num_tokens] + + if page_size == 1: + page_locs = kv_locs + else: + page_locs = kv_locs[::page_size] // page_size + + return page_locs.to(dtype=torch.int32).cpu().numpy() + + def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool: return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or ( req.bootstrap_host is None @@ -409,7 +431,11 @@ class DecodePreallocQueue: ) self.queue.append( - DecodeRequest(req=req, kv_receiver=kv_receiver, waiting_for_input=False) + DecodeRequest( + req=req, + kv_receiver=kv_receiver, + waiting_for_input=False, + ) ) def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool: @@ -675,16 +701,17 @@ class DecodePreallocQueue: continue allocatable_tokens -= required_tokens_for_request - self._pre_alloc(decode_req.req) - - kv_indices = ( - self.req_to_token_pool.req_to_token[decode_req.req.req_pool_idx][ - : len(decode_req.req.origin_input_ids) - ] - .cpu() - .numpy() + kv_loc = self._pre_alloc( + decode_req.req, + write_req_to_token=False, ) + page_size = self.token_to_kv_pool_allocator.page_size + page_indices = _kv_locs_to_page_indices_cpu( + kv_loc, + page_size, + num_tokens=origin_input_len, + ) # Prepare extra pool indices for hybrid models if isinstance(self.token_to_kv_pool, HybridLinearKVPool): @@ -703,9 +730,7 @@ class DecodePreallocQueue: window_start = max(0, seq_len - window_size) window_start = (window_start // page_size) * page_size - window_kv_indices_full = self.req_to_token_pool.req_to_token[ - decode_req.req.req_pool_idx, window_start:seq_len - ] + window_kv_indices_full = kv_loc[window_start:seq_len] # Translate to SWA pool indices window_kv_indices_swa = ( @@ -713,23 +738,23 @@ class DecodePreallocQueue: window_kv_indices_full ) ) - state_indices = window_kv_indices_swa.cpu().numpy() - state_indices = kv_to_page_indices(state_indices, page_size) + state_indices = _kv_locs_to_page_indices_cpu( + window_kv_indices_swa, + page_size, + ) elif isinstance(self.token_to_kv_pool, NSATokenToKVPool): - seq_len = len(decode_req.req.origin_input_ids) - kv_indices_full = self.req_to_token_pool.req_to_token[ - decode_req.req.req_pool_idx, :seq_len - ] - state_indices = kv_indices_full.cpu().numpy() - state_indices = kv_to_page_indices(state_indices, page_size) + state_indices = page_indices else: state_indices = None + self.req_to_token_pool.write( + (decode_req.req.req_pool_idx, slice(0, len(kv_loc))), kv_loc + ) + decode_req.metadata_buffer_index = ( self.req_to_metadata_buffer_idx_allocator.alloc() ) assert decode_req.metadata_buffer_index is not None - page_indices = kv_to_page_indices(kv_indices, page_size) decode_req.kv_receiver.init( page_indices, decode_req.metadata_buffer_index, state_indices ) @@ -840,7 +865,12 @@ class DecodePreallocQueue: last_loc, ) - def _pre_alloc(self, req: Req) -> torch.Tensor: + def _pre_alloc( + self, + req: Req, + *, + write_req_to_token: bool = True, + ) -> torch.Tensor: """Pre-allocate the memory for req_to_token and token_kv_pool""" req_pool_indices = self.req_to_token_pool.alloc([req]) @@ -875,7 +905,10 @@ class DecodePreallocQueue: "KV cache is full! There is a bug in memory estimation." ) - self.req_to_token_pool.write((req.req_pool_idx, slice(0, len(kv_loc))), kv_loc) + if write_req_to_token: + self.req_to_token_pool.write( + (req.req_pool_idx, slice(0, len(kv_loc))), kv_loc + ) # populate metadata req.fill_ids = req.origin_input_ids + req.output_ids diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 1a0a3ce2c..d4a5539df 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -37,7 +37,6 @@ from sglang.srt.disaggregation.utils import ( TransferBackend, get_kv_class, is_mla_backend, - kv_to_page_indices, kv_to_page_num, poll_and_all_reduce_attn_cp_tp_group, prepare_abort, @@ -62,6 +61,20 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +def _kv_locs_to_page_indices_cpu( + kv_locs: torch.Tensor, + page_size: int, +): + """Return int32 page indices without materializing token-level CPU indices.""" + + if page_size == 1: + page_locs = kv_locs + else: + page_locs = kv_locs[::page_size] // page_size + + return page_locs.to(dtype=torch.int32).cpu().numpy() + + def release_req_to_metadata_buffer( req: Req, allocator: ReqToMetadataIdxAllocator ) -> None: @@ -723,10 +736,9 @@ class SchedulerDisaggregationPrefillMixin: # if not the last chunk and the last page is partial, delay the last partial page to the next send end_idx = end_idx - end_idx % page_size - kv_indices = ( - self.req_to_token_pool.req_to_token[req.req_pool_idx, start_idx:end_idx] - .cpu() - .numpy() + page_indices = _kv_locs_to_page_indices_cpu( + self.req_to_token_pool.req_to_token[req.req_pool_idx, start_idx:end_idx], + page_size, ) req.start_send_idx = end_idx state_indices = None @@ -762,19 +774,19 @@ class SchedulerDisaggregationPrefillMixin: window_kv_indices_full ) ) - state_indices = window_kv_indices_swa.cpu().numpy() - state_indices = kv_to_page_indices(state_indices, page_size) + state_indices = _kv_locs_to_page_indices_cpu( + window_kv_indices_swa, + page_size, + ) elif isinstance( self.token_to_kv_pool_allocator.get_kvcache(), NSATokenToKVPool ): seq_len = len(req.fill_ids) - kv_indices_full = self.req_to_token_pool.req_to_token[ - req.req_pool_idx, :seq_len - ] - state_indices = kv_indices_full.cpu().numpy() - state_indices = kv_to_page_indices(state_indices, page_size) + state_indices = _kv_locs_to_page_indices_cpu( + self.req_to_token_pool.req_to_token[req.req_pool_idx, :seq_len], + page_size, + ) - page_indices = kv_to_page_indices(kv_indices, page_size) if len(page_indices) == 0: logger.info( f"Skip sending kv chunk for request {req.rid=} {req.bootstrap_room=} because page_indices is empty" diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 9526820b5..ce9cdd6ef 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -203,6 +203,8 @@ class Envs: SGLANG_FORCE_SHUTDOWN = EnvBool(False) SGLANG_DEBUG_MEMORY_POOL = EnvBool(False) SGLANG_DEBUG_CP_SHARED_KV = EnvBool(False) + SGLANG_DEBUG_SORT_NVTX = EnvBool(False) + SGLANG_DEBUG_MOE_SORT_NVTX = EnvBool(False) SGLANG_CP_SHARED_KV_CURRENT_REUSE = EnvBool(False) SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE = EnvBool(False) SGLANG_CP_SHARED_KV_ENABLE_MLA_PREFETCH = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 625f70d60..3df144bcd 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -15,6 +15,7 @@ logger = logging.getLogger(__name__) _DEBUG_LOG_COUNTS: dict[str, int] = {} _TAI_MATERIALIZE_FALLBACK_LOG_COUNTS: dict[str, int] = {} _MLA_PREFETCH_LOG_PROBE_LAYER = 2 +_SORT_NVTX_ENABLED = envs.SGLANG_DEBUG_SORT_NVTX.get() def cp_shared_kv_debug_enabled() -> bool: @@ -25,6 +26,10 @@ def cp_shared_kv_current_reuse_enabled() -> bool: return envs.SGLANG_CP_SHARED_KV_CURRENT_REUSE.get() +def cp_shared_kv_sort_nvtx_enabled() -> bool: + return _SORT_NVTX_ENABLED + + def cp_shared_kv_tai_materialize_enabled() -> bool: return envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get() @@ -725,6 +730,9 @@ def remap_logical_locs_to_dense_locs( def build_current_loc_remap( query_locs: torch.Tensor, current_locs: torch.Tensor, + *, + page_size: int | None = None, + logical_page_capacity: int | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Map logical locs into rows of the gathered current chunk tensor. @@ -732,6 +740,11 @@ def build_current_loc_remap( `current_locs` is `forward_batch.out_cache_loc`; its row order is the row order of the already CP-all-gathered current KV/index tensor. + When `page_size` and `logical_page_capacity` are provided, this uses a + page-inverse fast path for the single-request current-only case. That avoids + the PyTorch `torch.sort + searchsorted` path, which dispatches ATen/CCCL + radix-sort kernels in the attention hot path. + Returns: - is_current_mask: true where query_locs is a non-negative current loc. - compact_row_ids: row id into the current compact tensor where valid, @@ -745,7 +758,74 @@ def build_current_loc_remap( query_flat_long = query_locs.reshape(-1).to(torch.long) current_flat_long = current_locs.reshape(-1).to(torch.long) - sorted_current_locs, sorted_to_current_rows = torch.sort(current_flat_long) + + if page_size is not None and logical_page_capacity is not None: + if cp_shared_kv_sort_nvtx_enabled(): + torch.cuda.nvtx.range_push( + f"CP_SHARED_KV:current_loc_remap:page_inverse num_current={current_flat_long.numel()} num_query={query_flat_long.numel()} page_size={page_size}" + ) + try: + current_pages = torch.div( + current_flat_long[::page_size], page_size, rounding_mode="floor" + ) + page_inverse = torch.full( + (logical_page_capacity,), + -1, + device=current_flat_long.device, + dtype=torch.long, + ) + row_page_ids = torch.arange( + current_pages.numel(), device=current_flat_long.device, dtype=torch.long + ) + valid_current_pages = (current_pages >= 0) & ( + current_pages < logical_page_capacity + ) + safe_current_pages = torch.where( + valid_current_pages, current_pages, torch.zeros_like(current_pages) + ) + safe_row_page_ids = torch.where( + valid_current_pages, row_page_ids, torch.zeros_like(row_page_ids) + ) + page_inverse.scatter_(0, safe_current_pages, safe_row_page_ids) + + valid_query = query_flat_long >= 0 + safe_query_locs = torch.where( + valid_query, query_flat_long, torch.zeros_like(query_flat_long) + ) + query_pages = torch.div( + safe_query_locs, page_size, rounding_mode="floor" + ) + query_offsets = torch.remainder(safe_query_locs, page_size) + query_pages_in_range = query_pages < logical_page_capacity + safe_query_pages = torch.clamp(query_pages, max=logical_page_capacity - 1) + row_pages = page_inverse[safe_query_pages] + row_values = row_pages * page_size + query_offsets + matched = ( + valid_query + & query_pages_in_range + & (row_pages >= 0) + & (row_values < current_flat_long.numel()) + ) + compact_flat = torch.where( + matched, + row_values.to(compact_row_ids.dtype), + torch.full_like(compact_row_ids.reshape(-1), -1), + ) + return matched.reshape(query_locs.shape), compact_flat.reshape(query_locs.shape) + finally: + if cp_shared_kv_sort_nvtx_enabled(): + torch.cuda.nvtx.range_pop() + + if cp_shared_kv_sort_nvtx_enabled(): + torch.cuda.nvtx.range_push( + f"CP_SHARED_KV:current_loc_remap:torch.sort_fallback num_current={current_flat_long.numel()} num_query={query_flat_long.numel()}" + ) + try: + sorted_current_locs, sorted_to_current_rows = torch.sort(current_flat_long) + finally: + torch.cuda.nvtx.range_pop() + else: + sorted_current_locs, sorted_to_current_rows = torch.sort(current_flat_long) insert_positions = torch.searchsorted(sorted_current_locs, query_flat_long) safe_positions = torch.clamp(insert_positions, max=current_flat_long.numel() - 1) diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index bfd6d026f..fd7495faf 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -1667,9 +1667,25 @@ class NativeSparseAttnBackend( ) if can_reuse_current_kv: logical_page_table_1 = page_table_1 + current_remap_page_size = None + current_remap_logical_page_capacity = None + if len(getattr(forward_batch, "extend_seq_lens_cpu", []) or []) == 1: + current_remap_page_size = forward_batch.token_to_kv_pool.page_size + current_remap_logical_page_capacity = ( + max( + forward_batch.token_to_kv_pool.size + // current_remap_page_size + - 1, + 0, + ) + * forward_batch.cp_shared_kv_layout.cp_size + + 1 + ) current_mask, page_table_1 = build_current_loc_remap( logical_page_table_1, forward_batch.out_cache_loc, + page_size=current_remap_page_size, + logical_page_capacity=current_remap_logical_page_capacity, ) if cp_shared_kv_debug_enabled(): missing_current = (logical_page_table_1 >= 0) & (~current_mask) diff --git a/python/sglang/srt/layers/moe/ep_moe/kernels.py b/python/sglang/srt/layers/moe/ep_moe/kernels.py index 044c590f2..00c344ffe 100644 --- a/python/sglang/srt/layers/moe/ep_moe/kernels.py +++ b/python/sglang/srt/layers/moe/ep_moe/kernels.py @@ -3,10 +3,13 @@ import logging import torch import triton +from sglang.srt.environ import envs from sglang.srt.utils import ceil_div, is_cuda logger = logging.getLogger(__name__) +_MOE_SORT_NVTX_ENABLED = envs.SGLANG_DEBUG_MOE_SORT_NVTX.get() + _is_cuda = is_cuda() if _is_cuda: from sglang.srt.layers.quantization.fp8_kernel import ( @@ -16,6 +19,23 @@ if _is_cuda: import triton.language as tl +def _debug_sort_nvtx_enabled() -> bool: + return _MOE_SORT_NVTX_ENABLED + + +def _sort_with_optional_nvtx( + tensor: torch.Tensor, *, stable: bool, nvtx_name: str +): + if not _debug_sort_nvtx_enabled(): + return torch.sort(tensor, stable=stable) + + torch.cuda.nvtx.range_push(f"{nvtx_name} numel={tensor.numel()} dtype={tensor.dtype}") + try: + return torch.sort(tensor, stable=stable) + finally: + torch.cuda.nvtx.range_pop() + + def _get_launch_config_1d(device, numel): MAX_THREADS_PER_BLOCK = 1024 MIN_THREADS_PER_BLOCK = 512 @@ -158,7 +178,11 @@ def deepep_compute_src2dst_triton_kernel( def deepep_run_moe_deep_preprocess(topk_ids: torch.Tensor, num_experts: int): - reorder_topk_ids, reorder_ids = torch.sort(topk_ids.view(-1), stable=True) + reorder_topk_ids, reorder_ids = _sort_with_optional_nvtx( + topk_ids.view(-1), + stable=True, + nvtx_name="MOE_EP:deepep_preprocess:topk_ids:torch.sort", + ) seg_indptr = torch.empty(num_experts + 1, device=topk_ids.device, dtype=torch.int64) src2dst = torch.empty(topk_ids.numel(), device=topk_ids.device, dtype=torch.int64) @@ -197,7 +221,11 @@ def compute_seg_indptr_triton_kernel(reorder_topk_ids, seg_indptr, num_toks): def cutlass_w4_run_moe_ep_preproess(topk_ids: torch.Tensor): - _, reorder_ids = torch.sort(topk_ids.view(-1), stable=True) + _, reorder_ids = _sort_with_optional_nvtx( + topk_ids.view(-1), + stable=True, + nvtx_name="MOE_EP:cutlass_w4_preprocess:topk_ids:torch.sort", + ) BLOCK_SIZE = 512 grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) @@ -1046,7 +1074,11 @@ def moe_ep_deepgemm_preprocess( block_shape, output_dtype: torch.dtype = torch.float8_e4m3fn, ): - reorder_topk_ids, reorder_ids = torch.sort(topk_ids.view(-1), stable=True) + reorder_topk_ids, reorder_ids = _sort_with_optional_nvtx( + topk_ids.view(-1), + stable=True, + nvtx_name="MOE_EP:deepgemm_preprocess:topk_ids:torch.sort", + ) seg_indptr = torch.zeros( num_local_experts + 1, device=topk_ids.device, dtype=torch.int64 ) diff --git a/python/sglang/srt/mem_cache/allocator.py b/python/sglang/srt/mem_cache/allocator.py index 6d62a21fc..1731b8d08 100644 --- a/python/sglang/srt/mem_cache/allocator.py +++ b/python/sglang/srt/mem_cache/allocator.py @@ -24,14 +24,22 @@ from typing import TYPE_CHECKING, List, Optional import torch import triton + +from sglang.srt.environ import envs import triton.language as tl from sglang.srt.utils import get_bool_env_var, get_num_new_pages, next_power_of_2 +_SORT_NVTX_ENABLED = envs.SGLANG_DEBUG_SORT_NVTX.get() + if TYPE_CHECKING: from sglang.srt.mem_cache.memory_pool import KVCache +def _debug_sort_nvtx_enabled() -> bool: + return _SORT_NVTX_ENABLED + + class BaseTokenToKVPoolAllocator(abc.ABC): @abc.abstractmethod def __init__( @@ -81,8 +89,19 @@ class BaseTokenToKVPoolAllocator(abc.ABC): def merge_and_sort_free(self): if len(self.release_pages) > 0: + num_free_pages = len(self.free_pages) + num_release_pages = len(self.release_pages) self.free_pages = torch.cat((self.free_pages, self.release_pages)) - self.free_pages, _ = torch.sort(self.free_pages) + if _debug_sort_nvtx_enabled(): + torch.cuda.nvtx.range_push( + f"KV_ALLOCATOR:merge_and_sort_free:torch.sort free_pages={num_free_pages} release_pages={num_release_pages}" + ) + try: + self.free_pages, _ = torch.sort(self.free_pages) + finally: + torch.cuda.nvtx.range_pop() + else: + self.free_pages, _ = torch.sort(self.free_pages) self.release_pages = torch.empty( (0,), dtype=self.release_pages.dtype, device=self.device ) @@ -588,30 +607,69 @@ class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): def _select_compute_owner_pages( self, page_compute_owners: List[int], - ) -> Optional[torch.Tensor]: - selected_pages = [] - lane_offsets = [0 for _ in range(self.cp_size)] - lane_pages = [ - self.free_pages[ - torch.remainder(self.free_pages - 1, self.cp_size) == owner - ] - for owner in range(self.cp_size) - ] + ) -> Optional[tuple[torch.Tensor, torch.Tensor, torch.Tensor]]: + if not page_compute_owners: + return ( + torch.empty((0,), dtype=torch.int64, device=self.device), + torch.zeros_like(self.free_pages, dtype=torch.bool), + torch.zeros_like(self.release_pages, dtype=torch.bool), + ) + required_by_owner = [0 for _ in range(self.cp_size)] for owner in page_compute_owners: if owner < 0 or owner >= self.cp_size: raise ValueError( f"compute owner must be in [0, {self.cp_size}), got {owner}" ) + required_by_owner[owner] += 1 + + lane_pages = [None for _ in range(self.cp_size)] + selected_free_mask = torch.zeros_like(self.free_pages, dtype=torch.bool) + selected_release_mask = torch.zeros_like(self.release_pages, dtype=torch.bool) + for owner, required_count in enumerate(required_by_owner): + if required_count == 0: + continue + + owner_mask = torch.remainder(self.free_pages - 1, self.cp_size) == owner + selected_owner_free_mask = owner_mask & ( + torch.cumsum(owner_mask.to(torch.int64), dim=0) <= required_count + ) + selected_owner_pages = self.free_pages[selected_owner_free_mask] + + remaining_count = required_count - selected_owner_pages.numel() + if remaining_count > 0: + release_owner_mask = ( + torch.remainder(self.release_pages - 1, self.cp_size) == owner + ) + selected_owner_release_mask = release_owner_mask & ( + torch.cumsum(release_owner_mask.to(torch.int64), dim=0) + <= remaining_count + ) + selected_owner_release_pages = self.release_pages[ + selected_owner_release_mask + ] + if remaining_count > selected_owner_release_pages.numel(): + return None + selected_owner_pages = torch.cat( + (selected_owner_pages, selected_owner_release_pages) + ) + selected_release_mask |= selected_owner_release_mask + + lane_pages[owner] = selected_owner_pages + selected_free_mask |= selected_owner_free_mask + + selected_pages = [] + lane_offsets = [0 for _ in range(self.cp_size)] + for owner in page_compute_owners: lane_offset = lane_offsets[owner] - if lane_offset >= lane_pages[owner].numel(): - return None selected_pages.append(lane_pages[owner][lane_offset]) lane_offsets[owner] = lane_offset + 1 - if not selected_pages: - return torch.empty((0,), dtype=torch.int64, device=self.device) - return torch.stack(selected_pages).to(torch.int64) + return ( + torch.stack(selected_pages).to(torch.int64), + selected_free_mask, + selected_release_mask, + ) def alloc_extend_compute_owner( self, @@ -645,15 +703,10 @@ class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): f"{num_new_pages=} page_compute_owners={len(page_compute_owners)}" ) - if self.need_sort and num_new_pages > len(self.free_pages): - self.merge_and_sort_free() - - selected_pages = self._select_compute_owner_pages(page_compute_owners) - if selected_pages is None and self.need_sort and len(self.release_pages) > 0: - self.merge_and_sort_free() - selected_pages = self._select_compute_owner_pages(page_compute_owners) - if selected_pages is None: + selected = self._select_compute_owner_pages(page_compute_owners) + if selected is None: return None + selected_pages, selected_free_mask, selected_release_mask = selected out_indices = torch.empty( (extend_num_tokens,), dtype=torch.int64, device=self.device @@ -668,8 +721,8 @@ class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): self.device, ) - selected_mask = torch.isin(self.free_pages, selected_pages) - self.free_pages = self.free_pages[~selected_mask] + self.free_pages = self.free_pages[~selected_free_mask] + self.release_pages = self.release_pages[~selected_release_mask] if self.debug_mode: assert len(torch.unique(out_indices)) == len(out_indices) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 732d63909..129782935 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -364,13 +364,6 @@ def alloc_paged_token_slots_extend( "multi_batch" if len(prefix_lens_cpu) != 1 else "unknown" ) - if page_compute_owners is not None: - _evict_for_compute_owner_lanes( - tree_cache=tree_cache, - allocator=allocator, - page_compute_owners=page_compute_owners, - ) - state = None if backup_state: state = allocator.backup_state() @@ -385,6 +378,23 @@ def alloc_paged_token_slots_extend( extend_num_tokens, page_compute_owners, ) + if out_cache_loc is None: + _evict_for_compute_owner_lanes( + tree_cache=tree_cache, + allocator=allocator, + page_compute_owners=page_compute_owners, + ) + if backup_state: + state = allocator.backup_state() + out_cache_loc = alloc_extend_compute_owner( + prefix_lens, + prefix_lens_cpu, + seq_lens, + seq_lens_cpu, + last_loc, + extend_num_tokens, + page_compute_owners, + ) if out_cache_loc is None: required = available = deficits = None compute_owner_lane_stats = getattr( diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py b/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py index 5b569ae74..eebfd6b01 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py @@ -1,4 +1,5 @@ import unittest +from unittest.mock import patch import numpy as np import torch @@ -244,6 +245,170 @@ class TestCPSharedPagedAllocator(unittest.TestCase): ) self.assertEqual(allocator.available_size(), page_size * 16) + def test_compute_owner_alloc_does_not_use_torch_isin_for_page_removal(self): + from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator + + page_size = 64 + cp_size = 4 + page_compute_owners = [0, 2, 0, 2, 3] + allocator = CPSharedPagedTokenToKVPoolAllocator( + logical_size=page_size * 32, + physical_size=page_size * 8, + page_size=page_size, + dtype=torch.bfloat16, + device="cpu", + kvcache=None, + need_sort=False, + cp_size=cp_size, + cp_rank=0, + ) + allocator.free_pages = torch.tensor( + [9, 1, 3, 5, 7, 4] + + [page for page in range(2, 33) if page not in {3, 4, 5, 7, 9}], + dtype=torch.int64, + ) + + with patch( + "torch.isin", + side_effect=AssertionError("compute-owner alloc should not use torch.isin"), + ): + locs = allocator.alloc_extend_compute_owner( + prefix_lens=torch.tensor([0], dtype=torch.int64), + prefix_lens_cpu=torch.tensor([0], dtype=torch.int64), + seq_lens=torch.tensor( + [page_size * len(page_compute_owners)], dtype=torch.int64 + ), + seq_lens_cpu=torch.tensor( + [page_size * len(page_compute_owners)], dtype=torch.int64 + ), + last_loc=torch.tensor([-1], dtype=torch.int64), + extend_num_tokens=page_size * len(page_compute_owners), + page_compute_owners=page_compute_owners, + ) + + self.assertIsNotNone(locs) + logical_pages = locs.view(-1, page_size)[:, 0] // page_size + self.assertEqual(logical_pages.tolist(), [9, 3, 1, 7, 4]) + self.assertEqual( + ((logical_pages - 1) % cp_size).tolist(), + page_compute_owners, + ) + self.assertEqual( + allocator.available_size(), + page_size * (32 - len(page_compute_owners)), + ) + for selected_page in logical_pages.tolist(): + self.assertNotIn(selected_page, allocator.free_pages.tolist()) + + def test_compute_owner_alloc_can_select_release_pages_without_sort_merge(self): + from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator + + page_size = 64 + cp_size = 4 + allocator = CPSharedPagedTokenToKVPoolAllocator( + logical_size=page_size * 16, + physical_size=page_size * 4, + page_size=page_size, + dtype=torch.bfloat16, + device="cpu", + kvcache=None, + need_sort=True, + cp_size=cp_size, + cp_rank=0, + ) + allocator.free_pages = torch.tensor([2, 3, 4], dtype=torch.int64) + allocator.release_pages = torch.tensor([1, 5], dtype=torch.int64) + + with patch( + "torch.sort", + side_effect=AssertionError("compute-owner alloc should not sort-merge"), + ): + locs = allocator.alloc_extend_compute_owner( + prefix_lens=torch.tensor([0], dtype=torch.int64), + prefix_lens_cpu=torch.tensor([0], dtype=torch.int64), + seq_lens=torch.tensor([page_size * 2], dtype=torch.int64), + seq_lens_cpu=torch.tensor([page_size * 2], dtype=torch.int64), + last_loc=torch.tensor([-1], dtype=torch.int64), + extend_num_tokens=page_size * 2, + page_compute_owners=[0, 0], + ) + + self.assertIsNotNone(locs) + logical_pages = locs.view(-1, page_size)[:, 0] // page_size + self.assertEqual(logical_pages.tolist(), [1, 5]) + self.assertEqual(allocator.free_pages.tolist(), [2, 3, 4]) + self.assertEqual(allocator.release_pages.tolist(), []) + + def test_compute_owner_alloc_does_not_evict_lanes_when_first_try_succeeds(self): + from types import SimpleNamespace + + from sglang.srt.mem_cache import common + + page_size = 64 + + class FakeAllocator: + def __init__(self): + self.page_size = page_size + self.cp_size = 4 + self.calls = 0 + + def available_size(self): + return page_size * 100 + + def alloc_extend_compute_owner( + self, + _prefix_lens, + _prefix_lens_cpu, + _seq_lens, + _seq_lens_cpu, + _last_loc, + extend_num_tokens, + _page_compute_owners, + ): + self.calls += 1 + return torch.arange(extend_num_tokens, dtype=torch.int64) + + def alloc_extend(self, *_args, **_kwargs): + raise AssertionError("legacy allocation should not be used") + + class FakeTreeCache: + def __init__(self, allocator): + self.token_to_kv_pool_allocator = allocator + + def is_chunk_cache(self): + return False + + def evict(self, _params): + raise AssertionError("tree eviction should not be used") + + allocator = FakeAllocator() + server_args = SimpleNamespace( + enable_nsa_prefill_cp_shared_kv=True, + enable_nsa_prefill_context_parallel=True, + nsa_prefill_cp_mode="in-seq-split", + ) + + with ( + patch.object(common, "get_global_server_args", return_value=server_args), + patch.object( + common, + "_evict_for_compute_owner_lanes", + side_effect=AssertionError("lane eviction should be lazy"), + ), + ): + out = common.alloc_paged_token_slots_extend( + tree_cache=FakeTreeCache(allocator), + prefix_lens=torch.tensor([0], dtype=torch.int64), + prefix_lens_cpu=torch.tensor([0], dtype=torch.int64), + seq_lens=torch.tensor([page_size * 8], dtype=torch.int64), + seq_lens_cpu=torch.tensor([page_size * 8], dtype=torch.int64), + last_loc=torch.tensor([-1], dtype=torch.int64), + extend_num_tokens=page_size * 8, + ) + + self.assertEqual(allocator.calls, 1) + self.assertEqual(out.numel(), page_size * 8) + def test_compute_owner_lane_eviction_recovers_exhausted_owner_lane(self): from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import EvictResult diff --git a/test/registered/unit/mem_cache/test_req_to_token_pool.py b/test/registered/unit/mem_cache/test_req_to_token_pool.py index 752eb0801..3ac97ea8a 100644 --- a/test/registered/unit/mem_cache/test_req_to_token_pool.py +++ b/test/registered/unit/mem_cache/test_req_to_token_pool.py @@ -7,7 +7,14 @@ from unittest.mock import MagicMock, patch import torch -from sglang.srt.disaggregation.decode import DecodePreallocQueue, DecodeReqToTokenPool +from sglang.srt.disaggregation.decode import ( + DecodePreallocQueue, + DecodeReqToTokenPool, + _kv_locs_to_page_indices_cpu, +) +from sglang.srt.disaggregation.prefill import ( + _kv_locs_to_page_indices_cpu as _prefill_kv_locs_to_page_indices_cpu, +) from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.test.ci.ci_register import register_cpu_ci @@ -162,6 +169,58 @@ class TestDecodePreallocQueue(CustomTestCase): self.assertEqual(allocator_calls[1]["last_loc"].item(), -1) self.assertEqual(allocator_calls[1]["extend_num_tokens"], 5) + def test_kv_locs_to_page_indices_cpu_copies_only_page_starts(self): + kv_locs = torch.tensor( + [ + 128, + 129, + 130, + 131, + 448, + 449, + 450, + 451, + 640, + ], + dtype=torch.int64, + ) + + page_indices = _kv_locs_to_page_indices_cpu(kv_locs, page_size=4) + + self.assertEqual(page_indices.dtype.name, "int32") + self.assertEqual(page_indices.tolist(), [32, 112, 160]) + + def test_kv_locs_to_page_indices_cpu_respects_num_tokens(self): + kv_locs = torch.arange(128, 128 + 16, dtype=torch.int64) + + page_indices = _kv_locs_to_page_indices_cpu( + kv_locs, page_size=4, num_tokens=9 + ) + + self.assertEqual(page_indices.dtype.name, "int32") + self.assertEqual(page_indices.tolist(), [32, 33, 34]) + + def test_prefill_kv_locs_to_page_indices_cpu_copies_only_page_starts(self): + kv_locs = torch.tensor( + [ + 256, + 257, + 258, + 259, + 512, + 513, + 514, + 515, + 768, + ], + dtype=torch.int64, + ) + + page_indices = _prefill_kv_locs_to_page_indices_cpu(kv_locs, page_size=4) + + self.assertEqual(page_indices.dtype.name, "int32") + self.assertEqual(page_indices.tolist(), [64, 128, 192]) + if __name__ == "__main__": unittest.main()