From b4e608cf47a838cd20ec6f07341defba7de63884 Mon Sep 17 00:00:00 2001 From: Chen Xu Date: Tue, 20 Jun 2023 13:15:18 +0800 Subject: [PATCH] [Snippets] Implement shuffling based horizontal reduction emitter (#18099) --- .../src/emitters/x64/cpu_generator.cpp | 4 +- .../emitters/x64/jit_snippets_emitters.cpp | 100 ++++++++---------- .../emitters/x64/jit_snippets_emitters.hpp | 28 ++--- 3 files changed, 51 insertions(+), 81 deletions(-) diff --git a/src/plugins/intel_cpu/src/emitters/x64/cpu_generator.cpp b/src/plugins/intel_cpu/src/emitters/x64/cpu_generator.cpp index 1244fac99ad..6d776ab57eb 100644 --- a/src/plugins/intel_cpu/src/emitters/x64/cpu_generator.cpp +++ b/src/plugins/intel_cpu/src/emitters/x64/cpu_generator.cpp @@ -139,8 +139,8 @@ ov::intel_cpu::CPUTargetMachine::CPUTargetMachine(dnnl::impl::cpu::x64::cpu_isa_ jitters[ngraph::op::v7::Gelu::get_type_info_static()] = CREATE_EMITTER(ov::intel_cpu::jit_gelu_v7_emitter); jitters[snippets::op::Fill::get_type_info_static()] = CREATE_EMITTER(FillEmitter); - jitters[snippets::op::HorizonMax::get_type_info_static()] = CREATE_EMITTER(HorizonMaxEmitter); - jitters[snippets::op::HorizonSum::get_type_info_static()] = CREATE_EMITTER(HorizonSumEmitter); + jitters[snippets::op::HorizonMax::get_type_info_static()] = CREATE_EMITTER(HorizonEmitter); + jitters[snippets::op::HorizonSum::get_type_info_static()] = CREATE_EMITTER(HorizonEmitter); jitters[snippets::op::Kernel::get_type_info_static()] = CREATE_EMITTER(KernelEmitter); jitters[snippets::op::LoopBegin::get_type_info_static()] = CREATE_EMITTER(LoopBeginEmitter); diff --git a/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.cpp b/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.cpp index c339b72cfd1..09bdf0efd29 100644 --- a/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.cpp +++ b/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.cpp @@ -1334,10 +1334,18 @@ void BrgemmCopyBEmitter::execute(matmul::jit_brgemm_matmul_copy_b_t *kernel, con (*kernel)(&ctx); } -HorizonMaxEmitter::HorizonMaxEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n) : - jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) {} +HorizonEmitter::HorizonEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n) : + jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) { + if (ov::is_type(n)) { + m_op_type = OpType::max; + } else if (ov::is_type(n)) { + m_op_type = OpType::sum; + } else { + OPENVINO_THROW("HorizonEmitter exprects HorizonMax or HorizonSum ops"); + } +} -void HorizonMaxEmitter::emit_impl(const std::vector& in, +void HorizonEmitter::emit_impl(const std::vector& in, const std::vector& out) const { if (host_isa_ == dnnl::impl::cpu::x64::sse41) { emit_isa(in, out); @@ -1351,71 +1359,49 @@ void HorizonMaxEmitter::emit_impl(const std::vector& in, } template -void HorizonMaxEmitter::emit_isa(const std::vector &in, const std::vector &out) const { +void HorizonEmitter::emit_isa(const std::vector &in, const std::vector &out) const { using Vmm = typename dnnl::impl::utils::conditional3::type; Vmm src_vmm = Vmm(in[0]); - Xmm dst_xmm = Xmm(out[0]); - Xmm aux_xmm = Xmm(aux_vec_idxs[0]); + Vmm dst_vmm = Vmm(out[0]); + Vmm aux_vmm = Vmm(aux_vec_idxs[0]); - Reg64 aux_reg = Reg64(aux_gpr_idxs[0]); - - const size_t vlen = dnnl::impl::cpu::x64::cpu_isa_traits::vlen; - const size_t vec_size = vlen / sizeof(float); - h->sub(h->rsp, vlen); - h->uni_vmovups(h->ptr[h->rsp], src_vmm); - // Let the first value be the max - h->mov(aux_reg, h->ptr[h->rsp]); - h->vmovq(dst_xmm, aux_reg); - for (size_t i = 1; i < vec_size; i++) { - h->mov(aux_reg, h->ptr[h->rsp + i * sizeof(float)]); - h->vmovq(aux_xmm, aux_reg); - h->uni_vmaxps(dst_xmm, dst_xmm, aux_xmm); + if (in[0] != out[0]) + h->uni_vmovups(dst_vmm, src_vmm); + if (isa == dnnl::impl::cpu::x64::avx512_core) { + Zmm dst_zmm = Zmm(out[0]); + Zmm aux_zmm = Zmm(aux_vec_idxs[0]); + h->vshuff32x4(aux_zmm, dst_zmm, dst_zmm, 0x4E); + perform_op(dst_zmm, dst_zmm, aux_zmm); + h->vshuff32x4(aux_zmm, dst_zmm, dst_zmm, 0xB1); + perform_op(dst_zmm, dst_zmm, aux_zmm); + } else if (isa == dnnl::impl::cpu::x64::avx2) { + Ymm dst_ymm = Ymm(out[0]); + Ymm aux_ymm = Ymm(aux_vec_idxs[0]); + h->vperm2i128(aux_ymm, dst_ymm, dst_ymm, 0x01); + perform_op(dst_ymm, dst_ymm, aux_ymm); } - h->add(h->rsp, vlen); + h->uni_vshufps(aux_vmm, dst_vmm, dst_vmm, 0x4E); + perform_op(dst_vmm, dst_vmm, aux_vmm); + h->uni_vshufps(aux_vmm, dst_vmm, dst_vmm, 0xB1); + perform_op(dst_vmm, dst_vmm, aux_vmm); } -HorizonSumEmitter::HorizonSumEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n) : - jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) {} - -void HorizonSumEmitter::emit_impl(const std::vector& in, - const std::vector& out) const { - if (host_isa_ == dnnl::impl::cpu::x64::sse41) { - emit_isa(in, out); - } else if (host_isa_ == dnnl::impl::cpu::x64::avx2) { - emit_isa(in, out); - } else if (host_isa_ == dnnl::impl::cpu::x64::avx512_core) { - emit_isa(in, out); - } else { - IE_THROW() << "HorizonSum emitter doesn't support " << host_isa_; +template +void HorizonEmitter::perform_op(const Vmm &vmm1, const Vmm &vmm2, const Vmm &vmm3) const { + switch (m_op_type) { + case OpType::max: + h->uni_vmaxps(vmm1, vmm2, vmm3); + break; + case OpType::sum: + h->uni_vaddps(vmm1, vmm2, vmm3); + break; + default: + assert(!"Unsupported horizontal operation."); } } -template -void HorizonSumEmitter::emit_isa(const std::vector &in, const std::vector &out) const { - using Vmm = typename dnnl::impl::utils::conditional3::type; - - Vmm src_vmm = Vmm(in[0]); - Xmm dst_xmm = Xmm(out[0]); - Xmm aux_xmm = Xmm(aux_vec_idxs[0]); - - Reg64 aux_reg = Reg64(aux_gpr_idxs[0]); - - const size_t vlen = dnnl::impl::cpu::x64::cpu_isa_traits::vlen; - const size_t vec_size = vlen / sizeof(float); - h->sub(h->rsp, vlen); - h->uni_vmovups(h->ptr[h->rsp], src_vmm); - h->uni_vpxor(dst_xmm, dst_xmm, dst_xmm); - for (size_t i = 0; i < vec_size; i++) { - h->mov(aux_reg, h->ptr[h->rsp + i * sizeof(float)]); - h->vmovq(aux_xmm, aux_reg); - h->uni_vaddps(dst_xmm, dst_xmm, aux_xmm); - } - h->add(h->rsp, vlen); -} - VectorBufferEmitter::VectorBufferEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n) : jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) {} diff --git a/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.hpp b/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.hpp index cc4ba3a55f8..a4c3e1f835e 100644 --- a/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.hpp +++ b/src/plugins/intel_cpu/src/emitters/x64/jit_snippets_emitters.hpp @@ -417,9 +417,9 @@ private: size_t m_comp_offset = 0lu; }; -class HorizonMaxEmitter : public jit_emitter { +class HorizonEmitter : public jit_emitter { public: - HorizonMaxEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n); + HorizonEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n); size_t get_inputs_num() const override {return 1;} static std::set> get_supported_precisions(const std::shared_ptr& node = nullptr) { @@ -427,7 +427,6 @@ public: } protected: - size_t aux_gprs_count() const override {return 1;} size_t aux_vecs_count() const override {return 1;} private: @@ -436,27 +435,12 @@ private: template void emit_isa(const std::vector &in, const std::vector &out) const; -}; -class HorizonSumEmitter : public jit_emitter { -public: - HorizonSumEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr& n); + template + void perform_op(const Vmm &vmm1, const Vmm &vmm2, const Vmm &vmm3) const; - size_t get_inputs_num() const override {return 1;} - static std::set> get_supported_precisions(const std::shared_ptr& node = nullptr) { - return {{element::f32}}; - } - -protected: - size_t aux_gprs_count() const override {return 1;} - size_t aux_vecs_count() const override {return 1;} - -private: - void emit_impl(const std::vector& in, - const std::vector& out) const override; - - template - void emit_isa(const std::vector &in, const std::vector &out) const; + enum class OpType { max, sum }; + OpType m_op_type = OpType::max; }; class VectorBufferEmitter : public jit_emitter {