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
12 changes: 12 additions & 0 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,18 @@ void Train(const nn::parallel::Rank &rank) {
model_chunks, ddp_world_size, ddp_rank);
} else {
optimizer = optimizer_creator(params_to_optimize);
std::unordered_map<const Tensor *, std::string> parameter_name_by_tensor;
for (const auto &[name, parameter] : model->NamedParameters()) {
parameter_name_by_tensor.emplace(parameter.get(), name);
}
std::vector<std::string> parameter_names;
parameter_names.reserve(params_to_optimize.size());
for (const auto &parameter : params_to_optimize) {
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);
}
optimizer->set_parameter_names(parameter_names);
}

const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
Expand Down
12 changes: 12 additions & 0 deletions example/llama3/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -320,6 +320,18 @@ void Train(const nn::parallel::Rank &rank) {
model_chunks, ddp_world_size, ddp_rank);
} else {
optimizer = optimizer_creator(params_to_optimize);
std::unordered_map<const Tensor *, std::string> parameter_name_by_tensor;
for (const auto &[name, parameter] : model->NamedParameters()) {
parameter_name_by_tensor.emplace(parameter.get(), name);
}
std::vector<std::string> parameter_names;
parameter_names.reserve(params_to_optimize.size());
for (const auto &parameter : params_to_optimize) {
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);
}
optimizer->set_parameter_names(parameter_names);
}

const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
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
5 changes: 5 additions & 0 deletions infini_train/include/optimizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -37,8 +37,13 @@ class Optimizer {

void set_initial_learning_rate(float lr);

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

const std::vector<std::string> &parameter_names() const;

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 Down
35 changes: 35 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,41 @@ 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::function<void(const Module &, const std::string &)> collect
= [&](const Module &module, const std::string &module_prefix) {
for (const auto &[name, parameter] : module.parameters_) {
if (!parameter) {
continue;
}
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);
}

if (!recurse) {
return;
}

for (const auto &[name, child] : module.modules_) {
if (!child) {
continue;
}
const auto child_prefix = module_prefix.empty() ? name : module_prefix + "." + name;
collect(*child, child_prefix);
}
};

collect(*this, prefix);
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
19 changes: 14 additions & 5 deletions infini_train/src/optimizer.cc
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
#include "infini_train/include/optimizer.h"

#include <format>
#include <vector>

#include "infini_train/include/core/runtime/device_guard.h"
Expand Down Expand Up @@ -33,6 +32,14 @@ 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;
}

const std::vector<std::string> &Optimizer::parameter_names() const { return parameter_names_; }

namespace optimizers {

SGD::SGD(const std::vector<std::shared_ptr<Tensor>> &params, float learning_rate) : Optimizer(params, learning_rate) {}
Expand Down Expand Up @@ -96,8 +103,9 @@ OptimizerCreator Adam::Create(float learning_rate, float beta1, float beta2, flo
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 +116,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}
)
76 changes: 76 additions & 0 deletions tests/module/test_named_parameters.cc
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
#include <memory>
#include <string>
#include <unordered_map>
#include <unordered_set>
#include <vector>

#include "gtest/gtest.h"

#include "infini_train/include/nn/modules/container.h"
#include "infini_train/include/nn/modules/linear.h"
#include "infini_train/include/tensor.h"

#include "tests/common/test_utils.h"

using namespace infini_train;

class ModuleNamedParametersTest : public test::InfiniTrainTest {};

TEST_P(ModuleNamedParametersTest, SupportsPrefixAndNonRecursiveLookup) {
auto linear = std::make_shared<nn::Linear>(2, 3, /*bias=*/true, GetDevice());

const auto parameters = linear->NamedParameters("model", false);
const std::unordered_map<std::string, std::shared_ptr<Tensor>> by_name(parameters.begin(), parameters.end());

ASSERT_EQ(by_name.size(), 2);
EXPECT_EQ(by_name.at("model.weight"), linear->parameter(nn::Linear::kParamWeightName));
EXPECT_EQ(by_name.at("model.bias"), linear->parameter(nn::Linear::kParamBiasName));
}

TEST_P(ModuleNamedParametersTest, SupportsRecursionAndSharedParameterDeduplication) {
auto shared = std::make_shared<nn::Linear>(2, 3, /*bias=*/true, GetDevice());
auto root
= std::make_shared<nn::Sequential>(std::vector<std::shared_ptr<nn::Module>>{shared, shared});

const auto deduplicated = root->NamedParameters("model");
ASSERT_EQ(deduplicated.size(), 2);
std::unordered_set<const Tensor *> tensors;
for (const auto &[name, parameter] : deduplicated) {
EXPECT_TRUE(name == "model.0.weight" || name == "model.0.bias" || name == "model.1.weight"
|| name == "model.1.bias");
tensors.insert(parameter.get());
}
EXPECT_TRUE(tensors.contains(shared->parameter(nn::Linear::kParamWeightName).get()));
EXPECT_TRUE(tensors.contains(shared->parameter(nn::Linear::kParamBiasName).get()));

const auto aliases = root->NamedParameters("model", true, false);
const std::unordered_map<std::string, std::shared_ptr<Tensor>> by_name(aliases.begin(), aliases.end());
ASSERT_EQ(by_name.size(), 4);
EXPECT_EQ(by_name.at("model.0.weight"), by_name.at("model.1.weight"));
EXPECT_EQ(by_name.at("model.0.bias"), by_name.at("model.1.bias"));
}

