[HiCache] Batch PP write and load completion sync (#33473)
This commit is contained in:
@@ -2568,8 +2568,18 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
self._drain_async_work()
|
self._drain_async_work()
|
||||||
|
|
||||||
if self.pp_size != 1:
|
if self.pp_size != 1:
|
||||||
self.writing_check()
|
finish_counts = torch.zeros(2, dtype=torch.int, device="cpu")
|
||||||
self.loading_check()
|
if self.pp_rank == 0 and self.cache_controller is not None:
|
||||||
|
finish_counts[0] = self._count_ready_acks(
|
||||||
|
self.cache_controller.ack_write_queue
|
||||||
|
)
|
||||||
|
finish_counts[1] = self._count_ready_acks(
|
||||||
|
self.cache_controller.ack_load_queue
|
||||||
|
)
|
||||||
|
self._all_reduce(finish_counts, torch.distributed.ReduceOp.MIN)
|
||||||
|
write_finish_count, load_finish_count = map(int, finish_counts.tolist())
|
||||||
|
self.writing_check(finish_count=write_finish_count)
|
||||||
|
self.loading_check(finish_count=load_finish_count)
|
||||||
if self.enable_storage:
|
if self.enable_storage:
|
||||||
self.drain_storage_control_queues()
|
self.drain_storage_control_queues()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
"""Unit test for HiRadixCache._drain_async_work PP-sync backpressure."""
|
"""Unit tests for HiCache PP synchronization."""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
|
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
|
||||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||||
@@ -46,5 +48,57 @@ class TestPPSyncDrain(unittest.TestCase):
|
|||||||
self.assertEqual(holder.work_list, [])
|
self.assertEqual(holder.work_list, [])
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedPPSyncBatching(unittest.TestCase):
|
||||||
|
def _make_cache(self, pp_rank, write_ready, load_ready):
|
||||||
|
cache = object.__new__(UnifiedRadixCache)
|
||||||
|
cache.tree_core = SimpleNamespace(enable_storage=False)
|
||||||
|
cache.pp_rank = pp_rank
|
||||||
|
cache.pp_size = 2
|
||||||
|
cache.enable_storage_metrics = False
|
||||||
|
cache.storage_metrics_collector = None
|
||||||
|
cache._drain_async_work = MagicMock()
|
||||||
|
cache._all_reduce = MagicMock()
|
||||||
|
cache.writing_check = MagicMock()
|
||||||
|
cache.loading_check = MagicMock()
|
||||||
|
cache.drain_storage_control_queues = MagicMock()
|
||||||
|
cache.cache_controller = SimpleNamespace(
|
||||||
|
ack_write_queue=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_event=SimpleNamespace(query=MagicMock(return_value=ready))
|
||||||
|
)
|
||||||
|
for ready in write_ready
|
||||||
|
],
|
||||||
|
ack_load_queue=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_event=SimpleNamespace(query=MagicMock(return_value=ready))
|
||||||
|
)
|
||||||
|
for ready in load_ready
|
||||||
|
],
|
||||||
|
)
|
||||||
|
return cache
|
||||||
|
|
||||||
|
def test_pp_batches_write_and_load_counts_once(self):
|
||||||
|
leader = self._make_cache(0, [True, False], [True, True])
|
||||||
|
leader.check_hicache_events()
|
||||||
|
|
||||||
|
leader._all_reduce.assert_called_once()
|
||||||
|
self.assertEqual(leader._all_reduce.call_args.args[0].tolist(), [1, 2])
|
||||||
|
leader.writing_check.assert_called_once_with(finish_count=1)
|
||||||
|
leader.loading_check.assert_called_once_with(finish_count=2)
|
||||||
|
|
||||||
|
follower = self._make_cache(1, [True], [True])
|
||||||
|
follower._all_reduce.side_effect = lambda counts, _: counts.fill_(1)
|
||||||
|
follower.check_hicache_events()
|
||||||
|
|
||||||
|
for queue in (
|
||||||
|
follower.cache_controller.ack_write_queue,
|
||||||
|
follower.cache_controller.ack_load_queue,
|
||||||
|
):
|
||||||
|
queue[0].finish_event.query.assert_not_called()
|
||||||
|
follower._all_reduce.assert_called_once()
|
||||||
|
follower.writing_check.assert_called_once_with(finish_count=1)
|
||||||
|
follower.loading_check.assert_called_once_with(finish_count=1)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user