mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Cumulative update on optimizers and tests (on-device training) (#15499)
This commit is contained in:
parent
8a1a40ac63
commit
29d13cea42
9 changed files with 550 additions and 258 deletions
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.");
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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>());
|
||||
|
|
|
|||
Loading…
Reference in a new issue