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);