More C API changes (#259)

* More API changes, remove 'Inference' from function names. Remove enum values. Make Status match other types.

* Switch to bool instead of int, and remove stdbool
This commit is contained in:
Ryan Hill 2018-12-28 14:53:19 -08:00 committed by Changming Sun
parent 29d03ffb08
commit 6a090985fb
14 changed files with 154 additions and 172 deletions

View file

@ -44,7 +44,7 @@ namespace Microsoft.ML.OnnxRuntime
_nativeHandle = IntPtr.Zero;
try
{
NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateInferenceSession(envHandle, modelPath, options.NativeHandle, out _nativeHandle));
NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateSession(envHandle, modelPath, options.NativeHandle, out _nativeHandle));
// Initialize input/output metadata
_inputMetadata = new Dictionary<string, NodeMetadata>();
@ -52,7 +52,7 @@ namespace Microsoft.ML.OnnxRuntime
// get input count
ulong inputCount = 0;
NativeApiStatus.VerifySuccess(NativeMethods.OrtInferenceSessionGetInputCount(_nativeHandle, out inputCount));
NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetInputCount(_nativeHandle, out inputCount));
// get all the output names
for (ulong i = 0; i < inputCount; i++)
@ -62,7 +62,7 @@ namespace Microsoft.ML.OnnxRuntime
// get output count
ulong outputCount = 0;
NativeApiStatus.VerifySuccess(NativeMethods.OrtInferenceSessionGetOutputCount(_nativeHandle, out outputCount));
NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOutputCount(_nativeHandle, out outputCount));
// get all the output names
for (ulong i = 0; i < outputCount; i++)
@ -148,7 +148,7 @@ namespace Microsoft.ML.OnnxRuntime
string[] outputNamesArray = outputNames.ToArray();
IntPtr[] outputValueArray = new IntPtr[outputNames.Count];
IntPtr status = NativeMethods.OrtRunInference(
IntPtr status = NativeMethods.OrtRun(
this._nativeHandle,
IntPtr.Zero, // TODO: use Run options when Run options creation API is available
// Passing null uses the default run options in the C-api
@ -212,7 +212,7 @@ namespace Microsoft.ML.OnnxRuntime
IntPtr nameHandle = IntPtr.Zero;
string str = null;
IntPtr status = NativeMethods.OrtInferenceSessionGetOutputName(
IntPtr status = NativeMethods.OrtSessionGetOutputName(
_nativeHandle,
index,
NativeMemoryAllocator.DefaultInstance.Handle,
@ -238,7 +238,7 @@ namespace Microsoft.ML.OnnxRuntime
IntPtr nameHandle = IntPtr.Zero;
string str = null;
IntPtr status = NativeMethods.OrtInferenceSessionGetInputName(
IntPtr status = NativeMethods.OrtSessionGetInputName(
_nativeHandle,
index,
NativeMemoryAllocator.DefaultInstance.Handle,
@ -265,7 +265,7 @@ namespace Microsoft.ML.OnnxRuntime
IntPtr typeInfo = IntPtr.Zero;
try
{
NativeApiStatus.VerifySuccess(NativeMethods.OrtInferenceSessionGetInputTypeInfo(_nativeHandle, index, out typeInfo));
NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetInputTypeInfo(_nativeHandle, index, out typeInfo));
return GetMetadataFromTypeInfo(typeInfo);
}
finally
@ -282,7 +282,7 @@ namespace Microsoft.ML.OnnxRuntime
IntPtr typeInfo = IntPtr.Zero;
try
{
NativeApiStatus.VerifySuccess(NativeMethods.OrtInferenceSessionGetOutputTypeInfo(_nativeHandle, index, out typeInfo));
NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOutputTypeInfo(_nativeHandle, index, out typeInfo));
return GetMetadataFromTypeInfo(typeInfo);
}
finally

View file

@ -43,14 +43,14 @@ namespace Microsoft.ML.OnnxRuntime
#region InferenceSession API
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /* OrtStatus* */OrtCreateInferenceSession(
public static extern IntPtr /* OrtStatus* */OrtCreateSession(
IntPtr /* (OrtEnv*) */ environment,
[MarshalAs(UnmanagedType.LPWStr)]string modelPath, //the model path is consumed as a wchar* in the C-api
IntPtr /* (OrtSessionOptions*) */sessopnOptions,
out IntPtr /**/ session);
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(ONNStatus*)*/ OrtRunInference(
public static extern IntPtr /*(ONNStatus*)*/ OrtRun(
IntPtr /*(OrtSession*)*/ session,
IntPtr /*(OrtSessionRunOptions*)*/ runOptions, // can be null to use the default options
string[] inputNames,
@ -65,25 +65,25 @@ namespace Microsoft.ML.OnnxRuntime
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(OrtStatus*)*/ OrtInferenceSessionGetInputCount(
public static extern IntPtr /*(OrtStatus*)*/ OrtSessionGetInputCount(
IntPtr /*(OrtSession*)*/ session,
out ulong /* TODO: size_t */ count);
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(OrtStatus*)*/ OrtInferenceSessionGetOutputCount(
public static extern IntPtr /*(OrtStatus*)*/ OrtSessionGetOutputCount(
IntPtr /*(OrtSession*)*/ session,
out ulong /*TODO: size_t port*/ count);
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(OrtStatus*)*/OrtInferenceSessionGetInputName(
public static extern IntPtr /*(OrtStatus*)*/OrtSessionGetInputName(
IntPtr /*(OrtSession*)*/ session,
ulong index, //TODO: port size_t
IntPtr /*(OrtAllocator*)*/ allocator,
out IntPtr /*(char**)*/name);
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(OrtStatus*)*/OrtInferenceSessionGetOutputName(
public static extern IntPtr /*(OrtStatus*)*/OrtSessionGetOutputName(
IntPtr /*(OrtSession*)*/ session,
ulong index, //TODO: port size_t
IntPtr /*(OrtAllocator*)*/ allocator,
@ -91,14 +91,14 @@ namespace Microsoft.ML.OnnxRuntime
// release the typeinfo using OrtReleaseObject
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(OrtStatus*)*/OrtInferenceSessionGetInputTypeInfo(
public static extern IntPtr /*(OrtStatus*)*/OrtSessionGetInputTypeInfo(
IntPtr /*(const OrtSession*)*/ session,
ulong index, //TODO: port for size_t
out IntPtr /*(struct OrtTypeInfo**)*/ typeInfo);
// release the typeinfo using OrtReleaseObject
[DllImport(nativeLib, CharSet = charSet)]
public static extern IntPtr /*(OrtStatus*)*/OrtInferenceSessionGetOutputTypeInfo(
public static extern IntPtr /*(OrtStatus*)*/OrtSessionGetOutputTypeInfo(
IntPtr /*(const OrtSession*)*/ session,
ulong index, //TODO: port for size_t
out IntPtr /* (struct OrtTypeInfo**)*/ typeInfo);

View file

@ -6,7 +6,6 @@
// =====================================================================================================
#pragma once
#include <stdbool.h>
#include <stdlib.h>
#include <stdint.h>
#include <string.h>
@ -63,23 +62,23 @@ extern "C" {
// Copied from TensorProto::DataType
// Currently, Ort doesn't support complex64, complex128, bfloat16 types
typedef enum ONNXTensorElementDataType {
ONNX_TENSOR_ELEMENT_DATA_TYPE_UNDEFINED = 0,
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT = 1, // maps to c type float
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8 = 2, // maps to c type uint8_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8 = 3, // maps to c type int8_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16 = 4, // maps to c type uint16_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16 = 5, // maps to c type int16_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32 = 6, // maps to c type int32_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64 = 7, // maps to c type int64_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING = 8, // maps to c++ type std::string
ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL = 9, //
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16 = 10,
ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE = 11, // maps to c type double
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32 = 12, // maps to c type uint32_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64 = 13, // maps to c type uint64_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_COMPLEX64 = 14, // complex with float32 real and imaginary components
ONNX_TENSOR_ELEMENT_DATA_TYPE_COMPLEX128 = 15, // complex with float64 real and imaginary components
ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16 = 16, // Non-IEEE floating-point format based on IEEE754 single-precision
ONNX_TENSOR_ELEMENT_DATA_TYPE_UNDEFINED,
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, // maps to c type float
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8, // maps to c type uint8_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8, // maps to c type int8_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16, // maps to c type uint16_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16, // maps to c type int16_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, // maps to c type int32_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64, // maps to c type int64_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING, // maps to c++ type std::string
ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL, //
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16,
ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE, // maps to c type double
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32, // maps to c type uint32_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64, // maps to c type uint64_t
ONNX_TENSOR_ELEMENT_DATA_TYPE_COMPLEX64, // complex with float32 real and imaginary components
ONNX_TENSOR_ELEMENT_DATA_TYPE_COMPLEX128, // complex with float64 real and imaginary components
ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16, // Non-IEEE floating-point format based on IEEE754 single-precision
} ONNXTensorElementDataType;
// Synced with onnx TypeProto oneof
@ -93,32 +92,29 @@ typedef enum ONNXType {
} ONNXType;
typedef enum OrtLoggingLevel {
ORT_LOGGING_LEVEL_kVERBOSE = 0,
ORT_LOGGING_LEVEL_kINFO = 1,
ORT_LOGGING_LEVEL_kWARNING = 2,
ORT_LOGGING_LEVEL_kERROR = 3,
ORT_LOGGING_LEVEL_kFATAL = 4
ORT_LOGGING_LEVEL_kVERBOSE,
ORT_LOGGING_LEVEL_kINFO,
ORT_LOGGING_LEVEL_kWARNING,
ORT_LOGGING_LEVEL_kERROR,
ORT_LOGGING_LEVEL_kFATAL,
} OrtLoggingLevel;
typedef enum OrtErrorCode {
ORT_OK = 0,
ORT_FAIL = 1,
ORT_INVALID_ARGUMENT = 2,
ORT_NO_SUCHFILE = 3,
ORT_NO_MODEL = 4,
ORT_ENGINE_ERROR = 5,
ORT_RUNTIME_EXCEPTION = 6,
ORT_INVALID_PROTOBUF = 7,
ORT_MODEL_LOADED = 8,
ORT_NOT_IMPLEMENTED = 9,
ORT_INVALID_GRAPH = 10,
ORT_SHAPE_INFERENCE_NOT_REGISTERED = 11,
ORT_REQUIREMENT_NOT_REGISTERED = 12
ORT_OK,
ORT_FAIL,
ORT_INVALID_ARGUMENT,
ORT_NO_SUCHFILE,
ORT_NO_MODEL,
ORT_ENGINE_ERROR,
ORT_RUNTIME_EXCEPTION,
ORT_INVALID_PROTOBUF,
ORT_MODEL_LOADED,
ORT_NOT_IMPLEMENTED,
ORT_INVALID_GRAPH,
ORT_SHAPE_INFERENCE_NOT_REGISTERED,
ORT_REQUIREMENT_NOT_REGISTERED,
} OrtErrorCode;
// OrtStatus is always returned as a pointer. nullptr indicates success
typedef void OrtStatus;
// __VA_ARGS__ on Windows and Linux are different
#define ORT_API(RETURN_TYPE, NAME, ...) \
ORT_EXPORT RETURN_TYPE ORT_API_CALL NAME(__VA_ARGS__) NO_EXCEPTION
@ -130,22 +126,13 @@ typedef void OrtStatus;
#define ORT_API_STATUS_IMPL(NAME, ...) \
ORT_EXPORT OrtStatus* ORT_API_CALL NAME(__VA_ARGS__) NO_EXCEPTION
#define ORT_RUNTIME_CLASS2(NAME, TYPE) \
ORT_API(void, OrtRelease##NAME, _Frees_ptr_opt_ TYPE* input);
#define ORT_RUNTIME_CLASS(X) \
struct Ort##X; \
typedef struct Ort##X Ort##X; \
ORT_RUNTIME_CLASS2(X, Ort##X)
// OrtStatus* is pointer to something like this:
// struct OrtStatus {
// OrtErrorCode code;
// char msg[]; // a null-terminated string, var length
// }
ORT_RUNTIME_CLASS2(Status, void);
ORT_API(void, OrtRelease##X, _Frees_ptr_opt_ Ort##X* input);
// The actual types defined have an Ort prefix
ORT_RUNTIME_CLASS(Status); // nullptr for Status* indicates success
ORT_RUNTIME_CLASS(Provider);
ORT_RUNTIME_CLASS(AllocatorInfo);
ORT_RUNTIME_CLASS(Session);
@ -212,14 +199,14 @@ ORT_API_STATUS(OrtInitializeWithCustomLogger, OrtLoggingFunction logging_functio
// TODO: document the path separator convention? '/' vs '\'
// TODO: should specify the access characteristics of model_path. Is this read only during the
// execution of OrtCreateInferenceSession, or does the OrtSession retain a handle to the file/directory
// execution of OrtCreateSession, or does the OrtSession retain a handle to the file/directory
// and continue to access throughout the OrtSession lifetime?
// What sort of access is needed to model_path : read or read/write?
// TODO: allow loading from an in-memory byte-array
ORT_API_STATUS(OrtCreateInferenceSession, _In_ OrtEnv* env, _In_ const ORTCHAR_T* model_path,
ORT_API_STATUS(OrtCreateSession, _In_ OrtEnv* env, _In_ const ORTCHAR_T* model_path,
_In_ const OrtSessionOptions* options, _Out_ OrtSession** out);
ORT_API_STATUS(OrtRunInference, _Inout_ OrtSession* sess,
ORT_API_STATUS(OrtRun, _Inout_ OrtSession* sess,
_In_ OrtRunOptions* run_options,
_In_ const char* const* input_names, _In_ const OrtValue* const* input, size_t input_len,
_In_ const char* const* output_names, size_t output_names_len, _Out_ OrtValue** output);
@ -227,7 +214,7 @@ ORT_API_STATUS(OrtRunInference, _Inout_ OrtSession* sess,
/**
* \return A pointer of the newly created object. The pointer should be freed by OrtReleaseObject after use
*/
ORT_API(OrtSessionOptions*, OrtCreateSessionOptions, void);
ORT_API(OrtSessionOptions*, OrtCreateSessionOptions);
/// create a copy of an existing OrtSessionOptions
ORT_API(OrtSessionOptions*, OrtCloneSessionOptions, OrtSessionOptions*);
@ -245,7 +232,7 @@ ORT_API(void, OrtDisableProfiling, _In_ OrtSessionOptions* options);
ORT_API(void, OrtEnableMemPattern, _In_ OrtSessionOptions* options);
ORT_API(void, OrtDisableMemPattern, _In_ OrtSessionOptions* options);
// enable the memory arena on CPU
// Enable the memory arena on CPU
// Arena may pre-allocate memory for future usage.
// set this option to false if you don't want it.
ORT_API(void, OrtEnableCpuMemArena, _In_ OrtSessionOptions* options);
@ -267,24 +254,24 @@ ORT_API(int, OrtSetSessionThreadPoolSize, _In_ OrtSessionOptions* options, int s
*/
ORT_API(void, OrtSessionOptionsAppendExecutionProvider, _In_ OrtSessionOptions* options, _In_ OrtProviderFactoryInterface** f);
ORT_API(void, OrtAddCustomOp, _In_ OrtSessionOptions* options, const char* custom_op_path);
ORT_API(void, OrtAppendCustomOpLibPath, _In_ OrtSessionOptions* options, const char* lib_path);
ORT_API_STATUS(OrtInferenceSessionGetInputCount, _In_ const OrtSession* sess, _Out_ size_t* out);
ORT_API_STATUS(OrtInferenceSessionGetOutputCount, _In_ const OrtSession* sess, _Out_ size_t* out);
ORT_API_STATUS(OrtSessionGetInputCount, _In_ const OrtSession* sess, _Out_ size_t* out);
ORT_API_STATUS(OrtSessionGetOutputCount, _In_ const OrtSession* sess, _Out_ size_t* out);
/**
* \param out should be freed by OrtReleaseObject after use
*/
ORT_API_STATUS(OrtInferenceSessionGetInputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ OrtTypeInfo** out);
ORT_API_STATUS(OrtSessionGetInputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ OrtTypeInfo** out);
/**
* \param out should be freed by OrtReleaseObject after use
*/
ORT_API_STATUS(OrtInferenceSessionGetOutputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ OrtTypeInfo** out);
ORT_API_STATUS(OrtSessionGetOutputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ OrtTypeInfo** out);
ORT_API_STATUS(OrtInferenceSessionGetInputName, _In_ const OrtSession* sess, size_t index,
ORT_API_STATUS(OrtSessionGetInputName, _In_ const OrtSession* sess, size_t index,
_Inout_ OrtAllocator* allocator, _Out_ char** value);
ORT_API_STATUS(OrtInferenceSessionGetOutputName, _In_ const OrtSession* sess, size_t index,
ORT_API_STATUS(OrtSessionGetOutputName, _In_ const OrtSession* sess, size_t index,
_Inout_ OrtAllocator* allocator, _Out_ char** value);
/**
@ -298,9 +285,9 @@ ORT_API_STATUS(OrtRunOptionsSetRunTag, _In_ OrtRunOptions*, _In_ const char* run
ORT_API(unsigned int, OrtRunOptionsGetRunLogVerbosityLevel, _In_ OrtRunOptions*);
ORT_API(const char*, OrtRunOptionsGetRunTag, _In_ OrtRunOptions*);
// set a flag so that any running OrtRunInference* calls that are using this instance of ORtRunOptions
// Set a flag so that any running OrtRun* calls that are using this instance of OrtRunOptions
// will exit as soon as possible if the flag is true.
ORT_API(void, OrtRunOptionsSetTerminate, _In_ OrtRunOptions*, _In_ bool value);
ORT_API(void, OrtRunOptionsSetTerminate, _In_ OrtRunOptions*, _In_ int flag);
/**
* Create a tensor from an allocator. OrtReleaseValue will also release the buffer inside the output value
@ -326,8 +313,7 @@ ORT_API_STATUS(OrtCreateTensorWithDataAsOrtValue, _In_ const OrtAllocatorInfo* i
ORT_API_STATUS(OrtGetTensorMutableData, _Inout_ OrtValue* value, _Out_ void** out);
/**
* Test if an OrtValue is a tensor
* \return zero if false. non-zero if true
* \Return 1 iff an OrtValue is a tensor, 0 otherwise
*/
ORT_API(int, OrtIsTensor, _In_ const OrtValue* value);
@ -378,15 +364,13 @@ ORT_API(size_t, OrtGetNumOfDimensions, _In_ const OrtTensorTypeAndShapeInfo* inf
ORT_API(void, OrtGetDimensions, _In_ const OrtTensorTypeAndShapeInfo* info, _Out_ int64_t* dim_values, size_t dim_values_length);
/**
* How many elements does this tensor have.
* May return a negative value
* Return the number of elements specified by the tensor shape.
* Return a negative value if unknown (i.e., any dimension is negative.)
* e.g.
* [] -> 1
* [1,3,4] -> 12
* [2,0,4] -> 0
* [-1,3,4] -> -1
* return a negative value if unknown. (That this shape contains a symbolic variable which
* represents an unknown dimension.)
*/
ORT_API(int64_t, OrtGetTensorShapeElementCount, _In_ const OrtTensorTypeAndShapeInfo* info);

View file

@ -11,10 +11,10 @@
//TODO: encode error code in the message?
#define ORT_THROW_ON_ERROR(expr) \
do { \
OrtStatus* onnx_status = (expr); \
OrtStatus* onnx_status = (expr); \
if (onnx_status != nullptr) { \
std::string ort_error_message = OrtGetErrorMessage(onnx_status); \
OrtReleaseStatus(onnx_status); \
OrtReleaseStatus(onnx_status); \
throw std::runtime_error(ort_error_message); \
} \
} while (0);
@ -24,14 +24,14 @@
return Ort##NAME(value.get()); \
}
#define DECLARE_DEFAULT_DELETER_FOR_ONNX_OBJECT(TYPE_NAME) \
namespace std { \
template <> \
struct default_delete<Ort##TYPE_NAME> { \
void operator()(Ort##TYPE_NAME* ptr) { \
#define DECLARE_DEFAULT_DELETER_FOR_ONNX_OBJECT(TYPE_NAME) \
namespace std { \
template <> \
struct default_delete<Ort##TYPE_NAME> { \
void operator()(Ort##TYPE_NAME* ptr) { \
(*reinterpret_cast<OrtObject**>(ptr))->Release(ptr); \
} \
}; \
} \
}; \
}
DECLARE_DEFAULT_DELETER_FOR_ONNX_OBJECT(Env);
@ -89,20 +89,20 @@ class SessionOptionsWrapper {
return SessionOptionsWrapper(env_, p);
}
#ifdef _WIN32
OrtSession* OrtCreateInferenceSession(_In_ const wchar_t* model_path) {
OrtSession* OrtCreateSession(_In_ const wchar_t* model_path) {
OrtSession* ret;
ORT_THROW_ON_ERROR(::OrtCreateInferenceSession(env_, model_path, value.get(), &ret));
ORT_THROW_ON_ERROR(::OrtCreateSession(env_, model_path, value.get(), &ret));
return ret;
}
#else
OrtSession* OrtCreateInferenceSession(_In_ const char* model_path) {
OrtSession* OrtCreateSession(_In_ const char* model_path) {
OrtSession* ret;
ORT_THROW_ON_ERROR(::OrtCreateInferenceSession(env_, model_path, value.get(), &ret));
ORT_THROW_ON_ERROR(::OrtCreateSession(env_, model_path, value.get(), &ret));
return ret;
}
#endif
void AddCustomOp(_In_ const char* custom_op_path) {
OrtAddCustomOp(value.get(), custom_op_path);
void AppendCustomOpLibPath(_In_ const char* lib_path) {
OrtAppendCustomOpLibPath(value.get(), lib_path);
}
};
inline OrtValue* OrtCreateTensorAsOrtValue(_Inout_ OrtAllocator* env, const std::vector<size_t>& shape, ONNXTensorElementDataType type) {

View file

@ -7,39 +7,37 @@
#include <cassert>
using onnxruntime::common::Status;
struct OrtStatus {
OrtErrorCode code;
char msg[1]; // a null-terminated string
};
ORT_API(OrtStatus*, OrtCreateStatus, OrtErrorCode code, const char* msg) {
assert(!(code == 0 && msg != nullptr));
size_t clen = strlen(msg);
size_t len = clen + 1 + sizeof(int);
char* p = new char[len];
char* ret = p;
*reinterpret_cast<int*>(p) = static_cast<int>(code);
p += sizeof(int);
memcpy(p, msg, clen);
p += clen;
*p = '\0';
return ret;
OrtStatus* p = reinterpret_cast<OrtStatus*>(new char[sizeof(OrtStatus) + clen]);
p->code = code;
memcpy(p->msg, msg, clen);
p->msg[clen] = '\0';
return p;
}
namespace onnxruntime {
OrtStatus* ToOrtStatus(const Status& st) {
if (st.IsOK())
return nullptr;
size_t clen = st.ErrorMessage().length();
size_t len = clen + 1 + sizeof(int);
char* p = new char[len];
char* ret = p;
*reinterpret_cast<int*>(p) = static_cast<int>(st.Code());
p += sizeof(int);
memcpy(p, st.ErrorMessage().c_str(), clen);
p += clen;
*p = '\0';
return ret;
OrtStatus* p = reinterpret_cast<OrtStatus*>(new char[sizeof(OrtStatus) + clen]);
p->code = static_cast<OrtErrorCode>(st.Code());
memcpy(p->msg, st.ErrorMessage().c_str(), clen);
p->msg[clen] = '\0';
return p;
}
} // namespace onnxruntime
ORT_API(OrtErrorCode, OrtGetErrorCode, _In_ const OrtStatus* status) {
return *reinterpret_cast<OrtErrorCode*>(const_cast<OrtStatus*>(status));
return status->code;
}
ORT_API(const char*, OrtGetErrorMessage, _In_ const OrtStatus* status) {
return reinterpret_cast<const char*>(status) + sizeof(int);
return status->msg;
}

View file

@ -1,4 +1,3 @@
OrtAddCustomOp
OrtAddRefToObject
OrtAllocatorAlloc
OrtAllocatorFree
@ -7,6 +6,7 @@ OrtAllocatorInfoGetId
OrtAllocatorInfoGetMemType
OrtAllocatorInfoGetName
OrtAllocatorInfoGetType
OrtAppendCustomOpLibPath
OrtCastTypeInfoToTensorInfo
OrtCloneSessionOptions
OrtCompareAllocatorInfo
@ -14,8 +14,8 @@ OrtCreateAllocatorInfo
OrtCreateCpuAllocatorInfo
OrtCreateCpuExecutionProviderFactory
OrtCreateDefaultAllocator
OrtCreateInferenceSession
OrtCreateRunOptions
OrtCreateSession
OrtCreateSessionOptions
OrtCreateTensorAsOrtValue
OrtCreateTensorTypeAndShapeInfo
@ -41,12 +41,6 @@ OrtGetTensorShapeAndType
OrtGetTensorShapeElementCount
OrtGetTypeInfo
OrtGetValueType
OrtInferenceSessionGetInputCount
OrtInferenceSessionGetInputName
OrtInferenceSessionGetInputTypeInfo
OrtInferenceSessionGetOutputCount
OrtInferenceSessionGetOutputName
OrtInferenceSessionGetOutputTypeInfo
OrtInitialize
OrtInitializeWithCustomLogger
OrtIsTensor
@ -56,12 +50,18 @@ OrtReleaseObject
OrtReleaseSession
OrtReleaseStatus
OrtReleaseValue
OrtRunInference
OrtRun
OrtRunOptionsGetRunLogVerbosityLevel
OrtRunOptionsGetRunTag
OrtRunOptionsSetRunLogVerbosityLevel
OrtRunOptionsSetRunTag
OrtRunOptionsSetTerminate
OrtSessionGetInputCount
OrtSessionGetInputName
OrtSessionGetInputTypeInfo
OrtSessionGetOutputCount
OrtSessionGetOutputName
OrtSessionGetOutputTypeInfo
OrtSessionOptionsAppendExecutionProvider
OrtSetDims
OrtSetSessionLogId

View file

@ -97,6 +97,6 @@ ORT_API(int, OrtSetSessionThreadPoolSize, _In_ OrtSessionOptions* options, int s
return 0;
}
ORT_API(void, OrtAddCustomOp, _In_ OrtSessionOptions* options, const char* custom_op_path) {
options->custom_op_paths.emplace_back(custom_op_path);
ORT_API(void, OrtAppendCustomOpLibPath, _In_ OrtSessionOptions* options, const char* lib_path) {
options->custom_op_paths.emplace_back(lib_path);
}

View file

@ -357,9 +357,9 @@ ORT_API_STATUS_IMPL(OrtCreateTensorAsOrtValue, _Inout_ OrtAllocator* allocator,
}
template <typename T>
static OrtStatus* CreateInferenceSessionImpl(_In_ OrtEnv* env, _In_ T model_path,
_In_ const OrtSessionOptions* options,
_Out_ OrtSession** out) {
static OrtStatus* CreateSessionImpl(_In_ OrtEnv* env, _In_ T model_path,
_In_ const OrtSessionOptions* options,
_Out_ OrtSession** out) {
API_IMPL_BEGIN
auto sess = std::make_unique<::onnxruntime::InferenceSession>(options == nullptr ? onnxruntime::SessionOptions() : options->value, env->loggingManager);
Status status;
@ -389,22 +389,22 @@ static OrtStatus* CreateInferenceSessionImpl(_In_ OrtEnv* env, _In_ T model_path
}
#ifdef _WIN32
ORT_API_STATUS_IMPL(OrtCreateInferenceSession, _In_ OrtEnv* env, _In_ const wchar_t* model_path,
ORT_API_STATUS_IMPL(OrtCreateSession, _In_ OrtEnv* env, _In_ const wchar_t* model_path,
_In_ const OrtSessionOptions* options, _Out_ OrtSession** out) {
API_IMPL_BEGIN
return CreateInferenceSessionImpl(env, model_path, options, out);
return CreateSessionImpl(env, model_path, options, out);
API_IMPL_END
}
#else
ORT_API_STATUS_IMPL(OrtCreateInferenceSession, _In_ OrtEnv* env, _In_ const char* model_path,
ORT_API_STATUS_IMPL(OrtCreateSession, _In_ OrtEnv* env, _In_ const char* model_path,
_In_ const OrtSessionOptions* options, _Out_ OrtSession** out) {
API_IMPL_BEGIN
return CreateInferenceSessionImpl(env, model_path, options, out);
return CreateSessionImpl(env, model_path, options, out);
API_IMPL_END
}
#endif
ORT_API_STATUS_IMPL(OrtRunInference, _In_ OrtSession* sess,
ORT_API_STATUS_IMPL(OrtRun, _In_ OrtSession* sess,
_In_ OrtRunOptions* run_options,
_In_ const char* const* input_names, _In_ const OrtValue* const* input, size_t input_len,
_In_ const char* const* output_names1, size_t output_names_len, _Out_ OrtValue** output) {
@ -526,7 +526,7 @@ ORT_API_STATUS_IMPL(OrtTensorProtoToOrtValue, _Inout_ OrtAllocator* allocator,
delete[] reinterpret_cast<REAL_TYPE*>(value); \
}
ORT_API_STATUS_IMPL(OrtInferenceSessionGetInputCount, _In_ const OrtSession* sess, _Out_ size_t* out) {
ORT_API_STATUS_IMPL(OrtSessionGetInputCount, _In_ const OrtSession* sess, _Out_ size_t* out) {
API_IMPL_BEGIN
auto session = reinterpret_cast<const ::onnxruntime::InferenceSession*>(sess);
std::pair<Status, const InputDefList*> p = session->GetModelInputs();
@ -537,7 +537,7 @@ ORT_API_STATUS_IMPL(OrtInferenceSessionGetInputCount, _In_ const OrtSession* ses
API_IMPL_END
}
ORT_API_STATUS_IMPL(OrtInferenceSessionGetOutputCount, _In_ const OrtSession* sess, _Out_ size_t* out) {
ORT_API_STATUS_IMPL(OrtSessionGetOutputCount, _In_ const OrtSession* sess, _Out_ size_t* out) {
API_IMPL_BEGIN
auto session = reinterpret_cast<const ::onnxruntime::InferenceSession*>(sess);
std::pair<Status, const InputDefList*> p = session->GetModelOutputs();
@ -548,7 +548,7 @@ ORT_API_STATUS_IMPL(OrtInferenceSessionGetOutputCount, _In_ const OrtSession* se
API_IMPL_END
}
ORT_API_STATUS_IMPL(OrtInferenceSessionGetInputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ struct OrtTypeInfo** out) {
ORT_API_STATUS_IMPL(OrtSessionGetInputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ struct OrtTypeInfo** out) {
API_IMPL_BEGIN
auto session = reinterpret_cast<const ::onnxruntime::InferenceSession*>(sess);
std::pair<Status, const InputDefList*> p = session->GetModelInputs();
@ -560,7 +560,7 @@ ORT_API_STATUS_IMPL(OrtInferenceSessionGetInputTypeInfo, _In_ const OrtSession*
return OrtTypeInfo::FromDataTypeImpl(type_proto, out);
API_IMPL_END
}
ORT_API_STATUS_IMPL(OrtInferenceSessionGetOutputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ struct OrtTypeInfo** out) {
ORT_API_STATUS_IMPL(OrtSessionGetOutputTypeInfo, _In_ const OrtSession* sess, size_t index, _Out_ struct OrtTypeInfo** out) {
API_IMPL_BEGIN
auto session = reinterpret_cast<const ::onnxruntime::InferenceSession*>(sess);
std::pair<Status, const InputDefList*> p = session->GetModelOutputs();
@ -624,14 +624,14 @@ ORT_API(const struct OrtAllocatorInfo*, OrtAllocatorGetInfo, _In_ const OrtAlloc
}
}
ORT_API_STATUS_IMPL(OrtInferenceSessionGetInputName, _In_ const OrtSession* sess, size_t index,
ORT_API_STATUS_IMPL(OrtSessionGetInputName, _In_ const OrtSession* sess, size_t index,
_Inout_ OrtAllocator* allocator, _Out_ char** output) {
API_IMPL_BEGIN
return GetInputOutputNameImpl(sess, index, allocator, true, output);
API_IMPL_END
}
ORT_API_STATUS_IMPL(OrtInferenceSessionGetOutputName, _In_ const OrtSession* sess, size_t index,
ORT_API_STATUS_IMPL(OrtSessionGetOutputName, _In_ const OrtSession* sess, size_t index,
_Inout_ OrtAllocator* allocator, _Out_ char** output) {
API_IMPL_BEGIN
return GetInputOutputNameImpl(sess, index, allocator, false, output);

View file

@ -531,18 +531,18 @@ Status OnnxTestCase::ConvertTestData(OrtSession* session, const std::vector<onnx
if (!has_valid_names) {
size_t count;
if (is_input) {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetInputCount(session, &count));
ORT_THROW_ON_ERROR(OrtSessionGetInputCount(session, &count));
} else {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetOutputCount(session, &count));
ORT_THROW_ON_ERROR(OrtSessionGetOutputCount(session, &count));
}
if (count != test_data_pbs.size())
ORT_THROW("data count mismatch");
for (size_t i = 0; i != count; ++i) {
char* temp_name;
if (is_input) {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetInputName(session, i, allocator, &temp_name));
ORT_THROW_ON_ERROR(OrtSessionGetInputName(session, i, allocator, &temp_name));
} else {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetOutputName(session, i, allocator, &temp_name));
ORT_THROW_ON_ERROR(OrtSessionGetOutputName(session, i, allocator, &temp_name));
}
var_names[i] = temp_name;
(*allocator)->Free(allocator, temp_name);

View file

@ -322,11 +322,11 @@ EXECUTE_RESULT DataRunner::RunTaskImpl(size_t task_id) {
// Create output feed
size_t output_count;
ORT_THROW_ON_ERROR(OrtInferenceSessionGetOutputCount(session, &output_count));
ORT_THROW_ON_ERROR(OrtSessionGetOutputCount(session, &output_count));
std::vector<std::string> output_names(output_count);
for (size_t i = 0; i != output_count; ++i) {
char* output_name = nullptr;
ORT_THROW_ON_ERROR(OrtInferenceSessionGetOutputName(session, i, default_allocator, &output_name));
ORT_THROW_ON_ERROR(OrtSessionGetOutputName(session, i, default_allocator, &output_name));
assert(output_name != nullptr);
output_names[i] = output_name;
(*default_allocator)->Free(default_allocator, output_name);
@ -348,7 +348,7 @@ EXECUTE_RESULT DataRunner::RunTaskImpl(size_t task_id) {
for (size_t i = 0; i != output_count; ++i) {
output_names_raw_ptr[i] = output_names[i].c_str();
}
auto onnx_status = OrtRunInference(session, nullptr, input_names.data(), input_values.data(), input_index, output_names_raw_ptr.data(), output_count, output_values.data());
auto onnx_status = OrtRun(session, nullptr, input_names.data(), input_values.data(), input_index, output_names_raw_ptr.data(), output_count, output_values.data());
if (onnx_status != nullptr) {
std::string onnx_runtime_error_message = OrtGetErrorMessage(onnx_status);
OrtReleaseStatus(onnx_status);
@ -493,7 +493,7 @@ void RunSingleTestCase(ITestCase* info, const onnxruntime::SessionOptionsWrapper
auto sf2 = sf.clone();
sf2.SetSessionLogId(info->GetTestCaseName().c_str());
std::unique_ptr<OrtSession, decltype(&OrtReleaseSession)> session_object(
sf2.OrtCreateInferenceSession(info->GetModelUrl()), OrtReleaseSession);
sf2.OrtCreateSession(info->GetModelUrl()), OrtReleaseSession);
LOGF_DEFAULT(INFO, "testing %s\n", info->GetTestCaseName().c_str());
//temp hack. Because we have no resource control. We may not have enough memory to run this test in parallel
if (info->GetTestCaseName() == "coreml_FNS-Candy_ImageNet")

View file

@ -8,11 +8,11 @@
#define ORT_ABORT_ON_ERROR(expr) \
do { \
OrtStatus* onnx_status = (expr); \
OrtStatus* onnx_status = (expr); \
if (onnx_status != NULL) { \
const char* msg = OrtGetErrorMessage(onnx_status); \
fprintf(stderr, "%s\n", msg); \
OrtReleaseStatus(onnx_status); \
OrtReleaseStatus(onnx_status); \
abort(); \
} \
} while (0);
@ -154,14 +154,14 @@ int run_inference(OrtSession* session, const char* input_file, const char* outpu
OrtValue* input_tensor = NULL;
ORT_ABORT_ON_ERROR(OrtCreateTensorWithDataAsOrtValue(allocator_info, model_input, model_input_len, input_shape, input_shape_len, ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, &input_tensor));
assert(input_tensor != NULL);
assert(OrtIsTensor(input_tensor) != 0);
assert(OrtIsTensor(input_tensor));
OrtReleaseAllocatorInfo(allocator_info);
const char* input_names[] = {"inputImage"};
const char* output_names[] = {"outputImage"};
OrtValue* output_tensor = NULL;
ORT_ABORT_ON_ERROR(OrtRunInference(session, NULL, input_names, (const OrtValue* const*)&input_tensor, 1, output_names, 1, &output_tensor));
ORT_ABORT_ON_ERROR(OrtRun(session, NULL, input_names, (const OrtValue* const*)&input_tensor, 1, output_names, 1, &output_tensor));
assert(output_tensor != NULL);
assert(OrtIsTensor(output_tensor) != 0);
assert(OrtIsTensor(output_tensor));
int ret = 0;
if (write_tensor_to_png_file(output_tensor, output_file) != 0) {
ret = -1;
@ -174,9 +174,9 @@ int run_inference(OrtSession* session, const char* input_file, const char* outpu
void verify_input_output_count(OrtSession* session) {
size_t count;
ORT_ABORT_ON_ERROR(OrtInferenceSessionGetInputCount(session, &count));
ORT_ABORT_ON_ERROR(OrtSessionGetInputCount(session, &count));
assert(count == 1);
ORT_ABORT_ON_ERROR(OrtInferenceSessionGetOutputCount(session, &count));
ORT_ABORT_ON_ERROR(OrtSessionGetOutputCount(session, &count));
assert(count == 1);
}
@ -204,7 +204,7 @@ int main(int argc, char* argv[]) {
enable_cuda(session_option);
#endif
OrtSession* session;
ORT_ABORT_ON_ERROR(OrtCreateInferenceSession(env, model_path, session_option, &session));
ORT_ABORT_ON_ERROR(OrtCreateSession(env, model_path, session_option, &session));
verify_input_output_count(session);
int ret = run_inference(session, input_file, output_file);
OrtReleaseObject(session_option);

View file

@ -28,7 +28,7 @@ void RunSession(OrtAllocator* env, OrtSession* session_object,
std::vector<const char*> input_names{"X"};
OrtValue* output_tensor = nullptr;
const char* output_names[] = {"Y"};
ORT_THROW_ON_ERROR(OrtRunInference(session_object, NULL, input_names.data(), inputs.data(), inputs.size(), output_names, 1, &output_tensor));
ORT_THROW_ON_ERROR(OrtRun(session_object, NULL, input_names.data(), inputs.data(), inputs.size(), output_names, 1, &output_tensor));
ASSERT_NE(output_tensor, nullptr);
std::unique_ptr<OrtTensorTypeAndShapeInfo> shape_info;
{
@ -96,9 +96,9 @@ void TestInference(OrtEnv* env, T model_uri,
std::cout << "Running simple inference with default provider" << std::endl;
}
if (custom_op) {
sf.AddCustomOp("libonnxruntime_custom_op_shared_lib_test.so");
sf.AppendCustomOpLibPath("libonnxruntime_custom_op_shared_lib_test.so");
}
std::unique_ptr<OrtSession, decltype(&OrtReleaseSession)> inference_session(sf.OrtCreateInferenceSession(model_uri), OrtReleaseSession);
std::unique_ptr<OrtSession, decltype(&OrtReleaseSession)> inference_session(sf.OrtCreateSession(model_uri), OrtReleaseSession);
std::unique_ptr<OrtAllocator> default_allocator(MockedOrtAllocator::Create());
// Now run
RunSession(default_allocator.get(), inference_session.get(), dims_x, values_x, expected_dims_y, expected_values_y);
@ -148,7 +148,7 @@ TEST_F(CApiTest, DISABLED_custom_op) {
TEST_F(CApiTest, create_session_without_session_option) {
constexpr PATH_TYPE model_uri = TSTR("../models/opset8/test_squeezenet/model.onnx");
OrtSession* ret;
ORT_THROW_ON_ERROR(::OrtCreateInferenceSession(env, model_uri, nullptr, &ret));
ORT_THROW_ON_ERROR(::OrtCreateSession(env, model_uri, nullptr, &ret));
ASSERT_NE(nullptr, ret);
OrtReleaseSession(ret);
}

View file

@ -9,18 +9,18 @@ using namespace onnxruntime;
static void TestModelInfo(const OrtSession* inference_session, bool is_input, const std::vector<int64_t>& dims) {
size_t input_count;
if (is_input) {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetInputCount(inference_session, &input_count));
ORT_THROW_ON_ERROR(OrtSessionGetInputCount(inference_session, &input_count));
} else {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetOutputCount(inference_session, &input_count));
ORT_THROW_ON_ERROR(OrtSessionGetOutputCount(inference_session, &input_count));
}
ASSERT_EQ(1, input_count);
std::unique_ptr<OrtTypeInfo> input_type_info;
{
OrtTypeInfo* t;
if (is_input) {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetInputTypeInfo(inference_session, 0, &t));
ORT_THROW_ON_ERROR(OrtSessionGetInputTypeInfo(inference_session, 0, &t));
} else {
ORT_THROW_ON_ERROR(OrtInferenceSessionGetOutputTypeInfo(inference_session, 0, &t));
ORT_THROW_ON_ERROR(OrtSessionGetOutputTypeInfo(inference_session, 0, &t));
}
input_type_info.reset(t);
}
@ -39,7 +39,7 @@ static void TestModelInfo(const OrtSession* inference_session, bool is_input, co
TEST_F(CApiTest, input_output_type_info) {
SessionOptionsWrapper sf(env);
constexpr PATH_TYPE model_uri = TSTR("../models/opset8/test_squeezenet/model.onnx");
std::unique_ptr<OrtSession, decltype(&OrtReleaseSession)> inference_session(sf.OrtCreateInferenceSession(model_uri), OrtReleaseSession);
std::unique_ptr<OrtSession, decltype(&OrtReleaseSession)> inference_session(sf.OrtCreateSession(model_uri), OrtReleaseSession);
TestModelInfo(inference_session.get(), true, {1, 3, 224, 224});
TestModelInfo(inference_session.get(), false, {1, 1000, 1, 1});
}

View file

@ -358,7 +358,7 @@ std::pair<COMPARE_RESULT, std::string> CompareMLValue(const MLValue& o, const ML
std::pair<COMPARE_RESULT, std::string> VerifyValueInfo(const ONNX_NAMESPACE::ValueInfoProto& v, const OrtValue* o) {
if (!v.has_type()) return std::make_pair(COMPARE_RESULT::SUCCESS, "");
if (v.type().has_tensor_type()) {
if (OrtIsTensor(o) == 0) {
if (!OrtIsTensor(o)) {
return std::make_pair(COMPARE_RESULT::TYPE_MISMATCH, "");
}