[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:
co-authored by
Claude Opus 4.7
vuuihc
Zhiqiang Xie
parent
f8eac995aa
commit
095a817612
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user