[Snippets] Implement shuffling based horizontal reduction emitter (#18099)

This commit is contained in:
Chen Xu
2023-06-20 09:15:18 +04:00
committed by GitHub
parent a9c4e4ab56
commit b4e608cf47
3 changed files with 51 additions and 81 deletions
@@ -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);
@@ -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<ov::Node>& 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<ov::Node>& n) :
jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) {
if (ov::is_type<const snippets::op::HorizonMax>(n)) {
m_op_type = OpType::max;
} else if (ov::is_type<const snippets::op::HorizonSum>(n)) {
m_op_type = OpType::sum;
} else {
OPENVINO_THROW("HorizonEmitter exprects HorizonMax or HorizonSum ops");
}
}
void HorizonMaxEmitter::emit_impl(const std::vector<size_t>& in,
void HorizonEmitter::emit_impl(const std::vector<size_t>& in,
const std::vector<size_t>& out) const {
if (host_isa_ == dnnl::impl::cpu::x64::sse41) {
emit_isa<dnnl::impl::cpu::x64::sse41>(in, out);
@@ -1351,71 +1359,49 @@ void HorizonMaxEmitter::emit_impl(const std::vector<size_t>& in,
}
template <dnnl::impl::cpu::x64::cpu_isa_t isa>
void HorizonMaxEmitter::emit_isa(const std::vector<size_t> &in, const std::vector<size_t> &out) const {
void HorizonEmitter::emit_isa(const std::vector<size_t> &in, const std::vector<size_t> &out) const {
using Vmm = typename dnnl::impl::utils::conditional3<isa == dnnl::impl::cpu::x64::sse41,
Xmm, isa == dnnl::impl::cpu::x64::avx2, Ymm, Zmm>::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<isa>::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<Zmm>(dst_zmm, dst_zmm, aux_zmm);
h->vshuff32x4(aux_zmm, dst_zmm, dst_zmm, 0xB1);
perform_op<Zmm>(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<Ymm>(dst_ymm, dst_ymm, aux_ymm);
}
h->add(h->rsp, vlen);
h->uni_vshufps(aux_vmm, dst_vmm, dst_vmm, 0x4E);
perform_op<Xmm>(dst_vmm, dst_vmm, aux_vmm);
h->uni_vshufps(aux_vmm, dst_vmm, dst_vmm, 0xB1);
perform_op<Xmm>(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<ov::Node>& n) :
jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) {}
void HorizonSumEmitter::emit_impl(const std::vector<size_t>& in,
const std::vector<size_t>& out) const {
if (host_isa_ == dnnl::impl::cpu::x64::sse41) {
emit_isa<dnnl::impl::cpu::x64::sse41>(in, out);
} else if (host_isa_ == dnnl::impl::cpu::x64::avx2) {
emit_isa<dnnl::impl::cpu::x64::avx2>(in, out);
} else if (host_isa_ == dnnl::impl::cpu::x64::avx512_core) {
emit_isa<dnnl::impl::cpu::x64::avx512_core>(in, out);
} else {
IE_THROW() << "HorizonSum emitter doesn't support " << host_isa_;
template<typename Vmm>
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 <dnnl::impl::cpu::x64::cpu_isa_t isa>
void HorizonSumEmitter::emit_isa(const std::vector<size_t> &in, const std::vector<size_t> &out) const {
using Vmm = typename dnnl::impl::utils::conditional3<isa == dnnl::impl::cpu::x64::sse41,
Xmm, isa == dnnl::impl::cpu::x64::avx2, Ymm, Zmm>::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<isa>::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<ov::Node>& n) :
jit_emitter(h, isa, n, Precision::FP32, emitter_in_out_map::vec_to_vec) {}
@@ -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<ov::Node>& n);
HorizonEmitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const std::shared_ptr<ov::Node>& n);
size_t get_inputs_num() const override {return 1;}
static std::set<std::vector<element::Type>> get_supported_precisions(const std::shared_ptr<ngraph::Node>& 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 <dnnl::impl::cpu::x64::cpu_isa_t isa>
void emit_isa(const std::vector<size_t> &in, const std::vector<size_t> &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<ov::Node>& n);
template<typename Vmm>
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<std::vector<element::Type>> get_supported_precisions(const std::shared_ptr<ngraph::Node>& 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<size_t>& in,
const std::vector<size_t>& out) const override;
template <dnnl::impl::cpu::x64::cpu_isa_t isa>
void emit_isa(const std::vector<size_t> &in, const std::vector<size_t> &out) const;
enum class OpType { max, sum };
OpType m_op_type = OpType::max;
};
class VectorBufferEmitter : public jit_emitter {