执行摘要
- 一句话:集成 pplx-kernels 作为新的 MoE all-to-all 后端
- 推荐动作:
- 此 PR 值得精读,特别是
pplx.py 中 PplxAllToAllManager 如何管理 NVSHMEM 和 c10d 组生命周期,以及如何适配 BaseDispatcher 接口。
- 注意
server_args.py 中的启动校验(见关键源码片段)是处理第三方后端集成的典范,包括模式自动降级(auto→low_latency)和提前断言失败。
dp_attention.py 中强制 MAX_LEN 的修复也是一个重要的设计决策——对称集合通信要求所有 rank 一致 token 数。
功能与动机
支持 pplx-kernels 作为 MoE all-to-all 后端,利用 NVSHMEM 实现专家间高效同步,提升 MoE 推理性能。PR 正文中给出了详细基准数据:相比默认后端,output token throughput 从 6085 tok/s 提升到 7157 tok/s。
实现拆解
-
配置键扩展:在 python/sglang/srt/server_args.py 中将 "pplx" 加入 MOE_A2A_BACKEND_CHOICES 和 moe_a2a_backend 的 Literal 类型,并在 _handle_a2a_moe 中添加 pplx 专属校验:强制 deepep_mode="low_latency"、要求 enable_dp_attention 且 dp_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 比较,提供早期失败。
-
环境变量定义:在 python/sglang/srt/environ.py 中引入 SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK(默认 128),控制 pplx 内核预分配的最大 dispatch token 数,防止分配超出显存。
-
派发器实现:新增 python/sglang/srt/layers/moe/token_dispatcher/pplx.py,包含 PplxAllToAllManager 单例(管理 NVSHMEM 初始化和 c10d 组注册)、PplxDispatchOutput / PplxCombineInput 数据类型(复用 DeepEP 的格式),以及核心派发器类 _PplxDispatcherImpl。dispatch_a 方法负责将 token 按 top-1 路由分发到对应专家;combine_a 通过 AllToAll 操作收集专家输出并执行 combine。NVSHMEM 初始化与 EP group 的 rank/size 一致。
-
集成调整:修改 python/sglang/srt/layers/dp_attention.py 中 get_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 分派器实例化。
-
测试与文档:新增 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 调度器;类别 source;类型 dependency-wiring;符号 PplxDispatchOutput, format, PplxCombineInput, PplxAllToAllManager): 新增的 pplx MoE 派发器核心文件,包含 PplxAllToAllManager(管理 NVSHMEM 和 c10d 组)和 _PplxDispatcherImpl(实现 dispatch/combine 逻辑)。
test/manual/ep/test_pplx_small.py(模块 手动测试;类别 test;类型 test-coverage;符号 TestPureDP, setUpClass, tearDownClass, test_gsm8k): 手动端到端测试,验证 pplx 后端在纯 DP 和 Hybrid DP+TP 配置下的 GSM8K 精度。
python/sglang/srt/server_args.py(模块 服务配置;类别 source;类型 core-logic;符号 _required_pplx_dispatch_tokens_per_rank): 服务参数配置,添加 "pplx" 选项、验证逻辑和 dispatch token 上限检查。
python/sglang/srt/layers/moe/utils.py(模块 MoE 工具;类别 source;类型 core-logic;符号 is_pplx): MoE 工具函数,新增 PPLX 枚举值和 is_pplx() 方法,更新 is_deepep_class_backend(),控制 skip_post_experts_all_reduce。
python/sglang/srt/layers/dp_attention.py(模块 DP 注意力;类别 source;类型 dependency-wiring): DP attention padding 模式调整,强制 MAX_LEN 以避免空闲 rank 导致 NVSHMEM deadlock。
关键符号:_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
新增的 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
服务参数配置,添加 "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
评论区精华
风险与影响
- 风险:
- 测试覆盖风险:pplx-kernels 专有,CI 无法运行任何功能测试,回归需依赖作者手动执行。文档强调仅内部运行。
- NVSHMEM 初始化假设:当前假设 EP group 等于全局 world size,不支持 PP 或多副本布局。如果未来引入这些模式,代码会失败(有 assert 但可能阻塞)。
- 仅限特定硬件:pplx 仅支持 Hopper(sm_90a)并依赖 deep_gemm runner。配置错误(如使用 triton runner)会导致运行时深层断言,虽然 server_args 做了前置校验。
- 硬编码 params_bytes=2:
pplx.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_reduce、uses_per_rank_fused_shared_slots 等逻辑。
- 团队维护:代码主要集中在
pplx.py(527 行),与现有 token dispatcher 抽象(BaseDispatcher)耦合,通过工厂方法创建,侵入性较低。
- 风险标记:测试覆盖不足(CI 不可用), 仅支持 Hopper + deep_gemm, NVSHMEM 初始化假设, 硬编码 params_bytes=2
关联脉络
参与讨论