Prhub

#30756 Integrate pplx a2a backend

原始 PR 作者 trangdough 合并时间 2026-07-31 06:33 文件变更 16 提交数 16 评论 26 代码增减 +788 / -10

执行摘要

集成 pplx-kernels 作为新的 MoE all-to-all 后端

支持 pplx-kernels 作为 MoE all-to-all 后端,利用 NVSHMEM 实现专家间高效同步,提升 MoE 推理性能。PR 正文中给出了详细基准数据:相比默认后端,output token throughput 从 6085 tok/s 提升到 7157 tok/s。

  • 此 PR 值得精读,特别是 pplx.pyPplxAllToAllManager 如何管理 NVSHMEM 和 c10d 组生命周期,以及如何适配 BaseDispatcher 接口。
  • 注意 server_args.py 中的启动校验(见关键源码片段)是处理第三方后端集成的典范,包括模式自动降级(auto→low_latency)和提前断言失败。
  • dp_attention.py 中强制 MAX_LEN 的修复也是一个重要的设计决策——对称集合通信要求所有 rank 一致 token 数。
讨论亮点
  • NVSHMEM 初始化一致性:ch-wan 指出 nvshmem_init 使用默认 global rank/size 而 AllToAll 使用 EP group,当 EP 为 subset 时可能 hang。作者改为统一使用 group.rank()/size(),并断言 ep_size == world_size,暂时不支持 PP。
  • DP 组数校验:ch-wan 发现仅检查 ep_size >= 2 不足以确保 numDPGroups > 1,缺少 enable_dp_attention 时 AllToAll 构造会静默失败。最终添加 assert enable_dp_attention and dp_size >= 2
  • Per-rank token 上限:ch-wan 指出默认 SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK=128 远小于实际 chunked prefill 大小,会导致运行时断言失败。作者新增 _required_pplx_dispatch_tokens_per_rank 计算最坏情况并启动时断言。
  • Internode 检测:ch-wan 认为 world_size > torch.cuda.device_count() 不可靠。作者以当前配置(ep == world)下改用 server_args.nnodes > 1 作为 future work。
  • CUDA Graph 与 buffer 分配:gemini-code-assist 建议预分配输出缓冲区以减少开销,但作者指出 CUDA graph capture 需要固定地址,每次分配 torch.zeros 是 graph 兼容的预要求,因此保留现状。
  • 测试移至 manual:ch-wan 注意到 pplx-kernels 不在 CI 环境,建议将注册测试移至手动目录,作者已执行。

实现拆解

  1. 配置键扩展:在 python/sglang/srt/server_args.py 中将 "pplx" 加入 MOE_A2A_BACKEND_CHOICESmoe_a2a_backend 的 Literal 类型,并在 _handle_a2a_moe 中添加 pplx 专属校验:强制 deepep_mode="low_latency"、要求 enable_dp_attentiondp_size >= 2、仅允许 moe_runner_backend="deep_gemm"auto(auto 会重写为 deep_gemm)。同时新增 _required_pplx_dispatch_tokens_per_rank 方法计算最大每 rank 分发 token 数,并与环境变量 SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK 比较,提供早期失败。

  2. 环境变量定义:在 python/sglang/srt/environ.py 中引入 SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK(默认 128),控制 pplx 内核预分配的最大 dispatch token 数,防止分配超出显存。

  3. 派发器实现:新增 python/sglang/srt/layers/moe/token_dispatcher/pplx.py,包含 PplxAllToAllManager 单例(管理 NVSHMEM 初始化和 c10d 组注册)、PplxDispatchOutput / PplxCombineInput 数据类型(复用 DeepEP 的格式),以及核心派发器类 _PplxDispatcherImpldispatch_a 方法负责将 token 按 top-1 路由分发到对应专家;combine_a 通过 AllToAll 操作收集专家输出并执行 combine。NVSHMEM 初始化与 EP group 的 rank/size 一致。

  4. 集成调整:修改 python/sglang/srt/layers/dp_attention.pyget_dp_padding_mode,当使用 pplx 后端时强制返回 DpPaddingMode.MAX_LEN,避免空闲 DP rank 因 token 数不匹配导致 deadlock。修改 python/sglang/srt/layers/moe/utils.py:在 MoeA2ABackend 枚举中添加 PPLX 值和 is_pplx() 方法;更新 is_deepep_class_backend() 包含 pplx;在 should_skip_post_experts_all_reduce() 中为 pplx 返回 True(因其 combine 已包含求和)。修改 two_batch_overlap.py 增加 pplx 分派器实例化。

  5. 测试与文档:新增 test/manual/ep/test_pplx_small.py(手动运行,需要 4x H100 和 pplx-kernels 安装),覆盖纯 DP 和 Hybrid DP+TP 场景的 GSM8K 精度测试。更新 docs/advanced_features/expert_parallelism.md 在支持矩阵中加入 pplx 条目。删除多余测试文件 test/registered/unit/layers/moe/test_should_skip_post_experts_all_reduce.py

