Prhub

#21944 Fix:fix(timeout): fix timeout not propagated

原始 PR 作者 JasonHe-WQ 合并时间 2026-04-23 03:48 文件变更 1 提交数 2 评论 2 代码增减 +15 / -3

执行摘要

修复分布式超时参数未传递到模型并行子组的问题

根据关联Issue #21911的描述,当前--dist-timeout参数仅传递给torch.distributed.init_process_group,但SGLang后续为模型并行(如TP、ATTN_TP、MOE_DP等)创建的NCCL子组未继承此超时设置。这导致即使设置了超时(如3600秒),子组集合操作仍可能因使用PyTorch NCCL后端默认的10分钟超时而失败,引发类似“Watchdog caught collective operation timeout”的错误。

该PR是重要的bugfix,涉及分布式核心路径,值得精读以理解超时传递机制。关注_MODEL_PARALLEL_GROUP_TIMEOUT全局变量的引入和传递方式,以及不同后端(mooncake、默认NCCL)的统一处理。

讨论亮点

reviewer gemini-code-assist[bot]指出,在mooncake后端中,device_group已添加超时参数,但mooncake-cpu组尚未应用。作者在后续提交中已修复此问题,确保所有子组(包括mooncake-cpu)都接收subgroup_timeout。reviewer ch-wan简单批准(LGTM),表明变更被接受。

实现拆解

  1. 引入全局超时变量:在python/sglang/srt/distributed/parallel_state.py中新增模块级变量_MODEL_PARALLEL_GROUP_TIMEOUT,用于存储用户提供的超时值,类型为Optional[timedelta],初始化为None
  2. 在初始化时设置超时:在init_distributed_environment函数中,当初始化全局进程组时,将传入的timeout参数转换为timedelta后赋值给_MODEL_PARALLEL_GROUP_TIMEOUT,确保后续子组创建可复用。
  3. 在子组创建时传递超时:在_ModelParallelGroup类的__init__方法中,为每个子组(包括mooncake、mooncake-cpu及默认后端)的torch.distributed.new_group调用添加timeout=subgroup_timeout参数,其中subgroup_timeout从全局变量获取。
  4. 清理时重置变量:在destroy_distributed_environment函数中,将_MODEL_PARALLEL_GROUP_TIMEOUT重置为None,避免残留状态影响后续运行。
文件 模块 状态 重要度
python/sglang/srt/distributed/parallel_state.py 分布式 modified 6.51

关键符号

init_distributed_environment _ModelParallelGroup.__init__ destroy_distributed_environment

关键源码片段

python/sglang/srt/distributed/parallel_state.py core-logic

这是唯一变更的文件,包含分布式并行状态管理的核心逻辑,负责进程组初始化和子组创建。

# 在模块顶部定义全局超时变量,用于存储用户提供的超时值
_MODEL_PARALLEL_GROUP_TIMEOUT: Optional[timedelta] = Nonedef init_distributed_environment(
    world_size: int,
    rank: int,
    local_rank: int,
    distributed_init_method: Optional[str],
    backend: str,
    timeout: Optional[int] = None,
    # ... 其他参数
):
    # ... 其他初始化代码
    if not torch.distributed.is_initialized():
        global _MODEL_PARALLEL_GROUP_TIMEOUT # 声明全局变量以修改
        # ... 参数验证
        if timeout is not None:
            assert isinstance(timeout, (int)), "timeout must be a number"
            assert timeout > 0, "timeout must be positive"
            timeout = timedelta(seconds=timeout) # 转换为 timedelta
            _MODEL_PARALLEL_GROUP_TIMEOUT = timeout # 存储到全局变量
        # ... 初始化全局进程组class _ModelParallelGroup:
    def __init__(
        self,
        group_ranks: List[List[int]],
        torch_distributed_backend: str,
        gloo_timeout: timedelta,
        # ... 其他参数
    ):
        # ... 设备初始化代码
        for ranks in group_ranks:
            active_ranks = torch.ones(len(ranks), dtype=torch.int32, device=self.device)
            active_ranks_cpu = torch.ones(len(ranks), dtype=torch.int32)
            subgroup_timeout = _MODEL_PARALLEL_GROUP_TIMEOUT # 从全局变量获取超时
            if "mooncake" in torch_distributed_backend:
                from mooncake.ep import MooncakeBackendOptions
                # 为 mooncake 后端设备组和 CPU 组都传递超时
                device_group = torch.distributed.new_group(
                    ranks,
                    backend="mooncake",
                    pg_options=MooncakeBackendOptions(active_ranks),
                    timeout=subgroup_timeout, # 传递超时
                )
                cpu_group = torch.distributed.new_group(
                    ranks,
                    backend="mooncake-cpu",
                    pg_options=MooncakeBackendOptions(active_ranks_cpu),
                    timeout=subgroup_timeout, # 传递超时
                )
            else:
                pg_options = get_torch_distributed_pg_options(group_name)
                # 为默认后端设备组传递超时
                device_group = torch.distributed.new_group(
                    ranks,
                    backend=torch_distributed_backend,
                    pg_options=pg_options,
                    timeout=subgroup_timeout, # 传递超时
                )
                # CPU 组使用 gloo 后端,已有 gloo_timeout 参数
                cpu_group = torch.distributed.new_group(
                    ranks, backend="gloo", timeout=gloo_timeout
                )
            # ... 组属性设置def destroy_distributed_environment():
    global _WORLD, _MODEL_PARALLEL_GROUP_TIMEOUT # 清理时重置全局变量
    if _WORLD:
        _WORLD.destroy()
        _WORLD = None
    _MODEL_PARALLEL_GROUP_TIMEOUT = None # 重置超时变量
    if torch.distributed.is_initialized():
        torch.distributed.destroy_process_group()

评论区精华

超时参数传递到 mooncake-cpu 子组 正确性

reviewer gemini-code-assist[bot] 指出,在 mooncake 后端中,device_group 已添加超时参数,但 mooncake-cpu 组尚未应用,可能导致不一致。

结论:作者在后续提交中修复了此问题,确保 mooncake-cpu 组也接收 subgroup_timeout。 · 已解决

风险与影响

回归风险低:变更仅添加超时参数传递,未修改核心逻辑;但需确保超时值正确传递到所有后端(包括mooncake、gloo及默认NCCL)。
兼容性风险torch.distributed.new_grouptimeout参数支持timedelta类型,与现有代码中gloo_timeout(可能为timedelta)类型一致,但需确认所有PyTorch版本均支持此参数。
性能影响:超时设置本身不影响性能,仅改变超时行为;但若超时设置过短,可能增加因超时导致的失败频率。

对用户:修复后,用户通过--dist-timeout设置的超时将正确应用于所有分布式子组,避免因默认超时导致的意外失败,提升大规模分布式训练的稳定性。
对系统:确保超时一致性,减少因超时配置不一致引发的调试复杂度。
对团队:此修复针对核心分布式模块,需在涉及模型并行的测试中验证超时行为。

核心路径变更 缺少测试覆盖

关联 Issue

#21911 [Bug] --dist-timeout is only applied to init_process_group, but not propagated to NCCL subgroups

完整报告

参与讨论