diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index fd2522e3..976a7031 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -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 parameter_name_by_tensor; + for (const auto &[name, parameter] : model->NamedParameters()) { + parameter_name_by_tensor.emplace(parameter.get(), name); + } + std::vector parameter_names; + parameter_names.reserve(params_to_optimize.size()); + for (const auto ¶meter : 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; diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 05fad4a8..2a209614 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -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 parameter_name_by_tensor; + for (const auto &[name, parameter] : model->NamedParameters()) { + parameter_name_by_tensor.emplace(parameter.get(), name); + } + std::vector parameter_names; + parameter_names.reserve(params_to_optimize.size()); + for (const auto ¶meter : 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; diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 2b8cd6dc..ca32b4d9 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -49,6 +49,8 @@ class Module : public std::enable_shared_from_this { // TODO: Change return type to filterable iterator (like PyTorch's named_parameters with prefix matching) virtual std::vector> Parameters() const; + std::vector>> + NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const; bool has_parameter(const std::string &name) const; std::shared_ptr *mutable_parameter(const std::string &name); const std::shared_ptr ¶meter(const std::string &name) const; diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index 0a33f0b8..eb9e67bc 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -37,8 +37,13 @@ class Optimizer { void set_initial_learning_rate(float lr); + void set_parameter_names(const std::vector &names); + + const std::vector ¶meter_names() const; + protected: std::vector> params_; + std::vector parameter_names_; float learning_rate_ = 0.0f; float initial_learning_rate_ = 0.0f; bool initial_lr_set_ = false; diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 81068fe8..fb9e53e1 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -48,6 +48,41 @@ std::vector> Module::Parameters() const { return params; } +std::vector>> +Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { + std::vector>> named_parameters; + std::unordered_set visited; + + std::function 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 *Module::mutable_parameter(const std::string &name) { diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 1e97bfe0..5305b82e 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -1,6 +1,5 @@ #include "infini_train/include/optimizer.h" -#include #include #include "infini_train/include/core/runtime/device_guard.h" @@ -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 &names) { + CHECK_EQ(names.size(), params_.size()); + parameter_names_ = names; +} + +const std::vector &Optimizer::parameter_names() const { return parameter_names_; } + namespace optimizers { SGD::SGD(const std::vector> ¶ms, float learning_rate) : Optimizer(params, learning_rate) {} @@ -96,8 +103,9 @@ OptimizerCreator Adam::Create(float learning_rate, float beta1, float beta2, flo std::unordered_map> Adam::StateDict() const { std::unordered_map> 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(std::vector{}, DataType::kINT64, Device()); @@ -108,8 +116,9 @@ std::unordered_map> Adam::StateDict() const void Adam::LoadStateDict(const std::unordered_map> &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)); diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 96776585..c7dd49b9 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -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) diff --git a/tests/module/CMakeLists.txt b/tests/module/CMakeLists.txt new file mode 100644 index 00000000..84b78098 --- /dev/null +++ b/tests/module/CMakeLists.txt @@ -0,0 +1,5 @@ +file(GLOB MODULE_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/test_*.cc) + +infini_train_add_test_suite(test_module + SOURCES ${MODULE_SOURCES} +) diff --git a/tests/module/test_named_parameters.cc b/tests/module/test_named_parameters.cc new file mode 100644 index 00000000..5a23d62b --- /dev/null +++ b/tests/module/test_named_parameters.cc @@ -0,0 +1,76 @@ +#include +#include +#include +#include +#include + +#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(2, 3, /*bias=*/true, GetDevice()); + + const auto parameters = linear->NamedParameters("model", false); + const std::unordered_map> 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(2, 3, /*bias=*/true, GetDevice()); + auto root + = std::make_shared(std::vector>{shared, shared}); + + const auto deduplicated = root->NamedParameters("model"); + ASSERT_EQ(deduplicated.size(), 2); + std::unordered_set 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> 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(2, 3, /*bias=*/false, GetDevice()); + auto second = std::make_shared(3, 4, /*bias=*/false, GetDevice()); + auto nested + = std::make_shared(std::vector>{first, second}); + auto root = std::make_shared( + std::vector>{std::make_shared(2, 2, false, GetDevice()), nested}); + + const auto parameters = root->NamedParameters(); + const std::unordered_map> 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(std::vector>{nullptr}); + + EXPECT_TRUE(root->NamedParameters().empty()); +} + +INFINI_TRAIN_REGISTER_TEST(ModuleNamedParametersTest); diff --git a/tests/optimizer/test_optimizer_parameter_names.cc b/tests/optimizer/test_optimizer_parameter_names.cc new file mode 100644 index 00000000..9354f659 --- /dev/null +++ b/tests/optimizer/test_optimizer_parameter_names.cc @@ -0,0 +1,50 @@ +#include +#include + +#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(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); + auto second = std::make_shared(std::vector{3}, DataType::kFLOAT32, GetDevice()); + auto adam = std::make_shared(std::vector>{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(std::vector>{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(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); + auto adam = std::make_shared(std::vector>{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(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); + auto adam = std::make_shared(std::vector>{parameter}, 0.001); + + EXPECT_DEATH(adam->set_parameter_names({"first", "second"}), ""); +} + +INFINI_TRAIN_REGISTER_TEST(OptimizerParameterNamesTest);