[Bugfix][HiCache] measure load-back duration with CUDA events (#26411)

Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com>
Co-authored-by: vuuihc <vuuihc@users.noreply.github.com>
Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
Ethan ZHU
2026-07-16 00:11:04 -07:00
committed by GitHub
co-authored by Claude Opus 4.7 vuuihc Zhiqiang Xie
parent f8eac995aa
commit 095a817612
10 changed files with 253 additions and 71 deletions
@@ -217,9 +217,9 @@ class DecodeKVCacheOffloadManager:
def _check_offload_progress(self, finish_count): def _check_offload_progress(self, finish_count):
"""Check the progress of offload from device to host.""" """Check the progress of offload from device to host."""
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0) ack = self.cache_controller.ack_write_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
( (
req, req,
host_indices, host_indices,
+31 -2
View File
@@ -16,6 +16,7 @@ limitations under the License.
import logging import logging
import threading import threading
import time import time
from functools import cache
from queue import Empty, Queue from queue import Empty, Queue
from typing import TYPE_CHECKING, List, NamedTuple, Optional from typing import TYPE_CHECKING, List, NamedTuple, Optional
@@ -46,6 +47,26 @@ logger = logging.getLogger(__name__)
device_module = get_device_module() device_module = get_device_module()
@cache
def _timing_events_supported() -> bool:
try:
device_module.Event(enable_timing=True)
return True
except (TypeError, NotImplementedError):
logger.warning(
"%s.Event does not support enable_timing=True; load-back "
"duration metric will be skipped on this backend.",
device_module.__name__,
)
return False
def make_timing_event_pair():
timing_enabled = _timing_events_supported()
kwargs = {"enable_timing": True} if timing_enabled else {}
return device_module.Event(**kwargs), device_module.Event(**kwargs), timing_enabled
class LayerLoadingEvent: class LayerLoadingEvent:
def __init__(self, num_layers: int): def __init__(self, num_layers: int):
self._num_layers = num_layers self._num_layers = num_layers
@@ -140,6 +161,8 @@ class HiCacheAck(NamedTuple):
start_event: device_module.Event start_event: device_module.Event
finish_event: device_module.Event finish_event: device_module.Event
node_ids: List[int] node_ids: List[int]
num_tokens: int = 0
timing_enabled: bool = False
class StorageOperation: class StorageOperation:
@@ -767,8 +790,11 @@ class HiCacheController:
producer_event = self.layer_done_counter.events[producer_id] producer_event = self.layer_done_counter.events[producer_id]
producer_event.start_event.record() producer_event.start_event.record()
ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair()
with device_module.stream(self.load_stream): with device_module.stream(self.load_stream):
producer_event.start_event.wait(self.load_stream) producer_event.start_event.wait(self.load_stream)
ack_start_event.record()
for i in range(self.layer_num): for i in range(self.layer_num):
self.mem_pool_host.load_to_device_per_layer( self.mem_pool_host.load_to_device_per_layer(
self.mem_pool_device, self.mem_pool_device,
@@ -786,6 +812,7 @@ class HiCacheController:
self.io_backend, self.io_backend,
) )
producer_event.complete(i) producer_event.complete(i)
ack_finish_event.record()
# NOTE: We must save the host indices and device indices here, # NOTE: We must save the host indices and device indices here,
# this is because we need to guarantee that these tensors are # this is because we need to guarantee that these tensors are
# still alive when the load stream is executing. # still alive when the load stream is executing.
@@ -796,9 +823,11 @@ class HiCacheController:
self.ack_load_queue.append( self.ack_load_queue.append(
HiCacheAck( HiCacheAck(
start_event=producer_event.start_event, start_event=ack_start_event,
finish_event=producer_event.finish_event, finish_event=ack_finish_event,
node_ids=op.node_ids, node_ids=op.node_ids,
num_tokens=len(op.device_indices),
timing_enabled=timing_enabled,
) )
) )
return producer_id return producer_id
@@ -385,9 +385,9 @@ class HiMambaRadixCache(MambaRadixCache):
if write_back: if write_back:
# blocking till all write back complete # blocking till all write back complete
while len(self.ongoing_write_through) > 0: while len(self.ongoing_write_through) > 0:
for _, finish_event, ack_list in self.cache_controller.ack_write_queue: for ack in self.cache_controller.ack_write_queue:
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
backuped_node = self.ongoing_write_through.pop(ack_id) backuped_node = self.ongoing_write_through.pop(ack_id)
self._record_store_event( self._record_store_event(
backuped_node, medium=StorageMedium.CPU backuped_node, medium=StorageMedium.CPU
@@ -403,8 +403,8 @@ class HiMambaRadixCache(MambaRadixCache):
# independently (no cross-rank sync). # independently (no cross-rank sync).
finish_count = 0 finish_count = 0
if len(self.ongoing_write_through) > 0: if len(self.ongoing_write_through) > 0:
for _, finish_event, ack_list in self.cache_controller.ack_write_queue: for ack in self.cache_controller.ack_write_queue:
if not finish_event.query(): if not ack.finish_event.query():
break break
finish_count += 1 finish_count += 1
@@ -418,9 +418,9 @@ class HiMambaRadixCache(MambaRadixCache):
finish_count = int(queue_size.item()) finish_count = int(queue_size.item())
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0) ack = self.cache_controller.ack_write_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
backuped_node = self.ongoing_write_through.pop(ack_id) backuped_node = self.ongoing_write_through.pop(ack_id)
self._record_store_event(backuped_node, medium=StorageMedium.CPU) self._record_store_event(backuped_node, medium=StorageMedium.CPU)
self.dec_lock_ref(backuped_node) self.dec_lock_ref(backuped_node)
@@ -432,8 +432,8 @@ class HiMambaRadixCache(MambaRadixCache):
# Every rank must enter the all_reduce below; ongoing_load_back can # Every rank must enter the all_reduce below; ongoing_load_back can
# diverge across ranks. # diverge across ranks.
finish_count = 0 finish_count = 0
for _, finish_event, ack_list in self.cache_controller.ack_load_queue: for ack in self.cache_controller.ack_load_queue:
if not finish_event.query(): if not ack.finish_event.query():
break break
finish_count += 1 finish_count += 1
@@ -447,11 +447,19 @@ class HiMambaRadixCache(MambaRadixCache):
finish_count = int(queue_size.item()) finish_count = int(queue_size.item())
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_load_queue.pop(0) ack = self.cache_controller.ack_load_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
end_node = self.ongoing_load_back.pop(ack_id) end_node = self.ongoing_load_back.pop(ack_id)
self.dec_lock_ref(end_node) self.dec_lock_ref(end_node)
if self.metrics_collector is not None:
self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens)
if ack.timing_enabled:
duration_ms = ack.start_event.elapsed_time(ack.finish_event)
self.metrics_collector.observe_load_back_duration(
duration_ms / 1000.0
)
finish_count -= 1 finish_count -= 1
def ready_to_load_host_cache(self) -> int: def ready_to_load_host_cache(self) -> int:
+21 -20
View File
@@ -938,9 +938,9 @@ class HiRadixCache(RadixCache):
if write_back: if write_back:
# blocking till all write back complete # blocking till all write back complete
while len(self.ongoing_write_through) > 0: while len(self.ongoing_write_through) > 0:
for _, finish_event, ack_list in self.cache_controller.ack_write_queue: for ack in self.cache_controller.ack_write_queue:
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
self._finish_write_through_ack(ack_id, release_lock=False) self._finish_write_through_ack(ack_id, release_lock=False)
self.cache_controller.ack_write_queue.clear() self.cache_controller.ack_write_queue.clear()
assert len(self.ongoing_write_through) == 0 assert len(self.ongoing_write_through) == 0
@@ -952,8 +952,8 @@ class HiRadixCache(RadixCache):
# sequence and deadlocks under TP > 1. (Matches UnifiedRadixCache.) # sequence and deadlocks under TP > 1. (Matches UnifiedRadixCache.)
finish_count = 0 finish_count = 0
if self.pp_rank == 0: if self.pp_rank == 0:
for _, finish_event, ack_list in self.cache_controller.ack_write_queue: for ack in self.cache_controller.ack_write_queue:
if not finish_event.query(): if not ack.finish_event.query():
break break
finish_count += 1 finish_count += 1
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu") finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
@@ -963,17 +963,17 @@ class HiRadixCache(RadixCache):
if finish_count > 0: if finish_count > 0:
logger.debug(f"Process {finish_count} write back operations") logger.debug(f"Process {finish_count} write back operations")
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0) ack = self.cache_controller.ack_write_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
self._finish_write_through_ack(ack_id, release_lock=True) self._finish_write_through_ack(ack_id, release_lock=True)
finish_count -= 1 finish_count -= 1
def loading_check(self): def loading_check(self):
finish_count = 0 finish_count = 0
if self.pp_rank == 0: if self.pp_rank == 0:
for _, finish_event, ack_list in self.cache_controller.ack_load_queue: for ack in self.cache_controller.ack_load_queue:
if not finish_event.query(): if not ack.finish_event.query():
break break
finish_count += 1 finish_count += 1
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu") finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
@@ -983,11 +983,19 @@ class HiRadixCache(RadixCache):
if finish_count > 0: if finish_count > 0:
logger.debug(f"Process {finish_count} load operations") logger.debug(f"Process {finish_count} load operations")
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_load_queue.pop(0) ack = self.cache_controller.ack_load_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
end_node = self.ongoing_load_back.pop(ack_id) end_node = self.ongoing_load_back.pop(ack_id)
self.dec_lock_ref(end_node) self.dec_lock_ref(end_node)
if self.metrics_collector is not None:
self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens)
if ack.timing_enabled:
duration_ms = ack.start_event.elapsed_time(ack.finish_event)
self.metrics_collector.observe_load_back_duration(
duration_ms / 1000.0
)
finish_count -= 1 finish_count -= 1
def is_load_back_event_done(self, consumer_index: int) -> bool: def is_load_back_event_done(self, consumer_index: int) -> bool:
@@ -1243,7 +1251,6 @@ class HiRadixCache(RadixCache):
self, node: TreeNode, mem_quota: Optional[int] = None self, node: TreeNode, mem_quota: Optional[int] = None
) -> Optional[torch.Tensor]: ) -> Optional[torch.Tensor]:
start_time = time.perf_counter()
last_hit_node = node last_hit_node = node
nodes_to_load = [] nodes_to_load = []
while node.evicted: while node.evicted:
@@ -1311,12 +1318,6 @@ class HiRadixCache(RadixCache):
self.evictable_size_ += len(device_indices) self.evictable_size_ += len(device_indices)
self.inc_lock_ref(last_hit_node) self.inc_lock_ref(last_hit_node)
if self.metrics_collector is not None:
self.metrics_collector.observe_load_back_duration(
time.perf_counter() - start_time
)
self.metrics_collector.increment_load_back_num_tokens(len(device_indices))
return device_indices return device_indices
def init_load_back( def init_load_back(
@@ -23,6 +23,9 @@ from sglang.srt.managers.cache_controller import (
from sglang.srt.managers.cache_controller import ( from sglang.srt.managers.cache_controller import (
StorageOperation as BaseStorageOperation, StorageOperation as BaseStorageOperation,
) )
from sglang.srt.managers.cache_controller import (
make_timing_event_pair,
)
from sglang.srt.mem_cache.hicache_storage import ( from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorageExtraInfo, HiCacheStorageExtraInfo,
PoolHitPolicy, PoolHitPolicy,
@@ -490,8 +493,12 @@ class HybridCacheController(BaseHiCacheController):
self.load_queue.clear() self.load_queue.clear()
producer_event = self.layer_done_counter.events[producer_id] producer_event = self.layer_done_counter.events[producer_id]
producer_event.start_event.record() producer_event.start_event.record()
ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair()
with device_module.stream(self.load_stream): with device_module.stream(self.load_stream):
producer_event.start_event.wait(self.load_stream) producer_event.start_event.wait(self.load_stream)
ack_start_event.record()
for i in range(self.layer_num): for i in range(self.layer_num):
self.mem_pool_host.load_to_device_per_layer( self.mem_pool_host.load_to_device_per_layer(
self.mem_pool_device, self.mem_pool_device,
@@ -514,6 +521,7 @@ class HybridCacheController(BaseHiCacheController):
self.io_backend, self.io_backend,
) )
producer_event.complete(i) producer_event.complete(i)
ack_finish_event.record()
self._record_transfer_indices_on_stream( self._record_transfer_indices_on_stream(
self.load_stream, self.load_stream,
host_indices, host_indices,
@@ -522,9 +530,11 @@ class HybridCacheController(BaseHiCacheController):
) )
self.ack_load_queue.append( self.ack_load_queue.append(
HiCacheAck( HiCacheAck(
producer_event.start_event, ack_start_event,
producer_event.finish_event, ack_finish_event,
op.node_ids, op.node_ids,
num_tokens=len(op.device_indices),
timing_enabled=timing_enabled,
) )
) )
return producer_id return producer_id
@@ -1668,7 +1668,6 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
if self.cache_controller is None: if self.cache_controller is None:
return False return False
start_time = time.perf_counter()
host_anchor_params = self.inc_host_lock_ref(best_match_node).to_dec_params() host_anchor_params = self.inc_host_lock_ref(best_match_node).to_dec_params()
# Build KV transfer # Build KV transfer
kv_xfer = self.components[BASE_COMPONENT_TYPE].build_hicache_transfers( kv_xfer = self.components[BASE_COMPONENT_TYPE].build_hicache_transfers(
@@ -1753,12 +1752,6 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
host_anchor_params, host_anchor_params,
) )
if self.metrics_collector is not None:
self.metrics_collector.observe_load_back_duration(
time.perf_counter() - start_time
)
self.metrics_collector.increment_load_back_num_tokens(len(device_indices))
return True return True
def _build_sidecar_transfers( def _build_sidecar_transfers(
@@ -2350,9 +2343,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
if write_back: if write_back:
# Blocking: wait for all pending write-backs # Blocking: wait for all pending write-backs
while self.ongoing_write_through: while self.ongoing_write_through:
for _, finish_event, ack_list in cc.ack_write_queue: for ack in cc.ack_write_queue:
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
if ack_id in self.ongoing_write_through: if ack_id in self.ongoing_write_through:
self._finish_write_through_ack(ack_id) self._finish_write_through_ack(ack_id)
cc.ack_write_queue.clear() cc.ack_write_queue.clear()
@@ -2363,8 +2356,8 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
# diverge across ranks (e.g. write_backup returning 0 on a subset). # diverge across ranks (e.g. write_backup returning 0 on a subset).
finish_count = 0 finish_count = 0
if self.pp_rank == 0: if self.pp_rank == 0:
for _, finish_event, ack_list in cc.ack_write_queue: for ack in cc.ack_write_queue:
if not finish_event.query(): if not ack.finish_event.query():
break break
finish_count += 1 finish_count += 1
@@ -2374,9 +2367,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
# Process completed acks # Process completed acks
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = cc.ack_write_queue.pop(0) ack = cc.ack_write_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
self._finish_write_through_ack(ack_id) self._finish_write_through_ack(ack_id)
finish_count -= 1 finish_count -= 1
@@ -2389,8 +2382,8 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
# diverge across ranks. # diverge across ranks.
finish_count = 0 finish_count = 0
if self.pp_rank == 0: if self.pp_rank == 0:
for _, finish_event, ack_list in cc.ack_load_queue: for ack in cc.ack_load_queue:
if not finish_event.query(): if not ack.finish_event.query():
break break
finish_count += 1 finish_count += 1
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu") finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
@@ -2398,12 +2391,20 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
finish_count = finish_count_tensor.item() finish_count = finish_count_tensor.item()
while finish_count > 0: while finish_count > 0:
_, finish_event, ack_list = cc.ack_load_queue.pop(0) ack = cc.ack_load_queue.pop(0)
finish_event.synchronize() ack.finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack.node_ids:
node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id) node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id)
self.dec_lock_ref(node, lock_params) self.dec_lock_ref(node, lock_params)
self.dec_host_lock_ref(node, host_lock_params) self.dec_host_lock_ref(node, host_lock_params)
if self.metrics_collector is not None:
self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens)
if ack.timing_enabled:
duration_ms = ack.start_event.elapsed_time(ack.finish_event)
self.metrics_collector.observe_load_back_duration(
duration_ms / 1000.0
)
finish_count -= 1 finish_count -= 1
# ---- HiCache: Scheduler Entry Points ---- # ---- HiCache: Scheduler Entry Points ----
@@ -1934,7 +1934,7 @@ class RadixCacheMetricsCollector(_StatLoggerDIMixin):
self.load_back_duration_seconds = Histogram( self.load_back_duration_seconds = Histogram(
name="sglang:load_back_duration_seconds", name="sglang:load_back_duration_seconds",
documentation="Time taken to load memory from CPU to GPU in seconds.", documentation="Time taken to load KV cache from CPU back to GPU in seconds.",
labelnames=labels.keys(), labelnames=labels.keys(),
buckets=bucket_load_back_duration, buckets=bucket_load_back_duration,
) )
@@ -17,6 +17,7 @@ from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
DecodeKVCacheOffloadManager, DecodeKVCacheOffloadManager,
) )
from sglang.srt.disaggregation.kv_events import OffloadedState from sglang.srt.disaggregation.kv_events import OffloadedState
from sglang.srt.managers.cache_controller import HiCacheAck
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=8, stage="base-b", runner_config="1-gpu-small") register_cuda_ci(est_time=8, stage="base-b", runner_config="1-gpu-small")
@@ -280,7 +281,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
8, 8,
) )
manager.cache_controller = MagicMock() manager.cache_controller = MagicMock()
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [7])] manager.cache_controller.ack_write_queue = [
HiCacheAck(None, _FinishedEvent(), [7])
]
manager._trigger_backup = MagicMock(return_value="last_hash") manager._trigger_backup = MagicMock(return_value="last_hash")
manager._check_offload_progress(1) manager._check_offload_progress(1)
@@ -314,7 +317,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4) self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4)
manager.cache_controller.write.assert_called_once() manager.cache_controller.write.assert_called_once()
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [1])] manager.cache_controller.ack_write_queue = [
HiCacheAck(None, _FinishedEvent(), [1])
]
manager._trigger_backup = MagicMock(return_value="last_hash") manager._trigger_backup = MagicMock(return_value="last_hash")
manager._check_offload_progress(1) manager._check_offload_progress(1)
@@ -357,7 +362,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
8, 8,
) )
manager.cache_controller = MagicMock() manager.cache_controller = MagicMock()
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [8])] manager.cache_controller.ack_write_queue = [
HiCacheAck(None, _FinishedEvent(), [8])
]
manager._trigger_backup = MagicMock(return_value="last_hash") manager._trigger_backup = MagicMock(return_value="last_hash")
manager._check_offload_progress(1) manager._check_offload_progress(1)
@@ -387,7 +394,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
12, 12,
) )
manager.cache_controller = MagicMock() manager.cache_controller = MagicMock()
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [9])] manager.cache_controller.ack_write_queue = [
HiCacheAck(None, _FinishedEvent(), [9])
]
manager._trigger_backup = MagicMock(return_value="last_hash") manager._trigger_backup = MagicMock(return_value="last_hash")
manager._check_offload_progress(1) manager._check_offload_progress(1)
@@ -0,0 +1,124 @@
"""Unit tests for the HiCache load-back duration metric."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small")
@unittest.skipUnless(torch.cuda.is_available(), "CUDA required")
class TestLoadBackDurationMetric(CustomTestCase):
def setUp(self):
from sglang.srt.managers import cache_controller as cc
cc._timing_events_supported.cache_clear()
self.cc = cc
def _completed_pair(self, payload_floats=1024 * 1024):
start, finish, timing_enabled = self.cc.make_timing_event_pair()
self.assertTrue(timing_enabled)
stream = torch.cuda.Stream()
start.record()
with torch.cuda.stream(stream):
start.wait(stream)
torch.empty(payload_floats, device="cuda").fill_(0)
finish.record()
torch.cuda.synchronize()
return start, finish
def test_elapsed_time_works(self):
start, finish = self._completed_pair()
self.assertGreater(start.elapsed_time(finish), 0.0)
def test_timing_fallback_uses_dedicated_events(self):
events = []
def create_event(*, enable_timing=False):
if enable_timing:
raise TypeError
event = MagicMock()
events.append(event)
return event
with patch.object(self.cc.device_module, "Event", side_effect=create_event):
self.cc._timing_events_supported.cache_clear()
start, finish, timing_enabled = self.cc.make_timing_event_pair()
self.assertFalse(timing_enabled)
self.assertIs(start, events[0])
self.assertIs(finish, events[1])
self.assertIsNot(start, finish)
def test_loading_check_observes_duration_and_tokens(self):
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
start, finish = self._completed_pair()
ack = self.cc.HiCacheAck(
start,
finish,
node_ids=[1, 2],
num_tokens=1024,
timing_enabled=True,
)
stub = SimpleNamespace(
cache_controller=SimpleNamespace(ack_load_queue=[ack]),
ongoing_load_back={1: object(), 2: object()},
dec_lock_ref=MagicMock(),
metrics_collector=MagicMock(),
pp_rank=0,
_all_reduce=MagicMock(),
)
HiRadixCache.loading_check(stub)
stub.metrics_collector.increment_load_back_num_tokens.assert_called_once_with(
1024
)
stub.metrics_collector.observe_load_back_duration.assert_called_once()
(observed,), _ = stub.metrics_collector.observe_load_back_duration.call_args
self.assertGreater(observed, 0.0)
self.assertEqual(stub.cache_controller.ack_load_queue, [])
def test_loading_check_fallback_when_timing_unsupported(self):
"""On backends without enable_timing, count tokens but skip duration."""
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
start = torch.cuda.Event()
finish = torch.cuda.Event()
start.record()
finish.record()
torch.cuda.synchronize()
ack = self.cc.HiCacheAck(
start_event=start,
finish_event=finish,
node_ids=[7],
num_tokens=512,
timing_enabled=False,
)
stub = SimpleNamespace(
cache_controller=SimpleNamespace(ack_load_queue=[ack]),
ongoing_load_back={7: object()},
dec_lock_ref=MagicMock(),
metrics_collector=MagicMock(),
pp_rank=0,
_all_reduce=MagicMock(),
)
HiRadixCache.loading_check(stub)
stub.metrics_collector.increment_load_back_num_tokens.assert_called_once_with(
512
)
stub.metrics_collector.observe_load_back_duration.assert_not_called()
self.assertEqual(stub.cache_controller.ack_load_queue, [])
if __name__ == "__main__":
unittest.main()
@@ -499,8 +499,8 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
self.assertTrue(loaded) self.assertTrue(loaded)
producer_id = cache.ready_to_load_host_cache() producer_id = cache.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1) self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue): for ack in list(cache.cache_controller.ack_load_queue):
finish_event.synchronize() ack.finish_event.synchronize()
cache.loading_check() cache.loading_check()
def test_kv_events_store_and_remove_full_blocks(self): def test_kv_events_store_and_remove_full_blocks(self):
@@ -2673,8 +2673,8 @@ class UnifiedRadixCacheSuite:
self.assertTrue(loaded) self.assertTrue(loaded)
producer_id = cache.ready_to_load_host_cache() producer_id = cache.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1) self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue): for ack in list(cache.cache_controller.ack_load_queue):
finish_event.synchronize() ack.finish_event.synchronize()
cache.loading_check() cache.loading_check()
return node.component_data[ComponentType.FULL].value return node.component_data[ComponentType.FULL].value
@@ -3127,8 +3127,8 @@ class UnifiedRadixCacheSuite:
def _finish_pending_loads(self, cache): def _finish_pending_loads(self, cache):
producer_id = cache.ready_to_load_host_cache() producer_id = cache.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1) self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue): for ack in list(cache.cache_controller.ack_load_queue):
finish_event.synchronize() ack.finish_event.synchronize()
cache.loading_check() cache.loading_check()
def _match_tokens_for_chain(self, chain): def _match_tokens_for_chain(self, chain):