From e31e5319a0f7e3a2692c14ffeef630785c71b543 Mon Sep 17 00:00:00 2001 From: Chao Shi Date: Tue, 8 Sep 2026 10:13:22 +0800 Subject: [PATCH] HiCache: Reduce the number of `all_reduce` in `check_hicache_events` for PP (#37562) Co-authored-by: huangtingwei9988 --- .../srt/mem_cache/unified_radix_cache.py | 85 +++++++------------ .../mem_cache/test_hiradix_pp_sync_drain.py | 16 +++- 2 files changed, 45 insertions(+), 56 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 3eb51fa61..482df86df 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2742,18 +2742,20 @@ class UnifiedRadixCache(BasePrefixCache): self, ) -> tuple[int, int, tuple[int, ...], tuple[PoolName, ...]]: cc = self.cache_controller - if cc is None: + extra_release_queues = getattr(cc, "extra_host_mem_release_queues", {}) + extra_pool_names = tuple(extra_release_queues) if self.enable_storage else () + if cc is None or self.pp_rank > 0: write_acks = 0 load_acks = 0 - storage_queue_sizes = () - extra_pool_names = () + # Zero placeholders shaped like PP0's slots: _pp_sync hands the + # received tensor back in place, so all ranks must build the same + # length or PP1+ would recv into a mismatched buffer. + storage_queue_sizes = ( + (0,) * (4 + len(extra_pool_names)) if self.enable_storage else () + ) else: write_acks = self._count_ready_acks(cc.ack_write_queue) load_acks = self._count_ready_acks(cc.ack_load_queue) - extra_release_queues = getattr(cc, "extra_host_mem_release_queues", {}) - extra_pool_names = ( - tuple(extra_release_queues) if self.enable_storage else () - ) storage_queue_sizes = ( ( cc.prefetch_hit_queue.qsize(), @@ -2783,8 +2785,8 @@ class UnifiedRadixCache(BasePrefixCache): self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN) count_values = list(map(int, ready_counts.tolist())) - assert count_values[-2] == -count_values[-1], ( - "write_back duplicate-reclaim victims diverged across TP ranks" + assert digest == count_values[-2] and digest == -count_values[-1], ( + "write_back duplicate-reclaim victims diverged across PP/TP ranks" ) return ( count_values[0], @@ -2978,50 +2980,29 @@ class UnifiedRadixCache(BasePrefixCache): # Reap the previous round's PP-sync sends before issuing new ones. self._drain_async_work() - if self.pp_size != 1: - finish_counts = torch.zeros(2, dtype=torch.int, device="cpu") - 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: - self.drain_storage_control_queues() - else: - ( - write_finish_count, - load_finish_count, - storage_queue_sizes, - extra_pool_names, - ) = self._sync_hicache_ready_counts() - self.writing_check(finish_count=write_finish_count) - self.loading_check(finish_count=load_finish_count) + ( + write_finish_count, + load_finish_count, + storage_queue_sizes, + extra_pool_names, + ) = self._sync_hicache_ready_counts() + self.writing_check(finish_count=write_finish_count) + self.loading_check(finish_count=load_finish_count) - if self.enable_storage and storage_queue_sizes: - n_storage_hit, n_ack_prefetch, n_backup, n_release = ( - storage_queue_sizes[:4] - ) - extra_release_counts = { - pool_name: count - for pool_name, count in zip( - extra_pool_names, - storage_queue_sizes[4:], - ) - } - self._drain_storage_control_queues_impl( - n_storage_hit=n_storage_hit, - n_ack_prefetch=n_ack_prefetch, - n_backup=n_backup, - n_release=n_release, - extra_release_counts=extra_release_counts, - log_metrics=True, - ) + if self.enable_storage and storage_queue_sizes: + n_storage_hit, n_ack_prefetch, n_backup, n_release = storage_queue_sizes[:4] + extra_release_counts = { + pool_name: count + for pool_name, count in zip(extra_pool_names, storage_queue_sizes[4:]) + } + self._drain_storage_control_queues_impl( + n_storage_hit=n_storage_hit, + n_ack_prefetch=n_ack_prefetch, + n_backup=n_backup, + n_release=n_release, + extra_release_counts=extra_release_counts, + log_metrics=True, + ) if self.buffer_pipeline is not None: self.buffer_pipeline.flush_pending_writes() if self.enable_storage_metrics and self.storage_metrics_collector is not None: diff --git a/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py b/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py index 685028523..8a9da6c56 100644 --- a/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py +++ b/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py @@ -4,6 +4,8 @@ import unittest from types import SimpleNamespace from unittest.mock import MagicMock +import torch + from sglang.srt.mem_cache.hiradix_cache import HiRadixCache from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.test.ci.ci_register import register_cpu_ci @@ -51,7 +53,10 @@ class TestPPSyncDrain(unittest.TestCase): 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.tree_core = SimpleNamespace( + enable_storage=False, + write_back_duplicate_reclaim_digest=0, + ) cache.pp_rank = pp_rank cache.pp_size = 2 cache.enable_storage_metrics = False @@ -62,7 +67,6 @@ class TestUnifiedPPSyncBatching(unittest.TestCase): 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( @@ -84,12 +88,16 @@ class TestUnifiedPPSyncBatching(unittest.TestCase): leader.check_hicache_events() leader._all_reduce.assert_called_once() - self.assertEqual(leader._all_reduce.call_args.args[0].tolist(), [1, 2]) + self.assertEqual(leader._all_reduce.call_args.args[0].tolist(), [1, 2, 0, 0]) 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) + + def reduce_to_min(counts, _): + counts.copy_(torch.tensor([1, 1, 0, 0], dtype=torch.int64)) + + follower._all_reduce.side_effect = reduce_to_min follower.check_hicache_events() for queue in (