Skip to content
Merged
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
Comment thread
kilinchange marked this conversation as resolved.
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 AddShardParam = 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 AddShardParam &add_shard_param);

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_;
Comment thread
kilinchange marked this conversation as resolved.

// 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 {
Comment thread
kilinchange marked this conversation as resolved.
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 AddShardParam &add_shard_param) {
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, 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