From fb4707f76d9161b3f8fdc4ba294e3533382f30f3 Mon Sep 17 00:00:00 2001 From: Adam Louly Date: Thu, 8 Dec 2022 18:03:53 -0600 Subject: [PATCH] add cuda support to python bindings (#13700) ### Description Add cuda support to the on device training python bindings. ### Motivation and Context Now users can set the execution provider (cpu or cuda) when using python bindings for on device training apis. --- .../python/onnxruntime_pybind_state_common.h | 1 - .../python/orttraining_pybind_state.cc | 42 ++++++++++++++++--- .../orttraining/python/training/api/module.py | 19 ++++++--- .../python/training/api/optimizer.py | 2 +- .../orttraining_test_python_bindings.py | 17 ++++++++ 5 files changed, 68 insertions(+), 13 deletions(-) diff --git a/onnxruntime/python/onnxruntime_pybind_state_common.h b/onnxruntime/python/onnxruntime_pybind_state_common.h index 482ec9befa..9a97576949 100644 --- a/onnxruntime/python/onnxruntime_pybind_state_common.h +++ b/onnxruntime/python/onnxruntime_pybind_state_common.h @@ -506,6 +506,5 @@ std::shared_ptr CreateExecutionProviderFactory_Nnapi( uint32_t flags, const optional& partitioning_stop_ops_list); std::shared_ptr CreateExecutionProviderFactory_Rknpu(); std::shared_ptr CreateExecutionProviderFactory_CoreML(uint32_t flags); - constexpr const char* kDefaultExecutionProviderEntry = "GetProvider"; } // namespace onnxruntime diff --git a/orttraining/orttraining/python/orttraining_pybind_state.cc b/orttraining/orttraining/python/orttraining_pybind_state.cc index ca9b63f80a..3015feeab2 100644 --- a/orttraining/orttraining/python/orttraining_pybind_state.cc +++ b/orttraining/orttraining/python/orttraining_pybind_state.cc @@ -34,6 +34,7 @@ #ifdef ENABLE_TRAINING_ON_DEVICE #include "orttraining/training_api/include/checkpoint.h" +#include "core/providers/provider_factory_creators.h" #endif @@ -74,7 +75,33 @@ void ResolveExtraProviderOptions(const std::vector& provider_types, j += 1; } } +#ifdef ENABLE_TRAINING_ON_DEVICE +namespace { +// This function is used to create an execution provider to be passed to Module and Optimizer. +std::vector> +GetExecutionProvidersForTrainingApis(OrtDevice device) { + std::vector> provider; +#ifdef USE_CUDA + if (device.Type() == OrtDevice::GPU) { + OrtCUDAProviderOptions provider_options{}; + provider_options.device_id = device.Id(); + + if (auto factory = CudaProviderFactoryCreator::Create(&provider_options)) + provider.push_back(factory->CreateProvider()); + + return provider; + } +#endif + if (device.Type() == OrtDevice::CPU) { + provider = std::vector>(); + } else { + ORT_THROW("Unsupported device type: ", device.Type()); + } + return provider; +} +} // namespace +#endif struct TrainingParameters { std::string loss_output_name; std::unordered_set weights_to_train; @@ -823,18 +850,19 @@ void addObjectMethodsForTraining(py::module& m, ExecutionProviderRegistrationFn GradientDefinitionRegistry::Instance().SetStopGradientEdgesForNode(key, edges); }); #ifdef ENABLE_TRAINING_ON_DEVICE - // Python apis only supports CPU device for now. - // TODO(adamlouly) : Add support for CUDA device. py::class_ training_module(m, "Module", R"pbdoc(Training Module.)pbdoc"); training_module .def(py::init([](const std::string& model_uri, onnxruntime::training::api::CheckpointState& state, - std::optional eval_model_uri) { + std::optional eval_model_uri, + OrtDevice device) { onnxruntime::SessionOptions session_option; + std::vector> provider = GetExecutionProvidersForTrainingApis(device); + return std::make_unique( model_uri, state.module_checkpoint_state.named_parameters, session_option, - GetTrainingORTEnv(), std::vector>(), eval_model_uri); + GetTrainingORTEnv(), provider, eval_model_uri); })) .def("train_step", [](onnxruntime::training::api::Module* model, @@ -890,12 +918,14 @@ void addObjectMethodsForTraining(py::module& m, ExecutionProviderRegistrationFn training_optimizer(m, "Optimizer", R"pbdoc(Training Optimizer.)pbdoc"); training_optimizer.def(py::init([]( const std::string optimizer_model_uri, - onnxruntime::training::api::Module* model) { + onnxruntime::training::api::Module* model, + OrtDevice device) { onnxruntime::SessionOptions session_option; + std::vector> provider = GetExecutionProvidersForTrainingApis(device); return std::make_unique( optimizer_model_uri, model->NamedParameters(), session_option, - GetTrainingORTEnv(), std::vector>()); + GetTrainingORTEnv(), provider); })) .def("set_learning_rate", [](onnxruntime::training::api::Optimizer* optimizer, float lr) -> void { ORT_THROW_IF_ERROR(optimizer->SetLearningRate(lr)); diff --git a/orttraining/orttraining/python/training/api/module.py b/orttraining/orttraining/python/training/api/module.py index 2bcca528cb..d98fff2c5c 100644 --- a/orttraining/orttraining/python/training/api/module.py +++ b/orttraining/orttraining/python/training/api/module.py @@ -5,7 +5,7 @@ import numpy as np from onnxruntime.capi import _pybind_state as C -from onnxruntime.capi.onnxruntime_inference_collection import OrtValue +from onnxruntime.capi.onnxruntime_inference_collection import OrtValue, get_ort_device_type from onnxruntime.capi.onnxruntime_pybind11_state import OrtValueVector @@ -17,14 +17,23 @@ class Module: training: bool - def __init__(self, train_model_uri, state, eval_model_uri=None) -> None: + def __init__(self, train_model_uri, state, eval_model_uri=None, device: str = "cpu") -> None: """ Initializes Model for Training. __init__ will call an internatl function to create the model. """ # TODO : Add support for bytes on train_model_uri and eval_model_uri. self.training = True - self._model = C.Module(train_model_uri, state._state, eval_model_uri) + options = device.split(":") + self._device_type = options[0] + device_id = 0 if len(options) < 2 else int(options[1]) + + self._device = C.OrtDevice( + get_ort_device_type(self._device_type, device_id), + C.OrtDevice.default_memory(), + device_id, + ) + self._model = C.Module(train_model_uri, state._state, eval_model_uri, self._device) def __call__(self, user_inputs): """ @@ -94,8 +103,8 @@ class Module: self.get_parameters_size(trainable_only), ], np.float32, - "cpu", - 0, + self._device_type, + self._device.device_id(), )._ortvalue self._model.copy_parameters_to_buffer(parameters) diff --git a/orttraining/orttraining/python/training/api/optimizer.py b/orttraining/orttraining/python/training/api/optimizer.py index 65480a5e9e..3a89b27e9f 100644 --- a/orttraining/orttraining/python/training/api/optimizer.py +++ b/orttraining/orttraining/python/training/api/optimizer.py @@ -15,7 +15,7 @@ class Optimizer: """ Initializes Optimizer with the optimizer onnx and the parameters from the model. """ - self._optimizer = C.Optimizer(train_optimizer_uri, model._model) + self._optimizer = C.Optimizer(train_optimizer_uri, model._model, model._device) def step(self): """ diff --git a/orttraining/orttraining/test/python/orttraining_test_python_bindings.py b/orttraining/orttraining/test/python/orttraining_test_python_bindings.py index 378e58a3ac..78ae346064 100644 --- a/orttraining/orttraining/test/python/orttraining_test_python_bindings.py +++ b/orttraining/orttraining/test/python/orttraining_test_python_bindings.py @@ -259,3 +259,20 @@ def test_export_model_for_inferencing(): inference_model_file_path = os.path.join(temp_dir, "inference_model.onnx") model.export_model_for_inferencing(inference_model_file_path, ["output-0"]) assert os.path.exists(inference_model_file_path) + + +def test_cuda_execution_provider(): + # Initialize Models + simple_model, onnx_model, _, _, pt_model = _create_training_models() + + with tempfile.TemporaryDirectory() as temp_dir: + # Save models & checkpoint files to load them later. + checkpoint_file_path, model_file_path = _get_test_models_path(temp_dir, simple_model, onnx_model) + # Create Checkpoint State. + state = CheckpointState(checkpoint_file_path) + # Create a Module. + model = Module(model_file_path, state, device="cuda") + params = model.get_contiguous_parameters() + + # Check if parameters are moved to cuda. + assert params.device_name() == "Cuda"