文件 模块 状态 重要度
python/sglang/srt/layers/moe/token_dispatcher/pplx.py MoE 调度器 added 9.25
test/manual/ep/test_pplx_small.py 手动测试 added 7.89
python/sglang/srt/server_args.py 服务配置 modified 7.27
python/sglang/srt/layers/moe/utils.py MoE 工具 modified 6.54
python/sglang/srt/layers/dp_attention.py DP 注意力 modified 5.8

关键符号

_ensure_nvshmem _register_group get_all_to_all dispatch_a combine_a is_pplx _required_pplx_dispatch_tokens_per_rank

关键源码片段

python/sglang/srt/layers/moe/token_dispatcher/pplx.py dependency-wiring

新增的 pplx MoE 派发器核心文件,包含 PplxAllToAllManager(管理 NVSHMEM 和 c10d 组)和 _PplxDispatcherImpl(实现 dispatch/combine 逻辑)。

# pplx.py: PplxAllToAllManager 类,管理 NVSHMEM 初始化和 c10d 组注册class PplxAllToAllManager:
    _nvshmem_initialized = False
    _all_to_all: Optional[AllToAll] = None
    _key: Optional[tuple] = None
    _group_name: Optional[str] = None
    _GROUP_NAME = "pplx_ep" # 供 pplx-kernels 通过 resolve_process_group() 查找
​
    @classmethod
    def _ensure_nvshmem(cls, group: dist.ProcessGroup) -> None:
        """初始化 NVSHMEM,确保与 EP group 使用一致的 rank/size。"""
        if cls._nvshmem_initialized:
            return
        # 目前仅支持 EP group 等于全局 world size ( 无 PP 或 EP subset)
        assert group.size() == dist.get_world_size(), (
            "moe_a2a_backend='pplx' requires the EP group to span the whole "
            f"world (got ep_size={group.size()}, world_size={dist.get_world_size()}); "
            "pipeline parallelism and EP-subset layouts are not supported."
        )
        global_rank = group.rank()
        world_size = group.size()
        device = torch.device("cuda", torch.cuda.current_device())
        local_rank = torch.cuda.current_device()
        nvshmem_init(
            global_rank=global_rank,
            local_rank=local_rank,
            world_size=world_size,
            device=device,
        )
        cls._nvshmem_initialized = True
​
    @classmethod
    def _register_group(cls, group: dist.ProcessGroup) -> str:
        """注册 EP process group 到 c10d,使 pplx-kernels 可解析。"""
        if cls._group_name is not None:
            return cls._group_name
        ranks = dist.get_process_group_ranks(group)
        # 创建 hybrid group (cpu:gloo + cuda:nccl)
        combined = dist.new_group(ranks=ranks, backend="cpu:gloo,cuda:nccl")
        torch._C._distributed_c10d._register_process_group(cls._GROUP_NAME, combined)
        cls._group_name = cls._GROUP_NAME
        return cls._group_name
python/sglang/srt/server_args.py core-logic

服务参数配置,添加 "pplx" 选项、验证逻辑和 dispatch token 上限检查。

# server_args.py: pplx 后端验证和 dispatch token 计算方法# 在 _handle_a2a_moe 中的 pplx 分支 :
if a2a_backend == "pplx":
    # 强制 low_latency 模式
    if self.deepep_mode == "normal":
        raise ValueError(
            "moe_a2a_backend='pplx' only supports low-latency mode; "
            "set --deepep-mode to 'low_latency' or 'auto'."
        )
    if self.deepep_mode == "auto":
        self.deepep_mode = "low_latency"
        logger.warning("auto set deepep_mode=`low_latency` for PPLX EP")
​
    # 要求 DP attention 且至少 2 个 DP group
    assert resolved_view(self).enable_dp_attention and self.dp_size >= 2, (
        "moe_a2a_backend='pplx' requires --enable-dp-attention with at "
        "least 2 DP groups (--dp-size >= 2)."
    )
​
    # 仅支持 deep_gemm runner,自动降级
    assert resolved_view(self).moe_runner_backend in ("deep_gemm", "auto"), (
        "moe_a2a_backend='pplx' is only supported with --moe-runner-backend "
        "deep_gemm (or auto)."
    )
    if self.moe_runner_backend == "auto":
        self.moe_runner_backend = "deep_gemm"
        logger.warning("auto set moe_runner_backend=`deep_gemm` for PPLX EP")
