[UnifiedTree] Add CP sync (#25395)

Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
Vladislav Nosivskoy
2026-06-03 16:10:27 +08:00
committed by GitHub
co-authored by Zhangheng
parent 93173b27e8
commit 63dc20ae6c
2 changed files with 99 additions and 5 deletions
@@ -304,6 +304,8 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
self.session = StreamingSession(inner=self)
self.tp_group = params.tp_cache_group
self.attn_cp_group = params.attn_cp_cache_group
self.attn_tp_group = params.attn_tp_cache_group
self.tp_world_size = (
1
if self.tp_group is None
@@ -322,6 +324,24 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
self.reset()
logger.info(f"Init Unified RadixTree with components {self.tree_components}")
def _all_reduce_attn_groups(self, tensor: torch.Tensor, op):
reduced = False
for group in (self.attn_cp_group, self.attn_tp_group):
if group is not None and torch.distributed.get_world_size(group=group) > 1:
torch.distributed.all_reduce(tensor, op=op, group=group)
reduced = True
if not reduced and self.tp_world_size > 1:
torch.distributed.all_reduce(tensor, op=op, group=self.tp_group)
def _barrier_attn_groups(self):
waited = False
for group in (self.attn_cp_group, self.attn_tp_group):
if group is not None and torch.distributed.get_world_size(group=group) > 1:
torch.distributed.barrier(group=group)
waited = True
if not waited and self.tp_world_size > 1:
torch.distributed.barrier(group=self.tp_group)
def reset(self) -> None:
self._reset_full()
@@ -2142,12 +2162,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
break
finish_count += 1
# TP sync: MIN across all ranks for consistent tree updates
# Keep cache state transitions identical across CPxTP participants.
queue_size = torch.tensor(finish_count, dtype=torch.int, device="cpu")
if self.tp_world_size > 1:
torch.distributed.all_reduce(
queue_size, op=torch.distributed.ReduceOp.MIN, group=self.tp_group
)
self._all_reduce_attn_groups(queue_size, torch.distributed.ReduceOp.MIN)
finish_count = int(queue_size.item())
# Process completed acks
@@ -0,0 +1,77 @@
import unittest
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.unified_radix_cache_kit import UnifiedRadixTreeTestMixin
from sglang.test.kl_multiturn_utils import get_input_ids
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=400, stage="base-c", runner_config="4-gpu-h100")
QWEN3_30B_MODEL = "Qwen/Qwen3-30B-A3B-FP8"
class TestUnifiedQwen3HiCacheCP(UnifiedRadixTreeTestMixin, CustomTestCase):
"""Qwen3-30B-A3B-FP8 + HiCache + CP + UnifiedRadixCache."""
hicache_io_backend = "direct"
hicache_mem_layout = "page_first_direct"
max_running_requests = 32
kl_threshold = 0.005
gsm8k_threshold = 0.7
mmlu_threshold = 0.7
@classmethod
def setUpClass(cls):
cls.model = QWEN3_30B_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--tp-size",
"4",
"--moe-dp-size",
"1",
"--ep-size",
"4",
"--attn-cp-size",
"2",
"--enable-prefill-context-parallel",
"--mem-fraction-static",
"0.8",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
str(cls.max_running_requests),
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
"--enable-hierarchical-cache",
"--hicache-ratio",
"2",
"--hicache-write-policy",
"write_through",
"--hicache-io-backend",
cls.hicache_io_backend,
"--hicache-mem-layout",
cls.hicache_mem_layout,
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
if __name__ == "__main__":
unittest.main()