From 354f6a56e4eb45515154318858603929b243e24f Mon Sep 17 00:00:00 2001 From: rachguo Date: Tue, 11 Jul 2023 11:02:44 -0700 Subject: [PATCH] exclude some objc files --- objectivec/include/onnxruntime_training.h | 9 - objectivec/include/ort_checkpoint.h | 119 ---------- objectivec/include/ort_training_session.h | 263 --------------------- objectivec/ort_checkpoint.mm | 111 --------- objectivec/ort_checkpoint_internal.h | 16 -- objectivec/ort_training_session.mm | 224 ------------------ objectivec/ort_training_session_internal.h | 16 -- 7 files changed, 758 deletions(-) delete mode 100644 objectivec/include/onnxruntime_training.h delete mode 100644 objectivec/include/ort_checkpoint.h delete mode 100644 objectivec/include/ort_training_session.h delete mode 100644 objectivec/ort_checkpoint.mm delete mode 100644 objectivec/ort_checkpoint_internal.h delete mode 100644 objectivec/ort_training_session.mm delete mode 100644 objectivec/ort_training_session_internal.h diff --git a/objectivec/include/onnxruntime_training.h b/objectivec/include/onnxruntime_training.h deleted file mode 100644 index 504447e..0000000 --- a/objectivec/include/onnxruntime_training.h +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -// this header contains the entire ONNX Runtime training Objective-C API -// the headers below can also be imported individually - -#import "onnxruntime.h" -#import "ort_checkpoint.h" -#import "ort_training_session.h" diff --git a/objectivec/include/ort_checkpoint.h b/objectivec/include/ort_checkpoint.h deleted file mode 100644 index 85e5844..0000000 --- a/objectivec/include/ort_checkpoint.h +++ /dev/null @@ -1,119 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#import -#include - -NS_ASSUME_NONNULL_BEGIN - -/** - * An ORT checkpoint is a snapshot of the state of a model at a given point in time. - * - * This class holds the entire training session state that includes model parameters, - * their gradients, optimizer parameters, and user properties. The `ORTTrainingSession` leverages the - * `ORTCheckpoint` by accessing and updating the contained training state. - * - * Available since 1.16. - * - * @note This class is only available when the training APIs are enabled. - */ -@interface ORTCheckpoint : NSObject - -- (instancetype)init NS_UNAVAILABLE; - -/** - * Creates a checkpoint from directory on disk. - * - * @param path The path to the checkpoint directory. - * @param error Optional error information set if an error occurs. - * @return The instance, or nil if an error occurs. - * - * @warning The construction of the checkpoint state requires instantiation of `ORTEnv`. - * The intialization will fail if the `ORTEnv` is not properly initialized. - */ -- (nullable instancetype)initWithPath:(NSString*)path - error:(NSError**)error NS_DESIGNATED_INITIALIZER; - -/** - * Saves a checkpoint to directory on disk. - * - * @param path The path to the checkpoint directory. - * @param includeOptimizerState Flag to indicate whether to save the optimizer state or not. - * @param error Optional error information set if an error occurs. - * @return Whether the checkpoint was saved successfully. - */ -- (BOOL)saveCheckpointToPath:(NSString*)path - withOptimizerState:(BOOL)includeOptimizerState - error:(NSError**)error; - -/** - * Adds an int property to this checkpoint. - * - * @param name The name of the property. - * @param value The value of the property. - * @param error Optional error information set if an error occurs. - * @return Whether the property was added successfully. - */ -- (BOOL)addIntPropertyWithName:(NSString*)name - value:(int64_t)value - error:(NSError**)error; - -/** - * Adds a float property to this checkpoint. - * - * @param name The name of the property. - * @param value The value of the property. - * @param error Optional error information set if an error occurs. - * @return Whether the property was added successfully. - */ -- (BOOL)addFloatPropertyWithName:(NSString*)name - value:(float)value - error:(NSError**)error; - -/** - * Adds a string property to this checkpoint. - * - * @param name The name of the property. - * @param value The value of the property. - * @param error Optional error information set if an error occurs. - * @return Whether the property was added successfully. - */ - -- (BOOL)addStringPropertyWithName:(NSString*)name - value:(NSString*)value - error:(NSError**)error; - -/** - * Gets an int property from this checkpoint. - * - * @param name The name of the property. - * @param error Optional error information set if an error occurs. - * @return The value of the property or 0 if an error occurs. - */ -- (int64_t)getIntPropertyWithName:(NSString*)name - error:(NSError**)error __attribute__((swift_error(nonnull_error))); - -/** - * Gets a float property from this checkpoint. - * - * @param name The name of the property. - * @param error Optional error information set if an error occurs. - * @return The value of the property or 0.0f if an error occurs. - */ -- (float)getFloatPropertyWithName:(NSString*)name - error:(NSError**)error __attribute__((swift_error(nonnull_error))); - -/** - * - * Gets a string property from this checkpoint. - * - * @param name The name of the property. - * @param error Optional error information set if an error occurs. - * @return The value of the property. - */ -- (nullable NSString*)getStringPropertyWithName:(NSString*)name - error:(NSError**)error __attribute__((swift_error(nonnull_error))); - -@end - -NS_ASSUME_NONNULL_END diff --git a/objectivec/include/ort_training_session.h b/objectivec/include/ort_training_session.h deleted file mode 100644 index 15c0137..0000000 --- a/objectivec/include/ort_training_session.h +++ /dev/null @@ -1,263 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#import -#include - -NS_ASSUME_NONNULL_BEGIN - -@class ORTCheckpoint; -@class ORTEnv; -@class ORTValue; -@class ORTSessionOptions; - -/** - * Trainer class that provides methods to train, evaluate and optimize ONNX models. - * - * The training session requires four training artifacts: - * 1. Training onnx model - * 2. Evaluation onnx model (optional) - * 3. Optimizer onnx model - * 4. Checkpoint directory - * - * [onnxruntime-training python utility](https://github.com/microsoft/onnxruntime/blob/main/orttraining/orttraining/python/training/onnxblock/README.md) - * can be used to generate above training artifacts. - * - * Available since 1.16. - * - * @note This class is only available when the training APIs are enabled. - */ -@interface ORTTrainingSession : NSObject - -- (instancetype)init NS_UNAVAILABLE; - -/** - * Creates a training session from the training artifacts that can be used to begin or resume training. - * - * The initializer instantiates the training session based on provided env and session options, which can be used to - * begin or resume training from a given checkpoint state. The checkpoint state represents the parameters of training - * session which will be moved to the device specified in the session option if needed. - * - * @param env The `ORTEnv` instance to use for the training session. - * @param sessionOptions The `ORTSessionOptions` to use for the training session. - * @param checkpoint Training states that are used as a starting point for training. - * @param trainModelPath The path to the training onnx model. - * @param evalModelPath The path to the evaluation onnx model. - * @param optimizerModelPath The path to the optimizer onnx model used to perform gradient descent. - * @param error Optional error information set if an error occurs. - * @return The instance, or nil if an error occurs. - * - * @note Note that the training session created with a checkpoint state uses this state to store the entire training - * state (including model parameters, its gradients, the optimizer states and the properties). The training session - * keeps a strong (owning) pointer to the checkpoint state. - */ -- (nullable instancetype)initWithEnv:(ORTEnv*)env - sessionOptions:(ORTSessionOptions*)sessionOptions - checkpoint:(ORTCheckpoint*)checkpoint - trainModelPath:(NSString*)trainModelPath - evalModelPath:(nullable NSString*)evalModelPath - optimizerModelPath:(nullable NSString*)optimizerModelPath - error:(NSError**)error NS_DESIGNATED_INITIALIZER; - -/** - * Performs a training step, which is equivalent to a forward and backward propagation in a single step. - * - * The training step computes the outputs of the training model and the gradients of the trainable parameters - * for the given input values. The train step is performed based on the training model that was provided to the training session. - * It is equivalent to running forward and backward propagation in a single step. The computed gradients are stored inside - * the training session state so they can be later consumed by `optimizerStep`. The gradients can be lazily reset by - * calling `lazyResetGrad` method. - * - * @param inputs The input values to the training model. - * @param error Optional error information set if an error occurs. - * @return The output values of the training model. - */ -- (nullable NSArray*)trainStepWithInputValues:(NSArray*)inputs - error:(NSError**)error; - -/** - * Performs a evaluation step that computes the outputs of the evaluation model for the given inputs. - * The eval step is performed based on the evaluation model that was provided to the training session. - * - * @param inputs The input values to the eval model. - * @param error Optional error information set if an error occurs. - * @return The output values of the eval model. - * - */ -- (nullable NSArray*)evalStepWithInputValues:(NSArray*)inputs - error:(NSError**)error; - -/** - * Reset the gradients of all trainable parameters to zero lazily. - * - * Calling this method sets the internal state of the training session such that the gradients of the trainable parameters - * in the ORTCheckpoint will be scheduled to be reset just before the new gradients are computed on the next - * invocation of the `trainStep` method. - * - * @param error Optional error information set if an error occurs. - * @return YES if the gradients are set to reset successfully, NO otherwise. - */ -- (BOOL)lazyResetGradWithError:(NSError**)error; - -/** - * Performs the weight updates for the trainable parameters using the optimizer model. 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 state so that they can be used by the next `trainStep` method call. - * - * @param error Optional error information set if an error occurs. - * @return YES if the optimizer step was performed successfully, NO otherwise. - */ -- (BOOL)optimizerStepWithError:(NSError**)error; - -/** - * Returns the names of the user inputs for the training model that can be associated with - * the `ORTValue` provided to the `trainStep`. - * - * @param error Optional error information set if an error occurs. - * @return The names of the user inputs for the training model. - */ -- (nullable NSArray*)getTrainInputNamesWithError:(NSError**)error; - -/** - * Returns the names of the user inputs for the evaluation model that can be associated with - * the `ORTValue` provided to the `evalStep`. - * - * @param error Optional error information set if an error occurs. - * @return The names of the user inputs for the evaluation model. - */ -- (nullable NSArray*)getEvalInputNamesWithError:(NSError**)error; - -/** - * Returns the names of the user outputs for the training model that can be associated with - * the `ORTValue` returned by the `trainStep`. - * - * @param error Optional error information set if an error occurs. - * @return The names of the user outputs for the training model. - */ -- (nullable NSArray*)getTrainOutputNamesWithError:(NSError**)error; - -/** - * Returns the names of the user outputs for the evaluation model that can be associated with - * the `ORTValue` returned by the `evalStep`. - * - * @param error Optional error information set if an error occurs. - * @return The names of the user outputs for the evaluation model. - */ -- (nullable NSArray*)getEvalOutputNamesWithError:(NSError**)error; - -/** - * Registers a linear learning rate scheduler for the training session. - * - * The scheduler gradually decreases the learning rate from the initial value to zero over the course of the training. - * The decrease is performed by multiplying the current learning rate by a linearly updated factor. - * Before the decrease, the learning rate is gradually increased from zero to the initial value during a warmup phase. - * - * @param warmupStepCount The number of steps to perform the linear warmup. - * @param totalStepCount The total number of steps to perform the linear decay. - * @param initialLr The initial learning rate. - * @param error Optional error information set if an error occurs. - * @return YES if the scheduler was registered successfully, NO otherwise. - */ -- (BOOL)registerLinearLRSchedulerWithWarmupStepCount:(int64_t)warmupStepCount - totalStepCount:(int64_t)totalStepCount - initialLr:(float)initialLr - error:(NSError**)error; - -/** - * Update the learning rate based on the registered learning rate scheduler. - * - * Performs 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 necessary - * to update the learning rate being used by the training session. - * - * @note A valid predefined learning rate scheduler must be first registered to invoke this method. - * - * @param error Optional error information set if an error occurs. - * @return YES if the scheduler step was performed successfully, NO otherwise. - */ -- (BOOL)schedulerStepWithError:(NSError**)error; - -/** - * Returns the current learning rate being used by the training session. - * - * @param error Optional error information set if an error occurs. - * @return The current learning rate or 0.0f if an error occurs. - */ -- (float)getLearningRateWithError:(NSError**)error __attribute__((swift_error(nonnull_error))); - -/** - * Sets the learning rate being used by the training session. - * - * The current learning rate is maintained by the training session and can be overwritten by invoking this method - * 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 a constant - * learning rate to be used throughout the training session. - * - * @note It 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, use the `registerLinearLRScheduler` method. - * - * @param lr The learning rate to be used by the training session. - * @param error Optional error information set if an error occurs. - * @return YES if the learning rate was set successfully, NO otherwise. - */ -- (BOOL)setLearningRate:(float)lr - error:(NSError**)error; - -/** - * Loads the training session model parameters from a contiguous buffer. - * - * @param buffer Contiguous buffer to load the parameters from. - * @param error Optional error information set if an error occurs. - * @return YES if the parameters were loaded successfully, NO otherwise. - */ -- (BOOL)fromBufferWithValue:(ORTValue*)buffer - error:(NSError**)error; - -/** - * Returns a contiguous buffer that holds a copy of all training state parameters. - * - * @param onlyTrainable If YES, returns a buffer that holds only the trainable parameters, otherwise returns a buffer - * that holds all the parameters. - * @param error Optional error information set if an error occurs. - * @return A contiguous buffer that holds a copy of all training state parameters. - */ -- (nullable ORTValue*)toBufferWithTrainable:(BOOL)onlyTrainable - error:(NSError**)error; - -/** - * Exports the training session model that can be used for inference. - * - * 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 - * inference model's outputs align with the provided outputs. The exported model is saved at the path provided and - * can be used for inferencing with `ORTSession`. - * - * @note The method reloads the eval model from the path provided to the initializer and expects this path to be valid. - * - * @param inferenceModelPath The path to the serialized the inference model. - * @param graphOutputNames The names of the outputs that are needed in the inference model. - * @param error Optional error information set if an error occurs. - * @return YES if the inference model was exported successfully, NO otherwise. - */ -- (BOOL)exportModelForInferenceWithOutputPath:(NSString*)inferenceModelPath - graphOutputNames:(NSArray*)graphOutputNames - error:(NSError**)error; -@end - -#ifdef __cplusplus -extern "C" { -#endif - -/** - * This function sets the seed for generating random numbers. - * Use this function to generate reproducible results. It should be noted that completely reproducible results are not guaranteed. - * - * @param seed Manually set seed to use for random number generation. - */ -void ORTSetSeed(int64_t seed); - -#ifdef __cplusplus -} -#endif - -NS_ASSUME_NONNULL_END diff --git a/objectivec/ort_checkpoint.mm b/objectivec/ort_checkpoint.mm deleted file mode 100644 index 1238645..0000000 --- a/objectivec/ort_checkpoint.mm +++ /dev/null @@ -1,111 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#import "ort_checkpoint_internal.h" - -#include -#include -#include -#import "cxx_api.h" - -#import "error_utils.h" - -NS_ASSUME_NONNULL_BEGIN - -@implementation ORTCheckpoint { - std::optional _checkpoint; -} - -- (nullable instancetype)initWithPath:(NSString*)path - error:(NSError**)error { - if ((self = [super init]) == nil) { - return nil; - } - - try { - _checkpoint = Ort::CheckpointState::LoadCheckpoint(path.UTF8String); - return self; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (BOOL)saveCheckpointToPath:(NSString*)path - withOptimizerState:(BOOL)includeOptimizerState - error:(NSError**)error { - try { - Ort::CheckpointState::SaveCheckpoint([self CXXAPIOrtCheckpoint], path.UTF8String, includeOptimizerState); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (BOOL)addIntPropertyWithName:(NSString*)name - value:(int64_t)value - error:(NSError**)error { - try { - [self CXXAPIOrtCheckpoint].AddProperty(name.UTF8String, value); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (BOOL)addFloatPropertyWithName:(NSString*)name - value:(float)value - error:(NSError**)error { - try { - [self CXXAPIOrtCheckpoint].AddProperty(name.UTF8String, value); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (BOOL)addStringPropertyWithName:(NSString*)name - value:(NSString*)value - error:(NSError**)error { - try { - [self CXXAPIOrtCheckpoint].AddProperty(name.UTF8String, value.UTF8String); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (nullable NSString*)getStringPropertyWithName:(NSString*)name error:(NSError**)error { - try { - Ort::Property value = [self CXXAPIOrtCheckpoint].GetProperty(name.UTF8String); - if (std::string* str = std::get_if(&value)) { - return [NSString stringWithUTF8String:str->c_str()]; - } - ORT_CXX_API_THROW("Property is not a string.", ORT_INVALID_ARGUMENT); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (int64_t)getIntPropertyWithName:(NSString*)name error:(NSError**)error { - try { - Ort::Property value = [self CXXAPIOrtCheckpoint].GetProperty(name.UTF8String); - if (int64_t* i = std::get_if(&value)) { - return *i; - } - ORT_CXX_API_THROW("Property is not an integer.", ORT_INVALID_ARGUMENT); - } - ORT_OBJC_API_IMPL_CATCH(error, 0) -} - -- (float)getFloatPropertyWithName:(NSString*)name error:(NSError**)error { - try { - Ort::Property value = [self CXXAPIOrtCheckpoint].GetProperty(name.UTF8String); - if (float* f = std::get_if(&value)) { - return *f; - } - ORT_CXX_API_THROW("Property is not a float.", ORT_INVALID_ARGUMENT); - } - ORT_OBJC_API_IMPL_CATCH(error, 0.0f) -} - -- (Ort::CheckpointState&)CXXAPIOrtCheckpoint { - return *_checkpoint; -} - -@end - -NS_ASSUME_NONNULL_END diff --git a/objectivec/ort_checkpoint_internal.h b/objectivec/ort_checkpoint_internal.h deleted file mode 100644 index 3d1550c..0000000 --- a/objectivec/ort_checkpoint_internal.h +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#import "ort_checkpoint.h" - -#import "cxx_api.h" - -NS_ASSUME_NONNULL_BEGIN - -@interface ORTCheckpoint () - -- (Ort::CheckpointState&)CXXAPIOrtCheckpoint; - -@end - -NS_ASSUME_NONNULL_END diff --git a/objectivec/ort_training_session.mm b/objectivec/ort_training_session.mm deleted file mode 100644 index 285151b..0000000 --- a/objectivec/ort_training_session.mm +++ /dev/null @@ -1,224 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#import "ort_training_session_internal.h" - -#import -#import -#import - -#import "cxx_api.h" -#import "cxx_utils.h" -#import "error_utils.h" -#import "ort_checkpoint_internal.h" -#import "ort_session_internal.h" -#import "ort_enums_internal.h" -#import "ort_env_internal.h" -#import "ort_value_internal.h" - -NS_ASSUME_NONNULL_BEGIN - -@implementation ORTTrainingSession { - std::optional _session; - ORTCheckpoint* _checkpoint; -} - -- (Ort::TrainingSession&)CXXAPIOrtTrainingSession { - return *_session; -} - -- (nullable instancetype)initWithEnv:(ORTEnv*)env - sessionOptions:(ORTSessionOptions*)sessionOptions - checkpoint:(ORTCheckpoint*)checkpoint - trainModelPath:(NSString*)trainModelPath - evalModelPath:(nullable NSString*)evalModelPath - optimizerModelPath:(nullable NSString*)optimizerModelPath - error:(NSError**)error { - if ((self = [super init]) == nil) { - return nil; - } - - try { - std::optional evalPath = utils::toStdOptionalString(evalModelPath); - std::optional optimizerPath = utils::toStdOptionalString(optimizerModelPath); - - _checkpoint = checkpoint; - _session = Ort::TrainingSession{ - [env CXXAPIOrtEnv], - [sessionOptions CXXAPIOrtSessionOptions], - [checkpoint CXXAPIOrtCheckpoint], - trainModelPath.UTF8String, - evalPath, - optimizerPath}; - return self; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (nullable NSArray*)trainStepWithInputValues:(NSArray*)inputs - error:(NSError**)error { - try { - std::vector inputValues = utils::getWrappedCAPIOrtValues(inputs); - - size_t outputCount; - Ort::ThrowOnError(Ort::GetTrainingApi().TrainingSessionGetTrainingModelOutputCount(*_session, &outputCount)); - std::vector outputValues(outputCount, nullptr); - - Ort::RunOptions runOptions; - Ort::ThrowOnError(Ort::GetTrainingApi().TrainStep( - *_session, - runOptions, - inputValues.size(), - inputValues.data(), - outputValues.size(), - outputValues.data())); - - return utils::wrapUnownedCAPIOrtValues(outputValues, error); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} -- (nullable NSArray*)evalStepWithInputValues:(NSArray*)inputs - error:(NSError**)error { - try { - // create vector of OrtValue from NSArray with same size as inputValues - std::vector inputValues = utils::getWrappedCAPIOrtValues(inputs); - - size_t outputCount; - Ort::ThrowOnError(Ort::GetTrainingApi().TrainingSessionGetEvalModelOutputCount(*_session, &outputCount)); - std::vector outputValues(outputCount, nullptr); - - Ort::RunOptions runOptions; - Ort::ThrowOnError(Ort::GetTrainingApi().EvalStep( - *_session, - runOptions, - inputValues.size(), - inputValues.data(), - outputValues.size(), - outputValues.data())); - - return utils::wrapUnownedCAPIOrtValues(outputValues, error); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (BOOL)lazyResetGradWithError:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].LazyResetGrad(); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (BOOL)optimizerStepWithError:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].OptimizerStep(); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (nullable NSArray*)getTrainInputNamesWithError:(NSError**)error { - try { - std::vector inputNames = [self CXXAPIOrtTrainingSession].InputNames(true); - return utils::toNSStringNSArray(inputNames); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (nullable NSArray*)getTrainOutputNamesWithError:(NSError**)error { - try { - std::vector outputNames = [self CXXAPIOrtTrainingSession].OutputNames(true); - return utils::toNSStringNSArray(outputNames); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (nullable NSArray*)getEvalInputNamesWithError:(NSError**)error { - try { - std::vector inputNames = [self CXXAPIOrtTrainingSession].InputNames(false); - return utils::toNSStringNSArray(inputNames); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (nullable NSArray*)getEvalOutputNamesWithError:(NSError**)error { - try { - std::vector outputNames = [self CXXAPIOrtTrainingSession].OutputNames(false); - return utils::toNSStringNSArray(outputNames); - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (BOOL)registerLinearLRSchedulerWithWarmupStepCount:(int64_t)warmupStepCount - totalStepCount:(int64_t)totalStepCount - initialLr:(float)initialLr - error:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].RegisterLinearLRScheduler(warmupStepCount, totalStepCount, initialLr); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (BOOL)schedulerStepWithError:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].SchedulerStep(); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (float)getLearningRateWithError:(NSError**)error { - try { - return [self CXXAPIOrtTrainingSession].GetLearningRate(); - } - ORT_OBJC_API_IMPL_CATCH(error, 0.0f); -} - -- (BOOL)setLearningRate:(float)lr - error:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].SetLearningRate(lr); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (BOOL)fromBufferWithValue:(ORTValue*)buffer - error:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].FromBuffer([buffer CXXAPIOrtValue]); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -- (nullable ORTValue*)toBufferWithTrainable:(BOOL)onlyTrainable - error:(NSError**)error { - try { - Ort::Value val = [self CXXAPIOrtTrainingSession].ToBuffer(onlyTrainable); - return [[ORTValue alloc] initWithCXXAPIOrtValue:std::move(val) - externalTensorData:nil - error:error]; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) -} - -- (BOOL)exportModelForInferenceWithOutputPath:(NSString*)inferenceModelPath - graphOutputNames:(NSArray*)graphOutputNames - error:(NSError**)error { - try { - [self CXXAPIOrtTrainingSession].ExportModelForInferencing(utils::toStdString(inferenceModelPath), - utils::toStdStringVector(graphOutputNames)); - return YES; - } - ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) -} - -@end - -void ORTSetSeed(int64_t seed) { - Ort::SetSeed(seed); -} - -NS_ASSUME_NONNULL_END diff --git a/objectivec/ort_training_session_internal.h b/objectivec/ort_training_session_internal.h deleted file mode 100644 index 453c941..0000000 --- a/objectivec/ort_training_session_internal.h +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#import "ort_training_session.h" - -#import "cxx_api.h" - -NS_ASSUME_NONNULL_BEGIN - -@interface ORTTrainingSession () - -- (Ort::TrainingSession&)CXXAPIOrtTrainingSession; - -@end - -NS_ASSUME_NONNULL_END