Scope prefetch cache state to the request attempt (#39318)
This commit is contained in:
@@ -10,6 +10,7 @@ from sglang.srt.disaggregation.decode_hicache_mixin import (
|
||||
DecodeHiCachePreallocMixin,
|
||||
DecodePrefixMatch,
|
||||
)
|
||||
from sglang.srt.mem_cache.base_prefix_cache import CacheRequestHandle
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -26,6 +27,7 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
||||
tree_cache = SimpleNamespace(
|
||||
hicache_storage_pass_prefix_keys=True,
|
||||
ongoing_prefetch=ongoing_prefetch,
|
||||
has_ongoing_prefetch=ongoing_prefetch.__contains__,
|
||||
is_backuped=Mock(return_value=True),
|
||||
is_root=Mock(return_value=False),
|
||||
get_last_hash_value=Mock(return_value="h2"),
|
||||
@@ -39,6 +41,7 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
||||
)
|
||||
req = SimpleNamespace(
|
||||
rid="req-0",
|
||||
cache_request_handle=CacheRequestHandle("req-0", 0),
|
||||
origin_input_ids=[0, 1, 2, 3, 4, 5, 6, 7],
|
||||
extra_key="model",
|
||||
cache_salt="tenant-a",
|
||||
@@ -68,7 +71,7 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
||||
|
||||
self.assertTrue(prefix_match.prefetch_registered)
|
||||
tree_cache.prefetch_from_storage.assert_called_once_with(
|
||||
"req-0",
|
||||
req.cache_request_handle,
|
||||
22,
|
||||
[4, 5],
|
||||
"h2",
|
||||
@@ -88,6 +91,7 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
||||
harness = SimpleNamespace(tree_cache=tree_cache)
|
||||
req = SimpleNamespace(
|
||||
rid="req-0",
|
||||
cache_request_handle=CacheRequestHandle("req-0", 0),
|
||||
origin_input_ids=[0, 1, 2, 3, 4, 5],
|
||||
extra_key=None,
|
||||
cache_salt=None,
|
||||
|
||||
@@ -11,6 +11,10 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||
SchedulerBatchResultProcessor,
|
||||
)
|
||||
from sglang.srt.managers.utils import GenerationBatchResult
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
CacheRequestHandle,
|
||||
CacheRequestOutcome,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
@@ -20,6 +24,7 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
class _Req:
|
||||
def __init__(self, *, inflight_middle_chunks: int, allocated: bool = True):
|
||||
self.rid = "aborted-prefill"
|
||||
self.cache_request_handle = CacheRequestHandle(self.rid, 0)
|
||||
self.inflight_middle_chunks = inflight_middle_chunks
|
||||
self.kv = ReqKvInfo(
|
||||
req_pool_idx=1 if allocated else None,
|
||||
@@ -105,7 +110,9 @@ def test_aborted_final_result_releases_hybrid_cache(
|
||||
maybe_cache_unfinished_req.assert_not_called()
|
||||
req.disagg_kv_sender.abort.assert_called_once_with()
|
||||
scheduler.req_to_metadata_buffer_idx_allocator.free.assert_called_once_with(7)
|
||||
scheduler.tree_cache.release_aborted_request.assert_called_once_with(req.rid)
|
||||
scheduler.tree_cache.finish.assert_called_once_with(
|
||||
req.cache_request_handle, CacheRequestOutcome.ABORT
|
||||
)
|
||||
scheduler.output_streamer.stream_output.assert_called_once_with([req], False)
|
||||
scheduler.send_kv_chunk.assert_not_called()
|
||||
assert req.output_ids == []
|
||||
@@ -254,7 +261,9 @@ def test_sampling_mask_abort_preserves_error_and_releases_once(
|
||||
release_kv_cache.assert_called_once_with(req, scheduler.tree_cache, is_insert=False)
|
||||
req.disagg_kv_sender.abort.assert_called_once_with()
|
||||
scheduler.req_to_metadata_buffer_idx_allocator.free.assert_called_once_with(7)
|
||||
scheduler.tree_cache.release_aborted_request.assert_called_once_with(req.rid)
|
||||
scheduler.tree_cache.finish.assert_called_once_with(
|
||||
req.cache_request_handle, CacheRequestOutcome.ABORT
|
||||
)
|
||||
scheduler.output_streamer.stream_output.assert_called_once_with([req], False)
|
||||
scheduler.send_kv_chunk.assert_not_called()
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from sglang.srt.managers.schedule_policy import (
|
||||
estimate_prefill_extend_tile_metrics,
|
||||
)
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
CacheRequestHandle,
|
||||
DecLockRefResult,
|
||||
IncLockRefResult,
|
||||
)
|
||||
@@ -92,6 +93,7 @@ class TestPrefillAdder(CustomTestCase):
|
||||
def create_mock_req(self, rid, priority, max_new_tokens, output_len=0, wait_time=0):
|
||||
req = MagicMock(spec=Req)
|
||||
req.rid = str(rid)
|
||||
req.cache_request_handle = CacheRequestHandle(req.rid, 0)
|
||||
req.priority = priority
|
||||
req.prefix_indices = []
|
||||
req.full_untruncated_fill_ids = []
|
||||
@@ -139,7 +141,7 @@ class TestPrefillAdder(CustomTestCase):
|
||||
adder._account_prefill_cache_admission(req, prefix_len=12)
|
||||
|
||||
self.mock_tree_cache.finish_storage_prefetch_admission.assert_called_once_with(
|
||||
"storage-hit",
|
||||
req.cache_request_handle,
|
||||
fulfilled_tokens=8,
|
||||
reason=None,
|
||||
)
|
||||
@@ -153,7 +155,7 @@ class TestPrefillAdder(CustomTestCase):
|
||||
req.fulfilled_storage_hit_len.return_value = 0
|
||||
adder._account_prefill_cache_admission(req, prefix_len=0)
|
||||
self.mock_tree_cache.finish_storage_prefetch_admission.assert_called_once_with(
|
||||
"storage-hit", fulfilled_tokens=0, reason="device_capacity"
|
||||
req.cache_request_handle, fulfilled_tokens=0, reason="device_capacity"
|
||||
)
|
||||
|
||||
def test_retracted_storage_prefetch_accounting_is_omitted(self):
|
||||
@@ -166,7 +168,7 @@ class TestPrefillAdder(CustomTestCase):
|
||||
adder._account_prefill_cache_admission(req, prefix_len=8)
|
||||
|
||||
self.mock_tree_cache.discard_storage_prefetch_accounting.assert_called_once_with(
|
||||
"retracted-storage-hit"
|
||||
req.cache_request_handle
|
||||
)
|
||||
self.mock_tree_cache.finish_storage_prefetch_admission.assert_not_called()
|
||||
|
||||
|
||||
@@ -17,6 +17,10 @@ from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||
maybe_stub_sgl_kernel()
|
||||
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
CacheRequestHandle,
|
||||
CacheRequestOutcome,
|
||||
)
|
||||
|
||||
register_cpu_ci(est_time=12, suite="base-a-test-cpu")
|
||||
|
||||
@@ -24,6 +28,7 @@ register_cpu_ci(est_time=12, suite="base-a-test-cpu")
|
||||
class _FakeReq:
|
||||
def __init__(self, rid, wait_entry=0.0, forward_entry=0.0, is_finished=False):
|
||||
self.rid = rid
|
||||
self.cache_request_handle = CacheRequestHandle(rid, 0)
|
||||
self.to_finish = None
|
||||
self.beam_group = None
|
||||
self._finished = is_finished
|
||||
@@ -83,11 +88,13 @@ class TestQueuedLimitAbort(CustomTestCase):
|
||||
s.enable_priority_scheduling = True
|
||||
s.schedule_low_priority_values_first = False
|
||||
s.enable_hierarchical_cache = True
|
||||
s.tree_cache = MagicMock(spec=["release_aborted_request"])
|
||||
s.tree_cache = MagicMock(spec=["finish"])
|
||||
|
||||
self.assertFalse(s._abort_on_queued_limit(incoming))
|
||||
|
||||
s.tree_cache.release_aborted_request.assert_called_once_with("candidate")
|
||||
s.tree_cache.finish.assert_called_once_with(
|
||||
candidate.cache_request_handle, CacheRequestOutcome.ABORT
|
||||
)
|
||||
self.assertEqual(s.waiting_queue, [])
|
||||
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.base_prefix_cache import CacheRequestHandle
|
||||
from sglang.srt.mem_cache.buffer_mode.pipeline import (
|
||||
BufferModePipeline,
|
||||
_UnifiedBackupIntent,
|
||||
@@ -237,7 +238,7 @@ class TestBufferModeSidecar(unittest.TestCase):
|
||||
storage_start=0,
|
||||
)
|
||||
host_indices = torch.arange(4, dtype=torch.int64)
|
||||
req_id = "sidecar-prefetch"
|
||||
req_id = CacheRequestHandle("sidecar-prefetch", 0)
|
||||
|
||||
cache = MagicMock()
|
||||
cache.page_size = 2
|
||||
@@ -263,7 +264,7 @@ class TestBufferModeSidecar(unittest.TestCase):
|
||||
|
||||
self.assertTrue(
|
||||
pipeline.stage_completed_prefetch(
|
||||
req_id=req_id,
|
||||
request=req_id,
|
||||
num_tokens=len(host_indices),
|
||||
hash_value=["page-0", "page-1"],
|
||||
)
|
||||
|
||||
@@ -9,6 +9,7 @@ import torch
|
||||
|
||||
from sglang.srt.managers.cache_controller import CacheOperation, HiCacheController
|
||||
from sglang.srt.mem_cache import l2_transfer as transfer_module
|
||||
from sglang.srt.mem_cache.base_prefix_cache import CacheRequestHandle
|
||||
from sglang.srt.mem_cache.buffer_mode.pipeline import BufferModePipeline
|
||||
from sglang.srt.mem_cache.hicache_storage import (
|
||||
PoolHitPolicy,
|
||||
@@ -264,12 +265,13 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
||||
self.assertEqual(controller.ack_load_queue[0].node_ids, [7, 7])
|
||||
|
||||
def test_short_staged_swa_tail_resolves_device_covered_head(self):
|
||||
handle = CacheRequestHandle("r", 0)
|
||||
pipeline = BufferModePipeline.__new__(BufferModePipeline)
|
||||
pipeline._cache = mock.Mock()
|
||||
pipeline.release_staged_hold = mock.Mock(return_value=True)
|
||||
pipeline.staged_prefetches = {
|
||||
"r": SimpleNamespace(
|
||||
req_id="r",
|
||||
handle: SimpleNamespace(
|
||||
request=handle,
|
||||
key_tokens=list(range(8)),
|
||||
extra_key=None,
|
||||
cache_salt=None,
|
||||
@@ -288,9 +290,13 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
||||
)
|
||||
}
|
||||
|
||||
self.assertEqual(pipeline.plan_staged_splice("r", device_prefix_len=6), (0, 0))
|
||||
pipeline._cache._resolve_storage_prefetch_tokens.assert_called_once_with("r", 4)
|
||||
pipeline.release_staged_hold.assert_called_once_with("r", reason="shrunk")
|
||||
self.assertEqual(
|
||||
pipeline.plan_staged_splice(handle, device_prefix_len=6), (0, 0)
|
||||
)
|
||||
pipeline._cache._resolve_storage_prefetch_tokens.assert_called_once_with(
|
||||
handle, 4
|
||||
)
|
||||
pipeline.release_staged_hold.assert_called_once_with(handle, reason="shrunk")
|
||||
|
||||
def test_l2_transfer_maps_global_layers(self):
|
||||
host_pool = mock.Mock()
|
||||
|
||||
@@ -30,6 +30,8 @@ from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req, ReqKvInfo
|
||||
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
CacheRequestHandle,
|
||||
CacheRequestOutcome,
|
||||
DecLockRefParams,
|
||||
EvictParams,
|
||||
InitLoadBackParams,
|
||||
@@ -899,7 +901,9 @@ class TestUnifiedRadixCacheEagleHiCacheStorageKey(CustomTestCase):
|
||||
|
||||
controller = FakeCacheController()
|
||||
cache.cache_controller = controller
|
||||
cache.prefetch_from_storage("req", cache.root_node_handle(), tokens)
|
||||
cache.prefetch_from_storage(
|
||||
CacheRequestHandle("req", 0), cache.root_node_handle(), tokens
|
||||
)
|
||||
|
||||
_, storage_key, _, _, _ = controller.prefetch_args
|
||||
self.assertIsInstance(storage_key, RadixKey)
|
||||
@@ -954,7 +958,7 @@ class TestUnifiedRadixCacheEagleHiCacheStorageKey(CustomTestCase):
|
||||
)
|
||||
self.assertEqual(len(match.device_indices), len(prefix_tokens) - 1)
|
||||
|
||||
req_id = "bigram-anchor"
|
||||
req_id = CacheRequestHandle("bigram-anchor", 0)
|
||||
prefetch_key = RadixKey(
|
||||
array("q", [prefix_tokens[-1], 6, 7, 8, 9]),
|
||||
extra_key=extra_key,
|
||||
@@ -3193,7 +3197,8 @@ class UnifiedRadixCacheSuite:
|
||||
if prefix_len is None:
|
||||
prefix_len = f.matched_len
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
req.rid = req_id.rid
|
||||
req.cache_request_handle = req_id
|
||||
req.extra_key = extra_key
|
||||
req.cache_salt = cache_salt
|
||||
if prefix_indices is not None:
|
||||
@@ -3315,7 +3320,7 @@ class UnifiedRadixCacheSuite:
|
||||
storage_dir=storage_dir,
|
||||
prefetch_threshold=1,
|
||||
)
|
||||
req_id = "l3-prefetch-req"
|
||||
req_id = CacheRequestHandle("l3-prefetch-req", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -3387,7 +3392,7 @@ class UnifiedRadixCacheSuite:
|
||||
storage_dir=storage_dir,
|
||||
prefetch_threshold=1,
|
||||
)
|
||||
req_id = "abort-req"
|
||||
req_id = CacheRequestHandle("abort-req", 0)
|
||||
|
||||
cc = cons.cache_controller
|
||||
kv_pool_available_size_before = cc.mem_pool_host.get_pool(
|
||||
@@ -3449,7 +3454,7 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertFalse(op.pool_transfers_done)
|
||||
self.assertEqual(cc.host_mem_release_queue.qsize(), 0)
|
||||
self.assertEqual(swa_release_q.qsize(), 0)
|
||||
cons.release_aborted_request(req_id)
|
||||
cons.finish(req_id, CacheRequestOutcome.ABORT)
|
||||
self.assertEqual(cc.host_mem_release_queue.qsize(), 0)
|
||||
self.assertEqual(swa_release_q.qsize(), 0)
|
||||
|
||||
@@ -3542,7 +3547,7 @@ class UnifiedRadixCacheSuite:
|
||||
storage_dir=storage_dir,
|
||||
prefetch_threshold=1,
|
||||
)
|
||||
req_id = "abort-req"
|
||||
req_id = CacheRequestHandle("abort-req", 0)
|
||||
|
||||
cc = cons.cache_controller
|
||||
kv_pool_available_size_before = cc.mem_pool_host.get_pool(
|
||||
@@ -3594,7 +3599,7 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertEqual(swa_release_q.qsize(), 0)
|
||||
|
||||
# --- Act: abort without committing the prefetch. ---
|
||||
cons.release_aborted_request(req_id)
|
||||
cons.finish(req_id, CacheRequestOutcome.ABORT)
|
||||
|
||||
self.assertTrue(op.pool_transfers_done)
|
||||
self.assertGreater(swa_release_q.qsize(), 0)
|
||||
@@ -3815,7 +3820,7 @@ class UnifiedRadixCacheSuite:
|
||||
dev_avail0 = cons.token_to_kv_pool_allocator.available_size()
|
||||
stats = cons._prefetch_outcome_stats
|
||||
|
||||
req_id = "buffer-read-roundtrip"
|
||||
req_id = CacheRequestHandle("buffer-read-roundtrip", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -3865,7 +3870,8 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
held = cons.buffer_pipeline.staged_prefetches[req_id]
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
req.rid = req_id.rid
|
||||
req.cache_request_handle = req_id
|
||||
req.extra_key = None
|
||||
req.cache_salt = None
|
||||
req.last_node = cons.root_node_handle()
|
||||
@@ -3954,7 +3960,7 @@ class UnifiedRadixCacheSuite:
|
||||
),
|
||||
len(seq),
|
||||
)
|
||||
root_req = "salted-root-prefetch"
|
||||
root_req = CacheRequestHandle("salted-root-prefetch", 0)
|
||||
cons.prefetch_from_storage(
|
||||
root_req,
|
||||
cons.root_node_handle(),
|
||||
@@ -4021,7 +4027,7 @@ class UnifiedRadixCacheSuite:
|
||||
)
|
||||
anchor = prefix_match.last_device_node
|
||||
lock_ref = _device_lock_ref(cons2, anchor, ComponentType.FULL)
|
||||
anchored_req = "salted-mid-tree-prefetch"
|
||||
anchored_req = CacheRequestHandle("salted-mid-tree-prefetch", 0)
|
||||
cons2.prefetch_from_storage(
|
||||
anchored_req,
|
||||
anchor,
|
||||
@@ -4080,7 +4086,7 @@ class UnifiedRadixCacheSuite:
|
||||
stats = cons._prefetch_outcome_stats
|
||||
|
||||
# Query BEFORE any producer wrote the span: full miss -> revoked.
|
||||
req_id = "early-query-miss"
|
||||
req_id = CacheRequestHandle("early-query-miss", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4112,7 +4118,7 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertEqual(cons.pop_prefetch_loaded_tokens(req_id), len(seq))
|
||||
|
||||
# Unserved markers must not leak: abort cleanup ...
|
||||
aborted_rid = "aborted-miss"
|
||||
aborted_rid = CacheRequestHandle("aborted-miss", 0)
|
||||
cons.prefetch_from_storage(
|
||||
aborted_rid,
|
||||
cons.root_node_handle(),
|
||||
@@ -4125,15 +4131,21 @@ class UnifiedRadixCacheSuite:
|
||||
lambda: cons.check_prefetch_progress(aborted_rid),
|
||||
"aborted-rid miss did not resolve",
|
||||
)
|
||||
cons.release_aborted_request(aborted_rid)
|
||||
cons.finish(aborted_rid, CacheRequestOutcome.ABORT)
|
||||
self.assertFalse(cons.pop_storage_prefetch_miss(aborted_rid))
|
||||
|
||||
# A fully-device-matched (empty-suffix) decline also arms the retry:
|
||||
# the device match can evict while the request waits in the queue.
|
||||
cons.prefetch_from_storage(
|
||||
"fully-matched", cons.root_node_handle(), array("q", []), None, None
|
||||
CacheRequestHandle("fully-matched", 0),
|
||||
cons.root_node_handle(),
|
||||
array("q", []),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
self.assertTrue(
|
||||
cons.pop_storage_prefetch_miss(CacheRequestHandle("fully-matched", 0))
|
||||
)
|
||||
self.assertTrue(cons.pop_storage_prefetch_miss("fully-matched"))
|
||||
cons.sanity_check()
|
||||
|
||||
def test_buffer_only_anchor_lock_cap_clamped_by_context_headroom(self):
|
||||
@@ -4192,7 +4204,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cons, cons_alloc, _ = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
req_id = "buffer-swa-admission-oom"
|
||||
req_id = CacheRequestHandle("buffer-swa-admission-oom", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4226,7 +4238,8 @@ class UnifiedRadixCacheSuite:
|
||||
# Consume at admission (init_load_back + request lock).
|
||||
held = cons.buffer_pipeline.staged_prefetches[req_id]
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
req.rid = req_id.rid
|
||||
req.cache_request_handle = req_id
|
||||
req.extra_key = None
|
||||
req.cache_salt = None
|
||||
req.last_node = cons.root_node_handle()
|
||||
@@ -4297,7 +4310,7 @@ class UnifiedRadixCacheSuite:
|
||||
cons.storage_metrics_collector = mock.Mock()
|
||||
avail0 = self._host_avail_sizes(cons)
|
||||
|
||||
req_id = "sibling-publish"
|
||||
req_id = CacheRequestHandle("sibling-publish", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4371,7 +4384,7 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertEqual(masked.full_kv_hit_length, len(seq), "live FULL not resident")
|
||||
|
||||
avail0 = self._host_avail_sizes(cons)
|
||||
req_id = "masked-overlap"
|
||||
req_id = CacheRequestHandle("masked-overlap", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4414,7 +4427,7 @@ class UnifiedRadixCacheSuite:
|
||||
cons, cons_alloc, cons_rtp = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
|
||||
req_id = "post-check-overlap"
|
||||
req_id = CacheRequestHandle("post-check-overlap", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4439,7 +4452,8 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
f = cons.buffer_pipeline.staged_prefetches[req_id]
|
||||
req = mock.Mock()
|
||||
req.rid = req_id
|
||||
req.rid = req_id.rid
|
||||
req.cache_request_handle = req_id
|
||||
req.extra_key = None
|
||||
req.cache_salt = None
|
||||
req.prefix_indices = torch.zeros(
|
||||
@@ -4482,7 +4496,7 @@ class UnifiedRadixCacheSuite:
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
avail0 = self._host_avail_sizes(cons)
|
||||
|
||||
req_id = "growth-trim"
|
||||
req_id = CacheRequestHandle("growth-trim", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4551,7 +4565,7 @@ class UnifiedRadixCacheSuite:
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
avail0 = self._host_avail_sizes(cons)
|
||||
|
||||
req_id = "covered-hold"
|
||||
req_id = CacheRequestHandle("covered-hold", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4597,7 +4611,7 @@ class UnifiedRadixCacheSuite:
|
||||
avail0 = self._host_avail_sizes(cons)
|
||||
stats = cons._prefetch_outcome_stats
|
||||
|
||||
req_id = "covered-at-commit"
|
||||
req_id = CacheRequestHandle("covered-at-commit", 0)
|
||||
cons.prefetch_from_storage(
|
||||
req_id, cons.root_node_handle(), array("q", seq), None, None
|
||||
)
|
||||
@@ -4648,9 +4662,13 @@ class UnifiedRadixCacheSuite:
|
||||
cons, _, _ = build_fixture(self.cfg)
|
||||
self._init_buffer_hicache(cons, storage_dir)
|
||||
cons.prefetch_from_storage(
|
||||
"short-req", cons.root_node_handle(), array("q", seq), None, None
|
||||
CacheRequestHandle("short-req", 0),
|
||||
cons.root_node_handle(),
|
||||
array("q", seq),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
self._run_prefetch_to_completion(cons, "short-req")
|
||||
self._run_prefetch_to_completion(cons, CacheRequestHandle("short-req", 0))
|
||||
cons.drain_storage_control_queues()
|
||||
mc = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self.assertEqual(len(mc.device_indices), len(seq))
|
||||
@@ -4671,7 +4689,7 @@ class UnifiedRadixCacheSuite:
|
||||
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",
|
||||
CacheRequestHandle("subwin-req", 0),
|
||||
m.last_device_node,
|
||||
array("q", seq_ab[len(seq_a) :]),
|
||||
cons2.get_last_hash_value(m.last_device_node),
|
||||
@@ -4681,8 +4699,10 @@ class UnifiedRadixCacheSuite:
|
||||
self._pump_hicache_until(
|
||||
cons2,
|
||||
lambda: (
|
||||
cons2.check_prefetch_progress("subwin-req")
|
||||
and cons2.buffer_pipeline.has_staged("subwin-req")
|
||||
cons2.check_prefetch_progress(CacheRequestHandle("subwin-req", 0))
|
||||
and cons2.buffer_pipeline.has_staged(
|
||||
CacheRequestHandle("subwin-req", 0)
|
||||
)
|
||||
),
|
||||
"sub-window prefetch did not stage",
|
||||
)
|
||||
@@ -4690,13 +4710,15 @@ class UnifiedRadixCacheSuite:
|
||||
any(
|
||||
t.name == PoolName.SWA
|
||||
for t in cons2.buffer_pipeline.staged_prefetches[
|
||||
"subwin-req"
|
||||
CacheRequestHandle("subwin-req", 0)
|
||||
].aux_xfers
|
||||
),
|
||||
"sub-window fetch degraded to KV-only",
|
||||
)
|
||||
spliced = self._consume_staged_prefetch(
|
||||
cons2, "subwin-req", prefix_indices=m.device_indices
|
||||
cons2,
|
||||
CacheRequestHandle("subwin-req", 0),
|
||||
prefix_indices=m.device_indices,
|
||||
)
|
||||
self.assertEqual(int(spliced.numel()), len(seq_ab) - len(seq_a))
|
||||
self.assertEqual(
|
||||
@@ -4719,9 +4741,13 @@ class UnifiedRadixCacheSuite:
|
||||
self._init_buffer_hicache(cons3, storage_dir)
|
||||
avail3 = self._host_avail_sizes(cons3)
|
||||
cons3.prefetch_from_storage(
|
||||
"partial-req", cons3.root_node_handle(), array("q", full), None, None
|
||||
CacheRequestHandle("partial-req", 0),
|
||||
cons3.root_node_handle(),
|
||||
array("q", full),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
self._run_prefetch_to_completion(cons3, "partial-req")
|
||||
self._run_prefetch_to_completion(cons3, CacheRequestHandle("partial-req", 0))
|
||||
cons3.drain_storage_control_queues()
|
||||
self.assertEqual(
|
||||
len(
|
||||
@@ -4817,7 +4843,7 @@ class UnifiedRadixCacheSuite:
|
||||
# Baseline (single rank) must actually adopt SWA, else the TP assertions
|
||||
# below would be vacuous -> skip.
|
||||
base = self._l3_consumer(storage_dir)
|
||||
self._consume_prefetch(base, seq, "base")
|
||||
self._consume_prefetch(base, seq, CacheRequestHandle("base", 0))
|
||||
if not self._swa_host_on_path(base, seq):
|
||||
self.skipTest("fixture does not exercise SWA L3 prefetch")
|
||||
return storage_dir, seq
|
||||
@@ -4833,7 +4859,7 @@ class UnifiedRadixCacheSuite:
|
||||
cons = self._l3_consumer(storage_dir)
|
||||
cons.tp_world_size = 2
|
||||
self._patch_tp_prefetch_sync(cons, drop_swa=True)
|
||||
self._consume_prefetch(cons, seq, "drop")
|
||||
self._consume_prefetch(cons, seq, CacheRequestHandle("drop", 0))
|
||||
|
||||
m = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self.assertEqual(m.host_hit_length, 0)
|
||||
@@ -4853,7 +4879,7 @@ class UnifiedRadixCacheSuite:
|
||||
cons = self._l3_consumer(storage_dir)
|
||||
cons.tp_world_size = 2
|
||||
self._patch_tp_prefetch_sync(cons, drop_swa=False) # peer == local
|
||||
self._consume_prefetch(cons, seq, "keep")
|
||||
self._consume_prefetch(cons, seq, CacheRequestHandle("keep", 0))
|
||||
|
||||
m = cons.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
self.assertEqual(m.host_hit_length, len(seq))
|
||||
@@ -4875,7 +4901,7 @@ class UnifiedRadixCacheSuite:
|
||||
cons.tp_world_size = 2
|
||||
self._patch_tp_prefetch_sync(cons, drop_swa=True)
|
||||
avail_before = cons.swa_kv_pool_host.available_size()
|
||||
self._consume_prefetch(cons, seq, "drop")
|
||||
self._consume_prefetch(cons, seq, CacheRequestHandle("drop", 0))
|
||||
|
||||
self.assertEqual(
|
||||
cons.match_prefix(
|
||||
@@ -9003,10 +9029,11 @@ class TestPrefetchCommitOrdering(CustomTestCase):
|
||||
insert_result.host_insert_dropped = False
|
||||
cache.tree_core.insert_host.return_value = insert_result
|
||||
operation = mock.MagicMock()
|
||||
operation.handle = CacheRequestHandle("req", 0)
|
||||
operation.request_id = "req"
|
||||
operation.completed_tokens = 8
|
||||
cache.ongoing_prefetch = {
|
||||
operation.request_id: (
|
||||
operation.handle: (
|
||||
7,
|
||||
list(range(8)),
|
||||
list(range(100, 108)),
|
||||
@@ -9041,7 +9068,11 @@ class TestPrefetchCommitOrdering(CustomTestCase):
|
||||
|
||||
cache._handle_prefetch_result = _handle_prefetch_result
|
||||
|
||||
self.assertTrue(UnifiedRadixCache.check_prefetch_progress(cache, "req"))
|
||||
self.assertTrue(
|
||||
UnifiedRadixCache.check_prefetch_progress(
|
||||
cache, CacheRequestHandle("req", 0)
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual([c[0] for c in order.mock_calls], ["apply", "commit", "apply"])
|
||||
self.assertEqual(applied[0], [walk_action])
|
||||
@@ -9182,8 +9213,9 @@ class TestUnifiedRadixPrefetchCorruption(CustomTestCase):
|
||||
},
|
||||
)
|
||||
anchor_lock_params = cache.inc_host_lock_ref(parent_id).to_dec_params()
|
||||
req_id = "drop-all-resources"
|
||||
operation.request_id = req_id
|
||||
req_id = CacheRequestHandle("drop-all-resources", 0)
|
||||
operation.handle = req_id
|
||||
operation.request_id = req_id.rid
|
||||
cache.ongoing_prefetch[req_id] = _OngoingPrefetch(
|
||||
parent_id,
|
||||
prefetch_key,
|
||||
@@ -9496,7 +9528,7 @@ class TestAnchorLockOutcomePolicy(CustomTestCase):
|
||||
instead of gambling the read; cap_skip over budget (checked before the
|
||||
match walk)."""
|
||||
|
||||
_REQ = "req-1"
|
||||
_REQ = CacheRequestHandle("req-1", 0)
|
||||
_PREFIX = list(range(100, 100 + 8))
|
||||
|
||||
def _make_pipeline(self, cache, cap_tokens=10_000):
|
||||
|
||||
Reference in New Issue
Block a user