diff --git a/cmake/external/mkldnn.cmake b/cmake/external/mkldnn.cmake index e3a638cc7a..d2ccdd04d2 100644 --- a/cmake/external/mkldnn.cmake +++ b/cmake/external/mkldnn.cmake @@ -2,7 +2,7 @@ include (ExternalProject) set(MKLDNN_URL https://github.com/intel/mkl-dnn.git) # If MKLDNN_TAG is updated, check if MKLML_VERSION and platform.cmake.patch need to be updated. -set(MKLDNN_TAG v0.18.1) +set(MKLDNN_TAG v1.0.2) set(MKLML_VERSION 2019.0.5.20190502) if(WIN32) @@ -21,10 +21,10 @@ if(WIN32) else() set(MKLML_FILE_EXTENSION "tgz") if (APPLE) - set(MKLDNN_SHARED_LIB libmkldnn.0.dylib) + set(MKLDNN_SHARED_LIB libmkldnn.1.dylib) set(MKLML_OS_VERSION_STR "mac") else() - set(MKLDNN_SHARED_LIB libmkldnn.so.0) + set(MKLDNN_SHARED_LIB libmkldnn.so.1) set(MKLML_OS_VERSION_STR "lnx") endif() if(onnxruntime_USE_MKLML) @@ -62,7 +62,7 @@ if (onnxruntime_USE_MKLDNN) endif() set(MKLDNN_INCLUDE_DIR ${MKLDNN_INSTALL}/include) set(MKLDNN_CMAKE_EXTRA_ARGS) - set(MKLDNN_PATCH_COMMAND1 git apply ${CMAKE_SOURCE_DIR}/patches/mkldnn/mem-patch.cmake.patch) + set(MKLDNN_PATCH_COMMAND git apply ${CMAKE_SOURCE_DIR}/patches/mkldnn/constexpr.patch) # discard prior changes due to patching in mkldnn source to unblock incremental builds. set(MKLDNN_PATCH_DISCARD_COMMAND cd ${MKLDNN_SOURCE} && git checkout -- .) if(NOT onnxruntime_BUILD_FOR_NATIVE_MACHINE) @@ -75,7 +75,7 @@ if (onnxruntime_USE_MKLDNN) PREFIX mkl-dnn GIT_REPOSITORY ${MKLDNN_URL} GIT_TAG ${MKLDNN_TAG} - PATCH_COMMAND ${MKLDNN_PATCH_DISCARD_COMMAND} COMMAND ${MKLDNN_PATCH_COMMAND1} + PATCH_COMMAND ${MKLDNN_PATCH_DISCARD_COMMAND} COMMAND ${MKLDNN_PATCH_COMMAND} SOURCE_DIR ${MKLDNN_SOURCE} CMAKE_ARGS -DMKLDNN_PRODUCT_BUILD_MODE=OFF -DCMAKE_BUILD_TYPE=${CMAKE_BUILD_TYPE} -DCMAKE_INSTALL_PREFIX=${MKLDNN_INSTALL} -DMKLROOT=${MKML_DIR} ${MKLDNN_CMAKE_EXTRA_ARGS} ) diff --git a/cmake/patches/mkldnn/constexpr.patch b/cmake/patches/mkldnn/constexpr.patch new file mode 100644 index 0000000000..8b57c95c65 --- /dev/null +++ b/cmake/patches/mkldnn/constexpr.patch @@ -0,0 +1,43 @@ +diff --git a/src/cpu/rnn/rnn_reorders.hpp b/src/cpu/rnn/rnn_reorders.hpp +index fbfedd01..28a6dc1d 100644 +--- a/src/cpu/rnn/rnn_reorders.hpp ++++ b/src/cpu/rnn/rnn_reorders.hpp +@@ -128,7 +128,7 @@ struct rnn_weights_reorder_t : public cpu_primitive_t { + return status::success; + } + +- format_tag_t itag_; ++ format_tag_t itag_ = mkldnn_format_tag_undef; + + private: + void init_scratchpad() { +diff --git a/src/cpu/simple_concat.hpp b/src/cpu/simple_concat.hpp +index fe84eb14..f1b2bd17 100644 +--- a/src/cpu/simple_concat.hpp ++++ b/src/cpu/simple_concat.hpp +@@ -98,9 +98,9 @@ struct simple_concat_t: public cpu_primitive_t { + return status::success; + } + +- int perm_[MKLDNN_MAX_NDIMS]; +- int iperm_[MKLDNN_MAX_NDIMS]; +- dims_t blocks_; ++ int perm_[MKLDNN_MAX_NDIMS] {}; ++ int iperm_[MKLDNN_MAX_NDIMS] {}; ++ dims_t blocks_ {}; + + dim_t nelems_to_concat(const memory_desc_wrapper &data_d) const { + const int ndims = data_d.ndims(); +diff --git a/src/cpu/simple_sum.hpp b/src/cpu/simple_sum.hpp +index 46a24810..c47c4a64 100644 +--- a/src/cpu/simple_sum.hpp ++++ b/src/cpu/simple_sum.hpp +@@ -71,7 +71,7 @@ struct simple_sum_t: public cpu_primitive_t { + } + + sum_bf16_params_t bf16_p_; +- dim_t block_size_, nelems_, blocks_number_, tail_; ++ dim_t block_size_ = 0, nelems_= 0, blocks_number_= 0, tail_= 0; + + private: + const dim_t cacheline_size_ = 64; // bytes diff --git a/cmake/patches/mkldnn/mem-patch.cmake.patch b/cmake/patches/mkldnn/mem-patch.cmake.patch deleted file mode 100644 index 9d0a2543e6..0000000000 --- a/cmake/patches/mkldnn/mem-patch.cmake.patch +++ /dev/null @@ -1,107 +0,0 @@ - ---- - src/cpu/jit_avx2_1x1_convolution.cpp | 6 +++--- - src/cpu/jit_avx512_common_1x1_convolution.cpp | 9 ++++----- - src/cpu/jit_avx512_core_x8s8s32x_1x1_convolution.cpp | 6 ++++-- - src/cpu/jit_uni_1x1_conv_utils.hpp | 3 ++- - 4 files changed, 13 insertions(+), 11 deletions(-) - -diff --git a/src/cpu/jit_avx2_1x1_convolution.cpp b/src/cpu/jit_avx2_1x1_convolution.cpp -index 46362886..edb2b6fb 100644 ---- a/src/cpu/jit_avx2_1x1_convolution.cpp -+++ b/src/cpu/jit_avx2_1x1_convolution.cpp -@@ -50,7 +50,7 @@ void jit_avx2_1x1_convolution_fwd_t::execute_forward() const { - const memory_desc_wrapper weights_d(pd()->weights_pd(0)); - - const auto &jcp = kernel_->jcp; -- auto rtus_space = scratchpad().get(key_conv_rtus_space); -+ auto rtus_space = pd()->rtus_.reduce_src_?scratchpad().get(key_conv_rtus_space):NULL; - - const int work_amount = jcp.mb * jcp.ngroups * jcp.nb_bcast; - const int ndims = dst_d.ndims(); -@@ -180,7 +180,7 @@ void jit_avx2_1x1_convolution_bwd_data_t::execute_backward_data() const { - const memory_desc_wrapper diff_src_d(pd()->diff_src_pd()); - - const auto &jcp = kernel_->jcp; -- auto rtus_space = scratchpad().get(key_conv_rtus_space); -+ auto rtus_space = pd()->rtus_.reduce_src_?scratchpad().get(key_conv_rtus_space):NULL; - - // TODO (Roma): remove this restriction - assert(jcp.stride_w == 1 && jcp.stride_h == 1); -@@ -306,7 +306,7 @@ void jit_avx2_1x1_convolution_bwd_weights_t::execute_backward_weights() const { - const memory_desc_wrapper diff_bias_d(pd()->diff_weights_pd(1)); - - const auto &jcp = kernel_->jcp; -- auto rtus_space = scratchpad.get(key_conv_rtus_space); -+ auto rtus_space = pd()->rtus_.reduce_src_?scratchpad.get(key_conv_rtus_space):NULL; - - data_t *diff_bias = pd()->wants_padded_bias() - ? scratchpad.get(key_conv_padded_bias) : diff_bias_in; -diff --git a/src/cpu/jit_avx512_common_1x1_convolution.cpp b/src/cpu/jit_avx512_common_1x1_convolution.cpp -index 6879cd91..6a32aa49 100644 ---- a/src/cpu/jit_avx512_common_1x1_convolution.cpp -+++ b/src/cpu/jit_avx512_common_1x1_convolution.cpp -@@ -106,7 +106,7 @@ execute_forward_thr(const int ithr, const int nthr, const src_data_t *src, - const memory_desc_wrapper weights_d(pd()->weights_pd(0)); - - const auto &jcp = kernel_->jcp; -- auto rtus_space = scratchpad.get(key_conv_rtus_space); -+ auto rtus_space = pd()->rtus_.reduce_src_?scratchpad.get(key_conv_rtus_space):NULL; - - const int ndims = src_d.ndims(); - const int stride_h = (ndims == 3) ? 1 : pd()->desc()->strides[0]; -@@ -301,9 +301,8 @@ void jit_avx512_common_1x1_convolution_bwd_data_tdiff_src_pd()); - - const auto &jcp = kernel_->jcp; -- auto rtus_space = scratchpad().template get( -- key_conv_rtus_space); -- -+ auto rtus_space = pd()->rtus_.reduce_src_? scratchpad().template get(key_conv_rtus_space): NULL; -+ - const int ndims = diff_src_d.ndims(); - - // TODO (Roma): remove this restriction -@@ -470,7 +469,7 @@ void jit_avx512_common_1x1_convolution_bwd_weights_t::execute_backward_weights() - - const auto scratchpad = this->scratchpad(); - -- auto rtus_space = scratchpad.get(key_conv_rtus_space); -+ auto rtus_space = pd()->rtus_.reduce_src_?scratchpad.get(key_conv_rtus_space):NULL; - data_t *diff_bias = pd()->wants_padded_bias() - ? scratchpad.get(key_conv_padded_bias) : diff_bias_in; - auto wei_reduction = scratchpad.get(key_conv_wei_reduction); -diff --git a/src/cpu/jit_avx512_core_x8s8s32x_1x1_convolution.cpp b/src/cpu/jit_avx512_core_x8s8s32x_1x1_convolution.cpp -index de303cd2..ec0c54e7 100644 ---- a/src/cpu/jit_avx512_core_x8s8s32x_1x1_convolution.cpp -+++ b/src/cpu/jit_avx512_core_x8s8s32x_1x1_convolution.cpp -@@ -100,8 +100,10 @@ void jit_avx512_core_x8s8s32x_1x1_convolution_fwd_t - ? types::data_type_size(pd()->desc()->bias_desc.data_type) : 0; - - const auto &jcp = kernel_->jcp; -- auto rtus_space = scratchpad.get(key_conv_rtus_space); -- auto local_scales = scratchpad.get(key_conv_adjusted_scales); -+ -+ auto rtus_space = pd()->rtus_.reduce_src_?scratchpad.get(key_conv_rtus_space):NULL; -+ -+ auto local_scales = scratchpad.get(key_conv_adjusted_scales); - - const int work_amount = jcp.mb * jcp.ngroups * jcp.nb_bcast; - -diff --git a/src/cpu/jit_uni_1x1_conv_utils.hpp b/src/cpu/jit_uni_1x1_conv_utils.hpp -index a3ed769a..5a0e0635 100644 ---- a/src/cpu/jit_uni_1x1_conv_utils.hpp -+++ b/src/cpu/jit_uni_1x1_conv_utils.hpp -@@ -94,7 +94,8 @@ inline void rtus_prepare(conv_pd_t *self, const convolution_desc_t *&conv_d, - template - inline void rtus_prepare_space_info(conv_pd_t *self, - memory_tracking::registrar_t &scratchpad) { -- const auto &jcp = self->jcp_; -+ if (!self->rtus_.reduce_src_) return; -+ const auto &jcp = self->jcp_; - - const int max_threads = mkldnn_get_max_threads(); - const size_t factor = utils::pick_by_prop_kind(self->desc()->prop_kind, --- -2.17.0.windows.1 - diff --git a/onnxruntime/core/providers/mkldnn/activation/activations.cc b/onnxruntime/core/providers/mkldnn/activation/activations.cc deleted file mode 100644 index 1382c73a8b..0000000000 --- a/onnxruntime/core/providers/mkldnn/activation/activations.cc +++ /dev/null @@ -1,195 +0,0 @@ -// Copyright(C) 2018 Intel Corporation -// Licensed under the MIT License - -#ifdef _WIN32 -#pragma warning(disable : 4244) -#endif - -#include "core/providers/mkldnn/mkldnn_common.h" -#include "core/providers/mkldnn/activation/activations.h" -#include "core/providers/mkldnn/mkldnn_fwd.h" - -namespace onnxruntime { -namespace mkl_dnn { - -namespace { -// Struct which encapsulates parameters for MKLDNN Pool primitive. -struct ReluParams { - const mkldnn::memory::dims& src_dims; - const mkldnn::memory::dims& dst_dims; - - ReluParams(const mkldnn::memory::dims& src_dims, const mkldnn::memory::dims& dst_dims) - : src_dims(src_dims), - dst_dims(dst_dims) {} - - // Used as the key for Pool Primitive Reuse Pool. - std::string ToString() const { - std::string key; - key.reserve(64); - key.append("Relu_"); - AddDimsToKey(key, src_dims); - AddDimsToKey(key, dst_dims); - return key; - } -}; - -template -class ReluPrimitive final : public PrimitiveBase { - public: - explicit ReluPrimitive(const ReluParams& params) - : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); - if (context_.relu_fwd == nullptr) { - Initialize(params); - } - } - - ~ReluPrimitive() = default; - - void Compute(const T* src_data, T* dst_data) { - context_.src_mem->set_data_handle( - static_cast(const_cast(src_data))); - context_.dst_mem->set_data_handle( - static_cast(dst_data)); - context_.stream->submit(context_.net); - - context_.src_mem->set_data_handle(nullptr); - context_.dst_mem->set_data_handle(nullptr); - return; - } - - private: - struct ReluContext { - mkldnn::memory::format src_fmt; - - std::unique_ptr src_mem; - std::unique_ptr dst_mem; - - size_t src_size; - size_t dst_size; - - std::unique_ptr fwd_desc; - std::unique_ptr relu_fwd_pd; - std::unique_ptr relu_fwd; - - std::unique_ptr src_md; - std::unique_ptr dst_md; - - std::unique_ptr stream; - std::vector net; - }; - - void Initialize(const ReluParams& params) { - - mkldnn::memory::format fmt = mkldnn::memory::format::any; - switch (params.src_dims.size()) { - case 1: { fmt = mkldnn::memory::format::x; break; } - case 2: { fmt = mkldnn::memory::format::nc; break; } - case 3: { fmt = mkldnn::memory::format::ntc; break; } - case 4: { fmt = mkldnn::memory::format::nchw; break; } - case 5: { fmt = mkldnn::memory::format::ncdhw; break; } - default: { fmt = mkldnn::memory::format::any; break; } - } - - context_.src_md.reset(new mkldnn::memory::desc({ params.src_dims}, MklDnnType(), fmt)); - - mkldnn::algorithm algo = mkldnn::algorithm::eltwise_relu; - context_.fwd_desc.reset(new mkldnn::eltwise_forward::desc( - mkldnn::prop_kind::forward_inference, algo, *context_.src_md, 0)); - - context_.relu_fwd_pd.reset(new mkldnn::eltwise_forward::primitive_desc( - *context_.fwd_desc, cpu_engine_)); - - context_.src_fmt = static_cast( - context_.relu_fwd_pd.get()->src_primitive_desc().desc().data.format); - - context_.src_size = context_.relu_fwd_pd.get()->src_primitive_desc().get_size(); - context_.dst_size = context_.relu_fwd_pd.get()->dst_primitive_desc().get_size(); - - context_.src_mem.reset(new mkldnn::memory(context_.relu_fwd_pd.get()->src_primitive_desc(), nullptr)); - context_.dst_mem.reset(new mkldnn::memory(context_.relu_fwd_pd.get()->dst_primitive_desc(), nullptr)); - context_.relu_fwd.reset( - new mkldnn::eltwise_forward(*context_.relu_fwd_pd, *context_.src_mem, *context_.dst_mem)); - context_.net.push_back(*context_.relu_fwd); - } - - ReluContext context_; - mkldnn::engine& cpu_engine_; -}; - -// Pool which allows for reuse of MKLDNN Relu primitives which are expensive -// to instantiate. To address thread safety, the primitives are stored in a map -// on thread local storage. -template -class ReluPrimitivePool : public PrimitivePool { - public: - static ReluPrimitive* Get(const ReluParams& params) { - ReluPrimitive* primitive = dynamic_cast*>( - ReluPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - - if (primitive == nullptr) { - auto relu_primitive = onnxruntime::make_unique>(params); - primitive = relu_primitive.get(); - ReluPrimitivePool::GetInstance().SetPrimitive(params.ToString(), - std::move(relu_primitive)); - } - return primitive; - } - - private: - ReluPrimitivePool() = default; - ~ReluPrimitivePool() = default; - - static ReluPrimitivePool& GetInstance() { - static ReluPrimitivePool pool; - return pool; - } -}; -} // namespace - -template -Status Relu::Compute(OpKernelContext* context) const { - const Tensor* X = context->Input(0); - Tensor* Y = context->Output(0, X->Shape()); - - const TensorShape& x_shape = X->Shape(); - const auto& x_dims = x_shape.GetDims(); - - if (X->Shape().NumDimensions() > 5 ) { - // Fall Back to CPU implementation. - // mkldnn support up to dim of size 5 - return onnxruntime::Relu::Compute(context); - } - - const TensorShape& y_shape = Y->Shape(); - auto& y_dims = y_shape.GetDims(); - - const T* src_data = X->template Data(); - T* dst_data = Y->template MutableData(); - - mkldnn::memory::dims src_dims_mkl(x_dims.begin(), x_dims.end()); - mkldnn::memory::dims dst_dims_mkl(y_dims.begin(), y_dims.end()); - - try { - ReluParams pool_params(src_dims_mkl, dst_dims_mkl); - ReluPrimitive* relu_primitive = ReluPrimitivePool::Get(pool_params); - - relu_primitive->Compute(src_data, dst_data); - } catch (const mkldnn::error& e) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, - ", message: ", e.message.c_str()); - } - - return Status::OK(); -} - -ONNX_OPERATOR_KERNEL_EX( - Relu, - kOnnxDomain, - 6, - kMklDnnExecutionProvider, - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), - Relu); - -} // namespace mkl_dnn -} // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/providers/mkldnn/activation/activations.h b/onnxruntime/core/providers/mkldnn/activation/activations.h deleted file mode 100644 index 21d3a6d5e3..0000000000 --- a/onnxruntime/core/providers/mkldnn/activation/activations.h +++ /dev/null @@ -1,20 +0,0 @@ -// Copyright(C) 2018 Intel Corporation -// Licensed under the MIT License - -#pragma once -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/activation/activations.h" - -namespace onnxruntime { -namespace mkl_dnn { - -template -class Relu : public onnxruntime::Relu { - public: - explicit Relu(const OpKernelInfo& info) : onnxruntime::Relu(info) {} - - Status Compute(OpKernelContext* context) const override; -}; - -} // namespace mkl_dnn -} // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/providers/mkldnn/math/gemm.cc b/onnxruntime/core/providers/mkldnn/math/gemm.cc index 4ea9924940..fbe2ac024b 100644 --- a/onnxruntime/core/providers/mkldnn/math/gemm.cc +++ b/onnxruntime/core/providers/mkldnn/math/gemm.cc @@ -5,6 +5,7 @@ #include "core/providers/cpu/math/gemm_helper.h" #include "core/util/math_cpuonly.h" #include "mkldnn.h" +#include "mkldnn.hpp" #include "core/providers/mkldnn/mkldnn_fwd.h" namespace onnxruntime { @@ -28,11 +29,14 @@ Status Gemm::Compute(OpKernelContext* ctx) const { if (!helper.State().IsOK()) return helper.State(); - int M = gsl::narrow_cast(helper.M()); - int N = gsl::narrow_cast(helper.N()); - int K = gsl::narrow_cast(helper.K()); + mkldnn::memory::dim M = gsl::narrow_cast(helper.M()); + mkldnn::memory::dim N = gsl::narrow_cast(helper.N()); + mkldnn::memory::dim K = gsl::narrow_cast(helper.K()); auto Y = ctx->Output(0, TensorShape({M, N})); + if (M <= 0) + return ORT_MAKE_STATUS(ONNXRUNTIME, NOT_IMPLEMENTED, "Empty Tensor not supported"); + if (beta_ != 0) { auto output_mat = EigenMatrixMapRowMajor( Y->template MutableData(), @@ -77,15 +81,13 @@ Status Gemm::Compute(OpKernelContext* ctx) const { } } - // mkldnn_sgemm expects col major matrices, so we need to swap the operands A and B - auto status = mkldnn_sgemm(trans_B_ ? "T" : "N", - trans_A_ ? "T" : "N", - &N, &M, &K, - &alpha_, W->template Data(), - trans_B_ ? &K : &N, - X->template Data(), - trans_A_ ? &M : &K, - &beta_, Y->template MutableData(), &N); + // mkldnn_sgemm expects row major matrices, so no need to swap the operands A and B + auto status = mkldnn_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); if (status == mkldnn_success) { return Status::OK(); } else { diff --git a/onnxruntime/core/providers/mkldnn/math/sum.cc b/onnxruntime/core/providers/mkldnn/math/sum.cc deleted file mode 100644 index 82c4a92aba..0000000000 --- a/onnxruntime/core/providers/mkldnn/math/sum.cc +++ /dev/null @@ -1,235 +0,0 @@ -// Copyright(C) 2018 Intel Corporation -// Licensed under the MIT License - -#ifdef _WIN32 -#pragma warning(disable : 4244) -#endif - -#include "core/providers/mkldnn/mkldnn_common.h" -#include "core/providers/mkldnn/math/sum.h" -#include "core/providers/mkldnn/mkldnn_fwd.h" - -namespace onnxruntime { -namespace mkl_dnn { - -namespace { -// Struct which encapsulates parameters for MKLDNN Sum primitives. -struct SumParams { - const std::vector& src_dims; - const mkldnn::memory::dims& dst_dim; - const int num_inputs; - const int num_dimensions; - - SumParams(const std::vector& dims, - const mkldnn::memory::dims& dst_dims, int numinputs, - int dimensions) - : src_dims(dims), - dst_dim(dst_dims), - num_inputs(numinputs), - num_dimensions(dimensions) {} - - // Used as the key for Sum Primitive Reuse Sum. - std::string ToString() const { - std::string key; - key.reserve(64); - key.append("sum_"); - for (size_t i = 0; i < src_dims.size(); i++) { - AddDimsToKey(key, src_dims[i]); - } - AddDimsToKey(key, dst_dim); - return key; - } -}; - -template -class SumPrimitive final : public PrimitiveBase { - public: - explicit SumPrimitive(const SumParams& params) - : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); - if (context_.sum_pd == nullptr) { - Initialize(params); - } - } - - ~SumPrimitive() = default; - - void Compute(OpKernelContext* context, int numinputs) { - const Tensor* X1 = context->Input(0); - Tensor* Y = context->Output(0, X1->Shape()); - T* dst_data = Y->template MutableData(); - - context_.dst_mem->set_data_handle( - static_cast(dst_data)); - - for (int i = 0; i < numinputs; i++) { - const Tensor* X = context->Input(i); - const T* src_data = X->template Data(); - context_.srcs_memory[i].set_data_handle( - static_cast(const_cast(src_data))); - } - context_.stream->submit(context_.net); - - for (int i = 0; i < numinputs; i++) { - context_.srcs_memory[i].set_data_handle(nullptr); - } - - } - - std::unique_ptr GetDstMemoryDesc() const { - return context_.dst_md; - } - - std::unique_ptr - GetPrimitiveDesc() const { - return context_.sum_pd; - } - - private: - struct SumContext { - std::unique_ptr src_md; - std::unique_ptr dst_md; - - std::vector srcs_memory; - std::unique_ptr dst_mem; - - std::vector srcs_pd; - std::unique_ptr src_mpd; - std::unique_ptr dst_pd; - std::unique_ptr sum_pd; - - std::unique_ptr stream; - std::vector net; - }; - - void Initialize(const SumParams& params) { - std::vector coeff; - - mkldnn::memory::format fmt = mkldnn::memory::format::any; - switch (params.num_dimensions) { - case 1: { fmt = mkldnn::memory::format::x; break; } - case 2: { fmt = mkldnn::memory::format::nc; break; } - case 3: { fmt = mkldnn::memory::format::ntc; break; } - case 4: { fmt = mkldnn::memory::format::nchw; break; } - case 5: { fmt = mkldnn::memory::format::ncdhw; break; } - default: { fmt = mkldnn::memory::format::any; break; } - } - - for (int i = 0; i < params.num_inputs; i++) { - context_.src_md.reset( - new mkldnn::memory::desc({params.src_dims[i]}, MklDnnType(), fmt)); - auto mpd = mkldnn::memory::primitive_desc(*context_.src_md, cpu_engine_); - auto src_memory = mkldnn::memory(mpd, nullptr); - - context_.srcs_pd.push_back(mpd); - context_.srcs_memory.push_back(src_memory); - coeff.push_back(1.0); - } - - std::unique_ptr dst; - context_.dst_md.reset(new mkldnn::memory::desc( - {params.dst_dim}, MklDnnType(), mkldnn::memory::format::any)); - context_.sum_pd.reset(new mkldnn::sum::primitive_desc( - *context_.dst_md, coeff, context_.srcs_pd)); - context_.dst_mem.reset(new mkldnn::memory( - context_.sum_pd->dst_primitive_desc(), nullptr)); - - std::vector inputs; - for (int i = 0; i < params.num_inputs; i++) { - inputs.push_back(context_.srcs_memory[i]); - } - auto c = mkldnn::sum(*context_.sum_pd, inputs, *context_.dst_mem); - context_.net.push_back(c); - } - - SumContext context_; - mkldnn::engine& cpu_engine_; -}; - -// Pool which allows for reuse of MKLDNN Sum primitives which are -// expensive to instantiate. To address thread safety, the primitives -// are stored in a map on thread local storage. - -template -class SumPrimitivePool : public PrimitivePool { - public: - static SumPrimitive* Get(const SumParams& params) { - SumPrimitive* primitive = dynamic_cast*>( - SumPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - - if (primitive == nullptr) { - auto sum_primitive = onnxruntime::make_unique>(params); - primitive = sum_primitive.get(); - SumPrimitivePool::GetInstance().SetPrimitive( - params.ToString(), std::move(sum_primitive)); - } - return primitive; - } - - private: - SumPrimitivePool() = default; - ~SumPrimitivePool() = default; - - static SumPrimitivePool& GetInstance() { - static SumPrimitivePool pool; - return pool; - } -}; -} // namespace_ - -template -Status Sum::Compute(OpKernelContext* context) const { - int num_inputs = static_cast(OpKernel::Node().InputDefs().size()); - - ORT_ENFORCE(num_inputs > 0, "MKLDNN Sum kernel: Must have at least one input"); - - if (num_inputs == 1) { - // Fall Back to CPU implementation. For one input, use CPU implementation - // to copy input to output - return onnxruntime::Sum_6::Compute(context); - } - - std::vector src_dims; - - const Tensor* X1 = context->Input(0); - if (X1 == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "input count mismatch"); - Tensor* Y = context->Output(0, X1->Shape()); - int dimensions = static_cast(X1->Shape().NumDimensions()); - - const TensorShape& x_shape = X1->Shape(); - const auto& x_dims = x_shape.GetDims(); - mkldnn::memory::dims src_dim(x_dims.begin(), x_dims.end()); - - mkldnn::memory::dims dst_dims_mkl( - Y->Shape().GetDims().begin(), Y->Shape().GetDims().end()); - - for (int i = 0; i < num_inputs; i++) { - const Tensor* X = context->Input(i); - if (X == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "input count mismatch"); - mkldnn::memory::dims src_dims_mkl( - X->Shape().GetDims().begin(), X->Shape().GetDims().end()); - src_dims.push_back(src_dims_mkl); - } - try { - SumParams parameters(src_dims, dst_dims_mkl, num_inputs, dimensions); - SumPrimitive* sum_primitive = SumPrimitivePool::Get(parameters); - ORT_RETURN_IF_NOT(sum_primitive != nullptr); - sum_primitive->Compute(context, num_inputs); - } catch (const mkldnn::error& e) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, - ", message: ", e.message.c_str()); - } - - return Status::OK(); -} - -ONNX_OPERATOR_KERNEL_EX( - Sum, - kOnnxDomain, - 6, - kMklDnnExecutionProvider, - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), - Sum); - -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/math/sum.h b/onnxruntime/core/providers/mkldnn/math/sum.h deleted file mode 100644 index 9cf4aae12b..0000000000 --- a/onnxruntime/core/providers/mkldnn/math/sum.h +++ /dev/null @@ -1,22 +0,0 @@ -// Copyright(C) 2018 Intel Corporation -// Licensed under the MIT License - -#pragma once -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/math/element_wise_ops.h" -#include "../mkldnn_execution_provider.h" - -namespace onnxruntime { -namespace mkl_dnn { - -template -class Sum final : public onnxruntime::Sum_6 { - public: - explicit Sum(const OpKernelInfo& info) : onnxruntime::Sum_6(info) {} - - Status Compute(OpKernelContext* context) const override; - -private: -}; -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/mkldnn_common.h b/onnxruntime/core/providers/mkldnn/mkldnn_common.h index f5e8a20e31..10f4ac8ce8 100644 --- a/onnxruntime/core/providers/mkldnn/mkldnn_common.h +++ b/onnxruntime/core/providers/mkldnn/mkldnn_common.h @@ -5,6 +5,7 @@ #include "core/common/common.h" #include "mkldnn.hpp" #include +#include namespace onnxruntime { namespace mkl_dnn { @@ -19,7 +20,7 @@ mkldnn::memory::data_type MklDnnType() { } static mkldnn::engine& GetEngine() { - static mkldnn::engine cpu_engine = mkldnn::engine(mkldnn::engine::cpu, 0); + static mkldnn::engine cpu_engine = mkldnn::engine(mkldnn::engine::kind::cpu, 0); return cpu_engine; } @@ -68,108 +69,5 @@ class PrimitivePool { return map; } }; - -// Struct which encapsulates parameters for MKLDNN memory reorder primitive. -struct MemoryReorderParams { - const mkldnn::memory& src; - const mkldnn::memory& dst; - - MemoryReorderParams(const mkldnn::memory& src, const mkldnn::memory& dst) : src(src), dst(dst) {} - - // Used as the key for MemoryReorder primitive reuse pool. - std::string ToString() const { - std::string key; - key.reserve(64); - key.append("reorder_"); - const auto& src_desc = src.get_primitive_desc().desc().data; - const auto& dst_desc = dst.get_primitive_desc().desc().data; - mkldnn::memory::dims src_dims(src_desc.dims, &src_desc.dims[src_desc.ndims]); - mkldnn::memory::dims dst_dims(dst_desc.dims, &dst_desc.dims[dst_desc.ndims]); - key.append(std::to_string(src_desc.format)); - key.append(1, '_'); - key.append(std::to_string(dst_desc.data_type)); - AddDimsToKey(key, src_dims); - key.append(std::to_string(dst_desc.format)); - key.append(1, '_'); - key.append(std::to_string(dst_desc.data_type)); - AddDimsToKey(key, dst_dims); - return key; - } -}; - -// Encapsulates an MKLDNN memory reorder primitive. -// these are needed to convert the source/weight/destination memory layout -// to one that is identified by MKLDNN to be optimal for performance. -class MemoryReorderPrimitive : public PrimitiveBase { - public: - explicit MemoryReorderPrimitive(const MemoryReorderParams& params) : cpu_engine_(GetEngine()) { - Initialize(params); - } - ~MemoryReorderPrimitive() = default; - - std::shared_ptr GetPrimitive() { - return context_.primitive; - } - - void SetMemory(const MemoryReorderParams& params) { - context_.src_mem->set_data_handle(params.src.get_data_handle()); - context_.dst_mem->set_data_handle(params.dst.get_data_handle()); - } - - private: - struct MemoryReorderContext { - std::shared_ptr src_mem; - std::shared_ptr dst_mem; - std::shared_ptr primitive; - MemoryReorderContext() : src_mem(nullptr), dst_mem(nullptr), primitive(nullptr) { - } - } context_; - - mkldnn::engine& cpu_engine_; - - void Initialize(const MemoryReorderParams& params) { - context_.src_mem = std::make_shared( - mkldnn::memory({params.src.get_primitive_desc().desc(), cpu_engine_}, nullptr)); - context_.dst_mem = std::make_shared( - mkldnn::memory({params.dst.get_primitive_desc().desc(), cpu_engine_}, nullptr)); - context_.primitive = std::make_shared( - mkldnn::reorder(*context_.src_mem, *context_.dst_mem)); - } -}; - -// Pool which allows for reuse of MKLDNN memory reorder primitives which are expensive to instantiate. -// To address thread safety, the primitives are stored in a map on thread local storage. -template -class MemoryReorderPrimitivePool : public PrimitivePool { - public: - static MemoryReorderPrimitivePool& GetInstance() { - static MemoryReorderPrimitivePool pool; - return pool; - } - - static MemoryReorderPrimitive* Get(const MemoryReorderParams& params) { - MemoryReorderPrimitive* primitive = static_cast( - MemoryReorderPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - if (primitive == nullptr) { - auto reorder_primitive = onnxruntime::make_unique(params); - primitive = reorder_primitive.get(); - MemoryReorderPrimitivePool::GetInstance().SetPrimitive(params.ToString(), std::move(reorder_primitive)); - } - primitive->SetMemory(params); - return primitive; - } - - private: - MemoryReorderPrimitivePool() = default; - ~MemoryReorderPrimitivePool() = default; -}; - -template -static void DoReorder(const MemoryReorderParams& params) { - std::vector net; - net.push_back(*(MemoryReorderPrimitivePool::Get(params)->GetPrimitive())); - mkldnn::stream(mkldnn::stream::kind::eager).submit(net).wait(); -} - } // namespace mkl_dnn } // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.cc b/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.cc index 84c20531d6..cdbafe4fad 100644 --- a/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.cc +++ b/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.cc @@ -42,31 +42,11 @@ MKLDNNExecutionProvider::~MKLDNNExecutionProvider() { } namespace mkl_dnn { -class ONNX_OPERATOR_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 1, Conv); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 7, Gemm); -class ONNX_OPERATOR_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 6, Relu); -class ONNX_OPERATOR_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 6, Sum); -class ONNX_OPERATOR_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 7, BatchNormalization); -class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 7, 8, float, AveragePool); -class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 1, 8, float, GlobalAveragePool); -class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 1, 7, float, MaxPool); -class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 8, 8, float, MaxPool); -class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 1, 8, float, GlobalMaxPool); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kMklDnnExecutionProvider, kOnnxDomain, 1, float, LRN); void RegisterMKLDNNKernels(KernelRegistry& kernel_registry) { static const BuildKernelCreateInfoFn function_table[] = { - BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, }; for (auto& function_table_entry : function_table) { @@ -161,6 +141,7 @@ void MKLDNNExecutionProvider::CreateOrUpdateMklDnnNode(const Node* node, subgraph_ptr->mkldnn_nodes.push_back(mkldnn_node); output_to_source_node_map.insert(std::make_pair(node_outputs[0]->Name(), subgraph_ptr->mkldnn_nodes.size() - 1)); } else { + subgraph_ptr->mkldnn_nodes.back().num_inputs += static_cast(node->InputDefs().size() - 1); const auto& node_outputs = node->OutputDefs(); output_to_source_node_map.erase(subgraph_ptr->mkldnn_nodes.back().output_name); subgraph_ptr->mkldnn_nodes.back().output_name = node_outputs[0]->Name(); @@ -183,9 +164,17 @@ void MKLDNNExecutionProvider::CreateOrUpdateMklDnnNode(const Node* node, NodeAttributes attributes = node->GetAttributes(); if (attributes.size() > 0) { size_t index = subgraph_ptr->mkldnn_nodes.size(); + std::string op_name; + if (fused) { + for (auto iter = node->InputNodesBegin(); iter != node->InputNodesEnd(); ++iter) { + op_name = (*iter).OpType(); + } + } else { + op_name = node->OpType(); + } for (auto att_it = attributes.begin(); att_it != attributes.end(); ++att_it) { - std::string key = node->OpType() + "-" + std::to_string(index) + "-" + att_it->first; + 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; } @@ -245,8 +234,14 @@ std::vector> MKLDNNExecutionProvider::GetCapa // can we fuse (at mkldnn level) nodes? bool fused = false; + if (sub_var.subgraph_node_indexes.size() > 1 && node->OpType() == "BatchNormalization") { + if (subgraph_ptr->mkldnn_nodes.back().name == "Conv") { + subgraph_ptr->mkldnn_nodes.back().name += "-BatchNormalization"; + fused = true; + } + } if (sub_var.subgraph_node_indexes.size() > 1 && node->OpType() == "Relu") { - if (subgraph_ptr->mkldnn_nodes.back().name == "BatchNormalization" || subgraph_ptr->mkldnn_nodes.back().name == "Conv") { + if (subgraph_ptr->mkldnn_nodes.back().name == "Conv-BatchNormalization" || subgraph_ptr->mkldnn_nodes.back().name == "BatchNormalization" || subgraph_ptr->mkldnn_nodes.back().name == "Conv") { subgraph_ptr->mkldnn_nodes.back().name += "-Relu"; fused = true; } diff --git a/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.h b/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.h index a57f290689..cf5624e87f 100644 --- a/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.h +++ b/onnxruntime/core/providers/mkldnn/mkldnn_execution_provider.h @@ -58,6 +58,24 @@ class MKLDNNExecutionProvider : public IExecutionProvider { reordered_buffers_.push_back(std::move(buffer)); } + std::shared_ptr GetBiasMemoryBuffer(const std::string& key) { + auto iter = bias_mem_map_.find(key); + if (iter != bias_mem_map_.end()) + return iter->second; + return nullptr; + } + + // Conv+BathNorm fusion. save scaled bias memory. + void SetBiasMemoryBuffer(const std::string& key, + const std::shared_ptr& bias_mem) { + bias_mem_map_.insert(std::make_pair(key, bias_mem)); + } + + void SaveAllocatedBiasMemory(IAllocatorUniquePtr buffer) { + // keep reordered memory buffers in scope. + biass_buffers_.push_back(std::move(buffer)); + } + std::vector> GetCapability(const onnxruntime::GraphViewer& graph, const std::vector& /*kernel_registries*/) const override; @@ -71,6 +89,11 @@ class MKLDNNExecutionProvider : public IExecutionProvider { std::unordered_map> weights_mem_map_; // Save reordered memory buffers in list so that memory is not freed. std::vector> reordered_buffers_; + + // conv+batchnorm fusion. normalized bias memory blocks from first iteration + std::unordered_map> bias_mem_map_; + // Conv+BathNorm fusion bias memory buffer. + std::vector> biass_buffers_; OrtMutex mutex_; // SUBGRAPH @@ -79,13 +102,13 @@ class MKLDNNExecutionProvider : public IExecutionProvider { 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 graph_name; int opset = GetOnnxOpSet(graph_viewer); - - int index = 0; + + int index = 0; if (graph_viewer.MaxNodeIndex() > 0) { auto first_node = graph_viewer.GetNode(index); while (first_node == nullptr) { diff --git a/onnxruntime/core/providers/mkldnn/nn/batch_norm.cc b/onnxruntime/core/providers/mkldnn/nn/batch_norm.cc deleted file mode 100644 index 67cdf353c4..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/batch_norm.cc +++ /dev/null @@ -1,285 +0,0 @@ -// Copyright(C) 2018 Intel Corporation -// Licensed under the MIT License - -#ifdef _WIN32 -#pragma warning(disable : 4244) -#endif - -#include "core/providers/mkldnn/mkldnn_common.h" -#include "core/providers/mkldnn/nn/batch_norm.h" -#include "core/providers/mkldnn/mkldnn_fwd.h" -#include "core/providers/mkldnn/memcpy_s.h" -#include "core/providers/cpu/nn/batch_norm_helper.h" - -namespace onnxruntime { -namespace mkl_dnn { - -namespace { -// Struct which encapsulates parameters for MKLDNN BatchNorm primitive. -struct BatchNormParams { - const mkldnn::memory::dims& src_dims; - const mkldnn::memory::dims& scale_dims; - const mkldnn::memory::dims& b_dims; - const mkldnn::memory::dims& mean_dims; - const mkldnn::memory::dims& var_dims; - const mkldnn::memory::dims& dst_dims; - const float epsilon; - - BatchNormParams(const mkldnn::memory::dims& src_dims_mkl, - const mkldnn::memory::dims& scale_dims_mkl, - const mkldnn::memory::dims& b_dims_mkl, const mkldnn::memory::dims& mean_dims_mkl, - const mkldnn::memory::dims& var_dims_mkl, const mkldnn::memory::dims& dst_dims_mkl, - float eps) - : src_dims(src_dims_mkl), - scale_dims(scale_dims_mkl), - b_dims(b_dims_mkl), - mean_dims(mean_dims_mkl), - var_dims(var_dims_mkl), - dst_dims(dst_dims_mkl), - epsilon(eps) {} - - // Used as the key for BatchNorm Primitive Reuse Pool. - std::string ToString() const { - std::string key; - key.reserve(128); - key.append("BatchNorm_"); - AddDimsToKey(key, src_dims); - AddDimsToKey(key, scale_dims); - AddDimsToKey(key, b_dims); - AddDimsToKey(key, mean_dims); - AddDimsToKey(key, var_dims); - AddDimsToKey(key, dst_dims); - return key; - } -}; - -template -class BatchNormPrimitive final : public PrimitiveBase { - public: - explicit BatchNormPrimitive(const BatchNormParams& params) - : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); - if (context_.batchnorm_fwd == nullptr) { - Initialize(params); - } - } - - ~BatchNormPrimitive() = default; - - void Compute(const T* src_data, const T* scale_data, const T* b_data, - const T* mean_data, const T* var_data, T* dst_data, - int scale_dims_channels) { - context_.src_mem->set_data_handle( - static_cast(const_cast(src_data))); - context_.mean_mem->set_data_handle( - static_cast(const_cast(mean_data))); - context_.var_mem->set_data_handle( - static_cast(const_cast(var_data))); - context_.dst_mem->set_data_handle( - static_cast(dst_data)); - - T* scale_shift_buf = static_cast(context_.scale_shift_mem->get_data_handle()); - - size_t src_bytes = sizeof(T) * scale_dims_channels; - size_t dst_bytes = sizeof(T) * scale_dims_channels; - - MEMCPY_S(scale_shift_buf, scale_data, src_bytes, dst_bytes); - MEMCPY_S(&scale_shift_buf[scale_dims_channels], b_data, src_bytes, dst_bytes); - context_.stream->submit(context_.net); - return; - } - - std::unique_ptr - GetPrimitiveDesc() const { - return context_.conv_fwd_pd; - } - - private: - struct BatchNormContext { - std::unique_ptr src_mem; - std::unique_ptr scale_shift_mem; - std::unique_ptr mean_mem; - std::unique_ptr var_mem; - std::unique_ptr dst_mem; - - std::unique_ptr src_md; - std::unique_ptr scale_shift_md; - std::unique_ptr mean_md; - std::unique_ptr var_md; - std::unique_ptr dst_md; - - std::unique_ptr batchnorm_fwd; - std::unique_ptr - batchnorm_fwd_pd; - - std::unique_ptr stream; - std::vector net; - }; - - void Initialize(const BatchNormParams& params) { - mkldnn::memory::format fmt = mkldnn::memory::format::any; - switch (params.src_dims.size()) { - case 1: { fmt = mkldnn::memory::format::x; break; } - case 2: { fmt = mkldnn::memory::format::nc; break; } - case 3: { fmt = mkldnn::memory::format::ntc; break; } - case 4: { fmt = mkldnn::memory::format::nchw; break; } - case 5: { fmt = mkldnn::memory::format::ncdhw; break; } - default: { fmt = mkldnn::memory::format::any; break; } - } - context_.src_md.reset(new mkldnn::memory::desc( - { params.src_dims }, MklDnnType(), fmt)); - - context_.scale_shift_md.reset(new mkldnn::memory::desc( - { 2, params.scale_dims[0] }, MklDnnType(), mkldnn::memory::format::nc)); - - context_.mean_md.reset(new mkldnn::memory::desc( - { params.mean_dims }, MklDnnType(), mkldnn::memory::format::x)); - context_.var_md.reset(new mkldnn::memory::desc( - { params.var_dims }, MklDnnType(), mkldnn::memory::format::x)); - context_.dst_md.reset(new mkldnn::memory::desc( - { params.dst_dims }, MklDnnType(), fmt)); - - context_.src_mem.reset( - new mkldnn::memory({ *context_.src_md, cpu_engine_ }, nullptr)); - - // scale_shift_mem will allocate 2*C*sizeof(float) buffer - // - context_.scale_shift_mem.reset( - new mkldnn::memory({ *context_.scale_shift_md, cpu_engine_ })); - - context_.mean_mem.reset( - new mkldnn::memory({ *context_.mean_md, cpu_engine_ }, nullptr)); - context_.var_mem.reset( - new mkldnn::memory({ *context_.var_md, cpu_engine_ }, nullptr)); - - context_.batchnorm_fwd.reset(new mkldnn::batch_normalization_forward::desc( - mkldnn::prop_kind::forward_inference, *context_.src_md, params.epsilon, - mkldnn::batch_normalization_flag::use_scale_shift | - mkldnn::batch_normalization_flag::use_global_stats)); - - context_.batchnorm_fwd_pd.reset( - new mkldnn::batch_normalization_forward::primitive_desc( - *context_.batchnorm_fwd, cpu_engine_)); - - context_.dst_mem.reset( - new mkldnn::memory( - context_.batchnorm_fwd_pd->dst_primitive_desc(), nullptr)); - - auto bn = mkldnn::batch_normalization_forward( - *context_.batchnorm_fwd_pd, - (const mkldnn::primitive::at)*context_.src_mem, - (const mkldnn::primitive::at)*context_.mean_mem, - (const mkldnn::primitive::at)*context_.var_mem, - (const mkldnn::memory)*context_.scale_shift_mem, - (const mkldnn::memory) *context_.dst_mem); - - context_.net.push_back(bn); - } - - BatchNormContext context_; - mkldnn::engine& cpu_engine_; -}; - -// Pool which allows for reuse of MKLDNN BatchNorm primitives which are -// expensive to instantiate. To address thread safety, the primitives are -// stored in a map on thread local storage. -template -class BatchNormPrimitivePool : public PrimitivePool { - public: - static BatchNormPrimitive* Get(const BatchNormParams& params) { - BatchNormPrimitive* primitive = - dynamic_cast*>( - BatchNormPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - - if (primitive == nullptr) { - auto BatchNorm_primitive = onnxruntime::make_unique>(params); - primitive = BatchNorm_primitive.get(); - BatchNormPrimitivePool::GetInstance().SetPrimitive( - params.ToString(), std::move(BatchNorm_primitive)); - } - return primitive; - } - - private: - BatchNormPrimitivePool() = default; - ~BatchNormPrimitivePool() = default; - - static BatchNormPrimitivePool& GetInstance() { - static BatchNormPrimitivePool pool; - return pool; - } -}; -} // namespace - -template -Status BatchNorm::Compute(OpKernelContext* context) const { - const Tensor* X = context->Input(0); - - int num_dimensions = static_cast(X->Shape().NumDimensions()); - if (num_dimensions == 3) { - // Fall back CPU implementation - return onnxruntime::BatchNorm::Compute(context); - } - - const T* src_data = X->template Data(); - - const Tensor* scale = context->Input(1); - const T* scale_data = scale->template Data(); - - const Tensor* B = context->Input(2); - const T* b_data = B->template Data(); - - const Tensor* mean = context->Input(3); - const T* mean_data = mean->template Data(); - - const Tensor* var = context->Input(4); - const T* var_data = var->template Data(); - - Tensor* Y = context->Output(0, X->Shape()); - T* dst_data = Y->template MutableData(); - - ORT_RETURN_IF_ERROR( - BatchNormHelper::ValidateInputs(X, scale, B, mean, var)); - - mkldnn::memory::dims src_dims_mkl( - X->Shape().GetDims().begin(), X->Shape().GetDims().end()); - mkldnn::memory::dims scale_dims_mkl( - scale->Shape().GetDims().begin(), scale->Shape().GetDims().end()); - mkldnn::memory::dims b_dims_mkl( - B->Shape().GetDims().begin(), B->Shape().GetDims().end()); - mkldnn::memory::dims mean_dims_mkl( - mean->Shape().GetDims().begin(), mean->Shape().GetDims().end()); - mkldnn::memory::dims var_dims_mkl( - var->Shape().GetDims().begin(), var->Shape().GetDims().end()); - - mkldnn::memory::dims dst_dims_mkl( - Y->Shape().GetDims().begin(), Y->Shape().GetDims().end()); - - try { - BatchNormParams batchNorm_params(src_dims_mkl, scale_dims_mkl, - b_dims_mkl, mean_dims_mkl, var_dims_mkl, dst_dims_mkl, - onnxruntime::BatchNorm::epsilon_); - BatchNormPrimitive* batchNorm_primitive = - BatchNormPrimitivePool::Get(batchNorm_params); - ORT_RETURN_IF_NOT(batchNorm_primitive != nullptr); - batchNorm_primitive->Compute(src_data, scale_data, b_data, - mean_data, var_data, dst_data, scale_dims_mkl[0]); - - } catch (const mkldnn::error& e) { - return ORT_MAKE_STATUS( - ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); - } - - return Status::OK(); -} - -ONNX_OPERATOR_KERNEL_EX( - BatchNormalization, - kOnnxDomain, - 7, - kMklDnnExecutionProvider, - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), - BatchNorm); - -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/nn/batch_norm.h b/onnxruntime/core/providers/mkldnn/nn/batch_norm.h deleted file mode 100644 index eb4fb5a53e..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/batch_norm.h +++ /dev/null @@ -1,18 +0,0 @@ -// Copyright(C) 2018 Intel Corporation -// Licensed under the MIT License - -#pragma once -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/nn/batch_norm.h" - -namespace onnxruntime { -namespace mkl_dnn { - -template -class BatchNorm final : public onnxruntime::BatchNorm { - public: - explicit BatchNorm(const OpKernelInfo& info) : onnxruntime::BatchNorm(info) {} - Status Compute(OpKernelContext* context) const override; -}; -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/nn/conv.cc b/onnxruntime/core/providers/mkldnn/nn/conv.cc deleted file mode 100644 index b14d8a8519..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/conv.cc +++ /dev/null @@ -1,473 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#ifdef _WIN32 -#pragma warning(disable : 4244) -#endif -#include -#include - -#include "core/providers/mkldnn/nn/conv.h" -#include "core/providers/mkldnn/mkldnn_common.h" -#include "core/providers/mkldnn/mkldnn_fwd.h" - -namespace onnxruntime { -namespace mkl_dnn { - -namespace { -// Struct which encapsulates parameters for MKLDNN Conv primitive. -struct ConvParams { - const mkldnn::memory::dims& src_dims; - const mkldnn::memory::dims& filter_dims; - const mkldnn::memory::dims& bias_dims; - const mkldnn::memory::dims& dst_dims; - const mkldnn::memory::dims& strides; - const mkldnn::memory::dims& dilations; - const mkldnn::memory::dims& padding_left; - const mkldnn::memory::dims& padding_right; - - ConvParams(const mkldnn::memory::dims& src_dims, const mkldnn::memory::dims& filter_dims, - const mkldnn::memory::dims& bias_dims, const mkldnn::memory::dims& dst_dims, - const mkldnn::memory::dims& strides, const mkldnn::memory::dims& dilations, - const mkldnn::memory::dims& padding_left, const mkldnn::memory::dims& padding_right) - : src_dims(src_dims), - filter_dims(filter_dims), - bias_dims(bias_dims), - dst_dims(dst_dims), - strides(strides), - dilations(dilations), - padding_left(padding_left), - padding_right(padding_right) {} - - // Used as the key for Conv Primitive Reuse Pool. - std::string ToString() const { - std::string key; - key.reserve(128); - key.append("conv_"); - AddDimsToKey(key, src_dims); - AddDimsToKey(key, filter_dims); - AddDimsToKey(key, bias_dims); - AddDimsToKey(key, dst_dims); - AddDimsToKey(key, strides); - AddDimsToKey(key, dilations); - AddDimsToKey(key, padding_left); - AddDimsToKey(key, padding_right); - return key; - } -}; - -template -class ConvPrimitive : public PrimitiveBase { - public: - explicit ConvPrimitive(const ConvParams& params) - : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); - if (context_.conv_fwd == nullptr) { - Initialize(params); - } - } - - ~ConvPrimitive() = default; - - void Compute(const T* src_data, const T* filter_data, - T* dst_data, const T* bias_data = nullptr) { - context_.src_mem->set_data_handle( - static_cast(const_cast(src_data))); - context_.filter_mem->set_data_handle( - static_cast(const_cast(filter_data))); - if (bias_data != nullptr) { - context_.bias_mem->set_data_handle( - static_cast(const_cast(bias_data))); - } - context_.dst_mem->set_data_handle( - static_cast(dst_data)); - context_.stream->submit(context_.net); - - context_.src_mem->set_data_handle(nullptr); - context_.filter_mem->set_data_handle(nullptr); - if (bias_data != nullptr) { - context_.bias_mem->set_data_handle(nullptr); - } - context_.dst_mem->set_data_handle(nullptr); - return; - } - - mkldnn::memory::format GetSrcMemoryFormat() const { return context_.src_fmt; } - - mkldnn::memory::format GetFilterMemoryFormat() const { return context_.filter_fmt; } - - mkldnn::memory::format GetDstMemoryFormat() const { return context_.dst_fmt; } - - size_t GetSrcSize() const { return context_.src_size; } - - size_t GetFilterSize() const { return context_.filter_size; } - - size_t GetDstSize() const { return context_.dst_size; } - - mkldnn::convolution_forward::primitive_desc* GetPrimitiveDesc() const { - return context_.conv_fwd_pd.get(); - } - - private: - struct ConvContext { - mkldnn::memory::format src_fmt; - mkldnn::memory::format filter_fmt; - mkldnn::memory::format dst_fmt; - - size_t src_size; - size_t filter_size; - size_t dst_size; - - std::unique_ptr src_mem; - std::unique_ptr filter_mem; - std::unique_ptr bias_mem; - std::unique_ptr dst_mem; - - std::unique_ptr fwd_desc; - - std::unique_ptr src_md; - std::unique_ptr filter_md; - std::unique_ptr bias_md; - std::unique_ptr dst_md; - - std::unique_ptr conv_fwd_pd; - std::unique_ptr conv_fwd; - - std::unique_ptr stream; - std::vector net; - - ConvContext() - : src_fmt(mkldnn::memory::format::any), - filter_fmt(mkldnn::memory::format::any), - dst_fmt(mkldnn::memory::format::any), - src_size(0), - filter_size(0), - dst_size(0), - src_mem(nullptr), - filter_mem(nullptr), - bias_mem(nullptr), - dst_mem(nullptr), - fwd_desc(nullptr), - src_md(nullptr), - filter_md(nullptr), - bias_md(nullptr), - conv_fwd_pd(nullptr), - conv_fwd(nullptr), - stream(nullptr) {} - }; - - void Initialize(const ConvParams& params) { - // Set the memory descriptors to format::any to allow MKLDNN to decide what the optimal memory layout should be - // for the computation given the input params. - context_.src_md.reset(new mkldnn::memory::desc( - {params.src_dims}, MklDnnType(), mkldnn::memory::format::any)); - context_.filter_md.reset(new mkldnn::memory::desc( - {params.filter_dims}, MklDnnType(), mkldnn::memory::format::any)); - context_.dst_md.reset(new mkldnn::memory::desc( - {params.dst_dims}, MklDnnType(), mkldnn::memory::format::any)); - if (!params.bias_dims.empty()) - context_.bias_md.reset(new mkldnn::memory::desc( - {params.bias_dims}, MklDnnType(), mkldnn::memory::format::any)); - - if (!params.bias_dims.empty()) { - context_.fwd_desc.reset(new mkldnn::convolution_forward::desc( - mkldnn::prop_kind::forward_inference, mkldnn::convolution_direct, *context_.src_md, - *context_.filter_md, *context_.bias_md, *context_.dst_md, - params.strides, params.dilations, params.padding_left, - params.padding_right, mkldnn::padding_kind::zero)); - } else { - context_.fwd_desc.reset(new mkldnn::convolution_forward::desc( - mkldnn::prop_kind::forward_inference, mkldnn::convolution_direct, *context_.src_md, - *context_.filter_md, *context_.dst_md, params.strides, - params.dilations, params.padding_left, - params.padding_right, mkldnn::padding_kind::zero)); - } - - context_.conv_fwd_pd.reset(new mkldnn::convolution_forward::primitive_desc( - *context_.fwd_desc, cpu_engine_)); - - context_.src_fmt = static_cast( - context_.conv_fwd_pd.get()->src_primitive_desc().desc().data.format); - - context_.filter_fmt = static_cast( - context_.conv_fwd_pd.get()->weights_primitive_desc().desc().data.format); - - context_.dst_fmt = static_cast( - context_.conv_fwd_pd.get()->dst_primitive_desc().desc().data.format); - - context_.src_size = context_.conv_fwd_pd.get()->src_primitive_desc().get_size(); - - context_.filter_size = context_.conv_fwd_pd.get()->weights_primitive_desc().get_size(); - - context_.dst_size = context_.conv_fwd_pd.get()->dst_primitive_desc().get_size(); - - context_.src_mem.reset( - new mkldnn::memory(context_.conv_fwd_pd.get()->src_primitive_desc(), nullptr)); - context_.filter_mem.reset( - new mkldnn::memory(context_.conv_fwd_pd.get()->weights_primitive_desc(), nullptr)); - context_.dst_mem.reset( - new mkldnn::memory(context_.conv_fwd_pd.get()->dst_primitive_desc(), nullptr)); - - if (!params.bias_dims.empty()) { - context_.bias_mem.reset( - new mkldnn::memory(context_.conv_fwd_pd.get()->bias_primitive_desc(), nullptr)); - - context_.conv_fwd.reset(new mkldnn::convolution_forward( - *context_.conv_fwd_pd, *context_.src_mem, *context_.filter_mem, - *context_.bias_mem, *context_.dst_mem)); - } else { - context_.conv_fwd.reset( - new mkldnn::convolution_forward(*context_.conv_fwd_pd, *context_.src_mem, - *context_.filter_mem, *context_.dst_mem)); - } - - context_.net.push_back(*context_.conv_fwd); - } - - ConvContext context_; - mkldnn::engine& cpu_engine_; -}; - -// Pool which allows for reuse of MKLDNN Conv primitives which are expensive to instantiate. -// To address thread safety, the primitives are stored in a map on thread local storage. -template -class ConvPrimitivePool : public PrimitivePool { - public: - static ConvPrimitive* Get(const ConvParams& params) { - ConvPrimitive* primitive = dynamic_cast*>( - ConvPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - - if (primitive == nullptr) { - auto conv_primitive = onnxruntime::make_unique>(params); - primitive = conv_primitive.get(); - ConvPrimitivePool::GetInstance().SetPrimitive(params.ToString(), std::move(conv_primitive)); - } - return primitive; - } - - private: - ConvPrimitivePool() = default; - ~ConvPrimitivePool() = default; - - static ConvPrimitivePool& GetInstance() { - static ConvPrimitivePool pool; - return pool; - } -}; -} // namespace - -template -Status Conv::Compute(OpKernelContext* context) const { - size_t num_inputs = OpKernel::Node().InputDefs().size(); - - const Tensor* X = context->Input(0); - const Tensor* W = context->Input(1); - const Tensor* B = num_inputs == 3 ? context->Input(2) : nullptr; - const int64_t N = X->Shape()[0]; - const int64_t M = W->Shape()[0]; - const int group_mkl = static_cast(this->conv_attrs_.group); - - ORT_RETURN_IF_ERROR(this->conv_attrs_.ValidateInputShape(X, W)); - - std::vector kernel_shape; - ORT_RETURN_IF_ERROR(this->conv_attrs_.ComputeKernelShape(W->Shape(), kernel_shape)); - const size_t kernel_rank = kernel_shape.size(); - - if (kernel_rank > 3) { - // Fall Back to CPU implementation. - return onnxruntime::Conv::Compute(context); - } - - if (kernel_rank + 2 != W->Shape().NumDimensions()) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "kernel_shape num_dims is not compatible with W num_dims.", - " kernel_shape: ", TensorShape(kernel_shape).ToString().c_str(), - " W: ", W->Shape().ToString().c_str()); - } - - for (size_t i = 0; i < kernel_rank; ++i) { - if (kernel_shape[i] != W->Shape()[i + 2]) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "kernel_shape is not compatible with W shape.", - " kernel_shape: ", TensorShape(kernel_shape).ToString().c_str(), - " W: ", W->Shape().ToString().c_str()); - } - } - - std::vector pads(this->conv_attrs_.pads); - if (pads.empty()) { - pads.resize(kernel_rank * 2, 0); - } - std::vector dilations(this->conv_attrs_.dilations); - if (dilations.empty()) { - dilations.resize(kernel_rank, 1); - } - std::vector strides(this->conv_attrs_.strides); - if (strides.empty()) { - strides.resize(kernel_rank, 1); - } - - std::vector Y_dims; - Y_dims.insert(Y_dims.begin(), {N, M}); - TensorShape input_shape = X->Shape().Slice(2); - ORT_RETURN_IF_ERROR(this->conv_attrs_.InferOutputShape(input_shape, kernel_shape, - strides, dilations, &pads, &Y_dims)); - Tensor* Y = context->Output(0, TensorShape(Y_dims)); - TensorShape output_shape = Y->Shape().Slice(2); - - mkldnn::memory::dims src_dims_mkl(X->Shape().GetDims().begin(), X->Shape().GetDims().end()); - mkldnn::memory::dims filter_dims_mkl; - if (group_mkl == 1) { - filter_dims_mkl.assign(W->Shape().GetDims().begin(), W->Shape().GetDims().end()); - } else { - filter_dims_mkl.assign({group_mkl, - static_cast(W->Shape()[0] / group_mkl)}); - filter_dims_mkl.insert(filter_dims_mkl.end(), W->Shape().GetDims().begin() + 1, W->Shape().GetDims().end()); - } - mkldnn::memory::dims strides_mkl(strides.begin(), strides.end()); - mkldnn::memory::dims dilations_mkl(dilations.begin(), dilations.end()); - // mkldnn dilations start from 0 so we need to subtract 1 from each dim. - for (size_t dim = 0; dim < kernel_rank; dim++) { - dilations_mkl[dim] -= 1; - } - - mkldnn::memory::dims padding_left_mkl(pads.begin(), pads.begin() + kernel_rank); - mkldnn::memory::dims padding_right_mkl(pads.begin() + kernel_rank, pads.end()); - mkldnn::memory::dims dst_dims_mkl(Y_dims.begin(), Y_dims.end()); - mkldnn::memory::dims bias_dims_mkl; - if (B != nullptr) { - bias_dims_mkl.assign(B->Shape().GetDims().begin(), B->Shape().GetDims().end()); - } - - AllocatorPtr alloc; - ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&alloc)); - IAllocatorUniquePtr src_reorder_buffer; - IAllocatorUniquePtr dst_reorder_buffer; - - const T* src_data = X->template Data(); - const T* filter_data = W->template Data(); - T* dst_data = Y->template MutableData(); - const T* bias_data = nullptr; - if (B != nullptr) { - bias_data = B->template Data(); - } - - try { - ConvParams conv_params(src_dims_mkl, filter_dims_mkl, bias_dims_mkl, - dst_dims_mkl, strides_mkl, dilations_mkl, - padding_left_mkl, padding_right_mkl); - ConvPrimitive* conv_primitive = ConvPrimitivePool::Get(conv_params); - auto conv_fwd_pd = conv_primitive->GetPrimitiveDesc(); - - mkldnn::engine& cpu_engine = GetEngine(); - - enum mkldnn::memory::format src_format = mkldnn::memory::format::format_undef; - enum mkldnn::memory::format filter_format = mkldnn::memory::format::format_undef; - enum mkldnn::memory::format dst_format = mkldnn::memory::format::format_undef; - - if (kernel_rank == 1) { - src_format = mkldnn::memory::format::ncw; - if (group_mkl == 1) { - filter_format = mkldnn::memory::format::oiw; - } else { - filter_format = mkldnn::memory::format::goiw; - } - dst_format = mkldnn::memory::format::ncw; - } else if (kernel_rank == 2) { - src_format = mkldnn::memory::format::nchw; - if (group_mkl == 1) { - filter_format = mkldnn::memory::format::oihw; - } else { - filter_format = mkldnn::memory::format::goihw; - } - dst_format = mkldnn::memory::format::nchw; - } else { - src_format = mkldnn::memory::format::ncdhw; - if (group_mkl == 1) { - filter_format = mkldnn::memory::format::oidhw; - } else { - filter_format = mkldnn::memory::format::goidhw; - } - dst_format = mkldnn::memory::format::ncdhw; - } - - auto src_md = mkldnn::memory::desc(src_dims_mkl, MklDnnType(), src_format); - auto dst_md = mkldnn::memory::desc(dst_dims_mkl, MklDnnType(), dst_format); - - // Reorder src memory layout if necessary. - if (src_md.data.format != conv_primitive->GetSrcMemoryFormat()) { - auto pd = mkldnn::memory::primitive_desc(src_md, cpu_engine); - mkldnn::memory src = mkldnn::memory(pd, (void*)src_data); - // allocate the size queried from memory primitive desc. it may not match tensor logical size due to - // mkldnn using padding to allow use of blocked format. - src_reorder_buffer = IAllocator::MakeUniquePtr(alloc, conv_primitive->GetSrcSize()); - mkldnn::memory dst = mkldnn::memory(conv_fwd_pd->src_primitive_desc(), src_reorder_buffer.get()); - MemoryReorderParams params(src, dst); - DoReorder(params); - src_data = static_cast(dst.get_data_handle()); - } - - // Reorder filter memory layout if necessary - // Avoid data reordering. Save filter memory in mkldnn format from first iteration - // in execution provider mapped by weight name. - { - // lock to make sure reordering is done only once - std::lock_guard lock(provider_->GetMutex()); - auto weight_name = OpKernel::Node().InputDefs()[1]->Name(); - std::shared_ptr filter_dst_mem = provider_->GetWeightsMemoryBuffer(weight_name); - - if (filter_dst_mem == nullptr) { - if (filter_format != conv_primitive->GetFilterMemoryFormat()) { - auto pd = mkldnn::memory::primitive_desc(mkldnn::memory::desc( - filter_dims_mkl, MklDnnType(), filter_format), - cpu_engine); - mkldnn::memory src = mkldnn::memory(pd, (void*)filter_data); - IAllocatorUniquePtr filter_reorder_buffer = IAllocator::MakeUniquePtr(alloc, conv_primitive->GetFilterSize()); - filter_dst_mem.reset( - new mkldnn::memory(conv_fwd_pd->weights_primitive_desc(), filter_reorder_buffer.get())); - - MemoryReorderParams params(src, *filter_dst_mem); - DoReorder(params); - provider_->SaveAllocatedMemory(std::move(filter_reorder_buffer)); - - filter_data = static_cast(filter_dst_mem->get_data_handle()); - provider_->SetWeightsMemoryBuffer(weight_name, filter_dst_mem); - } - } else { - filter_data = static_cast(filter_dst_mem->get_data_handle()); - } - } - // Allocate dst buffer if reorder is necessary - if (dst_md.data.format != conv_primitive->GetDstMemoryFormat()) { - // allocate the size queried from memory primitive desc. it may not match tensor logical size due to - // mkldnn using padding to allow use of blocked format. - dst_reorder_buffer = IAllocator::MakeUniquePtr(alloc, conv_primitive->GetDstSize()); - dst_data = static_cast(dst_reorder_buffer.get()); - } - - conv_primitive->Compute(src_data, filter_data, dst_data, bias_data); - - // Reorder dst memory layout if necessary - if (dst_md.data.format != conv_primitive->GetDstMemoryFormat()) { - mkldnn::memory src = mkldnn::memory(conv_fwd_pd->dst_primitive_desc(), (void*)dst_data); - auto pd = mkldnn::memory::primitive_desc(dst_md, cpu_engine); - mkldnn::memory dst = mkldnn::memory(pd, Y->template MutableData()); - MemoryReorderParams params(src, dst); - DoReorder(params); - } - - } catch (const mkldnn::error& e) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); - } - - return Status::OK(); -} - -ONNX_OPERATOR_KERNEL_EX( - Conv, - kOnnxDomain, - 1, - kMklDnnExecutionProvider, - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), - Conv); - -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/nn/conv.h b/onnxruntime/core/providers/mkldnn/nn/conv.h deleted file mode 100644 index e66566262e..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/conv.h +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#pragma once -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/nn/conv.h" -#include "core/providers/mkldnn/mkldnn_execution_provider.h" - -namespace onnxruntime { -namespace mkl_dnn { - -template -class Conv final : public onnxruntime::Conv { - public: - explicit Conv(const OpKernelInfo& info) : onnxruntime::Conv(info) { - provider_ = (const_cast( - dynamic_cast(info.GetExecutionProvider()))); - } - - Status Compute(OpKernelContext* context) const override; - - private: - MKLDNNExecutionProvider* provider_; -}; -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/nn/lrn.cc b/onnxruntime/core/providers/mkldnn/nn/lrn.cc deleted file mode 100644 index abbda45575..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/lrn.cc +++ /dev/null @@ -1,241 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#ifdef _WIN32 -#pragma warning(disable : 4244) -#endif - -#include "core/providers/mkldnn/mkldnn_common.h" -#include "core/providers/mkldnn/nn/lrn.h" -#include "core/providers/mkldnn/mkldnn_fwd.h" - -namespace onnxruntime { -namespace mkl_dnn { - -ONNX_OPERATOR_TYPED_KERNEL_EX( - LRN, - kOnnxDomain, - 1, - float, - kMklDnnExecutionProvider, - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), - LRN); - -namespace { -// Struct which encapsulates parameters for MKLDNN LRN primitive. -struct LRNParams { - const mkldnn::memory::dims& dims_; - float alpha_; - float beta_; - float bias_; - int size_; - - LRNParams(const mkldnn::memory::dims& dims, float alpha, float beta, float bias, int size) - : dims_(dims), alpha_(alpha), beta_(beta), bias_(bias), size_(size) {} - - // Used as the key for LRN Primitive Reuse LRN. - std::string ToString() const { - std::string key; - key.reserve(128); - key.append("lrn"); - AddDimsToKey(key, dims_); - key.append('#' + std::to_string(alpha_) + '#'); - key.append('#' + std::to_string(beta_) + '#'); - key.append('#' + std::to_string(bias_) + '#'); - key.append('#' + std::to_string(size_) + '#'); - - return key; - } -}; - -template -class LRNPrimitive : public PrimitiveBase { - public: - explicit LRNPrimitive(const LRNParams& params) - : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); - if (context_.lrn_fwd == nullptr) { - Initialize(params); - } - } - - ~LRNPrimitive() = default; - - void Compute(const T* src_data, T* dst_data) { - context_.src_mem->set_data_handle(static_cast(const_cast(src_data))); - context_.dst_mem->set_data_handle(static_cast(dst_data)); - context_.stream->submit(context_.net); - - context_.src_mem->set_data_handle(nullptr); - context_.dst_mem->set_data_handle(nullptr); - return; - } - - mkldnn::memory::format GetSrcMemoryFormat() const { return context_.src_fmt; } - - size_t GetSrcSize() const { return context_.src_size; } - - size_t GetDstSize() const { return context_.dst_size; } - - mkldnn::lrn_forward::primitive_desc* GetPrimitiveDesc() const { - return context_.fwd_primitive_desc.get(); - } - - private: - struct LRNContext { - mkldnn::memory::format src_fmt; - std::unique_ptr src_md; - - size_t src_size; - size_t dst_size; - - std::unique_ptr src_mem; - std::unique_ptr dst_mem; - - std::unique_ptr fwd_desc; - std::unique_ptr fwd_primitive_desc; - std::unique_ptr lrn_fwd; - - std::unique_ptr stream; - std::vector net; - - LRNContext() - : src_fmt(mkldnn::memory::format::any), - src_md(nullptr), - src_size(0), - dst_size(0), - src_mem(nullptr), - dst_mem(nullptr), - fwd_desc(nullptr), - fwd_primitive_desc(nullptr), - lrn_fwd(nullptr), - stream(nullptr) {} - }; - - void Initialize(const LRNParams& params) { - context_.src_md.reset(new mkldnn::memory::desc({params.dims_}, MklDnnType(), mkldnn::memory::format::nchw)); - - mkldnn::algorithm algo = mkldnn::algorithm::lrn_across_channels; - context_.fwd_desc.reset(new mkldnn::lrn_forward::desc( - mkldnn::prop_kind::forward_scoring, algo, *context_.src_md, - params.size_, params.alpha_, params.beta_, params.bias_)); - - context_.fwd_primitive_desc.reset(new mkldnn::lrn_forward::primitive_desc( - *context_.fwd_desc, cpu_engine_)); - - context_.src_fmt = static_cast( - context_.fwd_primitive_desc.get()->src_primitive_desc().desc().data.format); - - context_.src_size = context_.fwd_primitive_desc.get()->src_primitive_desc().get_size(); - context_.dst_size = context_.fwd_primitive_desc.get()->dst_primitive_desc().get_size(); - - context_.src_mem.reset(new mkldnn::memory(context_.fwd_primitive_desc.get()->src_primitive_desc(), nullptr)); - context_.dst_mem.reset(new mkldnn::memory(context_.fwd_primitive_desc.get()->dst_primitive_desc(), nullptr)); - context_.lrn_fwd.reset( - new mkldnn::lrn_forward(*context_.fwd_primitive_desc, *context_.src_mem, *context_.dst_mem)); - context_.net.push_back(*context_.lrn_fwd); - } - - LRNContext context_; - mkldnn::engine& cpu_engine_; -}; - -// Pool which allows for reuse of MKLDNN Pool primitives which are expensive to instantiate. -// To address thread safety, the primitives are stored in a map on thread local storage. -template -class LRNPrimitivePool : public PrimitivePool { - public: - static LRNPrimitive* Get(const LRNParams& params) { - LRNPrimitive* primitive = dynamic_cast*>( - LRNPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - if (primitive == nullptr) { - auto pool_primitive = onnxruntime::make_unique>(params); - primitive = pool_primitive.get(); - LRNPrimitivePool::GetInstance().SetPrimitive(params.ToString(), std::move(pool_primitive)); - } - return primitive; - } - - private: - LRNPrimitivePool() = default; - ~LRNPrimitivePool() = default; - - static LRNPrimitivePool& GetInstance() { - static LRNPrimitivePool pool; - return pool; - } -}; -} // namespace - -template -Status LRN::Compute(OpKernelContext* context) const { - const Tensor* X = context->Input(0); - const T* src_data = X->template Data(); - - const TensorShape& x_shape = X->Shape(); - if (x_shape.NumDimensions() != 4) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Support NCHW image only."); - } - - const auto& x_dims = x_shape.GetDims(); - mkldnn::memory::dims dims_mkl(x_dims.begin(), x_dims.end()); - - Tensor* Y = context->Output(0, TensorShape(x_dims)); - T* dst_data = Y->template MutableData(); - - AllocatorPtr alloc; - ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&alloc)); - IAllocatorUniquePtr src_reorder_buffer; - IAllocatorUniquePtr dst_reorder_buffer; - - try { - LRNParams lrn_params(dims_mkl, this->alpha_, this->beta_, this->bias_, this->size_); - LRNPrimitive* lrn_primitive = LRNPrimitivePool::Get(lrn_params); - auto fwd_primitive_desc = lrn_primitive->GetPrimitiveDesc(); - - mkldnn::engine& cpu_engine = GetEngine(); - mkldnn::memory::format mem_format = dims_mkl.size() == 5 ? mkldnn::memory::format::ncdhw : mkldnn::memory::format::nchw; - // Per ONNX spec, X (src) is NCHW and Y (dst) is NCHW - auto src_md = mkldnn::memory::desc(dims_mkl, MklDnnType(), mem_format); - auto dst_md = mkldnn::memory::desc(dims_mkl, MklDnnType(), mem_format); - - // Reorder src memory layout if necessary. - if (src_md.data.format != lrn_primitive->GetSrcMemoryFormat()) { - auto pd = mkldnn::memory::primitive_desc(src_md, cpu_engine); - mkldnn::memory src = mkldnn::memory(pd, (void*)src_data); - // allocate the size queried from memory primitive desc. it may not match tensor logical size due to - // mkldnn using padding to allow use of blocked format. - src_reorder_buffer = IAllocator::MakeUniquePtr(alloc, lrn_primitive->GetSrcSize()); - mkldnn::memory dst = mkldnn::memory(fwd_primitive_desc->src_primitive_desc(), src_reorder_buffer.get()); - MemoryReorderParams params(src, dst); - DoReorder(params); - src_data = static_cast(dst.get_data_handle()); - } - - // Allocate dst buffer if reorder is necessary - if (src_md.data.format != lrn_primitive->GetSrcMemoryFormat()) { - // allocate the size queried from memory primitive desc. it may not match tensor logical size due to - // mkldnn using padding to allow use of blocked format. - dst_reorder_buffer = IAllocator::MakeUniquePtr(alloc, lrn_primitive->GetDstSize()); - dst_data = static_cast(dst_reorder_buffer.get()); - } - - lrn_primitive->Compute(src_data, dst_data); - - // Reorder dst memory layout if necessary - if (src_md.data.format != lrn_primitive->GetSrcMemoryFormat()) { - mkldnn::memory src = mkldnn::memory(fwd_primitive_desc->dst_primitive_desc(), (void*)dst_data); - auto pd = mkldnn::memory::primitive_desc(dst_md, cpu_engine); - mkldnn::memory dst = mkldnn::memory(pd, Y->template MutableData()); - MemoryReorderParams params(src, dst); - DoReorder(params); - } - } catch (const mkldnn::error& e) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); - } - - return Status::OK(); -} - -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/nn/lrn.h b/onnxruntime/core/providers/mkldnn/nn/lrn.h deleted file mode 100644 index 44ea04b444..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/lrn.h +++ /dev/null @@ -1,20 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#pragma once -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/nn/lrn.h" - -namespace onnxruntime { -namespace mkl_dnn { - -template -class LRN final : public onnxruntime::LRN { - public: - explicit LRN(const OpKernelInfo& info) : onnxruntime::LRN(info) {} - - Status Compute(OpKernelContext* p_op_kernel_context) const override; -}; - -} -} diff --git a/onnxruntime/core/providers/mkldnn/nn/pool.cc b/onnxruntime/core/providers/mkldnn/nn/pool.cc deleted file mode 100644 index 2c979145ae..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/pool.cc +++ /dev/null @@ -1,337 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#ifdef _WIN32 -#pragma warning(disable : 4244) -#endif - -#include "core/common/cpuid_info.h" -#include "core/providers/mkldnn/mkldnn_common.h" -#include "core/providers/mkldnn/nn/pool.h" -#include "core/providers/mkldnn/mkldnn_fwd.h" - -namespace onnxruntime { -namespace mkl_dnn { - -#define POOLING_KERNEL(op_name, data_type, pool_type, since_version, end_version) \ - ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \ - op_name, \ - kOnnxDomain, \ - since_version, \ - end_version, \ - data_type, \ - kMklDnnExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - Pool); - -POOLING_KERNEL(AveragePool, float, AveragePool, 7, 8) -POOLING_KERNEL(GlobalAveragePool, float, AveragePool, 1, 8) -POOLING_KERNEL(MaxPool, float, MaxPool<1>, 1, 7) -POOLING_KERNEL(MaxPool, float, MaxPool<8>, 8, 8) -POOLING_KERNEL(GlobalMaxPool, float, MaxPool<1>, 1, 8) - -namespace { -// Struct which encapsulates parameters for MKLDNN Pool primitive. -struct PoolParams { - std::string op_name; - // Pooling primitive needs version as part of key because there - // are multiple versions of mkldnn pool kernels. - std::string version; - mkldnn::memory::dims& src_dims; - mkldnn::memory::dims& dst_dims; - mkldnn::memory::dims& kernel; - mkldnn::memory::dims& strides; - mkldnn::memory::dims& padding_left; - mkldnn::memory::dims& padding_right; - bool count_include_pad; - - PoolParams(const std::string& op_name, const std::string& version, - mkldnn::memory::dims& src_dims, mkldnn::memory::dims& dst_dims, - mkldnn::memory::dims& kernel, mkldnn::memory::dims& strides, - mkldnn::memory::dims& padding_left, mkldnn::memory::dims& padding_right, - bool count_include_pad) - : op_name(op_name), - version(version), - src_dims(src_dims), - dst_dims(dst_dims), - kernel(kernel), - strides(strides), - padding_left(padding_left), - padding_right(padding_right), - count_include_pad(count_include_pad) {} - - // Used as the key for Pool Primitive Reuse Pool. - std::string ToString() const { - std::string key; - key.reserve(128); - key.append(op_name); - key.append(version); - AddDimsToKey(key, src_dims); - AddDimsToKey(key, dst_dims); - AddDimsToKey(key, kernel); - AddDimsToKey(key, strides); - AddDimsToKey(key, padding_left); - AddDimsToKey(key, padding_right); - key.append(count_include_pad ? "true" : "false"); - return key; - } -}; - -template -class PoolPrimitive : public PrimitiveBase { - public: - explicit PoolPrimitive(const PoolParams& params) - : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); - if (context_.pool_fwd == nullptr) { - Initialize(params); - } - } - - ~PoolPrimitive() = default; - - void Compute(const T* src_data, T* dst_data) { - context_.src_mem->set_data_handle(static_cast(const_cast(src_data))); - context_.dst_mem->set_data_handle(static_cast(dst_data)); - context_.stream->submit(context_.net); - - context_.src_mem->set_data_handle(nullptr); - context_.dst_mem->set_data_handle(nullptr); - return; - } - - mkldnn::memory::format GetSrcMemoryFormat() const { return context_.src_fmt; } - mkldnn::memory::format GetDstMemoryFormat() const { return context_.dst_fmt; } - - size_t GetSrcSize() const { return context_.src_size; } - size_t GetDstSize() const { return context_.dst_size; } - - // std::unique_ptr GetDstMemoryDesc() const { return context_.dst_md; } - mkldnn::pooling_forward::primitive_desc* GetPrimitiveDesc() const { - return context_.fwd_primitive_desc.get(); - } - - private: - struct PoolContext { - mkldnn::memory::format src_fmt; - mkldnn::memory::format dst_fmt; - - size_t src_size; - size_t dst_size; - - std::unique_ptr src_mem; - std::unique_ptr dst_mem; - - std::unique_ptr fwd_desc; - - std::unique_ptr src_md; - std::unique_ptr dst_md; - - std::unique_ptr fwd_primitive_desc; - - std::unique_ptr pool_fwd; - - std::unique_ptr stream; - std::vector net; - - PoolContext() - : src_fmt(mkldnn::memory::format::any), - dst_fmt(mkldnn::memory::format::any), - src_size(0), - dst_size(0), - src_mem(nullptr), - dst_mem(nullptr), - fwd_desc(nullptr), - src_md(nullptr), - fwd_primitive_desc(nullptr), - pool_fwd(nullptr), - stream(nullptr) {} - }; - - void Initialize(const PoolParams& params) { - bool is_2D = params.src_dims.size() == 4 ? true : false; - mkldnn::memory::format fmt = mkldnn::memory::format::any; - if (CPUIDInfo::GetCPUIDInfo().HasAVX512f()) { - fmt = is_2D ? mkldnn::memory::format::nChw16c : mkldnn::memory::format::nCdhw16c; - } else if (CPUIDInfo::GetCPUIDInfo().HasAVX2() && (params.src_dims[1] % 8 == 0)) { - fmt = is_2D ? mkldnn::memory::format::nChw8c : mkldnn::memory::format::ncdhw; - } else { - fmt = is_2D ? mkldnn::memory::format::nchw : mkldnn::memory::format::ncdhw; - } - context_.src_md.reset(new mkldnn::memory::desc( - {params.src_dims}, MklDnnType(), fmt)); - context_.dst_md.reset(new mkldnn::memory::desc( - {params.dst_dims}, MklDnnType(), mkldnn::memory::format::any)); - - mkldnn::algorithm algo = mkldnn::algorithm::pooling_max; - if (PoolType::type == onnxruntime::PoolType::kAveragePool) { - algo = mkldnn::algorithm::pooling_avg_exclude_padding; - if (params.count_include_pad) { - algo = mkldnn::algorithm::pooling_avg_include_padding; - } - } - context_.fwd_desc.reset(new mkldnn::pooling_forward::desc( - mkldnn::prop_kind::forward_inference, algo, - *context_.src_md, *context_.dst_md, - params.strides, params.kernel, - params.padding_left, params.padding_right, - mkldnn::padding_kind::zero)); - - context_.fwd_primitive_desc.reset(new mkldnn::pooling_forward::primitive_desc( - *context_.fwd_desc, cpu_engine_)); - - context_.src_fmt = static_cast( - context_.fwd_primitive_desc.get()->src_primitive_desc().desc().data.format); - - context_.dst_fmt = static_cast( - context_.fwd_primitive_desc.get()->dst_primitive_desc().desc().data.format); - - context_.src_size = context_.fwd_primitive_desc.get()->src_primitive_desc().get_size(); - context_.dst_size = context_.fwd_primitive_desc.get()->dst_primitive_desc().get_size(); - - context_.src_mem.reset( - new mkldnn::memory(context_.fwd_primitive_desc.get()->src_primitive_desc(), nullptr)); - context_.dst_mem.reset( - new mkldnn::memory(context_.fwd_primitive_desc.get()->dst_primitive_desc(), nullptr)); - context_.pool_fwd.reset( - new mkldnn::pooling_forward(*context_.fwd_primitive_desc, *context_.src_mem, *context_.dst_mem)); - context_.net.push_back(*context_.pool_fwd); - } - - PoolContext context_; - mkldnn::engine& cpu_engine_; -}; - -// Pool which allows for reuse of MKLDNN Pool primitives which are expensive to instantiate. -// To address thread safety, the primitives are stored in a map on thread local storage. -template -class PoolPrimitivePool : public PrimitivePool { - public: - static PoolPrimitive* Get(const PoolParams& params) { - PoolPrimitive* primitive = dynamic_cast*>( - PoolPrimitivePool::GetInstance().GetPrimitive(params.ToString())); - if (primitive == nullptr) { - auto pool_primitive = onnxruntime::make_unique>(params); - primitive = pool_primitive.get(); - PoolPrimitivePool::GetInstance().SetPrimitive(params.ToString(), std::move(pool_primitive)); - } - return primitive; - } - - private: - PoolPrimitivePool() = default; - ~PoolPrimitivePool() = default; - - static PoolPrimitivePool& GetInstance() { - static PoolPrimitivePool pool; - return pool; - } -}; -} // namespace - -template -Status Pool::Compute(OpKernelContext* context) const { - const Tensor* X = context->Input(0); - const TensorShape& x_shape = X->Shape(); - const auto& x_dims = x_shape.GetDims(); - - if (x_shape.NumDimensions() < 3) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Input dimension cannot be less than 3."); - } - - if (x_shape.NumDimensions() == 3) { - // Fall Back to CPU implementation. - return onnxruntime::Pool::Compute(context); - } - - std::vector kernel_shape = this->pool_attrs_.kernel_shape; - std::vector pads = this->pool_attrs_.pads; - std::vector strides = this->pool_attrs_.strides; - - if (this->pool_attrs_.global_pooling) { - kernel_shape.assign(x_dims.begin() + 2, x_dims.end()); - pads.assign(kernel_shape.size() * 2, 0); - strides.assign(kernel_shape.size(), 1); - } - - std::vector y_dims = this->pool_attrs_.SetOutputSize(x_shape, x_shape[1], &pads); - Tensor* Y = context->Output(0, TensorShape(y_dims)); - - size_t num_outputs = OpKernel::Node().OutputDefs().size(); - if (num_outputs == 2) { - Tensor* I = context->Output(1, TensorShape(y_dims)); - if (nullptr != I) { - return onnxruntime::Pool::Compute(context); - } - } - - const T* src_data = X->template Data(); - T* dst_data = Y->template MutableData(); - - mkldnn::memory::dims src_dims_mkl(x_dims.begin(), x_dims.end()); - mkldnn::memory::dims dst_dims_mkl(y_dims.begin(), y_dims.end()); - mkldnn::memory::dims kernel_mkl(kernel_shape.begin(), kernel_shape.end()); - mkldnn::memory::dims strides_mkl(strides.begin(), strides.end()); - mkldnn::memory::dims padding_left_mkl(pads.begin(), pads.begin() + (pads.size() / 2)); - mkldnn::memory::dims padding_right_mkl(pads.begin() + (pads.size() / 2), pads.end()); - - AllocatorPtr alloc; - ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&alloc)); - IAllocatorUniquePtr src_reorder_buffer; - IAllocatorUniquePtr dst_reorder_buffer; - - try { - PoolParams pool_params(this->op_name_, this->opset_version_, - src_dims_mkl, dst_dims_mkl, - kernel_mkl, strides_mkl, - padding_left_mkl, padding_right_mkl, - this->pool_attrs_.count_include_pad); - PoolPrimitive* pool_primitive = PoolPrimitivePool::Get(pool_params); - auto fwd_primitive_desc = pool_primitive->GetPrimitiveDesc(); - - mkldnn::engine& cpu_engine = GetEngine(); - mkldnn::memory::format mem_format = src_dims_mkl.size() == 5 ? mkldnn::memory::format::ncdhw : mkldnn::memory::format::nchw; - // Per ONNX spec, X (src) is NCHW and Y (dst) is NCHW - auto src_md = mkldnn::memory::desc(src_dims_mkl, MklDnnType(), mem_format); - auto dst_md = mkldnn::memory::desc(dst_dims_mkl, MklDnnType(), mem_format); - - // Reorder src memory layout if necessary. - if (src_md.data.format != pool_primitive->GetSrcMemoryFormat()) { - auto pd = mkldnn::memory::primitive_desc(src_md, cpu_engine); - mkldnn::memory src = mkldnn::memory(pd, (void*)src_data); - // allocate the size queried from memory primitive desc. it may not match tensor logical size due to - // mkldnn using padding to allow use of blocked format. - src_reorder_buffer = IAllocator::MakeUniquePtr(alloc, pool_primitive->GetSrcSize()); - mkldnn::memory dst = mkldnn::memory(fwd_primitive_desc->src_primitive_desc(), src_reorder_buffer.get()); - MemoryReorderParams params(src, dst); - DoReorder(params); - src_data = static_cast(dst.get_data_handle()); - } - - // Allocate dst buffer if reorder is necessary - if (dst_md.data.format != pool_primitive->GetDstMemoryFormat()) { - // allocate the size queried from memory primitive desc. it may not match tensor logical size due to - // mkldnn using padding to allow use of blocked format. - dst_reorder_buffer = IAllocator::MakeUniquePtr(alloc, pool_primitive->GetDstSize()); - dst_data = static_cast(dst_reorder_buffer.get()); - } - - pool_primitive->Compute(src_data, dst_data); - - // Reorder dst memory layout if necessary - if (dst_md.data.format != pool_primitive->GetDstMemoryFormat()) { - mkldnn::memory src = mkldnn::memory(fwd_primitive_desc->dst_primitive_desc(), (void*)dst_data); - auto pd = mkldnn::memory::primitive_desc(dst_md, cpu_engine); - mkldnn::memory dst = mkldnn::memory(pd, Y->template MutableData()); - MemoryReorderParams params(src, dst); - DoReorder(params); - } - } catch (const mkldnn::error& e) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); - } - - return Status::OK(); -} - -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/nn/pool.h b/onnxruntime/core/providers/mkldnn/nn/pool.h deleted file mode 100644 index 98ce4c16f1..0000000000 --- a/onnxruntime/core/providers/mkldnn/nn/pool.h +++ /dev/null @@ -1,29 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#pragma once -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/nn/pool.h" - -namespace onnxruntime { -namespace mkl_dnn { - -template -class Pool final : public onnxruntime::Pool { - public: - explicit Pool(const OpKernelInfo& info) : onnxruntime::Pool(info) { - // Since there are multiple versions of Pooling kernels, we need to use - // the opset version as part of the key for caching Pooling Primitives. - int start, end; - OpKernel::KernelDef().SinceVersion(&start, &end); - opset_version_ = std::to_string(start); - } - - Status Compute(OpKernelContext* context) const override; - - private: - std::string opset_version_; -}; - -} // namespace mkl_dnn -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_activations.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_activations.h index 406542fd15..cf704f4cf2 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_activations.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_activations.h @@ -27,7 +27,7 @@ class MklDnnRelu : public MklDnnKernel { OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& source_format) { + std::vector> &net_args) { Ort::CustomOpApi ort{*api}; int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; @@ -44,25 +44,32 @@ class MklDnnRelu : public MklDnnKernel { mkldnn::memory::dims dims(xdim); ort_source_format_ = GetSourceFormat(static_cast(xdim)); - source_format = ort_source_format_; - src_format_ = ort_source_format_; - x_shape = TensorShape(xshape, xdim); + x_shape = TensorShape(xshape, xdim); - mkldnn::memory::dims src_dims_mkl( + if (x_shape.NumDimensions() == 0) { + primitive_created_ = Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, "Shape of size zero " + x_shape.ToString()); + return primitive_created_; + } + + mkldnn::memory::dims src_dims( x_shape.GetDims().begin(), x_shape.GetDims().end()); + ort_source_desc_ = mkldnn::memory::desc( + {src_dims}, MklDnnType(), ort_source_format_); + source_desc_ = ort_source_desc_; src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), src_format_)); + {src_dims}, MklDnnType(), ort_source_format_)); src_mem_.reset( - new mkldnn::memory({*src_md_, cpu_engine}, nullptr)); + new mkldnn::memory({{src_dims}, MklDnnType(), ort_source_format_}, cpu_engine, nullptr)); } else { src_md_.reset( - new mkldnn::memory::desc(parents_[0].get()->primitive_dst_mem_.get()->get_primitive_desc().desc())); + new mkldnn::memory::desc(parents_[0].get()->primitive_dst_desc_)); src_mem_ = parents_[0].get()->primitive_dst_mem_; x_shape = parents_[0].get()->primitive_dst_shape_; - ort_source_format_ = source_format; - src_format_ = parents_[0].get()->primitive_dst_format_; + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; } primitive_dst_shape_ = TensorShape(x_shape); @@ -75,37 +82,38 @@ class MklDnnRelu : public MklDnnKernel { relu_fwd_pd_.reset(new mkldnn::eltwise_forward::primitive_desc( *fwd_desc_, cpu_engine)); - primitive_src_format_ = static_cast( - relu_fwd_pd_.get()->src_primitive_desc().desc().data.format); - primitive_dst_format_ = static_cast( - relu_fwd_pd_.get()->dst_primitive_desc().desc().data.format); + primitive_src_desc_ = relu_fwd_pd_.get()->src_desc(); + primitive_dst_desc_ = relu_fwd_pd_.get()->dst_desc(); if (mklnode_ptr_->output_index >= 0) { // last node of sub-graph. need to allocate memory for output_tensor - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { // reorder neded. Use primitive output as input to reorder and // allocate buffer for reorder output, final output of this subgraph - primitive_dst_mem_.reset(new mkldnn::memory(relu_fwd_pd_.get()->dst_primitive_desc())); + primitive_dst_mem_.reset(new mkldnn::memory(relu_fwd_pd_.get()->dst_desc(), cpu_engine)); } else { // Last node but re-order not needed. Allocate buffer to output of this node - primitive_dst_mem_.reset(new mkldnn::memory(relu_fwd_pd_.get()->dst_primitive_desc(), nullptr)); + primitive_dst_mem_.reset(new mkldnn::memory(relu_fwd_pd_.get()->dst_desc(), cpu_engine, nullptr)); } } else { // Intermediate node. Use mkldnn kernel internal memory for output and // use this as input to next node. - primitive_dst_mem_.reset(new mkldnn::memory(relu_fwd_pd_.get()->dst_primitive_desc())); + primitive_dst_mem_.reset(new mkldnn::memory(relu_fwd_pd_.get()->dst_desc(), cpu_engine)); } relu_fwd_.reset( - new mkldnn::eltwise_forward(*relu_fwd_pd_, *src_mem_, *primitive_dst_mem_)); + new mkldnn::eltwise_forward(*relu_fwd_pd_)); net.push_back(*relu_fwd_); + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); + if (mklnode_ptr_->output_index >= 0) { // one of the end nodes. Allocate output buffer memory and // reorder is necessary mkldnn::memory::data_type t = MklDnnType(); - InitDstReorderOutput(cpu_engine, t, net); + InitDstReorderOutput(cpu_engine, t, net, net_args); } return Status::OK(); @@ -114,7 +122,10 @@ class MklDnnRelu : public MklDnnKernel { Status Bind(const OrtCustomOpApi* api, OrtKernelContext* context) override { Ort::CustomOpApi ort{*api}; - int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; + if (primitive_created_ != Status::OK()) + return primitive_created_; + + int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; if (mklnode_ptr_->parent_nodes.empty()) { // Sub-graph's first node. Read input from input buffer @@ -129,7 +140,7 @@ class MklDnnRelu : public MklDnnKernel { OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); T* dst_data = ort.GetTensorMutableData(output); - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { reorder_dst_mem_to_->set_data_handle(dst_data); } else { primitive_dst_mem_->set_data_handle(dst_data); diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_batchnorm.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_batchnorm.h index 17728b927f..cfa843c320 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_batchnorm.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_batchnorm.h @@ -101,7 +101,7 @@ class MklDnnBatchNorm : public MklDnnKernel { OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& source_format) override { + std::vector>& net_args) override { Ort::CustomOpApi ort{*api}; int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; @@ -114,21 +114,25 @@ class MklDnnBatchNorm : public MklDnnKernel { auto xshape = tensor_shape.data(); auto xdim = tensor_shape.size(); mkldnn::memory::dims dims(xdim); - ort_source_format_ = GetSourceFormat(static_cast(xdim)); - source_format = ort_source_format_; - src_format_ = ort_source_format_; + x_shape = TensorShape(xshape, xdim); - mkldnn::memory::dims src_dims_mkl( + mkldnn::memory::dims src_dims( x_shape.GetDims().begin(), x_shape.GetDims().end()); + + ort_source_desc_ = mkldnn::memory::desc( + {src_dims}, MklDnnType(), ort_source_format_); + source_desc_ = ort_source_desc_; src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), src_format_)); + {src_dims}, MklDnnType(), ort_source_format_)); } else { - src_md_.reset(new mkldnn::memory::desc(parents_[0].get()->primitive_dst_mem_.get()->get_primitive_desc().desc())); + src_md_.reset( + new mkldnn::memory::desc(parents_[0].get()->primitive_dst_desc_)); x_shape = parents_[0].get()->primitive_dst_shape_; - ort_source_format_ = source_format; - src_format_ = parents_[0].get()->primitive_dst_format_; + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; } int num_dimensions = static_cast(x_shape.NumDimensions()); @@ -192,13 +196,13 @@ class MklDnnBatchNorm : public MklDnnKernel { primitive_dst_shape_.GetDims().begin(), primitive_dst_shape_.GetDims().end()); scale_shift_md_.reset(new mkldnn::memory::desc( - {2, scale_dims_mkl[0]}, MklDnnType(), mkldnn::memory::format::nc)); + {2, scale_dims_mkl[0]}, MklDnnType(), mkldnn::memory::format_tag::nc)); mean_md_.reset(new mkldnn::memory::desc( - {mean_dims_mkl}, MklDnnType(), mkldnn::memory::format::x)); + {mean_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::x)); var_md_.reset(new mkldnn::memory::desc( - {var_dims_mkl}, MklDnnType(), mkldnn::memory::format::x)); + {var_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::x)); primitive_dst_md_.reset(new mkldnn::memory::desc( - {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); // scale_shift_mem will allocate 2*C*sizeof(float) buffer // @@ -206,18 +210,18 @@ class MklDnnBatchNorm : public MklDnnKernel { new mkldnn::memory({*scale_shift_md_, cpu_engine})); mean_mem_.reset( - new mkldnn::memory({*mean_md_, cpu_engine}, nullptr)); + new mkldnn::memory(*mean_md_, cpu_engine, nullptr)); var_mem_.reset( - new mkldnn::memory({*var_md_, cpu_engine}, nullptr)); + new mkldnn::memory(*var_md_, cpu_engine, nullptr)); batchnorm_fwd_.reset(new mkldnn::batch_normalization_forward::desc( mkldnn::prop_kind::forward_inference, *src_md_, epsilon_, - mkldnn::batch_normalization_flag::use_scale_shift | - mkldnn::batch_normalization_flag::use_global_stats)); + mkldnn::normalization_flags::use_scale_shift | + mkldnn::normalization_flags::use_global_stats)); if (fuse_relu_) { mkldnn::primitive_attr attr; - attr.set_int_output_round_mode(mkldnn::round_mode::round_nearest); + // attr.set_int_output_round_mode(mkldnn::round_mode::round_nearest); // Execute RELU as Fuse PostOps const float ops_scale = 1.f; const float ops_alpha = 0.f; // relu negative slope @@ -235,44 +239,44 @@ class MklDnnBatchNorm : public MklDnnKernel { } // out format of this kernel - primitive_dst_format_ = static_cast( - batchnorm_fwd_pd_.get()->dst_primitive_desc().desc().data.format); - primitive_src_format_ = static_cast( - batchnorm_fwd_pd_.get()->dst_primitive_desc().desc().data.format); + primitive_dst_desc_ = static_cast( + batchnorm_fwd_pd_.get()->dst_desc()); + primitive_src_desc_ = static_cast( + batchnorm_fwd_pd_.get()->dst_desc()); if (mklnode_ptr_->parent_nodes.empty()) { src_mem_.reset( - new mkldnn::memory(batchnorm_fwd_pd_.get()->src_primitive_desc(), nullptr)); + new mkldnn::memory(batchnorm_fwd_pd_.get()->src_desc(), cpu_engine, nullptr)); } else { src_mem_ = parents_[0].get()->primitive_dst_mem_; } if (mklnode_ptr_->output_index >= 0) { // Use mkldnn's internal output buffer - if (primitive_dst_format_ != ort_source_format_) { - primitive_dst_mem_.reset(new mkldnn::memory(batchnorm_fwd_pd_->dst_primitive_desc())); + if (primitive_dst_desc_ != ort_source_desc_) { + primitive_dst_mem_.reset(new mkldnn::memory(batchnorm_fwd_pd_->dst_desc(), cpu_engine)); } else { - primitive_dst_mem_.reset(new mkldnn::memory(batchnorm_fwd_pd_->dst_primitive_desc(), nullptr)); + primitive_dst_mem_.reset(new mkldnn::memory(batchnorm_fwd_pd_->dst_desc(), cpu_engine, nullptr)); } } else { // last node of sub-graph. need to allocate memory for output_tensor - primitive_dst_mem_.reset(new mkldnn::memory(batchnorm_fwd_pd_->dst_primitive_desc())); + primitive_dst_mem_.reset(new mkldnn::memory(batchnorm_fwd_pd_->dst_desc(), cpu_engine)); } auto bn = mkldnn::batch_normalization_forward( - *batchnorm_fwd_pd_, - (const mkldnn::primitive::at)*src_mem_, - (const mkldnn::primitive::at)*mean_mem_, - (const mkldnn::primitive::at)*var_mem_, - (const mkldnn::memory)*scale_shift_mem_, - (const mkldnn::memory)*primitive_dst_mem_); + *batchnorm_fwd_pd_); net.push_back(bn); + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_MEAN, *mean_mem_}, + {MKLDNN_ARG_VARIANCE, *var_mem_}, + {MKLDNN_ARG_SCALE_SHIFT, *scale_shift_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); // Allocate dst buffer if reorder is necessary if (mklnode_ptr_->output_index >= 0) { // one of the end nodes. Allocate output buffer memory and // reorder is necessary mkldnn::memory::data_type t = MklDnnType(); - InitDstReorderOutput(cpu_engine, t, net); + InitDstReorderOutput(cpu_engine, t, net, net_args); } return Status::OK(); } @@ -330,7 +334,7 @@ class MklDnnBatchNorm : public MklDnnKernel { OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); T* dst_data = ort.GetTensorMutableData(output); - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { reorder_dst_mem_to_->set_data_handle(dst_data); } else { primitive_dst_mem_->set_data_handle(dst_data); diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv.h index 9e421ec365..2dd02ec6ad 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv.h @@ -74,8 +74,9 @@ class MklDnnConv : public MklDnnKernel { OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& source_format) override { + std::vector>& net_args) override { Ort::CustomOpApi ort{*api}; + stream_.reset(new mkldnn::stream(cpu_engine)); int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; const OrtValue* winput_tensor = ort.KernelContext_GetInput(context, input_index + 1); @@ -97,21 +98,18 @@ class MklDnnConv : public MklDnnKernel { ort.ReleaseTensorTypeAndShapeInfo(tensor_info); auto xshape = tensor_shape.data(); auto xdim = tensor_shape.size(); - + ort_source_format_ = mkldnn::memory::format_tag::any; x_shape = TensorShape(xshape, xdim); - - mkldnn::memory::dims src_dims_mkl(x_shape.GetDims().begin(), x_shape.GetDims().end()); - src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); } else { // get the output of previous node (mkldnn block propagation). // TODO Sourcenode will set src of this node. x_shape = parents_[0].get()->primitive_dst_shape_; - ort_source_format_ = source_format; - src_format_ = parents_[0].get()->primitive_dst_format_; - mkldnn::memory::dims src_dims_mkl(x_shape.GetDims().begin(), x_shape.GetDims().end()); - src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; + + mkldnn::memory::dims src_dims_mkl(x_shape.GetDims().begin(), x_shape.GetDims().end()); + src_md_.reset(new mkldnn::memory::desc(source_desc_)); } primitive_created_ = ValidateInputShape(x_shape, w_shape); @@ -168,7 +166,7 @@ class MklDnnConv : public MklDnnKernel { TensorShape output_shape = y_shape.Slice(2); mkldnn::memory::dims dst_dims_mkl(y_dims.begin(), y_dims.end()); primitive_dst_md_.reset(new mkldnn::memory::desc( - {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); mkldnn::memory::dims filter_dims_mkl; if (group_mkl == 1) { @@ -199,60 +197,65 @@ class MklDnnConv : public MklDnnKernel { bias_dims_mkl.assign(b_shape.GetDims().begin(), b_shape.GetDims().end()); } - auto fmt = mkldnn::memory::format::any; + auto src_format = mkldnn::memory::format_tag::any; if (kernel_rank == 1) { - fmt = mkldnn::memory::format::ncw; + src_format = mkldnn::memory::format_tag::ncw; if (group_mkl == 1) { - filter_format_ = mkldnn::memory::format::oiw; + filter_format_ = mkldnn::memory::format_tag::oiw; } else { - filter_format_ = mkldnn::memory::format::goiw; + filter_format_ = mkldnn::memory::format_tag::goiw; } } else if (kernel_rank == 2) { - fmt = mkldnn::memory::format::nchw; + src_format = mkldnn::memory::format_tag::nchw; if (group_mkl == 1) { - filter_format_ = mkldnn::memory::format::oihw; + filter_format_ = mkldnn::memory::format_tag::oihw; } else { - filter_format_ = mkldnn::memory::format::goihw; + filter_format_ = mkldnn::memory::format_tag::goihw; } } else { - fmt = mkldnn::memory::format::ncdhw; + src_format = mkldnn::memory::format_tag::ncdhw; if (group_mkl == 1) { - filter_format_ = mkldnn::memory::format::oidhw; + filter_format_ = mkldnn::memory::format_tag::oidhw; } else { - filter_format_ = mkldnn::memory::format::goidhw; + filter_format_ = mkldnn::memory::format_tag::goidhw; } } - if (src_format_ == mkldnn::memory::format::any) { - src_format_ = fmt; - ort_source_format_ = fmt; - source_format = fmt; + + if (mklnode_ptr_->parent_nodes.empty()) { + mkldnn::memory::dims src_dims_mkl(x_shape.GetDims().begin(), x_shape.GetDims().end()); + + ort_source_format_ = src_format; + src_md_.reset(new mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); + ort_source_desc_ = mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), src_format); + source_desc_ = mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), src_format); } // Set the memory descriptors to format::any to allow MKLDNN to decide what the optimal memory layout should be // for the computation given the input filter_md_.reset(new mkldnn::memory::desc( - {filter_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + {filter_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); if (!bias_dims_mkl.empty()) bias_md_.reset(new mkldnn::memory::desc( - {bias_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + {bias_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); + + mkldnn::memory::dims conv_zero_padding = {0, 0}; if (!bias_dims_mkl.empty()) { fwd_desc_.reset(new mkldnn::convolution_forward::desc( - mkldnn::prop_kind::forward_inference, mkldnn::convolution_direct, *src_md_, + mkldnn::prop_kind::forward_inference, mkldnn::algorithm::convolution_direct, *src_md_, *filter_md_, *bias_md_, *primitive_dst_md_, strides_mkl, dilations_mkl, padding_left_mkl, - padding_right_mkl, mkldnn::padding_kind::zero)); + padding_right_mkl)); } else { fwd_desc_.reset(new mkldnn::convolution_forward::desc( - mkldnn::prop_kind::forward_inference, mkldnn::convolution_direct, *src_md_, + mkldnn::prop_kind::forward_inference, mkldnn::algorithm::convolution_direct, *src_md_, *filter_md_, *primitive_dst_md_, strides_mkl, - dilations_mkl, padding_left_mkl, - padding_right_mkl, mkldnn::padding_kind::zero)); + dilations_mkl, padding_left_mkl, padding_right_mkl)); } if (fuse_relu_) { mkldnn::primitive_attr attr; - attr.set_int_output_round_mode(mkldnn::round_mode::round_nearest); + // attr.set_int_output_round_mode(mkldnn::round_mode::round_nearest); // Execute RELU as Fuse PostOps const float ops_scale = 1.f; const float ops_alpha = 0.f; // relu negative slope @@ -268,37 +271,38 @@ class MklDnnConv : public MklDnnKernel { *fwd_desc_, cpu_engine)); } - primitive_src_format_ = static_cast( - conv_fwd_pd_.get()->src_primitive_desc().desc().data.format); + primitive_src_desc_ = static_cast( + conv_fwd_pd_.get()->src_desc()); - mkldnn_filter_format_ = static_cast( - conv_fwd_pd_.get()->weights_primitive_desc().desc().data.format); + filter_desc_= static_cast( + conv_fwd_pd_.get()->weights_desc()); - primitive_dst_format_ = static_cast( - conv_fwd_pd_.get()->dst_primitive_desc().desc().data.format); + primitive_dst_desc_ = static_cast( + conv_fwd_pd_.get()->dst_desc()); - src_size_ = conv_fwd_pd_.get()->src_primitive_desc().get_size(); - filter_size_ = conv_fwd_pd_.get()->weights_primitive_desc().get_size(); - dst_size_ = conv_fwd_pd_.get()->dst_primitive_desc().get_size(); + src_size_ = conv_fwd_pd_.get()->src_desc().get_size(); + filter_size_ = conv_fwd_pd_.get()->weights_desc().get_size(); + dst_size_ = conv_fwd_pd_.get()->dst_desc().get_size(); filter_mem_.reset( - new mkldnn::memory(conv_fwd_pd_.get()->weights_primitive_desc(), nullptr)); + new mkldnn::memory(conv_fwd_pd_.get()->weights_desc(), cpu_engine, nullptr)); - if (primitive_src_format_ != src_format_) { - mkldnn::memory::dims src_dims_mkl(x_shape.GetDims().begin(), x_shape.GetDims().end()); - auto src_md = mkldnn::memory::desc(src_dims_mkl, MklDnnType(), src_format_); - auto pd = mkldnn::memory::primitive_desc(src_md, cpu_engine); + if (primitive_src_desc_ != source_desc_) { + mkldnn::memory::dims src_dims(x_shape.GetDims().begin(), x_shape.GetDims().end()); + auto pd = mkldnn::memory::desc({{src_dims}, MklDnnType(), ort_source_format_}); if (mklnode_ptr_->parent_nodes.empty()) - src_mem_from_.reset(new mkldnn::memory(pd, nullptr)); + src_mem_from_.reset(new mkldnn::memory(pd, cpu_engine, nullptr)); else src_mem_from_ = parents_[0].get()->primitive_dst_mem_; - src_mem_.reset(new mkldnn::memory(conv_fwd_pd_->src_primitive_desc(), nullptr)); + src_mem_.reset(new mkldnn::memory(conv_fwd_pd_->src_desc(), cpu_engine, nullptr)); net.push_back(mkldnn::reorder(*src_mem_from_, *src_mem_)); + net_args.push_back({{MKLDNN_ARG_FROM, *src_mem_from_}, + {MKLDNN_ARG_TO, *src_mem_}}); } else { if (mklnode_ptr_->parent_nodes.empty()) { - src_mem_.reset(new mkldnn::memory(conv_fwd_pd_->src_primitive_desc(), nullptr)); + src_mem_.reset(new mkldnn::memory(conv_fwd_pd_->src_desc(), cpu_engine, nullptr)); } else { src_mem_ = parents_[0].get()->primitive_dst_mem_; } @@ -306,31 +310,36 @@ class MklDnnConv : public MklDnnKernel { if (mklnode_ptr_->output_index >= 0) { // Use mkldnn's internal output buffer - if (primitive_dst_format_ != ort_source_format_) { - primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_primitive_desc())); + if (primitive_dst_desc_ != ort_source_desc_) { + primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_desc(), cpu_engine)); } else { - primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_primitive_desc(), nullptr)); + primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_desc(), cpu_engine, nullptr)); } } else { // last node of sub-graph. need to allocate memory for output_tensor - primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_primitive_desc())); + primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_desc(), cpu_engine)); } if (!bias_dims_mkl.empty()) { - bias_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->bias_primitive_desc(), nullptr)); - conv_fwd_.reset(new mkldnn::convolution_forward(*conv_fwd_pd_, *src_mem_, *filter_mem_, - *bias_mem_, *primitive_dst_mem_)); + bias_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->bias_desc(), cpu_engine, nullptr)); + conv_fwd_.reset(new mkldnn::convolution_forward(*conv_fwd_pd_)); + net.push_back(*conv_fwd_); + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_WEIGHTS, *filter_mem_}, + {MKLDNN_ARG_BIAS, *bias_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); } else { - conv_fwd_.reset(new mkldnn::convolution_forward(*conv_fwd_pd_, *src_mem_, - *filter_mem_, *primitive_dst_mem_)); + conv_fwd_.reset(new mkldnn::convolution_forward(*conv_fwd_pd_)); + net.push_back(*conv_fwd_); + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_WEIGHTS, *filter_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); } - net.push_back(*conv_fwd_); - if (mklnode_ptr_->output_index >= 0) { // one of the end nodes. Allocate output buffer memory and // reorder is necessary mkldnn::memory::data_type t = MklDnnType(); - InitDstReorderOutput(cpu_engine, t, net); + InitDstReorderOutput(cpu_engine, t, net, net_args); } primitive_created_ = Status::OK(); return primitive_created_; @@ -367,17 +376,16 @@ class MklDnnConv : public MklDnnKernel { std::shared_ptr filter_dst_mem = provider_->GetWeightsMemoryBuffer(mklnode_ptr_->weight_name); if (filter_dst_mem == nullptr) { - auto pd = mkldnn::memory::primitive_desc( - mkldnn::memory::desc(filter_dims_mkl, MklDnnType(), filter_format_), cpu_engine); - mkldnn::memory src = mkldnn::memory(pd, (void*)filter_data); + mkldnn::memory src = mkldnn::memory({{filter_dims_mkl}, MklDnnType(), filter_format_}, cpu_engine, (void*)filter_data); IAllocatorUniquePtr filter_reorder_buffer = IAllocator::MakeUniquePtr(alloc_, filter_size_); filter_dst_mem.reset( - new mkldnn::memory(conv_fwd_pd_->weights_primitive_desc(), filter_reorder_buffer.get())); + new mkldnn::memory(conv_fwd_pd_->weights_desc(), cpu_engine, filter_reorder_buffer.get())); + + mkldnn::reorder(src, *filter_dst_mem) + .execute(cpu_engine, src, *filter_dst_mem); - MemoryReorderParams params(src, *filter_dst_mem); - DoReorder(params); - provider_->SaveAllocatedMemory(std::move(filter_reorder_buffer)); + provider_->SaveAllocatedMemory(std::move(filter_reorder_buffer)); filter_data = static_cast(filter_dst_mem->get_data_handle()); provider_->SetWeightsMemoryBuffer(mklnode_ptr_->weight_name, filter_dst_mem); } @@ -417,7 +425,7 @@ class MklDnnConv : public MklDnnKernel { bias_mem_->set_data_handle(static_cast(const_cast(bias_data))); } - if (primitive_src_format_ != src_format_) { + if (primitive_src_desc_ != source_desc_) { if (mklnode_ptr_->parent_nodes.empty()) { const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index); const T* src_data = const_cast(ort.GetTensorData(input_tensor)); @@ -426,7 +434,7 @@ class MklDnnConv : public MklDnnKernel { src_mem_from_ = parents_[0].get()->primitive_dst_mem_; } - auto src_size = conv_fwd_pd_.get()->src_primitive_desc().get_size(); + auto src_size = conv_fwd_pd_.get()->src_desc().get_size(); src_reorder_buffer_ = IAllocator::MakeUniquePtr(alloc_, src_size); src_mem_->set_data_handle(src_reorder_buffer_.get()); } else { @@ -445,7 +453,7 @@ class MklDnnConv : public MklDnnKernel { OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); T* dst_data = ort.GetTensorMutableData(output); - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { reorder_dst_mem_to_->set_data_handle(dst_data); } else { primitive_dst_mem_->set_data_handle(dst_data); @@ -514,8 +522,8 @@ class MklDnnConv : public MklDnnKernel { } private: - mkldnn::memory::format mkldnn_filter_format_; - mkldnn::memory::format filter_format_; + mkldnn::memory::desc filter_desc_; + mkldnn::memory::format_tag filter_format_; std::shared_ptr src_mem_from_; std::unique_ptr src_mem_to_; @@ -623,6 +631,8 @@ class MklDnnConv : public MklDnnKernel { } private: + + std::unique_ptr stream_; std::vector kernel_shape_; // must use ComputeKernelShape(...), instead of kernel_shape_ AutoPadType auto_pad_; int64_t group_; diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv_batchnorm.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv_batchnorm.h new file mode 100644 index 0000000000..eb18c12f7e --- /dev/null +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_conv_batchnorm.h @@ -0,0 +1,696 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once +#include "mkldnn_types.h" +#include "core/framework/op_kernel.h" +#include "core/providers/mkldnn/mkldnn_fwd.h" +#include "core/providers/cpu/nn/autopad_type.h" +#include "core/providers/mkldnn/mkldnn_execution_provider.h" +#include "core/providers/mkldnn/subgraph/mkldnn_kernel.h" +#include "core/util/math.h" + +namespace onnxruntime { +namespace mkl_dnn { + +template +class MklDnnConvBatchNorm : public MklDnnKernel { + public: + MklDnnConvBatchNorm(const MklDnnNode& node, + MKLDNNExecutionProvider* provider, + const NodeAttributes& attributes, + const std::string attributes_prefix = "") : MklDnnKernel(node, provider) { + ReadAttributes(attributes, attributes_prefix); + } + + Status CreatePrimitives(const OrtCustomOpApi* api, + OrtKernelContext* context, + mkldnn::engine& cpu_engine, + std::vector& net, + std::vector>& net_args) override { + Ort::CustomOpApi ort{*api}; + stream_.reset(new mkldnn::stream(cpu_engine)); + int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; + const OrtValue* winput_tensor = ort.KernelContext_GetInput(context, input_index + 1); + auto wtensor_info = ort.GetTensorTypeAndShape(winput_tensor); + auto wtensor_shape = ort.GetTensorShape(wtensor_info); + ort.ReleaseTensorTypeAndShapeInfo(wtensor_info); + auto wshape = wtensor_shape.data(); + auto wdim = wtensor_shape.size(); + + TensorShape w_shape(wshape, wdim); + const int group_mkl = static_cast(group_); + + TensorShape x_shape; + // std::unique_ptr x_shape; + if (mklnode_ptr_->parent_nodes.empty()) { + const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index); + auto tensor_info = ort.GetTensorTypeAndShape(input_tensor); + auto tensor_shape = ort.GetTensorShape(tensor_info); + ort.ReleaseTensorTypeAndShapeInfo(tensor_info); + auto xshape = tensor_shape.data(); + auto xdim = tensor_shape.size(); + ort_source_format_ = mkldnn::memory::format_tag::any; + x_shape = TensorShape(xshape, xdim); + } else { + // get the output of previous node (mkldnn block propagation). + // TODO Sourcenode will set src of this node. + x_shape = parents_[0].get()->primitive_dst_shape_; + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; + } + + primitive_created_ = ValidateInputShape(x_shape, w_shape); + if (!primitive_created_.IsOK()) + return primitive_created_; + + std::vector kernel_shape; + primitive_created_ = ComputeKernelShape(w_shape, kernel_shape); + if (!primitive_created_.IsOK()) + return primitive_created_; + + const size_t kernel_rank = kernel_shape.size(); + + if (kernel_rank + 2 != wdim) { + primitive_created_ = ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "kernel_shape num_dims is not compatible with W num_dims.", + " kernel_shape: ", TensorShape(kernel_shape).ToString().c_str(), + " W: ", w_shape.ToString().c_str()); + return primitive_created_; + } + + for (size_t i = 0; i < kernel_rank; ++i) { + if (kernel_shape[i] != w_shape[i + 2]) { + primitive_created_ = ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "kernel_shape is not compatible with W shape.", + " kernel_shape: ", TensorShape(kernel_shape).ToString().c_str(), + " W: ", w_shape.ToString().c_str()); + return primitive_created_; + } + } + + std::vector pads(pads_); + if (pads.empty()) { + pads.resize(kernel_rank * 2, 0); + } + std::vector dilations(dilations_); + if (dilations.empty()) { + dilations.resize(kernel_rank, 1); + } + std::vector strides(strides_); + if (strides.empty()) { + strides.resize(kernel_rank, 1); + } + + const int64_t N = x_shape[0]; + const int64_t M = w_shape[0]; + std::vector y_dims; + y_dims.insert(y_dims.begin(), {N, M}); + TensorShape input_shape = x_shape.Slice(2); + primitive_created_ = InferOutputShape(input_shape, kernel_shape, strides, dilations, &pads, &y_dims); + if (!primitive_created_.IsOK()) + return primitive_created_; + + TensorShape y_shape(y_dims); + primitive_dst_shape_ = TensorShape(y_dims); + TensorShape output_shape = y_shape.Slice(2); + mkldnn::memory::dims dst_dims_mkl(y_dims.begin(), y_dims.end()); + primitive_dst_md_.reset(new mkldnn::memory::desc( + {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); + + mkldnn::memory::dims filter_dims_mkl; + if (group_mkl == 1) { + filter_dims_mkl.assign(w_shape.GetDims().begin(), w_shape.GetDims().end()); + } else { + filter_dims_mkl.assign({group_mkl, + static_cast(w_shape[0] / group_mkl)}); + filter_dims_mkl.insert(filter_dims_mkl.end(), w_shape.GetDims().begin() + 1, w_shape.GetDims().end()); + } + mkldnn::memory::dims strides_mkl(strides.begin(), strides.end()); + mkldnn::memory::dims dilations_mkl(dilations.begin(), dilations.end()); + // mkldnn dilations start from 0 so we need to subtract 1 from each dim. + for (size_t dim = 0; dim < kernel_rank; dim++) { + dilations_mkl[dim] -= 1; + } + + mkldnn::memory::dims padding_left_mkl(pads.begin(), pads.begin() + kernel_rank); + mkldnn::memory::dims padding_right_mkl(pads.begin() + kernel_rank, pads.end()); + mkldnn::memory::dims bias_dims_mkl; + + int batchNormIndix = (mklnode_ptr_->num_inputs == 7) ? input_index + 3 : input_index + 2; + if (mklnode_ptr_->num_inputs == 7) { + const OrtValue* binput_tensor = ort.KernelContext_GetInput(context, input_index + 2); + auto btensor_info = ort.GetTensorTypeAndShape(binput_tensor); + auto btensor_shape = ort.GetTensorShape(btensor_info); + ort.ReleaseTensorTypeAndShapeInfo(btensor_info); + auto bshape = btensor_shape.data(); + auto bdim = btensor_shape.size(); + TensorShape b_shape(bshape, bdim); + bias_dims_mkl.assign(b_shape.GetDims().begin(), b_shape.GetDims().end()); + } else { + const OrtValue* b_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix + 1); + auto b_tensor_info = ort.GetTensorTypeAndShape(b_input_tensor); + auto b_tensor_shape = ort.GetTensorShape(b_tensor_info); + ort.ReleaseTensorTypeAndShapeInfo(b_tensor_info); + auto bshape = b_tensor_shape.data(); + auto bdim = b_tensor_shape.size(); + TensorShape b_shape(bshape, bdim); + bias_dims_mkl.assign(b_shape.GetDims().begin(), b_shape.GetDims().end()); + } + + auto src_format = mkldnn::memory::format_tag::any; + if (kernel_rank == 1) { + src_format = mkldnn::memory::format_tag::ncw; + if (group_mkl == 1) { + filter_format_ = mkldnn::memory::format_tag::oiw; + } else { + filter_format_ = mkldnn::memory::format_tag::goiw; + } + } else if (kernel_rank == 2) { + src_format = mkldnn::memory::format_tag::nchw; + if (group_mkl == 1) { + filter_format_ = mkldnn::memory::format_tag::oihw; + } else { + filter_format_ = mkldnn::memory::format_tag::goihw; + } + } else { + src_format = mkldnn::memory::format_tag::ncdhw; + if (group_mkl == 1) { + filter_format_ = mkldnn::memory::format_tag::oidhw; + } else { + filter_format_ = mkldnn::memory::format_tag::goidhw; + } + } + + mkldnn::memory::dims src_dims_mkl(x_shape.GetDims().begin(), x_shape.GetDims().end()); + if (mklnode_ptr_->parent_nodes.empty()) { + ort_source_format_ = src_format; + ort_source_desc_ = mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), src_format); + source_desc_ = mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), src_format); + } + + src_md_.reset(new mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); + + // Set the memory descriptors to format::any to allow MKLDNN to decide what the optimal memory layout should be + // for the computation given the input + filter_md_.reset(new mkldnn::memory::desc( + {filter_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); + bias_md_.reset(new mkldnn::memory::desc( + {bias_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); + + mkldnn::memory::dims conv_zero_padding = {0, 0}; + + fwd_desc_.reset(new mkldnn::convolution_forward::desc( + mkldnn::prop_kind::forward_inference, mkldnn::algorithm::convolution_direct, *src_md_, + *filter_md_, *bias_md_, *primitive_dst_md_, + strides_mkl, dilations_mkl, padding_left_mkl, + padding_right_mkl)); + + + if (fuse_relu_) { + mkldnn::primitive_attr attr; + // attr.set_int_output_round_mode(mkldnn::round_mode::round_nearest); + // Execute RELU as Fuse PostOps + const float ops_scale = 1.f; + const float ops_alpha = 0.f; // relu negative slope + const float ops_beta = 0.f; + mkldnn::post_ops ops; + ops.append_eltwise(ops_scale, mkldnn::algorithm::eltwise_relu, ops_alpha, ops_beta); + attr.set_post_ops(ops); + + conv_fwd_pd_.reset(new mkldnn::convolution_forward::primitive_desc( + *fwd_desc_, attr, cpu_engine)); + } else { + conv_fwd_pd_.reset(new mkldnn::convolution_forward::primitive_desc( + *fwd_desc_, cpu_engine)); + } + + primitive_src_desc_ = static_cast( + conv_fwd_pd_.get()->src_desc()); + + filter_desc_ = static_cast( + conv_fwd_pd_.get()->weights_desc()); + + primitive_dst_desc_ = static_cast( + conv_fwd_pd_.get()->dst_desc()); + + src_size_ = conv_fwd_pd_.get()->src_desc().get_size(); + filter_size_ = conv_fwd_pd_.get()->weights_desc().get_size(); + dst_size_ = conv_fwd_pd_.get()->dst_desc().get_size(); + + filter_mem_.reset( + new mkldnn::memory(conv_fwd_pd_.get()->weights_desc(), cpu_engine, nullptr)); + + if (primitive_src_desc_ != source_desc_) { + mkldnn::memory::dims src_dims(x_shape.GetDims().begin(), x_shape.GetDims().end()); + auto pd = mkldnn::memory::desc({{src_dims}, MklDnnType(), ort_source_format_}); + + if (mklnode_ptr_->parent_nodes.empty()) + src_mem_from_.reset(new mkldnn::memory(pd, cpu_engine, nullptr)); + else + src_mem_from_ = parents_[0].get()->primitive_dst_mem_; + + src_mem_.reset(new mkldnn::memory(conv_fwd_pd_->src_desc(), cpu_engine, nullptr)); + net.push_back(mkldnn::reorder(*src_mem_from_, *src_mem_)); + net_args.push_back({{MKLDNN_ARG_FROM, *src_mem_from_}, + {MKLDNN_ARG_TO, *src_mem_}}); + } else { + if (mklnode_ptr_->parent_nodes.empty()) { + src_mem_.reset(new mkldnn::memory(conv_fwd_pd_->src_desc(), cpu_engine, nullptr)); + } else { + src_mem_ = parents_[0].get()->primitive_dst_mem_; + } + } + + if (mklnode_ptr_->output_index >= 0) { + // Use mkldnn's internal output buffer + if (primitive_dst_desc_ != ort_source_desc_) { + primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_desc(), cpu_engine)); + } else { + primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_desc(), cpu_engine, nullptr)); + } + } else { + // last node of sub-graph. need to allocate memory for output_tensor + primitive_dst_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->dst_desc(), cpu_engine)); + } + + bias_mem_.reset(new mkldnn::memory(conv_fwd_pd_.get()->bias_desc(), cpu_engine, nullptr)); + conv_fwd_.reset(new mkldnn::convolution_forward(*conv_fwd_pd_)); + net.push_back(*conv_fwd_); + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_WEIGHTS, *filter_mem_}, + {MKLDNN_ARG_BIAS, *bias_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); + + if (mklnode_ptr_->output_index >= 0) { + // one of the end nodes. Allocate output buffer memory and + // reorder is necessary + mkldnn::memory::data_type t = MklDnnType(); + InitDstReorderOutput(cpu_engine, t, net, net_args); + } + primitive_created_ = Status::OK(); + return primitive_created_; + } + + void GamaInverseVariance(const OrtCustomOpApi* api, OrtKernelContext* context, std::vector& inv_scale_factor, size_t O) { + Ort::CustomOpApi ort{*api}; + int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; + int batchNormIndix = (mklnode_ptr_->num_inputs == 7) ? input_index + 3 : input_index + 2; + const OrtValue* scale_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix); + const T* bn_scale_data = reinterpret_cast(ort.GetTensorData(scale_input_tensor)); + const OrtValue* var_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix + 3); + const T* bn_var_data = reinterpret_cast(ort.GetTensorData(var_input_tensor)); + + std::vector inv_scale; + inv_scale.assign(static_cast(O), 0.0f); + + float* data = inv_scale_factor.data(); + for (size_t i = 0; i < O; i++) { + data[i] = bn_scale_data[i] / std::sqrt(bn_var_data[i] + epsilon_); + } + } + + void WeightsScaleByAxix(const OrtCustomOpApi* api, OrtKernelContext* context, + std::vector& weights_scaled, std::vector inv_scale, + TensorShape& W) { + Ort::CustomOpApi ort{*api}; + int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; + const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index + 1); + const auto& w_dims = W.GetDims(); + const T* filter_data = const_cast(ort.GetTensorData(input_tensor)); + + int batchNormIndix = (mklnode_ptr_->num_inputs == 7) ? input_index + 3 : input_index + 2; + const OrtValue* scale_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix); + auto bn_tensor_info = ort.GetTensorTypeAndShape(scale_input_tensor); + auto bn_tensor_shape = ort.GetTensorShape(bn_tensor_info); + ort.ReleaseTensorTypeAndShapeInfo(bn_tensor_info); + auto bn_shape = bn_tensor_shape.data(); + auto bn_dim = bn_tensor_shape.size(); + TensorShape bn_scale_shape(bn_shape, bn_dim); + const auto& bn_dims = bn_scale_shape.GetDims(); + + int64_t num = 1; + int axis = 1; + for (size_t k = axis; k < w_dims.size(); k++) { + num *= w_dims[k]; + } + + int64_t w_size = std::accumulate(w_dims.begin(), w_dims.end(), static_cast(1), std::multiplies{}); + int64_t bn_scale_size = std::accumulate(bn_dims.begin(), bn_dims.end(), static_cast(1), std::multiplies{}); + int64_t n = w_size / num; + + float* w_scale_data = weights_scaled.data(); + for (auto i = 0; i < n; i++) { + int index = bn_scale_size == 1 ? 0 : i; + for (int64_t j = 0; j < num; j++) { + w_scale_data[i * num + j] = filter_data[i * num + j] * inv_scale[index]; + } + } + } + + virtual void ReorderWeights(const OrtCustomOpApi* api, OrtKernelContext* context, mkldnn::engine& cpu_engine) override { + Ort::CustomOpApi ort{*api}; + int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; + const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index + 1); + auto tensor_info = ort.GetTensorTypeAndShape(input_tensor); + auto tensor_shape = ort.GetTensorShape(tensor_info); + ort.ReleaseTensorTypeAndShapeInfo(tensor_info); + auto xshape = tensor_shape.data(); + auto xdim = tensor_shape.size(); + TensorShape W(xshape, xdim); + + const int group_mkl = static_cast(group_); + mkldnn::memory::dims filter_dims_mkl; + if (group_mkl == 1) { + filter_dims_mkl.assign(W.GetDims().begin(), W.GetDims().end()); + } else { + filter_dims_mkl.assign({group_mkl, + static_cast(W[0] / group_mkl)}); + filter_dims_mkl.insert(filter_dims_mkl.end(), W.GetDims().begin() + 1, W.GetDims().end()); + } + + std::vector inv_scale_factor; + const auto& w_dims = W.GetDims(); + const size_t O = w_dims[0]; + inv_scale_factor.assign(static_cast(O), 0.0f); + GamaInverseVariance(api, context, inv_scale_factor, O); + + std::vector weights_scaled_by_axis; + weights_scaled_by_axis.assign(static_cast(O), 0.0f); + auto w_size = std::accumulate(w_dims.begin(), w_dims.end(), static_cast(1), std::multiplies{}); + weights_scaled_by_axis.assign(static_cast(w_size), 0.0f); + WeightsScaleByAxix(api, context, weights_scaled_by_axis, inv_scale_factor, W); + + { + // lock to make sure reordering is done only once + std::lock_guard lock(provider_->GetMutex()); + std::shared_ptr filter_dst_mem = provider_->GetWeightsMemoryBuffer(mklnode_ptr_->weight_name); + + if (filter_dst_mem == nullptr) { + mkldnn::memory src = mkldnn::memory({{filter_dims_mkl}, MklDnnType(), filter_format_}, cpu_engine, (void*)weights_scaled_by_axis.data()); + IAllocatorUniquePtr filter_reorder_buffer = + IAllocator::MakeUniquePtr(alloc_, filter_size_); + filter_dst_mem.reset( + new mkldnn::memory(conv_fwd_pd_->weights_desc(), cpu_engine, filter_reorder_buffer.get())); + + mkldnn::reorder(src, *filter_dst_mem) + .execute(cpu_engine, src, *filter_dst_mem); + + provider_->SaveAllocatedMemory(std::move(filter_reorder_buffer)); + provider_->SetWeightsMemoryBuffer(mklnode_ptr_->weight_name, filter_dst_mem); + } + + std::shared_ptr bias_mem = provider_->GetBiasMemoryBuffer(mklnode_ptr_->weight_name); + if (bias_mem == nullptr) { + auto bias_size = conv_fwd_pd_.get()->bias_desc().get_size(); + IAllocatorUniquePtr bias_buffer = + IAllocator::MakeUniquePtr(alloc_, bias_size); + bias_mem.reset( + new mkldnn::memory(conv_fwd_pd_->bias_desc(), cpu_engine, bias_buffer.get())); + float* bias_buffer_data = static_cast(bias_buffer.get()); + if (mklnode_ptr_->num_inputs == 7) { + const OrtValue* conv_bias_tensor = ort.KernelContext_GetInput(context, input_index + 2); + const T* conv_bias_data = reinterpret_cast(ort.GetTensorData(conv_bias_tensor)); + + int batchNormIndix = input_index + 3; + const OrtValue* b_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix + 1); + const T* bn_b_data = reinterpret_cast(ort.GetTensorData(b_input_tensor)); + + const OrtValue* mean_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix + 2); + const T* bn_mean_data = reinterpret_cast(ort.GetTensorData(mean_input_tensor)); + + //norm_bias = (conv_bias_arr - mean_arr) * inv_std + bn_b_arr; + for (size_t j = 0; j < O; j++) { + bias_buffer_data[j] = (conv_bias_data[j] - bn_mean_data[j]) * inv_scale_factor[j] + bn_b_data[j]; + } + } else { + // norm_bias = inv_std * (mean_arr - bn_b_arr); + int batchNormIndix = (mklnode_ptr_->num_inputs == 7) ? input_index + 3 : input_index + 2; + const OrtValue* b_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix + 1); + const T* bn_b_data = reinterpret_cast(ort.GetTensorData(b_input_tensor)); + const OrtValue* mean_input_tensor = ort.KernelContext_GetInput(context, batchNormIndix + 2); + const T* bn_mean_data = reinterpret_cast(ort.GetTensorData(mean_input_tensor)); + + for (size_t j = 0; j < O; j++) { + bias_buffer_data[j] = bn_b_data[j] - (bn_mean_data[j] * inv_scale_factor[j]); + } + } + + provider_->SaveAllocatedBiasMemory(std::move(bias_buffer)); + provider_->SetBiasMemoryBuffer(mklnode_ptr_->weight_name, bias_mem); + } + } + } + + Status Bind(const OrtCustomOpApi* api, OrtKernelContext* context) override { + Ort::CustomOpApi ort{*api}; + + int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; + if (!primitive_created_.IsOK()) { + // abort as MKLDNN cannot execute this. but + // ORT try to delete output_tensor buffer data. allocate memory so that it can delete + // fix for test_averagepool_1d_default node test + //auto xshape = input_tensors[input_index].shape; + //auto xdim = input_tensors[input_index].ndim; + //AllocateOutputTensor(output_tensors, mklnode_ptr_->output_index, xshape, xdim, input_tensors[0].dtype); + return primitive_created_; + } + const OrtValue* winput_tensor = ort.KernelContext_GetInput(context, input_index + 1); + const T* filter_data = const_cast(ort.GetTensorData(winput_tensor)); + + std::shared_ptr filter_dst_mem = provider_->GetWeightsMemoryBuffer(mklnode_ptr_->weight_name); + if (filter_dst_mem == nullptr) { + ReorderWeights(api, context, GetEngine()); + filter_dst_mem = provider_->GetWeightsMemoryBuffer(mklnode_ptr_->weight_name); + } + filter_data = static_cast(filter_dst_mem->get_data_handle()); + filter_mem_->set_data_handle(static_cast(const_cast(filter_data))); + + std::shared_ptr bias_mem = provider_->GetBiasMemoryBuffer(mklnode_ptr_->weight_name); + const T* bias_data = static_cast(bias_mem->get_data_handle()); + bias_mem_->set_data_handle(static_cast(const_cast(bias_data))); + + if (primitive_src_desc_ != source_desc_) { + if (mklnode_ptr_->parent_nodes.empty()) { + const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index); + const T* src_data = const_cast(ort.GetTensorData(input_tensor)); + src_mem_from_->set_data_handle(static_cast(const_cast(src_data))); + } else { + src_mem_from_ = parents_[0].get()->primitive_dst_mem_; + } + + auto src_size = conv_fwd_pd_.get()->src_desc().get_size(); + src_reorder_buffer_ = IAllocator::MakeUniquePtr(alloc_, src_size); + src_mem_->set_data_handle(src_reorder_buffer_.get()); + } else { + if (mklnode_ptr_->parent_nodes.empty()) { + const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index); + const T* src_data = const_cast(ort.GetTensorData(input_tensor)); + src_mem_->set_data_handle(static_cast(const_cast(src_data))); + } else { + src_mem_ = parents_[0].get()->primitive_dst_mem_; + } + } + + if (mklnode_ptr_->output_index >= 0) { + auto& y_dims = primitive_dst_shape_.GetDims(); + // Allocate memory for output bufffer + OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); + T* dst_data = ort.GetTensorMutableData(output); + + if (primitive_dst_desc_ != ort_source_desc_) { + reorder_dst_mem_to_->set_data_handle(dst_data); + } else { + primitive_dst_mem_->set_data_handle(dst_data); + } + } + return Status::OK(); + } + + private: + void ReadAttributes(const 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(); + } + 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; + 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; + 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; + if (GetIntsAttr(proto, pads_) == Status::OK()) + attr_read = true; + } + if (!attr_read) { + pads_.resize(kernel_shape_.size() * 2, 0); + } + + attr_read = false; + attr = attributes.find(attributes_prefix + "dilations"); + if (attr != attributes.end()) { + ONNX_NAMESPACE::AttributeProto proto = attr->second; + if (GetIntsAttr(proto, dilations_) == Status::OK()) + attr_read = true; + } + if (!attr_read) { + dilations_.resize(kernel_shape_.size(), 1); + } + + attr_read = false; + attr = attributes.find(attributes_prefix + "group"); + if (attr != attributes.end()) { + ONNX_NAMESPACE::AttributeProto proto = attr->second; + if (GetIntAttr(proto, group_) == Status::OK()) + attr_read = true; + } + if (!attr_read) { + group_ = 1; + } + + attr = attributes.find(attributes_prefix + "epsilon"); + if (attr != attributes.end()) { + epsilon_ = attr->second.f(); + } + } + + private: + mkldnn::memory::desc filter_desc_; + mkldnn::memory::format_tag filter_format_; + + std::shared_ptr src_mem_from_; + std::unique_ptr src_mem_to_; + + size_t src_size_; + size_t filter_size_; + size_t dst_size_; + + std::shared_ptr src_mem_; + std::unique_ptr filter_mem_; + std::unique_ptr bias_mem_; + + std::unique_ptr fwd_desc_; + + std::unique_ptr src_md_; + std::unique_ptr filter_md_; + std::unique_ptr bias_md_; + + std::unique_ptr conv_fwd_pd_; + std::unique_ptr conv_fwd_; + + private: + IAllocatorUniquePtr src_reorder_buffer_; + IAllocatorUniquePtr dst_reorder_buffer_; + + private: + Status ComputeKernelShape(const TensorShape& weight_shape, std::vector& kernel_shape) const { + if (kernel_shape_specified_) { + kernel_shape = kernel_shape_; + if (kernel_shape.size() + 2 != weight_shape.NumDimensions()) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "kernel_shape num_dims is not compatible with W num_dims.", + " kernel_shape: ", TensorShape(kernel_shape).ToString().c_str(), + " W: ", weight_shape.ToString().c_str()); + } + for (size_t i = 0; i < kernel_shape.size(); ++i) { + if (kernel_shape[i] != weight_shape[i + 2]) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "kernel_shape is not compatible with W shape.", + " kernel_shape: ", TensorShape(kernel_shape).ToString().c_str(), + " W: ", weight_shape.ToString().c_str()); + } + } + } else { + auto& weight_dims = weight_shape.GetDims(); + kernel_shape = std::vector(weight_dims.begin() + 2, weight_dims.end()); + } + + return Status::OK(); + } + + Status ValidateInputShape(const TensorShape& X, const TensorShape& W) const { + const int64_t C = X[1]; + const int64_t M = W[0]; + + if (X.NumDimensions() != W.NumDimensions()) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "X num_dims does not match W num_dims.", + " X: ", X.ToString().c_str(), + " W: ", W.ToString().c_str()); + } + + if (C != W[1] * group_) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Input channels C is not equal to kernel channels * group.", + " C: ", C, + " kernel channels: ", W[1], + " group: ", group_); + } + + if (M % group_ != 0) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Output channels M is not divisible by group.", + " M: ", M, + " group: ", group_); + } + return Status::OK(); + } + + template + Status InferOutputShape(const TensorShape& input_shape, + const std::vector& kernel_shape, + const std::vector& strides, + const std::vector& dilations, + std::vector* pads, + std::vector* output_shape) const { + size_t rank = gsl::narrow_cast(input_shape.NumDimensions()); + for (size_t dim = 0; dim < rank; ++dim) { + if (dim >= strides.size() || dim >= kernel_shape.size() || + dim >= dilations.size() || dim >= pads->size() || + rank + dim >= pads->size()) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Out of bound access to array"); + } + int64_t dim_size = 0; + ORT_RETURN_IF_ERROR(ComputePadAndOutputShape( + input_shape[dim], + strides[dim], + kernel_shape[dim], + dilations[dim], + auto_pad_, + &pads->at(dim), + &pads->at(input_shape.NumDimensions() + dim), + &dim_size)); + if (dim_size <= 0) { + return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, "Invalid input shape: " + input_shape.ToString()); + } + output_shape->push_back(dim_size); + } + return Status::OK(); + } + + private: + std::unique_ptr stream_; + std::vector kernel_shape_; // must use ComputeKernelShape(...), instead of kernel_shape_ + AutoPadType auto_pad_; + int64_t group_; + bool kernel_shape_specified_; + std::vector strides_; + std::vector pads_; + std::vector dilations_; + std::string activation_; + float alpha_; + float epsilon_ = 1e-5f; // attribute of fused batchnorm. +}; +} // namespace mkl_dnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.cc b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.cc index 0c7c489445..f1f1a97f21 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.cc +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.cc @@ -10,6 +10,7 @@ #include "core/providers/mkldnn/mkldnn_common.h" #include "core/providers/mkldnn/subgraph/mkldnn_conv.h" #include "core/providers/mkldnn/subgraph/mkldnn_batchnorm.h" +#include "core/providers/mkldnn/subgraph/mkldnn_conv_batchnorm.h" #include "core/providers/mkldnn/subgraph/mkldnn_activations.h" #include "core/providers/mkldnn/subgraph/mkldnn_pool.h" #include "core/providers/mkldnn/subgraph/mkldnn_sum.h" @@ -27,7 +28,7 @@ class SubgraphPrimitive : public PrimitiveBase { OrtKernelContext* context, const SubgraphParams& params) : cpu_engine_(GetEngine()) { - context_.stream.reset(new mkldnn::stream(mkldnn::stream::kind::eager)); + context_.stream.reset(new mkldnn::stream(cpu_engine_)); if (context_.net.size() == 0) { CreateKernels(params); @@ -50,8 +51,10 @@ class SubgraphPrimitive : public PrimitiveBase { if (!status.IsOK()) break; } - if (status.IsOK()) - context_.stream->submit(context_.net); + if (status.IsOK()) { + for (size_t i = 0; i < context_.net.size(); ++i) + context_.net.at(i).execute(*context_.stream, context_.net_args.at(i)); + } return status; } @@ -107,6 +110,25 @@ class SubgraphPrimitive : public PrimitiveBase { kernel->parents_.push_back(context_.kernels[index]); } context_.kernels.push_back(kernel); + } else if (mkldnn_node.name == "Conv-BatchNormalization") { + std::ostringstream os; + os << "Conv-" << mkldnn_node.node_index << "-"; + std::shared_ptr> kernel; + kernel.reset(new MklDnnConvBatchNorm(mkldnn_node, params.provider, params.attributes, os.str())); + for (auto index : mkldnn_node.parent_nodes) { + kernel->parents_.push_back(context_.kernels[index]); + } + context_.kernels.push_back(kernel); + } else if (mkldnn_node.name == "Conv-BatchNormalization-Relu") { + std::ostringstream os; + os << "Conv-" << mkldnn_node.node_index << "-"; + std::shared_ptr> kernel; + kernel.reset(new MklDnnConvBatchNorm(mkldnn_node, params.provider, params.attributes, os.str())); + kernel->fuse_relu_ = true; + for (auto index : mkldnn_node.parent_nodes) { + kernel->parents_.push_back(context_.kernels[index]); + } + context_.kernels.push_back(kernel); } else if (mkldnn_node.name == "MaxPool") { std::ostringstream os; os << "MaxPool-" << mkldnn_node.node_index << "-"; @@ -143,15 +165,6 @@ class SubgraphPrimitive : public PrimitiveBase { kernel->parents_.push_back(context_.kernels[index]); } context_.kernels.push_back(kernel); - } else if (mkldnn_node.name == "Sum") { - std::ostringstream os; - os << "Sum-" << mkldnn_node.node_index << "-"; - std::shared_ptr> kernel; - kernel.reset(new MklDnnSum(mkldnn_node, params.provider, params.attributes, os.str())); - for (auto index : mkldnn_node.parent_nodes) { - kernel->parents_.push_back(context_.kernels[index]); - } - context_.kernels.push_back(kernel); } else if (mkldnn_node.name == "LRN") { std::ostringstream os; os << "LRN-" << mkldnn_node.node_index << "-"; @@ -161,14 +174,24 @@ class SubgraphPrimitive : public PrimitiveBase { kernel->parents_.push_back(context_.kernels[index]); } context_.kernels.push_back(kernel); + } else if (mkldnn_node.name == "Sum") { + std::ostringstream os; + os << "Sum-" << mkldnn_node.node_index << "-"; + std::shared_ptr> kernel; + kernel.reset(new MklDnnSum(mkldnn_node, params.provider, params.attributes, os.str())); + for (auto index : mkldnn_node.parent_nodes) { + kernel->parents_.push_back(context_.kernels[index]); + } + context_.kernels.push_back(kernel); } } } - + private: struct SubgraphContext { std::unique_ptr stream; std::vector net; + std::vector> net_args; std::vector> kernels; SubgraphContext() : stream(nullptr) {} @@ -177,9 +200,8 @@ class SubgraphPrimitive : public PrimitiveBase { Status Initialize(const OrtCustomOpApi* api, OrtKernelContext* context) { // Propagate mkldnn block format // dst format of current node to src format of next node - mkldnn::memory::format source_format = mkldnn::memory::format::any; // ONNXRuntime format for (auto& kernel : context_.kernels) { - Status status = kernel->CreatePrimitives(api, context, cpu_engine_, context_.net, source_format); + Status status = kernel->CreatePrimitives(api, context, cpu_engine_, context_.net, context_.net_args); if (status.IsOK()) kernel->ReorderWeights(api, context, cpu_engine_); else @@ -246,7 +268,7 @@ Status MkldnnFuncKernel::Compute(const OrtCustomOpApi* api, OrtKernelContext* status = primitive->Compute(api, context); } catch (const mkldnn::error& e) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, - ", message: ", e.message.c_str()); + ", message: ", e.what()); } return status; } diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.h index 59105a3b8c..b80f367cc9 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_func_kernel.h @@ -6,6 +6,7 @@ #include "core/providers/mkldnn/mkldnn_execution_provider.h" #include "core/session/onnxruntime_c_api.h" #include "core/framework/func_api.h" +#include "mkldnn_kernel.h" namespace onnxruntime { namespace mkl_dnn { @@ -37,8 +38,8 @@ class MkldnnFuncKernel { if (sub_it->second.type() == ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_STRING) { params_.subgraph_id = sub_it->second.s(); params_.subgraph = provider->GetMklDnnSubgraph(params_.subgraph_id); - - std::ostringstream key_os; + + std::ostringstream key_os; key_os << params_.subgraph->graph_name << "_" << params_.subgraph_id << "-"; key_os << params_.subgraph->mkldnn_nodes.back().ToString() << "-"; key_os << params_.subgraph->mkldnn_nodes.back().output_name; @@ -146,7 +147,7 @@ class MkldnnFuncKernel { if (attr != attributes.end()) { key.append(1, '#'); ONNX_NAMESPACE::AttributeProto proto = attr->second; - key.append(proto.s()); + key.append(proto.s()); key.append(1, '#'); } @@ -183,19 +184,19 @@ class MkldnnFuncKernel { key.append(1, '#'); } - attr = attributes.find(attributes_prefix + "group"); + attr = attributes.find(attributes_prefix + "group"); if (attr != attributes.end()) { key.append(1, '#'); ONNX_NAMESPACE::AttributeProto proto = attr->second; - key.append(std::to_string(proto.i())); + key.append(std::to_string(proto.i())); key.append(1, '#'); } return key; } - + std::string GetLrnAttributeKey(const NodeAttributes& attributes, - const std::string attributes_prefix = "") { + const std::string attributes_prefix = "") { std::string key; auto attr = attributes.find(attributes_prefix + "alpha"); @@ -214,7 +215,7 @@ class MkldnnFuncKernel { key.append(1, '#'); } - attr = attributes.find(attributes_prefix + "bias"); + attr = attributes.find(attributes_prefix + "bias"); if (attr != attributes.end()) { key.append(1, '#'); ONNX_NAMESPACE::AttributeProto proto = attr->second; diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.cc b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.cc index 19007e99ef..05533beceb 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.cc +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.cc @@ -8,44 +8,48 @@ namespace mkl_dnn { void MklDnnKernel::InitDstReorderOutput(mkldnn::engine& cpu_engine, mkldnn::memory::data_type& data_type, - std::vector& net) { + std::vector& net, + std::vector>& net_args) { // Allocate dst buffer if reorder is necessary - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) + { // reorder to ONNXRuntime format mkldnn::memory::dims dst_dims_mkl( primitive_dst_shape_.GetDims().begin(), primitive_dst_shape_.GetDims().end()); mkldnn::memory::desc dst_des = mkldnn::memory::desc(dst_dims_mkl, data_type, ort_source_format_); - reorder_dst_mem_to_.reset(new mkldnn::memory({dst_des, cpu_engine}, nullptr)); + reorder_dst_mem_to_.reset(new mkldnn::memory(dst_des, cpu_engine)); net.push_back(mkldnn::reorder(*primitive_dst_mem_, *reorder_dst_mem_to_)); + net_args.push_back({{MKLDNN_ARG_FROM, *primitive_dst_mem_}, + {MKLDNN_ARG_TO, *reorder_dst_mem_to_}}); } } -mkldnn::memory::format MklDnnKernel::GetSourceFormat(int dim_size) { - mkldnn::memory::format source_format = mkldnn::memory::format::any; +mkldnn::memory::format_tag MklDnnKernel::GetSourceFormat(int dim_size) { + mkldnn::memory::format_tag source_format = mkldnn::memory::format_tag::any; switch (dim_size) { case 1: { - source_format = mkldnn::memory::format::x; + source_format = mkldnn::memory::format_tag::x; break; } case 2: { - source_format = mkldnn::memory::format::nc; + source_format = mkldnn::memory::format_tag::nc; break; } case 3: { - source_format = mkldnn::memory::format::ntc; + source_format = mkldnn::memory::format_tag::ntc; break; } case 4: { - source_format = mkldnn::memory::format::nchw; + source_format = mkldnn::memory::format_tag::nchw; break; } case 5: { - source_format = mkldnn::memory::format::ncdhw; + source_format = mkldnn::memory::format_tag::ncdhw; break; } default: { - source_format = mkldnn::memory::format::any; + source_format = mkldnn::memory::format_tag::any; break; } } diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.h index 0a4bc68eea..7d6912d058 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_kernel.h @@ -19,6 +19,7 @@ class MklDnnKernel { public: MklDnnKernel(const MklDnnNode& node, MKLDNNExecutionProvider* provider) { + name_ = node.name; mklnode_ptr_ = std::make_shared(node); provider_ = provider; alloc_ = provider_->GetAllocator(0, OrtMemTypeDefault); @@ -29,7 +30,7 @@ class MklDnnKernel { OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& src_fmt) = 0; + std::vector> &net_args) = 0; virtual void ReorderWeights(const OrtCustomOpApi* api, OrtKernelContext* context, mkldnn::engine& cpu_engine) { ORT_UNUSED_PARAMETER(api); @@ -80,33 +81,35 @@ class MklDnnKernel { void InitDstReorderOutput(mkldnn::engine& cpu_engine, mkldnn::memory::data_type& data_type, - std::vector& net); + std::vector& net, + std::vector>& net_args); - //void AllocateMemoryAndReorderIfNeeded(const OrtCustomOpApi* api); - - //void AllocateOutputTensor(const OrtCustomOpApi* api, int index, const int64_t* shape, size_t dim); - - mkldnn::memory::format GetSourceFormat(int dim_size); + mkldnn::memory::format_tag GetSourceFormat(int dim_size); public: + std::string name_; std::vector> parents_; bool fuse_relu_ = false; bool fuse_sum_ = false; std::shared_ptr primitive_dst_mem_; std::unique_ptr primitive_dst_md_; TensorShape primitive_dst_shape_; - mkldnn::memory::format primitive_dst_format_ = mkldnn::memory::format::any; + mkldnn::memory::desc primitive_dst_desc_; + // ONNX Runtime format + mkldnn::memory::format_tag ort_source_format_; + mkldnn::memory::desc ort_source_desc_; + + // input format. + // It can be ORT format (nchw) or blocked memory format from parent node + // mkldnn::memory::format_tag source_format_ = mkldnn::memory::format_tag::any; + mkldnn::memory::desc source_desc_ = mkldnn::memory::desc(); protected: - // ONNX Runtime format - mkldnn::memory::format ort_source_format_ = mkldnn::memory::format::any; - // input format. - // It can be ORT format (nchw) or blocked memory format from parent nce - mkldnn::memory::format src_format_ = mkldnn::memory::format::any; + // Pointer to MklNode of subgraph IR std::shared_ptr mklnode_ptr_; // input format expected by primitive object - mkldnn::memory::format primitive_src_format_ = mkldnn::memory::format::any; + mkldnn::memory::desc primitive_src_desc_; // memory used for reorders std::unique_ptr reorder_dst_mem_to_; diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_lrn.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_lrn.h index c27914e79a..2b03717c14 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_lrn.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_lrn.h @@ -26,7 +26,7 @@ class MklDnnLrn : public MklDnnKernel { OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& source_format) { + std::vector>& net_args) override { Ort::CustomOpApi ort{*api}; int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; @@ -40,22 +40,25 @@ class MklDnnLrn : public MklDnnKernel { auto xdim = tensor_shape.size(); ort_source_format_ = GetSourceFormat(static_cast(xdim)); - source_format = ort_source_format_; x_shape = TensorShape(xshape, xdim); - mkldnn::memory::dims src_dims_mkl( + mkldnn::memory::dims src_dims( x_shape.GetDims().begin(), x_shape.GetDims().end()); + ort_source_desc_ = mkldnn::memory::desc( + {src_dims}, MklDnnType(), ort_source_format_); src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), source_format)); + {src_dims}, MklDnnType(), ort_source_format_)); src_mem_.reset( - new mkldnn::memory({*src_md_, cpu_engine}, nullptr)); + new mkldnn::memory(*src_md_, cpu_engine, nullptr)); } else { src_md_.reset( - new mkldnn::memory::desc(parents_[0].get()->primitive_dst_mem_.get()->get_primitive_desc().desc())); + new mkldnn::memory::desc(parents_[0].get()->primitive_dst_desc_)); src_mem_ = parents_[0].get()->primitive_dst_mem_; x_shape = parents_[0].get()->primitive_dst_shape_; - ort_source_format_ = source_format; + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; } primitive_dst_shape_ = TensorShape(x_shape); @@ -68,39 +71,39 @@ class MklDnnLrn : public MklDnnKernel { fwd_primitive_desc_.reset(new mkldnn::lrn_forward::primitive_desc( *fwd_desc_, cpu_engine)); - primitive_src_format_ = static_cast( - fwd_primitive_desc_.get()->src_primitive_desc().desc().data.format); - primitive_dst_format_ = static_cast( - fwd_primitive_desc_.get()->dst_primitive_desc().desc().data.format); + primitive_src_desc_ = fwd_primitive_desc_.get()->src_desc(); + primitive_dst_desc_ = fwd_primitive_desc_.get()->dst_desc(); if (mklnode_ptr_->output_index >= 0) { // last node of sub-graph. need to allocate memory for output_tensor - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { // reorder neded. Use primitive output as input to reorder and // allocate buffer for reorder output, final output of this subgraph primitive_dst_mem_.reset( - new mkldnn::memory(fwd_primitive_desc_.get()->dst_primitive_desc())); + new mkldnn::memory(fwd_primitive_desc_.get()->dst_desc(), cpu_engine)); } else { // Last node but re-order not needed. Allocate buffer to output of this node primitive_dst_mem_.reset( - new mkldnn::memory(fwd_primitive_desc_.get()->dst_primitive_desc(), nullptr)); + new mkldnn::memory(fwd_primitive_desc_.get()->dst_desc(), cpu_engine, nullptr)); } } else { // Intermediate node. Use mkldnn kernel internal memory for output and // use this as input to next node. primitive_dst_mem_.reset( - new mkldnn::memory(fwd_primitive_desc_.get()->dst_primitive_desc())); + new mkldnn::memory(fwd_primitive_desc_.get()->dst_desc(), cpu_engine)); } lrn_fwd_.reset( - new mkldnn::lrn_forward(*fwd_primitive_desc_, *src_mem_, *primitive_dst_mem_)); + new mkldnn::lrn_forward(*fwd_primitive_desc_)); net.push_back(*lrn_fwd_); + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); if (mklnode_ptr_->output_index >= 0) { // one of the end nodes. Allocate output buffer memory and // reorder is necessary mkldnn::memory::data_type t = MklDnnType(); - InitDstReorderOutput(cpu_engine, t, net); + InitDstReorderOutput(cpu_engine, t, net, net_args); } return Status::OK(); @@ -123,7 +126,7 @@ class MklDnnLrn : public MklDnnKernel { OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); T* dst_data = ort.GetTensorMutableData(output); - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_!= ort_source_desc_) { reorder_dst_mem_to_->set_data_handle(dst_data); } else { primitive_dst_mem_->set_data_handle(dst_data); diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_pool.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_pool.h index e4ac5099fc..268c108f3c 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_pool.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_pool.h @@ -24,7 +24,7 @@ class MklDnnPool : public MklDnnKernel { Status CreatePrimitives(const OrtCustomOpApi* api, OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& source_format) override { + std::vector>& net_args) override { Ort::CustomOpApi ort{*api}; int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; @@ -37,33 +37,36 @@ class MklDnnPool : public MklDnnKernel { auto xdim = tensor_shape.size(); mkldnn::memory::dims dims(xdim); - ort_source_format_ = GetSourceFormat(static_cast(xdim)); - source_format = ort_source_format_; - src_format_ = ort_source_format_; x_shape_ = TensorShape(xshape, xdim); mkldnn::memory::dims src_dims_mkl(x_shape_.GetDims().begin(), x_shape_.GetDims().end()); + ort_source_format_ = GetSourceFormat(static_cast(xdim)); + // ort_source_desc is the format of ONNX Runtime tensor format + ort_source_desc_ = mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), ort_source_format_); + // source_desc is propagating format. input to this op. + source_desc_ = mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), ort_source_format_); // reorder for better performance - mkldnn::memory::format fmt = GetAVXFormat(src_dims_mkl); - src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), fmt)); + mkldnn::memory::format_tag src_format = GetAVXFormat(src_dims_mkl); + src_md_.reset(new mkldnn::memory::desc({src_dims_mkl}, MklDnnType(), src_format)); } else { // get the output of previous node (mkldnn block propagation). // TODO Sourcenode will set src of this node. x_shape_ = parents_[0].get()->primitive_dst_shape_; - ort_source_format_ = source_format; - src_format_ = parents_[0].get()->primitive_dst_format_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; mkldnn::memory::dims src_dims_mkl(x_shape_.GetDims().begin(), x_shape_.GetDims().end()); - if (src_format_ == ort_source_format_) { + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; + + if (source_desc_ == ort_source_desc_) { // reorder for better performance - mkldnn::memory::format fmt = GetAVXFormat(src_dims_mkl); + mkldnn::memory::format_tag fmt = GetAVXFormat(src_dims_mkl); src_md_.reset(new mkldnn::memory::desc( {src_dims_mkl}, MklDnnType(), fmt)); } else { - src_md_.reset(new mkldnn::memory::desc( - parents_[0].get()->primitive_dst_mem_.get()->get_primitive_desc().desc())); + src_md_.reset(new mkldnn::memory::desc(parents_[0].get()->primitive_dst_mem_->get_desc())); } } @@ -93,7 +96,7 @@ class MklDnnPool : public MklDnnKernel { mkldnn::memory::dims padding_right_mkl(pads_.begin() + (pads_.size() / 2), pads_.end()); primitive_dst_md_.reset(new mkldnn::memory::desc( - {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); mkldnn::algorithm algo = mkldnn::algorithm::pooling_max; if (op_name_ == "AveragePool" || op_name_ == "GlobalAveragePool") { @@ -106,8 +109,7 @@ class MklDnnPool : public MklDnnKernel { mkldnn::prop_kind::forward_inference, algo, *src_md_, *primitive_dst_md_, strides_mkl, kernel_mkl, - padding_left_mkl, padding_right_mkl, - mkldnn::padding_kind::zero)); + padding_left_mkl, padding_right_mkl)); fwd_primitive_desc_.reset(new mkldnn::pooling_forward::primitive_desc( *fwd_desc_, cpu_engine)); @@ -115,37 +117,35 @@ class MklDnnPool : public MklDnnKernel { if (mklnode_ptr_->parent_nodes.empty()) { // Sub-graph's first node. Read input from input buffer src_mem_.reset(new mkldnn::memory( - fwd_primitive_desc_.get()->src_primitive_desc(), nullptr)); + fwd_primitive_desc_.get()->src_desc(), cpu_engine, nullptr)); } else { // Sub-graph's inner node. set input to parent's output src_mem_ = parents_[0].get()->primitive_dst_mem_; } - primitive_src_format_ = static_cast( - fwd_primitive_desc_.get()->src_primitive_desc().desc().data.format); + primitive_src_desc_ = fwd_primitive_desc_.get()->src_desc(); + primitive_dst_desc_ = fwd_primitive_desc_.get()->dst_desc(); - primitive_dst_format_ = static_cast( - fwd_primitive_desc_.get()->dst_primitive_desc().desc().data.format); + src_size_ = fwd_primitive_desc_.get()->src_desc().get_size(); + dst_size_ = fwd_primitive_desc_.get()->dst_desc().get_size(); - src_size_ = fwd_primitive_desc_.get()->src_primitive_desc().get_size(); - dst_size_ = fwd_primitive_desc_.get()->dst_primitive_desc().get_size(); - - // reorder source memory for best performance (AVX512); - if (primitive_src_format_ != src_format_) { - mkldnn::memory::dims src_dims_mkl(x_shape_.GetDims().begin(), x_shape_.GetDims().end()); - auto src_md = mkldnn::memory::desc(src_dims_mkl, MklDnnType(), src_format_); - auto pd = mkldnn::memory::primitive_desc(src_md, cpu_engine); + // reorder source memory for best performance (AVX512); + if (primitive_src_desc_ != source_desc_) { + mkldnn::memory::dims src_dims(x_shape_.GetDims().begin(), x_shape_.GetDims().end()); + auto pd = mkldnn::memory::desc(source_desc_); if (mklnode_ptr_->parent_nodes.empty()) - src_mem_from_.reset(new mkldnn::memory(pd, nullptr)); + src_mem_from_.reset(new mkldnn::memory(pd, cpu_engine, nullptr)); else src_mem_from_ = parents_[0].get()->primitive_dst_mem_; - src_mem_.reset(new mkldnn::memory(fwd_primitive_desc_->src_primitive_desc(), nullptr)); + src_mem_.reset(new mkldnn::memory(fwd_primitive_desc_->src_desc(), cpu_engine, nullptr)); net.push_back(mkldnn::reorder(*src_mem_from_, *src_mem_)); + net_args.push_back({{MKLDNN_ARG_FROM, *src_mem_from_}, + {MKLDNN_ARG_TO, *src_mem_}}); } else { if (mklnode_ptr_->parent_nodes.empty()) { - src_mem_.reset(new mkldnn::memory(fwd_primitive_desc_->src_primitive_desc(), nullptr)); + src_mem_.reset(new mkldnn::memory(fwd_primitive_desc_->src_desc(), cpu_engine, nullptr)); } else { src_mem_ = parents_[0].get()->primitive_dst_mem_; } @@ -153,32 +153,33 @@ class MklDnnPool : public MklDnnKernel { if (mklnode_ptr_->output_index >= 0) { // last node of sub-graph. need to allocate memory for output_tensor - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { // reorder neded. Use primitive output as input to reorder and // allocate buffer for reorder output, final output of this subgraph primitive_dst_mem_.reset( - new mkldnn::memory(fwd_primitive_desc_.get()->dst_primitive_desc())); + new mkldnn::memory(fwd_primitive_desc_.get()->dst_desc(), cpu_engine)); } else { // Last node but re-order not needed. Allocate buffer to output of this node primitive_dst_mem_.reset( - new mkldnn::memory(fwd_primitive_desc_.get()->dst_primitive_desc(), nullptr)); + new mkldnn::memory(fwd_primitive_desc_.get()->dst_desc(), cpu_engine, nullptr)); } } else { // Intermediate node. Use mkldnn kernel internal memory for output and // use this as input to next node. primitive_dst_mem_.reset( - new mkldnn::memory(fwd_primitive_desc_.get()->dst_primitive_desc())); + new mkldnn::memory(fwd_primitive_desc_.get()->dst_desc(), cpu_engine)); } pool_fwd_.reset( - new mkldnn::pooling_forward(*fwd_primitive_desc_, *src_mem_, *primitive_dst_mem_)); + new mkldnn::pooling_forward(*fwd_primitive_desc_)); net.push_back(*pool_fwd_); - + net_args.push_back({{MKLDNN_ARG_SRC, *src_mem_}, + {MKLDNN_ARG_DST, *primitive_dst_mem_}}); if (mklnode_ptr_->output_index >= 0) { // one of the end nodes. Allocate output buffer memory and // reorder is necessary mkldnn::memory::data_type t = MklDnnType(); - InitDstReorderOutput(cpu_engine, t, net); + InitDstReorderOutput(cpu_engine, t, net, net_args); } return Status::OK(); } @@ -200,7 +201,7 @@ class MklDnnPool : public MklDnnKernel { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Please call default CPU kernel."); } - if (primitive_src_format_ != src_format_) { + if (fwd_primitive_desc_.get()->src_desc() != source_desc_) { if (mklnode_ptr_->parent_nodes.empty()) { const OrtValue* input_tensor = ort.KernelContext_GetInput(context, input_index); const T* src_data = const_cast(ort.GetTensorData(input_tensor)); @@ -209,7 +210,7 @@ class MklDnnPool : public MklDnnKernel { src_mem_from_ = parents_[0].get()->primitive_dst_mem_; } - auto src_size = fwd_primitive_desc_.get()->src_primitive_desc().get_size(); + auto src_size = fwd_primitive_desc_.get()->src_desc().get_size(); src_reorder_buffer_ = IAllocator::MakeUniquePtr(alloc_, src_size); src_mem_->set_data_handle(src_reorder_buffer_.get()); } else { @@ -230,7 +231,7 @@ class MklDnnPool : public MklDnnKernel { OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); T* dst_data = ort.GetTensorMutableData(output); - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { reorder_dst_mem_to_->set_data_handle(dst_data); } else { primitive_dst_mem_->set_data_handle(dst_data); @@ -322,15 +323,15 @@ class MklDnnPool : public MklDnnKernel { std::unique_ptr dst_mem_to_; private: - mkldnn::memory::format GetAVXFormat(const mkldnn::memory::dims& src_dims_mkl) { + mkldnn::memory::format_tag GetAVXFormat(const mkldnn::memory::dims& src_dims_mkl) { bool is_2D = src_dims_mkl.size() == 4 ? true : false; - mkldnn::memory::format fmt = mkldnn::memory::format::any; + mkldnn::memory::format_tag fmt = mkldnn::memory::format_tag::any; if (CPUIDInfo::GetCPUIDInfo().HasAVX512f()) { - fmt = is_2D ? mkldnn::memory::format::nChw16c : mkldnn::memory::format::nCdhw16c; + fmt = is_2D ? mkldnn::memory::format_tag::nChw16c : mkldnn::memory::format_tag::nCdhw16c; } else if (CPUIDInfo::GetCPUIDInfo().HasAVX2() && (src_dims_mkl[1] % 8 == 0)) { - fmt = is_2D ? mkldnn::memory::format::nChw8c : mkldnn::memory::format::ncdhw; + fmt = is_2D ? mkldnn::memory::format_tag::nChw8c : mkldnn::memory::format_tag::ncdhw; } else { - fmt = is_2D ? mkldnn::memory::format::nchw : mkldnn::memory::format::ncdhw; + fmt = is_2D ? mkldnn::memory::format_tag::nchw : mkldnn::memory::format_tag::ncdhw; } return fmt; } diff --git a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_sum.h b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_sum.h index a45af46acc..3899e5b093 100644 --- a/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_sum.h +++ b/onnxruntime/core/providers/mkldnn/subgraph/mkldnn_sum.h @@ -25,7 +25,7 @@ class MklDnnSum : public MklDnnKernel { OrtKernelContext* context, mkldnn::engine& cpu_engine, std::vector& net, - mkldnn::memory::format& source_format) override { + std::vector>& net_args) override { Ort::CustomOpApi ort{*api}; int num_inputs = mklnode_ptr_->num_inputs; int input_index = mklnode_ptr_->input_start_index < 0 ? 0 : mklnode_ptr_->input_start_index; @@ -41,12 +41,18 @@ class MklDnnSum : public MklDnnKernel { auto xdim = tensor_shape.size(); ort_source_format_ = GetSourceFormat(static_cast(xdim)); - source_format = ort_source_format_; - src_format_ = ort_source_format_; + x_shape = TensorShape(xshape, xdim); + mkldnn::memory::dims src_dims( + x_shape.GetDims().begin(), x_shape.GetDims().end()); + ort_source_desc_ = mkldnn::memory::desc( + {src_dims}, MklDnnType(), ort_source_format_); + source_desc_ = ort_source_desc_; } else { x_shape = parents_[0].get()->primitive_dst_shape_; - src_format_ = parents_[0].get()->primitive_dst_format_; + ort_source_format_ = parents_[0].get()->ort_source_format_; + ort_source_desc_ = parents_[0].get()->ort_source_desc_; + source_desc_ = parents_[0].get()->primitive_dst_desc_; } primitive_dst_shape_ = TensorShape(x_shape); @@ -65,66 +71,60 @@ class MklDnnSum : public MklDnnKernel { auto xdim = tensor_shape.size(); mkldnn::memory::dims dims(xdim); - ort_source_format_ = GetSourceFormat(static_cast(xdim)); x_shape1 = TensorShape(xshape, xdim); - mkldnn::memory::dims src_dims_mkl( + mkldnn::memory::dims src_dims( x_shape1.GetDims().begin(), x_shape1.GetDims().end()); - - src_md_.reset(new mkldnn::memory::desc( - {src_dims_mkl}, MklDnnType(), src_format_)); - - auto mpd = mkldnn::memory::primitive_desc(*src_md_, cpu_engine); - auto src_memory = mkldnn::memory(mpd, nullptr); + auto mpd = mkldnn::memory::desc({src_dims}, MklDnnType(), ort_source_format_); + auto src_memory = mkldnn::memory(mpd, cpu_engine, nullptr); srcs_pd_.push_back(mpd); srcs_memory_.push_back(src_memory); coeff.push_back(1.0); } else { - src_md_.reset( - new mkldnn::memory::desc(parents_[i].get()->primitive_dst_mem_.get()->get_primitive_desc().desc())); - auto mpd = mkldnn::memory::primitive_desc(*src_md_, cpu_engine); - auto src_memory = *parents_[i].get()->primitive_dst_mem_; //mkldnn::memory(mpd); + x_shape = parents_[0].get()->primitive_dst_shape_; + auto mpd = mkldnn::memory::desc(parents_[i].get()->primitive_dst_desc_); + auto src_memory = *parents_[i].get()->primitive_dst_mem_; srcs_pd_.push_back(mpd); srcs_memory_.push_back(src_memory); coeff.push_back(1.0); - ort_source_format_ = source_format; } } - primitive_dst_md_.reset(new mkldnn::memory::desc( - {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format::any)); + primitive_dst_md_.reset(new mkldnn::memory::desc( + {dst_dims_mkl}, MklDnnType(), mkldnn::memory::format_tag::any)); sum_pd_.reset(new mkldnn::sum::primitive_desc( - *primitive_dst_md_, coeff, srcs_pd_)); - primitive_dst_format_ = static_cast(sum_pd_->dst_primitive_desc().desc().data.format); + *primitive_dst_md_, coeff, srcs_pd_, cpu_engine)); if (mklnode_ptr_->output_index >= 0) { // last node of sub-graph. need to allocate memory for output_tensor - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { // reorder neded. Use primitive output as input to reorder and // allocate buffer for reorder output, final output of this subgraph - primitive_dst_mem_.reset(new mkldnn::memory(sum_pd_->dst_primitive_desc())); + primitive_dst_mem_.reset(new mkldnn::memory(sum_pd_->dst_desc(), cpu_engine)); } else { // Last node but re-order not needed. Allocate buffer to output of this node - primitive_dst_mem_.reset(new mkldnn::memory(sum_pd_->dst_primitive_desc(), nullptr)); + primitive_dst_mem_.reset(new mkldnn::memory(sum_pd_->dst_desc(), cpu_engine, nullptr)); } } else { // Intermediate node. Use mkldnn kernel internal memory for output and // use this as input to next node. - primitive_dst_mem_.reset(new mkldnn::memory(sum_pd_->dst_primitive_desc())); + primitive_dst_mem_.reset(new mkldnn::memory(sum_pd_->dst_desc(), cpu_engine)); } - primitive_dst_format_ = static_cast(sum_pd_->dst_primitive_desc().desc().data.format); + primitive_dst_desc_ = sum_pd_->dst_desc(); - std::vector inputs; - for (int i = 0; i < num_inputs; i++) { - inputs.push_back(srcs_memory_[i]); - } - auto c = mkldnn::sum(*sum_pd_, inputs, *primitive_dst_mem_); + auto c = mkldnn::sum(*sum_pd_); net.push_back(c); + std::unordered_map args{ + {MKLDNN_ARG_DST, *primitive_dst_mem_}}; + for (int i = 0; i < (int)num_inputs; i++) { + args.insert({MKLDNN_ARG_MULTIPLE_SRC + i, srcs_memory_[i]}); + } + net_args.push_back(args); if (mklnode_ptr_->output_index >= 0) { // one of the end nodes. Allocate output buffer memory and // reorder is necessary mkldnn::memory::data_type t = MklDnnType(); - InitDstReorderOutput(cpu_engine, t, net); + InitDstReorderOutput(cpu_engine, t, net, net_args); } return Status::OK(); } @@ -150,7 +150,7 @@ class MklDnnSum : public MklDnnKernel { OrtValue* output = ort.KernelContext_GetOutput(context, mklnode_ptr_->output_index, &y_dims[0], static_cast(primitive_dst_shape_.GetDims().size())); T* dst_data = ort.GetTensorMutableData(output); - if (primitive_dst_format_ != ort_source_format_) { + if (primitive_dst_desc_ != ort_source_desc_) { reorder_dst_mem_to_->set_data_handle(dst_data); } else { primitive_dst_mem_->set_data_handle(dst_data); @@ -161,11 +161,12 @@ class MklDnnSum : public MklDnnKernel { private: std::unique_ptr src_md_; + std::vector src_mds_; std::vector srcs_memory_; - std::vector srcs_pd_; - std::unique_ptr src_mpd_; - std::unique_ptr dst_pd_; + std::vector srcs_pd_; + std::unique_ptr src_mpd_; + std::unique_ptr dst_pd_; std::unique_ptr sum_pd_; }; } // namespace mkl_dnn diff --git a/onnxruntime/test/providers/cpu/math/gemm_test.cc b/onnxruntime/test/providers/cpu/math/gemm_test.cc index 688fd7af46..0a17bfce59 100644 --- a/onnxruntime/test/providers/cpu/math/gemm_test.cc +++ b/onnxruntime/test/providers/cpu/math/gemm_test.cc @@ -3,7 +3,6 @@ #include "gtest/gtest.h" #include "test/providers/provider_test_utils.h" -#include "test/common/cuda_op_test_utils.h" namespace onnxruntime { namespace test { @@ -30,11 +29,6 @@ TEST(GemmOpTest, GemmNoTrans) { // Only CUDA kernel has float 16 support #ifdef USE_CUDA TEST(GemmOpTest, GemmNoTrans_f16) { - int min_cuda_architecture = 530; - if (!HasCudaEnvironment(min_cuda_architecture)) { - LOGS_DEFAULT(WARNING) << "Hardware NOT support FP16"; - return; - } OpTester test("Gemm"); test.AddAttribute("transA", (int64_t)0); @@ -235,7 +229,7 @@ TEST(GemmOpTest, GemmEmptyTensor) { test.AddInput("C", {3}, std::vector(3, 1.0f)); test.AddOutput("Y", {0, 3}, {}); - test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); //TensorRT: doesn't support dynamic shape yet + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider, kMklDnnExecutionProvider}); //TensorRT: doesn't support dynamic shape yet } TEST(GemmOpTest, GemmNoBiasOpset11) {