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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user