Skip to content

feat: add distopt AllGather - next step Forward overlap - #217

Open
Chamberlain0w0 wants to merge 1 commit into
masterfrom
feat/overlap-param-gather
Open

feat: add distopt AllGather - next step Forward overlap#217
Chamberlain0w0 wants to merge 1 commit into
masterfrom
feat/overlap-param-gather

Conversation

@Chamberlain0w0

Copy link
Copy Markdown
Contributor

背景

ZeRO-1/2 在 optimizer 更新本地参数分片后,需要通过 AllGather 恢复完整参数。

原实现会在 DistributedOptimizer::Step() 中启动并等待所有参数 AllGather 完成,下一轮 Forward 必须等全部参数同步结束后才能开始。实际上 Forward 只需要保证“当前即将使用的参数”已经同步,因此可以将后续 bucket 的 AllGather 与下一轮 Forward 计算重叠。

主要修改

1. 按参数实际使用时机同步 bucket

  • optimizer step 完成后,只启动每个 model chunk 的第一个参数 bucket AllGather。
  • 为包含直接参数的 Module 注册 Forward pre-hook。
  • Module 执行前等待其依赖的 bucket AllGather 完成。
  • 当前 bucket 完成后,链式启动下一个 bucket 的 AllGather。
  • 后续 bucket 通信可以与前面 Module 的 Forward 计算重叠。

Non-overlap 路径保持原有行为,在 optimizer step 后完成全部参数 AllGather。

2. 管理跨 iteration 的 AllGather 状态

  • 参数 buffer 初始已经是完整副本,第一次 Forward 前无需额外 AllGather。
  • 每次 optimizer 更新后重新发布参数分片。
  • 下一次 optimizer 写入参数前,确保上一轮仍在执行的 AllGather 已完成。
  • 增加重复 dispatch 检查,避免同一参数版本被重复同步。

3. 建立稳定的 Module/Parameter 注册顺序

Forward pre-hook 的 bucket 顺序依赖参数注册顺序,因此为 Module 增加与 PyTorch 语义对齐的注册接口:

  • RegisterParameter()
  • RegisterBuffer()
  • RegisterModule()

底层仍保留 unordered_map 用于按名称查询,同时使用 order vector 保存注册顺序。

以下接口改为按注册顺序遍历:

  • Parameters()
  • NamedParameters()
  • NamedModules()
  • Buffers()
  • StateDict()
  • To() / Apply()

同时支持:

  • Parameters(recurse=false) 获取 Module 的直接参数。
  • 重复注册同名对象时替换对象,但不改变原注册位置。
  • 非持久 buffer 不写入 StateDict()
  • 注册名称和跨 registry 名称冲突检查。

项目内原先直接写入 registry map 的代码已迁移至 RegisterXXX()

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant