diff --git a/orttraining/orttraining/test/training_api/core/checkpoint_test.cc b/orttraining/orttraining/test/training_api/core/checkpoint_test.cc index 4fca303b26..8edad98472 100644 --- a/orttraining/orttraining/test/training_api/core/checkpoint_test.cc +++ b/orttraining/orttraining/test/training_api/core/checkpoint_test.cc @@ -321,7 +321,7 @@ TEST(CheckpointApiTest, SaveOptimizerStateAsCheckpoint_ThenLoad_CUDA) { ASSERT_EQ(param_tensor.DataType(), restored_tensor.DataType()); std::vector 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); } diff --git a/orttraining/orttraining/test/training_api/core/data_utils.h b/orttraining/orttraining/test/training_api/core/data_utils.h index 815fbd1b8a..6e9fe5db73 100644 --- a/orttraining/orttraining/test/training_api/core/data_utils.h +++ b/orttraining/orttraining/test/training_api/core/data_utils.h @@ -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 -void OrtValueToVec(const OrtValue& val, std::vector& output) { - const Tensor& tensor = val.Get(); +void CpuOrtValueToVec(const OrtValue& src_cpu_ortvalue, std::vector& output) { + const Tensor& tensor = src_cpu_ortvalue.Get(); int64_t num_elem = tensor.Shape().Size(); const T* val_ptr = tensor.template Data(); output.assign(val_ptr, val_ptr + num_elem); } template -void CudaOrtValueToCpuVec(const OrtValue& val, std::vector& output, - std::shared_ptr cuda_provider, - std::shared_ptr cpu_provider) { - const Tensor& src_tensor = val.Get(); +void CudaOrtValueToCpuVec(const OrtValue& src_cuda_ortvalue, std::vector& output) { + std::unique_ptr cuda_provider = onnxruntime::test::DefaultCudaExecutionProvider(); + std::unique_ptr cpu_provider = onnxruntime::test::DefaultCpuExecutionProvider(); + + const Tensor& src_tensor = src_cuda_ortvalue.Get(); auto allocator = cpu_provider->GetAllocator(OrtMemTypeDefault); ORT_ENFORCE(allocator, "Cpu allocator is a nullptr."); diff --git a/orttraining/orttraining/test/training_api/core/training_api_tests.cc b/orttraining/orttraining/test/training_api/core/training_api_tests.cc index 1487c2b29d..e08687eb18 100644 --- a/orttraining/orttraining/test/training_api/core/training_api_tests.cc +++ b/orttraining/orttraining/test/training_api/core/training_api_tests.cc @@ -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>& named_parameters, + const std::vector& 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& 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(); + 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>& 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>& 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& 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 env; - ASSERT_STATUS_OK(Environment::Create(nullptr, env)); - const std::vector> providers{onnxruntime::test::DefaultCudaExecutionProvider()}; - auto model = std::make_unique( - ToUTF8String(model_uri), &state, - session_option, *env, providers); - auto optim = std::make_shared( - ToUTF8String(optim_uri), &state, session_option, - *env, providers); - - OrtValue input, target; - GenerateRandomInput(std::array{2, 784}, input); - onnxruntime::test::CreateInputOrtValueOnCPU( - std::array{2}, std::vector(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 > - typedef std::vector>> TestDataDictType; - TestDataDictType test_data; - const json j = json::parse(in); - j.get_to(test_data); - - int64_t resume_step = (*test_data.begin()).first; - ASSERT_EQ(total_step_count, static_cast(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(); - 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( - 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 inputs{input, target}; - std::vector 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& inputs = *data_loader.begin(); std::vector 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 run_cuda_list{false}; + // #ifdef USE_CUDA + // run_cuda_list.push_back(true); + // #endif + + for (auto run_cuda : run_cuda_list) { + std::vector> 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 env; + + ASSERT_STATUS_OK(Environment::Create(nullptr, env)); + + std::shared_ptr model = std::make_shared( + 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 optim = std::make_shared( + 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 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 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& test_file_name, + float initial_lr, + int64_t total_step_count, + int64_t warmup_step_count) { + std::vector run_cuda_list{false}; + // #ifdef USE_CUDA + // run_cuda_list.push_back(true); + // #endif + + for (auto run_cuda : run_cuda_list) { + std::vector> 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 env; + + ASSERT_STATUS_OK(Environment::Create(nullptr, env)); + + std::shared_ptr model = std::make_shared( + ToUTF8String(model_uri), &state, session_option, + *env, providers); + + OrtValue input, target; + GenerateRandomInput(std::array{2, 784}, input); + onnxruntime::test::CreateInputOrtValueOnCPU( + std::array{2}, std::vector(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 > + typedef std::vector>> TestDataDictType; + TestDataDictType test_data; + const json j = json::parse(in); + j.get_to(test_data); + + int64_t resume_step = (*test_data.begin()).first; + ASSERT_EQ(total_step_count, static_cast(test_data.size()) + resume_step); + + if (resume_step != 0) { + state.optimizer_checkpoint_state.group_named_optimizer_states.insert( + {"group0", std::make_shared()}); + 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 optim = std::make_shared( + 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( + 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 inputs{input, target}; + std::vector 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> providers{onnxruntime::test::DefaultCpuExecutionProvider()}; TestModuleExport(providers); @@ -316,7 +475,6 @@ TEST(TrainingApiTest, OptimStep) { std::unique_ptr env; std::vector> providers{onnxruntime::test::DefaultCudaExecutionProvider()}; std::shared_ptr cuda_provider = providers.front(); - std::shared_ptr cpu_provider = onnxruntime::test::DefaultCpuExecutionProvider(); ASSERT_STATUS_OK(Environment::Create(nullptr, env)); auto model = std::make_unique( 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 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 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 fetches; ASSERT_STATUS_OK(model->TrainStep(inputs, fetches)); std::vector 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 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 diff --git a/orttraining/orttraining/training_api/module.cc b/orttraining/orttraining/training_api/module.cc index 0987e87dd8..a5bbd1a6c4 100644 --- a/orttraining/orttraining/training_api/module.cc +++ b/orttraining/orttraining/training_api/module.cc @@ -157,8 +157,8 @@ Module::Module(const std::string& train_model_path_or_bytes, const std::optional& 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 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(); - // 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(); diff --git a/orttraining/orttraining/training_api/module.h b/orttraining/orttraining/training_api/module.h index 71f7659e8e..4cbddc9656 100644 --- a/orttraining/orttraining/training_api/module.h +++ b/orttraining/orttraining/training_api/module.h @@ -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 diff --git a/orttraining/orttraining/training_api/optimizer.cc b/orttraining/orttraining/training_api/optimizer.cc index 2008e4c7c3..7e0d21b4e5 100644 --- a/orttraining/orttraining/training_api/optimizer.cc +++ b/orttraining/orttraining/training_api/optimizer.cc @@ -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 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 actual_graph_inputs, - gsl::span expected_graph_inputs) { + gsl::span expected_graph_inputs) { const auto stringify = [](const auto& container) { if (container.empty()) { return std::string("[]"); @@ -71,86 +59,129 @@ Status GraphInputsAreExpected(gsl::span actual_graph_inputs, } // namespace -Status Optimizer::GenerateMomentumNamedStates() { +std::unique_ptr OptimizerAlorithmFactory::CreateInstance( + const std::string& optim_path_or_bytes, int32_t& group_count) { + std::shared_ptr model; + ORT_ENFORCE(Model::Load(ToWideString(optim_path_or_bytes), model, nullptr, + logging::LoggingManager::DefaultLogger()) + .IsOK()); + Graph& graph = model->MainGraph(); + std::map, 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(); + } else if (domain == kMSDomain && type == "SGDOptimizerV2") { + return std::make_unique(); + } 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 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(); - params.emplace_back( - Tensor(param_tensor->DataType(), param_tensor->Shape(), - param_tensor->MutableDataRaw(), param_tensor->Location())); + std::vector params, grads; + std::vector> 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(); - 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(); + 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(); - 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(); + 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(); - 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(); + 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(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(), - DataTypeImpl::GetType()->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(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(), + DataTypeImpl::GetType()->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>& providers) : optim_sess_(std::make_unique(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()}); + 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()}); + 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>& providers) { + optim_sess_ = std::make_unique(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 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(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(optimizer_state_->step + 1, &step_input); std::vector 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(outputs[0]) == 1LL) { + // Extract step output and update + if (utils::GetScalarFromOrtValue(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(*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& 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(); + 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(); + + // 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(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(); + param_momentum.Init(target_tensor.release(), ml_tensor_type, ml_tensor_type->GetDeleteFunc()); + } + } + } + } + + ORT_THROW_IF_ERROR(ConstructInputs()); + return Status::OK(); } diff --git a/orttraining/orttraining/training_api/optimizer.h b/orttraining/orttraining/training_api/optimizer.h index 4c28a600ba..a6769172e3 100644 --- a/orttraining/orttraining/training_api/optimizer.h +++ b/orttraining/orttraining/training_api/optimizer.h @@ -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 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& momentum_keys, + const std::vector& optimizer_states_inputs) + : momentum_keys(momentum_keys), optimizer_states_inputs(optimizer_states_inputs) {} + std::vector momentum_keys; + std::vector 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 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>& 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 optimizer_algo_ptr_; std::unique_ptr optim_sess_; CheckpointState* state_; // Non owning pointer to the state. std::shared_ptr optimizer_state_; std::vector input_names_; std::vector output_names_; std::vector inputs_; + + int32_t group_count_{0}; }; } // namespace api diff --git a/orttraining/orttraining/training_api/utils.cc b/orttraining/orttraining/training_api/utils.cc index e719c48bea..aa2c0173a7 100644 --- a/orttraining/orttraining/training_api/utils.cc +++ b/orttraining/orttraining/training_api/utils.cc @@ -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(); const TensorShape& shape = param_tensor.Shape(); auto& tensor_location = param_tensor.Location(); diff --git a/orttraining/orttraining/training_api/utils.h b/orttraining/orttraining/training_api/utils.h index 27391e596c..37ae4b6747 100644 --- a/orttraining/orttraining/training_api/utils.h +++ b/orttraining/orttraining/training_api/utils.h @@ -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 @@ -48,8 +48,12 @@ void WrapInOrtValue(T value, } template -T GetValue(OrtValue& ort_value) { +T GetScalarFromOrtValue(OrtValue& ort_value) { const Tensor& tensor = ort_value.Get(); + 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() == tensor.DataType()) { val = *(tensor.template Data());