fix: fix grad accumulation/overlap under ZeRO-2 - #213
Conversation
ae24cfc to
a459eff
Compare
| Module::Module(const std::string &type) : type_(type), device_(Device()) {} | ||
|
|
||
| std::unique_ptr<NoSyncGuard> Module::no_sync() { | ||
| return std::make_unique<NoSyncGuard>([] {}); |
|
|
||
| std::unique_ptr<nn::NoSyncGuard> DistributedDataParallel::no_sync() { | ||
| SetIsLastMicrobatch(false); | ||
| return std::make_unique<nn::NoSyncGuard>([this] { SetIsLastMicrobatch(true); }); |
There was a problem hiding this comment.
PyTorch 在这里和 Megatron 不太一样,进入 no_sync() 时保存旧状态,退出时会恢复旧状态,所以能支持嵌套作用域。目前我觉得不太需要,如果后续需要支持嵌套的话可以参考。
| (*mutable_chunks)[chunk_id] | ||
| = std::make_shared<DistributedDataParallel>(mutable_chunks->at(chunk_id), rank, ddp_config); | ||
| } | ||
| pipeline_model->SetNoSyncFunc([mutable_chunks] { |
There was a problem hiding this comment.
这里主要是给 Pipeline 每个 chunk 注册回调函数,但 PipelineSchedule 自己能访问 stage_->chunks,能不能让 PipelineSchedule 直接遍历 chunks 创建 guard 啊,这样相关的逻辑不用放在训练入口。
There was a problem hiding this comment.
stage 里维护的 chunk 是 Module 基类类型,需要根据 ddp_world_size 判断,才能安全将 module 转换成 DistributedDataParallel 类型。因此不建议目前在 pp 模块内部引入 ddp 相关的耦合逻辑,后续训练入口这块可以考虑统一整理成类似 megatron train.py 的形式,提供一个统一的训练入口。
| (*mutable_chunks)[chunk_id] | ||
| = std::make_shared<DistributedDataParallel>(mutable_chunks->at(chunk_id), rank, ddp_config); | ||
| } | ||
| pipeline_model->SetNoSyncFunc([mutable_chunks] { |
There was a problem hiding this comment.
stage 里维护的 chunk 是 Module 基类类型,需要根据 ddp_world_size 判断,才能安全将 module 转换成 DistributedDataParallel 类型。因此不建议目前在 pp 模块内部引入 ddp 相关的耦合逻辑,后续训练入口这块可以考虑统一整理成类似 megatron train.py 的形式,提供一个统一的训练入口。
背景
先前的实现
is_last_microbatch的处理过于草率,导致目前 ZeRO-2 同时开启梯度累积和overlap_grad_reduce时,当前实现会在第一个 microbatch backward 期间发起 reduce-scatter,并将grad_reduce_dispatched_设置为true。该状态直到
optimizer->step()调用FinishGradSync()后才会重置,因此存在两个问题:temp_full_grad_buffer,造成计算与通信之间的数据竞争。关闭
overlap_grad_reduce时不会触发该问题,因为梯度同步统一在所有 microbatch 完成后的optimizer->step()中执行。修改内容
参考 Megatron-LM 的
no_sync机制,引入真实的is_last_microbatch_控制:Module增加通用的no_sync()接口和 RAIINoSyncGuard。DistributedDataParallel::no_sync()在 guard 生命周期内将 bucket group 的is_last_microbatch_设置为false,退出时恢复为true。overlap_grad_reduce时保持原有行为,由optimizer->step()发起同步。no_sync_func_使用该机制,不直接依赖或包含 DDP 实现。NoSyncGuard。overlap_grad_reduce命令行参数,继续使用 DDP 配置中的默认行为。行为变化
开启梯度累积和
overlap_grad_reduce后,同一个 optimizer step 内的执行过程变为:optimizer->step():等待通信完成并更新参数。这样可以确保 reduce-scatter 读取的是所有 microbatch 累积后的完整梯度,同时避免通信期间继续修改 full gradient buffer。