Prhub

#2386 fix: drop context parallelism from the FSDP backend

原始 PR 作者 Zhichenzzz 合并时间 2026-08-12 06:05 文件变更 4 提交数 2 评论 0 代码增减 +23 / -67

执行摘要

FSDP 后端移除不可达的上下文并行代码

PR body 指出 miles/utils/arguments.py 在 FSDP 路径上断言 context_parallel_size == 1,因此 FSDP 后端名下所有 CP 分支都是不可达代码。其中一个分支 _get_model_inputs_args 已经腐烂:它引用 self.cp_group,而该属性在 actor.py 的任何地方都没有被赋值,一旦执行必然 AttributeError。作者希望通过删除整套 CP 机制来消除死代码、避免误导后续维护者以为 FSDP 支持上下文并行,并顺手修复这个隐藏缺陷。同时特意保留 cp.group 为单 rank 组和 context_parallel_size 参数,以兼容共享代码中的 all_reduce 调用和校验逻辑。

值得精读,尤其适合作为“删除从未支持的配置分支”的范式:先由共享校验证明分支不可达、再找出已腐烂的引用、删除时保留最小兼容面(cp.group 非 None、保留参数名)。重点关注 miles/backends/fsdp_utils/parallel.py 中 mesh 从二维降一维的写法,以及 PR body 对两个“刻意保留”的论证。

讨论亮点

该 PR 没有 review 评论线程,唯一审核意见是 Rockdu 的 APPROVED(“LGTM”)。核心设计讨论沉淀在 PR body 的两个“刻意保留”说明里:一是 cp.group 必须保留为真实单 rank 组而非 None,因为共享代码 training_utils/cp_utils.py 会把它直接传给 dist.nn.all_reduce,而 group=None 表示 WORLD,会把本意的 no-op 变成跨 rank 归约;二是 context_parallel_size 需保留在 FSDPArgs 上,因为共享校验无条件读取 args.context_parallel_size,保留它才能用清晰的断言报错代替 unrecognized arguments。

实现拆解

变更入口在 miles/backends/fsdp_utils/parallel.py,随后波及 actor.py、arguments.py 与测试 worker。

  1. mesh 构建简化:build_fsdp_meshes 签名去掉 context_parallel_size,init_device_mesh 直接使用 (world_size,) 一维形状与 ('dp',) 维度名,不再先建 (dp, cp) 二维 mesh 再切片;_unflatten 的分片因子从 data_parallel_size 改为 world_size。返回字典只保留 dp 与 fsdp 两个视图。
  2. 并行状态瘦身:create_fsdp_parallel_state 删除 cp_size/dp_rank/cp_rank 算术和 ring_flash_attn 的 substitute_hf_flash_attn 条件导入,新增 self_group = dist.new_group([rank]),让 cp 与 tp 共享该单 rank 组;日志中的 dp_rank/cp_rank 简化为 rank。
  3. 输入路径清理:actor.py 的 _get_model_inputs_args 删除 ring_flash_attn 的 update_ring_flash_attn_params 调用与 torch.chunk 序列切分,直接使用 batch['tokens'] 与 batch['position_ids'],消除了 self.cp_group 的隐晦引用。
  4. 参数校验简化:arguments.py 的 validate_hybrid_shard_args 删除 context_parallel_size 的范围与整除检查,只保留 world_size % dp_replicate_size;FSDPArgs.context_parallel_size 字段保留,注释明确它仅用于让共享校验拒绝 CP 运行并给出清晰报错。
  5. 测试配套:tests/fast-gpu/_fsdp_hybrid_shard_worker.py 更新 build_fsdp_meshes 调用(去掉 context_parallel_size=1);test_fsdp_hybrid_shard.py 继续覆盖 r1s4、r2s2、r4s1 三种拓扑的 FSDP2 梯度一致性,是 mesh 重写的主要守护。
文件 模块 状态 重要度
miles/backends/fsdp_utils/parallel.py 并行状态 modified 6.9
miles/backends/fsdp_utils/actor.py 执行器 modified 5.91
miles/backends/fsdp_utils/arguments.py 参数解析 modified 5.89
tests/fast-gpu/_fsdp_hybrid_shard_worker.py 混合分片 modified 2.88

关键符号

build_fsdp_meshes create_fsdp_parallel_state _get_model_inputs_args validate_hybrid_shard_args

关键源码片段

miles/backends/fsdp_utils/arguments.py core-logic

validate_hybrid_shard_args 删除 CP 整除检查,context_parallel_size 字段保留并注释其真实用途,是“保留兼容面”设计决策的落点,也呼应同仓库 #2384 的 FSDP 参数层清理。

class FSDPArgs:
    # ... 其余字段省略 ...
​
    # 该字段仅用于让共享参数校验拒绝 context-parallel 的 FSDP 运行
    # (共享代码会无条件读取 args.context_parallel_size)。
    context_parallel_size: int = 1
​
​
def validate_hybrid_shard_args(args) -> None:
    """校验训练拓扑能否构成请求的 FSDP2 mesh。"""
    replicate_size = args.dp_replicate_size
    if replicate_size < 1:
        raise ValueError(f'dp_replicate_size must be at least 1, got {replicate_size}')
​
    world_size = args.actor_num_nodes * args.actor_num_gpus_per_node
    # FSDP 后端为纯数据并行,data_parallel_size 就是 world_size,
    # 因此只需保证 world_size 能被 dp_replicate_size 整除。
    if world_size % replicate_size:
        raise ValueError(
            f'world_size({world_size}) must be divisible by dp_replicate_size({replicate_size})'
        )

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

主要风险来自接口签名变更与共享代码的隐式假设:

  • build_fsdp_meshes 返回值不再包含 'dp_cp' 与 'cp' 键,仓库内当前没有其他调用方(测试已同步),但任何外部或未来代码若按旧接口读取 meshes['cp'] 会立即 KeyError。
  • cp 与 tp 现在共用一个 dist.new_group([rank]) 进程组。若未来有共享代码在同一时刻对这组发起两路不同集合通信,可能存在组复用导致的串扰;目前 FSDP 路径上 tp 恒为单 rank,风险低。
  • actor.py 删除序列切分后,如果有人绕过校验强行以 context_parallel_size > 1 启动 FSDP,会静默得到错误样本而非显式异常。该防护依赖共享校验仍然断言 context_parallel_size == 1,属于跨文件隐式约束。
  • 改动集中在 FSDP 后端核心路径(mesh 构建与并行状态),但通过了 test_fsdp_hybrid_shard.py(4x H200)与两次端到端 rollout 验证,且所有可达路径行为不变,回归面可控。

对用户无感知:FSDP 后端从未支持 CP,删除后任何可达路径行为不变。对系统有一定正向影响:进程组数量减少(dp_cp mesh 内部不再创建 cp 组),ParallelState 更薄,日志中的 dp_rank/cp_rank 简化为单 rank。对团队的主要收益是维护成本下降——消除了一个必然触发 AttributeError 的腐烂分支和两处悬挂的 ring_flash_attn 导入,并明确标注了 context_parallel_size 在 FSDP 下的真实含义,避免后续开发者被误导。

核心路径变更 接口签名变更 跨文件隐式约束 无 review 讨论

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论