exclude some objc files
This commit is contained in:
@@ -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"
|
||||
@@ -1,119 +0,0 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#import <Foundation/Foundation.h>
|
||||
#include <stdint.h>
|
||||
|
||||
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
|
||||
@@ -1,263 +0,0 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#import <Foundation/Foundation.h>
|
||||
#include <stdint.h>
|
||||
|
||||
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<ORTValue*>*)trainStepWithInputValues:(NSArray<ORTValue*>*)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<ORTValue*>*)evalStepWithInputValues:(NSArray<ORTValue*>*)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<NSString*>*)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<NSString*>*)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<NSString*>*)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<NSString*>*)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<NSString*>*)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
|
||||
@@ -1,111 +0,0 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#import "ort_checkpoint_internal.h"
|
||||
|
||||
#include <optional>
|
||||
#include <string>
|
||||
#include <variant>
|
||||
#import "cxx_api.h"
|
||||
|
||||
#import "error_utils.h"
|
||||
|
||||
NS_ASSUME_NONNULL_BEGIN
|
||||
|
||||
@implementation ORTCheckpoint {
|
||||
std::optional<Ort::CheckpointState> _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<std::string>(&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<int64_t>(&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<float>(&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
|
||||
@@ -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
|
||||
@@ -1,224 +0,0 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#import "ort_training_session_internal.h"
|
||||
|
||||
#import <vector>
|
||||
#import <optional>
|
||||
#import <string>
|
||||
|
||||
#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<Ort::TrainingSession> _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<std::string> evalPath = utils::toStdOptionalString(evalModelPath);
|
||||
std::optional<std::string> 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<ORTValue*>*)trainStepWithInputValues:(NSArray<ORTValue*>*)inputs
|
||||
error:(NSError**)error {
|
||||
try {
|
||||
std::vector<const OrtValue*> inputValues = utils::getWrappedCAPIOrtValues(inputs);
|
||||
|
||||
size_t outputCount;
|
||||
Ort::ThrowOnError(Ort::GetTrainingApi().TrainingSessionGetTrainingModelOutputCount(*_session, &outputCount));
|
||||
std::vector<OrtValue*> 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<ORTValue*>*)evalStepWithInputValues:(NSArray<ORTValue*>*)inputs
|
||||
error:(NSError**)error {
|
||||
try {
|
||||
// create vector of OrtValue from NSArray<ORTValue*> with same size as inputValues
|
||||
std::vector<const OrtValue*> inputValues = utils::getWrappedCAPIOrtValues(inputs);
|
||||
|
||||
size_t outputCount;
|
||||
Ort::ThrowOnError(Ort::GetTrainingApi().TrainingSessionGetEvalModelOutputCount(*_session, &outputCount));
|
||||
std::vector<OrtValue*> 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<NSString*>*)getTrainInputNamesWithError:(NSError**)error {
|
||||
try {
|
||||
std::vector<std::string> inputNames = [self CXXAPIOrtTrainingSession].InputNames(true);
|
||||
return utils::toNSStringNSArray(inputNames);
|
||||
}
|
||||
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||
}
|
||||
|
||||
- (nullable NSArray<NSString*>*)getTrainOutputNamesWithError:(NSError**)error {
|
||||
try {
|
||||
std::vector<std::string> outputNames = [self CXXAPIOrtTrainingSession].OutputNames(true);
|
||||
return utils::toNSStringNSArray(outputNames);
|
||||
}
|
||||
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||
}
|
||||
|
||||
- (nullable NSArray<NSString*>*)getEvalInputNamesWithError:(NSError**)error {
|
||||
try {
|
||||
std::vector<std::string> inputNames = [self CXXAPIOrtTrainingSession].InputNames(false);
|
||||
return utils::toNSStringNSArray(inputNames);
|
||||
}
|
||||
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||
}
|
||||
|
||||
- (nullable NSArray<NSString*>*)getEvalOutputNamesWithError:(NSError**)error {
|
||||
try {
|
||||
std::vector<std::string> 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<NSString*>*)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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user