using System; using System.Buffers; using System.Collections; using System.Collections.Generic; using System.Diagnostics; using System.Linq; using System.Reflection; using System.Runtime.CompilerServices; using System.Runtime.InteropServices; using System.Security; using System.Security.Permissions; using System.Text; using System.Threading.Tasks; using Microsoft.CodeAnalysis; using Microsoft.ML.OnnxRuntime.Tensors; [assembly: CompilationRelaxations(8)] [assembly: RuntimeCompatibility(WrapNonExceptionThrows = true)] [assembly: Debuggable(DebuggableAttribute.DebuggingModes.Default | DebuggableAttribute.DebuggingModes.DisableOptimizations | DebuggableAttribute.DebuggingModes.IgnoreSymbolStoreSequencePoints | DebuggableAttribute.DebuggingModes.EnableEditAndContinue)] [assembly: InternalsVisibleTo("Microsoft.ML.OnnxRuntime.Tests.Common, PublicKey=002400000480000094000000060200000024000052534131000400000100010059013e94e4bc70136ca4c35f33acd6b62974536b698f9c7a21cee18d805c7ad860ad9eebfdc47a96ba2f8d03f4cf1c36b9d30787e276c7b9833b5bf2a6eba7e919e6b90083078a352262aed1d842e5f70a3085cbcf4c56ae851b161137920961c23fcc246598d61d258ccc615c927b2441359eea666a99ce1c3c07dca18fb0e1")] [assembly: InternalsVisibleTo("Microsoft.ML.OnnxRuntime.Tests.Droid, PublicKey=002400000480000094000000060200000024000052534131000400000100010059013e94e4bc70136ca4c35f33acd6b62974536b698f9c7a21cee18d805c7ad860ad9eebfdc47a96ba2f8d03f4cf1c36b9d30787e276c7b9833b5bf2a6eba7e919e6b90083078a352262aed1d842e5f70a3085cbcf4c56ae851b161137920961c23fcc246598d61d258ccc615c927b2441359eea666a99ce1c3c07dca18fb0e1")] [assembly: InternalsVisibleTo("Microsoft.ML.OnnxRuntime.Tests.iOS, PublicKey=002400000480000094000000060200000024000052534131000400000100010059013e94e4bc70136ca4c35f33acd6b62974536b698f9c7a21cee18d805c7ad860ad9eebfdc47a96ba2f8d03f4cf1c36b9d30787e276c7b9833b5bf2a6eba7e919e6b90083078a352262aed1d842e5f70a3085cbcf4c56ae851b161137920961c23fcc246598d61d258ccc615c927b2441359eea666a99ce1c3c07dca18fb0e1")] [assembly: InternalsVisibleTo("Microsoft.ML.OnnxRuntime.Tests.NetCoreApp, PublicKey=002400000480000094000000060200000024000052534131000400000100010059013e94e4bc70136ca4c35f33acd6b62974536b698f9c7a21cee18d805c7ad860ad9eebfdc47a96ba2f8d03f4cf1c36b9d30787e276c7b9833b5bf2a6eba7e919e6b90083078a352262aed1d842e5f70a3085cbcf4c56ae851b161137920961c23fcc246598d61d258ccc615c927b2441359eea666a99ce1c3c07dca18fb0e1")] [assembly: SecurityPermission(SecurityAction.RequestMinimum, SkipVerification = true)] [assembly: AssemblyVersion("0.0.0.0")] [module: UnverifiableCode] [module: RefSafetyRules(11)] namespace Microsoft.CodeAnalysis { [CompilerGenerated] [Embedded] internal sealed class EmbeddedAttribute : Attribute { } } namespace System.Runtime.CompilerServices { [CompilerGenerated] [Embedded] internal sealed class IsReadOnlyAttribute : Attribute { } [CompilerGenerated] [Embedded] internal sealed class IsUnmanagedAttribute : Attribute { } [CompilerGenerated] [Embedded] internal sealed class IsByRefLikeAttribute : Attribute { } [CompilerGenerated] [Embedded] [AttributeUsage(AttributeTargets.Module, AllowMultiple = false, Inherited = false)] internal sealed class RefSafetyRulesAttribute : Attribute { public readonly int Version; public RefSafetyRulesAttribute(int P_0) { Version = P_0; } } } namespace Microsoft.ML.OnnxRuntime { public interface IDisposableReadOnlyCollection : IReadOnlyCollection, IEnumerable, IEnumerable, IReadOnlyList, IDisposable { } internal class DisposableList : List, IDisposableReadOnlyCollection, IReadOnlyCollection, IEnumerable, IEnumerable, IReadOnlyList, IDisposable where T : IDisposable { private bool _disposed; public DisposableList() { } public DisposableList(int count) : base(count) { } public DisposableList(IEnumerable collection) : base(collection) { } protected virtual void Dispose(bool disposing) { if (_disposed || !disposing) { return; } for (int num = base.Count - 1; num >= 0; num--) { T val = base[num]; if (val != null) { val.Dispose(); } } Clear(); _disposed = true; } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } } public class DisposableNamedOnnxValue : NamedOnnxValue, IDisposable { private IOrtValueOwner _ortValueHolder; private bool _disposed; public TensorElementType ElementType { get; } private DisposableNamedOnnxValue(string name, object value, TensorElementType elementType, IOrtValueOwner ortValueHolder) : base(name, value, OnnxValueType.ONNX_TYPE_TENSOR) { _ortValueHolder = ortValueHolder; ElementType = elementType; } private DisposableNamedOnnxValue(string name, object value, OnnxValueType onnxValueType, IOrtValueOwner ortValueHolder) : base(name, value, onnxValueType) { _ortValueHolder = ortValueHolder; ElementType = TensorElementType.DataTypeMax; } private DisposableNamedOnnxValue(string name, object value, MapHelper mapHelper, IOrtValueOwner ortValueHolder) : base(name, value, mapHelper) { _ortValueHolder = ortValueHolder; ElementType = TensorElementType.DataTypeMax; } internal override IntPtr InputToOrtValueHandle(NodeMetadata metadata, out IDisposable memoryHolder) { if (_ortValueHolder == null) { throw new InvalidOperationException("The instance of this class does not own an OrtValue"); } memoryHolder = null; return _ortValueHolder.Value.Handle; } internal override IntPtr OutputToOrtValueHandle(NodeMetadata metadata, out IDisposable memoryOwner) { return InputToOrtValueHandle(metadata, out memoryOwner); } internal static DisposableNamedOnnxValue CreateFromOrtValue(string name, ref OrtValue ortValue) { return CreateFromOrtValue(name, ref ortValue, OrtAllocator.DefaultInstance); } internal static DisposableNamedOnnxValue CreateFromOrtValue(string name, ref OrtValue ortValue, OrtAllocator allocator) { OnnxValueType onnxType = ortValue.OnnxType; return onnxType switch { OnnxValueType.ONNX_TYPE_TENSOR => FromNativeTensor(name, ref ortValue), OnnxValueType.ONNX_TYPE_SEQUENCE => FromNativeSequence(name, ref ortValue, allocator), OnnxValueType.ONNX_TYPE_MAP => FromNativeMap(name, ref ortValue, allocator), _ => throw new NotSupportedException($"OnnxValueType : {onnxType} is not supported"), }; } private static DisposableNamedOnnxValue FromNativeTensor(string name, ref OrtValue ortValue) { OrtTensorTypeAndShapeInfo tensorTypeAndShape = ortValue.GetTensorTypeAndShape(); switch (tensorTypeAndShape.ElementDataType) { case TensorElementType.Float: return FromNativeTensor(name, ref ortValue); case TensorElementType.Double: return FromNativeTensor(name, ref ortValue); case TensorElementType.Int16: return FromNativeTensor(name, ref ortValue); case TensorElementType.UInt16: return FromNativeTensor(name, ref ortValue); case TensorElementType.Int32: return FromNativeTensor(name, ref ortValue); case TensorElementType.UInt32: return FromNativeTensor(name, ref ortValue); case TensorElementType.Int64: return FromNativeTensor(name, ref ortValue); case TensorElementType.UInt64: return FromNativeTensor(name, ref ortValue); case TensorElementType.UInt8: return FromNativeTensor(name, ref ortValue); case TensorElementType.Int8: return FromNativeTensor(name, ref ortValue); case TensorElementType.String: { int[] shape = Array.ConvertAll(tensorTypeAndShape.Shape, Convert.ToInt32); return FromNativeStringTensor(name, shape, ref ortValue); } case TensorElementType.Bool: return FromNativeTensor(name, ref ortValue); case TensorElementType.Float16: return FromNativeTensor(name, ref ortValue); case TensorElementType.BFloat16: return FromNativeTensor(name, ref ortValue); default: throw new NotSupportedException($"Tensor of element type: {tensorTypeAndShape.ElementDataType} is not supported"); } } private static DisposableNamedOnnxValue FromNativeStringTensor(string name, int[] shape, ref OrtValue ortValue) { DenseTensor value = new DenseTensor(ortValue.GetStringTensorAsArray(), shape); DisposableNamedOnnxValue result = new DisposableNamedOnnxValue(name, value, TensorElementType.String, ortValue); ortValue = null; return result; } private static DisposableNamedOnnxValue FromNativeTensor(string name, ref OrtValue ortValue) { OrtValueTensor ortValueTensor = new OrtValueTensor(ref ortValue); try { DenseTensor value = new DenseTensor(ortValueTensor.Memory, ortValueTensor.Dimensions); return new DisposableNamedOnnxValue(name, value, ortValueTensor.ElementType, ortValueTensor); } catch (Exception) { ortValueTensor.Dispose(); throw; } } private static DisposableNamedOnnxValue FromNativeSequence(string name, ref OrtValue ortValueSequence, OrtAllocator allocator) { int valueCount = ortValueSequence.GetValueCount(); DisposableList disposableList = new DisposableList(valueCount); try { for (int i = 0; i < valueCount; i++) { OrtValue ortValue = ortValueSequence.GetValue(i, allocator); try { disposableList.Add(CreateFromOrtValue(string.Empty, ref ortValue, allocator)); } finally { ortValue?.Dispose(); } } NativeOrtValueCollectionOwner ortValueHolder = new NativeOrtValueCollectionOwner(ref ortValueSequence, disposableList); return new DisposableNamedOnnxValue(name, disposableList, OnnxValueType.ONNX_TYPE_SEQUENCE, ortValueHolder); } catch (Exception) { disposableList.Dispose(); throw; } } private static DisposableNamedOnnxValue FromNativeMap(string name, ref OrtValue ortValueMap, OrtAllocator allocator) { DisposableNamedOnnxValue result = null; Span disposables = new OrtValue[2]; DisposableArray disposableArray = new DisposableArray(disposables); try { disposables[0] = ortValueMap.GetValue(0, allocator); disposables[1] = ortValueMap.GetValue(1, allocator); OrtTensorTypeAndShapeInfo tensorTypeAndShape = disposables[0].GetTensorTypeAndShape(); OrtTensorTypeAndShapeInfo tensorTypeAndShape2 = disposables[1].GetTensorTypeAndShape(); int[] keysShape = Array.ConvertAll(tensorTypeAndShape.Shape, Convert.ToInt32); int[] valsShape = Array.ConvertAll(tensorTypeAndShape2.Shape, Convert.ToInt32); result = tensorTypeAndShape.ElementDataType switch { TensorElementType.Int64 => tensorTypeAndShape2.ElementDataType switch { TensorElementType.Float => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), TensorElementType.Double => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), TensorElementType.Int64 => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), TensorElementType.String => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), _ => throw new NotSupportedException($"Map value type: {tensorTypeAndShape2.ElementDataType} is not supported"), }, TensorElementType.String => tensorTypeAndShape2.ElementDataType switch { TensorElementType.Float => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), TensorElementType.Double => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), TensorElementType.Int64 => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), TensorElementType.String => FromNativeMapElements(name, ref ortValueMap, keysShape, ref disposables[0], valsShape, ref disposables[1]), _ => throw new NotSupportedException($"Map value type: {tensorTypeAndShape2.ElementDataType} is not supported"), }, _ => throw new NotSupportedException($"Map key type: {tensorTypeAndShape.ElementDataType} is not supported"), }; } finally { disposableArray.Dispose(); } return result; } private static DisposableNamedOnnxValue FromNativeMapElements(string name, ref OrtValue ortValueMap, int[] keysShape, ref OrtValue ortValueTensorKeys, int[] valsShape, ref OrtValue ortValueTensorValues) { if (typeof(K) == typeof(string)) { DenseTensor denseTensorKeys = new DenseTensor(ortValueTensorKeys.GetStringTensorAsArray(), keysShape); if (typeof(V) == typeof(string)) { DenseTensor denseTensorValues = new DenseTensor(ortValueTensorValues.GetStringTensorAsArray(), valsShape); Dictionary value = Enumerable.Range(0, (int)denseTensorKeys.Length).ToDictionary((int i) => denseTensorKeys[new int[1] { i }], (int i) => denseTensorValues[new int[1] { i }]); MapHelper mapHelper = new MapHelper(denseTensorKeys, denseTensorValues); DisposableNamedOnnxValue result = new DisposableNamedOnnxValue(name, value, mapHelper, ortValueMap); ortValueMap = null; return result; } OrtValueTensor ortValueTensor = new OrtValueTensor(ref ortValueTensorValues); try { DenseTensor values = new DenseTensor(ortValueTensor.Memory, ortValueTensor.Dimensions); return FromMapDenseTensors(name, ref ortValueMap, denseTensorKeys, values, ortValueTensor); } catch (Exception) { ortValueTensor.Dispose(); throw; } } DisposableList disposableList = new DisposableList(2); try { OrtValueTensor ortValueTensor2 = new OrtValueTensor(ref ortValueTensorKeys); disposableList.Add(ortValueTensor2); DenseTensor keys = new DenseTensor(ortValueTensor2.Memory, ortValueTensor2.Dimensions); if (typeof(V) == typeof(string)) { DenseTensor values2 = new DenseTensor(ortValueTensorValues.GetStringTensorAsArray(), valsShape); return FromMapDenseTensors(name, ref ortValueMap, keys, values2, disposableList); } OrtValueTensor ortValueTensor3 = new OrtValueTensor(ref ortValueTensorValues); disposableList.Add(ortValueTensor3); DenseTensor values3 = new DenseTensor(ortValueTensor3.Memory, ortValueTensor3.Dimensions); return FromMapDenseTensors(name, ref ortValueMap, keys, values3, disposableList); } catch (Exception) { disposableList.Dispose(); throw; } } private static DisposableNamedOnnxValue FromMapDenseTensors(string name, ref OrtValue ortValueMap, DenseTensor keys, DenseTensor values, IDisposable disposables) { Dictionary value = Enumerable.Range(0, (int)keys.Length).ToDictionary((int i) => keys[new int[1] { i }], (int i) => values[new int[1] { i }]); MapHelper mapHelper = new MapHelper(keys, values); NativeOrtValueCollectionOwner ortValueHolder = new NativeOrtValueCollectionOwner(ref ortValueMap, disposables); return new DisposableNamedOnnxValue(name, value, mapHelper, ortValueHolder); } protected virtual void Dispose(bool disposing) { if (!_disposed) { if (disposing && _ortValueHolder != null) { _ortValueHolder.Dispose(); _ortValueHolder = null; } _disposed = true; } } public void Dispose() { Dispose(disposing: true); } } internal enum ErrorCode { Ok, Fail, InvalidArgument, NoSuchFile, NoModel, EngineError, RuntimeException, InvalidProtobuf, ModelLoaded, NotImplemented, InvalidGraph, ShapeInferenceNotRegistered, RequirementNotRegistered } public class OnnxRuntimeException : Exception { private static Dictionary errorCodeToString = new Dictionary { { ErrorCode.Ok, "Ok" }, { ErrorCode.Fail, "Fail" }, { ErrorCode.InvalidArgument, "InvalidArgument" }, { ErrorCode.NoSuchFile, "NoSuchFile" }, { ErrorCode.NoModel, "NoModel" }, { ErrorCode.EngineError, "EngineError" }, { ErrorCode.RuntimeException, "RuntimeException" }, { ErrorCode.InvalidProtobuf, "InvalidProtobuf" }, { ErrorCode.ModelLoaded, "ModelLoaded" }, { ErrorCode.NotImplemented, "NotImplemented" }, { ErrorCode.InvalidGraph, "InvalidGraph" }, { ErrorCode.ShapeInferenceNotRegistered, "ShapeInferenceNotRegistered" }, { ErrorCode.RequirementNotRegistered, "RequirementNotRegistered" } }; internal OnnxRuntimeException(ErrorCode errorCode, string message) : base("[ErrorCode:" + errorCodeToString[errorCode] + "] " + message) { } } public class FixedBufferOnnxValue : IDisposable { private bool _disposed = false; internal OrtValue Value { get; private set; } internal OnnxValueType OnnxValueType { get; private set; } internal TensorElementType ElementType { get; private set; } private FixedBufferOnnxValue(ref OrtValue ortValue, OnnxValueType onnxValueType, TensorElementType elementType) { Value = ortValue; ortValue = null; OnnxValueType = onnxValueType; ElementType = elementType; } public static FixedBufferOnnxValue CreateFromTensor(Tensor value) { TensorElementType elementType; OrtValue ortValue = OrtValue.CreateFromTensorObject(value, out elementType); return new FixedBufferOnnxValue(ref ortValue, OnnxValueType.ONNX_TYPE_TENSOR, elementType); } public static FixedBufferOnnxValue CreateFromMemory(OrtMemoryInfo memoryInfo, Memory memory, TensorElementType elementType, long[] shape, long bytesSize) where T : unmanaged { if (elementType == TensorElementType.String) { throw new ArgumentException("String data type is not supported"); } OrtValue ortValue = OrtValue.CreateTensorValueFromMemory(memoryInfo, memory, shape); try { return new FixedBufferOnnxValue(ref ortValue, OnnxValueType.ONNX_TYPE_TENSOR, elementType); } catch (Exception) { ortValue?.Dispose(); throw; } } protected virtual void Dispose(bool disposing) { if (!_disposed) { if (disposing) { Value.Dispose(); } _disposed = true; } } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } } public class InferenceSession : IDisposable { private delegate string NameExtractor(TInput input); private delegate IntPtr OrtValueHandleExtractor(NamedOnnxValue value, NodeMetadata metadata, out IDisposable memOwner); private delegate NodeMetadata MetadataLookup(string nodeName); [UnmanagedFunctionPointer(CallingConvention.Cdecl)] private delegate void OrtCallbackDelegate(IntPtr userData, IntPtr[] outputs, uint numOutputs, IntPtr status); private delegate void UserCallbackDelegate(IReadOnlyCollection outputs, IntPtr status); private class CallbackHost { public IReadOnlyCollection inputNames { get; } public IReadOnlyCollection inputValues { get; } public IReadOnlyCollection outputNames { get; } public IReadOnlyCollection outputValues { get; } public UserCallbackDelegate callback { get; } public IntPtr[] rawInputNames { get; } public IntPtr[] rawInputValues { get; } public IntPtr[] rawOutputNames { get; } public IntPtr[] rawOutputValues { get; } public CallbackHost(InferenceSession session, IReadOnlyCollection cbInputNames, IReadOnlyCollection cbinputValues, IReadOnlyCollection cbOutputNames, IReadOnlyCollection cbOutputValues, UserCallbackDelegate userCallback) { inputNames = cbInputNames; inputValues = cbinputValues; outputNames = cbOutputNames; outputValues = cbOutputValues; callback = userCallback; rawInputNames = LookupUtf8Names(inputNames, (string n) => n, session.LookupInputMetadata); rawInputValues = inputValues.Select((OrtValue v) => v.Handle).ToArray(); rawOutputNames = LookupUtf8Names(outputNames, (string n) => n, session.LookupOutputMetadata); rawOutputValues = outputValues.Select((OrtValue v) => v.Handle).ToArray(); } } private IntPtr _nativeHandle; private Dictionary _inputMetadata; private List _inputNames; private Dictionary _outputMetadata; private List _outputNames; private Dictionary _overridableInitializerMetadata; private List _namesMemoryPtrs; private SessionOptions _builtInSessionOptions = null; private RunOptions _builtInRunOptions = null; private ModelMetadata _modelMetadata = null; private bool _disposed = false; private ulong _profilingStartTimeNs = 0uL; private static OrtCallbackDelegate ortCallback = OrtCallback; public IReadOnlyDictionary InputMetadata => _inputMetadata; public IReadOnlyList InputNames => _inputNames; public IReadOnlyDictionary OutputMetadata => _outputMetadata; public IReadOnlyList OutputNames => _outputNames; public IReadOnlyDictionary OverridableInitializerMetadata => _overridableInitializerMetadata; public ModelMetadata ModelMetadata { get { if (_modelMetadata != null) { return _modelMetadata; } _modelMetadata = new ModelMetadata(this); return _modelMetadata; } } public ulong ProfilingStartTimeNs => _profilingStartTimeNs; internal IntPtr Handle => _nativeHandle; public InferenceSession(string modelPath) { _builtInSessionOptions = new SessionOptions(); Init(modelPath, _builtInSessionOptions); } public InferenceSession(string modelPath, PrePackedWeightsContainer prepackedWeightsContainer) { _builtInSessionOptions = new SessionOptions(); Init(modelPath, _builtInSessionOptions, prepackedWeightsContainer); } public InferenceSession(string modelPath, SessionOptions options) { Init(modelPath, options); } public InferenceSession(string modelPath, SessionOptions options, PrePackedWeightsContainer prepackedWeightsContainer) { Init(modelPath, options, prepackedWeightsContainer); } public InferenceSession(byte[] model) { _builtInSessionOptions = new SessionOptions(); Init(model, _builtInSessionOptions); } public InferenceSession(byte[] model, PrePackedWeightsContainer prepackedWeightsContainer) { _builtInSessionOptions = new SessionOptions(); Init(model, _builtInSessionOptions, prepackedWeightsContainer); } public InferenceSession(byte[] model, SessionOptions options) { Init(model, options); } public InferenceSession(byte[] model, SessionOptions options, PrePackedWeightsContainer prepackedWeightsContainer) { Init(model, options, prepackedWeightsContainer); } public IDisposableReadOnlyCollection Run(IReadOnlyCollection inputs) { return Run(inputs, _outputNames); } public IDisposableReadOnlyCollection Run(IReadOnlyCollection inputs, IReadOnlyCollection outputNames) { return Run(inputs, outputNames, _builtInRunOptions); } public IDisposableReadOnlyCollection Run(IReadOnlyCollection inputs, IReadOnlyCollection outputNames, RunOptions options) { IntPtr[] inputNames = LookupUtf8Names(inputs, (NamedOnnxValue v) => v.Name, LookupInputMetadata); IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); DisposableArray disposer; IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputs, LookupInputMetadata, ExtractOrtValueHandleForInput, out disposer); try { using DisposableOrtValueHandleArray disposableOrtValueHandleArray = RunImpl(options, inputNames, ortValuesHandles, outputNames2); return CreateDisposableResult(disposableOrtValueHandleArray.Span, outputNames); } finally { disposer.Dispose(); } } public IDisposableReadOnlyCollection Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues) { return Run(inputNames, inputValues, _outputNames, _builtInRunOptions); } public IDisposableReadOnlyCollection Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames) { return Run(inputNames, inputValues, outputNames, _builtInRunOptions); } public IDisposableReadOnlyCollection Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames, RunOptions options) { if (inputNames.Count != inputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "inputNames", inputNames.Count, "inputValues", inputValues.Count)); } IntPtr[] inputNames2 = LookupUtf8Names(inputNames, (string n) => n, LookupInputMetadata); IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); using DisposableOrtValueHandleArray disposableOrtValueHandleArray = RunImpl(options, inputNames2, ortValuesHandles, outputNames2); return CreateDisposableResult(disposableOrtValueHandleArray.Span, outputNames); } public void Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues) { Run(inputNames, inputValues, outputNames, outputValues, _builtInRunOptions); } public void Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues, RunOptions options) { if (inputNames.Count != inputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "inputNames", inputNames.Count, "inputValues", inputValues.Count)); } if (outputNames.Count != outputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "outputNames", outputNames.Count, "outputValues", outputValues.Count)); } IntPtr[] inputNames2 = LookupUtf8Names(inputNames, (string n) => n, LookupInputMetadata); IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); IntPtr[] ortValuesHandles2 = GetOrtValuesHandles(outputValues, input: false); NativeApiStatus.VerifySuccess(NativeMethods.OrtRun(_nativeHandle, options.Handle, inputNames2, ortValuesHandles, (UIntPtr)(ulong)inputNames.Count, outputNames2, (UIntPtr)(ulong)outputNames.Count, ortValuesHandles2)); } public void Run(IReadOnlyCollection inputs, IReadOnlyCollection outputs) { Run(inputs, outputs, _builtInRunOptions); } public void Run(IReadOnlyCollection inputs, IReadOnlyCollection outputs, RunOptions options) { IntPtr[] inputNames = LookupUtf8Names(inputs, (NamedOnnxValue i) => i.Name, LookupInputMetadata); IntPtr[] outputNames = LookupUtf8Names(outputs, (NamedOnnxValue o) => o.Name, LookupOutputMetadata); DisposableArray disposer; IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputs, LookupInputMetadata, ExtractOrtValueHandleForInput, out disposer); try { DisposableArray disposer2; IntPtr[] ortValuesHandles2 = GetOrtValuesHandles(outputs, LookupOutputMetadata, ExtractOrtValueHandleForOutput, out disposer2); try { NativeApiStatus.VerifySuccess(NativeMethods.OrtRun(_nativeHandle, options.Handle, inputNames, ortValuesHandles, (UIntPtr)(ulong)inputs.Count, outputNames, (UIntPtr)(ulong)outputs.Count, ortValuesHandles2)); } finally { disposer2.Dispose(); } } finally { disposer.Dispose(); } } public void Run(IReadOnlyCollection inputs, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues) { Run(inputs, outputNames, outputValues, _builtInRunOptions); } public void Run(IReadOnlyCollection inputs, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues, RunOptions options) { if (outputNames.Count != outputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "outputNames", outputNames.Count, "outputValues", outputValues.Count)); } IntPtr[] inputNames = LookupUtf8Names(inputs, (NamedOnnxValue i) => i.Name, LookupInputMetadata); IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); IntPtr[] ortValuesHandles = GetOrtValuesHandles(outputValues, input: false); DisposableArray disposer; IntPtr[] ortValuesHandles2 = GetOrtValuesHandles(inputs, LookupInputMetadata, ExtractOrtValueHandleForInput, out disposer); try { NativeApiStatus.VerifySuccess(NativeMethods.OrtRun(_nativeHandle, options.Handle, inputNames, ortValuesHandles2, (UIntPtr)(ulong)inputs.Count, outputNames2, (UIntPtr)(ulong)outputNames.Count, ortValuesHandles)); } finally { disposer.Dispose(); } } public void Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputs) { Run(inputNames, inputValues, outputs, _builtInRunOptions); } public void Run(IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputs, RunOptions options) { if (inputNames.Count != inputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "inputNames", inputNames.Count, "inputValues", inputValues.Count)); } IntPtr[] inputNames2 = LookupUtf8Names(inputNames, (string n) => n, LookupInputMetadata); IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] outputNames = LookupUtf8Names(outputs, (NamedOnnxValue o) => o.Name, LookupOutputMetadata); DisposableArray disposer; IntPtr[] ortValuesHandles2 = GetOrtValuesHandles(outputs, LookupOutputMetadata, ExtractOrtValueHandleForOutput, out disposer); try { NativeApiStatus.VerifySuccess(NativeMethods.OrtRun(_nativeHandle, options.Handle, inputNames2, ortValuesHandles, (UIntPtr)(ulong)inputNames.Count, outputNames, (UIntPtr)(ulong)outputs.Count, ortValuesHandles2)); } finally { disposer.Dispose(); } } public IDisposableReadOnlyCollection Run(RunOptions runOptions, IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames) { if (inputNames.Count != inputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "inputNames", inputNames.Count, "inputValues", inputValues.Count)); } IntPtr[] inputNames2 = LookupUtf8Names(inputNames, (string n) => n, LookupInputMetadata); IntPtr[] inputValues2 = inputValues.Select((OrtValue v) => v.Handle).ToArray(); IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); using DisposableOrtValueHandleArray disposableHandles = RunImpl(runOptions, inputNames2, inputValues2, outputNames2); return CreateDisposableResult(disposableHandles); } public IDisposableReadOnlyCollection Run(RunOptions runOptions, IReadOnlyDictionary inputs, IReadOnlyCollection outputNames) { IntPtr[] array = new IntPtr[inputs.Count]; IntPtr[] array2 = new IntPtr[inputs.Count]; int num = 0; foreach (KeyValuePair input in inputs) { array[num] = LookupInputMetadata(input.Key).ZeroTerminatedName; array2[num] = input.Value.Handle; num++; } IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); using DisposableOrtValueHandleArray disposableHandles = RunImpl(runOptions, array, array2, outputNames2); return CreateDisposableResult(disposableHandles); } private static IDisposableReadOnlyCollection CreateDisposableResult(DisposableOrtValueHandleArray disposableHandles) { DisposableList disposableList = new DisposableList(disposableHandles.Span.Length); try { for (int i = 0; i < disposableHandles.Span.Length; i++) { disposableList.Add(new OrtValue(disposableHandles.Span[i])); disposableHandles.Span[i] = IntPtr.Zero; } return disposableList; } catch (Exception) { disposableList.Dispose(); throw; } } public void Run(RunOptions runOptions, IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues) { if (inputNames.Count != inputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "inputNames", inputNames.Count, "inputValues", inputValues.Count)); } if (outputNames.Count != outputValues.Count) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of {2} ({3}).", "outputNames", outputNames.Count, "outputValues", outputValues.Count)); } if (runOptions == null) { runOptions = _builtInRunOptions; } IntPtr[] inputNames2 = LookupUtf8Names(inputNames, (string n) => n, LookupInputMetadata); IntPtr[] inputValues2 = inputValues.Select((OrtValue v) => v.Handle).ToArray(); IntPtr[] outputNames2 = LookupUtf8Names(outputNames, (string n) => n, LookupOutputMetadata); IntPtr[] outputValues2 = outputValues.Select((OrtValue v) => v.Handle).ToArray(); NativeApiStatus.VerifySuccess(NativeMethods.OrtRun(_nativeHandle, runOptions.Handle, inputNames2, inputValues2, (UIntPtr)(ulong)inputNames.Count, outputNames2, (UIntPtr)(ulong)outputNames.Count, outputValues2)); } public OrtIoBinding CreateIoBinding() { return new OrtIoBinding(this); } public void RunWithBinding(RunOptions runOptions, OrtIoBinding ioBinding) { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunWithBinding(Handle, runOptions.Handle, ioBinding.Handle)); } public IDisposableReadOnlyCollection RunWithBoundResults(RunOptions runOptions, OrtIoBinding ioBinding) { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunWithBinding(Handle, runOptions.Handle, ioBinding.Handle)); return ioBinding.GetOutputValues(); } public IDisposableReadOnlyCollection RunWithBindingAndNames(RunOptions runOptions, OrtIoBinding ioBinding, string[] names = null) { string[] array = names; if (array == null || names.Length == 0) { array = ioBinding.GetOutputNames(); } NativeApiStatus.VerifySuccess(NativeMethods.OrtRunWithBinding(Handle, runOptions.Handle, ioBinding.Handle)); OrtValue[] outputOrtValues = ioBinding.GetOutputOrtValues(); DisposableArray disposableArray = new DisposableArray(outputOrtValues); try { DisposableList disposableList = new DisposableList(outputOrtValues.Length); try { for (int i = 0; i < array.Length; i++) { disposableList.Add(DisposableNamedOnnxValue.CreateFromOrtValue(array[i], ref outputOrtValues[i])); } return disposableList; } catch (Exception) { disposableList.Dispose(); throw; } } finally { disposableArray.Dispose(); } } public string EndProfiling() { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionEndProfiling(_nativeHandle, defaultInstance.Pointer, out var profile_file)); return NativeOnnxValueHelper.StringFromNativeUtf8(profile_file, defaultInstance); } private NodeMetadata LookupInputMetadata(string nodeName) { if (!_inputMetadata.TryGetValue(nodeName, out var value) && !_overridableInitializerMetadata.TryGetValue(nodeName, out value)) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Input name: '" + nodeName + "' is not in the metadata"); } return value; } private NodeMetadata LookupOutputMetadata(string nodeName) { if (!_outputMetadata.TryGetValue(nodeName, out var value)) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Output name: '" + nodeName + "' is not in the metadata"); } return value; } private static IntPtr ExtractOrtValueHandleForInput(NamedOnnxValue input, NodeMetadata metadata, out IDisposable memOwner) { return input.InputToOrtValueHandle(metadata, out memOwner); } private static IntPtr ExtractOrtValueHandleForOutput(NamedOnnxValue output, NodeMetadata metadata, out IDisposable memOwner) { return output.OutputToOrtValueHandle(metadata, out memOwner); } private static IntPtr[] LookupUtf8Names(IReadOnlyCollection values, NameExtractor nameExtractor, MetadataLookup metaLookup) { IntPtr[] array = new IntPtr[values.Count]; for (int i = 0; i < values.Count; i++) { string nodeName = nameExtractor(values.ElementAt(i)); NodeMetadata nodeMetadata = metaLookup(nodeName); array[i] = nodeMetadata.ZeroTerminatedName; } return array; } private static IntPtr[] GetOrtValuesHandles(IReadOnlyCollection values, MetadataLookup metaLookup, OrtValueHandleExtractor ortValueExtractor, out DisposableArray disposer) { IDisposable[] array = new IDisposable[values.Count]; DisposableArray disposableArray = new DisposableArray(array); try { IntPtr[] array2 = new IntPtr[values.Count]; for (int i = 0; i < values.Count; i++) { NamedOnnxValue namedOnnxValue = values.ElementAt(i); NodeMetadata metadata = metaLookup(namedOnnxValue.Name); array2[i] = ortValueExtractor(namedOnnxValue, metadata, out var memOwner); if (memOwner != null) { array[i] = memOwner; } } disposer = disposableArray; return array2; } catch (Exception) { disposableArray.Dispose(); throw; } } private static IntPtr[] GetOrtValuesHandles(IReadOnlyCollection values, bool input) { IntPtr[] array = new IntPtr[values.Count]; for (int i = 0; i < values.Count; i++) { FixedBufferOnnxValue fixedBufferOnnxValue = values.ElementAt(i); if (!input && fixedBufferOnnxValue.ElementType == TensorElementType.String) { throw new NotSupportedException("Using string type FixedBufferOnnxValue in outputs is not supported."); } array[i] = fixedBufferOnnxValue.Value.Handle; } return array; } private DisposableOrtValueHandleArray RunImpl(RunOptions options, IntPtr[] inputNames, IntPtr[] inputValues, IntPtr[] outputNames) { IntPtr[] array = new IntPtr[outputNames.Length]; NativeApiStatus.VerifySuccess(NativeMethods.OrtRun(_nativeHandle, options.Handle, inputNames, inputValues, (UIntPtr)(ulong)inputNames.Length, outputNames, (UIntPtr)(ulong)outputNames.Length, array)); return new DisposableOrtValueHandleArray(array); } private static IDisposableReadOnlyCollection CreateDisposableResult(Span valueHandles, IReadOnlyCollection outputNames) { DisposableList disposableList = new DisposableList(valueHandles.Length); try { for (int i = 0; i < valueHandles.Length; i++) { OrtValue ortValue = new OrtValue(valueHandles[i]); disposableList.Add(DisposableNamedOnnxValue.CreateFromOrtValue(outputNames.ElementAt(i), ref ortValue)); valueHandles[i] = IntPtr.Zero; } } catch (OnnxRuntimeException) { disposableList.Dispose(); throw; } return disposableList; } private static void OrtCallback(IntPtr userData, IntPtr[] ouputs, uint numOutputs, IntPtr status) { GCHandle gCHandle = GCHandle.FromIntPtr(userData); CallbackHost callbackHost = (CallbackHost)gCHandle.Target; try { callbackHost.callback(callbackHost.outputValues, status); } finally { gCHandle.Free(); } } private unsafe void RunAsyncInternal(RunOptions options, IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues, UserCallbackDelegate callback) { CallbackHost callbackHost = new CallbackHost(this, inputNames, inputValues, outputNames, outputValues, callback); GCHandle value = GCHandle.Alloc(callbackHost, GCHandleType.Normal); try { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunAsync(_nativeHandle, options?.Handle ?? ((IntPtr)(void*)null), callbackHost.rawInputNames, callbackHost.rawInputValues, (UIntPtr)(ulong)callbackHost.rawInputNames.Length, callbackHost.rawOutputNames, (UIntPtr)(ulong)callbackHost.rawOutputNames.Length, callbackHost.rawOutputValues, Marshal.GetFunctionPointerForDelegate(ortCallback), GCHandle.ToIntPtr(value))); } catch (OnnxRuntimeException) { value.Free(); throw; } } public async Task> RunAsync(RunOptions options, IReadOnlyCollection inputNames, IReadOnlyCollection inputValues, IReadOnlyCollection outputNames, IReadOnlyCollection outputValues) { TaskCompletionSource> promise = new TaskCompletionSource>(); RunAsyncInternal(options, inputNames, inputValues, outputNames, outputValues, delegate(IReadOnlyCollection outputs, IntPtr status) { try { NativeApiStatus.VerifySuccess(status); promise.SetResult(outputs); } catch (Exception exception) { promise.SetException(exception); } }); return await promise.Task; } private void Init(string modelPath, SessionOptions options, PrePackedWeightsContainer prepackedWeightsContainer = null) { IntPtr handle = OrtEnv.Instance().Handle; IntPtr session; if (prepackedWeightsContainer == null) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateSession(handle, NativeOnnxValueHelper.GetPlatformSerializedString(modelPath), options.Handle, out session)); } else { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateSessionWithPrepackedWeightsContainer(handle, NativeOnnxValueHelper.GetPlatformSerializedString(modelPath), options.Handle, prepackedWeightsContainer.Pointer, out session)); } InitWithSessionHandle(session); } private void Init(byte[] modelData, SessionOptions options, PrePackedWeightsContainer prepackedWeightsContainer = null) { IntPtr handle = OrtEnv.Instance().Handle; IntPtr session; if (prepackedWeightsContainer == null) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateSessionFromArray(handle, modelData, (UIntPtr)(ulong)modelData.Length, options.Handle, out session)); } else { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateSessionFromArrayWithPrepackedWeightsContainer(handle, modelData, (UIntPtr)(ulong)modelData.Length, options.Handle, prepackedWeightsContainer.Pointer, out session)); } InitWithSessionHandle(session); } private void InitWithSessionHandle(IntPtr session) { _nativeHandle = session; try { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetInputCount(_nativeHandle, out var count)); NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOutputCount(_nativeHandle, out var count2)); NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOverridableInitializerCount(_nativeHandle, out var count3)); int capacity = (int)((uint)count + (uint)count2 + (uint)count3); _namesMemoryPtrs = new List(capacity); _inputMetadata = new Dictionary((int)(uint)count); _inputNames = new List((int)(uint)count); for (ulong num = 0uL; num < (ulong)count; num++) { NodeMetadata inputMetadata = GetInputMetadata(num); IntPtr utf; string inputName = GetInputName(num, out utf); _namesMemoryPtrs.Add(utf); inputMetadata.ZeroTerminatedName = utf; _inputNames.Add(inputName); _inputMetadata[inputName] = inputMetadata; } _outputMetadata = new Dictionary((int)(uint)count2); _outputNames = new List((int)(uint)count2); for (ulong num2 = 0uL; num2 < (ulong)count2; num2++) { NodeMetadata outputMetadata = GetOutputMetadata(num2); IntPtr utf2; string outputName = GetOutputName(num2, out utf2); _namesMemoryPtrs.Add(utf2); outputMetadata.ZeroTerminatedName = utf2; _outputNames.Add(outputName); _outputMetadata[outputName] = outputMetadata; } _overridableInitializerMetadata = new Dictionary((int)(uint)count3); for (ulong num3 = 0uL; num3 < (ulong)count3; num3++) { NodeMetadata overridableInitializerMetadata = GetOverridableInitializerMetadata(num3); IntPtr utf3; string overridableInitializerName = GetOverridableInitializerName(num3, out utf3); _namesMemoryPtrs.Add(utf3); overridableInitializerMetadata.ZeroTerminatedName = utf3; _overridableInitializerMetadata[overridableInitializerName] = overridableInitializerMetadata; } NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetProfilingStartTimeNs(_nativeHandle, out var startTime)); _profilingStartTimeNs = (ulong)startTime; } catch (Exception) { DisposeImpl(disposing: true); throw; } _builtInRunOptions = new RunOptions(); } private string GetOutputName(ulong index, out IntPtr utf8) { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOutputName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out var name)); NativeOnnxValueHelper.StringAndUtf8FromNative(defaultInstance, name, out var str, out utf8); return str; } private string GetInputName(ulong index, out IntPtr utf8) { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetInputName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out var name)); NativeOnnxValueHelper.StringAndUtf8FromNative(defaultInstance, name, out var str, out utf8); return str; } private string GetOverridableInitializerName(ulong index, out IntPtr utf8) { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOverridableInitializerName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out var name)); NativeOnnxValueHelper.StringAndUtf8FromNative(defaultInstance, name, out var str, out utf8); return str; } private NodeMetadata GetInputMetadata(ulong index) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetInputTypeInfo(_nativeHandle, (UIntPtr)index, out var typeInfo)); try { return GetMetadataFromTypeInfo(typeInfo); } finally { NativeMethods.OrtReleaseTypeInfo(typeInfo); } } private NodeMetadata GetOutputMetadata(ulong index) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOutputTypeInfo(_nativeHandle, (UIntPtr)index, out var typeInfo)); try { return GetMetadataFromTypeInfo(typeInfo); } finally { NativeMethods.OrtReleaseTypeInfo(typeInfo); } } private NodeMetadata GetOverridableInitializerMetadata(ulong index) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetOverridableInitializerTypeInfo(_nativeHandle, (UIntPtr)index, out var typeInfo)); try { return GetMetadataFromTypeInfo(typeInfo); } finally { NativeMethods.OrtReleaseTypeInfo(typeInfo); } } internal static NodeMetadata GetMetadataFromTypeInfo(IntPtr typeInfo) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetOnnxTypeFromTypeInfo(typeInfo, out var onnxtype)); OnnxValueType onnxValueType = (OnnxValueType)(int)onnxtype; switch (onnxValueType) { case OnnxValueType.ONNX_TYPE_TENSOR: case OnnxValueType.ONNX_TYPE_SPARSETENSOR: return GetTensorNodeMetadata(onnxValueType, typeInfo); case OnnxValueType.ONNX_TYPE_SEQUENCE: return GetSequenceMetadataFromTypeInfo(typeInfo); case OnnxValueType.ONNX_TYPE_MAP: return GetMapMetadataFromTypeInfo(typeInfo); case OnnxValueType.ONNX_TYPE_OPTIONAL: return GetOptionalMetadataFromTypeInfo(typeInfo); default: throw new OnnxRuntimeException(ErrorCode.NotImplemented, $"Value type: '{onnxValueType}' not supported in this code"); } } internal static NodeMetadata GetSequenceMetadataFromTypeInfo(IntPtr typeInfo) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToSequenceTypeInfo(typeInfo, out var sequenceTypeInfo)); if (sequenceTypeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.Fail, "TypeInfo cast to SequenceTypeInfo failed. The object does not represent a sequence"); } NativeApiStatus.VerifySuccess(NativeMethods.OrtGetSequenceElementType(sequenceTypeInfo, out var elementTypeInfo)); try { NodeMetadata metadataFromTypeInfo = GetMetadataFromTypeInfo(elementTypeInfo); SequenceMetadata sequenceMetadata = new SequenceMetadata(metadataFromTypeInfo); return new NodeMetadata(sequenceMetadata); } finally { NativeMethods.OrtReleaseTypeInfo(elementTypeInfo); } } internal static NodeMetadata GetMapMetadataFromTypeInfo(IntPtr typeInfo) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToMapTypeInfo(typeInfo, out var mapTypeInfo)); if (mapTypeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.Fail, "TypeInfo cast to MapTypeInfo failed. The object does not represent a map"); } NativeApiStatus.VerifySuccess(NativeMethods.OrtGetMapKeyType(mapTypeInfo, out var tensorElementType)); NativeApiStatus.VerifySuccess(NativeMethods.OrtGetMapValueType(mapTypeInfo, out var type_info)); try { NodeMetadata metadataFromTypeInfo = GetMetadataFromTypeInfo(type_info); MapMetadata mapMetadata = new MapMetadata((TensorElementType)(int)tensorElementType, metadataFromTypeInfo); return new NodeMetadata(mapMetadata); } finally { NativeMethods.OrtReleaseTypeInfo(type_info); } } internal static NodeMetadata GetOptionalMetadataFromTypeInfo(IntPtr typeInfo) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToOptionalTypeInfo(typeInfo, out var optionalTypeInfo)); if (optionalTypeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.Fail, "TypeInfo cast to OptionalTypeInfo failed. The object does not represent a optional value"); } NativeApiStatus.VerifySuccess(NativeMethods.OrtGetOptionalContainedTypeInfo(optionalTypeInfo, out var containedTypeInfo)); try { NodeMetadata metadataFromTypeInfo = GetMetadataFromTypeInfo(containedTypeInfo); OptionalMetadata optMetadata = new OptionalMetadata(metadataFromTypeInfo); return new NodeMetadata(optMetadata); } finally { NativeMethods.OrtReleaseTypeInfo(containedTypeInfo); } } internal static NodeMetadata GetTensorNodeMetadata(OnnxValueType valueType, IntPtr typeInfo) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToTensorInfo(typeInfo, out var typeAndShapeInfo)); if (typeAndShapeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.Fail, "TypeInfo cast to TensorTypeInfo failed. The object does not represent a tensor"); } NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorElementType(typeAndShapeInfo, out var output)); TensorElementType elementType = (TensorElementType)(int)output; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetDimensionsCount(typeAndShapeInfo, out var output2)); long[] array = new long[(uint)output2]; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetDimensions(typeAndShapeInfo, array, output2)); int[] array2 = new int[(uint)output2]; for (int i = 0; (long)i < (long)(ulong)output2; i++) { array2[i] = (int)array[i]; } IntPtr[] array3 = new IntPtr[(uint)output2]; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetSymbolicDimensions(typeAndShapeInfo, array3, output2)); string[] array4 = new string[(uint)output2]; for (int j = 0; j < (int)(uint)output2; j++) { array4[j] = NativeOnnxValueHelper.StringFromNativeUtf8(array3[j]); } TensorTypeAndShape typeAndShape = new TensorTypeAndShape(elementType, array2, array4); return new NodeMetadata(valueType, typeAndShape); } ~InferenceSession() { Dispose(disposing: false); } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } protected virtual void Dispose(bool disposing) { if (!_disposed) { DisposeImpl(disposing); } } private void DisposeImpl(bool disposing) { if (disposing) { if (_namesMemoryPtrs != null) { foreach (IntPtr namesMemoryPtr in _namesMemoryPtrs) { Marshal.FreeHGlobal(namesMemoryPtr); } _namesMemoryPtrs = null; } if (_builtInSessionOptions != null) { _builtInSessionOptions.Dispose(); _builtInSessionOptions = null; } if (_builtInRunOptions != null) { _builtInRunOptions.Dispose(); _builtInRunOptions = null; } } if (_nativeHandle != IntPtr.Zero) { NativeMethods.OrtReleaseSession(_nativeHandle); _nativeHandle = IntPtr.Zero; } _disposed = true; } } public class TensorTypeAndShape { public TensorElementType ElementDataType { get; } public int[] Dimensions { get; } public string[] SymbolicDimensions { get; } public TensorElementTypeInfo ElementTypeInfo { get; } internal TensorTypeAndShape(TensorElementType elementType, int[] dimensions, string[] symbolicDimensions) { ElementTypeInfo = TensorBase.GetElementTypeInfo(elementType); if (ElementTypeInfo == null) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Unregistered TensorElementType value of: " + elementType); } ElementDataType = elementType; Dimensions = dimensions; SymbolicDimensions = symbolicDimensions; } } public class SequenceMetadata { public NodeMetadata ElementMeta { get; } internal SequenceMetadata(NodeMetadata elementData) { ElementMeta = elementData; } } public class OptionalMetadata { public NodeMetadata ElementMeta { get; } internal OptionalMetadata(NodeMetadata elementData) { ElementMeta = elementData; } } public class MapMetadata { public TensorElementType KeyDataType { get; } public NodeMetadata ValueMetadata { get; } internal MapMetadata(TensorElementType keyDataType, NodeMetadata valueMetadata) { KeyDataType = keyDataType; ValueMetadata = valueMetadata; } } public class NodeMetadata { private readonly object _metadata; public OnnxValueType OnnxValueType { get; } internal IntPtr ZeroTerminatedName { get; set; } public int[] Dimensions { get { CheckTensor(); return (_metadata as TensorTypeAndShape).Dimensions; } } public string[] SymbolicDimensions { get { CheckTensor(); return (_metadata as TensorTypeAndShape).SymbolicDimensions; } } public Type ElementType { get { CheckTensor(); return (_metadata as TensorTypeAndShape).ElementTypeInfo.TensorType; } } public TensorElementType ElementDataType { get { CheckTensor(); return (_metadata as TensorTypeAndShape).ElementDataType; } } public bool IsString { get { CheckTensor(); return (_metadata as TensorTypeAndShape).ElementTypeInfo.IsString; } } public bool IsTensor => OnnxValueType == OnnxValueType.ONNX_TYPE_TENSOR || OnnxValueType == OnnxValueType.ONNX_TYPE_SPARSETENSOR; internal NodeMetadata(OnnxValueType onnxValueType, TensorTypeAndShape typeAndShape) { OnnxValueType = onnxValueType; CheckTensor(); _metadata = typeAndShape; } internal NodeMetadata(MapMetadata mapMetadata) { OnnxValueType = OnnxValueType.ONNX_TYPE_MAP; _metadata = mapMetadata; } internal NodeMetadata(SequenceMetadata sequenceMetadata) { OnnxValueType = OnnxValueType.ONNX_TYPE_SEQUENCE; _metadata = sequenceMetadata; } internal NodeMetadata(OptionalMetadata optMetadata) { OnnxValueType = OnnxValueType.ONNX_TYPE_OPTIONAL; _metadata = optMetadata; } private void CheckTensor() { if (!IsTensor) { throw new OnnxRuntimeException(ErrorCode.Fail, "OnnxValueType must either be a tensor or sparse tensor"); } } public MapMetadata AsMapMetadata() { if (OnnxValueType != OnnxValueType.ONNX_TYPE_MAP) { throw new OnnxRuntimeException(ErrorCode.Fail, "Instance does not contain Map metadata"); } return _metadata as MapMetadata; } public SequenceMetadata AsSequenceMetadata() { if (OnnxValueType != OnnxValueType.ONNX_TYPE_SEQUENCE) { throw new OnnxRuntimeException(ErrorCode.Fail, "Instance does not contain Sequence metadata"); } return _metadata as SequenceMetadata; } public OptionalMetadata AsOptionalMetadata() { if (OnnxValueType != OnnxValueType.ONNX_TYPE_OPTIONAL) { throw new OnnxRuntimeException(ErrorCode.Fail, "Instance does not contain Optional metadata"); } return _metadata as OptionalMetadata; } } public class ModelMetadata { private string _producerName; private string _graphName; private string _domain; private string _description; private string _graphDescription; private long _version; private Dictionary _customMetadataMap = new Dictionary(); public string ProducerName => _producerName; public string GraphName => _graphName; public string Domain => _domain; public string Description => _description; public string GraphDescription => _graphDescription; public long Version => _version; public Dictionary CustomMetadataMap => _customMetadataMap; internal unsafe ModelMetadata(InferenceSession session) { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionGetModelMetadata(session.Handle, out var modelMetadata)); try { NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetProducerName(modelMetadata, defaultInstance.Pointer, out var value)); _producerName = NativeOnnxValueHelper.StringFromNativeUtf8(value, defaultInstance); NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetGraphName(modelMetadata, defaultInstance.Pointer, out var value2)); _graphName = NativeOnnxValueHelper.StringFromNativeUtf8(value2, defaultInstance); NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetDomain(modelMetadata, defaultInstance.Pointer, out var value3)); _domain = NativeOnnxValueHelper.StringFromNativeUtf8(value3, defaultInstance); NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetDescription(modelMetadata, defaultInstance.Pointer, out var value4)); _description = NativeOnnxValueHelper.StringFromNativeUtf8(value4, defaultInstance); NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetGraphDescription(modelMetadata, defaultInstance.Pointer, out var value5)); _graphDescription = NativeOnnxValueHelper.StringFromNativeUtf8(value5, defaultInstance); NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetVersion(modelMetadata, out _version)); NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataGetCustomMetadataMapKeys(modelMetadata, defaultInstance.Pointer, out var keys, out var numKeys)); try { Span span = new Span(keys.ToPointer(), (int)numKeys); using DisposableList disposableList = new DisposableList((int)numKeys); Span span2 = span; for (int i = 0; i < span2.Length; i++) { IntPtr pointer = span2[i]; disposableList.Add(new OrtMemoryAllocation(defaultInstance, pointer, 0u)); } Span span3 = span; for (int j = 0; j < span3.Length; j++) { IntPtr intPtr = span3[j]; NativeApiStatus.VerifySuccess(NativeMethods.OrtModelMetadataLookupCustomMetadataMap(modelMetadata, defaultInstance.Pointer, intPtr, out var value6)); string value7 = NativeOnnxValueHelper.StringFromNativeUtf8(value6, defaultInstance); string key = NativeOnnxValueHelper.StringFromNativeUtf8(intPtr); _customMetadataMap[key] = value7; } } finally { defaultInstance.FreeMemory(keys); } } finally { NativeMethods.OrtReleaseModelMetadata(modelMetadata); } } } internal class ManagedTypeProjection { internal static OrtValue CreateProjection(NamedOnnxValue namedOnnxValue, NodeMetadata metadata) { NodeMetadata nodeMetadata = metadata; if (metadata.OnnxValueType == OnnxValueType.ONNX_TYPE_OPTIONAL) { nodeMetadata = metadata.AsOptionalMetadata().ElementMeta; } if (namedOnnxValue.ValueType != nodeMetadata.OnnxValueType) { throw new OnnxRuntimeException(ErrorCode.RuntimeException, $"NamedOnnxValue: {namedOnnxValue.Name} has value type: {namedOnnxValue.ValueType}" + $" expected: {nodeMetadata.OnnxValueType} after optional type adjustment"); } return namedOnnxValue.ValueType switch { OnnxValueType.ONNX_TYPE_TENSOR => CreateTensorProjection(namedOnnxValue, nodeMetadata), OnnxValueType.ONNX_TYPE_SEQUENCE => CreateSequenceProjection(namedOnnxValue, nodeMetadata), OnnxValueType.ONNX_TYPE_MAP => CreateMapProjection(namedOnnxValue, nodeMetadata), _ => throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "ManagedTypeProjection can only project tensors, sequences, maps and optional types"), }; } private static OrtValue CreateSequenceProjection(NamedOnnxValue namedOnnxValue, NodeMetadata metadata) { NodeMetadata elementMeta = metadata.AsSequenceMetadata().ElementMeta; OnnxValueType onnxValueType = elementMeta.OnnxValueType; IEnumerable enumerable = namedOnnxValue.AsEnumerable() ?? throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "NamedOnnxValue: " + namedOnnxValue.Name + " sequence does not contain NamedOnnxValue elements"); int count = 0; if (enumerable is ICollection collection) { count = collection.Count; } DisposableList compositeMembers = new DisposableList(count); try { foreach (NamedOnnxValue item in enumerable) { if (onnxValueType != item.ValueType) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"NamedOnnxValue: {namedOnnxValue.Name} sequence element expected to be {onnxValueType}, received {item.ValueType}"); } compositeMembers.Add(CreateProjection(item, elementMeta)); } return OrtValue.CreateSequence(ref compositeMembers); } catch (Exception) { compositeMembers?.Dispose(); throw; } } private static OrtValue CreateMapProjection(NamedOnnxValue node, NodeMetadata elementMeta) { MapMetadata mapMetadata = elementMeta.AsMapMetadata(); NodeMetadata valueMetadata = mapMetadata.ValueMetadata; if (valueMetadata.OnnxValueType != OnnxValueType.ONNX_TYPE_TENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Node: " + node.Name + " onnxruntime only supports maps with primitive types values"); } Span disposables = new OrtValue[2]; DisposableArray disposableArray = new DisposableArray(disposables); try { TensorBase dictionaryKeys = node.GetDictionaryKeys(); disposables[0] = OrtValue.CreateFromTensorObject(dictionaryKeys, out var elementType); if (elementType != mapMetadata.KeyDataType) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Map key data type supplied: {elementType} metadata expected: {mapMetadata.KeyDataType}"); } TensorBase dictionaryValues = node.GetDictionaryValues(); disposables[1] = OrtValue.CreateFromTensorObject(dictionaryValues, out var elementType2); if (elementType2 != valueMetadata.ElementDataType) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Map value data type supplied: {elementType2} metadata expected: {valueMetadata.ElementDataType}"); } return OrtValue.CreateMap(ref disposables[0], ref disposables[1]); } catch (Exception) { disposableArray.Dispose(); throw; } } private static OrtValue CreateTensorProjection(NamedOnnxValue node, NodeMetadata elementMeta) { if (!(node.Value is TensorBase)) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"NamedOnnxValue contains: {node.Value.GetType()}, expecting a Tensor"); } TensorElementType elementType; OrtValue ortValue = OrtValue.CreateFromTensorObject(node.Value as TensorBase, out elementType); try { if (elementType != elementMeta.ElementDataType) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Tensor element data type discovered: {elementType} metadata expected: {elementMeta.ElementDataType}"); } } catch (Exception) { ortValue.Dispose(); throw; } return ortValue; } } internal class MapHelper { internal TensorBase Keys { get; } internal TensorBase Values { get; } internal MapHelper(TensorBase keys, TensorBase values) { Keys = keys; Values = values; } } public class NamedOnnxValue { private object _value; private string _name; private MapHelper _mapHelper; public OnnxValueType ValueType { get; internal set; } public string Name { get { return _name; } set { _name = value; } } public object Value { get { return _value; } set { _value = value; } } [Obsolete("Use constructors with valueType or static factory methods")] protected NamedOnnxValue(string name, object value) { _name = name; _value = value; ValueType = OnnxValueType.ONNX_TYPE_UNKNOWN; } internal NamedOnnxValue(string name, object value, OnnxValueType valueType) { _name = name; _value = value; ValueType = valueType; if (valueType == OnnxValueType.ONNX_TYPE_MAP) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Use another __ctor for maps"); } } internal NamedOnnxValue(string name, object value, MapHelper helper) { _name = name; _value = value; ValueType = OnnxValueType.ONNX_TYPE_MAP; _mapHelper = helper; } public static NamedOnnxValue CreateFromTensor(string name, Tensor value) { return new NamedOnnxValue(name, value, OnnxValueType.ONNX_TYPE_TENSOR); } public static NamedOnnxValue CreateFromSequence(string name, IEnumerable value) { return new NamedOnnxValue(name, value, OnnxValueType.ONNX_TYPE_SEQUENCE); } public static NamedOnnxValue CreateFromMap(string name, IDictionary value) { return CreateFromMap(name, value.Keys, value.Values); } internal static NamedOnnxValue CreateFromMap(string name, ICollection keys, ICollection values) { DenseTensor keys2 = new DenseTensor(keys.ToArray(), new int[1] { keys.Count }); DenseTensor values2 = new DenseTensor(values.ToArray(), new int[1] { values.Count }); return new NamedOnnxValue(name, values, new MapHelper(keys2, values2)); } public Tensor AsTensor() { return _value as Tensor; } public IEnumerable AsEnumerable() { return _value as IEnumerable; } public IDictionary AsDictionary() { return _value as IDictionary; } internal virtual IntPtr InputToOrtValueHandle(NodeMetadata metadata, out IDisposable memoryOwner) { return ((OrtValue)(memoryOwner = ManagedTypeProjection.CreateProjection(this, metadata))).Handle; } internal virtual IntPtr OutputToOrtValueHandle(NodeMetadata metadata, out IDisposable memoryOwner) { if (metadata.OnnxValueType == OnnxValueType.ONNX_TYPE_TENSOR) { return ((OrtValue)(memoryOwner = ManagedTypeProjection.CreateProjection(this, metadata))).Handle; } if (metadata.OnnxValueType == OnnxValueType.ONNX_TYPE_OPTIONAL) { NodeMetadata elementMeta = metadata.AsOptionalMetadata().ElementMeta; if (elementMeta.OnnxValueType == OnnxValueType.ONNX_TYPE_TENSOR) { return ((OrtValue)(memoryOwner = ManagedTypeProjection.CreateProjection(this, metadata))).Handle; } } throw new OnnxRuntimeException(ErrorCode.NotImplemented, $"Can not create output OrtValue for NamedOnnxValue '{metadata.OnnxValueType}' type." + " Only tensors can be pre-allocated for outputs Use Run() overloads that return DisposableNamedOnnxValue to get access to all Onnx value types that may be returned as output."); } internal TensorBase GetDictionaryKeys() { if (ValueType != OnnxValueType.ONNX_TYPE_MAP) { throw new OnnxRuntimeException(ErrorCode.Fail, "This NamedOnnxValue instance does not contain a dictionary"); } return _mapHelper.Keys; } internal TensorBase GetDictionaryValues() { if (ValueType != OnnxValueType.ONNX_TYPE_MAP) { throw new OnnxRuntimeException(ErrorCode.Fail, "This NamedOnnxValue instance does not contain a dictionary"); } return _mapHelper.Values; } } internal class NativeApiStatus { private static string GetErrorMessage(IntPtr status) { IntPtr nativeUtf = NativeMethods.OrtGetErrorMessage(status); return NativeOnnxValueHelper.StringFromNativeUtf8(nativeUtf); } [MethodImpl(MethodImplOptions.AggressiveInlining)] public static void VerifySuccess(IntPtr nativeStatus) { if (nativeStatus != IntPtr.Zero) { try { ErrorCode errorCode = NativeMethods.OrtGetErrorCode(nativeStatus); string errorMessage = GetErrorMessage(nativeStatus); throw new OnnxRuntimeException(errorCode, errorMessage); } finally { NativeMethods.OrtReleaseStatus(nativeStatus); } } } } public struct OrtApiBase { public IntPtr GetApi; public IntPtr GetVersionString; } public struct OrtApi { public IntPtr CreateStatus; public IntPtr GetErrorCode; public IntPtr GetErrorMessage; public IntPtr CreateEnv; public IntPtr CreateEnvWithCustomLogger; public IntPtr EnableTelemetryEvents; public IntPtr DisableTelemetryEvents; public IntPtr CreateSession; public IntPtr CreateSessionFromArray; public IntPtr Run; public IntPtr CreateSessionOptions; public IntPtr SetOptimizedModelFilePath; public IntPtr CloneSessionOptions; public IntPtr SetSessionExecutionMode; public IntPtr EnableProfiling; public IntPtr DisableProfiling; public IntPtr EnableMemPattern; public IntPtr DisableMemPattern; public IntPtr EnableCpuMemArena; public IntPtr DisableCpuMemArena; public IntPtr SetSessionLogId; public IntPtr SetSessionLogVerbosityLevel; public IntPtr SetSessionLogSeverityLevel; public IntPtr SetSessionGraphOptimizationLevel; public IntPtr SetIntraOpNumThreads; public IntPtr SetInterOpNumThreads; public IntPtr CreateCustomOpDomain; public IntPtr CustomOpDomain_Add; public IntPtr AddCustomOpDomain; public IntPtr RegisterCustomOpsLibrary; public IntPtr SessionGetInputCount; public IntPtr SessionGetOutputCount; public IntPtr SessionGetOverridableInitializerCount; public IntPtr SessionGetInputTypeInfo; public IntPtr SessionGetOutputTypeInfo; public IntPtr SessionGetOverridableInitializerTypeInfo; public IntPtr SessionGetInputName; public IntPtr SessionGetOutputName; public IntPtr SessionGetOverridableInitializerName; public IntPtr CreateRunOptions; public IntPtr RunOptionsSetRunLogVerbosityLevel; public IntPtr RunOptionsSetRunLogSeverityLevel; public IntPtr RunOptionsSetRunTag; public IntPtr RunOptionsGetRunLogVerbosityLevel; public IntPtr RunOptionsGetRunLogSeverityLevel; public IntPtr RunOptionsGetRunTag; public IntPtr RunOptionsSetTerminate; public IntPtr RunOptionsUnsetTerminate; public IntPtr CreateTensorAsOrtValue; public IntPtr CreateTensorWithDataAsOrtValue; public IntPtr IsTensor; public IntPtr GetTensorMutableData; public IntPtr FillStringTensor; public IntPtr GetStringTensorDataLength; public IntPtr GetStringTensorContent; public IntPtr CastTypeInfoToTensorInfo; public IntPtr GetOnnxTypeFromTypeInfo; public IntPtr CreateTensorTypeAndShapeInfo; public IntPtr SetTensorElementType; public IntPtr SetDimensions; public IntPtr GetTensorElementType; public IntPtr GetDimensionsCount; public IntPtr GetDimensions; public IntPtr GetSymbolicDimensions; public IntPtr GetTensorShapeElementCount; public IntPtr GetTensorTypeAndShape; public IntPtr GetTypeInfo; public IntPtr GetValueType; public IntPtr CreateMemoryInfo; public IntPtr CreateCpuMemoryInfo; public IntPtr CompareMemoryInfo; public IntPtr MemoryInfoGetName; public IntPtr MemoryInfoGetId; public IntPtr MemoryInfoGetMemType; public IntPtr MemoryInfoGetType; public IntPtr AllocatorAlloc; public IntPtr AllocatorFree; public IntPtr AllocatorGetInfo; public IntPtr GetAllocatorWithDefaultOptions; public IntPtr AddFreeDimensionOverride; public IntPtr GetValue; public IntPtr GetValueCount; public IntPtr CreateValue; public IntPtr CreateOpaqueValue; public IntPtr GetOpaqueValue; public IntPtr KernelInfoGetAttribute_float; public IntPtr KernelInfoGetAttribute_int64; public IntPtr KernelInfoGetAttribute_string; public IntPtr KernelContext_GetInputCount; public IntPtr KernelContext_GetOutputCount; public IntPtr KernelContext_GetInput; public IntPtr KernelContext_GetOutput; public IntPtr ReleaseEnv; public IntPtr ReleaseStatus; public IntPtr ReleaseMemoryInfo; public IntPtr ReleaseSession; public IntPtr ReleaseValue; public IntPtr ReleaseRunOptions; public IntPtr ReleaseTypeInfo; public IntPtr ReleaseTensorTypeAndShapeInfo; public IntPtr ReleaseSessionOptions; public IntPtr ReleaseCustomOpDomain; public IntPtr GetDenotationFromTypeInfo; public IntPtr CastTypeInfoToMapTypeInfo; public IntPtr CastTypeInfoToSequenceTypeInfo; public IntPtr GetMapKeyType; public IntPtr GetMapValueType; public IntPtr GetSequenceElementType; public IntPtr ReleaseMapTypeInfo; public IntPtr ReleaseSequenceTypeInfo; public IntPtr SessionEndProfiling; public IntPtr SessionGetModelMetadata; public IntPtr ModelMetadataGetProducerName; public IntPtr ModelMetadataGetGraphName; public IntPtr ModelMetadataGetDomain; public IntPtr ModelMetadataGetDescription; public IntPtr ModelMetadataLookupCustomMetadataMap; public IntPtr ModelMetadataGetVersion; public IntPtr ReleaseModelMetadata; public IntPtr CreateEnvWithGlobalThreadPools; public IntPtr DisablePerSessionThreads; public IntPtr CreateThreadingOptions; public IntPtr ReleaseThreadingOptions; public IntPtr ModelMetadataGetCustomMetadataMapKeys; public IntPtr AddFreeDimensionOverrideByName; public IntPtr GetAvailableProviders; public IntPtr ReleaseAvailableProviders; public IntPtr GetStringTensorElementLength; public IntPtr GetStringTensorElement; public IntPtr FillStringTensorElement; public IntPtr AddSessionConfigEntry; public IntPtr CreateAllocator; public IntPtr ReleaseAllocator; public IntPtr RunWithBinding; public IntPtr CreateIoBinding; public IntPtr ReleaseIoBinding; public IntPtr BindInput; public IntPtr BindOutput; public IntPtr BindOutputToDevice; public IntPtr GetBoundOutputNames; public IntPtr GetBoundOutputValues; public IntPtr ClearBoundInputs; public IntPtr ClearBoundOutputs; public IntPtr TensorAt; public IntPtr CreateAndRegisterAllocator; public IntPtr SetLanguageProjection; public IntPtr SessionGetProfilingStartTimeNs; public IntPtr SetGlobalIntraOpNumThreads; public IntPtr SetGlobalInterOpNumThreads; public IntPtr SetGlobalSpinControl; public IntPtr AddInitializer; public IntPtr CreateEnvWithCustomLoggerAndGlobalThreadPools; public IntPtr SessionOptionsAppendExecutionProvider_CUDA; public IntPtr SessionOptionsAppendExecutionProvider_ROCM; public IntPtr SessionOptionsAppendExecutionProvider_OpenVINO; public IntPtr SetGlobalDenormalAsZero; public IntPtr CreateArenaCfg; public IntPtr ReleaseArenaCfg; public IntPtr ModelMetadataGetGraphDescription; public IntPtr SessionOptionsAppendExecutionProvider_TensorRT; public IntPtr SetCurrentGpuDeviceId; public IntPtr GetCurrentGpuDeviceId; public IntPtr KernelInfoGetAttributeArray_float; public IntPtr KernelInfoGetAttributeArray_int64; public IntPtr CreateArenaCfgV2; public IntPtr AddRunConfigEntry; public IntPtr CreatePrepackedWeightsContainer; public IntPtr ReleasePrepackedWeightsContainer; public IntPtr CreateSessionWithPrepackedWeightsContainer; public IntPtr CreateSessionFromArrayWithPrepackedWeightsContainer; public IntPtr SessionOptionsAppendExecutionProvider_TensorRT_V2; public IntPtr CreateTensorRTProviderOptions; public IntPtr UpdateTensorRTProviderOptions; public IntPtr GetTensorRTProviderOptionsAsString; public IntPtr ReleaseTensorRTProviderOptions; public IntPtr EnableOrtCustomOps; public IntPtr RegisterAllocator; public IntPtr UnregisterAllocator; public IntPtr IsSparseTensor; public IntPtr CreateSparseTensorAsOrtValue; public IntPtr FillSparseTensorCoo; public IntPtr FillSparseTensorCsr; public IntPtr FillSparseTensorBlockSparse; public IntPtr CreateSparseTensorWithValuesAsOrtValue; public IntPtr UseCooIndices; public IntPtr UseCsrIndices; public IntPtr UseBlockSparseIndices; public IntPtr GetSparseTensorFormat; public IntPtr GetSparseTensorValuesTypeAndShape; public IntPtr GetSparseTensorValues; public IntPtr GetSparseTensorIndicesTypeShape; public IntPtr GetSparseTensorIndices; public IntPtr HasValue; public IntPtr KernelContext_GetGPUComputeStream; public IntPtr GetTensorMemoryInfo; public IntPtr GetExecutionProviderApi; public IntPtr SessionOptionsSetCustomCreateThreadFn; public IntPtr SessionOptionsSetCustomThreadCreationOptions; public IntPtr SessionOptionsSetCustomJoinThreadFn; public IntPtr SetGlobalCustomCreateThreadFn; public IntPtr SetGlobalCustomThreadCreationOptions; public IntPtr SetGlobalCustomJoinThreadFn; public IntPtr SynchronizeBoundInputs; public IntPtr SynchronizeBoundOutputs; public IntPtr SessionOptionsAppendExecutionProvider_CUDA_V2; public IntPtr CreateCUDAProviderOptions; public IntPtr UpdateCUDAProviderOptions; public IntPtr GetCUDAProviderOptionsAsString; public IntPtr ReleaseCUDAProviderOptions; public IntPtr SessionOptionsAppendExecutionProvider_MIGraphX; public IntPtr AddExternalInitializers; public IntPtr CreateOpAttr; public IntPtr ReleaseOpAttr; public IntPtr CreateOp; public IntPtr InvokeOp; public IntPtr ReleaseOp; public IntPtr SessionOptionsAppendExecutionProvider; public IntPtr CopyKernelInfo; public IntPtr ReleaseKernelInfo; public IntPtr GetTrainingApi; public IntPtr SessionOptionsAppendExecutionProvider_CANN; public IntPtr CreateCANNProviderOptions; public IntPtr UpdateCANNProviderOptions; public IntPtr GetCANNProviderOptionsAsString; public IntPtr ReleaseCANNProviderOptions; public IntPtr MemoryInfoGetDeviceType; public IntPtr UpdateEnvWithCustomLogLevel; public IntPtr SetGlobalIntraOpThreadAffinity; public IntPtr RegisterCustomOpsLibrary_V2; public IntPtr RegisterCustomOpsUsingFunction; public IntPtr KernelInfo_GetInputCount; public IntPtr KernelInfo_GetOutputCount; public IntPtr KernelInfo_GetInputName; public IntPtr KernelInfo_GetOutputName; public IntPtr KernelInfo_GetInputTypeInfo; public IntPtr KernelInfo_GetOutputTypeInfo; public IntPtr KernelInfoGetAttribute_tensor; public IntPtr HasSessionConfigEntry; public IntPtr GetSessionConfigEntry; public IntPtr SessionOptionsAppendExecutionProvider_Dnnl; public IntPtr CreateDnnlProviderOptions; public IntPtr UpdateDnnlProviderOptions; public IntPtr GetDnnlProviderOptionsAsString; public IntPtr ReleaseDnnlProviderOptions; public IntPtr KernelInfo_GetNodeName; public IntPtr KernelInfo_GetLogger; public IntPtr KernelContext_GetLogger; public IntPtr Logger_LogMessage; public IntPtr Logger_GetLoggingSeverityLevel; public IntPtr KernelInfoGetConstantInput_tensor; public IntPtr CastTypeInfoToOptionalTypeInfo; public IntPtr GetOptionalContainedTypeInfo; public IntPtr GetResizedStringTensorElementBuffer; public IntPtr KernelContext_GetAllocator; public IntPtr GetBuildInfoString; public IntPtr CreateROCMProviderOptions; public IntPtr UpdateROCMProviderOptions; public IntPtr GetROCMProviderOptionsAsString; public IntPtr ReleaseROCMProviderOptions; public IntPtr CreateAndRegisterAllocatorV2; public IntPtr RunAsync; } internal static class NativeMethods { [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate ref OrtApi DOrtGetApi(uint version); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetVersionString(); internal class NativeLib { internal const string DllName = "onnxruntime"; } [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateEnv(OrtLoggingLevel defaultLoggingLevel, byte[] logId, out IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateEnvWithCustomLogger(IntPtr loggingFunction, IntPtr loggerParam, OrtLoggingLevel defaultLoggingLevel, byte[] logId, out IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateEnvWithGlobalThreadPools(OrtLoggingLevel defaultWarningLevel, byte[] logId, IntPtr threadingOptions, out IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateEnvWithCustomLoggerAndGlobalThreadPools(IntPtr loggingFunction, IntPtr loggerParam, OrtLoggingLevel logSeverityLevel, byte[] logId, IntPtr threadingOptions, out IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseEnv(IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtEnableTelemetryEvents(IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtDisableTelemetryEvents(IntPtr env); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtUpdateEnvWithCustomLogLevel(IntPtr env, OrtLoggingLevel custom_log_level); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateTensorRTProviderOptions(out IntPtr trtProviderOptionsInstance); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtUpdateTensorRTProviderOptions(IntPtr trtProviderOptionsInstance, IntPtr[] providerOptionsKeys, IntPtr[] providerOptionsValues, UIntPtr numKeys); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTensorRTProviderOptionsAsString(IntPtr trtProviderOptionsInstance, IntPtr allocator, out IntPtr ptr); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseTensorRTProviderOptions(IntPtr trtProviderOptionsInstance); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateCUDAProviderOptions(out IntPtr cudaProviderOptionsInstance); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtUpdateCUDAProviderOptions(IntPtr cudaProviderOptionsInstance, IntPtr[] providerOptionsKeys, IntPtr[] providerOptionsValues, UIntPtr numKeys); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetCUDAProviderOptionsAsString(IntPtr cudaProviderOptionsInstance, IntPtr allocator, out IntPtr ptr); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseCUDAProviderOptions(IntPtr cudaProviderOptionsInstance); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateROCMProviderOptions(out IntPtr rocmProviderOptionsInstance); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtUpdateROCMProviderOptions(IntPtr rocmProviderOptionsInstance, IntPtr[] providerOptionsKeys, IntPtr[] providerOptionsValues, UIntPtr numKeys); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetROCMProviderOptionsAsString(IntPtr rocmProviderOptionsInstance, IntPtr allocator, out IntPtr ptr); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseROCMProviderOptions(IntPtr rocmProviderOptionsInstance); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate ErrorCode DOrtGetErrorCode(IntPtr status); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetErrorMessage(IntPtr status); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseStatus(IntPtr statusPtr); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateSession(IntPtr environment, byte[] modelPath, IntPtr sessopnOptions, out IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateSessionWithPrepackedWeightsContainer(IntPtr environment, byte[] modelPath, IntPtr sessionOptions, IntPtr prepackedWeightsContainer, out IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateSessionFromArray(IntPtr environment, byte[] modelData, UIntPtr modelSize, IntPtr sessionOptions, out IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateSessionFromArrayWithPrepackedWeightsContainer(IntPtr environment, byte[] modelData, UIntPtr modelSize, IntPtr sessionOptions, IntPtr prepackedWeightsContainer, out IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRun(IntPtr session, IntPtr runOptions, IntPtr[] inputNames, IntPtr[] inputValues, UIntPtr inputCount, IntPtr[] outputNames, UIntPtr outputCount, IntPtr[] outputValues); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunWithBinding(IntPtr session, IntPtr runOptions, IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetInputCount(IntPtr session, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetOutputCount(IntPtr session, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetOverridableInitializerCount(IntPtr session, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetInputName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetOutputName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionEndProfiling(IntPtr session, IntPtr allocator, out IntPtr profile_file); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetOverridableInitializerName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetInputTypeInfo(IntPtr session, UIntPtr index, out IntPtr typeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetOutputTypeInfo(IntPtr session, UIntPtr index, out IntPtr typeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetOverridableInitializerTypeInfo(IntPtr session, UIntPtr index, out IntPtr typeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseTypeInfo(IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseSession(IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetProfilingStartTimeNs(IntPtr session, out UIntPtr startTime); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DCreateAndRegisterAllocatorV2(IntPtr environment, IntPtr provider_type, IntPtr mem_info, IntPtr arena_cfg, IntPtr provider_options_keys, IntPtr provider_options_values, UIntPtr num_keys); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunAsync(IntPtr session, IntPtr runOptions, IntPtr[] inputNames, IntPtr[] inputValues, UIntPtr inputCount, IntPtr[] outputNames, UIntPtr outputCount, IntPtr[] outputValues, IntPtr callback, IntPtr user_data); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateSessionOptions(out IntPtr sessionOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseSessionOptions(IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCloneSessionOptions(IntPtr sessionOptions, out IntPtr output); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetSessionExecutionMode(IntPtr options, ExecutionMode execution_mode); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetOptimizedModelFilePath(IntPtr options, byte[] optimizedModelFilepath); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtEnableProfiling(IntPtr options, byte[] profilePathPrefix); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtDisableProfiling(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtEnableMemPattern(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtDisableMemPattern(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtEnableCpuMemArena(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtDisableCpuMemArena(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetSessionLogId(IntPtr options, byte[] logId); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetSessionLogVerbosityLevel(IntPtr options, int sessionLogVerbosityLevel); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetSessionLogSeverityLevel(IntPtr options, OrtLoggingLevel sessionLogSeverityLevel); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetIntraOpNumThreads(IntPtr options, int intraOpNumThreads); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetInterOpNumThreads(IntPtr options, int interOpNumThreads); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetSessionGraphOptimizationLevel(IntPtr options, GraphOptimizationLevel graphOptimizationLevel); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAddSessionConfigEntry(IntPtr options, byte[] configKey, byte[] configValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DSessionOptionsAppendExecutionProvider_TensorRT(IntPtr options, IntPtr trtProviderOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DSessionOptionsAppendExecutionProvider_TensorRT_V2(IntPtr options, IntPtr trtProviderOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DSessionOptionsAppendExecutionProvider_CUDA(IntPtr options, IntPtr cudaProviderOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DSessionOptionsAppendExecutionProvider_CUDA_V2(IntPtr options, IntPtr cudaProviderOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DSessionOptionsAppendExecutionProvider_ROCM(IntPtr options, IntPtr rocmProviderOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAddFreeDimensionOverride(IntPtr options, byte[] dimDenotation, long dimValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAddFreeDimensionOverrideByName(IntPtr options, byte[] dimName, long dimValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRegisterCustomOpsLibrary(IntPtr options, byte[] libraryPath, out IntPtr libraryHandle); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRegisterCustomOpsLibrary_V2(IntPtr options, byte[] libraryPath); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAddInitializer(IntPtr options, byte[] name, IntPtr ortValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DSessionOptionsAppendExecutionProvider(IntPtr options, byte[] providerName, IntPtr[] providerOptionsKeys, IntPtr[] providerOptionsValues, UIntPtr numKeys); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateRunOptions(out IntPtr runOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseRunOptions(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsSetRunLogVerbosityLevel(IntPtr options, int value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsSetRunLogSeverityLevel(IntPtr options, OrtLoggingLevel value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsSetRunTag(IntPtr options, byte[] runTag); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsGetRunLogVerbosityLevel(IntPtr options, out int verbosityLevel); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsGetRunLogSeverityLevel(IntPtr options, out OrtLoggingLevel severityLevel); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsGetRunTag(IntPtr options, out IntPtr runtag); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsSetTerminate(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRunOptionsUnsetTerminate(IntPtr options); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAddRunConfigEntry(IntPtr options, byte[] configKey, byte[] configValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateThreadingOptions(out IntPtr threadingOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtReleaseThreadingOptions(IntPtr threadingOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtThreadingOptionsSetGlobalInterOpNumThreads(IntPtr threadingOptions, int numThreads); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtThreadingOptionsSetGlobalIntraOpNumThreads(IntPtr threadingOptions, int numThreads); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtThreadingOptionsSetGlobalDenormalAsZero(IntPtr threadingOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtThreadingOptionsSetGlobalSpinControl(IntPtr threadingOptions, int allowSpinning); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateMemoryInfo(byte[] name, OrtAllocatorType allocatorType, int identifier, OrtMemType memType, out IntPtr allocatorInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateCpuMemoryInfo(OrtAllocatorType allocatorType, OrtMemType memoryType, out IntPtr allocatorInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseMemoryInfo(IntPtr allocatorInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCompareMemoryInfo(IntPtr info1, IntPtr info2, out int result); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtMemoryInfoGetName(IntPtr mem_info, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtMemoryInfoGetId(IntPtr mem_info, out int id); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtMemoryInfoGetMemType(IntPtr mem_info, out OrtMemType mem_type); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtMemoryInfoGetType(IntPtr mem_info, out OrtAllocatorType alloc_type); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetAllocatorWithDefaultOptions(out IntPtr allocator); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAllocatorGetInfo(IntPtr ptr, out IntPtr info); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateArenaCfg(UIntPtr maxMemory, int arenaExtendStrategy, int initialChunkSizeBytes, int maxDeadBytesPerChunk, out IntPtr arenaCfg); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseArenaCfg(IntPtr arenaCfg); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateAllocator(IntPtr session, IntPtr info, out IntPtr allocator); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseAllocator(IntPtr allocator); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAllocatorAlloc(IntPtr allocator, UIntPtr size, out IntPtr p); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAllocatorFree(IntPtr allocator, IntPtr p); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateIoBinding(IntPtr session, out IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseIoBinding(IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtBindInput(IntPtr io_binding, byte[] name, IntPtr ort_value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSynchronizeBoundInputs(IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtBindOutput(IntPtr io_binding, byte[] name, IntPtr ort_value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtBindOutputToDevice(IntPtr io_binding, byte[] name, IntPtr mem_info); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSynchronizeBoundOutputs(IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetBoundOutputNames(IntPtr io_binding, IntPtr allocator, out IntPtr buffer, out IntPtr lengths, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetBoundOutputValues(IntPtr io_binding, IntPtr allocator, out IntPtr ortvalues, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtClearBoundInputs(IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtClearBoundOutputs(IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtTensorAt(IntPtr io_binding); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateAndRegisterAllocator(IntPtr env, IntPtr memInfo, IntPtr arenaCfg); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetLanguageProjection(IntPtr environment, int projection); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSessionGetModelMetadata(IntPtr session, out IntPtr modelMetadata); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetProducerName(IntPtr modelMetadata, IntPtr allocator, out IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetGraphName(IntPtr modelMetadata, IntPtr allocator, out IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetDomain(IntPtr modelMetadata, IntPtr allocator, out IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetDescription(IntPtr modelMetadata, IntPtr allocator, out IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetGraphDescription(IntPtr modelMetadata, IntPtr allocator, out IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetVersion(IntPtr modelMetadata, out long value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataGetCustomMetadataMapKeys(IntPtr modelMetadata, IntPtr allocator, out IntPtr keys, out long numKeys); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtModelMetadataLookupCustomMetadataMap(IntPtr modelMetadata, IntPtr allocator, IntPtr key, out IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseModelMetadata(IntPtr modelMetadata); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtHasValue(IntPtr value, out IntPtr hasValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetValue(IntPtr value, int index, IntPtr allocator, out IntPtr outputValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetValueType(IntPtr value, out IntPtr onnxtype); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetOnnxTypeFromTypeInfo(IntPtr typeinfo, out IntPtr onnxtype); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetValueCount(IntPtr value, out IntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateValue(IntPtr[] values, UIntPtr num_values, IntPtr onnxValueType, out IntPtr ortValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTypeInfo(IntPtr value, out IntPtr typeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateTensorAsOrtValue(IntPtr allocator, long[] shape, UIntPtr shape_len, TensorElementType type, out IntPtr outputValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateTensorWithDataAsOrtValue(IntPtr allocatorInfo, IntPtr dataBufferHandle, UIntPtr dataLength, long[] shape, UIntPtr shapeLength, TensorElementType type, out IntPtr outputValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtValueIsTensor(IntPtr ortValue, out IntPtr val); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtValueIsSparseTensor(IntPtr ortValue, out IntPtr val); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTensorMutableData(IntPtr value, out IntPtr dataBufferHandle); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtFillStringTensor(IntPtr value, IntPtr[] s, UIntPtr s_len); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetResizedStringTensorElementBuffer(IntPtr value, UIntPtr index, UIntPtr length_in_bytes, out IntPtr buffer); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetStringTensorContent(IntPtr value, byte[] dst_buffer, UIntPtr dst_buffer_len, UIntPtr[] offsets, UIntPtr offsets_len); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetStringTensorDataLength(IntPtr value, out UIntPtr len); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetStringTensorElementLength(IntPtr value, UIntPtr index, out UIntPtr len); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetStringTensorElement(IntPtr value, UIntPtr bufferLength, UIntPtr elementIndex, byte[] buffer); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCastTypeInfoToTensorInfo(IntPtr typeInfo, out IntPtr typeAndShapeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTensorTypeAndShape(IntPtr value, out IntPtr typeAndShapeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseTensorTypeAndShapeInfo(IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTensorElementType(IntPtr typeAndShapeInfo, out IntPtr output); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetDimensionsCount(IntPtr typeAndShapeInfo, out UIntPtr output); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetDimensions(IntPtr typeAndShapeInfo, long[] dim_values, UIntPtr dim_values_length); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetSymbolicDimensions(IntPtr typeAndShapeInfo, IntPtr[] dim_params, UIntPtr dim_params_length); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTensorShapeElementCount(IntPtr typeAndShapeInfo, out UIntPtr output); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTensorMemoryInfo(IntPtr ortValue, out IntPtr ortMemoryInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DCastTypeInfoToMapTypeInfo(IntPtr typeInfo, out IntPtr mapTypeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DGetMapKeyType(IntPtr mapTypeInfo, out IntPtr tensorElementType); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DGetMapValueType(IntPtr map_type_info, out IntPtr type_info); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DCastTypeInfoToSequenceTypeInfo(IntPtr typeInfo, out IntPtr sequenceTypeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DGetSequenceElementType(IntPtr sequenceTypeInfo, out IntPtr elementTypeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCastTypeInfoToOptionalTypeInfo(IntPtr typeInfo, out IntPtr optionalTypeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DGetOptionalContainedTypeInfo(IntPtr optTypeInfo, out IntPtr containedTypeInfo); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseValue(IntPtr value); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetAvailableProviders(out IntPtr providers, out int numProviders); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtReleaseAvailableProviders(IntPtr providers, int numProviders); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreatePrepackedWeightsContainer(out IntPtr prepackedWeightsContainer); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleasePrepackedWeightsContainer(IntPtr prepackedWeightsContainer); private static OrtApi api_; public static DOrtGetVersionString OrtGetVersionString; public static DOrtCreateEnv OrtCreateEnv; public static DOrtCreateEnvWithCustomLogger OrtCreateEnvWithCustomLogger; public static DOrtCreateEnvWithGlobalThreadPools OrtCreateEnvWithGlobalThreadPools; public static DOrtCreateEnvWithCustomLoggerAndGlobalThreadPools OrtCreateEnvWithCustomLoggerAndGlobalThreadPools; public static DOrtReleaseEnv OrtReleaseEnv; public static DOrtEnableTelemetryEvents OrtEnableTelemetryEvents; public static DOrtDisableTelemetryEvents OrtDisableTelemetryEvents; public static DOrtUpdateEnvWithCustomLogLevel OrtUpdateEnvWithCustomLogLevel; public static DOrtCreateTensorRTProviderOptions OrtCreateTensorRTProviderOptions; public static DOrtUpdateTensorRTProviderOptions OrtUpdateTensorRTProviderOptions; public static DOrtGetTensorRTProviderOptionsAsString OrtGetTensorRTProviderOptionsAsString; public static DOrtReleaseTensorRTProviderOptions OrtReleaseTensorRTProviderOptions; public static DOrtCreateCUDAProviderOptions OrtCreateCUDAProviderOptions; public static DOrtUpdateCUDAProviderOptions OrtUpdateCUDAProviderOptions; public static DOrtGetCUDAProviderOptionsAsString OrtGetCUDAProviderOptionsAsString; public static DOrtReleaseCUDAProviderOptions OrtReleaseCUDAProviderOptions; public static DOrtCreateROCMProviderOptions OrtCreateROCMProviderOptions; public static DOrtUpdateROCMProviderOptions OrtUpdateROCMProviderOptions; public static DOrtGetROCMProviderOptionsAsString OrtGetROCMProviderOptionsAsString; public static DOrtReleaseROCMProviderOptions OrtReleaseROCMProviderOptions; public static DOrtGetErrorCode OrtGetErrorCode; public static DOrtGetErrorMessage OrtGetErrorMessage; public static DOrtReleaseStatus OrtReleaseStatus; public static DOrtCreateSession OrtCreateSession; public static DOrtCreateSessionWithPrepackedWeightsContainer OrtCreateSessionWithPrepackedWeightsContainer; public static DOrtCreateSessionFromArray OrtCreateSessionFromArray; public static DOrtCreateSessionFromArrayWithPrepackedWeightsContainer OrtCreateSessionFromArrayWithPrepackedWeightsContainer; public static DOrtRun OrtRun; public static DOrtRunWithBinding OrtRunWithBinding; public static DOrtSessionGetInputCount OrtSessionGetInputCount; public static DOrtSessionGetOutputCount OrtSessionGetOutputCount; public static DOrtSessionGetOverridableInitializerCount OrtSessionGetOverridableInitializerCount; public static DOrtSessionGetInputName OrtSessionGetInputName; public static DOrtSessionGetOutputName OrtSessionGetOutputName; public static DOrtSessionEndProfiling OrtSessionEndProfiling; public static DOrtSessionGetOverridableInitializerName OrtSessionGetOverridableInitializerName; public static DOrtSessionGetInputTypeInfo OrtSessionGetInputTypeInfo; public static DOrtSessionGetOutputTypeInfo OrtSessionGetOutputTypeInfo; public static DOrtSessionGetOverridableInitializerTypeInfo OrtSessionGetOverridableInitializerTypeInfo; public static DOrtReleaseTypeInfo OrtReleaseTypeInfo; public static DOrtReleaseSession OrtReleaseSession; public static DOrtSessionGetProfilingStartTimeNs OrtSessionGetProfilingStartTimeNs; public static DCreateAndRegisterAllocatorV2 OrtCreateAndRegisterAllocatorV2; public static DOrtRunAsync OrtRunAsync; public static DOrtCreateSessionOptions OrtCreateSessionOptions; public static DOrtReleaseSessionOptions OrtReleaseSessionOptions; public static DOrtCloneSessionOptions OrtCloneSessionOptions; public static DOrtSetSessionExecutionMode OrtSetSessionExecutionMode; public static DOrtSetOptimizedModelFilePath OrtSetOptimizedModelFilePath; public static DOrtEnableProfiling OrtEnableProfiling; public static DOrtDisableProfiling OrtDisableProfiling; public static DOrtEnableMemPattern OrtEnableMemPattern; public static DOrtDisableMemPattern OrtDisableMemPattern; public static DOrtEnableCpuMemArena OrtEnableCpuMemArena; public static DOrtDisableCpuMemArena OrtDisableCpuMemArena; public static DOrtSetSessionLogId OrtSetSessionLogId; public static DOrtSetSessionLogVerbosityLevel OrtSetSessionLogVerbosityLevel; public static DOrtSetSessionLogSeverityLevel OrtSetSessionLogSeverityLevel; public static DOrtSetIntraOpNumThreads OrtSetIntraOpNumThreads; public static DOrtSetInterOpNumThreads OrtSetInterOpNumThreads; public static DOrtSetSessionGraphOptimizationLevel OrtSetSessionGraphOptimizationLevel; public static DOrtAddSessionConfigEntry OrtAddSessionConfigEntry; public static DSessionOptionsAppendExecutionProvider_TensorRT SessionOptionsAppendExecutionProvider_TensorRT; public static DSessionOptionsAppendExecutionProvider_TensorRT_V2 SessionOptionsAppendExecutionProvider_TensorRT_V2; public static DSessionOptionsAppendExecutionProvider_CUDA SessionOptionsAppendExecutionProvider_CUDA; public static DSessionOptionsAppendExecutionProvider_CUDA_V2 SessionOptionsAppendExecutionProvider_CUDA_V2; public static DSessionOptionsAppendExecutionProvider_ROCM SessionOptionsAppendExecutionProvider_ROCM; public static DOrtAddFreeDimensionOverride OrtAddFreeDimensionOverride; public static DOrtAddFreeDimensionOverrideByName OrtAddFreeDimensionOverrideByName; public static DOrtRegisterCustomOpsLibrary OrtRegisterCustomOpsLibrary; public static DOrtRegisterCustomOpsLibrary_V2 OrtRegisterCustomOpsLibrary_V2; public static DOrtAddInitializer OrtAddInitializer; public static DSessionOptionsAppendExecutionProvider SessionOptionsAppendExecutionProvider; public static DOrtCreateRunOptions OrtCreateRunOptions; public static DOrtReleaseRunOptions OrtReleaseRunOptions; public static DOrtRunOptionsSetRunLogVerbosityLevel OrtRunOptionsSetRunLogVerbosityLevel; public static DOrtRunOptionsSetRunLogSeverityLevel OrtRunOptionsSetRunLogSeverityLevel; public static DOrtRunOptionsSetRunTag OrtRunOptionsSetRunTag; public static DOrtRunOptionsGetRunLogVerbosityLevel OrtRunOptionsGetRunLogVerbosityLevel; public static DOrtRunOptionsGetRunLogSeverityLevel OrtRunOptionsGetRunLogSeverityLevel; public static DOrtRunOptionsGetRunTag OrtRunOptionsGetRunTag; public static DOrtRunOptionsSetTerminate OrtRunOptionsSetTerminate; public static DOrtRunOptionsUnsetTerminate OrtRunOptionsUnsetTerminate; public static DOrtAddRunConfigEntry OrtAddRunConfigEntry; public static DOrtCreateThreadingOptions OrtCreateThreadingOptions; public static DOrtReleaseThreadingOptions OrtReleaseThreadingOptions; public static DOrtThreadingOptionsSetGlobalInterOpNumThreads OrtThreadingOptionsSetGlobalInterOpNumThreads; public static DOrtThreadingOptionsSetGlobalIntraOpNumThreads OrtThreadingOptionsSetGlobalIntraOpNumThreads; public static DOrtThreadingOptionsSetGlobalDenormalAsZero OrtThreadingOptionsSetGlobalDenormalAsZero; public static DOrtThreadingOptionsSetGlobalSpinControl OrtThreadingOptionsSetGlobalSpinControl; public static DOrtCreateMemoryInfo OrtCreateMemoryInfo; public static DOrtCreateCpuMemoryInfo OrtCreateCpuMemoryInfo; public static DOrtReleaseMemoryInfo OrtReleaseMemoryInfo; public static DOrtCompareMemoryInfo OrtCompareMemoryInfo; public static DOrtMemoryInfoGetName OrtMemoryInfoGetName; public static DOrtMemoryInfoGetId OrtMemoryInfoGetId; public static DOrtMemoryInfoGetMemType OrtMemoryInfoGetMemType; public static DOrtMemoryInfoGetType OrtMemoryInfoGetType; public static DOrtGetAllocatorWithDefaultOptions OrtGetAllocatorWithDefaultOptions; public static DOrtAllocatorGetInfo OrtAllocatorGetInfo; public static DOrtCreateArenaCfg OrtCreateArenaCfg; public static DOrtReleaseArenaCfg OrtReleaseArenaCfg; public static DOrtCreateAllocator OrtCreateAllocator; public static DOrtReleaseAllocator OrtReleaseAllocator; public static DOrtAllocatorAlloc OrtAllocatorAlloc; public static DOrtAllocatorFree OrtAllocatorFree; public static DOrtCreateIoBinding OrtCreateIoBinding; public static DOrtReleaseIoBinding OrtReleaseIoBinding; public static DOrtBindInput OrtBindInput; public static DOrtSynchronizeBoundInputs OrtSynchronizeBoundInputs; public static DOrtBindOutput OrtBindOutput; public static DOrtBindOutputToDevice OrtBindOutputToDevice; public static DOrtSynchronizeBoundOutputs OrtSynchronizeBoundOutputs; public static DOrtGetBoundOutputNames OrtGetBoundOutputNames; public static DOrtGetBoundOutputValues OrtGetBoundOutputValues; public static DOrtClearBoundInputs OrtClearBoundInputs; public static DOrtClearBoundOutputs OrtClearBoundOutputs; public static DOrtTensorAt OrtTensorAt; public static DOrtCreateAndRegisterAllocator OrtCreateAndRegisterAllocator; public static DOrtSetLanguageProjection OrtSetLanguageProjection; public static DOrtSessionGetModelMetadata OrtSessionGetModelMetadata; public static DOrtModelMetadataGetProducerName OrtModelMetadataGetProducerName; public static DOrtModelMetadataGetGraphName OrtModelMetadataGetGraphName; public static DOrtModelMetadataGetDomain OrtModelMetadataGetDomain; public static DOrtModelMetadataGetDescription OrtModelMetadataGetDescription; public static DOrtModelMetadataGetGraphDescription OrtModelMetadataGetGraphDescription; public static DOrtModelMetadataGetVersion OrtModelMetadataGetVersion; public static DOrtModelMetadataGetCustomMetadataMapKeys OrtModelMetadataGetCustomMetadataMapKeys; public static DOrtModelMetadataLookupCustomMetadataMap OrtModelMetadataLookupCustomMetadataMap; public static DOrtReleaseModelMetadata OrtReleaseModelMetadata; public static DOrtHasValue OrtHasValue; public static DOrtGetValue OrtGetValue; public static DOrtGetValueType OrtGetValueType; public static DOrtGetOnnxTypeFromTypeInfo OrtGetOnnxTypeFromTypeInfo; public static DOrtGetValueCount OrtGetValueCount; public static DOrtCreateValue OrtCreateValue; public static DOrtGetTypeInfo OrtGetTypeInfo; public static DOrtCreateTensorAsOrtValue OrtCreateTensorAsOrtValue; public static DOrtCreateTensorWithDataAsOrtValue OrtCreateTensorWithDataAsOrtValue; public static DOrtValueIsTensor OrtValueIsTensor; public static DOrtValueIsSparseTensor OrtValueIsSparseTensor; public static DOrtGetTensorMutableData OrtGetTensorMutableData; public static DOrtFillStringTensor OrtFillStringTensor; public static DOrtGetResizedStringTensorElementBuffer OrtGetResizedStringTensorElementBuffer; public static DOrtGetStringTensorContent OrtGetStringTensorContent; public static DOrtGetStringTensorDataLength OrtGetStringTensorDataLength; public static DOrtGetStringTensorElementLength OrtGetStringTensorElementLength; public static DOrtGetStringTensorElement OrtGetStringTensorElement; public static DOrtCastTypeInfoToTensorInfo OrtCastTypeInfoToTensorInfo; public static DOrtGetTensorTypeAndShape OrtGetTensorTypeAndShape; public static DOrtReleaseTensorTypeAndShapeInfo OrtReleaseTensorTypeAndShapeInfo; public static DOrtGetTensorElementType OrtGetTensorElementType; public static DOrtGetDimensionsCount OrtGetDimensionsCount; public static DOrtGetDimensions OrtGetDimensions; public static DOrtGetSymbolicDimensions OrtGetSymbolicDimensions; public static DOrtGetTensorShapeElementCount OrtGetTensorShapeElementCount; public static DOrtGetTensorMemoryInfo OrtGetTensorMemoryInfo; public static DCastTypeInfoToMapTypeInfo OrtCastTypeInfoToMapTypeInfo; public static DGetMapKeyType OrtGetMapKeyType; public static DGetMapValueType OrtGetMapValueType; public static DCastTypeInfoToSequenceTypeInfo OrtCastTypeInfoToSequenceTypeInfo; public static DGetSequenceElementType OrtGetSequenceElementType; public static DOrtCastTypeInfoToOptionalTypeInfo OrtCastTypeInfoToOptionalTypeInfo; public static DGetOptionalContainedTypeInfo OrtGetOptionalContainedTypeInfo; public static DOrtReleaseValue OrtReleaseValue; public static DOrtGetAvailableProviders OrtGetAvailableProviders; public static DOrtReleaseAvailableProviders OrtReleaseAvailableProviders; public static DOrtCreatePrepackedWeightsContainer OrtCreatePrepackedWeightsContainer; public static DOrtReleasePrepackedWeightsContainer OrtReleasePrepackedWeightsContainer; static NativeMethods() { DOrtGetApi dOrtGetApi = (DOrtGetApi)Marshal.GetDelegateForFunctionPointer(OrtGetApiBase().GetApi, typeof(DOrtGetApi)); api_ = dOrtGetApi(14u); OrtGetVersionString = (DOrtGetVersionString)Marshal.GetDelegateForFunctionPointer(OrtGetApiBase().GetVersionString, typeof(DOrtGetVersionString)); OrtCreateEnv = (DOrtCreateEnv)Marshal.GetDelegateForFunctionPointer(api_.CreateEnv, typeof(DOrtCreateEnv)); OrtCreateEnvWithCustomLogger = (DOrtCreateEnvWithCustomLogger)Marshal.GetDelegateForFunctionPointer(api_.CreateEnvWithCustomLogger, typeof(DOrtCreateEnvWithCustomLogger)); OrtCreateEnvWithGlobalThreadPools = (DOrtCreateEnvWithGlobalThreadPools)Marshal.GetDelegateForFunctionPointer(api_.CreateEnvWithGlobalThreadPools, typeof(DOrtCreateEnvWithGlobalThreadPools)); OrtCreateEnvWithCustomLoggerAndGlobalThreadPools = (DOrtCreateEnvWithCustomLoggerAndGlobalThreadPools)Marshal.GetDelegateForFunctionPointer(api_.CreateEnvWithCustomLoggerAndGlobalThreadPools, typeof(DOrtCreateEnvWithCustomLoggerAndGlobalThreadPools)); OrtReleaseEnv = (DOrtReleaseEnv)Marshal.GetDelegateForFunctionPointer(api_.ReleaseEnv, typeof(DOrtReleaseEnv)); OrtEnableTelemetryEvents = (DOrtEnableTelemetryEvents)Marshal.GetDelegateForFunctionPointer(api_.EnableTelemetryEvents, typeof(DOrtEnableTelemetryEvents)); OrtDisableTelemetryEvents = (DOrtDisableTelemetryEvents)Marshal.GetDelegateForFunctionPointer(api_.DisableTelemetryEvents, typeof(DOrtDisableTelemetryEvents)); OrtGetErrorCode = (DOrtGetErrorCode)Marshal.GetDelegateForFunctionPointer(api_.GetErrorCode, typeof(DOrtGetErrorCode)); OrtGetErrorMessage = (DOrtGetErrorMessage)Marshal.GetDelegateForFunctionPointer(api_.GetErrorMessage, typeof(DOrtGetErrorMessage)); OrtReleaseStatus = (DOrtReleaseStatus)Marshal.GetDelegateForFunctionPointer(api_.ReleaseStatus, typeof(DOrtReleaseStatus)); OrtCreateSession = (DOrtCreateSession)Marshal.GetDelegateForFunctionPointer(api_.CreateSession, typeof(DOrtCreateSession)); OrtCreateSessionWithPrepackedWeightsContainer = (DOrtCreateSessionWithPrepackedWeightsContainer)Marshal.GetDelegateForFunctionPointer(api_.CreateSessionWithPrepackedWeightsContainer, typeof(DOrtCreateSessionWithPrepackedWeightsContainer)); OrtCreateSessionFromArray = (DOrtCreateSessionFromArray)Marshal.GetDelegateForFunctionPointer(api_.CreateSessionFromArray, typeof(DOrtCreateSessionFromArray)); OrtCreateSessionFromArrayWithPrepackedWeightsContainer = (DOrtCreateSessionFromArrayWithPrepackedWeightsContainer)Marshal.GetDelegateForFunctionPointer(api_.CreateSessionFromArrayWithPrepackedWeightsContainer, typeof(DOrtCreateSessionFromArrayWithPrepackedWeightsContainer)); OrtRun = (DOrtRun)Marshal.GetDelegateForFunctionPointer(api_.Run, typeof(DOrtRun)); OrtRunWithBinding = (DOrtRunWithBinding)Marshal.GetDelegateForFunctionPointer(api_.RunWithBinding, typeof(DOrtRunWithBinding)); OrtSessionGetInputCount = (DOrtSessionGetInputCount)Marshal.GetDelegateForFunctionPointer(api_.SessionGetInputCount, typeof(DOrtSessionGetInputCount)); OrtSessionGetOutputCount = (DOrtSessionGetOutputCount)Marshal.GetDelegateForFunctionPointer(api_.SessionGetOutputCount, typeof(DOrtSessionGetOutputCount)); OrtSessionGetOverridableInitializerCount = (DOrtSessionGetOverridableInitializerCount)Marshal.GetDelegateForFunctionPointer(api_.SessionGetOverridableInitializerCount, typeof(DOrtSessionGetOverridableInitializerCount)); OrtSessionGetInputName = (DOrtSessionGetInputName)Marshal.GetDelegateForFunctionPointer(api_.SessionGetInputName, typeof(DOrtSessionGetInputName)); OrtSessionGetOutputName = (DOrtSessionGetOutputName)Marshal.GetDelegateForFunctionPointer(api_.SessionGetOutputName, typeof(DOrtSessionGetOutputName)); OrtSessionEndProfiling = (DOrtSessionEndProfiling)Marshal.GetDelegateForFunctionPointer(api_.SessionEndProfiling, typeof(DOrtSessionEndProfiling)); OrtSessionGetOverridableInitializerName = (DOrtSessionGetOverridableInitializerName)Marshal.GetDelegateForFunctionPointer(api_.SessionGetOverridableInitializerName, typeof(DOrtSessionGetOverridableInitializerName)); OrtSessionGetInputTypeInfo = (DOrtSessionGetInputTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.SessionGetInputTypeInfo, typeof(DOrtSessionGetInputTypeInfo)); OrtSessionGetOutputTypeInfo = (DOrtSessionGetOutputTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.SessionGetOutputTypeInfo, typeof(DOrtSessionGetOutputTypeInfo)); OrtSessionGetOverridableInitializerTypeInfo = (DOrtSessionGetOverridableInitializerTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.SessionGetOverridableInitializerTypeInfo, typeof(DOrtSessionGetOverridableInitializerTypeInfo)); OrtReleaseTypeInfo = (DOrtReleaseTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.ReleaseTypeInfo, typeof(DOrtReleaseTypeInfo)); OrtReleaseSession = (DOrtReleaseSession)Marshal.GetDelegateForFunctionPointer(api_.ReleaseSession, typeof(DOrtReleaseSession)); OrtSessionGetProfilingStartTimeNs = (DOrtSessionGetProfilingStartTimeNs)Marshal.GetDelegateForFunctionPointer(api_.SessionGetProfilingStartTimeNs, typeof(DOrtSessionGetProfilingStartTimeNs)); OrtCreateSessionOptions = (DOrtCreateSessionOptions)Marshal.GetDelegateForFunctionPointer(api_.CreateSessionOptions, typeof(DOrtCreateSessionOptions)); OrtReleaseSessionOptions = (DOrtReleaseSessionOptions)Marshal.GetDelegateForFunctionPointer(api_.ReleaseSessionOptions, typeof(DOrtReleaseSessionOptions)); OrtCloneSessionOptions = (DOrtCloneSessionOptions)Marshal.GetDelegateForFunctionPointer(api_.CloneSessionOptions, typeof(DOrtCloneSessionOptions)); OrtSetSessionExecutionMode = (DOrtSetSessionExecutionMode)Marshal.GetDelegateForFunctionPointer(api_.SetSessionExecutionMode, typeof(DOrtSetSessionExecutionMode)); OrtSetOptimizedModelFilePath = (DOrtSetOptimizedModelFilePath)Marshal.GetDelegateForFunctionPointer(api_.SetOptimizedModelFilePath, typeof(DOrtSetOptimizedModelFilePath)); OrtEnableProfiling = (DOrtEnableProfiling)Marshal.GetDelegateForFunctionPointer(api_.EnableProfiling, typeof(DOrtEnableProfiling)); OrtDisableProfiling = (DOrtDisableProfiling)Marshal.GetDelegateForFunctionPointer(api_.DisableProfiling, typeof(DOrtDisableProfiling)); OrtEnableMemPattern = (DOrtEnableMemPattern)Marshal.GetDelegateForFunctionPointer(api_.EnableMemPattern, typeof(DOrtEnableMemPattern)); OrtDisableMemPattern = (DOrtDisableMemPattern)Marshal.GetDelegateForFunctionPointer(api_.DisableMemPattern, typeof(DOrtDisableMemPattern)); OrtEnableCpuMemArena = (DOrtEnableCpuMemArena)Marshal.GetDelegateForFunctionPointer(api_.EnableCpuMemArena, typeof(DOrtEnableCpuMemArena)); OrtDisableCpuMemArena = (DOrtDisableCpuMemArena)Marshal.GetDelegateForFunctionPointer(api_.DisableCpuMemArena, typeof(DOrtDisableCpuMemArena)); OrtSetSessionLogId = (DOrtSetSessionLogId)Marshal.GetDelegateForFunctionPointer(api_.SetSessionLogId, typeof(DOrtSetSessionLogId)); OrtSetSessionLogVerbosityLevel = (DOrtSetSessionLogVerbosityLevel)Marshal.GetDelegateForFunctionPointer(api_.SetSessionLogVerbosityLevel, typeof(DOrtSetSessionLogVerbosityLevel)); OrtSetSessionLogSeverityLevel = (DOrtSetSessionLogSeverityLevel)Marshal.GetDelegateForFunctionPointer(api_.SetSessionLogSeverityLevel, typeof(DOrtSetSessionLogSeverityLevel)); OrtSetInterOpNumThreads = (DOrtSetInterOpNumThreads)Marshal.GetDelegateForFunctionPointer(api_.SetInterOpNumThreads, typeof(DOrtSetInterOpNumThreads)); OrtSetIntraOpNumThreads = (DOrtSetIntraOpNumThreads)Marshal.GetDelegateForFunctionPointer(api_.SetIntraOpNumThreads, typeof(DOrtSetIntraOpNumThreads)); OrtSetSessionGraphOptimizationLevel = (DOrtSetSessionGraphOptimizationLevel)Marshal.GetDelegateForFunctionPointer(api_.SetSessionGraphOptimizationLevel, typeof(DOrtSetSessionGraphOptimizationLevel)); OrtRegisterCustomOpsLibrary = (DOrtRegisterCustomOpsLibrary)Marshal.GetDelegateForFunctionPointer(api_.RegisterCustomOpsLibrary, typeof(DOrtRegisterCustomOpsLibrary)); OrtRegisterCustomOpsLibrary_V2 = (DOrtRegisterCustomOpsLibrary_V2)Marshal.GetDelegateForFunctionPointer(api_.RegisterCustomOpsLibrary_V2, typeof(DOrtRegisterCustomOpsLibrary_V2)); OrtAddSessionConfigEntry = (DOrtAddSessionConfigEntry)Marshal.GetDelegateForFunctionPointer(api_.AddSessionConfigEntry, typeof(DOrtAddSessionConfigEntry)); OrtAddInitializer = (DOrtAddInitializer)Marshal.GetDelegateForFunctionPointer(api_.AddInitializer, typeof(DOrtAddInitializer)); SessionOptionsAppendExecutionProvider_TensorRT = (DSessionOptionsAppendExecutionProvider_TensorRT)Marshal.GetDelegateForFunctionPointer(api_.SessionOptionsAppendExecutionProvider_TensorRT, typeof(DSessionOptionsAppendExecutionProvider_TensorRT)); OrtCreateRunOptions = (DOrtCreateRunOptions)Marshal.GetDelegateForFunctionPointer(api_.CreateRunOptions, typeof(DOrtCreateRunOptions)); OrtReleaseRunOptions = (DOrtReleaseRunOptions)Marshal.GetDelegateForFunctionPointer(api_.ReleaseRunOptions, typeof(DOrtReleaseRunOptions)); OrtRunOptionsSetRunLogVerbosityLevel = (DOrtRunOptionsSetRunLogVerbosityLevel)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsSetRunLogVerbosityLevel, typeof(DOrtRunOptionsSetRunLogVerbosityLevel)); OrtRunOptionsSetRunLogSeverityLevel = (DOrtRunOptionsSetRunLogSeverityLevel)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsSetRunLogSeverityLevel, typeof(DOrtRunOptionsSetRunLogSeverityLevel)); OrtRunOptionsSetRunTag = (DOrtRunOptionsSetRunTag)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsSetRunTag, typeof(DOrtRunOptionsSetRunTag)); OrtRunOptionsGetRunLogVerbosityLevel = (DOrtRunOptionsGetRunLogVerbosityLevel)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsGetRunLogVerbosityLevel, typeof(DOrtRunOptionsGetRunLogVerbosityLevel)); OrtRunOptionsGetRunLogSeverityLevel = (DOrtRunOptionsGetRunLogSeverityLevel)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsGetRunLogSeverityLevel, typeof(DOrtRunOptionsGetRunLogSeverityLevel)); OrtRunOptionsGetRunTag = (DOrtRunOptionsGetRunTag)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsGetRunTag, typeof(DOrtRunOptionsGetRunTag)); OrtRunOptionsSetTerminate = (DOrtRunOptionsSetTerminate)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsSetTerminate, typeof(DOrtRunOptionsSetTerminate)); OrtRunOptionsUnsetTerminate = (DOrtRunOptionsUnsetTerminate)Marshal.GetDelegateForFunctionPointer(api_.RunOptionsUnsetTerminate, typeof(DOrtRunOptionsUnsetTerminate)); OrtCreateThreadingOptions = (DOrtCreateThreadingOptions)Marshal.GetDelegateForFunctionPointer(api_.CreateThreadingOptions, typeof(DOrtCreateThreadingOptions)); OrtReleaseThreadingOptions = (DOrtReleaseThreadingOptions)Marshal.GetDelegateForFunctionPointer(api_.ReleaseThreadingOptions, typeof(DOrtReleaseThreadingOptions)); OrtThreadingOptionsSetGlobalInterOpNumThreads = (DOrtThreadingOptionsSetGlobalInterOpNumThreads)Marshal.GetDelegateForFunctionPointer(api_.SetGlobalInterOpNumThreads, typeof(DOrtThreadingOptionsSetGlobalInterOpNumThreads)); OrtThreadingOptionsSetGlobalIntraOpNumThreads = (DOrtThreadingOptionsSetGlobalIntraOpNumThreads)Marshal.GetDelegateForFunctionPointer(api_.SetGlobalIntraOpNumThreads, typeof(DOrtThreadingOptionsSetGlobalIntraOpNumThreads)); OrtThreadingOptionsSetGlobalDenormalAsZero = (DOrtThreadingOptionsSetGlobalDenormalAsZero)Marshal.GetDelegateForFunctionPointer(api_.SetGlobalDenormalAsZero, typeof(DOrtThreadingOptionsSetGlobalDenormalAsZero)); OrtThreadingOptionsSetGlobalSpinControl = (DOrtThreadingOptionsSetGlobalSpinControl)Marshal.GetDelegateForFunctionPointer(api_.SetGlobalSpinControl, typeof(DOrtThreadingOptionsSetGlobalSpinControl)); OrtAddRunConfigEntry = (DOrtAddRunConfigEntry)Marshal.GetDelegateForFunctionPointer(api_.AddRunConfigEntry, typeof(DOrtAddRunConfigEntry)); OrtCreateArenaCfg = (DOrtCreateArenaCfg)Marshal.GetDelegateForFunctionPointer(api_.CreateArenaCfg, typeof(DOrtCreateArenaCfg)); OrtReleaseArenaCfg = (DOrtReleaseArenaCfg)Marshal.GetDelegateForFunctionPointer(api_.ReleaseArenaCfg, typeof(DOrtReleaseArenaCfg)); OrtReleaseAllocator = (DOrtReleaseAllocator)Marshal.GetDelegateForFunctionPointer(api_.ReleaseAllocator, typeof(DOrtReleaseAllocator)); OrtCreateMemoryInfo = (DOrtCreateMemoryInfo)Marshal.GetDelegateForFunctionPointer(api_.CreateMemoryInfo, typeof(DOrtCreateMemoryInfo)); OrtCreateCpuMemoryInfo = (DOrtCreateCpuMemoryInfo)Marshal.GetDelegateForFunctionPointer(api_.CreateCpuMemoryInfo, typeof(DOrtCreateCpuMemoryInfo)); OrtReleaseMemoryInfo = (DOrtReleaseMemoryInfo)Marshal.GetDelegateForFunctionPointer(api_.ReleaseMemoryInfo, typeof(DOrtReleaseMemoryInfo)); OrtCompareMemoryInfo = (DOrtCompareMemoryInfo)Marshal.GetDelegateForFunctionPointer(api_.CompareMemoryInfo, typeof(DOrtCompareMemoryInfo)); OrtMemoryInfoGetName = (DOrtMemoryInfoGetName)Marshal.GetDelegateForFunctionPointer(api_.MemoryInfoGetName, typeof(DOrtMemoryInfoGetName)); OrtMemoryInfoGetId = (DOrtMemoryInfoGetId)Marshal.GetDelegateForFunctionPointer(api_.MemoryInfoGetId, typeof(DOrtMemoryInfoGetId)); OrtMemoryInfoGetMemType = (DOrtMemoryInfoGetMemType)Marshal.GetDelegateForFunctionPointer(api_.MemoryInfoGetMemType, typeof(DOrtMemoryInfoGetMemType)); OrtMemoryInfoGetType = (DOrtMemoryInfoGetType)Marshal.GetDelegateForFunctionPointer(api_.MemoryInfoGetType, typeof(DOrtMemoryInfoGetType)); OrtGetAllocatorWithDefaultOptions = (DOrtGetAllocatorWithDefaultOptions)Marshal.GetDelegateForFunctionPointer(api_.GetAllocatorWithDefaultOptions, typeof(DOrtGetAllocatorWithDefaultOptions)); OrtCreateAllocator = (DOrtCreateAllocator)Marshal.GetDelegateForFunctionPointer(api_.CreateAllocator, typeof(DOrtCreateAllocator)); OrtReleaseAllocator = (DOrtReleaseAllocator)Marshal.GetDelegateForFunctionPointer(api_.ReleaseAllocator, typeof(DOrtReleaseAllocator)); OrtAllocatorAlloc = (DOrtAllocatorAlloc)Marshal.GetDelegateForFunctionPointer(api_.AllocatorAlloc, typeof(DOrtAllocatorAlloc)); OrtAllocatorFree = (DOrtAllocatorFree)Marshal.GetDelegateForFunctionPointer(api_.AllocatorFree, typeof(DOrtAllocatorFree)); OrtAllocatorGetInfo = (DOrtAllocatorGetInfo)Marshal.GetDelegateForFunctionPointer(api_.AllocatorGetInfo, typeof(DOrtAllocatorGetInfo)); OrtAddFreeDimensionOverride = (DOrtAddFreeDimensionOverride)Marshal.GetDelegateForFunctionPointer(api_.AddFreeDimensionOverride, typeof(DOrtAddFreeDimensionOverride)); OrtAddFreeDimensionOverrideByName = (DOrtAddFreeDimensionOverrideByName)Marshal.GetDelegateForFunctionPointer(api_.AddFreeDimensionOverrideByName, typeof(DOrtAddFreeDimensionOverrideByName)); OrtCreateIoBinding = (DOrtCreateIoBinding)Marshal.GetDelegateForFunctionPointer(api_.CreateIoBinding, typeof(DOrtCreateIoBinding)); OrtReleaseIoBinding = (DOrtReleaseIoBinding)Marshal.GetDelegateForFunctionPointer(api_.ReleaseIoBinding, typeof(DOrtReleaseIoBinding)); OrtBindInput = (DOrtBindInput)Marshal.GetDelegateForFunctionPointer(api_.BindInput, typeof(DOrtBindInput)); OrtSynchronizeBoundInputs = (DOrtSynchronizeBoundInputs)Marshal.GetDelegateForFunctionPointer(api_.SynchronizeBoundInputs, typeof(DOrtSynchronizeBoundInputs)); OrtBindOutput = (DOrtBindOutput)Marshal.GetDelegateForFunctionPointer(api_.BindOutput, typeof(DOrtBindOutput)); OrtBindOutputToDevice = (DOrtBindOutputToDevice)Marshal.GetDelegateForFunctionPointer(api_.BindOutputToDevice, typeof(DOrtBindOutputToDevice)); OrtSynchronizeBoundOutputs = (DOrtSynchronizeBoundOutputs)Marshal.GetDelegateForFunctionPointer(api_.SynchronizeBoundOutputs, typeof(DOrtSynchronizeBoundOutputs)); OrtGetBoundOutputNames = (DOrtGetBoundOutputNames)Marshal.GetDelegateForFunctionPointer(api_.GetBoundOutputNames, typeof(DOrtGetBoundOutputNames)); OrtGetBoundOutputValues = (DOrtGetBoundOutputValues)Marshal.GetDelegateForFunctionPointer(api_.GetBoundOutputValues, typeof(DOrtGetBoundOutputValues)); OrtClearBoundInputs = (DOrtClearBoundInputs)Marshal.GetDelegateForFunctionPointer(api_.ClearBoundInputs, typeof(DOrtClearBoundInputs)); OrtClearBoundOutputs = (DOrtClearBoundOutputs)Marshal.GetDelegateForFunctionPointer(api_.ClearBoundOutputs, typeof(DOrtClearBoundOutputs)); OrtTensorAt = (DOrtTensorAt)Marshal.GetDelegateForFunctionPointer(api_.TensorAt, typeof(DOrtTensorAt)); OrtCreateAndRegisterAllocator = (DOrtCreateAndRegisterAllocator)Marshal.GetDelegateForFunctionPointer(api_.CreateAndRegisterAllocator, typeof(DOrtCreateAndRegisterAllocator)); OrtSetLanguageProjection = (DOrtSetLanguageProjection)Marshal.GetDelegateForFunctionPointer(api_.SetLanguageProjection, typeof(DOrtSetLanguageProjection)); OrtHasValue = (DOrtHasValue)Marshal.GetDelegateForFunctionPointer(api_.HasValue, typeof(DOrtHasValue)); OrtGetValue = (DOrtGetValue)Marshal.GetDelegateForFunctionPointer(api_.GetValue, typeof(DOrtGetValue)); OrtGetValueCount = (DOrtGetValueCount)Marshal.GetDelegateForFunctionPointer(api_.GetValueCount, typeof(DOrtGetValueCount)); OrtCreateValue = (DOrtCreateValue)Marshal.GetDelegateForFunctionPointer(api_.CreateValue, typeof(DOrtCreateValue)); OrtGetValueType = (DOrtGetValueType)Marshal.GetDelegateForFunctionPointer(api_.GetValueType, typeof(DOrtGetValueType)); OrtGetOnnxTypeFromTypeInfo = (DOrtGetOnnxTypeFromTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.GetOnnxTypeFromTypeInfo, typeof(DOrtGetOnnxTypeFromTypeInfo)); OrtGetTypeInfo = (DOrtGetTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.GetTypeInfo, typeof(DOrtGetTypeInfo)); OrtCreateTensorAsOrtValue = (DOrtCreateTensorAsOrtValue)Marshal.GetDelegateForFunctionPointer(api_.CreateTensorAsOrtValue, typeof(DOrtCreateTensorAsOrtValue)); OrtCreateTensorWithDataAsOrtValue = (DOrtCreateTensorWithDataAsOrtValue)Marshal.GetDelegateForFunctionPointer(api_.CreateTensorWithDataAsOrtValue, typeof(DOrtCreateTensorWithDataAsOrtValue)); OrtValueIsTensor = (DOrtValueIsTensor)Marshal.GetDelegateForFunctionPointer(api_.IsTensor, typeof(DOrtValueIsTensor)); OrtValueIsSparseTensor = (DOrtValueIsSparseTensor)Marshal.GetDelegateForFunctionPointer(api_.IsSparseTensor, typeof(DOrtValueIsSparseTensor)); OrtGetTensorMutableData = (DOrtGetTensorMutableData)Marshal.GetDelegateForFunctionPointer(api_.GetTensorMutableData, typeof(DOrtGetTensorMutableData)); OrtFillStringTensor = (DOrtFillStringTensor)Marshal.GetDelegateForFunctionPointer(api_.FillStringTensor, typeof(DOrtFillStringTensor)); OrtGetResizedStringTensorElementBuffer = (DOrtGetResizedStringTensorElementBuffer)Marshal.GetDelegateForFunctionPointer(api_.GetResizedStringTensorElementBuffer, typeof(DOrtGetResizedStringTensorElementBuffer)); OrtGetStringTensorContent = (DOrtGetStringTensorContent)Marshal.GetDelegateForFunctionPointer(api_.GetStringTensorContent, typeof(DOrtGetStringTensorContent)); OrtGetStringTensorDataLength = (DOrtGetStringTensorDataLength)Marshal.GetDelegateForFunctionPointer(api_.GetStringTensorDataLength, typeof(DOrtGetStringTensorDataLength)); OrtGetStringTensorElementLength = (DOrtGetStringTensorElementLength)Marshal.GetDelegateForFunctionPointer(api_.GetStringTensorElementLength, typeof(DOrtGetStringTensorElementLength)); OrtGetStringTensorElement = (DOrtGetStringTensorElement)Marshal.GetDelegateForFunctionPointer(api_.GetStringTensorElement, typeof(DOrtGetStringTensorElement)); OrtCastTypeInfoToTensorInfo = (DOrtCastTypeInfoToTensorInfo)Marshal.GetDelegateForFunctionPointer(api_.CastTypeInfoToTensorInfo, typeof(DOrtCastTypeInfoToTensorInfo)); OrtGetTensorTypeAndShape = (DOrtGetTensorTypeAndShape)Marshal.GetDelegateForFunctionPointer(api_.GetTensorTypeAndShape, typeof(DOrtGetTensorTypeAndShape)); OrtReleaseTensorTypeAndShapeInfo = (DOrtReleaseTensorTypeAndShapeInfo)Marshal.GetDelegateForFunctionPointer(api_.ReleaseTensorTypeAndShapeInfo, typeof(DOrtReleaseTensorTypeAndShapeInfo)); OrtGetTensorElementType = (DOrtGetTensorElementType)Marshal.GetDelegateForFunctionPointer(api_.GetTensorElementType, typeof(DOrtGetTensorElementType)); OrtGetDimensionsCount = (DOrtGetDimensionsCount)Marshal.GetDelegateForFunctionPointer(api_.GetDimensionsCount, typeof(DOrtGetDimensionsCount)); OrtGetDimensions = (DOrtGetDimensions)Marshal.GetDelegateForFunctionPointer(api_.GetDimensions, typeof(DOrtGetDimensions)); OrtGetSymbolicDimensions = (DOrtGetSymbolicDimensions)Marshal.GetDelegateForFunctionPointer(api_.GetSymbolicDimensions, typeof(DOrtGetSymbolicDimensions)); OrtGetTensorShapeElementCount = (DOrtGetTensorShapeElementCount)Marshal.GetDelegateForFunctionPointer(api_.GetTensorShapeElementCount, typeof(DOrtGetTensorShapeElementCount)); OrtGetTensorMemoryInfo = (DOrtGetTensorMemoryInfo)Marshal.GetDelegateForFunctionPointer(api_.GetTensorMemoryInfo, typeof(DOrtGetTensorMemoryInfo)); OrtGetMapKeyType = (DGetMapKeyType)Marshal.GetDelegateForFunctionPointer(api_.GetMapKeyType, typeof(DGetMapKeyType)); OrtCastTypeInfoToMapTypeInfo = (DCastTypeInfoToMapTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.CastTypeInfoToMapTypeInfo, typeof(DCastTypeInfoToMapTypeInfo)); OrtGetMapValueType = (DGetMapValueType)Marshal.GetDelegateForFunctionPointer(api_.GetMapValueType, typeof(DGetMapValueType)); OrtCastTypeInfoToSequenceTypeInfo = (DCastTypeInfoToSequenceTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.CastTypeInfoToSequenceTypeInfo, typeof(DCastTypeInfoToSequenceTypeInfo)); OrtGetSequenceElementType = (DGetSequenceElementType)Marshal.GetDelegateForFunctionPointer(api_.GetSequenceElementType, typeof(DGetSequenceElementType)); OrtCastTypeInfoToOptionalTypeInfo = (DOrtCastTypeInfoToOptionalTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.CastTypeInfoToOptionalTypeInfo, typeof(DOrtCastTypeInfoToOptionalTypeInfo)); OrtGetOptionalContainedTypeInfo = (DGetOptionalContainedTypeInfo)Marshal.GetDelegateForFunctionPointer(api_.GetOptionalContainedTypeInfo, typeof(DGetOptionalContainedTypeInfo)); OrtReleaseValue = (DOrtReleaseValue)Marshal.GetDelegateForFunctionPointer(api_.ReleaseValue, typeof(DOrtReleaseValue)); OrtSessionGetModelMetadata = (DOrtSessionGetModelMetadata)Marshal.GetDelegateForFunctionPointer(api_.SessionGetModelMetadata, typeof(DOrtSessionGetModelMetadata)); OrtModelMetadataGetProducerName = (DOrtModelMetadataGetProducerName)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetProducerName, typeof(DOrtModelMetadataGetProducerName)); OrtModelMetadataGetGraphName = (DOrtModelMetadataGetGraphName)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetGraphName, typeof(DOrtModelMetadataGetGraphName)); OrtModelMetadataGetDomain = (DOrtModelMetadataGetDomain)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetDomain, typeof(DOrtModelMetadataGetDomain)); OrtModelMetadataGetDescription = (DOrtModelMetadataGetDescription)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetDescription, typeof(DOrtModelMetadataGetDescription)); OrtModelMetadataGetGraphDescription = (DOrtModelMetadataGetGraphDescription)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetGraphDescription, typeof(DOrtModelMetadataGetGraphDescription)); OrtModelMetadataGetVersion = (DOrtModelMetadataGetVersion)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetVersion, typeof(DOrtModelMetadataGetVersion)); OrtModelMetadataGetCustomMetadataMapKeys = (DOrtModelMetadataGetCustomMetadataMapKeys)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataGetCustomMetadataMapKeys, typeof(DOrtModelMetadataGetCustomMetadataMapKeys)); OrtModelMetadataLookupCustomMetadataMap = (DOrtModelMetadataLookupCustomMetadataMap)Marshal.GetDelegateForFunctionPointer(api_.ModelMetadataLookupCustomMetadataMap, typeof(DOrtModelMetadataLookupCustomMetadataMap)); OrtReleaseModelMetadata = (DOrtReleaseModelMetadata)Marshal.GetDelegateForFunctionPointer(api_.ReleaseModelMetadata, typeof(DOrtReleaseModelMetadata)); OrtGetAvailableProviders = (DOrtGetAvailableProviders)Marshal.GetDelegateForFunctionPointer(api_.GetAvailableProviders, typeof(DOrtGetAvailableProviders)); OrtReleaseAvailableProviders = (DOrtReleaseAvailableProviders)Marshal.GetDelegateForFunctionPointer(api_.ReleaseAvailableProviders, typeof(DOrtReleaseAvailableProviders)); OrtCreatePrepackedWeightsContainer = (DOrtCreatePrepackedWeightsContainer)Marshal.GetDelegateForFunctionPointer(api_.CreatePrepackedWeightsContainer, typeof(DOrtCreatePrepackedWeightsContainer)); OrtReleasePrepackedWeightsContainer = (DOrtReleasePrepackedWeightsContainer)Marshal.GetDelegateForFunctionPointer(api_.ReleasePrepackedWeightsContainer, typeof(DOrtReleasePrepackedWeightsContainer)); SessionOptionsAppendExecutionProvider_TensorRT_V2 = (DSessionOptionsAppendExecutionProvider_TensorRT_V2)Marshal.GetDelegateForFunctionPointer(api_.SessionOptionsAppendExecutionProvider_TensorRT_V2, typeof(DSessionOptionsAppendExecutionProvider_TensorRT_V2)); OrtCreateTensorRTProviderOptions = (DOrtCreateTensorRTProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.CreateTensorRTProviderOptions, typeof(DOrtCreateTensorRTProviderOptions)); OrtUpdateTensorRTProviderOptions = (DOrtUpdateTensorRTProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.UpdateTensorRTProviderOptions, typeof(DOrtUpdateTensorRTProviderOptions)); OrtGetTensorRTProviderOptionsAsString = (DOrtGetTensorRTProviderOptionsAsString)Marshal.GetDelegateForFunctionPointer(api_.GetTensorRTProviderOptionsAsString, typeof(DOrtGetTensorRTProviderOptionsAsString)); OrtReleaseTensorRTProviderOptions = (DOrtReleaseTensorRTProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.ReleaseTensorRTProviderOptions, typeof(DOrtReleaseTensorRTProviderOptions)); SessionOptionsAppendExecutionProvider_CUDA = (DSessionOptionsAppendExecutionProvider_CUDA)Marshal.GetDelegateForFunctionPointer(api_.SessionOptionsAppendExecutionProvider_CUDA, typeof(DSessionOptionsAppendExecutionProvider_CUDA)); SessionOptionsAppendExecutionProvider_CUDA_V2 = (DSessionOptionsAppendExecutionProvider_CUDA_V2)Marshal.GetDelegateForFunctionPointer(api_.SessionOptionsAppendExecutionProvider_CUDA_V2, typeof(DSessionOptionsAppendExecutionProvider_CUDA_V2)); OrtCreateCUDAProviderOptions = (DOrtCreateCUDAProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.CreateCUDAProviderOptions, typeof(DOrtCreateCUDAProviderOptions)); OrtUpdateCUDAProviderOptions = (DOrtUpdateCUDAProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.UpdateCUDAProviderOptions, typeof(DOrtUpdateCUDAProviderOptions)); OrtGetCUDAProviderOptionsAsString = (DOrtGetCUDAProviderOptionsAsString)Marshal.GetDelegateForFunctionPointer(api_.GetCUDAProviderOptionsAsString, typeof(DOrtGetCUDAProviderOptionsAsString)); OrtReleaseCUDAProviderOptions = (DOrtReleaseCUDAProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.ReleaseCUDAProviderOptions, typeof(DOrtReleaseCUDAProviderOptions)); SessionOptionsAppendExecutionProvider = (DSessionOptionsAppendExecutionProvider)Marshal.GetDelegateForFunctionPointer(api_.SessionOptionsAppendExecutionProvider, typeof(DSessionOptionsAppendExecutionProvider)); OrtUpdateEnvWithCustomLogLevel = (DOrtUpdateEnvWithCustomLogLevel)Marshal.GetDelegateForFunctionPointer(api_.UpdateEnvWithCustomLogLevel, typeof(DOrtUpdateEnvWithCustomLogLevel)); SessionOptionsAppendExecutionProvider_ROCM = (DSessionOptionsAppendExecutionProvider_ROCM)Marshal.GetDelegateForFunctionPointer(api_.SessionOptionsAppendExecutionProvider_ROCM, typeof(DSessionOptionsAppendExecutionProvider_ROCM)); OrtCreateROCMProviderOptions = (DOrtCreateROCMProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.CreateROCMProviderOptions, typeof(DOrtCreateROCMProviderOptions)); OrtUpdateROCMProviderOptions = (DOrtUpdateROCMProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.UpdateROCMProviderOptions, typeof(DOrtUpdateROCMProviderOptions)); OrtGetROCMProviderOptionsAsString = (DOrtGetROCMProviderOptionsAsString)Marshal.GetDelegateForFunctionPointer(api_.GetROCMProviderOptionsAsString, typeof(DOrtGetROCMProviderOptionsAsString)); OrtReleaseROCMProviderOptions = (DOrtReleaseROCMProviderOptions)Marshal.GetDelegateForFunctionPointer(api_.ReleaseROCMProviderOptions, typeof(DOrtReleaseROCMProviderOptions)); OrtCreateAndRegisterAllocatorV2 = (DCreateAndRegisterAllocatorV2)Marshal.GetDelegateForFunctionPointer(api_.CreateAndRegisterAllocatorV2, typeof(DCreateAndRegisterAllocatorV2)); OrtRunAsync = (DOrtRunAsync)Marshal.GetDelegateForFunctionPointer(api_.RunAsync, typeof(DOrtRunAsync)); } [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern ref OrtApiBase OrtGetApiBase(); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_CPU(IntPtr options, int use_arena); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_Dnnl(IntPtr options, int use_arena); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_CUDA(IntPtr options, int device_id); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_ROCM(IntPtr options, int device_id); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_DML(IntPtr options, int device_id); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_OpenVINO(IntPtr options, byte[] device_id); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_Tensorrt(IntPtr options, int device_id); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_MIGraphX(IntPtr options, int device_id); [DllImport("onnxruntime", CharSet = CharSet.Ansi)] public static extern IntPtr OrtSessionOptionsAppendExecutionProvider_Tvm(IntPtr options, byte[] settings); } internal static class OrtExtensionsNativeMethods { internal const string ExtensionsDllName = "ortextensions"; [DllImport("ortextensions", CharSet = CharSet.Ansi)] public static extern IntPtr RegisterCustomOps(IntPtr sessionOptions, ref OrtApiBase ortApiBase); } internal static class NativeOnnxValueHelper { internal static byte[] StringToZeroTerminatedUtf8(string s) { int byteCount = Encoding.UTF8.GetByteCount(s); byte[] array = new byte[byteCount + 1]; int bytes = Encoding.UTF8.GetBytes(s, 0, s.Length, array, 0); array[^1] = 0; return array; } internal unsafe static void StringToUtf8NativeMemory(char* strPtr, int strLength, IntPtr ptr, int nativeBufferSize) { int bytes = Encoding.UTF8.GetBytes(strPtr, strLength, (byte*)(void*)ptr, nativeBufferSize); if (bytes != nativeBufferSize) { throw new OnnxRuntimeException(ErrorCode.RuntimeException, $"Failed to convert to UTF8. Expected bytes: {nativeBufferSize}, written: {bytes}"); } } internal unsafe static string StringFromNativeUtf8(IntPtr nativeUtf8, OrtAllocator allocator = null) { try { int i; for (i = 0; *(bool*)(void*)(nativeUtf8 + i); i++) { } if (i == 0) { return string.Empty; } byte* bytes = (byte*)(void*)nativeUtf8; return Encoding.UTF8.GetString(bytes, i); } finally { allocator?.FreeMemory(nativeUtf8); } } internal unsafe static void StringAndUtf8FromNative(OrtAllocator allocator, IntPtr nativeUtf8, out string str, out IntPtr utf8) { try { int i; for (i = 0; *(bool*)(void*)(nativeUtf8 + i); i++) { } if (i == 0) { str = string.Empty; utf8 = IntPtr.Zero; return; } Span span = new Span(nativeUtf8.ToPointer(), i); utf8 = Marshal.AllocHGlobal(i + 1); try { Span destination = new Span(utf8.ToPointer(), i + 1); span.CopyTo(destination); destination[i] = 0; byte* bytes = (byte*)(void*)nativeUtf8; str = Encoding.UTF8.GetString(bytes, i); } catch (Exception) { Marshal.FreeHGlobal(utf8); throw; } } finally { allocator.FreeMemory(nativeUtf8); } } internal static byte[] GetPlatformSerializedString(string str) { if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) { return Encoding.Unicode.GetBytes(str + "\0"); } return StringToZeroTerminatedUtf8(str); } } internal ref struct DisposableArray where T : IDisposable { internal Span Span { get; private set; } internal DisposableArray(Span disposables) { Span = disposables; } public void Dispose() { Span span = Span; for (int num = span.Length - 1; num >= 0; num--) { span = Span; ref T reference = ref span[num]; T val = default(T); if (val == null) { val = reference; reference = ref val; if (val == null) { continue; } } reference.Dispose(); } } } internal ref struct DisposableOrtValueHandleArray { internal Span Span { get; private set; } internal DisposableOrtValueHandleArray(Span handles) { Span = handles; } public void Dispose() { for (int num = Span.Length - 1; num >= 0; num--) { if (Span[num] != IntPtr.Zero) { NativeMethods.OrtReleaseValue(Span[num]); } } } } public struct MarshaledString : IDisposable { internal int Length { get; private set; } internal IntPtr Value { get; private set; } internal unsafe MarshaledString(string input) { int num; IntPtr value; if (input == null) { num = 0; value = IntPtr.Zero; } else { byte[] array = ((input.Length != 0) ? Encoding.UTF8.GetBytes(input) : ArrayUtilities.GetEmpty()); num = array.Length; value = Marshal.AllocHGlobal(num + 1); Span destination = new Span(value.ToPointer(), num + 1); array.AsSpan(0, num).CopyTo(destination); destination[num] = 0; } Length = num; Value = value; } public void Dispose() { if (Value != IntPtr.Zero) { Marshal.FreeHGlobal(Value); Value = IntPtr.Zero; Length = 0; } } } public ref struct MarshaledStringArray { private MarshaledString[] _values; internal ReadOnlySpan Values => _values; internal MarshaledStringArray(Tensor inputs) { if (inputs.Length == 0) { _values = null; return; } _values = new MarshaledString[inputs.Length]; for (int i = 0; i < inputs.Length; i++) { _values[i] = new MarshaledString(inputs.GetValue(i)); } } internal MarshaledStringArray(IEnumerable inputs) { if (inputs == null) { _values = null; return; } _values = new MarshaledString[inputs.Count()]; int num = 0; foreach (string input in inputs) { _values[num++] = new MarshaledString(input); } } internal void Fill(IntPtr[] pDestination) { if (_values != null) { for (int i = 0; i < _values.Length; i++) { pDestination[i] = Values[i].Value; } } } public void Dispose() { if (_values != null) { for (int i = 0; i < _values.Length; i++) { _values[i].Dispose(); } _values = null; } } } internal class ProviderOptionsUpdater { internal static void Update(Dictionary providerOptions, IntPtr handle, Func updateFunc) { string[] array = providerOptions.Keys.ToArray(); string[] array2 = providerOptions.Values.ToArray(); MarshaledStringArray marshaledStringArray = default(MarshaledStringArray); MarshaledStringArray marshaledStringArray2 = default(MarshaledStringArray); try { marshaledStringArray = new MarshaledStringArray(array); marshaledStringArray2 = new MarshaledStringArray(array2); IntPtr[] array3 = new IntPtr[array.Length]; marshaledStringArray.Fill(array3); IntPtr[] array4 = new IntPtr[array2.Length]; marshaledStringArray2.Fill(array4); NativeApiStatus.VerifySuccess(updateFunc(handle, array3, array4, (UIntPtr)(ulong)providerOptions.Count)); } finally { marshaledStringArray.Dispose(); marshaledStringArray2.Dispose(); } } } public enum OrtAllocatorType { DeviceAllocator, ArenaAllocator } public enum OrtMemType { CpuInput = -2, CpuOutput = -1, Cpu = -1, Default = 0 } public class OrtArenaCfg : SafeHandle { internal IntPtr Pointer => handle; public override bool IsInvalid => handle == IntPtr.Zero; public OrtArenaCfg(uint maxMemory, int arenaExtendStrategy, int initialChunkSizeBytes, int maxDeadBytesPerChunk) : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateArenaCfg((UIntPtr)maxMemory, arenaExtendStrategy, initialChunkSizeBytes, maxDeadBytesPerChunk, out handle)); } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseArenaCfg(handle); handle = IntPtr.Zero; return true; } } public class OrtMemoryInfo : SafeHandle { private static readonly Lazy _defaultCpuAllocInfo = new Lazy(CreateCpuMemoryInfo); private readonly bool _owned; public static readonly byte[] allocatorCPU = Encoding.UTF8.GetBytes("Cpu\0"); public static readonly byte[] allocatorCUDA = Encoding.UTF8.GetBytes("Cuda\0"); public static readonly byte[] allocatorCUDA_PINNED = Encoding.UTF8.GetBytes("CudaPinned\0"); public static readonly byte[] allocatorHIP = Encoding.UTF8.GetBytes("Hip\0"); public static readonly byte[] allocatorHIP_PINNED = Encoding.UTF8.GetBytes("HipPinned\0"); public static OrtMemoryInfo DefaultInstance => _defaultCpuAllocInfo.Value; internal IntPtr Pointer => handle; public override bool IsInvalid => handle == IntPtr.Zero; public string Name { get { NativeApiStatus.VerifySuccess(NativeMethods.OrtMemoryInfoGetName(handle, out var name)); return NativeOnnxValueHelper.StringFromNativeUtf8(name); } } public int Id { get { NativeApiStatus.VerifySuccess(NativeMethods.OrtMemoryInfoGetId(handle, out var id)); return id; } } private static OrtMemoryInfo CreateCpuMemoryInfo() { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateCpuMemoryInfo(OrtAllocatorType.DeviceAllocator, OrtMemType.CpuOutput, out var allocatorInfo)); return new OrtMemoryInfo(allocatorInfo, owned: true); } internal OrtMemoryInfo(IntPtr allocInfo, bool owned) : base(allocInfo, ownsHandle: true) { _owned = owned; } public OrtMemoryInfo(byte[] utf8AllocatorName, OrtAllocatorType allocatorType, int deviceId, OrtMemType memoryType) : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateMemoryInfo(utf8AllocatorName, allocatorType, deviceId, memoryType, out handle)); _owned = true; } public OrtMemoryInfo(string allocatorName, OrtAllocatorType allocatorType, int deviceId, OrtMemType memoryType) : this(NativeOnnxValueHelper.StringToZeroTerminatedUtf8(allocatorName), allocatorType, deviceId, memoryType) { } public OrtMemType GetMemoryType() { NativeApiStatus.VerifySuccess(NativeMethods.OrtMemoryInfoGetMemType(handle, out var mem_type)); return mem_type; } public OrtAllocatorType GetAllocatorType() { NativeApiStatus.VerifySuccess(NativeMethods.OrtMemoryInfoGetType(handle, out var alloc_type)); return alloc_type; } public override bool Equals(object obj) { if (!(obj is OrtMemoryInfo other)) { return false; } return Equals(other); } public bool Equals(OrtMemoryInfo other) { if (this == other) { return true; } NativeApiStatus.VerifySuccess(NativeMethods.OrtCompareMemoryInfo(handle, other.Pointer, out var result)); return result == 0; } public override int GetHashCode() { return Pointer.ToInt32(); } protected override bool ReleaseHandle() { if (_owned) { NativeMethods.OrtReleaseMemoryInfo(handle); } handle = IntPtr.Zero; return true; } } [Obsolete("Create OrtValue over an arbitrary piece of memory and use it where appropriate.")] public class OrtExternalAllocation { public OrtMemoryInfo Info { get; private set; } public long[] Shape { get; private set; } public TensorElementType ElementType { get; private set; } public IntPtr Pointer { get; private set; } public long Size { get; private set; } public OrtExternalAllocation(OrtMemoryInfo memInfo, long[] shape, TensorElementType elementType, IntPtr pointer, long sizeInBytes) { TensorElementTypeInfo tensorElementTypeInfo = TensorBase.GetElementTypeInfo(elementType) ?? throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Tensor element type: {elementType} is not supported"); if (tensorElementTypeInfo.IsString) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Strings are not supported by this API"); } long sizeForShape = ShapeUtils.GetSizeForShape(shape); long num = sizeForShape * tensorElementTypeInfo.TypeSize; if (num > sizeInBytes) { string message = $"Shape of {sizeForShape} elements requires a buffer of at least {num} bytes. Provided: {sizeInBytes} bytes"; throw new OnnxRuntimeException(ErrorCode.InvalidArgument, message); } Info = memInfo; Shape = shape; ElementType = elementType; Pointer = pointer; Size = sizeInBytes; } } public class OrtMemoryAllocation : SafeHandle { private OrtAllocator _allocator; internal IntPtr Pointer => handle; public override bool IsInvalid => handle == IntPtr.Zero; public uint Size { get; private set; } public OrtMemoryInfo Info => _allocator.Info; internal OrtMemoryAllocation(OrtAllocator allocator, IntPtr pointer, uint size) : base(pointer, ownsHandle: true) { _allocator = allocator; Size = size; } protected override bool ReleaseHandle() { _allocator.FreeMemory(handle); handle = IntPtr.Zero; return true; } } public class OrtAllocator : SafeHandle { private static readonly Lazy _defaultInstance = new Lazy(GetDefaultCpuAllocator); private readonly bool _owned; public static OrtAllocator DefaultInstance => _defaultInstance.Value; internal IntPtr Pointer => handle; public override bool IsInvalid => handle == IntPtr.Zero; public OrtMemoryInfo Info { get { NativeApiStatus.VerifySuccess(NativeMethods.OrtAllocatorGetInfo(handle, out var info)); return new OrtMemoryInfo(info, owned: false); } } private static OrtAllocator GetDefaultCpuAllocator() { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetAllocatorWithDefaultOptions(out var allocator)); return new OrtAllocator(allocator, owned: false); } internal OrtAllocator(IntPtr allocator, bool owned) : base(allocator, ownsHandle: true) { _owned = owned; } public OrtAllocator(InferenceSession session, OrtMemoryInfo memInfo) : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateAllocator(session.Handle, memInfo.Pointer, out handle)); _owned = true; } public OrtMemoryAllocation Allocate(uint size) { NativeApiStatus.VerifySuccess(NativeMethods.OrtAllocatorAlloc(handle, (UIntPtr)size, out var p)); return new OrtMemoryAllocation(this, p, size); } internal void FreeMemory(IntPtr allocation) { NativeApiStatus.VerifySuccess(NativeMethods.OrtAllocatorFree(handle, allocation)); } protected override bool ReleaseHandle() { if (_owned) { NativeMethods.OrtReleaseAllocator(handle); } handle = IntPtr.Zero; return true; } } public delegate void DOrtLoggingFunction(IntPtr param, OrtLoggingLevel severity, string category, string logId, string codeLocation, string message); public struct EnvironmentCreationOptions { public string logId; public OrtLoggingLevel? logLevel; public OrtThreadingOptions threadOptions; public IntPtr? loggingParam; public DOrtLoggingFunction loggingFunction; } public sealed class OrtEnv : SafeHandle { private delegate void DOrtLoggingFunctionInternal(IntPtr param, IntPtr severity, IntPtr category, IntPtr logid, IntPtr codeLocation, IntPtr message); private static readonly int ORT_PROJECTION_CSHARP = 2; private static readonly byte[] _defaultLogId = NativeOnnxValueHelper.StringToZeroTerminatedUtf8("CSharpOnnxRuntime"); private static EnvironmentCreationOptions? _createOptions; private static Lazy _instance = new Lazy(CreateInstance); private static readonly DOrtLoggingFunctionInternal _loggingFunctionInternal = LoggingFunctionThunk; private static DOrtLoggingFunction _userLoggingFunction; private OrtLoggingLevel _envLogLevel; public static bool IsCreated => _instance.IsValueCreated; public OrtLoggingLevel EnvLogLevel { get { return _envLogLevel; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtUpdateEnvWithCustomLogLevel(Handle, value)); _envLogLevel = value; } } internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; private OrtEnv(IntPtr handle, OrtLoggingLevel logLevel) : base(handle, ownsHandle: true) { _envLogLevel = logLevel; } private static void LoggingFunctionThunk(IntPtr param, IntPtr severity, IntPtr category, IntPtr logid, IntPtr codeLocation, IntPtr message) { string category2 = NativeOnnxValueHelper.StringFromNativeUtf8(category); string logId = NativeOnnxValueHelper.StringFromNativeUtf8(logid); string codeLocation2 = NativeOnnxValueHelper.StringFromNativeUtf8(codeLocation); string message2 = NativeOnnxValueHelper.StringFromNativeUtf8(message); _userLoggingFunction(param, (OrtLoggingLevel)(int)severity, category2, logId, codeLocation2, message2); } private static OrtEnv CreateInstance() { OrtEnv ortEnv = null; if (!_createOptions.HasValue) { return CreateDefaultEnv(OrtLoggingLevel.ORT_LOGGING_LEVEL_WARNING, _defaultLogId); } EnvironmentCreationOptions value = _createOptions.Value; byte[] logIdUtf = (string.IsNullOrEmpty(value.logId) ? _defaultLogId : NativeOnnxValueHelper.StringToZeroTerminatedUtf8(value.logId)); OrtLoggingLevel logLevel = value.logLevel ?? OrtLoggingLevel.ORT_LOGGING_LEVEL_WARNING; OrtThreadingOptions threadOptions = value.threadOptions; DOrtLoggingFunction loggingFunction = value.loggingFunction; IntPtr intPtr = value.loggingParam ?? IntPtr.Zero; if (threadOptions == null && loggingFunction == null) { return CreateDefaultEnv(logLevel, logIdUtf); } if (threadOptions == null) { return CreateWithCustomLogger(logLevel, logIdUtf, intPtr, loggingFunction); } if (loggingFunction == null) { return CreateWithThreadingOptions(logLevel, logIdUtf, threadOptions); } return CreateEnvWithCustomLoggerAndGlobalThreadPools(logLevel, logIdUtf, intPtr, threadOptions, loggingFunction); } private static OrtEnv CreateDefaultEnv(OrtLoggingLevel logLevel, byte[] logIdUtf8) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateEnv(logLevel, logIdUtf8, out var env)); OrtEnv ortEnv = new OrtEnv(env, logLevel); SetLanguageProjection(ortEnv); return ortEnv; } private static OrtEnv CreateWithCustomLogger(OrtLoggingLevel logLevel, byte[] logIdUtf8, IntPtr loggerParam, DOrtLoggingFunction loggingFunction) { _userLoggingFunction = loggingFunction; IntPtr functionPointerForDelegate = Marshal.GetFunctionPointerForDelegate(_loggingFunctionInternal); NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateEnvWithCustomLogger(functionPointerForDelegate, loggerParam, logLevel, logIdUtf8, out var env)); OrtEnv ortEnv = new OrtEnv(env, logLevel); SetLanguageProjection(ortEnv); return ortEnv; } private static OrtEnv CreateWithThreadingOptions(OrtLoggingLevel logLevel, byte[] logIdUtf8, OrtThreadingOptions threadingOptions) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateEnvWithGlobalThreadPools(logLevel, logIdUtf8, threadingOptions.Handle, out var env)); OrtEnv ortEnv = new OrtEnv(env, logLevel); SetLanguageProjection(ortEnv); return ortEnv; } private static OrtEnv CreateEnvWithCustomLoggerAndGlobalThreadPools(OrtLoggingLevel logLevel, byte[] logIdUtf8, IntPtr logParam, OrtThreadingOptions threadingOptions, DOrtLoggingFunction loggingFunction) { _userLoggingFunction = loggingFunction; IntPtr functionPointerForDelegate = Marshal.GetFunctionPointerForDelegate(_loggingFunctionInternal); NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateEnvWithCustomLoggerAndGlobalThreadPools(functionPointerForDelegate, logParam, logLevel, logIdUtf8, threadingOptions.Handle, out var env)); OrtEnv ortEnv = new OrtEnv(env, logLevel); SetLanguageProjection(ortEnv); return ortEnv; } private static void SetLanguageProjection(OrtEnv env) { try { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetLanguageProjection(env.Handle, ORT_PROJECTION_CSHARP)); } catch (Exception) { env.Dispose(); throw; } } public static OrtEnv Instance() { return _instance.Value; } public static OrtEnv CreateInstanceWithOptions(ref EnvironmentCreationOptions options) { if (_instance.IsValueCreated) { throw new OnnxRuntimeException(ErrorCode.RuntimeException, "OrtEnv singleton instance already exists, supplied options would not have effect"); } _createOptions = options; return _instance.Value; } public void EnableTelemetryEvents() { NativeApiStatus.VerifySuccess(NativeMethods.OrtEnableTelemetryEvents(Handle)); } public void DisableTelemetryEvents() { NativeApiStatus.VerifySuccess(NativeMethods.OrtDisableTelemetryEvents(Handle)); } public void CreateAndRegisterAllocator(OrtMemoryInfo memInfo, OrtArenaCfg arenaCfg) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateAndRegisterAllocator(Handle, memInfo.Pointer, arenaCfg.Pointer)); } public string GetVersionString() { IntPtr nativeUtf = NativeMethods.OrtGetVersionString(); return NativeOnnxValueHelper.StringFromNativeUtf8(nativeUtf); } public string[] GetAvailableProviders() { IntPtr providers = IntPtr.Zero; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetAvailableProviders(out providers, out var numProviders)); try { string[] array = new string[numProviders]; for (int i = 0; i < numProviders; i++) { array[i] = NativeOnnxValueHelper.StringFromNativeUtf8(Marshal.ReadIntPtr(providers, IntPtr.Size * i)); } return array; } finally { NativeApiStatus.VerifySuccess(NativeMethods.OrtReleaseAvailableProviders(providers, numProviders)); } } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseEnv(handle); handle = IntPtr.Zero; _instance = new Lazy(CreateInstance); return true; } } internal class BitOpsUtils { internal const uint SingleBiasedExponentMask = 2139095040u; internal const int SingleBiasedExponentShift = 23; internal const uint SingleSignMask = 2147483648u; internal const int SingleSignShift = 31; internal const uint SingleMostSignificantSigBit = 4194304u; internal const uint SingleTrailingSignificandMask = 8388607u; internal static int LeadingZeroCount(uint num) { if (num == 0) { return 32; } int num2 = 0; while ((num & 0xF0000000u) == 0) { num2 += 4; num <<= 4; } while ((num & 0x80000000u) == 0) { num2++; num <<= 1; } return num2; } internal unsafe static uint SingleToUInt32Bits(float single) { uint result = default(uint); Buffer.MemoryCopy(&single, &result, 4L, 4L); return result; } internal unsafe static float UInt32BitsToSingle(uint singleBits) { float result = default(float); Buffer.MemoryCopy(&singleBits, &result, 4L, 4L); return result; } internal static ushort SingleBitsToBFloat16Bits(uint singleBits) { if (!BitConverter.IsLittleEndian) { return (ushort)(singleBits & 0xFFFF); } return (ushort)(singleBits >> 16); } internal static uint BFloat16BitsToSingleBits(ushort bfloatBits) { if (!BitConverter.IsLittleEndian) { return bfloatBits; } return (uint)(bfloatBits << 16); } internal static float CreateSingleNaN(bool sign, ulong significand) { uint num = (uint)((sign ? 1 : 0) << 31); uint num2 = (uint)(significand >> 41); uint singleBits = num | 0x7FC00000 | num2; return UInt32BitsToSingle(singleBits); } internal static float CreateSingle(bool sign, byte exponent, uint significand) { uint num = (uint)((sign ? 1 : 0) << 31); uint num2 = (uint)(exponent << 23) + significand; uint singleBits = num + num2; return UInt32BitsToSingle(singleBits); } } public readonly struct Float16 : IComparable, IComparable, IEquatable { internal const ushort SignMask = 32768; internal const int SignShift = 15; internal const byte ShiftedSignMask = 1; internal const ushort BiasedExponentMask = 31744; internal const int BiasedExponentShift = 10; internal const byte ShiftedBiasedExponentMask = 31; internal const ushort TrailingSignificandMask = 1023; internal const byte MinSign = 0; internal const byte MaxSign = 1; internal const byte MinBiasedExponent = 0; internal const byte MaxBiasedExponent = 31; internal const byte ExponentBias = 15; internal const sbyte MinExponent = -14; internal const sbyte MaxExponent = 15; private const ushort PositiveZeroBits = 0; private const ushort NegativeZeroBits = 32768; private const ushort OneBits = 15360; private const ushort EpsilonBits = 1024; private const ushort PositiveInfinityBits = 31744; private const ushort NegativeInfinityBits = 64512; private const ushort PositiveQNaNBits = 32256; private const ushort NegativeQNaNBits = 65024; private const ushort MinValueBits = 64511; private const ushort MaxValueBits = 31743; private const ushort PositiveOneBits = 15360; private const ushort NegativeOneBits = 48128; private const ushort EBits = 16752; private const ushort PiBits = 16968; private const ushort TauBits = 17992; public readonly ushort value; public static Float16 Epsilon => new Float16(1024); public static Float16 Pi => new Float16(16968); public static Float16 PositiveInfinity => new Float16(31744); public static Float16 NegativeInfinity => new Float16(64512); public static Float16 NaN => new Float16(65024); public static Float16 Zero => new Float16(0); public static Float16 One => new Float16(15360); public static Float16 NegativeZero => new Float16(32768); public static Float16 MinValue => new Float16(64511); public static Float16 MaxValue => new Float16(31743); internal byte BiasedExponent { get { ushort bits = value; return ExtractBiasedExponentFromBits(bits); } } internal sbyte Exponent => (sbyte)(BiasedExponent - 15); internal ushort Significand => (ushort)(TrailingSignificand | ((BiasedExponent != 0) ? 1024 : 0)); internal ushort TrailingSignificand { get { ushort bits = value; return ExtractTrailingSignificandFromBits(bits); } } public Float16(ushort v) { value = v; } private Float16(bool sign, ushort exp, ushort sig) { value = (ushort)(((sign ? 1 : 0) << 15) + (exp << 10) + sig); } internal static byte ExtractBiasedExponentFromBits(ushort bits) { return (byte)((bits >> 10) & 0x1F); } internal static ushort ExtractTrailingSignificandFromBits(ushort bits) { return (ushort)(bits & 0x3FF); } public static bool operator <(Float16 left, Float16 right) { if (IsNaN(left) || IsNaN(right)) { return false; } bool flag = IsNegative(left); if (flag != IsNegative(right)) { return flag && !AreZero(left, right); } return left.value != right.value && ((left.value < right.value) ^ flag); } public static bool operator >(Float16 left, Float16 right) { return right < left; } public static bool operator <=(Float16 left, Float16 right) { if (IsNaN(left) || IsNaN(right)) { return false; } bool flag = IsNegative(left); if (flag != IsNegative(right)) { return flag || AreZero(left, right); } return left.value == right.value || ((left.value < right.value) ^ flag); } public static bool operator >=(Float16 left, Float16 right) { return right <= left; } public static bool operator ==(Float16 left, Float16 right) { if (IsNaN(left) || IsNaN(right)) { return false; } return left.value == right.value; } public static bool operator !=(Float16 left, Float16 right) { return !(left == right); } public static bool IsFinite(Float16 value) { return StripSign(value) < 31744; } public static bool IsInfinity(Float16 value) { return StripSign(value) == 31744; } public static bool IsNaN(Float16 value) { return StripSign(value) > 31744; } public static bool IsNegative(Float16 value) { return (short)value.value < 0; } public static bool IsNegativeInfinity(Float16 value) { return value.value == 64512; } public static bool IsNormal(Float16 value) { uint num = StripSign(value); return num < 31744 && num != 0 && (num & 0x7C00) != 0; } public static bool IsPositiveInfinity(Float16 value) { return value.value == 31744; } public static bool IsSubnormal(Float16 value) { uint num = StripSign(value); return num < 31744 && num != 0 && (num & 0x7C00) == 0; } public int CompareTo(object obj) { if (!(obj is Float16)) { if (obj != null) { throw new ArgumentException("Object must be of type Float16"); } return 1; } return CompareTo((Float16)obj); } public int CompareTo(Float16 other) { if (this < other) { return -1; } if (this > other) { return 1; } if (this == other) { return 0; } if (IsNaN(this)) { return (!IsNaN(other)) ? (-1) : 0; } return 1; } public bool Equals(Float16 other) { return value == other.value || AreZero(this, other) || (IsNaN(this) && IsNaN(other)); } public override bool Equals(object obj) { return obj is Float16 other && Equals(other); } public override int GetHashCode() { if (IsNaNOrZero(this)) { return value & 0x7C00; } return value; } public override string ToString() { return $"{value} : {ToFloat()}"; } public float ToFloat() { return (float)this; } public static explicit operator Float16(float value) { uint num = BitOpsUtils.SingleToUInt32Bits(value); bool flag = (num & 0x80000000u) >> 31 != 0; int num2 = (int)(num & 0x7F800000) >> 23; uint num3 = num & 0x7FFFFF; if (num2 == 255) { if (num3 != 0) { return CreateFloat16NaN(flag, (ulong)num3 << 41); } return flag ? NegativeInfinity : PositiveInfinity; } uint num4 = (num3 >> 9) | (uint)(((num3 & 0x1FF) != 0) ? 1 : 0); if (((uint)num2 | num4) == 0) { return new Float16(flag, 0, 0); } return new Float16(RoundPackToFloat16(flag, (short)(num2 - 113), (ushort)(num4 | 0x4000))); } public static explicit operator float(Float16 value) { bool flag = IsNegative(value); int num = value.BiasedExponent; uint num2 = value.TrailingSignificand; switch (num) { case 31: if (num2 != 0) { return BitOpsUtils.CreateSingleNaN(flag, (ulong)num2 << 54); } return flag ? float.NegativeInfinity : float.PositiveInfinity; case 0: { if (num2 == 0) { return flag ? -0f : 0f; } (int Exp, uint Sig) tuple = NormSubnormalF16Sig(num2); num = tuple.Exp; num2 = tuple.Sig; num--; break; } } return BitOpsUtils.CreateSingle(flag, (byte)(num + 112), num2 << 13); } public static Float16 Negate(Float16 value) { return IsNaN(value) ? value : new Float16((ushort)(value.value ^ 0x8000)); } private static bool AreZero(Float16 left, Float16 right) { return (ushort)((left.value | right.value) & -32769) == 0; } public static bool IsNaNOrZero(Float16 value) { uint num = StripSign(value); return num == 0 || num > 31744; } private static uint StripSign(Float16 value) { return (ushort)(value.value & -32769); } private static (int Exp, uint Sig) NormSubnormalF16Sig(uint sig) { int num = BitOpsUtils.LeadingZeroCount(sig) - 16 - 5; return (Exp: 1 - num, Sig: sig << num); } private static Float16 CreateFloat16NaN(bool sign, ulong significand) { uint num = (uint)((sign ? 1 : 0) << 15); ushort num2 = (ushort)(significand >> 54); ushort v = (ushort)(num | 0x7E00 | num2); return new Float16(v); } private static ushort RoundPackToFloat16(bool sign, short exp, ushort sig) { int num = sig & 0xF; if ((uint)exp >= 29u) { if (exp < 0) { sig = (ushort)ShiftRightJam(sig, -exp); exp = 0; num = sig & 0xF; } else if (exp > 29 || sig + 8 >= 32768) { return (ushort)(sign ? 64512 : 31744); } } sig = (ushort)(sig + 8 >> 4); sig &= (ushort)(~((((num ^ 8) == 0) ? 1u : 0u) & 1u)); if (sig == 0) { exp = 0; } return new Float16(sign, (ushort)exp, sig).value; } private static uint ShiftRightJam(uint i, int dist) { return (dist < 31) ? ((i >> dist) | (uint)((i << -dist != 0) ? 1 : 0)) : ((i != 0) ? 1u : 0u); } private static ulong ShiftRightJam(ulong l, int dist) { return (dist < 63) ? ((l >> dist) | (ulong)((l << -dist != 0L) ? 1 : 0)) : ((ulong)((l != 0L) ? 1 : 0)); } } public readonly struct BFloat16 : IComparable, IComparable, IEquatable { internal const ushort SignMask = 32768; internal const int SignShift = 15; internal const byte ShiftedSignMask = 1; internal const ushort BiasedExponentMask = 32640; internal const int BiasedExponentShift = 7; internal const byte ShiftedBiasedExponentMask = byte.MaxValue; internal const ushort TrailingSignificandMask = 127; internal const byte MinSign = 0; internal const byte MaxSign = 1; internal const byte MinBiasedExponent = 0; internal const byte MaxBiasedExponent = byte.MaxValue; internal const byte ExponentBias = 127; internal const sbyte MinExponent = -126; internal const sbyte MaxExponent = sbyte.MaxValue; private const ushort PositiveZeroBits = 0; private const ushort NegativeZeroBits = 32768; private const ushort OneBits = 16256; private const ushort PositiveInfinityBits = 32640; private const ushort NegativeInfinityBits = 65408; private const ushort PositiveQNaNBits = 32705; private const ushort NegativeQNaNBits = 65473; private const ushort MinValueBits = 65407; private const ushort MaxValueBits = 32639; private const ushort EpsilonBits = 128; private const ushort PiBits = 16457; private const uint RoundingBase = 32767u; public readonly ushort value; public static BFloat16 Epsilon => new BFloat16(128); public static BFloat16 Pi => new BFloat16(16457); public static BFloat16 PositiveInfinity => new BFloat16(32640); public static BFloat16 NegativeInfinity => new BFloat16(65408); public static BFloat16 NaN => new BFloat16(65473); public static BFloat16 Zero => new BFloat16(0); public static BFloat16 One => new BFloat16(16256); public static BFloat16 NegativeZero => new BFloat16(32768); public static BFloat16 MinValue => new BFloat16(65407); public static BFloat16 MaxValue => new BFloat16(32639); internal byte BiasedExponent { get { ushort bits = value; return ExtractBiasedExponentFromBits(bits); } } internal ushort TrailingSignificand { get { ushort bits = value; return ExtractTrailingSignificandFromBits(bits); } } public BFloat16(ushort v) { value = v; } internal static byte ExtractBiasedExponentFromBits(ushort bits) { return (byte)((bits >> 7) & 0xFF); } internal static ushort ExtractTrailingSignificandFromBits(ushort bits) { return (ushort)(bits & 0x7F); } public static bool operator <(BFloat16 left, BFloat16 right) { if (IsNaN(left) || IsNaN(right)) { return false; } bool flag = IsNegative(left); if (flag != IsNegative(right)) { return flag && !AreZero(left, right); } return left.value != right.value && ((left.value < right.value) ^ flag); } public static bool operator >(BFloat16 left, BFloat16 right) { return right < left; } public static bool operator <=(BFloat16 left, BFloat16 right) { if (IsNaN(left) || IsNaN(right)) { return false; } bool flag = IsNegative(left); if (flag != IsNegative(right)) { return flag || AreZero(left, right); } return left.value == right.value || ((left.value < right.value) ^ flag); } public static bool operator >=(BFloat16 left, BFloat16 right) { return right <= left; } public static bool operator ==(BFloat16 left, BFloat16 right) { if (IsNaN(left) || IsNaN(right)) { return false; } return left.value == right.value; } public static bool operator !=(BFloat16 left, BFloat16 right) { return !(left == right); } public static bool IsFinite(BFloat16 value) { return StripSign(value) < 32640; } public static bool IsInfinity(BFloat16 value) { return StripSign(value) == 32640; } public static bool IsNaN(BFloat16 value) { return StripSign(value) > 32640; } public static bool IsNegative(BFloat16 value) { return (short)value.value < 0; } public static bool IsNegativeInfinity(BFloat16 value) { return value.value == 65408; } public static bool IsNormal(BFloat16 value) { uint num = StripSign(value); return num < 32640 && num != 0 && (num & 0x7F80) != 0; } public static bool IsPositiveInfinity(BFloat16 value) { return value.value == 32640; } public static bool IsSubnormal(BFloat16 value) { uint num = StripSign(value); return num < 32640 && num != 0 && (num & 0x7F80) == 0; } public int CompareTo(object obj) { if (!(obj is BFloat16)) { if (obj != null) { throw new ArgumentException("Object must be of type BFloat16"); } return 1; } return CompareTo((BFloat16)obj); } public int CompareTo(BFloat16 other) { if (this < other) { return -1; } if (this > other) { return 1; } if (this == other) { return 0; } if (IsNaN(this)) { return (!IsNaN(other)) ? (-1) : 0; } return 1; } public bool Equals(BFloat16 other) { return value == other.value || AreZero(this, other) || (IsNaN(this) && IsNaN(other)); } public override bool Equals(object obj) { return obj is BFloat16 other && Equals(other); } public override int GetHashCode() { if (IsNaNOrZero(this)) { return value & 0x7F80; } return value; } public override string ToString() { return $"{value} : {ToFloat()}"; } public float ToFloat() { return (float)this; } public static explicit operator BFloat16(float value) { if (float.IsNaN(value)) { return NaN; } uint num = BitOpsUtils.SingleToUInt32Bits(value); ushort num2 = BitOpsUtils.SingleBitsToBFloat16Bits(num); num += (uint)((num2 & 1) + 32767); num2 = BitOpsUtils.SingleBitsToBFloat16Bits(num); return new BFloat16(num2); } public static explicit operator float(BFloat16 value) { bool flag = IsNegative(value); int biasedExponent = value.BiasedExponent; uint trailingSignificand = value.TrailingSignificand; switch (biasedExponent) { case 255: if (trailingSignificand != 0) { return BitOpsUtils.CreateSingleNaN(flag, (ulong)trailingSignificand << 56); } return flag ? float.NegativeInfinity : float.PositiveInfinity; case 0: if (trailingSignificand == 0) { return flag ? -0f : 0f; } break; } uint singleBits = BitOpsUtils.BFloat16BitsToSingleBits(value.value); return BitOpsUtils.UInt32BitsToSingle(singleBits); } public static BFloat16 Negate(BFloat16 value) { return IsNaN(value) ? value : new BFloat16((ushort)(value.value ^ 0x8000)); } public static bool IsNaNOrZero(BFloat16 value) { uint num = StripSign(value); return num == 0 || num > 32640; } private static bool AreZero(BFloat16 left, BFloat16 right) { return (ushort)((left.value | right.value) & -32769) == 0; } private static uint StripSign(BFloat16 value) { return (ushort)(value.value & -32769); } } public class OrtIoBinding : SafeHandle { internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; internal OrtIoBinding(InferenceSession session) : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateIoBinding(session.Handle, out handle)); } public void BindInput(string name, OrtValue ortValue) { BindInputOrOutput(name, ortValue.Handle, isInput: true); } public void BindInput(string name, TensorElementType elementType, long[] shape, OrtMemoryAllocation allocation) { BindOrtAllocation(name, elementType, shape, allocation, isInput: true); } [Obsolete("This BindInput overload is deprecated. Create OrtValue over an arbitrary piece of memory.")] public void BindInput(string name, OrtExternalAllocation allocation) { BindExternalAllocation(name, allocation, isInput: true); } [Obsolete("This BindInput overload is deprecated. Use of OrtValue based overload is recommended.")] public void BindInput(string name, FixedBufferOnnxValue fixedValue) { if (fixedValue.OnnxValueType != OnnxValueType.ONNX_TYPE_TENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Binding works only with Tensors"); } BindInputOrOutput(name, fixedValue.Value.Handle, isInput: true); } public void SynchronizeBoundInputs() { NativeApiStatus.VerifySuccess(NativeMethods.OrtSynchronizeBoundInputs(handle)); } public void BindOutput(string name, OrtValue ortValue) { BindInputOrOutput(name, ortValue.Handle, isInput: false); } public void BindOutput(string name, TensorElementType elementType, long[] shape, OrtMemoryAllocation allocation) { BindOrtAllocation(name, elementType, shape, allocation, isInput: false); } [Obsolete("This BindOutput overload is deprecated. Create OrtValue over an arbitrary piece of memory.")] public void BindOutput(string name, OrtExternalAllocation allocation) { BindExternalAllocation(name, allocation, isInput: false); } [Obsolete("This BindOutput overload is deprecated. Use of OrtValue based overload is recommended.")] public void BindOutput(string name, FixedBufferOnnxValue fixedValue) { if (fixedValue.OnnxValueType != OnnxValueType.ONNX_TYPE_TENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Binding works only with Tensors"); } BindInputOrOutput(name, fixedValue.Value.Handle, isInput: false); } public void BindOutputToDevice(string name, OrtMemoryInfo memInfo) { byte[] name2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(name); NativeApiStatus.VerifySuccess(NativeMethods.OrtBindOutputToDevice(handle, name2, memInfo.Pointer)); } public void SynchronizeBoundOutputs() { NativeApiStatus.VerifySuccess(NativeMethods.OrtSynchronizeBoundOutputs(handle)); } private void BindOrtAllocation(string name, TensorElementType elementType, long[] shape, OrtMemoryAllocation allocation, bool isInput) { using OrtValue ortValue = OrtValue.CreateTensorValueWithData(allocation.Info, elementType, shape, allocation.Pointer, allocation.Size); BindInputOrOutput(name, ortValue.Handle, isInput); } private void BindExternalAllocation(string name, OrtExternalAllocation allocation, bool isInput) { using OrtValue ortValue = OrtValue.CreateTensorValueWithData(allocation.Info, allocation.ElementType, allocation.Shape, allocation.Pointer, allocation.Size); BindInputOrOutput(name, ortValue.Handle, isInput); } private void BindInputOrOutput(string name, IntPtr ortValue, bool isInput) { byte[] name2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(name); if (isInput) { NativeApiStatus.VerifySuccess(NativeMethods.OrtBindInput(handle, name2, ortValue)); } else { NativeApiStatus.VerifySuccess(NativeMethods.OrtBindOutput(handle, name2, ortValue)); } } public unsafe string[] GetOutputNames() { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetBoundOutputNames(handle, defaultInstance.Pointer, out var buffer, out var lengths, out var count)); if ((ulong)count == 0) { return Array.Empty(); } int num = (int)(uint)count; Span span = new Span(lengths.ToPointer(), num); try { string[] array = new string[num]; int num2 = 0; for (int i = 0; i < num; i++) { int num3 = (int)span[i]; IntPtr intPtr = new IntPtr(buffer.ToInt64() + num2); array[i] = Encoding.UTF8.GetString((byte*)intPtr.ToPointer(), num3); num2 += num3; } return array; } finally { defaultInstance.FreeMemory(lengths); defaultInstance.FreeMemory(buffer); } } public IDisposableReadOnlyCollection GetOutputValues() { OrtValue[] outputOrtValues = GetOutputOrtValues(); return new DisposableList(outputOrtValues); } internal unsafe OrtValue[] GetOutputOrtValues() { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetBoundOutputValues(handle, defaultInstance.Pointer, out var ortvalues, out var count)); if ((ulong)count == 0) { return Array.Empty(); } int num = (int)(uint)count; Span span = new Span(ortvalues.ToPointer(), num); try { OrtValue[] array = new OrtValue[num]; for (int i = 0; i < num; i++) { array[i] = new OrtValue(span[i]); } return array; } catch (Exception) { for (int j = 0; j < span.Length; j++) { NativeMethods.OrtReleaseValue(span[j]); } throw; } finally { defaultInstance.FreeMemory(ortvalues); } } public void ClearBoundInputs() { NativeMethods.OrtClearBoundInputs(handle); } public void ClearBoundOutputs() { NativeMethods.OrtClearBoundOutputs(handle); } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseIoBinding(handle); handle = IntPtr.Zero; return true; } } public enum OrtLoggingLevel { ORT_LOGGING_LEVEL_VERBOSE, ORT_LOGGING_LEVEL_INFO, ORT_LOGGING_LEVEL_WARNING, ORT_LOGGING_LEVEL_ERROR, ORT_LOGGING_LEVEL_FATAL } public class OrtThreadingOptions : SafeHandle { internal IntPtr Handle => handle; public int GlobalInterOpNumThreads { set { NativeApiStatus.VerifySuccess(NativeMethods.OrtThreadingOptionsSetGlobalInterOpNumThreads(handle, value)); } } public int GlobalIntraOpNumThreads { set { NativeApiStatus.VerifySuccess(NativeMethods.OrtThreadingOptionsSetGlobalIntraOpNumThreads(handle, value)); } } public bool GlobalSpinControl { set { NativeApiStatus.VerifySuccess(NativeMethods.OrtThreadingOptionsSetGlobalSpinControl(handle, value ? 1 : 0)); } } public override bool IsInvalid => handle == IntPtr.Zero; public OrtThreadingOptions() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateThreadingOptions(out handle)); } public void SetGlobalDenormalAsZero() { NativeApiStatus.VerifySuccess(NativeMethods.OrtThreadingOptionsSetGlobalDenormalAsZero(handle)); } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseThreadingOptions(handle); handle = IntPtr.Zero; return true; } } public class OrtTypeInfo { private OrtTensorTypeAndShapeInfo? _tensorTypeAndShape; private OrtSequenceOrOptionalTypeInfo? _sequenceOrOptional; private OrtMapTypeInfo? _mapTypeInfo; public OnnxValueType OnnxType { get; private set; } public OrtTensorTypeAndShapeInfo TensorTypeAndShapeInfo { get { if (OnnxType != OnnxValueType.ONNX_TYPE_TENSOR && OnnxType != OnnxValueType.ONNX_TYPE_SPARSETENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo does not represent a tensor/sparsetensor"); } return _tensorTypeAndShape.Value; } } public OrtSequenceOrOptionalTypeInfo SequenceTypeInfo { get { if (OnnxType != OnnxValueType.ONNX_TYPE_SEQUENCE) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo does not represent a sequence"); } return _sequenceOrOptional.Value; } } public OrtMapTypeInfo MapTypeInfo { get { if (OnnxType != OnnxValueType.ONNX_TYPE_MAP) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo does not represent a map"); } return _mapTypeInfo.Value; } } public OrtSequenceOrOptionalTypeInfo OptionalTypeInfo { get { if (OnnxType != OnnxValueType.ONNX_TYPE_OPTIONAL) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo does not represent a optional"); } return _sequenceOrOptional.Value; } } internal OrtTypeInfo(IntPtr handle) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetOnnxTypeFromTypeInfo(handle, out var onnxtype)); OnnxType = (OnnxValueType)(int)onnxtype; switch (OnnxType) { case OnnxValueType.ONNX_TYPE_TENSOR: case OnnxValueType.ONNX_TYPE_SPARSETENSOR: { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToTensorInfo(handle, out var typeAndShapeInfo)); if (typeAndShapeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.Fail, "Type Information indicates a tensor, but casting to TensorInfo fails"); } _tensorTypeAndShape = new OrtTensorTypeAndShapeInfo(typeAndShapeInfo); break; } case OnnxValueType.ONNX_TYPE_SEQUENCE: { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToSequenceTypeInfo(handle, out var sequenceTypeInfo)); if (sequenceTypeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo cast to SequenceTypeInfo failed. The object does not represent a sequence"); } _sequenceOrOptional = new OrtSequenceOrOptionalTypeInfo(sequenceTypeInfo); break; } case OnnxValueType.ONNX_TYPE_MAP: { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToMapTypeInfo(handle, out var mapTypeInfo)); if (mapTypeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo cast to MapTypeInfo failed. The object does not represent a map"); } _mapTypeInfo = new OrtMapTypeInfo(mapTypeInfo); break; } case OnnxValueType.ONNX_TYPE_OPTIONAL: { NativeApiStatus.VerifySuccess(NativeMethods.OrtCastTypeInfoToOptionalTypeInfo(handle, out var optionalTypeInfo)); if (optionalTypeInfo == IntPtr.Zero) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "TypeInfo cast to OptionalTypeInfo failed. The object does not represent a optional"); } _sequenceOrOptional = new OrtSequenceOrOptionalTypeInfo(optionalTypeInfo); break; } default: throw new OnnxRuntimeException(ErrorCode.NotImplemented, $"OnnxValueType: {OnnxType} is not supported here"); } } } public struct OrtTensorTypeAndShapeInfo { public TensorElementType ElementDataType { get; private set; } public bool IsString => ElementDataType == TensorElementType.String; public long ElementCount { get; private set; } public int DimensionsCount => Shape.Length; public long[] Shape { get; private set; } internal OrtTensorTypeAndShapeInfo(IntPtr handle) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorElementType(handle, out var output)); ElementDataType = (TensorElementType)(int)output; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorShapeElementCount(handle, out var output2)); ElementCount = (long)(ulong)output2; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetDimensionsCount(handle, out var output3)); Shape = new long[(uint)output3]; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetDimensions(handle, Shape, output3)); } } public struct OrtSequenceOrOptionalTypeInfo { public OrtTypeInfo ElementType { get; private set; } internal OrtSequenceOrOptionalTypeInfo(IntPtr handle) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetSequenceElementType(handle, out var elementTypeInfo)); try { ElementType = new OrtTypeInfo(elementTypeInfo); } finally { NativeMethods.OrtReleaseTypeInfo(elementTypeInfo); } } } public struct OrtMapTypeInfo { public TensorElementType KeyType { get; private set; } public OrtTypeInfo ValueType { get; private set; } internal OrtMapTypeInfo(IntPtr handle) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetMapKeyType(handle, out var tensorElementType)); KeyType = (TensorElementType)(int)tensorElementType; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetMapValueType(handle, out var type_info)); try { ValueType = new OrtTypeInfo(type_info); } finally { NativeMethods.OrtReleaseTypeInfo(type_info); } } } public enum OnnxValueType { ONNX_TYPE_UNKNOWN, ONNX_TYPE_TENSOR, ONNX_TYPE_SEQUENCE, ONNX_TYPE_MAP, ONNX_TYPE_OPAQUE, ONNX_TYPE_SPARSETENSOR, ONNX_TYPE_OPTIONAL } public class OrtValue : IOrtValueOwner, IDisposable { public delegate void SequenceElementVisitor(OrtValue ortValue, int index); public delegate void MapVisitor(OrtValue keys, OrtValue values); private DisposableList _compositeMembers; private IntPtr _handle; private MemoryHandle? _memHandle; private bool _disposed; internal IntPtr Handle => _handle; public OrtValue Value => this; public OnnxValueType OnnxType { get; private set; } public bool IsTensor => OnnxType == OnnxValueType.ONNX_TYPE_TENSOR; public bool IsSparseTensor => OnnxType == OnnxValueType.ONNX_TYPE_SPARSETENSOR; internal OrtValue(IntPtr handle) { _handle = handle; InitOnnxType(); } internal OrtValue(IntPtr handle, OnnxValueType onnxValueType) { if (onnxValueType == OnnxValueType.ONNX_TYPE_UNKNOWN) { throw new ArgumentException("onnxValueType argument is passed as unknown"); } _handle = handle; OnnxType = onnxValueType; } internal OrtValue(IntPtr handle, OnnxValueType onnxValueType, ref DisposableList compositeMembers) { if (onnxValueType == OnnxValueType.ONNX_TYPE_UNKNOWN) { throw new ArgumentException("onnxValueType argument is passed as unknown"); } _handle = handle; OnnxType = onnxValueType; _compositeMembers = compositeMembers; compositeMembers = null; } private OrtValue(IntPtr handle, MemoryHandle memHandle) { _handle = handle; _memHandle = memHandle; OnnxType = OnnxValueType.ONNX_TYPE_TENSOR; } private void InitOnnxType() { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetValueType(Handle, out var onnxtype)); OnnxType = (OnnxValueType)(int)onnxtype; } public int GetValueCount() { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetValueCount(Handle, out var count)); return (int)count; } public OrtValue GetValue(int index, OrtAllocator allocator) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetValue(Handle, index, allocator.Pointer, out var outputValue)); return new OrtValue(outputValue); } public ReadOnlySpan GetTensorDataAsSpan() where T : unmanaged { Span tensorBufferRawData = GetTensorBufferRawData(typeof(T)); return MemoryMarshal.Cast(tensorBufferRawData); } public Span GetTensorMutableDataAsSpan() where T : unmanaged { Span tensorBufferRawData = GetTensorBufferRawData(typeof(T)); return MemoryMarshal.Cast(tensorBufferRawData); } public Span GetTensorMutableRawData() { return GetTensorBufferRawData(typeof(byte)); } public ReadOnlyMemory GetStringElementAsMemory(int index) { char[] stringTensorElementChars = GetStringTensorElementChars(index); if (stringTensorElementChars.Length == 0) { return ReadOnlyMemory.Empty; } return new ReadOnlyMemory(stringTensorElementChars); } public unsafe string GetStringElement(int index) { GetStringTensorElementBuffer((UIntPtr)(ulong)index, out var bytesLen, out var bufferPtr); if (bytesLen == 0) { return string.Empty; } return Encoding.UTF8.GetString((byte*)bufferPtr.ToPointer(), (int)bytesLen); } public unsafe ReadOnlySpan GetStringElementAsSpan(int index) { GetStringTensorElementBuffer((UIntPtr)(ulong)index, out var bytesLen, out var bufferPtr); if (bytesLen == 0) { return ReadOnlySpan.Empty; } return new ReadOnlySpan(bufferPtr.ToPointer(), (int)bytesLen); } public string[] GetStringTensorAsArray() { GetTensorElementTypeAndCount(out var count, out var elementType); if (elementType != TensorElementType.String) { throw new OnnxRuntimeException(ErrorCode.Fail, $"GetStringTensorAsArray() is only supported for string tensors. This OrtValue contains a {elementType} tensor."); } string[] array = new string[count]; for (int i = 0; i < count; i++) { array[i] = GetStringElement(i); } return array; } public OrtTypeInfo GetTypeInfo() { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTypeInfo(Handle, out var typeInfo)); try { return new OrtTypeInfo(typeInfo); } finally { NativeMethods.OrtReleaseTypeInfo(typeInfo); } } public OrtTensorTypeAndShapeInfo GetTensorTypeAndShape() { OnnxValueType onnxType = OnnxType; if (onnxType != OnnxValueType.ONNX_TYPE_TENSOR && onnxType != OnnxValueType.ONNX_TYPE_SPARSETENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"This OrtValue type contains: {onnxType}, not a tensor or sparse tensor"); } NativeMethods.OrtGetTensorTypeAndShape(Handle, out var typeAndShapeInfo); try { return new OrtTensorTypeAndShapeInfo(typeAndShapeInfo); } finally { NativeMethods.OrtReleaseTensorTypeAndShapeInfo(typeAndShapeInfo); } } public OrtMemoryInfo GetTensorMemoryInfo() { OnnxValueType onnxType = OnnxType; if (onnxType != OnnxValueType.ONNX_TYPE_TENSOR && onnxType != OnnxValueType.ONNX_TYPE_SPARSETENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"This OrtValue type contains: {onnxType}, not a tensor or sparse tensor"); } NativeMethods.OrtGetTensorMemoryInfo(Handle, out var ortMemoryInfo); return new OrtMemoryInfo(ortMemoryInfo, owned: false); } private void GetTensorElementTypeAndCount(out long count, out TensorElementType elementType) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorTypeAndShape(Handle, out var typeAndShapeInfo)); try { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorElementType(typeAndShapeInfo, out var output)); NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorShapeElementCount(typeAndShapeInfo, out var output2)); elementType = (TensorElementType)(int)output; count = (long)(ulong)output2; } finally { NativeMethods.OrtReleaseTensorTypeAndShapeInfo(typeAndShapeInfo); } } private unsafe char[] GetStringTensorElementChars(int index) { GetStringTensorElementBuffer((UIntPtr)(ulong)index, out var bytesLen, out var bufferPtr); if (bytesLen == 0) { return Array.Empty(); } int charCount = Encoding.UTF8.GetCharCount((byte*)bufferPtr.ToPointer(), (int)bytesLen); char[] array = new char[charCount]; fixed (char* chars = array) { Encoding.UTF8.GetChars((byte*)bufferPtr.ToPointer(), (int)bytesLen, chars, charCount); } return array; } private void GetStringTensorElementBuffer(UIntPtr index, out uint bytesLen, out IntPtr bufferPtr) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetStringTensorElementLength(Handle, index, out var len)); bytesLen = (uint)len; if (bytesLen == 0) { bufferPtr = IntPtr.Zero; } else { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetResizedStringTensorElementBuffer(Handle, index, len, out bufferPtr)); } } private unsafe Span GetTensorBufferRawData(Type requestedType) { if (OnnxType != OnnxValueType.ONNX_TYPE_TENSOR) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"This OrtValue type contains: {OnnxType}, not a tensor"); } GetTensorElementTypeAndCount(out var count, out var elementType); if (elementType == TensorElementType.String) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Strings are not supported by this API"); } TensorElementTypeInfo tensorElementTypeInfo = TensorBase.GetElementTypeInfo(elementType) ?? throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Element type: {elementType} is not registered type."); if (requestedType != typeof(byte) && requestedType != tensorElementTypeInfo.TensorType) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Requested type: {requestedType} does not match the actual type: {tensorElementTypeInfo.TensorType}"); } if (count == 0) { return Span.Empty; } NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorMutableData(Handle, out var dataBufferHandle)); long num = count * tensorElementTypeInfo.TypeSize; return new Span(dataBufferHandle.ToPointer(), (int)num); } public static OrtValue CreateTensorValueWithData(OrtMemoryInfo memInfo, TensorElementType elementType, long[] shape, IntPtr dataBufferPtr, long bufferLengthInBytes) { TensorElementTypeInfo tensorElementTypeInfo = TensorBase.GetElementTypeInfo(elementType) ?? throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Tensor element type: {elementType} is not supported"); if (tensorElementTypeInfo.IsString) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Cannot map managed strings buffer to native OrtValue. Use string specific interfaces"); } long sizeForShape = ShapeUtils.GetSizeForShape(shape); long num = sizeForShape * tensorElementTypeInfo.TypeSize; if (num > bufferLengthInBytes) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Shape: {shape} has: {sizeForShape} elements requires a buffer of at least {num} bytes. Provided: {bufferLengthInBytes} bytes"); } NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateTensorWithDataAsOrtValue(memInfo.Pointer, dataBufferPtr, (UIntPtr)(ulong)bufferLengthInBytes, shape, (UIntPtr)(ulong)shape.Length, elementType, out var outputValue)); return new OrtValue(outputValue, OnnxValueType.ONNX_TYPE_TENSOR); } public unsafe static OrtValue CreateTensorValueFromMemory(OrtMemoryInfo memoryInfo, Memory memory, long[] shape) where T : unmanaged { TensorTypeInfo tensorTypeInfo = TensorBase.GetTypeInfo(typeof(T)) ?? throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Tensor of type: {typeof(T)} is not supported"); if (tensorTypeInfo.IsString) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Cannot map managed strings buffer to native OrtValue. Use string specific interfaces."); } long sizeForShape = ShapeUtils.GetSizeForShape(shape); if (sizeForShape > memory.Length) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Managed memory size: {memory.Length} elements is less than shape size: {sizeForShape} elements"); } int num = memory.Length * tensorTypeInfo.TypeSize; MemoryHandle memHandle = memory.Pin(); try { IntPtr dataBufferHandle = new IntPtr(memHandle.Pointer); NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateTensorWithDataAsOrtValue(memoryInfo.Pointer, dataBufferHandle, (UIntPtr)(ulong)num, shape, (UIntPtr)(ulong)shape.Length, tensorTypeInfo.ElementType, out var outputValue)); return new OrtValue(outputValue, memHandle); } catch (Exception) { memHandle.Dispose(); throw; } } public static OrtValue CreateTensorValueFromMemory(T[] data, long[] shape) where T : unmanaged { return CreateTensorValueFromMemory(OrtMemoryInfo.DefaultInstance, new Memory(data), shape); } public static OrtValue CreateAllocatedTensorValue(OrtAllocator allocator, TensorElementType elementType, long[] shape) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateTensorAsOrtValue(allocator.Pointer, shape, (UIntPtr)(ulong)shape.Length, elementType, out var outputValue)); return new OrtValue(outputValue, OnnxValueType.ONNX_TYPE_TENSOR); } internal unsafe static OrtValue CreateFromTensorObject(TensorBase value, out TensorElementType elementType) { TensorTypeInfo typeInfo = value.GetTypeInfo(); OrtValue ortValue = null; TensorElementType elementType2 = typeInfo.ElementType; int typeSize = typeInfo.TypeSize; if (typeInfo.IsString) { ortValue = CreateFromStringTensor(value as Tensor); } else { MemoryHandle pinnedHandle; int dataBufferLength; long[] shape; int rank; switch (elementType2) { case TensorElementType.Float: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Double: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Int32: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.UInt32: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Int64: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.UInt64: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Int16: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.UInt16: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.UInt8: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Int8: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Bool: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.Float16: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; case TensorElementType.BFloat16: PinAsTensor(value as Tensor, typeSize, out pinnedHandle, out dataBufferLength, out shape, out rank); break; default: throw new NotSupportedException("Element type: " + elementType2.ToString() + " is not of a supported type"); } try { IntPtr zero = IntPtr.Zero; zero = (IntPtr)pinnedHandle.Pointer; NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateTensorWithDataAsOrtValue(OrtMemoryInfo.DefaultInstance.Pointer, zero, (UIntPtr)(ulong)dataBufferLength, shape, (UIntPtr)(ulong)rank, elementType2, out var outputValue)); ortValue = new OrtValue(outputValue, pinnedHandle); } catch (Exception) { pinnedHandle.Dispose(); throw; } } elementType = elementType2; return ortValue; } public static OrtValue CreateTensorWithEmptyStrings(OrtAllocator allocator, long[] shape) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateTensorAsOrtValue(allocator.Pointer, shape, (UIntPtr)(ulong)shape.Length, TensorElementType.String, out var outputValue)); return new OrtValue(outputValue, OnnxValueType.ONNX_TYPE_TENSOR); } public unsafe void StringTensorSetElementAt(ReadOnlySpan str, int index) { fixed (char* strPtr = str) { FillStringTensorElement(strPtr, str.Length, index); } } public void StringTensorSetElementAt(ReadOnlyMemory rom, int index) { StringTensorSetElementAt(rom.Span, index); } public unsafe void StringTensorSetElementAt(ReadOnlySpan utf8Bytes, int index) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetResizedStringTensorElementBuffer(Handle, (UIntPtr)(ulong)index, (UIntPtr)(ulong)utf8Bytes.Length, out var buffer)); if (utf8Bytes.Length != 0) { Span destination = new Span(buffer.ToPointer(), utf8Bytes.Length); utf8Bytes.CopyTo(destination); } } public unsafe static OrtValue CreateFromStringTensor(Tensor tensor) { if (tensor == null) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "Expecting a valid string tensor"); } long[] shape = Array.ConvertAll(tensor.Dimensions.ToArray(), Convert.ToInt64); OrtValue ortValue = CreateTensorWithEmptyStrings(OrtAllocator.DefaultInstance, shape); try { long length = tensor.Length; for (int i = 0; i < length; i++) { string text = tensor.GetValue(i) ?? throw new ArgumentNullException($"Tensor contains null reference at index:{i}"); fixed (char* strPtr = text) { ortValue.FillStringTensorElement(strPtr, text.Length, i); } } } catch (Exception) { ortValue.Dispose(); throw; } return ortValue; } public static OrtValue CreateSequence(ICollection ortValues) { if (ortValues == null) { throw new ArgumentNullException("ortValues"); } if (ortValues.IsReadOnly) { throw new ArgumentException("ortValues argument can not be a readonly collection"); } DisposableList compositeMembers = new DisposableList(ortValues); try { OrtValue result = CreateSequence(ref compositeMembers); ortValues.Clear(); return result; } catch (Exception) { compositeMembers?.Clear(); throw; } } internal static OrtValue CreateSequence(ref DisposableList compositeMembers) { IntPtr[] array = new IntPtr[compositeMembers.Count]; for (int i = 0; i < compositeMembers.Count; i++) { array[i] = compositeMembers[i].Handle; } NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateValue(array, (UIntPtr)(ulong)array.Length, (IntPtr)2, out var ortValue)); return new OrtValue(ortValue, OnnxValueType.ONNX_TYPE_SEQUENCE, ref compositeMembers); } public void ProcessSequence(SequenceElementVisitor visitor, OrtAllocator allocator) { if (OnnxType != OnnxValueType.ONNX_TYPE_SEQUENCE) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"OrtValue.OnnxType of {OnnxType} is not a sequence"); } int valueCount = GetValueCount(); for (int i = 0; i < valueCount; i++) { using OrtValue ortValue = GetValue(i, allocator); visitor(ortValue, i); } } public static OrtValue CreateMap(ref OrtValue keys, ref OrtValue values) { if (keys == null || values == null) { throw new ArgumentNullException("keys or/and values are null"); } IntPtr[] array = new IntPtr[2] { keys.Handle, values.Handle }; NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateValue(array, (UIntPtr)(ulong)array.Length, (IntPtr)3, out var ortValue)); DisposableList compositeMembers = new DisposableList { keys, values }; keys = null; values = null; return new OrtValue(ortValue, OnnxValueType.ONNX_TYPE_MAP, ref compositeMembers); } public static OrtValue CreateMap(K[] keys, V[] values) where K : unmanaged where V : unmanaged { if (keys == null || values == null) { throw new ArgumentNullException("Keys or/and values are null"); } if (keys.Length != values.Length) { throw new ArgumentException("Expecting keys and values same len. " + $"Received keys: {keys.Length}, Values: {values.Length}"); } long[] shape = new long[1] { keys.Length }; Span disposables = new OrtValue[2]; DisposableArray disposableArray = new DisposableArray(disposables); try { disposables[0] = CreateTensorValueFromMemory(keys, shape); disposables[1] = CreateTensorValueFromMemory(values, shape); return CreateMap(ref disposables[0], ref disposables[1]); } catch (Exception) { disposableArray.Dispose(); throw; } } public static OrtValue CreateMapWithStringKeys(IReadOnlyCollection keys, V[] values) where V : unmanaged { if (keys == null || values == null) { throw new ArgumentNullException("Keys or/and values are null"); } if (keys.Count != values.Length) { throw new ArgumentException("Expecting keys and values same len. " + $"Received keys: {keys.Count}, Values: {values.Length}"); } long[] shape = new long[1] { keys.Count }; Span disposables = new OrtValue[2]; DisposableArray disposableArray = new DisposableArray(disposables); try { disposables[0] = CreateTensorWithEmptyStrings(OrtAllocator.DefaultInstance, shape); int num = 0; foreach (string key in keys) { disposables[0].StringTensorSetElementAt(key.AsSpan(), num++); } disposables[1] = CreateTensorValueFromMemory(values, shape); return CreateMap(ref disposables[0], ref disposables[1]); } catch (Exception) { disposableArray.Dispose(); throw; } } public static OrtValue CreateMapWithStringValues(K[] keys, IReadOnlyCollection values) where K : unmanaged { if (keys == null || values == null) { throw new ArgumentNullException("Keys or/and values are null"); } if (keys.Length != values.Count) { throw new ArgumentException("Expecting keys and values same len. " + $"Received keys: {keys.Length}, Values: {values.Count}"); } long[] shape = new long[1] { keys.Length }; Span disposables = new OrtValue[2]; DisposableArray disposableArray = new DisposableArray(disposables); try { disposables[0] = CreateTensorValueFromMemory(keys, shape); disposables[1] = CreateTensorWithEmptyStrings(OrtAllocator.DefaultInstance, shape); int num = 0; foreach (string value in values) { disposables[1].StringTensorSetElementAt(value.AsSpan(), num++); } return CreateMap(ref disposables[0], ref disposables[1]); } catch (Exception) { disposableArray.Dispose(); throw; } } public void ProcessMap(MapVisitor visitor, OrtAllocator allocator) { if (OnnxType != OnnxValueType.ONNX_TYPE_MAP) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, "This OrtValue does not represent a map"); } using OrtValue keys = GetValue(0, allocator); using OrtValue values = GetValue(1, allocator); visitor(keys, values); } private unsafe void FillStringTensorElement(char* strPtr, int strLength, int index) { IntPtr buffer; if (strLength == 0) { NativeApiStatus.VerifySuccess(NativeMethods.OrtGetResizedStringTensorElementBuffer(Handle, (UIntPtr)(ulong)index, UIntPtr.Zero, out buffer)); return; } int byteCount = Encoding.UTF8.GetByteCount(strPtr, strLength); NativeApiStatus.VerifySuccess(NativeMethods.OrtGetResizedStringTensorElementBuffer(Handle, (UIntPtr)(ulong)index, (UIntPtr)(ulong)byteCount, out buffer)); NativeOnnxValueHelper.StringToUtf8NativeMemory(strPtr, strLength, buffer, byteCount); } private static void PinAsTensor(Tensor tensor, int elementSize, out MemoryHandle pinnedHandle, out int dataBufferLength, out long[] shape, out int rank) { if (tensor == null) { throw new OnnxRuntimeException(ErrorCode.Fail, "Cast to Tensor failed. BUG check!"); } if (tensor.IsReversedStride) { throw new NotSupportedException("Tensor of reverseStride is not supported"); } DenseTensor denseTensor = (tensor as DenseTensor) ?? tensor.ToDenseTensor(); shape = Array.ConvertAll(denseTensor.Dimensions.ToArray(), Convert.ToInt64); rank = denseTensor.Rank; dataBufferLength = denseTensor.Buffer.Length * elementSize; pinnedHandle = denseTensor.Buffer.Pin(); } ~OrtValue() { Dispose(disposing: false); } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } protected virtual void Dispose(bool disposing) { if (!_disposed) { if (disposing) { _memHandle?.Dispose(); _memHandle = null; _compositeMembers?.Dispose(); _compositeMembers = null; } NativeMethods.OrtReleaseValue(_handle); _handle = IntPtr.Zero; _disposed = true; } } } internal interface IOrtValueOwner : IDisposable { OrtValue Value { get; } } internal class NativeOrtValueCollectionOwner : IOrtValueOwner, IDisposable { private OrtValue _ortValue; private IDisposable _disposables; private bool _disposed = false; public OrtValue Value => _ortValue; internal NativeOrtValueCollectionOwner(ref OrtValue ortValue, IDisposable disposables) { _ortValue = ortValue; ortValue = null; _disposables = disposables; } protected virtual void Dispose(bool disposing) { if (!_disposed && disposing) { if (_disposables != null) { _disposables.Dispose(); _disposables = null; } if (_ortValue != null) { _ortValue.Dispose(); _ortValue = null; } _disposed = true; } } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } } internal class OrtValueTensor : MemoryManager, IOrtValueOwner, IDisposable { private OrtValue _ortValue; private readonly IntPtr _dataBufferPointer; public OrtValue Value => _ortValue; public bool IsDisposed { get; private set; } = false; public int[] Dimensions { get; } public int Rank => Dimensions.Length; public int Count { get; } public int ElementWidth { get; } public TensorElementType ElementType { get; } public OrtValueTensor(ref OrtValue ortValue) { OrtTensorTypeAndShapeInfo tensorTypeAndShape = ortValue.GetTensorTypeAndShape(); TensorElementType elementDataType = tensorTypeAndShape.ElementDataType; TensorElementTypeInfo tensorElementTypeInfo = TensorBase.GetElementTypeInfo(elementDataType) ?? throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"Unable to query type information for data type: {elementDataType}"); if (typeof(T) != tensorElementTypeInfo.TensorType) { throw new OnnxRuntimeException(ErrorCode.InvalidArgument, $"The OrtValueTensor type being instantiated for T = [{typeof(T)}] while supplied OrtValue contains T = [{tensorElementTypeInfo.TensorType}]"); } ElementType = elementDataType; ElementWidth = tensorElementTypeInfo.TypeSize; Count = (int)tensorTypeAndShape.ElementCount; Dimensions = Array.ConvertAll(tensorTypeAndShape.Shape, Convert.ToInt32); NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorMutableData(ortValue.Handle, out _dataBufferPointer)); _ortValue = ortValue; ortValue = null; } public unsafe override Span GetSpan() { Span span = null; return new Span((void*)_dataBufferPointer, Count); } public unsafe override MemoryHandle Pin(int elementIndex = 0) { if (elementIndex >= Count) { throw new ArgumentOutOfRangeException("elementIndex"); } IntPtr dataBufferPointer = _dataBufferPointer; return new MemoryHandle(new IntPtr(dataBufferPointer.ToInt64() + (long)elementIndex * (long)ElementWidth).ToPointer()); } public override void Unpin() { } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } protected override void Dispose(bool disposing) { if (!IsDisposed) { if (_ortValue != null) { _ortValue.Dispose(); _ortValue = null; } IsDisposed = true; } } } public class PrePackedWeightsContainer : SafeHandle { internal IntPtr Pointer => handle; public override bool IsInvalid => handle == IntPtr.Zero; public PrePackedWeightsContainer() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreatePrepackedWeightsContainer(out handle)); } protected override bool ReleaseHandle() { NativeMethods.OrtReleasePrepackedWeightsContainer(handle); handle = IntPtr.Zero; return true; } } public class OrtTensorRTProviderOptions : SafeHandle { private int _deviceId = 0; private string _deviceIdStr = "device_id"; internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; public OrtTensorRTProviderOptions() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateTensorRTProviderOptions(out handle)); } public string GetOptions() { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetTensorRTProviderOptionsAsString(handle, defaultInstance.Pointer, out var ptr)); return NativeOnnxValueHelper.StringFromNativeUtf8(ptr, defaultInstance); } private static IntPtr UpdateTRTOptions(IntPtr handle, IntPtr[] keys, IntPtr[] values, UIntPtr count) { return NativeMethods.OrtUpdateTensorRTProviderOptions(handle, keys, values, count); } public void UpdateOptions(Dictionary providerOptions) { ProviderOptionsUpdater.Update(providerOptions, handle, UpdateTRTOptions); if (providerOptions.ContainsKey(_deviceIdStr)) { _deviceId = int.Parse(providerOptions[_deviceIdStr]); } } public int GetDeviceId() { return _deviceId; } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseTensorRTProviderOptions(handle); handle = IntPtr.Zero; return true; } } public class OrtCUDAProviderOptions : SafeHandle { internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; public OrtCUDAProviderOptions() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateCUDAProviderOptions(out handle)); } public string GetOptions() { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetCUDAProviderOptionsAsString(handle, defaultInstance.Pointer, out var ptr)); return NativeOnnxValueHelper.StringFromNativeUtf8(ptr, defaultInstance); } private static IntPtr UpdateCUDAProviderOptions(IntPtr handle, IntPtr[] keys, IntPtr[] values, UIntPtr count) { return NativeMethods.OrtUpdateCUDAProviderOptions(handle, keys, values, count); } public void UpdateOptions(Dictionary providerOptions) { ProviderOptionsUpdater.Update(providerOptions, handle, UpdateCUDAProviderOptions); } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseCUDAProviderOptions(handle); handle = IntPtr.Zero; return true; } } public class OrtROCMProviderOptions : SafeHandle { internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; public OrtROCMProviderOptions() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateROCMProviderOptions(out handle)); } public string GetOptions() { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; NativeApiStatus.VerifySuccess(NativeMethods.OrtGetROCMProviderOptionsAsString(handle, defaultInstance.Pointer, out var ptr)); return NativeOnnxValueHelper.StringFromNativeUtf8(ptr, defaultInstance); } private static IntPtr UpdateROCMProviderOptions(IntPtr handle, IntPtr[] keys, IntPtr[] values, UIntPtr count) { return NativeMethods.OrtUpdateROCMProviderOptions(handle, keys, values, count); } public void UpdateOptions(Dictionary providerOptions) { ProviderOptionsUpdater.Update(providerOptions, handle, UpdateROCMProviderOptions); } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseROCMProviderOptions(handle); handle = IntPtr.Zero; return true; } } public class ProviderOptionsValueHelper { public static void StringToDict(string s, Dictionary dict) { string[] array = s.Split(new char[1] { ';' }); string[] array2 = array; foreach (string text in array2) { string[] array3 = text.Split(new char[1] { '=' }); if (array3.Length != 2) { throw new ArgumentException("Make sure input string contains key-value paris, e.g. key1=value1;key2=value2...", "s"); } dict.Add(array3[0], array3[1]); } } } [Flags] public enum CoreMLFlags : uint { COREML_FLAG_USE_NONE = 0u, COREML_FLAG_USE_CPU_ONLY = 1u, COREML_FLAG_ENABLE_ON_SUBGRAPH = 2u, COREML_FLAG_ONLY_ENABLE_DEVICE_WITH_ANE = 4u, COREML_FLAG_LAST = 4u } [Flags] public enum NnapiFlags { NNAPI_FLAG_USE_NONE = 0, NNAPI_FLAG_USE_FP16 = 1, NNAPI_FLAG_USE_NCHW = 2, NNAPI_FLAG_CPU_DISABLED = 4, NNAPI_FLAG_CPU_ONLY = 8, NNAPI_FLAG_LAST = 8 } public class RunOptions : SafeHandle { private OrtLoggingLevel _logSeverityLevel = OrtLoggingLevel.ORT_LOGGING_LEVEL_WARNING; private int _logVerbosityLevel = 0; private string _logId = ""; private bool _terminate = false; internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; public OrtLoggingLevel LogSeverityLevel { get { return _logSeverityLevel; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunOptionsSetRunLogSeverityLevel(handle, value)); _logSeverityLevel = value; } } public int LogVerbosityLevel { get { return _logVerbosityLevel; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunOptionsSetRunLogVerbosityLevel(handle, value)); _logVerbosityLevel = value; } } public string LogId { get { return _logId; } set { byte[] runTag = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(value); NativeApiStatus.VerifySuccess(NativeMethods.OrtRunOptionsSetRunTag(handle, runTag)); _logId = value; } } public bool Terminate { get { return _terminate; } set { if (!_terminate && value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunOptionsSetTerminate(handle)); _terminate = true; } else if (_terminate && !value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtRunOptionsUnsetTerminate(handle)); _terminate = false; } } } public RunOptions() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateRunOptions(out handle)); } public void AddRunConfigEntry(string configKey, string configValue) { byte[] configKey2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(configKey); byte[] configValue2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(configValue); NativeApiStatus.VerifySuccess(NativeMethods.OrtAddRunConfigEntry(handle, configKey2, configValue2)); } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseRunOptions(handle); handle = IntPtr.Zero; return true; } } public enum GraphOptimizationLevel { ORT_DISABLE_ALL = 0, ORT_ENABLE_BASIC = 1, ORT_ENABLE_EXTENDED = 2, ORT_ENABLE_ALL = 99 } public enum ExecutionMode { ORT_SEQUENTIAL, ORT_PARALLEL } public class SessionOptions : SafeHandle { private class ExecutionProviderAppender { private byte[] _utf8ProviderName; internal ExecutionProviderAppender(byte[] providerName) { _utf8ProviderName = providerName; } public IntPtr Appender(IntPtr handle, IntPtr[] optKeys, IntPtr[] optValues, UIntPtr optCount) { return NativeMethods.SessionOptionsAppendExecutionProvider(handle, _utf8ProviderName, optKeys, optValues, optCount); } } private static string[] cudaDelayLoadedLibs = new string[0]; private static string[] trtDelayLoadedLibs = new string[0]; private bool _enableMemoryPattern = true; private bool _enableProfiling = false; private string _optimizedModelFilePath = ""; private bool _enableCpuMemArena = true; private string _logId = string.Empty; private OrtLoggingLevel _logSeverityLevel = OrtLoggingLevel.ORT_LOGGING_LEVEL_WARNING; private int _logVerbosityLevel = 0; private int _intraOpNumThreads = 0; private int _interOpNumThreads = 0; private GraphOptimizationLevel _graphOptimizationLevel = GraphOptimizationLevel.ORT_ENABLE_ALL; private ExecutionMode _executionMode = ExecutionMode.ORT_SEQUENTIAL; internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; public bool EnableMemoryPattern { get { return _enableMemoryPattern; } set { if (!_enableMemoryPattern && value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtEnableMemPattern(handle)); _enableMemoryPattern = true; } else if (_enableMemoryPattern && !value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtDisableMemPattern(handle)); _enableMemoryPattern = false; } } } public string ProfileOutputPathPrefix { get; set; } = "onnxruntime_profile_"; public bool EnableProfiling { get { return _enableProfiling; } set { if (!_enableProfiling && value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtEnableProfiling(handle, NativeOnnxValueHelper.GetPlatformSerializedString(ProfileOutputPathPrefix))); _enableProfiling = true; } else if (_enableProfiling && !value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtDisableProfiling(handle)); _enableProfiling = false; } } } public string OptimizedModelFilePath { get { return _optimizedModelFilePath; } set { if (value != _optimizedModelFilePath) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetOptimizedModelFilePath(handle, NativeOnnxValueHelper.GetPlatformSerializedString(value))); _optimizedModelFilePath = value; } } } public bool EnableCpuMemArena { get { return _enableCpuMemArena; } set { if (!_enableCpuMemArena && value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtEnableCpuMemArena(handle)); _enableCpuMemArena = true; } else if (_enableCpuMemArena && !value) { NativeApiStatus.VerifySuccess(NativeMethods.OrtDisableCpuMemArena(handle)); _enableCpuMemArena = false; } } } public string LogId { get { return _logId; } set { byte[] logId = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(value); NativeApiStatus.VerifySuccess(NativeMethods.OrtSetSessionLogId(handle, logId)); _logId = value; } } public OrtLoggingLevel LogSeverityLevel { get { return _logSeverityLevel; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetSessionLogSeverityLevel(handle, value)); _logSeverityLevel = value; } } public int LogVerbosityLevel { get { return _logVerbosityLevel; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetSessionLogVerbosityLevel(handle, value)); _logVerbosityLevel = value; } } public int IntraOpNumThreads { get { return _intraOpNumThreads; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetIntraOpNumThreads(handle, value)); _intraOpNumThreads = value; } } public int InterOpNumThreads { get { return _interOpNumThreads; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetInterOpNumThreads(handle, value)); _interOpNumThreads = value; } } public GraphOptimizationLevel GraphOptimizationLevel { get { return _graphOptimizationLevel; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetSessionGraphOptimizationLevel(handle, value)); _graphOptimizationLevel = value; } } public ExecutionMode ExecutionMode { get { return _executionMode; } set { NativeApiStatus.VerifySuccess(NativeMethods.OrtSetSessionExecutionMode(handle, value)); _executionMode = value; } } public SessionOptions() : base(IntPtr.Zero, ownsHandle: true) { NativeApiStatus.VerifySuccess(NativeMethods.OrtCreateSessionOptions(out handle)); OrtEnv.Instance(); } public static SessionOptions MakeSessionOptionWithCudaProvider(int deviceId = 0) { CheckCudaExecutionProviderDLLs(); SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_CUDA(deviceId); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public static SessionOptions MakeSessionOptionWithCudaProvider(OrtCUDAProviderOptions cudaProviderOptions) { CheckCudaExecutionProviderDLLs(); SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_CUDA(cudaProviderOptions); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public static SessionOptions MakeSessionOptionWithTensorrtProvider(int deviceId = 0) { CheckTensorrtExecutionProviderDLLs(); SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_Tensorrt(deviceId); sessionOptions.AppendExecutionProvider_CUDA(deviceId); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public static SessionOptions MakeSessionOptionWithTensorrtProvider(OrtTensorRTProviderOptions trtProviderOptions) { CheckTensorrtExecutionProviderDLLs(); SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_Tensorrt(trtProviderOptions); sessionOptions.AppendExecutionProvider_CUDA(trtProviderOptions.GetDeviceId()); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public static SessionOptions MakeSessionOptionWithTvmProvider(string settings = "") { SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_Tvm(settings); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public static SessionOptions MakeSessionOptionWithRocmProvider(int deviceId = 0) { CheckRocmExecutionProviderDLLs(); SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_ROCm(deviceId); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public static SessionOptions MakeSessionOptionWithRocmProvider(OrtROCMProviderOptions rocmProviderOptions) { CheckRocmExecutionProviderDLLs(); SessionOptions sessionOptions = new SessionOptions(); try { sessionOptions.AppendExecutionProvider_ROCm(rocmProviderOptions); return sessionOptions; } catch (Exception) { sessionOptions.Dispose(); throw; } } public void AppendExecutionProvider_CPU(int useArena = 1) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_CPU(handle, useArena)); } public void AppendExecutionProvider_Dnnl(int useArena = 1) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_Dnnl(handle, useArena)); } public void AppendExecutionProvider_CUDA(int deviceId = 0) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_CUDA(handle, deviceId)); } public void AppendExecutionProvider_CUDA(OrtCUDAProviderOptions cudaProviderOptions) { NativeApiStatus.VerifySuccess(NativeMethods.SessionOptionsAppendExecutionProvider_CUDA_V2(handle, cudaProviderOptions.Handle)); } public void AppendExecutionProvider_DML(int deviceId = 0) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_DML(handle, deviceId)); } public void AppendExecutionProvider_OpenVINO(string deviceId = "") { byte[] device_id = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(deviceId); NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_OpenVINO(handle, device_id)); } public void AppendExecutionProvider_Tensorrt(int deviceId = 0) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_Tensorrt(handle, deviceId)); } public void AppendExecutionProvider_Tensorrt(OrtTensorRTProviderOptions trtProviderOptions) { NativeApiStatus.VerifySuccess(NativeMethods.SessionOptionsAppendExecutionProvider_TensorRT_V2(handle, trtProviderOptions.Handle)); } public void AppendExecutionProvider_ROCm(int deviceId = 0) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_ROCM(handle, deviceId)); } public void AppendExecutionProvider_ROCm(OrtROCMProviderOptions rocmProviderOptions) { NativeApiStatus.VerifySuccess(NativeMethods.SessionOptionsAppendExecutionProvider_ROCM(handle, rocmProviderOptions.Handle)); } public void AppendExecutionProvider_MIGraphX(int deviceId = 0) { NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_MIGraphX(handle, deviceId)); } public void AppendExecutionProvider_Nnapi(NnapiFlags nnapiFlags = NnapiFlags.NNAPI_FLAG_USE_NONE) { throw new NotSupportedException("The NNAPI Execution Provider is not supported in this build"); } public void AppendExecutionProvider_CoreML(CoreMLFlags coremlFlags = CoreMLFlags.COREML_FLAG_USE_NONE) { throw new NotSupportedException("The CoreML Execution Provider is not supported in this build"); } public void AppendExecutionProvider_Tvm(string settings = "") { byte[] settings2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(settings); NativeApiStatus.VerifySuccess(NativeMethods.OrtSessionOptionsAppendExecutionProvider_Tvm(handle, settings2)); } public void AppendExecutionProvider(string providerName, Dictionary providerOptions = null) { if (providerName != "SNPE" && providerName != "XNNPACK" && providerName != "QNN" && providerName != "AZURE") { throw new NotSupportedException("Only QNN, SNPE, XNNPACK and AZURE execution providers can be enabled by this method."); } if (providerOptions == null) { providerOptions = new Dictionary(); } byte[] providerName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(providerName); ExecutionProviderAppender executionProviderAppender = new ExecutionProviderAppender(providerName2); ProviderOptionsUpdater.Update(providerOptions, handle, executionProviderAppender.Appender); } public void RegisterCustomOpLibrary(string libraryPath) { NativeApiStatus.VerifySuccess(NativeMethods.OrtRegisterCustomOpsLibrary_V2(handle, NativeOnnxValueHelper.GetPlatformSerializedString(libraryPath))); } public void RegisterCustomOpLibraryV2(string libraryPath, out IntPtr libraryHandle) { byte[] libraryPath2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(libraryPath); NativeApiStatus.VerifySuccess(NativeMethods.OrtRegisterCustomOpsLibrary(handle, libraryPath2, out libraryHandle)); } public void RegisterOrtExtensions() { try { OrtApiBase ortApiBase = NativeMethods.OrtGetApiBase(); NativeApiStatus.VerifySuccess(OrtExtensionsNativeMethods.RegisterCustomOps(handle, ref ortApiBase)); } catch (DllNotFoundException) { throw new OnnxRuntimeException(ErrorCode.NoSuchFile, "The ONNX Runtime extensions library was not found. The Microsoft.ML.OnnxRuntime.Extensions NuGet package must be referenced by the project to use 'OrtExtensions.RegisterCustomOps."); } } public void AddInitializer(string name, OrtValue ortValue) { byte[] name2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(name); NativeApiStatus.VerifySuccess(NativeMethods.OrtAddInitializer(handle, name2, ortValue.Handle)); } public void AddSessionConfigEntry(string configKey, string configValue) { byte[] configKey2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(configKey); byte[] configValue2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(configValue); NativeApiStatus.VerifySuccess(NativeMethods.OrtAddSessionConfigEntry(handle, configKey2, configValue2)); } public void AddFreeDimensionOverride(string dimDenotation, long dimValue) { byte[] dimDenotation2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(dimDenotation); NativeApiStatus.VerifySuccess(NativeMethods.OrtAddFreeDimensionOverride(handle, dimDenotation2, dimValue)); } public void AddFreeDimensionOverrideByName(string dimName, long dimValue) { byte[] dimName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(dimName); NativeApiStatus.VerifySuccess(NativeMethods.OrtAddFreeDimensionOverrideByName(handle, dimName2, dimValue)); } [DllImport("kernel32.dll")] private static extern IntPtr LoadLibrary(string dllToLoad); [DllImport("kernel32.dll")] private static extern uint GetSystemDirectory([Out] StringBuilder lpBuffer, uint uSize); private static bool CheckCudaExecutionProviderDLLs() { if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) { string[] array = cudaDelayLoadedLibs; foreach (string text in array) { IntPtr intPtr = LoadLibrary(text); if (!(intPtr != IntPtr.Zero)) { StringBuilder stringBuilder = new StringBuilder(string.Empty, 2048); GetSystemDirectory(stringBuilder, (uint)stringBuilder.Capacity); throw new OnnxRuntimeException(ErrorCode.NoSuchFile, "kernel32.LoadLibrary():'" + text + "' not found. CUDA is required for GPU execution. " + $". Verify it is available in the system directory={stringBuilder}. Else copy it to the output folder."); } } } return true; } private static bool CheckTensorrtExecutionProviderDLLs() { if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) { string[] array = trtDelayLoadedLibs; foreach (string text in array) { IntPtr intPtr = LoadLibrary(text); if (!(intPtr != IntPtr.Zero)) { StringBuilder stringBuilder = new StringBuilder(string.Empty, 2048); GetSystemDirectory(stringBuilder, (uint)stringBuilder.Capacity); throw new OnnxRuntimeException(ErrorCode.NoSuchFile, "kernel32.LoadLibrary():'" + text + "' not found. TensorRT/CUDA are required for GPU execution. " + $". Verify it is available in the system directory={stringBuilder}. Else copy it to the output folder."); } } } return true; } private static bool CheckRocmExecutionProviderDLLs() { if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) { throw new NotSupportedException("ROCm Execution Provider is not currently supported on Windows."); } return true; } protected override bool ReleaseHandle() { NativeMethods.OrtReleaseSessionOptions(handle); handle = IntPtr.Zero; return true; } } public static class SessionOptionsContainer { private static Lazy> _defaultHandler; private static readonly Dictionary>> _configurationHandlers = new Dictionary>>(); private static Lazy> DefaultHandler => (_defaultHandler != null) ? _defaultHandler : (_defaultHandler = new Lazy>(() => delegate { })); public static void Register(Action defaultHandler) { _defaultHandler = new Lazy>(() => defaultHandler); } public static void Register(string configuration, Action handler) { _configurationHandlers[configuration] = new Lazy>(() => handler); } public static SessionOptions Create(string configuration = null, bool useDefaultAsFallback = true) { return new SessionOptions().ApplyConfiguration(configuration, useDefaultAsFallback); } public static void Reset() { _defaultHandler = null; _configurationHandlers.Clear(); } public static SessionOptions ApplyConfiguration(this SessionOptions options, string configuration = null, bool useDefaultAsFallback = true) { Action action = Resolve(configuration, useDefaultAsFallback); action(options); return options; } private static Action Resolve(string configuration = null, bool useDefaultAsFallback = true) { if (string.IsNullOrWhiteSpace(configuration)) { return DefaultHandler.Value; } if (_configurationHandlers.TryGetValue(configuration, out var value)) { return value.Value; } if (useDefaultAsFallback) { return DefaultHandler.Value; } throw new KeyNotFoundException("Configuration not found for '" + configuration + "'"); } } public class CheckpointState : SafeHandle { internal enum PropertyType : long { Int, Float, String } internal IntPtr Handle => handle; public override bool IsInvalid => handle == IntPtr.Zero; private CheckpointState(IntPtr checkpointHandle) : base(checkpointHandle, ownsHandle: true) { } private unsafe void AddPropertyImpl(string propertyName, PropertyType propertyType, T propertyValue) where T : unmanaged { byte[] propertyName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(propertyName); fixed (T* ptr = new T[1] { propertyValue }) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtAddProperty(handle, propertyName2, propertyType, (IntPtr)ptr)); } } public static CheckpointState LoadCheckpoint(string checkpointPath) { if (!NativeTrainingMethods.TrainingEnabled()) { throw new InvalidOperationException("This package does not contain the training API. Please install the Microsoft.ML.OnnxRuntime.Training NuGet package.\n"); } IntPtr intPtr = OrtEnv.Instance().Handle; IntPtr checkpointState = IntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtLoadCheckpoint(NativeOnnxValueHelper.GetPlatformSerializedString(checkpointPath), out checkpointState)); return new CheckpointState(checkpointState); } public static void SaveCheckpoint(CheckpointState state, string checkpointPath, bool includeOptimizerState = false) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtSaveCheckpoint(state.Handle, NativeOnnxValueHelper.GetPlatformSerializedString(checkpointPath), includeOptimizerState)); } public void AddProperty(string propertyName, long propertyValue) { AddPropertyImpl(propertyName, PropertyType.Int, propertyValue); } public void AddProperty(string propertyName, float propertyValue) { AddPropertyImpl(propertyName, PropertyType.Float, propertyValue); } public unsafe void AddProperty(string propertyName, string propertyValue) { byte[] propertyName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(propertyName); fixed (byte* ptr = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(propertyValue)) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtAddProperty(handle, propertyName2, PropertyType.String, (IntPtr)ptr)); } } public unsafe object GetProperty(string propertyName) { byte[] propertyName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(propertyName); OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; IntPtr propertyValue = IntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetProperty(handle, propertyName2, defaultInstance.Pointer, out var propertyType, out propertyValue)); try { switch (propertyType) { case PropertyType.Int: { long num2 = *(long*)(void*)propertyValue; return num2; } case PropertyType.Float: { float num = *(float*)(void*)propertyValue; return num; } case PropertyType.String: return NativeOnnxValueHelper.StringFromNativeUtf8(propertyValue); default: throw new ArgumentException("Expected the property type to be one of long, float or string. Unknown type retrieved " + propertyValue); } } finally { defaultInstance.FreeMemory(propertyValue); } } public void UpdateParameter(string parameterName, OrtValue parameter) { if (parameter.OnnxType != OnnxValueType.ONNX_TYPE_TENSOR) { throw new ArgumentException("Incorrect buffer received. Expected a tensor parameter."); } byte[] parameterName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(parameterName); NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtUpdateParameter(handle, parameterName2, parameter.Handle)); } public OrtValue GetParameter(string parameterName) { byte[] parameterName2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(parameterName); NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetParameter(handle, parameterName2, OrtAllocator.DefaultInstance.Pointer, out var parameter)); return new OrtValue(parameter); } protected override bool ReleaseHandle() { NativeTrainingMethods.OrtReleaseCheckpointState(handle); handle = IntPtr.Zero; return true; } } public struct OrtTrainingApi { public IntPtr LoadCheckpoint; public IntPtr SaveCheckpoint; public IntPtr CreateTrainingSession; public IntPtr CreateTrainingSessionFromBuffer; public IntPtr TrainingSessionGetTrainingModelOutputCount; public IntPtr TrainingSessionGetEvalModelOutputCount; public IntPtr TrainingSessionGetTrainingModelOutputName; public IntPtr TrainingSessionGetEvalModelOutputName; public IntPtr LazyResetGrad; public IntPtr TrainStep; public IntPtr EvalStep; public IntPtr SetLearningRate; public IntPtr GetLearningRate; public IntPtr OptimizerStep; public IntPtr RegisterLinearLRScheduler; public IntPtr SchedulerStep; public IntPtr GetParametersSize; public IntPtr CopyParametersToBuffer; public IntPtr CopyBufferToParameters; public IntPtr ReleaseTrainingSession; public IntPtr ReleaseCheckpointState; public IntPtr ExportModelForInferencing; public IntPtr SetSeed; public IntPtr TrainingSessionGetTrainingModelInputCount; public IntPtr TrainingSessionGetEvalModelInputCount; public IntPtr TrainingSessionGetTrainingModelInputName; public IntPtr TrainingSessionGetEvalModelInputName; public IntPtr AddProperty; public IntPtr GetProperty; public IntPtr LoadCheckpointFromBuffer; public IntPtr GetParameterTypeAndShape; public IntPtr UpdateParameter; public IntPtr GetParameter; } internal static class NativeTrainingMethods { [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate ref OrtApi DOrtGetApi(uint version); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTrainingApi(uint version); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtLoadCheckpoint(byte[] checkpointPath, out IntPtr checkpointState); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSaveCheckpoint(IntPtr checkpointState, byte[] checkpointPath, bool includeOptimizerState); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCreateTrainingSession(IntPtr environment, IntPtr sessionOptions, IntPtr checkpointState, byte[] trainModelPath, byte[] evalModelPath, byte[] optimizerModelPath, out IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTrainingModelOutputCount(IntPtr session, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetEvalModelOutputCount(IntPtr session, out UIntPtr count); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTrainingModelOutputName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetEvalModelOutputName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtLazyResetGrad(IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtTrainStep(IntPtr session, IntPtr runOptions, UIntPtr inputCount, IntPtr[] inputValues, UIntPtr outputCount, IntPtr[] outputValues); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtEvalStep(IntPtr session, IntPtr runOptions, UIntPtr inputCount, IntPtr[] inputValues, UIntPtr outputCount, IntPtr[] outputValues); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtOptimizerStep(IntPtr session, IntPtr runOptions); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetLearningRate(IntPtr session, float learningRate); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetLearningRate(IntPtr session, out float learningRate); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtRegisterLinearLRScheduler(IntPtr session, long warmupStepCount, long totalStepCount, float learningRate); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSchedulerStep(IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetParametersSize(IntPtr session, out UIntPtr buffer_size, bool only_trainable); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCopyParametersToBuffer(IntPtr session, IntPtr buffer, bool only_trainable); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtCopyBufferToParameters(IntPtr session, IntPtr buffer, bool only_trainable); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseTrainingSession(IntPtr session); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate void DOrtReleaseCheckpointState(IntPtr checkpointState); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtExportModelForInferencing(IntPtr session, byte[] inferenceModelPath, UIntPtr graphOutputCount, IntPtr[] graphOutputNames); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtSetSeed(long seed); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTrainingModelInputCount(IntPtr session, out UIntPtr inputCount); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetEvalModelInputCount(IntPtr session, out UIntPtr inputCount); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetTrainingModelInputName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetEvalModelInputName(IntPtr session, UIntPtr index, IntPtr allocator, out IntPtr name); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtAddProperty(IntPtr checkpointState, byte[] propertyName, CheckpointState.PropertyType propertyType, IntPtr propertyValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetProperty(IntPtr checkpointState, byte[] propertyName, IntPtr allocator, out CheckpointState.PropertyType propertyType, out IntPtr propertyValue); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetParameterTypeAndShape(IntPtr checkpointState, byte[] parameterName, out IntPtr parameterTypeAndShape); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtUpdateParameter(IntPtr checkpointState, byte[] parameterName, IntPtr parameter); [UnmanagedFunctionPointer(CallingConvention.Winapi)] public delegate IntPtr DOrtGetParameter(IntPtr checkpointState, byte[] parameterName, IntPtr allocator, out IntPtr parameter); private static OrtApi api_; private static OrtTrainingApi trainingApi_; private static IntPtr trainingApiPtr; public static DOrtGetTrainingApi OrtGetTrainingApi; public static DOrtLoadCheckpoint OrtLoadCheckpoint; public static DOrtSaveCheckpoint OrtSaveCheckpoint; public static DOrtCreateTrainingSession OrtCreateTrainingSession; public static DOrtGetTrainingModelOutputCount OrtGetTrainingModelOutputCount; public static DOrtGetEvalModelOutputCount OrtGetEvalModelOutputCount; public static DOrtGetTrainingModelOutputName OrtGetTrainingModelOutputName; public static DOrtGetEvalModelOutputName OrtGetEvalModelOutputName; public static DOrtLazyResetGrad OrtLazyResetGrad; public static DOrtTrainStep OrtTrainStep; public static DOrtEvalStep OrtEvalStep; public static DOrtOptimizerStep OrtOptimizerStep; public static DOrtSetLearningRate OrtSetLearningRate; public static DOrtGetLearningRate OrtGetLearningRate; public static DOrtRegisterLinearLRScheduler OrtRegisterLinearLRScheduler; public static DOrtSchedulerStep OrtSchedulerStep; public static DOrtGetParametersSize OrtGetParametersSize; public static DOrtCopyParametersToBuffer OrtCopyParametersToBuffer; public static DOrtCopyBufferToParameters OrtCopyBufferToParameters; public static DOrtReleaseTrainingSession OrtReleaseTrainingSession; public static DOrtReleaseCheckpointState OrtReleaseCheckpointState; public static DOrtExportModelForInferencing OrtExportModelForInferencing; public static DOrtSetSeed OrtSetSeed; public static DOrtGetTrainingModelInputCount OrtGetTrainingModelInputCount; public static DOrtGetEvalModelInputCount OrtGetEvalModelInputCount; public static DOrtGetTrainingModelInputName OrtGetTrainingModelInputName; public static DOrtGetEvalModelInputName OrtGetEvalModelInputName; public static DOrtAddProperty OrtAddProperty; public static DOrtGetProperty OrtGetProperty; public static DOrtGetParameterTypeAndShape OrtGetParameterTypeAndShape; public static DOrtUpdateParameter OrtUpdateParameter; public static DOrtGetParameter OrtGetParameter; static NativeTrainingMethods() { DOrtGetApi dOrtGetApi = (DOrtGetApi)Marshal.GetDelegateForFunctionPointer(NativeMethods.OrtGetApiBase().GetApi, typeof(DOrtGetApi)); api_ = dOrtGetApi(17u); OrtGetTrainingApi = (DOrtGetTrainingApi)Marshal.GetDelegateForFunctionPointer(api_.GetTrainingApi, typeof(DOrtGetTrainingApi)); trainingApiPtr = OrtGetTrainingApi(17u); if (trainingApiPtr != IntPtr.Zero) { trainingApi_ = (OrtTrainingApi)Marshal.PtrToStructure(trainingApiPtr, typeof(OrtTrainingApi)); OrtLoadCheckpoint = (DOrtLoadCheckpoint)Marshal.GetDelegateForFunctionPointer(trainingApi_.LoadCheckpoint, typeof(DOrtLoadCheckpoint)); OrtSaveCheckpoint = (DOrtSaveCheckpoint)Marshal.GetDelegateForFunctionPointer(trainingApi_.SaveCheckpoint, typeof(DOrtSaveCheckpoint)); OrtCreateTrainingSession = (DOrtCreateTrainingSession)Marshal.GetDelegateForFunctionPointer(trainingApi_.CreateTrainingSession, typeof(DOrtCreateTrainingSession)); OrtGetTrainingModelOutputCount = (DOrtGetTrainingModelOutputCount)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetTrainingModelOutputCount, typeof(DOrtGetTrainingModelOutputCount)); OrtGetEvalModelOutputCount = (DOrtGetEvalModelOutputCount)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetEvalModelOutputCount, typeof(DOrtGetEvalModelOutputCount)); OrtGetTrainingModelOutputName = (DOrtGetTrainingModelOutputName)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetTrainingModelOutputName, typeof(DOrtGetTrainingModelOutputName)); OrtGetEvalModelOutputName = (DOrtGetEvalModelOutputName)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetEvalModelOutputName, typeof(DOrtGetEvalModelOutputName)); OrtLazyResetGrad = (DOrtLazyResetGrad)Marshal.GetDelegateForFunctionPointer(trainingApi_.LazyResetGrad, typeof(DOrtLazyResetGrad)); OrtTrainStep = (DOrtTrainStep)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainStep, typeof(DOrtTrainStep)); OrtEvalStep = (DOrtEvalStep)Marshal.GetDelegateForFunctionPointer(trainingApi_.EvalStep, typeof(DOrtEvalStep)); OrtSetLearningRate = (DOrtSetLearningRate)Marshal.GetDelegateForFunctionPointer(trainingApi_.SetLearningRate, typeof(DOrtSetLearningRate)); OrtGetLearningRate = (DOrtGetLearningRate)Marshal.GetDelegateForFunctionPointer(trainingApi_.GetLearningRate, typeof(DOrtGetLearningRate)); OrtOptimizerStep = (DOrtOptimizerStep)Marshal.GetDelegateForFunctionPointer(trainingApi_.OptimizerStep, typeof(DOrtOptimizerStep)); OrtRegisterLinearLRScheduler = (DOrtRegisterLinearLRScheduler)Marshal.GetDelegateForFunctionPointer(trainingApi_.RegisterLinearLRScheduler, typeof(DOrtRegisterLinearLRScheduler)); OrtSchedulerStep = (DOrtSchedulerStep)Marshal.GetDelegateForFunctionPointer(trainingApi_.SchedulerStep, typeof(DOrtSchedulerStep)); OrtGetParametersSize = (DOrtGetParametersSize)Marshal.GetDelegateForFunctionPointer(trainingApi_.GetParametersSize, typeof(DOrtGetParametersSize)); OrtCopyParametersToBuffer = (DOrtCopyParametersToBuffer)Marshal.GetDelegateForFunctionPointer(trainingApi_.CopyParametersToBuffer, typeof(DOrtCopyParametersToBuffer)); OrtCopyBufferToParameters = (DOrtCopyBufferToParameters)Marshal.GetDelegateForFunctionPointer(trainingApi_.CopyBufferToParameters, typeof(DOrtCopyBufferToParameters)); OrtReleaseTrainingSession = (DOrtReleaseTrainingSession)Marshal.GetDelegateForFunctionPointer(trainingApi_.ReleaseTrainingSession, typeof(DOrtReleaseTrainingSession)); OrtReleaseCheckpointState = (DOrtReleaseCheckpointState)Marshal.GetDelegateForFunctionPointer(trainingApi_.ReleaseCheckpointState, typeof(DOrtReleaseCheckpointState)); OrtExportModelForInferencing = (DOrtExportModelForInferencing)Marshal.GetDelegateForFunctionPointer(trainingApi_.ExportModelForInferencing, typeof(DOrtExportModelForInferencing)); OrtSetSeed = (DOrtSetSeed)Marshal.GetDelegateForFunctionPointer(trainingApi_.SetSeed, typeof(DOrtSetSeed)); OrtGetTrainingModelInputCount = (DOrtGetTrainingModelInputCount)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetTrainingModelInputCount, typeof(DOrtGetTrainingModelInputCount)); OrtGetEvalModelInputCount = (DOrtGetEvalModelInputCount)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetEvalModelInputCount, typeof(DOrtGetEvalModelInputCount)); OrtGetTrainingModelInputName = (DOrtGetTrainingModelInputName)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetTrainingModelInputName, typeof(DOrtGetTrainingModelInputName)); OrtGetEvalModelInputName = (DOrtGetEvalModelInputName)Marshal.GetDelegateForFunctionPointer(trainingApi_.TrainingSessionGetEvalModelInputName, typeof(DOrtGetEvalModelInputName)); OrtAddProperty = (DOrtAddProperty)Marshal.GetDelegateForFunctionPointer(trainingApi_.AddProperty, typeof(DOrtAddProperty)); OrtGetProperty = (DOrtGetProperty)Marshal.GetDelegateForFunctionPointer(trainingApi_.GetProperty, typeof(DOrtGetProperty)); OrtGetParameterTypeAndShape = (DOrtGetParameterTypeAndShape)Marshal.GetDelegateForFunctionPointer(trainingApi_.GetParameterTypeAndShape, typeof(DOrtGetParameterTypeAndShape)); OrtUpdateParameter = (DOrtUpdateParameter)Marshal.GetDelegateForFunctionPointer(trainingApi_.UpdateParameter, typeof(DOrtUpdateParameter)); OrtGetParameter = (DOrtGetParameter)Marshal.GetDelegateForFunctionPointer(trainingApi_.GetParameter, typeof(DOrtGetParameter)); } } public static bool TrainingEnabled() { if (trainingApiPtr == IntPtr.Zero) { return false; } return true; } } public class TrainingUtils { public static void SetSeed(long seed) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtSetSeed(seed)); } } internal enum LRScheduler { None, Constant, Linear } public class TrainingSession : IDisposable { private IntPtr _nativeHandle; private ulong _trainOutputCount; private ulong _evalOutputCount; private List _trainOutputNames; private List _evalOutputNames; private List _trainInputNames; private List _evalInputNames; private SessionOptions _builtInSessionOptions = null; private RunOptions _builtInRunOptions = null; private LRScheduler _scheduler = LRScheduler.None; private bool _disposed = false; internal IntPtr Handle => _nativeHandle; public TrainingSession(CheckpointState state, string trainModelPath, string evalModelPath, string optimizerModelPath) { Init(null, state, NativeOnnxValueHelper.GetPlatformSerializedString(trainModelPath), NativeOnnxValueHelper.GetPlatformSerializedString(evalModelPath), NativeOnnxValueHelper.GetPlatformSerializedString(optimizerModelPath)); } public TrainingSession(CheckpointState state, string trainModelPath, string optimizerModelPath) { Init(null, state, NativeOnnxValueHelper.GetPlatformSerializedString(trainModelPath), null, NativeOnnxValueHelper.GetPlatformSerializedString(optimizerModelPath)); } public TrainingSession(SessionOptions options, CheckpointState state, string trainModelPath, string evalModelPath, string optimizerModelPath) { Init(options, state, NativeOnnxValueHelper.GetPlatformSerializedString(trainModelPath), NativeOnnxValueHelper.GetPlatformSerializedString(evalModelPath), NativeOnnxValueHelper.GetPlatformSerializedString(optimizerModelPath)); } public void TrainStep(IReadOnlyCollection inputValues, IReadOnlyCollection outputValues) { TrainStep(_builtInRunOptions, inputValues, outputValues); } public void TrainStep(RunOptions options, IReadOnlyCollection inputValues, IReadOnlyCollection outputValues) { if (_trainOutputCount != (ulong)outputValues.Count()) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of train model ({2}).", "outputValues", outputValues.Count, _trainOutputCount)); } IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] ortValuesHandles2 = GetOrtValuesHandles(outputValues, input: false); NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtTrainStep(_nativeHandle, options.Handle, (UIntPtr)(ulong)inputValues.Count, ortValuesHandles, (UIntPtr)(ulong)outputValues.Count, ortValuesHandles2)); } public IDisposableReadOnlyCollection TrainStep(IReadOnlyCollection inputValues) { IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] array = new IntPtr[(uint)_trainOutputCount]; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtTrainStep(_nativeHandle, _builtInRunOptions.Handle, (UIntPtr)(ulong)inputValues.Count, ortValuesHandles, (UIntPtr)_trainOutputCount, array)); DisposableArray disposableArray = ConvertNativeHandlesToOrtValues(array); try { DisposableList disposableList = new DisposableList(_trainOutputNames.Count); try { for (int i = 0; i < disposableArray.Span.Length; i++) { disposableList.Add(DisposableNamedOnnxValue.CreateFromOrtValue(_trainOutputNames[i], ref disposableArray.Span[i])); } } catch (OnnxRuntimeException) { disposableList.Dispose(); throw; } return disposableList; } finally { disposableArray.Dispose(); } } public IDisposableReadOnlyCollection TrainStep(RunOptions options, IReadOnlyCollection inputValues) { IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] array = new IntPtr[(uint)_trainOutputCount]; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtTrainStep(_nativeHandle, options.Handle, (UIntPtr)(ulong)inputValues.Count, ortValuesHandles, (UIntPtr)_trainOutputCount, array)); DisposableArray disposableArray = ConvertNativeHandlesToOrtValues(array); try { DisposableList disposableList = new DisposableList(_trainOutputNames.Count); try { for (int i = 0; i < disposableArray.Span.Length; i++) { disposableList.Add(DisposableNamedOnnxValue.CreateFromOrtValue(_trainOutputNames[i], ref disposableArray.Span[i])); } } catch (OnnxRuntimeException) { disposableList.Dispose(); throw; } return disposableList; } finally { disposableArray.Dispose(); } } private DisposableArray ConvertNativeHandlesToOrtValues(IntPtr[] nativeHandles) { DisposableOrtValueHandleArray disposableOrtValueHandleArray = new DisposableOrtValueHandleArray(nativeHandles); try { OrtValue[] array = new OrtValue[nativeHandles.Length]; DisposableArray result = new DisposableArray(array); try { for (int i = 0; i < nativeHandles.Length; i++) { array[i] = new OrtValue(nativeHandles[i]); nativeHandles[i] = IntPtr.Zero; } return result; } catch (Exception) { result.Dispose(); throw; } } catch (Exception) { disposableOrtValueHandleArray.Dispose(); throw; } } public void LazyResetGrad() { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtLazyResetGrad(_nativeHandle)); } public void EvalStep(IReadOnlyCollection inputValues, IReadOnlyCollection outputValues) { EvalStep(_builtInRunOptions, inputValues, outputValues); } public void EvalStep(RunOptions options, IReadOnlyCollection inputValues, IReadOnlyCollection outputValues) { if (_evalOutputCount != (ulong)outputValues.Count()) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match that of eval model ({2}).", "outputValues", outputValues.Count, _evalOutputCount)); } IntPtr[] ortValuesHandles = GetOrtValuesHandles(inputValues, input: true); IntPtr[] ortValuesHandles2 = GetOrtValuesHandles(outputValues, input: false); NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtEvalStep(_nativeHandle, options.Handle, (UIntPtr)(ulong)inputValues.Count, ortValuesHandles, (UIntPtr)(ulong)outputValues.Count, ortValuesHandles2)); } public void SetLearningRate(float learningRate) { if (_scheduler != LRScheduler.None && _scheduler != LRScheduler.Constant) { throw new InvalidOperationException("Cannot set constant LR while using LR scheduler."); } NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtSetLearningRate(_nativeHandle, learningRate)); _scheduler = LRScheduler.Constant; } public float GetLearningRate() { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetLearningRate(_nativeHandle, out var learningRate)); return learningRate; } public void RegisterLinearLRScheduler(long warmupStepCount, long totalStepCount, float initialLearningRate) { if (_scheduler != LRScheduler.None && _scheduler != LRScheduler.Constant) { throw new InvalidOperationException("Cannot set LR scheduler while using constant LR."); } NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtRegisterLinearLRScheduler(_nativeHandle, warmupStepCount, totalStepCount, initialLearningRate)); _scheduler = LRScheduler.Linear; } public void SchedulerStep() { if (_scheduler == LRScheduler.Constant || _scheduler == LRScheduler.None) { throw new InvalidOperationException("Cannot take scheduler step without registering a valid LR scheduler."); } NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtSchedulerStep(_nativeHandle)); } public void OptimizerStep() { OptimizerStep(_builtInRunOptions); } public void OptimizerStep(RunOptions options) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtOptimizerStep(_nativeHandle, options.Handle)); } public void ExportModelForInferencing(string inferenceModelPath, IReadOnlyCollection graphOutputNames) { using DisposableList cleanupList = new DisposableList(); IntPtr[] graphOutputNames2 = ConvertNamesToUtf8(graphOutputNames, cleanupList); NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtExportModelForInferencing(_nativeHandle, NativeOnnxValueHelper.GetPlatformSerializedString(inferenceModelPath), (UIntPtr)(ulong)graphOutputNames.Count, graphOutputNames2)); } public OrtValue ToBuffer(bool onlyTrainable) { UIntPtr buffer_size = UIntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetParametersSize(_nativeHandle, out buffer_size, onlyTrainable)); float[] array = new float[buffer_size.ToUInt64()]; long[] shape = new long[1] { (long)(ulong)buffer_size }; OrtValue ortValue = OrtValue.CreateAllocatedTensorValue(OrtAllocator.DefaultInstance, TensorElementType.Float, shape); NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtCopyParametersToBuffer(_nativeHandle, ortValue.Handle, onlyTrainable)); return ortValue; } public void FromBuffer(OrtValue ortValue, bool onlyTrainable) { if (ortValue.OnnxType != OnnxValueType.ONNX_TYPE_TENSOR) { throw new ArgumentException("Incorrect buffer received. Expected a tensor buffer."); } OrtTensorTypeAndShapeInfo tensorTypeAndShape = ortValue.GetTensorTypeAndShape(); if (tensorTypeAndShape.ElementDataType != TensorElementType.Float) { throw new ArgumentException("Incorrect buffer received. Expected a tensor buffer of type float."); } UIntPtr buffer_size = UIntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetParametersSize(_nativeHandle, out buffer_size, onlyTrainable)); if (tensorTypeAndShape.ElementCount != (long)(ulong)buffer_size) { string message = "Incorrect buffer size received. Expected size to be " + buffer_size + ". Actual size: " + tensorTypeAndShape.ElementCount; throw new ArgumentException(message); } NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtCopyBufferToParameters(_nativeHandle, ortValue.Handle, onlyTrainable)); } public List OutputNames(bool training) { return training ? _trainOutputNames : _evalOutputNames; } public List InputNames(bool training) { return training ? _trainInputNames : _evalInputNames; } private void Init(SessionOptions sessOptions, CheckpointState state, byte[] trainModelPath, byte[] evalModelPath, byte[] optimizerModelPath) { if (!NativeTrainingMethods.TrainingEnabled()) { throw new InvalidOperationException("This package does not contain the training API. Please install the Microsoft.ML.OnnxRuntime.Training NuGet package.\n"); } SessionOptions sessionOptions = sessOptions; if (sessOptions == null) { _builtInSessionOptions = new SessionOptions(); sessionOptions = _builtInSessionOptions; } IntPtr handle = OrtEnv.Instance().Handle; try { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtCreateTrainingSession(handle, sessionOptions.Handle, state.Handle, trainModelPath, evalModelPath, optimizerModelPath, out _nativeHandle)); UIntPtr count = UIntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetTrainingModelOutputCount(_nativeHandle, out count)); _trainOutputCount = count.ToUInt64(); _trainOutputNames = new List(); for (ulong num = 0uL; num < _trainOutputCount; num++) { _trainOutputNames.Add(GetOutputName(num, training: true)); } _trainInputNames = new List(); UIntPtr inputCount = UIntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetTrainingModelInputCount(_nativeHandle, out inputCount)); for (ulong num2 = 0uL; num2 < inputCount.ToUInt64(); num2++) { _trainInputNames.Add(GetInputName(num2, training: true)); } if (evalModelPath != null) { count = UIntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetEvalModelOutputCount(_nativeHandle, out count)); _evalOutputCount = count.ToUInt64(); _evalOutputNames = new List(); for (ulong num3 = 0uL; num3 < _evalOutputCount; num3++) { _evalOutputNames.Add(GetOutputName(num3, training: false)); } _evalInputNames = new List(); inputCount = UIntPtr.Zero; NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetEvalModelInputCount(_nativeHandle, out inputCount)); for (ulong num4 = 0uL; num4 < inputCount.ToUInt64(); num4++) { _evalInputNames.Add(GetInputName(num4, training: false)); } } _builtInRunOptions = new RunOptions(); } catch (Exception) { CleanupHelper(disposing: true); throw; } } private string GetOutputName(ulong index, bool training) { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; IntPtr name; if (training) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetTrainingModelOutputName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out name)); } else { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetEvalModelOutputName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out name)); } return NativeOnnxValueHelper.StringFromNativeUtf8(name, defaultInstance); } private string GetInputName(ulong index, bool training) { OrtAllocator defaultInstance = OrtAllocator.DefaultInstance; IntPtr name; if (training) { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetTrainingModelInputName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out name)); } else { NativeApiStatus.VerifySuccess(NativeTrainingMethods.OrtGetEvalModelInputName(_nativeHandle, (UIntPtr)index, defaultInstance.Pointer, out name)); } return NativeOnnxValueHelper.StringFromNativeUtf8(name, defaultInstance); } private IntPtr[] GetOrtValuesHandles(IReadOnlyCollection values, bool input) { IntPtr[] array = new IntPtr[values.Count]; for (int i = 0; i < values.Count; i++) { FixedBufferOnnxValue fixedBufferOnnxValue = values.ElementAt(i); if (!input && fixedBufferOnnxValue.ElementType == TensorElementType.String) { throw new NotSupportedException("Using string type FixedBufferOnnxValue in outputs is not supported."); } array[i] = fixedBufferOnnxValue.Value.Handle; } return array; } private unsafe IntPtr[] ConvertNamesToUtf8(IReadOnlyCollection names, DisposableList cleanupList) { cleanupList.Capacity += names.Count; IntPtr[] array = new IntPtr[names.Count]; for (int i = 0; i < names.Count; i++) { string s = names.ElementAt(i); byte[] array2 = NativeOnnxValueHelper.StringToZeroTerminatedUtf8(s); MemoryHandle memoryHandle = new Memory(array2).Pin(); array[i] = (IntPtr)memoryHandle.Pointer; cleanupList.Add(memoryHandle); } return array; } ~TrainingSession() { Dispose(disposing: false); } public void Dispose() { Dispose(disposing: true); GC.SuppressFinalize(this); } protected virtual void Dispose(bool disposing) { if (!_disposed) { CleanupHelper(disposing); _disposed = true; } } private void CleanupHelper(bool disposing) { if (disposing) { if (_builtInRunOptions != null) { _builtInRunOptions.Dispose(); _builtInRunOptions = null; } if (_builtInSessionOptions != null) { _builtInSessionOptions.Dispose(); _builtInSessionOptions = null; } } if (_nativeHandle != IntPtr.Zero) { NativeTrainingMethods.OrtReleaseTrainingSession(_nativeHandle); _nativeHandle = IntPtr.Zero; } } } } namespace Microsoft.ML.OnnxRuntime.Tensors { public static class ShapeUtils { public static long GetSizeForShape(ReadOnlySpan shape) { long num = 1L; ReadOnlySpan readOnlySpan = shape; for (int i = 0; i < readOnlySpan.Length; i++) { long num2 = readOnlySpan[i]; if (num2 < 0) { throw new ArgumentOutOfRangeException($"Shape must not have negative elements: {num2}"); } num = checked(num * num2); } return num; } public static long[] GetStrides(ReadOnlySpan dimensions) { long[] array = new long[dimensions.Length]; if (dimensions.Length == 0) { return array; } long num = 1L; for (int num2 = array.Length - 1; num2 >= 0; num2--) { array[num2] = num; if (dimensions[num2] < 0) { throw new ArgumentException($"Dimension {num2} is negative"); } num *= dimensions[num2]; } return array; } public static long GetIndex(ReadOnlySpan strides, ReadOnlySpan indices, int startFromDimension = 0) { long num = 0L; for (int i = startFromDimension; i < indices.Length; i++) { num += strides[i] * indices[i]; } return num; } } public static class ArrayTensorExtensions { public static DenseTensor ToTensor(this T[] array) { int[] array2 = new int[1] { array.Length }; T[] array3 = new T[array.Length]; array.CopyTo(array3, 0); return new DenseTensor(new Memory(array3), array2); } public static DenseTensor ToTensor(this T[,] array, bool reverseStride = false) { if (reverseStride) { return new DenseTensor(array, reverseStride); } T[] array2 = new T[array.Length]; int[] array3 = new int[2] { array.GetLength(0), array.GetLength(1) }; long num = 0L; foreach (T val in array) { array2[num++] = val; } return new DenseTensor(new Memory(array2), array3); } public static DenseTensor ToTensor(this T[,,] array, bool reverseStride = false) { if (reverseStride) { return new DenseTensor(array, reverseStride); } T[] array2 = new T[array.Length]; int[] array3 = new int[3] { array.GetLength(0), array.GetLength(1), array.GetLength(2) }; long num = 0L; foreach (T val in array) { array2[num++] = val; } return new DenseTensor(new Memory(array2), array3); } public static DenseTensor ToTensor(this T[,,,] array, bool reverseStride = false) { if (reverseStride) { return new DenseTensor(array, reverseStride); } T[] array2 = new T[array.Length]; int[] array3 = new int[4] { array.GetLength(0), array.GetLength(1), array.GetLength(2), array.GetLength(3) }; long num = 0L; foreach (T val in array) { array2[num++] = val; } return new DenseTensor(new Memory(array2), array3); } public static DenseTensor ToTensor(this Array array, bool reverseStride = false) { return new DenseTensor(array, reverseStride); } } internal static class ArrayUtilities { private static class EmptyArray { public static readonly T[] Value = new T[0]; } public const int StackallocMax = 16; public static long GetProduct(ReadOnlySpan dimensions, int startIndex = 0) { long num = 1L; for (int i = startIndex; i < dimensions.Length; i++) { if (dimensions[i] < 0) { throw new ArgumentOutOfRangeException(string.Format("{0}[{1}]", "dimensions", i)); } num = checked(num * dimensions[i]); } return num; } public static bool IsAscending(ReadOnlySpan values) { for (int i = 1; i < values.Length; i++) { if (values[i] < values[i - 1]) { return false; } } return true; } public static bool IsDescending(ReadOnlySpan values) { for (int i = 1; i < values.Length; i++) { if (values[i] > values[i - 1]) { return false; } } return true; } public static int[] GetStrides(ReadOnlySpan dimensions, bool reverseStride = false) { int[] array = new int[dimensions.Length]; if (dimensions.Length == 0) { return array; } int num = 1; if (reverseStride) { for (int i = 0; i < array.Length; i++) { array[i] = num; num *= dimensions[i]; } } else { for (int num2 = array.Length - 1; num2 >= 0; num2--) { array[num2] = num; num *= dimensions[num2]; } } return array; } public static void SplitStrides(int[] strides, int[] splitAxes, int[] newStrides, int stridesOffset, int[] splitStrides, int splitStridesOffset) { int num = 0; for (int i = 0; i < strides.Length; i++) { int num2 = strides[i]; bool flag = false; for (int j = 0; j < splitAxes.Length; j++) { if (splitAxes[j] == i) { splitStrides[splitStridesOffset + j] = num2; flag = true; break; } } if (!flag) { newStrides[stridesOffset + num++] = num2; } } } public static int GetIndex(int[] strides, ReadOnlySpan indices, int startFromDimension = 0) { int num = 0; for (int i = startFromDimension; i < indices.Length; i++) { num += strides[i] * indices[i]; } return num; } public static void GetIndices(ReadOnlySpan strides, bool reverseStride, int index, int[] indices, int startFromDimension = 0) { if (indices.Length != 0) { int num = index; for (int i = startFromDimension; i < strides.Length; i++) { int num2 = (reverseStride ? (strides.Length - 1 - i) : i); int num3 = strides[num2]; indices[num2] = num / num3; num %= num3; } } } public static void GetIndices(ReadOnlySpan strides, bool reverseStride, int index, Span indices, int startFromDimension = 0) { if (indices.Length != 0) { int num = index; for (int i = startFromDimension; i < strides.Length; i++) { int index2 = (reverseStride ? (strides.Length - 1 - i) : i); int num2 = strides[index2]; indices[index2] = num / num2; num %= num2; } } } public static int TransformIndexByStrides(int index, int[] sourceStrides, bool sourceReverseStride, int[] transformStrides) { if (sourceStrides.Length == 0) { return 0; } int num = 0; int num2 = index; for (int i = 0; i < sourceStrides.Length; i++) { int num3 = (sourceReverseStride ? (sourceStrides.Length - 1 - i) : i); int num4 = sourceStrides[num3]; int num5 = transformStrides[num3]; num += num5 * (num2 / num4); num2 %= num4; } return num; } public static T[] GetEmpty() { return EmptyArray.Value; } } public class DenseTensor : Tensor { private readonly Memory memory; public Memory Buffer => memory; internal DenseTensor(Array fromArray, bool reverseStride = false) : base(fromArray, reverseStride) { T[] array = new T[fromArray.Length]; int num = 0; if (reverseStride) { int[] sourceStrides = ArrayUtilities.GetStrides(dimensions); foreach (object item in fromArray) { int num2 = ArrayUtilities.TransformIndexByStrides(num++, sourceStrides, sourceReverseStride: false, strides); array[num2] = (T)item; } } else { foreach (object item2 in fromArray) { array[num++] = (T)item2; } } memory = array; } public DenseTensor(int length) : base(length) { memory = new T[length]; } public DenseTensor(ReadOnlySpan dimensions, bool reverseStride = false) : base(dimensions, reverseStride) { memory = new T[base.Length]; } public DenseTensor(Memory memory, ReadOnlySpan dimensions, bool reverseStride = false) : base(dimensions, reverseStride) { this.memory = memory; if (base.Length != memory.Length) { throw new ArgumentException(string.Format("Length of {0} ({1}) must match product of ", "memory", memory.Length) + string.Format("{0} ({1}).", "dimensions", base.Length)); } } public override T GetValue(int index) { return Buffer.Span[index]; } public override void SetValue(int index, T value) { Buffer.Span[index] = value; } protected override void CopyTo(T[] array, int arrayIndex) { if (array == null) { throw new ArgumentNullException("array"); } if (array.Length < arrayIndex + base.Length) { throw new ArgumentException("The number of elements in the Tensor is greater than the available space from index to the end of the destination array.", "array"); } Buffer.Span.CopyTo(array.AsSpan(arrayIndex)); } protected override int IndexOf(T item) { if (MemoryMarshal.TryGetArray((ReadOnlyMemory)Buffer, out ArraySegment segment)) { int num = Array.IndexOf(segment.Array, item, segment.Offset, segment.Count); if (num != -1) { num -= segment.Offset; } return num; } return base.IndexOf(item); } public override Tensor Clone() { return new DenseTensor(new Memory(memory.ToArray()), dimensions, base.IsReversedStride); } public override Tensor CloneEmpty(ReadOnlySpan dimensions) { return new DenseTensor(dimensions, base.IsReversedStride); } public override Tensor Reshape(ReadOnlySpan dimensions) { long product = ArrayUtilities.GetProduct(dimensions); if (product != base.Length) { throw new ArgumentException("Cannot reshape array due to mismatch in lengths, currently {Length} would become {newSize}.", "dimensions"); } return new DenseTensor(Buffer, dimensions, base.IsReversedStride); } } public enum TensorElementType { Float = 1, UInt8, Int8, UInt16, Int16, Int32, Int64, String, Bool, Float16, Double, UInt32, UInt64, Complex64, Complex128, BFloat16, DataTypeMax } public class TensorTypeInfo { public TensorElementType ElementType { get; private set; } public int TypeSize { get; private set; } public bool IsString => ElementType == TensorElementType.String; public TensorTypeInfo(TensorElementType elementType, int typeSize) { ElementType = elementType; TypeSize = typeSize; } } public class TensorElementTypeInfo { public Type TensorType { get; private set; } public int TypeSize { get; private set; } public bool IsString { get; private set; } public TensorElementTypeInfo(Type type, int typeSize) { TensorType = type; TypeSize = typeSize; IsString = type == typeof(string); } } public class TensorBase { private static readonly Dictionary typeInfoMap; private static readonly Dictionary tensorElementTypeInfoMap; private readonly Type _primitiveType; static TensorBase() { typeInfoMap = new Dictionary { { typeof(float), new TensorTypeInfo(TensorElementType.Float, 4) }, { typeof(byte), new TensorTypeInfo(TensorElementType.UInt8, 1) }, { typeof(sbyte), new TensorTypeInfo(TensorElementType.Int8, 1) }, { typeof(ushort), new TensorTypeInfo(TensorElementType.UInt16, 2) }, { typeof(short), new TensorTypeInfo(TensorElementType.Int16, 2) }, { typeof(int), new TensorTypeInfo(TensorElementType.Int32, 4) }, { typeof(long), new TensorTypeInfo(TensorElementType.Int64, 8) }, { typeof(string), new TensorTypeInfo(TensorElementType.String, -1) }, { typeof(bool), new TensorTypeInfo(TensorElementType.Bool, 1) }, { typeof(Float16), new TensorTypeInfo(TensorElementType.Float16, 2) }, { typeof(double), new TensorTypeInfo(TensorElementType.Double, 8) }, { typeof(uint), new TensorTypeInfo(TensorElementType.UInt32, 4) }, { typeof(ulong), new TensorTypeInfo(TensorElementType.UInt64, 8) }, { typeof(BFloat16), new TensorTypeInfo(TensorElementType.BFloat16, 2) } }; tensorElementTypeInfoMap = new Dictionary(); foreach (KeyValuePair item in typeInfoMap) { tensorElementTypeInfoMap.Add(item.Value.ElementType, new TensorElementTypeInfo(item.Key, item.Value.TypeSize)); } } protected TensorBase(Type primitiveType) { _primitiveType = primitiveType; } public static TensorTypeInfo GetTypeInfo(Type type) { TensorTypeInfo value = null; typeInfoMap.TryGetValue(type, out value); return value; } public static TensorElementTypeInfo GetElementTypeInfo(TensorElementType elementType) { TensorElementTypeInfo value = null; tensorElementTypeInfoMap.TryGetValue(elementType, out value); return value; } public TensorTypeInfo GetTypeInfo() { return GetTypeInfo(_primitiveType); } } public static class Tensor { public static Tensor CreateIdentity(int size) { return CreateIdentity(size, columMajor: false, Tensor.One); } public static Tensor CreateIdentity(int size, bool columMajor) { return CreateIdentity(size, columMajor, Tensor.One); } public static Tensor CreateIdentity(int size, bool columMajor, T oneValue) { Span span = stackalloc int[2]; span[0] = (span[1] = size); DenseTensor denseTensor = new DenseTensor(span, columMajor); for (int i = 0; i < size; i++) { denseTensor.SetValue(i * size + i, oneValue); } return denseTensor; } public static Tensor CreateFromDiagonal(Tensor diagonal) { return CreateFromDiagonal(diagonal, 0); } public static Tensor CreateFromDiagonal(Tensor diagonal, int offset) { if (diagonal.Rank < 1) { throw new ArgumentException("Tensor diagonal must have at least one dimension.", "diagonal"); } int num = diagonal.dimensions[0]; int num2 = diagonal.dimensions.Length + 1; Span span = ((num2 >= 16) ? ((Span)new int[num2]) : stackalloc int[num2]); Span span2 = span; int num3 = num + Math.Abs(offset); span2[0] = (span2[1] = num3); for (int i = 1; i < diagonal.dimensions.Length; i++) { span2[i + 1] = diagonal.dimensions[i]; } Tensor tensor = diagonal.CloneEmpty(span2); long num4 = diagonal.Length / num; int num5 = ((!diagonal.IsReversedStride || diagonal.Rank <= 1) ? 1 : diagonal.strides[1]); int num6 = ((!tensor.IsReversedStride || tensor.Rank <= 2) ? 1 : tensor.strides[2]); for (int j = 0; j < num; j++) { int num7 = ((offset < 0) ? (j - offset) : j); int num8 = ((offset > 0) ? (j + offset) : j); int num9 = num7 * tensor.strides[0] + num8 * tensor.strides[1]; int num10 = j * diagonal.strides[0]; for (int k = 0; k < num4; k++) { tensor.SetValue(num9 + k * num6, diagonal.GetValue(num10 + k * num5)); } } return tensor; } } [DebuggerDisplay("{GetArrayString(false)}")] public abstract class Tensor : TensorBase, IList, ICollection, IEnumerable, IList, ICollection, IEnumerable, IReadOnlyList, IReadOnlyCollection, IStructuralComparable, IStructuralEquatable { internal readonly int[] dimensions; internal readonly int[] strides; private readonly bool isReversedStride; private readonly long length; internal static T Zero { get { if (typeof(T) == typeof(bool)) { return (T)(object)false; } if (typeof(T) == typeof(byte)) { return (T)(object)(byte)0; } if (typeof(T) == typeof(char)) { return (T)(object)'\0'; } if (typeof(T) == typeof(decimal)) { return (T)(object)0m; } if (typeof(T) == typeof(double)) { return (T)(object)0.0; } if (typeof(T) == typeof(float)) { return (T)(object)0f; } if (typeof(T) == typeof(int)) { return (T)(object)0; } if (typeof(T) == typeof(long)) { return (T)(object)0L; } if (typeof(T) == typeof(sbyte)) { return (T)(object)(sbyte)0; } if (typeof(T) == typeof(short)) { return (T)(object)(short)0; } if (typeof(T) == typeof(uint)) { return (T)(object)0u; } if (typeof(T) == typeof(ulong)) { return (T)(object)0uL; } if (typeof(T) == typeof(ushort)) { return (T)(object)(ushort)0; } if (typeof(T) == typeof(Float16)) { return (T)(object)(ushort)0; } if (typeof(T) == typeof(BFloat16)) { return (T)(object)(ushort)0; } if (typeof(T) == typeof(string)) { return (T)(object)"0"; } throw new NotSupportedException(); } } internal static T One { get { if (typeof(T) == typeof(bool)) { return (T)(object)true; } if (typeof(T) == typeof(byte)) { return (T)(object)(byte)1; } if (typeof(T) == typeof(char)) { return (T)(object)'\u0001'; } if (typeof(T) == typeof(decimal)) { return (T)(object)1m; } if (typeof(T) == typeof(double)) { return (T)(object)1.0; } if (typeof(T) == typeof(float)) { return (T)(object)1f; } if (typeof(T) == typeof(int)) { return (T)(object)1; } if (typeof(T) == typeof(long)) { return (T)(object)1L; } if (typeof(T) == typeof(sbyte)) { return (T)(object)(sbyte)1; } if (typeof(T) == typeof(short)) { return (T)(object)(short)1; } if (typeof(T) == typeof(uint)) { return (T)(object)1u; } if (typeof(T) == typeof(ulong)) { return (T)(object)1uL; } if (typeof(T) == typeof(ushort)) { return (T)(object)(ushort)1; } if (typeof(T) == typeof(Float16)) { return (T)(object)(ushort)15360; } if (typeof(T) == typeof(BFloat16)) { return (T)(object)(ushort)16256; } if (typeof(T) == typeof(string)) { return (T)(object)"1"; } throw new NotSupportedException(); } } public long Length => length; public int Rank => dimensions.Length; public bool IsReversedStride => isReversedStride; public ReadOnlySpan Dimensions => dimensions; public ReadOnlySpan Strides => strides; public virtual T this[params int[] indices] { get { if (indices == null) { throw new ArgumentNullException("indices"); } ReadOnlySpan indices2 = new ReadOnlySpan(indices); return this[indices2]; } set { if (indices == null) { throw new ArgumentNullException("indices"); } ReadOnlySpan indices2 = new ReadOnlySpan(indices); this[indices2] = value; } } public virtual T this[ReadOnlySpan indices] { get { return GetValue(ArrayUtilities.GetIndex(strides, indices)); } set { SetValue(ArrayUtilities.GetIndex(strides, indices), value); } } int ICollection.Count => (int)Length; bool ICollection.IsSynchronized => false; object ICollection.SyncRoot => this; object IList.this[int index] { get { return GetValue(index); } set { try { SetValue(index, (T)value); } catch (InvalidCastException) { throw new ArgumentException($"The value \"{value}\" is not of type \"{typeof(T)}\" and cannot be used in this generic collection."); } } } public bool IsFixedSize => true; public bool IsReadOnly => false; int ICollection.Count => (int)Length; int IReadOnlyCollection.Count => (int)Length; T IList.this[int index] { get { return GetValue(index); } set { SetValue(index, value); } } T IReadOnlyList.this[int index] => GetValue(index); protected Tensor(int length) : base(typeof(T)) { dimensions = new int[1] { length }; strides = new int[1] { 1 }; isReversedStride = false; this.length = length; } protected Tensor(ReadOnlySpan dimensions, bool reverseStride) : base(typeof(T)) { this.dimensions = new int[dimensions.Length]; long num = 1L; for (int i = 0; i < dimensions.Length; i++) { if (dimensions[i] < 0) { throw new ArgumentOutOfRangeException("dimensions", "Dimensions must be non-negative"); } this.dimensions[i] = dimensions[i]; num *= dimensions[i]; } strides = ArrayUtilities.GetStrides(dimensions, reverseStride); isReversedStride = reverseStride; length = num; } protected Tensor(Array fromArray, bool reverseStride) : base(typeof(T)) { if (fromArray == null) { throw new ArgumentNullException("fromArray"); } dimensions = new int[fromArray.Rank]; long num = 1L; for (int i = 0; i < dimensions.Length; i++) { dimensions[i] = fromArray.GetLength(i); num *= dimensions[i]; } strides = ArrayUtilities.GetStrides(dimensions, reverseStride); isReversedStride = reverseStride; length = num; } public virtual void Fill(T value) { for (int i = 0; i < Length; i++) { SetValue(i, value); } } public abstract Tensor Clone(); public virtual Tensor CloneEmpty() { return CloneEmpty(dimensions); } public virtual Tensor CloneEmpty(ReadOnlySpan dimensions) { return CloneEmpty(dimensions); } public virtual Tensor CloneEmpty() { return CloneEmpty(dimensions); } public abstract Tensor CloneEmpty(ReadOnlySpan dimensions); public Tensor GetDiagonal() { return GetDiagonal(0); } public Tensor GetDiagonal(int offset) { if (Rank < 2) { throw new InvalidOperationException("Cannot compute diagonal of Tensor with Rank less than 2."); } int num = dimensions[0]; int num2 = dimensions[1]; int val = ((offset < 0) ? (num + offset) : num); int val2 = ((offset > 0) ? (num2 - offset) : num2); int num3 = Math.Min(val, val2); if (num3 <= 0) { throw new ArgumentException($"Cannot compute diagonal with offset {offset}", "offset"); } int num4 = Rank - 1; Span span = ((num4 >= 16) ? ((Span)new int[num4]) : stackalloc int[num4]); Span span2 = span; span2[0] = num3; for (int i = 2; i < dimensions.Length; i++) { span2[i - 1] = dimensions[i]; } Tensor tensor = CloneEmpty(span2); long num5 = tensor.Length / tensor.Dimensions[0]; int num6 = ((!tensor.IsReversedStride || tensor.Rank <= 1) ? 1 : tensor.strides[1]); int num7 = ((!IsReversedStride || Rank <= 2) ? 1 : strides[2]); for (int j = 0; j < num3; j++) { int num8 = ((offset < 0) ? (j - offset) : j); int num9 = ((offset > 0) ? (j + offset) : j); int num10 = num8 * strides[0] + num9 * strides[1]; int num11 = j * tensor.strides[0]; for (int k = 0; k < num5; k++) { tensor.SetValue(num11 + k * num6, GetValue(num10 + k * num7)); } } return tensor; } public Tensor GetTriangle() { return GetTriangle(0, upper: false); } public Tensor GetTriangle(int offset) { return GetTriangle(offset, upper: false); } public Tensor GetUpperTriangle() { return GetTriangle(0, upper: true); } public Tensor GetUpperTriangle(int offset) { return GetTriangle(offset, upper: true); } public Tensor GetTriangle(int offset, bool upper) { if (Rank < 2) { throw new InvalidOperationException("Cannot compute triangle of Tensor with Rank less than 2."); } int num = dimensions[0]; int num2 = dimensions[1]; int num3 = Math.Max(num, num2); Tensor tensor = CloneEmpty(); long num4 = Length / (num * num2); int num5 = ((!IsReversedStride || Rank <= 2) ? 1 : strides[2]); for (int i = 0; i < num3; i++) { int num6 = ((offset > 0) ? (i - offset) : i); int num7 = ((offset > 0) ? i : (i + offset)); if (num6 < 0) { if (upper) { continue; } num6 = 0; } if (num7 < 0) { if (!upper) { continue; } num7 = 0; } while (num7 < num2 && num6 < num) { int num8 = num6 * strides[0] + num7 * tensor.strides[1]; for (int j = 0; j < num4; j++) { int index = num8 + j * num5; tensor.SetValue(index, GetValue(index)); } if (upper) { num7++; } else { num6++; } } } return tensor; } public abstract Tensor Reshape(ReadOnlySpan dimensions); public abstract T GetValue(int index); public abstract void SetValue(int index, T value); public static int Compare(Tensor left, Tensor right) { return StructuralComparisons.StructuralComparer.Compare(left, right); } public static bool Equals(Tensor left, Tensor right) { return StructuralComparisons.StructuralEqualityComparer.Equals(left, right); } IEnumerator IEnumerable.GetEnumerator() { return ((IEnumerable)this).GetEnumerator(); } void ICollection.CopyTo(Array array, int index) { if (array is T[] array2) { CopyTo(array2, index); return; } if (array == null) { throw new ArgumentNullException("array"); } if (array.Rank != 1) { throw new ArgumentException("Only single dimensional arrays are supported for the requested action.", "array"); } if (array.Length < index + Length) { throw new ArgumentException("The number of elements in the Tensor is greater than the available space from index to the end of the destination array.", "array"); } for (int i = 0; i < length; i++) { array.SetValue(GetValue(i), index + i); } } int IList.Add(object value) { throw new InvalidOperationException(); } void IList.Clear() { Fill(default(T)); } bool IList.Contains(object value) { if (IsCompatibleObject(value)) { return Contains((T)value); } return false; } int IList.IndexOf(object value) { if (IsCompatibleObject(value)) { return IndexOf((T)value); } return -1; } void IList.Insert(int index, object value) { throw new InvalidOperationException(); } void IList.Remove(object value) { throw new InvalidOperationException(); } void IList.RemoveAt(int index) { throw new InvalidOperationException(); } IEnumerator IEnumerable.GetEnumerator() { for (int i = 0; i < Length; i++) { yield return GetValue(i); } } void ICollection.Add(T item) { throw new InvalidOperationException(); } void ICollection.Clear() { Fill(default(T)); } bool ICollection.Contains(T item) { return Contains(item); } protected virtual bool Contains(T item) { return Length != 0L && IndexOf(item) != -1; } void ICollection.CopyTo(T[] array, int arrayIndex) { CopyTo(array, arrayIndex); } protected virtual void CopyTo(T[] array, int arrayIndex) { if (array == null) { throw new ArgumentNullException("array"); } if (array.Length < arrayIndex + Length) { throw new ArgumentException("The number of elements in the Tensor is greater than the available space from index to the end of the destination array.", "array"); } for (int i = 0; i < length; i++) { array[arrayIndex + i] = GetValue(i); } } bool ICollection.Remove(T item) { throw new InvalidOperationException(); } int IList.IndexOf(T item) { return IndexOf(item); } protected virtual int IndexOf(T item) { for (int i = 0; i < Length; i++) { if (GetValue(i).Equals(item)) { return i; } } return -1; } void IList.Insert(int index, T item) { throw new InvalidOperationException(); } void IList.RemoveAt(int index) { throw new InvalidOperationException(); } int IStructuralComparable.CompareTo(object other, IComparer comparer) { if (other == null) { return 1; } if (other is Tensor) { return CompareTo((Tensor)other, comparer); } if (other is Array other2) { return CompareTo(other2, comparer); } throw new ArgumentException(string.Format("Cannot compare {0} to {1}.", "Tensor", other.GetType()), "other"); } private int CompareTo(Tensor other, IComparer comparer) { if (Rank != other.Rank) { throw new ArgumentException(string.Format("Cannot compare {0} with Rank {1} to {2} with Rank {3}.", "Tensor", Rank, "other", other.Rank), "other"); } for (int i = 0; i < dimensions.Length; i++) { if (dimensions[i] != other.dimensions[i]) { throw new ArgumentException(string.Format("Cannot compare {0}s with differning dimension {1}, {2} != {3}.", "Tensor", i, dimensions[i], other.dimensions[i]), "other"); } } int num = 0; if (IsReversedStride == other.IsReversedStride) { for (int j = 0; j < Length; j++) { num = comparer.Compare(GetValue(j), other.GetValue(j)); if (num != 0) { break; } } } else { Span span = ((Rank >= 16) ? ((Span)new int[Rank]) : stackalloc int[Rank]); Span span2 = span; for (int k = 0; k < Length; k++) { ArrayUtilities.GetIndices(strides, IsReversedStride, k, span2); num = comparer.Compare(this[span2], other[span2]); if (num != 0) { break; } } } return num; } private int CompareTo(Array other, IComparer comparer) { if (Rank != other.Rank) { throw new ArgumentException(string.Format("Cannot compare {0} with Rank {1} to {2} with rank {3}.", "Tensor", Rank, "Array", other.Rank), "other"); } for (int i = 0; i < dimensions.Length; i++) { int num = other.GetLength(i); if (dimensions[i] != num) { throw new ArgumentException(string.Format("Cannot compare {0} to {1} with differning dimension {2}, {3} != {4}.", "Tensor", "Array", i, dimensions[i], num), "other"); } } int num2 = 0; int[] indices = new int[Rank]; for (int j = 0; j < Length; j++) { ArrayUtilities.GetIndices(strides, IsReversedStride, j, indices); num2 = comparer.Compare(GetValue(j), other.GetValue(indices)); if (num2 != 0) { break; } } return num2; } bool IStructuralEquatable.Equals(object other, IEqualityComparer comparer) { if (other == null) { return false; } if (other is Tensor) { return Equals((Tensor)other, comparer); } if (other is Array other2) { return Equals(other2, comparer); } throw new ArgumentException(string.Format("Cannot compare {0} to {1}.", "Tensor", other.GetType()), "other"); } private bool Equals(Tensor other, IEqualityComparer comparer) { if (Rank != other.Rank) { throw new ArgumentException(string.Format("Cannot compare {0} with Rank {1} to {2} with Rank {3}.", "Tensor", Rank, "other", other.Rank), "other"); } for (int i = 0; i < dimensions.Length; i++) { if (dimensions[i] != other.dimensions[i]) { throw new ArgumentException(string.Format("Cannot compare {0}s with differning dimension {1}, {2} != {3}.", "Tensor", i, dimensions[i], other.dimensions[i]), "other"); } } if (IsReversedStride == other.IsReversedStride) { for (int j = 0; j < Length; j++) { if (!comparer.Equals(GetValue(j), other.GetValue(j))) { return false; } } } else { Span span = ((Rank >= 16) ? ((Span)new int[Rank]) : stackalloc int[Rank]); Span span2 = span; for (int k = 0; k < Length; k++) { ArrayUtilities.GetIndices(strides, IsReversedStride, k, span2); if (!comparer.Equals(this[span2], other[span2])) { return false; } } } return true; } private bool Equals(Array other, IEqualityComparer comparer) { if (Rank != other.Rank) { throw new ArgumentException(string.Format("Cannot compare {0} with Rank {1} to {2} with rank {3}.", "Tensor", Rank, "Array", other.Rank), "other"); } for (int i = 0; i < dimensions.Length; i++) { int num = other.GetLength(i); if (dimensions[i] != num) { throw new ArgumentException(string.Format("Cannot compare {0} to {1} with differning dimension {2}, {3} != {4}.", "Tensor", "Array", i, dimensions[i], num), "other"); } } int[] indices = new int[Rank]; for (int j = 0; j < Length; j++) { ArrayUtilities.GetIndices(strides, IsReversedStride, j, indices); if (!comparer.Equals(GetValue(j), other.GetValue(indices))) { return false; } } return true; } int IStructuralEquatable.GetHashCode(IEqualityComparer comparer) { int num = 0; for (int i = 0; i < Length; i++) { num ^= comparer.GetHashCode(GetValue(i)); } return num; } public virtual DenseTensor ToDenseTensor() { DenseTensor denseTensor = new DenseTensor(Dimensions, IsReversedStride); for (int i = 0; i < Length; i++) { denseTensor.SetValue(i, GetValue(i)); } return denseTensor; } public string GetArrayString(bool includeWhitespace = true) { StringBuilder stringBuilder = new StringBuilder(); int[] array = ArrayUtilities.GetStrides(dimensions); int[] array2 = new int[Rank]; int num = Rank - 1; int num2 = dimensions[num]; long num3 = Length / num2; int num4 = 0; for (int i = 0; i < Length; i += num2) { ArrayUtilities.GetIndices(array, reverseStride: false, i, array2); while (num4 < num && array2[num4] == 0) { if (includeWhitespace) { Indent(stringBuilder, num4); } num4++; stringBuilder.Append('{'); if (includeWhitespace) { stringBuilder.AppendLine(); } } for (int j = 0; j < num2; j++) { array2[num] = j; if (j == 0) { if (includeWhitespace) { Indent(stringBuilder, num4); } stringBuilder.Append('{'); } else { stringBuilder.Append(','); } stringBuilder.Append(this[array2]); } stringBuilder.Append('}'); int num5 = Rank - 2; while (num5 >= 0) { int num6 = dimensions[num5] - 1; if (array2[num5] == num6) { num4--; if (includeWhitespace) { stringBuilder.AppendLine(); Indent(stringBuilder, num4); } stringBuilder.Append('}'); num5--; continue; } stringBuilder.Append(','); if (includeWhitespace) { stringBuilder.AppendLine(); } break; } } return stringBuilder.ToString(); } private static void Indent(StringBuilder builder, int tabs, int spacesPerTab = 4) { for (int i = 0; i < tabs; i++) { for (int j = 0; j < spacesPerTab; j++) { builder.Append(' '); } } } private static bool IsCompatibleObject(object value) { return value is T || (value == null && default(T) == null); } } }