[HiCache] Buffer-only mode for HiCache host memory layer (#34798)
This commit is contained in:
@@ -92,8 +92,8 @@ from sglang.srt.utils import get_device
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=16, stage="base-b", runner_config="1-gpu-small")
|
||||
register_amd_ci(est_time=16, suite="stage-b-test-1-gpu-small-amd")
|
||||
register_cuda_ci(est_time=50, stage="base-b", runner_config="1-gpu-small")
|
||||
register_amd_ci(est_time=50, suite="stage-b-test-1-gpu-small-amd")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -595,39 +595,20 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
return leaf
|
||||
|
||||
def _init_hicache(self, cache, *, write_policy: str = "write_through"):
|
||||
import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler
|
||||
|
||||
# Wrap the host-pool factory (not MHATokenToKVPoolHost directly)
|
||||
# because the assembler picks between MHATokenToKVPoolHost and
|
||||
# AsymmetricMHATokenToKVPoolHost via get_mha_host_pool_cls(device_pool).
|
||||
orig_get_mha_host_pool_cls = assembler.get_mha_host_pool_cls
|
||||
|
||||
def get_mha_host_pool_cls_wrapper(device_pool):
|
||||
host_pool_cls = orig_get_mha_host_pool_cls(device_pool)
|
||||
|
||||
def kv_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return host_pool_cls(*args, **kwargs)
|
||||
|
||||
return kv_host_pool_wrapper
|
||||
|
||||
patcher = mock.patch.object(
|
||||
assembler,
|
||||
"get_mha_host_pool_cls",
|
||||
side_effect=get_mha_host_pool_cls_wrapper,
|
||||
)
|
||||
patcher.start()
|
||||
self.addCleanup(patcher.stop)
|
||||
|
||||
# Production config: kernel IO backend + page_first layout with
|
||||
# PINNED host pools (kernel GPU DMA requires them). Pools track their
|
||||
# cudaHostRegister'd pointers and unregister on destroy()/GC, so the
|
||||
# many fixtures sharing this process cannot collide on recycled
|
||||
# address ranges (rc=712).
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
page_size=self.cfg.page_size,
|
||||
hicache_io_backend="direct",
|
||||
hicache_mem_layout="page_first_direct",
|
||||
hicache_io_backend="kernel",
|
||||
hicache_write_policy=write_policy,
|
||||
)
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
cache.init_hicache(server_args, cache.cache_init_params)
|
||||
self.addCleanup(cache.release_host_resources)
|
||||
cache.write_through_threshold = 1 << 30
|
||||
cache.load_back_threshold = 0
|
||||
|
||||
@@ -2468,28 +2449,87 @@ class UnifiedRadixCacheSuite:
|
||||
for n in self._path_chain(cache, node):
|
||||
cache.write_backup_storage(n.id)
|
||||
|
||||
def _ongoing_l3_backups(self, cache):
|
||||
"""Storage writes in flight (buffer mode tracks them on the pipeline)."""
|
||||
if cache.buffer_pipeline is not None:
|
||||
return cache.buffer_pipeline.ongoing_backup
|
||||
return cache.ongoing_backup
|
||||
|
||||
def _flush_l3_backups(self, cache, timeout: float = 10.0):
|
||||
"""Wait for backup threads to finish, then drain acks (release locks)."""
|
||||
deadline = time.time() + timeout
|
||||
while cache.ongoing_backup and time.time() < deadline:
|
||||
while self._ongoing_l3_backups(cache) and time.time() < deadline:
|
||||
cache.drain_storage_control_queues()
|
||||
if cache.ongoing_backup:
|
||||
if self._ongoing_l3_backups(cache):
|
||||
time.sleep(0.01)
|
||||
cache.drain_storage_control_queues()
|
||||
self.assertFalse(cache.ongoing_backup, "L3 backups did not complete in time")
|
||||
self.assertFalse(
|
||||
self._ongoing_l3_backups(cache), "L3 backups did not complete in time"
|
||||
)
|
||||
|
||||
def _run_prefetch_to_completion(self, cache, req_id, timeout: float = 10.0):
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
# Host memory is reserved (and IO started) by the scheduler-thread
|
||||
# drain once the L3 hit count is known, so pump it like the real
|
||||
# scheduler loop does (check_hicache_events before progress checks).
|
||||
cache.drain_storage_control_queues()
|
||||
# drain once the L3 hit count is known, and the buffer-mode fill
|
||||
# commits at its H2D ack in loading_check — pump the full event
|
||||
# round like the real scheduler loop does.
|
||||
cache.check_hicache_events()
|
||||
if cache.check_prefetch_progress(req_id):
|
||||
# Buffer mode parks a completed fetch as a staged prefetch;
|
||||
# consume it like the PrefillAdder would so callers see the
|
||||
# span tree-resident.
|
||||
if (
|
||||
cache.buffer_pipeline is not None
|
||||
and cache.buffer_pipeline.has_staged(req_id)
|
||||
):
|
||||
self._consume_staged_prefetch(cache, req_id, timeout=timeout)
|
||||
return
|
||||
time.sleep(0.01)
|
||||
self.fail(f"prefetch {req_id} did not complete in time")
|
||||
|
||||
def _consume_staged_prefetch(
|
||||
self, cache, req_id, prefix_len=None, prefix_indices=None, timeout: float = 10.0
|
||||
):
|
||||
"""Simulate the PrefillAdder consuming a staged prefetch at admission:
|
||||
init_load_back (buffer dispatch: device alloc + queued H2D), the batch
|
||||
start_loading flush, then pump until the ack commit lands. Returns
|
||||
the spliced device indices (empty on degrade)."""
|
||||
from sglang.srt.mem_cache.base_prefix_cache import InitLoadBackParams
|
||||
|
||||
f = cache.buffer_pipeline.staged_prefetches[req_id]
|
||||
if prefix_len is None:
|
||||
prefix_len = f.matched_len
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
if prefix_indices is not None:
|
||||
# Spliceable mid-anchor consumption publishes value=cat(prefix,
|
||||
# fill) — the real device prefix is required (zeros would insert
|
||||
# bogus slots into the tree).
|
||||
assert len(prefix_indices) == prefix_len
|
||||
req.prefix_indices = prefix_indices
|
||||
else:
|
||||
req.prefix_indices = torch.zeros(
|
||||
prefix_len,
|
||||
dtype=torch.int64,
|
||||
device=cache.tree_core.empty_match_result.device_indices.device,
|
||||
)
|
||||
req.last_node = cache.root_node.id
|
||||
new_indices, _last_node = cache.init_load_back(
|
||||
InitLoadBackParams(
|
||||
best_match_node=None, host_hit_length=f.num_tokens, req=req
|
||||
)
|
||||
)
|
||||
# Batch formation flushes the queued load into the batch's producer.
|
||||
cache.ready_to_load_host_cache()
|
||||
self._pump_hicache_until(
|
||||
cache,
|
||||
lambda: not cache.buffer_pipeline.ongoing_buffer_load_back,
|
||||
"staged-prefetch consumption did not commit",
|
||||
timeout=timeout,
|
||||
)
|
||||
return new_indices
|
||||
|
||||
def _all_page_hashes(self, cache, node):
|
||||
hashes = []
|
||||
for n in self._path_chain(cache, node):
|
||||
@@ -2606,6 +2646,494 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertTrue(torch.equal(loaded_v, expected_v))
|
||||
cons.sanity_check()
|
||||
|
||||
# ================================================================
|
||||
# Buffer-only host memory mode (host = transient staging, L3 = cache)
|
||||
# ================================================================
|
||||
|
||||
def _init_buffer_hicache(
|
||||
self,
|
||||
cache,
|
||||
storage_dir,
|
||||
prefetch_policy: str = "wait_complete",
|
||||
storage_extra: Optional[dict] = None,
|
||||
):
|
||||
if self.cfg.has_mamba:
|
||||
self.skipTest(
|
||||
"buffer_only is FULL/SWA-only (no Mamba state-handoff channel "
|
||||
"on the admission-time load-back read path)"
|
||||
)
|
||||
self._init_hicache(
|
||||
cache,
|
||||
storage_backend="file",
|
||||
storage_dir=storage_dir,
|
||||
prefetch_threshold=1,
|
||||
host_memory_mode="buffer_only",
|
||||
prefetch_policy=prefetch_policy,
|
||||
storage_extra=storage_extra,
|
||||
)
|
||||
|
||||
def _pump_hicache_until(self, cache, cond, msg, timeout: float = 10.0):
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
cache.check_hicache_events()
|
||||
if cond():
|
||||
return
|
||||
time.sleep(0.01)
|
||||
self.fail(msg)
|
||||
|
||||
def _host_avail_sizes(self, cache):
|
||||
group = cache.cache_controller.mem_pool_host
|
||||
return {entry.name: entry.host_pool.available_size() for entry in group.entries}
|
||||
|
||||
def _storage_exists_count(self, cache, page_hashes, pool_transfers=None):
|
||||
"""Ground-truth longest-prefix existence count from the backend,
|
||||
folded across pools like the prefetch hit query."""
|
||||
from sglang.srt.mem_cache.hicache_storage import HiCacheStorageExtraInfo
|
||||
|
||||
backend = cache.cache_controller.storage_backend
|
||||
extra_info = HiCacheStorageExtraInfo(prefix_keys=None)
|
||||
if pool_transfers:
|
||||
result = backend.batch_exists_v2(page_hashes, pool_transfers, extra_info)
|
||||
return min(result.kv_hit_pages, len(page_hashes))
|
||||
return backend.batch_exists(page_hashes, extra_info)
|
||||
|
||||
def _buffer_backup_and_wait(self, cache, node):
|
||||
# Parent-first over the whole path, mirroring the production trigger
|
||||
# (_inc_hit_count fires per matched node on the insert walk): the SWA
|
||||
# component may have split the leaf at the window boundary.
|
||||
pipeline = cache.buffer_pipeline
|
||||
for n in self._path_chain(cache, node):
|
||||
pipeline.enqueue_backup_intent(n)
|
||||
self.assertIn(n.id, pipeline.inflight_backup_node_ids)
|
||||
self._pump_hicache_until(
|
||||
cache,
|
||||
lambda: not pipeline.inflight_backup_node_ids
|
||||
and not pipeline.ongoing_backup,
|
||||
"buffer backup pipeline did not drain",
|
||||
)
|
||||
|
||||
def _produce_buffer_l3(self, storage_dir, seq, marker=None):
|
||||
"""Producer tree in buffer mode: insert seq and push it to L3."""
|
||||
prod, prod_alloc, prod_rtp = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(prod, storage_dir)
|
||||
self._insert(prod, prod_alloc, prod_rtp, seq)
|
||||
leaf = prod.resolve_node_handle(
|
||||
prod.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq)))
|
||||
).last_device_node
|
||||
)
|
||||
expected = None
|
||||
if marker is not None:
|
||||
m = prod.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self._fill_full_kv(prod_alloc, m.device_indices, marker=marker)
|
||||
expected = self._snapshot_full_kv(prod_alloc, m.device_indices)
|
||||
self._buffer_backup_and_wait(prod, leaf)
|
||||
return leaf, expected
|
||||
|
||||
def _buffer_swa_seq(self, min_pages=4):
|
||||
"""Sequence long enough for SWA prefetch (one full window + 1)."""
|
||||
num_pages = min_pages
|
||||
if self.cfg.has_swa:
|
||||
sw_pages = (
|
||||
self.cfg.sliding_window_size + self.cfg.page_size - 1
|
||||
) // self.cfg.page_size
|
||||
num_pages = max(num_pages, sw_pages + 1)
|
||||
return self._make_seq(1, num_pages)
|
||||
|
||||
def test_buffer_only_write_path_roundtrip(self):
|
||||
"""Write path end to end: admission -> D2H staging -> storage write
|
||||
-> free. Staging and locks fully released, pages stored under every
|
||||
pool namespace, beliefs registered, re-hits skipped."""
|
||||
self._skip_unsupported_hicache_test()
|
||||
storage_dir = tempfile.mkdtemp()
|
||||
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cache, storage_dir)
|
||||
avail0 = self._host_avail_sizes(cache)
|
||||
|
||||
seq_a = self._make_seq(1, 2)
|
||||
seq_ab = seq_a + self._make_seq(500, 2)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq_a)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq_ab)
|
||||
leaf = cache.resolve_node_handle(
|
||||
cache.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq_ab)))
|
||||
).last_device_node
|
||||
)
|
||||
chain = self._path_chain(cache, leaf)
|
||||
|
||||
self._buffer_backup_and_wait(cache, leaf)
|
||||
self.assertFalse(leaf.backuped)
|
||||
self.assertEqual(leaf.component_data[ComponentType.FULL].lock_ref, 0)
|
||||
self.assertEqual(self._host_avail_sizes(cache), avail0)
|
||||
self.assertEqual(cache.buffer_pipeline.write_backlog_tokens_, 0)
|
||||
page_hashes = self._all_page_hashes(cache, leaf)
|
||||
self.assertEqual(
|
||||
self._storage_exists_count(
|
||||
cache,
|
||||
page_hashes,
|
||||
cache.buffer_pipeline._build_aux_staging_transfers(leaf),
|
||||
),
|
||||
len(page_hashes),
|
||||
)
|
||||
for n in chain:
|
||||
self.assertTrue(
|
||||
cache.storage_existence_cache.contains_all(PoolName.KV, n.hash_value)
|
||||
)
|
||||
# Re-hit absorbed by the (FULL-focused) belief skip.
|
||||
cache.buffer_pipeline.enqueue_backup_intent(leaf)
|
||||
self.assertNotIn(leaf.id, cache.buffer_pipeline.inflight_backup_node_ids)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_buffer_only_read_path_roundtrip(self):
|
||||
"""Read path end to end: prefetch -> staged (host bounce only,
|
||||
nothing device-side, unmatchable, stable readiness, counters fed) ->
|
||||
admission-time load-back against a saturated-but-evictable pool
|
||||
(evict-before-alloc) publishing pre-ack -> ack frees the bounce.
|
||||
Data bytes match the producer's; no CPU-tier KV events anywhere;
|
||||
declines feed the outcome counters."""
|
||||
self._skip_unsupported_hicache_test()
|
||||
storage_dir = tempfile.mkdtemp()
|
||||
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||
|
||||
seq = self._buffer_swa_seq()
|
||||
_, (expected_k, expected_v) = self._produce_buffer_l3(
|
||||
storage_dir, seq, marker=7
|
||||
)
|
||||
|
||||
cons, cons_alloc, cons_rtp = build_fixture(
|
||||
self.cfg, enable_kv_cache_events=True
|
||||
)
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
cons.take_events()
|
||||
avail0 = self._host_avail_sizes(cons)
|
||||
dev_avail0 = cons.token_to_kv_pool_allocator.available_size()
|
||||
stats = cons._prefetch_outcome_stats
|
||||
|
||||
req_id = "buffer-read-roundtrip"
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node.id, array("q", seq), None, None
|
||||
)
|
||||
self.assertEqual((stats["attempts"], stats["issued"]), (1, 1))
|
||||
self._pump_hicache_until(
|
||||
cons,
|
||||
lambda: cons.check_prefetch_progress(req_id)
|
||||
and cons.buffer_pipeline.has_staged(req_id),
|
||||
"prefetch did not stage",
|
||||
)
|
||||
# Staged: bounce occupies host staging; nothing device-side; span
|
||||
# unmatchable; readiness stable; hit accounting reported once.
|
||||
self.assertNotEqual(self._host_avail_sizes(cons), avail0)
|
||||
self.assertEqual(cons.token_to_kv_pool_allocator.available_size(), dev_avail0)
|
||||
self.assertEqual(
|
||||
len(
|
||||
cons.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq)))
|
||||
).device_indices
|
||||
),
|
||||
0,
|
||||
)
|
||||
self.assertTrue(cons.check_prefetch_progress(req_id))
|
||||
self.assertEqual(cons.pop_prefetch_loaded_tokens(req_id), len(seq))
|
||||
self.assertEqual(stats["l3_demand_requests"], 1)
|
||||
self.assertEqual(stats["l3_miss_tokens"], 0)
|
||||
|
||||
# Saturate the device pool with unrelated evictable spans: the
|
||||
# load-back must evict, not degrade to recompute (run-2 regression).
|
||||
def _avail():
|
||||
if cons.supports_swa():
|
||||
return cons.token_to_kv_pool_allocator.full_available_size()
|
||||
return cons.token_to_kv_pool_allocator.available_size()
|
||||
|
||||
filler_base = 90000
|
||||
while _avail() >= self.cfg.page_size:
|
||||
pages = max(1, min(2048, _avail()) // self.cfg.page_size)
|
||||
self._insert(cons, cons_alloc, cons_rtp, self._make_seq(filler_base, pages))
|
||||
filler_base += 1000
|
||||
self.assertLess(_avail(), len(seq), "pool not saturated")
|
||||
|
||||
# Consume WITHOUT pumping the ack: published (matchable, same slot
|
||||
# ids) before any ack lands.
|
||||
from sglang.srt.mem_cache.base_prefix_cache import InitLoadBackParams
|
||||
|
||||
held = cons.buffer_pipeline.staged_prefetches[req_id]
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
req.last_node = cons.root_node.id
|
||||
req.prefix_indices = torch.zeros(
|
||||
held.matched_len,
|
||||
dtype=torch.int64,
|
||||
device=cons.tree_core.empty_match_result.device_indices.device,
|
||||
)
|
||||
spliced, _last = cons.init_load_back(
|
||||
InitLoadBackParams(
|
||||
best_match_node=None, host_hit_length=held.num_tokens, req=req
|
||||
)
|
||||
)
|
||||
self.assertEqual(int(spliced.numel()), len(seq))
|
||||
cons.ready_to_load_host_cache()
|
||||
m = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self.assertTrue(torch.equal(m.device_indices, spliced))
|
||||
|
||||
# Ack: bounce freed, beliefs fed, tree holds no host values, and the
|
||||
# loaded KV bytes equal the producer's.
|
||||
self._pump_hicache_until(
|
||||
cons,
|
||||
lambda: not cons.buffer_pipeline.ongoing_buffer_load_back
|
||||
and self._host_avail_sizes(cons) == avail0,
|
||||
"load-back ack did not free the bounce",
|
||||
)
|
||||
mc = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self.assertEqual(mc.host_hit_length, 0)
|
||||
self.assertEqual(len(mc.device_indices), len(seq))
|
||||
leaf = cons.resolve_node_handle(mc.last_device_node)
|
||||
for cur in self._path_chain(cons, leaf):
|
||||
for cd in cur.component_data:
|
||||
self.assertIsNone(cd.host_value)
|
||||
self.assertTrue(
|
||||
cons.storage_existence_cache.contains_all(
|
||||
PoolName.KV, self._all_page_hashes(cons, leaf)
|
||||
)
|
||||
)
|
||||
loaded_k, loaded_v = self._snapshot_full_kv(cons_alloc, mc.device_indices)
|
||||
self.assertTrue(torch.equal(loaded_k, expected_k))
|
||||
self.assertTrue(torch.equal(loaded_v, expected_v))
|
||||
self.assertEqual(cons.cache_controller.prefetch_tokens_occupied, 0)
|
||||
cpu_events = [
|
||||
e
|
||||
for e in cons.take_events()
|
||||
if isinstance(e, (BlockStored, BlockRemoved))
|
||||
and e.medium == StorageMedium.CPU
|
||||
]
|
||||
self.assertEqual(cpu_events, [])
|
||||
|
||||
self.assertIn("occupancy_ratio", cons.prefetch_outcome_stats_snapshot())
|
||||
cons.sanity_check()
|
||||
|
||||
def test_buffer_load_back_swa_window_charged_at_admission(self):
|
||||
"""Admission contract: a request the SWA budget gate accepts must be
|
||||
allocatable at batch time (_swa_reserved_tokens: "an admitted request
|
||||
cannot OOM"). Regression: buffer mode surfaced a staged prefetch as
|
||||
host_hit_length only, so the gate never charged the SWA window that
|
||||
consumption (init_load_back -> cc.load) allocates and the request
|
||||
lock pins; with the rest of the SWA pool batch-held, the batch alloc
|
||||
fell short by up to one window and raised the fail-loud prefill OOM
|
||||
(prod scheduler crash, 2026-08-17)."""
|
||||
self._skip_unsupported_hicache_test()
|
||||
if not self.cfg.has_swa:
|
||||
self.skipTest("SWA-specific admission accounting")
|
||||
from sglang.srt.mem_cache.allocation import (
|
||||
alloc_paged_token_slots_extend,
|
||||
alloc_token_slots,
|
||||
)
|
||||
from sglang.srt.mem_cache.base_prefix_cache import InitLoadBackParams
|
||||
|
||||
storage_dir = tempfile.mkdtemp()
|
||||
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||
ps = self.cfg.page_size
|
||||
window = self.cfg.sliding_window_size
|
||||
seq = self._buffer_swa_seq() # one full window + 1 page
|
||||
self._produce_buffer_l3(storage_dir, seq)
|
||||
|
||||
cons, cons_alloc, _ = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
req_id = "buffer-swa-admission-oom"
|
||||
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",
|
||||
)
|
||||
|
||||
# Batch-held SWA (chunk allocs, decode windows) is neither free nor
|
||||
# evictable: leave one token less than window + extend_need, enough
|
||||
# for an un-charged gate to accept.
|
||||
extend_need = 2 * ps + 1
|
||||
max_new = 8
|
||||
self.assertIsNotNone(
|
||||
cons_alloc.swa_attn_allocator.alloc(
|
||||
cons_alloc.swa_available_size() - (window + extend_need - 1)
|
||||
)
|
||||
)
|
||||
|
||||
# The adder's SWA gate for this request (_swa_budget_for_req).
|
||||
surfaced_swa_hit = cons.staged_prefetch_swa_tokens(req_id)
|
||||
reserved = (
|
||||
max(extend_need - window, 0)
|
||||
+ min(extend_need + max_new, window)
|
||||
+ ps
|
||||
+ (surfaced_swa_hit + ps - 1) // ps * ps
|
||||
)
|
||||
budget = cons_alloc.swa_available_size() + cons.swa_evictable_size()
|
||||
|
||||
# Consume at admission (init_load_back + request lock).
|
||||
held = cons.buffer_pipeline.staged_prefetches[req_id]
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
req.last_node = cons.root_node.id
|
||||
req.prefix_indices = torch.zeros(
|
||||
held.matched_len,
|
||||
dtype=torch.int64,
|
||||
device=cons.tree_core.empty_match_result.device_indices.device,
|
||||
)
|
||||
spliced, last_node = cons.init_load_back(
|
||||
InitLoadBackParams(
|
||||
best_match_node=None, host_hit_length=held.num_tokens, req=req
|
||||
)
|
||||
)
|
||||
self.assertEqual(int(spliced.numel()), len(seq), "load-back degraded")
|
||||
cons.ready_to_load_host_cache()
|
||||
cons.inc_lock_ref(last_node) # _req_inc_lock_ref
|
||||
|
||||
self.assertEqual(cons.swa_evictable_size(), 0) # window is protected
|
||||
# FULL stays roomy: only SWA can fail below.
|
||||
self.assertGreaterEqual(cons_alloc.full_available_size(), extend_need + ps)
|
||||
|
||||
if reserved <= budget:
|
||||
try:
|
||||
if ps == 1:
|
||||
alloc_token_slots(cons, extend_need)
|
||||
else: # paged batch path — the production crash site
|
||||
prefix_len = int(spliced.numel())
|
||||
prefix_t = torch.tensor(
|
||||
[prefix_len], dtype=torch.int64, device=spliced.device
|
||||
)
|
||||
seq_t = torch.tensor(
|
||||
[prefix_len + extend_need],
|
||||
dtype=torch.int64,
|
||||
device=spliced.device,
|
||||
)
|
||||
alloc_paged_token_slots_extend(
|
||||
cons,
|
||||
prefix_t,
|
||||
prefix_t.cpu(),
|
||||
seq_t,
|
||||
seq_t.cpu(),
|
||||
spliced[-1:],
|
||||
extend_need,
|
||||
)
|
||||
except RuntimeError as e:
|
||||
self.fail(
|
||||
f"gate admitted (reserved={reserved} <= budget={budget}) "
|
||||
f"but the batch alloc OOMed: {e}"
|
||||
)
|
||||
else:
|
||||
# Rejection must come from the surfaced window charge.
|
||||
self.assertGreaterEqual(surfaced_swa_hit, window)
|
||||
|
||||
def test_buffer_only_swa_window_semantics(self):
|
||||
"""SWA window handling across the three partial-window cases:
|
||||
root-anchored sub-window sequence (the sequence IS its window),
|
||||
mid-tree sub-window continuation (head = device ring state), and a
|
||||
storage hit shorter than the requested window (shrunk, tail
|
||||
released). Each was a zero-L3-reuse regression on Llama-4-Scout."""
|
||||
self._skip_unsupported_hicache_test()
|
||||
if not self.cfg.has_swa:
|
||||
self.skipTest("requires an SWA component")
|
||||
window = self.cfg.sliding_window_size
|
||||
if window <= self.cfg.page_size:
|
||||
self.skipTest("window fits in one page")
|
||||
sw_pages = (window + self.cfg.page_size - 1) // self.cfg.page_size
|
||||
storage_dir = tempfile.mkdtemp()
|
||||
self.addCleanup(shutil.rmtree, storage_dir, ignore_errors=True)
|
||||
|
||||
# 1. Root-anchored sequence shorter than the window.
|
||||
seq = self._make_seq(1, (window // self.cfg.page_size) - 1)
|
||||
self._produce_buffer_l3(storage_dir, seq, marker=5)
|
||||
cons, _, _ = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
cons.prefetch_from_storage(
|
||||
"short-req", cons.root_node.id, array("q", seq), None, None
|
||||
)
|
||||
self._run_prefetch_to_completion(cons, "short-req")
|
||||
cons.drain_storage_control_queues()
|
||||
mc = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self.assertEqual(len(mc.device_indices), len(seq))
|
||||
self.assertIsNotNone(
|
||||
cons.resolve_node_handle(mc.last_device_node)
|
||||
.component_data[ComponentType.SWA]
|
||||
.value
|
||||
)
|
||||
cons.sanity_check()
|
||||
|
||||
# 2. Mid-tree continuation shorter than the window: the staged
|
||||
# prefetch must carry an SWA transfer (not a KV-only degrade).
|
||||
if sw_pages >= 2:
|
||||
seq_a = self._make_seq(1, max(2, sw_pages))
|
||||
seq_ab = seq_a + self._make_seq(900, sw_pages - 1)
|
||||
self._produce_buffer_l3(storage_dir, seq_ab, marker=14)
|
||||
cons2, cons2_alloc, cons2_rtp = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons2, storage_dir)
|
||||
avail2 = self._host_avail_sizes(cons2)
|
||||
self._insert(cons2, cons2_alloc, cons2_rtp, seq_a)
|
||||
m = cons2.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
|
||||
cons2.prefetch_from_storage(
|
||||
"subwin-req",
|
||||
m.last_device_node,
|
||||
array("q", seq_ab[len(seq_a) :]),
|
||||
cons2.get_last_hash_value(m.last_device_node),
|
||||
None,
|
||||
matched_prefix_tokens=list(seq_a),
|
||||
)
|
||||
self._pump_hicache_until(
|
||||
cons2,
|
||||
lambda: cons2.check_prefetch_progress("subwin-req")
|
||||
and cons2.buffer_pipeline.has_staged("subwin-req"),
|
||||
"sub-window prefetch did not stage",
|
||||
)
|
||||
self.assertTrue(
|
||||
any(
|
||||
t.name == PoolName.SWA
|
||||
for t in cons2.buffer_pipeline.staged_prefetches[
|
||||
"subwin-req"
|
||||
].aux_xfers
|
||||
),
|
||||
"sub-window fetch degraded to KV-only",
|
||||
)
|
||||
spliced = self._consume_staged_prefetch(
|
||||
cons2, "subwin-req", prefix_indices=m.device_indices
|
||||
)
|
||||
self.assertEqual(int(spliced.numel()), len(seq_ab) - len(seq_a))
|
||||
self.assertEqual(
|
||||
len(
|
||||
cons2.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq_ab)))
|
||||
).device_indices
|
||||
),
|
||||
len(seq_ab),
|
||||
)
|
||||
self.assertEqual(self._host_avail_sizes(cons2), avail2)
|
||||
cons2.sanity_check()
|
||||
|
||||
# 3. Hit one page short of the requested window: the shrunk window
|
||||
# is kept (its own trailing window) and the buffer tail released.
|
||||
full = self._buffer_swa_seq()
|
||||
stored = full[: window - self.cfg.page_size]
|
||||
self._produce_buffer_l3(storage_dir, stored, marker=6)
|
||||
cons3, _, _ = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons3, storage_dir)
|
||||
avail3 = self._host_avail_sizes(cons3)
|
||||
cons3.prefetch_from_storage(
|
||||
"partial-req", cons3.root_node.id, array("q", full), None, None
|
||||
)
|
||||
self._run_prefetch_to_completion(cons3, "partial-req")
|
||||
cons3.drain_storage_control_queues()
|
||||
self.assertEqual(
|
||||
len(
|
||||
cons3.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", stored)))
|
||||
).device_indices
|
||||
),
|
||||
len(stored),
|
||||
"partial hit lost its SWA window",
|
||||
)
|
||||
self.assertEqual(self._host_avail_sizes(cons3), avail3)
|
||||
cons3.sanity_check()
|
||||
|
||||
# ---------- TP consistency for SWA prefetch (all-or-nothing) ----------
|
||||
|
||||
def _patch_tp_all_reduce(self, cache, drop_swa: bool):
|
||||
@@ -2898,44 +3426,9 @@ class UnifiedRadixCacheSuite:
|
||||
storage_dir: Optional[str] = None,
|
||||
prefetch_threshold: Optional[int] = None,
|
||||
prefetch_policy: str = "wait_complete",
|
||||
host_memory_mode: str = "cache",
|
||||
storage_extra: Optional[dict] = None,
|
||||
):
|
||||
import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler
|
||||
|
||||
# See _init_hicache: wrap the factory rather than MHATokenToKVPoolHost
|
||||
# directly so the pin_memory=False override applies to both
|
||||
# MHATokenToKVPoolHost and AsymmetricMHATokenToKVPoolHost.
|
||||
orig_get_mha_host_pool_cls = assembler.get_mha_host_pool_cls
|
||||
orig_mamba_host_pool = assembler.MambaPoolHost
|
||||
|
||||
def get_mha_host_pool_cls_wrapper(device_pool):
|
||||
host_pool_cls = orig_get_mha_host_pool_cls(device_pool)
|
||||
|
||||
def kv_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return host_pool_cls(*args, **kwargs)
|
||||
|
||||
return kv_host_pool_wrapper
|
||||
|
||||
def mamba_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return orig_mamba_host_pool(*args, **kwargs)
|
||||
|
||||
patchers = [
|
||||
mock.patch.object(
|
||||
assembler,
|
||||
"get_mha_host_pool_cls",
|
||||
side_effect=get_mha_host_pool_cls_wrapper,
|
||||
),
|
||||
mock.patch.object(
|
||||
assembler,
|
||||
"MambaPoolHost",
|
||||
side_effect=mamba_host_pool_wrapper,
|
||||
),
|
||||
]
|
||||
for patcher in patchers:
|
||||
patcher.start()
|
||||
self.addCleanup(patcher.stop)
|
||||
|
||||
storage_extra_config = None
|
||||
if storage_backend == "file":
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
@@ -2957,22 +3450,26 @@ class UnifiedRadixCacheSuite:
|
||||
extra = {}
|
||||
if prefetch_threshold is not None:
|
||||
extra["prefetch_threshold"] = prefetch_threshold
|
||||
if storage_extra:
|
||||
extra.update(storage_extra)
|
||||
storage_extra_config = json.dumps(extra) if extra else None
|
||||
|
||||
# See _init_hicache: production kernel IO backend, pinned pools.
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
page_size=self.cfg.page_size,
|
||||
hicache_io_backend="direct",
|
||||
hicache_mem_layout="page_first_direct",
|
||||
hicache_io_backend="kernel",
|
||||
hicache_write_policy=write_policy,
|
||||
hicache_storage_backend=storage_backend,
|
||||
hicache_storage_backend_extra_config=storage_extra_config,
|
||||
hicache_storage_prefetch_policy=prefetch_policy,
|
||||
hicache_host_memory_mode=host_memory_mode,
|
||||
)
|
||||
# See build_fixture for why _mamba_cache_chunk_size is preset.
|
||||
server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, self.cfg.page_size)
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
cache.init_hicache(server_args, cache.cache_init_params)
|
||||
self.addCleanup(cache.release_host_resources)
|
||||
cache.write_through_threshold = 1 << 30
|
||||
cache.load_back_threshold = 0
|
||||
if storage_backend is not None:
|
||||
@@ -3020,8 +3517,8 @@ class UnifiedRadixCacheSuite:
|
||||
if node is not cache.root_node:
|
||||
self._backup_node(cache, node)
|
||||
|
||||
def _load_back_node(self, cache, node):
|
||||
loaded = cache.load_back(node.id)
|
||||
def _load_back_node(self, cache, node, req=None):
|
||||
loaded = cache.load_back(node.id, req=req)
|
||||
self.assertTrue(loaded)
|
||||
producer_id = cache.ready_to_load_host_cache()
|
||||
self.assertNotEqual(producer_id, -1)
|
||||
@@ -3729,10 +4226,14 @@ class UnifiedRadixCacheSuite:
|
||||
cache, _, _ = self._build_hicache_fixture()
|
||||
sw = cache.sliding_window_size
|
||||
swa = cache.components[ComponentType.SWA]
|
||||
# below a full window -> does not participate, no alloc
|
||||
prep = swa.prepare_prefetch(cache.root_node.id, prefetch_tokens=sw - 1)
|
||||
# zero-length prefetch -> does not participate, no alloc
|
||||
prep = swa.prepare_prefetch(cache.root_node.id, prefetch_tokens=0)
|
||||
self.assertFalse(prep.alloc_failed)
|
||||
self.assertIsNone(prep.host_indices)
|
||||
# below a full window at the ROOT anchor -> the whole sequence is its
|
||||
# own trailing window (sub-window prompts stay reusable via storage)
|
||||
prep = swa.prepare_prefetch(cache.root_node.id, prefetch_tokens=sw - 1)
|
||||
self.assertEqual(int(prep.host_indices.numel()), sw - 1)
|
||||
# a full window available -> participates, allocs one window of host pages
|
||||
prep = swa.prepare_prefetch(cache.root_node.id, prefetch_tokens=sw)
|
||||
self.assertEqual(int(prep.host_indices.numel()), sw)
|
||||
@@ -6293,6 +6794,7 @@ class TestPrefetchCommitOrdering(CustomTestCase):
|
||||
cache = mock.MagicMock()
|
||||
cache.page_size = 1
|
||||
cache.enable_storage_metrics = False
|
||||
cache.buffer_pipeline = None # cache-mode commit path
|
||||
walk_action = object()
|
||||
insert_result = mock.MagicMock()
|
||||
insert_result.cache_actions = [walk_action]
|
||||
|
||||
Reference in New Issue
Block a user