From a7bc927b4b1f4234b44f902b213db09fde0e8d3b Mon Sep 17 00:00:00 2001 From: Ashwini Khade Date: Fri, 9 Dec 2022 16:01:11 -0800 Subject: [PATCH] fix typos in training apis (#13908) ### Description This PR fixes some typos in the training apis. We need to add more tests and make sure they are all run on the CIs to capture such issues. These changes are out of scope of this PR. ### Motivation and Context Co-authored-by: Ashwini Khade --- .../Training/NativeTrainingMethods.shared.cs | 14 +++++++------- .../include/onnxruntime_training_c_api.h | 4 ++-- .../training_api/include/ort_training_apis.h | 6 +++--- .../training_api/onnxruntime_training_c_api.cc | 4 ++-- 4 files changed, 14 insertions(+), 14 deletions(-) diff --git a/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs b/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs index 3054b229af..2e9a9c5498 100644 --- a/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs +++ b/csharp/src/Microsoft.ML.OnnxRuntime/Training/NativeTrainingMethods.shared.cs @@ -133,21 +133,21 @@ namespace Microsoft.ML.OnnxRuntime [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr /*(OrtStatus*)*/ DOrtGetTrainingModelOutputCount( - IntPtr /*(OrtSession*)*/ session, + IntPtr /*(OrtTrainingSession*)*/ session, out UIntPtr count); public static DOrtGetTrainingModelOutputCount OrtGetTrainingModelOutputCount; [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr /*(OrtStatus*)*/ DOrtGetEvalModelOutputCount( - IntPtr /*(OrtSession*)*/ session, + IntPtr /*(OrtTrainingSession*)*/ session, out UIntPtr count); public static DOrtGetEvalModelOutputCount OrtGetEvalModelOutputCount; [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr /*(OrtStatus*)*/ DOrtGetTrainingModelOutputName( - IntPtr /*(OrtSession*)*/ session, + IntPtr /*(OrtTrainingSession*)*/ session, UIntPtr index, IntPtr /*(OrtAllocator*)*/ allocator, out IntPtr /*(char**)*/name); @@ -156,7 +156,7 @@ namespace Microsoft.ML.OnnxRuntime [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr /*(OrtStatus*)*/ DOrtGetEvalModelOutputName( - IntPtr /*(OrtSession*)*/ session, + IntPtr /*(OrtTrainingSession*)*/ session, UIntPtr index, IntPtr /*(OrtAllocator*)*/ allocator, out IntPtr /*(char**)*/name); @@ -165,7 +165,7 @@ namespace Microsoft.ML.OnnxRuntime [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr /*(OrtStatus*)*/ DOrtResetGrad( - IntPtr /*(OrtSession*)*/ session); + IntPtr /*(OrtTrainingSession*)*/ session); public static DOrtResetGrad OrtResetGrad; @@ -233,11 +233,11 @@ namespace Microsoft.ML.OnnxRuntime public static DOrtSchedulerStep OrtSchedulerStep; [UnmanagedFunctionPointer(CallingConvention.Winapi)] - public delegate void DOrtReleaseTrainingSession(IntPtr /*(OrtSession*)*/session); + public delegate void DOrtReleaseTrainingSession(IntPtr /*(OrtTrainingSession*)*/session); public static DOrtReleaseTrainingSession OrtReleaseTrainingSession; [UnmanagedFunctionPointer(CallingConvention.Winapi)] - public delegate void DOrtReleaseCheckpointState(IntPtr /*(OrtSession*)*/session); + public delegate void DOrtReleaseCheckpointState(IntPtr /*(OrtCheckpointState*)*/checkpointState); public static DOrtReleaseCheckpointState OrtReleaseCheckpointState; #endregion TrainingSession 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 2e14d082e9..20c0dc0bd8 100644 --- a/orttraining/orttraining/training_api/include/onnxruntime_training_c_api.h +++ b/orttraining/orttraining/training_api/include/onnxruntime_training_c_api.h @@ -91,9 +91,9 @@ struct OrtTrainingApi { */ ORT_API2_STATUS(TrainingSessionGetEvalModelOutputCount, _In_ const OrtTrainingSession* sess, _Out_ size_t* out); - ORT_API2_STATUS(TrainingSessionGetTrainingModelOutputName, _In_ const OrtSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); + ORT_API2_STATUS(TrainingSessionGetTrainingModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); - ORT_API2_STATUS(TrainingSessionGetEvalModelOutputName, _In_ const OrtSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); + 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. * diff --git a/orttraining/orttraining/training_api/include/ort_training_apis.h b/orttraining/orttraining/training_api/include/ort_training_apis.h index 833989a1a8..778c892269 100644 --- a/orttraining/orttraining/training_api/include/ort_training_apis.h +++ b/orttraining/orttraining/training_api/include/ort_training_apis.h @@ -14,10 +14,10 @@ ORT_API_STATUS_IMPL(TrainingSessionGetTrainingModelOutputCount, _In_ const OrtTr ORT_API_STATUS_IMPL(TrainingSessionGetEvalModelOutputCount, _In_ const OrtTrainingSession* sess, _Out_ size_t* out); -ORT_API_STATUS_IMPL(TrainingSessionGetTrainingModelOutputName, _In_ const OrtSession* sess, size_t index, +ORT_API_STATUS_IMPL(TrainingSessionGetTrainingModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); -ORT_API_STATUS_IMPL(TrainingSessionGetEvalModelOutputName, _In_ const OrtSession* sess, size_t index, +ORT_API_STATUS_IMPL(TrainingSessionGetEvalModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output); ORT_API_STATUS_IMPL(ResetGrad, _Inout_ OrtTrainingSession* session); @@ -56,7 +56,7 @@ ORT_API_STATUS_IMPL(CopyParametersToBuffer, _Inout_ OrtTrainingSession* sess, ORT_API_STATUS_IMPL(CopyBufferToParameters, _Inout_ OrtTrainingSession* sess, _Inout_ OrtValue* parameters_buffer, bool trainable_only); -ORT_API(void, ReleaseCheckpointState, _Frees_ptr_opt_ OrtCheckpointState* session); +ORT_API(void, ReleaseCheckpointState, _Frees_ptr_opt_ OrtCheckpointState* checkpoint_state); ORT_API(void, ReleaseTrainingSession, _Frees_ptr_opt_ OrtTrainingSession* session); diff --git a/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc b/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc index c6eff3a8f6..5cf3d990b3 100644 --- a/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc +++ b/orttraining/orttraining/training_api/onnxruntime_training_c_api.cc @@ -80,7 +80,7 @@ ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetEvalModelOutputCount, _In API_IMPL_END } -ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetTrainingModelOutputName, _In_ const OrtSession* sess, size_t index, +ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetTrainingModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output) { API_IMPL_BEGIN auto session = reinterpret_cast(sess); @@ -90,7 +90,7 @@ ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetTrainingModelOutputName, API_IMPL_END } -ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetEvalModelOutputName, _In_ const OrtSession* sess, size_t index, +ORT_API_STATUS_IMPL(OrtTrainingApis::TrainingSessionGetEvalModelOutputName, _In_ const OrtTrainingSession* sess, size_t index, _Inout_ OrtAllocator* allocator, _Outptr_ char** output) { API_IMPL_BEGIN auto session = reinterpret_cast(sess);