diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index ff79a3775..8e5d92c02 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -498,6 +498,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 && 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 344247e2c..19620e993 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -477,6 +477,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 && 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/autograd/function_hook.h b/infini_train/include/autograd/function_hook.h index 734a0930b..8a34d190e 100644 --- a/infini_train/include/autograd/function_hook.h +++ b/infini_train/include/autograd/function_hook.h @@ -1,5 +1,6 @@ #pragma once +#include #include #include "infini_train/include/nn/parallel/reduce_op_type.h" @@ -36,12 +37,14 @@ class PostAccumulateGradHook { class AllReducePostAccumulateHook : public PostAccumulateGradHook { public: AllReducePostAccumulateHook(infini_train::nn::parallel::function::ReduceOpType reduce_op, - const infini_train::nn::parallel::ProcessGroup *pg = nullptr); + const infini_train::nn::parallel::ProcessGroup *pg = nullptr, + std::shared_ptr enabled = nullptr); void operator()(const std::shared_ptr &tensor) override; private: infini_train::nn::parallel::function::ReduceOpType reduce_op_; const infini_train::nn::parallel::ProcessGroup *pg_ = nullptr; + std::shared_ptr enabled_; }; } // namespace infini_train::autograd diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 1d42b2acc..8570c4768 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -20,6 +20,18 @@ template class HookHandleImpl; namespace infini_train::nn { class Module; +class NoSyncGuard { +public: + explicit NoSyncGuard(std::function exit_func); + ~NoSyncGuard(); + + 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 +94,8 @@ class Module : public std::enable_shared_from_this { return 0.0f; }; + virtual std::unique_ptr no_sync(); + 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 823ae82b5..905816b56 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h +++ b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h @@ -1,5 +1,6 @@ #pragma once +#include #include #include "infini_train/include/nn/modules/module.h" @@ -31,6 +32,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,9 +44,12 @@ 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; + // Whether to enable grad sync on last microbatch (DDP naive path) + std::shared_ptr is_last_microbatch_ = std::make_shared(true); DistributedDataParallelConfig ddp_config_; const ProcessGroup *ddp_pg_ = 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 4af99d818..2c572984d 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); + // Start grad reduce void StartGradSync(); @@ -150,6 +152,7 @@ class ParamAndGradBucketGroup { std::vector>> param_buffer_shard_list_; std::vector>> grad_buffer_shard_list_; + // Whether to enable grad sync on last microbatch (DDP + ZeRO path) bool is_last_microbatch_ = true; bool grad_reduce_dispatched_ = false; diff --git a/infini_train/include/nn/parallel/ddp/reducer.h b/infini_train/include/nn/parallel/ddp/reducer.h index 5507b2f64..514f6c427 100644 --- a/infini_train/include/nn/parallel/ddp/reducer.h +++ b/infini_train/include/nn/parallel/ddp/reducer.h @@ -63,6 +63,8 @@ class Reducer : public std::enable_shared_from_this { // Prepare bucket info for next step void PrepareForBackward(); + void SetIsLastMicrobatch(bool is_last_microbatch); + // For custom DDP hook to overwrite the default AllReduce. // This can be used for algorithms like Gradient Compression/GossipGrad. // Hook is registered using `Reducer::RegisterCommHook()`. @@ -149,10 +151,12 @@ class Reducer : public std::enable_shared_from_this { std::vector ready_seen_this_iter_; // Whether to rebuild buckets on next train step bool need_rebuild_ = false; - // Whether to buckets have already been rebuilt on the second step + // Whether buckets have already been rebuilt on the second step bool has_rebuilt_bucket_ = false; // Whether all buckets are ready and backward can be finalized bool all_buckets_ready_this_iter_ = false; + // Whether to enable grad sync on last microbatch (DDP gradient bucketing path) + bool is_last_microbatch_ = true; }; } // namespace infini_train::nn::parallel diff --git a/infini_train/include/nn/parallel/pp/pipeline_schedule.h b/infini_train/include/nn/parallel/pp/pipeline_schedule.h index 053650d7c..cae190f82 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_schedule.h +++ b/infini_train/include/nn/parallel/pp/pipeline_schedule.h @@ -10,7 +10,7 @@ class Tensor; class Optimizer; namespace nn { class Module; -} +} // namespace nn } // namespace infini_train namespace infini_train::nn::parallel { diff --git a/infini_train/src/autograd/function_hook.cc b/infini_train/src/autograd/function_hook.cc index 84094069a..9ab5d1142 100644 --- a/infini_train/src/autograd/function_hook.cc +++ b/infini_train/src/autograd/function_hook.cc @@ -1,16 +1,23 @@ #include "infini_train/include/autograd/function_hook.h" +#include + #include "infini_train/include/nn/parallel/parallel_functional.h" #include "infini_train/include/nn/parallel/process_group.h" #include "infini_train/include/tensor.h" namespace infini_train::autograd { AllReducePostAccumulateHook::AllReducePostAccumulateHook(infini_train::nn::parallel::function::ReduceOpType reduce_op, - const infini_train::nn::parallel::ProcessGroup *pg) + const infini_train::nn::parallel::ProcessGroup *pg, + std::shared_ptr enabled) : reduce_op_(reduce_op), - pg_(pg ? pg : infini_train::nn::parallel::ProcessGroupFactory::Instance()->GetDefaultProcessGroup()) {} + pg_(pg ? pg : infini_train::nn::parallel::ProcessGroupFactory::Instance()->GetDefaultProcessGroup()), + enabled_(std::move(enabled)) {} void AllReducePostAccumulateHook::operator()(const std::shared_ptr &tensor) { + if (enabled_ && !enabled_->load(std::memory_order_relaxed)) { + return; + } infini_train::nn::parallel::function::AllReduce(tensor, reduce_op_, pg_); } } // namespace infini_train::autograd diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 9475d49fe..498bcc075 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -21,10 +21,20 @@ namespace infini_train::nn { +NoSyncGuard::NoSyncGuard(std::function exit_func) : exit_func_(std::move(exit_func)) {} + +NoSyncGuard::~NoSyncGuard() { + if (exit_func_) { + exit_func_(); + } +} + Module::Module() : Module(kUndefinedType) {} Module::Module(const std::string &type) : type_(type), device_(Device()) {} +std::unique_ptr Module::no_sync() { return nullptr; } + const std::string &Module::type() const { return type_; } std::vector> Module::Parameters() const { 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 19361e960..dd05a8d71 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -47,7 +47,8 @@ DistributedDataParallel::DistributedDataParallel(std::shared_ptr mod if (!ddp_config.gradient_bucketing_enabled && ddp_config.zero_stage < 1) { const auto reduce_op = ddp_config.average_in_collective ? function::ReduceOpType::kAvg : function::ReduceOpType::kSum; - auto hook = std::make_unique(reduce_op, ddp_pg_); + auto hook = std::make_unique(reduce_op, ddp_pg_, + is_last_microbatch_); param->RegisterPostAccumulateGradHook(std::move(hook)); } } @@ -223,4 +224,18 @@ DistributedDataParallel::Forward(const std::vector> &inp } std::shared_ptr DistributedDataParallel::module() const { return modules_.at(kModuleName); } + +std::unique_ptr DistributedDataParallel::no_sync() { + const bool previous = is_last_microbatch_->load(std::memory_order_relaxed); + SetIsLastMicrobatch(false); + return std::make_unique([this, previous] { SetIsLastMicrobatch(previous); }); +} + +void DistributedDataParallel::SetIsLastMicrobatch(bool is_last_microbatch) { + is_last_microbatch_->store(is_last_microbatch, std::memory_order_relaxed); + if (reducer_) { + reducer_->SetIsLastMicrobatch(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 ab3a80024..8a3c2d052 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 @@ -146,6 +146,8 @@ void ParamAndGradBucketGroup::Reset() { } } +void ParamAndGradBucketGroup::SetIsLastMicrobatch(bool is_last_microbatch) { is_last_microbatch_ = is_last_microbatch; } + void ParamAndGradBucketGroup::RegisterGradReady(const std::shared_ptr ¶meter) { if (!ddp_config_.overlap_grad_reduce) { LOG(WARNING) @@ -154,19 +156,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,8 +303,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(); } void ParamAndGradBucketGroup::FinishGradSync() { diff --git a/infini_train/src/nn/parallel/ddp/reducer.cc b/infini_train/src/nn/parallel/ddp/reducer.cc index 80115b8e3..bbf2ca5e1 100644 --- a/infini_train/src/nn/parallel/ddp/reducer.cc +++ b/infini_train/src/nn/parallel/ddp/reducer.cc @@ -304,6 +304,11 @@ void Reducer::PrepareForBackward() { } } +void Reducer::SetIsLastMicrobatch(bool is_last_microbatch) { + std::lock_guard lock(mutex_); + is_last_microbatch_ = is_last_microbatch; +} + void Reducer::AttachHooksToParameters() { for (size_t param_idx = 0; param_idx < params_.size(); ++param_idx) { class BucketHook final : public autograd::PostAccumulateGradHook { @@ -332,6 +337,10 @@ void Reducer::AttachHooksToParameters() { void Reducer::MarkVariableReadyDense(size_t variable_index) { std::unique_lock lock(mutex_); + if (!is_last_microbatch_) { + return; + } + const auto loc = locators_.at(variable_index); auto &bucket = buckets_.at(loc.bucket_index); diff --git a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc index b702a3016..6578e628b 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc @@ -212,6 +212,11 @@ float PipelineSchedule::StepMicroBatches(const std::vector>>> activations( vpp_size, std::vector>>(n)); + std::vector> no_sync_guards; + no_sync_guards.reserve(stage_->chunks().size()); + for (const auto &chunk : stage_->chunks()) { no_sync_guards.push_back(chunk->no_sync()); } + 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 +249,10 @@ float PipelineSchedule::StepMicroBatches(const std::vector loss;