Fix buffer-mode HiCache load-back ownership races; add optional prefetch anchor lock (#35769)
Signed-off-by: Zhiqiang Xie <zqx@meta.com>
This commit is contained in:
@@ -684,6 +684,10 @@ class Envs:
|
|||||||
SGLANG_HICACHE_FILE_BACKEND_ENABLE_METADATA_CACHE = EnvBool(False)
|
SGLANG_HICACHE_FILE_BACKEND_ENABLE_METADATA_CACHE = EnvBool(False)
|
||||||
# Positive cache TTL for filesystem metadata lookups (-1 disables positive expiration)
|
# Positive cache TTL for filesystem metadata lookups (-1 disables positive expiration)
|
||||||
SGLANG_HICACHE_FILE_BACKEND_METADATA_TTL = EnvFloat(5.0)
|
SGLANG_HICACHE_FILE_BACKEND_METADATA_TTL = EnvFloat(5.0)
|
||||||
|
# Buffer mode: pin a staged prefetch's device anchor from IO commit to
|
||||||
|
# consumption so eviction cannot waste the fetch; cap = fraction of pool.
|
||||||
|
SGLANG_ENABLE_HICACHE_BUFFER_ANCHOR_LOCK = EnvBool(False)
|
||||||
|
SGLANG_HICACHE_BUFFER_ANCHOR_LOCK_CAP = EnvFloat(0.5)
|
||||||
SGLANG_HICACHE_NIXL_BACKEND_STORAGE_DIR = EnvStr(None)
|
SGLANG_HICACHE_NIXL_BACKEND_STORAGE_DIR = EnvStr(None)
|
||||||
# Enable O_DIRECT when opening NIXL POSIX backend files (bypasses OS page cache).
|
# Enable O_DIRECT when opening NIXL POSIX backend files (bypasses OS page cache).
|
||||||
# Disable with SGLANG_HICACHE_NIXL_USE_DIRECT_IO=0 or via the
|
# Disable with SGLANG_HICACHE_NIXL_USE_DIRECT_IO=0 or via the
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import msgspec
|
import msgspec
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.cache_controller import HICACHE_WRITE_STAGING_POOL_FRACTION
|
from sglang.srt.managers.cache_controller import HICACHE_WRITE_STAGING_POOL_FRACTION
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
DecLockRefParams,
|
DecLockRefParams,
|
||||||
@@ -120,6 +121,14 @@ class _OngoingBufferLoadBack(msgspec.Struct):
|
|||||||
hash_values: list[str]
|
hash_values: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
class _AnchorLock(msgspec.Struct):
|
||||||
|
"""Pins a staged prefetch's device anchor from IO commit to consumption."""
|
||||||
|
|
||||||
|
node_id: NodeId
|
||||||
|
lock_params: DecLockRefParams
|
||||||
|
tokens: int
|
||||||
|
|
||||||
|
|
||||||
def _track_content_refs(refs: dict[str, int], hash_values: list[str]) -> None:
|
def _track_content_refs(refs: dict[str, int], hash_values: list[str]) -> None:
|
||||||
"""Add one content ref per page hash (at D2H launch). Refcounted,
|
"""Add one content ref per page hash (at D2H launch). Refcounted,
|
||||||
not a flag: several launched entries can carry the same content
|
not a flag: several launched entries can carry the same content
|
||||||
@@ -199,6 +208,15 @@ class BufferModePipeline:
|
|||||||
# Metadata-only pending-write backlog cap; beyond it new intents
|
# Metadata-only pending-write backlog cap; beyond it new intents
|
||||||
# are dropped at admission (re-trigger on a later hit).
|
# are dropped at admission (re-trigger on a later hit).
|
||||||
self.write_backlog_cap = write_backlog_cap
|
self.write_backlog_cap = write_backlog_cap
|
||||||
|
# Anchor-lock knobs; the cap keeps queued holds from pinning the pool.
|
||||||
|
self.anchor_lock_enabled = envs.SGLANG_ENABLE_HICACHE_BUFFER_ANCHOR_LOCK.get()
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
|
|
||||||
|
kvcache = cache.token_to_kv_pool_allocator.get_kvcache()
|
||||||
|
full_pool = kvcache.full_kv_pool if isinstance(kvcache, SWAKVPool) else kvcache
|
||||||
|
self.anchor_lock_cap_tokens = int(
|
||||||
|
envs.SGLANG_HICACHE_BUFFER_ANCHOR_LOCK_CAP.get() * full_pool.size
|
||||||
|
)
|
||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
@@ -226,6 +244,10 @@ class BufferModePipeline:
|
|||||||
self.write_staged_tokens_ = 0
|
self.write_staged_tokens_ = 0
|
||||||
self.write_backlog_tokens_ = 0
|
self.write_backlog_tokens_ = 0
|
||||||
self._backlog_cap_hits = 0
|
self._backlog_cap_hits = 0
|
||||||
|
# rid-keyed anchor locks; released idempotently at every exit.
|
||||||
|
self.anchor_locks: dict[str, _AnchorLock] = {}
|
||||||
|
self.anchor_locked_tokens_ = 0
|
||||||
|
self._anchor_lock_cap_skips = 0
|
||||||
|
|
||||||
def is_idle(self) -> bool:
|
def is_idle(self) -> bool:
|
||||||
"""No queued writes, staged prefetches, or storage writes in flight
|
"""No queued writes, staged prefetches, or storage writes in flight
|
||||||
@@ -586,6 +608,61 @@ class BufferModePipeline:
|
|||||||
|
|
||||||
# ---- load back pipeline (storage -> staging -> device) ----
|
# ---- load back pipeline (storage -> staging -> device) ----
|
||||||
|
|
||||||
|
def try_lock_anchor(self, req_id: str, anchor_node_id: NodeId) -> None:
|
||||||
|
"""Pin the device anchor at IO commit so eviction cannot invalidate
|
||||||
|
the splice; drift resolves at consumption, and the cap keeps queued
|
||||||
|
holds from making the pool unevictable (over-cap launches unlocked)."""
|
||||||
|
if not self.anchor_lock_enabled or req_id in self.anchor_locks:
|
||||||
|
return
|
||||||
|
prefix_tokens = self._prefetch_prefix_ctx.get(req_id)
|
||||||
|
if not prefix_tokens:
|
||||||
|
return # root anchor: nothing to pin
|
||||||
|
matched_len = len(prefix_tokens)
|
||||||
|
if self.anchor_locked_tokens_ + matched_len > self.anchor_lock_cap_tokens:
|
||||||
|
self._anchor_lock_cap_skips += 1
|
||||||
|
if (
|
||||||
|
self._anchor_lock_cap_skips <= 3
|
||||||
|
or self._anchor_lock_cap_skips % 1000 == 0
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"HiCache anchor-lock cap reached (skip %d): locked=%d "
|
||||||
|
"want=%d cap=%d; launching unlocked.",
|
||||||
|
self._anchor_lock_cap_skips,
|
||||||
|
self.anchor_locked_tokens_,
|
||||||
|
matched_len,
|
||||||
|
self.anchor_lock_cap_tokens,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
cache = self._cache
|
||||||
|
try:
|
||||||
|
node = cache.tree_core.node_by_id(anchor_node_id)
|
||||||
|
except KeyError:
|
||||||
|
return # anchor deleted; fetch unlocked
|
||||||
|
if node.component_data[BASE_COMPONENT_TYPE].value is None:
|
||||||
|
# Evicted since enqueue; fetch unlocked.
|
||||||
|
logger.warning("HiCache anchor evicted before IO commit req=%s", req_id)
|
||||||
|
return
|
||||||
|
lock_params = cache.inc_lock_ref(anchor_node_id).to_dec_params()
|
||||||
|
self.anchor_locks[req_id] = _AnchorLock(
|
||||||
|
node_id=anchor_node_id,
|
||||||
|
lock_params=lock_params,
|
||||||
|
tokens=matched_len,
|
||||||
|
)
|
||||||
|
self.anchor_locked_tokens_ += matched_len
|
||||||
|
|
||||||
|
def release_anchor_lock(self, req_id: str) -> None:
|
||||||
|
"""Drop a staged prefetch's anchor lock (idempotent; called at every
|
||||||
|
consume/drop/abort exit)."""
|
||||||
|
lock = self.anchor_locks.pop(req_id, None)
|
||||||
|
if lock is None:
|
||||||
|
return
|
||||||
|
self._cache.dec_lock_ref(lock.node_id, lock.lock_params)
|
||||||
|
self.anchor_locked_tokens_ -= lock.tokens
|
||||||
|
assert self.anchor_locked_tokens_ >= 0, (
|
||||||
|
f"anchor-lock accounting corrupted: locked={self.anchor_locked_tokens_} "
|
||||||
|
f"after releasing {req_id}"
|
||||||
|
)
|
||||||
|
|
||||||
def set_prefix_ctx(self, req_id: str, matched_prefix_tokens) -> None:
|
def set_prefix_ctx(self, req_id: str, matched_prefix_tokens) -> None:
|
||||||
"""Record the device-matched prefix at prefetch enqueue; consumed at
|
"""Record the device-matched prefix at prefetch enqueue; consumed at
|
||||||
staging commit to build the full-span tree key."""
|
staging commit to build the full-span tree key."""
|
||||||
@@ -627,6 +704,7 @@ class BufferModePipeline:
|
|||||||
|
|
||||||
if num_tokens == 0 or prefix_tokens is None:
|
if num_tokens == 0 or prefix_tokens is None:
|
||||||
# Nothing usable fetched: recompute.
|
# Nothing usable fetched: recompute.
|
||||||
|
self.release_anchor_lock(req_id)
|
||||||
cc.append_host_mem_release(
|
cc.append_host_mem_release(
|
||||||
host_indices[:num_tokens], extra_pools=aux_xfers or None
|
host_indices[:num_tokens], extra_pools=aux_xfers or None
|
||||||
)
|
)
|
||||||
@@ -678,10 +756,14 @@ class BufferModePipeline:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_load_back(self, params: InitLoadBackParams) -> tuple[torch.Tensor, NodeId]:
|
def init_load_back(self, params: InitLoadBackParams) -> tuple[torch.Tensor, NodeId]:
|
||||||
"""Buffer-mode branch of init_load_back: consume the staged prefetch
|
"""Consume the staged prefetch at prefill admission: device alloc,
|
||||||
at prefill admission — device alloc (evict-before-alloc), layer-gated
|
layer-gated H2D, and a plain insert so downstream sees ordinary tree
|
||||||
H2D, and a plain insert so downstream sees ordinary tree state.
|
state; invalid holds drop and the request recomputes.
|
||||||
Misaligned or alloc-failed holds drop; the request recomputes."""
|
|
||||||
|
Ownership contract: cc.load queues the H2D before insert adjudicates
|
||||||
|
ownership, so the live pre-checks below must prove the insert can
|
||||||
|
only ADD nodes — a dedup would free slots the in-flight copy still
|
||||||
|
targets (queued use-after-free)."""
|
||||||
cache = self._cache
|
cache = self._cache
|
||||||
req = params.req
|
req = params.req
|
||||||
assert req is not None
|
assert req is not None
|
||||||
@@ -689,10 +771,12 @@ class BufferModePipeline:
|
|||||||
unchanged = (empty, req.last_node)
|
unchanged = (empty, req.last_node)
|
||||||
f = self.staged_prefetches.pop(req.rid, None)
|
f = self.staged_prefetches.pop(req.rid, None)
|
||||||
if f is None:
|
if f is None:
|
||||||
|
self.release_anchor_lock(req.rid)
|
||||||
return unchanged
|
return unchanged
|
||||||
cc = cache.cache_controller
|
cc = cache.cache_controller
|
||||||
|
|
||||||
def _drop() -> tuple[torch.Tensor, NodeId]:
|
def _drop() -> tuple[torch.Tensor, NodeId]:
|
||||||
|
self.release_anchor_lock(req.rid)
|
||||||
self._free_staging_now(f.host_indices, f.aux_xfers)
|
self._free_staging_now(f.host_indices, f.aux_xfers)
|
||||||
cc.prefetch_tokens_occupied -= f.occupied_tokens
|
cc.prefetch_tokens_occupied -= f.occupied_tokens
|
||||||
return unchanged
|
return unchanged
|
||||||
@@ -700,8 +784,44 @@ class BufferModePipeline:
|
|||||||
# Splice-validity: the span only fits if the device prefix still
|
# Splice-validity: the span only fits if the device prefix still
|
||||||
# ends exactly at the enqueue-time matched_len.
|
# ends exactly at the enqueue-time matched_len.
|
||||||
if len(req.prefix_indices) != f.matched_len:
|
if len(req.prefix_indices) != f.matched_len:
|
||||||
# Prefix moved while held (leaf eviction or sibling extension):
|
logger.warning(
|
||||||
# drop and recompute.
|
"HiCache staged prefetch dropped req=%s reason=%s matched=%d "
|
||||||
|
"now=%d tokens_wasted=%d locked=%s",
|
||||||
|
req.rid,
|
||||||
|
"growth" if len(req.prefix_indices) > f.matched_len else "shrink",
|
||||||
|
f.matched_len,
|
||||||
|
len(req.prefix_indices),
|
||||||
|
f.num_tokens,
|
||||||
|
req.rid in self.anchor_locks,
|
||||||
|
)
|
||||||
|
return _drop()
|
||||||
|
|
||||||
|
key = RadixKey(
|
||||||
|
array("q", f.key_tokens),
|
||||||
|
extra_key=f.extra_key,
|
||||||
|
is_bigram=cache.tree_core.is_eagle,
|
||||||
|
).page_aligned(cache.page_size)
|
||||||
|
span_end = f.matched_len + f.num_tokens
|
||||||
|
|
||||||
|
# Live ownership pre-check: the unified length detects anchor drift,
|
||||||
|
# full_kv_hit_length detects FULL overlap the insert would dedup-free
|
||||||
|
# (an SWA tombstone can mask live FULL from the unified match alone).
|
||||||
|
live = cache.match_prefix(MatchPrefixParams(key=key))
|
||||||
|
if (
|
||||||
|
len(live.device_indices) != f.matched_len
|
||||||
|
or live.full_kv_hit_length != f.matched_len
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"HiCache staged prefetch dropped req=%s reason=overlap "
|
||||||
|
"matched=%d live_unified=%d live_full=%d tokens_wasted=%d "
|
||||||
|
"locked=%s",
|
||||||
|
req.rid,
|
||||||
|
f.matched_len,
|
||||||
|
len(live.device_indices),
|
||||||
|
live.full_kv_hit_length,
|
||||||
|
f.num_tokens,
|
||||||
|
req.rid in self.anchor_locks,
|
||||||
|
)
|
||||||
return _drop()
|
return _drop()
|
||||||
|
|
||||||
# Evict-before-alloc (mirrors _load_back_transfers): the budget gate
|
# Evict-before-alloc (mirrors _load_back_transfers): the budget gate
|
||||||
@@ -748,13 +868,7 @@ class BufferModePipeline:
|
|||||||
|
|
||||||
# Publish via a plain insert under the admission lock choreography;
|
# Publish via a plain insert under the admission lock choreography;
|
||||||
# the caller's request lock then pins the span (load_back pattern).
|
# the caller's request lock then pins the span (load_back pattern).
|
||||||
key = RadixKey(
|
insert_result = cache.insert(
|
||||||
array("q", f.key_tokens),
|
|
||||||
extra_key=f.extra_key,
|
|
||||||
is_bigram=cache.tree_core.is_eagle,
|
|
||||||
).page_aligned(cache.page_size)
|
|
||||||
span_end = f.matched_len + f.num_tokens
|
|
||||||
cache.insert(
|
|
||||||
InsertParams(
|
InsertParams(
|
||||||
key=key,
|
key=key,
|
||||||
value=torch.cat([req.prefix_indices, device_indices]),
|
value=torch.cat([req.prefix_indices, device_indices]),
|
||||||
@@ -773,12 +887,24 @@ class BufferModePipeline:
|
|||||||
hash_values=f.hash_values,
|
hash_values=f.hash_values,
|
||||||
)
|
)
|
||||||
m = cache.match_prefix(MatchPrefixParams(key=key))
|
m = cache.match_prefix(MatchPrefixParams(key=key))
|
||||||
if len(m.device_indices) < span_end:
|
self.release_anchor_lock(req.rid)
|
||||||
# The insert walk did not adopt the full span (should not happen
|
canonical = m.device_indices[f.matched_len : span_end]
|
||||||
# for a locked prefix); the slots are tree-owned/evictable — do
|
if len(m.device_indices) < span_end or not torch.equal(
|
||||||
# not splice, the request recomputes.
|
canonical, device_indices
|
||||||
return unchanged
|
):
|
||||||
return device_indices, m.last_device_node
|
# Fail-stop: the insert freed or replaced slots the in-flight H2D
|
||||||
|
# still targets; continuing risks silent KV corruption.
|
||||||
|
raise RuntimeError(
|
||||||
|
f"HiCache buffer load-back ownership violation req={f.req_id}: "
|
||||||
|
f"insert prefix_len={insert_result.prefix_len} "
|
||||||
|
f"expected={f.matched_len}, adopted={len(m.device_indices)} "
|
||||||
|
f"span_end={span_end}, canonical_matches_incoming="
|
||||||
|
f"{len(m.device_indices) >= span_end and torch.equal(canonical, device_indices)}; "
|
||||||
|
f"in-flight H2D targets freed slots"
|
||||||
|
)
|
||||||
|
# Canonical ownership: return the post-insert tree slice, never the
|
||||||
|
# raw cc.load allocation (torch.equal here; the tree slice is truth).
|
||||||
|
return canonical, m.last_device_node
|
||||||
|
|
||||||
def try_finish_load_back(self, ack_id: int) -> bool:
|
def try_finish_load_back(self, ack_id: int) -> bool:
|
||||||
"""Fill ack: free the host bounce and return True when the ack id is
|
"""Fill ack: free the host bounce and return True when the ack id is
|
||||||
@@ -796,10 +922,11 @@ class BufferModePipeline:
|
|||||||
|
|
||||||
cc.prefetch_tokens_occupied -= f.occupied_tokens
|
cc.prefetch_tokens_occupied -= f.occupied_tokens
|
||||||
logger.info(
|
logger.info(
|
||||||
"HiCache prefetch fill committed req=%s filled=%d occupied=%d",
|
"HiCache prefetch fill committed req=%s filled=%d occupied=%d locked=%d",
|
||||||
f.req_id,
|
f.req_id,
|
||||||
f.num_tokens,
|
f.num_tokens,
|
||||||
cc.prefetch_tokens_occupied,
|
cc.prefetch_tokens_occupied,
|
||||||
|
self.anchor_locked_tokens_,
|
||||||
)
|
)
|
||||||
if cache.enable_storage_metrics and cache.storage_metrics_collector is not None:
|
if cache.enable_storage_metrics and cache.storage_metrics_collector is not None:
|
||||||
cache.storage_metrics_collector.log_prefetched_tokens(f.num_tokens)
|
cache.storage_metrics_collector.log_prefetched_tokens(f.num_tokens)
|
||||||
@@ -808,6 +935,7 @@ class BufferModePipeline:
|
|||||||
def release_aborted_staged(self, rid: str) -> bool:
|
def release_aborted_staged(self, rid: str) -> bool:
|
||||||
"""Free an aborted request's staged prefetch (nothing device-side
|
"""Free an aborted request's staged prefetch (nothing device-side
|
||||||
exists yet — only the bounce). Returns True when a hold existed."""
|
exists yet — only the bounce). Returns True when a hold existed."""
|
||||||
|
self.release_anchor_lock(rid)
|
||||||
staged = self.staged_prefetches.pop(rid, None)
|
staged = self.staged_prefetches.pop(rid, None)
|
||||||
if staged is None:
|
if staged is None:
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -1932,6 +1932,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
del self.ongoing_prefetch[req_id]
|
del self.ongoing_prefetch[req_id]
|
||||||
if self.buffer_pipeline is not None:
|
if self.buffer_pipeline is not None:
|
||||||
self.buffer_pipeline.pop_prefix_ctx(req_id)
|
self.buffer_pipeline.pop_prefix_ctx(req_id)
|
||||||
|
self.buffer_pipeline.release_anchor_lock(req_id)
|
||||||
self.cache_controller.prefetch_tokens_occupied -= (
|
self.cache_controller.prefetch_tokens_occupied -= (
|
||||||
self._prefetch_occupied_span(prefetch_key, host_indices)
|
self._prefetch_occupied_span(prefetch_key, host_indices)
|
||||||
)
|
)
|
||||||
@@ -1999,6 +2000,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
del self.ongoing_prefetch[rid]
|
del self.ongoing_prefetch[rid]
|
||||||
if self.buffer_pipeline is not None:
|
if self.buffer_pipeline is not None:
|
||||||
self.buffer_pipeline.pop_prefix_ctx(rid)
|
self.buffer_pipeline.pop_prefix_ctx(rid)
|
||||||
|
self.buffer_pipeline.release_anchor_lock(rid)
|
||||||
self.cache_controller.append_host_mem_release(
|
self.cache_controller.append_host_mem_release(
|
||||||
host_indices=host_indices[:completed_tokens],
|
host_indices=host_indices[:completed_tokens],
|
||||||
extra_pools=[x for xfers in comp_xfers.values() for x in xfers],
|
extra_pools=[x for xfers in comp_xfers.values() for x in xfers],
|
||||||
@@ -2077,6 +2079,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
self._invalidate_absent_from_hit_query(operation)
|
self._invalidate_absent_from_hit_query(operation)
|
||||||
if self.buffer_pipeline is not None:
|
if self.buffer_pipeline is not None:
|
||||||
self.buffer_pipeline.pop_prefix_ctx(req_id)
|
self.buffer_pipeline.pop_prefix_ctx(req_id)
|
||||||
|
self.buffer_pipeline.release_anchor_lock(req_id)
|
||||||
cc = self.cache_controller
|
cc = self.cache_controller
|
||||||
cc.append_host_mem_release(
|
cc.append_host_mem_release(
|
||||||
extra_pools=[x for xfers in comp_xfers.values() for x in xfers]
|
extra_pools=[x for xfers in comp_xfers.values() for x in xfers]
|
||||||
@@ -2159,6 +2162,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
self.ongoing_prefetch[req_id] = info._replace(host_indices=host_indices)
|
self.ongoing_prefetch[req_id] = info._replace(host_indices=host_indices)
|
||||||
if buffer_mode:
|
if buffer_mode:
|
||||||
cc.prefetch_tokens_occupied += alloc_len
|
cc.prefetch_tokens_occupied += alloc_len
|
||||||
|
# IO commit: pin the anchor until consumption. Do not read
|
||||||
|
# attributes off `operation` here — alternative cache
|
||||||
|
# controllers may expose a narrower surface.
|
||||||
|
self.buffer_pipeline.try_lock_anchor(req_id, info.anchor_node_id)
|
||||||
cc.prefetch_buffer.put(operation)
|
cc.prefetch_buffer.put(operation)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|||||||
@@ -3026,6 +3026,171 @@ class UnifiedRadixCacheSuite:
|
|||||||
# Rejection must come from the surfaced window charge.
|
# Rejection must come from the surfaced window charge.
|
||||||
self.assertGreaterEqual(surfaced_swa_hit, window)
|
self.assertGreaterEqual(surfaced_swa_hit, window)
|
||||||
|
|
||||||
|
def test_buffer_only_load_back_drops_on_sibling_published_span(self):
|
||||||
|
"""Queued-UAF regression: a sibling publishes the staged span between
|
||||||
|
staging and consumption; the live pre-check must drop the hold before
|
||||||
|
any device allocation or H2D exists."""
|
||||||
|
self._skip_unsupported_hicache_test()
|
||||||
|
storage_dir = tempfile.mkdtemp()
|
||||||
|
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||||
|
|
||||||
|
seq = self._buffer_swa_seq()
|
||||||
|
self._produce_buffer_l3(storage_dir, seq)
|
||||||
|
|
||||||
|
cons, cons_alloc, cons_rtp = build_fixture(self.cfg)
|
||||||
|
self._init_buffer_hicache(cons, storage_dir)
|
||||||
|
avail0 = self._host_avail_sizes(cons)
|
||||||
|
|
||||||
|
req_id = "sibling-publish"
|
||||||
|
cons.prefetch_from_storage(
|
||||||
|
req_id, cons.root_node.id, array("q", seq), None, None
|
||||||
|
)
|
||||||
|
self._pump_hicache_until(
|
||||||
|
cons,
|
||||||
|
lambda: cons.check_prefetch_progress(req_id)
|
||||||
|
and cons.buffer_pipeline.has_staged(req_id),
|
||||||
|
"prefetch did not stage",
|
||||||
|
)
|
||||||
|
cons.pop_prefetch_loaded_tokens(req_id)
|
||||||
|
|
||||||
|
# Sibling publishes the identical span (live FULL + SWA).
|
||||||
|
self._insert(cons, cons_alloc, cons_rtp, seq)
|
||||||
|
sib = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||||
|
self.assertEqual(len(sib.device_indices), len(seq))
|
||||||
|
self._fill_full_kv(cons_alloc, sib.device_indices, marker=3)
|
||||||
|
sib_kv = self._snapshot_full_kv(cons_alloc, sib.device_indices)
|
||||||
|
dev_avail0 = cons.token_to_kv_pool_allocator.available_size()
|
||||||
|
|
||||||
|
# Consume with the batch-stale empty prefix view: the live unified
|
||||||
|
# check must drop it.
|
||||||
|
spliced = self._consume_staged_prefetch(cons, req_id, prefix_len=0)
|
||||||
|
self.assertEqual(int(spliced.numel()), 0)
|
||||||
|
|
||||||
|
# Nothing device-side happened; sibling slots and staging intact.
|
||||||
|
self.assertEqual(cons.token_to_kv_pool_allocator.available_size(), dev_avail0)
|
||||||
|
self.assertEqual(cons.buffer_pipeline.ongoing_buffer_load_back, {})
|
||||||
|
self.assertFalse(cons.buffer_pipeline.has_staged(req_id))
|
||||||
|
self.assertEqual(cons.cache_controller.prefetch_tokens_occupied, 0)
|
||||||
|
self.assertEqual(self._host_avail_sizes(cons), avail0)
|
||||||
|
m = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||||
|
self.assertTrue(torch.equal(m.device_indices, sib.device_indices))
|
||||||
|
k, v = self._snapshot_full_kv(cons_alloc, m.device_indices)
|
||||||
|
self.assertTrue(torch.equal(k, sib_kv[0]))
|
||||||
|
self.assertTrue(torch.equal(v, sib_kv[1]))
|
||||||
|
cons.sanity_check()
|
||||||
|
|
||||||
|
def test_buffer_only_load_back_drops_on_full_overlap_masked_by_swa_tombstone(
|
||||||
|
self,
|
||||||
|
):
|
||||||
|
"""Queued-UAF regression: live FULL under an SWA tombstone is invisible
|
||||||
|
to the unified match but still dedup-freed by insert; only the
|
||||||
|
full_kv_hit_length pre-check can drop the hold."""
|
||||||
|
self._skip_unsupported_hicache_test()
|
||||||
|
if not self.cfg.has_swa:
|
||||||
|
self.skipTest("masked overlap requires an SWA component")
|
||||||
|
storage_dir = tempfile.mkdtemp()
|
||||||
|
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||||
|
|
||||||
|
seq = self._buffer_swa_seq()
|
||||||
|
self._produce_buffer_l3(storage_dir, seq)
|
||||||
|
|
||||||
|
cons, cons_alloc, cons_rtp = build_fixture(self.cfg)
|
||||||
|
self._init_buffer_hicache(cons, storage_dir)
|
||||||
|
|
||||||
|
# Masked state: nodes born with live FULL under SWA tombstones
|
||||||
|
# (sibling insert whose SWA ring had slid past the span).
|
||||||
|
value = self._alloc(cons_alloc, len(seq))
|
||||||
|
cons.insert(
|
||||||
|
InsertParams(
|
||||||
|
key=RadixKey(array("q", seq)),
|
||||||
|
value=value[: len(seq)],
|
||||||
|
swa_evicted_seqlen=len(seq),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
masked = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||||
|
self.assertEqual(len(masked.device_indices), 0, "unified match not masked")
|
||||||
|
self.assertEqual(masked.full_kv_hit_length, len(seq), "live FULL not resident")
|
||||||
|
|
||||||
|
avail0 = self._host_avail_sizes(cons)
|
||||||
|
req_id = "masked-overlap"
|
||||||
|
cons.prefetch_from_storage(
|
||||||
|
req_id, cons.root_node.id, array("q", seq), None, None
|
||||||
|
)
|
||||||
|
self._pump_hicache_until(
|
||||||
|
cons,
|
||||||
|
lambda: cons.check_prefetch_progress(req_id)
|
||||||
|
and cons.buffer_pipeline.has_staged(req_id),
|
||||||
|
"prefetch did not stage",
|
||||||
|
)
|
||||||
|
cons.pop_prefetch_loaded_tokens(req_id)
|
||||||
|
dev_avail0 = cons.token_to_kv_pool_allocator.available_size()
|
||||||
|
|
||||||
|
# Every unified-length guard passes at 0 == 0; only the Full-only
|
||||||
|
# pre-check drops.
|
||||||
|
spliced = self._consume_staged_prefetch(cons, req_id, prefix_len=0)
|
||||||
|
self.assertEqual(int(spliced.numel()), 0)
|
||||||
|
|
||||||
|
self.assertEqual(cons.token_to_kv_pool_allocator.available_size(), dev_avail0)
|
||||||
|
self.assertEqual(cons.buffer_pipeline.ongoing_buffer_load_back, {})
|
||||||
|
self.assertEqual(cons.cache_controller.prefetch_tokens_occupied, 0)
|
||||||
|
self.assertEqual(self._host_avail_sizes(cons), avail0)
|
||||||
|
# The masked FULL is still intact (nothing dedup-freed it).
|
||||||
|
after = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||||
|
self.assertEqual(after.full_kv_hit_length, len(seq))
|
||||||
|
cons.sanity_check()
|
||||||
|
|
||||||
|
def test_buffer_only_load_back_fail_stops_on_post_check_overlap(self):
|
||||||
|
"""If the tree mutates between the pre-check and the insert (simulated
|
||||||
|
via cc.load), consumption must fail-stop rather than hand out slots a
|
||||||
|
queued H2D no longer owns."""
|
||||||
|
self._skip_unsupported_hicache_test()
|
||||||
|
storage_dir = tempfile.mkdtemp()
|
||||||
|
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||||
|
|
||||||
|
seq = self._buffer_swa_seq()
|
||||||
|
self._produce_buffer_l3(storage_dir, seq)
|
||||||
|
|
||||||
|
cons, cons_alloc, cons_rtp = build_fixture(self.cfg)
|
||||||
|
self._init_buffer_hicache(cons, storage_dir)
|
||||||
|
|
||||||
|
req_id = "post-check-overlap"
|
||||||
|
cons.prefetch_from_storage(
|
||||||
|
req_id, cons.root_node.id, array("q", seq), None, None
|
||||||
|
)
|
||||||
|
self._pump_hicache_until(
|
||||||
|
cons,
|
||||||
|
lambda: cons.check_prefetch_progress(req_id)
|
||||||
|
and cons.buffer_pipeline.has_staged(req_id),
|
||||||
|
"prefetch did not stage",
|
||||||
|
)
|
||||||
|
cons.pop_prefetch_loaded_tokens(req_id)
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import InitLoadBackParams
|
||||||
|
|
||||||
|
real_load = cons.cache_controller.load
|
||||||
|
|
||||||
|
def adversarial_load(*args, **kwargs):
|
||||||
|
# Mutate the tree after the pre-check has already passed.
|
||||||
|
self._insert(cons, cons_alloc, cons_rtp, seq)
|
||||||
|
return real_load(*args, **kwargs)
|
||||||
|
|
||||||
|
f = cons.buffer_pipeline.staged_prefetches[req_id]
|
||||||
|
req = mock.Mock()
|
||||||
|
req.rid = req_id
|
||||||
|
req.prefix_indices = torch.zeros(
|
||||||
|
0,
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=cons.tree_core.empty_match_result.device_indices.device,
|
||||||
|
)
|
||||||
|
req.last_node = cons.root_node.id
|
||||||
|
with mock.patch.object(cons.cache_controller, "load", adversarial_load):
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "ownership violation"):
|
||||||
|
cons.init_load_back(
|
||||||
|
InitLoadBackParams(
|
||||||
|
best_match_node=None, host_hit_length=f.num_tokens, req=req
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def test_buffer_only_swa_window_semantics(self):
|
def test_buffer_only_swa_window_semantics(self):
|
||||||
"""SWA window handling across the three partial-window cases:
|
"""SWA window handling across the three partial-window cases:
|
||||||
root-anchored sub-window sequence (the sequence IS its window),
|
root-anchored sub-window sequence (the sequence IS its window),
|
||||||
|
|||||||
Reference in New Issue
Block a user