[mem_cache] Require page-aligned starts in free_segment and drop the boundary trim (#37729)

This commit is contained in:
Liangsheng Yin
2026-09-03 13:28:33 -07:00
committed by GitHub
parent 3ffacf949b
commit 2a980cbf10
7 changed files with 113 additions and 168 deletions
@@ -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()
+19 -12
View File
@@ -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)
+12 -15
View File
@@ -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))