diff --git a/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.cc b/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.cc index 84014da74a2..70466010dac 100644 --- a/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.cc +++ b/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.cc @@ -21,6 +21,7 @@ #include #include "utils/ms_utils.h" +#include "utils/profile.h" namespace mindspore { namespace kernel { @@ -105,28 +106,67 @@ void PoolingGradCpuKernelMod::InitKernel(const CNodePtr &kernel_node) { // Pooling_avg forward description const auto desc = CreateDesc(dnnl::prop_kind::forward_training, algorithm_, src_desc_, dst_desc_, strides, kernel, padding_l, padding_r); - forward_prim_desc_ = CreateDesc(desc, engine_); + auto forward_prim_desc = CreateDesc(desc, engine_); // Pooling_avg backward description const auto backward_desc = CreateDesc(algorithm_, src_desc_, dst_desc_, strides, kernel, padding_l, padding_r); const auto backward_prim_desc = - CreateDesc(backward_desc, engine_, forward_prim_desc_); + CreateDesc(backward_desc, engine_, forward_prim_desc); primitive_ = CreatePrimitive(backward_prim_desc); AddArgument(DNNL_ARG_DIFF_SRC, src_desc_); AddArgument(DNNL_ARG_DIFF_DST, dst_desc_); // For pooling_max, need a workspace that generated in forward and stored the max value indexes to compute grad. if (algorithm_ == dnnl::algorithm::pooling_max) { - workspace_desc_ = GetWorkspaceDesc(forward_prim_desc_); + primitive_forward_ = CreatePrimitive(forward_prim_desc); + workspace_desc_ = GetWorkspaceDesc(forward_prim_desc); AddArgument(DNNL_ARG_WORKSPACE, workspace_desc_); } } -void PoolingGradCpuKernelMod::ComputeMaxValueIndex(void *src, void *dst, void *work_array) const { +#ifdef USE_MS_THREADPOOL_FOR_DNNL +void PoolingGradCpuKernelMod::ExecuteForwardByMSThreadPool(const std::unordered_map &arguments) { + const size_t MAX_POW = 6; + const size_t AVG_COUNT = 5; + const size_t DIFF = 2; + size_t current_pow = forward_parallel_info_.search_count / AVG_COUNT; + int current_thread_nums = static_cast(std::pow(2.0f, current_pow)); + auto mkl_pool = dynamic_cast(mkl_threadpool_.get()); + if (current_pow >= MAX_POW) { + int best_thread_nums = static_cast(std::pow(2.0f, forward_parallel_info_.best_pow)); + mkl_pool->set_num_threads(best_thread_nums); + MS_LOG(DEBUG) << "begin to invoke primitive::execute"; + primitive_forward_->execute(stream_, arguments); + MS_LOG(DEBUG) << "end to invoke primitive::execute"; + return; + } + + if (forward_parallel_info_.search_count % AVG_COUNT == 0) { + forward_parallel_info_.tmp_sum_cost_time = 0; + } + double start_time = GetTime(); + mkl_pool->set_num_threads(current_thread_nums); + MS_LOG(DEBUG) << "begin to invoke primitive::execute"; + primitive_forward_->execute(stream_, arguments); + MS_LOG(DEBUG) << "end to invoke primitive::execute"; + double cost_time = GetTime() - start_time; + forward_parallel_info_.tmp_sum_cost_time += cost_time; + forward_parallel_info_.search_count++; + if (forward_parallel_info_.search_count % AVG_COUNT == 0) { + if (forward_parallel_info_.min_cost_time > forward_parallel_info_.tmp_sum_cost_time) { + forward_parallel_info_.min_cost_time = forward_parallel_info_.tmp_sum_cost_time; + forward_parallel_info_.best_pow = current_pow; + } else if (current_pow - forward_parallel_info_.best_pow >= DIFF) { + forward_parallel_info_.search_count = AVG_COUNT * MAX_POW; + } + } +} +#endif + +void PoolingGradCpuKernelMod::ComputeMaxValueIndex(void *src, void *dst, void *work_array) { // Compute maxvalue index for pooling_backward_max. MS_LOG(INFO) << "Compute maxvalue index for " << kernel_name_; - auto primitive_forward = CreatePrimitive(forward_prim_desc_); std::unordered_map arguments; dnnl::memory src_mem = dnnl::memory(src_desc_, engine_, nullptr); dnnl::memory dst_mem = dnnl::memory(dst_desc_, engine_, nullptr); @@ -137,8 +177,15 @@ void PoolingGradCpuKernelMod::ComputeMaxValueIndex(void *src, void *dst, void *w arguments[DNNL_ARG_SRC] = src_mem; arguments[DNNL_ARG_DST] = dst_mem; arguments[DNNL_ARG_WORKSPACE] = work_mem; - dnnl::stream stream(engine_); - primitive_forward->execute(stream, arguments); + +#ifdef USE_MS_THREADPOOL_FOR_DNNL + ExecuteForwardByMSThreadPool(arguments); +#else + MS_LOG(DEBUG) << "begin to invoke primitive::execute"; + primitive_forward_->execute(stream_, arguments); + MS_LOG(DEBUG) << "end to invoke primitive::execute"; +#endif + (void)stream_.wait(); } bool PoolingGradCpuKernelMod::Launch(const std::vector &inputs, diff --git a/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.h b/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.h index 1ced3cd49d2..4a61d5afe20 100644 --- a/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.h +++ b/mindspore/ccsrc/plugin/device/cpu/kernel/mkldnn/pooling_grad_cpu_kernel.h @@ -21,7 +21,7 @@ #include #include #include -#include +#include #include "plugin/device/cpu/kernel/mkldnn/pooling_cpu_kernel.h" @@ -45,7 +45,7 @@ class PoolingGradCpuKernelMod : public PoolingCpuKernelMod { protected: std::vector GetOpSupport() override { - static std::map> support_list = { + static std::unordered_map> support_list = { {kAvgPoolGrad, {{KernelAttr() .AddInputAttr(kNumberTypeFloat32) @@ -76,13 +76,17 @@ class PoolingGradCpuKernelMod : public PoolingCpuKernelMod { private: void InitPoolingGradFields(const CNodePtr &kernel_node); void InitInputOutputSize(const CNodePtr &kernel_node) override; - void ComputeMaxValueIndex(void *src, void *dst, void *work_array) const; + void ComputeMaxValueIndex(void *src, void *dst, void *work_array); +#ifdef USE_MS_THREADPOOL_FOR_DNNL + void ExecuteForwardByMSThreadPool(const std::unordered_map &arguments); +#endif + size_t grad_index_{0}; dnnl::memory::desc src_desc_{}; dnnl::memory::desc dst_desc_{}; dnnl::memory::desc workspace_desc_{}; - dnnl::pooling_forward::primitive_desc forward_prim_desc_{}; - size_t grad_index_{0}; + std::shared_ptr primitive_forward_{nullptr}; + ParallelSearchInfo forward_parallel_info_{}; std::string kernel_type_{kUnknown}; }; } // namespace kernel