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:
co-authored by
huangtingwei9988
parent
fba967ed9c
commit
e31e5319a0
@@ -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,22 +2980,6 @@ 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")
|
|
||||||
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,
|
write_finish_count,
|
||||||
load_finish_count,
|
load_finish_count,
|
||||||
@@ -3004,15 +2990,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
self.loading_check(finish_count=load_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 = {
|
extra_release_counts = {
|
||||||
pool_name: count
|
pool_name: count
|
||||||
for pool_name, count in zip(
|
for pool_name, count in zip(extra_pool_names, storage_queue_sizes[4:])
|
||||||
extra_pool_names,
|
|
||||||
storage_queue_sizes[4:],
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
self._drain_storage_control_queues_impl(
|
self._drain_storage_control_queues_impl(
|
||||||
n_storage_hit=n_storage_hit,
|
n_storage_hit=n_storage_hit,
|
||||||
|
|||||||
@@ -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 (
|
||||||
|
|||||||
Reference in New Issue
Block a user