Prhub

#27285 [HiCache] Fix crash when using PP + HiCache L2

原始 PR 作者 stepinto 合并时间 2026-06-06 16:57 文件变更 10 提交数 6 评论 20 代码增减 +313 / -76

执行摘要

引入 pp_sync 机制修复 HiCache + PP 崩溃

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 场景。

本 PR 核心设计值得精读,特别是 _pp_sync 的设计:利用 isend/recv 避免全局 barrier,减少同步开销。关注 writing_check 中 ack 队列消费的同步策略,这是一个典型的分布式状态一致性问题。如果团队计划支持 HiCache + PP + L3,建议以此 PR 为基础进行扩展。

讨论亮点

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 同意并回退了相关改动。

实现拆解

  1. 引入 pp_group 分布式通信组:在 CacheInitParamscache_init_params.py)中添加 pp_cache_group 字段,并在 hiradix_cache.pyunified_radix_cache.py 的初始化中保存 self.pp_group。同时将 HiCacheController 等构造函数的 pp_rank/pp_size 参数替换为 pp_group,以集中管理 PP 通信。
  2. 实现 PP 同步核心机制:在 hiradix_cache.pyunified_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 最终获得相同数据。
  3. 改造 ack 队列消费逻辑:在 unified_radix_cache.pywriting_check 方法中,只在 PP rank 0 上统计 ack_write_queue 的完成事件数,然后通过 _all_reduce_attn_groups 广播到 CP/TP 参与方,确保所有 rank 的状态一致。
  4. 配套配置调整:在 hybrid_pool_assembler.py 中,build_kv_only_stackbuild_hybrid_swa_stack 等函数的签名从 pp_rank/pp_size 改为接受 pp_group 参数,并将该对象传递给 HybridCacheControllerscheduler.pykv_cache_builder.py 中添加了 pp_cache_group 的获取和透传。
  5. 添加端到端测试:新增测试文件 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 层级缓存 modified 8.15
python/sglang/srt/mem_cache/unified_radix_cache.py 统一缓存 modified 8.05
test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_hicache_pp_kl.py PP 测试 added 7.48
python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py 缓存组装 modified 6.28
python/sglang/srt/managers/cache_controller.py 缓存控制器 modified 5.77

关键符号

_reap_completed_async_work _all_reduce _pp_sync _assert_pp_decode_cached_tokens test_gsm8k

关键源码片段

python/sglang/srt/mem_cache/unified_radix_cache.py core-logic

统一缓存层同样新增相同的 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 命名讨论 设计

reviewer hzh0425 提议将新增的 `pp_cache_group` 字段重命名为 `attn_pp_cache_group` 以与现有 `attn_cp_group` 等保持命名一致性。

结论:作者 stepinto 与 hzh0425、ShangmingCai 讨论后决定保持原样,因为该组用于 cache 层面的 PP 同步,不限于 attention。 · 已解决

DeepSeek V4 + PP 修复分离 设计

hzh0425 指出 DeepSeek V4 的 PP 支持应放在单独的 PR 中处理,不应混入本 PR。

结论: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 未覆盖 构造签名变更

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论