[mem_cache] Require page-aligned starts in free_segment and drop the boundary trim (#37729)
This commit is contained in:
@@ -27,7 +27,6 @@ from sglang.srt.runtime_context import (
|
||||
get_schedule,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.utils.common import ceil_align
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
@@ -245,25 +244,9 @@ class DecodeKVCacheOffloadManager:
|
||||
if req.kv.req_pool_idx is None or req.kv.req_pool_idx == -1:
|
||||
return
|
||||
|
||||
kv_committed_len = req.effective_kv_committed_len()
|
||||
|
||||
# Prefill-aligned slots are freed only here, at request finish; freeing
|
||||
# them mid-decode races with concurrent admission over live slots.
|
||||
prefill_len = self._prefill_offloaded_len(req)
|
||||
ranges = []
|
||||
if prefill_len > 0:
|
||||
ranges.append((0, prefill_len))
|
||||
# The incremental part of the request (DSA-aware)
|
||||
ranges.append((prefill_len, kv_committed_len))
|
||||
|
||||
# Over-allocated KV cache slots (e.g. from speculative decoding v2).
|
||||
# Without spec v2, start_p == end_p so this contributes nothing.
|
||||
start_p, end_p = kv_committed_len, req.kv.kv_allocated_len
|
||||
if self.page_size > 1:
|
||||
start_p = ceil_align(start_p, self.page_size)
|
||||
if start_p < end_p:
|
||||
ranges.append((start_p, end_p))
|
||||
self.tree_cache.free_kv_row(req.kv, ranges)
|
||||
# Released only at request finish; a mid-decode free races with
|
||||
# concurrent admission over live slots.
|
||||
self.tree_cache.free_kv_row(req.kv, [(0, req.kv.kv_allocated_len)])
|
||||
|
||||
self.req_to_token_pool.free(req)
|
||||
req.kv.mark_kv_released()
|
||||
|
||||
@@ -174,25 +174,32 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
||||
self.free(free_index)
|
||||
|
||||
def free_segment(self, free_index: torch.Tensor, *, start_pos: int):
|
||||
"""Free ``kv_row[start_pos : start_pos + n]`` of one request (or a
|
||||
page-aligned copy); subclasses may use ``start_pos`` to skip the
|
||||
data-dependent dedup. Default: plain free()."""
|
||||
"""Free ``kv_row[start_pos : start_pos + n]`` of one request.
|
||||
|
||||
In page units the segment is ``[start_pos // ps, ceil(end / ps))``:
|
||||
``start_pos`` sits on a page boundary, the end may fall mid-page, and
|
||||
the whole last page is released. Default: plain free()."""
|
||||
assert start_pos % self.page_size == 0, (
|
||||
f"segment start {start_pos} is not page-aligned"
|
||||
)
|
||||
self.free(free_index)
|
||||
|
||||
def free_segments(self, segments):
|
||||
"""Free disjoint ascending ``(free_index, start_pos)`` segments of one
|
||||
request's kv row; a boundary page shared by consecutive segments is
|
||||
emitted once (the later segment's head is trimmed)."""
|
||||
"""Free several ``(free_index, start_pos)`` segments of one request's
|
||||
kv row.
|
||||
|
||||
Each segment covers the pages ``[start_pos // ps, ceil(end / ps))``.
|
||||
Starts sit on page boundaries, ends may fall mid-page, and the page
|
||||
ranges of consecutive segments do not overlap -- so in page units the
|
||||
segments are aligned and disjoint, and every page is released once."""
|
||||
ps = self.page_size
|
||||
prev_end = None
|
||||
for free_index, start_pos in segments:
|
||||
n = free_index.numel()
|
||||
if n == 0:
|
||||
continue
|
||||
seg_end = start_pos + n
|
||||
if prev_end is not None and start_pos // ps == (prev_end - 1) // ps:
|
||||
boundary = (start_pos // ps + 1) * ps
|
||||
free_index = free_index[boundary - start_pos :]
|
||||
start_pos = boundary
|
||||
prev_end = seg_end
|
||||
assert prev_end is None or start_pos // ps > (prev_end - 1) // ps, (
|
||||
f"segment at {start_pos} shares a page with the one ending at {prev_end}"
|
||||
)
|
||||
prev_end = start_pos + n
|
||||
self.free_segment(free_index, start_pos=start_pos)
|
||||
|
||||
@@ -282,36 +282,33 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self._debug_check_no_duplicate_pages()
|
||||
|
||||
def free_segment(self, free_index: torch.Tensor, *, start_pos: int):
|
||||
"""Fixed-shape counterpart of free(): a page's tokens sit consecutively
|
||||
in the kv row, so page representatives are stride slices -- no
|
||||
torch.unique, whose data-dependent output shape forces a device sync.
|
||||
Contract: see base; a page must be freed by only one call per group."""
|
||||
"""Fixed-shape counterpart of free().
|
||||
|
||||
The segment starts on a page boundary and a page's tokens sit
|
||||
consecutively in the kv row, so ``free_index[::page_size]`` is one
|
||||
token from each page the segment covers -- including a partial last
|
||||
page. No torch.unique, whose data-dependent output shape forces a
|
||||
device sync. Contract: see base."""
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
|
||||
ps = self.page_size
|
||||
offset = start_pos % ps
|
||||
if offset == 0:
|
||||
pieces = (free_index[::ps],)
|
||||
else:
|
||||
pieces = (free_index[:1], free_index[ps - offset :: ps])
|
||||
assert start_pos % ps == 0, f"segment start {start_pos} is not page-aligned"
|
||||
reps = free_index[::ps]
|
||||
|
||||
if self.debug_mode:
|
||||
# reference unique on CPU: the NPU subclass deliberately avoids device unique
|
||||
page_ids = torch.cat([p // ps for p in pieces])
|
||||
assert torch.equal(
|
||||
torch.sort(page_ids.cpu())[0],
|
||||
torch.sort(reps.cpu() // ps)[0],
|
||||
torch.unique(free_index.cpu() // ps),
|
||||
)
|
||||
|
||||
if self.free_group is None:
|
||||
self._release_page_ids(*(p // ps for p in pieces))
|
||||
self._release_page_ids(reps // ps)
|
||||
if self.debug_mode:
|
||||
self._debug_check_no_duplicate_pages()
|
||||
else:
|
||||
self.free_page_reps_group.extend(
|
||||
self._copy_for_free_group(piece) for piece in pieces
|
||||
)
|
||||
self.free_page_reps_group.append(self._copy_for_free_group(reps))
|
||||
|
||||
def _debug_check_no_duplicate_pages(self):
|
||||
pages = self.get_all_free_pages()
|
||||
|
||||
@@ -1393,27 +1393,15 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.free_virtual_ids = torch.cat([self.free_virtual_ids, free_v_pages])
|
||||
self._compact_pending(freed_p_pages)
|
||||
|
||||
def _page_reps_pieces(
|
||||
self, free_index: torch.Tensor, start_pos: int
|
||||
) -> Tuple[torch.Tensor, ...]:
|
||||
"""Page-representative TOKEN slices of one kv-row segment.
|
||||
|
||||
Mirrors `PagedTokenToKVPoolAllocator.free_segment`: a page's tokens sit
|
||||
consecutively in the kv row, so with `start_pos` known on the host the
|
||||
representatives are stride slices -- no `torch.unique`, whose
|
||||
data-dependent output shape forces a device sync.
|
||||
|
||||
Exact for any segment shape: a partial head page is the `[:1]` term, a
|
||||
partial tail page the final stride step.
|
||||
"""
|
||||
def _page_reps(self, free_index: torch.Tensor, start_pos: int) -> torch.Tensor:
|
||||
"""One token of every page touched by a page-aligned kv-row segment:
|
||||
the fixed-shape stand-in for `unique(free_index // page_size)`."""
|
||||
ps = self.page_size
|
||||
offset = start_pos % ps
|
||||
if offset == 0:
|
||||
return (free_index[::ps],)
|
||||
return (free_index[:1], free_index[ps - offset :: ps])
|
||||
assert start_pos % ps == 0, f"segment start {start_pos} is not page-aligned"
|
||||
return free_index[::ps]
|
||||
|
||||
def free_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None:
|
||||
"""Fixed-shape counterpart of `free()`; see `_page_reps_pieces`.
|
||||
"""Fixed-shape counterpart of `free()`; see `_page_reps`.
|
||||
|
||||
Contract: see base; a page must be freed by only one call per group.
|
||||
"""
|
||||
@@ -1423,12 +1411,11 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
# token == page: nothing to dedup, the plain path is already exact.
|
||||
self.free(free_index)
|
||||
return
|
||||
pieces = self._page_reps_pieces(free_index.detach().to(torch.int64), start_pos)
|
||||
reps = self._page_reps(free_index.detach().to(torch.int64), start_pos)
|
||||
if self.free_page_reps_group is None:
|
||||
reps = pieces[0] if len(pieces) == 1 else torch.cat(pieces)
|
||||
self.free(reps, _pages=reps // self.page_size)
|
||||
else:
|
||||
self.free_page_reps_group.extend(pieces)
|
||||
self.free_page_reps_group.append(reps)
|
||||
|
||||
def _free_lazy(
|
||||
self, free_index: torch.Tensor, pages: Optional[torch.Tensor] = None
|
||||
@@ -3054,7 +3041,7 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
|
||||
def free_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None:
|
||||
"""Fixed-shape counterpart of `free()`; see
|
||||
`MultiEndedAllocator._page_reps_pieces`. The mamba sub-pool is
|
||||
`MultiEndedAllocator._page_reps`. The mamba sub-pool is
|
||||
slot-granular and untouched by a token free, so only the full side
|
||||
needs the representatives.
|
||||
"""
|
||||
@@ -3063,13 +3050,13 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if self.page_size == 1:
|
||||
self.free(free_index)
|
||||
return
|
||||
pieces = self.full_attn_allocator._page_reps_pieces(
|
||||
reps = self.full_attn_allocator._page_reps(
|
||||
free_index.detach().to(torch.int64), start_pos
|
||||
)
|
||||
if self.free_page_reps_group is None:
|
||||
self._release_page_reps(pieces)
|
||||
self._release_page_reps((reps,))
|
||||
else:
|
||||
self.free_page_reps_group.extend(pieces)
|
||||
self.free_page_reps_group.append(reps)
|
||||
|
||||
def _release_page_reps(self, pieces: Sequence[torch.Tensor]) -> None:
|
||||
reps = pieces[0] if len(pieces) == 1 else torch.cat(tuple(pieces))
|
||||
@@ -3637,8 +3624,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
v = free_index.detach().to(torch.int64)
|
||||
ps = self.page_size
|
||||
if start_pos is not None and ps > 1:
|
||||
pieces = self.swa_attn_allocator._page_reps_pieces(v, start_pos)
|
||||
reps = pieces[0] if len(pieces) == 1 else torch.cat(pieces)
|
||||
reps = self.swa_attn_allocator._page_reps(v, start_pos)
|
||||
# Keep only pages still bound on swa (freeing a tombstoned one
|
||||
# would corrupt the hole list). `> 0` strict: -1 = tombstoned,
|
||||
# page 0 = padding sink (never freeable).
|
||||
@@ -3702,7 +3688,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
|
||||
def free_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None:
|
||||
"""Fixed-shape counterpart of `free()`; see
|
||||
`MultiEndedAllocator._page_reps_pieces`. Both sides share one
|
||||
`MultiEndedAllocator._page_reps`. Both sides share one
|
||||
derivation -- neither repeats the position-less dedup.
|
||||
"""
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
@@ -3710,13 +3696,13 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
if self.page_size == 1:
|
||||
self.free(free_index)
|
||||
return
|
||||
pieces = self.full_attn_allocator._page_reps_pieces(
|
||||
reps = self.full_attn_allocator._page_reps(
|
||||
free_index.detach().to(torch.int64), start_pos
|
||||
)
|
||||
if self.free_page_reps_group is None:
|
||||
self._release_page_reps(pieces)
|
||||
self._release_page_reps((reps,))
|
||||
else:
|
||||
self.free_page_reps_group.extend(pieces)
|
||||
self.free_page_reps_group.append(reps)
|
||||
|
||||
def _release_page_reps(self, pieces: Sequence[torch.Tensor]) -> None:
|
||||
reps = pieces[0] if len(pieces) == 1 else torch.cat(tuple(pieces))
|
||||
|
||||
Reference in New Issue
Block a user