diff --git a/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs b/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs index 2e9a9c5498..576804b6a0 100644 --- a/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs +++ b/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs @@ -19,7 +19,7 @@ namespace Microsoft.ML.OnnxRuntime public IntPtr TrainingSessionGetEvalModelOutputCount; public IntPtr TrainingSessionGetTrainingModelOutputName; public IntPtr TrainingSessionGetEvalModelOutputName; - public IntPtr ResetGrad; + public IntPtr LazyResetGrad; public IntPtr TrainStep; public IntPtr EvalStep; public IntPtr SetLearningRate; @@ -67,7 +67,7 @@ namespace Microsoft.ML.OnnxRuntime OrtGetEvalModelOutputCount = (DOrtGetEvalModelOutputCount)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetEvalModelOutputCount, typeof(DOrtGetEvalModelOutputCount)); OrtGetTrainingModelOutputName = (DOrtGetTrainingModelOutputName)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetTrainingModelOutputName, typeof(DOrtGetTrainingModelOutputName)); OrtGetEvalModelOutputName = (DOrtGetEvalModelOutputName)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetEvalModelOutputName, typeof(DOrtGetEvalModelOutputName)); - OrtResetGrad = (DOrtResetGrad)Marshal.GetDelegateForFunctionPointer(trainingApi_.ResetGrad, typeof(DOrtResetGrad)); + OrtLazyResetGrad = (DOrtLazyResetGrad)Marshal.GetDelegateForFunctionPointer(trainingApi_.LazyResetGrad, typeof(DOrtLazyResetGrad)); OrtTrainStep = (DOrtTrainStep)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainStep, typeof(DOrtTrainStep)); OrtEvalStep = (DOrtEvalStep)Marshal.GetDelegateForFunctionPointer(trainingApi_.EvalStep, typeof(DOrtEvalStep)); OrtSetLearningRate = (DOrtSetLearningRate)Marshal.GetDelegateForFunctionPointer(trainingApi_.SetLearningRate, typeof(DOrtSetLearningRate)); @@ -81,7 +81,7 @@ namespace Microsoft.ML.OnnxRuntime } -#region TrainingSession API + #region TrainingSession API /// /// Creates an instance of OrtSession with provided parameters @@ -164,10 +164,10 @@ namespace Microsoft.ML.OnnxRuntime public static DOrtGetEvalModelOutputName OrtGetEvalModelOutputName; [UnmanagedFunctionPointer(CallingConvention.Winapi)] - public delegate IntPtr /*(OrtStatus*)*/ DOrtResetGrad( + public delegate IntPtr /*(OrtStatus*)*/ DOrtLazyResetGrad( IntPtr /*(OrtTrainingSession*)*/ session); - public static DOrtResetGrad OrtResetGrad; + public static DOrtLazyResetGrad OrtLazyResetGrad; [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr /*(ONNStatus*)*/ DOrtTrainStep( @@ -240,7 +240,7 @@ namespace Microsoft.ML.OnnxRuntime public delegate void DOrtReleaseCheckpointState(IntPtr /*(OrtCheckpointState*)*/checkpointState); public static DOrtReleaseCheckpointState OrtReleaseCheckpointState; -#endregion TrainingSession API + #endregion TrainingSession API public static bool TrainingEnabled() { diff --git a/csharp/src/Microsoft.ML.OnnxRuntime/Training/TrainingSession.shared.cs b/csharp/src/Microsoft.ML.OnnxRuntime/Training/TrainingSession.shared.cs index 8d804628af..7ccfdb30a5 100644 --- a/csharp/src/Microsoft.ML.OnnxRuntime/Training/TrainingSession.shared.cs +++ b/csharp/src/Microsoft.ML.OnnxRuntime/Training/TrainingSession.shared.cs @@ -205,9 +205,9 @@ namespace Microsoft.ML.OnnxRuntime /// Sets the reset grad flag on the training graph. The gradient buffers will be reset while executing the /// next train step. /// - public void ResetGrad() + public void LazyResetGrad() { - NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtResetGrad(_nativeHandle)); + NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtLazyResetGrad(_nativeHandle)); } /// @@ -388,15 +388,15 @@ namespace Microsoft.ML.OnnxRuntime IntPtr nameHandle; string str = null; if (training) - { + { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetTrainingModelOutputName( _nativeHandle, (UIntPtr)index, allocator.Pointer, out nameHandle)); - } + } else - { + { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetEvalModelOutputName( _nativeHandle, (UIntPtr)index, @@ -498,7 +498,7 @@ namespace Microsoft.ML.OnnxRuntime } } - #endregion + #endregion } #endif } diff --git a/csharp/test/Microsoft.ML.OnnxRuntime.Tests.Common/TrainingTest.cs b/csharp/test/Microsoft.ML.OnnxRuntime.Tests.Common/TrainingTest.cs index f643efa720..3c1ba092c0 100644 --- a/csharp/test/Microsoft.ML.OnnxRuntime.Tests.Common/TrainingTest.cs +++ b/csharp/test/Microsoft.ML.OnnxRuntime.Tests.Common/TrainingTest.cs @@ -129,7 +129,7 @@ namespace Microsoft.ML.OnnxRuntime.Tests TensorElementType.Int32, labelsShape, labels.Length * sizeof(Int32))); var outputs = trainingSession.TrainStep(pinnedInputs); - trainingSession.ResetGrad(); + trainingSession.LazyResetGrad(); outputs = trainingSession.TrainStep(pinnedInputs); var outputBuffer = outputs.ElementAtOrDefault(0); @@ -233,7 +233,7 @@ namespace Microsoft.ML.OnnxRuntime.Tests trainingSession.TrainStep(pinnedInputs, pinnedOutputs); Assert.Equal(expectedOutput_1, outputBuffer, new FloatComparer()); - trainingSession.ResetGrad(); + trainingSession.LazyResetGrad(); trainingSession.TrainStep(pinnedInputs, pinnedOutputs); Assert.Equal(expectedOutput_1, outputBuffer, new FloatComparer()); @@ -316,5 +316,5 @@ namespace Microsoft.ML.OnnxRuntime.Tests } } #endif - } + } } diff --git a/orttraining/orttraining/python/orttraining_pybind_state.cc b/orttraining/orttraining/python/orttraining_pybind_state.cc index 3015feeab2..a5548ff2ee 100644 --- a/orttraining/orttraining/python/orttraining_pybind_state.cc +++ b/orttraining/orttraining/python/orttraining_pybind_state.cc @@ -874,9 +874,9 @@ void addObjectMethodsForTraining(py::module& m, ExecutionProviderRegistrationFn const std::vector& user_inputs, std::vector& user_outputs) -> void { ORT_THROW_IF_ERROR(model->EvalStep(user_inputs, user_outputs)); }) - .def("reset_grad", + .def("lazy_reset_grad", [](onnxruntime::training::api::Module* model) -> void { - ORT_THROW_IF_ERROR(model->ResetGrad()); + ORT_THROW_IF_ERROR(model->LazyResetGrad()); }) .def("copy_parameters_to_buffer", [](onnxruntime::training::api::Module* model, OrtValue& output) -> void { diff --git a/orttraining/orttraining/python/training/api/module.py b/orttraining/orttraining/python/training/api/module.py index d98fff2c5c..433e190924 100644 --- a/orttraining/orttraining/python/training/api/module.py +++ b/orttraining/orttraining/python/training/api/module.py @@ -80,11 +80,14 @@ class Module: """ return self.train(False) - def reset_grad(self): + def lazy_reset_grad(self): + """Lazily resets the training gradients. + + This function sets the internal state of the module such that the module gradients + will be scheduled to be reset just before the new gradients are computed on the next invocation + of train(). """ - Resets the gradient of the parameters. - """ - return self._model.reset_grad() + return self._model.lazy_reset_grad() def save_checkpoint(self, ckpt_uri): """ @@ -127,3 +130,4 @@ class Module: Exports the model for inferencing. """ self._model.export_model_for_inferencing(inference_model_uri, graph_output_names) + self._model.export_model_for_inferencing(inference_model_uri, graph_output_names) diff --git a/orttraining/orttraining/test/training_api/core/training_api_tests.cc b/orttraining/orttraining/test/training_api/core/training_api_tests.cc index 877b62cbfe..fc6bf282df 100644 --- a/orttraining/orttraining/test/training_api/core/training_api_tests.cc +++ b/orttraining/orttraining/test/training_api/core/training_api_tests.cc @@ -300,7 +300,7 @@ TEST(TrainingApiTest, ModuleTrainStep) { } } // reset grad - ASSERT_STATUS_OK(model->ResetGrad()); + ASSERT_STATUS_OK(model->LazyResetGrad()); // run a single step std::vector& inputs = *data_loader.begin(); diff --git a/orttraining/orttraining/test/training_api/trainer/trainer.cc b/orttraining/orttraining/test/training_api/trainer/trainer.cc index 646fc91a7b..fffd34aa1b 100644 --- a/orttraining/orttraining/test/training_api/trainer/trainer.cc +++ b/orttraining/orttraining/test/training_api/trainer/trainer.cc @@ -301,12 +301,12 @@ int RunTraining(const TestRunnerParameters& params) { #if defined(USE_CUDA) && defined(ENABLE_NVTX_PROFILE) onnxruntime::profile::NvtxRangeCreator resetgrad_range( - "ResetGrad", + "LazyResetGrad", onnxruntime::profile::Color::Red); resetgrad_range.Begin(); #endif - session.ResetGrad(); + session.LazyResetGrad(); #if defined(USE_CUDA) && defined(ENABLE_NVTX_PROFILE) resetgrad_range.End(); diff --git a/orttraining/orttraining/training_api/include/module.h b/orttraining/orttraining/training_api/include/module.h index 738530f54e..12ce79d7a2 100644 --- a/orttraining/orttraining/training_api/include/module.h +++ b/orttraining/orttraining/training_api/include/module.h @@ -72,7 +72,7 @@ struct Module { } // Reset and release the gradient buffer of all trainable params lazily. - Status ResetGrad(); + Status LazyResetGrad(); // Train Step – does forward and backward computation. The outputs will be the forward’s outputs. // Gradients will be accumulated within the Parameter object @@ -111,6 +111,18 @@ struct Module { Status ExportModelForInferencing(const std::string& inference_model_path, gsl::span graph_output_names) const; + // Returns the user input count for training graph + size_t GetTrainingModelInputCount() const noexcept; + + // Returns the user input count for eval graph + size_t GetEvalModelInputCount() const noexcept; + + // Returns the user input name for train graph at given index + std::string GetTrainingModelInputName(size_t index) const; + + // Returns the user input name for eval graph at given index + std::string GetEvalModelInputName(size_t index) const; + private: std::unique_ptr train_sess_{nullptr}; std::unique_ptr eval_sess_{nullptr}; @@ -124,6 +136,8 @@ struct Module { bool accumulate_gradient_ = false; const std::unordered_map>& named_parameters_; std::string eval_model_path_; + size_t train_user_input_count_; + size_t eval_user_input_count_; }; } // namespace api diff --git a/orttraining/orttraining/training_api/include/onnxruntime_training_c_api.h b/orttraining/orttraining/training_api/include/onnxruntime_training_c_api.h index 20c0dc0bd8..c3cd5b692e 100644 --- a/orttraining/orttraining/training_api/include/onnxruntime_training_c_api.h +++ b/orttraining/orttraining/training_api/include/onnxruntime_training_c_api.h @@ -11,84 +11,84 @@ ORT_RUNTIME_CLASS(CheckpointState); /// Type that holds the training states for struct OrtTrainingApi { /** \brief Load a checkpoint state from directory on disk into checkpoint_state. - * - * This function will parse a checkpoint directory, pull relevant files and load the training - * states into the checkpoint_state. This checkpoint state can then be used to create the - * training session by invoking CreateTrainingSession. By doing so, the training session will resume - * training from the given checkpoint. - * - * \param[in] checkpoint_path Path to the checkpoint directory - * \param[out] checkpoint_state Checkpoint states that contains the states of the training session. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function will parse a checkpoint directory, pull relevant files and load the training + * states into the checkpoint_state. This checkpoint state can then be used to create the + * training session by invoking CreateTrainingSession. By doing so, the training session will resume + * training from the given checkpoint. + * + * \param[in] checkpoint_path Path to the checkpoint directory + * \param[out] checkpoint_state Checkpoint states that contains the states of the training session. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(LoadCheckpoint, _In_ const ORTCHAR_T* checkpoint_path, _Outptr_ OrtCheckpointState** checkpoint_state); /** \brief Save the training session states to a checkpoint directory on disk. - * - * This function retrieves the training session states from the training session and serializes them - * to a checkpoint directory on disk. This checkpoint can later be loaded by invoking LoadCheckpoint - * to continue the training with the same states. - * - * \param[in] checkpoint_path Path to the checkpoint directory - * \param[in] session The training session from where the checkpoint states are to be retrieved. - * \param[in] save_optimizer_state Boolean flag indicating whether or not to save the optimizer states to the checkpoint. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function retrieves the training session states from the training session and serializes them + * to a checkpoint directory on disk. This checkpoint can later be loaded by invoking LoadCheckpoint + * to continue the training with the same states. + * + * \param[in] checkpoint_path Path to the checkpoint directory + * \param[in] session The training session from where the checkpoint states are to be retrieved. + * \param[in] save_optimizer_state Boolean flag indicating whether or not to save the optimizer states to the checkpoint. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(SaveCheckpoint, _In_ const ORTCHAR_T* checkpoint_path, _In_ const OrtTrainingSession* session, bool save_optimizer_state); /** \brief Create a training session that can be used to begin or resume training. - * - * This function creates a training session based on the env and session options provided that can - * begin or resume training from a given checkpoint state for the given onnx models. - * The checkpoint state represents the parameters of the training session which will be moved - * to the device specified by the user through the session options (if necessary). - * - * \param[in] env Environment to be used for the training session. - * \param[in] options Session options that the user can customize for this training session. - * \param[in] checkpoint_state Training states that the training session uses as a starting point for training. - * \param[in] train_model_path Model to be used to perform training that can be generated using the offline tooling library. - * \param[in] eval_model_path Model to be used to perform evaluation that can be generated using the offline tooling library. - * \param[in] optimizer_model_path Model to be used to the optimizer step for weight updates. The model can be generated using the offline tooling library. - * \param[out] out Created training session. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function creates a training session based on the env and session options provided that can + * begin or resume training from a given checkpoint state for the given onnx models. + * The checkpoint state represents the parameters of the training session which will be moved + * to the device specified by the user through the session options (if necessary). + * + * \param[in] env Environment to be used for the training session. + * \param[in] options Session options that the user can customize for this training session. + * \param[in] checkpoint_state Training states that the training session uses as a starting point for training. + * \param[in] train_model_path Model to be used to perform training that can be generated using the offline tooling library. + * \param[in] eval_model_path Model to be used to perform evaluation that can be generated using the offline tooling library. + * \param[in] optimizer_model_path Model to be used to the optimizer step for weight updates. The model can be generated using the offline tooling library. + * \param[out] out Created training session. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(CreateTrainingSession, _In_ const OrtEnv* env, _In_ const OrtSessionOptions* options, _Inout_ OrtCheckpointState* checkpoint_state, _In_ const ORTCHAR_T* train_model_path, _In_ const ORTCHAR_T* eval_model_path, _In_ const ORTCHAR_T* optimizer_model_path, _Outptr_ OrtTrainingSession** out); /** \brief Retrieves the number of user outputs in the training model. - * - * This function returns the number of outputs of the training model so that the user can - * allocate space for the number of outputs when TrainStep is invoked. - * - * \param[in] sess The training session which has working knowledge of the training model. - * \param[out] out Number of user outputs in the training model. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function returns the number of outputs of the training model so that the user can + * allocate space for the number of outputs when TrainStep is invoked. + * + * \param[in] sess The training session which owns the training model. + * \param[out] out Number of user outputs in the training model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(TrainingSessionGetTrainingModelOutputCount, _In_ const OrtTrainingSession* sess, _Out_ size_t* out); /** \brief Retrieves the number of user outputs in the eval model. - * - * This function returns the number of outputs of the eval model so that the user can - * allocate space for the number of outputs when EvalStep is invoked. - * - * \param[in] sess The training session which has working knowledge of the eval model. - * \param[out] out Number of user outputs in the eval model. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function returns the number of outputs of the eval model so that the user can + * allocate space for the number of outputs when EvalStep is invoked. + * + * \param[in] sess The training session which owns the eval model. + * \param[out] out Number of user outputs in the eval model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(TrainingSessionGetEvalModelOutputCount, _In_ const OrtTrainingSession* sess, _Out_ size_t* out); ORT_API2_STATUS(TrainingSessionGetTrainingModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); @@ -96,225 +96,280 @@ struct OrtTrainingApi { ORT_API2_STATUS(TrainingSessionGetEvalModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); /** \brief Reset the training model gradients to zero lazily. - * - * This function sets the internal state of the training session such that the training model gradients - * will be reset just before the new gradients are computed on the next invocation of TrainStep. - * - * \param[in] session The training session which has working knowledge of the eval model. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ - ORT_API2_STATUS(ResetGrad, _Inout_ OrtTrainingSession* session); + * + * This function sets the internal state of the training session such that the training model gradients + * will be scheduled to be reset just before the new gradients are computed on the next invocation + * of TrainStep. + * + * \param[in] session The training session which owns the eval model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ + ORT_API2_STATUS(LazyResetGrad, _Inout_ OrtTrainingSession* session); /** \brief Computes the outputs and the gradients for the training model for the given inputs - * - * This function performs a training step that computes the outputs and the gradients of the training model - * for the given inputs. The train step is performed based on the training model that was provided - * to the training session. - * The gradients computed are stored inside the training session so they can be later consumed - * by the OptimizerStep function. - * - * \param[in] sess The training session which has working knowledge of the eval model. - * \param[in] run_options Run options for this training step. - * \param[in] inputs_len Number of user inputs to the training model. - * \param[in] inputs The user inputs to the training model. - * \param[in] outputs_len Number of user outputs expected from this training step. - * \param[out] outputs User outputs computed by train step. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function performs a training step that computes the outputs and the gradients of the training model + * for the given inputs. The train step is performed based on the training model that was provided + * to the training session. + * The gradients computed are stored inside the training session so they can be later consumed + * by the OptimizerStep function. + * + * \param[in] sess The training session which owns the eval model. + * \param[in] run_options Run options for this training step. + * \param[in] inputs_len Number of user inputs to the training model. + * \param[in] inputs The user inputs to the training model. + * \param[in] outputs_len Number of user outputs expected from this training step. + * \param[out] outputs User outputs computed by train step. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(TrainStep, _Inout_ OrtTrainingSession* sess, _In_opt_ const OrtRunOptions* run_options, size_t inputs_len, _In_reads_(inputs_len) const OrtValue* const* inputs, size_t outputs_len, _Inout_updates_all_(outputs_len) OrtValue** outputs); /** \brief Computes the outputs for the eval model for the given inputs - * - * This function performs an eval step that computes the outputs of the eval model for the given inputs. - * The eval step is performed based on the eval model that was provided to the training session. - * - * \param[in] sess The training session which has working knowledge of the eval model. - * \param[in] run_options Run options for this eval step. - * \param[in] inputs_len Number of user inputs to the eval model. - * \param[in] inputs The user inputs to the eval model. - * \param[in] outputs_len Number of user outputs expected from this eval step. - * \param[out] outputs User outputs computed by eval step. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function performs an eval step that computes the outputs of the eval model for the given inputs. + * The eval step is performed based on the eval model that was provided to the training session. + * + * \param[in] sess The training session which owns the eval model. + * \param[in] run_options Run options for this eval step. + * \param[in] inputs_len Number of user inputs to the eval model. + * \param[in] inputs The user inputs to the eval model. + * \param[in] outputs_len Number of user outputs expected from this eval step. + * \param[out] outputs User outputs computed by eval step. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(EvalStep, _In_ const OrtTrainingSession* sess, _In_opt_ const OrtRunOptions* run_options, size_t inputs_len, _In_reads_(inputs_len) const OrtValue* const* inputs, size_t outputs_len, _Inout_updates_all_(outputs_len) OrtValue** outputs); /** \brief Sets the learning rate for this training session. - * - * This function allows users to set the learning rate for the training session. The current - * learning rate is maintained by the training session and can be overwritten by invoking - * this function with the desired learning rate. This function should not be used when a valid - * learning rate scheduler is registered. It should be used either to set the learning rate - * derived from a custom learning rate scheduler or to set the learning rate constant to be used - * throughout the training session. - * Please note that this function does not set the initial learning rate that may be needed - * by the predefined learning rate schedulers. To set the initial learning rate for learning - * rate schedulers, please look at the function `RegisterLinearLRScheduler`. - * - * \param[in] sess The training session on which the learning rate needs to be set. - * \param[in] learning_rate Desired learning rate to set. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function allows users to set the learning rate for the training session. The current + * learning rate is maintained by the training session and can be overwritten by invoking + * this function with the desired learning rate. This function should not be used when a valid + * learning rate scheduler is registered. It should be used either to set the learning rate + * derived from a custom learning rate scheduler or to set the learning rate constant to be used + * throughout the training session. + * Please note that this function does not set the initial learning rate that may be needed + * by the predefined learning rate schedulers. To set the initial learning rate for learning + * rate schedulers, please look at the function `RegisterLinearLRScheduler`. + * + * \param[in] sess The training session on which the learning rate needs to be set. + * \param[in] learning_rate Desired learning rate to set. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(SetLearningRate, _Inout_ OrtTrainingSession* sess, _In_ float learning_rate); /** \brief Gets the current learning rate for this training session. - * - * This function allows users to get the learning rate for the training session. The current - * learning rate is maintained by the training session - * - * \param[in] sess The training session on which the learning rate needs to be set. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function allows users to get the learning rate for the training session. The current + * learning rate is maintained by the training session + * + * \param[in] sess The training session on which the learning rate needs to be set. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(GetLearningRate, _Inout_ OrtTrainingSession* sess, _Out_ float* learning_rate); /** \brief Performs the weight updates for the trainable parameters using the optimizer model. - * - * This function performs the weight update step that updates the trainable parameters such that they - * take a step in the direction of their gradients. The optimizer step is performed based on the optimizer - * model that was provided to the training session. - * The updated parameters are stored inside the training session so that they can be used by the next - * TrainStep function call. - * - * \param[in] sess The training session which has working knowledge of the optimizer model. - * \param[in] run_options Run options for this eval step. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * This function performs the weight update step that updates the trainable parameters such that they + * take a step in the direction of their gradients. The optimizer step is performed based on the optimizer + * model that was provided to the training session. + * The updated parameters are stored inside the training session so that they can be used by the next + * TrainStep function call. + * + * \param[in] sess The training session which owns the optimizer model. + * \param[in] run_options Run options for this eval step. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(OptimizerStep, _Inout_ OrtTrainingSession* sess, _In_opt_ const OrtRunOptions* run_options); - /** \brief Registers the use of the Linear learning rate scheduler for the training session. - * - * Register a Linear learning rate scheduler with the given - * learning rate scheduler parameters. Optionally specify the initial learning rate - * that should be used with this learning rate scheduler and training session. - * - * \param[in] sess The training session that should use the linear learning rate scheduler. - * \param[in] warmup_step_count Warmup steps for LR warmup. - * \param[in] total_step_count Total step count. - * \param[in] initial_lr The initial learning rate to be used by the training session. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + /** \brief Registers a Linear learning rate scheduler for the training session. + * + * Register a linear learning rate scheduler that decays the learning rate by linearly updated + * multiplicative factor from the initial learning rate set on the training session to 0. The decay + * is performed after the initial warm up phase where the learning rate is linearly incremented + * from 0 to the initial learning rate provided. + * + * \param[in] sess The training session that should use the linear learning rate scheduler. + * \param[in] warmup_step_count Warmup steps for LR warmup. + * \param[in] total_step_count Total step count. + * \param[in] initial_lr The initial learning rate to be used by the training session. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(RegisterLinearLRScheduler, _Inout_ OrtTrainingSession* sess, _In_ const int64_t warmup_step_count, - _In_ const int64_t total_step_count, _In_ const float initial_lr); + _In_ const int64_t total_step_count, _In_ const float initial_lr); /** \brief Update the learning rate based on the registered learing rate scheduler. - * - * Takes a scheduler step that updates the learning rate that is being used by the training session. - * This function should typically be called before invoking the optimizer step for each round, - * or as determined necessary to update the learning rate being used by the training session. - * Please note that a valid predefined learning rate scheduler must be first registered to invoke this - * function. - * - * \param[in] sess The training session that has the registered learning rate scheduler. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * Takes a scheduler step that updates the learning rate that is being used by the training session. + * This function should typically be called before invoking the optimizer step for each round, + * or as determined necessary to update the learning rate being used by the training session. + * Please note that a valid predefined learning rate scheduler must be first registered to invoke this + * function. + * + * \param[in] sess The training session that has the registered learning rate scheduler. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(SchedulerStep, _Inout_ OrtTrainingSession* sess); /** \brief Retrieves the size of all the parameters. - * - * Calculates the size of all the parameters for the training session. - * When 'trainable_only' is true, the size is calculated for trainable params only. - * - * \param[in] sess The training session. - * \param[in] trainable_only Whether to skip non-trainable parameters - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * Calculates the size of all the parameters for the training session. + * When 'trainable_only' is true, the size is calculated for trainable params only. + * + * \param[in] sess The training session. + * \param[in] trainable_only Whether to skip non-trainable parameters + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(GetParametersSize, _Inout_ OrtTrainingSession* sess, _Out_ size_t* out, bool trainable_only); /** \brief Copy parameters onto contiguous buffer held by parameters_buffer - * - * The parameters_buffer has to be of the size given by GetParametersSize api call, - * with matching setting for 'trainable_only'. All the target parameters must be of the same - * datatype. The OrtValue must be pre-allocated onto - * the desired device. This is a complementary function to 'CopyBufferToParameters'. - * Parameter ordering is preserved. - * User is responsible for allocating/freeing the 'parameters_buffer'. - * - * \param[in] sess The training session. - * \param[in] trainable_only Whether to skip non-trainable parameters - * \param[out] parameters_buffer The pre-allocated OrtValue buffer to copy onto. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * The parameters_buffer has to be of the size given by GetParametersSize api call, + * with matching setting for 'trainable_only'. All the target parameters must be of the same + * datatype. The OrtValue must be pre-allocated onto + * the desired device. This is a complementary function to 'CopyBufferToParameters'. + * Parameter ordering is preserved. + * User is responsible for allocating/freeing the 'parameters_buffer'. + * + * \param[in] sess The training session. + * \param[in] trainable_only Whether to skip non-trainable parameters + * \param[out] parameters_buffer The pre-allocated OrtValue buffer to copy onto. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(CopyParametersToBuffer, _Inout_ OrtTrainingSession* sess, - _Inout_ OrtValue* parameters_buffer, bool trainable_only); + _Inout_ OrtValue* parameters_buffer, bool trainable_only); /** \brief Copy parameter values from contiguous buffer held by parameters_buffer onto parameters - * - * The parameters_buffer has to be of the size given by GetParametersSize api call, - * with matching setting for 'trainable_only'. All the target parameters must be of the same - * datatype. This is a complementary function to 'CopyBufferToParameters' - * and can be used to load updated buffer values onto the parameters. - * Parameter ordering is preserved. - * User is responsible for allocating/freeing the 'parameters_buffer'. - * - * \param[in] sess The training session. - * \param[in] trainable_only Whether to skip non-trainable parameters - * \param[out] parameters_buffer The pre-allocated OrtValue buffer to copy from. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * The parameters_buffer has to be of the size given by GetParametersSize api call, + * with matching setting for 'trainable_only'. All the target parameters must be of the same + * datatype. This is a complementary function to 'CopyBufferToParameters' + * and can be used to load updated buffer values onto the parameters. + * Parameter ordering is preserved. + * User is responsible for allocating/freeing the 'parameters_buffer'. + * + * \param[in] sess The training session. + * \param[in] trainable_only Whether to skip non-trainable parameters + * \param[out] parameters_buffer The pre-allocated OrtValue buffer to copy from. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(CopyBufferToParameters, _Inout_ OrtTrainingSession* sess, - _Inout_ OrtValue* parameters_buffer, bool trainable_only); + _Inout_ OrtValue* parameters_buffer, bool trainable_only); /** \brief Frees up the memory used up by the training session. - * - * This function frees up any memory that was allocated in the training session. The training - * session can no longer be used after this call. - * - */ + * + * This function frees up any memory that was allocated in the training session. The training + * session can no longer be used after this call. + * + */ ORT_CLASS_RELEASE(TrainingSession); /** \brief Frees up the memory used up by the checkpoint state. - * - * This function frees up any memory that was allocated in the checkpoint state. The checkpoint - * state can no longer be used after this call. - * - */ + * + * This function frees up any memory that was allocated in the checkpoint state. The checkpoint + * state can no longer be used after this call. + * + */ ORT_CLASS_RELEASE(CheckpointState); /** \brief Export a model that can be used for inferencing. - * - * If the training session was provided with an eval model, the training session can generate - * an inference model if it knows the inference graph outputs. The input inference graph outputs - * are used to prune the eval model so that the output model's outputs align with the provided outputs. - * The exported model is saved at the path provided and can be used for inferencing with InferenceSession. - * Note that the function re-loads the eval model from the path provided to CreateTrainingSession and expects - * that this path still be valid. - * - * \param[in] sess The training session. - * \param[in] inference_model_path Path where the inference model should be serialized to. - * - * \snippet{doc} snippets.dox OrtStatus Return Value - * - */ + * + * If the training session was provided with an eval model, the training session can generate + * an inference model if it knows the inference graph outputs. The input inference graph outputs + * are used to prune the eval model so that the output model's outputs align with the provided outputs. + * The exported model is saved at the path provided and can be used for inferencing with InferenceSession. + * Note that the function re-loads the eval model from the path provided to CreateTrainingSession and expects + * that this path still be valid. + * + * \param[in] sess The training session. + * \param[in] inference_model_path Path where the inference model should be serialized to. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ ORT_API2_STATUS(ExportModelForInferencing, _Inout_ OrtTrainingSession* sess, _In_ const ORTCHAR_T* inference_model_path, size_t graph_outputs_len, _In_reads_(graph_outputs_len) const char* const* graph_output_names); + + /** \brief Sets the seed used for random number generation in Onnxruntime. + * + * Use this to get deterministic results. + * + * \param[in] seed The seed to be set. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ + ORT_API2_STATUS(SetSeed, _In_ const int64_t seed); + + /** \brief Retrieves the number of user inputs in the training model. + * + * \param[in] sess The training session which owns the training model. + * \param[out] out Number of user inputs in the training model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ + ORT_API2_STATUS(TrainingSessionGetTrainingModelInputCount, _In_ const OrtTrainingSession* sess, _Out_ size_t* out); + + /** \brief Retrieves the number of user inputs in the eval model. + * + * \param[in] sess The training session which owns the eval model. + * \param[out] out Number of user inputs in the eval model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ + ORT_API2_STATUS(TrainingSessionGetEvalModelInputCount, _In_ const OrtTrainingSession* sess, _Out_ size_t* out); + + /** \brief Retrieves the name of the user input at given index in the training model. + * + * \param[in] sess The training session which owns the training model. + * \param[out] out Number of user inputs in the eval model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ + ORT_API2_STATUS(TrainingSessionGetTrainingModelInputName, _In_ const OrtTrainingSession* sess, size_t index, + _In_ OrtAllocator* allocator, _Outptr_ char** output); + + /** \brief Retrieves the name of the user input at given index in the eval model. + * + * \param[in] sess The training session which owns the training model. + * \param[out] out Number of user inputs in the eval model. + * + * \snippet{doc} snippets.dox OrtStatus Return Value + * + */ + ORT_API2_STATUS(TrainingSessionGetEvalModelInputName, _In_ const OrtTrainingSession* sess, size_t index, + _In_ OrtAllocator* allocator, _Outptr_ char** output); }; typedef struct OrtTrainingApi OrtTrainingApi; diff --git a/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_api.h b/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_api.h index d810085646..58a9b5ade8 100644 --- a/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_api.h +++ b/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_api.h @@ -89,10 +89,10 @@ class TrainingSession : public detail::Base { /** \brief Lazily resets the gradients of the trainable parameters. * - * Wraps OrtTrainingApi::ResetGrad + * Wraps OrtTrainingApi::LazyResetGrad * */ - void ResetGrad(); + void LazyResetGrad(); /** \brief Run the evaluation step returning results in an Ort allocated vector. * @@ -153,6 +153,8 @@ class TrainingSession : public detail::Base { const std::vector& graph_output_names); }; +void SetSeed(const int64_t seed); + } // namespace Ort #include "onnxruntime_training_cxx_inline.h" diff --git a/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_inline.h b/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_inline.h index 8eaa1dbc83..75452eab46 100644 --- a/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_inline.h +++ b/orttraining/orttraining/training_api/include/onnxruntime_training_cxx_inline.h @@ -39,8 +39,8 @@ inline std::vector TrainingSession::TrainStep(const std::vector& i return output_values; } -inline void TrainingSession::ResetGrad() { - ThrowOnError(GetTrainingApi().ResetGrad(p_)); +inline void TrainingSession::LazyResetGrad() { + ThrowOnError(GetTrainingApi().LazyResetGrad(p_)); } inline std::vector TrainingSession::EvalStep(const std::vector& input_values) { @@ -105,4 +105,8 @@ inline void TrainingSession::ExportModelForInferencing(const std::basic_string(std::shared_ptr)>& get_scheduler, - std::optional initial_lr); + float initial_lr); size_t GetTrainingModelOutputCount() const noexcept; @@ -43,6 +43,14 @@ class TrainingSession { std::string GetEvalModelOutputName(size_t index) const noexcept; + size_t GetTrainingModelInputCount() const noexcept; + + size_t GetEvalModelInputCount() const noexcept; + + std::string GetTrainingModelInputName(size_t index) const noexcept; + + std::string GetEvalModelInputName(size_t index) const noexcept; + Status TrainStep(const RunOptions& run_options, const std::vector& inputs, std::vector& fetches); @@ -51,7 +59,7 @@ class TrainingSession { const std::vector& inputs, std::vector& fetches) const; - Status ResetGrad(); + Status LazyResetGrad(); Status OptimizerStep(const RunOptions& run_options); diff --git a/orttraining/orttraining/training_api/module.cc b/orttraining/orttraining/training_api/module.cc index cbf90e5472..d4afd9ea31 100644 --- a/orttraining/orttraining/training_api/module.cc +++ b/orttraining/orttraining/training_api/module.cc @@ -181,6 +181,7 @@ Module::Module(const std::string& train_model_path_or_bytes, gradients_.resize(grad_input_names.size()); train_input_names_ = user_input_names; + train_user_input_count_ = user_input_names.size(); train_input_names_.insert(train_input_names_.end(), param_input_names.begin(), param_input_names.end()); train_input_names_.insert(train_input_names_.end(), grad_input_names.begin(), grad_input_names.end()); train_input_names_.insert(train_input_names_.end(), reset_grad_name.begin(), reset_grad_name.end()); @@ -276,6 +277,7 @@ Module::Module(const std::string& train_model_path_or_bytes, } } eval_input_names_ = eval_user_input_names; + eval_user_input_count_ = eval_user_input_names.size(); eval_input_names_.insert(eval_input_names_.end(), eval_param_input_names.begin(), eval_param_input_names.end()); // Keep a copy of the eval model path to be able to later export the model for inferencing. @@ -407,7 +409,7 @@ Status Module::CopyBufferToParameters(OrtValue& parameters_buffer, const bool tr return Status::OK(); } -Status Module::ResetGrad() { +Status Module::LazyResetGrad() { accumulate_gradient_ = false; return Status::OK(); } @@ -479,6 +481,28 @@ Status Module::ExportModelForInferencing(const std::string& inference_model_path return Status::OK(); } +size_t Module::GetTrainingModelInputCount() const noexcept { + return train_user_input_count_; +} + +size_t Module::GetEvalModelInputCount() const noexcept { + return eval_user_input_count_; +} + +std::string Module::GetTrainingModelInputName(size_t index) const { + ORT_ENFORCE(index < train_user_input_count_, + "Train input name index out of range. Expected in range [0-", train_user_input_count_, "). Actual: ", + index); + return train_input_names_.at(index); +} + +std::string Module::GetEvalModelInputName(size_t index) const { + ORT_ENFORCE(index < eval_user_input_count_, + "Eval input name index out of range. Expected in range [0-", eval_user_input_count_, "). Actual: ", + index); + return eval_input_names_.at(index); +} + } // namespace api } // namespace training } // namespace onnxruntime diff --git a/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc b/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc index 5cf3d990b3..35d3b86a9c 100644 --- a/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc +++ b/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc @@ -2,14 +2,15 @@ // Licensed under the MIT License. #include "orttraining/training_api/include/onnxruntime_training_c_api.h" +#include "core/common/string_helper.h" #include "core/framework/error_code_helper.h" +#include "core/framework/random_seed.h" +#include "core/session/abi_session_options_impl.h" #include "core/session/ort_apis.h" #include "core/session/ort_env.h" -#include "core/session/abi_session_options_impl.h" #include "orttraining/training_api/include/checkpoint.h" -#include "orttraining/training_api/include/training_session.h" #include "orttraining/training_api/include/ort_training_apis.h" -#include "core/common/string_helper.h" +#include "orttraining/training_api/include/training_session.h" namespace { @@ -100,10 +101,10 @@ ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetEvalModelOutputName, _In_ API_IMPL_END } -ORT_API_STATUS_IMPL(OrtTrainingApis::ResetGrad, _Inout_ OrtTrainingSession* session) { +ORT_API_STATUS_IMPL(OrtTrainingApis::LazyResetGrad, _Inout_ OrtTrainingSession* session) { API_IMPL_BEGIN auto train_session = reinterpret_cast(session); - ORT_API_RETURN_IF_STATUS_NOT_OK(train_session->ResetGrad()); + ORT_API_RETURN_IF_STATUS_NOT_OK(train_session->LazyResetGrad()); return nullptr; API_IMPL_END @@ -256,7 +257,7 @@ ORT_API_STATUS_IMPL(OrtTrainingApis::RegisterLinearLRScheduler, _Inout_ OrtTrain return std::make_unique( optimizer, warmup_step_count, total_step_count); }, - std::optional(initial_lr))); + initial_lr)); return status; API_IMPL_END @@ -370,6 +371,53 @@ ORT_API_STATUS_IMPL(OrtTrainingApis::ExportModelForInferencing, _Inout_ OrtTrain API_IMPL_END } +ORT_API_STATUS_IMPL(OrtTrainingApis::SetSeed, _In_ const int64_t seed) { + API_IMPL_BEGIN + + onnxruntime::utils::SetRandomSeed(seed); + return nullptr; + + API_IMPL_END +} + +ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetTrainingModelInputCount, _In_ const OrtTrainingSession* sess, + _Out_ size_t* out) { + API_IMPL_BEGIN + auto session = reinterpret_cast(sess); + *out = session->GetTrainingModelInputCount(); + return nullptr; + API_IMPL_END +} + +ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetEvalModelInputCount, _In_ const OrtTrainingSession* sess, + _Out_ size_t* out) { + API_IMPL_BEGIN + auto session = reinterpret_cast(sess); + *out = session->GetEvalModelInputCount(); + return nullptr; + API_IMPL_END +} + +ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetTrainingModelInputName, _In_ const OrtTrainingSession* sess, + size_t index, _In_ OrtAllocator* allocator, _Outptr_ char** output) { + API_IMPL_BEGIN + auto session = reinterpret_cast(sess); + std::string name = session->GetTrainingModelInputName(index); + *output = onnxruntime::StrDup(name, allocator); + return nullptr; + API_IMPL_END +} + +ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetEvalModelInputName, _In_ const OrtTrainingSession* sess, + size_t index, _In_ OrtAllocator* allocator, _Outptr_ char** output) { + API_IMPL_BEGIN + auto session = reinterpret_cast(sess); + std::string name = session->GetEvalModelInputName(index); + *output = onnxruntime::StrDup(name, allocator); + return nullptr; + API_IMPL_END +} + static constexpr OrtTrainingApi ort_training_api = { // NOTE: The C# bindings depend on the API order within this struct. Since Training APIs are not officially // released, it is OK to change the order here, however a corresponding matching change should also be done in the @@ -381,7 +429,7 @@ static constexpr OrtTrainingApi ort_training_api = { &OrtTrainingApis::TrainingSessionGetEvalModelOutputCount, &OrtTrainingApis::TrainingSessionGetTrainingModelOutputName, &OrtTrainingApis::TrainingSessionGetEvalModelOutputName, - &OrtTrainingApis::ResetGrad, + &OrtTrainingApis::LazyResetGrad, &OrtTrainingApis::TrainStep, &OrtTrainingApis::EvalStep, &OrtTrainingApis::SetLearningRate, @@ -395,6 +443,11 @@ static constexpr OrtTrainingApi ort_training_api = { &OrtTrainingApis::ReleaseTrainingSession, &OrtTrainingApis::ReleaseCheckpointState, &OrtTrainingApis::ExportModelForInferencing, + &OrtTrainingApis::SetSeed, + &OrtTrainingApis::TrainingSessionGetTrainingModelInputCount, + &OrtTrainingApis::TrainingSessionGetEvalModelInputCount, + &OrtTrainingApis::TrainingSessionGetTrainingModelInputName, + &OrtTrainingApis::TrainingSessionGetEvalModelInputName, }; ORT_API(const OrtTrainingApi*, OrtTrainingApis::GetTrainingApi, uint32_t) { diff --git a/orttraining/orttraining/training_api/training_session.cc b/orttraining/orttraining/training_api/training_session.cc index 654e94df34..e69a2e20b2 100644 --- a/orttraining/orttraining/training_api/training_session.cc +++ b/orttraining/orttraining/training_api/training_session.cc @@ -23,14 +23,12 @@ TrainingSession::TrainingSession(const Environment& session_env, Status TrainingSession::RegisterScheduler( const std::function(std::shared_ptr)>& get_scheduler, - std::optional initial_lr) { + float initial_lr) { ORT_RETURN_IF_NOT(optimizer_, "No optimizer session initialized."); scheduler_ = get_scheduler(optimizer_); ORT_RETURN_IF_NOT(scheduler_, "The provided instance of the learning rate scheduler is a nullptr."); - if (initial_lr.has_value()) { - ORT_RETURN_IF_ERROR(optimizer_->SetInitialLearningRate(initial_lr.value())); - } + ORT_RETURN_IF_ERROR(optimizer_->SetInitialLearningRate(initial_lr)); return Status::OK(); } @@ -51,6 +49,22 @@ std::string TrainingSession::GetEvalModelOutputName(size_t index) const noexcept return module_->GetEvalModelOutputName(index); } +size_t TrainingSession::GetTrainingModelInputCount() const noexcept { + return module_->GetTrainingModelInputCount(); +} + +size_t TrainingSession::GetEvalModelInputCount() const noexcept { + return module_->GetEvalModelInputCount(); +} + +std::string TrainingSession::GetTrainingModelInputName(size_t index) const noexcept { + return module_->GetTrainingModelInputName(index); +} + +std::string TrainingSession::GetEvalModelInputName(size_t index) const noexcept { + return module_->GetEvalModelInputName(index); +} + Status TrainingSession::TrainStep(const RunOptions&, const std::vector& inputs, std::vector& fetches) { @@ -63,8 +77,8 @@ Status TrainingSession::EvalStep(const RunOptions&, return module_->EvalStep(inputs, fetches); } -Status TrainingSession::ResetGrad() { - return module_->ResetGrad(); +Status TrainingSession::LazyResetGrad() { + return module_->LazyResetGrad(); } Status TrainingSession::OptimizerStep(const RunOptions&) {