Reduce CP shared-KV request-boundary stalls
CP shared KV now avoids the PyTorch sort/search remap for the single-request current-only path by deriving compact rows from page-level inverse mapping. The same change keeps sort NVTX attribution gated and splits high-frequency MoE sort markers behind a separate env var so profiling does not perturb normal runs. Decode-side disaggregation prealloc also avoids rebuilding large token index tensors and records finer allocation timing, while compute-owner allocation/free tests cover the shared-KV page-lane behavior. Constraint: The runtime tree used for validation is the remote /sgl-workspace/sglang-tai mount, which is not itself a Git repository, so these tracked files were synchronized into the local repo before commit. Rejected: Keep torch.sort/searchsorted for current remap | it emits ATen/CCCL radixSortKVInPlace kernels in the attention hot path. Rejected: Enable MoE sort NVTX under the generic sort env | the MoE preprocess sort is too frequent and can make profiling look like a hang. Confidence: medium Scope-risk: moderate Directive: Do not reintroduce token-level torch.sort/searchsorted in CP shared-KV current remap without profiling the attention hot path under Nsight. Tested: Remote container py_compile for modified runtime files; git diff --cached --check. Not-tested: Full multi-node GLM5 PD throughput/profile rerun after the page-inverse current remap.
This commit is contained in:
@@ -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。**
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user