[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
@@ -17,6 +17,7 @@ from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
DecodeKVCacheOffloadManager,
)
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
register_cuda_ci(est_time=8, stage="base-b", runner_config="1-gpu-small")
@@ -280,7 +281,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
8,
)
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._check_offload_progress(1)
@@ -314,7 +317,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4)
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._check_offload_progress(1)
@@ -357,7 +362,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
8,
)
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._check_offload_progress(1)
@@ -387,7 +394,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
12,
)
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._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)
producer_id = cache.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue):
finish_event.synchronize()
for ack in list(cache.cache_controller.ack_load_queue):
ack.finish_event.synchronize()
cache.loading_check()
def test_kv_events_store_and_remove_full_blocks(self):
@@ -2673,8 +2673,8 @@ class UnifiedRadixCacheSuite:
self.assertTrue(loaded)
producer_id = cache.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue):
finish_event.synchronize()
for ack in list(cache.cache_controller.ack_load_queue):
ack.finish_event.synchronize()
cache.loading_check()
return node.component_data[ComponentType.FULL].value
@@ -3127,8 +3127,8 @@ class UnifiedRadixCacheSuite:
def _finish_pending_loads(self, cache):
producer_id = cache.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue):
finish_event.synchronize()
for ack in list(cache.cache_controller.ack_load_queue):
ack.finish_event.synchronize()
cache.loading_check()
def _match_tokens_for_chain(self, chain):