From e8f6bf2764a51a98bb2802a1e8f9d8c051608eaa Mon Sep 17 00:00:00 2001 From: bolunz Date: Thu, 17 Sep 2026 14:04:53 +0800 Subject: [PATCH] fix: fix param grad sync under SP --- example/gpt2/main.cc | 1 + example/llama3/main.cc | 1 + .../nn/parallel/ddp/distributed_optimizer.h | 3 +- infini_train/include/nn/parallel/utils.h | 2 ++ infini_train/include/optimizer.h | 16 +++++++-- infini_train/include/tensor.h | 4 +++ infini_train/src/nn/modules/normalization.cc | 3 ++ .../src/nn/modules/transformer/moe/router.cc | 7 ++++ .../src/nn/modules/transformer/transformer.cc | 3 ++ .../nn/parallel/ddp/distributed_optimizer.cc | 12 ++++--- .../src/nn/parallel/tensor_parallel.cc | 3 ++ infini_train/src/nn/parallel/utils.cc | 36 +++++++++++++++++++ infini_train/src/optimizer.cc | 19 ++++++++-- infini_train/src/tensor.cc | 3 ++ 14 files changed, 103 insertions(+), 10 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 8e5d92c02..46982a328 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -351,6 +351,7 @@ void Train(const nn::parallel::Rank &rank) { } else { optimizer = optimizer_creator(named_parameters); } + optimizer->set_model_grad_finalizer(nn::parallel::FinalizeModelGrads); const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration; TrainingLRSchedulerConfig sched_config; diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 19620e993..200a41c72 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -333,6 +333,7 @@ void Train(const nn::parallel::Rank &rank) { } else { optimizer = optimizer_creator(named_parameters); } + optimizer->set_model_grad_finalizer(nn::parallel::FinalizeModelGrads); const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration; TrainingLRSchedulerConfig sched_config; diff --git a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h index d7cea198e..3e8c7fea2 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h +++ b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h @@ -29,7 +29,8 @@ class DistributedOptimizer final : public infini_train::Optimizer { const std::vector> &model_chunks, size_t ddp_world_size, size_t ddp_rank); - void Step() override; + void FinalizeModelGrads() override; + void StepImpl() override; void ZeroGrad(bool set_to_none = true) override; diff --git a/infini_train/include/nn/parallel/utils.h b/infini_train/include/nn/parallel/utils.h index 8aa11856a..696b45949 100644 --- a/infini_train/include/nn/parallel/utils.h +++ b/infini_train/include/nn/parallel/utils.h @@ -31,4 +31,6 @@ std::vector> GatherFromSPRegionFunc(const std::shared_pt std::vector> ScatterToTPRegionFunc(const std::shared_ptr &input); std::vector> ReduceFromTPRegionFunc(const std::shared_ptr &input); std::vector> CopyToTPRegionFunc(const std::shared_ptr &input); + +void FinalizeModelGrads(const std::vector> ¶ms); } // namespace infini_train::nn::parallel diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index d85b1acea..362df76a8 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -18,6 +18,7 @@ using NamedParameter = std::pair>; using NamedParameterList = std::vector; using OptimizerCreator = std::function(const std::vector> ¶ms)>; using OptimizerCreatorNamed = std::function(const NamedParameterList &named_params)>; +using ModelGradFinalizer = std::function> ¶ms)>; class Optimizer { public: @@ -27,7 +28,11 @@ class Optimizer { virtual void ZeroGrad(bool set_to_none = true); - virtual void Step() = 0; + virtual void Step() final; + + void set_model_grad_finalizer(ModelGradFinalizer finalizer) { model_grad_finalizer_ = std::move(finalizer); } + + const std::vector> ¶meters() const { return params_; } virtual std::unordered_map> StateDict() const { return {}; }; @@ -44,11 +49,16 @@ class Optimizer { void set_initial_learning_rate(float lr); protected: + virtual void FinalizeModelGrads(); + virtual void StepImpl() = 0; + std::vector> params_; std::vector parameter_names_; float learning_rate_ = 0.0f; float initial_learning_rate_ = 0.0f; bool initial_lr_set_ = false; + + ModelGradFinalizer model_grad_finalizer_; }; namespace optimizers { @@ -57,7 +67,7 @@ class SGD : public Optimizer { SGD(const std::vector> ¶ms, float learning_rate); SGD(const NamedParameterList &named_params, float learning_rate); - void Step() override; + void StepImpl() override; static OptimizerCreator Create(float learning_rate); static OptimizerCreatorNamed CreateNamed(float learning_rate); @@ -70,7 +80,7 @@ class Adam : public Optimizer { Adam(const NamedParameterList &named_params, float learning_rate = 1e-3, float beta1 = 0.9, float beta2 = 0.999, float eps = 1e-8); - void Step() override; + void StepImpl() override; std::unordered_map> StateDict() const override; diff --git a/infini_train/include/tensor.h b/infini_train/include/tensor.h index dcfd8927f..dbc3860c5 100644 --- a/infini_train/include/tensor.h +++ b/infini_train/include/tensor.h @@ -80,6 +80,9 @@ class Tensor : public std::enable_shared_from_this { size_t NumElements() const; DataType Dtype() const; + void set_sequence_parallel(bool enabled) { sequence_parallel_ = enabled; } + bool sequence_parallel() const { return sequence_parallel_; } + std::shared_ptr Detach() const; void Fill(Scalar value); @@ -242,6 +245,7 @@ class Tensor : public std::enable_shared_from_this { private: std::shared_ptr grad_ = nullptr; bool requires_grad_ = false; + bool sequence_parallel_ = false; bool is_leaf_ = true; std::shared_ptr grad_fn_ = nullptr; int output_idx_ = 0; diff --git a/infini_train/src/nn/modules/normalization.cc b/infini_train/src/nn/modules/normalization.cc index 388b04de5..9791c56a7 100644 --- a/infini_train/src/nn/modules/normalization.cc +++ b/infini_train/src/nn/modules/normalization.cc @@ -18,6 +18,8 @@ LayerNorm::LayerNorm(const std::vector &normalized_shape, float eps, De = std::make_shared(normalized_shape, DataType::kFLOAT32, device_)->RequiresGrad(); parameters_[kParamBiasName] = std::make_shared(normalized_shape, DataType::kFLOAT32, device_)->RequiresGrad(); + parameters_[kParamWeightName]->set_sequence_parallel(true); + parameters_[kParamBiasName]->set_sequence_parallel(true); ResetParameters(); } @@ -35,6 +37,7 @@ void LayerNorm::ResetParameters() { RMSNorm::RMSNorm(int64_t dim, float eps, Device device) : CloneableModule(kType), eps_(eps) { parameters_[kParamWeightName] = std::make_shared(std::vector{dim}, DataType::kFLOAT32, device)->RequiresGrad(); + parameters_[kParamWeightName]->set_sequence_parallel(true); nn::init::Ones(parameters_[kParamWeightName]); } diff --git a/infini_train/src/nn/modules/transformer/moe/router.cc b/infini_train/src/nn/modules/transformer/moe/router.cc index 252086846..b045400aa 100644 --- a/infini_train/src/nn/modules/transformer/moe/router.cc +++ b/infini_train/src/nn/modules/transformer/moe/router.cc @@ -11,6 +11,7 @@ #include "infini_train/include/nn/functional.h" #include "infini_train/include/nn/init.h" #include "infini_train/include/nn/modules/transformer/moe/moe_utils.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/tensor.h" namespace infini_train::nn::moe { @@ -24,12 +25,18 @@ TopKRouter::TopKRouter(const TransformerConfig &config) : CloneableModule(kType) = std::make_shared(std::vector{moe_config.num_experts, config_.n_embd}, DataType::kFLOAT32, device_) ->RequiresGrad(); + if (parallel::global::GetSequenceParallelEnabled()) { + parameters_[kParamWeightName]->set_sequence_parallel(true); + } init::KaimingUniform(parameters_[kParamWeightName]); if (config_.add_bias_linear) { parameters_[kParamBiasName] = std::make_shared(std::vector{moe_config.num_experts}, DataType::kFLOAT32, device_) ->RequiresGrad(); + if (parallel::global::GetSequenceParallelEnabled()) { + parameters_[kParamBiasName]->set_sequence_parallel(true); + } parameters_[kParamBiasName]->Fill(0.0f); } } diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index 99a739d2d..f1048058a 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -32,6 +32,9 @@ TransformerFirstStage::TransformerFirstStage(const TransformerConfig &config) // Only learned absolute position embedding uses a trainable WPE table. if (config_.position_embedding_type == PositionEmbeddingType::kLearnedAbsolute) { modules_[kWPELayerName] = std::make_shared(config_.block_size, config_.n_embd); + if (parallel::global::GetSequenceParallelEnabled()) { + modules_[kWPELayerName]->parameter(Embedding::kParamWeightName)->set_sequence_parallel(true); + } } else if (config_.position_embedding_type != PositionEmbeddingType::kRoPE) { LOG(FATAL) << "Unsupported position embedding type"; } diff --git a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc index 523bcf2d7..d1ba03511 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc @@ -107,6 +107,7 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads(const AddShardParam &add auto param_piece = std::make_shared(*bucket_param, param_piece_offset_bytes, std::vector{static_cast(piece_numel)}); + param_piece->set_sequence_parallel(param->sequence_parallel()); auto grad_piece = std::make_shared(*bucket_grad, grad_piece_offset_bytes, std::vector{static_cast(piece_numel)}); @@ -170,15 +171,18 @@ float DistributedOptimizer::learning_rate() const { return Optimizer::learning_rate(); } -void DistributedOptimizer::Step() { - // 1. Ensure grads are synced +void DistributedOptimizer::FinalizeModelGrads() { FinishGradSync(); + if (model_grad_finalizer_) { + model_grad_finalizer_(base_optimizer_->parameters()); + } +} - // 2. Base optimizer step on owned param pieces +void DistributedOptimizer::StepImpl() { CHECK(base_optimizer_) << "DistributedOptimizer: base optimizer is null."; base_optimizer_->Step(); - // 3. Gather updated param shards back to full params + // Gather updated param shards back to full params StartParamSync(/*force_sync=*/false); // TODO(zbl): Delay sync call until param is actually used in next step FinishParamSync(/*skip_next_bucket_dispatch=*/true); diff --git a/infini_train/src/nn/parallel/tensor_parallel.cc b/infini_train/src/nn/parallel/tensor_parallel.cc index b16c526e6..65c55bbba 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -305,6 +305,9 @@ RowParallelLinear::RowParallelLinear(int64_t in_features, int64_t out_features, if (bias) { parameters_[kParamBiasName] = std::make_shared(std::vector{out_features}, DataType::kFLOAT32, device_)->RequiresGrad(); + if (sequence_parallel_) { + parameters_[kParamBiasName]->set_sequence_parallel(true); + } } LinearResetParameters(parameters_[kParamWeightName], bias ? parameters_[kParamBiasName] : nullptr); diff --git a/infini_train/src/nn/parallel/utils.cc b/infini_train/src/nn/parallel/utils.cc index ee28a1694..7f6908ada 100644 --- a/infini_train/src/nn/parallel/utils.cc +++ b/infini_train/src/nn/parallel/utils.cc @@ -5,9 +5,38 @@ #include "infini_train/include/nn/functional.h" #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/nn/parallel/process_group.h" +#include "infini_train/include/nn/parallel/reduce_op_type.h" #include "infini_train/include/tensor.h" namespace infini_train::nn::parallel { +namespace { +const ProcessGroup *GetTensorParallelGroup(const Tensor &tensor) { + const int global_rank = tensor.GetDevice().Rank().GlobalRank(); + return ProcessGroupFactory::Instance(tensor.GetDevice().type()) + ->Get(GetTensorParallelProcessGroupName(global_rank)); +} + +void FinalizeSequenceParallelGradients(const std::vector> ¶ms) { + // SP replicas see different sequence shards, so replicated parameter grads + // must be summed across TP before the optimizer consumes them. + if (!global::GetSequenceParallelEnabled() || global::GetTensorParallelSize() <= 1 || params.empty()) { + return; + } + + const ProcessGroup *tp_group = nullptr; + for (const auto ¶m : params) { + if (!param || !param->sequence_parallel() || !param->grad()) { + continue; + } + + if (tp_group == nullptr) { + tp_group = GetTensorParallelGroup(*param); + CHECK_NOTNULL(tp_group); + } + tp_group->AllReduce(param->grad(), function::ReduceOpType::kSum, /*async_op=*/false); + } +} +} // namespace std::string GetDataParallelProcessGroupName(int global_rank) { return "DP" + std::to_string(global::GetGroupId(global::DP, global_rank)); @@ -60,4 +89,11 @@ std::shared_ptr GatherTensorParallelShard(const std::shared_ptr return nn::function::Concat(rank_major_shards, dim)->Contiguous(); } +void FinalizeModelGrads(const std::vector> ¶ms) { + FinalizeSequenceParallelGradients(params); + + // TODO(zbl): Future model-gradient finalization goes here, such as PP tied embeddings, + // MoE shared parameters, and loss normalization. +} + } // namespace infini_train::nn::parallel diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 39b999c77..2b0322020 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -29,6 +29,21 @@ void Optimizer::ZeroGrad(bool set_to_none) { for (auto param : params_) { param->ZeroGrad(set_to_none); } } +void Optimizer::Step() { + FinalizeModelGrads(); + StepImpl(); +} + +// NOTE(zbl): Centralize model-gradient finalization here. This is the boundary +// for gradients that a regular DDP reducer cannot express automatically, +// such as SP norms, PP tied embeddings, MoE shared parameters, and loss +// normalization. +void Optimizer::FinalizeModelGrads() { + if (model_grad_finalizer_) { + model_grad_finalizer_(params_); + } +} + void Optimizer::set_learning_rate(float lr) { learning_rate_ = lr; } float Optimizer::learning_rate() const { return learning_rate_; } @@ -53,7 +68,7 @@ SGD::SGD(const std::vector> ¶ms, float learning_rate SGD::SGD(const NamedParameterList &named_params, float learning_rate) : Optimizer(named_params, learning_rate) {} -void SGD::Step() { +void SGD::StepImpl() { for (auto param : params_) { if (!param->grad()) { LOG(INFO) << "Skipping param with null grad."; @@ -99,7 +114,7 @@ Adam::Adam(const NamedParameterList &named_params, float learning_rate, float be } } -void Adam::Step() { +void Adam::StepImpl() { ++t_; for (size_t i = 0; i < params_.size(); ++i) { diff --git a/infini_train/src/tensor.cc b/infini_train/src/tensor.cc index 4e61e221b..e9a3f4626 100644 --- a/infini_train/src/tensor.cc +++ b/infini_train/src/tensor.cc @@ -57,6 +57,7 @@ Tensor::Tensor(const Tensor &tensor, size_t offset, const std::vector & : buffer_(tensor.buffer_), offset_(tensor.offset_ + offset), dims_(dims), num_elements_(std::accumulate(dims.begin(), dims.end(), 1, std::multiplies())), dtype_(tensor.dtype_) { CHECK_LE(offset_ + kDataTypeToSize.at(dtype_) * num_elements_, buffer_->Size()); + sequence_parallel_ = tensor.sequence_parallel_; } Tensor::Tensor(const float *data, const std::vector &dims, DataType dtype, Device device) @@ -170,6 +171,7 @@ Tensor Tensor::To(Device device) { } new_tensor.requires_grad_ = requires_grad_; + new_tensor.sequence_parallel_ = sequence_parallel_; return new_tensor; } @@ -194,6 +196,7 @@ Tensor Tensor::To(DataType dtype) { } new_tensor.requires_grad_ = requires_grad_; + new_tensor.sequence_parallel_ = sequence_parallel_; return new_tensor; }