diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 2551880e..d5e9e968 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -290,11 +290,20 @@ void Train(const nn::parallel::Rank &rank) { pp_rank, device, model_config.GetChunkSize()); if (ddp_world_size > 1) { auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; - auto *mutable_chunks = dynamic_cast(model.get())->mutable_chunks(); + auto *pipeline_model = dynamic_cast(model.get()); + auto *mutable_chunks = pipeline_model->mutable_chunks(); for (int chunk_id = 0; chunk_id < mutable_chunks->size(); ++chunk_id) { (*mutable_chunks)[chunk_id] = std::make_shared(mutable_chunks->at(chunk_id), rank, ddp_config); } + if (FLAGS_zero_stage >= 1) { + pipeline_model->SetNoSyncFunc([mutable_chunks] { + std::vector> guards; + guards.reserve(mutable_chunks->size()); + for (const auto &chunk : *mutable_chunks) { guards.push_back(chunk->no_sync()); } + return guards; + }); + } } } else if (ddp_world_size > 1) { // NOTE(dcj): Complete all device (.to(device)) and dtype (.to(dtype)) conversions @@ -486,6 +495,10 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": finish loss forward"; LOG(INFO) << "Rank " << rank.GlobalRank() << ": start backward"; + std::unique_ptr no_sync_guard; + if (ddp_world_size > 1 && FLAGS_zero_stage >= 1 && micro_step != grad_accum_steps - 1) { + no_sync_guard = model->no_sync(); + } loss->Backward(); // Defer the loss D2H copy until after backward; reading it earlier would synchronize CUDA // between forward and backward. diff --git a/example/llama3/main.cc b/example/llama3/main.cc index ccfca86a..b069d73f 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -263,11 +263,20 @@ void Train(const nn::parallel::Rank &rank) { pp_rank, device, model_config.GetChunkSize()); if (ddp_world_size > 1) { auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; - auto *mutable_chunks = dynamic_cast(model.get())->mutable_chunks(); + auto *pipeline_model = dynamic_cast(model.get()); + auto *mutable_chunks = pipeline_model->mutable_chunks(); for (int chunk_id = 0; chunk_id < mutable_chunks->size(); ++chunk_id) { (*mutable_chunks)[chunk_id] = std::make_shared(mutable_chunks->at(chunk_id), rank, ddp_config); } + if (FLAGS_zero_stage >= 1) { + pipeline_model->SetNoSyncFunc([mutable_chunks] { + std::vector> guards; + guards.reserve(mutable_chunks->size()); + for (const auto &chunk : *mutable_chunks) { guards.push_back(chunk->no_sync()); } + return guards; + }); + } } } else if (ddp_world_size > 1) { // NOTE(dcj): Complete all device (.to(device)) and dtype (.to(dtype)) conversions @@ -465,6 +474,10 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": finish loss forward"; LOG(INFO) << "Rank " << rank.GlobalRank() << ": start backward"; + std::unique_ptr no_sync_guard; + if (ddp_world_size > 1 && FLAGS_zero_stage >= 1 && micro_step != grad_accum_steps - 1) { + no_sync_guard = model->no_sync(); + } loss->Backward(); // Defer the loss D2H copy until after backward; reading it earlier would synchronize CUDA // between forward and backward. diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 1d42b2ac..08c9b8aa 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -20,6 +20,22 @@ template class HookHandleImpl; namespace infini_train::nn { class Module; +class NoSyncGuard { +public: + explicit NoSyncGuard(std::function exit_func) : exit_func_(std::move(exit_func)) {} + ~NoSyncGuard() { + if (exit_func_) { + exit_func_(); + } + } + + NoSyncGuard(const NoSyncGuard &) = delete; + NoSyncGuard &operator=(const NoSyncGuard &) = delete; + +private: + std::function exit_func_; +}; + namespace parallel::function { std::vector> Replicate(const std::shared_ptr &network, const std::vector &devices); @@ -82,6 +98,10 @@ class Module : public std::enable_shared_from_this { return 0.0f; }; + virtual std::unique_ptr no_sync() { + return std::make_unique([] {}); + } + virtual void To(Device device); virtual void To(DataType dtype); diff --git a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h index 823ae82b..8162b126 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h +++ b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h @@ -31,6 +31,8 @@ class DistributedDataParallel : public nn::Module { std::shared_ptr module() const; + std::unique_ptr no_sync() override; + DistributedDataParallelConfig ddp_config() const { return ddp_config_; } const std::vector> ¶m_grad_buffers() const { return param_grad_buffers_; } @@ -41,6 +43,7 @@ class DistributedDataParallel : public nn::Module { void BuildParamAndGradBuffers(); void RegisterBackwardHooks(); void OnGradReady(const std::shared_ptr ¶m); + void SetIsLastMicrobatch(bool is_last_microbatch); private: std::shared_ptr reducer_ = nullptr; diff --git a/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h b/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h index 4af99d81..523831ca 100644 --- a/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h +++ b/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h @@ -97,6 +97,8 @@ class ParamAndGradBucketGroup { // When all params in a bucket group are ready, will call StartGradSync() void RegisterGradReady(const std::shared_ptr ¶meter); + void SetIsLastMicrobatch(bool is_last_microbatch) { is_last_microbatch_ = is_last_microbatch; } + // Start grad reduce void StartGradSync(); diff --git a/infini_train/include/nn/parallel/pp/pipeline_parallel.h b/infini_train/include/nn/parallel/pp/pipeline_parallel.h index 25939bdc..948d1f9b 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_parallel.h +++ b/infini_train/include/nn/parallel/pp/pipeline_parallel.h @@ -1,6 +1,7 @@ // pipeline_parallel.h #pragma once +#include #include #include @@ -40,6 +41,8 @@ class PipelineParallel : public Module { std::vector> *mutable_chunks(); + void SetNoSyncFunc(std::function>()> func); + private: void BuildPipelineStage(const std::vector> &recv_shape, Device device, std::vector> &&chunks); diff --git a/infini_train/include/nn/parallel/pp/pipeline_schedule.h b/infini_train/include/nn/parallel/pp/pipeline_schedule.h index 053650d7..e4484cbb 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_schedule.h +++ b/infini_train/include/nn/parallel/pp/pipeline_schedule.h @@ -1,9 +1,11 @@ #pragma once +#include #include #include #include "infini_train/include/datatype.h" +#include "infini_train/include/nn/modules/module.h" namespace infini_train { class Tensor; @@ -31,12 +33,17 @@ class PipelineSchedule { const std::vector> &target_mbs, const std::shared_ptr &loss_fn, DataType dtype); + using NoSyncFunc = std::function>()>; + + void SetNoSyncFunc(NoSyncFunc func) { no_sync_func_ = std::move(func); } + std::vector> ReceiveFromPrev(int peer_rank); std::vector> SendToNext(const std::vector> &tensors, int peer_rank); protected: int num_micro_batches_ = -1; std::shared_ptr stage_ = nullptr; + NoSyncFunc no_sync_func_; }; class PipelineParallelScheduler { diff --git a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc index 57460ba0..a14782fd 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -216,4 +216,13 @@ DistributedDataParallel::Forward(const std::vector> &inp } std::shared_ptr DistributedDataParallel::module() const { return modules_.at(kModuleName); } + +std::unique_ptr DistributedDataParallel::no_sync() { + SetIsLastMicrobatch(false); + return std::make_unique([this] { SetIsLastMicrobatch(true); }); +} + +void DistributedDataParallel::SetIsLastMicrobatch(bool is_last_microbatch) { + for (auto &group : bucket_groups_) { group->SetIsLastMicrobatch(is_last_microbatch); } +} } // namespace infini_train::nn::parallel diff --git a/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc b/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc index ab3a8002..6a913f2a 100644 --- a/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc +++ b/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc @@ -154,19 +154,23 @@ void ParamAndGradBucketGroup::RegisterGradReady(const std::shared_ptr &p return; } - // TODO(zbl): Only register grads as ready and trigger grad sync when processing the last microbatch - // For now, is_last_microbatch_ is always true + // Only the last microbatch registers ready grads so the reduce can overlap with its backward pass. if (is_last_microbatch_) { if (!parameter || params_.find(parameter.get()) == params_.end()) { return; } - params_with_grad_.insert(parameter.get()); - // TODO(zbl): check this if sync is only done in last mircobatch - // if (!inserted) { - // LOG(FATAL) << "ParamAndGradBucketGroup: RegisterGradReady() was called twice for the same parameter in a - // bucket group."; return; - // } + if (grad_reduce_dispatched_) { + LOG(FATAL) << "ParamAndGradBucketGroup: RegisterGradReady() was called after grad sync was dispatched."; + return; + } + + auto [_, inserted] = params_with_grad_.insert(parameter.get()); + if (!inserted) { + LOG(FATAL) << "ParamAndGradBucketGroup: RegisterGradReady() was called twice for the same parameter in a " + "bucket group."; + return; + } if (params_with_grad_.size() == params_.size()) { // All param grads are ready in this group, trigger grad sync @@ -297,7 +301,6 @@ void ParamAndGradBucketGroup::StartGradSync() { } grad_reduce_dispatched_ = true; - // TODO(zbl): no need to clear params_with_grad_ here if grad sync is only done on last microbatch params_with_grad_.clear(); } diff --git a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc index c0369cde..a8563c7d 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc @@ -104,4 +104,8 @@ PipelineParallel::PipelineParallel(const std::shared_ptr module, int num } std::vector> *PipelineParallel::mutable_chunks() { return pipeline_stage_->mutable_chunks(); } + +void PipelineParallel::SetNoSyncFunc(std::function>()> func) { + schedule_->SetNoSyncFunc(std::move(func)); +} } // namespace infini_train::nn::parallel diff --git a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc index b702a301..865f3659 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc @@ -212,6 +212,12 @@ float PipelineSchedule::StepMicroBatches(const std::vector>>> activations( vpp_size, std::vector>>(n)); + std::vector> no_sync_guards; + if (no_sync_func_) { + no_sync_guards = no_sync_func_(); + } + std::vector backward_counts(vpp_size, 0); + for (size_t i = 0; i < schedule.size(); ++i) { const auto &task = schedule[i]; if (task.stage_id != stage_idx) { @@ -244,6 +250,10 @@ float PipelineSchedule::StepMicroBatches(const std::vector loss;