HiCache: Reduce the number of all_reduce in check_hicache_events for PP (#37562)

Co-authored-by: huangtingwei9988 <huangtingwei.htw@antgroup.com>
This commit is contained in:
Chao Shi
2026-09-08 10:13:22 +08:00
committed by GitHub
co-authored by huangtingwei9988
parent fba967ed9c
commit e31e5319a0
2 changed files with 45 additions and 56 deletions
@@ -2742,18 +2742,20 @@ class UnifiedRadixCache(BasePrefixCache):
self, self,
) -> tuple[int, int, tuple[int, ...], tuple[PoolName, ...]]: ) -> tuple[int, int, tuple[int, ...], tuple[PoolName, ...]]:
cc = self.cache_controller 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 write_acks = 0
load_acks = 0 load_acks = 0
storage_queue_sizes = () # Zero placeholders shaped like PP0's slots: _pp_sync hands the
extra_pool_names = () # 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: else:
write_acks = self._count_ready_acks(cc.ack_write_queue) write_acks = self._count_ready_acks(cc.ack_write_queue)
load_acks = self._count_ready_acks(cc.ack_load_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 = ( storage_queue_sizes = (
( (
cc.prefetch_hit_queue.qsize(), cc.prefetch_hit_queue.qsize(),
@@ -2783,8 +2785,8 @@ class UnifiedRadixCache(BasePrefixCache):
self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN) self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN)
count_values = list(map(int, ready_counts.tolist())) count_values = list(map(int, ready_counts.tolist()))
assert count_values[-2] == -count_values[-1], ( assert digest == count_values[-2] and digest == -count_values[-1], (
"write_back duplicate-reclaim victims diverged across TP ranks" "write_back duplicate-reclaim victims diverged across PP/TP ranks"
) )
return ( return (
count_values[0], count_values[0],
@@ -2978,50 +2980,29 @@ class UnifiedRadixCache(BasePrefixCache):
# Reap the previous round's PP-sync sends before issuing new ones. # Reap the previous round's PP-sync sends before issuing new ones.
self._drain_async_work() self._drain_async_work()
if self.pp_size != 1: (
finish_counts = torch.zeros(2, dtype=torch.int, device="cpu") write_finish_count,
if self.pp_rank == 0 and self.cache_controller is not None: load_finish_count,
finish_counts[0] = self._count_ready_acks( storage_queue_sizes,
self.cache_controller.ack_write_queue extra_pool_names,
) ) = self._sync_hicache_ready_counts()
finish_counts[1] = self._count_ready_acks( self.writing_check(finish_count=write_finish_count)
self.cache_controller.ack_load_queue self.loading_check(finish_count=load_finish_count)
)
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)
if self.enable_storage and storage_queue_sizes: if self.enable_storage and storage_queue_sizes:
n_storage_hit, n_ack_prefetch, n_backup, n_release = ( n_storage_hit, n_ack_prefetch, n_backup, n_release = storage_queue_sizes[:4]
storage_queue_sizes[:4] extra_release_counts = {
) pool_name: count
extra_release_counts = { for pool_name, count in zip(extra_pool_names, storage_queue_sizes[4:])
pool_name: count }
for pool_name, count in zip( self._drain_storage_control_queues_impl(
extra_pool_names, n_storage_hit=n_storage_hit,
storage_queue_sizes[4:], n_ack_prefetch=n_ack_prefetch,
) n_backup=n_backup,
} n_release=n_release,
self._drain_storage_control_queues_impl( extra_release_counts=extra_release_counts,
n_storage_hit=n_storage_hit, log_metrics=True,
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: if self.buffer_pipeline is not None:
self.buffer_pipeline.flush_pending_writes() self.buffer_pipeline.flush_pending_writes()
if self.enable_storage_metrics and self.storage_metrics_collector is not None: if self.enable_storage_metrics and self.storage_metrics_collector is not None:
@@ -4,6 +4,8 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock
import torch
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
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -51,7 +53,10 @@ class TestPPSyncDrain(unittest.TestCase):
class TestUnifiedPPSyncBatching(unittest.TestCase): class TestUnifiedPPSyncBatching(unittest.TestCase):
def _make_cache(self, pp_rank, write_ready, load_ready): def _make_cache(self, pp_rank, write_ready, load_ready):
cache = object.__new__(UnifiedRadixCache) 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_rank = pp_rank
cache.pp_size = 2 cache.pp_size = 2
cache.enable_storage_metrics = False cache.enable_storage_metrics = False
@@ -62,7 +67,6 @@ class TestUnifiedPPSyncBatching(unittest.TestCase):
cache._all_reduce = MagicMock() cache._all_reduce = MagicMock()
cache.writing_check = MagicMock() cache.writing_check = MagicMock()
cache.loading_check = MagicMock() cache.loading_check = MagicMock()
cache.drain_storage_control_queues = MagicMock()
cache.cache_controller = SimpleNamespace( cache.cache_controller = SimpleNamespace(
ack_write_queue=[ ack_write_queue=[
SimpleNamespace( SimpleNamespace(
@@ -84,12 +88,16 @@ class TestUnifiedPPSyncBatching(unittest.TestCase):
leader.check_hicache_events() leader.check_hicache_events()
leader._all_reduce.assert_called_once() 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.writing_check.assert_called_once_with(finish_count=1)
leader.loading_check.assert_called_once_with(finish_count=2) leader.loading_check.assert_called_once_with(finish_count=2)
follower = self._make_cache(1, [True], [True]) 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() follower.check_hicache_events()
for queue in ( for queue in (