​
    # 检查每 rank dispatch token 上限是否足够
    if self.chunked_prefill_size > 0 and self.disaggregation_mode != "decode":
        assert self._required_pplx_dispatch_tokens_per_rank() <= \
            envs.SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get(), (
                "SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK (default 128) "
                "must be >= the per-rank pplx dispatch tokens "
                "(chunked_prefill_size, or the decode cuda-graph batch size)"
            )def _required_pplx_dispatch_tokens_per_rank(self) -> int:
    """计算每 rank 可能的最大 dispatch token 数: max(chunked_prefill_size, cuda_graph_max_bs_decode)。"""
    required = self.chunked_prefill_size
    if self.cuda_graph_max_bs_decode is not None:
        required = max(required, self.cuda_graph_max_bs_decode)
    return required

评论区精华

NVSHMEM 初始化与 EP 组一致性 正确性

ch-wan 指出 nvshmem_init 使用默认 world rank/size,而 AllToAll 使用 EP group,当 EP 子集时可能 hang。

结论:作者改为使用 group.rank()/size() 初始化 NVSHMEM,使两者一致,并临时断言 ep_size == world_size。 · 已解决

DP 组数校验 正确性

ch-wan 发现仅检查 ep_size >= 2 不足以确保 numDPGroups > 1,缺少 enable_dp_attention 时 AllToAll 构造静默失败。

结论:添加 assert enable_dp_attention and dp_size >= 2。 · 已解决

Per-rank dispatch token 上限检查 正确性

ch-wan 指出默认 env 值 128 远小于 chunked_prefill_size,会导致运行时断言。

结论:新增 _required_pplx_dispatch_tokens_per_rank 方法并在 server_args 中启动时断言。 · 已解决

CUDA Graph 与 buffer 分配策略 性能

gemini-code-assist 建议预分配输出缓冲区减少开销,但作者指出 CUDA graph capture 需要分配固定地址,不能复用。

结论:保留每次分配 torch.zeros 的方式,因为 graph capture 期间 buffer 地址可变不可用。 · resolved (won't fix)

Internode vs intranode 检测 正确性

ch-wan 认为 world_size > torch.cuda.device_count() 不可靠。

结论:作者改用 server_args.nnodes > 1,当前配置 ep==world 时有效,未来改进。 · acknowledged

测试移至 manual 目录 测试

ch-wan 建议将 test_pplx_small.py 移至 manual 因为 pplx-kernels 不在 CI。

结论:执行移动,并删除一个无关测试文件。 · 已解决

风险与影响

  • 测试覆盖风险:pplx-kernels 专有,CI 无法运行任何功能测试,回归需依赖作者手动执行。文档强调仅内部运行。
  • NVSHMEM 初始化假设:当前假设 EP group 等于全局 world size,不支持 PP 或多副本布局。如果未来引入这些模式,代码会失败(有 assert 但可能阻塞)。
  • 仅限特定硬件:pplx 仅支持 Hopper(sm_90a)并依赖 deep_gemm runner。配置错误(如使用 triton runner)会导致运行时深层断言,虽然 server_args 做了前置校验。
  • 硬编码 params_bytes=2pplx.py 中 hard-code 为 2,若 params_dtype 变为 float32 则计算错误。虽当前 BF16 路径正确,但类型不泛化。
  • CUDA graph 兼容性:由于每次分配新的 buffer,CUDA graph capture 可能会尝试捕获分配操作?但原码静态分配了固定大小的 buffer 通过 torch.zeros 并且开启了 graph,所以应该没问题。但 buffer 尺寸固定,超过则失败;且分配代价集中在首次 capture,后续 replay 不分配。因此风险低。
  • 用户影响:对于使用 DP attention(--enable-dp-attention--dp >= 2)且 MoE runner 为 deep_gemm 的部署,启用 --moe-a2a-backend pplx 可获得 15-18% 的 output throughput 提升(基于 PR 中 4x H100 实测)。但需要手动安装 pplx-kernels 库,否则导入失败并降级为 use_pplx=False
  • 系统架构:新增的 PplxAllToAllManager 单例模式可供未来其他 NVSHMEM-based 后端复用。结合 is_deepep_class_backend() 等多态检查,pplx 被视为 DeepEP 家族成员,共享 should_skip_post_experts_all_reduceuses_per_rank_fused_shared_slots 等逻辑。
  • 团队维护:代码主要集中在 pplx.py(527 行),与现有 token dispatcher 抽象(BaseDispatcher)耦合,通过工厂方法创建,侵入性较低。
测试覆盖不足(CI 不可用) 仅支持 Hopper + deep_gemm NVSHMEM 初始化假设 硬编码 params_bytes=2

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论