Prhub

#7095 [fsdp, perf] feat: defer gradient sync during accumulation

原始 PR 作者 zhangxin81 合并时间 2026-07-22 14:49 文件变更 3 提交数 5 评论 5 代码增减 +282 / -4

执行摘要

FSDP 梯度延迟同步,减少微批通信开销

FSDP 在 FSDPEngine.forward_backward_batch 中每个微批后都执行梯度同步,但 PPO 训练中多个微批组成一个 mini-batch,优化器只在最终微批后更新。这导致不必要的通信开销。关联 Issue #6010 报告 Qwen3-4B + FSDP2 在 Ascend 910B 上 update_actor 极慢,本 PR 部分解决该瓶颈。

此 PR 值得精读,尤其是设计权衡(通信延迟 vs 显存增加)的实现方式。通过上下文管理器优雅地控制同步边界,测试方法(模拟模块事件记录)也是良好的单元测试范例。建议在引入新硬件后端时验证 FSDP2 的兼容性。

讨论亮点
  1. FSDP2 根模块方法兼容性:gemini-code-assist[bot] 指出 FSDP2 中根模块可能不是 FSDPModule(如 PEFT 包装),直接调用 set_requires_gradient_sync 可能引发 AttributeError。作者 zhangxin81 验证生产路径中 apply_fsdp2 始终对根模块执行 fully_shard,确保暴露该方法,无需额外防护。
  2. 配置开关去留:wuxibin89 认为配置不必要,建议默认启用。zhangxin81 采纳并移除配置,使推迟同步成为默认行为。

实现拆解

  1. 新增 _gradient_sync_context 上下文管理器verl/workers/engine/fsdp/transformer_impl.py):根据 is_last_micro_batch 标志决定是否推迟同步。FSDP1 使用 module.no_sync(),FSDP2 调用 set_requires_gradient_sync(False) 并在退出或异常时恢复,非 FSDP 模块无操作。
  2. 修改 forward_backward_batch 循环:引入 micro_batch_idx 枚举,为每个非最终微批(且非前向推理)应用 _gradient_sync_context,最终微批保持默认同步。
  3. 移除配置开关(根据 review):去掉 FSDPEngineConfig.use_no_sync_for_gradient_accumulation,默认启用推迟同步。
  4. 配套测试tests/workers/test_fsdp_gradient_accumulation_sync_on_cpu.py):新增 221 行 CPU 测试,覆盖 FSDP1/FSDP2 的上下文行为、异常恢复、最终微批保持同步,以及 3 微批调度序列的端到端验证。
  5. 更新文档docs/perf/perf_tuning.rst):新增“Reduce FSDP gradient synchronization during gradient accumulation”小节,说明机制和内存权衡。
文件 模块 状态 重要度
tests/workers/test_fsdp_gradient_accumulation_sync_on_cpu.py 梯度同步测试 added 7.72
verl/workers/engine/fsdp/transformer_impl.py FSDP 引擎 modified 7.25
docs/perf/perf_tuning.rst 性能文档 modified 3.05

关键符号

_gradient_sync_context forward_backward_batch

关键源码片段

tests/workers/test_fsdp_gradient_accumulation_sync_on_cpu.py test-coverage

新增 221 行测试,全面覆盖 FSDP1/FSDP2 的延迟同步行为,包括最终微批、异常恢复、端到端微批调度等。

from contextlib import contextmanager
import pytest
from verl.workers.engine.fsdp.transformer_impl import FSDPEngine, transformer_impl# 模拟 FSDP1 模块:记录 no_sync 进入 / 退出事件
class _FSDP1Module:
    def __init__(self): self.events = []
    @contextmanager
    def no_sync(self):
        self.events.append("enter")
        try: yield
        finally: self.events.append("exit")# 模拟 FSDP2 模块:记录 set_requires_gradient_sync 调用
class _FSDP2Module:
    def __init__(self): self.events = []
    def set_requires_gradient_sync(self, enabled):
        self.events.append(enabled)# 快速构造 FSDPEngine 实例,仅设置模块
def _make_engine(module):
    engine = object.__new__(FSDPEngine)
    engine.module = module
    return engine# 验证非最终微批时,FSDP1 进入 no_sync 上下文,FSDP2 先禁用同步再恢复
@pytest.mark.parametrize("version,module_cls", [(1, _FSDP1Module), (2, _FSDP2Module)])
def test_gradient_sync_context_skips_non_final_micro_batch(monkeypatch, version, module_cls):
    module = module_cls()
    engine = _make_engine(module)
    monkeypatch.setattr(transformer_impl, "fsdp_version", lambda _: version)
    with engine._gradient_sync_context(is_last_micro_batch=False):
        module.events.append("backward")
    expected = ["enter", "backward", "exit"] if version == 1 else [False, "backward", True]
    assert module.events == expected# 验证异常退出时 FSDP2 仍能正确恢复梯度同步状态
def test_gradient_sync_context_restores_fsdp2_after_error(monkeypatch):
    module = _FSDP2Module()
    engine = _make_engine(module)
    monkeypatch.setattr(transformer_impl, "fsdp_version", lambda _: 2)
    with pytest.raises(RuntimeError, match="backward failed"):
        with engine._gradient_sync_context(is_last_micro_batch=False):
            raise RuntimeError("backward failed")
    # finally 块应调用 set_requires_gradient_sync(True)
    assert module.events == [False, True]

评论区精华

FSDP2 根模块 `set_requires_gradient_sync` 方法兼容性 正确性

gemini-code-assist[bot] 指出 FSDP2 中根模块可能不是 FSDPModule(如 PEFT 包装),直接调用 `set_requires_gradient_sync` 可能引发 AttributeError。

结论:作者验证生产路径中 `apply_fsdp2` 始终对根模块执行 `fully_shard`,确保暴露该方法,无需额外防护。 · 已解决

配置开关去留 设计

wuxibin89 认为 `use_no_sync_for_gradient_accumulation` 配置不必要,建议默认启用。

结论:作者采纳,移除该配置,使推迟同步成为默认行为。 · 已解决

风险与影响

  1. 峰值显存增加:推迟同步使非最终微批的梯度保持未分片状态,可能增加显存占用。特别是 micro-batch 数多或模型大时需关注。前向推理不受影响。
  2. FSDP2 子模块兼容性:如果未来引入未通过 fully_shard 包装根模块的路径,可能缺少 set_requires_gradient_sync 方法。当前生产路径和 PEFT 测试已验证安全,但新硬件后端需注意。
  3. 核心路径变更:修改了 forward_backward_batch 循环,可能影响 VeOmniEngine 等子类,但通过 getattr(self, 'scaler', None) 等 fallback 保持兼容。

对用户:所有使用 FSDP 且进行梯度累积的训练任务自动受益,减少通信延迟,可能提升训练吞吐。但显存敏感场景需评估峰值增长。对系统:无配置侵入,默认行为改变,但前向推理和非 FSDP 模块不受影响。对团队:新增测试和文档维护成本,但测试覆盖全面,回归风险低。

峰值内存增加 FSDP2 子模块兼容性 核心路径变更

关联 Issue

#6010 Extremely slow update_actor in GRPO with Qwen3-4B + FSDP2 on Ascend 910B

完整报告

参与讨论