Cumulative update on optimizers and tests (on-device training) (#15499)

This commit is contained in:
pengwa 2023-04-29 00:55:39 +08:00 committed by GitHub
parent 8a1a40ac63
commit 29d13cea42
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
9 changed files with 550 additions and 258 deletions

View file

@ -321,7 +321,7 @@ TEST(CheckpointApiTest, SaveOptimizerStateAsCheckpoint_ThenLoad_CUDA) {
ASSERT_EQ(param_tensor.DataType(), restored_tensor.DataType());
std::vector<float> state_vect;
OrtValueToVec(restored_ort_value, state_vect);
CpuOrtValueToVec(restored_ort_value, state_vect);
for (size_t i = 0; i < state_vect.size(); i++) {
ASSERT_EQ(state_vect[i], 0.0f);
}

View file

@ -8,24 +8,26 @@
#include "core/framework/ort_value.h"
#include "core/framework/tensor.h"
#include "default_providers.h"
#include "test/framework/test_utils.h"
#include "test/util/include/test_utils.h"
namespace onnxruntime::training::test {
template <typename T>
void OrtValueToVec(const OrtValue& val, std::vector<T>& output) {
const Tensor& tensor = val.Get<Tensor>();
void CpuOrtValueToVec(const OrtValue& src_cpu_ortvalue, std::vector<T>& output) {
const Tensor& tensor = src_cpu_ortvalue.Get<Tensor>();
int64_t num_elem = tensor.Shape().Size();
const T* val_ptr = tensor.template Data<T>();
output.assign(val_ptr, val_ptr + num_elem);
}
template <typename T>
void CudaOrtValueToCpuVec(const OrtValue& val, std::vector<T>& output,
std::shared_ptr<IExecutionProvider> cuda_provider,
std::shared_ptr<IExecutionProvider> cpu_provider) {
const Tensor& src_tensor = val.Get<Tensor>();
void CudaOrtValueToCpuVec(const OrtValue& src_cuda_ortvalue, std::vector<T>& output) {
std::unique_ptr<IExecutionProvider> cuda_provider = onnxruntime::test::DefaultCudaExecutionProvider();
std::unique_ptr<IExecutionProvider> cpu_provider = onnxruntime::test::DefaultCpuExecutionProvider();
const Tensor& src_tensor = src_cuda_ortvalue.Get<Tensor>();
auto allocator = cpu_provider->GetAllocator(OrtMemTypeDefault);
ORT_ENFORCE(allocator, "Cpu allocator is a nullptr.");

View file

@ -19,6 +19,7 @@
#include "default_providers.h"
using json = nlohmann::json;
using namespace onnxruntime::training::api;
namespace onnxruntime {
namespace training {
@ -28,6 +29,43 @@ namespace {
#define MODEL_FOLDER ORT_TSTR("testdata/training_api/")
constexpr int64_t TOTAL_STEP_COUNT = 100;
constexpr float INITIAL_LR = 1e-3f;
/**
* @brief Create a Fake Optimizer Checkpoint State On CPU.
*
* @param named_parameters Parameter list
* @param momentum_keys Optimizer momentum keys.
* @param optimizer_checkpoint_state Used as output to store the state containing faked data.
* @return Status
*/
Status CreateFakeOptimizerCheckpointStateOnCPU(
const std::unordered_map<std::string, std::shared_ptr<Parameter>>& named_parameters,
const std::vector<std::string>& momentum_keys,
OptimizerCheckpointState& optimizer_checkpoint_state) {
auto& grouped_optimizer_states = optimizer_checkpoint_state.group_named_optimizer_states;
grouped_optimizer_states.insert({"group0", std::make_shared<GroupOptimizerState>()});
GroupOptimizerState& group_optimizer_state = *(grouped_optimizer_states["group0"]);
auto& param_named_optimizer_states = group_optimizer_state.param_named_optimizer_states;
for (auto& pair : named_parameters) {
if (pair.second->RequiresGrad()) {
param_named_optimizer_states.insert({pair.first, ParameterOptimizerState()});
ParameterOptimizerState& cur_param_optimizer_states = param_named_optimizer_states[pair.first];
for (auto& state_name : momentum_keys) {
OrtValue param_moment_state;
OrtValue param = pair.second->Data();
const auto& param_tensor = param.template Get<Tensor>();
GenerateRandomInput(param_tensor.Shape().GetDims(), param_moment_state);
cur_param_optimizer_states.momentum_named_states.insert({state_name, std::move(param_moment_state)});
}
}
}
return Status::OK();
}
void TestModuleExport(const std::vector<std::shared_ptr<IExecutionProvider>>& providers) {
auto training_model_uri = MODEL_FOLDER "training_model.onnx";
auto eval_model_uri = MODEL_FOLDER "eval_model.onnx";
@ -91,88 +129,6 @@ void TestModuleExport(const std::vector<std::shared_ptr<IExecutionProvider>>& pr
ASSERT_EQ(outputs.size(), 1U);
}
#if defined(USE_CUDA)
constexpr int64_t total_step_count = 100;
constexpr float initial_lr = 1e-3f;
constexpr int64_t resume_step = total_step_count / 2;
void CompareValue(float expected, float output, float rtol = 1e-4, float atol = 1e-5) {
ASSERT_NEAR(expected, output, atol);
ASSERT_NEAR(expected, output, rtol * std::abs(expected));
}
void TestLRSchduler(const std::basic_string<ORTCHAR_T>& test_file_name, float initial_lr, int64_t total_step_count,
int64_t warmup_step_count) {
/// Load model and optimizer graph, create Module, Optimizer and LRScheduler instances.
auto model_uri = MODEL_FOLDER "training_model.onnx";
auto optim_uri = MODEL_FOLDER "adamw.onnx";
onnxruntime::training::api::CheckpointState state;
auto checkpoint_to_load_path = MODEL_FOLDER "checkpoint.ckpt";
ASSERT_STATUS_OK(LoadCheckpoint(checkpoint_to_load_path, state));
onnxruntime::SessionOptions session_option;
std::unique_ptr<Environment> env;
ASSERT_STATUS_OK(Environment::Create(nullptr, env));
const std::vector<std::shared_ptr<IExecutionProvider>> providers{onnxruntime::test::DefaultCudaExecutionProvider()};
auto model = std::make_unique<onnxruntime::training::api::Module>(
ToUTF8String(model_uri), &state,
session_option, *env, providers);
auto optim = std::make_shared<onnxruntime::training::api::Optimizer>(
ToUTF8String(optim_uri), &state, session_option,
*env, providers);
OrtValue input, target;
GenerateRandomInput(std::array<int64_t, 2>{2, 784}, input);
onnxruntime::test::CreateInputOrtValueOnCPU<int32_t>(
std::array<int64_t, 1>{2}, std::vector<int32_t>(2, 1), &target);
/// Load test data for learning rate schedulers.
auto data_uri = ORT_TSTR("testdata/test_data_generation/lr_scheduler/" + test_file_name);
std::ifstream in{data_uri};
// Element of vector represent a pair of <step_count, list of learning rates>>
typedef std::vector<std::pair<int64_t, std::vector<float>>> TestDataDictType;
TestDataDictType test_data;
const json j = json::parse(in);
j.get_to<TestDataDictType>(test_data);
int64_t resume_step = (*test_data.begin()).first;
ASSERT_EQ(total_step_count, static_cast<int64_t>(test_data.size()) + resume_step);
if (resume_step != 0) {
/// Reset optimizer states to match the initial state we want to test.
onnxruntime::training::api::OptimizerCheckpointState optimizer_checkpoint_states;
auto group_opt_state =
optimizer_checkpoint_states.group_named_optimizer_states["group0"] =
std::make_shared<onnxruntime::training::api::GroupOptimizerState>();
group_opt_state->step = resume_step;
group_opt_state->initial_lr = initial_lr;
ASSERT_STATUS_OK(optim->LoadStateDict(optimizer_checkpoint_states));
}
// KNOWN ISSUE: LinearLRScheduler by default use optim's states to calculate the first step's learning rate.
// If we restored it after creation, it will only affect the learning rate from the second step.
auto scheduler = std::make_unique<onnxruntime::training::api::LinearLRScheduler>(
optim, warmup_step_count, total_step_count);
for (auto it = test_data.begin(); it != test_data.end(); ++it) {
onnxruntime::training::api::OptimizerCheckpointState optimizer_states;
ASSERT_STATUS_OK(optim->GetStateDict(optimizer_states));
auto group_optimizer_state = optimizer_states.group_named_optimizer_states["group0"];
CompareValue(it->second[0], group_optimizer_state->learning_rate);
ASSERT_EQ(it->first, group_optimizer_state->step);
std::vector<OrtValue> inputs{input, target};
std::vector<OrtValue> fetches;
ASSERT_STATUS_OK(model->TrainStep(inputs, fetches));
ASSERT_STATUS_OK(optim->Step());
ASSERT_STATUS_OK(scheduler->Step());
}
}
#endif
} // namespace
TEST(TrainingApiTest, ModuleParametersSize) {
@ -271,12 +227,12 @@ TEST(TrainingApiTest, ModuleTrainStep) {
bias_grad = bias_param->Gradient();
if (step > 1) {
OrtValueToVec(bias_grad, current_bias_grad_vec);
CpuOrtValueToVec(bias_grad, current_bias_grad_vec);
for (size_t i = 0; i < current_bias_grad_vec.size(); i++) {
ASSERT_EQ(current_bias_grad_vec[i], single_bias_grad_vec[i] * step);
}
} else {
OrtValueToVec(bias_grad, single_bias_grad_vec);
CpuOrtValueToVec(bias_grad, single_bias_grad_vec);
}
}
// reset grad
@ -286,12 +242,215 @@ TEST(TrainingApiTest, ModuleTrainStep) {
std::vector<OrtValue>& inputs = *data_loader.begin();
std::vector<OrtValue> fetches;
ASSERT_STATUS_OK(model->TrainStep(inputs, fetches));
OrtValueToVec(bias_grad, current_bias_grad_vec);
CpuOrtValueToVec(bias_grad, current_bias_grad_vec);
for (size_t i = 0; i < current_bias_grad_vec.size(); i++) {
ASSERT_EQ(current_bias_grad_vec[i], single_bias_grad_vec[i]);
}
}
TEST(TrainingApiTest, OptimizerCreatedWithOptimizerCheckpointState) {
std::vector<bool> run_cuda_list{false};
// #ifdef USE_CUDA
// run_cuda_list.push_back(true);
// #endif
for (auto run_cuda : run_cuda_list) {
std::vector<std::shared_ptr<IExecutionProvider>> providers;
if (run_cuda) {
providers = {onnxruntime::test::DefaultCudaExecutionProvider()};
} else {
providers = {onnxruntime::test::DefaultCpuExecutionProvider()};
}
auto model_uri = MODEL_FOLDER "training_model.onnx";
auto optim_uri = MODEL_FOLDER "adamw.onnx";
CheckpointState state;
auto checkpoint_to_load_path = MODEL_FOLDER "checkpoint.ckpt";
ASSERT_STATUS_OK(onnxruntime::training::api::LoadCheckpoint(checkpoint_to_load_path, state));
onnxruntime::SessionOptions session_option;
std::unique_ptr<Environment> env;
ASSERT_STATUS_OK(Environment::Create(nullptr, env));
std::shared_ptr<Module> model = std::make_shared<Module>(
ToUTF8String(model_uri), &state, session_option,
*env, providers);
// Load state dict from faked optimizer checkpoint state.
CheckpointState new_state = state;
OptimizerCheckpointState& external_optimizer_checkpoint_state = new_state.optimizer_checkpoint_state;
ASSERT_STATUS_OK(CreateFakeOptimizerCheckpointStateOnCPU(model->NamedParameters(),
{"momentum0", "momentum1"},
external_optimizer_checkpoint_state));
std::shared_ptr<Optimizer> optim = std::make_shared<Optimizer>(
ToUTF8String(optim_uri), &new_state, session_option, *env, providers);
// After loading state dict, check if optim state is updated to new states.
OptimizerCheckpointState optimizer_states;
ASSERT_STATUS_OK(optim->GetStateDict(optimizer_states));
for (auto& p : model->NamedParameters()) {
auto param_name = p.first;
ParameterOptimizerState& param_state =
optimizer_states.group_named_optimizer_states["group0"]->param_named_optimizer_states.at(param_name);
ParameterOptimizerState& external_param_state =
external_optimizer_checkpoint_state.group_named_optimizer_states["group0"]
->param_named_optimizer_states.at(param_name);
for (auto& param_p : param_state.momentum_named_states) {
std::vector<float> moment_vec;
if (run_cuda) {
CudaOrtValueToCpuVec(param_state.momentum_named_states.at(param_p.first), moment_vec);
} else {
CpuOrtValueToVec(param_state.momentum_named_states.at(param_p.first), moment_vec);
}
std::vector<float> external_moment_vect;
if (run_cuda) {
CudaOrtValueToCpuVec(external_param_state.momentum_named_states.at(param_p.first), external_moment_vect);
} else {
CpuOrtValueToVec(external_param_state.momentum_named_states.at(param_p.first), external_moment_vect);
}
ASSERT_EQ(moment_vec.size(), external_moment_vect.size());
for (size_t i = 0; i < moment_vec.size(); i++) {
ASSERT_EQ(moment_vec[i], external_moment_vect[i]);
}
}
}
}
}
void TestLRSchduler(const std::basic_string<ORTCHAR_T>& test_file_name,
float initial_lr,
int64_t total_step_count,
int64_t warmup_step_count) {
std::vector<bool> run_cuda_list{false};
// #ifdef USE_CUDA
// run_cuda_list.push_back(true);
// #endif
for (auto run_cuda : run_cuda_list) {
std::vector<std::shared_ptr<IExecutionProvider>> providers;
if (run_cuda) {
providers = {onnxruntime::test::DefaultCudaExecutionProvider()};
} else {
providers = {onnxruntime::test::DefaultCpuExecutionProvider()};
}
auto model_uri = MODEL_FOLDER "training_model.onnx";
auto optim_uri = MODEL_FOLDER "adamw.onnx";
CheckpointState state;
auto checkpoint_to_load_path = MODEL_FOLDER "checkpoint.ckpt";
ASSERT_STATUS_OK(onnxruntime::training::api::LoadCheckpoint(checkpoint_to_load_path, state));
onnxruntime::SessionOptions session_option;
std::unique_ptr<Environment> env;
ASSERT_STATUS_OK(Environment::Create(nullptr, env));
std::shared_ptr<Module> model = std::make_shared<Module>(
ToUTF8String(model_uri), &state, session_option,
*env, providers);
OrtValue input, target;
GenerateRandomInput(std::array<int64_t, 2>{2, 784}, input);
onnxruntime::test::CreateInputOrtValueOnCPU<int32_t>(
std::array<int64_t, 1>{2}, std::vector<int32_t>(2, 1), &target);
/// Load test data for learning rate schedulers.
auto data_uri = ORT_TSTR("testdata/test_data_generation/lr_scheduler/" + test_file_name);
std::ifstream in{data_uri};
// Element of vector represent a pair of <step_count, list of learning rates>>
typedef std::vector<std::pair<int64_t, std::vector<float>>> TestDataDictType;
TestDataDictType test_data;
const json j = json::parse(in);
j.get_to<TestDataDictType>(test_data);
int64_t resume_step = (*test_data.begin()).first;
ASSERT_EQ(total_step_count, static_cast<int64_t>(test_data.size()) + resume_step);
if (resume_step != 0) {
state.optimizer_checkpoint_state.group_named_optimizer_states.insert(
{"group0", std::make_shared<GroupOptimizerState>()});
auto& group_opt_state = state.optimizer_checkpoint_state.group_named_optimizer_states["group0"];
/// Reset optimizer states to match the initial state we want to test.
group_opt_state->step = resume_step;
group_opt_state->initial_lr = initial_lr;
}
std::shared_ptr<Optimizer> optim = std::make_shared<Optimizer>(
ToUTF8String(optim_uri), &state, session_option,
*env, providers);
// KNOWN ISSUE: LinearLRScheduler by default use optim's states to calculate the first step's learning rate.
// If we restored it after creation, it will only affect the learning rate from the second step.
auto scheduler = std::make_unique<LinearLRScheduler>(
optim, warmup_step_count, total_step_count);
for (auto it = test_data.begin(); it != test_data.end(); ++it) {
OptimizerCheckpointState optimizer_states;
ASSERT_STATUS_OK(optim->GetStateDict(optimizer_states));
auto group_optimizer_state = optimizer_states.group_named_optimizer_states["group0"];
constexpr const float rtol = 1e-4f, atol = 1e-5f;
ASSERT_NEAR(it->second[0], group_optimizer_state->learning_rate, atol);
ASSERT_NEAR(it->second[0], group_optimizer_state->learning_rate, rtol * std::abs(it->second[0]));
ASSERT_EQ(it->first, group_optimizer_state->step);
std::vector<OrtValue> inputs{input, target};
std::vector<OrtValue> fetches;
ASSERT_STATUS_OK(model->TrainStep(inputs, fetches));
ASSERT_STATUS_OK(optim->Step());
ASSERT_STATUS_OK(scheduler->Step());
}
}
}
TEST(TrainingApiTest, LinearLRScheduler_NoWarmUp_Test) {
// No warm up.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-0.json"), INITIAL_LR, TOTAL_STEP_COUNT, 0);
}
TEST(TrainingApiTest, LinearLRScheduler_NoWarmUp_ResumeFromCheckpoint_Test) {
// No warm up.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-0_restored.json"), INITIAL_LR, TOTAL_STEP_COUNT, 0);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp30Step_Test) {
// Warmp up completed before saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-30.json"), INITIAL_LR, TOTAL_STEP_COUNT, 30);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp30Step_ResumeFromCheckpoint_Test) {
// Warmp up completed before saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-30_restored.json"), INITIAL_LR, TOTAL_STEP_COUNT, 30);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp70Step_Test) {
// Warmp up completed after saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-70.json"), INITIAL_LR, TOTAL_STEP_COUNT, 70);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp70Step_ResumeFromCheckpoint_Test) {
// Warmp up completed after saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-70_restored.json"), INITIAL_LR, TOTAL_STEP_COUNT, 70);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp200Step_Test) {
// All steps are in warm-up phase.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-200.json"), INITIAL_LR, TOTAL_STEP_COUNT, 200);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp200Step_ResumeFromCheckpoint_Test) {
// All steps are in warm-up phase.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-200_restored.json"), INITIAL_LR, TOTAL_STEP_COUNT, 200);
}
TEST(TrainingApiTest, ModuleExportModelForInferencingCPU) {
std::vector<std::shared_ptr<IExecutionProvider>> providers{onnxruntime::test::DefaultCpuExecutionProvider()};
TestModuleExport(providers);
@ -316,7 +475,6 @@ TEST(TrainingApiTest, OptimStep) {
std::unique_ptr<Environment> env;
std::vector<std::shared_ptr<IExecutionProvider>> providers{onnxruntime::test::DefaultCudaExecutionProvider()};
std::shared_ptr<IExecutionProvider> cuda_provider = providers.front();
std::shared_ptr<IExecutionProvider> cpu_provider = onnxruntime::test::DefaultCpuExecutionProvider();
ASSERT_STATUS_OK(Environment::Create(nullptr, env));
auto model = std::make_unique<onnxruntime::training::api::Module>(
ToUTF8String(model_uri), &state, session_option,
@ -342,10 +500,9 @@ TEST(TrainingApiTest, OptimStep) {
OrtValue& moment_1 = param_state.momentum_named_states.at("momentum0");
std::vector<float> param_vec_before_optimizer_step;
CudaOrtValueToCpuVec(model->NamedParameters().at(param_name)->Data(), param_vec_before_optimizer_step,
cuda_provider, cpu_provider);
CudaOrtValueToCpuVec(model->NamedParameters().at(param_name)->Data(), param_vec_before_optimizer_step);
std::vector<float> moment_1_vec;
CudaOrtValueToCpuVec(moment_1, moment_1_vec, cuda_provider, cpu_provider);
CudaOrtValueToCpuVec(moment_1, moment_1_vec);
for (size_t i = 0; i < moment_1_vec.size(); i++) {
ASSERT_EQ(moment_1_vec[i], 0.0f);
}
@ -356,12 +513,11 @@ TEST(TrainingApiTest, OptimStep) {
std::vector<OrtValue> fetches;
ASSERT_STATUS_OK(model->TrainStep(inputs, fetches));
std::vector<float> grads;
CudaOrtValueToCpuVec(model->NamedParameters().at(param_name)->Gradient(), grads,
cuda_provider, cpu_provider);
CudaOrtValueToCpuVec(model->NamedParameters().at(param_name)->Gradient(), grads);
ASSERT_STATUS_OK(optim->Step());
// get optim state and check if it is updated
CudaOrtValueToCpuVec(moment_1, moment_1_vec, cuda_provider, cpu_provider);
CudaOrtValueToCpuVec(moment_1, moment_1_vec);
for (size_t i = 0; i < moment_1_vec.size(); i++) {
if (grads[i] != 0.0f) {
ASSERT_NE(moment_1_vec[i], 0.0f);
@ -369,8 +525,7 @@ TEST(TrainingApiTest, OptimStep) {
}
std::vector<float> param_vec_after_optimizer_step;
CudaOrtValueToCpuVec(model->NamedParameters().at(param_name)->Data(), param_vec_after_optimizer_step,
cuda_provider, cpu_provider);
CudaOrtValueToCpuVec(model->NamedParameters().at(param_name)->Data(), param_vec_after_optimizer_step);
for (size_t i = 0; i < param_vec_after_optimizer_step.size(); ++i) {
if (grads[i] != 0.0f && moment_1_vec[i] != 0.0f) {
ASSERT_NE(param_vec_after_optimizer_step[i], param_vec_before_optimizer_step[i]);
@ -379,46 +534,6 @@ TEST(TrainingApiTest, OptimStep) {
}
}
TEST(TrainingApiTest, LinearLRScheduler_NoWarmUp_Test) {
// No warm up.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-0.json"), initial_lr, total_step_count, 0);
}
TEST(TrainingApiTest, LinearLRScheduler_NoWarmUp_ResumeFromCheckpoint_Test) {
// No warm up.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-0_restored.json"), initial_lr, total_step_count, 0);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp30Step_Test) {
// Warmp up completed before saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-30.json"), initial_lr, total_step_count, 30);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp30Step_ResumeFromCheckpoint_Test) {
// Warmp up completed before saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-30_restored.json"), initial_lr, total_step_count, 30);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp70Step_Test) {
// Warmp up completed after saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-70.json"), initial_lr, total_step_count, 70);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp70Step_ResumeFromCheckpoint_Test) {
// Warmp up completed after saving checkpoint.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-70_restored.json"), initial_lr, total_step_count, 70);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp200Step_Test) {
// All steps are in warm-up phase.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-200.json"), initial_lr, total_step_count, 200);
}
TEST(TrainingApiTest, LinearLRScheduler_WarmUp200Step_ResumeFromCheckpoint_Test) {
// All steps are in warm-up phase.
TestLRSchduler(ORT_TSTR("warmup_linear_scheduler_warmupstep-200_restored.json"), initial_lr, total_step_count, 200);
}
#endif
} // namespace test

View file

@ -157,8 +157,8 @@ Module::Module(const std::string& train_model_path_or_bytes,
const std::optional<std::string>& eval_model_path_or_bytes)
: state_{state} {
// Enforce weight prepacking is disabled
// If user explicitly enabled weight prepacking then return error.
// Default value is enabled. Therefore, explicitly disable it if the value is not set by user.
// If the user explicitly enabled weight prepacking then return an error.
// Default value is enabled. Therefore, explicitly disable it if the value is not set by the user.
std::string disable_prepacking = "";
if (session_options.config_options.TryGetConfigEntry(kOrtSessionOptionsConfigDisablePrepacking, disable_prepacking)) {
ORT_ENFORCE(disable_prepacking == "1", "Prepacking is not supported for training scenarios.");
@ -214,13 +214,13 @@ Module::Module(const std::string& train_model_path_or_bytes,
}
}
// Loop each parameter, allocate it's memory based on user specified device.
// Loop each parameter, and allocate its memory based on the user-specified device.
auto& train_sess_state = train_sess_->GetSessionState();
for (auto& param_name : param_input_names) {
auto params_iter = state_->module_checkpoint_state.named_parameters.find(param_name);
ORT_ENFORCE(params_iter != state_->module_checkpoint_state.named_parameters.end());
// Retrieve the target device for "param_name"
// Retrieve the target device for "param_name".
InlinedVector<SessionState::NodeInfo> node_info_vec;
ORT_THROW_IF_ERROR(train_sess_state.GetInputNodeInfo(param_name, node_info_vec));
const auto& node_info = node_info_vec.front();
@ -229,14 +229,13 @@ Module::Module(const std::string& train_model_path_or_bytes,
ORT_ENFORCE(target_device == *(it->device), "Inconsistent device requirements found for input: ", param_name);
}
// TODO(pengwa): consider whether we should alloc contiguous buffer for parameters or gradients.
// Copy ortvalue buffer from CPU to target_device for this "param_name" (based on graph partitioning)
// Only copies data if target device is not the same as the current device the buffer is placed on
// Only copies data if the target device is not the same as the current device the buffer is placed on
OrtValue& param_data = params_iter->second->Data();
ORT_ENFORCE(param_data.IsTensor());
const Tensor& param_data_tensor = param_data.Get<Tensor>();
// If the source device type is already same as target device skip copy
// If the source device type is already the same as target device skip copy
if (param_data_tensor.Location().device.Type() != target_device.Type()) {
// TODO: move this outside of the for loop?
auto target_allocator = train_sess_state.GetAllocator(target_device);
@ -257,7 +256,7 @@ Module::Module(const std::string& train_model_path_or_bytes,
if (params_iter->second->RequiresGrad()) {
// Create gradient accumulation buffer.
auto it = param_name_to_grad_input_index_map.find(param_name);
ORT_ENFORCE(it != param_name_to_grad_input_index_map.end(), "Gradient buffer input not providered for param: ",
ORT_ENFORCE(it != param_name_to_grad_input_index_map.end(), "Gradient buffer input not provided for param: ",
param_name);
const size_t grad_input_index = it->second;
@ -265,7 +264,7 @@ Module::Module(const std::string& train_model_path_or_bytes,
// TODO: don't pre-allocate the gradient buffer.
// Gradient usually stays on the same device of its parameter.
OrtValue param_grad;
ORT_THROW_IF_ERROR(utils::OrtValueLike(train_sess_state, param_data, param_grad));
ORT_THROW_IF_ERROR(utils::CreateZeroValuedOrtValueLike(train_sess_state, param_data, param_grad));
ORT_THROW_IF_ERROR(params_iter->second->SetGrad(param_grad_name, param_grad));
gradients_[grad_input_index] = params_iter->second->Gradient();
}
@ -292,9 +291,9 @@ Module::Module(const std::string& train_model_path_or_bytes,
eval_param_input_names.emplace_back(input_name);
continue;
} else {
// It is a user input. We handle user inputs separately in eval
// because eval graph might have different user inputs.
// Eg if loss is not a part of eval graph, it won't have
// It is user input. We handle user inputs separately in the eval
// because the eval graph might have different user inputs.
// Eg if loss is not a part of the eval graph, it won't have
// certain inputs like targets
eval_user_input_names.emplace_back(input_name);
}
@ -492,7 +491,7 @@ Status Module::ExportModelForInferencing(const std::string& inference_model_path
ORT_RETURN_IF_ERROR(TransformModelInputsForInference(inference_model->MainGraph(), state_->module_checkpoint_state.named_parameters,
eval_sess_->GetDataTransferManager()));
// Save the model at desired location.
// Save the model at the desired location.
ORT_THROW_IF_ERROR(Model::Save(*inference_model, inference_model_path));
return Status::OK();

View file

@ -54,6 +54,21 @@ struct ModuleCheckpointState {
struct CheckpointState;
/**
* @brief Module class for running training forward and backward.
*
* This class is responsible for running forward and backward.
* It does NOT own the parameters but only holds a weak reference to the passed
* 'CheckpointState' in the constructor.
*
* During initialization, if the Parameter (stored in `CheckpointState`)'s
* device does not match the target device, it will re-create the tensor on the
* target device and update the Parameter's data in place. The 'target device'
* is extracted from node placement.
*
* Currently, we only support load checkpoints from the constructor;
* no public API to load state dict after Module instance is created.
*/
struct Module {
public:
// Initialize a module from an ORT inference session with loaded

View file

@ -17,23 +17,11 @@ namespace api {
namespace {
// Currently all parameters are in a single group, so we hardcode group0 here.
constexpr char GROUP_ZERO_NAME[] = "group0";
// TODO: don't hard code the state names, should get the state names according to the optimizer types.
// TODO: Consolidate with frontend tooling
const std::vector<std::string> MOMENT_STATE_NAMES{"momentum0", "momentum1"};
constexpr std::array AdamWOptimizerInputs = {
"learning_rate",
"step",
"params",
"gradients",
"first_order_moments",
"second_order_moments"};
static constexpr std::array CommonOptimizerInputs{"learning_rate", "step", "params", "gradients"};
Status GraphInputsAreExpected(gsl::span<std::string> actual_graph_inputs,
gsl::span<const char* const> expected_graph_inputs) {
gsl::span<std::string> expected_graph_inputs) {
const auto stringify = [](const auto& container) {
if (container.empty()) {
return std::string("[]");
@ -71,86 +59,129 @@ Status GraphInputsAreExpected(gsl::span<std::string> actual_graph_inputs,
} // namespace
Status Optimizer::GenerateMomentumNamedStates() {
std::unique_ptr<OptimizerAlgorithmBase> OptimizerAlorithmFactory::CreateInstance(
const std::string& optim_path_or_bytes, int32_t& group_count) {
std::shared_ptr<Model> model;
ORT_ENFORCE(Model::Load(ToWideString(optim_path_or_bytes), model, nullptr,
logging::LoggingManager::DefaultLogger())
.IsOK());
Graph& graph = model->MainGraph();
std::map<std::pair<std::string, std::string>, int32_t> opt_type_to_freq_map;
for (auto& node : graph.Nodes()) {
if (node.Domain() == kMSDomain && (node.OpType() == "AdamWOptimizer" || node.OpType() == "SGDOptimizerV2")) {
auto domain_type_pair = std::make_pair(node.Domain(), node.OpType());
if (opt_type_to_freq_map.find(domain_type_pair) == opt_type_to_freq_map.end()) {
opt_type_to_freq_map[domain_type_pair] = 0;
}
opt_type_to_freq_map[domain_type_pair] += 1;
}
}
ORT_ENFORCE(opt_type_to_freq_map.size() == 1U, "Only support one type of optimizer algorithm, but got: " +
std::to_string(opt_type_to_freq_map.size()));
auto opt_it = opt_type_to_freq_map.begin();
group_count = opt_it->second;
auto& domain = opt_it->first.first;
auto& type = opt_it->first.second;
// TODO: to support multiple groups, need to create a mapping between each group to its parameter list.
if (domain == kMSDomain && type == "AdamWOptimizer") {
return std::make_unique<AdamWOptimizerAlgorithm>();
} else if (domain == kMSDomain && type == "SGDOptimizerV2") {
return std::make_unique<SGDOptimizerV2Algorithm>();
} else {
ORT_NOT_IMPLEMENTED("Not implemented for optimizer algo: " + opt_it->first.second);
}
}
Status Optimizer::GenerateMomentumNamedStates(OptimizerCheckpointState& optimizer_checkpoint_states) {
auto group_optimizer_state_it =
optimizer_checkpoint_states.group_named_optimizer_states.find(GROUP_ZERO_NAME);
ORT_ENFORCE(group_optimizer_state_it != optimizer_checkpoint_states.group_named_optimizer_states.end(),
"Group 0 not found in the optimizer checkpoint states.");
optimizer_state_ = group_optimizer_state_it->second;
auto& param_named_optimizer_states = optimizer_state_->param_named_optimizer_states;
auto& optim_sess_state = optim_sess_->GetSessionState();
for (auto& pair : state_->module_checkpoint_state.named_parameters) {
if (pair.second->RequiresGrad()) {
param_named_optimizer_states.insert({pair.first, ParameterOptimizerState()});
ParameterOptimizerState& cur_param_optimizer_states = param_named_optimizer_states[pair.first];
for (auto& state_name : MOMENT_STATE_NAMES) {
for (auto& state_name : optimizer_algo_ptr_->momentum_keys) {
OrtValue param_state;
ORT_ENFORCE(utils::OrtValueLike(optim_sess_state, pair.second->Data(), param_state).IsOK(),
ORT_ENFORCE(utils::CreateZeroValuedOrtValueLike(optim_sess_state, pair.second->Data(), param_state).IsOK(),
"Error generating moment state for ", pair.first);
cur_param_optimizer_states.momentum_named_states.insert({state_name, std::move(param_state)});
}
}
}
return Status::OK();
}
// Constructs the ortvalue inputs to be fed to the graph at each step
Status Optimizer::ConstructInputs() {
if (optimizer_type_ == OptimizerType::AdamW) {
auto& param_named_optimizer_states = optimizer_state_->param_named_optimizer_states;
inputs_.clear();
std::vector<Tensor> params, grads, first_order_moments, second_order_moments;
auto& param_named_optimizer_states = optimizer_state_->param_named_optimizer_states;
// Collect all the non user defined inputs from the state_->module_checkpoint_state.named_parameters.
for (auto& [parameter_name, parameter] : state_->module_checkpoint_state.named_parameters) {
if (parameter->RequiresGrad()) {
// Collect parameters and prepare for tensorseq creation
auto* param_tensor = parameter->Data().GetMutable<Tensor>();
params.emplace_back(
Tensor(param_tensor->DataType(), param_tensor->Shape(),
param_tensor->MutableDataRaw(), param_tensor->Location()));
std::vector<Tensor> params, grads;
std::vector<std::vector<Tensor>> list_of_momentums;
list_of_momentums.resize(optimizer_algo_ptr_->momentum_keys.size());
// Collect gradients and prepare for tensorseq creation
auto* grad_tensor = parameter->Gradient().GetMutable<Tensor>();
grads.emplace_back(
Tensor(grad_tensor->DataType(), grad_tensor->Shape(),
grad_tensor->MutableDataRaw(), grad_tensor->Location()));
// Collect all the non-user-defined inputs from the named_parameters_.
for (auto& [parameter_name, parameter] : state_->module_checkpoint_state.named_parameters) {
if (parameter->RequiresGrad()) {
// Collect parameters and prepare for tensorseq creation
auto* param_tensor = parameter->Data().GetMutable<Tensor>();
params.emplace_back(
Tensor(param_tensor->DataType(), param_tensor->Shape(),
param_tensor->MutableDataRaw(), param_tensor->Location()));
// Collect first order moments and prepare for tensorseq creation
auto* first_order_moment_tensor = param_named_optimizer_states.at(parameter_name)
.momentum_named_states.at(MOMENT_STATE_NAMES[0])
.GetMutable<Tensor>();
first_order_moments.emplace_back(
Tensor(first_order_moment_tensor->DataType(), first_order_moment_tensor->Shape(),
first_order_moment_tensor->MutableDataRaw(), first_order_moment_tensor->Location()));
// Collect gradients and prepare for tensorseq creation
auto* grad_tensor = parameter->Gradient().GetMutable<Tensor>();
grads.emplace_back(
Tensor(grad_tensor->DataType(), grad_tensor->Shape(),
grad_tensor->MutableDataRaw(), grad_tensor->Location()));
// Collect second order moments and prepare for tensorseq creation
auto* second_order_moment_tensor = param_named_optimizer_states.at(parameter_name)
.momentum_named_states.at(MOMENT_STATE_NAMES[1])
.GetMutable<Tensor>();
second_order_moments.emplace_back(
Tensor(second_order_moment_tensor->DataType(), second_order_moment_tensor->Shape(),
second_order_moment_tensor->MutableDataRaw(), second_order_moment_tensor->Location()));
// Collect moments and prepare for tensorseq creation
for (size_t m_index = 0; m_index < optimizer_algo_ptr_->momentum_keys.size(); ++m_index) {
auto* moment_tensor =
param_named_optimizer_states.at(parameter_name)
.momentum_named_states.at(optimizer_algo_ptr_->momentum_keys[m_index])
.GetMutable<Tensor>();
list_of_momentums[m_index].emplace_back(
Tensor(moment_tensor->DataType(), moment_tensor->Shape(),
moment_tensor->MutableDataRaw(), moment_tensor->Location()));
}
}
const auto tensorseq_inserter = [](auto& tensors, auto* inputs) {
ORT_ENFORCE(!tensors.empty(), "Tensors vector cannot be empty while building a tensor sequence.");
auto tensor_seq = std::make_unique<TensorSeq>(tensors.front().DataType());
tensor_seq->Reserve(tensors.size());
for (auto& tensor : tensors) {
tensor_seq->Add(std::move(tensor));
}
inputs->emplace_back(
OrtValue(tensor_seq.release(), DataTypeImpl::GetType<TensorSeq>(),
DataTypeImpl::GetType<TensorSeq>()->GetDeleteFunc()));
};
// Add the params and moments as tensorseq ortvalues to inputs
tensorseq_inserter(params, &inputs_);
tensorseq_inserter(grads, &inputs_);
tensorseq_inserter(first_order_moments, &inputs_);
tensorseq_inserter(second_order_moments, &inputs_);
}
// Add other optimizer reordering logic here
const auto tensorseq_inserter = [](auto& tensors, auto* inputs) {
ORT_ENFORCE(!tensors.empty(), "Tensors vector cannot be empty while building a tensor sequence.");
auto tensor_seq = std::make_unique<TensorSeq>(tensors.front().DataType());
tensor_seq->Reserve(tensors.size());
for (auto& tensor : tensors) {
tensor_seq->Add(std::move(tensor));
}
inputs->emplace_back(
OrtValue(tensor_seq.release(), DataTypeImpl::GetType<TensorSeq>(),
DataTypeImpl::GetType<TensorSeq>()->GetDeleteFunc()));
};
// Add the params/grads as tensorseq ortvalues to inputs
tensorseq_inserter(params, &inputs_);
tensorseq_inserter(grads, &inputs_);
// Add all other momentums as tensorseq ortvalues to inputs.
for (auto& m : list_of_momentums) {
tensorseq_inserter(m, &inputs_);
}
return Status::OK();
}
} // namespace api
Optimizer::Optimizer(const std::string& optim_path_or_bytes,
CheckpointState* state,
@ -158,11 +189,28 @@ Optimizer::Optimizer(const std::string& optim_path_or_bytes,
const Environment& env,
const std::vector<std::shared_ptr<IExecutionProvider>>& providers)
: optim_sess_(std::make_unique<InferenceSession>(session_options, env)), state_(state) {
if (state_->optimizer_checkpoint_state.group_named_optimizer_states.empty()) {
state_->optimizer_checkpoint_state.group_named_optimizer_states.insert(
{GROUP_ZERO_NAME, std::make_shared<GroupOptimizerState>()});
Initialize(optim_path_or_bytes, session_options, env, providers);
ORT_ENFORCE(state != nullptr, "Checkpoint state cannot be null.");
auto g_it = state_->optimizer_checkpoint_state.group_named_optimizer_states.find(GROUP_ZERO_NAME);
bool find_group_zero = g_it != state_->optimizer_checkpoint_state.group_named_optimizer_states.end();
if (!find_group_zero || g_it->second->param_named_optimizer_states.empty()) {
if (!find_group_zero)
state_->optimizer_checkpoint_state.group_named_optimizer_states.insert(
{GROUP_ZERO_NAME, std::make_shared<GroupOptimizerState>()});
ORT_THROW_IF_ERROR(GenerateMomentumNamedStates(state_->optimizer_checkpoint_state));
ORT_THROW_IF_ERROR(ConstructInputs());
} else {
ORT_THROW_IF_ERROR(LoadStateDict(state_->optimizer_checkpoint_state));
}
optimizer_state_ = state_->optimizer_checkpoint_state.group_named_optimizer_states.at(GROUP_ZERO_NAME);
}
void Optimizer::Initialize(const std::string& optim_path_or_bytes,
const onnxruntime::SessionOptions& session_options,
const Environment& env,
const std::vector<std::shared_ptr<IExecutionProvider>>& providers) {
optim_sess_ = std::make_unique<InferenceSession>(session_options, env);
for (const auto& execution_provider : providers) {
ORT_THROW_IF_ERROR(optim_sess_->RegisterExecutionProvider(execution_provider));
}
@ -175,14 +223,17 @@ Optimizer::Optimizer(const std::string& optim_path_or_bytes,
utils::GetGraphInputOutputNames(optim_sess_, input_names_, output_names_);
if (optimizer_type_ == OptimizerType::AdamW) {
ORT_THROW_IF_ERROR(GraphInputsAreExpected(input_names_, AdamWOptimizerInputs));
optimizer_algo_ptr_ = OptimizerAlorithmFactory::CreateInstance(optim_path_or_bytes, group_count_);
ORT_ENFORCE(group_count_ == 1, "Group count can only be 1, but got: " + std::to_string(group_count_));
ORT_ENFORCE(optimizer_algo_ptr_, "optimizer_algo_ptr_ should not be nullptr.");
ORT_THROW_IF_ERROR(GenerateMomentumNamedStates());
} else {
ORT_THROW("Unsupported optimizer type");
}
ORT_THROW_IF_ERROR(ConstructInputs());
InlinedVector<std::string> all_input_names;
all_input_names.reserve(CommonOptimizerInputs.size() + optimizer_algo_ptr_->optimizer_states_inputs.size());
all_input_names.insert(all_input_names.end(), CommonOptimizerInputs.begin(),
CommonOptimizerInputs.end());
all_input_names.insert(all_input_names.end(), optimizer_algo_ptr_->optimizer_states_inputs.begin(),
optimizer_algo_ptr_->optimizer_states_inputs.end());
ORT_THROW_IF_ERROR(GraphInputsAreExpected(input_names_, all_input_names));
}
Status Optimizer::Step() {
@ -190,7 +241,7 @@ Status Optimizer::Step() {
utils::WrapInOrtValue<float>(optimizer_state_->learning_rate, &learning_rate_input);
// Use step count + 1 before running optimizer step.
// This is necessary since bias correction uses the step
// as a power. Using power of 0 is wrong.
// as a power. Using the power of 0 is wrong.
utils::WrapInOrtValue<int64_t>(optimizer_state_->step + 1, &step_input);
std::vector<OrtValue> feeds({learning_rate_input, step_input});
feeds.insert(feeds.end(), inputs_.begin(), inputs_.end());
@ -199,8 +250,8 @@ Status Optimizer::Step() {
auto status = optim_sess_->Run(RunOptions(), input_names_, feeds, output_names_, &outputs);
ORT_THROW_IF_ERROR(status);
// extract step output and update
if (utils::GetValue<int64_t>(outputs[0]) == 1LL) {
// Extract step output and update
if (utils::GetScalarFromOrtValue<int64_t>(outputs[0]) == 1LL) {
optimizer_state_->step++;
}
@ -210,7 +261,7 @@ Status Optimizer::Step() {
Status Optimizer::GetStateDict(OptimizerCheckpointState& optimizer_checkpoint_state) {
auto& grouped_optimizer_states = optimizer_checkpoint_state.group_named_optimizer_states;
// To support multiple groups, Optimizer constructor need accept informations for groupping.
// To support multiple groups, the Optimizer constructor needs to accept information for grouping.
grouped_optimizer_states.insert({GROUP_ZERO_NAME, std::make_shared<GroupOptimizerState>(*optimizer_state_)});
// Pass the optimizer session data transfer manager for data copying when saving.
@ -221,15 +272,59 @@ Status Optimizer::GetStateDict(OptimizerCheckpointState& optimizer_checkpoint_st
return Status::OK();
}
Status Optimizer::LoadStateDict(const OptimizerCheckpointState& optimizer_checkpoint_states) {
Status Optimizer::LoadStateDict(OptimizerCheckpointState& optimizer_checkpoint_states) {
auto group_optimizer_state_it =
optimizer_checkpoint_states.group_named_optimizer_states.find(GROUP_ZERO_NAME);
ORT_ENFORCE(group_optimizer_state_it != optimizer_checkpoint_states.group_named_optimizer_states.cend(),
ORT_ENFORCE(group_optimizer_state_it != optimizer_checkpoint_states.group_named_optimizer_states.end(),
"Group 0 not found in the optimizer checkpoint states.");
optimizer_state_->initial_lr = group_optimizer_state_it->second->initial_lr;
optimizer_state_->step = group_optimizer_state_it->second->step;
// TODO(pengwa): restore the momentums state from checkpoint.
optimizer_state_ = group_optimizer_state_it->second;
constexpr bool strict_match = true;
ORT_RETURN_IF_NOT(optim_sess_, "optimizer session not initialized");
auto& optim_sess_state = optim_sess_->GetSessionState();
auto& param_named_optimizer_states = optimizer_state_->param_named_optimizer_states;
for (auto& params_iter : state_->module_checkpoint_state.named_parameters) {
if (params_iter.second->RequiresGrad()) {
bool src_exist = param_named_optimizer_states.find(params_iter.first) !=
param_named_optimizer_states.cend();
ORT_ENFORCE(src_exist || !strict_match, "Parameter ", params_iter.first,
" not found in the source optimizer checkpoint states.");
std::unordered_map<std::string, OrtValue>& momentum_named_states =
param_named_optimizer_states.at(params_iter.first).momentum_named_states;
OrtValue& param_data = params_iter.second->Data();
ORT_ENFORCE(param_data.IsTensor());
const Tensor& param_data_tensor = param_data.Get<Tensor>();
const auto& param_data_device = param_data_tensor.Location().device;
auto target_allocator = optim_sess_state.GetAllocator(param_data_device);
ORT_ENFORCE(target_allocator != nullptr);
for (auto& momentum_state_pair : momentum_named_states) {
OrtValue& param_momentum = momentum_state_pair.second;
ORT_ENFORCE(param_momentum.IsTensor());
const Tensor& param_momentum_tensor = param_momentum.Get<Tensor>();
// If the source device type is already the same as the target device skip copy.
if (param_momentum_tensor.Location().device.Type() != param_data_device.Type()) {
// Create a new tensor on the target_device and switch the source_ortvalue to point to this new tensor
auto target_tensor = std::make_unique<Tensor>(param_momentum_tensor.DataType(),
param_momentum_tensor.Shape(),
target_allocator);
ORT_THROW_IF_ERROR(optim_sess_state.GetDataTransferMgr().CopyTensor(param_momentum_tensor,
*target_tensor.get()));
auto ml_tensor_type = DataTypeImpl::GetType<Tensor>();
param_momentum.Init(target_tensor.release(), ml_tensor_type, ml_tensor_type->GetDeleteFunc());
}
}
}
}
ORT_THROW_IF_ERROR(ConstructInputs());
return Status::OK();
}

View file

@ -2,6 +2,7 @@
// Licensed under the MIT License.
#pragma once
#include "core/providers/cpu/cpu_execution_provider.h"
#include "core/session/inference_session.h"
#include "core/session/environment.h"
@ -27,8 +28,11 @@ struct ParameterOptimizerState {
*/
struct GroupOptimizerState {
int64_t step = 0;
float initial_lr = 0.001f; // Default value used in torch AdamW
float learning_rate{initial_lr}; // Adaptive learning rate as training proceeds.
float initial_lr = 0.001f; // Default value used in torch AdamW
// Adaptive learning rate as training proceeds. Be noted, learning_rate can be
// restored by lr scheduler from given step and initial_lr, though, we still save/load this in checkpoint.
float learning_rate{initial_lr};
std::unordered_map<std::string, ParameterOptimizerState> param_named_optimizer_states;
};
@ -43,14 +47,50 @@ struct OptimizerCheckpointState {
const DataTransferManager* optimizer_session_data_transfer_mgr;
};
enum class OptimizerType {
AdamW,
// More optimizers can be added later as:
// Lamb,
struct OptimizerAlgorithmBase {
OptimizerAlgorithmBase(const std::vector<std::string>& momentum_keys,
const std::vector<std::string>& optimizer_states_inputs)
: momentum_keys(momentum_keys), optimizer_states_inputs(optimizer_states_inputs) {}
std::vector<std::string> momentum_keys;
std::vector<std::string> optimizer_states_inputs;
};
struct AdamWOptimizerAlgorithm : public OptimizerAlgorithmBase {
AdamWOptimizerAlgorithm() : OptimizerAlgorithmBase({"momentum0", "momentum1"},
{"first_order_moments", "second_order_moments"}) {}
};
struct SGDOptimizerV2Algorithm : public OptimizerAlgorithmBase {
SGDOptimizerV2Algorithm() : OptimizerAlgorithmBase({"momentum0"},
{"first_order_moments"}) {}
};
struct OptimizerAlorithmFactory {
static std::unique_ptr<OptimizerAlgorithmBase> CreateInstance(const std::string& optim_path_or_bytes,
int32_t& group_count);
};
struct CheckpointState;
/**
* @brief Optimizer class for running gradient updates.
*
* This class is responsible for running gradient updates on the parameters.
* > It does NOT own the parameters, and will not modify the "named_parameters" in `CheckpointState`
* passed from the constructor.
* A tensor sequence is created based on the "named_parameters" to construct parameter input (of type tensorseq).
* > If 'optimizer_checkpoint_states' is provided in the constructor as part of `CheckpointState`.
* Optimizer will reuse the data buffer from the passed in 'optimizer_checkpoint_states'.
* >> If the device of momentums in 'optimizer_checkpoint_states' is not
* matching its parameter device, a copy will be done during the 'LoadStateDict',
* but reserving using the original OrtValue with the copied data buffer;
* >> Otherwise, it generates the optimizer state initialized as all zeros and owns them.
* > If 'optimizer_checkpoint_states' is not provided in the constructor as part of `CheckpointState`.
* The optimizer states are initialized as all zeros on the same device of corresponding parameters.
*
* Currently, we only support load checkpoints from the constructor;
* no public API to load state dict after Optimizer instance is created.
*/
struct Optimizer {
friend struct LRSchedulerBase;
@ -66,10 +106,15 @@ struct Optimizer {
Status Step();
/**
* @brief Get the current optimizer state.
*
* Be noted the returned optimizer_checkpoint_states will hold new references to
* original momentum states.
* @return Status
*/
Status GetStateDict(OptimizerCheckpointState& optimizer_checkpoint_states);
Status LoadStateDict(const OptimizerCheckpointState& optimizer_checkpoint_states);
Status SetLearningRate(float lr) {
optimizer_state_->learning_rate = lr;
return Status::OK();
@ -86,24 +131,41 @@ struct Optimizer {
}
private:
void Initialize(const std::string& optim_path_or_bytes,
const onnxruntime::SessionOptions& session_options,
const Environment& env,
const std::vector<std::shared_ptr<IExecutionProvider>>& providers);
int64_t GetStep() const {
return optimizer_state_->step;
}
// Generates optimizer momentum states for applicable optimizer types
Status GenerateMomentumNamedStates();
// Generates optimizer momentum states for parameters that require grad.
Status GenerateMomentumNamedStates(OptimizerCheckpointState& optimizer_checkpoint_states);
// Constructs the ortvalue inputs to be fed to the graph
// at each step
// at each step.
Status ConstructInputs();
// TODO: load this info from checkpoint
OptimizerType optimizer_type_ = OptimizerType::AdamW;
/**
* @brief Load states from optimizer_checkpoint_states into current optimizer state.
*
* Be noted Optimizer will reuse the data buffer of passed in optimizer_checkpoint_states.
* If the device of momentums in optimizer_checkpoint_states is not matching its parameter device,
* an implicit copy will be done during the LoadStateDict, but reserving using the original OrtValue
* with the copied data buffer.
* @return Status
*/
Status LoadStateDict(OptimizerCheckpointState& optimizer_checkpoint_states);
std::unique_ptr<OptimizerAlgorithmBase> optimizer_algo_ptr_;
std::unique_ptr<onnxruntime::InferenceSession> optim_sess_;
CheckpointState* state_; // Non owning pointer to the state.
std::shared_ptr<GroupOptimizerState> optimizer_state_;
std::vector<std::string> input_names_;
std::vector<std::string> output_names_;
std::vector<OrtValue> inputs_;
int32_t group_count_{0};
};
} // namespace api

View file

@ -57,7 +57,7 @@ bool GetParamNameFromGradient(const std::string& grad_name, std::string& param_n
return false;
}
Status OrtValueLike(const SessionState& sess_state, const OrtValue& input_val, OrtValue& output_val) {
Status CreateZeroValuedOrtValueLike(const SessionState& sess_state, const OrtValue& input_val, OrtValue& output_val) {
const auto& param_tensor = input_val.template Get<Tensor>();
const TensorShape& shape = param_tensor.Shape();
auto& tensor_location = param_tensor.Location();

View file

@ -25,7 +25,7 @@ bool GetParamNameFromSuffix(const std::string& name, const std::string& suffix,
bool GetParamNameFromGradient(const std::string& grad_name, std::string& param_name);
// Allocate OrtValue like the input ortvalue on the same device
Status OrtValueLike(const SessionState& sess_state, const OrtValue& input_val, OrtValue& output_val);
Status CreateZeroValuedOrtValueLike(const SessionState& sess_state, const OrtValue& input_val, OrtValue& output_val);
// Create OrtValue from a single value of type T
template <typename T>
@ -48,8 +48,12 @@ void WrapInOrtValue(T value,
}
template <typename T>
T GetValue(OrtValue& ort_value) {
T GetScalarFromOrtValue(OrtValue& ort_value) {
const Tensor& tensor = ort_value.Get<Tensor>();
const TensorShape& shape = tensor.Shape();
size_t dim_count = shape.NumDimensions();
// Be noted: TensorShape returns 1 for rank 0 tensor.
ORT_ENFORCE(shape.Size() == 1 && (dim_count == 0 || dim_count == 1));
T val;
if (DataTypeImpl::GetType<T>() == tensor.DataType()) {
val = *(tensor.template Data<T>());