From 76116641bcfc72409aaef8ffcc39a84c727e4989 Mon Sep 17 00:00:00 2001 From: bolunz Date: Thu, 27 Aug 2026 17:58:46 +0800 Subject: [PATCH 1/3] fix: fix grad accu under zero2, add NoSyncGuard --- example/gpt2/main.cc | 25 ++++++++++++++++--- example/llama3/main.cc | 25 ++++++++++++++++--- infini_train/include/nn/modules/module.h | 20 +++++++++++++++ .../parallel/ddp/distributed_data_parallel.h | 3 +++ .../nn/parallel/ddp/param_and_grad_buffer.h | 2 ++ .../nn/parallel/pp/pipeline_parallel.h | 3 +++ .../nn/parallel/pp/pipeline_schedule.h | 7 ++++++ .../parallel/ddp/distributed_data_parallel.cc | 9 +++++++ .../nn/parallel/ddp/param_and_grad_buffer.cc | 10 +------- .../src/nn/parallel/pp/pipeline_parallel.cc | 4 +++ .../src/nn/parallel/pp/pipeline_schedule.cc | 23 +++++++++++++++++ 11 files changed, 116 insertions(+), 15 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 2551880e..c1051a36 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -280,6 +280,7 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Optimizing " << params_to_optimize.size() << " model parameters"; } + std::shared_ptr ddp_model = nullptr; if (pp_world_size > 1) { // NOTE(dcj): To ensure that the tensor shapes at the pipeline stage boundaries remain correct // when sequence parallelism (SP) is enabled, we need to divide by sp_world_size. @@ -290,10 +291,23 @@ 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(); + std::vector> ddp_chunks; + ddp_chunks.reserve(mutable_chunks->size()); for (int chunk_id = 0; chunk_id < mutable_chunks->size(); ++chunk_id) { - (*mutable_chunks)[chunk_id] + auto ddp_chunk = std::make_shared(mutable_chunks->at(chunk_id), rank, ddp_config); + (*mutable_chunks)[chunk_id] = ddp_chunk; + ddp_chunks.push_back(std::move(ddp_chunk)); + } + if (FLAGS_zero_stage >= 1) { + pipeline_model->SetNoSyncFunc([ddp_chunks = std::move(ddp_chunks)] { + std::vector> guards; + guards.reserve(ddp_chunks.size()); + for (const auto &chunk : ddp_chunks) { guards.push_back(chunk->no_sync()); } + return guards; + }); } } } else if (ddp_world_size > 1) { @@ -302,7 +316,8 @@ void Train(const nn::parallel::Rank &rank) { // Otherwise, DDP’s gradient hooks may be lost because new parameter tensors // are created during the conversion. auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; - model = std::make_shared(model, rank, ddp_config); + ddp_model = std::make_shared(model, rank, ddp_config); + model = ddp_model; } const size_t train_loader_batch_size = pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size; @@ -486,6 +501,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_model && FLAGS_zero_stage >= 1 && micro_step != grad_accum_steps - 1) { + no_sync_guard = ddp_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..a13568ce 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -253,6 +253,7 @@ void Train(const nn::parallel::Rank &rank) { auto num_micro_batches = FLAGS_total_batch_size / (FLAGS_batch_size * FLAGS_sequence_length * ddp_world_size); + std::shared_ptr ddp_model = nullptr; if (pp_world_size > 1) { // NOTE(dcj): To ensure that the tensor shapes at the pipeline stage boundaries remain correct // when sequence parallelism (SP) is enabled, we need to divide by sp_world_size. @@ -263,10 +264,23 @@ 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(); + std::vector> ddp_chunks; + ddp_chunks.reserve(mutable_chunks->size()); for (int chunk_id = 0; chunk_id < mutable_chunks->size(); ++chunk_id) { - (*mutable_chunks)[chunk_id] + auto ddp_chunk = std::make_shared(mutable_chunks->at(chunk_id), rank, ddp_config); + (*mutable_chunks)[chunk_id] = ddp_chunk; + ddp_chunks.push_back(std::move(ddp_chunk)); + } + if (FLAGS_zero_stage >= 1) { + pipeline_model->SetNoSyncFunc([ddp_chunks = std::move(ddp_chunks)] { + std::vector> guards; + guards.reserve(ddp_chunks.size()); + for (const auto &chunk : ddp_chunks) { guards.push_back(chunk->no_sync()); } + return guards; + }); } } } else if (ddp_world_size > 1) { @@ -276,7 +290,8 @@ void Train(const nn::parallel::Rank &rank) { // are created during the conversion. auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; - model = std::make_shared(model, rank, ddp_config); + ddp_model = std::make_shared(model, rank, ddp_config); + model = ddp_model; } const size_t train_loader_batch_size = pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size; @@ -465,6 +480,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_model && FLAGS_zero_stage >= 1 && micro_step != grad_accum_steps - 1) { + no_sync_guard = ddp_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..cf96bc10 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,20 +154,13 @@ 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 (params_with_grad_.size() == params_.size()) { // All param grads are ready in this group, trigger grad sync StartGradSync(); @@ -297,7 +290,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..c10fb603 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc @@ -201,6 +201,15 @@ float PipelineSchedule::StepMicroBatches(const std::vector last_backward_task(vpp_size, schedule.size()); + for (size_t i = 0; i < schedule.size(); ++i) { + const auto &task = schedule[i]; + if (task.stage_id == stage_idx && !task.is_forward) { + last_backward_task[task.local_chunk_idx] = i; + } + } + static bool has_printed = false; if (!has_printed && stage_idx == 0) { PrintScheduleTable(schedule, n, num_stages, vpp_size); @@ -212,6 +221,11 @@ 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_(); + } + for (size_t i = 0; i < schedule.size(); ++i) { const auto &task = schedule[i]; if (task.stage_id != stage_idx) { @@ -244,6 +258,10 @@ float PipelineSchedule::StepMicroBatches(const std::vector loss; @@ -267,9 +285,14 @@ float PipelineSchedule::StepMicroBatches(const std::vectorBackward(dummy_gradient); } + if (is_last_microbatch && no_sync_func_) { + no_sync_guards = no_sync_func_(); + } } } + no_sync_guards.clear(); + return total_loss; } From df14b83e38d46d76193d21dd8a97fd2b2244d277 Mon Sep 17 00:00:00 2001 From: bolunz Date: Mon, 31 Aug 2026 14:10:37 +0800 Subject: [PATCH 2/3] fix: remove redundant code --- example/gpt2/main.cc | 20 +++++++----------- example/llama3/main.cc | 20 +++++++----------- .../src/nn/parallel/pp/pipeline_schedule.cc | 21 ++++--------------- 3 files changed, 18 insertions(+), 43 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index c1051a36..d5e9e968 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -280,7 +280,6 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Optimizing " << params_to_optimize.size() << " model parameters"; } - std::shared_ptr ddp_model = nullptr; if (pp_world_size > 1) { // NOTE(dcj): To ensure that the tensor shapes at the pipeline stage boundaries remain correct // when sequence parallelism (SP) is enabled, we need to divide by sp_world_size. @@ -293,19 +292,15 @@ void Train(const nn::parallel::Rank &rank) { auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; auto *pipeline_model = dynamic_cast(model.get()); auto *mutable_chunks = pipeline_model->mutable_chunks(); - std::vector> ddp_chunks; - ddp_chunks.reserve(mutable_chunks->size()); for (int chunk_id = 0; chunk_id < mutable_chunks->size(); ++chunk_id) { - auto ddp_chunk + (*mutable_chunks)[chunk_id] = std::make_shared(mutable_chunks->at(chunk_id), rank, ddp_config); - (*mutable_chunks)[chunk_id] = ddp_chunk; - ddp_chunks.push_back(std::move(ddp_chunk)); } if (FLAGS_zero_stage >= 1) { - pipeline_model->SetNoSyncFunc([ddp_chunks = std::move(ddp_chunks)] { + pipeline_model->SetNoSyncFunc([mutable_chunks] { std::vector> guards; - guards.reserve(ddp_chunks.size()); - for (const auto &chunk : ddp_chunks) { guards.push_back(chunk->no_sync()); } + guards.reserve(mutable_chunks->size()); + for (const auto &chunk : *mutable_chunks) { guards.push_back(chunk->no_sync()); } return guards; }); } @@ -316,8 +311,7 @@ void Train(const nn::parallel::Rank &rank) { // Otherwise, DDP’s gradient hooks may be lost because new parameter tensors // are created during the conversion. auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; - ddp_model = std::make_shared(model, rank, ddp_config); - model = ddp_model; + model = std::make_shared(model, rank, ddp_config); } const size_t train_loader_batch_size = pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size; @@ -502,8 +496,8 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": start backward"; std::unique_ptr no_sync_guard; - if (ddp_model && FLAGS_zero_stage >= 1 && micro_step != grad_accum_steps - 1) { - no_sync_guard = ddp_model->no_sync(); + 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 diff --git a/example/llama3/main.cc b/example/llama3/main.cc index a13568ce..b069d73f 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -253,7 +253,6 @@ void Train(const nn::parallel::Rank &rank) { auto num_micro_batches = FLAGS_total_batch_size / (FLAGS_batch_size * FLAGS_sequence_length * ddp_world_size); - std::shared_ptr ddp_model = nullptr; if (pp_world_size > 1) { // NOTE(dcj): To ensure that the tensor shapes at the pipeline stage boundaries remain correct // when sequence parallelism (SP) is enabled, we need to divide by sp_world_size. @@ -266,19 +265,15 @@ void Train(const nn::parallel::Rank &rank) { auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; auto *pipeline_model = dynamic_cast(model.get()); auto *mutable_chunks = pipeline_model->mutable_chunks(); - std::vector> ddp_chunks; - ddp_chunks.reserve(mutable_chunks->size()); for (int chunk_id = 0; chunk_id < mutable_chunks->size(); ++chunk_id) { - auto ddp_chunk + (*mutable_chunks)[chunk_id] = std::make_shared(mutable_chunks->at(chunk_id), rank, ddp_config); - (*mutable_chunks)[chunk_id] = ddp_chunk; - ddp_chunks.push_back(std::move(ddp_chunk)); } if (FLAGS_zero_stage >= 1) { - pipeline_model->SetNoSyncFunc([ddp_chunks = std::move(ddp_chunks)] { + pipeline_model->SetNoSyncFunc([mutable_chunks] { std::vector> guards; - guards.reserve(ddp_chunks.size()); - for (const auto &chunk : ddp_chunks) { guards.push_back(chunk->no_sync()); } + guards.reserve(mutable_chunks->size()); + for (const auto &chunk : *mutable_chunks) { guards.push_back(chunk->no_sync()); } return guards; }); } @@ -290,8 +285,7 @@ void Train(const nn::parallel::Rank &rank) { // are created during the conversion. auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; - ddp_model = std::make_shared(model, rank, ddp_config); - model = ddp_model; + model = std::make_shared(model, rank, ddp_config); } const size_t train_loader_batch_size = pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size; @@ -481,8 +475,8 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": start backward"; std::unique_ptr no_sync_guard; - if (ddp_model && FLAGS_zero_stage >= 1 && micro_step != grad_accum_steps - 1) { - no_sync_guard = ddp_model->no_sync(); + 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 diff --git a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc index c10fb603..865f3659 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc @@ -201,15 +201,6 @@ float PipelineSchedule::StepMicroBatches(const std::vector last_backward_task(vpp_size, schedule.size()); - for (size_t i = 0; i < schedule.size(); ++i) { - const auto &task = schedule[i]; - if (task.stage_id == stage_idx && !task.is_forward) { - last_backward_task[task.local_chunk_idx] = i; - } - } - static bool has_printed = false; if (!has_printed && stage_idx == 0) { PrintScheduleTable(schedule, n, num_stages, vpp_size); @@ -225,6 +216,7 @@ float PipelineSchedule::StepMicroBatches(const std::vector backward_counts(vpp_size, 0); for (size_t i = 0; i < schedule.size(); ++i) { const auto &task = schedule[i]; @@ -258,9 +250,9 @@ float PipelineSchedule::StepMicroBatches(const std::vectorBackward(dummy_gradient); } - if (is_last_microbatch && no_sync_func_) { - no_sync_guards = no_sync_func_(); - } } } - no_sync_guards.clear(); - return total_loss; } From b55969a4f2d49102e1e12c8ed80235e3bd112290 Mon Sep 17 00:00:00 2001 From: bolunz Date: Mon, 31 Aug 2026 17:57:04 +0800 Subject: [PATCH 3/3] fix: fix bucket group sync behaviors --- .../src/nn/parallel/ddp/param_and_grad_buffer.cc | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) 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 cf96bc10..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 @@ -160,7 +160,18 @@ void ParamAndGradBucketGroup::RegisterGradReady(const std::shared_ptr &p return; } - params_with_grad_.insert(parameter.get()); + 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 StartGradSync();