From 025efc78849d26dc81c5fc623837139a79bee8e8 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Thu, 30 Jul 2026 23:12:14 +0800 Subject: [PATCH 1/4] feat: add named parameters API --- infini_train/include/nn/modules/module.h | 2 + infini_train/src/nn/modules/module.cc | 49 +++++++++++++ tests/CMakeLists.txt | 2 + tests/module/CMakeLists.txt | 5 ++ tests/module/test_named_parameters.cc | 93 ++++++++++++++++++++++++ 5 files changed, 151 insertions(+) create mode 100644 tests/module/CMakeLists.txt create mode 100644 tests/module/test_named_parameters.cc diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 2b8cd6dc..27629f1c 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; + virtual 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/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 81068fe8..47038904 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -48,6 +48,55 @@ 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) { + std::vector>> parameters; + parameters.reserve(module.parameters_.size()); + for (const auto &[name, parameter] : module.parameters_) { + if (parameter) { + parameters.emplace_back(name, parameter); + } + } + std::sort(parameters.begin(), parameters.end(), + [](const auto &left, const auto &right) { return left.first < right.first; }); + + for (auto &[name, parameter] : 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, std::move(parameter)); + } + + if (!recurse) { + return; + } + + std::vector>> children; + children.reserve(module.modules_.size()); + for (const auto &[name, child] : module.modules_) { + if (child) { + children.emplace_back(name, child); + } + } + std::sort(children.begin(), children.end(), + [](const auto &left, const auto &right) { return left.first < right.first; }); + + for (const auto &[name, child] : children) { + 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/tests/CMakeLists.txt b/tests/CMakeLists.txt index f7ced03d..3bfaa548 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -9,6 +9,8 @@ add_subdirectory(common) # Distributed tests add_subdirectory(distributed) +# 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..85dcf6bd --- /dev/null +++ b/tests/module/test_named_parameters.cc @@ -0,0 +1,93 @@ +#include +#include + +#include "gtest/gtest.h" + +#include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/tensor.h" + +#include "tests/common/test_utils.h" + +using namespace infini_train; + +namespace { + +class NamedParameterModule final : public nn::Module { +public: + void AddParameter(const std::string &name, const std::shared_ptr ¶meter) { + parameters_[name] = parameter; + } + + void AddModule(const std::string &name, const std::shared_ptr &module) { modules_[name] = module; } +}; + +std::shared_ptr MakeParameter(Device device) { + return std::make_shared(std::vector{1}, DataType::kFLOAT32, device); +} + +} // namespace + +class ModuleNamedParametersTest : public test::InfiniTrainTest {}; + +TEST_P(ModuleNamedParametersTest, SupportsPrefixRecursionAndSharedParameterDeduplication) { + auto root = std::make_shared(); + auto child = std::make_shared(); + auto grandchild = std::make_shared(); + auto shared = MakeParameter(GetDevice()); + auto child_weight = MakeParameter(GetDevice()); + auto grandchild_weight = MakeParameter(GetDevice()); + + root->AddParameter("root_weight", shared); + child->AddParameter("alias", shared); + child->AddParameter("weight", child_weight); + grandchild->AddParameter("weight", grandchild_weight); + child->AddModule("grandchild", grandchild); + root->AddModule("child", child); + + const auto local = root->NamedParameters("model", false); + ASSERT_EQ(local.size(), 1); + EXPECT_EQ(local[0].first, "model.root_weight"); + EXPECT_EQ(local[0].second, shared); + + const auto deduplicated = root->NamedParameters("model"); + ASSERT_EQ(deduplicated.size(), 3); + EXPECT_EQ(deduplicated[0].first, "model.root_weight"); + EXPECT_EQ(deduplicated[1].first, "model.child.weight"); + EXPECT_EQ(deduplicated[2].first, "model.child.grandchild.weight"); + + const auto aliases = root->NamedParameters("model", true, false); + ASSERT_EQ(aliases.size(), 4); + EXPECT_EQ(aliases[0].first, "model.root_weight"); + EXPECT_EQ(aliases[1].first, "model.child.alias"); + EXPECT_EQ(aliases[1].second, shared); +} + +TEST_P(ModuleNamedParametersTest, ProducesDeterministicLexicalOrder) { + auto root = std::make_shared(); + auto first_child = std::make_shared(); + auto second_child = std::make_shared(); + + root->AddParameter("z", MakeParameter(GetDevice())); + root->AddParameter("a", MakeParameter(GetDevice())); + first_child->AddParameter("weight", MakeParameter(GetDevice())); + second_child->AddParameter("weight", MakeParameter(GetDevice())); + root->AddModule("z_child", second_child); + root->AddModule("a_child", first_child); + + const auto parameters = root->NamedParameters(); + ASSERT_EQ(parameters.size(), 4); + EXPECT_EQ(parameters[0].first, "a"); + EXPECT_EQ(parameters[1].first, "z"); + EXPECT_EQ(parameters[2].first, "a_child.weight"); + EXPECT_EQ(parameters[3].first, "z_child.weight"); +} + +TEST_P(ModuleNamedParametersTest, SkipsNullEntries) { + auto root = std::make_shared(); + root->AddParameter("missing", nullptr); + root->AddModule("missing_child", nullptr); + + EXPECT_TRUE(root->NamedParameters().empty()); +} + +INFINI_TRAIN_REGISTER_TEST(ModuleNamedParametersTest); From 53775fd026ebcb61b38b89dce062e3ebdfa20174 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Tue, 11 Aug 2026 08:10:21 +0000 Subject: [PATCH 2/4] feat: use stable names for optimizer state --- example/gpt2/main.cc | 7 +- example/llama3/main.cc | 7 +- example/mixtral/main.cc | 4 +- infini_train/include/nn/modules/module.h | 2 +- .../nn/parallel/ddp/distributed_optimizer.h | 3 + infini_train/include/optimizer.h | 17 ++- infini_train/src/nn/modules/module.cc | 70 +++++------ .../nn/parallel/ddp/distributed_optimizer.cc | 16 ++- infini_train/src/optimizer.cc | 59 +++++++--- tests/module/test_named_parameters.cc | 111 ++++++++---------- tests/optimizer/CMakeLists.txt | 14 +++ .../test_optimizer_parameter_names.cc | 105 +++++++++++++++++ 12 files changed, 283 insertions(+), 132 deletions(-) create mode 100644 tests/optimizer/test_optimizer_parameter_names.cc diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 068d1e7c..36bae7ec 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -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 = nullptr; + const auto named_parameters = model->NamedParameters(); if (FLAGS_zero_stage >= 1) { auto model_chunks = (pp_world_size > 1) ? *(dynamic_cast(model.get())->mutable_chunks()) : std::vector>{model}; - optimizer = std::make_shared(optimizer_creator, params_to_optimize, - model_chunks, ddp_world_size, ddp_rank); + optimizer = std::make_shared( + 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; diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 9ce5a7e7..38622716 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -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(model.get())->mutable_chunks()) : std::vector>{model}; - optimizer = std::make_shared(optimizer_creator, params_to_optimize, - model_chunks, ddp_world_size, ddp_rank); + optimizer = std::make_shared( + 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; diff --git a/example/mixtral/main.cc b/example/mixtral/main.cc index ceabfad9..fd1556bd 100644 --- a/example/mixtral/main.cc +++ b/example/mixtral/main.cc @@ -104,8 +104,8 @@ int main(int argc, char *argv[]) { } auto loss_fn = std::make_shared(); - auto optimizer - = infini_train::optimizers::Adam::Create(static_cast(FLAGS_learning_rate))(model->Parameters()); + auto optimizer = infini_train::optimizers::Adam::Create(static_cast(FLAGS_learning_rate))( + model->Parameters(), model->NamedParameters()); auto device_impl = infini_train::core::GetDeviceGuardImpl(train_device.type()); std::vector step_duration_ms; diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 27629f1c..ca32b4d9 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -49,7 +49,7 @@ 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; - virtual std::vector>> + 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); diff --git a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h index d694ab2a..850bbb1c 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h +++ b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h @@ -21,6 +21,7 @@ class DistributedOptimizer final : public infini_train::Optimizer { public: DistributedOptimizer(OptimizerCreator base_optimizer_creator, const std::vector> &full_params, + const NamedParameterList &named_parameters, const std::vector> &model_chunks, size_t ddp_world_size, size_t ddp_rank); @@ -55,6 +56,8 @@ class DistributedOptimizer final : public infini_train::Optimizer { // shard params std::vector> shard_params_; + NamedParameterList shard_named_parameters_; + std::unordered_map parameter_name_by_tensor_; // Base optimizer (SGD, Adam and etc.) std::shared_ptr base_optimizer_; diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index 0a33f0b8..aa2b6ce4 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -5,6 +5,7 @@ #include #include #include +#include #include namespace infini_train { @@ -13,11 +14,15 @@ class Tensor; namespace infini_train { class Optimizer; -using OptimizerCreator = std::function(const std::vector> ¶ms)>; +using NamedParameter = std::pair>; +using NamedParameterList = std::vector; +using OptimizerCreator = std::function(const std::vector> ¶ms, + const NamedParameterList &named_parameters)>; class Optimizer { public: - explicit Optimizer(const std::vector> ¶ms, float learning_rate = 0.0f); + explicit Optimizer(const std::vector> ¶ms, float learning_rate = 0.0f, + const NamedParameterList &named_parameters = {}); virtual void ZeroGrad(bool set_to_none = true); @@ -37,8 +42,11 @@ class Optimizer { void set_initial_learning_rate(float lr); + void set_parameter_names(const std::vector &names); + protected: std::vector> params_; + std::vector parameter_names_; float learning_rate_ = 0.0f; float initial_learning_rate_ = 0.0f; bool initial_lr_set_ = false; @@ -47,7 +55,8 @@ class Optimizer { namespace optimizers { class SGD : public Optimizer { public: - SGD(const std::vector> ¶ms, float learning_rate); + SGD(const std::vector> ¶ms, float learning_rate, + const NamedParameterList &named_parameters = {}); void Step() override; @@ -57,7 +66,7 @@ class SGD : public Optimizer { class Adam : public Optimizer { public: Adam(const std::vector> ¶ms, 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; diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 47038904..e5b1187f 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -53,47 +53,35 @@ Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_dup std::vector>> named_parameters; std::unordered_set visited; - std::function collect - = [&](const Module &module, const std::string &module_prefix) { - std::vector>> parameters; - parameters.reserve(module.parameters_.size()); - for (const auto &[name, parameter] : module.parameters_) { - if (parameter) { - parameters.emplace_back(name, parameter); - } - } - std::sort(parameters.begin(), parameters.end(), - [](const auto &left, const auto &right) { return left.first < right.first; }); - - for (auto &[name, parameter] : 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, std::move(parameter)); - } - - if (!recurse) { - return; - } - - std::vector>> children; - children.reserve(module.modules_.size()); - for (const auto &[name, child] : module.modules_) { - if (child) { - children.emplace_back(name, child); - } - } - std::sort(children.begin(), children.end(), - [](const auto &left, const auto &right) { return left.first < right.first; }); - - for (const auto &[name, child] : children) { - const auto child_prefix = module_prefix.empty() ? name : module_prefix + "." + name; - collect(*child, child_prefix); - } - }; - - collect(*this, prefix); + std::vector>> 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(this)->NamedModules(/*memory=*/nullptr, prefix, /*remove_duplicate=*/false); + } else { + named_modules.emplace_back(prefix, std::const_pointer_cast(shared_from_this())); + } + + for (const auto &[module_prefix, module] : named_modules) { + std::vector>> 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; } diff --git a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc index 022a4758..ab44ac78 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc @@ -8,11 +8,18 @@ namespace infini_train::nn::parallel { DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator, const std::vector> &full_params, + const NamedParameterList &named_parameters, const std::vector> &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(model_chunks[i]); @@ -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; @@ -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); } } } diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 1e97bfe0..42708d71 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -1,6 +1,6 @@ #include "infini_train/include/optimizer.h" -#include +#include #include #include "infini_train/include/core/runtime/device_guard.h" @@ -9,8 +9,27 @@ #include "infini_train/include/tensor.h" namespace infini_train { -Optimizer::Optimizer(const std::vector> ¶ms, float learning_rate) - : params_(params), learning_rate_(learning_rate) {} +Optimizer::Optimizer(const std::vector> ¶ms, float learning_rate, + const NamedParameterList &named_parameters) + : params_(params), learning_rate_(learning_rate) { + if (named_parameters.empty()) { + return; + } + + std::unordered_map 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 ¶meter : 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); } @@ -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 &names) { + CHECK_EQ(names.size(), params_.size()); + parameter_names_ = names; +} + namespace optimizers { -SGD::SGD(const std::vector> ¶ms, float learning_rate) : Optimizer(params, learning_rate) {} +SGD::SGD(const std::vector> ¶ms, float learning_rate, + const NamedParameterList &named_parameters) + : Optimizer(params, learning_rate, named_parameters) {} void SGD::Step() { for (auto param : params_) { @@ -51,13 +78,15 @@ void SGD::Step() { } OptimizerCreator SGD::Create(float learning_rate) { - return [learning_rate](const std::vector> ¶ms) { - return std::make_shared(params, learning_rate); + return [learning_rate](const std::vector> ¶ms, + const NamedParameterList &named_parameters) { + return std::make_shared(params, learning_rate, named_parameters); }; } -Adam::Adam(const std::vector> ¶ms, 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> ¶ms, 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 ¶m : params_) { m_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); @@ -88,16 +117,17 @@ void Adam::Step() { } OptimizerCreator Adam::Create(float learning_rate, float beta1, float beta2, float eps) { - return [=](const std::vector> ¶ms) { - return std::make_shared(params, learning_rate, beta1, beta2, eps); + return [=](const std::vector> ¶ms, const NamedParameterList &named_parameters) { + return std::make_shared(params, learning_rate, beta1, beta2, eps, named_parameters); }; } 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 +138,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/module/test_named_parameters.cc b/tests/module/test_named_parameters.cc index 85dcf6bd..1a7a07e7 100644 --- a/tests/module/test_named_parameters.cc +++ b/tests/module/test_named_parameters.cc @@ -1,91 +1,78 @@ #include #include +#include +#include +#include #include "gtest/gtest.h" -#include "infini_train/include/nn/modules/module.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; -namespace { +class ModuleNamedParametersTest : public test::InfiniTrainTest {}; -class NamedParameterModule final : public nn::Module { -public: - void AddParameter(const std::string &name, const std::shared_ptr ¶meter) { - parameters_[name] = parameter; - } +TEST_P(ModuleNamedParametersTest, SupportsPrefixAndNonRecursiveLookup) { + auto linear = std::make_shared(2, 3, /*bias=*/true, GetDevice()); - void AddModule(const std::string &name, const std::shared_ptr &module) { modules_[name] = module; } -}; + const auto parameters = linear->NamedParameters("model", false); + const std::unordered_map> by_name(parameters.begin(), parameters.end()); -std::shared_ptr MakeParameter(Device device) { - return std::make_shared(std::vector{1}, DataType::kFLOAT32, device); + 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)); } -} // namespace - -class ModuleNamedParametersTest : public test::InfiniTrainTest {}; - -TEST_P(ModuleNamedParametersTest, SupportsPrefixRecursionAndSharedParameterDeduplication) { - auto root = std::make_shared(); - auto child = std::make_shared(); - auto grandchild = std::make_shared(); - auto shared = MakeParameter(GetDevice()); - auto child_weight = MakeParameter(GetDevice()); - auto grandchild_weight = MakeParameter(GetDevice()); - - root->AddParameter("root_weight", shared); - child->AddParameter("alias", shared); - child->AddParameter("weight", child_weight); - grandchild->AddParameter("weight", grandchild_weight); - child->AddModule("grandchild", grandchild); - root->AddModule("child", child); - - const auto local = root->NamedParameters("model", false); - ASSERT_EQ(local.size(), 1); - EXPECT_EQ(local[0].first, "model.root_weight"); - EXPECT_EQ(local[0].second, shared); +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(), 3); - EXPECT_EQ(deduplicated[0].first, "model.root_weight"); - EXPECT_EQ(deduplicated[1].first, "model.child.weight"); - EXPECT_EQ(deduplicated[2].first, "model.child.grandchild.weight"); + ASSERT_EQ(deduplicated.size(), 2); + EXPECT_EQ(deduplicated[0].first, "model.0.bias"); + EXPECT_EQ(deduplicated[1].first, "model.0.weight"); + 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); - ASSERT_EQ(aliases.size(), 4); - EXPECT_EQ(aliases[0].first, "model.root_weight"); - EXPECT_EQ(aliases[1].first, "model.child.alias"); - EXPECT_EQ(aliases[1].second, shared); + 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, ProducesDeterministicLexicalOrder) { - auto root = std::make_shared(); - auto first_child = std::make_shared(); - auto second_child = std::make_shared(); - - root->AddParameter("z", MakeParameter(GetDevice())); - root->AddParameter("a", MakeParameter(GetDevice())); - first_child->AddParameter("weight", MakeParameter(GetDevice())); - second_child->AddParameter("weight", MakeParameter(GetDevice())); - root->AddModule("z_child", second_child); - root->AddModule("a_child", first_child); +TEST_P(ModuleNamedParametersTest, ReturnsNestedParametersInStableNameOrder) { + 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(); - ASSERT_EQ(parameters.size(), 4); - EXPECT_EQ(parameters[0].first, "a"); - EXPECT_EQ(parameters[1].first, "z"); - EXPECT_EQ(parameters[2].first, "a_child.weight"); - EXPECT_EQ(parameters[3].first, "z_child.weight"); + const std::unordered_map> by_name(parameters.begin(), parameters.end()); + + ASSERT_EQ(by_name.size(), 3); + ASSERT_EQ(parameters.size(), 3); + EXPECT_EQ(parameters[0].first, "0.weight"); + EXPECT_EQ(parameters[1].first, "1.0.weight"); + EXPECT_EQ(parameters[2].first, "1.1.weight"); + 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, SkipsNullEntries) { - auto root = std::make_shared(); - root->AddParameter("missing", nullptr); - root->AddModule("missing_child", nullptr); +TEST_P(ModuleNamedParametersTest, SkipsNullSubmodules) { + auto root = std::make_shared(std::vector>{nullptr}); EXPECT_TRUE(root->NamedParameters().empty()); } diff --git a/tests/optimizer/CMakeLists.txt b/tests/optimizer/CMakeLists.txt index c0bfbd50..bce88694 100644 --- a/tests/optimizer/CMakeLists.txt +++ b/tests/optimizer/CMakeLists.txt @@ -7,3 +7,17 @@ file(GLOB OPTIMIZER_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/test_*.cc) infini_train_add_test_suite(test_optimizer SOURCES ${OPTIMIZER_SOURCES} ) + +add_test( + NAME OptimizerParameterNamesTest.DistributedOptimizerPropagatesNamesToShardOptimizer + COMMAND ${CMAKE_COMMAND} -E env + "PROC_WORLD_SIZE=2" + $ + --gtest_filter=CUDA/OptimizerParameterNamesTest.DistributedOptimizerPropagatesNamesToShardOptimizer/* +) +set_tests_properties( + OptimizerParameterNamesTest.DistributedOptimizerPropagatesNamesToShardOptimizer + PROPERTIES + LABELS "cuda;distributed" + TIMEOUT 30 +) diff --git a/tests/optimizer/test_optimizer_parameter_names.cc b/tests/optimizer/test_optimizer_parameter_names.cc new file mode 100644 index 00000000..7329b7d7 --- /dev/null +++ b/tests/optimizer/test_optimizer_parameter_names.cc @@ -0,0 +1,105 @@ +#include +#include + +#include "gtest/gtest.h" + +#include "infini_train/include/nn/modules/linear.h" +#include "infini_train/include/nn/parallel/ddp/distributed_data_parallel.h" +#include "infini_train/include/nn/parallel/ddp/distributed_data_parallel_config.h" +#include "infini_train/include/nn/parallel/ddp/distributed_optimizer.h" +#include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/process_group.h" +#include "infini_train/include/nn/parallel/rank.h" +#include "infini_train/include/nn/parallel/utils.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, ConstructorMatchesNamesToOptimizerParameterOrder) { + auto first = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); + auto second = std::make_shared(std::vector{3}, DataType::kFLOAT32, GetDevice()); + const NamedParameterList named_parameters{{"first", first}, {"second", second}}; + + auto adam = optimizers::Adam::Create(0.001)({second, first}, named_parameters); + const auto state = adam->StateDict(); + + EXPECT_TRUE(state.contains("adam.m.second")); + EXPECT_TRUE(state.contains("adam.v.second")); + EXPECT_TRUE(state.contains("adam.m.first")); + EXPECT_TRUE(state.contains("adam.v.first")); +} + +TEST_P(OptimizerParameterNamesTest, DistributedOptimizerPropagatesNamesToShardOptimizer) { + ONLY_CUDA(); + REQUIRE_MIN_DEVICES(2); + if (nn::parallel::global::GetDataParallelSize() != 2) { + GTEST_SKIP() << "requires PROC_WORLD_SIZE=2"; + } + + const nn::parallel::Rank rank(/*process_rank=*/0, /*thread_rank=*/0, /*process_size=*/1, /*thread_size=*/2); + auto *pg_factory = nn::parallel::ProcessGroupFactory::Instance(Device::DeviceType::kCUDA); + pg_factory->GetOrCreate(nn::parallel::GetDataParallelProcessGroupName(rank.GlobalRank()), + nn::parallel::GetDataParallelGroupRanks(rank.GlobalRank())); + + auto model = std::make_shared(4, 4, /*bias=*/false, GetDevice()); + const auto params = model->Parameters(); + const auto named_parameters = model->NamedParameters(); + + nn::parallel::DistributedDataParallelConfig ddp_config; + ddp_config.zero_stage = 1; + ddp_config.overlap_grad_reduce = false; + ddp_config.overlap_param_gather = false; + auto ddp_model = std::make_shared(model, rank, ddp_config); + + nn::parallel::DistributedOptimizer optimizer(optimizers::Adam::Create(0.001), params, named_parameters, + std::vector>{ddp_model}, + /*ddp_world_size=*/2, /*ddp_rank=*/0); + const auto state = optimizer.StateDict(); + + EXPECT_TRUE(state.contains("adam.m.weight")); + EXPECT_TRUE(state.contains("adam.v.weight")); + EXPECT_FALSE(state.contains("adam.m.0")); + EXPECT_FALSE(state.contains("adam.v.0")); +} + +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); From f95b66d3ab5e02f98e6ea81e4ef3d1b56680512b Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Wed, 12 Aug 2026 22:51:02 +0800 Subject: [PATCH 3/4] refactor: address named parameter review comments --- docs/lora_usage_guide.md | 2 +- .../nn/parallel/ddp/distributed_optimizer.h | 9 +-- infini_train/include/optimizer.h | 2 - infini_train/src/nn/modules/module.cc | 58 ++++++++++--------- .../nn/parallel/ddp/distributed_optimizer.cc | 38 ++++++------ infini_train/src/optimizer.cc | 5 -- .../test_optimizer_parameter_names.cc | 15 ++--- 7 files changed, 58 insertions(+), 71 deletions(-) diff --git a/docs/lora_usage_guide.md b/docs/lora_usage_guide.md index f7512b4c..2a3470db 100644 --- a/docs/lora_usage_guide.md +++ b/docs/lora_usage_guide.md @@ -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(); diff --git a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h index 850bbb1c..9ef503de 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h +++ b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h @@ -43,7 +43,9 @@ class DistributedOptimizer final : public infini_train::Optimizer { virtual float learning_rate() const override; private: - void BuildShardParamsAndBindGrads(); + void BuildShardParamsAndBindGrads(const NamedParameterList &named_parameters, + std::vector> &shard_params, + NamedParameterList &shard_named_parameters); private: // Inherit from DDP model @@ -54,11 +56,6 @@ class DistributedOptimizer final : public infini_train::Optimizer { size_t ddp_world_size_; size_t ddp_rank_; - // shard params - std::vector> shard_params_; - NamedParameterList shard_named_parameters_; - std::unordered_map parameter_name_by_tensor_; - // Base optimizer (SGD, Adam and etc.) std::shared_ptr base_optimizer_; }; diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index aa2b6ce4..8fbe4947 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -42,8 +42,6 @@ class Optimizer { void set_initial_learning_rate(float lr); - void set_parameter_names(const std::vector &names); - protected: std::vector> params_; std::vector parameter_names_; diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index e5b1187f..0d2492b2 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -51,36 +51,38 @@ std::vector> Module::Parameters() const { std::vector>> Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { std::vector>> named_parameters; - std::unordered_set visited; - std::vector>> 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(this)->NamedModules(/*memory=*/nullptr, prefix, /*remove_duplicate=*/false); - } else { - named_modules.emplace_back(prefix, std::const_pointer_cast(shared_from_this())); - } + std::function collect + = [&](const Module &module, const std::string &module_prefix) { + for (const auto &[name, parameter] : module.parameters_) { + if (!parameter) { + 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); + std::sort(named_parameters.begin(), named_parameters.end(), + [](const auto &lhs, const auto &rhs) { return lhs.first < rhs.first; }); - for (const auto &[module_prefix, module] : named_modules) { - std::vector>> 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); - } + if (remove_duplicate) { + std::unordered_set visited; + std::erase_if(named_parameters, [&](const auto &named_parameter) { + return !visited.insert(named_parameter.second.get()).second; + }); } return named_parameters; } diff --git a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc index ab44ac78..893ad7a1 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc @@ -11,15 +11,9 @@ DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator, const NamedParameterList &named_parameters, const std::vector> &model_chunks, size_t ddp_world_size, size_t ddp_rank) - : Optimizer(full_params, /*learning_rate=*/0.0f, named_parameters), 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) { 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(model_chunks[i]); @@ -31,16 +25,24 @@ DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator, ddp_chunk->bucket_groups().end()); } - BuildShardParamsAndBindGrads(); + std::vector> shard_params; + NamedParameterList shard_named_parameters; + BuildShardParamsAndBindGrads(named_parameters, shard_params, shard_named_parameters); // Build base optimizer - base_optimizer_ = creator(shard_params_, shard_named_parameters_); + 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(); +void DistributedOptimizer::BuildShardParamsAndBindGrads(const NamedParameterList &named_parameters, + std::vector> &shard_params, + NamedParameterList &shard_named_parameters) { + std::unordered_map 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); + } for (const auto &group : bucket_groups_) { const bool use_grad_shard = group->config().zero_stage >= 2; @@ -90,17 +92,17 @@ 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); - const auto name_it = parameter_name_by_tensor_.find(param.get()); - CHECK(name_it != parameter_name_by_tensor_.end()) + 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); + shard_named_parameters.emplace_back(name_it->second, param_piece); } } } - CHECK(!shard_params_.empty()) << "DistributedOptimizer: this DP rank owns no param pieces. " - << "Check bucket padding/divisibility and param bucketing order."; + CHECK(!shard_params.empty()) << "DistributedOptimizer: this DP rank owns no param pieces. " + << "Check bucket padding/divisibility and param bucketing order."; } void DistributedOptimizer::StartGradSync() { diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 42708d71..01f61df4 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -53,11 +53,6 @@ void Optimizer::set_initial_learning_rate(float lr) { initial_lr_set_ = true; } -void Optimizer::set_parameter_names(const std::vector &names) { - CHECK_EQ(names.size(), params_.size()); - parameter_names_ = names; -} - namespace optimizers { SGD::SGD(const std::vector> ¶ms, float learning_rate, diff --git a/tests/optimizer/test_optimizer_parameter_names.cc b/tests/optimizer/test_optimizer_parameter_names.cc index 7329b7d7..afb0f187 100644 --- a/tests/optimizer/test_optimizer_parameter_names.cc +++ b/tests/optimizer/test_optimizer_parameter_names.cc @@ -23,8 +23,9 @@ 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 NamedParameterList named_parameters{{"transformer.h.0.weight", first}, + {"transformer.h.0.bias", second}}; + auto adam = optimizers::Adam::Create(0.001)({first, second}, named_parameters); const auto state = adam->StateDict(); EXPECT_TRUE(state.contains("adam.m.transformer.h.0.weight")); @@ -33,8 +34,7 @@ TEST_P(OptimizerParameterNamesTest, AdamStateDictUsesStableParameterNames) { 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"}); + auto restored = optimizers::Adam::Create(0.001)({first, second}, named_parameters); restored->LoadStateDict(state); EXPECT_EQ(restored->StateDict().size(), state.size()); } @@ -95,11 +95,4 @@ TEST_P(OptimizerParameterNamesTest, PreservesNumericKeysWhenNamesAreNotSet) { 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); From 77b4a88e645a7876bcb157ea5a03003002bf5318 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Fri, 14 Aug 2026 09:14:30 +0000 Subject: [PATCH 4/4] fix: address named parameter review feedback --- example/gpt2/main.cc | 20 +++-- example/llama3/main.cc | 20 +++-- example/mixtral/main.cc | 4 +- infini_train/include/nn/modules/module.h | 4 +- .../nn/parallel/ddp/distributed_optimizer.h | 13 ++- infini_train/include/optimizer.h | 20 +++-- infini_train/src/nn/modules/module.cc | 80 +++++++++---------- .../nn/parallel/ddp/distributed_optimizer.cc | 70 ++++++++++------ infini_train/src/optimizer.cc | 68 +++++++++------- tests/module/test_named_parameters.cc | 18 ++--- .../test_optimizer_parameter_names.cc | 14 ++-- 11 files changed, 191 insertions(+), 140 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 36bae7ec..5a5cfc65 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -327,18 +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 = nullptr; - const auto named_parameters = model->NamedParameters(); + std::unordered_set params_to_optimize_set; + params_to_optimize_set.reserve(params_to_optimize.size()); + for (const auto ¶m : 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(model.get())->mutable_chunks()) : std::vector>{model}; - optimizer = std::make_shared( - optimizer_creator, params_to_optimize, named_parameters, model_chunks, ddp_world_size, ddp_rank); + optimizer = std::make_shared(optimizer_creator, named_parameters, + model_chunks, ddp_world_size, ddp_rank); } else { - optimizer = optimizer_creator(params_to_optimize, named_parameters); + optimizer = optimizer_creator(named_parameters); } 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 38622716..c4642cc2 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -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 = nullptr; std::vector> params_to_optimize; @@ -311,16 +311,26 @@ 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(); + std::unordered_set params_to_optimize_set; + params_to_optimize_set.reserve(params_to_optimize.size()); + for (const auto ¶m : 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(model.get())->mutable_chunks()) : std::vector>{model}; - optimizer = std::make_shared( - optimizer_creator, params_to_optimize, named_parameters, model_chunks, ddp_world_size, ddp_rank); + optimizer = std::make_shared(optimizer_creator, named_parameters, + model_chunks, ddp_world_size, ddp_rank); } else { - optimizer = optimizer_creator(params_to_optimize, named_parameters); + optimizer = optimizer_creator(named_parameters); } const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration; diff --git a/example/mixtral/main.cc b/example/mixtral/main.cc index fd1556bd..553d7a62 100644 --- a/example/mixtral/main.cc +++ b/example/mixtral/main.cc @@ -104,8 +104,8 @@ int main(int argc, char *argv[]) { } auto loss_fn = std::make_shared(); - auto optimizer = infini_train::optimizers::Adam::Create(static_cast(FLAGS_learning_rate))( - model->Parameters(), model->NamedParameters()); + auto optimizer = infini_train::optimizers::Adam::CreateNamed(static_cast(FLAGS_learning_rate))( + model->NamedParameters()); auto device_impl = infini_train::core::GetDeviceGuardImpl(train_device.type()); std::vector step_duration_ms; diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index ca32b4d9..1d42b2ac 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -47,8 +47,10 @@ class Module : public std::enable_shared_from_this { const std::string &type() const; - // TODO: Change return type to filterable iterator (like PyTorch's named_parameters with prefix matching) virtual std::vector> Parameters() const; + + // InfiniTrain's NamedParameters returns results ordered by full parameter name. + // TODO: Align with PyTorch's ordering in the future. std::vector>> NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const; bool has_parameter(const std::string &name) const; diff --git a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h index 9ef503de..fa9592ae 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h +++ b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h @@ -1,6 +1,7 @@ #pragma once #include +#include #include #include #include @@ -21,7 +22,10 @@ class DistributedOptimizer final : public infini_train::Optimizer { public: DistributedOptimizer(OptimizerCreator base_optimizer_creator, const std::vector> &full_params, - const NamedParameterList &named_parameters, + const std::vector> &model_chunks, size_t ddp_world_size, + size_t ddp_rank); + + DistributedOptimizer(OptimizerCreatorNamed base_optimizer_creator, const NamedParameterList &named_parameters, const std::vector> &model_chunks, size_t ddp_world_size, size_t ddp_rank); @@ -43,9 +47,10 @@ class DistributedOptimizer final : public infini_train::Optimizer { virtual float learning_rate() const override; private: - void BuildShardParamsAndBindGrads(const NamedParameterList &named_parameters, - std::vector> &shard_params, - NamedParameterList &shard_named_parameters); + using AddShard = std::function &, const std::shared_ptr &)>; + + void InitializeModelChunks(const std::vector> &model_chunks); + void BuildShardParamsAndBindGrads(const AddShard &add_shard); private: // Inherit from DDP model diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index 8fbe4947..d85b1ace 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -16,13 +16,14 @@ class Optimizer; using NamedParameter = std::pair>; using NamedParameterList = std::vector; -using OptimizerCreator = std::function(const std::vector> ¶ms, - const NamedParameterList &named_parameters)>; +using OptimizerCreator = std::function(const std::vector> ¶ms)>; +using OptimizerCreatorNamed = std::function(const NamedParameterList &named_params)>; class Optimizer { public: - explicit Optimizer(const std::vector> ¶ms, float learning_rate = 0.0f, - const NamedParameterList &named_parameters = {}); + explicit Optimizer(const std::vector> ¶ms, float learning_rate); + + Optimizer(const NamedParameterList &named_params, float learning_rate); virtual void ZeroGrad(bool set_to_none = true); @@ -53,18 +54,21 @@ class Optimizer { namespace optimizers { class SGD : public Optimizer { public: - SGD(const std::vector> ¶ms, float learning_rate, - const NamedParameterList &named_parameters = {}); + SGD(const std::vector> ¶ms, 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> ¶ms, float learning_rate = 1e-3, float beta1 = 0.9, - float beta2 = 0.999, float eps = 1e-8, const NamedParameterList &named_parameters = {}); + 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; @@ -73,6 +77,8 @@ class Adam : public Optimizer { void LoadStateDict(const std::unordered_map> &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_; diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 0d2492b2..9475d49f 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -28,22 +28,12 @@ Module::Module(const std::string &type) : type_(type), device_(Device()) {} const std::string &Module::type() const { return type_; } std::vector> Module::Parameters() const { - std::vector> params; - std::unordered_set visited; - - auto AddIfUnvisited = [&](const std::shared_ptr ¶m) { - if (visited.insert(param.get()).second) { - params.push_back(param); - } - }; + const auto &named_parameters = NamedParameters(); - // Add parameters of this module - for (const auto &[_, param] : parameters_) { AddIfUnvisited(param); } + std::vector> params; + params.reserve(named_parameters.size()); - // Recursively add parameters of submodules - for (const auto &[_, module] : modules_) { - for (const auto ¶m : module->Parameters()) { AddIfUnvisited(param); } - } + for (const auto &[_, param] : named_parameters) { params.emplace_back(param); } return params; } @@ -51,39 +41,41 @@ std::vector> Module::Parameters() const { std::vector>> Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { std::vector>> named_parameters; + std::unordered_set visited_parameters; - std::function collect - = [&](const Module &module, const std::string &module_prefix) { - for (const auto &[name, parameter] : module.parameters_) { - if (!parameter) { - 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); - std::sort(named_parameters.begin(), named_parameters.end(), - [](const auto &lhs, const auto &rhs) { return lhs.first < rhs.first; }); + std::vector>> named_modules; - if (remove_duplicate) { - std::unordered_set visited; - std::erase_if(named_parameters, [&](const auto &named_parameter) { - return !visited.insert(named_parameter.second.get()).second; - }); + if (recurse) { + named_modules = const_cast(this)->NamedModules( + /*memory=*/nullptr, prefix, remove_duplicate); + } else { + named_modules.emplace_back(prefix, std::const_pointer_cast(shared_from_this())); } + + for (const auto &[module_prefix, module] : named_modules) { + std::vector>> local_parameters; + local_parameters.reserve(module->parameters_.size()); + + for (const auto &[name, parameter] : module->parameters_) { + if (parameter != nullptr) { + 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_parameters.insert(parameter.get()).second) { + continue; + } + + const std::string full_name = module_prefix.empty() ? name : module_prefix + "." + name; + + named_parameters.emplace_back(full_name, parameter); + } + } + return named_parameters; } diff --git a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc index 893ad7a1..02d0425b 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc @@ -8,11 +8,49 @@ namespace infini_train::nn::parallel { DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator, const std::vector> &full_params, - const NamedParameterList &named_parameters, const std::vector> &model_chunks, size_t ddp_world_size, size_t ddp_rank) : Optimizer(full_params, /*learning_rate=*/0.0f), ddp_world_size_(ddp_world_size), ddp_rank_(ddp_rank) { + InitializeModelChunks(model_chunks); + std::vector> shard_params; + BuildShardParamsAndBindGrads( + [&shard_params](const std::shared_ptr &, const std::shared_ptr ¶m_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> &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 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( + [¶meter_name_by_tensor, &shard_named_parameters](const std::shared_ptr ¶meter, + const std::shared_ptr ¶m_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> &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) { @@ -24,25 +62,10 @@ DistributedOptimizer::DistributedOptimizer(OptimizerCreator creator, bucket_groups_.insert(bucket_groups_.end(), ddp_chunk->bucket_groups().begin(), ddp_chunk->bucket_groups().end()); } - - std::vector> shard_params; - NamedParameterList shard_named_parameters; - BuildShardParamsAndBindGrads(named_parameters, shard_params, shard_named_parameters); - - // Build base optimizer - base_optimizer_ = creator(shard_params, shard_named_parameters); - CHECK(base_optimizer_) << "DistributedOptimizer: failed to create base optimizer."; } -void DistributedOptimizer::BuildShardParamsAndBindGrads(const NamedParameterList &named_parameters, - std::vector> &shard_params, - NamedParameterList &shard_named_parameters) { - std::unordered_map 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); - } +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; @@ -92,17 +115,14 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads(const NamedParameterList // 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); - 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); + add_shard(param, param_piece); + ++num_shard_params; } } } - CHECK(!shard_params.empty()) << "DistributedOptimizer: this DP rank owns no param pieces. " - << "Check bucket padding/divisibility and param bucketing order."; + CHECK_GT(num_shard_params, 0) << "DistributedOptimizer: this DP rank owns no param pieces. " + << "Check bucket padding/divisibility and param bucketing order."; } void DistributedOptimizer::StartGradSync() { diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 01f61df4..39b999c7 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -9,25 +9,19 @@ #include "infini_train/include/tensor.h" namespace infini_train { -Optimizer::Optimizer(const std::vector> ¶ms, float learning_rate, - const NamedParameterList &named_parameters) - : params_(params), learning_rate_(learning_rate) { - if (named_parameters.empty()) { - return; - } +Optimizer::Optimizer(const std::vector> ¶ms, float learning_rate) + : params_(params), learning_rate_(learning_rate) {} - std::unordered_map 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); +Optimizer::Optimizer(const NamedParameterList &named_params, float learning_rate) : learning_rate_(learning_rate) { + if (named_params.empty()) { + return; } - parameter_names_.reserve(params_.size()); - for (const auto ¶meter : 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); + params_.reserve(named_params.size()); + parameter_names_.reserve(named_params.size()); + for (const auto &[name, parameter] : named_params) { + params_.push_back(parameter); + parameter_names_.push_back(name); } } @@ -55,9 +49,9 @@ void Optimizer::set_initial_learning_rate(float lr) { namespace optimizers { -SGD::SGD(const std::vector> ¶ms, float learning_rate, - const NamedParameterList &named_parameters) - : Optimizer(params, learning_rate, named_parameters) {} +SGD::SGD(const std::vector> ¶ms, float learning_rate) : Optimizer(params, learning_rate) {} + +SGD::SGD(const NamedParameterList &named_params, float learning_rate) : Optimizer(named_params, learning_rate) {} void SGD::Step() { for (auto param : params_) { @@ -73,15 +67,19 @@ void SGD::Step() { } OptimizerCreator SGD::Create(float learning_rate) { - return [learning_rate](const std::vector> ¶ms, - const NamedParameterList &named_parameters) { - return std::make_shared(params, learning_rate, named_parameters); + return [learning_rate](const std::vector> ¶ms) { + return std::make_shared(params, learning_rate); }; } -Adam::Adam(const std::vector> ¶ms, 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) { +OptimizerCreatorNamed SGD::CreateNamed(float learning_rate) { + return [learning_rate](const NamedParameterList &named_params) { + return std::make_shared(named_params, learning_rate); + }; +} + +Adam::Adam(const std::vector> ¶ms, float learning_rate, float beta1, float beta2, float eps) + : Optimizer(params, learning_rate), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) { for (const auto ¶m : params_) { m_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); @@ -91,6 +89,16 @@ Adam::Adam(const std::vector> ¶ms, float learning_ra } } +Adam::Adam(const NamedParameterList &named_params, float learning_rate, float beta1, float beta2, float eps) + : Optimizer(named_params, learning_rate), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) { + for (const auto &[name, param] : named_params) { + m_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); + v_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); + m_.back()->Fill(0.0); + v_.back()->Fill(0.0); + } +} + void Adam::Step() { ++t_; @@ -112,8 +120,14 @@ void Adam::Step() { } OptimizerCreator Adam::Create(float learning_rate, float beta1, float beta2, float eps) { - return [=](const std::vector> ¶ms, const NamedParameterList &named_parameters) { - return std::make_shared(params, learning_rate, beta1, beta2, eps, named_parameters); + return [=](const std::vector> ¶ms) { + return std::make_shared(params, learning_rate, beta1, beta2, eps); + }; +} + +OptimizerCreatorNamed Adam::CreateNamed(float learning_rate, float beta1, float beta2, float eps) { + return [=](const NamedParameterList &named_params) { + return std::make_shared(named_params, learning_rate, beta1, beta2, eps); }; } diff --git a/tests/module/test_named_parameters.cc b/tests/module/test_named_parameters.cc index 1a7a07e7..c1a4a621 100644 --- a/tests/module/test_named_parameters.cc +++ b/tests/module/test_named_parameters.cc @@ -36,11 +36,7 @@ TEST_P(ModuleNamedParametersTest, SupportsRecursionAndSharedParameterDeduplicati EXPECT_EQ(deduplicated[0].first, "model.0.bias"); EXPECT_EQ(deduplicated[1].first, "model.0.weight"); 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()); - } + for (const auto &[name, parameter] : deduplicated) { tensors.insert(parameter.get()); } EXPECT_TRUE(tensors.contains(shared->parameter(nn::Linear::kParamWeightName).get())); EXPECT_TRUE(tensors.contains(shared->parameter(nn::Linear::kParamBiasName).get())); @@ -55,20 +51,18 @@ TEST_P(ModuleNamedParametersTest, ReturnsNestedParametersInStableNameOrder) { 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}); + auto root_linear = std::make_shared(2, 2, false, GetDevice()); + auto root = std::make_shared(std::vector>{root_linear, nested}); const auto parameters = root->NamedParameters(); - const std::unordered_map> by_name(parameters.begin(), parameters.end()); - ASSERT_EQ(by_name.size(), 3); ASSERT_EQ(parameters.size(), 3); EXPECT_EQ(parameters[0].first, "0.weight"); EXPECT_EQ(parameters[1].first, "1.0.weight"); EXPECT_EQ(parameters[2].first, "1.1.weight"); - EXPECT_TRUE(by_name.contains("0.weight")); - EXPECT_TRUE(by_name.contains("1.0.weight")); - EXPECT_TRUE(by_name.contains("1.1.weight")); + EXPECT_EQ(parameters[0].second, root_linear->parameter(nn::Linear::kParamWeightName)); + EXPECT_EQ(parameters[1].second, first->parameter(nn::Linear::kParamWeightName)); + EXPECT_EQ(parameters[2].second, second->parameter(nn::Linear::kParamWeightName)); } TEST_P(ModuleNamedParametersTest, SkipsNullSubmodules) { diff --git a/tests/optimizer/test_optimizer_parameter_names.cc b/tests/optimizer/test_optimizer_parameter_names.cc index afb0f187..31943454 100644 --- a/tests/optimizer/test_optimizer_parameter_names.cc +++ b/tests/optimizer/test_optimizer_parameter_names.cc @@ -23,9 +23,8 @@ 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()); - const NamedParameterList named_parameters{{"transformer.h.0.weight", first}, - {"transformer.h.0.bias", second}}; - auto adam = optimizers::Adam::Create(0.001)({first, second}, named_parameters); + const NamedParameterList named_parameters{{"transformer.h.0.weight", first}, {"transformer.h.0.bias", second}}; + auto adam = optimizers::Adam::CreateNamed(0.001)(named_parameters); const auto state = adam->StateDict(); EXPECT_TRUE(state.contains("adam.m.transformer.h.0.weight")); @@ -34,7 +33,7 @@ TEST_P(OptimizerParameterNamesTest, AdamStateDictUsesStableParameterNames) { EXPECT_TRUE(state.contains("adam.v.transformer.h.0.bias")); EXPECT_TRUE(state.contains("adam.t")); - auto restored = optimizers::Adam::Create(0.001)({first, second}, named_parameters); + auto restored = optimizers::Adam::CreateNamed(0.001)(named_parameters); restored->LoadStateDict(state); EXPECT_EQ(restored->StateDict().size(), state.size()); } @@ -42,9 +41,9 @@ TEST_P(OptimizerParameterNamesTest, AdamStateDictUsesStableParameterNames) { TEST_P(OptimizerParameterNamesTest, ConstructorMatchesNamesToOptimizerParameterOrder) { auto first = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); auto second = std::make_shared(std::vector{3}, DataType::kFLOAT32, GetDevice()); - const NamedParameterList named_parameters{{"first", first}, {"second", second}}; + const NamedParameterList named_parameters{{"second", second}, {"first", first}}; - auto adam = optimizers::Adam::Create(0.001)({second, first}, named_parameters); + auto adam = optimizers::Adam::CreateNamed(0.001)(named_parameters); const auto state = adam->StateDict(); EXPECT_TRUE(state.contains("adam.m.second")); @@ -66,7 +65,6 @@ TEST_P(OptimizerParameterNamesTest, DistributedOptimizerPropagatesNamesToShardOp nn::parallel::GetDataParallelGroupRanks(rank.GlobalRank())); auto model = std::make_shared(4, 4, /*bias=*/false, GetDevice()); - const auto params = model->Parameters(); const auto named_parameters = model->NamedParameters(); nn::parallel::DistributedDataParallelConfig ddp_config; @@ -75,7 +73,7 @@ TEST_P(OptimizerParameterNamesTest, DistributedOptimizerPropagatesNamesToShardOp ddp_config.overlap_param_gather = false; auto ddp_model = std::make_shared(model, rank, ddp_config); - nn::parallel::DistributedOptimizer optimizer(optimizers::Adam::Create(0.001), params, named_parameters, + nn::parallel::DistributedOptimizer optimizer(optimizers::Adam::CreateNamed(0.001), named_parameters, std::vector>{ddp_model}, /*ddp_world_size=*/2, /*ddp_rank=*/0); const auto state = optimizer.StateDict();