diff --git a/include/onnxruntime/core/framework/kernel_registry.h b/include/onnxruntime/core/framework/kernel_registry.h index 9f391daa31..224ce00bc7 100644 --- a/include/onnxruntime/core/framework/kernel_registry.h +++ b/include/onnxruntime/core/framework/kernel_registry.h @@ -45,6 +45,11 @@ class KernelRegistry { Status TryFindKernel(const Node& node, ProviderType exec_provider, const KernelCreateInfo** out) const; + // Find KernelCreateInfo in instant mode + Status TryFindKernel(const std::string& op_name, const std::string& domain, const int& version, + const std::unordered_map& type_constraints, + ProviderType exec_provider, const KernelCreateInfo** out) const; + #endif // Try to find the kernel given a kernel def hash. diff --git a/include/onnxruntime/core/framework/op_kernel_context.h b/include/onnxruntime/core/framework/op_kernel_context.h index df3363a5eb..cf9cd3ac68 100644 --- a/include/onnxruntime/core/framework/op_kernel_context.h +++ b/include/onnxruntime/core/framework/op_kernel_context.h @@ -21,10 +21,10 @@ class OpKernelContext { @param arg_num The operator argument number. @returns Number of inputs the argument has. */ - int NumVariadicInputs(size_t arg_num) const; + virtual int NumVariadicInputs(size_t arg_num) const; - MLDataType InputType(int index) const; - MLDataType OutputType(int index) const; + virtual MLDataType InputType(int index) const; + virtual MLDataType OutputType(int index) const; const OrtValue* GetInputOrtValue(int index) const { return GetInputMLValue(index); @@ -98,29 +98,29 @@ class OpKernelContext { // Retrieve indexed shape obtained from memory planning before actual // computation. If the indexed shape cannot be inferred, this function returns // false. - bool TryGetInferredInputShape(int index, TensorShape& shape) const; + virtual bool TryGetInferredInputShape(int index, TensorShape& shape) const; // Retrieve indexed shape obtained from memory planning before actual // computation. If the indexed shape cannot be inferred, this function returns // false. - bool TryGetInferredOutputShape(int index, TensorShape& shape) const; + virtual bool TryGetInferredOutputShape(int index, TensorShape& shape) const; const logging::Logger& Logger() const { return *logger_; } // always >= 0 - int InputCount() const { + virtual int InputCount() const { return static_cast(kernel_->Node().InputDefs().size()); } // always >= 0 - int ImplicitInputCount() const { + virtual int ImplicitInputCount() const { return static_cast(kernel_->Node().ImplicitInputDefs().size()); } // always >= 0 - int OutputCount() const { + virtual int OutputCount() const { return static_cast(kernel_->Node().OutputDefs().size()); } @@ -128,7 +128,7 @@ class OpKernelContext { Return an allocator on device 0, with memtype of OrtMemTypeDefault. @remarks Use SafeInt when calculating the size of memory to allocate using AllocatorPtr->Alloc. */ - Status GetTempSpaceAllocator(AllocatorPtr* output) const ORT_MUST_USE_RESULT; + virtual Status GetTempSpaceAllocator(AllocatorPtr* output) const ORT_MUST_USE_RESULT; /** Return the allocator associated with the CPU EP with memtype of OrtMemTypeDefault. @@ -142,7 +142,7 @@ class OpKernelContext { @returns Point to the Fence of the input OrtValue. It is null if the input OrtValue doesn't have fence or the input is optional. */ - Fence_t InputFence(int index) const; + virtual Fence_t InputFence(int index) const; /** Return the fence of current node's implicit input. @@ -150,7 +150,7 @@ class OpKernelContext { @returns Point to the Fence of the implicit input OrtValue. It is null if the input OrtValue doesn't have fence or the input is optional. */ - Fence_t ImplicitInputFence(int index) const; + virtual Fence_t ImplicitInputFence(int index) const; /** Return the fence of current node's output identifed by index. @@ -158,12 +158,12 @@ class OpKernelContext { @returns Point to the Fence of the output OrtValue. It is null if the output OrtValue doesn't have fence or the output is optional. */ - Fence_t OutputFence(int index) const; + virtual Fence_t OutputFence(int index) const; /** Return the device id that current kernel runs on. */ - int GetDeviceId() const { + virtual int GetDeviceId() const { return kernel_->Info().GetExecutionProvider()->GetDeviceId(); } @@ -171,7 +171,7 @@ class OpKernelContext { Return the compute stream associated with the EP that the kernel is partitioned to. For EPs that do not have a compute stream (e.g. CPU EP), a nullptr is returned. */ - void* GetComputeStream() const { + virtual void* GetComputeStream() const { return kernel_->Info().GetExecutionProvider()->GetComputeStream(); } @@ -203,9 +203,12 @@ class OpKernelContext { } protected: + + OpKernelContext(concurrency::ThreadPool* threadpool, const logging::Logger& logger); + onnxruntime::NodeIndex GetNodeIndex() const; - const OrtValue* GetInputMLValue(int index) const; + virtual const OrtValue* GetInputMLValue(int index) const; const OrtValue* GetImplicitInputMLValue(int index) const; OrtValue* GetOutputMLValue(int index); @@ -214,21 +217,20 @@ class OpKernelContext { #endif // Creates the OrtValue* based on the shape, if it does not exist - OrtValue* OutputMLValue(int index, const TensorShape& shape); + virtual OrtValue* OutputMLValue(int index, const TensorShape& shape); + + virtual OrtValue* GetOrCreateOutputMLValue(int index); private: ORT_DISALLOW_COPY_AND_ASSIGNMENT(OpKernelContext); - - OrtValue* GetOrCreateOutputMLValue(int index); - int GetInputArgIndex(int index) const; int GetImplicitInputArgIndex(int index) const; int GetOutputArgIndex(int index) const; - IExecutionFrame* const execution_frame_; - const OpKernel* const kernel_; - concurrency::ThreadPool* const threadpool_; - const logging::Logger* const logger_; + IExecutionFrame* const execution_frame_{}; + const OpKernel* const kernel_{}; + concurrency::ThreadPool* const threadpool_{}; + const logging::Logger* const logger_{}; // The argument starting index in ExecutionFrame. int node_input_start_index_{-1}; diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index c762d48297..3c26379826 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -79,6 +79,7 @@ class Node { Fused = 1, ///< The node refers to a function. }; + explicit Node() = default; ~Node() = default; /** @@ -620,7 +621,7 @@ class Node { NodeAttributes attributes_; // Graph that contains this Node - Graph* graph_; + Graph* graph_ = nullptr; // Map of attribute name to the Graph instance created from the GraphProto attribute std::unordered_map> attr_to_subgraph_map_; diff --git a/include/onnxruntime/core/session/onnxruntime_c_api.h b/include/onnxruntime/core/session/onnxruntime_c_api.h index 922adca4ab..9c0080a467 100644 --- a/include/onnxruntime/core/session/onnxruntime_c_api.h +++ b/include/onnxruntime/core/session/onnxruntime_c_api.h @@ -227,6 +227,16 @@ typedef enum OrtErrorCode { ORT_EP_FAIL, } OrtErrorCode; +typedef enum OrtOpAttrType { + ORT_OP_ATTR_UNDEFINED = 0, + ORT_OP_ATTR_INT, + ORT_OP_ATTR_INTS, + ORT_OP_ATTR_FLOAT, + ORT_OP_ATTR_FLOATS, + ORT_OP_ATTR_STRING, + ORT_OP_ATTR_STRINGS, +} OrtOpAttrType; + //! @} #define ORT_RUNTIME_CLASS(X) \ struct Ort##X; \ @@ -257,6 +267,8 @@ ORT_RUNTIME_CLASS(ArenaCfg); ORT_RUNTIME_CLASS(PrepackedWeightsContainer); ORT_RUNTIME_CLASS(TensorRTProviderOptionsV2); ORT_RUNTIME_CLASS(CUDAProviderOptionsV2); +ORT_RUNTIME_CLASS(Op); +ORT_RUNTIME_CLASS(OpAttr); #ifdef _WIN32 typedef _Return_type_success_(return == 0) OrtStatus* OrtStatusPtr; @@ -3329,6 +3341,86 @@ struct OrtApi { ORT_API2_STATUS(AddExternalInitializers, _In_ OrtSessionOptions* options, _In_reads_(input_len) const char* const* initializer_names, _In_reads_(input_len) const OrtValue* const* initializers, size_t initializers_num); + + /** \brief: Create attribute of onnxruntime operator + * + * \param[in] name of the attribute + * \param[in] data of the attribute + * \param[in] data length + * \param[in] data type + * \param[out] attribute that has been created, which must be released by OrtApi::ReleaseOpAttr + * + * \since Version 1.12. + */ + ORT_API2_STATUS(CreateOpAttr, + _In_ const char* name, + _In_ const void* data, + _In_ int len, + _In_ OrtOpAttrType type, + _Outptr_ OrtOpAttr** op_attr); + + /* \brief: Release op attribute + * + * \param[in] attribute created by OrtApi::CreateOpAttr + * + * \since Version 1.12. + */ + ORT_CLASS_RELEASE(OpAttr); + + /** \brief: Create onnxruntime native operator + * + * \param[in] kernel info + * \param[in] operator name + * \param[in] operator domain + * \param[in] operator opset + * \param[in] name of the type contraints, such as "T" or "T1" + * \param[in] type of each contraints + * \param[in] number of contraints + * \param[in] attributes used to initialize the operator + * \param[in] number of the attributes + * \param[out] operator that has been created + * + * \since Version 1.12. + */ + ORT_API2_STATUS(CreateOp, + _In_ const OrtKernelInfo* info, + _In_ const char* op_name, + _In_ const char* domain, + _In_ int version, + _In_opt_ const char** type_constraint_names, + _In_opt_ const ONNXTensorElementDataType* type_constraint_values, + _In_opt_ int type_constraint_count, + _In_opt_ const OrtOpAttr* const* attr_values, + _In_opt_ int attr_count, + _Outptr_ OrtOp** ort_op); + + /** \brief: Invoke the operator created by OrtApi::CreateOp + * The inputs must follow the order as specified in onnx specification + * + * \param[in] kernel context + * \param[in] operator that has been created + * \param[in] inputs + * \param[in] number of inputs + * \param[in] outputs + * \param[in] number of outputs + * + * \since Version 1.12. + */ + ORT_API2_STATUS(InvokeOp, + _In_ const OrtKernelContext* context, + _In_ const OrtOp* ort_op, + _In_ const OrtValue* const* input_values, + _In_ int input_count, + _Inout_ OrtValue* const* output_values, + _In_ int output_count); + + /* \brief: Release an onnxruntime operator + * + * \param[in] operator created by OrtApi::CreateOp + * + * \since Version 1.12. + */ + ORT_CLASS_RELEASE(Op); }; /* @@ -3398,6 +3490,7 @@ ORT_API_STATUS(OrtSessionOptionsAppendExecutionProvider_CUDA, _In_ OrtSessionOpt * \param device_id HIP device id, starts from zero. */ ORT_API_STATUS(OrtSessionOptionsAppendExecutionProvider_MIGraphX, _In_ OrtSessionOptions* options, int device_id); + #ifdef __cplusplus } #endif diff --git a/include/onnxruntime/core/session/onnxruntime_cxx_api.h b/include/onnxruntime/core/session/onnxruntime_cxx_api.h index 7ae2a91ace..c6cb6ed36a 100644 --- a/include/onnxruntime/core/session/onnxruntime_cxx_api.h +++ b/include/onnxruntime/core/session/onnxruntime_cxx_api.h @@ -973,6 +973,34 @@ struct CustomOpApi { void ThrowOnError(OrtStatus* result); + void CreateOpAttr(_In_ const char* name, + _In_ const void* data, + _In_ int len, + _In_ OrtOpAttrType type, + _Outptr_ OrtOpAttr** op_attr); + + void ReleaseOpAttr(_In_ OrtOpAttr* op_attr); + + void CreateOp(_In_ const OrtKernelInfo* info, + _In_ const char* op_name, + _In_ const char* domain, + _In_ int version, + _In_opt_ const char** type_constraint_names, + _In_opt_ const ONNXTensorElementDataType* type_constraint_values, + _In_opt_ int type_constraint_count, + _In_opt_ const OrtOpAttr* const* attr_values, + _In_opt_ int attr_count, + _Outptr_ OrtOp** ort_op); + + void InvokeOp(_In_ const OrtKernelContext* context, + _In_ const OrtOp* ort_op, + _In_ const OrtValue* const* input_values, + _In_ int input_count, + _Inout_ OrtValue* const* output_values, + _In_ int output_count); + + void ReleaseOp(_In_ OrtOp* ort_op); + private: const OrtApi& api_; }; diff --git a/include/onnxruntime/core/session/onnxruntime_cxx_inline.h b/include/onnxruntime/core/session/onnxruntime_cxx_inline.h index 32ade6f1ff..f944590c28 100644 --- a/include/onnxruntime/core/session/onnxruntime_cxx_inline.h +++ b/include/onnxruntime/core/session/onnxruntime_cxx_inline.h @@ -1206,6 +1206,44 @@ inline void* CustomOpApi::KernelContext_GetGPUComputeStream(const OrtKernelConte return out; } +inline void CustomOpApi::CreateOpAttr(_In_ const char* name, + _In_ const void* data, + _In_ int len, + _In_ OrtOpAttrType type, + _Outptr_ OrtOpAttr** op_attr) { + ThrowOnError(api_.CreateOpAttr(name, data, len, type, op_attr)); +} + +inline void CustomOpApi::ReleaseOpAttr(_Frees_ptr_opt_ OrtOpAttr* op_attr) { + api_.ReleaseOpAttr(op_attr); +} + +inline void CustomOpApi::CreateOp(_In_ const OrtKernelInfo* info, + _In_ const char* op_name, + _In_ const char* domain, + _In_ int version, + _In_opt_ const char** type_constraint_names, + _In_opt_ const ONNXTensorElementDataType* type_constraint_values, + _In_opt_ int type_constraint_count, + _In_opt_ const OrtOpAttr* const* attr_values, + _In_opt_ int attr_count, + _Outptr_ OrtOp** ort_op) { + ThrowOnError(api_.CreateOp(info, op_name, domain, version, type_constraint_names, type_constraint_values, type_constraint_count, attr_values, attr_count, ort_op)); +} + +inline void CustomOpApi::InvokeOp(_In_ const OrtKernelContext* context, + _In_ const OrtOp* ort_op, + _In_ const OrtValue* const* input_values, + _In_ int input_count, + _Inout_ OrtValue* const* output_values, + _In_ int output_count) { + ThrowOnError(api_.InvokeOp(context, ort_op, input_values, input_count, output_values, output_count)); +} + +inline void CustomOpApi::ReleaseOp(_Frees_ptr_opt_ OrtOp* ort_op) { + api_.ReleaseOp(ort_op); +} + inline SessionOptions& SessionOptions::DisablePerSessionThreads() { ThrowOnError(GetApi().DisablePerSessionThreads(p_)); return *this; diff --git a/onnxruntime/core/framework/kernel_registry.cc b/onnxruntime/core/framework/kernel_registry.cc index e77eadc010..90b0b190a8 100644 --- a/onnxruntime/core/framework/kernel_registry.cc +++ b/onnxruntime/core/framework/kernel_registry.cc @@ -315,6 +315,35 @@ Status KernelRegistry::TryFindKernel(const Node& node, return Status(common::ONNXRUNTIME, common::FAIL, "Kernel not found"); } + +Status KernelRegistry::TryFindKernel(const std::string& op_name, const std::string& domain, const int& version, + const std::unordered_map& type_constraints, + ProviderType exec_provider, const KernelCreateInfo** out) const { + *out = nullptr; + auto range = kernel_creator_fn_map_.equal_range(GetMapKey(op_name, domain, exec_provider)); + for (auto i = range.first; i != range.second; ++i) { //loop through all kernels + const KernelCreateInfo& kci = i->second; + int start_ver{}; + int end_ver{}; + kci.kernel_def->SinceVersion(&start_ver, &end_ver); + if (start_ver <= version && end_ver >= version) { //try match the version + auto& kci_constraints = kci.kernel_def->TypeConstraints(); + bool match = true; + for (auto& constraint : type_constraints) { //try match type constraints + auto iter = kci_constraints.find(constraint.first); + if (iter == kci_constraints.end() || find(iter->second.begin(), iter->second.end(), constraint.second) == iter->second.end()) { + match = false; + break; + } + } //for + if (match) { + *out = &kci; //found match, exit loop + break; + } + } //if + } //for + return *out == nullptr ? Status(common::ONNXRUNTIME, common::FAIL, "Kernel not found") : Status::OK(); +} #endif // !defined(ORT_MINIMAL_BUILD) bool KernelRegistry::TryFindKernelByHash(HashValue kernel_def_hash, const KernelCreateInfo** out) const { diff --git a/onnxruntime/core/framework/op_kernel.cc b/onnxruntime/core/framework/op_kernel.cc index f2f63c947e..a1b0f2fc26 100644 --- a/onnxruntime/core/framework/op_kernel.cc +++ b/onnxruntime/core/framework/op_kernel.cc @@ -36,6 +36,9 @@ OpKernelContext::OpKernelContext(_Inout_ IExecutionFrame* frame, _In_ const OpKe node_output_start_index_ = node_implicit_input_start_index_ + ImplicitInputCount(); } +OpKernelContext::OpKernelContext(concurrency::ThreadPool* threadpool, + const logging::Logger& logger) : threadpool_(threadpool), logger_(&logger) {} + Tensor* OpKernelContext::Output(int index, const TensorShape& shape) { auto p_ml_value = OutputMLValue(index, shape); return p_ml_value ? p_ml_value->GetMutable() : nullptr; diff --git a/onnxruntime/core/framework/op_kernel_context_internal.h b/onnxruntime/core/framework/op_kernel_context_internal.h index 4f889a23d7..8e3a0747c0 100644 --- a/onnxruntime/core/framework/op_kernel_context_internal.h +++ b/onnxruntime/core/framework/op_kernel_context_internal.h @@ -45,7 +45,7 @@ class OpKernelContextInternal : public OpKernelContext { return session_state_.GetSubgraphSessionState(GetNodeIndex(), attribute_name); } - const OrtValue* GetInputMLValue(int index) const { + const OrtValue* GetInputMLValue(int index) const override { return OpKernelContext::GetInputMLValue(index); } @@ -59,7 +59,7 @@ class OpKernelContextInternal : public OpKernelContext { } #endif - OrtValue* OutputMLValue(int index, const TensorShape& shape) { + OrtValue* OutputMLValue(int index, const TensorShape& shape) override { return OpKernelContext::OutputMLValue(index, shape); } diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 35fbb9398b..a35b98f943 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -896,9 +896,10 @@ void Node::CreateSubgraph(const std::string& attr_name) { void Node::AddAttributeProto(AttributeProto value) { utils::SetNodeAttribute(std::move(value), attributes_); - - graph_->SetGraphResolveNeeded(); - graph_->SetGraphProtoSyncNeeded(); + if (graph_) { + graph_->SetGraphResolveNeeded(); + graph_->SetGraphProtoSyncNeeded(); + } } #define ADD_ATTR_SINGLE_IMPL(Type) \ diff --git a/onnxruntime/core/session/onnxruntime_c_api.cc b/onnxruntime/core/session/onnxruntime_c_api.cc index fe8580f1a0..cf114077a9 100644 --- a/onnxruntime/core/session/onnxruntime_c_api.cc +++ b/onnxruntime/core/session/onnxruntime_c_api.cc @@ -2522,6 +2522,11 @@ static constexpr OrtApi ort_api_1_to_12 = { &OrtApis::SessionOptionsAppendExecutionProvider_MIGraphX, // End of Version 11 - DO NOT MODIFY ABOVE (see above text for more information) &OrtApis::AddExternalInitializers, + &OrtApis::CreateOpAttr, + &OrtApis::ReleaseOpAttr, + &OrtApis::CreateOp, + &OrtApis::InvokeOp, + &OrtApis::ReleaseOp, }; // Asserts to do a some checks to ensure older Versions of the OrtApi never change (will detect an addition or deletion but not if they cancel out each other) diff --git a/onnxruntime/core/session/ort_apis.h b/onnxruntime/core/session/ort_apis.h index 98a95c8941..4673260085 100644 --- a/onnxruntime/core/session/ort_apis.h +++ b/onnxruntime/core/session/ort_apis.h @@ -339,8 +339,40 @@ ORT_API_STATUS_IMPL(UpdateCUDAProviderOptions, _Inout_ OrtCUDAProviderOptionsV2* size_t num_keys); ORT_API_STATUS_IMPL(GetCUDAProviderOptionsAsString, _In_ const OrtCUDAProviderOptionsV2* cuda_options, _Inout_ OrtAllocator* allocator, _Outptr_ char** ptr); ORT_API(void, ReleaseCUDAProviderOptions, _Frees_ptr_opt_ OrtCUDAProviderOptionsV2*); + ORT_API_STATUS_IMPL(AddExternalInitializers, _In_ OrtSessionOptions* options, _In_reads_(initializers_num) const char* const* initializer_names, _In_reads_(initializers_num) const OrtValue* const* initializers, size_t initializers_num); +ORT_API_STATUS_IMPL(CreateOpAttr, + _In_ const char* name, + _In_ const void* data, + _In_ int len, + _In_ OrtOpAttrType type, + _Outptr_ OrtOpAttr** op_attr); + +ORT_API(void, ReleaseOpAttr, _Frees_ptr_opt_ OrtOpAttr* op_attr); + +ORT_API_STATUS_IMPL(CreateOp, + _In_ const OrtKernelInfo* info, + _In_ const char* op_name, + _In_ const char* domain, + _In_ int version, + _In_opt_ const char** type_constraint_names, + _In_opt_ const ONNXTensorElementDataType* type_constraint_values, + _In_opt_ int type_constraint_count, + _In_opt_ const OrtOpAttr* const* attr_values, + _In_opt_ int attr_count, + _Outptr_ OrtOp** ort_op); + +ORT_API_STATUS_IMPL(InvokeOp, + _In_ const OrtKernelContext* context, + _In_ const OrtOp* ort_op, + _In_ const OrtValue* const* input_values, + _In_ int input_count, + _Inout_ OrtValue* const* output_values, + _In_ int output_count); + +ORT_API(void, ReleaseOp, _Frees_ptr_opt_ OrtOp* op); + } // namespace OrtApis diff --git a/onnxruntime/core/session/standalone_op_invoker.cc b/onnxruntime/core/session/standalone_op_invoker.cc new file mode 100644 index 0000000000..c39541b3d9 --- /dev/null +++ b/onnxruntime/core/session/standalone_op_invoker.cc @@ -0,0 +1,407 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/session/inference_session.h" +#include "core/framework/kernel_registry.h" +#include "core/framework/error_code_helper.h" +#include "core/session/ort_apis.h" +#include + +#ifdef ORT_MINIMAL_BUILD + +ORT_API_STATUS_IMPL(OrtApis::CreateOpAttr, + _In_ const char*, + _In_ const void*, + _In_ int, + _In_ OrtOpAttrType, + _Out_ OrtOpAttr**) { + API_IMPL_BEGIN + return CreateStatus(ORT_NOT_IMPLEMENTED, "CreateOpAttr is not implemented for minimal build."); + API_IMPL_END +} + +ORT_API(void, OrtApis::ReleaseOpAttr, _Frees_ptr_opt_ OrtOpAttr*) { +} + +ORT_API_STATUS_IMPL(OrtApis::CreateOp, + _In_ const OrtKernelInfo*, + _In_ const char*, + _In_ const char*, + _In_ int, + _In_ const char**, + _In_ const ONNXTensorElementDataType*, + _In_ int, + _In_ const OrtOpAttr* const*, + _In_ int, + _Out_ OrtOp**) { + API_IMPL_BEGIN + return CreateStatus(ORT_NOT_IMPLEMENTED, "CreateOp is not implemented for minimal build."); + API_IMPL_END +} + +ORT_API_STATUS_IMPL(OrtApis::InvokeOp, + _In_ const OrtKernelContext*, + _In_ const OrtOp*, + _In_ const OrtValue* const*, + _In_ int, + _Inout_ OrtValue* const*, + _In_ int) { + API_IMPL_BEGIN + return CreateStatus(ORT_NOT_IMPLEMENTED, "InvokeOp is not implemented for minimal build."); + API_IMPL_END +} + +ORT_API(void, OrtApis::ReleaseOp, _Frees_ptr_opt_ OrtOp*) { +} + +#else + +namespace onnxruntime { +namespace standalone { + +// For invoking kernels without a graph +class StandAloneKernelContext : public OpKernelContext { + public: + StandAloneKernelContext(const OrtValue* const* input_values, + int input_count, + OrtValue* const* output_values, + int output_count, + AllocatorPtr allocator, + onnxruntime::concurrency::ThreadPool* threadpool, + const logging::Logger& logger) : OpKernelContext(threadpool, logger), + input_values_(input_values), + input_count_(input_count), + output_values_(output_values), + output_count_(output_count), + allocator_(allocator) {} + + int NumVariadicInputs(size_t arg_num) const override { + ORT_ENFORCE(arg_num < static_cast(input_count_), "invalid arg_num."); + auto ort_value = input_values_[arg_num]; + if (ort_value->IsTensor()) { + return static_cast(ort_value->Get().Shape().Size()); + } else if (ort_value->IsTensorSequence()) { + return static_cast(ort_value->Get().Size()); + } else if (ort_value->IsSparseTensor()) { +#ifdef DISABLE_SPARSE_TENSORS + ORT_THROW("sparse tensor is not supported in this build."); +#else + return static_cast(ort_value->Get().Values().Shape().Size()); +#endif + } else { + return 0; + } + } + + MLDataType InputType(int index) const override { + if (index >= input_count_) { + return nullptr; + } else { + return input_values_[index]->Type(); + } + } + + MLDataType OutputType(int index) const override { + if (index >= output_count_) { + return nullptr; + } else { + return output_values_[index]->Type(); + } + } + + bool TryGetInferredInputShape(int, TensorShape&) const override { + return false; + } + + bool TryGetInferredOutputShape(int, TensorShape&) const override { + return false; + } + + int InputCount() const override { + return input_count_; + } + + int ImplicitInputCount() const override { + return 0; + } + + int OutputCount() const override { + return static_cast(output_count_); + } + + Status GetTempSpaceAllocator(AllocatorPtr* output) const override ORT_MUST_USE_RESULT { + *output = allocator_; + return Status::OK(); + } + + Fence_t InputFence(int index) const override { + if (index >= input_count_) { + return nullptr; + } else { + return input_values_[index]->Fence(); + } + } + + Fence_t ImplicitInputFence(int) const override { + return nullptr; + } + + Fence_t OutputFence(int index) const override { + if (index >= output_count_) { + return nullptr; + } else { + return output_values_[index]->Fence(); + } + } + + int GetDeviceId() const override { + return 0; + } + + void* GetComputeStream() const override { + return nullptr; + } + + protected: + const OrtValue* GetInputMLValue(int index) const override { + if (index >= input_count_) { + return nullptr; + } else { + return input_values_[index]; + } + } + + OrtValue* OutputMLValue(int index, const TensorShape& shape) override { + if (index >= output_count_) { + return nullptr; + } + OrtValue& ort_value = *output_values_[index]; + if (!ort_value.IsAllocated()) { + if (ort_value.IsTensor()) { + Tensor::InitOrtValue(ort_value.Type(), shape, allocator_, ort_value); + } else if (ort_value.IsTensorSequence()) { + auto ml_type = ort_value.Type(); + auto element_type = ml_type->AsSequenceTensorType()->GetElementType(); + auto p_sequence = std::make_unique(element_type); + auto ml_tensor_sequence = DataTypeImpl::GetType(); + ort_value.Init(p_sequence.release(), ml_tensor_sequence, ml_tensor_sequence->GetDeleteFunc()); + } else if (ort_value.IsSparseTensor()) { +#ifdef DISABLE_SPARSE_TENSORS + ORT_THROW("sparse tensor is not supported in this build."); +#else + auto ml_type = ort_value.Type(); + auto element_type = ml_type->AsSparseTensorType()->GetElementType(); + SparseTensor::InitOrtValue(element_type, shape, allocator_, ort_value); +#endif + } + } + return &ort_value; + } + + OrtValue* GetOrCreateOutputMLValue(int index) override { + if (index >= output_count_) { + return nullptr; + } else { + return output_values_[index]; + } + } + + const OrtValue* const* input_values_; + const int input_count_; + OrtValue* const* output_values_; + const int output_count_; + AllocatorPtr allocator_; +}; // StandAloneKernelContext + +onnxruntime::Status CreateOpAttr(const char* name, const void* data, int len, OrtOpAttrType type, OrtOpAttr** op_attr) { + auto attr = std::make_unique(); + onnxruntime::Status status = onnxruntime::Status::OK(); + attr->set_name(std::string{name}); + const int* ints = reinterpret_cast(data); + const float* floats = reinterpret_cast(data); + auto str = reinterpret_cast(data); + auto strs = reinterpret_cast(data); + switch (type) { + case OrtOpAttrType::ORT_OP_ATTR_INT: + attr->set_i(ints[0]); + attr->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INT); + break; + case OrtOpAttrType::ORT_OP_ATTR_INTS: + for (int j = 0; j < len; ++j) { + attr->add_ints(ints[j]); + } + attr->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INTS); + break; + case OrtOpAttrType::ORT_OP_ATTR_FLOAT: + attr->set_f(floats[0]); + attr->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT); + break; + case OrtOpAttrType::ORT_OP_ATTR_FLOATS: + for (int j = 0; j < len; ++j) { + attr->add_floats(floats[j]); + } + attr->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOATS); + break; + case OrtOpAttrType::ORT_OP_ATTR_STRING: + attr->set_s(std::string{str}); + attr->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING); + break; + case OrtOpAttrType::ORT_OP_ATTR_STRINGS: + for (int j = 0; j < len; ++j) { + attr->add_strings(std::string{strs[j]}); + } + attr->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRINGS); + break; + default: + status = Status(common::ONNXRUNTIME, common::FAIL, "Attribute type not supported yet."); + break; + } + *op_attr = reinterpret_cast(attr.release()); + return status; +} + +onnxruntime::Status CreateOp(const OrtKernelInfo* info, + const char* op_name, + const char* domain, + int version, + const char** type_constraint_names, + const ONNXTensorElementDataType* type_constraint_values, + int type_constraint_count, + const OrtOpAttr* const* attr_values, + int attr_count, + OrtOp** op) { + *op = nullptr; + auto kernel_info = reinterpret_cast(info); + auto ep = reinterpret_cast(kernel_info->GetExecutionProvider()); + auto kernel_registry = ep->GetKernelRegistry(); + const KernelCreateInfo* kernel_create_info{}; + std::unordered_map type_constraint_map; + for (int i = 0; i < type_constraint_count; ++i) { + ONNX_NAMESPACE::TypeProto proto; + proto.mutable_tensor_type()->set_elem_type(type_constraint_values[i]); + type_constraint_map[type_constraint_names[i]] = DataTypeImpl::TypeFromProto(proto); + } + auto status = kernel_registry->TryFindKernel(op_name, + domain, + version, + type_constraint_map, + ep->Type(), + &kernel_create_info); + ORT_RETURN_IF_ERROR(status); + onnxruntime::Node node; + for (int i = 0; i < attr_count; ++i) { + auto attr_proto = reinterpret_cast(attr_values[i]); + node.AddAttributeProto(*attr_proto); + } + auto kernel_def_builder = KernelDefBuilder::Create(); + kernel_def_builder->SetName(op_name); + kernel_def_builder->SetDomain(domain); + kernel_def_builder->SinceVersion(version); + OpKernelInfo instant_kernel_info(node, *kernel_def_builder->Build(), *ep, {}, {}, {}); + std::unique_ptr op_kernel; + FuncManager func_mgr; + status = kernel_create_info->kernel_create_func(func_mgr, instant_kernel_info, op_kernel); + ORT_RETURN_IF_ERROR(status); + *op = reinterpret_cast(op_kernel.release()); + return status; +} + +onnxruntime::Status InvokeOp(_In_ const OrtKernelContext* context, + _In_ const OrtOp* ort_op, + _In_ const OrtValue* const* input_values, + _In_ int input_count, + _Inout_ OrtValue* const* output_values, + _In_ int output_count) { + auto ctx = reinterpret_cast(context); + AllocatorPtr allocator{}; + ORT_RETURN_IF_ERROR(ctx->GetTempSpaceAllocator(&allocator)); + StandAloneKernelContext standalone_kernel_ctx(input_values, + input_count, + output_values, + output_count, + allocator, + ctx->GetOperatorThreadPool(), + ctx->Logger()); + auto kernel = reinterpret_cast(ort_op); + return kernel->Compute(&standalone_kernel_ctx); +} + +} // namespace standalone +} // namespace onnxruntime + +ORT_API_STATUS_IMPL(OrtApis::CreateOpAttr, + _In_ const char* name, + _In_ const void* data, + _In_ int len, + _In_ OrtOpAttrType type, + _Outptr_ OrtOpAttr** op_attr) { + API_IMPL_BEGIN + auto status = onnxruntime::standalone::CreateOpAttr(name, data, len, type, op_attr); + if (status.IsOK()) { + return nullptr; + } else { + return CreateStatus(static_cast(status.Code()), status.ErrorMessage().c_str()); + } + API_IMPL_END +} + +ORT_API(void, OrtApis::ReleaseOpAttr, _Frees_ptr_opt_ OrtOpAttr* op_attr) { + if (op_attr) { + delete reinterpret_cast(op_attr); + } +} + +ORT_API_STATUS_IMPL(OrtApis::CreateOp, + _In_ const OrtKernelInfo* info, + _In_ const char* op_name, + _In_ const char* domain, + _In_ int version, + _In_opt_ const char** type_constraint_names, + _In_opt_ const ONNXTensorElementDataType* type_constraint_values, + _In_opt_ int type_constraint_count, + _In_opt_ const OrtOpAttr* const* attr_values, + _In_opt_ int attr_count, + _Outptr_ OrtOp** ort_op) { + API_IMPL_BEGIN + auto status = onnxruntime::standalone::CreateOp(info, + op_name, + domain, + version, + type_constraint_names, + type_constraint_values, + type_constraint_count, + attr_values, + attr_count, + ort_op); + if (status.IsOK()) { + return nullptr; + } else { + return CreateStatus(static_cast(status.Code()), status.ErrorMessage().c_str()); + } + API_IMPL_END +} + +ORT_API_STATUS_IMPL(OrtApis::InvokeOp, + _In_ const OrtKernelContext* context, + _In_ const OrtOp* ort_op, + _In_ const OrtValue* const* input_values, + _In_ int input_count, + _Inout_ OrtValue* const* output_values, + _In_ int output_count) { + API_IMPL_BEGIN + auto status = onnxruntime::standalone::InvokeOp(context, ort_op, input_values, input_count, output_values, output_count); + if (status.IsOK()) { + return nullptr; + } else { + return CreateStatus(static_cast(status.Code()), status.ErrorMessage().c_str()); + } + API_IMPL_END +} + +ORT_API(void, OrtApis::ReleaseOp, _Frees_ptr_opt_ OrtOp* op) { + if (op) { + delete reinterpret_cast(op); + } +} + +#endif diff --git a/onnxruntime/test/shared_lib/custom_op_utils.cc b/onnxruntime/test/shared_lib/custom_op_utils.cc index 1bda3af84d..39abe25581 100644 --- a/onnxruntime/test/shared_lib/custom_op_utils.cc +++ b/onnxruntime/test/shared_lib/custom_op_utils.cc @@ -206,3 +206,220 @@ void SliceCustomOpKernel::Compute(OrtKernelContext* context) { ORT_THROW("Unsupported input type"); } } + +InstantCustomKernel::InstantCustomKernel(Ort::CustomOpApi ort, const OrtKernelInfo* info, void*) : ort_(ort) { + const char* add_type_constrait_names[1] = {"T"}; + int add_type_constrait_values[1] = {1}; + ort.CreateOp(info, "Add", "", 14, + (const char**)add_type_constrait_names, + (const ONNXTensorElementDataType*)add_type_constrait_values, + 1, nullptr, 0, &op_add); + ORT_ENFORCE(op_add, "op_add not initialzied"); + InitTopK(ort, info); + ORT_ENFORCE(op_topk, "op_add not initialzied"); + InitGru(ort, info); + ORT_ENFORCE(op_gru, "op_add not initialzied"); +} + +void InstantCustomKernel::InitTopK(Ort::CustomOpApi ort, const OrtKernelInfo* info) { + const char* type_constrait_names[2] = {"T", "I"}; + int type_constrait_values[2] = {1, 7}; + + int axis_value = -1; + OrtOpAttr* axis{}; + ort.CreateOpAttr("axis", &axis_value, 1, OrtOpAttrType::ORT_OP_ATTR_INT, &axis); + + int largest_value = 0; // return in ascending order + OrtOpAttr* largest{}; + ort.CreateOpAttr("largest", &largest_value, 1, OrtOpAttrType::ORT_OP_ATTR_INT, &largest); + + int sorted_value = 1; + OrtOpAttr* sorted{}; + ort.CreateOpAttr("sorted", &sorted_value, 1, OrtOpAttrType::ORT_OP_ATTR_INT, &sorted); + + if (!axis || !largest || !sorted) { + ORT_THROW("Failed to create attributes for topk."); + } + + OrtOpAttr* top_attrs[3] = {axis, largest, sorted}; + ort.CreateOp(info, "TopK", "", 14, + (const char**)type_constrait_names, + (const ONNXTensorElementDataType*)type_constrait_values, + 2, top_attrs, 3, &op_topk); + + ort.ReleaseOpAttr(axis); + ort.ReleaseOpAttr(largest); + ort.ReleaseOpAttr(sorted); +} + +void InstantCustomKernel::InvokeTopK(OrtKernelContext* context) { + auto mem_info = Ort::MemoryInfo::CreateCpu(OrtAllocatorType::OrtArenaAllocator, OrtMemType::OrtMemTypeCPU); + + float raw_x[10] = {6., 3., 4., 8., 7., 1., 9., 0., 5., 2.}; + int64_t raw_x_shape[1] = {10}; + auto topk_x = Ort::Value::CreateTensor(mem_info, raw_x, 10, raw_x_shape, 1); + + int64_t raw_k[1] = {2}; + int64_t raw_k_shape[1] = {1}; + auto topk_k = Ort::Value::CreateTensor(mem_info, raw_k, 1, raw_k_shape, 1); + + float raw_values[2] = {}; + int64_t raw_values_shape[1] = {2}; + auto topk_values = Ort::Value::CreateTensor(mem_info, raw_values, 2, raw_values_shape, 1); + + int64_t raw_indices[2] = {}; + int64_t raw_indices_shape[1] = {2}; + auto topk_indices = Ort::Value::CreateTensor(mem_info, raw_indices, 2, raw_indices_shape, 1); + + const OrtValue* topk_inputs[2] = {(OrtValue*)topk_x, (OrtValue*)topk_k}; + OrtValue* topk_outputs[2] = {(OrtValue*)topk_values, (OrtValue*)topk_indices}; + ort_.InvokeOp(context, op_topk, topk_inputs, 2, topk_outputs, 2); + + if (std::abs(raw_values[0] - 0.) > 1e-6 || std::abs(raw_values[1] - 1.) > 1e-6) { + ORT_THROW("topk instant operator returns wrong values"); + } + if (raw_indices[0] != 7 || raw_indices[1] != 5) { + ORT_THROW("topk instant operator returns wrong indices"); + } +} + +void InstantCustomKernel::InitGru(Ort::CustomOpApi ort, const OrtKernelInfo* info) { + const char* type_constrait_names[2] = {"T", "T1"}; + int type_constrait_values[2] = {1, 6}; + + const char* activition_names[4] = {"LeakyRelu", "Tanh", "Sigmoid", "ScaledTanh"}; + OrtOpAttr* activations{}; + ort.CreateOpAttr("activations", activition_names, 4, OrtOpAttrType::ORT_OP_ATTR_STRINGS, &activations); + + float alphas[2] = {0.5f, 2.f}; + OrtOpAttr* activation_alpha{}; + ort.CreateOpAttr("activation_alpha ", alphas, 2, OrtOpAttrType::ORT_OP_ATTR_FLOATS, &activation_alpha); + + float betas[1] = {2.f}; + OrtOpAttr* activation_beta{}; + ort.CreateOpAttr("activation_beta ", betas, 1, OrtOpAttrType::ORT_OP_ATTR_FLOATS, &activation_beta); + + const char* direction_string = "bidirectional"; + OrtOpAttr* direction{}; + ort.CreateOpAttr("direction", direction_string, 1, OrtOpAttrType::ORT_OP_ATTR_STRING, &direction); + + int linear_before_reset_value = 0; + OrtOpAttr* linear_before_reset{}; + ort.CreateOpAttr("linear_before_reset", &linear_before_reset_value, 1, OrtOpAttrType::ORT_OP_ATTR_INT, &linear_before_reset); + + int hidden_size_value = 2; + OrtOpAttr* hidden_size{}; + ort.CreateOpAttr("hidden_size", &hidden_size_value, 1, OrtOpAttrType::ORT_OP_ATTR_INT, &hidden_size); + + if (!activations || !activation_alpha || !activation_beta || !direction || !linear_before_reset || !hidden_size) { + ORT_THROW("failed to create attributes for gru."); + } + + OrtOpAttr* gru_attrs[6] = {activations, activation_alpha, activation_beta, direction, linear_before_reset, hidden_size}; + ort.CreateOp(info, "GRU", "", 14, + (const char**)type_constrait_names, + (const ONNXTensorElementDataType*)type_constrait_values, + 2, gru_attrs, 6, &op_gru); + + ort.ReleaseOpAttr(activations); + ort.ReleaseOpAttr(activation_alpha); + ort.ReleaseOpAttr(activation_beta); + ort.ReleaseOpAttr(direction); + ort.ReleaseOpAttr(linear_before_reset); + ort.ReleaseOpAttr(hidden_size); +} + +void InstantCustomKernel::InvokeGru(OrtKernelContext* context) { + auto mem_info = Ort::MemoryInfo::CreateCpu(OrtAllocatorType::OrtArenaAllocator, OrtMemType::OrtMemTypeCPU); + + float raw_x[2] = {1.0f, 2.0f}; + int64_t raw_x_shape[3] = {1, 1, 2}; + auto X = Ort::Value::CreateTensor(mem_info, raw_x, 2, raw_x_shape, 3); + + float raw_w[24] = { + -0.494659f, 0.0453352f, -0.487793f, 0.417264f, // Wz + -0.0091708f, -0.255364f, -0.106952f, -0.266717f, // Wr + -0.0888852f, -0.428709f, -0.283349f, 0.208792f, // Wh + -0.494659f, 0.0453352f, -0.487793f, 0.417264f, // WBz + -0.0091708f, -0.255364f, -0.106952f, -0.266717f, // WBr + -0.0888852f, -0.428709f, -0.283349f, 0.208792f // WBh + }; + int64_t raw_w_shape[3] = {2, 6, 2}; + auto W = Ort::Value::CreateTensor(mem_info, raw_w, 24, raw_w_shape, 3); + + float raw_r[24] = { + 0.146626f, -0.0620289f, -0.0815302f, 0.100482f, // Rz + -0.228172f, 0.405972f, 0.31576f, 0.281487f, // Rr + -0.394864f, 0.42111f, -0.386624f, -0.390225f, // Rh + 0.146626f, -0.0620289f, -0.0815302f, 0.100482f, // RBz + -0.228172f, 0.405972f, 0.31576f, 0.281487f, // RBr + -0.394864f, 0.42111f, -0.386624f, -0.390225f}; // RBh + int64_t raw_r_shape[3] = {2, 6, 2}; + auto R = Ort::Value::CreateTensor(mem_info, raw_r, 24, raw_r_shape, 3); + + float raw_b[24] = { + 0.381619f, 0.0323954f, // Wbz + -0.258721f, 0.45056f, // Wbr + -0.250755f, 0.0967895f, // Wbh + 0.0f, 0.0f, // Rbz + -0.0f, 0.0f, // Rbr + -0.0f, 0.0f, // Rbh + 0.381619f, 0.0323954f, // WBbz + -0.258721f, 0.45056f, // WBbr + -0.250755f, 0.0967895f, // WBbh + 0.0f, 0.0f, // RBbz + -0.0f, 0.0f, // RBbr + -0.0f, 0.0f}; // RBbh + int64_t raw_b_shape[2] = {2, 12}; + auto B = Ort::Value::CreateTensor(mem_info, raw_b, 24, raw_b_shape, 2); + + int32_t raw_seq_lens = 1; + int64_t seq_lens_shape[1] = {1}; + auto sequence_lens = Ort::Value::CreateTensor(mem_info, &raw_seq_lens, 1, seq_lens_shape, 1); + + std::vector raw_initial_h(4, 0.25f); + int64_t initial_h_shape[3] = {2, 1, 2}; + auto initial_h = Ort::Value::CreateTensor(mem_info, raw_initial_h.data(), 4, initial_h_shape, 3); + + float raw_y[4] = {}; + int64_t raw_y_shape[64] = {1, 2, 1, 2}; + auto Y = Ort::Value::CreateTensor(mem_info, raw_y, 4, raw_y_shape, 4); + + float raw_yh[4] = {}; + int64_t raw_yh_shape[64] = {2, 1, 2}; + auto YH = Ort::Value::CreateTensor(mem_info, raw_yh, 4, raw_yh_shape, 3); + + const OrtValue* inputs[6] = {(OrtValue*)X, (OrtValue*)W, (OrtValue*)R, (OrtValue*)B, (OrtValue*)sequence_lens, (OrtValue*)initial_h}; + OrtValue* outputs[2] = {(OrtValue*)Y, (OrtValue*)YH}; + + const float expected_y[4] = {-0.832559f, + 0.236267f, + 0.124924f, + 0.148701f}; + + ort_.InvokeOp(context, op_gru, inputs, 6, outputs, 2); + + for (int i = 0; i < 4; ++i) { + if (std::abs(raw_y[i] - expected_y[i]) > 1e-6) { + ORT_THROW("GRU op give unexpected output."); + } + } +} + +void InstantCustomKernel::Compute(OrtKernelContext* context) { + const OrtValue* input_X = ort_.KernelContext_GetInput(context, 0); + const OrtValue* input_Y = ort_.KernelContext_GetInput(context, 1); + OrtTensorDimensions dimensions(ort_, input_X); + OrtValue* output = ort_.KernelContext_GetOutput(context, 0, dimensions.data(), dimensions.size()); + const OrtValue* inputs[2] = {input_X, input_Y}; + OrtValue* outputs[1] = {output}; + ort_.InvokeOp(context, op_add, inputs, 2, outputs, 1); + InvokeTopK(context); + InvokeGru(context); +} + +InstantCustomKernel::~InstantCustomKernel() { + ort_.ReleaseOp(op_add); + ort_.ReleaseOp(op_topk); + ort_.ReleaseOp(op_gru); +} diff --git a/onnxruntime/test/shared_lib/custom_op_utils.h b/onnxruntime/test/shared_lib/custom_op_utils.h index d3e879aa80..7bf3ccb4cd 100644 --- a/onnxruntime/test/shared_lib/custom_op_utils.h +++ b/onnxruntime/test/shared_lib/custom_op_utils.h @@ -207,3 +207,42 @@ struct SliceCustomOp : Ort::CustomOpBase { private: const char* provider_; }; + +struct InstantCustomKernel { + InstantCustomKernel(Ort::CustomOpApi ort, const OrtKernelInfo* info, void*); + + ~InstantCustomKernel(); + void Compute(OrtKernelContext* context); + + private: + void InitTopK(Ort::CustomOpApi ort, const OrtKernelInfo* info); + void InvokeTopK(OrtKernelContext* context); + + void InitGru(Ort::CustomOpApi ort, const OrtKernelInfo* info); + void InvokeGru(OrtKernelContext* context); + + Ort::CustomOpApi ort_; + OrtOp* op_add{}; + OrtOp* op_topk{}; + OrtOp* op_gru{}; +}; + +struct InstantCustomOp : Ort::CustomOpBase { + explicit InstantCustomOp(const char* provider, void* compute_stream) : provider_(provider), compute_stream_(compute_stream) {} + + void* CreateKernel(Ort::CustomOpApi api, const OrtKernelInfo* info) const { return new InstantCustomKernel(api, info, compute_stream_); }; + const char* GetName() const { return "Foo"; }; + const char* GetExecutionProviderType() const { return provider_; }; + + size_t GetInputTypeCount() const { return 2; }; + ONNXTensorElementDataType GetInputType(size_t /*index*/) const { + return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT; + }; + + size_t GetOutputTypeCount() const { return 1; }; + ONNXTensorElementDataType GetOutputType(size_t /*index*/) const { return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT; }; + + private: + const char* provider_; + void* compute_stream_; +}; diff --git a/onnxruntime/test/shared_lib/test_inference.cc b/onnxruntime/test/shared_lib/test_inference.cc index 7c413cb4b8..d20cd97f9a 100644 --- a/onnxruntime/test/shared_lib/test_inference.cc +++ b/onnxruntime/test/shared_lib/test_inference.cc @@ -405,6 +405,26 @@ TEST(CApiTest, custom_op_handler) { #endif } +#if !defined(ORT_MINIMAL_BUILD) +TEST(CApiTest, instant_op_handler) { + std::vector inputs(1); + Input& input = inputs[0]; + input.name = "X"; + input.dims = {3, 2}; + input.values = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f}; + + std::vector expected_dims_y = {3, 2}; + std::vector expected_values_y = {2.0f, 4.0f, 6.0f, 8.0f, 10.0f, 12.0f}; + + InstantCustomOp instant_op{onnxruntime::kCpuExecutionProvider, nullptr}; + Ort::CustomOpDomain custom_op_domain(""); + custom_op_domain.Add(&instant_op); + + TestInference(*ort_env, CUSTOM_OP_MODEL_URI, inputs, "Y", expected_dims_y, expected_values_y, 0, + custom_op_domain, nullptr); +} +#endif + #ifdef ENABLE_EXTENSION_CUSTOM_OPS // test enabled ort-customops negpos TEST(CApiTest, test_enable_ort_customops_negpos) {