TEST_P(ModuleNamedParametersTest, ReturnsNestedParametersWithoutAnOrderingContract) {
auto first = std::make_shared<nn::Linear>(2, 3, /*bias=*/false, GetDevice());
auto second = std::make_shared<nn::Linear>(3, 4, /*bias=*/false, GetDevice());
auto nested
= std::make_shared<nn::Sequential>(std::vector<std::shared_ptr<nn::Module>>{first, second});
auto root = std::make_shared<nn::Sequential>(
std::vector<std::shared_ptr<nn::Module>>{std::make_shared<nn::Linear>(2, 2, false, GetDevice()), nested});

const auto parameters = root->NamedParameters();
const std::unordered_map<std::string, std::shared_ptr<Tensor>> by_name(parameters.begin(), parameters.end());

ASSERT_EQ(by_name.size(), 3);
EXPECT_TRUE(by_name.contains("0.weight"));
EXPECT_TRUE(by_name.contains("1.0.weight"));
EXPECT_TRUE(by_name.contains("1.1.weight"));
}

TEST_P(ModuleNamedParametersTest, SkipsNullSubmodules) {
auto root = std::make_shared<nn::Sequential>(std::vector<std::shared_ptr<nn::Module>>{nullptr});

EXPECT_TRUE(root->NamedParameters().empty());
}

INFINI_TRAIN_REGISTER_TEST(ModuleNamedParametersTest);
50 changes: 50 additions & 0 deletions tests/optimizer/test_optimizer_parameter_names.cc
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
#include <memory>
#include <vector>

#include "gtest/gtest.h"

#include "infini_train/include/optimizer.h"
#include "infini_train/include/tensor.h"

#include "tests/common/test_utils.h"

using namespace infini_train;

class OptimizerParameterNamesTest : public test::InfiniTrainTest {};

TEST_P(OptimizerParameterNamesTest, AdamStateDictUsesStableParameterNames) {
auto first = std::make_shared<Tensor>(std::vector<int64_t>{2, 2}, DataType::kFLOAT32, GetDevice());
auto second = std::make_shared<Tensor>(std::vector<int64_t>{3}, DataType::kFLOAT32, GetDevice());
auto adam = std::make_shared<optimizers::Adam>(std::vector<std::shared_ptr<Tensor>>{first, second}, 0.001);
adam->set_parameter_names({"transformer.h.0.weight", "transformer.h.0.bias"});

const auto state = adam->StateDict();
EXPECT_TRUE(state.contains("adam.m.transformer.h.0.weight"));
EXPECT_TRUE(state.contains("adam.v.transformer.h.0.weight"));
EXPECT_TRUE(state.contains("adam.m.transformer.h.0.bias"));
EXPECT_TRUE(state.contains("adam.v.transformer.h.0.bias"));
EXPECT_TRUE(state.contains("adam.t"));

auto restored = std::make_shared<optimizers::Adam>(std::vector<std::shared_ptr<Tensor>>{first, second}, 0.001);
restored->set_parameter_names({"transformer.h.0.weight", "transformer.h.0.bias"});
restored->LoadStateDict(state);
EXPECT_EQ(restored->StateDict().size(), state.size());
}

TEST_P(OptimizerParameterNamesTest, PreservesNumericKeysWhenNamesAreNotSet) {
auto parameter = std::make_shared<Tensor>(std::vector<int64_t>{2, 2}, DataType::kFLOAT32, GetDevice());
auto adam = std::make_shared<optimizers::Adam>(std::vector<std::shared_ptr<Tensor>>{parameter}, 0.001);

const auto state = adam->StateDict();
EXPECT_TRUE(state.contains("adam.m.0"));
EXPECT_TRUE(state.contains("adam.v.0"));
}

TEST_P(OptimizerParameterNamesTest, RejectsWrongNumberOfParameterNames) {
auto parameter = std::make_shared<Tensor>(std::vector<int64_t>{2, 2}, DataType::kFLOAT32, GetDevice());
auto adam = std::make_shared<optimizers::Adam>(std::vector<std::shared_ptr<Tensor>>{parameter}, 0.001);

EXPECT_DEATH(adam->set_parameter_names({"first", "second"}), "");
}

INFINI_TRAIN_REGISTER_TEST(OptimizerParameterNamesTest);
Loading