[Snippets] Implement shuffling based horizontal reduction emitter (#18099)
This commit is contained in:
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user