Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

FinalizeModelGrads 绑定到 optimizer 上感觉职责不太合适。optimizer 接收到的 grad 应该已经是 finalize 后可用于更新的完整梯度,本身只负责后续的 clip / step 等操作。

另外,将 FinalizeModelGrads 放到 optimizer->Step() 内部也容易影响 clip 执行顺序(理论上应该保证:
FinalizeModelGrads -> clip -> step,因为 clip 需要基于最终同步完成的梯度计算)。

建议目前在 optimizer->Step() 之前显式调用 nn::parallel::FinalizeModelGrads(named_parameters);后续如果抽象统一的 training/schedule 入口,再将这部分逻辑收进去。Megatron 也是在 forward/backward schedule 结束后调用 finalize_model_grads,而不是绑定到 optimizer:

(no pipeline schedule 情况,带 pipeline 时类似)https://github.com/NVIDIA/Megatron-LM/blob/0ac6ffd3859fef41bbfcd92a67d56861a79d4b34/megatron/core/pipeline_parallel/schedules.py#L866


const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
TrainingLRSchedulerConfig sched_config;
Expand Down
1 change: 1 addition & 0 deletions example/llama3/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
3 changes: 2 additions & 1 deletion infini_train/include/nn/parallel/ddp/distributed_optimizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,8 @@ class DistributedOptimizer final : public infini_train::Optimizer {
const std::vector<std::shared_ptr<Module>> &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;

Expand Down
2 changes: 2 additions & 0 deletions infini_train/include/nn/parallel/utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -31,4 +31,6 @@ std::vector<std::shared_ptr<Tensor>> GatherFromSPRegionFunc(const std::shared_pt
std::vector<std::shared_ptr<Tensor>> ScatterToTPRegionFunc(const std::shared_ptr<Tensor> &input);
std::vector<std::shared_ptr<Tensor>> ReduceFromTPRegionFunc(const std::shared_ptr<Tensor> &input);
std::vector<std::shared_ptr<Tensor>> CopyToTPRegionFunc(const std::shared_ptr<Tensor> &input);

void FinalizeModelGrads(const std::vector<std::shared_ptr<Tensor>> &params);
} // namespace infini_train::nn::parallel
16 changes: 13 additions & 3 deletions infini_train/include/optimizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ using NamedParameter = std::pair<std::string, std::shared_ptr<Tensor>>;
using NamedParameterList = std::vector<NamedParameter>;
using OptimizerCreator = std::function<std::shared_ptr<Optimizer>(const std::vector<std::shared_ptr<Tensor>> &params)>;
using OptimizerCreatorNamed = std::function<std::shared_ptr<Optimizer>(const NamedParameterList &named_params)>;
using ModelGradFinalizer = std::function<void(const std::vector<std::shared_ptr<Tensor>> &params)>;

class Optimizer {
public:
Expand All @@ -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<std::shared_ptr<Tensor>> &parameters() const { return params_; }

virtual std::unordered_map<std::string, std::shared_ptr<Tensor>> StateDict() const { return {}; };

Expand All @@ -44,11 +49,16 @@ class Optimizer {
void set_initial_learning_rate(float lr);

protected:
virtual void FinalizeModelGrads();
virtual void StepImpl() = 0;

std::vector<std::shared_ptr<Tensor>> params_;
std::vector<std::string> parameter_names_;
float learning_rate_ = 0.0f;
float initial_learning_rate_ = 0.0f;
bool initial_lr_set_ = false;

ModelGradFinalizer model_grad_finalizer_;
};

namespace optimizers {
Expand All @@ -57,7 +67,7 @@ class SGD : public Optimizer {
SGD(const std::vector<std::shared_ptr<Tensor>> &params, 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);
Expand All @@ -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<std::string, std::shared_ptr<Tensor>> StateDict() const override;

Expand Down
4 changes: 4 additions & 0 deletions infini_train/include/tensor.h
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,9 @@ class Tensor : public std::enable_shared_from_this<Tensor> {
size_t NumElements() const;
DataType Dtype() const;

void set_sequence_parallel(bool enabled) { sequence_parallel_ = enabled; }
bool sequence_parallel() const { return sequence_parallel_; }

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

实现放 .cc 里。


std::shared_ptr<Tensor> Detach() const;

void Fill(Scalar value);
Expand Down Expand Up @@ -242,6 +245,7 @@ class Tensor : public std::enable_shared_from_this<Tensor> {
private:
std::shared_ptr<Tensor> grad_ = nullptr;
bool requires_grad_ = false;
bool sequence_parallel_ = false;
bool is_leaf_ = true;
std::shared_ptr<autograd::Function> grad_fn_ = nullptr;
int output_idx_ = 0;
Expand Down
3 changes: 3 additions & 0 deletions infini_train/src/nn/modules/normalization.cc

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里得根据 parallel::global::GetSequenceParallelEnabled() 决定是否 set 吧

Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@ LayerNorm::LayerNorm(const std::vector<int64_t> &normalized_shape, float eps, De
= std::make_shared<Tensor>(normalized_shape, DataType::kFLOAT32, device_)->RequiresGrad();
parameters_[kParamBiasName]
= std::make_shared<Tensor>(normalized_shape, DataType::kFLOAT32, device_)->RequiresGrad();
parameters_[kParamWeightName]->set_sequence_parallel(true);
parameters_[kParamBiasName]->set_sequence_parallel(true);
ResetParameters();
}

Expand All @@ -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<Tensor>(std::vector<int64_t>{dim}, DataType::kFLOAT32, device)->RequiresGrad();
parameters_[kParamWeightName]->set_sequence_parallel(true);
nn::init::Ones(parameters_[kParamWeightName]);
}

Expand Down
7 changes: 7 additions & 0 deletions infini_train/src/nn/modules/transformer/moe/router.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -24,12 +25,18 @@ TopKRouter::TopKRouter(const TransformerConfig &config) : CloneableModule(kType)
= std::make_shared<Tensor>(std::vector<int64_t>{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<Tensor>(std::vector<int64_t>{moe_config.num_experts}, DataType::kFLOAT32, device_)
->RequiresGrad();
if (parallel::global::GetSequenceParallelEnabled()) {
parameters_[kParamBiasName]->set_sequence_parallel(true);
}
parameters_[kParamBiasName]->Fill(0.0f);
}
}
Expand Down
3 changes: 3 additions & 0 deletions infini_train/src/nn/modules/transformer/transformer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<Embedding>(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";
}
Expand Down
12 changes: 8 additions & 4 deletions infini_train/src/nn/parallel/ddp/distributed_optimizer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads(const AddShardParam &add

auto param_piece = std::make_shared<Tensor>(*bucket_param, param_piece_offset_bytes,
std::vector<int64_t>{static_cast<int64_t>(piece_numel)});
param_piece->set_sequence_parallel(param->sequence_parallel());

auto grad_piece = std::make_shared<Tensor>(*bucket_grad, grad_piece_offset_bytes,
std::vector<int64_t>{static_cast<int64_t>(piece_numel)});
Expand Down Expand Up @@ -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);
Expand Down
3 changes: 3 additions & 0 deletions infini_train/src/nn/parallel/tensor_parallel.cc
Original file line number Diff line number Diff line change
Expand Up @@ -305,6 +305,9 @@ RowParallelLinear::RowParallelLinear(int64_t in_features, int64_t out_features,
if (bias) {
parameters_[kParamBiasName]
= std::make_shared<Tensor>(std::vector<int64_t>{out_features}, DataType::kFLOAT32, device_)->RequiresGrad();
if (sequence_parallel_) {
parameters_[kParamBiasName]->set_sequence_parallel(true);
}
}

LinearResetParameters(parameters_[kParamWeightName], bias ? parameters_[kParamBiasName] : nullptr);
Expand Down
36 changes: 36 additions & 0 deletions infini_train/src/nn/parallel/utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<std::shared_ptr<Tensor>> &params) {
// 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 &param : 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));
Expand Down Expand Up @@ -60,4 +89,11 @@ std::shared_ptr<Tensor> GatherTensorParallelShard(const std::shared_ptr<Tensor>
return nn::function::Concat(rank_major_shards, dim)->Contiguous();
}

void FinalizeModelGrads(const std::vector<std::shared_ptr<Tensor>> &params) {
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
19 changes: 17 additions & 2 deletions infini_train/src/optimizer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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_; }
Expand All @@ -53,7 +68,7 @@ SGD::SGD(const std::vector<std::shared_ptr<Tensor>> &params, 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.";
Expand Down Expand Up @@ -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) {
Expand Down
3 changes: 3 additions & 0 deletions infini_train/src/tensor.cc
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@ Tensor::Tensor(const Tensor &tensor, size_t offset, const std::vector<int64_t> &
: buffer_(tensor.buffer_), offset_(tensor.offset_ + offset), dims_(dims),
num_elements_(std::accumulate(dims.begin(), dims.end(), 1, std::multiplies<int64_t>())), 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<int64_t> &dims, DataType dtype, Device device)
Expand Down Expand Up @@ -170,6 +171,7 @@ Tensor Tensor::To(Device device) {
}

new_tensor.requires_grad_ = requires_grad_;
new_tensor.sequence_parallel_ = sequence_parallel_;

return new_tensor;
}
Expand All @@ -194,6 +196,7 @@ Tensor Tensor::To(DataType dtype) {
}

new_tensor.requires_grad_ = requires_grad_;
new_tensor.sequence_parallel_ = sequence_parallel_;

return new_tensor;
}
Expand Down
Loading