onnxruntime/orttraining/orttraining/python/orttraining_pybind_common.h
Adam Louly e49f358686
expose lr scheduler python bindings for on device training. (#13882)
### Description
Exposing LR Scheduler python bindings for on device training.

Co-authored-by: Baiju Meswani <bmeswani@microsoft.com>
2022-12-22 18:44:04 -08:00

54 lines
1.8 KiB
C++

// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "python/onnxruntime_pybind_exceptions.h"
#include "python/onnxruntime_pybind_mlvalue.h"
#include "python/onnxruntime_pybind_state_common.h"
#include "core/platform/env.h"
#include <unordered_map>
#include <cstdlib>
namespace onnxruntime {
namespace python {
namespace py = pybind11;
using namespace onnxruntime::logging;
using ExecutionProviderMap = std::unordered_map<std::string, std::shared_ptr<IExecutionProvider>>;
using ExecutionProviderLibInfoMap = std::unordered_map<std::string, std::pair<std::string, ProviderOptions>>;
class ORTTrainingPythonEnv {
public:
ORTTrainingPythonEnv();
Environment& GetORTEnv();
std::shared_ptr<IExecutionProvider> GetExecutionProviderInstance(const std::string& provider_type,
size_t hash);
void AddExecutionProvider(const std::string& provider_type,
size_t hash,
std::unique_ptr<IExecutionProvider> execution_provider);
void RegisterExtExecutionProviderInfo(const std::string& provider_type,
const std::string& provider_lib_path,
const ProviderOptions& default_options);
const std::vector<std::string>& GetAvailableTrainingExecutionProviderTypes();
ExecutionProviderLibInfoMap ext_execution_provider_info_map_;
void ClearExecutionProviderInstances();
private:
std::string GetExecutionProviderMapKey(const std::string& provider_type,
size_t hash);
std::unique_ptr<Environment> ort_env_;
ExecutionProviderMap execution_provider_instances_map_;
std::vector<std::string> available_training_eps_;
};
} // namespace python
} // namespace onnxruntime