mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
CKPT API Implementation (On Device Training) (#11261)
* Checkpoint API Implementation * fix build issues * fix undefined reference for ParseData of type string. * refinements * resolve some comments * expose python api * make save and load test pass * some clean up * make optimizer save/load test pass * make custom property save/load test pass * formatting * fix comments - fix wave - code placement, remove legacy ckpt logic dependency, remove external data support * fix comment - wave 2 - Remove ParseData/ParseStringData, Use UnpackTensor, Simplify CheckpointProperty usage * fix comment - wave 3 - rename all api_test namespace to api * fix comment - wave 4 - load/save trainable/nontrainable param seperately. * Rename Load/SaveORTCheckpoint * renaming API && remove CheckpointUntils. api::LoadCheckpoint/SaveCheckpoint is the exposed interfaces. * revert unnecessary format change for onnxruntime/core/framework/tensorprotoutils.h/cc * formatting * re-org the class folders for better dependency managerment * save_checkpoint accpeting TensorProto as inputs * More clean up * clean up the naming * refactor a bit type constraints on custom property * fix comment - file read/write && report error when file read/write failed * extract LoopDir to FilterFilesFromDirectory * fix build
This commit is contained in:
parent
8ee8fdd59b
commit
d8a1531c37
18 changed files with 1838 additions and 356 deletions
|
|
@ -17,6 +17,15 @@ file(GLOB_RECURSE onnxruntime_training_srcs
|
|||
"${ORTTRAINING_SOURCE_DIR}/core/agent/*.cc"
|
||||
)
|
||||
|
||||
if (onnxruntime_ENABLE_TRAINING_ON_DEVICE)
|
||||
file(GLOB_RECURSE onnxruntime_training_api_srcs CONFIGURE_DEPENDS
|
||||
"${ORTTRAINING_SOURCE_DIR}/training_api/*.h"
|
||||
"${ORTTRAINING_SOURCE_DIR}/training_api/*.cc"
|
||||
)
|
||||
|
||||
list(APPEND onnxruntime_training_srcs ${onnxruntime_training_api_srcs})
|
||||
endif()
|
||||
|
||||
# This needs to be built in framework.cmake
|
||||
file(GLOB_RECURSE onnxruntime_training_framework_excluded_srcs CONFIGURE_DEPENDS
|
||||
"${ORTTRAINING_SOURCE_DIR}/core/framework/torch/*.h"
|
||||
|
|
@ -233,7 +242,8 @@ if (onnxruntime_BUILD_UNIT_TESTS)
|
|||
# Training API Tests
|
||||
# Currently disable it by default for internal development usage.
|
||||
if (onnxruntime_ENABLE_TRAINING_ON_DEVICE)
|
||||
file(GLOB_RECURSE training_api_test_runner_src
|
||||
# Only files in the direct folder will be compiled into test runner.
|
||||
file(GLOB training_api_test_runner_src
|
||||
"${ORTTRAINING_SOURCE_DIR}/test/training_api/*.h"
|
||||
"${ORTTRAINING_SOURCE_DIR}/test/training_api/*.cc"
|
||||
)
|
||||
|
|
@ -245,10 +255,28 @@ if (onnxruntime_BUILD_UNIT_TESTS)
|
|||
endif()
|
||||
endif()
|
||||
|
||||
onnxruntime_add_include_to_target(onnxruntime_training_api_test_runner onnxruntime_common onnx onnx_proto ${PROTOBUF_LIB} onnxruntime_training flatbuffers)
|
||||
target_include_directories(onnxruntime_training_api_test_runner PUBLIC ${CMAKE_CURRENT_BINARY_DIR} ${ONNXRUNTIME_ROOT} ${ORTTRAINING_ROOT} ${MPI_CXX_INCLUDE_DIRS} ${eigen_INCLUDE_DIRS} ${CXXOPTS} ${extra_includes} ${onnxruntime_graph_header} ${onnxruntime_exec_src_dir} ${CMAKE_CURRENT_BINARY_DIR} ${CMAKE_CURRENT_BINARY_DIR}/onnx)
|
||||
onnxruntime_add_include_to_target(onnxruntime_training_api_test_runner onnxruntime_training
|
||||
onnxruntime_framework onnxruntime_common onnx onnx_proto ${PROTOBUF_LIB} flatbuffers)
|
||||
|
||||
target_include_directories(onnxruntime_training_api_test_runner PUBLIC
|
||||
${CMAKE_CURRENT_BINARY_DIR}
|
||||
${ONNXRUNTIME_ROOT}
|
||||
${ORTTRAINING_ROOT}
|
||||
${MPI_CXX_INCLUDE_DIRS}
|
||||
${eigen_INCLUDE_DIRS}
|
||||
${CXXOPTS}
|
||||
${extra_includes}
|
||||
${onnxruntime_graph_header}
|
||||
${onnxruntime_exec_src_dir}
|
||||
${CMAKE_CURRENT_BINARY_DIR}
|
||||
${CMAKE_CURRENT_BINARY_DIR}/onnx
|
||||
)
|
||||
|
||||
target_link_libraries(onnxruntime_training_api_test_runner PRIVATE onnxruntime_training ${ONNXRUNTIME_LIBS} ${onnxruntime_EXTERNAL_LIBRARIES})
|
||||
target_link_libraries(onnxruntime_training_api_test_runner PRIVATE
|
||||
onnxruntime_training
|
||||
${ONNXRUNTIME_LIBS}
|
||||
${onnxruntime_EXTERNAL_LIBRARIES}
|
||||
)
|
||||
set_target_properties(onnxruntime_training_api_test_runner PROPERTIES FOLDER "ONNXRuntimeTest")
|
||||
endif()
|
||||
|
||||
|
|
|
|||
|
|
@ -252,6 +252,11 @@ file(GLOB onnxruntime_test_training_src
|
|||
"${ORTTRAINING_SOURCE_DIR}/test/distributed/*.cc"
|
||||
)
|
||||
|
||||
if (onnxruntime_ENABLE_TRAINING_ON_DEVICE)
|
||||
file(GLOB onnxruntime_test_training_on_device_src
|
||||
"${ORTTRAINING_SOURCE_DIR}/test/training_api/checkpoint/*.cc")
|
||||
endif()
|
||||
|
||||
if(WIN32)
|
||||
list(APPEND onnxruntime_test_framework_src_patterns
|
||||
"${TEST_SRC_DIR}/platform/windows/*.cc"
|
||||
|
|
@ -645,6 +650,9 @@ set(all_dependencies ${onnxruntime_test_providers_dependencies} )
|
|||
|
||||
if (onnxruntime_ENABLE_TRAINING)
|
||||
list(APPEND all_tests ${onnxruntime_test_training_src})
|
||||
if (onnxruntime_ENABLE_TRAINING_ON_DEVICE)
|
||||
list(APPEND all_tests ${onnxruntime_test_training_on_device_src})
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if (onnxruntime_USE_NUPHAR)
|
||||
|
|
|
|||
46
orttraining/orttraining/core/framework/checkpoint_common.cc
Normal file
46
orttraining/orttraining/core/framework/checkpoint_common.cc
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "core/common/logging/logging.h"
|
||||
#include "core/common/logging/sinks/clog_sink.h"
|
||||
#include "core/common/status.h"
|
||||
#include "core/framework/data_types.h"
|
||||
#include "core/framework/framework_common.h"
|
||||
#include "core/framework/tensorprotoutils.h"
|
||||
#include "core/platform/env.h"
|
||||
#include "core/platform/path_lib.h"
|
||||
#include "core/providers/cpu/cpu_execution_provider.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
|
||||
/**
|
||||
* @brief Create OrtValues From TensorProto objects
|
||||
*
|
||||
* @param tensor_protos vector of TensorProto
|
||||
* @param name_to_ort_value saved results.
|
||||
* @return Status
|
||||
*/
|
||||
Status CreateOrtValuesFromTensorProtos(
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& tensor_protos,
|
||||
NameMLValMap& name_to_ort_value) {
|
||||
static CPUExecutionProviderInfo info;
|
||||
static CPUExecutionProvider cpu_provider(info);
|
||||
static AllocatorPtr cpu_allocator = cpu_provider.GetAllocator(0, OrtMemTypeDefault);
|
||||
|
||||
for (const auto& tensor_proto : tensor_protos) {
|
||||
TensorShape tensor_shape{utils::GetTensorShapeFromTensorProto(tensor_proto)};
|
||||
const DataTypeImpl* tensor_dtype = DataTypeImpl::TensorTypeFromONNXEnum(tensor_proto.data_type())->GetElementType();
|
||||
auto p_tensor = std::make_unique<Tensor>(tensor_dtype, tensor_shape, cpu_allocator);
|
||||
ORT_RETURN_IF_ERROR(utils::TensorProtoToTensor(Env::Default(), nullptr, tensor_proto, *p_tensor));
|
||||
|
||||
OrtValue ort_value;
|
||||
ort_value.Init(p_tensor.release(), DataTypeImpl::GetType<Tensor>(), DataTypeImpl::GetType<Tensor>()->GetDeleteFunc());
|
||||
name_to_ort_value.emplace(tensor_proto.name(), ort_value);
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
59
orttraining/orttraining/core/framework/checkpoint_common.h
Normal file
59
orttraining/orttraining/core/framework/checkpoint_common.h
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "core/common/logging/logging.h"
|
||||
#include "core/common/logging/sinks/clog_sink.h"
|
||||
#include "core/common/path.h"
|
||||
#include "core/common/path_string.h"
|
||||
#include "core/common/status.h"
|
||||
#include "core/framework/framework_common.h"
|
||||
#include "core/platform/env.h"
|
||||
#include "core/platform/path_lib.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
|
||||
/**
|
||||
* @brief Open file descriptor and call use_fn
|
||||
*
|
||||
* @tparam TUseFileFn
|
||||
* @param path file path
|
||||
* @param readonly open mode.
|
||||
* @param use_fn function taking file descriptor as inputs.
|
||||
* @return common::Status
|
||||
*/
|
||||
template <typename TUseFileFn>
|
||||
common::Status WithOpenFile(const PathString& path, bool readonly, TUseFileFn use_fn) {
|
||||
int fd;
|
||||
if (readonly) {
|
||||
ORT_RETURN_IF_ERROR(Env::Default().FileOpenRd(path, fd));
|
||||
} else {
|
||||
ORT_RETURN_IF_ERROR(Env::Default().FileOpenWr(path, fd));
|
||||
}
|
||||
|
||||
Status use_fn_status{};
|
||||
try {
|
||||
use_fn_status = use_fn(fd);
|
||||
} catch (std::exception& e) {
|
||||
use_fn_status = ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, e.what());
|
||||
}
|
||||
|
||||
Status close_status = Env::Default().FileClose(fd);
|
||||
return !use_fn_status.IsOK() ? use_fn_status : close_status;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Create OrtValues From TensorProto objects
|
||||
*
|
||||
* @param tensor_protos vector of TensorProto
|
||||
* @param name_to_ort_value saved results.
|
||||
* @return Status
|
||||
*/
|
||||
Status CreateOrtValuesFromTensorProtos(
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& tensor_protos,
|
||||
NameMLValMap& name_to_ort_value);
|
||||
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -20,6 +20,7 @@
|
|||
#include "core/platform/path_lib.h"
|
||||
#include "orttraining/core/framework/protobuf_message_sequence.h"
|
||||
#include "core/util/protobuf_parsing_utils.h"
|
||||
#include "orttraining/core/framework/checkpoint_common.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
|
|
@ -93,28 +94,6 @@ Status SaveRuntimeTensor(
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
// opens file descriptor and calls use_fn
|
||||
// use_fn should have this signature: Status use_fn(int file_descriptor)
|
||||
template <typename TUseFileFn>
|
||||
Status WithOpenFile(const PathString& path, bool readonly, TUseFileFn use_fn) {
|
||||
int fd;
|
||||
if (readonly) {
|
||||
ORT_RETURN_IF_ERROR(Env::Default().FileOpenRd(path, fd));
|
||||
} else {
|
||||
ORT_RETURN_IF_ERROR(Env::Default().FileOpenWr(path, fd));
|
||||
}
|
||||
|
||||
Status use_fn_status{};
|
||||
try {
|
||||
use_fn_status = use_fn(fd);
|
||||
} catch (std::exception& e) {
|
||||
use_fn_status = ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, e.what());
|
||||
}
|
||||
|
||||
Status close_status = Env::Default().FileClose(fd);
|
||||
return !use_fn_status.IsOK() ? use_fn_status : close_status;
|
||||
}
|
||||
|
||||
std::vector<std::string> GetOrderedOrtValueNames(const NameMLValMap& name_to_value) {
|
||||
std::vector<std::string> ordered_names{};
|
||||
ordered_names.reserve(name_to_value.size());
|
||||
|
|
|
|||
|
|
@ -30,6 +30,11 @@
|
|||
#include "orttraining/core/framework/torch/custom_function_register.h"
|
||||
#endif
|
||||
|
||||
#if defined(ENABLE_TRAINING) && defined(ENABLE_TRAINING_ON_DEVICE)
|
||||
#include "orttraining/training_api/include/checkpoint.h"
|
||||
#include <google/protobuf/io/zero_copy_stream_impl.h>
|
||||
#endif
|
||||
|
||||
PYBIND11_MAKE_OPAQUE(std::vector<OrtValue>);
|
||||
PYBIND11_MAKE_OPAQUE(onnxruntime::OrtValueCache);
|
||||
|
||||
|
|
@ -142,7 +147,7 @@ struct PyGradientGraphBuilder {
|
|||
// TODO: this method does not handle parallel optimization.
|
||||
TrainingConfigurationResult ConfigureSessionForTraining(
|
||||
training::PipelineTrainingSession* sess, TrainingParameters& parameters) {
|
||||
//TODO tix, refactor the mpi related code to populate all fields correctly by default.
|
||||
// TODO tix, refactor the mpi related code to populate all fields correctly by default.
|
||||
ORT_ENFORCE(parameters.data_parallel_size <= parameters.world_size, "data_parallel_size: ", parameters.data_parallel_size, ", world_size: ", parameters.world_size);
|
||||
ORT_ENFORCE(parameters.horizontal_parallel_size <= parameters.world_size, "horizontal_parallel_size: ", parameters.horizontal_parallel_size, ", world_size: ", parameters.world_size);
|
||||
ORT_ENFORCE(parameters.pipeline_parallel_size <= parameters.world_size, "pipeline_parallel_size: ", parameters.pipeline_parallel_size, ", world_size: ", parameters.world_size);
|
||||
|
|
@ -370,84 +375,90 @@ void addObjectMethodsForTraining(py::module& m, ExecutionProviderRegistrationFn
|
|||
.def("reserve", [](std::vector<OrtValue>* v, const size_t len) { v->reserve(len); })
|
||||
.def("shrink_to_fit", [](std::vector<OrtValue>* v) { v->shrink_to_fit(); })
|
||||
.def("__len__", [](const std::vector<OrtValue>& v) { return v.size(); })
|
||||
.def("__iter__", [](const std::vector<OrtValue>& v) {
|
||||
return py::make_iterator(v.cbegin(), v.cend());
|
||||
}, py::keep_alive<0, 1>())
|
||||
.def(
|
||||
"__iter__", [](const std::vector<OrtValue>& v) {
|
||||
return py::make_iterator(v.cbegin(), v.cend());
|
||||
},
|
||||
py::keep_alive<0, 1>())
|
||||
.def("__getitem__", [](const std::vector<OrtValue>& v, const size_t idx) {
|
||||
return v.at(idx);
|
||||
})
|
||||
.def("bool_tensor_indices", [](std::vector<OrtValue>* v) -> std::vector<int64_t> {
|
||||
std::vector<int64_t> indices;
|
||||
for (size_t i = 0; i < v->size(); ++i) {
|
||||
if (GetTensorProtoType((*v)[i]) == ONNX_NAMESPACE::TensorProto_DataType_BOOL) {
|
||||
indices.push_back(static_cast<int64_t>(i));
|
||||
}
|
||||
}
|
||||
return indices;
|
||||
}, "Returns the indices of every boolean tensor in this vector of OrtValue. "
|
||||
"In case of a boolean tensor, method to_dlpacks returns a uint8 tensor instead of a boolean tensor. "
|
||||
"If torch consumes the dlpack structure, `.to(torch.bool)` must be applied to the torch tensor "
|
||||
"to get a boolean tensor.")
|
||||
.def(
|
||||
"bool_tensor_indices", [](std::vector<OrtValue>* v) -> std::vector<int64_t> {
|
||||
std::vector<int64_t> indices;
|
||||
for (size_t i = 0; i < v->size(); ++i) {
|
||||
if (GetTensorProtoType((*v)[i]) == ONNX_NAMESPACE::TensorProto_DataType_BOOL) {
|
||||
indices.push_back(static_cast<int64_t>(i));
|
||||
}
|
||||
}
|
||||
return indices;
|
||||
},
|
||||
"Returns the indices of every boolean tensor in this vector of OrtValue. "
|
||||
"In case of a boolean tensor, method to_dlpacks returns a uint8 tensor instead of a boolean tensor. "
|
||||
"If torch consumes the dlpack structure, `.to(torch.bool)` must be applied to the torch tensor "
|
||||
"to get a boolean tensor.")
|
||||
.def("dlpack_at", [](std::vector<OrtValue>* v, const size_t idx) {
|
||||
return py::reinterpret_steal<py::object>(ToDlpack(v->at(idx)));
|
||||
})
|
||||
.def("element_type_at", [](std::vector<OrtValue>* v, const size_t idx) -> int32_t {
|
||||
return GetTensorProtoType(v->at(idx));
|
||||
}, "Returns an integer equal to the ONNX proto type of the tensor at position i. "
|
||||
"This integer is one type defined by ONNX TensorProto_DataType "
|
||||
"(such as onnx.TensorProto.FLOAT)."
|
||||
.def(
|
||||
"element_type_at", [](std::vector<OrtValue>* v, const size_t idx) -> int32_t {
|
||||
return GetTensorProtoType(v->at(idx));
|
||||
},
|
||||
"Returns an integer equal to the ONNX proto type of the tensor at position i. "
|
||||
"This integer is one type defined by ONNX TensorProto_DataType "
|
||||
"(such as onnx.TensorProto.FLOAT)."
|
||||
"Raises an exception in any other case.")
|
||||
.def("to_dlpacks", [](const std::vector<OrtValue>& v, py::object to_tensor) -> py::list {
|
||||
.def(
|
||||
"to_dlpacks", [](const std::vector<OrtValue>& v, py::object to_tensor) -> py::list {
|
||||
if (v.size() == 0)
|
||||
return py::list();
|
||||
|
||||
if (v.size() == 0)
|
||||
return py::list();
|
||||
py::list list_dlpacks;
|
||||
PyObject* obj;
|
||||
|
||||
py::list list_dlpacks;
|
||||
PyObject* obj;
|
||||
py::gil_scoped_acquire acquire;
|
||||
|
||||
py::gil_scoped_acquire acquire;
|
||||
if (to_tensor.is_none()) {
|
||||
DLManagedTensor* dlmanaged_tensor;
|
||||
|
||||
if (to_tensor.is_none()) {
|
||||
DLManagedTensor* dlmanaged_tensor;
|
||||
|
||||
for (auto it : v) {
|
||||
dlmanaged_tensor = dlpack::OrtValueToDlpack(it);
|
||||
py::capsule capsule(dlmanaged_tensor, "dltensor", DlpackCapsuleDestructor);
|
||||
list_dlpacks.append(capsule);
|
||||
}
|
||||
} else {
|
||||
DLManagedTensor* dlmanaged_tensor;
|
||||
PyObject* capsule = NULL;
|
||||
PyObject* handle = to_tensor.ptr();
|
||||
|
||||
for (auto it : v) {
|
||||
// A new instance of dlpack needs to be created. The object which consumes it
|
||||
// is responsible for its deletion.
|
||||
dlmanaged_tensor = dlpack::OrtValueToDlpack(it);
|
||||
if (capsule == NULL) {
|
||||
capsule = PyCapsule_New(dlmanaged_tensor, "dltensor", NULL);
|
||||
if (capsule == NULL)
|
||||
throw std::runtime_error("Unexpected error: empty capsule returned.");
|
||||
for (auto it : v) {
|
||||
dlmanaged_tensor = dlpack::OrtValueToDlpack(it);
|
||||
py::capsule capsule(dlmanaged_tensor, "dltensor", DlpackCapsuleDestructor);
|
||||
list_dlpacks.append(capsule);
|
||||
}
|
||||
} else {
|
||||
// The same capsule is reused but FromDLPack rename the capsule into used_dltensor.
|
||||
PyCapsule_SetName(capsule, "dltensor");
|
||||
PyCapsule_SetPointer(capsule, dlmanaged_tensor);
|
||||
DLManagedTensor* dlmanaged_tensor;
|
||||
PyObject* capsule = NULL;
|
||||
PyObject* handle = to_tensor.ptr();
|
||||
|
||||
for (auto it : v) {
|
||||
// A new instance of dlpack needs to be created. The object which consumes it
|
||||
// is responsible for its deletion.
|
||||
dlmanaged_tensor = dlpack::OrtValueToDlpack(it);
|
||||
if (capsule == NULL) {
|
||||
capsule = PyCapsule_New(dlmanaged_tensor, "dltensor", NULL);
|
||||
if (capsule == NULL)
|
||||
throw std::runtime_error("Unexpected error: empty capsule returned.");
|
||||
} else {
|
||||
// The same capsule is reused but FromDLPack rename the capsule into used_dltensor.
|
||||
PyCapsule_SetName(capsule, "dltensor");
|
||||
PyCapsule_SetPointer(capsule, dlmanaged_tensor);
|
||||
}
|
||||
obj = PyObject_CallFunctionObjArgs(handle, capsule, NULL);
|
||||
if (obj == NULL)
|
||||
throw std::runtime_error("to_tensor returned a null pointer. This is usually caused by an error during the conversion.");
|
||||
list_dlpacks.append(obj);
|
||||
Py_DECREF(obj);
|
||||
}
|
||||
if (capsule != NULL) {
|
||||
// This test is never wrong because v is not empty if the execution goes through that path.
|
||||
// If not present, Guardian detects a potential failure.
|
||||
Py_DECREF(capsule);
|
||||
}
|
||||
}
|
||||
obj = PyObject_CallFunctionObjArgs(handle, capsule, NULL);
|
||||
if (obj == NULL)
|
||||
throw std::runtime_error("to_tensor returned a null pointer. This is usually caused by an error during the conversion.");
|
||||
list_dlpacks.append(obj);
|
||||
Py_DECREF(obj);
|
||||
}
|
||||
if (capsule != NULL) {
|
||||
// This test is never wrong because v is not empty if the execution goes through that path.
|
||||
// If not present, Guardian detects a potential failure.
|
||||
Py_DECREF(capsule);
|
||||
}
|
||||
}
|
||||
return list_dlpacks;
|
||||
},
|
||||
R"pbdoc(Converts all OrtValue into tensors through DLPack protocol, the method creates
|
||||
return list_dlpacks;
|
||||
},
|
||||
R"pbdoc(Converts all OrtValue into tensors through DLPack protocol, the method creates
|
||||
a DLPack structure for every tensors, then calls python function `to_tensor` to a new object
|
||||
consuming the DLPack structure or return a list of capsule if this function is None.
|
||||
|
||||
|
|
@ -675,7 +686,7 @@ for every transfered tensor.
|
|||
NameMLValMap state_tensors;
|
||||
ORT_THROW_IF_ERROR(static_cast<PipelineTrainingSession*>(sess->GetSessionHandle())->GetStateTensors(state_tensors));
|
||||
auto& data_transfer_manager = sess->GetSessionHandle()->GetDataTransferManager();
|
||||
//convert to numpy array
|
||||
// convert to numpy array
|
||||
std::map<std::string, py::object> rmap;
|
||||
for (auto& kv : state_tensors) {
|
||||
if (kv.second.IsTensor()) {
|
||||
|
|
@ -866,28 +877,28 @@ for every transfered tensor.
|
|||
const std::unordered_set<std::string>& y_node_arg_names,
|
||||
const std::unordered_set<std::string>& x_node_arg_names,
|
||||
const std::string loss_node_arg_name) {
|
||||
std::shared_ptr<Model> model;
|
||||
auto logger_ptr = std::make_unique<logging::Logger>(logging::LoggingManager::DefaultLogger());
|
||||
logger_ptr->SetSeverity(logging::Severity::kINFO);
|
||||
ONNX_NAMESPACE::ModelProto model_proto;
|
||||
std::istringstream model_istream(serialized_model);
|
||||
ORT_THROW_IF_ERROR(Model::Load(model_istream, &model_proto));
|
||||
ORT_THROW_IF_ERROR(Model::Load(model_proto, model, nullptr, *logger_ptr));
|
||||
GradientGraphConfiguration gradient_graph_config{};
|
||||
gradient_graph_config.set_gradients_as_graph_outputs = true;
|
||||
// Save some objects, otherwise they get lost.
|
||||
auto gradient_graph_config_ptr = std::make_unique<GradientGraphConfiguration>(gradient_graph_config);
|
||||
std::shared_ptr<Model> model;
|
||||
auto logger_ptr = std::make_unique<logging::Logger>(logging::LoggingManager::DefaultLogger());
|
||||
logger_ptr->SetSeverity(logging::Severity::kINFO);
|
||||
ONNX_NAMESPACE::ModelProto model_proto;
|
||||
std::istringstream model_istream(serialized_model);
|
||||
ORT_THROW_IF_ERROR(Model::Load(model_istream, &model_proto));
|
||||
ORT_THROW_IF_ERROR(Model::Load(model_proto, model, nullptr, *logger_ptr));
|
||||
GradientGraphConfiguration gradient_graph_config{};
|
||||
gradient_graph_config.set_gradients_as_graph_outputs = true;
|
||||
// Save some objects, otherwise they get lost.
|
||||
auto gradient_graph_config_ptr = std::make_unique<GradientGraphConfiguration>(gradient_graph_config);
|
||||
|
||||
auto builder = std::make_unique<GradientGraphBuilder>(
|
||||
&model->MainGraph(),
|
||||
y_node_arg_names,
|
||||
x_node_arg_names,
|
||||
loss_node_arg_name,
|
||||
*gradient_graph_config_ptr,
|
||||
*logger_ptr);
|
||||
auto builder = std::make_unique<GradientGraphBuilder>(
|
||||
&model->MainGraph(),
|
||||
y_node_arg_names,
|
||||
x_node_arg_names,
|
||||
loss_node_arg_name,
|
||||
*gradient_graph_config_ptr,
|
||||
*logger_ptr);
|
||||
|
||||
return std::make_unique<PyGradientGraphBuilder>(std::move(builder), std::move(model), std::move(logger_ptr), std::move(gradient_graph_config_ptr));
|
||||
}))
|
||||
return std::make_unique<PyGradientGraphBuilder>(std::move(builder), std::move(model), std::move(logger_ptr), std::move(gradient_graph_config_ptr));
|
||||
}))
|
||||
.def("build", [](PyGradientGraphBuilder* gradient_graph_builder) {
|
||||
ORT_THROW_IF_ERROR(gradient_graph_builder->builder->Build());
|
||||
})
|
||||
|
|
@ -928,6 +939,32 @@ for every transfered tensor.
|
|||
[](const std::string& key, const std::unordered_set<size_t> edges) -> void {
|
||||
GradientDefinitionRegistry::Instance().SetStopGradientEdgesForNode(key, edges);
|
||||
});
|
||||
|
||||
#if defined(ENABLE_TRAINING) && defined(ENABLE_TRAINING_ON_DEVICE)
|
||||
m.def("save_checkpoint",
|
||||
[](const std::vector<py::bytes>& trainable_tensor_protos_pybytes,
|
||||
const std::vector<py::bytes>& non_trainable_tensor_protos_pybytes,
|
||||
const std::string& checkpoint_path) {
|
||||
std::vector<TensorProto> trainable_tensor_protos(trainable_tensor_protos_pybytes.size());
|
||||
std::vector<TensorProto> non_trainable_tensor_protos(non_trainable_tensor_protos_pybytes.size());
|
||||
|
||||
auto parse_pybytes_to_tensor_proto =
|
||||
[](const std::vector<py::bytes>& tensor_protos_pybytes, std::vector<TensorProto>& tensor_protos) {
|
||||
for (size_t i = 0; i < tensor_protos_pybytes.size(); ++i) {
|
||||
std::istringstream tensor_proto_istream(tensor_protos_pybytes[i]);
|
||||
ORT_ENFORCE(tensor_proto_istream.good(), "Broken tensor proto istream to read.");
|
||||
google::protobuf::io::IstreamInputStream zero_copy_input(&tensor_proto_istream);
|
||||
const bool result = tensor_protos[i].ParseFromZeroCopyStream(&zero_copy_input) && tensor_proto_istream.eof();
|
||||
ORT_ENFORCE(result, "Parse tensor proto failed.");
|
||||
}
|
||||
};
|
||||
|
||||
parse_pybytes_to_tensor_proto(trainable_tensor_protos_pybytes, trainable_tensor_protos);
|
||||
parse_pybytes_to_tensor_proto(non_trainable_tensor_protos_pybytes, non_trainable_tensor_protos);
|
||||
|
||||
ORT_THROW_IF_ERROR(onnxruntime::training::api::SaveCheckpoint(trainable_tensor_protos, non_trainable_tensor_protos, checkpoint_path));
|
||||
});
|
||||
#endif
|
||||
}
|
||||
|
||||
} // namespace python
|
||||
|
|
|
|||
|
|
@ -0,0 +1,362 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
|
||||
#include "gtest/gtest.h"
|
||||
|
||||
#include "core/common/common.h"
|
||||
#include "core/common/logging/logging.h"
|
||||
#include "core/common/logging/sinks/clog_sink.h"
|
||||
#include "core/framework/framework_common.h"
|
||||
#include "core/framework/data_transfer.h"
|
||||
#include "core/framework/ort_value.h"
|
||||
#include "core/framework/tensor.h"
|
||||
#include "core/framework/tensorprotoutils.h"
|
||||
#include "core/graph/graph_viewer.h"
|
||||
#include "core/graph/model.h"
|
||||
#include "core/platform/path_lib.h"
|
||||
|
||||
#include "orttraining/core/framework/checkpoint_common.h"
|
||||
#include "orttraining/training_api/include/interfaces.h"
|
||||
|
||||
#include "test/test_environment.h"
|
||||
#include "test/util/include/asserts.h"
|
||||
#include "test/util/include/temp_dir.h"
|
||||
#include "test/util/include/test/test_environment.h"
|
||||
|
||||
using onnxruntime::test::TemporaryDirectory;
|
||||
using namespace onnxruntime::training::api;
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace test {
|
||||
namespace training_api {
|
||||
|
||||
#define MODEL_FOLDER ORT_TSTR("testdata/")
|
||||
|
||||
/**
|
||||
* Load ONNX model from file path, save into ORT checkpoint files,
|
||||
* Then load it into ORT, compare with the initial parameter values.
|
||||
*/
|
||||
TEST(CheckpointApiTest, SaveOnnxModelAsCheckpoint_ThenLoad_CPU) {
|
||||
/// Phase 1 - Test Preparison
|
||||
/// Prepare the data and dest folder for saving checkpoint.
|
||||
/// Also cooked the data for test result comparision.
|
||||
|
||||
// Model path and trainable parameter name definitions.
|
||||
auto model_uri = MODEL_FOLDER "transform/computation_reduction/e2e.onnx";
|
||||
std::vector<std::string> expected_trainable_param_names{
|
||||
"bert.encoder.layer.2.output.LayerNorm.weight",
|
||||
"bert.encoder.layer.2.output.LayerNorm.bias",
|
||||
"add1_initializerr",
|
||||
"cls.predictions.transform.LayerNorm.weight",
|
||||
"cls.predictions.transform.LayerNorm.bias",
|
||||
"bert.embeddings.word_embeddings.weight_transposed",
|
||||
"cls.predictions.bias",
|
||||
};
|
||||
|
||||
// Extract a weight value baseline to compare.
|
||||
// expected_trainable_param_name_to_ort_value is used to compare with the values after restoring from checkpoint.
|
||||
auto logger_ptr = std::make_unique<logging::Logger>(logging::LoggingManager::DefaultLogger());
|
||||
std::shared_ptr<Model> p_model;
|
||||
ORT_ENFORCE(Model::Load(model_uri, p_model, nullptr, *logger_ptr).IsOK());
|
||||
Graph& graph = p_model->MainGraph();
|
||||
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> trainable_param_values;
|
||||
trainable_param_values.reserve(expected_trainable_param_names.size());
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> non_trainable_param_values;
|
||||
const auto& initializer_tensors = graph.GetAllInitializedTensors();
|
||||
for (const std::pair<std::string, const ONNX_NAMESPACE::TensorProto*>& pair : initializer_tensors) {
|
||||
if (std::find(expected_trainable_param_names.begin(), expected_trainable_param_names.end(), pair.first) != expected_trainable_param_names.end()) {
|
||||
trainable_param_values.emplace_back(static_cast<ONNX_NAMESPACE::TensorProto>(*pair.second));
|
||||
} else {
|
||||
non_trainable_param_values.emplace_back(static_cast<ONNX_NAMESPACE::TensorProto>(*pair.second));
|
||||
}
|
||||
}
|
||||
|
||||
std::unordered_map<std::string, OrtValue> expected_trainable_param_name_to_ort_value;
|
||||
ORT_ENFORCE(CreateOrtValuesFromTensorProtos(trainable_param_values, expected_trainable_param_name_to_ort_value).IsOK());
|
||||
|
||||
// Remove the tempoprary directory if it already exists.
|
||||
auto ckpt_test_root_dir = ORT_TSTR("checkpointing_api_test_dir");
|
||||
if (Env::Default().FolderExists(ckpt_test_root_dir)) {
|
||||
ORT_ENFORCE(Env::Default().DeleteFolder(ckpt_test_root_dir).IsOK());
|
||||
}
|
||||
TemporaryDirectory tmp_dir{ckpt_test_root_dir};
|
||||
|
||||
/// Phase 2 - Run save checkpoint APIs.
|
||||
/// And check the result checkpoint files.
|
||||
|
||||
// Call Save APIs.
|
||||
PathString checkpoint_path{
|
||||
ConcatPathComponent<PathChar>(tmp_dir.Path(), ORT_TSTR("e2e_ckpt_save_cpu"))};
|
||||
ASSERT_STATUS_OK(SaveCheckpoint(trainable_param_values, non_trainable_param_values, checkpoint_path));
|
||||
|
||||
// Check the ckpt files in the directory.
|
||||
std::set<PathString> expected_file_names{"paramfrozen_tensors.pbseq", "paramtrain_tensors.pbseq"};
|
||||
std::set<PathString> valid_file_names;
|
||||
LoopDir(checkpoint_path,
|
||||
[&valid_file_names, &checkpoint_path](const PathChar* filename, OrtFileType file_type) -> bool {
|
||||
PathString filename_str = filename;
|
||||
bool is_valid_ckpt_file_exts = HasExtensionOf(filename_str, ORT_TSTR("pbseq"));
|
||||
if (filename_str[0] == '.' || file_type == OrtFileType::TYPE_DIR || !is_valid_ckpt_file_exts) {
|
||||
return true;
|
||||
}
|
||||
valid_file_names.emplace(filename_str);
|
||||
return true;
|
||||
});
|
||||
|
||||
ASSERT_EQ(expected_file_names, valid_file_names);
|
||||
|
||||
/// Phase 3 - Run load checkpoint APIs.
|
||||
/// And check the result comparible with initial parameter values.
|
||||
|
||||
// Call Load APIs
|
||||
CheckpointState checkpoint_state_to_load;
|
||||
ASSERT_STATUS_OK(LoadCheckpoint(checkpoint_path, checkpoint_state_to_load));
|
||||
ModuleCheckpointState module_state = checkpoint_state_to_load.module_checkpoint_state;
|
||||
const auto& param_states = module_state.named_parameters;
|
||||
std::unordered_map<std::string, OrtValue> restored_param_name_to_ort_values;
|
||||
std::vector<std::string> restored_trainable_param_names;
|
||||
for (auto it = param_states.begin(); it != param_states.end(); ++it) {
|
||||
restored_param_name_to_ort_values.insert({it->first, it->second->Data()});
|
||||
if (it->second->RequiresGrad()) {
|
||||
restored_trainable_param_names.emplace_back(it->first);
|
||||
}
|
||||
}
|
||||
|
||||
// Check loaded parameter's values are same with original ones.
|
||||
ASSERT_EQ(expected_trainable_param_name_to_ort_value.size(), restored_trainable_param_names.size());
|
||||
ASSERT_EQ(expected_trainable_param_name_to_ort_value.size(), 7);
|
||||
ASSERT_EQ(restored_param_name_to_ort_values.size(), 9);
|
||||
|
||||
std::sort(expected_trainable_param_names.begin(), expected_trainable_param_names.end());
|
||||
std::sort(restored_trainable_param_names.begin(), restored_trainable_param_names.end());
|
||||
ASSERT_EQ(expected_trainable_param_names, restored_trainable_param_names);
|
||||
|
||||
for (const auto& name : restored_trainable_param_names) {
|
||||
const auto& restored_ort_value = restored_param_name_to_ort_values[name];
|
||||
const auto& expected_ort_value = expected_trainable_param_name_to_ort_value.at(name);
|
||||
|
||||
ASSERT_TRUE(restored_ort_value.IsTensor() && expected_ort_value.IsTensor());
|
||||
const Tensor& restored_tensor = restored_ort_value.Get<Tensor>();
|
||||
const Tensor& expected_tensor = expected_ort_value.Get<Tensor>();
|
||||
ASSERT_EQ(expected_tensor.DataType(), restored_tensor.DataType());
|
||||
ASSERT_EQ(expected_tensor.SizeInBytes(), restored_tensor.SizeInBytes());
|
||||
ASSERT_EQ(expected_tensor.DataType(), restored_tensor.DataType());
|
||||
|
||||
ASSERT_TRUE(std::memcmp(expected_tensor.DataRaw(), restored_tensor.DataRaw(), expected_tensor.SizeInBytes()) == 0);
|
||||
}
|
||||
}
|
||||
|
||||
const OrtMemoryInfo cpu_alloc_info(onnxruntime::CPU, OrtDeviceAllocator);
|
||||
class OrtValueTensorData {
|
||||
public:
|
||||
OrtValueTensorData(TensorShape shape, std::vector<float> data) {
|
||||
ORT_ENFORCE(shape.Size() == static_cast<int64_t>(data.size()));
|
||||
shape_ = std::move(shape);
|
||||
data_ = std::move(data);
|
||||
}
|
||||
|
||||
OrtValue GetOrtValue() {
|
||||
return OrtValue(new Tensor(DataTypeImpl::GetType<float>(), shape_, data_.data(), cpu_alloc_info),
|
||||
DataTypeImpl::GetType<Tensor>(), DataTypeImpl::GetType<Tensor>()->GetDeleteFunc());
|
||||
}
|
||||
|
||||
private:
|
||||
TensorShape shape_;
|
||||
std::vector<float> data_;
|
||||
};
|
||||
|
||||
/**
|
||||
* Create Optimizer with sets of parameters,
|
||||
* Save Optimizer states into ORT checkpoint files,
|
||||
* Then load it into ORT, compare with the initial optimizer states values.
|
||||
*/
|
||||
TEST(CheckpointApiTest, SaveOptimizerStateAsCheckpoint_ThenLoad_CPU) {
|
||||
/// Phase 1 - Test Preparison
|
||||
/// Prepare the data and dest folder for saving checkpoint.
|
||||
/// Also cooked the data for test result comparision.
|
||||
|
||||
auto model_uri = MODEL_FOLDER "transform/computation_reduction/e2e.onnx";
|
||||
std::unordered_map<std::string, OrtValueTensorData> name_to_ort_value_data{
|
||||
{"param1", {{3}, {1.0f, 2.0f, 3.0f}}},
|
||||
{"param2", {{2, 2}, {1.0f, 2.0f, 3.0f, 4.0f}}},
|
||||
{"param3", {{3}, {1.0f, 2.0f, 3.0f}}},
|
||||
{"param4", {{2, 2}, {1.0f, 2.0f, 3.0f, 4.0f}}},
|
||||
};
|
||||
|
||||
std::vector<std::string> trainable_param_names{"param1", "param4"};
|
||||
NameMLValMap name_to_ort_value{};
|
||||
for (auto& name_and_ort_value_data : name_to_ort_value_data) {
|
||||
name_to_ort_value.emplace(
|
||||
name_and_ort_value_data.first, name_and_ort_value_data.second.GetOrtValue());
|
||||
}
|
||||
|
||||
// Optimizer creation and trainable parameter name definitions.
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>> named_parameters;
|
||||
for (auto it = name_to_ort_value.begin(); it != name_to_ort_value.end(); ++it) {
|
||||
auto param = std::make_shared<Parameter>(it->first, it->second);
|
||||
bool is_trainable =
|
||||
std::find(trainable_param_names.begin(), trainable_param_names.end(), param->Name()) != trainable_param_names.end();
|
||||
ASSERT_STATUS_OK(param->SetRequiresGrad(is_trainable));
|
||||
named_parameters.insert({it->first, param});
|
||||
}
|
||||
auto optimizer = Optimizer(model_uri, named_parameters);
|
||||
|
||||
/// Phase 2 - Run Optimizer.GetStateDict and call save checkpoint APIs.
|
||||
/// And check the result checkpoint files.
|
||||
|
||||
CheckpointState checkpoint_state;
|
||||
ORT_ENFORCE(optimizer.GetStateDict(checkpoint_state.optimizer_checkpoint_state).IsOK());
|
||||
|
||||
// Remove the tempoprary directory if it already exists.
|
||||
auto ckpt_test_root_dir = ORT_TSTR("checkpointing_api_test_dir");
|
||||
if (Env::Default().FolderExists(ckpt_test_root_dir)) {
|
||||
ORT_ENFORCE(Env::Default().DeleteFolder(ckpt_test_root_dir).IsOK());
|
||||
}
|
||||
TemporaryDirectory tmp_dir{ckpt_test_root_dir};
|
||||
|
||||
// Call Save APIs.
|
||||
PathString checkpoint_path{
|
||||
ConcatPathComponent<PathChar>(tmp_dir.Path(), ORT_TSTR("e2e_ckpt_save_cpu"))};
|
||||
ASSERT_STATUS_OK(SaveCheckpoint(checkpoint_state, checkpoint_path));
|
||||
|
||||
// Check the ckpt files in the directory.
|
||||
std::set<PathString> expected_file_names{
|
||||
"optim_group0_momentum0_tensors.pbseq",
|
||||
"optim_group0_momentum1_tensors.pbseq",
|
||||
"optim_group0_properties.pbseq",
|
||||
};
|
||||
|
||||
std::set<PathString> valid_file_names;
|
||||
LoopDir(checkpoint_path,
|
||||
[&valid_file_names, &checkpoint_path](const PathChar* filename, OrtFileType file_type) -> bool {
|
||||
PathString filename_str = filename;
|
||||
bool is_valid_ckpt_file_exts =
|
||||
HasExtensionOf(filename_str, ORT_TSTR("pbseq")) || HasExtensionOf(filename_str, ORT_TSTR("bin"));
|
||||
if (filename_str[0] == '.' || file_type == OrtFileType::TYPE_DIR || !is_valid_ckpt_file_exts) {
|
||||
return true;
|
||||
}
|
||||
valid_file_names.emplace(filename_str);
|
||||
return true;
|
||||
});
|
||||
|
||||
ASSERT_EQ(expected_file_names, valid_file_names);
|
||||
|
||||
/// Phase 3 - Run load checkpoint APIs.
|
||||
/// And check the result comparible with initial optimizer state values.
|
||||
|
||||
// Call Load APIs
|
||||
CheckpointState checkpoint_state_to_load;
|
||||
ASSERT_STATUS_OK(LoadCheckpoint(checkpoint_path, checkpoint_state_to_load));
|
||||
OptimizerCheckpointState optimizer_state = checkpoint_state_to_load.optimizer_checkpoint_state;
|
||||
std::unordered_map<std::string, std::shared_ptr<GroupOptimizerState>>&
|
||||
group_optimizer_states = optimizer_state.group_named_optimizer_states;
|
||||
|
||||
ASSERT_EQ(group_optimizer_states.size(), 1);
|
||||
ASSERT_EQ(group_optimizer_states.begin()->first, "group0");
|
||||
|
||||
std::unordered_map<std::string, ParameterOptimizerState>&
|
||||
param_named_optimizer_states = group_optimizer_states["group0"]->param_named_optimizer_states;
|
||||
|
||||
ASSERT_EQ(param_named_optimizer_states.size(), 2);
|
||||
auto it = param_named_optimizer_states.begin();
|
||||
ASSERT_EQ(it->first, "param1");
|
||||
std::advance(it, 1);
|
||||
ASSERT_EQ(it->first, "param4");
|
||||
|
||||
for (auto it = param_named_optimizer_states.begin(); it != param_named_optimizer_states.end(); ++it) {
|
||||
for (auto& state_pair : it->second.momentum_named_states) {
|
||||
ASSERT_TRUE(state_pair.first == "momentum0" || state_pair.first == "momentum1");
|
||||
const OrtValue& restored_ort_value = *(state_pair.second);
|
||||
const OrtValue& expected_ort_value = name_to_ort_value[it->first];
|
||||
ASSERT_TRUE(restored_ort_value.IsTensor() && expected_ort_value.IsTensor());
|
||||
const Tensor& restored_tensor = restored_ort_value.Get<Tensor>();
|
||||
const Tensor& expected_tensor = expected_ort_value.Get<Tensor>();
|
||||
|
||||
ASSERT_EQ(expected_tensor.DataType(), restored_tensor.DataType());
|
||||
ASSERT_EQ(expected_tensor.SizeInBytes(), restored_tensor.SizeInBytes());
|
||||
ASSERT_EQ(expected_tensor.DataType(), restored_tensor.DataType());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Create PropertyBag with sets of properties,
|
||||
* Save properties into ORT checkpoint files,
|
||||
* Then load it into ORT, compare with the initial properties' values.
|
||||
*/
|
||||
TEST(CheckpointApiTest, SaveCustomPropertyAsCheckpoint_ThenLoad_CPU) {
|
||||
/// Phase 1 - Test Preparison
|
||||
/// Prepare the data and dest folder for saving checkpoint.
|
||||
|
||||
CheckpointState checkpoint_state;
|
||||
PropertyBag& property_bag = checkpoint_state.property_bag;
|
||||
|
||||
float f_data = 0.5f;
|
||||
std::string f_property_name("float_number");
|
||||
property_bag.AddProperty<float>(f_property_name, f_data);
|
||||
|
||||
int64_t i_data = 400;
|
||||
std::string i_property_name("dataset_epoch_index");
|
||||
property_bag.AddProperty<int64_t>(i_property_name, i_data);
|
||||
|
||||
std::string s_data("/data/path/train.bin");
|
||||
std::string s_property_name("train_data_path");
|
||||
property_bag.AddProperty<std::string>(s_property_name, s_data);
|
||||
|
||||
// Remove the tempoprary directory if it already exists.
|
||||
auto ckpt_test_root_dir = ORT_TSTR("checkpointing_api_test_dir");
|
||||
if (Env::Default().FolderExists(ckpt_test_root_dir)) {
|
||||
ORT_ENFORCE(Env::Default().DeleteFolder(ckpt_test_root_dir).IsOK());
|
||||
}
|
||||
TemporaryDirectory tmp_dir{ckpt_test_root_dir};
|
||||
|
||||
/// Phase 2 - Call save checkpoint APIs.
|
||||
/// And check the result checkpoint files.
|
||||
|
||||
// Call Save APIs.
|
||||
PathString checkpoint_path{
|
||||
ConcatPathComponent<PathChar>(tmp_dir.Path(), ORT_TSTR("e2e_ckpt_save_cpu"))};
|
||||
ASSERT_STATUS_OK(SaveCheckpoint(checkpoint_state, checkpoint_path));
|
||||
|
||||
// Check the ckpt files in the directory.
|
||||
std::set<PathString> expected_file_names{
|
||||
"custom_properties.pbseq",
|
||||
};
|
||||
|
||||
std::set<PathString> valid_file_names;
|
||||
LoopDir(checkpoint_path,
|
||||
[&valid_file_names, &checkpoint_path](const PathChar* filename, OrtFileType file_type) -> bool {
|
||||
PathString filename_str = filename;
|
||||
bool is_valid_ckpt_file_exts =
|
||||
HasExtensionOf(filename_str, ORT_TSTR("pbseq")) || HasExtensionOf(filename_str, ORT_TSTR("bin"));
|
||||
if (filename_str[0] == '.' || file_type == OrtFileType::TYPE_DIR || !is_valid_ckpt_file_exts) {
|
||||
return true;
|
||||
}
|
||||
valid_file_names.emplace(filename_str);
|
||||
return true;
|
||||
});
|
||||
|
||||
ASSERT_EQ(expected_file_names, valid_file_names);
|
||||
|
||||
// Call Load APIs
|
||||
CheckpointState checkpoint_state_to_load;
|
||||
ASSERT_STATUS_OK(LoadCheckpoint(checkpoint_path, checkpoint_state_to_load));
|
||||
PropertyBag& restored_property_bag = checkpoint_state_to_load.property_bag;
|
||||
ASSERT_EQ(restored_property_bag.Size(), 3);
|
||||
float restored_f_data = restored_property_bag.GetProperty<float>(f_property_name);
|
||||
ASSERT_FLOAT_EQ(f_data, restored_f_data);
|
||||
int64_t restored_i_data = restored_property_bag.GetProperty<int64_t>(i_property_name);
|
||||
ASSERT_EQ(i_data, restored_i_data);
|
||||
std::string restored_s_data = restored_property_bag.GetProperty<std::string>(s_property_name);
|
||||
ASSERT_EQ(s_data, restored_s_data);
|
||||
}
|
||||
|
||||
} // namespace training_api
|
||||
} // namespace test
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -11,13 +11,15 @@
|
|||
#include "core/session/inference_session.h"
|
||||
#include "core/providers/cpu/cpu_provider_factory_creator.h"
|
||||
#include "orttraining/core/framework/tensorboard/event_writer.h"
|
||||
#include "orttraining/training_api/interfaces.h"
|
||||
|
||||
// ORT training C++ API includes
|
||||
#include "orttraining/training_api/include/interfaces.h"
|
||||
|
||||
using namespace onnxruntime;
|
||||
using namespace onnxruntime::common;
|
||||
using namespace onnxruntime::training;
|
||||
using namespace onnxruntime::training::tensorboard;
|
||||
using namespace onnxruntime::training::api_test;
|
||||
using namespace onnxruntime::training::api;
|
||||
using namespace std;
|
||||
|
||||
#ifdef USE_CUDA
|
||||
|
|
@ -201,18 +203,18 @@ Status RunTraining(const TestRunnerParameters& params) {
|
|||
std::string tensorboard_file = params.output_dir + "/tb.event";
|
||||
std::shared_ptr<EventWriter> tensorboard = std::make_shared<EventWriter>(tensorboard_file);
|
||||
|
||||
api_test::utils::CheckpointStates state_dicts;
|
||||
ORT_ENFORCE(api_test::utils::Ort_Load(params.checkpoint_to_load_path, state_dicts).IsOK());
|
||||
CheckpointState state;
|
||||
ORT_ENFORCE(LoadCheckpoint(params.checkpoint_to_load_path, state).IsOK());
|
||||
|
||||
Module module(params.model_training_graph_path,
|
||||
state_dicts.named_parameters,
|
||||
state.module_checkpoint_state.named_parameters,
|
||||
params.model_evaluation_graph_path);
|
||||
|
||||
Optimizer optimizer(params.optimizer_training_graph_path,
|
||||
state_dicts.named_parameters);
|
||||
state.module_checkpoint_state.named_parameters);
|
||||
|
||||
#ifdef USE_CUDA
|
||||
api_test::utils::SetExecutionProvider(module, optimizer, params.provider.get());
|
||||
api::SetExecutionProvider(module, optimizer, params.provider.get());
|
||||
#endif
|
||||
|
||||
auto scheduler = std::make_unique<LinearScheduler>(optimizer, 0.3333f, 1.0f, 5);
|
||||
|
|
@ -251,11 +253,12 @@ Status RunTraining(const TestRunnerParameters& params) {
|
|||
|
||||
if (batch_idx % SAVE_STEPS == 0) {
|
||||
// save trained weights
|
||||
api_test::utils::CheckpointStates state_dicts_to_save;
|
||||
ORT_ENFORCE(module.GetStateDict(state_dicts_to_save.named_parameters).IsOK());
|
||||
ORT_ENFORCE(optimizer.GetStateDict(state_dicts_to_save.optimizer_states).IsOK());
|
||||
CheckpointState state_to_save;
|
||||
ORT_ENFORCE(module.GetStateDict(state_to_save.module_checkpoint_state).IsOK());
|
||||
ORT_ENFORCE(optimizer.GetStateDict(state_to_save.optimizer_checkpoint_state).IsOK());
|
||||
state_to_save.property_bag.AddProperty<int64_t>(std::string("epoch"), static_cast<int64_t>(epoch));
|
||||
std::string ckpt_file = params.output_dir + "/ckpt_" + params.model_name + std::to_string(batch_idx);
|
||||
ORT_ENFORCE(api_test::utils::Ort_Save(state_dicts_to_save, ckpt_file).IsOK());
|
||||
ORT_ENFORCE(SaveCheckpoint(state_to_save, ckpt_file).IsOK());
|
||||
}
|
||||
|
||||
batch_idx++;
|
||||
|
|
|
|||
541
orttraining/orttraining/training_api/checkpoint.cc
Normal file
541
orttraining/orttraining/training_api/checkpoint.cc
Normal file
|
|
@ -0,0 +1,541 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "core/common/logging/logging.h"
|
||||
#include "core/common/logging/sinks/clog_sink.h"
|
||||
#include "core/common/path.h"
|
||||
#include "core/framework/framework_common.h"
|
||||
#include "core/framework/tensorprotoutils.h"
|
||||
#include "core/graph/graph_viewer.h"
|
||||
#include "core/graph/model.h"
|
||||
#include "core/platform/env.h"
|
||||
#include "core/platform/path_lib.h"
|
||||
#include "core/util/protobuf_parsing_utils.h"
|
||||
|
||||
#include "orttraining/core/framework/checkpoint_common.h"
|
||||
#include "orttraining/core/framework/protobuf_message_sequence.h"
|
||||
#include "orttraining/training_api/include/checkpoint.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
namespace {
|
||||
|
||||
constexpr const char* k_tensor_proto_file_name = "tensors.pbseq";
|
||||
constexpr const char* k_tensor_proto_properties_file_name = "properties.pbseq";
|
||||
constexpr const char* k_trainable_param_root_prefix = "paramtrain";
|
||||
constexpr const char* k_non_trainable_param_root_prefix = "paramfrozen";
|
||||
constexpr const char* k_optimizer_root_prefix = "optim";
|
||||
constexpr const char* k_property_root_prefix = "custom";
|
||||
constexpr const char* k_name_seperator = "_";
|
||||
|
||||
const std::string builtin_lr_property_name("builtin.learning_rate");
|
||||
const std::string builtin_step_property_name("builtin.step");
|
||||
|
||||
/**
|
||||
* @brief Create TensorProtos From OrtValue objects
|
||||
*
|
||||
* @param name_to_ort_value name to OrtValue mapping.
|
||||
* @param data_transfer_manager data transfer manager to copy the tensor in OrtValue.
|
||||
* @param saved_tensor_protos saved results.
|
||||
* @return Status
|
||||
*/
|
||||
Status CreateTensorProtosFromOrtValues(
|
||||
const NameMLValMap& name_to_ort_value,
|
||||
const DataTransferManager& data_transfer_manager,
|
||||
std::vector<ONNX_NAMESPACE::TensorProto>& saved_tensor_protos) {
|
||||
// Order the tensors by name.
|
||||
std::vector<std::string> ordered_tensor_names{};
|
||||
ordered_tensor_names.reserve(name_to_ort_value.size());
|
||||
std::transform(name_to_ort_value.begin(), name_to_ort_value.end(), std::back_inserter(ordered_tensor_names),
|
||||
[](const NameMLValMap::value_type& v) { return v.first; });
|
||||
std::sort(ordered_tensor_names.begin(), ordered_tensor_names.end());
|
||||
|
||||
// Copy the tensor data and create TensorProto storing the data.
|
||||
std::vector<char> tensor_data_buffer{};
|
||||
static const OrtMemoryInfo cpu_alloc_info{onnxruntime::CPU, OrtDeviceAllocator};
|
||||
|
||||
saved_tensor_protos.reserve(ordered_tensor_names.size());
|
||||
|
||||
unsigned long total_bytes = 0;
|
||||
constexpr unsigned long PROTOBUF_UPPER_LIMIT = 2 * 1000 * 1000 * 1000;
|
||||
for (const auto& tensor_name : ordered_tensor_names) {
|
||||
const OrtValue& ort_value = name_to_ort_value.at(tensor_name);
|
||||
ORT_RETURN_IF_NOT(ort_value.IsTensor(), "ort_value.IsTensor() was false");
|
||||
const Tensor& src_tensor = ort_value.Get<Tensor>();
|
||||
tensor_data_buffer.resize(src_tensor.SizeInBytes());
|
||||
|
||||
// Currently large model size not considered, so exception thrown here
|
||||
// when protobuf upper limit hit.
|
||||
total_bytes += src_tensor.SizeInBytes();
|
||||
if (total_bytes >= PROTOBUF_UPPER_LIMIT) {
|
||||
ORT_THROW("checkpoint file size hit upper limit.");
|
||||
}
|
||||
|
||||
auto& tensor_location = src_tensor.Location();
|
||||
if (tensor_location.device.Type() == OrtDevice::CPU ||
|
||||
tensor_location.mem_type == OrtMemTypeCPUInput ||
|
||||
tensor_location.mem_type == OrtMemTypeCPUOutput ||
|
||||
tensor_location.device.Type() == OrtDevice::GPU) {
|
||||
gsl::span<char> dst_span = gsl::make_span(tensor_data_buffer);
|
||||
ORT_RETURN_IF_NOT(src_tensor.SizeInBytes() == static_cast<size_t>(dst_span.size_bytes()), "src size != dst size");
|
||||
Tensor dst_tensor{src_tensor.DataType(), src_tensor.Shape(), dst_span.data(), cpu_alloc_info};
|
||||
ORT_RETURN_IF_ERROR(data_transfer_manager.CopyTensor(src_tensor, dst_tensor));
|
||||
|
||||
// Convert Tensor to TensorProto.
|
||||
ONNX_NAMESPACE::TensorProto tensor_proto;
|
||||
tensor_proto = utils::TensorToTensorProto(dst_tensor, tensor_name);
|
||||
saved_tensor_protos.emplace_back(tensor_proto);
|
||||
} else {
|
||||
ORT_THROW("Unsupported device type for saving tensors");
|
||||
}
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
PathString GetTensorProtoFilePath(const PathString& checkpoint_directory, const std::string& filename_prefix) {
|
||||
return ConcatPathComponent<PathChar>(checkpoint_directory, ORT_TSTR(filename_prefix + k_name_seperator) + k_tensor_proto_file_name);
|
||||
}
|
||||
|
||||
PathString GetTensorProtoPropertiesFilePath(const PathString& checkpoint_directory, const std::string& filename_prefix) {
|
||||
return ConcatPathComponent<PathChar>(checkpoint_directory, ORT_TSTR(filename_prefix + k_name_seperator) + k_tensor_proto_properties_file_name);
|
||||
}
|
||||
|
||||
std::string StringConcat(const std::string& s_a, const std::string& s_b, const std::string& del = k_name_seperator) {
|
||||
return s_a + del + s_b;
|
||||
}
|
||||
|
||||
void StringSplit(const std::string& s, std::vector<std::string>& results, const std::string& del = k_name_seperator) {
|
||||
ORT_ENFORCE(!s.empty(), "String to split is empty");
|
||||
int start = 0;
|
||||
int end = s.find(del);
|
||||
while (end != -1) {
|
||||
results.push_back(s.substr(start, end - start));
|
||||
start = end + del.size();
|
||||
end = s.find(del, start);
|
||||
}
|
||||
results.push_back(s.substr(start, end - start));
|
||||
}
|
||||
|
||||
bool StringStartsWith(std::string const& s, std::string const& p) {
|
||||
return s.rfind(p, 0) == 0;
|
||||
}
|
||||
|
||||
bool StringEndsWith(std::string const& s, std::string const& p) {
|
||||
if (p.size() > s.size()) return false;
|
||||
return std::equal(p.rbegin(), p.rend(), s.rbegin());
|
||||
}
|
||||
|
||||
void WriteTensorProtoToFile(const PathString& file_path,
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& tensor_protos,
|
||||
std::string caller_context) {
|
||||
auto file_write_status = WithOpenFile(
|
||||
file_path, false,
|
||||
[&tensor_protos](int fd) {
|
||||
google::protobuf::io::FileOutputStream output{fd};
|
||||
ORT_RETURN_IF_ERROR(WriteProtoMessageSequence(tensor_protos, output));
|
||||
return Status::OK();
|
||||
});
|
||||
|
||||
ORT_ENFORCE(file_write_status.IsOK(), caller_context, " write file failed: ", ToUTF8String(file_path));
|
||||
}
|
||||
|
||||
void LoadTensorProtoFromFile(const PathString& file_path,
|
||||
std::vector<ONNX_NAMESPACE::TensorProto>& tensor_protos,
|
||||
std::string caller_context) {
|
||||
auto file_read_status = WithOpenFile(
|
||||
file_path, true,
|
||||
[&tensor_protos](int fd) {
|
||||
google::protobuf::io::FileInputStream input{fd};
|
||||
ORT_RETURN_IF_ERROR(ReadProtoMessageSequence(tensor_protos, input));
|
||||
return Status::OK();
|
||||
});
|
||||
|
||||
ORT_ENFORCE(file_read_status.IsOK(), caller_context, " load file failed: ", ToUTF8String(file_path));
|
||||
}
|
||||
|
||||
template <typename Func>
|
||||
void FilterFilesFromDirectory(const PathString& folder_path, Func func) {
|
||||
LoopDir(folder_path, [&func](const PathChar* filename, OrtFileType file_type) -> bool {
|
||||
std::string filename_str = filename;
|
||||
if (filename_str[0] == '.' || file_type == OrtFileType::TYPE_DIR) {
|
||||
return true;
|
||||
}
|
||||
|
||||
return func(filename_str);
|
||||
});
|
||||
}
|
||||
|
||||
Status OrtSaveInternal(
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& trainable_tensor_protos,
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& non_trainable_tensor_protos,
|
||||
const PathString& checkpoint_path) {
|
||||
// Make sure name unique across trainable and non-trainable lists.
|
||||
std::unordered_set<std::string> trainable_unique_names;
|
||||
std::unordered_set<std::string> non_trainable_unique_names;
|
||||
std::vector<std::string> inter_sec;
|
||||
auto check_unique = [](const std::vector<ONNX_NAMESPACE::TensorProto>& tensor_protos,
|
||||
std::unordered_set<std::string>& unique_names) {
|
||||
for (auto& tensor_proto : tensor_protos) {
|
||||
ORT_ENFORCE(unique_names.find(tensor_proto.name()) == unique_names.end(),
|
||||
"Duplicated tensor proto named ", tensor_proto.name());
|
||||
unique_names.emplace(tensor_proto.name());
|
||||
}
|
||||
};
|
||||
check_unique(trainable_tensor_protos, trainable_unique_names);
|
||||
check_unique(non_trainable_tensor_protos, non_trainable_unique_names);
|
||||
std::set_intersection(trainable_unique_names.begin(), trainable_unique_names.end(),
|
||||
non_trainable_unique_names.begin(), non_trainable_unique_names.end(),
|
||||
std::back_inserter(inter_sec));
|
||||
ORT_RETURN_IF_NOT(inter_sec.empty(), "Tensor name exists in both trainable param list and non-trainable param list.");
|
||||
|
||||
// Keep following saving logic aligned with OrtSaveModuleStatesInternal.
|
||||
LOGS_DEFAULT(INFO)
|
||||
<< "Saving model checkpoint files to " << ToUTF8String(checkpoint_path);
|
||||
LOGS_DEFAULT_IF(Env::Default().FolderExists(checkpoint_path), WARNING)
|
||||
<< "Checkpoint directory exists - data may be overwritten.";
|
||||
ORT_RETURN_IF_ERROR(Env::Default().CreateFolder(checkpoint_path));
|
||||
|
||||
// Save TensorProto to file.
|
||||
if (trainable_tensor_protos.size() > 0) {
|
||||
WriteTensorProtoToFile(
|
||||
GetTensorProtoFilePath(checkpoint_path, k_trainable_param_root_prefix),
|
||||
trainable_tensor_protos, "[trainable_param]");
|
||||
}
|
||||
|
||||
if (non_trainable_tensor_protos.size() > 0) {
|
||||
WriteTensorProtoToFile(
|
||||
GetTensorProtoFilePath(checkpoint_path, k_non_trainable_param_root_prefix),
|
||||
non_trainable_tensor_protos, "[non_trainable_param]");
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtSaveModuleStatesInternal(ModuleCheckpointState& module_state,
|
||||
const PathString& parameter_folder_path) {
|
||||
// Write weight tensors files.
|
||||
const auto& param_states = module_state.named_parameters;
|
||||
if (!param_states.empty()) {
|
||||
ORT_ENFORCE(module_state.train_session_data_transfer_mgr,
|
||||
"module checkpoint state has null train_session_data_transfer_mgr.");
|
||||
|
||||
std::unordered_map<std::string, std::unordered_map<std::string, OrtValue>> parameter_ort_values;
|
||||
parameter_ort_values[k_trainable_param_root_prefix] = {};
|
||||
parameter_ort_values[k_non_trainable_param_root_prefix] = {};
|
||||
for (auto it = param_states.begin(); it != param_states.end(); ++it) {
|
||||
if (it->second->RequiresGrad()) {
|
||||
parameter_ort_values[k_trainable_param_root_prefix].insert({it->first, it->second->Data()});
|
||||
} else {
|
||||
parameter_ort_values[k_non_trainable_param_root_prefix].insert({it->first, it->second->Data()});
|
||||
}
|
||||
}
|
||||
|
||||
// Parameters saving.
|
||||
for (auto& pair : parameter_ort_values) {
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> param_tensor_protos;
|
||||
ORT_RETURN_IF_ERROR(CreateTensorProtosFromOrtValues(
|
||||
pair.second,
|
||||
*module_state.train_session_data_transfer_mgr,
|
||||
param_tensor_protos));
|
||||
|
||||
// Save TensorProto to file.
|
||||
WriteTensorProtoToFile(
|
||||
GetTensorProtoFilePath(parameter_folder_path, pair.first),
|
||||
param_tensor_protos, "[param]");
|
||||
}
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtSaveOptimizerStatesInternal(OptimizerCheckpointState& optimizer_state,
|
||||
const PathString& checkpoint_path) {
|
||||
if (optimizer_state.group_named_optimizer_states.empty()) {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
ORT_ENFORCE(optimizer_state.optimizer_session_data_transfer_mgr,
|
||||
"optimizer checkpoint state has null optimizer_session_data_transfer_mgr.");
|
||||
|
||||
// Write optimizer state tensors files.
|
||||
for (auto& group_named_optimizer_state : optimizer_state.group_named_optimizer_states) {
|
||||
const std::string& group_name = group_named_optimizer_state.first;
|
||||
const std::shared_ptr<GroupOptimizerState>& group_optimizer_state_ptr = group_named_optimizer_state.second;
|
||||
const std::string& cur_group_filename_prefix = StringConcat(k_optimizer_root_prefix, group_name);
|
||||
|
||||
// Re-organize optimizer_state_ort_values mapping
|
||||
// Firstly indexed by momentum names; Secondly indexed by parameter names.
|
||||
std::unordered_map<std::string, std::unordered_map<std::string, OrtValue>> optimizer_state_ort_values;
|
||||
for (const std::pair<std::string, ParameterOptimizerState>&
|
||||
param_named_optimizer_state : group_optimizer_state_ptr->param_named_optimizer_states) {
|
||||
const std::string& param_name = param_named_optimizer_state.first;
|
||||
const auto& param_optimizer_state = param_named_optimizer_state.second;
|
||||
|
||||
for (const std::pair<std::string, std::shared_ptr<OrtValue>>&
|
||||
momentum_named_state : param_optimizer_state.momentum_named_states) {
|
||||
const std::string& momentum_name = momentum_named_state.first;
|
||||
const std::shared_ptr<OrtValue>& m_state_val = momentum_named_state.second;
|
||||
|
||||
if (optimizer_state_ort_values.find(momentum_name) == optimizer_state_ort_values.end()) {
|
||||
std::unordered_map<std::string, OrtValue> param_name_to_ortvalue{{param_name, *(m_state_val)}};
|
||||
optimizer_state_ort_values.insert({momentum_name, param_name_to_ortvalue});
|
||||
} else {
|
||||
optimizer_state_ort_values[momentum_name].insert({param_name, *(m_state_val)});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Save each optimizer state (of all parameters) into single file.
|
||||
// For example: save "momentum_1" of all parameters into one file.
|
||||
for (auto& pair : optimizer_state_ort_values) {
|
||||
const auto& momentum_name = pair.first;
|
||||
const std::unordered_map<std::string, OrtValue>& param_name_to_ortvalue = pair.second;
|
||||
const std::string& cur_state_filename_prefix = StringConcat(cur_group_filename_prefix, momentum_name);
|
||||
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> saved_tensor_protos;
|
||||
ORT_RETURN_IF_ERROR(CreateTensorProtosFromOrtValues(
|
||||
param_name_to_ortvalue,
|
||||
*optimizer_state.optimizer_session_data_transfer_mgr,
|
||||
saved_tensor_protos));
|
||||
|
||||
// Save TensorProto to file.
|
||||
WriteTensorProtoToFile(
|
||||
GetTensorProtoFilePath(checkpoint_path, cur_state_filename_prefix),
|
||||
saved_tensor_protos, "[optimizer_state]");
|
||||
}
|
||||
|
||||
// Storing group-wise properties.
|
||||
PropertyBag properties;
|
||||
properties.AddProperty<float>(builtin_lr_property_name, group_optimizer_state_ptr->learning_rate);
|
||||
properties.AddProperty<int64_t>(builtin_step_property_name, group_optimizer_state_ptr->step);
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> group_wise_properties_tensor_protos;
|
||||
properties.ToTensorProtos(group_wise_properties_tensor_protos);
|
||||
|
||||
WriteTensorProtoToFile(
|
||||
GetTensorProtoPropertiesFilePath(checkpoint_path, cur_group_filename_prefix),
|
||||
group_wise_properties_tensor_protos, "[param_group_properties]");
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtSaveInternal(
|
||||
CheckpointState& state, const PathString& checkpoint_path) {
|
||||
LOGS_DEFAULT(INFO) << "Saving model checkpoint files to " << ToUTF8String(checkpoint_path);
|
||||
LOGS_DEFAULT_IF(Env::Default().FolderExists(checkpoint_path), WARNING)
|
||||
<< "Checkpoint directory exists - data may be overwritten.";
|
||||
ORT_RETURN_IF_ERROR(Env::Default().CreateFolder(checkpoint_path));
|
||||
|
||||
// Write weight tensors files.
|
||||
ORT_RETURN_IF_ERROR(OrtSaveModuleStatesInternal(state.module_checkpoint_state, checkpoint_path));
|
||||
|
||||
// Write optimizer state tensors files.
|
||||
ORT_RETURN_IF_ERROR(OrtSaveOptimizerStatesInternal(state.optimizer_checkpoint_state, checkpoint_path));
|
||||
|
||||
// Write properties file
|
||||
const PropertyBag& property_bag = state.property_bag;
|
||||
if (property_bag.Size() > 0) {
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> properties_tensor_protos;
|
||||
property_bag.ToTensorProtos(properties_tensor_protos);
|
||||
|
||||
WriteTensorProtoToFile(
|
||||
GetTensorProtoPropertiesFilePath(checkpoint_path, k_property_root_prefix),
|
||||
properties_tensor_protos, "[custom_properties]");
|
||||
}
|
||||
|
||||
LOGS_DEFAULT(INFO) << "Checkpoint saved successfully.";
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtLoadModuleStatesInternal(
|
||||
const PathString& parameter_folder_path, ModuleCheckpointState& module_state) {
|
||||
// Find parameter files.
|
||||
std::vector<std::pair<std::string, bool>> param_filenames;
|
||||
FilterFilesFromDirectory(
|
||||
parameter_folder_path,
|
||||
[¶m_filenames](const std::string& filename_str) -> bool {
|
||||
if (StringStartsWith(filename_str, k_trainable_param_root_prefix)) {
|
||||
param_filenames.push_back(std::make_pair(filename_str, true));
|
||||
} else if (StringStartsWith(filename_str, k_non_trainable_param_root_prefix)) {
|
||||
param_filenames.push_back(std::make_pair(filename_str, false));
|
||||
}
|
||||
return true;
|
||||
});
|
||||
|
||||
if (param_filenames.empty()) {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Parameter parsing.
|
||||
auto& named_parameters = module_state.named_parameters;
|
||||
auto load_model_proto_into_module =
|
||||
[&named_parameters](const PathString module_state_file_path, bool is_trainable) -> Status {
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> param_tensor_protos{};
|
||||
|
||||
LoadTensorProtoFromFile(module_state_file_path, param_tensor_protos, "[params]");
|
||||
|
||||
std::unordered_map<std::string, OrtValue> name_to_ort_values;
|
||||
ORT_RETURN_IF_ERROR(CreateOrtValuesFromTensorProtos(param_tensor_protos, name_to_ort_values));
|
||||
for (auto it = name_to_ort_values.begin(); it != name_to_ort_values.end(); ++it) {
|
||||
auto param = std::make_shared<Parameter>(it->first, it->second);
|
||||
ORT_RETURN_IF_ERROR(param->SetRequiresGrad(is_trainable));
|
||||
named_parameters.insert({it->first, param});
|
||||
}
|
||||
return Status::OK();
|
||||
};
|
||||
|
||||
for (auto& pair : param_filenames) {
|
||||
auto param_file_path = ConcatPathComponent<PathChar>(parameter_folder_path, pair.first);
|
||||
ORT_RETURN_IF_ERROR(load_model_proto_into_module(param_file_path, pair.second));
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtLoadOptimizerStatesInternal(const PathString& optimizer_folder_path,
|
||||
OptimizerCheckpointState& optimizer_state) {
|
||||
// Optimizer states parsing.
|
||||
std::vector<std::string> optim_state_filenames;
|
||||
std::vector<std::string> optim_property_filenames;
|
||||
FilterFilesFromDirectory(
|
||||
optimizer_folder_path,
|
||||
[&optim_state_filenames, &optim_property_filenames](const std::string& filename_str) -> bool {
|
||||
if (StringStartsWith(filename_str, k_optimizer_root_prefix)) {
|
||||
if (StringEndsWith(filename_str, k_tensor_proto_file_name)) {
|
||||
optim_state_filenames.push_back(filename_str);
|
||||
} else if (StringEndsWith(filename_str, k_tensor_proto_properties_file_name)) {
|
||||
optim_property_filenames.push_back(filename_str);
|
||||
} else {
|
||||
ORT_THROW("Unexpected file extension.");
|
||||
}
|
||||
}
|
||||
return true;
|
||||
});
|
||||
|
||||
auto& grouped_optimizer_states = optimizer_state.group_named_optimizer_states;
|
||||
// For each optimizer state files, parse the data and feed into grouped_optimizer_states.
|
||||
for (auto& filename : optim_state_filenames) {
|
||||
std::vector<std::string> results;
|
||||
StringSplit(filename, results);
|
||||
const std::string& group_name = results[1];
|
||||
const std::string& momentum_name = results[2];
|
||||
const std::string& cur_group_filename_prefix = StringConcat(k_optimizer_root_prefix, group_name);
|
||||
std::string cur_momentum_state_filename_prefix = StringConcat(cur_group_filename_prefix, momentum_name);
|
||||
|
||||
ORT_ENFORCE(filename.compare(StringConcat(cur_momentum_state_filename_prefix, k_tensor_proto_file_name)) == 0);
|
||||
|
||||
if (grouped_optimizer_states.find(group_name) == grouped_optimizer_states.end()) {
|
||||
grouped_optimizer_states.insert({group_name, std::make_shared<GroupOptimizerState>()});
|
||||
}
|
||||
|
||||
auto& group_optimizer_state = grouped_optimizer_states[group_name];
|
||||
std::unordered_map<std::string, ParameterOptimizerState>&
|
||||
param_optimizer_states = group_optimizer_state->param_named_optimizer_states;
|
||||
|
||||
const PathString& tensor_file_path = GetTensorProtoFilePath(optimizer_folder_path, cur_momentum_state_filename_prefix);
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> param_optimizer_state_tensor_protos{};
|
||||
LoadTensorProtoFromFile(tensor_file_path, param_optimizer_state_tensor_protos, "[optimizer_state]");
|
||||
|
||||
std::unordered_map<std::string, OrtValue> name_to_ort_values;
|
||||
ORT_RETURN_IF_ERROR(CreateOrtValuesFromTensorProtos(param_optimizer_state_tensor_protos, name_to_ort_values));
|
||||
for (auto& pair : name_to_ort_values) {
|
||||
auto& param_name = pair.first;
|
||||
if (param_optimizer_states.find(param_name) == param_optimizer_states.end()) {
|
||||
ParameterOptimizerState param_state;
|
||||
param_optimizer_states.insert({param_name, param_state});
|
||||
}
|
||||
param_optimizer_states[param_name].momentum_named_states.insert({momentum_name, std::make_shared<OrtValue>(pair.second)});
|
||||
}
|
||||
}
|
||||
|
||||
// For each optimizer properties files, parse the data and feed into grouped_optimizer_states.
|
||||
for (auto& filename : optim_property_filenames) {
|
||||
std::vector<std::string> results;
|
||||
StringSplit(filename, results);
|
||||
const std::string& group_name = results[1];
|
||||
|
||||
if (grouped_optimizer_states.find(group_name) == grouped_optimizer_states.end()) {
|
||||
grouped_optimizer_states.insert({group_name, std::make_shared<GroupOptimizerState>()});
|
||||
}
|
||||
|
||||
auto& group_optimizer_state = grouped_optimizer_states[group_name];
|
||||
|
||||
// Parse group-wise properties.
|
||||
const std::string& cur_group_filename_prefix = StringConcat(k_optimizer_root_prefix, group_name);
|
||||
const PathString& tensor_file_path = GetTensorProtoPropertiesFilePath(optimizer_folder_path, cur_group_filename_prefix);
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> group_wise_property_protos{};
|
||||
LoadTensorProtoFromFile(tensor_file_path, group_wise_property_protos, "[optimizer_groupwise_property]");
|
||||
|
||||
PropertyBag properties;
|
||||
for (auto& property_proto : group_wise_property_protos) {
|
||||
properties.AddProperty(property_proto);
|
||||
}
|
||||
|
||||
group_optimizer_state->learning_rate = properties.GetProperty<float>(builtin_lr_property_name);
|
||||
group_optimizer_state->step = properties.GetProperty<int64_t>(builtin_step_property_name);
|
||||
grouped_optimizer_states.insert({group_name, group_optimizer_state});
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtLoadCustomPropertyInternal(const PathString& property_folder_path,
|
||||
PropertyBag& property_bag) {
|
||||
// Find custom property files.
|
||||
std::vector<std::string> custom_property_filenames;
|
||||
FilterFilesFromDirectory(
|
||||
property_folder_path,
|
||||
[&custom_property_filenames](const std::string& filename_str) -> bool {
|
||||
if (StringStartsWith(filename_str, k_property_root_prefix)) {
|
||||
custom_property_filenames.push_back(filename_str);
|
||||
}
|
||||
return true;
|
||||
});
|
||||
|
||||
if (custom_property_filenames.empty()) {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
for (auto& property_file_path : custom_property_filenames) {
|
||||
std::vector<ONNX_NAMESPACE::TensorProto> property_protos{};
|
||||
auto property_file_full_path = ConcatPathComponent<PathChar>(property_folder_path, property_file_path);
|
||||
LoadTensorProtoFromFile(property_file_full_path, property_protos, "[custom_property]");
|
||||
|
||||
for (auto& property_proto : property_protos) {
|
||||
property_bag.AddProperty(property_proto);
|
||||
}
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status OrtLoadInternal(const PathString& checkpoint_path, CheckpointState& state) {
|
||||
ORT_ENFORCE(Env::Default().FolderExists(checkpoint_path), "Checkpoint folder not exit");
|
||||
ORT_RETURN_IF_ERROR(OrtLoadModuleStatesInternal(checkpoint_path, state.module_checkpoint_state));
|
||||
ORT_RETURN_IF_ERROR(OrtLoadOptimizerStatesInternal(checkpoint_path, state.optimizer_checkpoint_state));
|
||||
ORT_RETURN_IF_ERROR(OrtLoadCustomPropertyInternal(checkpoint_path, state.property_bag));
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
Status SaveCheckpoint(const std::vector<ONNX_NAMESPACE::TensorProto>& trainable_tensor_protos,
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& non_trainable_tensor_protos,
|
||||
const PathString& checkpoint_path) {
|
||||
return OrtSaveInternal(trainable_tensor_protos, non_trainable_tensor_protos, checkpoint_path);
|
||||
}
|
||||
|
||||
Status SaveCheckpoint(CheckpointState& states, const PathString& checkpoint_path) {
|
||||
return OrtSaveInternal(states, checkpoint_path);
|
||||
}
|
||||
|
||||
Status LoadCheckpoint(const PathString& checkpoint_path, CheckpointState& checkpoint_states) {
|
||||
return OrtLoadInternal(checkpoint_path, checkpoint_states);
|
||||
}
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
77
orttraining/orttraining/training_api/checkpoint_property.cc
Normal file
77
orttraining/orttraining/training_api/checkpoint_property.cc
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "onnx/defs/tensor_proto_util.h"
|
||||
#include "core/platform/path_lib.h"
|
||||
#include "core/platform/env.h"
|
||||
#include "core/framework/tensorprotoutils.h"
|
||||
#include "orttraining/training_api/include/checkpoint_property.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
template <typename T>
|
||||
TypedCheckpointProperty<T>::TypedCheckpointProperty(const ONNX_NAMESPACE::TensorProto& tensor_proto) {
|
||||
std::vector<int64_t> tensor_shape_vec = utils::GetTensorShapeFromTensorProto(tensor_proto);
|
||||
int64_t expected_num_elements = 1;
|
||||
for (auto& d : tensor_shape_vec) {
|
||||
expected_num_elements *= d;
|
||||
}
|
||||
ORT_ENFORCE(expected_num_elements == 1, "Only scalar value support for checkpoint property.");
|
||||
Path model_path;
|
||||
std::vector<T> data_vector(1);
|
||||
T* p = data_vector.data();
|
||||
ORT_ENFORCE(utils::UnpackTensor<T>(tensor_proto, model_path, p, expected_num_elements).IsOK());
|
||||
prop_name_ = tensor_proto.name();
|
||||
prop_value_ = data_vector[0];
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
ONNX_NAMESPACE::TensorProto TypedCheckpointProperty<T>::ToTensorProto() {
|
||||
auto t_proto = ONNX_NAMESPACE::ToTensor<T>(prop_value_);
|
||||
t_proto.set_name(prop_name_);
|
||||
return t_proto;
|
||||
}
|
||||
|
||||
namespace {
|
||||
|
||||
std::shared_ptr<CheckpointProperty> CreateCheckpointPropertyFromTensorProto(
|
||||
const ONNX_NAMESPACE::TensorProto& tensor_proto) {
|
||||
auto data_type = tensor_proto.data_type();
|
||||
switch (data_type) {
|
||||
case ONNX_NAMESPACE::TensorProto::FLOAT: {
|
||||
return std::static_pointer_cast<CheckpointProperty>(
|
||||
std::make_shared<TypedCheckpointProperty<float>>(tensor_proto));
|
||||
break;
|
||||
}
|
||||
case ONNX_NAMESPACE::TensorProto::STRING: {
|
||||
return std::static_pointer_cast<CheckpointProperty>(
|
||||
std::make_shared<TypedCheckpointProperty<std::string>>(tensor_proto));
|
||||
break;
|
||||
}
|
||||
case ONNX_NAMESPACE::TensorProto::INT64: {
|
||||
return std::static_pointer_cast<CheckpointProperty>(
|
||||
std::make_shared<TypedCheckpointProperty<int64_t>>(tensor_proto));
|
||||
break;
|
||||
}
|
||||
default:
|
||||
ORT_THROW("Unsupported input data type of ", data_type);
|
||||
}
|
||||
}
|
||||
} // namespace
|
||||
|
||||
void PropertyBag::AddProperty(const ONNX_NAMESPACE::TensorProto& tensor_proto) {
|
||||
ORT_ENFORCE(named_properties.find(tensor_proto.name()) == named_properties.end(),
|
||||
"Duplicated property named ", tensor_proto.name());
|
||||
|
||||
if (!IsSupportedDataType(tensor_proto.data_type())) {
|
||||
ORT_THROW("Failed to add property from tensorproto: float, int64_t and std::string data types supported only.");
|
||||
}
|
||||
|
||||
named_properties.insert({tensor_proto.name(), CreateCheckpointPropertyFromTensorProto(tensor_proto)});
|
||||
}
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
86
orttraining/orttraining/training_api/include/checkpoint.h
Normal file
86
orttraining/orttraining/training_api/include/checkpoint.h
Normal file
|
|
@ -0,0 +1,86 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "core/platform/path_lib.h"
|
||||
#include "core/platform/env.h"
|
||||
#include "onnx/defs/tensor_proto_util.h"
|
||||
|
||||
#include "orttraining/training_api/include/module.h"
|
||||
#include "orttraining/training_api/include/optimizer.h"
|
||||
#include "orttraining/training_api/include/checkpoint_property.h"
|
||||
|
||||
/**
|
||||
* There are two representation for checkpoint respectively in memory and files:
|
||||
*
|
||||
* 1. CheckpointState. A data class representing traing states in memory, which include:
|
||||
* i. module state:
|
||||
* a instance of data class `ModuleCheckpointState` managed along with Module/Parameter classes,
|
||||
* ii. optimizer state:
|
||||
* a instance of data class `OptimizerCheckpointState` managed along with Optimizer class,
|
||||
* iii. user defined training properties, for example 'epoch', 'best_score':
|
||||
* a instance of data class `PropertyBag` managed along with CheckpointProperty classes.
|
||||
*
|
||||
* In terms of class dependencies, Checkpoint implementations are dependent on (and on top of)
|
||||
* Parameter/Module/Optimizer/CheckpointProperty, NOT vice versa.
|
||||
*
|
||||
* 2. A directory of files:
|
||||
* checkpoint/
|
||||
* paramtrain_tensors.pbseq - trainable parameter tensor protobuf messages
|
||||
* paramfrozen_tensors.pbseq - non_trainable parameter tensor protobuf messages
|
||||
* optim_group0_momentum0_tensors.pbseq - optimizer momentum state tensor protobuf messages
|
||||
* optim_group0_momentum1_tensors.pbseq - optimizer momentum state tensor protobuf messages
|
||||
* optim_group0_properties.pbseq - group-wise optimizer property tensor protobuf messages
|
||||
* custom_properties.pbseq - custom property protobuf messages
|
||||
*
|
||||
* LoadCheckpoint takes CheckpointState as outputs, loading from a directory of checkpoint.
|
||||
* SaveCheckpoint takes CheckpointState as inputs, saving checkpoint files into a directory.
|
||||
*/
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
struct CheckpointState {
|
||||
public:
|
||||
ModuleCheckpointState module_checkpoint_state;
|
||||
OptimizerCheckpointState optimizer_checkpoint_state;
|
||||
PropertyBag property_bag;
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Save ONNX initializers as ORT checkpoint.
|
||||
*
|
||||
* @param trainable_tensor_protos trainable parameters in TensorProto format.
|
||||
* @param non_trainable_tensor_protos non-trainable parameters in TensorProto format.
|
||||
* @param checkpoint_path folder where checkpoint is saved.
|
||||
* @return Status
|
||||
*/
|
||||
Status SaveCheckpoint(const std::vector<ONNX_NAMESPACE::TensorProto>& trainable_tensor_protos,
|
||||
const std::vector<ONNX_NAMESPACE::TensorProto>& non_trainable_tensor_protos,
|
||||
const PathString& checkpoint_path);
|
||||
|
||||
/**
|
||||
* @brief Save training states as ORT checkpoint.
|
||||
*
|
||||
* @param state parameter/optimizer and other user defined training states.
|
||||
* @param checkpoint_path folder where checkpoint is saved.
|
||||
* @return Status
|
||||
*/
|
||||
Status SaveCheckpoint(CheckpointState& state,
|
||||
const PathString& checkpoint_path);
|
||||
|
||||
/**
|
||||
* @brief Load training states from ORT checkpoint.
|
||||
*
|
||||
* @param checkpoint_path folder where checkpoint is stored.
|
||||
* @param checkpoint_states parameter/optimizer and other user defined training states.
|
||||
* @return Status
|
||||
*/
|
||||
Status LoadCheckpoint(const PathString& checkpoint_path,
|
||||
CheckpointState& checkpoint_state);
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -0,0 +1,130 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <type_traits>
|
||||
#include "onnx/defs/tensor_proto_util.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
template <typename T>
|
||||
struct TypedCheckpointProperty;
|
||||
|
||||
/**
|
||||
* @brief Base class for user defined checkpoint property.
|
||||
*/
|
||||
struct CheckpointProperty {
|
||||
public:
|
||||
CheckpointProperty() {}
|
||||
|
||||
CheckpointProperty(const std::string& prop_name)
|
||||
: prop_name_(prop_name) {
|
||||
}
|
||||
|
||||
virtual ~CheckpointProperty() {}
|
||||
virtual ONNX_NAMESPACE::TensorProto ToTensorProto() = 0;
|
||||
|
||||
std::string GetName() const {
|
||||
return prop_name_;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
T GetData() {
|
||||
auto ptr = dynamic_cast<TypedCheckpointProperty<T>*>(this);
|
||||
ORT_ENFORCE(ptr);
|
||||
return ptr->GetData();
|
||||
}
|
||||
|
||||
protected:
|
||||
std::string prop_name_;
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief User defined checkpoint property.
|
||||
*/
|
||||
template <typename T>
|
||||
struct TypedCheckpointProperty : public CheckpointProperty {
|
||||
public:
|
||||
TypedCheckpointProperty(const std::string& prop_name, const T& prop_value)
|
||||
: CheckpointProperty(prop_name), prop_value_(prop_value) {
|
||||
}
|
||||
TypedCheckpointProperty(const ONNX_NAMESPACE::TensorProto& tensor_proto);
|
||||
|
||||
ONNX_NAMESPACE::TensorProto ToTensorProto() override;
|
||||
|
||||
T GetData() const {
|
||||
return prop_value_;
|
||||
}
|
||||
|
||||
private:
|
||||
T prop_value_;
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Collection of user defined properties.
|
||||
* Currently supported scalar value of type int64_t, float, and std::string only.
|
||||
*/
|
||||
struct PropertyBag {
|
||||
public:
|
||||
PropertyBag() {}
|
||||
|
||||
template <typename T>
|
||||
void AddProperty(std::string name, T val) {
|
||||
ORT_ENFORCE(named_properties.find(name) == named_properties.end(),
|
||||
"Duplicated property named ", name);
|
||||
|
||||
if (!IsSupportedDataType<T>()) {
|
||||
ORT_THROW("Failed to add property: float, int64_t and std::string data types supported only.");
|
||||
}
|
||||
|
||||
named_properties.insert({name, std::make_shared<TypedCheckpointProperty<T>>(name, val)});
|
||||
}
|
||||
|
||||
void AddProperty(const ONNX_NAMESPACE::TensorProto& tensor_proto);
|
||||
|
||||
template <typename T>
|
||||
T GetProperty(const std::string& name) const {
|
||||
if (!IsSupportedDataType<T>()) {
|
||||
ORT_THROW("Failed to get property: float, int64_t and std::string data types supported only.");
|
||||
}
|
||||
|
||||
auto it = named_properties.find(name);
|
||||
ORT_ENFORCE(it != named_properties.end(), "No property named ", name);
|
||||
return it->second->GetData<T>();
|
||||
}
|
||||
|
||||
void ToTensorProtos(std::vector<ONNX_NAMESPACE::TensorProto>& properties_tensor_protos) const {
|
||||
for (auto it = named_properties.begin(); it != named_properties.end(); ++it) {
|
||||
properties_tensor_protos.emplace_back((it->second)->ToTensorProto());
|
||||
}
|
||||
}
|
||||
|
||||
int Size() const {
|
||||
return named_properties.size();
|
||||
}
|
||||
|
||||
private:
|
||||
const std::vector<int32_t> supported_data_types{
|
||||
ONNX_NAMESPACE::TensorProto::FLOAT,
|
||||
ONNX_NAMESPACE::TensorProto::INT64,
|
||||
ONNX_NAMESPACE::TensorProto::STRING};
|
||||
|
||||
bool IsSupportedDataType(int32_t data_type) const {
|
||||
return std::find(supported_data_types.begin(), supported_data_types.end(), data_type) != supported_data_types.end();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
bool IsSupportedDataType() const {
|
||||
return (std::is_same<T, float>::value || std::is_same<T, int64_t>::value ||
|
||||
std::is_same<T, std::string>::value);
|
||||
}
|
||||
|
||||
std::unordered_map<std::string, std::shared_ptr<CheckpointProperty>> named_properties;
|
||||
};
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
25
orttraining/orttraining/training_api/include/interfaces.h
Normal file
25
orttraining/orttraining/training_api/include/interfaces.h
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "orttraining/training_api/include/module.h"
|
||||
#include "orttraining/training_api/include/optimizer.h"
|
||||
#include "orttraining/training_api/include/checkpoint_property.h"
|
||||
#include "orttraining/training_api/include/checkpoint.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
/*
|
||||
module.train_sess.RegisterExecutionProvider(provider);
|
||||
module.eval_sess.RegisterExecutionProvider(provider);
|
||||
optimizer.optim_sess.RegisterExecutionProvider(provider);
|
||||
*/
|
||||
void SetExecutionProvider(const Module& /*module*/, const Optimizer& /*optimizer*/, IExecutionProvider* /*provider*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
104
orttraining/orttraining/training_api/include/module.h
Normal file
104
orttraining/orttraining/training_api/include/module.h
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
#include "core/session/inference_session.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
struct Parameter {
|
||||
public:
|
||||
// Create parameter
|
||||
Parameter(std::string name, const OrtValue& data)
|
||||
: name_(name), data_(data) {
|
||||
}
|
||||
|
||||
// Return the mutable data.
|
||||
OrtValue& Data() { return data_; }
|
||||
std::string Name() const { return name_; }
|
||||
|
||||
// Return if trainable. The trainable property of a param
|
||||
// cannot change over the lifetime of the on-device training
|
||||
// session since the gradient graph is prebuilt for this setting.
|
||||
bool RequiresGrad() const { return requires_grad_; }
|
||||
|
||||
// Return the mutable gradient for trainable parameter.
|
||||
OrtValue& Gradient() { return gradient_; }
|
||||
std::string GradientName() const { return gradient_name_; }
|
||||
|
||||
// Reset and release the gradient buffer of this Parameter.
|
||||
Status ResetGrad() {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status SetRequiresGrad(bool requires_grad) {
|
||||
requires_grad_ = requires_grad;
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// need to set grad but not public api
|
||||
private:
|
||||
std::string name_;
|
||||
OrtValue data_;
|
||||
|
||||
OrtValue gradient_;
|
||||
std::string gradient_name_;
|
||||
|
||||
// Whether the param is trainable. The optimizer state is
|
||||
// only created for a trainable param
|
||||
bool requires_grad_{true};
|
||||
};
|
||||
|
||||
struct ModuleCheckpointState {
|
||||
public:
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>> named_parameters;
|
||||
const DataTransferManager* train_session_data_transfer_mgr;
|
||||
};
|
||||
|
||||
struct Module {
|
||||
public:
|
||||
// Initialize a module from an ORT inference session with loaded
|
||||
// training ONNX model and load parameters
|
||||
Module(const std::string& /*train_model_path_or_bytes*/,
|
||||
const std::unordered_map<std::string, std::shared_ptr<Parameter>>& /*parameters*/,
|
||||
const std::optional<std::string>& /*eval_model_path_or_bytes*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
// Return the trainable/nontrainable parameters
|
||||
std::vector<std::shared_ptr<Parameter>> parameters() const {
|
||||
return parameters_;
|
||||
}
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>> named_parameters() const {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return {};
|
||||
}
|
||||
|
||||
// Train Step – does forward and backward computation. The outputs will be the forward’s outputs.
|
||||
// Gradients will be accumulated within the Parameter object
|
||||
Status TrainStep(const std::vector<OrtValue>& /*inputs*/, std::vector<OrtValue>& /*outputs*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Eval Step – does forward computation. This will use a separate inference session
|
||||
// and take in a separate inference graph, while sharing the parameters
|
||||
Status EvalStep(const std::vector<OrtValue>& /*inputs*/, std::vector<OrtValue>& /*outputs*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Return the states of the module as a map.
|
||||
Status GetStateDict(ModuleCheckpointState& module_checkpoint_states);
|
||||
|
||||
private:
|
||||
std::unique_ptr<onnxruntime::InferenceSession> train_sess_;
|
||||
std::unique_ptr<onnxruntime::InferenceSession> eval_sess_;
|
||||
std::vector<std::shared_ptr<Parameter>> parameters_;
|
||||
};
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
125
orttraining/orttraining/training_api/include/optimizer.h
Normal file
125
orttraining/orttraining/training_api/include/optimizer.h
Normal file
|
|
@ -0,0 +1,125 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// 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"
|
||||
|
||||
#include "orttraining/training_api/include/module.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
/**
|
||||
* @brief States belong to one specific trainable Parameter.
|
||||
* Momentum states for each Parameter.
|
||||
* For Adam optimizer, it looks like:
|
||||
* {
|
||||
* "moment_0": shared_ptr<OrtValue>,
|
||||
* "moment_1": shared_ptr<OrtValue>,
|
||||
* }.
|
||||
*/
|
||||
struct ParameterOptimizerState {
|
||||
std::unordered_map<std::string, std::shared_ptr<OrtValue>> momentum_named_states;
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief States belong to one specific group of trainable Parameters.
|
||||
*/
|
||||
struct GroupOptimizerState {
|
||||
int64_t step;
|
||||
float learning_rate;
|
||||
std::unordered_map<std::string, ParameterOptimizerState> param_named_optimizer_states;
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief States belong to all groups of trainable Parameters.
|
||||
* Besides, also maintain a pointer of DataTransferManager* that is owned by InferenceSession.
|
||||
* This is used to do Tensor copy in the file saving stage.
|
||||
*/
|
||||
struct OptimizerCheckpointState {
|
||||
public:
|
||||
std::unordered_map<std::string, std::shared_ptr<GroupOptimizerState>> group_named_optimizer_states;
|
||||
const DataTransferManager* optimizer_session_data_transfer_mgr;
|
||||
};
|
||||
|
||||
struct Optimizer {
|
||||
public:
|
||||
// Initialize an optimizer module from an ORT inference session with loaded
|
||||
// training ONNX model For each parameter, initialize the OptimizerState based
|
||||
// on the graph input's ValueInfoProto if the parameter doesn't have it already.
|
||||
Optimizer(const std::string& optim_path_or_bytes,
|
||||
const std::unordered_map<std::string, std::shared_ptr<Parameter>>& parameters);
|
||||
|
||||
// Reset and release the gradient buffer of all trainable params
|
||||
Status ResetGrad() {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Optimizer Step.
|
||||
Status Step() {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status GetStateDict(OptimizerCheckpointState& optimizer_checkpoint_states);
|
||||
|
||||
protected:
|
||||
int64_t GetStep() const {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return 0;
|
||||
}
|
||||
Status SetLearningRate(float /*lr*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
private:
|
||||
std::unique_ptr<onnxruntime::InferenceSession> optim_sess_;
|
||||
std::vector<std::shared_ptr<Parameter>> parameters_;
|
||||
GroupOptimizerState optimizer_state_;
|
||||
};
|
||||
|
||||
class LearningRateScheduler {
|
||||
public:
|
||||
LearningRateScheduler(const Optimizer& optim)
|
||||
: optim_(optim) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
virtual ~LearningRateScheduler() = default;
|
||||
|
||||
// Modify the current learning rate based on current step
|
||||
virtual Status Step(/*int64_t step*/) = 0;
|
||||
|
||||
const Optimizer& optim_;
|
||||
};
|
||||
|
||||
class LinearScheduler : public LearningRateScheduler {
|
||||
public:
|
||||
explicit LinearScheduler(const Optimizer& optim, float start_factor, float end_factor, int64_t total_iters)
|
||||
: LearningRateScheduler(optim),
|
||||
start_factor_(start_factor),
|
||||
end_factor_(end_factor),
|
||||
total_iters_(total_iters) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
// Fetch the step, calculate next value and set lr in optimizer
|
||||
Status Step(/*int64_t step*/) override {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
private:
|
||||
float start_factor_;
|
||||
float end_factor_;
|
||||
int64_t total_iters_;
|
||||
};
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -1,232 +0,0 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#if defined(ENABLE_TRAINING) && defined(ENABLE_TRAINING_ON_DEVICE)
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api_test {
|
||||
|
||||
class Parameter {
|
||||
public:
|
||||
// create parameter
|
||||
Parameter(std::string /*name*/, const OrtValue& /*data*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
// Return the mutable data
|
||||
OrtValue& data() { return data_; }
|
||||
std::string name() const { return name_; }
|
||||
|
||||
// Return if trainable. The trainable property of a param
|
||||
// cannot change over the lifetime of the on-device training
|
||||
// session since the gradient graph is prebuilt for this setting.
|
||||
bool requires_grad() const { return requires_grad_; }
|
||||
|
||||
// Return the mutable gradient for trainable parameter
|
||||
OrtValue& gradient() { return gradient_; }
|
||||
std::string gradient_name() const { return gradient_name_; }
|
||||
|
||||
// Reset and release the gradient buffer of this Parameter
|
||||
Status ResetGrad() {
|
||||
return Status::OK();
|
||||
}
|
||||
// need to set grad but not public api
|
||||
private:
|
||||
OrtValue data_;
|
||||
std::string name_;
|
||||
|
||||
OrtValue gradient_;
|
||||
std::string gradient_name_;
|
||||
|
||||
// Whether the param is trainable. The optimizer state is
|
||||
// only created for a trainable param
|
||||
bool requires_grad_{true};
|
||||
};
|
||||
|
||||
class Module {
|
||||
public:
|
||||
// Initialize a module from an ORT inference session with loaded
|
||||
// training ONNX model and load parameters
|
||||
Module(const std::string& /*train_model_path_or_bytes*/,
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>>& /*parameters*/,
|
||||
const std::optional<std::string>& /*eval_model_path_or_bytes*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
// Return the trainable/nontrainable parameters
|
||||
std::vector<std::shared_ptr<Parameter>> parameters() const {
|
||||
return parameters_;
|
||||
}
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>> named_parameters() const {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return {};
|
||||
}
|
||||
|
||||
// Train Step – does forward and backward computation. The outputs will be the forward’s outputs. Gradients will be accumulated within the Parameter object
|
||||
Status TrainStep(const std::vector<OrtValue>& /*inputs*/, std::vector<OrtValue>& /*outputs*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Eval Step – does forward computation. This will use a separate inference session
|
||||
// and take in a separate inference graph, while sharing the parameters
|
||||
Status EvalStep(const std::vector<OrtValue>& /*inputs*/, std::vector<OrtValue>& /*outputs*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Return the states of the module as a map.
|
||||
Status GetStateDict(const std::unordered_map<std::string, std::shared_ptr<Parameter>>& /*module_state_dict*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
private:
|
||||
std::unique_ptr<onnxruntime::InferenceSession> train_sess_;
|
||||
std::unique_ptr<onnxruntime::InferenceSession> eval_sess_;
|
||||
std::vector<std::shared_ptr<Parameter>> parameters_;
|
||||
};
|
||||
|
||||
// Internal state
|
||||
struct ParameterOptimizerState {
|
||||
int64_t step_;
|
||||
float learning_rate_;
|
||||
// Per param optimizer state. E.g. For Adam and param_0, this would contain
|
||||
// {“Moment_1_param_0”:<value>, …},
|
||||
// It should be noted that the names should only be maintained to correlate with
|
||||
// the graph inputs for the optimizer graph
|
||||
std::map<std::string, OrtValue> states_;
|
||||
};
|
||||
|
||||
struct OptimizerState {
|
||||
// overall state related to optimizer
|
||||
int64_t step_;
|
||||
float learning_rate_;
|
||||
std::unordered_map<std::string, ParameterOptimizerState> optimizer_states_;
|
||||
};
|
||||
|
||||
class Optimizer {
|
||||
public:
|
||||
// Initialize an optimizer module from an ORT inference session with loaded
|
||||
// training ONNX model For each parameter, initialize the OptimizerState based
|
||||
// on the graph input’s ValueInfoProto if the parameter doesn’t have it already.
|
||||
Optimizer(const std::string& /*optim_path_or_bytes*/,
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>>& /*parameters*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
// Reset and release the gradient buffer of all trainable params
|
||||
Status ResetGrad() {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Optimizer Step.
|
||||
Status Step() {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Return the states of the optimizer as a map.
|
||||
Status GetStateDict(const OptimizerState& /*optimizer_state_dict*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
protected:
|
||||
int64_t GetStep() const {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return 0;
|
||||
}
|
||||
Status SetLearningRate(float /*lr*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
private:
|
||||
std::unique_ptr<onnxruntime::InferenceSession> optim_sess_;
|
||||
std::vector<std::shared_ptr<Parameter>> parameters_;
|
||||
OptimizerState optimizer_state_;
|
||||
};
|
||||
|
||||
class LearningRateScheduler {
|
||||
public:
|
||||
LearningRateScheduler(const Optimizer& optim)
|
||||
: optim_(optim) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
virtual ~LearningRateScheduler() = default;
|
||||
|
||||
// Modify the current learning rate based on current step
|
||||
virtual Status Step(/*int64_t step*/) = 0;
|
||||
|
||||
const Optimizer& optim_;
|
||||
};
|
||||
|
||||
class LinearScheduler : public LearningRateScheduler {
|
||||
public:
|
||||
explicit LinearScheduler(const Optimizer& optim, float start_factor, float end_factor, int64_t total_iters)
|
||||
: LearningRateScheduler(optim),
|
||||
start_factor_(start_factor),
|
||||
end_factor_(end_factor),
|
||||
total_iters_(total_iters) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
|
||||
// Fetch the step, calculate next value and set lr in optimizer
|
||||
Status Step(/*int64_t step*/) override {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
private:
|
||||
float start_factor_;
|
||||
float end_factor_;
|
||||
int64_t total_iters_;
|
||||
};
|
||||
|
||||
namespace utils {
|
||||
|
||||
struct CheckpointProperty {
|
||||
int value;
|
||||
// Support primitive types like int, float, string leveraging type trait.
|
||||
};
|
||||
|
||||
struct CheckpointStates {
|
||||
CheckpointStates() {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
std::unordered_map<std::string, std::shared_ptr<Parameter>> named_parameters;
|
||||
OptimizerState optimizer_states;
|
||||
std::unordered_map<std::string, CheckpointProperty> named_properties;
|
||||
};
|
||||
|
||||
// Save properties into a checkpoint property file (with postfix .prop).
|
||||
Status Ort_Save(CheckpointStates& /*state_dicts*/, const PathString& /*checkpoint_path*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Load properties file having postfix being '.prop'.
|
||||
Status Ort_Load(const PathString& /*checkpoint_path*/, CheckpointStates& /*state_dicts*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
/*
|
||||
module.train_sess.RegisterExecutionProvider(provider);
|
||||
module.eval_sess.RegisterExecutionProvider(provider);
|
||||
optimizer.optim_sess.RegisterExecutionProvider(provider);
|
||||
*/
|
||||
void SetExecutionProvider(const Module& /*module*/, const Optimizer& /*optimizer*/, IExecutionProvider* /*provider*/) {
|
||||
ORT_NOT_IMPLEMENTED("Not implemented.");
|
||||
}
|
||||
} // namespace utils
|
||||
|
||||
} // namespace api_test
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
|
||||
#endif
|
||||
24
orttraining/orttraining/training_api/module.cc
Normal file
24
orttraining/orttraining/training_api/module.cc
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "core/session/inference_session.h"
|
||||
#include "orttraining/training_api/include/module.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
Status Module::GetStateDict(ModuleCheckpointState& module_checkpoint_state) {
|
||||
module_checkpoint_state.named_parameters = named_parameters();
|
||||
|
||||
// Pass the training session data transfer manager for data copying when saving.
|
||||
// An alternative is, we can do copy at this stage.
|
||||
ORT_RETURN_IF_NOT(train_sess_, "training session not initialized");
|
||||
const DataTransferManager& sess_data_transfer_manager = train_sess_->GetDataTransferManager();
|
||||
module_checkpoint_state.train_session_data_transfer_mgr = &sess_data_transfer_manager;
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
80
orttraining/orttraining/training_api/optimizer.cc
Normal file
80
orttraining/orttraining/training_api/optimizer.cc
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "core/providers/cpu/cpu_execution_provider.h"
|
||||
#include "core/session/inference_session.h"
|
||||
#include "core/session/environment.h"
|
||||
|
||||
#include "orttraining/training_api/include/optimizer.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
namespace api {
|
||||
|
||||
namespace {
|
||||
|
||||
Status CreateOrtValueFromOrtValue(
|
||||
const OrtValue& src_ort_value,
|
||||
OrtValue& dest_ort_value,
|
||||
onnxruntime::InferenceSession* sess) {
|
||||
const Tensor& tensor = src_ort_value.Get<Tensor>();
|
||||
AllocatorPtr allocator = sess->GetAllocator(tensor.Location());
|
||||
const TensorShape& tensor_shape = tensor.Shape();
|
||||
MLDataType element_type = tensor.DataType();
|
||||
|
||||
auto p_tensor = std::make_unique<Tensor>(element_type, tensor_shape, allocator);
|
||||
dest_ort_value.Init(p_tensor.release(), DataTypeImpl::GetType<Tensor>(), DataTypeImpl::GetType<Tensor>()->GetDeleteFunc());
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
Optimizer::Optimizer(const std::string& optim_path_or_bytes,
|
||||
const std::unordered_map<std::string, std::shared_ptr<Parameter>>& parameters) {
|
||||
std::unordered_map<std::string, ParameterOptimizerState>&
|
||||
param_named_optimizer_states = optimizer_state_.param_named_optimizer_states;
|
||||
|
||||
const SessionOptions session_options;
|
||||
std::unique_ptr<Environment> env;
|
||||
ORT_ENFORCE(Environment::Create(nullptr, env) == Status::OK(), "Enviroment creation fails.");
|
||||
optim_sess_ = std::move(std::make_unique<InferenceSession>(session_options, *env));
|
||||
|
||||
ORT_ENFORCE(optim_sess_->Load(optim_path_or_bytes).IsOK());
|
||||
ORT_ENFORCE(optim_sess_->Initialize().IsOK());
|
||||
|
||||
// TODO: don't hard code the state names, should get the state names according to the optimizer types.
|
||||
std::vector<std::string> state_names{"momentum0", "momentum1"};
|
||||
for (auto& pair : 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 : state_names) {
|
||||
OrtValue param_state;
|
||||
// TODO: should reset the state to zero (for both CPU or CUDA Tensors.)
|
||||
ORT_ENFORCE(CreateOrtValueFromOrtValue(pair.second->Data(), param_state, optim_sess_.get()).IsOK());
|
||||
cur_param_optimizer_states.momentum_named_states.insert({state_name, std::make_shared<OrtValue>(param_state)});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Status Optimizer::GetStateDict(OptimizerCheckpointState& optimizer_checkpoint_state) {
|
||||
auto& grouped_optimizer_states = optimizer_checkpoint_state.group_named_optimizer_states;
|
||||
|
||||
// Currently all parameters are in a single group, so we hardcode group0 here.
|
||||
// To support multiple groups, Optimizer constructor need accept informations for groupping.
|
||||
const std::string group_zero_name = "group0";
|
||||
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.
|
||||
// An alternative is, we can do copy at this stage.
|
||||
ORT_RETURN_IF_NOT(optim_sess_, "optimizer session not initialized");
|
||||
const DataTransferManager& sess_data_transfer_manager = optim_sess_->GetDataTransferManager();
|
||||
optimizer_checkpoint_state.optimizer_session_data_transfer_mgr = &sess_data_transfer_manager;
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace api
|
||||
} // namespace training
|
||||
} // namespace onnxruntime
|
||||
Loading…
Reference in a new issue