From d5ec353e58114e18fa654c59262b5a3e63abd4c1 Mon Sep 17 00:00:00 2001 From: Ryan Hill <38674843+RyanUnderhill@users.noreply.github.com> Date: Wed, 6 May 2020 00:57:09 -0700 Subject: [PATCH] Ryanunderhill/mkldnn dll (#3314) First version of allowing providers to work as DLLs, only implemented for DNNL so far. More improvements to come next! --- cmake/onnxruntime_providers.cmake | 18 +- .../onnxruntime/core/common/logging/logging.h | 1 + .../core/framework/execution_provider.h | 5 +- include/onnxruntime/core/framework/func_api.h | 3 +- .../core/framework/provider_bridge_ort.cc | 565 ++++++++++++++++++ onnxruntime/core/providers/dnnl/dnnl_common.h | 9 +- .../providers/dnnl/dnnl_execution_provider.cc | 109 ++-- .../providers/dnnl/dnnl_execution_provider.h | 46 +- onnxruntime/core/providers/dnnl/dnnl_fwd.h | 4 +- .../providers/dnnl/dnnl_provider_factory.cc | 36 +- .../core/providers/dnnl/exported_symbols.lst | 1 + onnxruntime/core/providers/dnnl/math/gemm.cc | 98 ++- onnxruntime/core/providers/dnnl/math/gemm.h | 7 +- onnxruntime/core/providers/dnnl/memcpy_s.h | 2 +- .../dnnl/subgraph/dnnl_activations.h | 9 +- .../providers/dnnl/subgraph/dnnl_batchnorm.h | 14 +- .../core/providers/dnnl/subgraph/dnnl_conv.h | 31 +- .../dnnl/subgraph/dnnl_conv_batchnorm.h | 35 +- .../dnnl/subgraph/dnnl_func_kernel.cc | 11 +- .../dnnl/subgraph/dnnl_func_kernel.h | 51 +- .../providers/dnnl/subgraph/dnnl_kernel.cc | 11 +- .../providers/dnnl/subgraph/dnnl_kernel.h | 18 +- .../core/providers/dnnl/subgraph/dnnl_lrn.h | 29 +- .../core/providers/dnnl/subgraph/dnnl_pool.h | 30 +- .../core/providers/dnnl/subgraph/dnnl_sum.h | 8 +- .../core/providers/dnnl/subgraph/subgraph.h | 9 +- onnxruntime/core/providers/dnnl/symbols.def | 2 + .../core/providers/dnnl/version_script.lds | 9 + .../providers/shared_library/provider_api.h | 260 ++++++++ .../provider_bridge_provider.cc | 273 +++++++++ .../shared_library/provider_interfaces.h | 445 ++++++++++++++ .../python/onnxruntime_pybind_state.cc | 164 +++-- onnxruntime/test/onnx/main.cc | 1 + .../dnnl/dnnl_execution_provider_test.cc | 4 +- 34 files changed, 1970 insertions(+), 348 deletions(-) create mode 100644 onnxruntime/core/framework/provider_bridge_ort.cc create mode 100644 onnxruntime/core/providers/dnnl/exported_symbols.lst create mode 100644 onnxruntime/core/providers/dnnl/symbols.def create mode 100644 onnxruntime/core/providers/dnnl/version_script.lds create mode 100644 onnxruntime/core/providers/shared_library/provider_api.h create mode 100644 onnxruntime/core/providers/shared_library/provider_bridge_provider.cc create mode 100644 onnxruntime/core/providers/shared_library/provider_interfaces.h diff --git a/cmake/onnxruntime_providers.cmake b/cmake/onnxruntime_providers.cmake index 5d8d42a16b..96cc54a059 100644 --- a/cmake/onnxruntime_providers.cmake +++ b/cmake/onnxruntime_providers.cmake @@ -265,16 +265,30 @@ if (onnxruntime_USE_DNNL) file(GLOB_RECURSE onnxruntime_providers_dnnl_cc_srcs CONFIGURE_DEPENDS "${ONNXRUNTIME_ROOT}/core/providers/dnnl/*.h" "${ONNXRUNTIME_ROOT}/core/providers/dnnl/*.cc" + "${ONNXRUNTIME_ROOT}/core/providers/shared_library/*.h" + "${ONNXRUNTIME_ROOT}/core/providers/shared_library/*.cc" ) source_group(TREE ${ONNXRUNTIME_ROOT}/core FILES ${onnxruntime_providers_dnnl_cc_srcs}) - add_library(onnxruntime_providers_dnnl ${onnxruntime_providers_dnnl_cc_srcs}) - onnxruntime_add_include_to_target(onnxruntime_providers_dnnl onnxruntime_common onnxruntime_framework onnx onnx_proto protobuf::libprotobuf) + add_library(onnxruntime_providers_dnnl SHARED ${onnxruntime_providers_dnnl_cc_srcs}) + onnxruntime_add_include_to_target(onnxruntime_providers_dnnl onnxruntime_common onnx) # onnx needed for stl_backports.h add_dependencies(onnxruntime_providers_dnnl ${onnxruntime_EXTERNAL_DEPENDENCIES}) set_target_properties(onnxruntime_providers_dnnl PROPERTIES FOLDER "ONNXRuntime") target_include_directories(onnxruntime_providers_dnnl PRIVATE ${ONNXRUNTIME_ROOT} ${eigen_INCLUDE_DIRS} ${DNNL_INCLUDE_DIR}) + target_link_libraries(onnxruntime_providers_dnnl PRIVATE dnnl) install(DIRECTORY ${PROJECT_SOURCE_DIR}/../include/onnxruntime/core/providers/dnnl DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}/onnxruntime/core/providers) set_target_properties(onnxruntime_providers_dnnl PROPERTIES LINKER_LANGUAGE CXX) + + if(APPLE) + set_property(TARGET onnxruntime_providers_dnnl APPEND_STRING PROPERTY LINK_FLAGS "-Xlinker -exported_symbols_list ${ONNXRUNTIME_ROOT}/core/providers/dnnl/exported_symbols.lst") + target_link_libraries(onnxruntime_providers_dnnl PRIVATE nsync_cpp) + elseif(UNIX) + set_property(TARGET onnxruntime_providers_dnnl APPEND_STRING PROPERTY LINK_FLAGS "-Xlinker --version-script=${ONNXRUNTIME_ROOT}/core/providers/dnnl/version_script.lds -Xlinker --gc-sections") + target_link_libraries(onnxruntime_providers_dnnl PRIVATE nsync_cpp) + else() + set_property(TARGET onnxruntime_providers_dnnl APPEND_STRING PROPERTY LINK_FLAGS "-DEF:${ONNXRUNTIME_ROOT}/core/providers/dnnl/symbols.def") + endif() + endif() if (onnxruntime_USE_TENSORRT) diff --git a/include/onnxruntime/core/common/logging/logging.h b/include/onnxruntime/core/common/logging/logging.h index 10869779d4..0c8919f8dc 100644 --- a/include/onnxruntime/core/common/logging/logging.h +++ b/include/onnxruntime/core/common/logging/logging.h @@ -219,6 +219,7 @@ class LoggingManager final { const bool default_filter_user_data_; const int default_max_vlog_level_; bool owns_default_logger_; + static Logger* s_default_logger_; struct Epochs { diff --git a/include/onnxruntime/core/framework/execution_provider.h b/include/onnxruntime/core/framework/execution_provider.h index 877788a9bd..e66eb83a3d 100644 --- a/include/onnxruntime/core/framework/execution_provider.h +++ b/include/onnxruntime/core/framework/execution_provider.h @@ -15,9 +15,6 @@ namespace onnxruntime { class GraphViewer; class Node; -} // namespace onnxruntime -namespace onnxruntime { - struct ComputeCapability; class KernelRegistry; class KernelRegistryManager; @@ -27,7 +24,7 @@ class KernelRegistryManager; */ typedef std::map AllocatorMap; -// if we are export the fused function to dll, the function will still in the same binary as lotus +// if we are export the fused function to dll, the function will still in the same binary as onnxruntime // use std function to give execution provider some chance to capture some state. using CreateFunctionStateFunc = std::function; using ComputeFunc = std::function; diff --git a/include/onnxruntime/core/framework/func_api.h b/include/onnxruntime/core/framework/func_api.h index 99f48b6326..f12a94da02 100644 --- a/include/onnxruntime/core/framework/func_api.h +++ b/include/onnxruntime/core/framework/func_api.h @@ -1,5 +1,6 @@ #pragma once -#include "core/common/common.h" +#include "core/common/status.h" +using onnxruntime::common::Status; namespace onnxruntime { diff --git a/onnxruntime/core/framework/provider_bridge_ort.cc b/onnxruntime/core/framework/provider_bridge_ort.cc new file mode 100644 index 0000000000..a69831aea7 --- /dev/null +++ b/onnxruntime/core/framework/provider_bridge_ort.cc @@ -0,0 +1,565 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// This is the Onnxruntime side of the bridge to allow providers to be built as a DLL +// It implements onnxruntime::ProviderHost + +#include "core/framework/data_types.h" +#include "core/framework/allocatormgr.h" +#include "core/providers/dnnl/dnnl_provider_factory.h" +#include "core/session/abi_session_options_impl.h" +#include "core/session/ort_apis.h" +#include "core/platform/env.h" +#include "core/framework/execution_provider.h" +#include "core/framework/compute_capability.h" +#define PROVIDER_BRIDGE_ORT +#include "core/providers/shared_library/provider_interfaces.h" +#include "onnx/common/stl_backports.h" +#include "core/common/logging/logging.h" +#include "core/common/cpuid_info.h" + +namespace onnxruntime { + +struct Provider_OrtDevice_Impl : Provider_OrtDevice { + OrtDevice v_; +}; + +struct Provider_OrtMemoryInfo_Impl : Provider_OrtMemoryInfo { + Provider_OrtMemoryInfo_Impl(const char* name_, OrtAllocatorType type_, OrtDevice device_, int id_, OrtMemType mem_type_) : info_{onnxruntime::make_unique(name_, type_, device_, id_, mem_type_)} {} + + std::unique_ptr info_; +}; + +struct Provider_IAllocator_Impl : Provider_IAllocator { + Provider_IAllocator_Impl(AllocatorPtr p) : p_{p} {} + + void* Alloc(size_t size) override { return p_->Alloc(size); } + void Free(void* p) override { return p_->Free(p); } + + AllocatorPtr p_; +}; + +struct Provider_IDeviceAllocator_Impl : Provider_IDeviceAllocator { + Provider_IDeviceAllocator_Impl(std::unique_ptr p) : p_{std::move(p)} {} + + void* Alloc(size_t size) override { return p_->Alloc(size); } + void Free(void* p) override { return p_->Free(p); } + + bool AllowsArena() const override { return p_->AllowsArena(); } + + std::unique_ptr p_; +}; + +struct Provider_TensorProto_Impl : ONNX_NAMESPACE::Provider_TensorProto { + Provider_TensorProto_Impl(ONNX_NAMESPACE::TensorProto* p) : p_{p} {} + + void CopyFrom(const Provider_TensorProto& v) override { + *p_ = *static_cast(&v)->p_; + } + + ONNX_NAMESPACE::TensorProto* p_{}; +}; + +struct Provider_AttributeProto_Impl : ONNX_NAMESPACE::Provider_AttributeProto { + Provider_AttributeProto_Impl() = default; + Provider_AttributeProto_Impl(const ONNX_NAMESPACE::AttributeProto& copy) : v_{copy} {} + + std::unique_ptr Clone() const override { + return onnxruntime::make_unique(v_); + } + + ::onnx::AttributeProto_AttributeType type() const override { return v_.type(); } + + int ints_size() const override { + return v_.ints_size(); + } + + int64_t ints(int i) const override { return v_.ints(i); } + int64_t i() const override { return v_.i(); } + float f() const override { return v_.f(); } + void set_s(const ::std::string& value) override { v_.set_s(value); } + const ::std::string& s() const override { return v_.s(); } + void set_name(const ::std::string& value) override { v_.set_name(value); } + void set_type(::onnx::AttributeProto_AttributeType value) override { v_.set_type(value); } + ::onnx::Provider_TensorProto* add_tensors() override { + // Kind of a hack, but the pointer is only valid until the next add_tensors call + tensors_ = onnxruntime::make_unique(v_.add_tensors()); + return tensors_.get(); + } + + ONNX_NAMESPACE::AttributeProto v_; + std::unique_ptr tensors_; +}; + +struct Provider_KernelDef_Impl : Provider_KernelDef { + Provider_KernelDef_Impl(std::unique_ptr p) : p_(std::move(p)) {} + std::unique_ptr p_; +}; + +struct Provider_KernelDefBuilder_Impl : Provider_KernelDefBuilder { + Provider_KernelDefBuilder& SetName(const char* op_name) override { + v_.SetName(op_name); + return *this; + } + Provider_KernelDefBuilder& SetDomain(const char* domain) override { + v_.SetDomain(domain); + return *this; + } + + Provider_KernelDefBuilder& SinceVersion(int since_version) override { + v_.SinceVersion(since_version); + return *this; + } + Provider_KernelDefBuilder& Provider(const char* provider_type) override { + v_.Provider(provider_type); + return *this; + } + Provider_KernelDefBuilder& TypeConstraint(const char* arg_name, MLDataType supported_type) override { + v_.TypeConstraint(arg_name, supported_type); + return *this; + } + + std::unique_ptr Build() override { + return onnxruntime::make_unique(v_.Build()); + } + + KernelDefBuilder v_; +}; + +struct Provider_NodeArg_Impl : Provider_NodeArg { + Provider_NodeArg_Impl(const NodeArg* p) : p_{p} { + if (p_->Shape()) + tensor_shape_proto_.dim_size_ = p_->Shape()->dim_size(); + } + + const std::string& Name() const noexcept override { return p_->Name(); } + const ONNX_NAMESPACE::Provider_TensorShapeProto* Shape() const override { return &tensor_shape_proto_; } + virtual ONNX_NAMESPACE::DataType Type() const noexcept override { return p_->Type(); } + + const NodeArg* p_; + ONNX_NAMESPACE::Provider_TensorShapeProto tensor_shape_proto_; +}; + +struct Provider_Node_Impl : Provider_Node { + Provider_Node_Impl(const Node* p) : p_{p} {} + ~Provider_Node_Impl() override { + for (auto p : input_defs_) + delete p; + for (auto p : output_defs_) + delete p; + } + + const std::string& OpType() const noexcept override { return p_->OpType(); } + // const ONNX_NAMESPACE::OpSchema* Op() const noexcept + + ConstPointerContainer> InputDefs() const noexcept override { + if (input_defs_.empty()) { + for (auto p : p_->InputDefs()) + input_defs_.push_back(new Provider_NodeArg_Impl(p)); + } + + return ConstPointerContainer>(input_defs_); + } + + ConstPointerContainer> OutputDefs() const noexcept override { + if (output_defs_.empty()) { + for (auto p : p_->OutputDefs()) + output_defs_.push_back(new Provider_NodeArg_Impl(p)); + } + + return ConstPointerContainer>(output_defs_); + } + + NodeIndex Index() const noexcept override { return p_->Index(); } + + const Provider_NodeAttributes& GetAttributes() const noexcept override { + if (attributes_.empty()) { + for (auto& v : p_->GetAttributes()) + attributes_[v.first] = onnxruntime::make_unique(v.second); + } + return attributes_; + } + + size_t GetInputEdgesCount() const noexcept override { + return p_->GetInputEdgesCount(); + } + size_t GetOutputEdgesCount() const noexcept override { return p_->GetOutputEdgesCount(); } + + std::unique_ptr InputNodesBegin_internal() const noexcept override; + std::unique_ptr InputNodesEnd_internal() const noexcept override; + + const Node* p_; + mutable std::vector input_defs_; + mutable std::vector output_defs_; + mutable Provider_NodeAttributes attributes_; +}; + +struct Provider_NodeIterator_Impl : Provider_Node::Provider_NodeIterator { + Provider_NodeIterator_Impl(Node::NodeConstIterator&& v) : v_{std::move(v)} {} + + bool operator!=(const Provider_NodeIterator& p) const override { return v_ != static_cast(&p)->v_; } + + void operator++() override { return v_.operator++(); } + const Provider_Node& operator*() override { + node_ = Provider_Node_Impl(&*v_); + return node_; + } + + Node::NodeConstIterator v_; + Provider_Node_Impl node_{nullptr}; +}; + +std::unique_ptr Provider_Node_Impl::InputNodesBegin_internal() const noexcept { + return onnxruntime::make_unique(p_->InputNodesBegin()); +} + +std::unique_ptr Provider_Node_Impl::InputNodesEnd_internal() const noexcept { + return onnxruntime::make_unique(p_->InputNodesEnd()); +} + +struct Provider_IndexedSubGraph_Impl : Provider_IndexedSubGraph { + Provider_IndexedSubGraph_Impl() = default; + Provider_IndexedSubGraph_Impl(std::unique_ptr p) : p_{std::move(p)} {} + + void SetMetaDef(std::unique_ptr& def_) override { + auto real = onnxruntime::make_unique(); + + real->name = std::move(def_->name); + real->domain = std::move(def_->domain); + real->since_version = def_->since_version; + real->status = def_->status; + real->inputs = std::move(def_->inputs); + real->outputs = std::move(def_->outputs); + + for (const auto& v : def_->attributes) + real->attributes.emplace(v.first, static_cast(v.second.p_.get())->v_); + + real->doc_string = std::move(def_->doc_string); + + p_->SetMetaDef(real); + } + + std::vector& Nodes() override { return p_->nodes; } + + std::unique_ptr p_{onnxruntime::make_unique()}; +}; + +struct Provider_GraphViewer_Impl : Provider_GraphViewer { + Provider_GraphViewer_Impl(const GraphViewer& v) : v_(v) { + for (int i = 0; i < v_.MaxNodeIndex(); i++) + provider_nodes_.emplace_back(v_.GetNode(i)); + } + + const std::string& Name() const noexcept override { return v_.Name(); } + + const Provider_Node* GetNode(NodeIndex node_index) const override { + auto& node = provider_nodes_[node_index]; + if (node.p_) + return &node; + return nullptr; + } + + int MaxNodeIndex() const noexcept override { return v_.MaxNodeIndex(); } + + const Provider_InitializedTensorSet& GetAllInitializedTensors() const noexcept override { + if (initialized_tensor_set_.empty()) { + initialized_tensors_.reserve(v_.GetAllInitializedTensors().size()); + + for (auto& v : v_.GetAllInitializedTensors()) { + initialized_tensors_.emplace_back(const_cast(v.second)); + initialized_tensor_set_.emplace(v.first, &initialized_tensors_.back()); + } + } + + return initialized_tensor_set_; + } + + const std::unordered_map& DomainToVersionMap() const noexcept override { return v_.DomainToVersionMap(); } + + const GraphViewer& v_; + + std::vector provider_nodes_; + + mutable std::vector initialized_tensors_; + mutable Provider_InitializedTensorSet initialized_tensor_set_; +}; + +struct Provider_OpKernelInfo_Impl : Provider_OpKernelInfo { + Provider_OpKernelInfo_Impl(const OpKernelInfo& info) : info_(info) {} + + Status GetAttr(const std::string& name, int64_t* value) const override { + return info_.GetAttr(name, value); + } + + Status GetAttr(const std::string& name, float* value) const override { + return info_.GetAttr(name, value); + } + + const OpKernelInfo& info_; +}; + +struct Provider_Tensor_Impl final : Provider_Tensor { + Provider_Tensor_Impl(const Tensor* p) : p_(const_cast(p)) {} + + float* MutableData_float() override { return p_->MutableData(); } + const float* Data_float() const override { return p_->Data(); } + + const TensorShape& Shape() const override { return p_->Shape(); } + + Tensor* p_; +}; + +struct Provider_OpKernelContext_Impl : Provider_OpKernelContext { + Provider_OpKernelContext_Impl(OpKernelContext* context) : p_(context) {} + + const Provider_Tensor* Input_Tensor(int index) const override { + tensors_.push_back(onnxruntime::make_unique(p_->Input(index))); + return tensors_.back().get(); + } + + Provider_Tensor* Output(int index, const TensorShape& shape) override { + tensors_.push_back(onnxruntime::make_unique(p_->Output(index, shape))); + return tensors_.back().get(); + } + + OpKernelContext* p_; + mutable std::vector> tensors_; +}; + +struct Provider_OpKernel_Impl : Provider_OpKernel { + OpKernelInfo op_kernel_info_; +}; + +struct OpKernel_Translator : OpKernel { + OpKernel_Translator(Provider_OpKernelInfo_Impl& info, Provider_OpKernel* p) : OpKernel(info.info_), p_(p) {} + ~OpKernel_Translator() { + delete p_; + } + + Status Compute(OpKernelContext* context) const override { + Provider_OpKernelContext_Impl provider_context(context); + return p_->Compute(&provider_context); + } + + Provider_OpKernel* p_; +}; + +struct Provider_KernelRegistry_Impl : Provider_KernelRegistry { + Provider_KernelRegistry_Impl(std::shared_ptr p) : p_owned_(p) {} + Provider_KernelRegistry_Impl(KernelRegistry* p) : p_(p) {} + Provider_KernelRegistry_Impl() : p_owned_(std::make_shared()) {} + + Status Register(Provider_KernelCreateInfo&& create_info) override { + KernelCreateInfo info_real(std::move(static_cast(create_info.kernel_def.get())->p_), + [kernel_create_func = create_info.kernel_create_func](const OpKernelInfo& info) -> OpKernel* { + Provider_OpKernelInfo_Impl provider_info(info); + return new OpKernel_Translator(provider_info, kernel_create_func(provider_info)); + }); + + return p_->Register(std::move(info_real)); + } + + std::shared_ptr p_owned_; + KernelRegistry* p_{&*p_owned_}; +}; + +struct Provider_IExecutionProvider_Router_Impl : Provider_IExecutionProvider_Router, IExecutionProvider { + Provider_IExecutionProvider_Router_Impl(Provider_IExecutionProvider* outer, const std::string& type) : IExecutionProvider(type), outer_(outer) { + } + + virtual ~Provider_IExecutionProvider_Router_Impl() {} + + std::shared_ptr Provider_GetKernelRegistry() const override { + return std::make_shared(GetKernelRegistry()); + } + + std::shared_ptr GetKernelRegistry() const override { + return static_cast(&*outer_->Provider_GetKernelRegistry())->p_owned_; + } + + std::vector> Provider_GetCapability(const onnxruntime::Provider_GraphViewer& graph, + const std::vector& kernel_registries) const override { + std::vector kernel_registries_internal; + for (auto& v : kernel_registries) + kernel_registries_internal.emplace_back(static_cast(v)->p_); + + auto capabilities_internal = IExecutionProvider::GetCapability(static_cast(&graph)->v_, kernel_registries_internal); + + std::vector> capabilities; + for (auto& v : capabilities_internal) + capabilities.emplace_back(onnxruntime::make_unique(onnxruntime::make_unique(std::move(v->sub_graph)))); + return capabilities; + } + + std::vector> GetCapability(const onnxruntime::GraphViewer& graph, + const std::vector& kernel_registries) const override { + std::vector registries; + for (auto p : kernel_registries) + registries.push_back(new Provider_KernelRegistry_Impl(const_cast(p))); + + auto provider_result = outer_->Provider_GetCapability(Provider_GraphViewer_Impl(graph), registries); + std::vector> result; + + for (auto& p : provider_result) + result.emplace_back(onnxruntime::make_unique(std::move(static_cast(p->t_sub_graph_.get())->p_))); + + for (auto p : registries) + delete p; + + return result; + } + + common::Status Compile(const std::vector& fused_nodes, std::vector& node_compute_funcs) override { + std::vector provider_fused_nodes_values; + std::vector provider_fused_nodes; + provider_fused_nodes_values.reserve(fused_nodes.size()); + for (auto& p : fused_nodes) { + provider_fused_nodes_values.emplace_back(p); + provider_fused_nodes.emplace_back(&provider_fused_nodes_values.back()); + } + + return outer_->Provider_Compile(provider_fused_nodes, node_compute_funcs); + } + + Provider_AllocatorPtr Provider_GetAllocator(int id, OrtMemType mem_type) const override { + return std::make_shared(IExecutionProvider::GetAllocator(id, mem_type)); + } + + void Provider_InsertAllocator(Provider_AllocatorPtr allocator) override { + IExecutionProvider::InsertAllocator(static_cast(allocator.get())->p_); + } + + std::unique_ptr outer_; +}; + +struct ProviderHostImpl : ProviderHost { + ProviderHostImpl() { + DataTypeImpl_GetType_Tensor = &DataTypeImpl::GetType; + DataTypeImpl_GetType_float = &DataTypeImpl::GetType; + DataTypeImpl_GetTensorType_float = &DataTypeImpl::GetTensorType; + } + + std::unique_ptr AttributeProto_Create() override { + return onnxruntime::make_unique(); + } + + std::unique_ptr OrtMemoryInfo_Create(const char* name_, OrtAllocatorType type_, Provider_OrtDevice* device_, int id_, OrtMemType mem_type_) override { + return onnxruntime::make_unique(name_, type_, device_ ? static_cast(device_)->v_ : OrtDevice(), id_, mem_type_); + } + + std::unique_ptr KernelDefBuilder_Create() override { + return onnxruntime::make_unique(); + } + + std::shared_ptr KernelRegistry_Create() override { + return std::make_shared(); + } + + std::unique_ptr IndexedSubGraph_Create() override { + return onnxruntime::make_unique(); + } + + Provider_AllocatorPtr CreateAllocator(Provider_DeviceAllocatorRegistrationInfo& info, OrtDevice::DeviceId device_id = 0) override { + DeviceAllocatorRegistrationInfo info_real{ + info.mem_type, [&info](int value) { return std::move(static_cast(&*info.factory(value))->p_); }, + info.max_mem}; + + return std::make_shared(onnxruntime::CreateAllocator(info_real, device_id)); + } + + std::unique_ptr + CreateCPUAllocator(std::unique_ptr memory_info) override { + return onnxruntime::make_unique(onnxruntime::make_unique(std::move(static_cast(memory_info.get())->info_))); + }; + + Provider_AllocatorPtr + CreateDummyArenaAllocator(std::unique_ptr resource_allocator) override { + return std::make_shared(onnxruntime::make_unique(std::move(static_cast(resource_allocator.get())->p_))); + }; + + std::unique_ptr Create_IExecutionProvider_Router(Provider_IExecutionProvider* outer, const std::string& type) override { + return onnxruntime::make_unique(outer, type); + }; + + logging::Logger* LoggingManager_GetDefaultLogger() override { + return const_cast(&logging::LoggingManager::DefaultLogger()); + } + + void* HeapAllocate(size_t size) override { return new uint8_t[size]; } + void HeapFree(void* p) override { delete reinterpret_cast(p); } + + bool CPU_HasAVX2() override { + return CPUIDInfo::GetCPUIDInfo().HasAVX2(); + } + + bool CPU_HasAVX512f() override { + return CPUIDInfo::GetCPUIDInfo().HasAVX512f(); + } + + void LogRuntimeError(uint32_t session_id, const common::Status& status, const char* file, const char* function, uint32_t line) override { + return ::onnxruntime::LogRuntimeError(session_id, status, file, function, line); + } + +} provider_host_; + +struct ProviderLibrary { + ProviderLibrary(const char* filename) { + Env::Default().LoadDynamicLibrary(filename, &handle_); + if (!handle_) + return; + + Provider* (*PGetProvider)(); + Env::Default().GetSymbolFromLibrary(handle_, "GetProvider", (void**)&PGetProvider); + + provider_ = PGetProvider(); + provider_->SetProviderHost(provider_host_); + } + + ~ProviderLibrary() { + Env::Default().UnloadDynamicLibrary(handle_); + } + + Provider* provider_{}; + void* handle_{}; +}; + +// This class translates the IExecutionProviderFactory interface to work with the interface providers implement +struct IExecutionProviderFactory_Translator : IExecutionProviderFactory { + IExecutionProviderFactory_Translator(std::shared_ptr p) : p_{p} {} + + std::unique_ptr CreateProvider() override { + auto provider = p_->CreateProvider(); + return std::unique_ptr(static_cast(provider.release()->p_)); + } + + std::shared_ptr p_; +}; + +std::shared_ptr CreateExecutionProviderFactory_Dnnl(int device_id) { +#ifdef _WIN32 + static ProviderLibrary library("onnxruntime_providers_dnnl.dll"); +#else + static ProviderLibrary library("libonnxruntime_providers_dnnl.so"); +#endif + if (!library.provider_) { + LOGS_DEFAULT(ERROR) << "Failed to load provider shared library"; + return nullptr; + } + + //return std::make_shared(device_id); + //TODO: This is apparently a bug. The constructor parameter is create-arena-flag, not the device-id + return std::make_shared(library.provider_->CreateExecutionProviderFactory(device_id)); +} + +} // namespace onnxruntime + +// TODO: Right now Dnnl is the only provider in here, but this will be made more generic and support more providers in the future +ORT_API_STATUS_IMPL(OrtSessionOptionsAppendExecutionProvider_Dnnl, _In_ OrtSessionOptions* options, int use_arena) { + auto factory = onnxruntime::CreateExecutionProviderFactory_Dnnl(use_arena); + if (!factory) { + return OrtApis::CreateStatus(ORT_FAIL, "OrtSessionOptionsAppendExecutionProvider_Dnnl: Failed to load shared library"); + } + + options->provider_factories.push_back(factory); + return nullptr; +} diff --git a/onnxruntime/core/providers/dnnl/dnnl_common.h b/onnxruntime/core/providers/dnnl/dnnl_common.h index ee901506ff..90c22a9702 100644 --- a/onnxruntime/core/providers/dnnl/dnnl_common.h +++ b/onnxruntime/core/providers/dnnl/dnnl_common.h @@ -2,7 +2,8 @@ // Licensed under the MIT License. #pragma once -#include "core/common/common.h" +#include "core/providers/shared_library/provider_api.h" +#include "gsl/gsl-lite.hpp" #include "dnnl.hpp" #include #include @@ -65,9 +66,11 @@ class PrimitivePool { private: // For thread safety, the map needs to be kept in thread local storage. static inline std::unordered_map>& GetMap() { - static thread_local std::unordered_map> map; - return map; + using MapType = std::unordered_map>; + static thread_local DeleteOnUnloadPtr map(new MapType()); + return *map; } }; + } // namespace ort_dnnl } // namespace onnxruntime diff --git a/onnxruntime/core/providers/dnnl/dnnl_execution_provider.cc b/onnxruntime/core/providers/dnnl/dnnl_execution_provider.cc index fb17342e7f..96c50e3798 100644 --- a/onnxruntime/core/providers/dnnl/dnnl_execution_provider.cc +++ b/onnxruntime/core/providers/dnnl/dnnl_execution_provider.cc @@ -5,42 +5,43 @@ #pragma warning(disable : 4996) #endif -#include "core/framework/allocator.h" -#include "core/framework/compute_capability.h" -#include "core/framework/kernel_registry.h" -#include "core/providers/dnnl/subgraph/dnnl_func_kernel.h" +#include "core/providers/shared_library/provider_api.h" +#include +#include "subgraph/dnnl_func_kernel.h" #include "dnnl_execution_provider.h" #include "dnnl_fwd.h" namespace { + struct KernelRegistryAndStatus { - std::shared_ptr kernel_registry = std::make_shared(); + std::shared_ptr kernel_registry{onnxruntime::Provider_KernelRegistry::Create()}; + Status st; }; + } // namespace + namespace onnxruntime { constexpr const char* DNNL = "Dnnl"; constexpr const char* DNNL_CPU = "DnnlCpu"; DNNLExecutionProvider::DNNLExecutionProvider(const DNNLExecutionProviderInfo& info) - : IExecutionProvider{onnxruntime::kDnnlExecutionProvider} { - DeviceAllocatorRegistrationInfo default_memory_info({OrtMemTypeDefault, - [](int) { return onnxruntime::make_unique(onnxruntime::make_unique(DNNL, OrtAllocatorType::OrtDeviceAllocator)); }, std::numeric_limits::max()}); + : Provider_IExecutionProvider{onnxruntime::kDnnlExecutionProvider} { + Provider_DeviceAllocatorRegistrationInfo default_memory_info({OrtMemTypeDefault, + [](int) { return onnxruntime::CreateCPUAllocator(onnxruntime::Provider_OrtMemoryInfo::Create(DNNL, OrtAllocatorType::OrtDeviceAllocator)); }, std::numeric_limits::max()}); - DeviceAllocatorRegistrationInfo cpu_memory_info({OrtMemTypeCPUOutput, - [](int) { return onnxruntime::make_unique(onnxruntime::make_unique(DNNL_CPU, OrtAllocatorType::OrtDeviceAllocator, OrtDevice(), 0, OrtMemTypeCPUOutput)); }, std::numeric_limits::max()}); + Provider_DeviceAllocatorRegistrationInfo cpu_memory_info({OrtMemTypeCPUOutput, + [](int) { return onnxruntime::CreateCPUAllocator(onnxruntime::Provider_OrtMemoryInfo::Create(DNNL_CPU, OrtAllocatorType::OrtDeviceAllocator, nullptr, 0, OrtMemTypeCPUOutput)); }, std::numeric_limits::max()}); if (info.create_arena) { - InsertAllocator(CreateAllocator(default_memory_info)); + Provider_InsertAllocator(CreateAllocator(default_memory_info)); - InsertAllocator(CreateAllocator(cpu_memory_info)); + Provider_InsertAllocator(CreateAllocator(cpu_memory_info)); } else { - InsertAllocator(std::shared_ptr( - onnxruntime::make_unique(default_memory_info.factory(0)))); + Provider_InsertAllocator(onnxruntime::CreateDummyArenaAllocator(default_memory_info.factory(0))); - InsertAllocator(std::shared_ptr( - onnxruntime::make_unique(cpu_memory_info.factory(0)))); + Provider_InsertAllocator(onnxruntime::CreateDummyArenaAllocator(cpu_memory_info.factory(0))); } } // namespace onnxruntime @@ -50,8 +51,8 @@ DNNLExecutionProvider::~DNNLExecutionProvider() { namespace ort_dnnl { class ONNX_OPERATOR_KERNEL_CLASS_NAME(kDnnlExecutionProvider, kOnnxDomain, 7, Gemm); -Status RegisterDNNLKernels(KernelRegistry& kernel_registry) { - static const BuildKernelCreateInfoFn function_table[] = { +Status RegisterDNNLKernels(Provider_KernelRegistry& kernel_registry) { + static const Provider_BuildKernelCreateInfoFn function_table[] = { BuildKernelCreateInfo, }; @@ -68,14 +69,14 @@ KernelRegistryAndStatus GetDnnlKernelRegistry() { } } // namespace ort_dnnl -std::shared_ptr DNNLExecutionProvider::GetKernelRegistry() const { +std::shared_ptr DNNLExecutionProvider::Provider_GetKernelRegistry() const { static KernelRegistryAndStatus k = onnxruntime::ort_dnnl::GetDnnlKernelRegistry(); // throw if the registry failed to initialize ORT_THROW_IF_ERROR(k.st); return k.kernel_registry; } -bool DNNLExecutionProvider::UseSubgraph(const onnxruntime::GraphViewer& graph_viewer) const { +bool DNNLExecutionProvider::UseSubgraph(const onnxruntime::Provider_GraphViewer& graph_viewer) const { bool use_subgraph = true; bool FP16_graph = false; @@ -120,12 +121,12 @@ bool DNNLExecutionProvider::UseSubgraph(const onnxruntime::GraphViewer& graph_vi return use_subgraph; } -void DNNLExecutionProvider::CreateOrUpdateDnnlNode(const Node* node, - std::shared_ptr& subgraph_ptr, - ort_dnnl::Subgraph::SubgraphVariables& sub_var, - bool fused, - std::map& output_to_source_node_map, - NodeAttributes& subgraph_attributes) const { +void DNNLExecutionProvider::CreateOrUpdateDnnlNode(const Provider_Node* node, + std::shared_ptr& subgraph_ptr, + ort_dnnl::Subgraph::SubgraphVariables& sub_var, + bool fused, + std::map& output_to_source_node_map, + Provider_NodeAttributes& subgraph_attributes) const { const auto& node_inputs = node->InputDefs(); sub_var.outputs.push_back(node->OutputDefs()[0]->Name()); @@ -168,7 +169,7 @@ void DNNLExecutionProvider::CreateOrUpdateDnnlNode(const Node* node, } } - NodeAttributes attributes = node->GetAttributes(); + const Provider_NodeAttributes& attributes = node->GetAttributes(); if (attributes.size() > 0) { size_t index = subgraph_ptr->dnnl_nodes.size(); std::string op_name; @@ -182,24 +183,23 @@ void DNNLExecutionProvider::CreateOrUpdateDnnlNode(const Node* node, for (auto att_it = attributes.begin(); att_it != attributes.end(); ++att_it) { std::string key = op_name + "-" + std::to_string(index) + "-" + att_it->first; - std::pair att(key, att_it->second); subgraph_attributes[key] = att_it->second; } } } -std::vector> DNNLExecutionProvider::GetCapability( - const onnxruntime::GraphViewer& graph_viewer, - const std::vector& kernel_registries) const { +std::vector> DNNLExecutionProvider::Provider_GetCapability( + const onnxruntime::Provider_GraphViewer& graph_viewer, + const std::vector& kernel_registries) const { ORT_UNUSED_PARAMETER(kernel_registries); if (UseSubgraph(graph_viewer) == false) { - return IExecutionProvider::GetCapability(graph_viewer, kernel_registries); + return Provider_IExecutionProvider::Provider_GetCapability(graph_viewer, kernel_registries); } LOGS_DEFAULT(INFO) << "Using DNNL Subgraph"; // use sub-graph implementation - std::vector> result; + std::vector> result; ort_dnnl::Subgraph::SubgraphVariables sub_var; std::shared_ptr subgraph_ptr; @@ -213,7 +213,7 @@ std::vector> DNNLExecutionProvider::GetCapabi // output name to node index map. Using it to find sub-graph end nodes // if output of a node is not an input to any node in a sub-graph is end node std::map output_to_source_node_map; - NodeAttributes subgraph_attributes; + Provider_NodeAttributes subgraph_attributes; int node_index = 0; while (node_index < graph_viewer.MaxNodeIndex()) { @@ -353,11 +353,11 @@ std::vector> DNNLExecutionProvider::GetCapabi return result; } -void DNNLExecutionProvider::CreateMetaDef(const onnxruntime::GraphViewer& graph_viewer, - const NodeAttributes& subgraph_attributes, - std::shared_ptr& subgraph_ptr, - ort_dnnl::Subgraph::SubgraphVariables& sub_var, - std::vector>& result) const { +void DNNLExecutionProvider::CreateMetaDef(const onnxruntime::Provider_GraphViewer& graph_viewer, + const Provider_NodeAttributes& subgraph_attributes, + std::shared_ptr& subgraph_ptr, + ort_dnnl::Subgraph::SubgraphVariables& sub_var, + std::vector>& result) const { std::string graph_fused_nodes; std::string node_list; std::string subgraph_id = std::to_string(subgraph_index_); @@ -368,19 +368,19 @@ void DNNLExecutionProvider::CreateMetaDef(const onnxruntime::GraphViewer& graph_ std::unordered_set input_initializers; // Create ng_required_initializers attribute of NGraphCustomOp - ONNX_NAMESPACE::AttributeProto initializers; - initializers.set_name("initializers"); - initializers.set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_TENSORS); + auto initializers = ONNX_NAMESPACE::Provider_AttributeProto::Create(); + initializers->set_name("initializers"); + initializers->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_TENSORS); for (const auto& init : sub_var.inputs) { if (graph_viewer.GetAllInitializedTensors().count(init)) { - auto tensor = initializers.add_tensors(); + auto tensor = initializers->add_tensors(); *tensor = *(graph_viewer.GetAllInitializedTensors().at(init)); } } - auto meta_def = onnxruntime::make_unique<::onnxruntime::IndexedSubGraph::MetaDef>(); - meta_def->attributes["initializers"] = initializers; + auto meta_def = onnxruntime::make_unique<::onnxruntime::Provider_IndexedSubGraph::MetaDef>(); + meta_def->attributes["initializers"] = std::move(initializers); meta_def->name = "DnnlCustomOp" + std::to_string(subgraph_index_); meta_def->domain = kMSDomain; meta_def->since_version = 1; @@ -398,22 +398,22 @@ void DNNLExecutionProvider::CreateMetaDef(const onnxruntime::GraphViewer& graph_ } } - ONNX_NAMESPACE::AttributeProto ap; - ap.set_s(subgraph_id); - ap.set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING); - meta_def->attributes["subgraph_id"] = ap; - std::unique_ptr sub_graph = onnxruntime::make_unique(); - sub_graph->nodes = sub_var.subgraph_node_indexes; + auto ap = ONNX_NAMESPACE::Provider_AttributeProto::Create(); + ap->set_s(subgraph_id); + ap->set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING); + meta_def->attributes["subgraph_id"] = std::move(ap); + auto sub_graph = onnxruntime::Provider_IndexedSubGraph::Create(); + sub_graph->Nodes() = sub_var.subgraph_node_indexes; sub_graph->SetMetaDef(meta_def); - result.push_back(onnxruntime::make_unique(std::move(sub_graph))); + result.push_back(onnxruntime::make_unique(std::move(sub_graph))); mkl_subgraphs_.insert(std::make_pair(subgraph_id, subgraph_ptr)); // Reset subgraph and meta_Def sub_var.Reset(); } -Status DNNLExecutionProvider::Compile(const std::vector& fused_nodes, - std::vector& node_compute_funcs) { +Status DNNLExecutionProvider::Provider_Compile(const std::vector& fused_nodes, + std::vector& node_compute_funcs) { for (const auto* fused_node : fused_nodes) { auto attributes = fused_node->GetAttributes(); NodeComputeInfo compute_info; @@ -438,4 +438,5 @@ Status DNNLExecutionProvider::Compile(const std::vector& fus } return Status::OK(); } + } // namespace onnxruntime diff --git a/onnxruntime/core/providers/dnnl/dnnl_execution_provider.h b/onnxruntime/core/providers/dnnl/dnnl_execution_provider.h index a551a5dc66..77ebdd9e5f 100644 --- a/onnxruntime/core/providers/dnnl/dnnl_execution_provider.h +++ b/onnxruntime/core/providers/dnnl/dnnl_execution_provider.h @@ -8,9 +8,7 @@ #include #include -#include "core/graph/constants.h" -#include "core/framework/allocatormgr.h" -#include "core/framework/execution_provider.h" +#include "core/platform/ort_mutex.h" #include "core/providers/dnnl/subgraph/subgraph.h" #include "core/platform/ort_mutex.h" @@ -30,12 +28,12 @@ struct DNNLExecutionProviderInfo { }; // Logical device representation. -class DNNLExecutionProvider : public IExecutionProvider { +class DNNLExecutionProvider : public Provider_IExecutionProvider { public: explicit DNNLExecutionProvider(const DNNLExecutionProviderInfo& info); virtual ~DNNLExecutionProvider(); - virtual std::shared_ptr GetKernelRegistry() const override; + virtual std::shared_ptr Provider_GetKernelRegistry() const override; std::shared_ptr GetWeightsMemoryBuffer(const std::string& weight_key) { auto iter = weights_mem_map_.find(weight_key); @@ -76,12 +74,12 @@ class DNNLExecutionProvider : public IExecutionProvider { biass_buffers_.push_back(std::move(buffer)); } - std::vector> - GetCapability(const onnxruntime::GraphViewer& graph, - const std::vector& /*kernel_registries*/) const override; + std::vector> + Provider_GetCapability(const onnxruntime::Provider_GraphViewer& graph, + const std::vector& /*kernel_registries*/) const override; - common::Status Compile(const std::vector& fused_nodes, - std::vector& node_compute_funcs) override; + common::Status Provider_Compile(const std::vector& fused_nodes, + std::vector& node_compute_funcs) override; private: // dnnl weights(filer data) memory blocks from first iteration @@ -98,12 +96,12 @@ class DNNLExecutionProvider : public IExecutionProvider { // SUBGRAPH private: - static int GetOnnxOpSet(const GraphViewer& graph_viewer) { + static int GetOnnxOpSet(const Provider_GraphViewer& graph_viewer) { const auto& dm_to_ver = graph_viewer.DomainToVersionMap(); return dm_to_ver.at(kOnnxDomain); } - std::string GetGraphName(const onnxruntime::GraphViewer& graph_viewer) const { + std::string GetGraphName(const onnxruntime::Provider_GraphViewer& graph_viewer) const { std::string graph_name; int opset = GetOnnxOpSet(graph_viewer); @@ -121,12 +119,12 @@ class DNNLExecutionProvider : public IExecutionProvider { return graph_name; } - bool UseSubgraph(const onnxruntime::GraphViewer& graph_viewer) const; + bool UseSubgraph(const onnxruntime::Provider_GraphViewer& graph_viewer) const; // Some dimensions are not supported by DNNL // example: Pool with NumDimensions <= 3 is not supported // Fall back to CPU implementation - bool IsDimensionSupported(const Node* node) const { + bool IsDimensionSupported(const Provider_Node* node) const { bool supported = true; if (node->OpType() == "BatchNormalization") { auto node_inputs = node->InputDefs(); @@ -146,20 +144,20 @@ class DNNLExecutionProvider : public IExecutionProvider { return supported; } - void CreateOrUpdateDnnlNode(const Node* node, - std::shared_ptr& subgraph_ptr, - ort_dnnl::Subgraph::SubgraphVariables& sub_var, - bool fused, - std::map& output_to_source_node_map, - NodeAttributes& subgraph_attributes) const; + void CreateOrUpdateDnnlNode(const Provider_Node* node, + std::shared_ptr& subgraph_ptr, + ort_dnnl::Subgraph::SubgraphVariables& sub_var, + bool fused, + std::map& output_to_source_node_map, + Provider_NodeAttributes& subgraph_attributes) const; // Create Dnnl node, update inputs, outputs and parent nodes // collect attribtes - void CreateMetaDef(const onnxruntime::GraphViewer& graph_viewer, - const NodeAttributes& subgraph_attributes, + void CreateMetaDef(const onnxruntime::Provider_GraphViewer& graph_viewer, + const Provider_NodeAttributes& subgraph_attributes, std::shared_ptr& subgraph_ptr, ort_dnnl::Subgraph::SubgraphVariables& sub_var, - std::vector>& result) const; + std::vector>& result) const; public: const std::shared_ptr GetDnnlSubgraph(const std::string& subgraph_id) { @@ -171,7 +169,7 @@ class DNNLExecutionProvider : public IExecutionProvider { // supported Dnnl Operators std::set dnnl_ops_ = {"Conv", "BatchNormalization", "Relu", "Sum", - "AveragePool", "GlobalMaxPool", "GlobalAveragePool", "MaxPool", "LRN"}; + "AveragePool", "GlobalMaxPool", "GlobalAveragePool", "MaxPool", "LRN"}; mutable std::unordered_map> mkl_subgraphs_; }; diff --git a/onnxruntime/core/providers/dnnl/dnnl_fwd.h b/onnxruntime/core/providers/dnnl/dnnl_fwd.h index dc9417c3fe..d0a42737fa 100644 --- a/onnxruntime/core/providers/dnnl/dnnl_fwd.h +++ b/onnxruntime/core/providers/dnnl/dnnl_fwd.h @@ -6,6 +6,6 @@ namespace onnxruntime { namespace ort_dnnl { template -KernelCreateInfo BuildKernelCreateInfo(); -} +Provider_KernelCreateInfo BuildKernelCreateInfo(); } +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/dnnl/dnnl_provider_factory.cc b/onnxruntime/core/providers/dnnl/dnnl_provider_factory.cc index 65760d8c48..28c6f76f63 100644 --- a/onnxruntime/core/providers/dnnl/dnnl_provider_factory.cc +++ b/onnxruntime/core/providers/dnnl/dnnl_provider_factory.cc @@ -1,38 +1,52 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#include "core/providers/shared_library/provider_api.h" #include "core/providers/dnnl/dnnl_provider_factory.h" #include #include "dnnl_execution_provider.h" -#include "core/session/abi_session_options_impl.h" using namespace onnxruntime; +extern onnxruntime::ProviderHost* g_host; + namespace onnxruntime { -struct DnnlProviderFactory : IExecutionProviderFactory { + +void SetProviderHost(ProviderHost& host); + +struct DnnlProviderFactory : Provider_IExecutionProviderFactory { DnnlProviderFactory(bool create_arena) : create_arena_(create_arena) {} ~DnnlProviderFactory() override {} - std::unique_ptr CreateProvider() override; + std::unique_ptr CreateProvider() override; private: bool create_arena_; }; -std::unique_ptr DnnlProviderFactory::CreateProvider() { +std::unique_ptr DnnlProviderFactory::CreateProvider() { DNNLExecutionProviderInfo info; info.create_arena = create_arena_; return onnxruntime::make_unique(info); } -std::shared_ptr CreateExecutionProviderFactory_Dnnl(int device_id) { - return std::make_shared(device_id); - //TODO: This is apparently a bug. The consructor parameter is create-arena-flag, not the device-id -} +struct Dnnl_Provider : Provider { + std::shared_ptr CreateExecutionProviderFactory(int device_id) override { + //TODO: This is apparently a bug. The consructor parameter is create-arena-flag, not the device-id + // Will be fixed by PR #2850 + return std::make_shared(device_id); + } + + void SetProviderHost(ProviderHost& host) { + onnxruntime::SetProviderHost(host); + } +} g_provider; } // namespace onnxruntime -ORT_API_STATUS_IMPL(OrtSessionOptionsAppendExecutionProvider_Dnnl, _In_ OrtSessionOptions* options, int use_arena) { - options->provider_factories.push_back(onnxruntime::CreateExecutionProviderFactory_Dnnl(use_arena)); - return nullptr; +extern "C" { + +ORT_API(onnxruntime::Provider*, GetProvider) { + return &onnxruntime::g_provider; +} } diff --git a/onnxruntime/core/providers/dnnl/exported_symbols.lst b/onnxruntime/core/providers/dnnl/exported_symbols.lst new file mode 100644 index 0000000000..f4c4141259 --- /dev/null +++ b/onnxruntime/core/providers/dnnl/exported_symbols.lst @@ -0,0 +1 @@ +_GetProvider diff --git a/onnxruntime/core/providers/dnnl/math/gemm.cc b/onnxruntime/core/providers/dnnl/math/gemm.cc index faa405750b..c8239b9e14 100644 --- a/onnxruntime/core/providers/dnnl/math/gemm.cc +++ b/onnxruntime/core/providers/dnnl/math/gemm.cc @@ -1,12 +1,13 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#include "core/providers/shared_library/provider_api.h" #include "gemm.h" -#include "core/providers/cpu/math/gemm_helper.h" -#include "core/util/math_cpuonly.h" #include "dnnl.h" #include "dnnl.hpp" #include "core/providers/dnnl/dnnl_fwd.h" +#include "gsl/gsl" +#include "Eigen/Core" namespace onnxruntime { namespace ort_dnnl { @@ -16,14 +17,85 @@ ONNX_OPERATOR_KERNEL_EX( kOnnxDomain, 7, kDnnlExecutionProvider, - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + Provider_KernelDefBuilder::Create()->TypeConstraint("T", DataTypeImpl::GetTensorType()), Gemm); +class GemmHelper { + public: + GemmHelper(const TensorShape& left, bool trans_left, const TensorShape& right, bool trans_right, const TensorShape& bias) { + ORT_ENFORCE(left.NumDimensions() == 2 || left.NumDimensions() == 1); + ORT_ENFORCE(right.NumDimensions() == 2); + + if (trans_left) { + M_ = left.NumDimensions() == 2 ? left[1] : left[0]; + K_ = left.NumDimensions() == 2 ? left[0] : 1; + } else { + M_ = left.NumDimensions() == 2 ? left[0] : 1; + K_ = left.NumDimensions() == 2 ? left[1] : left[0]; + } + + int k_dim; + if (trans_right) { + N_ = right[0]; + k_dim = 1; + } else { + N_ = right[1]; + k_dim = 0; + } + + if (right[k_dim] != K_) + status_ = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "GEMM: Dimension mismatch, W: ", + right.ToString(), + " K: " + std::to_string(K_), + " N:" + std::to_string(N_)); + + if (!IsValidBroadcast(bias, M_, N_)) + status_ = common::Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, "Gemm: Invalid bias shape for broadcast"); + + // it is possible the input is empty tensor, for example the output of roipool in fast rcnn. + ORT_ENFORCE(M_ >= 0 && K_ > 0 && N_ >= 0); + } + + int64_t M() const { return M_; } + int64_t N() const { return N_; } + int64_t K() const { return K_; } + Status State() const { return status_; } + + private: + bool IsValidBroadcast(const TensorShape& bias_shape, int64_t M, int64_t N) { + // valid shapes are (,) , (1, N) , (M, 1) , (M, N) + if (bias_shape.NumDimensions() > 2) + return false; + // shape is (1,) or (1, 1), or (,) + if (bias_shape.Size() == 1) + return true; + // valid bias_shape (s) are (N,) or (1, N) or (M, 1) or (M, N), + // In last case no broadcasting needed, so don't fail it + return ((bias_shape.NumDimensions() == 1 && bias_shape[0] == N) || + (bias_shape.NumDimensions() == 2 && bias_shape[0] == M && (bias_shape[1] == 1 || bias_shape[1] == N)) || + (bias_shape.NumDimensions() == 2 && bias_shape[0] == 1 && bias_shape[1] == N)); + } + + private: + int64_t M_; + int64_t K_; + int64_t N_; + Status status_; +}; + +template +using EigenMatrixMapRowMajor = Eigen::Map>; +template +using ConstEigenVectorMap = Eigen::Map>; +template +using ConstEigenMatrixMapRowMajor = Eigen::Map>; + template <> -Status Gemm::Compute(OpKernelContext* ctx) const { - const auto X = ctx->Input(0); - const auto W = ctx->Input(1); - const auto B = ctx->Input(2); +Status Gemm::Compute(Provider_OpKernelContext* ctx) const { + const auto X = ctx->Input(0); + const auto W = ctx->Input(1); + const auto B = ctx->Input(2); GemmHelper helper(X->Shape(), trans_A_, W->Shape(), trans_B_, B->Shape()); if (!helper.State().IsOK()) @@ -35,7 +107,7 @@ Status Gemm::Compute(OpKernelContext* ctx) const { auto Y = ctx->Output(0, TensorShape({M, N})); if (M <= 0) - return ORT_MAKE_STATUS(ONNXRUNTIME, NOT_IMPLEMENTED, "Empty Tensor not supported"); + return ORT_MAKE_STATUS(ONNXRUNTIME, NOT_IMPLEMENTED, "Empty Tensor not supported"); if (beta_ != 0) { auto output_mat = EigenMatrixMapRowMajor( @@ -83,11 +155,11 @@ Status Gemm::Compute(OpKernelContext* ctx) const { // dnnl_sgemm expects row major matrices, so no need to swap the operands A and B auto status = dnnl_sgemm(trans_A_ ? 'T' : 'N', - trans_B_ ? 'T' : 'N', - M, N, K, - alpha_, X->template Data() , trans_A_ ? M : K, - W->template Data(), trans_B_ ? K : N, - beta_, Y->template MutableData(), N); + trans_B_ ? 'T' : 'N', + M, N, K, + alpha_, X->template Data(), trans_A_ ? M : K, + W->template Data(), trans_B_ ? K : N, + beta_, Y->template MutableData(), N); if (status == dnnl_success) { return Status::OK(); } else { diff --git a/onnxruntime/core/providers/dnnl/math/gemm.h b/onnxruntime/core/providers/dnnl/math/gemm.h index 9503ed6468..d570a21e0f 100644 --- a/onnxruntime/core/providers/dnnl/math/gemm.h +++ b/onnxruntime/core/providers/dnnl/math/gemm.h @@ -2,14 +2,13 @@ // Licensed under the MIT License. #pragma once -#include "core/framework/op_kernel.h" namespace onnxruntime { namespace ort_dnnl { template -class Gemm final : public OpKernel { +class Gemm final : public Provider_OpKernel { public: - Gemm(const OpKernelInfo& info) : OpKernel(info) { + Gemm(const Provider_OpKernelInfo& info) : Provider_OpKernel(info) { int64_t temp; ORT_ENFORCE(info.GetAttr("transA", &temp).IsOK()); trans_A_ = (temp != 0); @@ -21,7 +20,7 @@ class Gemm final : public OpKernel { ORT_ENFORCE(info.GetAttr("beta", &beta_).IsOK()); } - Status Compute(OpKernelContext* context) const override; + Status Compute(Provider_OpKernelContext* context) const override; private: bool trans_A_; diff --git a/onnxruntime/core/providers/dnnl/memcpy_s.h b/onnxruntime/core/providers/dnnl/memcpy_s.h index 3db9715711..46fd81f3fb 100644 --- a/onnxruntime/core/providers/dnnl/memcpy_s.h +++ b/onnxruntime/core/providers/dnnl/memcpy_s.h @@ -3,7 +3,7 @@ #pragma once -#ifdef _WIN32 +#ifndef MIN #define MIN(a, b) ((a) < (b) ? (a) : (b)) #endif diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_activations.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_activations.h index a2382b13ff..e19d795c1f 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_activations.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_activations.h @@ -2,9 +2,6 @@ // Licensed under the MIT License #pragma once -#include "core/util/math.h" -#include "core/util/math_cpuonly.h" -#include "core/framework/op_kernel.h" #include "core/providers/dnnl/dnnl_fwd.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" @@ -16,9 +13,9 @@ template class DnnlRelu : public DnnlKernel { public: DnnlRelu(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { ORT_UNUSED_PARAMETER(attributes); ORT_UNUSED_PARAMETER(attributes_prefix); } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_batchnorm.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_batchnorm.h index 0f54f6cf46..f2eee1cce9 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_batchnorm.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_batchnorm.h @@ -2,12 +2,10 @@ // Licensed under the MIT License #pragma once -#include "core/framework/op_kernel.h" #include "core/providers/dnnl/dnnl_fwd.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" #include "core/providers/dnnl/memcpy_s.h" -#include "core/util/math.h" namespace onnxruntime { namespace ort_dnnl { @@ -83,17 +81,17 @@ template class DnnlBatchNorm : public DnnlKernel { public: explicit DnnlBatchNorm(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { ReadAttributes(attributes, attributes_prefix); } - void ReadAttributes(const NodeAttributes& attributes, + void ReadAttributes(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") override { auto attr = attributes.find(attributes_prefix + "epsilon"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { - epsilon_ = attr->second.f(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { + epsilon_ = attr->second->f(); } } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv.h index 09169ec3df..45ac68db1c 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv.h @@ -2,13 +2,10 @@ // Licensed under the MIT License. #pragma once -#include "dnnl_types.h" -#include "core/framework/op_kernel.h" +#include "mkldnn_types.h" #include "core/providers/dnnl/dnnl_fwd.h" -#include "core/providers/cpu/nn/autopad_type.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" -#include "core/util/math.h" namespace onnxruntime { namespace ort_dnnl { @@ -64,9 +61,9 @@ template class DnnlConv : public DnnlKernel { public: DnnlConv(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { ReadAttributes(attributes, attributes_prefix); } @@ -390,7 +387,7 @@ class DnnlConv : public DnnlKernel { if (filter_dst_mem == nullptr) { dnnl::memory src = dnnl::memory({{filter_dims_mkl}, DnnnType(), filter_format_}, cpu_engine, (void*)filter_data); IAllocatorUniquePtr filter_reorder_buffer = - IAllocator::MakeUniquePtr(alloc_, filter_size_); + Provider_IAllocator::MakeUniquePtr(alloc_, filter_size_); filter_dst_mem = onnxruntime::make_unique( dnnl::memory(conv_fwd_pd_->weights_desc(), cpu_engine, filter_reorder_buffer.get())); @@ -440,7 +437,7 @@ class DnnlConv : public DnnlKernel { } auto src_size = conv_fwd_pd_.get()->src_desc().get_size(); - src_reorder_buffer_ = IAllocator::MakeUniquePtr(alloc_, src_size); + src_reorder_buffer_ = Provider_IAllocator::MakeUniquePtr(alloc_, src_size); src_mem_->set_data_handle(src_reorder_buffer_.get()); } else { if (mklnode_ptr_->parent_nodes.empty()) { @@ -468,34 +465,34 @@ class DnnlConv : public DnnlKernel { } private: - void ReadAttributes(const NodeAttributes& attributes, + void ReadAttributes(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") override { std::string auto_pad; auto attr = attributes.find(attributes_prefix + "auto_pad"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { - auto_pad = attr->second.s(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { + auto_pad = attr->second->s(); } auto_pad_ = (auto_pad != "") ? StringToAutoPadType(auto_pad) : AutoPadType::NOTSET; kernel_shape_specified_ = false; attr = attributes.find(attributes_prefix + "kernel_shape"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; Status status = GetIntsAttr(proto, kernel_shape_); kernel_shape_specified_ = true; } attr = attributes.find(attributes_prefix + "strides"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; Status status = GetIntsAttr(proto, strides_); } bool attr_read = false; attr = attributes.find(attributes_prefix + "pads"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntsAttr(proto, pads_) == Status::OK()) attr_read = true; } @@ -506,7 +503,7 @@ class DnnlConv : public DnnlKernel { attr_read = false; attr = attributes.find(attributes_prefix + "dilations"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntsAttr(proto, dilations_) == Status::OK()) attr_read = true; } @@ -517,7 +514,7 @@ class DnnlConv : public DnnlKernel { attr_read = false; attr = attributes.find(attributes_prefix + "group"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntAttr(proto, group_) == Status::OK()) attr_read = true; } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv_batchnorm.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv_batchnorm.h index 087518f354..93f2e052bc 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv_batchnorm.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_conv_batchnorm.h @@ -3,13 +3,10 @@ #pragma once #include "dnnl_types.h" -#include "core/framework/op_kernel.h" #include "core/providers/dnnl/dnnl_fwd.h" -#include "core/providers/cpu/nn/autopad_type.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" -#include "core/util/math.h" -#include + namespace onnxruntime { namespace ort_dnnl { @@ -17,9 +14,9 @@ template class DnnlConvBatchNorm : public DnnlKernel { public: DnnlConvBatchNorm(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { ReadAttributes(attributes, attributes_prefix); } @@ -398,7 +395,7 @@ class DnnlConvBatchNorm : public DnnlKernel { if (filter_dst_mem == nullptr) { dnnl::memory src = dnnl::memory({{filter_dims_mkl}, DnnnType(), filter_format_}, cpu_engine, (void*)weights_scaled_by_axis.data()); IAllocatorUniquePtr filter_reorder_buffer = - IAllocator::MakeUniquePtr(alloc_, filter_size_); + Provider_IAllocator::MakeUniquePtr(alloc_, filter_size_); filter_dst_mem = onnxruntime::make_unique( dnnl::memory(conv_fwd_pd_->weights_desc(), cpu_engine, filter_reorder_buffer.get())); @@ -413,7 +410,7 @@ class DnnlConvBatchNorm : public DnnlKernel { if (bias_mem == nullptr) { auto bias_size = conv_fwd_pd_.get()->bias_desc().get_size(); IAllocatorUniquePtr bias_buffer = - IAllocator::MakeUniquePtr(alloc_, bias_size); + Provider_IAllocator::MakeUniquePtr(alloc_, bias_size); bias_mem = onnxruntime::make_unique( dnnl::memory(conv_fwd_pd_->bias_desc(), cpu_engine, bias_buffer.get())); float* bias_buffer_data = static_cast(bias_buffer.get()); @@ -482,7 +479,7 @@ class DnnlConvBatchNorm : public DnnlKernel { } auto src_size = conv_fwd_pd_.get()->src_desc().get_size(); - src_reorder_buffer_ = IAllocator::MakeUniquePtr(alloc_, src_size); + src_reorder_buffer_ = Provider_IAllocator::MakeUniquePtr(alloc_, src_size); src_mem_->set_data_handle(src_reorder_buffer_.get()); } else { if (mklnode_ptr_->parent_nodes.empty()) { @@ -510,34 +507,34 @@ class DnnlConvBatchNorm : public DnnlKernel { } private: - void ReadAttributes(const NodeAttributes& attributes, + void ReadAttributes(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") override { std::string auto_pad; auto attr = attributes.find(attributes_prefix + "auto_pad"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { - auto_pad = attr->second.s(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { + auto_pad = attr->second->s(); } auto_pad_ = (auto_pad != "") ? StringToAutoPadType(auto_pad) : AutoPadType::NOTSET; kernel_shape_specified_ = false; attr = attributes.find(attributes_prefix + "kernel_shape"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; Status status = GetIntsAttr(proto, kernel_shape_); kernel_shape_specified_ = true; } attr = attributes.find(attributes_prefix + "strides"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; Status status = GetIntsAttr(proto, strides_); } bool attr_read = false; attr = attributes.find(attributes_prefix + "pads"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntsAttr(proto, pads_) == Status::OK()) attr_read = true; } @@ -548,7 +545,7 @@ class DnnlConvBatchNorm : public DnnlKernel { attr_read = false; attr = attributes.find(attributes_prefix + "dilations"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntsAttr(proto, dilations_) == Status::OK()) attr_read = true; } @@ -559,7 +556,7 @@ class DnnlConvBatchNorm : public DnnlKernel { attr_read = false; attr = attributes.find(attributes_prefix + "group"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntAttr(proto, group_) == Status::OK()) attr_read = true; } @@ -569,7 +566,7 @@ class DnnlConvBatchNorm : public DnnlKernel { attr = attributes.find(attributes_prefix + "epsilon"); if (attr != attributes.end()) { - epsilon_ = attr->second.f(); + epsilon_ = attr->second->f(); } } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.cc index fc469ae142..fac0a73866 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.cc +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.cc @@ -5,7 +5,7 @@ #endif #include "dnnl_func_kernel.h" -#include "core/common/exceptions.h" +#define EXCLUDE_REFERENCE_TO_ORT_DLL #include "core/session/onnxruntime_cxx_api.h" #include "core/providers/dnnl/dnnl_common.h" #include "core/providers/dnnl/subgraph/dnnl_conv.h" @@ -15,7 +15,6 @@ #include "core/providers/dnnl/subgraph/dnnl_pool.h" #include "core/providers/dnnl/subgraph/dnnl_sum.h" #include "core/providers/dnnl/subgraph/dnnl_lrn.h" -#include "core/session/onnxruntime_cxx_api.h" namespace onnxruntime { namespace ort_dnnl { @@ -223,11 +222,8 @@ class SubgraphPrimitivePool : public PrimitivePool { auto tensor_info = ort.GetTensorTypeAndShape(input_tensor); auto tensor_shape = ort.GetTensorShape(tensor_info); ort.ReleaseTensorTypeAndShapeInfo(tensor_info); - auto shape = tensor_shape.data(); - auto dim = tensor_shape.size(); - TensorShape x_shape(shape, dim); - dnnl::memory::dims src_dims(x_shape.GetDims().begin(), x_shape.GetDims().end()); + dnnl::memory::dims src_dims(tensor_shape); AddDimsToKey(dims_str, src_dims); } @@ -261,8 +257,7 @@ Status DnnlFuncKernel::Compute(const OrtCustomOpApi* api, OrtKernelContext* c primitive->UpdateProvider(params_); status = primitive->Compute(api, context); } catch (const dnnl::error& e) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, - ", message: ", e.what()); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.what()); } return status; } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.h index ee4aac0f76..2e3cc6a22d 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_func_kernel.h @@ -2,10 +2,9 @@ // Licensed under the MIT License #pragma once -#include "core/graph/onnx_protobuf.h" +#include "core/providers/shared_library/provider_api.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/session/onnxruntime_c_api.h" -#include "core/framework/func_api.h" #include "dnnl_kernel.h" namespace onnxruntime { @@ -13,7 +12,7 @@ namespace ort_dnnl { namespace { struct SubgraphParams { - NodeAttributes attributes; + Provider_NodeAttributes attributes; DNNLExecutionProvider* provider; std::shared_ptr subgraph; std::string subgraph_id; @@ -27,16 +26,16 @@ template class DnnlFuncKernel { public: explicit DnnlFuncKernel(const ComputeContext* context, - const NodeAttributes& attributes, - DNNLExecutionProvider* provider) { + const Provider_NodeAttributes& attributes, + DNNLExecutionProvider* provider) { ORT_UNUSED_PARAMETER(context); params_.provider = provider; params_.attributes = attributes; auto sub_it = attributes.find("subgraph_id"); - if (sub_it->second.type() == ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { - params_.subgraph_id = sub_it->second.s(); + if (sub_it->second->type() == ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { + params_.subgraph_id = sub_it->second->s(); params_.subgraph = provider->GetDnnlSubgraph(params_.subgraph_id); std::ostringstream key_os; @@ -66,14 +65,14 @@ class DnnlFuncKernel { } } - std::string GetPoolAttributesKey(const NodeAttributes& attributes, + std::string GetPoolAttributesKey(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") { std::string key; auto attr = attributes.find(attributes_prefix + "kernel_shape"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -83,13 +82,13 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "auto_pad"); if (attr != attributes.end()) { - key.append(attr->second.s()); + key.append(attr->second->s()); } attr = attributes.find(attributes_prefix + "pads"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -100,7 +99,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "strides"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -111,7 +110,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "count_include_pad"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(std::to_string(proto.i())); key.append(1, '_'); key.append(1, '#'); @@ -120,7 +119,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "ceil_mode"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(std::to_string(proto.i())); key.append(1, '_'); key.append(1, '#'); @@ -128,14 +127,14 @@ class DnnlFuncKernel { return key; } - std::string GetConvAttributeKey(const NodeAttributes& attributes, + std::string GetConvAttributeKey(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") { std::string key; auto attr = attributes.find(attributes_prefix + "dilations"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -146,7 +145,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "auto_pad"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(proto.s()); key.append(1, '#'); } @@ -154,7 +153,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "pads"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -165,7 +164,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "strides"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -176,7 +175,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "kernel_shape"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; for (int i = 0; i < proto.ints_size(); i++) { key.append(std::to_string(proto.ints(i))); key.append(1, '_'); @@ -187,7 +186,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "group"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(std::to_string(proto.i())); key.append(1, '#'); } @@ -195,14 +194,14 @@ class DnnlFuncKernel { return key; } - std::string GetLrnAttributeKey(const NodeAttributes& attributes, + std::string GetLrnAttributeKey(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") { std::string key; auto attr = attributes.find(attributes_prefix + "alpha"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(std::to_string(proto.f())); key.append(1, '#'); } @@ -210,7 +209,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "beta"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(std::to_string(proto.f())); key.append(1, '#'); } @@ -218,7 +217,7 @@ class DnnlFuncKernel { attr = attributes.find(attributes_prefix + "bias"); if (attr != attributes.end()) { key.append(1, '#'); - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; key.append(std::to_string(proto.f())); key.append(1, '#'); } @@ -232,4 +231,4 @@ class DnnlFuncKernel { SubgraphParams params_; }; } // namespace ort_dnnl -} // namespace onnxruntime \ No newline at end of file +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.cc index f627d7d779..9c07aab299 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.cc +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.cc @@ -1,22 +1,23 @@ // Copyright(C) 2019 Intel Corporation // Licensed under the MIT License +#include "core/providers/shared_library/provider_api.h" #include "dnnl_kernel.h" namespace onnxruntime { namespace ort_dnnl { void DnnlKernel::InitDstReorderOutput(dnnl::engine& cpu_engine, - dnnl::memory::data_type& data_type, - std::vector& net, - std::vector>& net_args) { + dnnl::memory::data_type& data_type, + std::vector& net, + std::vector>& net_args) { // Allocate dst buffer if reorder is necessary if (primitive_dst_desc_ != ort_source_desc_) { // reorder to ONNXRuntime format dnnl::memory::dims dst_dims_mkl( primitive_dst_shape_.GetDims().begin(), primitive_dst_shape_.GetDims().end()); dnnl::memory::desc dst_des = dnnl::memory::desc(dst_dims_mkl, - data_type, ort_source_format_); + data_type, ort_source_format_); reorder_dst_mem_to_ = onnxruntime::make_unique( dnnl::memory(dst_des, cpu_engine)); net.push_back(dnnl::reorder(*primitive_dst_mem_, *reorder_dst_mem_to_)); @@ -58,4 +59,4 @@ dnnl::memory::format_tag DnnlKernel::GetSourceFormat(int dim_size) { } } // namespace ort_dnnl -} // namespace onnxruntime \ No newline at end of file +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.h index ef7a0a1bdd..eb9832ecf9 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_kernel.h @@ -7,8 +7,6 @@ #endif #include "dnnl.hpp" -#include "core/common/cpuid_info.h" -#include "core/session/onnxruntime_cxx_api.h" #include "core/providers/dnnl/subgraph/subgraph.h" #include "core/providers/dnnl/dnnl_execution_provider.h" @@ -18,11 +16,11 @@ namespace ort_dnnl { class DnnlKernel { public: DnnlKernel(const DnnlNode& node, - DNNLExecutionProvider* provider) { + DNNLExecutionProvider* provider) { name_ = node.name; mklnode_ptr_ = std::make_shared(node); provider_ = provider; - alloc_ = provider_->GetAllocator(0, OrtMemTypeDefault); + alloc_ = provider_->Provider_GetAllocator(0, OrtMemTypeDefault); } virtual ~DnnlKernel(){}; @@ -47,13 +45,13 @@ class DnnlKernel { virtual Status Bind(const OrtCustomOpApi* api, OrtKernelContext* context) = 0; protected: - virtual void ReadAttributes(const NodeAttributes& attributes, + virtual void ReadAttributes(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") { ORT_UNUSED_PARAMETER(attributes); ORT_UNUSED_PARAMETER(attributes_prefix); } - Status GetIntsAttr(ONNX_NAMESPACE::AttributeProto& proto, std::vector& values) { + Status GetIntsAttr(ONNX_NAMESPACE::Provider_AttributeProto& proto, std::vector& values) { ORT_RETURN_IF_NOT(proto.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INTS); values.reserve(proto.ints_size()); for (int i = 0; i < proto.ints_size(); i++) { @@ -62,18 +60,18 @@ class DnnlKernel { return Status::OK(); } - Status GetIntAttr(ONNX_NAMESPACE::AttributeProto& proto, int64_t& value) { + Status GetIntAttr(ONNX_NAMESPACE::Provider_AttributeProto& proto, int64_t& value) { ORT_RETURN_IF_NOT(proto.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INT); value = proto.i(); return Status::OK(); } - Status GetFloatAttr(ONNX_NAMESPACE::AttributeProto& proto, float& value) { + Status GetFloatAttr(ONNX_NAMESPACE::Provider_AttributeProto& proto, float& value) { ORT_RETURN_IF_NOT(proto.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT); value = proto.f(); return Status::OK(); } - Status GetStringAttr(ONNX_NAMESPACE::AttributeProto& proto, std::string& value) { + Status GetStringAttr(ONNX_NAMESPACE::Provider_AttributeProto& proto, std::string& value) { ORT_RETURN_IF_NOT(proto.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING); value = proto.s(); return Status::OK(); @@ -113,7 +111,7 @@ class DnnlKernel { // memory used for reorders std::unique_ptr reorder_dst_mem_to_; - AllocatorPtr alloc_; + Provider_AllocatorPtr alloc_; DNNLExecutionProvider* provider_; }; diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_lrn.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_lrn.h index 4c856a2d0a..5f19954d29 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_lrn.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_lrn.h @@ -2,9 +2,6 @@ // Licensed under the MIT License. #pragma once -#include "core/util/math.h" -#include "core/util/math_cpuonly.h" -#include "core/framework/op_kernel.h" #include "core/providers/dnnl/dnnl_fwd.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" @@ -16,9 +13,9 @@ template class DnnlLrn : public DnnlKernel { public: DnnlLrn(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { ReadAttributes(attributes, attributes_prefix); } @@ -66,7 +63,7 @@ class DnnlLrn : public DnnlKernel { dnnl::algorithm algo = dnnl::algorithm::lrn_across_channels; fwd_desc_ = onnxruntime::make_unique( dnnl::lrn_forward::desc(dnnl::prop_kind::forward_scoring, algo, *src_md_, - size_, alpha_, beta_, bias_)); + size_, alpha_, beta_, bias_)); fwd_primitive_desc_ = onnxruntime::make_unique( dnnl::lrn_forward::primitive_desc(*fwd_desc_, cpu_engine)); @@ -135,33 +132,33 @@ class DnnlLrn : public DnnlKernel { } private: - void ReadAttributes(const NodeAttributes& attributes, + void ReadAttributes(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") override { auto attr = attributes.find(attributes_prefix + "size"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INT) { - size_ = attr->second.i(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INT) { + size_ = attr->second->i(); } ORT_ENFORCE(size_ > 0); ORT_ENFORCE(size_ % 2 == 1); attr = attributes.find(attributes_prefix + "alpha"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { - alpha_ = attr->second.f(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { + alpha_ = attr->second->f(); } attr = attributes.find(attributes_prefix + "beta"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { - beta_ = attr->second.f(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { + beta_ = attr->second->f(); } bias_ = 1.0f; attr = attributes.find(attributes_prefix + "bias"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { - bias_ = attr->second.f(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_FLOAT) { + bias_ = attr->second->f(); } } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_pool.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_pool.h index 369471fceb..b33438a05b 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_pool.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_pool.h @@ -3,10 +3,8 @@ #pragma once #include "core/providers/dnnl/dnnl_fwd.h" -#include "core/providers/cpu/nn/autopad_type.h" #include "core/providers/dnnl/dnnl_execution_provider.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" -#include "core/util/math.h" namespace onnxruntime { namespace ort_dnnl { @@ -14,9 +12,9 @@ template class DnnlPool : public DnnlKernel { public: DnnlPool(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { op_name_ = node.name; ReadAttributes(attributes, attributes_prefix); } @@ -105,9 +103,9 @@ class DnnlPool : public DnnlKernel { } fwd_desc_ = onnxruntime::make_unique( dnnl::pooling_forward::desc(dnnl::prop_kind::forward_inference, algo, - *src_md_, *primitive_dst_md_, - strides_mkl, kernel_mkl, - padding_left_mkl, padding_right_mkl)); + *src_md_, *primitive_dst_md_, + strides_mkl, kernel_mkl, + padding_left_mkl, padding_right_mkl)); fwd_primitive_desc_ = onnxruntime::make_unique( dnnl::pooling_forward::primitive_desc(*fwd_desc_, cpu_engine)); @@ -200,7 +198,7 @@ class DnnlPool : public DnnlKernel { } auto src_size = fwd_primitive_desc_.get()->src_desc().get_size(); - src_reorder_buffer_ = IAllocator::MakeUniquePtr(alloc_, src_size); + src_reorder_buffer_ = Provider_IAllocator::MakeUniquePtr(alloc_, src_size); src_mem_->set_data_handle(src_reorder_buffer_.get()); } else { if (mklnode_ptr_->parent_nodes.empty()) { @@ -230,7 +228,7 @@ class DnnlPool : public DnnlKernel { } private: - void ReadAttributes(const NodeAttributes& attributes, + void ReadAttributes(const Provider_NodeAttributes& attributes, const std::string attributes_prefix = "") override { global_pooling_ = (op_name_ == "GlobalAveragePool" || op_name_ == "GlobalMaxPool" || op_name_ == "GlobalLpPool"); global_pooling_ = (op_name_ == "GlobalAveragePool" || op_name_ == "GlobalMaxPool" || op_name_ == "GlobalLpPool"); @@ -239,7 +237,7 @@ class DnnlPool : public DnnlKernel { bool attr_read = false; auto attr = attributes.find(attributes_prefix + "kernel_shape"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; GetIntsAttr(proto, kernel_shape_); attr_read = true; } @@ -248,15 +246,15 @@ class DnnlPool : public DnnlKernel { std::string auto_padding; attr = attributes.find(attributes_prefix + "auto_pad"); if (attr != attributes.end() && - attr->second.type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { - auto_padding = attr->second.s(); + attr->second->type() == ::ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { + auto_padding = attr->second->s(); } auto_pad_ = StringToAutoPadType(auto_padding); attr_read = false; attr = attributes.find(attributes_prefix + "pads"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntsAttr(proto, pads_) == Status::OK()) attr_read = true; } @@ -267,7 +265,7 @@ class DnnlPool : public DnnlKernel { attr_read = false; attr = attributes.find(attributes_prefix + "strides"); if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; if (GetIntsAttr(proto, strides_) == Status::OK()) attr_read = true; } @@ -278,7 +276,7 @@ class DnnlPool : public DnnlKernel { attr = attributes.find(attributes_prefix + "count_include_pad"); int64_t temp = 0; if (attr != attributes.end()) { - ONNX_NAMESPACE::AttributeProto proto = attr->second; + auto& proto = *attr->second; GetIntAttr(proto, temp); } count_include_pad_ = (temp != 0); diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_sum.h b/onnxruntime/core/providers/dnnl/subgraph/dnnl_sum.h index 1eb645bcf4..e1edd95d9f 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_sum.h +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_sum.h @@ -2,11 +2,9 @@ // Licensed under the MIT License #pragma once -#include "core/framework/op_kernel.h" #include "core/providers/dnnl/dnnl_fwd.h" #include "core/providers/dnnl/dnnl_common.h" #include "core/providers/dnnl/subgraph/dnnl_kernel.h" -#include "core/util/math.h" namespace onnxruntime { namespace ort_dnnl { @@ -15,9 +13,9 @@ template class DnnlSum : public DnnlKernel { public: explicit DnnlSum(const DnnlNode& node, - DNNLExecutionProvider* provider, - const NodeAttributes& attributes, - const std::string attributes_prefix = "") : DnnlKernel(node, provider) { + DNNLExecutionProvider* provider, + const Provider_NodeAttributes& attributes, + const std::string attributes_prefix = "") : DnnlKernel(node, provider) { ReadAttributes(attributes, attributes_prefix); } diff --git a/onnxruntime/core/providers/dnnl/subgraph/subgraph.h b/onnxruntime/core/providers/dnnl/subgraph/subgraph.h index d63fe251be..a633f94f90 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/subgraph.h +++ b/onnxruntime/core/providers/dnnl/subgraph/subgraph.h @@ -6,8 +6,7 @@ #include #include #include -#include "core/framework/op_node_proto_helper.h" -#include "core/graph/graph.h" +#include "core/providers/providers.h" namespace onnxruntime { namespace ort_dnnl { @@ -20,7 +19,7 @@ struct DnnlNode { int output_index = -1; // index in output() std::string weight_name; std::string output_name; - std::vector parent_nodes; // index to parents in vector mklnodes + std::vector parent_nodes; // index to parents in vector mklnodes std::string ToString() const { std::string key; @@ -49,7 +48,7 @@ struct Subgraph { std::vector outputs_as_input_other_node; std::vector subgraph_node_indexes; - void Reset() { + void Reset() { subgraph_node_indexes.clear(); inputs.clear(); outputs.clear(); @@ -66,4 +65,4 @@ struct Subgraph { std::vector dnnl_nodes; }; } // namespace ort_dnnl -} // namespace onnxruntime \ No newline at end of file +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/dnnl/symbols.def b/onnxruntime/core/providers/dnnl/symbols.def new file mode 100644 index 0000000000..4ec2f7914c --- /dev/null +++ b/onnxruntime/core/providers/dnnl/symbols.def @@ -0,0 +1,2 @@ +EXPORTS + GetProvider diff --git a/onnxruntime/core/providers/dnnl/version_script.lds b/onnxruntime/core/providers/dnnl/version_script.lds new file mode 100644 index 0000000000..094abb3329 --- /dev/null +++ b/onnxruntime/core/providers/dnnl/version_script.lds @@ -0,0 +1,9 @@ +#_init and _fini should be local +VERS_1.0 { + global: + GetProvider; + + # Hide everything else. + local: + *; +}; diff --git a/onnxruntime/core/providers/shared_library/provider_api.h b/onnxruntime/core/providers/shared_library/provider_api.h new file mode 100644 index 0000000000..ec199b7679 --- /dev/null +++ b/onnxruntime/core/providers/shared_library/provider_api.h @@ -0,0 +1,260 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// Provider implementations include this file + +// NOTE: This is still in development so there are many parts that will be fixed in the future. This is just the first version of +// switching providers to be runnable as shared libraries. The interfaces will become more tightly integrated into the core code. + +#pragma once + +#include +#include +#include +#include +#include "onnx/common/stl_backports.h" +#include "core/common/common.h" +#include "core/common/const_pointer_container.h" +#include "core/session/onnxruntime_c_api.h" +#include "provider_interfaces.h" + +namespace ONNX_NAMESPACE { + +enum AttributeProto_AttributeType : int { + AttributeProto_AttributeType_UNDEFINED = 0, + AttributeProto_AttributeType_FLOAT = 1, + AttributeProto_AttributeType_INT = 2, + AttributeProto_AttributeType_STRING = 3, + AttributeProto_AttributeType_TENSOR = 4, + AttributeProto_AttributeType_GRAPH = 5, + AttributeProto_AttributeType_SPARSE_TENSOR = 11, + AttributeProto_AttributeType_FLOATS = 6, + AttributeProto_AttributeType_INTS = 7, + AttributeProto_AttributeType_STRINGS = 8, + AttributeProto_AttributeType_TENSORS = 9, + AttributeProto_AttributeType_GRAPHS = 10, + AttributeProto_AttributeType_SPARSE_TENSORS = 12 +}; + +enum OperatorStatus : int { + EXPERIMENTAL = 0, + STABLE = 1 +}; + +} // namespace ONNX_NAMESPACE + +namespace onnxruntime { + +// The function passed in will be run on provider DLL unload. This is used to free thread_local variables that are in threads we don't own +// Since these are not destroyed when the DLL unloads we have to do it manually. Search for usage for an example. +void RunOnUnload(std::function function); + +// A pointer stored in here will be deleted when the DLL gets unloaded, this is really only useful for thread_locals which don't get cleaned up properly otherwise +template +struct DeleteOnUnloadPtr { + DeleteOnUnloadPtr(T* p) : p_(p) { + RunOnUnload([p = p_]() { + delete p; + }); + } + + operator T*() { + return p_; + } + + private: + T* p_; +}; + +constexpr const char* kOnnxDomain = ""; +constexpr const char* kDnnlExecutionProvider = "DnnlExecutionProvider"; + +class DataTypeImpl { + public: + virtual ~DataTypeImpl() = default; + + template + static MLDataType GetType(); + template + static MLDataType GetTensorType(); +}; + +class TensorShape : private std::vector { + public: + TensorShape() = default; + + TensorShape(const TensorShape& /*other*/) = default; + TensorShape& operator=(const TensorShape& /*other*/) = default; + + TensorShape(TensorShape&& /*other*/) = default; + TensorShape& operator=(TensorShape&& /*other*/) = default; + + TensorShape(const std::vector& dims) : std::vector{dims} {} + TensorShape(std::vector&& dims) : std::vector{dims} {} + TensorShape(const std::initializer_list& dims) : std::vector{dims} {} + + TensorShape(const int64_t* dimension_sizes, size_t dimension_count); + TensorShape(const std::vector& dims, size_t start, size_t end); + + using std::vector::operator[]; + + size_t NumDimensions() const noexcept { + return size(); + } + + const std::vector& GetDims() const { return *this; } + + int64_t Size() const; + + /** + Return a new TensorShape of the dimensions from dimstart to dimend. + */ + TensorShape Slice(size_t dimstart, size_t dimend) const; + + /** + Return a new TensorShape of the dimensions from dimstart to end. + */ + TensorShape Slice(size_t dimstart) const; + + /** + output dimensions nicely formatted + */ + std::string ToString() const; + + /** + Calculate size between start and end. + Assumes start and end are between 0 and this->NumDimensions(), inclusive, and that + start < end. + */ + int64_t SizeHelper(size_t start, size_t end) const; +}; + +constexpr const char* kMSDomain = "com.microsoft"; +constexpr const char* kMklDnnExecutionProvider = "MKLDNNExecutionProvider"; + +template +using IAllocatorUniquePtr = std::unique_ptr>; + +std::unique_ptr CreateCPUAllocator(std::unique_ptr memory_info); +Provider_AllocatorPtr CreateDummyArenaAllocator(std::unique_ptr resource_allocator); +Provider_AllocatorPtr CreateAllocator(Provider_DeviceAllocatorRegistrationInfo& info, int16_t device_id = 0); + +class CPUIDInfo { + public: + static const CPUIDInfo& GetCPUIDInfo(); + + bool HasAVX2() const; + bool HasAVX512f() const; +}; + +namespace logging { + +enum class Severity { + kVERBOSE = 0, + kINFO = 1, + kWARNING = 2, + kERROR = 3, + kFATAL = 4 +}; + +enum class DataType { + SYSTEM = 0, ///< System data. + USER = 1 ///< Contains potentially sensitive user data. +}; + +struct Category { + static const char* onnxruntime; ///< General output + static const char* System; ///< Log output regarding interactions with the host system + // TODO: What other high level categories are meaningful? Model? Optimizer? Execution? +}; + +constexpr const char* SEVERITY_PREFIX = "VIWEF"; + +class Logger { + public: + bool OutputIsEnabled(Severity severity, DataType data_type) const noexcept; +}; + +class LoggingManager { + public: + static const Logger& DefaultLogger(); +}; + +class Capture { + public: + Capture(const Logger& logger, logging::Severity severity, const char* category, + logging::DataType dataType, const CodeLocation& location); + + std::ostream& Stream() noexcept; +}; +} // namespace logging + +enum class AutoPadType { + NOTSET = 0, + VALID = 1, + SAME_UPPER = 2, + SAME_LOWER = 3, +}; + +// TODO(RyanHill): Move this to a host function +inline AutoPadType StringToAutoPadType(const std::string& str) { + if (str.empty()) { + return AutoPadType::NOTSET; + } + if (str == "NOTSET") { // in onnx spec, default value is "NOTSET" + return AutoPadType::NOTSET; + } + if (str == "VALID") { + return AutoPadType::VALID; + } + if (str == "SAME_UPPER") { + return AutoPadType::SAME_UPPER; + } + if (str == "SAME_LOWER") { + return AutoPadType::SAME_LOWER; + } + ORT_ENFORCE(false, "Unknown AutoPadType String"); +} + +namespace math { + +// Rounds a up to the next highest multiple of b, which is power-of-2. User must be careful +// to ensure that there is no overflow or underflow in the calculation +// of divUp. +template +constexpr T roundUpPow2(T a) { + return (a + (b - 1)) & (~(b - 1)); +} +} // namespace math +} // namespace onnxruntime + +#define ONNX_OPERATOR_KERNEL_CLASS_NAME(provider, domain, ver, name) \ + provider##_##name##_##domain##_ver##ver + +#define ONNX_OPERATOR_KERNEL_EX(name, domain, ver, provider, builder, ...) \ + class ONNX_OPERATOR_KERNEL_CLASS_NAME(provider, domain, ver, name); \ + template <> \ + Provider_KernelCreateInfo \ + BuildKernelCreateInfo() { \ + return Provider_KernelCreateInfo( \ + builder.SetName(#name) \ + .SetDomain(domain) \ + .SinceVersion(ver) \ + .Provider(provider) \ + .Build(), \ + static_cast([](const Provider_OpKernelInfo& info) -> Provider_OpKernel* { return new __VA_ARGS__(info); })); \ + } + +#define CREATE_MESSAGE(logger, severity, category, datatype) \ + ::onnxruntime::logging::Capture(logger, ::onnxruntime::logging::Severity::k##severity, category, datatype, ORT_WHERE) + +// iostream style logging. Capture log info in Message, and push to the logger in ~Message. +#define LOGS_CATEGORY(logger, severity, category) \ + if ((logger).OutputIsEnabled(::onnxruntime::logging::Severity::k##severity, ::onnxruntime::logging::DataType::SYSTEM)) \ + CREATE_MESSAGE(logger, severity, category, ::onnxruntime::logging::DataType::SYSTEM).Stream() + +#define LOGS_DEFAULT_CATEGORY(severity, category) \ + LOGS_CATEGORY(::onnxruntime::logging::LoggingManager::DefaultLogger(), severity, category) + +#define LOGS_DEFAULT(severity) \ + LOGS_DEFAULT_CATEGORY(severity, ::onnxruntime::logging::Category::onnxruntime) diff --git a/onnxruntime/core/providers/shared_library/provider_bridge_provider.cc b/onnxruntime/core/providers/shared_library/provider_bridge_provider.cc new file mode 100644 index 0000000000..3c95f526a2 --- /dev/null +++ b/onnxruntime/core/providers/shared_library/provider_bridge_provider.cc @@ -0,0 +1,273 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// This is the provider DLL side of the provider API to let providers be built as a DLL + +#include "provider_api.h" +#include +#include +#include // For std::cout used in a stub + +onnxruntime::ProviderHost* g_host{}; + +#define PROVIDER_NOT_IMPLEMENTED ORT_THROW("Unimplemented shared library provider method"); + +namespace onnxruntime { + +void SetProviderHost(ProviderHost& host) { + g_host = &host; +} + +static std::unique_ptr>> s_run_on_unload_; + +void RunOnUnload(std::function function) { + static std::mutex mutex; + std::lock_guard guard{mutex}; + if (!s_run_on_unload_) + s_run_on_unload_ = onnxruntime::make_unique>>(); + s_run_on_unload_->push_back(std::move(function)); +} + +// This object is destroyed as part of the DLL unloading code and handles running all of the RunOnLoad functions +struct OnUnload { + ~OnUnload() { + if (!s_run_on_unload_) + return; + + for (auto& function : *s_run_on_unload_) + function(); + + s_run_on_unload_.reset(); + } + +} g_on_unload; + +} // namespace onnxruntime + +// Override default new/delete so that we match the host's allocator +void* operator new(size_t n) { return g_host->HeapAllocate(n); } +void operator delete(void* p) { return g_host->HeapFree(p); } +void operator delete(void* p, size_t /*size*/) { return g_host->HeapFree(p); } + +namespace onnx { +std::unique_ptr Provider_AttributeProto::Create() { + return g_host->AttributeProto_Create(); +} +} // namespace onnx + +namespace onnxruntime { + +Provider_AllocatorPtr CreateAllocator(Provider_DeviceAllocatorRegistrationInfo& info, int16_t device_id) { + return g_host->CreateAllocator(info, device_id); +} + +std::unique_ptr Provider_KernelDefBuilder::Create() { + return g_host->KernelDefBuilder_Create(); +} + +std::shared_ptr Provider_KernelRegistry::Create() { + return g_host->KernelRegistry_Create(); +} + +std::unique_ptr Provider_OrtMemoryInfo::Create(const char* name_, OrtAllocatorType type_, Provider_OrtDevice* device_, int id_, OrtMemType mem_type_) { + return g_host->OrtMemoryInfo_Create(name_, type_, device_, id_, mem_type_); +} + +std::unique_ptr Provider_IndexedSubGraph::Create() { + return g_host->IndexedSubGraph_Create(); +} + +template <> +MLDataType DataTypeImpl::GetType() { + return g_host->DataTypeImpl_GetType_float(); +} + +template <> +MLDataType DataTypeImpl::GetTensorType() { + return g_host->DataTypeImpl_GetTensorType_float(); +} + +TensorShape::TensorShape(const int64_t* dimension_sizes, size_t dimension_count) + : std::vector(dimension_count) { + for (size_t i = 0; i < dimension_count; ++i) { + (*this)[i] = dimension_sizes[i]; + } +} + +TensorShape::TensorShape(const std::vector& dims, size_t start, size_t end) { + assign(dims.begin() + start, dims.begin() + end); +} + +int64_t TensorShape::Size() const { + size_t arraySize = size(); + int64_t size = SizeHelper(0, arraySize); + //should we cache the size? as multiple operation may be expensive. + return size; +} + +int64_t TensorShape::SizeHelper(size_t start, size_t end) const { + // Must return 1 for an empty sequence + int64_t size = 1; + for (size_t i = start; i < end; i++) { + if ((*this)[i] < 0) return -1; + size *= (*this)[i]; + } + return size; +} + +TensorShape TensorShape::Slice(size_t dimstart, size_t dimend) const { + assert(dimstart <= dimend && dimend <= size()); // "Invalid tensor shape slice argument." + return TensorShape(*this, dimstart, dimend); +} + +TensorShape TensorShape::Slice(size_t dimstart) const { + return Slice(dimstart, size()); +} + +std::string TensorShape::ToString() const { + std::string result; + + result.append("{"); + bool first = true; + for (auto dim : (*this)) { + if (!first) { + result.append(","); + } + + result.append(std::to_string(dim)); + first = false; + } + result.append("}"); + + return result; +} + +CPUIDInfo g_info; + +const CPUIDInfo& CPUIDInfo::GetCPUIDInfo() { + return g_info; +} + +bool CPUIDInfo::HasAVX2() const { + return g_host->CPU_HasAVX2(); +} + +bool CPUIDInfo::HasAVX512f() const { + return g_host->CPU_HasAVX512f(); +} + +Provider_AllocatorPtr CreateAllocator(Provider_DeviceAllocatorRegistrationInfo info, int16_t device_id) { + return g_host->CreateAllocator(info, device_id); +} + +std::unique_ptr CreateCPUAllocator(std::unique_ptr info) { + return g_host->CreateCPUAllocator(std::move(info)); +} + +Provider_AllocatorPtr CreateDummyArenaAllocator(std::unique_ptr resource_allocator) { + return g_host->CreateDummyArenaAllocator(std::move(resource_allocator)); +} + +Provider_IExecutionProvider::Provider_IExecutionProvider(const std::string& type) { + p_ = g_host->Create_IExecutionProvider_Router(this, type).release(); +} + +namespace logging { + +bool Logger::OutputIsEnabled(Severity severity, DataType data_type) const noexcept { + ORT_UNUSED_PARAMETER(severity); + ORT_UNUSED_PARAMETER(data_type); + return false; + // TODO: Logging not essential to make it work initially, do later +} + +static Logger g_default_logger; + +const Logger& LoggingManager::DefaultLogger() { + return g_default_logger; +} + +Capture::Capture(const Logger& logger, logging::Severity severity, const char* category, + logging::DataType dataType, const CodeLocation& location) { + PROVIDER_NOT_IMPLEMENTED + ORT_UNUSED_PARAMETER(logger); + ORT_UNUSED_PARAMETER(severity); + ORT_UNUSED_PARAMETER(category); + ORT_UNUSED_PARAMETER(dataType); + ORT_UNUSED_PARAMETER(location); +} + +std::ostream& Capture::Stream() noexcept { + // PROVIDER_NOT_IMPLEMENTED + return std::cout; +} + +const char* Category::onnxruntime = "foo"; + +} // namespace logging + +namespace common { + +Status::Status(StatusCategory category, int code, const std::string& msg) { + // state_ will be allocated here causing the status to be treated as a failure + ORT_ENFORCE(code != static_cast(common::OK)); + + state_ = onnxruntime::make_unique(category, code, msg); +} + +Status::Status(StatusCategory category, int code, const char* msg) { + // state_ will be allocated here causing the status to be treated as a failure + ORT_ENFORCE(code != static_cast(common::OK)); + + state_ = onnxruntime::make_unique(category, code, msg); +} + +int Status::Code() const noexcept { + return IsOK() ? static_cast(common::OK) : state_->code; +} + +const std::string& Status::ErrorMessage() const noexcept { + return IsOK() ? EmptyString() : state_->msg; +} + +std::string Status::ToString() const { + if (state_ == nullptr) { + return std::string("OK"); + } + + std::string result; + + if (common::SYSTEM == state_->category) { + result += "SystemError"; + result += " : "; + result += std::to_string(errno); + } else if (common::ONNXRUNTIME == state_->category) { + result += "[ONNXRuntimeError]"; + result += " : "; + result += std::to_string(Code()); + result += " : "; + result += StatusCodeToString(static_cast(Code())); + result += " : "; + result += state_->msg; + } + + return result; +} + +const std::string& Status::EmptyString() noexcept { + static std::string s_empty; + return s_empty; +} + +} // namespace common + +std::vector GetStackTrace() { + // PROVIDER_NOT_IMPLEMENTED + return {}; +} + +void LogRuntimeError(uint32_t session_id, const common::Status& status, const char* file, const char* function, uint32_t line) { + return g_host->LogRuntimeError(session_id, status, file, function, line); +} + +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/shared_library/provider_interfaces.h b/onnxruntime/core/providers/shared_library/provider_interfaces.h new file mode 100644 index 0000000000..bb2daa1182 --- /dev/null +++ b/onnxruntime/core/providers/shared_library/provider_interfaces.h @@ -0,0 +1,445 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// Public wrappers around internal ort interfaces (currently) +// In the future the internal implementations could derive from these to remove the need for the wrapper implementations + +#include "core/framework/func_api.h" + +namespace ONNX_NAMESPACE { +enum AttributeProto_AttributeType : int; +enum OperatorStatus : int; + +// String pointer as unique TypeProto identifier. +using DataType = const std::string*; + +struct Provider_TensorProto { + virtual ~Provider_TensorProto() = default; + + virtual void CopyFrom(const Provider_TensorProto& v) = 0; + + void operator=(const Provider_TensorProto& v) { CopyFrom(v); } +}; + +struct Provider_AttributeProto { + static std::unique_ptr Create(); + + virtual ~Provider_AttributeProto() = default; + virtual std::unique_ptr Clone() const = 0; + + virtual AttributeProto_AttributeType type() const = 0; + virtual int ints_size() const = 0; + virtual int64_t ints(int i) const = 0; + virtual int64_t i() const = 0; + virtual float f() const = 0; + virtual void set_s(const ::std::string& value) = 0; + virtual const ::std::string& s() const = 0; + virtual void set_name(const ::std::string& value) = 0; + virtual void set_type(AttributeProto_AttributeType value) = 0; + virtual Provider_TensorProto* add_tensors() = 0; + + void operator=(const Provider_AttributeProto& v) = delete; +}; + +// This is needed since Provider_NodeAttributes is a map of unique_ptr to Provider_AttributeProto and that won't work since unique_ptrs are not copyable +// (supposedly this should work in the latest C++ STL but it didn't for me so I used this to make it copyable) +struct Provider_AttributeProto_Copyable { + Provider_AttributeProto_Copyable() = default; + Provider_AttributeProto_Copyable(const Provider_AttributeProto_Copyable& copy) : p_{copy->Clone()} {} + + void operator=(std::unique_ptr&& p) { p_ = std::move(p); } + void operator=(const Provider_AttributeProto_Copyable& p) { p_ = p->Clone(); } + + Provider_AttributeProto& operator*() const { return *p_.get(); } + Provider_AttributeProto* operator->() const { return p_.get(); } + + std::unique_ptr p_; +}; + +struct Provider_TensorShapeProto { + int dim_size() const { return dim_size_; } + + int dim_size_{}; +}; + +} // namespace ONNX_NAMESPACE + +namespace onnxruntime { + +struct ProviderHost; +struct Provider_IExecutionProvider; + +struct Provider_IExecutionProviderFactory { + virtual ~Provider_IExecutionProviderFactory() = default; + virtual std::unique_ptr CreateProvider() = 0; +}; + +//struct KernelCreateInfo; + +class DataTypeImpl; +using MLDataType = const DataTypeImpl*; + +struct Provider_OrtDevice { + virtual ~Provider_OrtDevice() {} +}; + +struct Provider_OrtMemoryInfo { + static std::unique_ptr Create(const char* name_, OrtAllocatorType type_, Provider_OrtDevice* device_ = nullptr, int id_ = 0, OrtMemType mem_type_ = OrtMemTypeDefault); + virtual ~Provider_OrtMemoryInfo() {} + + void operator=(const Provider_OrtMemoryInfo& v) = delete; +}; + +template +using Provider_IAllocatorUniquePtr = std::unique_ptr>; + +struct Provider_IAllocator { + virtual ~Provider_IAllocator() {} + + virtual void* Alloc(size_t size) = 0; + virtual void Free(void* p) = 0; + + template + static Provider_IAllocatorUniquePtr MakeUniquePtr(std::shared_ptr allocator, size_t count_or_bytes) { + if (allocator == nullptr) return nullptr; + + size_t alloc_size = count_or_bytes; + + // if T is not void, 'count_or_bytes' == number of items so allow for that + if (!std::is_void::value) { + // TODO: Use internal implementation to get correct sizes + return nullptr; + } + return Provider_IAllocatorUniquePtr{ + static_cast(allocator->Alloc(alloc_size)), // allocate + [=](T* ptr) { allocator->Free(ptr); }}; // capture IAllocator so it's always valid, and use as deleter + } + + void operator=(const Provider_IAllocator& v) = delete; +}; + +struct Provider_IDeviceAllocator : Provider_IAllocator { + virtual bool AllowsArena() const = 0; +}; + +using Provider_AllocatorPtr = std::shared_ptr; +using Provider_DeviceAllocatorFactory = std::function(int)>; + +struct Provider_DeviceAllocatorRegistrationInfo { + OrtMemType mem_type; + Provider_DeviceAllocatorFactory factory; + size_t max_mem; +}; + +class TensorShape; + +struct Provider_Tensor { + virtual float* MutableData_float() = 0; + virtual const float* Data_float() const = 0; + + template + T* MutableData(); + + template + const T* Data() const; + + virtual const TensorShape& Shape() const = 0; +}; + +template <> +inline float* Provider_Tensor::MutableData() { return MutableData_float(); } + +template <> +inline const float* Provider_Tensor::Data() const { return Data_float(); } + +struct Provider_OpKernelInfo { + virtual Status GetAttr(const std::string& name, int64_t* value) const = 0; + virtual Status GetAttr(const std::string& name, float* value) const = 0; + + template + Status GetAttr(const std::string& name, T* value) const; +}; + +template <> +inline Status Provider_OpKernelInfo::GetAttr(const std::string& name, int64_t* value) const { + return GetAttr(name, value); +} + +template <> +inline Status Provider_OpKernelInfo::GetAttr(const std::string& name, float* value) const { + return GetAttr(name, value); +} + +struct Provider_OpKernelContext { + virtual const Provider_Tensor* Input_Tensor(int index) const = 0; + + template + const T* Input(int index) const; + + virtual Provider_Tensor* Output(int index, const TensorShape& shape) = 0; +}; + +template <> +inline const Provider_Tensor* Provider_OpKernelContext::Input(int index) const { + return Input_Tensor(index); +} + +struct Provider_OpKernel { + Provider_OpKernel(const Provider_OpKernelInfo& /*info*/) {} + virtual ~Provider_OpKernel() = default; + + virtual Status Compute(Provider_OpKernelContext* context) const = 0; +}; + +struct Provider_KernelDef { + virtual ~Provider_KernelDef() = default; +}; + +using Provider_KernelCreateFn = std::function; +using Provider_KernelCreatePtrFn = std::add_pointer::type; + +struct Provider_KernelCreateInfo { + std::unique_ptr kernel_def; // Owned and stored in the global kernel registry. + Provider_KernelCreateFn kernel_create_func; + + Provider_KernelCreateInfo(std::unique_ptr definition, + Provider_KernelCreateFn create_func) + : kernel_def(std::move(definition)), + kernel_create_func(create_func) {} + + Provider_KernelCreateInfo(Provider_KernelCreateInfo&& other) noexcept + : kernel_def(std::move(other.kernel_def)), + kernel_create_func(std::move(other.kernel_create_func)) {} +}; + +using Provider_BuildKernelCreateInfoFn = Provider_KernelCreateInfo (*)(); + +struct Provider_KernelDefBuilder { + static std::unique_ptr Create(); + + virtual ~Provider_KernelDefBuilder() = default; + virtual Provider_KernelDefBuilder& SetName(const char* op_name) = 0; + virtual Provider_KernelDefBuilder& SetDomain(const char* domain) = 0; + virtual Provider_KernelDefBuilder& SinceVersion(int since_version) = 0; + virtual Provider_KernelDefBuilder& Provider(const char* provider_type) = 0; + virtual Provider_KernelDefBuilder& TypeConstraint(const char* arg_name, MLDataType supported_type) = 0; + + virtual std::unique_ptr Build() = 0; + + void operator=(const Provider_KernelDefBuilder& v) = delete; +}; + +using NodeIndex = size_t; +using Provider_NodeAttributes = std::unordered_map; + +using Provider_InitializedTensorSet = std::unordered_map; + +struct Provider_NodeArg { + virtual ~Provider_NodeArg() = default; + virtual const std::string& Name() const noexcept = 0; + virtual const ONNX_NAMESPACE::Provider_TensorShapeProto* Shape() const = 0; + virtual ONNX_NAMESPACE::DataType Type() const noexcept = 0; + + void operator=(const Provider_NodeArg& v) = delete; +}; + +struct Provider_Node { + virtual ~Provider_Node() = default; + + virtual const std::string& OpType() const noexcept = 0; + + virtual ConstPointerContainer> InputDefs() const noexcept = 0; + virtual ConstPointerContainer> OutputDefs() const noexcept = 0; + virtual NodeIndex Index() const noexcept = 0; + + virtual const Provider_NodeAttributes& GetAttributes() const noexcept = 0; + virtual size_t GetInputEdgesCount() const noexcept = 0; + virtual size_t GetOutputEdgesCount() const noexcept = 0; + + struct Provider_NodeIterator { + virtual ~Provider_NodeIterator() {} + virtual bool operator!=(const Provider_NodeIterator& p) const = 0; + + virtual void operator++() = 0; + virtual const Provider_Node& operator*() = 0; + }; + + struct NodeConstIterator { + NodeConstIterator(std::unique_ptr p) : impl_{std::move(p)} {} + + bool operator==(const NodeConstIterator& p_other) const; + bool operator!=(const NodeConstIterator& p_other) const { + return *impl_ != *p_other.impl_; + } + + void operator++() { + impl_->operator++(); + } + void operator--(); + + const Provider_Node& operator*() const { + return impl_->operator*(); + } + const Provider_Node* operator->() const; + + std::unique_ptr impl_; + }; + + NodeConstIterator InputNodesBegin() const noexcept { return NodeConstIterator(InputNodesBegin_internal()); } + NodeConstIterator InputNodesEnd() const noexcept { return NodeConstIterator(InputNodesEnd_internal()); } + + virtual std::unique_ptr InputNodesBegin_internal() const noexcept = 0; + virtual std::unique_ptr InputNodesEnd_internal() const noexcept = 0; +}; + +#ifndef PROVIDER_BRIDGE_ORT +// TODO: These are from execution_provider.h and should be factored out in the future into a common header +using CreateFunctionStateFunc = std::function; +using ComputeFunc = std::function; +using DestroyFunctionStateFunc = std::function; + +struct NodeComputeInfo { + CreateFunctionStateFunc create_state_func; + ComputeFunc compute_func; + DestroyFunctionStateFunc release_state_func; +}; +#endif + +struct Provider_GraphViewer { + virtual ~Provider_GraphViewer() = default; + virtual const std::string& Name() const noexcept = 0; + + virtual const Provider_Node* GetNode(NodeIndex node_index) const = 0; + + virtual int MaxNodeIndex() const noexcept = 0; + + virtual const Provider_InitializedTensorSet& GetAllInitializedTensors() const noexcept = 0; + + virtual const std::unordered_map& DomainToVersionMap() const noexcept = 0; + + void operator=(const Provider_GraphViewer& v) = delete; +}; + +struct Provider_IndexedSubGraph { + static std::unique_ptr Create(); + virtual ~Provider_IndexedSubGraph() = default; + + struct MetaDef { + std::string name; ///< Name of customized SubGraph/FunctionProto + std::string domain; ///< Domain of customized SubGraph/FunctionProto + int since_version; ///< Since version of customized SubGraph/FunctionProto. + + ONNX_NAMESPACE::OperatorStatus status; ///< Status of customized SubGraph/FunctionProto. + + std::vector inputs; ///< Inputs of customized SubGraph/FunctionProto. + std::vector outputs; ///< Outputs of customized SubGraph/FunctionProto. + Provider_NodeAttributes attributes; ///< Attributes of customized SubGraph/FunctionProto. + + std::string doc_string; ///< Doc string of customized SubGraph/FunctionProto. + }; + + /** Nodes covered by this subgraph. The NodeIndex values are from the parent Graph.*/ + virtual std::vector& Nodes() = 0; + + virtual void SetMetaDef(std::unique_ptr& meta_def_) = 0; + + void operator=(const Provider_IndexedSubGraph& v) = delete; +}; + +struct Provider_KernelRegistry { + static std::shared_ptr Create(); + + virtual ~Provider_KernelRegistry() = default; + virtual Status Register(Provider_KernelCreateInfo&& create_info) = 0; + + void operator=(const Provider_KernelRegistry& v) = delete; +}; + +struct Provider_ComputeCapability { + Provider_ComputeCapability(std::unique_ptr t_sub_graph) : t_sub_graph_{std::move(t_sub_graph)} {} + + std::unique_ptr t_sub_graph_; + + void operator=(const Provider_ComputeCapability& v) = delete; +}; + +// Provides the base class implementations, since Provider_IExecutionProvider is just an interface. This is to fake the C++ inheritance used by internal IExecutionProvider implementations +struct Provider_IExecutionProvider_Router { + virtual ~Provider_IExecutionProvider_Router() {} + + virtual std::shared_ptr Provider_GetKernelRegistry() const = 0; + + virtual std::vector> Provider_GetCapability(const onnxruntime::Provider_GraphViewer& graph, + const std::vector& kernel_registries) const = 0; + + virtual Provider_AllocatorPtr Provider_GetAllocator(int id, OrtMemType mem_type) const = 0; + virtual void Provider_InsertAllocator(Provider_AllocatorPtr allocator) = 0; + + void operator=(const Provider_IExecutionProvider_Router& v) = delete; +}; + +struct Provider_IExecutionProvider { + Provider_IExecutionProvider(const std::string& type); + virtual ~Provider_IExecutionProvider() {} + + virtual std::shared_ptr Provider_GetKernelRegistry() const { return p_->Provider_GetKernelRegistry(); } + + virtual std::vector> Provider_GetCapability(const onnxruntime::Provider_GraphViewer& graph, + const std::vector& kernel_registries) const { return p_->Provider_GetCapability(graph, kernel_registries); } + + virtual common::Status Provider_Compile(const std::vector& fused_nodes, std::vector& node_compute_funcs) = 0; + + virtual Provider_AllocatorPtr Provider_GetAllocator(int id, OrtMemType mem_type) const { return p_->Provider_GetAllocator(id, mem_type); } + virtual void Provider_InsertAllocator(Provider_AllocatorPtr allocator) { return p_->Provider_InsertAllocator(allocator); } + + Provider_IExecutionProvider_Router* p_; + + void operator=(const Provider_IExecutionProvider& v) = delete; +}; + +namespace logging { +class Logger; +} + +struct Provider { + virtual std::shared_ptr CreateExecutionProviderFactory(int device_id) = 0; + virtual void SetProviderHost(ProviderHost& host) = 0; +}; + +// There are two ways to route a function, one is a virtual method and the other is a function pointer (or pointer to member function) +// The function pointers are nicer in that they directly call the target function, but they cannot be used in cases where we're calling +// a specific implementation of a virtual class member. Trying to get a pointer to member of a virtual function will return a thunk that +// calls the virtual function (which will lead to infinite recursion in the bridge). There is no known way to get the non virtual member +// function pointer implementation in this case. +struct ProviderHost { + virtual Provider_AllocatorPtr CreateAllocator(Provider_DeviceAllocatorRegistrationInfo& info, int16_t device_id = 0) = 0; + + virtual logging::Logger* LoggingManager_GetDefaultLogger() = 0; + + virtual std::unique_ptr AttributeProto_Create() = 0; + + virtual std::unique_ptr OrtMemoryInfo_Create(const char* name_, OrtAllocatorType type_, Provider_OrtDevice* device_, int id_, OrtMemType mem_type_) = 0; + virtual std::unique_ptr KernelDefBuilder_Create() = 0; + + virtual std::shared_ptr KernelRegistry_Create() = 0; + + virtual std::unique_ptr IndexedSubGraph_Create() = 0; + + virtual std::unique_ptr CreateCPUAllocator(std::unique_ptr memory_info) = 0; + virtual Provider_AllocatorPtr CreateDummyArenaAllocator(std::unique_ptr resource_allocator) = 0; + virtual std::unique_ptr Create_IExecutionProvider_Router(Provider_IExecutionProvider* outer, const std::string& type) = 0; + + MLDataType (*DataTypeImpl_GetType_Tensor)(); + MLDataType (*DataTypeImpl_GetType_float)(); + MLDataType (*DataTypeImpl_GetTensorType_float)(); + + virtual void* HeapAllocate(size_t size) = 0; + virtual void HeapFree(void*) = 0; + + virtual void LogRuntimeError(uint32_t session_id, const common::Status& status, const char* file, const char* function, uint32_t line) = 0; + + virtual bool CPU_HasAVX2() = 0; + virtual bool CPU_HasAVX512f() = 0; +}; + +} // namespace onnxruntime diff --git a/onnxruntime/python/onnxruntime_pybind_state.cc b/onnxruntime/python/onnxruntime_pybind_state.cc index a67cd8aa76..bafe3abb71 100644 --- a/onnxruntime/python/onnxruntime_pybind_state.cc +++ b/onnxruntime/python/onnxruntime_pybind_state.cc @@ -35,7 +35,6 @@ #if USE_DNNL #define BACKEND_DNNL "-DNNL" -#include "core/providers/dnnl/dnnl_execution_provider.h" #else #define BACKEND_DNNL "" #endif @@ -54,26 +53,26 @@ #endif #ifdef USE_OPENVINO - #if OPENVINO_CONFIG_CPU_FP32 - #define BACKEND_OPENVINO "-OPENVINO_CPU_FP32" +#if OPENVINO_CONFIG_CPU_FP32 +#define BACKEND_OPENVINO "-OPENVINO_CPU_FP32" - #elif OPENVINO_CONFIG_GPU_FP32 - #define BACKEND_OPENVINO "-OPENVINO_GPU_FP32" +#elif OPENVINO_CONFIG_GPU_FP32 +#define BACKEND_OPENVINO "-OPENVINO_GPU_FP32" - #elif OPENVINO_CONFIG_GPU_FP16 - #define BACKEND_OPENVINO "-OPENVINO_GPU_FP16" +#elif OPENVINO_CONFIG_GPU_FP16 +#define BACKEND_OPENVINO "-OPENVINO_GPU_FP16" - #elif OPENVINO_CONFIG_MYRIAD - #define BACKEND_OPENVINO "-OPENVINO_MYRIAD" +#elif OPENVINO_CONFIG_MYRIAD +#define BACKEND_OPENVINO "-OPENVINO_MYRIAD" - #elif OPENVINO_CONFIG_VAD_M - #define BACKEND_OPENVINO "-OPENVINO_VAD_M" +#elif OPENVINO_CONFIG_VAD_M +#define BACKEND_OPENVINO "-OPENVINO_VAD_M" - #elif OPENVINO_CONFIG_VAD_F - #define BACKEND_OPENVINO "-OPENVINO_VAD_F" - #endif +#elif OPENVINO_CONFIG_VAD_F +#define BACKEND_OPENVINO "-OPENVINO_VAD_F" +#endif #else - #define BACKEND_OPENVINO "" +#define BACKEND_OPENVINO "" #endif #ifdef USE_NUPHAR @@ -103,9 +102,6 @@ onnxruntime::ArenaExtendStrategy arena_extend_strategy = onnxruntime::ArenaExten #ifdef USE_TENSORRT #include "core/providers/tensorrt/tensorrt_provider_factory.h" #endif -#ifdef USE_DNNL -#include "core/providers/dnnl/dnnl_provider_factory.h" -#endif #ifdef USE_NGRAPH #include "core/providers/ngraph/ngraph_provider_factory.h" #endif @@ -384,14 +380,14 @@ void addGlobalMethods(py::module& m, const Environment& env) { #endif #ifdef USE_OPENVINO - m.def("set_openvino_device", [](const std::string& device) { - openvino_device = device;} , - "Set the prefered OpenVINO device(s) to be used. If left unset, all available devices will be used." - ); - m.def("get_openvino_device", []() -> std::string { - return openvino_device; - }, "" - ); + m.def( + "set_openvino_device", [](const std::string& device) { openvino_device = device; }, + "Set the prefered OpenVINO device(s) to be used. If left unset, all available devices will be used."); + m.def( + "get_openvino_device", []() -> std::string { + return openvino_device; + }, + ""); #endif #ifdef onnxruntime_PYBIND_EXPORT_OPSCHEMA @@ -416,7 +412,7 @@ void addGlobalMethods(py::module& m, const Environment& env) { onnxruntime::CreateExecutionProviderFactory_NGraph("CPU"), #endif #ifdef USE_OPENVINO - onnxruntime::CreateExecutionProviderFactory_OpenVINO(openvino_device), + onnxruntime::CreateExecutionProviderFactory_OpenVINO(openvino_device), #endif #ifdef USE_TENSORRT onnxruntime::CreateExecutionProviderFactory_Tensorrt(0) @@ -647,18 +643,10 @@ Set this option to false if you don't want it. Default is True.)pbdoc") .def_readwrite("log_verbosity_level", &SessionOptions::session_log_verbosity_level, R"pbdoc(VLOG level if DEBUG build and session_log_verbosity_level is 0. Applies to session load, initialization, etc. Default is 0.)pbdoc") - .def_property( - "intra_op_num_threads", [](const SessionOptions* options) -> int { - return options->intra_op_param.thread_pool_size; - }, [](SessionOptions* options, int value) -> void { - options->intra_op_param.thread_pool_size = value; - },R"pbdoc(Sets the number of threads used to parallelize the execution within nodes. Default is 0 to let onnxruntime choose.)pbdoc") - .def_property( - "inter_op_num_threads", [](const SessionOptions* options) -> int { - return options->inter_op_param.thread_pool_size; - }, [](SessionOptions* options, int value) -> void { - options->inter_op_param.thread_pool_size = value; - },R"pbdoc(Sets the number of threads used to parallelize the execution of the graph (across nodes). Default is 0 to let onnxruntime choose.)pbdoc") + .def_property( + "intra_op_num_threads", [](const SessionOptions* options) -> int { return options->intra_op_param.thread_pool_size; }, [](SessionOptions* options, int value) -> void { options->intra_op_param.thread_pool_size = value; }, R"pbdoc(Sets the number of threads used to parallelize the execution within nodes. Default is 0 to let onnxruntime choose.)pbdoc") + .def_property( + "inter_op_num_threads", [](const SessionOptions* options) -> int { return options->inter_op_param.thread_pool_size; }, [](SessionOptions* options, int value) -> void { options->inter_op_param.thread_pool_size = value; }, R"pbdoc(Sets the number of threads used to parallelize the execution of the graph (across nodes). Default is 0 to let onnxruntime choose.)pbdoc") .def_readwrite("execution_mode", &SessionOptions::execution_mode, R"pbdoc(Sets the execution mode. Default is sequential.)pbdoc") .def_property( @@ -739,55 +727,57 @@ including arg name, arg type (contains both type and shape).)pbdoc") return *(na.Type()); }, "node type") - .def("__str__", [](const onnxruntime::NodeArg& na) -> std::string { - std::ostringstream res; - res << "NodeArg(name='" << na.Name() << "', type='" << *(na.Type()) << "', shape="; - auto shape = na.Shape(); - std::vector arr; - if (shape == nullptr || shape->dim_size() == 0) { - res << "[]"; - } else { - res << "["; - for (int i = 0; i < shape->dim_size(); ++i) { - if (utils::HasDimValue(shape->dim(i))) { - res << shape->dim(i).dim_value(); - } else if (utils::HasDimParam(shape->dim(i))) { - res << "'" << shape->dim(i).dim_param() << "'"; + .def( + "__str__", [](const onnxruntime::NodeArg& na) -> std::string { + std::ostringstream res; + res << "NodeArg(name='" << na.Name() << "', type='" << *(na.Type()) << "', shape="; + auto shape = na.Shape(); + std::vector arr; + if (shape == nullptr || shape->dim_size() == 0) { + res << "[]"; } else { - res << "None"; + res << "["; + for (int i = 0; i < shape->dim_size(); ++i) { + if (utils::HasDimValue(shape->dim(i))) { + res << shape->dim(i).dim_value(); + } else if (utils::HasDimParam(shape->dim(i))) { + res << "'" << shape->dim(i).dim_param() << "'"; + } else { + res << "None"; + } + + if (i < shape->dim_size() - 1) { + res << ", "; + } + } + res << "]"; + } + res << ")"; + + return std::string(res.str()); + }, + "converts the node into a readable string") + .def_property_readonly( + "shape", [](const onnxruntime::NodeArg& na) -> std::vector { + auto shape = na.Shape(); + std::vector arr; + if (shape == nullptr || shape->dim_size() == 0) { + return arr; } - if (i < shape->dim_size() - 1) { - res << ", "; + arr.resize(shape->dim_size()); + for (int i = 0; i < shape->dim_size(); ++i) { + if (utils::HasDimValue(shape->dim(i))) { + arr[i] = py::cast(shape->dim(i).dim_value()); + } else if (utils::HasDimParam(shape->dim(i))) { + arr[i] = py::cast(shape->dim(i).dim_param()); + } else { + arr[i] = py::none(); + } } - } - res << "]"; - } - res << ")"; - - return std::string(res.str()); - }, - "converts the node into a readable string") - .def_property_readonly("shape", [](const onnxruntime::NodeArg& na) -> std::vector { - auto shape = na.Shape(); - std::vector arr; - if (shape == nullptr || shape->dim_size() == 0) { - return arr; - } - - arr.resize(shape->dim_size()); - for (int i = 0; i < shape->dim_size(); ++i) { - if (utils::HasDimValue(shape->dim(i))) { - arr[i] = py::cast(shape->dim(i).dim_value()); - } else if (utils::HasDimParam(shape->dim(i))) { - arr[i] = py::cast(shape->dim(i).dim_param()); - } else { - arr[i] = py::none(); - } - } - return arr; - }, - "node shape (assuming the node holds a tensor)"); + return arr; + }, + "node shape (assuming the node holds a tensor)"); py::class_(m, "SessionObjectInitializer"); py::class_(m, "InferenceSession", R"pbdoc(This is the main class used to run a model.)pbdoc") @@ -988,7 +978,7 @@ PYBIND11_MODULE(onnxruntime_pybind11_state, m) { // static variable used to create inference session and training session. static std::unique_ptr session_env; -void initialize_env(){ +void initialize_env() { auto initialize = [&]() { // Initialization of the module ([]() -> void { @@ -1011,8 +1001,8 @@ void initialize_env(){ initialize(); } -onnxruntime::Environment& get_env(){ - if (!session_env){ +onnxruntime::Environment& get_env() { + if (!session_env) { initialize_env(); } return *session_env; diff --git a/onnxruntime/test/onnx/main.cc b/onnxruntime/test/onnx/main.cc index 1a81e88fbe..3baa1dc51b 100644 --- a/onnxruntime/test/onnx/main.cc +++ b/onnxruntime/test/onnx/main.cc @@ -855,5 +855,6 @@ int main(int argc, char* argv[]) { retval = -1; } ::google::protobuf::ShutdownProtobufLibrary(); + std::cout << "*** Exiting Test Runner\r\n"; return retval; } diff --git a/onnxruntime/test/providers/dnnl/dnnl_execution_provider_test.cc b/onnxruntime/test/providers/dnnl/dnnl_execution_provider_test.cc index 69eb6b7844..0f88af2b84 100644 --- a/onnxruntime/test/providers/dnnl/dnnl_execution_provider_test.cc +++ b/onnxruntime/test/providers/dnnl/dnnl_execution_provider_test.cc @@ -1,17 +1,19 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "core/providers/dnnl/dnnl_execution_provider.h" +//#include "core/providers/dnnl/dnnl_execution_provider.h" #include "gtest/gtest.h" namespace onnxruntime { namespace test { TEST(DNNLExecutionProviderTest, MetadataTest) { + #if 0 // With DNNL as a DLL this can't be tested here TODO(pranav) DNNLExecutionProviderInfo info; info.create_arena = false; auto provider = onnxruntime::make_unique(info); EXPECT_TRUE(provider != nullptr); ASSERT_STREQ(provider->GetAllocator(0, OrtMemTypeCPUOutput)->Info().name, "DnnlCpu"); + #endif } } // namespace test } // namespace onnxruntime