执行摘要
- 一句话:引入 pp_sync 机制修复 HiCache + PP 崩溃
- 推荐动作:本 PR 核心设计值得精读,特别是
_pp_sync 的设计:利用 isend/recv 避免全局 barrier,减少同步开销。关注 writing_check 中 ack 队列消费的同步策略,这是一个典型的分布式状态一致性问题。如果团队计划支持 HiCache + PP + L3,建议以此 PR 为基础进行扩展。
功能与动机
HiCache 与 PP 同时启用时,由于 PP0 和 PP1 的 scheduler 线程的异步事件队列(ack_load_queue, ack_prefetch_queue 等)缺乏同步,导致消费事件数量不一致,引发 RuntimeError: shape '[3013, -1, 128]' is invalid for input of size 8192000 崩溃(见 PR #22607 讨论)。本 PR 旨在修复该崩溃以支持 PP + HiCache L2 场景。
实现拆解
- 引入
pp_group 分布式通信组:在 CacheInitParams(cache_init_params.py)中添加 pp_cache_group 字段,并在 hiradix_cache.py 和 unified_radix_cache.py 的初始化中保存 self.pp_group。同时将 HiCacheController 等构造函数的 pp_rank/pp_size 参数替换为 pp_group,以集中管理 PP 通信。
- 实现 PP 同步核心机制:在
hiradix_cache.py 和 unified_radix_cache.py 中新增三个私有方法:
_reap_completed_async_work():轮询 work_list 中已完成的 torch.distributed.Work 对象并清理。
_all_reduce(data, tp_reduce_op):仅在 PP rank 0 上执行 TP 组的 all-reduce,然后通过 _pp_sync 将结果传播到后续 PP rank。
_pp_sync(data):使用 torch.distributed.isend/recv 在 PP 管道中逐级传递数据,所有 rank 最终获得相同数据。
- 改造 ack 队列消费逻辑:在
unified_radix_cache.py 的 writing_check 方法中,只在 PP rank 0 上统计 ack_write_queue 的完成事件数,然后通过 _all_reduce_attn_groups 广播到 CP/TP 参与方,确保所有 rank 的状态一致。
- 配套配置调整:在
hybrid_pool_assembler.py 中,build_kv_only_stack、build_hybrid_swa_stack 等函数的签名从 pp_rank/pp_size 改为接受 pp_group 参数,并将该对象传递给 HybridCacheController。scheduler.py 和 kv_cache_builder.py 中添加了 pp_cache_group 的获取和透传。
- 添加端到端测试:新增测试文件
test_unified_radix_cache_hicache_pp_kl.py,在 TP=2、PP=2 配置下启动 Qwen3‑30B 模型,运行 GSM8K 评估并验证 cached_tokens 的合理性。同时修改现有单元测试 test_unified_radix_cache_unittest.py 以适配新的参数。
关键文件:
python/sglang/srt/mem_cache/hiradix_cache.py(模块 层级缓存;类别 source;类型 core-logic;符号 _reap_completed_async_work, _all_reduce, _pp_sync): 核心变更文件,新增 _pp_sync、_all_reduce、_reap_completed_async_work 方法,并保存 pp_group 和 work_list。
python/sglang/srt/mem_cache/unified_radix_cache.py(模块 统一缓存;类别 source;类型 core-logic;符号 _reap_completed_async_work, _all_reduce, _pp_sync): 统一缓存层同样新增相同的 PP 同步方法,并修改 writing_check 中的 ack 消费逻辑,确保 PP rank 0 只在本地统计后广播。
test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_hicache_pp_kl.py(模块 PP测试;类别 test;类型 test-coverage;符号 _assert_pp_decode_cached_tokens, TestUnifiedQwen3HiCachePP, test_gsm8k, setUpClass): 新增端到端测试文件,验证 PP + HiCache + UnifiedRadixCache 下的 GSM8K 准确率和 cached_tokens 正确性。
python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py(模块 缓存组装;类别 source;类型 core-logic): 修改构建函数签名,从 pp_rank/pp_size 改为传递 pp_group,是配置传递的重要环节。
python/sglang/srt/managers/cache_controller.py(模块 缓存控制器;类别 source;类型 entrypoint): HiCacheController 构造函数新增 pp_group 参数,替换 pp_rank/pp_size,同时内部通过 get_pipeline_model_parallel_rank 获取 rank。
关键符号:_reap_completed_async_work, _all_reduce, _pp_sync, _assert_pp_decode_cached_tokens, test_gsm8k
关键源码片段
python/sglang/srt/mem_cache/unified_radix_cache.py
统一缓存层同样新增相同的 PP 同步方法,并修改 writing_check 中的 ack 消费逻辑,确保 PP rank 0 只在本地统计后广播。
# python/sglang/srt/mem_cache/unified_radix_cache.py — PP 同步及 ack 消费改造
def _reap_completed_async_work(self):
"""
轮询 outstanding 异步 work 并清理已完成的项。
"""
count = 0
while count < len(self.work_list) and self.work_list[count].is_completed():
count += 1
if count > 0:
logger.debug(f"Reap {count} completed async work")
self.work_list = self.work_list[count:]
def _all_reduce(self, data: torch.Tensor, tp_reduce_op: torch.distributed.ReduceOp):
if self.pp_rank == 0:
self._all_reduce_attn_groups(data, tp_reduce_op)
self._pp_sync(data)
def _pp_sync(self, data: torch.Tensor) -> None:
if self.pp_size <= 1 or self.pp_group is None:
return
if self.pp_rank > 0:
torch.distributed.recv(
data, group_src=self.pp_rank - 1, group=self.pp_group, tag=2
)
if self.pp_rank + 1 < self.pp_size:
copy_of_data = data.clone()
send_work = torch.distributed.isend(
copy_of_data, group_dst=self.pp_rank + 1, group=self.pp_group, tag=2
)
self.work_list.append(send_work)
# writing_check 中的关键改动:只在 rank 0 统计 ack 完成数
if self.pp_rank == 0:
for _, finish_event, ack_list in cc.ack_write_queue:
if not finish_event.query():
break
finish_count += 1
queue_size = torch.tensor(finish_count, dtype=torch.int, device="cpu")
self._all_reduce_attn_groups(queue_size, torch.distributed.ReduceOp.MAX)
评论区精华
pp_cache_group 命名讨论:hzh0425 提议将新增字段重命名为 attn_pp_cache_group 以与 attn_cp_group 等保持命名一致。stepinto 与 hzh0425、ShangmingCai 讨论后决定保持原样,因为该组用于 cache 层面的 PP 同步,而不限于 attention。
DeepSeek V4 + PP 分离:hzh0425 指出 dsv4 + PP 的修复应在单独的 PR 中处理。stepinto 同意并回退了相关改动。
- pp_cache_group 命名讨论 (design): 作者 stepinto 与 hzh0425、ShangmingCai 讨论后决定保持原样,因为该组用于 cache 层面的 PP 同步,不限于 attention。
- DeepSeek V4 + PP 修复分离 (design): stepinto 同意并回退了 dsv4 相关改动。
风险与影响
-
风险:死锁风险:_pp_sync 使用阻塞的 recv 与非阻塞的 isend,如果调用次序错误(例如 PP1 在 PP0 未发送时提前 recv),可能导致死锁。代码通过仅在 pp_rank > 0 时 recv、仅在 pp_rank + 1 < pp_size 时 isend 保证顺序,但仍需谨慎维护。
性能影响:每次 _all_reduce 会在 PP 间增加一次数据传播(isend/recv + clone),可能增加调度延迟。但该操作仅在 Hierarchical Cache 的路径上调用,频次不高。
L3 未覆盖:本 PR 仅解决 L2,L3(存储后端)与 PP 的组合可能仍有问题,用户需注意。
回归风险:修改了多个类的构造签名(pp_rank/pp_size 替换为 pp_group),依赖这些签名的外部代码(如自定义 cache 组件)需要同步更新。
-
影响:用户:现在可以安全启用 HiCache + PP(L2)而不会崩溃,Multi-turn 场景的 cache hit rate 提升显著(实验显示 Round 4-7 的 hit rate 从 <35% 提升至 >70%)。
系统:引入 PP 同步机制作为新的基础设施,后续可扩展至 L3 支持。调度器线程增加了 _reap_completed_async_work 轮询,但开销极小。
团队:代码可读性良好,但需注意后续开发中保持同步机制一致。
-
风险标记:同步死锁风险, L3 未覆盖, 构造签名变更
关联脉络
参与讨论