From 3875511f9e6bc2c7d16ca4d78e4af001d5caea07 Mon Sep 17 00:00:00 2001 From: Pranav Sharma Date: Fri, 21 Dec 2018 15:12:23 -0800 Subject: [PATCH] Fix inefficiencies in the mkldnn kernels. Some of these were (unfortunately) getting replicated in the new kernels. (#241) --- onnxruntime/core/providers/mkldnn/nn/conv.cc | 6 +++--- onnxruntime/core/providers/mkldnn/nn/conv.h | 2 +- onnxruntime/core/providers/mkldnn/nn/lrn.cc | 11 +++++------ onnxruntime/core/providers/mkldnn/nn/lrn.h | 4 ++-- onnxruntime/core/providers/mkldnn/nn/pool.cc | 8 ++++---- onnxruntime/core/providers/mkldnn/nn/pool.h | 2 +- 6 files changed, 16 insertions(+), 17 deletions(-) diff --git a/onnxruntime/core/providers/mkldnn/nn/conv.cc b/onnxruntime/core/providers/mkldnn/nn/conv.cc index 355562bf5a..f960c26fd9 100644 --- a/onnxruntime/core/providers/mkldnn/nn/conv.cc +++ b/onnxruntime/core/providers/mkldnn/nn/conv.cc @@ -68,7 +68,7 @@ class ConvPrimitive : public PrimitiveBase { ~ConvPrimitive() = default; void Compute(const T* src_data, const T* filter_data, - const T* dst_data, const T* bias_data = nullptr) { + 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( @@ -78,7 +78,7 @@ class ConvPrimitive : public PrimitiveBase { static_cast(const_cast(bias_data))); } context_.dst_mem->set_data_handle( - static_cast(const_cast(dst_data))); + static_cast(dst_data)); context_.stream->submit(context_.net); context_.src_mem->set_data_handle(nullptr); @@ -437,7 +437,7 @@ Status Conv::Compute(OpKernelContext* context) const { DoReorder(params); } - } catch (mkldnn::error& e) { + } catch (const mkldnn::error& e) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); } diff --git a/onnxruntime/core/providers/mkldnn/nn/conv.h b/onnxruntime/core/providers/mkldnn/nn/conv.h index 8189a7ca6b..76bd48a4c8 100644 --- a/onnxruntime/core/providers/mkldnn/nn/conv.h +++ b/onnxruntime/core/providers/mkldnn/nn/conv.h @@ -10,7 +10,7 @@ namespace mkl_dnn { template class Conv final : public onnxruntime::Conv { public: - Conv(const OpKernelInfo& info) : onnxruntime::Conv(info) { + explicit Conv(const OpKernelInfo& info) : onnxruntime::Conv(info) { } Status Compute(OpKernelContext* context) const override; diff --git a/onnxruntime/core/providers/mkldnn/nn/lrn.cc b/onnxruntime/core/providers/mkldnn/nn/lrn.cc index e1a7b6147d..eb72ad4bbe 100644 --- a/onnxruntime/core/providers/mkldnn/nn/lrn.cc +++ b/onnxruntime/core/providers/mkldnn/nn/lrn.cc @@ -1,4 +1,3 @@ - // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. @@ -25,13 +24,13 @@ ONNX_OPERATOR_TYPED_KERNEL_EX( namespace { // Struct which encapsulates parameters for MKLDNN LRN primitive. struct LRNParams { - mkldnn::memory::dims& dims_; + const mkldnn::memory::dims& dims_; float alpha_; float beta_; float bias_; int size_; - LRNParams(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. @@ -62,9 +61,9 @@ class LRNPrimitive : public PrimitiveBase { ~LRNPrimitive() = default; - void Compute(const T* src_data, const T* dst_data) { + 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(const_cast(dst_data))); + context_.dst_mem->set_data_handle(static_cast(dst_data)); context_.stream->submit(context_.net); context_.src_mem->set_data_handle(nullptr); @@ -231,7 +230,7 @@ Status LRN::Compute(OpKernelContext* context) const { MemoryReorderParams params(src, dst); DoReorder(params); } - } catch (mkldnn::error& e) { + } catch (const mkldnn::error& e) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); } diff --git a/onnxruntime/core/providers/mkldnn/nn/lrn.h b/onnxruntime/core/providers/mkldnn/nn/lrn.h index fe2a4c0373..44ea04b444 100644 --- a/onnxruntime/core/providers/mkldnn/nn/lrn.h +++ b/onnxruntime/core/providers/mkldnn/nn/lrn.h @@ -11,10 +11,10 @@ namespace mkl_dnn { template class LRN final : public onnxruntime::LRN { public: - LRN(const OpKernelInfo& info) : onnxruntime::LRN(info) {} + explicit LRN(const OpKernelInfo& info) : onnxruntime::LRN(info) {} Status Compute(OpKernelContext* p_op_kernel_context) const override; }; } -} \ No newline at end of file +} diff --git a/onnxruntime/core/providers/mkldnn/nn/pool.cc b/onnxruntime/core/providers/mkldnn/nn/pool.cc index 460c7c7aef..694f1bd641 100644 --- a/onnxruntime/core/providers/mkldnn/nn/pool.cc +++ b/onnxruntime/core/providers/mkldnn/nn/pool.cc @@ -45,7 +45,7 @@ struct PoolParams { mkldnn::memory::dims& padding_right; bool count_include_pad; - PoolParams(std::string op_name, std::string version, + 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, @@ -90,9 +90,9 @@ class PoolPrimitive : public PrimitiveBase { ~PoolPrimitive() = default; - void Compute(const T* src_data, const T* dst_data) { + 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(const_cast(dst_data))); + context_.dst_mem->set_data_handle(static_cast(dst_data)); context_.stream->submit(context_.net); context_.src_mem->set_data_handle(nullptr); @@ -326,7 +326,7 @@ Status Pool::Compute(OpKernelContext* context) const { MemoryReorderParams params(src, dst); DoReorder(params); } - } catch (mkldnn::error& e) { + } catch (const mkldnn::error& e) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Status: ", e.status, ", message: ", e.message.c_str()); } diff --git a/onnxruntime/core/providers/mkldnn/nn/pool.h b/onnxruntime/core/providers/mkldnn/nn/pool.h index bf8fc5f802..98ce4c16f1 100644 --- a/onnxruntime/core/providers/mkldnn/nn/pool.h +++ b/onnxruntime/core/providers/mkldnn/nn/pool.h @@ -11,7 +11,7 @@ namespace mkl_dnn { template class Pool final : public onnxruntime::Pool { public: - Pool(const OpKernelInfo& info) : onnxruntime::Pool(info) { + 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;