pull swift/ and objc/ folder contents
This commit is contained in:
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,2 @@
|
|||||||
|
NOTE: Flat directory structure to work with both the Objective-C build and the Swift Package Manager build which is done
|
||||||
|
via ../Package.swift
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
static_assert(__has_feature(objc_arc), "Objective-C ARC must be enabled.");
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
// wrapper for ORT C/C++ API headers
|
||||||
|
|
||||||
|
#if defined(__clang__)
|
||||||
|
#pragma clang diagnostic push
|
||||||
|
// ignore clang documentation-related warnings
|
||||||
|
// instead, we will rely on Doxygen warnings for the C/C++ API headers
|
||||||
|
#pragma clang diagnostic ignored "-Wdocumentation"
|
||||||
|
#endif // defined(__clang__)
|
||||||
|
|
||||||
|
// paths are different when building the Swift Package Manager package as the headers come from the iOS pod archive
|
||||||
|
// clang-format off
|
||||||
|
#define STRINGIFY(x) #x
|
||||||
|
#ifdef SPM_BUILD
|
||||||
|
#define ORT_C_CXX_HEADER_FILE_PATH(x) STRINGIFY(onnxruntime/x)
|
||||||
|
#else
|
||||||
|
#define ORT_C_CXX_HEADER_FILE_PATH(x) STRINGIFY(x)
|
||||||
|
#endif
|
||||||
|
// clang-format on
|
||||||
|
|
||||||
|
#if __has_include(ORT_C_CXX_HEADER_FILE_PATH(onnxruntime_training_c_api.h))
|
||||||
|
#include ORT_C_CXX_HEADER_FILE_PATH(onnxruntime_training_c_api.h)
|
||||||
|
#include ORT_C_CXX_HEADER_FILE_PATH(onnxruntime_training_cxx_api.h)
|
||||||
|
#else
|
||||||
|
#include ORT_C_CXX_HEADER_FILE_PATH(onnxruntime_c_api.h)
|
||||||
|
#include ORT_C_CXX_HEADER_FILE_PATH(onnxruntime_cxx_api.h)
|
||||||
|
#endif
|
||||||
|
|
||||||
|
#if __has_include(ORT_C_CXX_HEADER_FILE_PATH(coreml_provider_factory.h))
|
||||||
|
#define ORT_OBJC_API_COREML_EP_AVAILABLE 1
|
||||||
|
#include ORT_C_CXX_HEADER_FILE_PATH(coreml_provider_factory.h)
|
||||||
|
#else
|
||||||
|
#define ORT_OBJC_API_COREML_EP_AVAILABLE 0
|
||||||
|
#endif
|
||||||
|
|
||||||
|
#if defined(__clang__)
|
||||||
|
#pragma clang diagnostic pop
|
||||||
|
#endif // defined(__clang__)
|
||||||
@@ -0,0 +1,32 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#include <optional>
|
||||||
|
#include <string>
|
||||||
|
#include <variant>
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
@class ORTValue;
|
||||||
|
|
||||||
|
namespace utils {
|
||||||
|
|
||||||
|
NSString* toNSString(const std::string& str);
|
||||||
|
NSString* _Nullable toNullableNSString(const std::optional<std::string>& str);
|
||||||
|
|
||||||
|
std::string toStdString(NSString* str);
|
||||||
|
std::optional<std::string> toStdOptionalString(NSString* _Nullable str);
|
||||||
|
|
||||||
|
std::vector<std::string> toStdStringVector(NSArray<NSString*>* strs);
|
||||||
|
NSArray<NSString*>* toNSStringNSArray(const std::vector<std::string>& strs);
|
||||||
|
|
||||||
|
NSArray<ORTValue*>* _Nullable wrapUnownedCAPIOrtValues(const std::vector<OrtValue*>& values, NSError** error);
|
||||||
|
|
||||||
|
std::vector<const OrtValue*> getWrappedCAPIOrtValues(NSArray<ORTValue*>* values);
|
||||||
|
|
||||||
|
} // namespace utils
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "cxx_utils.h"
|
||||||
|
|
||||||
|
#include <vector>
|
||||||
|
#include <optional>
|
||||||
|
#include <string>
|
||||||
|
|
||||||
|
#import "error_utils.h"
|
||||||
|
|
||||||
|
#import "ort_value_internal.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
namespace utils {
|
||||||
|
|
||||||
|
NSString* toNSString(const std::string& str) {
|
||||||
|
NSString* nsStr = [NSString stringWithUTF8String:str.c_str()];
|
||||||
|
if (!nsStr) {
|
||||||
|
ORT_CXX_API_THROW("Failed to convert std::string to NSString", ORT_INVALID_ARGUMENT);
|
||||||
|
}
|
||||||
|
|
||||||
|
return nsStr;
|
||||||
|
}
|
||||||
|
|
||||||
|
NSString* _Nullable toNullableNSString(const std::optional<std::string>& str) {
|
||||||
|
if (str.has_value()) {
|
||||||
|
return toNSString(*str);
|
||||||
|
}
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
std::string toStdString(NSString* str) {
|
||||||
|
return std::string([str UTF8String]);
|
||||||
|
}
|
||||||
|
|
||||||
|
std::optional<std::string> toStdOptionalString(NSString* _Nullable str) {
|
||||||
|
if (str) {
|
||||||
|
return std::optional<std::string>([str UTF8String]);
|
||||||
|
}
|
||||||
|
return std::nullopt;
|
||||||
|
}
|
||||||
|
|
||||||
|
std::vector<std::string> toStdStringVector(NSArray<NSString*>* strs) {
|
||||||
|
std::vector<std::string> result;
|
||||||
|
result.reserve(strs.count);
|
||||||
|
for (NSString* str in strs) {
|
||||||
|
result.push_back([str UTF8String]);
|
||||||
|
}
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
NSArray<NSString*>* toNSStringNSArray(const std::vector<std::string>& strs) {
|
||||||
|
NSMutableArray<NSString*>* result = [NSMutableArray arrayWithCapacity:strs.size()];
|
||||||
|
for (const std::string& str : strs) {
|
||||||
|
[result addObject:toNSString(str)];
|
||||||
|
}
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
NSArray<ORTValue*>* _Nullable wrapUnownedCAPIOrtValues(const std::vector<OrtValue*>& CAPIValues, NSError** error) {
|
||||||
|
NSMutableArray<ORTValue*>* result = [NSMutableArray arrayWithCapacity:CAPIValues.size()];
|
||||||
|
for (size_t i = 0; i < CAPIValues.size(); ++i) {
|
||||||
|
// Wrap the C OrtValue in a C++ Ort::Value to automatically handle its release.
|
||||||
|
// Then, transfer that C++ Ort::Value to a new ORTValue.
|
||||||
|
Ort::Value CXXAPIValue{CAPIValues[i]};
|
||||||
|
ORTValue* val = [[ORTValue alloc] initWithCXXAPIOrtValue:std::move(CXXAPIValue)
|
||||||
|
externalTensorData:nil
|
||||||
|
error:error];
|
||||||
|
if (!val) {
|
||||||
|
// clean up remaining C OrtValues which haven't been wrapped by a C++ Ort::Value yet
|
||||||
|
for (size_t j = i + 1; j < CAPIValues.size(); ++j) {
|
||||||
|
Ort::GetApi().ReleaseValue(CAPIValues[j]);
|
||||||
|
}
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
[result addObject:val];
|
||||||
|
}
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
std::vector<const OrtValue*> getWrappedCAPIOrtValues(NSArray<ORTValue*>* values) {
|
||||||
|
std::vector<const OrtValue*> result;
|
||||||
|
result.reserve(values.count);
|
||||||
|
for (ORTValue* val in values) {
|
||||||
|
result.push_back(static_cast<const OrtValue*>([val CXXAPIOrtValue]));
|
||||||
|
}
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace utils
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#include <exception>
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
void ORTSaveCodeAndDescriptionToError(int code, const char* description, NSError** error);
|
||||||
|
void ORTSaveCodeAndDescriptionToError(int code, NSString* description, NSError** error);
|
||||||
|
void ORTSaveOrtExceptionToError(const Ort::Exception& e, NSError** error);
|
||||||
|
void ORTSaveExceptionToError(const std::exception& e, NSError** error);
|
||||||
|
|
||||||
|
// helper macros to catch and handle C++ exceptions
|
||||||
|
#define ORT_OBJC_API_IMPL_CATCH(error, failure_return_value) \
|
||||||
|
catch (const Ort::Exception& e) { \
|
||||||
|
ORTSaveOrtExceptionToError(e, (error)); \
|
||||||
|
return (failure_return_value); \
|
||||||
|
} \
|
||||||
|
catch (const std::exception& e) { \
|
||||||
|
ORTSaveExceptionToError(e, (error)); \
|
||||||
|
return (failure_return_value); \
|
||||||
|
}
|
||||||
|
|
||||||
|
#define ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error) \
|
||||||
|
ORT_OBJC_API_IMPL_CATCH(error, NO)
|
||||||
|
|
||||||
|
#define ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error) \
|
||||||
|
ORT_OBJC_API_IMPL_CATCH(error, nil)
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "error_utils.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
static NSString* const kOrtErrorDomain = @"onnxruntime";
|
||||||
|
|
||||||
|
void ORTSaveCodeAndDescriptionToError(int code, const char* descriptionCstr, NSError** error) {
|
||||||
|
if (!error) return;
|
||||||
|
|
||||||
|
NSString* description = [NSString stringWithCString:descriptionCstr
|
||||||
|
encoding:NSASCIIStringEncoding];
|
||||||
|
|
||||||
|
*error = [NSError errorWithDomain:kOrtErrorDomain
|
||||||
|
code:code
|
||||||
|
userInfo:@{NSLocalizedDescriptionKey : description}];
|
||||||
|
}
|
||||||
|
|
||||||
|
void ORTSaveCodeAndDescriptionToError(int code, NSString* description, NSError** error) {
|
||||||
|
if (!error) return;
|
||||||
|
|
||||||
|
*error = [NSError errorWithDomain:kOrtErrorDomain
|
||||||
|
code:code
|
||||||
|
userInfo:@{NSLocalizedDescriptionKey : description}];
|
||||||
|
}
|
||||||
|
|
||||||
|
void ORTSaveOrtExceptionToError(const Ort::Exception& e, NSError** error) {
|
||||||
|
ORTSaveCodeAndDescriptionToError(e.GetOrtErrorCode(), e.what(), error);
|
||||||
|
}
|
||||||
|
|
||||||
|
void ORTSaveExceptionToError(const std::exception& e, NSError** error) {
|
||||||
|
ORTSaveCodeAndDescriptionToError(ORT_RUNTIME_EXCEPTION, e.what(), error);
|
||||||
|
}
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
Executable
+9
@@ -0,0 +1,9 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
# formats Objective-C/C++ code
|
||||||
|
|
||||||
|
set -e
|
||||||
|
|
||||||
|
SCRIPT_DIR="$( cd "$( dirname "${BASH_SOURCE[0]}" )" && pwd )"
|
||||||
|
|
||||||
|
clang-format -i $(find ${SCRIPT_DIR} -name "*.h" -o -name "*.m" -o -name "*.mm")
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
// this header contains the entire ONNX Runtime Objective-C API
|
||||||
|
// the headers below can also be imported individually
|
||||||
|
|
||||||
|
#import "ort_coreml_execution_provider.h"
|
||||||
|
#import "ort_custom_op_registration.h"
|
||||||
|
#import "ort_enums.h"
|
||||||
|
#import "ort_env.h"
|
||||||
|
#import "ort_session.h"
|
||||||
|
#import "ort_value.h"
|
||||||
|
#import "ort_xnnpack_execution_provider.h"
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
// 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"
|
||||||
@@ -0,0 +1,119 @@
|
|||||||
|
// 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
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#import "ort_session.h"
|
||||||
|
|
||||||
|
#ifdef __cplusplus
|
||||||
|
extern "C" {
|
||||||
|
#endif
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets whether the CoreML execution provider is available.
|
||||||
|
*/
|
||||||
|
BOOL ORTIsCoreMLExecutionProviderAvailable(void);
|
||||||
|
|
||||||
|
#ifdef __cplusplus
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Options for configuring the CoreML execution provider.
|
||||||
|
*/
|
||||||
|
@interface ORTCoreMLExecutionProviderOptions : NSObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Whether the CoreML execution provider should run on CPU only.
|
||||||
|
*/
|
||||||
|
@property BOOL useCPUOnly;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Whether the CoreML execution provider is enabled on subgraphs.
|
||||||
|
*/
|
||||||
|
@property BOOL enableOnSubgraphs;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Whether the CoreML execution provider is only enabled for devices with Apple
|
||||||
|
* Neural Engine (ANE).
|
||||||
|
*/
|
||||||
|
@property BOOL onlyEnableForDevicesWithANE;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@interface ORTSessionOptions (ORTSessionOptionsCoreMLEP)
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Enables the CoreML execution provider in the session configuration options.
|
||||||
|
* It is appended to the execution provider list which is ordered by
|
||||||
|
* decreasing priority.
|
||||||
|
*
|
||||||
|
* @param options The CoreML execution provider configuration options.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the provider was enabled successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)appendCoreMLExecutionProviderWithOptions:(ORTCoreMLExecutionProviderOptions*)options
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
/** C API type forward declaration. */
|
||||||
|
struct OrtStatus;
|
||||||
|
|
||||||
|
/** C API type forward declaration. */
|
||||||
|
struct OrtApiBase;
|
||||||
|
|
||||||
|
/** C API type forward declaration. */
|
||||||
|
struct OrtSessionOptions;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Pointer to a custom op registration function that uses the ONNX Runtime C API.
|
||||||
|
*
|
||||||
|
* The signature is defined in the ONNX Runtime C API:
|
||||||
|
* https://github.com/microsoft/onnxruntime/blob/67f4cd54fab321d83e4a75a40efeee95a6a17079/include/onnxruntime/core/session/onnxruntime_c_api.h#L697
|
||||||
|
*
|
||||||
|
* This is a low-level type intended for interoperating with libraries which provide such a function for custom op
|
||||||
|
* registration, such as [ONNX Runtime Extensions](https://github.com/microsoft/onnxruntime-extensions).
|
||||||
|
*/
|
||||||
|
typedef struct OrtStatus* (*ORTCAPIRegisterCustomOpsFnPtr)(struct OrtSessionOptions* /*options*/,
|
||||||
|
const struct OrtApiBase* /*api*/);
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
/**
|
||||||
|
* The ORT logging verbosity levels.
|
||||||
|
*/
|
||||||
|
typedef NS_ENUM(int32_t, ORTLoggingLevel) {
|
||||||
|
ORTLoggingLevelVerbose,
|
||||||
|
ORTLoggingLevelInfo,
|
||||||
|
ORTLoggingLevelWarning,
|
||||||
|
ORTLoggingLevelError,
|
||||||
|
ORTLoggingLevelFatal,
|
||||||
|
};
|
||||||
|
|
||||||
|
/**
|
||||||
|
* The ORT value types.
|
||||||
|
* Currently, a subset of all types is supported.
|
||||||
|
*/
|
||||||
|
typedef NS_ENUM(int32_t, ORTValueType) {
|
||||||
|
ORTValueTypeUnknown,
|
||||||
|
ORTValueTypeTensor,
|
||||||
|
};
|
||||||
|
|
||||||
|
/**
|
||||||
|
* The ORT tensor element data types.
|
||||||
|
* Currently, a subset of all types is supported.
|
||||||
|
*/
|
||||||
|
typedef NS_ENUM(int32_t, ORTTensorElementDataType) {
|
||||||
|
ORTTensorElementDataTypeUndefined,
|
||||||
|
ORTTensorElementDataTypeFloat,
|
||||||
|
ORTTensorElementDataTypeInt8,
|
||||||
|
ORTTensorElementDataTypeUInt8,
|
||||||
|
ORTTensorElementDataTypeInt32,
|
||||||
|
ORTTensorElementDataTypeUInt32,
|
||||||
|
ORTTensorElementDataTypeInt64,
|
||||||
|
ORTTensorElementDataTypeUInt64,
|
||||||
|
};
|
||||||
|
|
||||||
|
/**
|
||||||
|
* The ORT graph optimization levels.
|
||||||
|
* See here for more details:
|
||||||
|
* https://onnxruntime.ai/docs/performance/graph-optimizations.html
|
||||||
|
*/
|
||||||
|
typedef NS_ENUM(int32_t, ORTGraphOptimizationLevel) {
|
||||||
|
ORTGraphOptimizationLevelNone,
|
||||||
|
ORTGraphOptimizationLevelBasic,
|
||||||
|
ORTGraphOptimizationLevelExtended,
|
||||||
|
ORTGraphOptimizationLevelAll,
|
||||||
|
};
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#import "ort_enums.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
#ifdef __cplusplus
|
||||||
|
extern "C" {
|
||||||
|
#endif
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the ORT version string in format major.minor.patch.
|
||||||
|
*
|
||||||
|
* Available since 1.15.
|
||||||
|
*/
|
||||||
|
NSString* _Nullable ORTVersion(void);
|
||||||
|
|
||||||
|
#ifdef __cplusplus
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
/**
|
||||||
|
* The ORT environment.
|
||||||
|
*/
|
||||||
|
@interface ORTEnv : NSObject
|
||||||
|
|
||||||
|
- (instancetype)init NS_UNAVAILABLE;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates an ORT Environment.
|
||||||
|
*
|
||||||
|
* @param loggingLevel The environment logging level.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The instance, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable instancetype)initWithLoggingLevel:(ORTLoggingLevel)loggingLevel
|
||||||
|
error:(NSError**)error NS_DESIGNATED_INITIALIZER;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,300 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#import "ort_custom_op_registration.h"
|
||||||
|
#import "ort_enums.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@class ORTEnv;
|
||||||
|
@class ORTRunOptions;
|
||||||
|
@class ORTSessionOptions;
|
||||||
|
@class ORTValue;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* An ORT session loads and runs a model.
|
||||||
|
*/
|
||||||
|
@interface ORTSession : NSObject
|
||||||
|
|
||||||
|
- (instancetype)init NS_UNAVAILABLE;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates a session.
|
||||||
|
*
|
||||||
|
* @param env The ORT Environment instance.
|
||||||
|
* @param path The path to the ONNX model.
|
||||||
|
* @param sessionOptions Optional session configuration options.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The instance, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable instancetype)initWithEnv:(ORTEnv*)env
|
||||||
|
modelPath:(NSString*)path
|
||||||
|
sessionOptions:(nullable ORTSessionOptions*)sessionOptions
|
||||||
|
error:(NSError**)error NS_DESIGNATED_INITIALIZER;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Runs the model.
|
||||||
|
* The inputs and outputs are pre-allocated.
|
||||||
|
*
|
||||||
|
* @param inputs Dictionary of input names to input ORT values.
|
||||||
|
* @param outputs Dictionary of output names to output ORT values.
|
||||||
|
* @param runOptions Optional run configuration options.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the model was run successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)runWithInputs:(NSDictionary<NSString*, ORTValue*>*)inputs
|
||||||
|
outputs:(NSDictionary<NSString*, ORTValue*>*)outputs
|
||||||
|
runOptions:(nullable ORTRunOptions*)runOptions
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Runs the model.
|
||||||
|
* The inputs are pre-allocated and the outputs are allocated by ORT.
|
||||||
|
*
|
||||||
|
* @param inputs Dictionary of input names to input ORT values.
|
||||||
|
* @param outputNames Set of output names.
|
||||||
|
* @param runOptions Optional run configuration options.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return A dictionary of output names to output ORT values with the outputs
|
||||||
|
* requested in `outputNames`, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable NSDictionary<NSString*, ORTValue*>*)runWithInputs:(NSDictionary<NSString*, ORTValue*>*)inputs
|
||||||
|
outputNames:(NSSet<NSString*>*)outputNames
|
||||||
|
runOptions:(nullable ORTRunOptions*)runOptions
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the model's input names.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return An array of input names, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable NSArray<NSString*>*)inputNamesWithError:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the model's overridable initializer names.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return An array of overridable initializer names, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable NSArray<NSString*>*)overridableInitializerNamesWithError:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the model's output names.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return An array of output names, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable NSArray<NSString*>*)outputNamesWithError:(NSError**)error;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Options for configuring a session.
|
||||||
|
*/
|
||||||
|
@interface ORTSessionOptions : NSObject
|
||||||
|
|
||||||
|
- (instancetype)init NS_UNAVAILABLE;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates session configuration options.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The instance, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable instancetype)initWithError:(NSError**)error NS_SWIFT_NAME(init());
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Appends an execution provider to the session options to enable the execution provider to be used when running
|
||||||
|
* the model.
|
||||||
|
*
|
||||||
|
* Available since 1.14.
|
||||||
|
*
|
||||||
|
* The execution provider list is ordered by decreasing priority.
|
||||||
|
* i.e. the first provider registered has the highest priority.
|
||||||
|
*
|
||||||
|
* @param providerName Provider name. For example, "xnnpack".
|
||||||
|
* @param providerOptions Provider-specific options. For example, for provider "xnnpack", {"intra_op_num_threads": "2"}.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the execution provider was appended successfully
|
||||||
|
*/
|
||||||
|
- (BOOL)appendExecutionProvider:(NSString*)providerName
|
||||||
|
providerOptions:(NSDictionary<NSString*, NSString*>*)providerOptions
|
||||||
|
error:(NSError**)error;
|
||||||
|
/**
|
||||||
|
* Sets the number of threads used to parallelize the execution within nodes.
|
||||||
|
* A value of 0 means ORT will pick a default value.
|
||||||
|
*
|
||||||
|
* @param intraOpNumThreads The number of threads.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setIntraOpNumThreads:(int)intraOpNumThreads
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets the graph optimization level.
|
||||||
|
*
|
||||||
|
* @param graphOptimizationLevel The graph optimization level.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setGraphOptimizationLevel:(ORTGraphOptimizationLevel)graphOptimizationLevel
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets the path to which the optimized model file will be saved.
|
||||||
|
*
|
||||||
|
* @param optimizedModelFilePath The optimized model file path.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setOptimizedModelFilePath:(NSString*)optimizedModelFilePath
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets the session log ID.
|
||||||
|
*
|
||||||
|
* @param logID The log ID.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setLogID:(NSString*)logID
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets the session log severity level.
|
||||||
|
*
|
||||||
|
* @param loggingLevel The log severity level.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setLogSeverityLevel:(ORTLoggingLevel)loggingLevel
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets a session configuration key-value pair.
|
||||||
|
* Any value for a previously set key will be overwritten.
|
||||||
|
* The session configuration keys and values are documented here:
|
||||||
|
* https://github.com/microsoft/onnxruntime/blob/main/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h
|
||||||
|
*
|
||||||
|
* @param key The key.
|
||||||
|
* @param value The value.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)addConfigEntryWithKey:(NSString*)key
|
||||||
|
value:(NSString*)value
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Registers custom ops for use with `ORTSession`s using this SessionOptions by calling the specified
|
||||||
|
* native function name. The custom ops library must either be linked against, or have previously been loaded
|
||||||
|
* by the user.
|
||||||
|
*
|
||||||
|
* Available since 1.14.
|
||||||
|
*
|
||||||
|
* The registration function must have the signature:
|
||||||
|
* `OrtStatus* (*fn)(OrtSessionOptions* options, const OrtApiBase* api);`
|
||||||
|
*
|
||||||
|
* The signature is defined in the ONNX Runtime C API:
|
||||||
|
* https://github.com/microsoft/onnxruntime/blob/67f4cd54fab321d83e4a75a40efeee95a6a17079/include/onnxruntime/core/session/onnxruntime_c_api.h#L697
|
||||||
|
*
|
||||||
|
* See https://onnxruntime.ai/docs/reference/operators/add-custom-op.html for more information on custom ops.
|
||||||
|
* See https://github.com/microsoft/onnxruntime/blob/342a5bf2b756d1a1fc6fdc582cfeac15182632fe/onnxruntime/test/testdata/custom_op_library/custom_op_library.cc#L115
|
||||||
|
* for an example of a custom op library registration function.
|
||||||
|
*
|
||||||
|
* @note The caller must ensure that `registrationFuncName` names a valid function that is visible to the native ONNX
|
||||||
|
* Runtime code and has the correct signature.
|
||||||
|
* They must ensure that the function does what they expect it to do because this method will just call it.
|
||||||
|
*
|
||||||
|
* @param registrationFuncName The name of the registration function to call.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the registration function was successfully called.
|
||||||
|
*/
|
||||||
|
- (BOOL)registerCustomOpsUsingFunction:(NSString*)registrationFuncName
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Registers custom ops for use with `ORTSession`s using this SessionOptions by calling the specified function
|
||||||
|
* pointed to by `registerCustomOpsFn`.
|
||||||
|
*
|
||||||
|
* Available since 1.16.
|
||||||
|
*
|
||||||
|
* The registration function must have the signature:
|
||||||
|
* `OrtStatus* (*fn)(OrtSessionOptions* options, const OrtApiBase* api);`
|
||||||
|
*
|
||||||
|
* The signature is defined in the ONNX Runtime C API:
|
||||||
|
* https://github.com/microsoft/onnxruntime/blob/67f4cd54fab321d83e4a75a40efeee95a6a17079/include/onnxruntime/core/session/onnxruntime_c_api.h#L697
|
||||||
|
*
|
||||||
|
* See https://onnxruntime.ai/docs/reference/operators/add-custom-op.html for more information on custom ops.
|
||||||
|
* See https://github.com/microsoft/onnxruntime/blob/342a5bf2b756d1a1fc6fdc582cfeac15182632fe/onnxruntime/test/testdata/custom_op_library/custom_op_library.cc#L115
|
||||||
|
* for an example of a custom op library registration function.
|
||||||
|
*
|
||||||
|
* @note The caller must ensure that `registerCustomOpsFn` is a valid function pointer and has the correct signature.
|
||||||
|
* They must ensure that the function does what they expect it to do because this method will just call it.
|
||||||
|
*
|
||||||
|
* @param registerCustomOpsFn A pointer to the registration function to call.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the registration function was successfully called.
|
||||||
|
*/
|
||||||
|
- (BOOL)registerCustomOpsUsingFunctionPointer:(ORTCAPIRegisterCustomOpsFnPtr)registerCustomOpsFn
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Options for configuring a run.
|
||||||
|
*/
|
||||||
|
@interface ORTRunOptions : NSObject
|
||||||
|
|
||||||
|
- (instancetype)init NS_UNAVAILABLE;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates run configuration options.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The instance, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable instancetype)initWithError:(NSError**)error NS_SWIFT_NAME(init());
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets the run log tag.
|
||||||
|
*
|
||||||
|
* @param logTag The log tag.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setLogTag:(NSString*)logTag
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets the run log severity level.
|
||||||
|
*
|
||||||
|
* @param loggingLevel The log severity level.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)setLogSeverityLevel:(ORTLoggingLevel)loggingLevel
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sets a run configuration key-value pair.
|
||||||
|
* Any value for a previously set key will be overwritten.
|
||||||
|
* The run configuration keys and values are documented here:
|
||||||
|
* https://github.com/microsoft/onnxruntime/blob/main/include/onnxruntime/core/session/onnxruntime_run_options_config_keys.h
|
||||||
|
*
|
||||||
|
* @param key The key.
|
||||||
|
* @param value The value.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the option was set successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)addConfigEntryWithKey:(NSString*)key
|
||||||
|
value:(NSString*)value
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,263 @@
|
|||||||
|
// 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
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#import "ort_enums.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@class ORTValueTypeInfo;
|
||||||
|
@class ORTTensorTypeAndShapeInfo;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* An ORT value encapsulates data used as an input or output to a model at runtime.
|
||||||
|
*/
|
||||||
|
@interface ORTValue : NSObject
|
||||||
|
|
||||||
|
- (instancetype)init NS_UNAVAILABLE;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates a value that is a tensor.
|
||||||
|
* The tensor data is allocated by the caller.
|
||||||
|
*
|
||||||
|
* @param tensorData The tensor data.
|
||||||
|
* @param elementType The tensor element data type.
|
||||||
|
* @param shape The tensor shape.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The instance, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable instancetype)initWithTensorData:(NSMutableData*)tensorData
|
||||||
|
elementType:(ORTTensorElementDataType)elementType
|
||||||
|
shape:(NSArray<NSNumber*>*)shape
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the type information.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The type information, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable ORTValueTypeInfo*)typeInfoWithError:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the tensor type and shape information.
|
||||||
|
* This assumes that the value is a tensor.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The tensor type and shape information, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable ORTTensorTypeAndShapeInfo*)tensorTypeAndShapeInfoWithError:(NSError**)error;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the tensor data.
|
||||||
|
* This assumes that the value is a tensor.
|
||||||
|
*
|
||||||
|
* This returns the value's underlying data directly, not a copy of it.
|
||||||
|
* The memory's lifetime may be tied to this value, i.e., if it was allocated
|
||||||
|
* by ORT. On the other hand, the memory's lifetime is independent of the value
|
||||||
|
* if the value was created with user-provided data.
|
||||||
|
*
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The tensor data, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable NSMutableData*)tensorDataWithError:(NSError**)error;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
/**
|
||||||
|
* A value's type information.
|
||||||
|
*/
|
||||||
|
@interface ORTValueTypeInfo : NSObject
|
||||||
|
|
||||||
|
/** The value type. */
|
||||||
|
@property(nonatomic) ORTValueType type;
|
||||||
|
|
||||||
|
/** The tensor type and shape information, if the value is a tensor. */
|
||||||
|
@property(nonatomic, nullable) ORTTensorTypeAndShapeInfo* tensorTypeAndShapeInfo;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
/**
|
||||||
|
* A tensor's type and shape information.
|
||||||
|
*/
|
||||||
|
@interface ORTTensorTypeAndShapeInfo : NSObject
|
||||||
|
|
||||||
|
/** The tensor element data type. */
|
||||||
|
@property(nonatomic) ORTTensorElementDataType elementType;
|
||||||
|
|
||||||
|
/** The tensor shape. */
|
||||||
|
@property(nonatomic) NSArray<NSNumber*>* shape;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
|
||||||
|
#import "ort_session.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Options for configuring the Xnnpack execution provider.
|
||||||
|
*/
|
||||||
|
@interface ORTXnnpackExecutionProviderOptions : NSObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* How many threads used for the Xnnpack execution provider.
|
||||||
|
*/
|
||||||
|
@property int intra_op_num_threads;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@interface ORTSessionOptions (ORTSessionOptionsXnnpackEP)
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Available since 1.14.
|
||||||
|
* Enables the Xnnpack execution provider in the session configuration options.
|
||||||
|
* It is appended to the execution provider list which is ordered by
|
||||||
|
* decreasing priority.
|
||||||
|
*
|
||||||
|
* @param options The Xnnpack execution provider configuration options.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return Whether the provider was enabled successfully.
|
||||||
|
*/
|
||||||
|
- (BOOL)appendXnnpackExecutionProviderWithOptions:(ORTXnnpackExecutionProviderOptions*)options
|
||||||
|
error:(NSError**)error;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
// 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
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
// 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
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_coreml_execution_provider.h"
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
#import "error_utils.h"
|
||||||
|
#import "ort_session_internal.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
BOOL ORTIsCoreMLExecutionProviderAvailable() {
|
||||||
|
return ORT_OBJC_API_COREML_EP_AVAILABLE ? YES : NO;
|
||||||
|
}
|
||||||
|
|
||||||
|
@implementation ORTCoreMLExecutionProviderOptions
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTSessionOptions (ORTSessionOptionsCoreMLEP)
|
||||||
|
|
||||||
|
- (BOOL)appendCoreMLExecutionProviderWithOptions:(ORTCoreMLExecutionProviderOptions*)options
|
||||||
|
error:(NSError**)error {
|
||||||
|
#if ORT_OBJC_API_COREML_EP_AVAILABLE
|
||||||
|
try {
|
||||||
|
const uint32_t flags =
|
||||||
|
(options.useCPUOnly ? COREML_FLAG_USE_CPU_ONLY : 0) |
|
||||||
|
(options.enableOnSubgraphs ? COREML_FLAG_ENABLE_ON_SUBGRAPH : 0) |
|
||||||
|
(options.onlyEnableForDevicesWithANE ? COREML_FLAG_ONLY_ENABLE_DEVICE_WITH_ANE : 0);
|
||||||
|
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CoreML(
|
||||||
|
[self CXXAPIOrtSessionOptions], flags));
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error);
|
||||||
|
#else // !ORT_OBJC_API_COREML_EP_AVAILABLE
|
||||||
|
static_cast<void>(options);
|
||||||
|
ORTSaveCodeAndDescriptionToError(ORT_FAIL, "CoreML execution provider is not enabled.", error);
|
||||||
|
return NO;
|
||||||
|
#endif
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_enums_internal.h"
|
||||||
|
|
||||||
|
#include <algorithm>
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
|
||||||
|
struct LoggingLevelInfo {
|
||||||
|
ORTLoggingLevel logging_level;
|
||||||
|
OrtLoggingLevel capi_logging_level;
|
||||||
|
};
|
||||||
|
|
||||||
|
// supported ORT logging levels
|
||||||
|
// define the mapping from ORTLoggingLevel to C API OrtLoggingLevel here
|
||||||
|
constexpr LoggingLevelInfo kLoggingLevelInfos[]{
|
||||||
|
{ORTLoggingLevelVerbose, ORT_LOGGING_LEVEL_VERBOSE},
|
||||||
|
{ORTLoggingLevelInfo, ORT_LOGGING_LEVEL_INFO},
|
||||||
|
{ORTLoggingLevelWarning, ORT_LOGGING_LEVEL_WARNING},
|
||||||
|
{ORTLoggingLevelError, ORT_LOGGING_LEVEL_ERROR},
|
||||||
|
{ORTLoggingLevelFatal, ORT_LOGGING_LEVEL_FATAL},
|
||||||
|
};
|
||||||
|
|
||||||
|
struct ValueTypeInfo {
|
||||||
|
ORTValueType type;
|
||||||
|
ONNXType capi_type;
|
||||||
|
};
|
||||||
|
|
||||||
|
// supported ORT value types
|
||||||
|
// define the mapping from ORTValueType to C API ONNXType here
|
||||||
|
constexpr ValueTypeInfo kValueTypeInfos[]{
|
||||||
|
{ORTValueTypeUnknown, ONNX_TYPE_UNKNOWN},
|
||||||
|
{ORTValueTypeTensor, ONNX_TYPE_TENSOR},
|
||||||
|
};
|
||||||
|
|
||||||
|
struct TensorElementTypeInfo {
|
||||||
|
ORTTensorElementDataType type;
|
||||||
|
ONNXTensorElementDataType capi_type;
|
||||||
|
size_t element_size;
|
||||||
|
};
|
||||||
|
|
||||||
|
// supported ORT tensor element data types
|
||||||
|
// define the mapping from ORTTensorElementDataType to C API ONNXTensorElementDataType here
|
||||||
|
constexpr TensorElementTypeInfo kElementTypeInfos[]{
|
||||||
|
{ORTTensorElementDataTypeUndefined, ONNX_TENSOR_ELEMENT_DATA_TYPE_UNDEFINED, 0},
|
||||||
|
{ORTTensorElementDataTypeFloat, ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, sizeof(float)},
|
||||||
|
{ORTTensorElementDataTypeInt8, ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8, sizeof(int8_t)},
|
||||||
|
{ORTTensorElementDataTypeUInt8, ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8, sizeof(uint8_t)},
|
||||||
|
{ORTTensorElementDataTypeInt32, ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, sizeof(int32_t)},
|
||||||
|
{ORTTensorElementDataTypeUInt32, ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32, sizeof(uint32_t)},
|
||||||
|
{ORTTensorElementDataTypeInt64, ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64, sizeof(int64_t)},
|
||||||
|
{ORTTensorElementDataTypeUInt64, ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64, sizeof(uint64_t)},
|
||||||
|
};
|
||||||
|
|
||||||
|
struct GraphOptimizationLevelInfo {
|
||||||
|
ORTGraphOptimizationLevel opt_level;
|
||||||
|
GraphOptimizationLevel capi_opt_level;
|
||||||
|
};
|
||||||
|
|
||||||
|
// ORT graph optimization levels
|
||||||
|
// define the mapping from ORTGraphOptimizationLevel to C API GraphOptimizationLevel here
|
||||||
|
constexpr GraphOptimizationLevelInfo kGraphOptimizationLevelInfos[]{
|
||||||
|
{ORTGraphOptimizationLevelNone, ORT_DISABLE_ALL},
|
||||||
|
{ORTGraphOptimizationLevelBasic, ORT_ENABLE_BASIC},
|
||||||
|
{ORTGraphOptimizationLevelExtended, ORT_ENABLE_EXTENDED},
|
||||||
|
{ORTGraphOptimizationLevelAll, ORT_ENABLE_ALL},
|
||||||
|
};
|
||||||
|
|
||||||
|
template <typename Container, typename SelectFn, typename TransformFn>
|
||||||
|
auto SelectAndTransform(
|
||||||
|
const Container& container, SelectFn select_fn, TransformFn transform_fn,
|
||||||
|
const char* not_found_msg)
|
||||||
|
-> decltype(transform_fn(*std::begin(container))) {
|
||||||
|
const auto it = std::find_if(
|
||||||
|
std::begin(container), std::end(container), select_fn);
|
||||||
|
if (it == std::end(container)) {
|
||||||
|
ORT_CXX_API_THROW(not_found_msg, ORT_NOT_IMPLEMENTED);
|
||||||
|
}
|
||||||
|
return transform_fn(*it);
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace
|
||||||
|
|
||||||
|
OrtLoggingLevel PublicToCAPILoggingLevel(ORTLoggingLevel logging_level) {
|
||||||
|
return SelectAndTransform(
|
||||||
|
kLoggingLevelInfos,
|
||||||
|
[logging_level](const auto& logging_level_info) { return logging_level_info.logging_level == logging_level; },
|
||||||
|
[](const auto& logging_level_info) { return logging_level_info.capi_logging_level; },
|
||||||
|
"unsupported logging level");
|
||||||
|
}
|
||||||
|
|
||||||
|
ORTValueType CAPIToPublicValueType(ONNXType capi_type) {
|
||||||
|
return SelectAndTransform(
|
||||||
|
kValueTypeInfos,
|
||||||
|
[capi_type](const auto& type_info) { return type_info.capi_type == capi_type; },
|
||||||
|
[](const auto& type_info) { return type_info.type; },
|
||||||
|
"unsupported value type");
|
||||||
|
}
|
||||||
|
|
||||||
|
ONNXTensorElementDataType PublicToCAPITensorElementType(ORTTensorElementDataType type) {
|
||||||
|
return SelectAndTransform(
|
||||||
|
kElementTypeInfos,
|
||||||
|
[type](const auto& type_info) { return type_info.type == type; },
|
||||||
|
[](const auto& type_info) { return type_info.capi_type; },
|
||||||
|
"unsupported tensor element type");
|
||||||
|
}
|
||||||
|
|
||||||
|
ORTTensorElementDataType CAPIToPublicTensorElementType(ONNXTensorElementDataType capi_type) {
|
||||||
|
return SelectAndTransform(
|
||||||
|
kElementTypeInfos,
|
||||||
|
[capi_type](const auto& type_info) { return type_info.capi_type == capi_type; },
|
||||||
|
[](const auto& type_info) { return type_info.type; },
|
||||||
|
"unsupported tensor element type");
|
||||||
|
}
|
||||||
|
|
||||||
|
size_t SizeOfCAPITensorElementType(ONNXTensorElementDataType capi_type) {
|
||||||
|
return SelectAndTransform(
|
||||||
|
kElementTypeInfos,
|
||||||
|
[capi_type](const auto& type_info) { return type_info.capi_type == capi_type; },
|
||||||
|
[](const auto& type_info) { return type_info.element_size; },
|
||||||
|
"unsupported tensor element type");
|
||||||
|
}
|
||||||
|
|
||||||
|
GraphOptimizationLevel PublicToCAPIGraphOptimizationLevel(ORTGraphOptimizationLevel opt_level) {
|
||||||
|
return SelectAndTransform(
|
||||||
|
kGraphOptimizationLevelInfos,
|
||||||
|
[opt_level](const auto& opt_level_info) { return opt_level_info.opt_level == opt_level; },
|
||||||
|
[](const auto& opt_level_info) { return opt_level_info.capi_opt_level; },
|
||||||
|
"unsupported graph optimization level");
|
||||||
|
}
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_enums.h"
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
OrtLoggingLevel PublicToCAPILoggingLevel(ORTLoggingLevel logging_level);
|
||||||
|
|
||||||
|
ORTValueType CAPIToPublicValueType(ONNXType capi_type);
|
||||||
|
|
||||||
|
ONNXTensorElementDataType PublicToCAPITensorElementType(ORTTensorElementDataType type);
|
||||||
|
ORTTensorElementDataType CAPIToPublicTensorElementType(ONNXTensorElementDataType capi_type);
|
||||||
|
|
||||||
|
size_t SizeOfCAPITensorElementType(ONNXTensorElementDataType capi_type);
|
||||||
|
|
||||||
|
GraphOptimizationLevel PublicToCAPIGraphOptimizationLevel(ORTGraphOptimizationLevel opt_level);
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_env_internal.h"
|
||||||
|
|
||||||
|
#include <optional>
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
#import "error_utils.h"
|
||||||
|
#import "ort_enums_internal.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
NSString* _Nullable ORTVersion(void) {
|
||||||
|
return [NSString stringWithUTF8String:OrtGetApiBase()->GetVersionString()];
|
||||||
|
}
|
||||||
|
|
||||||
|
@implementation ORTEnv {
|
||||||
|
std::optional<Ort::Env> _env;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable instancetype)initWithLoggingLevel:(ORTLoggingLevel)loggingLevel
|
||||||
|
error:(NSError**)error {
|
||||||
|
if ((self = [super init]) == nil) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
const auto CAPILoggingLevel = PublicToCAPILoggingLevel(loggingLevel);
|
||||||
|
_env = Ort::Env{CAPILoggingLevel};
|
||||||
|
return self;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (Ort::Env&)CXXAPIOrtEnv {
|
||||||
|
return *_env;
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_env.h"
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTEnv ()
|
||||||
|
|
||||||
|
- (Ort::Env&)CXXAPIOrtEnv;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,387 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_session_internal.h"
|
||||||
|
|
||||||
|
#include <optional>
|
||||||
|
#include <vector>
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
#import "error_utils.h"
|
||||||
|
#import "ort_enums_internal.h"
|
||||||
|
#import "ort_env_internal.h"
|
||||||
|
#import "ort_value_internal.h"
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
enum class NamedValueType {
|
||||||
|
Input,
|
||||||
|
OverridableInitializer,
|
||||||
|
Output,
|
||||||
|
};
|
||||||
|
} // namespace
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@implementation ORTSession {
|
||||||
|
std::optional<Ort::Session> _session;
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Public
|
||||||
|
|
||||||
|
- (nullable instancetype)initWithEnv:(ORTEnv*)env
|
||||||
|
modelPath:(NSString*)path
|
||||||
|
sessionOptions:(nullable ORTSessionOptions*)sessionOptions
|
||||||
|
error:(NSError**)error {
|
||||||
|
if ((self = [super init]) == nil) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
if (!sessionOptions) {
|
||||||
|
sessionOptions = [[ORTSessionOptions alloc] initWithError:error];
|
||||||
|
if (!sessionOptions) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_session = Ort::Session{[env CXXAPIOrtEnv],
|
||||||
|
path.UTF8String,
|
||||||
|
[sessionOptions CXXAPIOrtSessionOptions]};
|
||||||
|
|
||||||
|
return self;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)runWithInputs:(NSDictionary<NSString*, ORTValue*>*)inputs
|
||||||
|
outputs:(NSDictionary<NSString*, ORTValue*>*)outputs
|
||||||
|
runOptions:(nullable ORTRunOptions*)runOptions
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
if (!runOptions) {
|
||||||
|
runOptions = [[ORTRunOptions alloc] initWithError:error];
|
||||||
|
if (!runOptions) {
|
||||||
|
return NO;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
std::vector<const char*> inputNames, outputNames;
|
||||||
|
std::vector<const OrtValue*> inputCAPIValues;
|
||||||
|
std::vector<OrtValue*> outputCAPIValues;
|
||||||
|
|
||||||
|
inputNames.reserve(inputs.count);
|
||||||
|
inputCAPIValues.reserve(inputs.count);
|
||||||
|
for (NSString* inputName in inputs) {
|
||||||
|
inputNames.push_back(inputName.UTF8String);
|
||||||
|
inputCAPIValues.push_back(static_cast<const OrtValue*>([inputs[inputName] CXXAPIOrtValue]));
|
||||||
|
}
|
||||||
|
|
||||||
|
outputNames.reserve(outputs.count);
|
||||||
|
outputCAPIValues.reserve(outputs.count);
|
||||||
|
for (NSString* outputName in outputs) {
|
||||||
|
outputNames.push_back(outputName.UTF8String);
|
||||||
|
outputCAPIValues.push_back(static_cast<OrtValue*>([outputs[outputName] CXXAPIOrtValue]));
|
||||||
|
}
|
||||||
|
|
||||||
|
Ort::ThrowOnError(Ort::GetApi().Run(*_session, [runOptions CXXAPIOrtRunOptions],
|
||||||
|
inputNames.data(), inputCAPIValues.data(), inputNames.size(),
|
||||||
|
outputNames.data(), outputNames.size(), outputCAPIValues.data()));
|
||||||
|
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable NSDictionary<NSString*, ORTValue*>*)runWithInputs:(NSDictionary<NSString*, ORTValue*>*)inputs
|
||||||
|
outputNames:(NSSet<NSString*>*)outputNameSet
|
||||||
|
runOptions:(nullable ORTRunOptions*)runOptions
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
if (!runOptions) {
|
||||||
|
runOptions = [[ORTRunOptions alloc] initWithError:error];
|
||||||
|
if (!runOptions) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
NSArray<NSString*>* outputNameArray = outputNameSet.allObjects;
|
||||||
|
|
||||||
|
std::vector<const char*> inputNames, outputNames;
|
||||||
|
std::vector<const OrtValue*> inputCAPIValues;
|
||||||
|
std::vector<OrtValue*> outputCAPIValues;
|
||||||
|
|
||||||
|
inputNames.reserve(inputs.count);
|
||||||
|
inputCAPIValues.reserve(inputs.count);
|
||||||
|
for (NSString* inputName in inputs) {
|
||||||
|
inputNames.push_back(inputName.UTF8String);
|
||||||
|
inputCAPIValues.push_back(static_cast<const OrtValue*>([inputs[inputName] CXXAPIOrtValue]));
|
||||||
|
}
|
||||||
|
|
||||||
|
outputNames.reserve(outputNameArray.count);
|
||||||
|
outputCAPIValues.reserve(outputNameArray.count);
|
||||||
|
for (NSString* outputName in outputNameArray) {
|
||||||
|
outputNames.push_back(outputName.UTF8String);
|
||||||
|
outputCAPIValues.push_back(nullptr);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ort::ThrowOnError(Ort::GetApi().Run(*_session, [runOptions CXXAPIOrtRunOptions],
|
||||||
|
inputNames.data(), inputCAPIValues.data(), inputNames.size(),
|
||||||
|
outputNames.data(), outputNames.size(), outputCAPIValues.data()));
|
||||||
|
|
||||||
|
NSMutableDictionary<NSString*, ORTValue*>* outputs = [[NSMutableDictionary alloc] init];
|
||||||
|
for (NSUInteger i = 0; i < outputNameArray.count; ++i) {
|
||||||
|
// Wrap the C OrtValue in a C++ Ort::Value to automatically handle its release.
|
||||||
|
// Then, transfer that C++ Ort::Value to a new ORTValue.
|
||||||
|
Ort::Value outputCXXAPIValue{outputCAPIValues[i]};
|
||||||
|
ORTValue* outputValue = [[ORTValue alloc] initWithCXXAPIOrtValue:std::move(outputCXXAPIValue)
|
||||||
|
externalTensorData:nil
|
||||||
|
error:error];
|
||||||
|
if (!outputValue) {
|
||||||
|
// clean up remaining C OrtValues which haven't been wrapped by a C++ Ort::Value yet
|
||||||
|
for (NSUInteger j = i + 1; j < outputNameArray.count; ++j) {
|
||||||
|
Ort::GetApi().ReleaseValue(outputCAPIValues[j]);
|
||||||
|
}
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
outputs[outputNameArray[i]] = outputValue;
|
||||||
|
}
|
||||||
|
|
||||||
|
return outputs;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable NSArray<NSString*>*)inputNamesWithError:(NSError**)error {
|
||||||
|
return [self namesWithType:NamedValueType::Input error:error];
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable NSArray<NSString*>*)overridableInitializerNamesWithError:(NSError**)error {
|
||||||
|
return [self namesWithType:NamedValueType::OverridableInitializer error:error];
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable NSArray<NSString*>*)outputNamesWithError:(NSError**)error {
|
||||||
|
return [self namesWithType:NamedValueType::Output error:error];
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Private
|
||||||
|
|
||||||
|
- (nullable NSArray<NSString*>*)namesWithType:(NamedValueType)namedValueType
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
auto getCount = [&session = *_session, namedValueType]() {
|
||||||
|
if (namedValueType == NamedValueType::Input) {
|
||||||
|
return session.GetInputCount();
|
||||||
|
} else if (namedValueType == NamedValueType::OverridableInitializer) {
|
||||||
|
return session.GetOverridableInitializerCount();
|
||||||
|
} else {
|
||||||
|
return session.GetOutputCount();
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
auto getName = [&session = *_session, namedValueType](size_t i, OrtAllocator* allocator) {
|
||||||
|
if (namedValueType == NamedValueType::Input) {
|
||||||
|
return session.GetInputNameAllocated(i, allocator);
|
||||||
|
} else if (namedValueType == NamedValueType::OverridableInitializer) {
|
||||||
|
return session.GetOverridableInitializerNameAllocated(i, allocator);
|
||||||
|
} else {
|
||||||
|
return session.GetOutputNameAllocated(i, allocator);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const size_t nameCount = getCount();
|
||||||
|
|
||||||
|
Ort::AllocatorWithDefaultOptions allocator;
|
||||||
|
NSMutableArray<NSString*>* result = [NSMutableArray arrayWithCapacity:nameCount];
|
||||||
|
|
||||||
|
for (size_t i = 0; i < nameCount; ++i) {
|
||||||
|
auto name = getName(i, allocator);
|
||||||
|
NSString* nameNsstr = [NSString stringWithUTF8String:name.get()];
|
||||||
|
NSAssert(nameNsstr != nil, @"nameNsstr must not be nil");
|
||||||
|
[result addObject:nameNsstr];
|
||||||
|
}
|
||||||
|
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTSessionOptions {
|
||||||
|
std::optional<Ort::SessionOptions> _sessionOptions;
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Public
|
||||||
|
|
||||||
|
- (nullable instancetype)initWithError:(NSError**)error {
|
||||||
|
if ((self = [super init]) == nil) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
_sessionOptions = Ort::SessionOptions{};
|
||||||
|
return self;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)appendExecutionProvider:(NSString*)providerName
|
||||||
|
providerOptions:(NSDictionary<NSString*, NSString*>*)providerOptions
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
std::unordered_map<std::string, std::string> options;
|
||||||
|
NSArray* keys = [providerOptions allKeys];
|
||||||
|
|
||||||
|
for (NSString* key in keys) {
|
||||||
|
NSString* value = [providerOptions objectForKey:key];
|
||||||
|
options.emplace(key.UTF8String, value.UTF8String);
|
||||||
|
}
|
||||||
|
|
||||||
|
_sessionOptions->AppendExecutionProvider(providerName.UTF8String, options);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setIntraOpNumThreads:(int)intraOpNumThreads
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->SetIntraOpNumThreads(intraOpNumThreads);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setGraphOptimizationLevel:(ORTGraphOptimizationLevel)graphOptimizationLevel
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->SetGraphOptimizationLevel(
|
||||||
|
PublicToCAPIGraphOptimizationLevel(graphOptimizationLevel));
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setOptimizedModelFilePath:(NSString*)optimizedModelFilePath
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->SetOptimizedModelFilePath(optimizedModelFilePath.UTF8String);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setLogID:(NSString*)logID
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->SetLogId(logID.UTF8String);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setLogSeverityLevel:(ORTLoggingLevel)loggingLevel
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->SetLogSeverityLevel(PublicToCAPILoggingLevel(loggingLevel));
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)addConfigEntryWithKey:(NSString*)key
|
||||||
|
value:(NSString*)value
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->AddConfigEntry(key.UTF8String, value.UTF8String);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)registerCustomOpsUsingFunction:(NSString*)registrationFuncName
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_sessionOptions->RegisterCustomOpsUsingFunction(registrationFuncName.UTF8String);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)registerCustomOpsUsingFunctionPointer:(ORTCAPIRegisterCustomOpsFnPtr)registerCustomOpsFn
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
if (!registerCustomOpsFn) {
|
||||||
|
ORT_CXX_API_THROW("registerCustomOpsFn must not be null", ORT_INVALID_ARGUMENT);
|
||||||
|
}
|
||||||
|
Ort::ThrowOnError((*registerCustomOpsFn)(static_cast<OrtSessionOptions*>(*_sessionOptions),
|
||||||
|
OrtGetApiBase()));
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Internal
|
||||||
|
|
||||||
|
- (Ort::SessionOptions&)CXXAPIOrtSessionOptions {
|
||||||
|
return *_sessionOptions;
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTRunOptions {
|
||||||
|
std::optional<Ort::RunOptions> _runOptions;
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Public
|
||||||
|
|
||||||
|
- (nullable instancetype)initWithError:(NSError**)error {
|
||||||
|
if ((self = [super init]) == nil) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
_runOptions = Ort::RunOptions{};
|
||||||
|
return self;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setLogTag:(NSString*)logTag
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_runOptions->SetRunTag(logTag.UTF8String);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)setLogSeverityLevel:(ORTLoggingLevel)loggingLevel
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_runOptions->SetRunLogSeverityLevel(PublicToCAPILoggingLevel(loggingLevel));
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (BOOL)addConfigEntryWithKey:(NSString*)key
|
||||||
|
value:(NSString*)value
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
_runOptions->AddConfigEntry(key.UTF8String, value.UTF8String);
|
||||||
|
return YES;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Internal
|
||||||
|
|
||||||
|
- (Ort::RunOptions&)CXXAPIOrtRunOptions {
|
||||||
|
return *_runOptions;
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_session.h"
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTSessionOptions ()
|
||||||
|
|
||||||
|
- (Ort::SessionOptions&)CXXAPIOrtSessionOptions;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@interface ORTRunOptions ()
|
||||||
|
|
||||||
|
- (Ort::RunOptions&)CXXAPIOrtRunOptions;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,224 @@
|
|||||||
|
// 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
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
// 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
|
||||||
@@ -0,0 +1,162 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_value_internal.h"
|
||||||
|
|
||||||
|
#include <optional>
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
#import "error_utils.h"
|
||||||
|
#import "ort_enums_internal.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
|
||||||
|
ORTTensorTypeAndShapeInfo* CXXAPIToPublicTensorTypeAndShapeInfo(
|
||||||
|
const Ort::ConstTensorTypeAndShapeInfo& CXXAPITensorTypeAndShapeInfo) {
|
||||||
|
auto* result = [[ORTTensorTypeAndShapeInfo alloc] init];
|
||||||
|
const auto elementType = CXXAPITensorTypeAndShapeInfo.GetElementType();
|
||||||
|
const std::vector<int64_t> shape = CXXAPITensorTypeAndShapeInfo.GetShape();
|
||||||
|
|
||||||
|
result.elementType = CAPIToPublicTensorElementType(elementType);
|
||||||
|
auto* shapeArray = [[NSMutableArray alloc] initWithCapacity:shape.size()];
|
||||||
|
for (size_t i = 0; i < shape.size(); ++i) {
|
||||||
|
shapeArray[i] = @(shape[i]);
|
||||||
|
}
|
||||||
|
result.shape = shapeArray;
|
||||||
|
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
ORTValueTypeInfo* CXXAPIToPublicValueTypeInfo(
|
||||||
|
const Ort::TypeInfo& CXXAPITypeInfo) {
|
||||||
|
auto* result = [[ORTValueTypeInfo alloc] init];
|
||||||
|
const auto valueType = CXXAPITypeInfo.GetONNXType();
|
||||||
|
|
||||||
|
result.type = CAPIToPublicValueType(valueType);
|
||||||
|
|
||||||
|
if (valueType == ONNX_TYPE_TENSOR) {
|
||||||
|
const auto tensorTypeAndShapeInfo = CXXAPITypeInfo.GetTensorTypeAndShapeInfo();
|
||||||
|
result.tensorTypeAndShapeInfo = CXXAPIToPublicTensorTypeAndShapeInfo(tensorTypeAndShapeInfo);
|
||||||
|
}
|
||||||
|
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
// out = a * b
|
||||||
|
// returns true iff the result does not overflow
|
||||||
|
bool SafeMultiply(size_t a, size_t b, size_t& out) {
|
||||||
|
return !__builtin_mul_overflow(a, b, &out);
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace
|
||||||
|
|
||||||
|
@interface ORTValue ()
|
||||||
|
|
||||||
|
// pointer to any external tensor data to keep alive for the lifetime of the ORTValue
|
||||||
|
@property(nonatomic, nullable) NSMutableData* externalTensorData;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTValue {
|
||||||
|
std::optional<Ort::Value> _value;
|
||||||
|
std::optional<Ort::TypeInfo> _typeInfo;
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Public
|
||||||
|
|
||||||
|
- (nullable instancetype)initWithTensorData:(NSMutableData*)tensorData
|
||||||
|
elementType:(ORTTensorElementDataType)elementType
|
||||||
|
shape:(NSArray<NSNumber*>*)shape
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
const auto memoryInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
|
||||||
|
const auto ONNXElementType = PublicToCAPITensorElementType(elementType);
|
||||||
|
const auto shapeVector = [shape]() {
|
||||||
|
std::vector<int64_t> result{};
|
||||||
|
result.reserve(shape.count);
|
||||||
|
for (NSNumber* dim in shape) {
|
||||||
|
result.push_back(dim.longLongValue);
|
||||||
|
}
|
||||||
|
return result;
|
||||||
|
}();
|
||||||
|
Ort::Value ortValue = Ort::Value::CreateTensor(
|
||||||
|
memoryInfo, tensorData.mutableBytes, tensorData.length,
|
||||||
|
shapeVector.data(), shapeVector.size(), ONNXElementType);
|
||||||
|
|
||||||
|
return [self initWithCXXAPIOrtValue:std::move(ortValue)
|
||||||
|
externalTensorData:tensorData
|
||||||
|
error:error];
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable ORTValueTypeInfo*)typeInfoWithError:(NSError**)error {
|
||||||
|
try {
|
||||||
|
return CXXAPIToPublicValueTypeInfo(*_typeInfo);
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable ORTTensorTypeAndShapeInfo*)tensorTypeAndShapeInfoWithError:(NSError**)error {
|
||||||
|
try {
|
||||||
|
const auto tensorTypeAndShapeInfo = _typeInfo->GetTensorTypeAndShapeInfo();
|
||||||
|
return CXXAPIToPublicTensorTypeAndShapeInfo(tensorTypeAndShapeInfo);
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
- (nullable NSMutableData*)tensorDataWithError:(NSError**)error {
|
||||||
|
try {
|
||||||
|
const auto tensorTypeAndShapeInfo = _typeInfo->GetTensorTypeAndShapeInfo();
|
||||||
|
const size_t elementCount = tensorTypeAndShapeInfo.GetElementCount();
|
||||||
|
const size_t elementSize = SizeOfCAPITensorElementType(tensorTypeAndShapeInfo.GetElementType());
|
||||||
|
size_t rawDataLength;
|
||||||
|
if (!SafeMultiply(elementCount, elementSize, rawDataLength)) {
|
||||||
|
ORT_CXX_API_THROW("failed to compute tensor data length", ORT_RUNTIME_EXCEPTION);
|
||||||
|
}
|
||||||
|
|
||||||
|
void* rawData;
|
||||||
|
Ort::ThrowOnError(Ort::GetApi().GetTensorMutableData(*_value, &rawData));
|
||||||
|
|
||||||
|
return [NSMutableData dataWithBytesNoCopy:rawData
|
||||||
|
length:rawDataLength
|
||||||
|
freeWhenDone:NO];
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error)
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma mark - Internal
|
||||||
|
|
||||||
|
- (nullable instancetype)initWithCXXAPIOrtValue:(Ort::Value&&)existingCXXAPIOrtValue
|
||||||
|
externalTensorData:(nullable NSMutableData*)externalTensorData
|
||||||
|
error:(NSError**)error {
|
||||||
|
if ((self = [super init]) == nil) {
|
||||||
|
return nil;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
_typeInfo = existingCXXAPIOrtValue.GetTypeInfo();
|
||||||
|
_externalTensorData = externalTensorData;
|
||||||
|
|
||||||
|
// transfer C++ Ort::Value ownership to this instance
|
||||||
|
_value = std::move(existingCXXAPIOrtValue);
|
||||||
|
return self;
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_NULLABLE(error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (Ort::Value&)CXXAPIOrtValue {
|
||||||
|
return *_value;
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTValueTypeInfo
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTTensorTypeAndShapeInfo
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_value.h"
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTValue ()
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates a value from an existing C++ API Ort::Value and takes ownership from it.
|
||||||
|
* Note: Ownership is guaranteed to be transferred on success but not otherwise.
|
||||||
|
*
|
||||||
|
* @param existingCXXAPIOrtValue The existing C++ API Ort::Value.
|
||||||
|
* @param externalTensorData Any external tensor data referenced by `existingCXXAPIOrtValue`.
|
||||||
|
* @param error Optional error information set if an error occurs.
|
||||||
|
* @return The instance, or nil if an error occurs.
|
||||||
|
*/
|
||||||
|
- (nullable instancetype)initWithCXXAPIOrtValue:(Ort::Value&&)existingCXXAPIOrtValue
|
||||||
|
externalTensorData:(nullable NSMutableData*)externalTensorData
|
||||||
|
error:(NSError**)error NS_DESIGNATED_INITIALIZER;
|
||||||
|
|
||||||
|
- (Ort::Value&)CXXAPIOrtValue;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "ort_xnnpack_execution_provider.h"
|
||||||
|
|
||||||
|
#import "cxx_api.h"
|
||||||
|
#import "error_utils.h"
|
||||||
|
#import "ort_session_internal.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@implementation ORTXnnpackExecutionProviderOptions
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTSessionOptions (ORTSessionOptionsXnnpackEP)
|
||||||
|
|
||||||
|
- (BOOL)appendXnnpackExecutionProviderWithOptions:(ORTXnnpackExecutionProviderOptions*)options
|
||||||
|
error:(NSError**)error {
|
||||||
|
try {
|
||||||
|
NSDictionary* provider_options = @{
|
||||||
|
@"intra_op_num_threads" : [NSString stringWithFormat:@"%d", options.intra_op_num_threads]
|
||||||
|
};
|
||||||
|
return [self appendExecutionProvider:@"XNNPACK" providerOptions:provider_options error:error];
|
||||||
|
}
|
||||||
|
ORT_OBJC_API_IMPL_CATCH_RETURNING_BOOL(error);
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
static_assert(__has_feature(objc_arc), "Objective-C ARC must be enabled.");
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
#define ORTAssertNullableResultSuccessful(result, error) \
|
||||||
|
do { \
|
||||||
|
XCTAssertNotNil(result, @"Expected non-nil result but got nil. Error: %@", error); \
|
||||||
|
XCTAssertNil(error); \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
#define ORTAssertBoolResultSuccessful(result, error) \
|
||||||
|
do { \
|
||||||
|
XCTAssertTrue(result, @"Expected true result but got false. Error: %@", error); \
|
||||||
|
XCTAssertNil(error); \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
#define ORTAssertNullableResultUnsuccessful(result, error) \
|
||||||
|
do { \
|
||||||
|
XCTAssertNil(result); \
|
||||||
|
XCTAssertNotNil(error); \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
#define ORTAssertBoolResultUnsuccessful(result, error) \
|
||||||
|
do { \
|
||||||
|
XCTAssertFalse(result); \
|
||||||
|
XCTAssertNotNil(error); \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
#define ORTAssertEqualFloatAndNoError(expected, result, error) \
|
||||||
|
do { \
|
||||||
|
XCTAssertEqualWithAccuracy(expected, result, 1e-3f, @"Expected %f but got %f. Error:%@", expected, result, error); \
|
||||||
|
XCTAssertNil(error); \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
#define ORTAssertEqualFloatArrays(expected, result) \
|
||||||
|
do { \
|
||||||
|
XCTAssertEqual(expected.count, result.count); \
|
||||||
|
for (size_t i = 0; i < expected.count; ++i) { \
|
||||||
|
XCTAssertEqualWithAccuracy([expected[i] floatValue], [result[i] floatValue], 1e-3f); \
|
||||||
|
} \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,117 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
#import "ort_checkpoint.h"
|
||||||
|
#import "ort_training_session.h"
|
||||||
|
#import "ort_env.h"
|
||||||
|
#import "ort_session.h"
|
||||||
|
|
||||||
|
#import "test/test_utils.h"
|
||||||
|
#import "test/assertion_utils.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTCheckpointTest : XCTestCase
|
||||||
|
@property(readonly, nullable) ORTEnv* ortEnv;
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTCheckpointTest
|
||||||
|
|
||||||
|
- (void)setUp {
|
||||||
|
[super setUp];
|
||||||
|
|
||||||
|
self.continueAfterFailure = NO;
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
_ortEnv = [[ORTEnv alloc] initWithLoggingLevel:ORTLoggingLevelWarning
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(_ortEnv, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (NSString*)getCheckpointPath {
|
||||||
|
NSBundle* bundle = [NSBundle bundleForClass:[ORTCheckpointTest class]];
|
||||||
|
NSString* path = [[bundle resourcePath] stringByAppendingPathComponent:@"checkpoint.ckpt"];
|
||||||
|
return path;
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (NSString*)getTrainingModelPath {
|
||||||
|
NSBundle* bundle = [NSBundle bundleForClass:[ORTCheckpointTest class]];
|
||||||
|
NSString* path = [[bundle resourcePath] stringByAppendingPathComponent:@"training_model.onnx"];
|
||||||
|
return path;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testSaveCheckpoint {
|
||||||
|
NSError* error = nil;
|
||||||
|
ORTCheckpoint* checkpoint = [[ORTCheckpoint alloc] initWithPath:[ORTCheckpointTest getCheckpointPath] error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(checkpoint, error);
|
||||||
|
|
||||||
|
// save checkpoint
|
||||||
|
NSString* path = [test_utils::createTemporaryDirectory(self) stringByAppendingPathComponent:@"save_checkpoint.ckpt"];
|
||||||
|
XCTAssertNotNil(path);
|
||||||
|
BOOL result = [checkpoint saveCheckpointToPath:path withOptimizerState:NO error:&error];
|
||||||
|
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitCheckpoint {
|
||||||
|
NSError* error = nil;
|
||||||
|
ORTCheckpoint* checkpoint = [[ORTCheckpoint alloc] initWithPath:[ORTCheckpointTest getCheckpointPath] error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(checkpoint, error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testIntProperty {
|
||||||
|
NSError* error = nil;
|
||||||
|
// Load checkpoint
|
||||||
|
ORTCheckpoint* checkpoint = [[ORTCheckpoint alloc] initWithPath:[ORTCheckpointTest getCheckpointPath] error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(checkpoint, error);
|
||||||
|
|
||||||
|
// Add property
|
||||||
|
BOOL result = [checkpoint addIntPropertyWithName:@"test" value:314 error:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
// Get property
|
||||||
|
int64_t value = [checkpoint getIntPropertyWithName:@"test" error:&error];
|
||||||
|
XCTAssertEqual(value, 314);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testFloatProperty {
|
||||||
|
NSError* error = nil;
|
||||||
|
// Load checkpoint
|
||||||
|
ORTCheckpoint* checkpoint = [[ORTCheckpoint alloc] initWithPath:[ORTCheckpointTest getCheckpointPath] error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(checkpoint, error);
|
||||||
|
|
||||||
|
// Add property
|
||||||
|
BOOL result = [checkpoint addFloatPropertyWithName:@"test" value:3.14f error:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
// Get property
|
||||||
|
float value = [checkpoint getFloatPropertyWithName:@"test" error:&error];
|
||||||
|
XCTAssertEqual(value, 3.14f);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testStringProperty {
|
||||||
|
NSError* error = nil;
|
||||||
|
// Load checkpoint
|
||||||
|
ORTCheckpoint* checkpoint = [[ORTCheckpoint alloc] initWithPath:[ORTCheckpointTest getCheckpointPath] error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(checkpoint, error);
|
||||||
|
|
||||||
|
// Add property
|
||||||
|
BOOL result = [checkpoint addStringPropertyWithName:@"test" value:@"hello" error:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
// Get property
|
||||||
|
NSString* value = [checkpoint getStringPropertyWithName:@"test" error:&error];
|
||||||
|
XCTAssertEqualObjects(value, @"hello");
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)tearDown {
|
||||||
|
_ortEnv = nil;
|
||||||
|
|
||||||
|
[super tearDown];
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
#import "ort_env.h"
|
||||||
|
|
||||||
|
#import "test/assertion_utils.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTEnvTest : XCTestCase
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTEnvTest
|
||||||
|
- (void)testGetOrtVersion {
|
||||||
|
NSString* ver = ORTVersion();
|
||||||
|
XCTAssertNotNil(ver);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitOk {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTEnv* env = [[ORTEnv alloc] initWithLoggingLevel:ORTLoggingLevelWarning
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(env, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,264 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
#import "ort_coreml_execution_provider.h"
|
||||||
|
#import "ort_xnnpack_execution_provider.h"
|
||||||
|
#import "ort_env.h"
|
||||||
|
#import "ort_session.h"
|
||||||
|
#import "ort_value.h"
|
||||||
|
|
||||||
|
#import "test/assertion_utils.h"
|
||||||
|
|
||||||
|
#include <vector>
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTSessionTest : XCTestCase
|
||||||
|
|
||||||
|
@property(readonly, nullable) ORTEnv* ortEnv;
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTSessionTest
|
||||||
|
|
||||||
|
- (void)setUp {
|
||||||
|
[super setUp];
|
||||||
|
|
||||||
|
self.continueAfterFailure = NO;
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
_ortEnv = [[ORTEnv alloc] initWithLoggingLevel:ORTLoggingLevelWarning
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(_ortEnv, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)tearDown {
|
||||||
|
_ortEnv = nil;
|
||||||
|
|
||||||
|
[super tearDown];
|
||||||
|
}
|
||||||
|
|
||||||
|
// model with an Add op
|
||||||
|
// inputs: A, B
|
||||||
|
// output: C = A + B
|
||||||
|
+ (NSString*)getAddModelPath {
|
||||||
|
NSBundle* bundle = [NSBundle bundleForClass:[ORTSessionTest class]];
|
||||||
|
NSString* path = [bundle pathForResource:@"single_add.basic"
|
||||||
|
ofType:@"ort"];
|
||||||
|
return path;
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (NSMutableData*)dataWithScalarFloat:(float)value {
|
||||||
|
NSMutableData* data = [[NSMutableData alloc] initWithBytes:&value length:sizeof(value)];
|
||||||
|
return data;
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (ORTValue*)ortValueWithScalarFloatData:(NSMutableData*)data {
|
||||||
|
NSArray<NSNumber*>* shape = @[ @1 ];
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTValue* ortValue = [[ORTValue alloc] initWithTensorData:data
|
||||||
|
elementType:ORTTensorElementDataTypeFloat
|
||||||
|
shape:shape
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(ortValue, err);
|
||||||
|
return ortValue;
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (ORTSessionOptions*)makeSessionOptions {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSessionOptions* sessionOptions = [[ORTSessionOptions alloc] initWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(sessionOptions, err);
|
||||||
|
return sessionOptions;
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (ORTRunOptions*)makeRunOptions {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTRunOptions* runOptions = [[ORTRunOptions alloc] initWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(runOptions, err);
|
||||||
|
return runOptions;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitAndRunWithPreallocatedOutputOk {
|
||||||
|
NSMutableData* aData = [ORTSessionTest dataWithScalarFloat:1.0f];
|
||||||
|
NSMutableData* bData = [ORTSessionTest dataWithScalarFloat:2.0f];
|
||||||
|
NSMutableData* cData = [ORTSessionTest dataWithScalarFloat:0.0f];
|
||||||
|
|
||||||
|
ORTValue* a = [ORTSessionTest ortValueWithScalarFloatData:aData];
|
||||||
|
ORTValue* b = [ORTSessionTest ortValueWithScalarFloatData:bData];
|
||||||
|
ORTValue* c = [ORTSessionTest ortValueWithScalarFloatData:cData];
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:[ORTSessionTest getAddModelPath]
|
||||||
|
sessionOptions:[ORTSessionTest makeSessionOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(session, err);
|
||||||
|
|
||||||
|
BOOL runResult = [session runWithInputs:@{@"A" : a, @"B" : b}
|
||||||
|
outputs:@{@"C" : c}
|
||||||
|
runOptions:[ORTSessionTest makeRunOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertBoolResultSuccessful(runResult, err);
|
||||||
|
|
||||||
|
const float cExpected = 3.0f;
|
||||||
|
float cActual;
|
||||||
|
memcpy(&cActual, cData.bytes, sizeof(float));
|
||||||
|
XCTAssertEqual(cActual, cExpected);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitAndRunOk {
|
||||||
|
NSMutableData* aData = [ORTSessionTest dataWithScalarFloat:1.0f];
|
||||||
|
NSMutableData* bData = [ORTSessionTest dataWithScalarFloat:2.0f];
|
||||||
|
|
||||||
|
ORTValue* a = [ORTSessionTest ortValueWithScalarFloatData:aData];
|
||||||
|
ORTValue* b = [ORTSessionTest ortValueWithScalarFloatData:bData];
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:[ORTSessionTest getAddModelPath]
|
||||||
|
sessionOptions:[ORTSessionTest makeSessionOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(session, err);
|
||||||
|
|
||||||
|
NSDictionary<NSString*, ORTValue*>* outputs =
|
||||||
|
[session runWithInputs:@{@"A" : a, @"B" : b}
|
||||||
|
outputNames:[NSSet setWithArray:@[ @"C" ]]
|
||||||
|
runOptions:[ORTSessionTest makeRunOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(outputs, err);
|
||||||
|
|
||||||
|
ORTValue* cOutput = outputs[@"C"];
|
||||||
|
XCTAssertNotNil(cOutput);
|
||||||
|
|
||||||
|
NSData* cData = [cOutput tensorDataWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(cData, err);
|
||||||
|
|
||||||
|
const float cExpected = 3.0f;
|
||||||
|
float cActual;
|
||||||
|
memcpy(&cActual, cData.bytes, sizeof(float));
|
||||||
|
XCTAssertEqual(cActual, cExpected);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testGetNamesOk {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:[ORTSessionTest getAddModelPath]
|
||||||
|
sessionOptions:[ORTSessionTest makeSessionOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(session, err);
|
||||||
|
|
||||||
|
NSArray<NSString*>* inputNames = [session inputNamesWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(inputNames, err);
|
||||||
|
XCTAssertEqualObjects(inputNames, (@[ @"A", @"B" ]));
|
||||||
|
|
||||||
|
NSArray<NSString*>* overridableInitializerNames = [session overridableInitializerNamesWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(overridableInitializerNames, err);
|
||||||
|
XCTAssertEqualObjects(overridableInitializerNames, (@[]));
|
||||||
|
|
||||||
|
NSArray<NSString*>* outputNames = [session outputNamesWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(outputNames, err);
|
||||||
|
XCTAssertEqualObjects(outputNames, (@[ @"C" ]));
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitFailsWithInvalidPath {
|
||||||
|
NSString* invalidModelPath = @"invalid/path/to/model.ort";
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:invalidModelPath
|
||||||
|
sessionOptions:[ORTSessionTest makeSessionOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultUnsuccessful(session, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testRunFailsWithInvalidInput {
|
||||||
|
NSMutableData* dData = [ORTSessionTest dataWithScalarFloat:1.0f];
|
||||||
|
NSMutableData* cData = [ORTSessionTest dataWithScalarFloat:0.0f];
|
||||||
|
|
||||||
|
ORTValue* d = [ORTSessionTest ortValueWithScalarFloatData:dData];
|
||||||
|
ORTValue* c = [ORTSessionTest ortValueWithScalarFloatData:cData];
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:[ORTSessionTest getAddModelPath]
|
||||||
|
sessionOptions:[ORTSessionTest makeSessionOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(session, err);
|
||||||
|
|
||||||
|
BOOL runResult = [session runWithInputs:@{@"D" : d}
|
||||||
|
outputs:@{@"C" : c}
|
||||||
|
runOptions:[ORTSessionTest makeRunOptions]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertBoolResultUnsuccessful(runResult, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testAppendCoreMLEP {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSessionOptions* sessionOptions = [ORTSessionTest makeSessionOptions];
|
||||||
|
ORTCoreMLExecutionProviderOptions* coreMLOptions = [[ORTCoreMLExecutionProviderOptions alloc] init];
|
||||||
|
coreMLOptions.enableOnSubgraphs = YES; // set an arbitrary option
|
||||||
|
|
||||||
|
BOOL appendResult = [sessionOptions appendCoreMLExecutionProviderWithOptions:coreMLOptions
|
||||||
|
error:&err];
|
||||||
|
|
||||||
|
if (!ORTIsCoreMLExecutionProviderAvailable()) {
|
||||||
|
ORTAssertBoolResultUnsuccessful(appendResult, err);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
ORTAssertBoolResultSuccessful(appendResult, err);
|
||||||
|
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:[ORTSessionTest getAddModelPath]
|
||||||
|
sessionOptions:sessionOptions
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(session, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testAppendXnnpackEP {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSessionOptions* sessionOptions = [ORTSessionTest makeSessionOptions];
|
||||||
|
ORTXnnpackExecutionProviderOptions* XnnpackOptions = [[ORTXnnpackExecutionProviderOptions alloc] init];
|
||||||
|
XnnpackOptions.intra_op_num_threads = 2;
|
||||||
|
|
||||||
|
BOOL appendResult = [sessionOptions appendXnnpackExecutionProviderWithOptions:XnnpackOptions
|
||||||
|
error:&err];
|
||||||
|
// Without xnnpack EP in building also can pass the test
|
||||||
|
NSString* err_msg = [err localizedDescription];
|
||||||
|
if (!appendResult && [err_msg containsString:@"XNNPACK execution provider is not supported in this build. "]) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
ORTAssertBoolResultSuccessful(appendResult, err);
|
||||||
|
|
||||||
|
ORTSession* session = [[ORTSession alloc] initWithEnv:self.ortEnv
|
||||||
|
modelPath:[ORTSessionTest getAddModelPath]
|
||||||
|
sessionOptions:sessionOptions
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(session, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
static bool gDummyRegisterCustomOpsFnCalled = false;
|
||||||
|
|
||||||
|
static OrtStatus* _Nullable DummyRegisterCustomOpsFn(OrtSessionOptions* /*session_options*/,
|
||||||
|
const OrtApiBase* /*api*/) {
|
||||||
|
gDummyRegisterCustomOpsFnCalled = true;
|
||||||
|
return nullptr;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testRegisterCustomOpsUsingFunctionPointer {
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTSessionOptions* sessionOptions = [ORTSessionTest makeSessionOptions];
|
||||||
|
|
||||||
|
gDummyRegisterCustomOpsFnCalled = false;
|
||||||
|
BOOL registerResult = [sessionOptions registerCustomOpsUsingFunctionPointer:&DummyRegisterCustomOpsFn
|
||||||
|
error:&err];
|
||||||
|
ORTAssertBoolResultSuccessful(registerResult, err);
|
||||||
|
|
||||||
|
XCTAssertEqual(gDummyRegisterCustomOpsFnCalled, true);
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,359 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
#import "ort_checkpoint.h"
|
||||||
|
#import "ort_training_session.h"
|
||||||
|
#import "ort_env.h"
|
||||||
|
#import "ort_session.h"
|
||||||
|
#import "ort_value.h"
|
||||||
|
|
||||||
|
#import "test/test_utils.h"
|
||||||
|
#import "test/assertion_utils.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTTrainingSessionTest : XCTestCase
|
||||||
|
@property(readonly, nullable) ORTEnv* ortEnv;
|
||||||
|
@property(readonly, nullable) ORTCheckpoint* checkpoint;
|
||||||
|
@property(readonly, nullable) ORTTrainingSession* session;
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTTrainingSessionTest
|
||||||
|
|
||||||
|
- (void)setUp {
|
||||||
|
[super setUp];
|
||||||
|
|
||||||
|
self.continueAfterFailure = NO;
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
_ortEnv = [[ORTEnv alloc] initWithLoggingLevel:ORTLoggingLevelWarning
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(_ortEnv, err);
|
||||||
|
_checkpoint = [[ORTCheckpoint alloc] initWithPath:[ORTTrainingSessionTest
|
||||||
|
getFilePathFromName:@"checkpoint.ckpt"]
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(_checkpoint, err);
|
||||||
|
_session = [self makeTrainingSessionWithCheckpoint:_checkpoint];
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (NSString*)getFilePathFromName:(NSString*)name {
|
||||||
|
NSBundle* bundle = [NSBundle bundleForClass:[ORTTrainingSessionTest class]];
|
||||||
|
NSString* path = [[bundle resourcePath] stringByAppendingPathComponent:name];
|
||||||
|
return path;
|
||||||
|
}
|
||||||
|
|
||||||
|
+ (NSMutableData*)loadTensorDataFromFile:(NSString*)filePath skipHeader:(BOOL)skipHeader {
|
||||||
|
NSError* error = nil;
|
||||||
|
NSString* fileContents = [NSString stringWithContentsOfFile:filePath
|
||||||
|
encoding:NSUTF8StringEncoding
|
||||||
|
error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(fileContents, error);
|
||||||
|
|
||||||
|
NSArray<NSString*>* lines = [fileContents componentsSeparatedByCharactersInSet:[NSCharacterSet newlineCharacterSet]];
|
||||||
|
|
||||||
|
if (skipHeader) {
|
||||||
|
lines = [lines subarrayWithRange:NSMakeRange(1, lines.count - 1)];
|
||||||
|
}
|
||||||
|
|
||||||
|
NSArray<NSString*>* dataArray = [lines[0] componentsSeparatedByCharactersInSet:
|
||||||
|
[NSCharacterSet characterSetWithCharactersInString:@",[] "]];
|
||||||
|
NSMutableData* tensorData = [NSMutableData data];
|
||||||
|
|
||||||
|
for (NSString* str in dataArray) {
|
||||||
|
if (str.length > 0) {
|
||||||
|
float value = [str floatValue];
|
||||||
|
[tensorData appendBytes:&value length:sizeof(float)];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return tensorData;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (ORTTrainingSession*)makeTrainingSessionWithCheckpoint:(ORTCheckpoint*)checkpoint {
|
||||||
|
NSError* error = nil;
|
||||||
|
ORTSessionOptions* sessionOptions = [[ORTSessionOptions alloc] initWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(sessionOptions, error);
|
||||||
|
|
||||||
|
ORTTrainingSession* session = [[ORTTrainingSession alloc]
|
||||||
|
initWithEnv:self.ortEnv
|
||||||
|
sessionOptions:sessionOptions
|
||||||
|
checkpoint:checkpoint
|
||||||
|
trainModelPath:[ORTTrainingSessionTest getFilePathFromName:@"training_model.onnx"]
|
||||||
|
evalModelPath:[ORTTrainingSessionTest getFilePathFromName:@"eval_model.onnx"]
|
||||||
|
optimizerModelPath:[ORTTrainingSessionTest getFilePathFromName:@"adamw.onnx"]
|
||||||
|
error:&error];
|
||||||
|
|
||||||
|
ORTAssertNullableResultSuccessful(session, error);
|
||||||
|
return session;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitTrainingSession {
|
||||||
|
NSError* error = nil;
|
||||||
|
|
||||||
|
// check that inputNames contains input-0
|
||||||
|
NSArray<NSString*>* inputNames = [self.session getTrainInputNamesWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(inputNames, error);
|
||||||
|
|
||||||
|
XCTAssertTrue(inputNames.count > 0);
|
||||||
|
XCTAssertTrue([inputNames containsObject:@"input-0"]);
|
||||||
|
|
||||||
|
// check that outNames contains onnx::loss::21273
|
||||||
|
NSArray<NSString*>* outputNames = [self.session getTrainOutputNamesWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputNames, error);
|
||||||
|
|
||||||
|
XCTAssertTrue(outputNames.count > 0);
|
||||||
|
XCTAssertTrue([outputNames containsObject:@"onnx::loss::21273"]);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitTrainingSessionWithEval {
|
||||||
|
NSError* error = nil;
|
||||||
|
|
||||||
|
// check that inputNames contains input-0
|
||||||
|
NSArray<NSString*>* inputNames = [self.session getEvalInputNamesWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(inputNames, error);
|
||||||
|
|
||||||
|
XCTAssertTrue(inputNames.count > 0);
|
||||||
|
XCTAssertTrue([inputNames containsObject:@"input-0"]);
|
||||||
|
|
||||||
|
// check that outNames contains onnx::loss::21273
|
||||||
|
NSArray<NSString*>* outputNames = [self.session getEvalOutputNamesWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputNames, error);
|
||||||
|
|
||||||
|
XCTAssertTrue(outputNames.count > 0);
|
||||||
|
XCTAssertTrue([outputNames containsObject:@"onnx::loss::21273"]);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)runTrainStep {
|
||||||
|
// load input and expected output
|
||||||
|
NSError* error = nil;
|
||||||
|
NSMutableData* expectedOutput = [ORTTrainingSessionTest loadTensorDataFromFile:[ORTTrainingSessionTest
|
||||||
|
getFilePathFromName:@"loss_1.out"]
|
||||||
|
skipHeader:YES];
|
||||||
|
|
||||||
|
NSMutableData* input = [ORTTrainingSessionTest loadTensorDataFromFile:[ORTTrainingSessionTest
|
||||||
|
getFilePathFromName:@"input-0.in"]
|
||||||
|
skipHeader:YES];
|
||||||
|
|
||||||
|
int32_t labels[] = {1, 1};
|
||||||
|
|
||||||
|
// create ORTValue array for input and labels
|
||||||
|
NSMutableArray<ORTValue*>* inputValues = [NSMutableArray array];
|
||||||
|
|
||||||
|
ORTValue* inputTensor = [[ORTValue alloc] initWithTensorData:input
|
||||||
|
elementType:ORTTensorElementDataTypeFloat
|
||||||
|
shape:@[ @2, @784 ]
|
||||||
|
error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(inputTensor, error);
|
||||||
|
[inputValues addObject:inputTensor];
|
||||||
|
|
||||||
|
ORTValue* labelTensor = [[ORTValue alloc] initWithTensorData:[NSMutableData dataWithBytes:labels
|
||||||
|
length:sizeof(labels)]
|
||||||
|
elementType:ORTTensorElementDataTypeInt32
|
||||||
|
shape:@[ @2 ]
|
||||||
|
error:&error];
|
||||||
|
|
||||||
|
ORTAssertNullableResultSuccessful(labelTensor, error);
|
||||||
|
[inputValues addObject:labelTensor];
|
||||||
|
|
||||||
|
NSArray<ORTValue*>* outputs = [self.session trainStepWithInputValues:inputValues error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputs, error);
|
||||||
|
XCTAssertTrue(outputs.count > 0);
|
||||||
|
|
||||||
|
BOOL result = [self.session lazyResetGradWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
outputs = [self.session trainStepWithInputValues:inputValues error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputs, error);
|
||||||
|
XCTAssertTrue(outputs.count > 0);
|
||||||
|
|
||||||
|
ORTValue* outputValue = outputs[0];
|
||||||
|
ORTValueTypeInfo* typeInfo = [outputValue typeInfoWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(typeInfo, error);
|
||||||
|
XCTAssertEqual(typeInfo.type, ORTValueTypeTensor);
|
||||||
|
XCTAssertNotNil(typeInfo.tensorTypeAndShapeInfo);
|
||||||
|
|
||||||
|
ORTTensorTypeAndShapeInfo* tensorInfo = [outputValue tensorTypeAndShapeInfoWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(tensorInfo, error);
|
||||||
|
XCTAssertEqual(tensorInfo.elementType, ORTTensorElementDataTypeFloat);
|
||||||
|
|
||||||
|
NSMutableData* tensorData = [outputValue tensorDataWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(tensorData, error);
|
||||||
|
ORTAssertEqualFloatArrays(test_utils::getFloatArrayFromData(tensorData),
|
||||||
|
test_utils::getFloatArrayFromData(expectedOutput));
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testTrainStepOutput {
|
||||||
|
[self runTrainStep];
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testOptimizerStep {
|
||||||
|
// load input and expected output
|
||||||
|
NSError* error = nil;
|
||||||
|
NSMutableData* expectedOutput1 = [ORTTrainingSessionTest loadTensorDataFromFile:[ORTTrainingSessionTest
|
||||||
|
getFilePathFromName:@"loss_1.out"]
|
||||||
|
skipHeader:YES];
|
||||||
|
|
||||||
|
NSMutableData* expectedOutput2 = [ORTTrainingSessionTest loadTensorDataFromFile:[ORTTrainingSessionTest
|
||||||
|
getFilePathFromName:@"loss_2.out"]
|
||||||
|
skipHeader:YES];
|
||||||
|
|
||||||
|
NSMutableData* input = [ORTTrainingSessionTest loadTensorDataFromFile:[ORTTrainingSessionTest
|
||||||
|
getFilePathFromName:@"input-0.in"]
|
||||||
|
skipHeader:YES];
|
||||||
|
|
||||||
|
int32_t labels[] = {1, 1};
|
||||||
|
|
||||||
|
// create ORTValue array for input and labels
|
||||||
|
NSMutableArray<ORTValue*>* inputValues = [NSMutableArray array];
|
||||||
|
|
||||||
|
ORTValue* inputTensor = [[ORTValue alloc] initWithTensorData:input
|
||||||
|
elementType:ORTTensorElementDataTypeFloat
|
||||||
|
shape:@[ @2, @784 ]
|
||||||
|
error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(inputTensor, error);
|
||||||
|
[inputValues addObject:inputTensor];
|
||||||
|
|
||||||
|
ORTValue* labelTensor = [[ORTValue alloc] initWithTensorData:[NSMutableData dataWithBytes:labels
|
||||||
|
length:sizeof(labels)]
|
||||||
|
elementType:ORTTensorElementDataTypeInt32
|
||||||
|
shape:@[ @2 ]
|
||||||
|
error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(labelTensor, error);
|
||||||
|
[inputValues addObject:labelTensor];
|
||||||
|
|
||||||
|
// run train step, optimizer steps and check loss
|
||||||
|
NSArray<ORTValue*>* outputs = [self.session trainStepWithInputValues:inputValues error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputs, error);
|
||||||
|
|
||||||
|
NSMutableData* loss = [outputs[0] tensorDataWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(loss, error);
|
||||||
|
ORTAssertEqualFloatArrays(test_utils::getFloatArrayFromData(loss),
|
||||||
|
test_utils::getFloatArrayFromData(expectedOutput1));
|
||||||
|
|
||||||
|
BOOL result = [self.session lazyResetGradWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
outputs = [self.session trainStepWithInputValues:inputValues error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputs, error);
|
||||||
|
|
||||||
|
loss = [outputs[0] tensorDataWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(loss, error);
|
||||||
|
ORTAssertEqualFloatArrays(test_utils::getFloatArrayFromData(loss),
|
||||||
|
test_utils::getFloatArrayFromData(expectedOutput1));
|
||||||
|
|
||||||
|
result = [self.session optimizerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
outputs = [self.session trainStepWithInputValues:inputValues error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(outputs, error);
|
||||||
|
loss = [outputs[0] tensorDataWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(loss, error);
|
||||||
|
ORTAssertEqualFloatArrays(test_utils::getFloatArrayFromData(loss),
|
||||||
|
test_utils::getFloatArrayFromData(expectedOutput2));
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testSetLearningRate {
|
||||||
|
NSError* error = nil;
|
||||||
|
|
||||||
|
float learningRate = 0.1f;
|
||||||
|
BOOL result = [self.session setLearningRate:learningRate error:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
float actualLearningRate = [self.session getLearningRateWithError:&error];
|
||||||
|
ORTAssertEqualFloatAndNoError(learningRate, actualLearningRate, error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testLinearLRScheduler {
|
||||||
|
NSError* error = nil;
|
||||||
|
|
||||||
|
float learningRate = 0.1f;
|
||||||
|
BOOL result = [self.session registerLinearLRSchedulerWithWarmupStepCount:2
|
||||||
|
totalStepCount:4
|
||||||
|
initialLr:learningRate
|
||||||
|
error:&error];
|
||||||
|
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
|
||||||
|
[self runTrainStep];
|
||||||
|
|
||||||
|
result = [self.session optimizerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
result = [self.session schedulerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
ORTAssertEqualFloatAndNoError(0.05f, [self.session getLearningRateWithError:&error], error);
|
||||||
|
|
||||||
|
result = [self.session optimizerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
result = [self.session schedulerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
ORTAssertEqualFloatAndNoError(0.1f, [self.session getLearningRateWithError:&error], error);
|
||||||
|
|
||||||
|
result = [self.session optimizerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
result = [self.session schedulerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
ORTAssertEqualFloatAndNoError(0.05f, [self.session getLearningRateWithError:&error], error);
|
||||||
|
|
||||||
|
result = [self.session optimizerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
result = [self.session schedulerStepWithError:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
ORTAssertEqualFloatAndNoError(0.0f, [self.session getLearningRateWithError:&error], error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testExportModelForInference {
|
||||||
|
NSError* error = nil;
|
||||||
|
|
||||||
|
NSString* inferenceModelPath = [test_utils::createTemporaryDirectory(self)
|
||||||
|
stringByAppendingPathComponent:@"inference_model.onnx"];
|
||||||
|
XCTAssertNotNil(inferenceModelPath);
|
||||||
|
|
||||||
|
NSArray<NSString*>* graphOutputNames = [NSArray arrayWithObjects:@"output-0", nil];
|
||||||
|
|
||||||
|
BOOL result = [self.session exportModelForInferenceWithOutputPath:inferenceModelPath
|
||||||
|
graphOutputNames:graphOutputNames
|
||||||
|
error:&error];
|
||||||
|
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
XCTAssertTrue([[NSFileManager defaultManager] fileExistsAtPath:inferenceModelPath]);
|
||||||
|
|
||||||
|
[self addTeardownBlock:^{
|
||||||
|
NSError* error = nil;
|
||||||
|
[[NSFileManager defaultManager] removeItemAtPath:inferenceModelPath error:&error];
|
||||||
|
}];
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testToBuffer {
|
||||||
|
NSError* error = nil;
|
||||||
|
ORTValue* buffer = [self.session toBufferWithTrainable:YES error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(buffer, error);
|
||||||
|
|
||||||
|
ORTValueTypeInfo* typeInfo = [buffer typeInfoWithError:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(typeInfo, error);
|
||||||
|
XCTAssertEqual(typeInfo.type, ORTValueTypeTensor);
|
||||||
|
XCTAssertNotNil(typeInfo.tensorTypeAndShapeInfo);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testFromBuffer {
|
||||||
|
NSError* error = nil;
|
||||||
|
|
||||||
|
ORTValue* buffer = [self.session toBufferWithTrainable:YES error:&error];
|
||||||
|
ORTAssertNullableResultSuccessful(buffer, error);
|
||||||
|
|
||||||
|
BOOL result = [self.session fromBufferWithValue:buffer error:&error];
|
||||||
|
ORTAssertBoolResultSuccessful(result, error);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)tearDown {
|
||||||
|
_session = nil;
|
||||||
|
_checkpoint = nil;
|
||||||
|
_ortEnv = nil;
|
||||||
|
|
||||||
|
[super tearDown];
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
#import "ort_training_session.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTTrainingUtilsTest : XCTestCase
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTTrainingUtilsTest
|
||||||
|
|
||||||
|
- (void)setUp {
|
||||||
|
[super setUp];
|
||||||
|
|
||||||
|
self.continueAfterFailure = NO;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testSetSeed {
|
||||||
|
ORTSetSeed(2718);
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
#import "ort_value.h"
|
||||||
|
|
||||||
|
#include <vector>
|
||||||
|
|
||||||
|
#import "test/assertion_utils.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
@interface ORTValueTest : XCTestCase
|
||||||
|
@end
|
||||||
|
|
||||||
|
@implementation ORTValueTest
|
||||||
|
|
||||||
|
- (void)setUp {
|
||||||
|
[super setUp];
|
||||||
|
|
||||||
|
self.continueAfterFailure = NO;
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitTensorOk {
|
||||||
|
int32_t value = 42;
|
||||||
|
NSMutableData* data = [[NSMutableData alloc] initWithBytes:&value
|
||||||
|
length:sizeof(int32_t)];
|
||||||
|
NSArray<NSNumber*>* shape = @[ @1 ];
|
||||||
|
|
||||||
|
const ORTTensorElementDataType elementType = ORTTensorElementDataTypeInt32;
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTValue* ortValue = [[ORTValue alloc] initWithTensorData:data
|
||||||
|
elementType:elementType
|
||||||
|
shape:shape
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(ortValue, err);
|
||||||
|
|
||||||
|
auto checkTensorInfo = [&](ORTTensorTypeAndShapeInfo* tensorInfo) {
|
||||||
|
XCTAssertEqual(tensorInfo.elementType, elementType);
|
||||||
|
XCTAssertEqualObjects(tensorInfo.shape, shape);
|
||||||
|
};
|
||||||
|
|
||||||
|
ORTValueTypeInfo* typeInfo = [ortValue typeInfoWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(typeInfo, err);
|
||||||
|
XCTAssertEqual(typeInfo.type, ORTValueTypeTensor);
|
||||||
|
XCTAssertNotNil(typeInfo.tensorTypeAndShapeInfo);
|
||||||
|
checkTensorInfo(typeInfo.tensorTypeAndShapeInfo);
|
||||||
|
|
||||||
|
ORTTensorTypeAndShapeInfo* tensorInfo = [ortValue tensorTypeAndShapeInfoWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(tensorInfo, err);
|
||||||
|
checkTensorInfo(tensorInfo);
|
||||||
|
|
||||||
|
NSData* actualData = [ortValue tensorDataWithError:&err];
|
||||||
|
ORTAssertNullableResultSuccessful(actualData, err);
|
||||||
|
XCTAssertEqual(actualData.length, sizeof(int32_t));
|
||||||
|
int32_t actualValue;
|
||||||
|
memcpy(&actualValue, actualData.bytes, sizeof(int32_t));
|
||||||
|
XCTAssertEqual(actualValue, value);
|
||||||
|
}
|
||||||
|
|
||||||
|
- (void)testInitTensorFailsWithDataSmallerThanShape {
|
||||||
|
std::vector<int32_t> values{1, 2, 3, 4};
|
||||||
|
NSMutableData* data = [[NSMutableData alloc] initWithBytes:values.data()
|
||||||
|
length:values.size() * sizeof(int32_t)];
|
||||||
|
NSArray<NSNumber*>* shape = @[ @2, @3 ]; // too large
|
||||||
|
|
||||||
|
NSError* err = nil;
|
||||||
|
ORTValue* ortValue = [[ORTValue alloc] initWithTensorData:data
|
||||||
|
elementType:ORTTensorElementDataTypeInt32
|
||||||
|
shape:shape
|
||||||
|
error:&err];
|
||||||
|
ORTAssertNullableResultUnsuccessful(ortValue, err);
|
||||||
|
}
|
||||||
|
|
||||||
|
@end
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import <Foundation/Foundation.h>
|
||||||
|
#import <XCTest/XCTest.h>
|
||||||
|
|
||||||
|
#import "ort_session.h"
|
||||||
|
#import "ort_env.h"
|
||||||
|
#import "ort_value.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
namespace test_utils {
|
||||||
|
|
||||||
|
NSString* _Nullable createTemporaryDirectory(XCTestCase* testCase);
|
||||||
|
|
||||||
|
NSArray<NSNumber*>* getFloatArrayFromData(NSData* data);
|
||||||
|
|
||||||
|
} // namespace test_utils
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
#import "test_utils.h"
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_BEGIN
|
||||||
|
|
||||||
|
namespace test_utils {
|
||||||
|
|
||||||
|
NSString* createTemporaryDirectory(XCTestCase* testCase) {
|
||||||
|
NSString* temporaryDirectory = NSTemporaryDirectory();
|
||||||
|
NSString* directoryPath = [temporaryDirectory stringByAppendingPathComponent:@"ort-objective-c-test"];
|
||||||
|
|
||||||
|
NSError* error = nil;
|
||||||
|
[[NSFileManager defaultManager] createDirectoryAtPath:directoryPath
|
||||||
|
withIntermediateDirectories:YES
|
||||||
|
attributes:nil
|
||||||
|
error:&error];
|
||||||
|
|
||||||
|
XCTAssertNil(error, @"Error creating temporary directory: %@", error.localizedDescription);
|
||||||
|
|
||||||
|
// add teardown block to delete the temporary directory
|
||||||
|
[testCase addTeardownBlock:^{
|
||||||
|
NSError* error = nil;
|
||||||
|
[[NSFileManager defaultManager] removeItemAtPath:directoryPath error:&error];
|
||||||
|
XCTAssertNil(error, @"Error removing temporary directory: %@", error.localizedDescription);
|
||||||
|
}];
|
||||||
|
|
||||||
|
return directoryPath;
|
||||||
|
}
|
||||||
|
|
||||||
|
NSArray<NSNumber*>* getFloatArrayFromData(NSData* data) {
|
||||||
|
NSMutableArray<NSNumber*>* array = [NSMutableArray array];
|
||||||
|
float value;
|
||||||
|
for (size_t i = 0; i < data.length / sizeof(float); ++i) {
|
||||||
|
[data getBytes:&value range:NSMakeRange(i * sizeof(float), sizeof(float))];
|
||||||
|
[array addObject:[NSNumber numberWithFloat:value]];
|
||||||
|
}
|
||||||
|
return array;
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace test_utils
|
||||||
|
|
||||||
|
NS_ASSUME_NONNULL_END
|
||||||
+12
@@ -0,0 +1,12 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
set -e
|
||||||
|
|
||||||
|
# Get directory this script is in
|
||||||
|
DIR="$( cd "$( dirname "${BASH_SOURCE[0]}" )" && pwd )"
|
||||||
|
|
||||||
|
cd ${DIR}
|
||||||
|
|
||||||
|
python3 ./single_add_gen.py
|
||||||
|
|
||||||
|
ORT_CONVERT_ONNX_MODELS_TO_ORT_OPTIMIZATION_LEVEL=basic python3 -m onnxruntime.tools.convert_onnx_models_to_ort --optimization_style=Fixed .
|
||||||
BIN
Binary file not shown.
BIN
Binary file not shown.
+19
@@ -0,0 +1,19 @@
|
|||||||
|
import onnx
|
||||||
|
from onnx import TensorProto, helper
|
||||||
|
|
||||||
|
graph = helper.make_graph(
|
||||||
|
[ # nodes
|
||||||
|
helper.make_node("Add", ["A", "B"], ["C"], "Add"),
|
||||||
|
],
|
||||||
|
"SingleAdd", # name
|
||||||
|
[ # inputs
|
||||||
|
helper.make_tensor_value_info("A", TensorProto.FLOAT, [1]),
|
||||||
|
helper.make_tensor_value_info("B", TensorProto.FLOAT, [1]),
|
||||||
|
],
|
||||||
|
[ # outputs
|
||||||
|
helper.make_tensor_value_info("C", TensorProto.FLOAT, [1]),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 12)])
|
||||||
|
onnx.save(model, r"single_add.onnx")
|
||||||
Vendored
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,62 @@
|
|||||||
|
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||||
|
// Licensed under the MIT License.
|
||||||
|
|
||||||
|
import XCTest
|
||||||
|
import Foundation
|
||||||
|
@testable import OnnxRuntimeBindings
|
||||||
|
|
||||||
|
final class SwiftOnnxRuntimeBindingsTests: XCTestCase {
|
||||||
|
let modelPath: String = Bundle.module.url(forResource: "single_add.basic", withExtension: "ort")!.path
|
||||||
|
|
||||||
|
func testGetVersionString() throws {
|
||||||
|
do {
|
||||||
|
let version = ORTVersion()
|
||||||
|
XCTAssertNotNil(version)
|
||||||
|
} catch let error {
|
||||||
|
XCTFail(error.localizedDescription)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func testCreateSession() throws {
|
||||||
|
do {
|
||||||
|
let env = try ORTEnv(loggingLevel: ORTLoggingLevel.verbose)
|
||||||
|
let options = try ORTSessionOptions()
|
||||||
|
try options.setLogSeverityLevel(ORTLoggingLevel.verbose)
|
||||||
|
try options.setIntraOpNumThreads(1)
|
||||||
|
// Create the ORTSession
|
||||||
|
_ = try ORTSession(env: env, modelPath: modelPath, sessionOptions: options)
|
||||||
|
} catch let error {
|
||||||
|
XCTFail(error.localizedDescription)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func testAppendCoreMLEP() throws {
|
||||||
|
do {
|
||||||
|
let env = try ORTEnv(loggingLevel: ORTLoggingLevel.verbose)
|
||||||
|
let sessionOptions: ORTSessionOptions = try ORTSessionOptions()
|
||||||
|
let coreMLOptions: ORTCoreMLExecutionProviderOptions = ORTCoreMLExecutionProviderOptions()
|
||||||
|
coreMLOptions.enableOnSubgraphs = true
|
||||||
|
try sessionOptions.appendCoreMLExecutionProvider(with: coreMLOptions)
|
||||||
|
|
||||||
|
XCTAssertTrue(ORTIsCoreMLExecutionProviderAvailable())
|
||||||
|
_ = try ORTSession(env: env, modelPath: modelPath, sessionOptions: sessionOptions)
|
||||||
|
} catch let error {
|
||||||
|
XCTFail(error.localizedDescription)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func testAppendXnnpackEP() throws {
|
||||||
|
do {
|
||||||
|
let env = try ORTEnv(loggingLevel: ORTLoggingLevel.verbose)
|
||||||
|
let sessionOptions: ORTSessionOptions = try ORTSessionOptions()
|
||||||
|
let XnnpackOptions: ORTXnnpackExecutionProviderOptions = ORTXnnpackExecutionProviderOptions()
|
||||||
|
XnnpackOptions.intra_op_num_threads = 2
|
||||||
|
try sessionOptions.appendXnnpackExecutionProvider(with: XnnpackOptions)
|
||||||
|
|
||||||
|
XCTAssertTrue(ORTIsCoreMLExecutionProviderAvailable())
|
||||||
|
_ = try ORTSession(env: env, modelPath: modelPath, sessionOptions: sessionOptions)
|
||||||
|
} catch let error {
|
||||||
|
XCTFail(error.localizedDescription)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user