执行摘要
- 一句话:FSDP 后端移除不可达的上下文并行代码
- 推荐动作:值得精读,尤其适合作为“删除从未支持的配置分支”的范式:先由共享校验证明分支不可达、再找出已腐烂的引用、删除时保留最小兼容面(cp.group 非 None、保留参数名)。重点关注 miles/backends/fsdp_utils/parallel.py 中 mesh 从二维降一维的写法,以及 PR body 对两个“刻意保留”的论证。
功能与动机
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 调用和校验逻辑。
实现拆解
变更入口在 miles/backends/fsdp_utils/parallel.py,随后波及 actor.py、arguments.py 与测试 worker。
- mesh 构建简化:build_fsdp_meshes 签名去掉 context_parallel_size,init_device_mesh 直接使用 (world_size,) 一维形状与 ('dp',) 维度名,不再先建 (dp, cp) 二维 mesh 再切片;_unflatten 的分片因子从 data_parallel_size 改为 world_size。返回字典只保留 dp 与 fsdp 两个视图。
- 并行状态瘦身: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。
- 输入路径清理:actor.py 的 _get_model_inputs_args 删除 ring_flash_attn 的 update_ring_flash_attn_params 调用与 torch.chunk 序列切分,直接使用 batch['tokens'] 与 batch['position_ids'],消除了 self.cp_group 的隐晦引用。
- 参数校验简化:arguments.py 的 validate_hybrid_shard_args 删除 context_parallel_size 的范围与整除检查,只保留 world_size % dp_replicate_size;FSDPArgs.context_parallel_size 字段保留,注释明确它仅用于让共享校验拒绝 CP 运行并给出清晰报错。
- 测试配套: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(模块 并行状态;类别 source;类型 core-logic;符号 build_fsdp_meshes, create_fsdp_parallel_state): FSDP 后端 mesh 与并行状态构建的核心文件,改动最大(+16/-34):build_fsdp_meshes 从 (dp, cp) 二维 mesh 降为一维 dp mesh,create_fsdp_parallel_state 删除 CP 算术与 ring_flash_attn 条件导入,cp 与 tp 共享单 rank 进程组。
miles/backends/fsdp_utils/actor.py(模块 执行器;类别 source;类型 dependency-wiring;符号 _get_model_inputs_args): _get_model_inputs_args 删除引用了未定义 self.cp_group 的腐烂分支和序列切分逻辑,修掉潜在 AttributeError,是该 PR 除 mesh 重写外最重要的行为清理。
miles/backends/fsdp_utils/arguments.py(模块 参数解析;类别 source;类型 core-logic;符号 FSDPArgs, validate_hybrid_shard_args): validate_hybrid_shard_args 删除 CP 整除检查,context_parallel_size 字段保留并注释其真实用途,是“保留兼容面”设计决策的落点,也呼应同仓库 #2384 的 FSDP 参数层清理。
tests/fast-gpu/_fsdp_hybrid_shard_worker.py(模块 混合分片;类别 test;类型 test-coverage;符号 main): 同步 build_fsdp_meshes 新签名,去掉 context_parallel_size=1 参数,是 mesh 重写后的测试配套;对应的 test_fsdp_hybrid_shard.py 是 FSDP2 梯度一致性的主要守护。
关键符号:build_fsdp_meshes, create_fsdp_parallel_state, _get_model_inputs_args, validate_hybrid_shard_args
关键源码片段
miles/backends/fsdp_utils/arguments.py
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})'
)
评论区精华
该 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。
风险与影响
关联脉络
- PR #2384 fix(fsdp): stop store_true from shadowing bool defaults in FSDPArgs: 同一条 FSDP 参数层清理线,都改到 miles/backends/fsdp_utils/arguments.py,且都伴随 FSDP 相关测试调整。
- PR #2382 fix: drop duplicated rematerialize validation call: 同类“删除冗余校验”的清理,位于共享参数校验模块 miles/utils/arguments.py,与本 PR 的 validate_hybrid_shard_args 简化同属参数校验瘦身脉络。
参与讨论