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
2 changes: 1 addition & 1 deletion docs/lora_usage_guide.md
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ model = GetLoRAModel(model, config);
PrintLoRASummary(model);

auto params = GetLoRAParameters(model);
auto optimizer = infini_train::optimizers::Adam::Create(/*learning_rate=*/1e-4)(params);
auto optimizer = infini_train::optimizers::Adam::Create(/*learning_rate=*/1e-4)(params, model->NamedParameters());

for (int step = 0; step < num_steps; ++step) {
optimizer->ZeroGrad();
Expand Down
17 changes: 14 additions & 3 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -327,17 +327,28 @@ void Train(const nn::parallel::Rank &rank) {

// TODO(dcj): support more complex optimizer later
// auto optimizer = optimizers::SGD(model->Parameters(), FLAGS_learning_rate);
auto optimizer_creator = optimizers::SGD::Create(FLAGS_learning_rate);
auto optimizer_creator = optimizers::SGD::CreateNamed(FLAGS_learning_rate);
std::shared_ptr<Optimizer> optimizer = nullptr;
std::unordered_set<const Tensor *> params_to_optimize_set;
params_to_optimize_set.reserve(params_to_optimize.size());
for (const auto &param : params_to_optimize) { params_to_optimize_set.insert(param.get()); }

NamedParameterList named_parameters;
for (const auto &[name, param] : model->NamedParameters()) {
if (params_to_optimize_set.contains(param.get())) {
named_parameters.emplace_back(name, param);
}
}
CHECK_EQ(named_parameters.size(), params_to_optimize.size());

if (FLAGS_zero_stage >= 1) {
auto model_chunks = (pp_world_size > 1)
? *(dynamic_cast<nn::parallel::PipelineParallel *>(model.get())->mutable_chunks())
: std::vector<std::shared_ptr<nn::Module>>{model};
optimizer = std::make_shared<nn::parallel::DistributedOptimizer>(optimizer_creator, params_to_optimize,
optimizer = std::make_shared<nn::parallel::DistributedOptimizer>(optimizer_creator, named_parameters,
model_chunks, ddp_world_size, ddp_rank);
} else {
optimizer = optimizer_creator(params_to_optimize);
optimizer = optimizer_creator(named_parameters);
}

const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
Expand Down
17 changes: 14 additions & 3 deletions example/llama3/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -300,7 +300,7 @@ void Train(const nn::parallel::Rank &rank) {

// TODO(dcj): support more complex optimizer later
// auto optimizer = optimizers::Adam(model->Parameters(), FLAGS_learning_rate);
auto optimizer_creator = optimizers::Adam::Create(FLAGS_learning_rate);
auto optimizer_creator = optimizers::Adam::CreateNamed(FLAGS_learning_rate);
std::shared_ptr<Optimizer> optimizer = nullptr;

std::vector<std::shared_ptr<Tensor>> params_to_optimize;
Expand All @@ -311,15 +311,26 @@ void Train(const nn::parallel::Rank &rank) {
params_to_optimize = model->Parameters();
LOG(INFO) << "Optimizing " << params_to_optimize.size() << " model parameters";
}
std::unordered_set<const Tensor *> params_to_optimize_set;
params_to_optimize_set.reserve(params_to_optimize.size());
for (const auto &param : params_to_optimize) { params_to_optimize_set.insert(param.get()); }

NamedParameterList named_parameters;
for (const auto &[name, param] : model->NamedParameters()) {
if (params_to_optimize_set.contains(param.get())) {
named_parameters.emplace_back(name, param);
}
}
CHECK_EQ(named_parameters.size(), params_to_optimize.size());

if (FLAGS_zero_stage >= 1) {
auto model_chunks = (pp_world_size > 1)
? *(dynamic_cast<nn::parallel::PipelineParallel *>(model.get())->mutable_chunks())
: std::vector<std::shared_ptr<nn::Module>>{model};
optimizer = std::make_shared<nn::parallel::DistributedOptimizer>(optimizer_creator, params_to_optimize,
optimizer = std::make_shared<nn::parallel::DistributedOptimizer>(optimizer_creator, named_parameters,
model_chunks, ddp_world_size, ddp_rank);
} else {
optimizer = optimizer_creator(params_to_optimize);
optimizer = optimizer_creator(named_parameters);
}

const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
Expand Down
4 changes: 2 additions & 2 deletions example/mixtral/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -104,8 +104,8 @@ int main(int argc, char *argv[]) {
}

auto loss_fn = std::make_shared<infini_train::nn::CrossEntropyLoss>();
auto optimizer
= infini_train::optimizers::Adam::Create(static_cast<float>(FLAGS_learning_rate))(model->Parameters());
auto optimizer = infini_train::optimizers::Adam::CreateNamed(static_cast<float>(FLAGS_learning_rate))(
model->NamedParameters());

auto device_impl = infini_train::core::GetDeviceGuardImpl(train_device.type());
std::vector<double> step_duration_ms;
Expand Down
6 changes: 5 additions & 1 deletion infini_train/include/nn/modules/module.h
Original file line number Diff line number Diff line change
Expand Up @@ -47,8 +47,12 @@ class Module : public std::enable_shared_from_this<Module> {

const std::string &type() const;

// TODO: Change return type to filterable iterator (like PyTorch's named_parameters with prefix matching)
virtual std::vector<std::shared_ptr<Tensor>> Parameters() const;

// InfiniTrain's NamedParameters returns results ordered by full parameter name.
// TODO: Align with PyTorch's ordering in the future.
std::vector<std::pair<std::string, std::shared_ptr<Tensor>>>
NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const;
bool has_parameter(const std::string &name) const;
std::shared_ptr<Tensor> *mutable_parameter(const std::string &name);
const std::shared_ptr<Tensor> &parameter(const std::string &name) const;
Expand Down
13 changes: 9 additions & 4 deletions infini_train/include/nn/parallel/ddp/distributed_optimizer.h

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.

DistributedOptimizer 的改动麻烦 @Chamberlain0w0 也看一下

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

正确性上没问题,但是构造函数里面同时传入 full_params 和 named_parameters 有点奇怪,这两者并不会同时用到的话,感觉跟普通 Optimizer 构造方式一样,提供两种方式 overload 比较合适。

DistributedOptimizer(
    OptimizerCreator creator,
    const std::vector<std::shared_ptr<Tensor>> &params,
    const std::vector<std::shared_ptr<Module>> &model_chunks,
    size_t ddp_world_size,
    size_t ddp_rank);

DistributedOptimizer(
    OptimizerCreator creator,
    const NamedParameterList &named_params,
    const std::vector<std::shared_ptr<Module>> &model_chunks,
    size_t ddp_world_size,
    size_t ddp_rank);

Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
#pragma once

#include <cstdint>
#include <functional>
#include <memory>
#include <unordered_map>
#include <vector>
Expand All @@ -24,6 +25,10 @@ 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);

DistributedOptimizer(OptimizerCreatorNamed base_optimizer_creator, const NamedParameterList &named_parameters,
const std::vector<std::shared_ptr<Module>> &model_chunks, size_t ddp_world_size,
size_t ddp_rank);

void Step() override;

void ZeroGrad(bool set_to_none = true) override;
Expand All @@ -42,7 +47,10 @@ class DistributedOptimizer final : public infini_train::Optimizer {
virtual float learning_rate() const override;

private:
void BuildShardParamsAndBindGrads();
using AddShard = std::function<void(const std::shared_ptr<Tensor> &, const std::shared_ptr<Tensor> &)>;

void InitializeModelChunks(const std::vector<std::shared_ptr<Module>> &model_chunks);
void BuildShardParamsAndBindGrads(const AddShard &add_shard);

private:
// Inherit from DDP model
Expand All @@ -53,9 +61,6 @@ class DistributedOptimizer final : public infini_train::Optimizer {
size_t ddp_world_size_;
size_t ddp_rank_;

// shard params
std::vector<std::shared_ptr<Tensor>> shard_params_;

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.

当时加这个成员变量时好像就讨论过,麻烦 @Chamberlain0w0 确认下这样修改是否合适。

我看目前是将 shard_params 作为了 BuildShardParamsAndBindGrads 参数传入,就不需要 DistributedOptimizer 维护了,似乎也更合理,因为 base_optimizer_ 本身已经维护了分片参数,没必要在 DistributedOptimizer 里额外维护一份。

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

确实,这里可以删掉了。


// Base optimizer (SGD, Adam and etc.)
std::shared_ptr<Optimizer> base_optimizer_;
};
Expand Down
15 changes: 14 additions & 1 deletion infini_train/include/optimizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
#include <memory>
#include <string>
#include <unordered_map>
#include <utility>
#include <vector>

namespace infini_train {
Expand All @@ -13,11 +14,16 @@ class Tensor;
namespace infini_train {
class Optimizer;

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)>;

class Optimizer {
public:
explicit Optimizer(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate = 0.0f);
explicit Optimizer(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate);

Optimizer(const NamedParameterList &named_params, float learning_rate);

virtual void ZeroGrad(bool set_to_none = true);

Expand All @@ -39,6 +45,7 @@ class Optimizer {

protected:
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;
Expand All @@ -48,16 +55,20 @@ namespace optimizers {
class SGD : public Optimizer {
public:
SGD(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate);
SGD(const NamedParameterList &named_params, float learning_rate);

void Step() override;

static OptimizerCreator Create(float learning_rate);
static OptimizerCreatorNamed CreateNamed(float learning_rate);
};

class Adam : public Optimizer {
public:
Adam(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate = 1e-3, float beta1 = 0.9,
float beta2 = 0.999, float eps = 1e-8);
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;

Expand All @@ -66,6 +77,8 @@ class Adam : public Optimizer {
void LoadStateDict(const std::unordered_map<std::string, std::shared_ptr<Tensor>> &state_dict) override;
static OptimizerCreator Create(float learning_rate = 1e-3, float beta1 = 0.9, float beta2 = 0.999,
float eps = 1e-8);
static OptimizerCreatorNamed CreateNamed(float learning_rate = 1e-3, float beta1 = 0.9, float beta2 = 0.999,
float eps = 1e-8);

private:
int64_t t_;
Expand Down
53 changes: 42 additions & 11 deletions infini_train/src/nn/modules/module.cc
Original file line number Diff line number Diff line change
Expand Up @@ -28,24 +28,55 @@ Module::Module(const std::string &type) : type_(type), device_(Device()) {}
const std::string &Module::type() const { return type_; }

std::vector<std::shared_ptr<Tensor>> Module::Parameters() const {
const auto &named_parameters = NamedParameters();

std::vector<std::shared_ptr<Tensor>> params;
std::unordered_set<const Tensor *> visited;
params.reserve(named_parameters.size());

for (const auto &[_, param] : named_parameters) { params.emplace_back(param); }

return params;
}

auto AddIfUnvisited = [&](const std::shared_ptr<Tensor> &param) {
if (visited.insert(param.get()).second) {
params.push_back(param);
std::vector<std::pair<std::string, std::shared_ptr<Tensor>>>
Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const {

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.

std::vector<std::pair<std::string, std::shared_ptr<Tensor>>> named_parameters;
std::unordered_set<const Tensor *> visited_parameters;

std::vector<std::pair<std::string, std::shared_ptr<Module>>> named_modules;

if (recurse) {
named_modules = const_cast<Module *>(this)->NamedModules(
/*memory=*/nullptr, prefix, remove_duplicate);
} else {
named_modules.emplace_back(prefix, std::const_pointer_cast<Module>(shared_from_this()));
}

for (const auto &[module_prefix, module] : named_modules) {
std::vector<std::pair<std::string, std::shared_ptr<Tensor>>> local_parameters;
local_parameters.reserve(module->parameters_.size());

for (const auto &[name, parameter] : module->parameters_) {
if (parameter != nullptr) {
local_parameters.emplace_back(name, parameter);
}
}
};

// Add parameters of this module
for (const auto &[_, param] : parameters_) { AddIfUnvisited(param); }
std::sort(local_parameters.begin(), local_parameters.end(),
[](const auto &lhs, const auto &rhs) { return lhs.first < rhs.first; });

for (const auto &[name, parameter] : local_parameters) {
if (remove_duplicate && !visited_parameters.insert(parameter.get()).second) {
continue;
}

const std::string full_name = module_prefix.empty() ? name : module_prefix + "." + name;

// Recursively add parameters of submodules
for (const auto &[_, module] : modules_) {
for (const auto &param : module->Parameters()) { AddIfUnvisited(param); }
named_parameters.emplace_back(full_name, parameter);
}
}

return params;
return named_parameters;
}

bool Module::has_parameter(const std::string &name) const { return parameters_.find(name) != parameters_.end(); }
Expand Down
56 changes: 45 additions & 11 deletions infini_train/src/nn/parallel/ddp/distributed_optimizer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,47 @@ DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator,
const std::vector<std::shared_ptr<Tensor>> &full_params,
const std::vector<std::shared_ptr<Module>> &model_chunks,
size_t ddp_world_size, size_t ddp_rank)
: Optimizer(full_params), ddp_world_size_(ddp_world_size), ddp_rank_(ddp_rank) {
: Optimizer(full_params, /*learning_rate=*/0.0f), ddp_world_size_(ddp_world_size), ddp_rank_(ddp_rank) {
InitializeModelChunks(model_chunks);

std::vector<std::shared_ptr<Tensor>> shard_params;
BuildShardParamsAndBindGrads(
[&shard_params](const std::shared_ptr<Tensor> &, const std::shared_ptr<Tensor> &param_piece) {
shard_params.push_back(param_piece);
});

base_optimizer_ = creator(shard_params);
CHECK(base_optimizer_) << "DistributedOptimizer: failed to create base optimizer.";
}

DistributedOptimizer::DistributedOptimizer(OptimizerCreatorNamed creator, const NamedParameterList &named_parameters,
const std::vector<std::shared_ptr<Module>> &model_chunks,
size_t ddp_world_size, size_t ddp_rank)
: Optimizer(named_parameters, /*learning_rate=*/0.0f), ddp_world_size_(ddp_world_size), ddp_rank_(ddp_rank) {
InitializeModelChunks(model_chunks);

std::unordered_map<const Tensor *, std::string> parameter_name_by_tensor;
parameter_name_by_tensor.reserve(named_parameters.size());
for (const auto &[name, parameter] : named_parameters) {
CHECK(parameter);
parameter_name_by_tensor.emplace(parameter.get(), name);
}

NamedParameterList shard_named_parameters;
BuildShardParamsAndBindGrads(
[&parameter_name_by_tensor, &shard_named_parameters](const std::shared_ptr<Tensor> &parameter,
const std::shared_ptr<Tensor> &param_piece) {
const auto name_it = parameter_name_by_tensor.find(parameter.get());
CHECK(name_it != parameter_name_by_tensor.end())
<< "DistributedOptimizer parameter is not registered in the model";
shard_named_parameters.emplace_back(name_it->second, param_piece);
});

base_optimizer_ = creator(shard_named_parameters);
CHECK(base_optimizer_) << "DistributedOptimizer: failed to create base optimizer.";
}

void DistributedOptimizer::InitializeModelChunks(const std::vector<std::shared_ptr<Module>> &model_chunks) {
CHECK(ddp_world_size_ > 1) << "DistributedOptimizer: ddp_world_size must be greater than 1.";

for (size_t i = 0; i < model_chunks.size(); ++i) {
Expand All @@ -23,16 +62,10 @@ DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator,
bucket_groups_.insert(bucket_groups_.end(), ddp_chunk->bucket_groups().begin(),
ddp_chunk->bucket_groups().end());
}

BuildShardParamsAndBindGrads();

// Build base optimizer
base_optimizer_ = creator(shard_params_);
CHECK(base_optimizer_) << "DistributedOptimizer: failed to create base optimizer.";
}

void DistributedOptimizer::BuildShardParamsAndBindGrads() {
shard_params_.clear();
void DistributedOptimizer::BuildShardParamsAndBindGrads(const AddShard &add_shard) {
size_t num_shard_params = 0;

for (const auto &group : bucket_groups_) {
const bool use_grad_shard = group->config().zero_stage >= 2;
Expand Down Expand Up @@ -82,12 +115,13 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads() {
// NOTE(zbl): Do not call `param->set_grad(grad_piece);` under ZeRO-2.
// The base optimizer updates param_piece views only; original param->grad()
// would be a partial flattened shard and does not represent the full parameter grad.
shard_params_.push_back(param_piece);
add_shard(param, param_piece);
++num_shard_params;
}
}
}

CHECK(!shard_params_.empty()) << "DistributedOptimizer: this DP rank owns no param pieces. "
CHECK_GT(num_shard_params, 0) << "DistributedOptimizer: this DP rank owns no param pieces. "
<< "Check bucket padding/divisibility and param bucketing order.";
}

Expand Down
Loading
Loading