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
7 changes: 4 additions & 3 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -329,15 +329,16 @@ void Train(const nn::parallel::Rank &rank) {
// auto optimizer = optimizers::SGD(model->Parameters(), FLAGS_learning_rate);
auto optimizer_creator = optimizers::SGD::Create(FLAGS_learning_rate);
std::shared_ptr<Optimizer> optimizer = nullptr;
const auto named_parameters = model->NamedParameters();

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,
model_chunks, ddp_world_size, ddp_rank);
optimizer = std::make_shared<nn::parallel::DistributedOptimizer>(
optimizer_creator, params_to_optimize, named_parameters, model_chunks, ddp_world_size, ddp_rank);
} else {
optimizer = optimizer_creator(params_to_optimize);
optimizer = optimizer_creator(params_to_optimize, named_parameters);
}

const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
Expand Down
7 changes: 4 additions & 3 deletions example/llama3/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -311,15 +311,16 @@ void Train(const nn::parallel::Rank &rank) {
params_to_optimize = model->Parameters();
LOG(INFO) << "Optimizing " << params_to_optimize.size() << " model parameters";
}
const auto named_parameters = model->NamedParameters();

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,
model_chunks, ddp_world_size, ddp_rank);
optimizer = std::make_shared<nn::parallel::DistributedOptimizer>(
optimizer_creator, params_to_optimize, named_parameters, model_chunks, ddp_world_size, ddp_rank);
} else {
optimizer = optimizer_creator(params_to_optimize);
optimizer = optimizer_creator(params_to_optimize, 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::Create(static_cast<float>(FLAGS_learning_rate))(
model->Parameters(), model->NamedParameters());

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

// TODO: Change return type to filterable iterator (like PyTorch's named_parameters with prefix matching)
virtual std::vector<std::shared_ptr<Tensor>> Parameters() const;
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
3 changes: 3 additions & 0 deletions infini_train/include/nn/parallel/ddp/distributed_optimizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ class DistributedOptimizer final : public infini_train::Optimizer {
public:
DistributedOptimizer(OptimizerCreator base_optimizer_creator,
const std::vector<std::shared_ptr<Tensor>> &full_params,
const NamedParameterList &named_parameters,
const std::vector<std::shared_ptr<Module>> &model_chunks, size_t ddp_world_size,
size_t ddp_rank);

Expand Down Expand Up @@ -55,6 +56,8 @@ class DistributedOptimizer final : public infini_train::Optimizer {

// shard params
std::vector<std::shared_ptr<Tensor>> shard_params_;
NamedParameterList shard_named_parameters_;
std::unordered_map<const Tensor *, std::string> parameter_name_by_tensor_;

// Base optimizer (SGD, Adam and etc.)
std::shared_ptr<Optimizer> base_optimizer_;
Expand Down
17 changes: 13 additions & 4 deletions 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,15 @@ class Tensor;
namespace infini_train {
class Optimizer;

using OptimizerCreator = std::function<std::shared_ptr<Optimizer>(const std::vector<std::shared_ptr<Tensor>> &params)>;
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,
const NamedParameterList &named_parameters)>;

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 = 0.0f,
const NamedParameterList &named_parameters = {});

virtual void ZeroGrad(bool set_to_none = true);

Expand All @@ -37,8 +42,11 @@ class Optimizer {

void set_initial_learning_rate(float lr);

void set_parameter_names(const std::vector<std::string> &names);

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 @@ -47,7 +55,8 @@ class Optimizer {
namespace optimizers {
class SGD : public Optimizer {
public:
SGD(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate);
SGD(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate,
const NamedParameterList &named_parameters = {});

void Step() override;

Expand All @@ -57,7 +66,7 @@ class SGD : public Optimizer {
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);
float beta2 = 0.999, float eps = 1e-8, const NamedParameterList &named_parameters = {});

void Step() override;

Expand Down
37 changes: 37 additions & 0 deletions infini_train/src/nn/modules/module.cc
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,43 @@ std::vector<std::shared_ptr<Tensor>> Module::Parameters() const {
return params;
}

std::vector<std::pair<std::string, std::shared_ptr<Tensor>>>
Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const {
std::vector<std::pair<std::string, std::shared_ptr<Tensor>>> named_parameters;
std::unordered_set<const Tensor *> visited;

std::vector<std::pair<std::string, std::shared_ptr<Module>>> named_modules;
if (recurse) {
// NamedModules only reads the hierarchy and provides its stable, name-sorted traversal order. Keep all module
// aliases here so parameter-level deduplication deterministically selects the first full parameter name.
named_modules
= const_cast<Module *>(this)->NamedModules(/*memory=*/nullptr, prefix, /*remove_duplicate=*/false);
} 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) {
local_parameters.emplace_back(name, parameter);
}
}
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.insert(parameter.get()).second) {
continue;
}
const auto full_name = module_prefix.empty() ? name : module_prefix + "." + name;
named_parameters.emplace_back(full_name, parameter);
}
}
return named_parameters;
}

bool Module::has_parameter(const std::string &name) const { return parameters_.find(name) != parameters_.end(); }

std::shared_ptr<Tensor> *Module::mutable_parameter(const std::string &name) {
Expand Down
16 changes: 14 additions & 2 deletions infini_train/src/nn/parallel/ddp/distributed_optimizer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,18 @@
namespace infini_train::nn::parallel {
DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator,
const std::vector<std::shared_ptr<Tensor>> &full_params,
const NamedParameterList &named_parameters,
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, named_parameters), ddp_world_size_(ddp_world_size),
ddp_rank_(ddp_rank) {

CHECK(ddp_world_size_ > 1) << "DistributedOptimizer: ddp_world_size must be greater than 1.";
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);
}

for (size_t i = 0; i < model_chunks.size(); ++i) {
auto ddp_chunk = std::dynamic_pointer_cast<DistributedDataParallel>(model_chunks[i]);
Expand All @@ -27,12 +34,13 @@ DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator,
BuildShardParamsAndBindGrads();

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

void DistributedOptimizer::BuildShardParamsAndBindGrads() {
shard_params_.clear();
shard_named_parameters_.clear();

for (const auto &group : bucket_groups_) {
const bool use_grad_shard = group->config().zero_stage >= 2;
Expand Down Expand Up @@ -83,6 +91,10 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads() {
// 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);
const auto name_it = parameter_name_by_tensor_.find(param.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);
}
}
}
Expand Down
59 changes: 45 additions & 14 deletions infini_train/src/optimizer.cc
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
#include "infini_train/include/optimizer.h"

#include <format>
#include <unordered_map>
#include <vector>

#include "infini_train/include/core/runtime/device_guard.h"
Expand All @@ -9,8 +9,27 @@
#include "infini_train/include/tensor.h"

namespace infini_train {
Optimizer::Optimizer(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate)
: params_(params), learning_rate_(learning_rate) {}
Optimizer::Optimizer(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate,
const NamedParameterList &named_parameters)
: params_(params), learning_rate_(learning_rate) {
if (named_parameters.empty()) {
return;
}

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

parameter_names_.reserve(params_.size());
for (const auto &parameter : params_) {
const auto it = parameter_name_by_tensor.find(parameter.get());
CHECK(it != parameter_name_by_tensor.end()) << "Optimizer parameter is not registered in the model";
parameter_names_.push_back(it->second);
}
}

void Optimizer::ZeroGrad(bool set_to_none) {
for (auto param : params_) { param->ZeroGrad(set_to_none); }
Expand All @@ -33,9 +52,17 @@ void Optimizer::set_initial_learning_rate(float lr) {
initial_learning_rate_ = lr;
initial_lr_set_ = true;
}

void Optimizer::set_parameter_names(const std::vector<std::string> &names) {
CHECK_EQ(names.size(), params_.size());
parameter_names_ = names;
}

namespace optimizers {

SGD::SGD(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate) : Optimizer(params, learning_rate) {}
SGD::SGD(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate,
const NamedParameterList &named_parameters)
: Optimizer(params, learning_rate, named_parameters) {}

void SGD::Step() {
for (auto param : params_) {
Expand All @@ -51,13 +78,15 @@ void SGD::Step() {
}

OptimizerCreator SGD::Create(float learning_rate) {
return [learning_rate](const std::vector<std::shared_ptr<Tensor>> &params) {
return std::make_shared<SGD>(params, learning_rate);
return [learning_rate](const std::vector<std::shared_ptr<Tensor>> &params,
const NamedParameterList &named_parameters) {
return std::make_shared<SGD>(params, learning_rate, named_parameters);
};
}

Adam::Adam(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate, float beta1, float beta2, float eps)
: Optimizer(params, learning_rate), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) {
Adam::Adam(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate, float beta1, float beta2, float eps,
const NamedParameterList &named_parameters)
: Optimizer(params, learning_rate, named_parameters), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) {

for (const auto &param : params_) {
m_.emplace_back(std::make_shared<Tensor>(param->Dims(), param->Dtype(), param->GetDevice()));
Expand Down Expand Up @@ -88,16 +117,17 @@ void Adam::Step() {
}

OptimizerCreator Adam::Create(float learning_rate, float beta1, float beta2, float eps) {
return [=](const std::vector<std::shared_ptr<Tensor>> &params) {
return std::make_shared<Adam>(params, learning_rate, beta1, beta2, eps);
return [=](const std::vector<std::shared_ptr<Tensor>> &params, const NamedParameterList &named_parameters) {
return std::make_shared<Adam>(params, learning_rate, beta1, beta2, eps, named_parameters);
};
}

std::unordered_map<std::string, std::shared_ptr<Tensor>> Adam::StateDict() const {
std::unordered_map<std::string, std::shared_ptr<Tensor>> state;
for (size_t i = 0; i < m_.size(); ++i) {
state.emplace(std::format("adam.m.{}", i), m_[i]);
state.emplace(std::format("adam.v.{}", i), v_[i]);
const auto suffix = parameter_names_.empty() ? std::to_string(i) : parameter_names_[i];
state.emplace("adam.m." + suffix, m_[i]);
state.emplace("adam.v." + suffix, v_[i]);
}

auto t_tensor = std::make_shared<Tensor>(std::vector<int64_t>{}, DataType::kINT64, Device());
Expand All @@ -108,8 +138,9 @@ std::unordered_map<std::string, std::shared_ptr<Tensor>> Adam::StateDict() const

void Adam::LoadStateDict(const std::unordered_map<std::string, std::shared_ptr<Tensor>> &state_dict) {
for (size_t i = 0; i < m_.size(); ++i) {
const auto m_key = std::format("adam.m.{}", i);
const auto v_key = std::format("adam.v.{}", i);
const auto suffix = parameter_names_.empty() ? std::to_string(i) : parameter_names_[i];
const auto m_key = "adam.m." + suffix;
const auto v_key = "adam.v." + suffix;
CHECK(state_dict.contains(m_key)) << "Missing optimizer state: " << m_key;
CHECK(state_dict.contains(v_key)) << "Missing optimizer state: " << v_key;
m_[i]->CopyFrom(state_dict.at(m_key));
Expand Down
3 changes: 3 additions & 0 deletions tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,9 @@ include(${CMAKE_SOURCE_DIR}/cmake/test_macros.cmake)
# Common test utilities
add_subdirectory(common)

# Module tests
add_subdirectory(module)

# Tensor tests
add_subdirectory(tensor)

Expand Down
5 changes: 5 additions & 0 deletions tests/module/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
file(GLOB MODULE_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/test_*.cc)

infini_train_add_test_suite(test_module
SOURCES ${MODULE_SOURCES}
)
Loading
Loading