4#include "../../InternalHeaderCheck.h"
9#if EIGEN_ARCH_ARM && EIGEN_COMP_CLANG
15struct gebp_traits<float, float, false, false, Architecture::NEON, GEBPPacketFull>
16 : gebp_traits<float, float, false, false, Architecture::Generic, GEBPPacketFull> {
17 EIGEN_STRONG_INLINE
void acc(
const AccPacket& c,
const ResPacket& alpha, ResPacket& r)
const {
20 asm volatile(
"vmla.f32 %q[r], %q[c], %q[alpha]" : [r]
"+w"(r) : [c]
"w"(c), [alpha]
"w"(alpha) :);
23 template <
typename LaneIdType>
24 EIGEN_STRONG_INLINE
void madd(
const Packet4f& a,
const Packet4f& b, Packet4f& c, Packet4f&,
const LaneIdType&)
const {
28 template <
typename LaneIdType>
29 EIGEN_STRONG_INLINE
void madd(
const Packet4f& a,
const QuadPacket<Packet4f>& b, Packet4f& c, Packet4f& tmp,
30 const LaneIdType& lane)
const {
31 madd(a, b.get(lane), c, tmp, lane);
39#ifndef EIGEN_NEON_GEBP_NR
40#define EIGEN_NEON_GEBP_NR 8
44struct gebp_traits<float, float, false, false, Architecture::NEON, GEBPPacketFull>
45 : gebp_traits<float, float, false, false, Architecture::Generic, GEBPPacketFull> {
46 typedef float RhsPacket;
47 typedef float32x4_t RhsPacketx4;
48 enum { nr = EIGEN_NEON_GEBP_NR };
49 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacket& dest)
const { dest = *b; }
51 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const { dest = vld1q_f32(b); }
53 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, RhsPacket& dest)
const { dest = *b; }
55 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
57 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, RhsPacket& dest)
const { loadRhs(b, dest); }
59 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacket& b, AccPacket& c, RhsPacket& ,
60 const FixedInt<0>&)
const {
61 c = vfmaq_n_f32(c, a, b);
66 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
67 const FixedInt<0>&)
const {
68 madd_helper<0>(a, b, c);
70 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
71 const FixedInt<1>&)
const {
72 madd_helper<1>(a, b, c);
74 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
75 const FixedInt<2>&)
const {
76 madd_helper<2>(a, b, c);
78 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
79 const FixedInt<3>&)
const {
80 madd_helper<3>(a, b, c);
85 EIGEN_STRONG_INLINE
void madd_helper(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c)
const {
86#if EIGEN_GNUC_STRICT_LESS_THAN(9, 0, 0)
91 asm(
"fmla %0.4s, %1.4s, %2.s[0]\n" :
"+w"(c) :
"w"(a),
"w"(b) :);
93 asm(
"fmla %0.4s, %1.4s, %2.s[1]\n" :
"+w"(c) :
"w"(a),
"w"(b) :);
95 asm(
"fmla %0.4s, %1.4s, %2.s[2]\n" :
"+w"(c) :
"w"(a),
"w"(b) :);
97 asm(
"fmla %0.4s, %1.4s, %2.s[3]\n" :
"+w"(c) :
"w"(a),
"w"(b) :);
99 c = vfmaq_laneq_f32(c, a, b, LaneID);
105struct gebp_traits<double, double, false, false, Architecture::NEON>
106 : gebp_traits<double, double, false, false, Architecture::Generic> {
107 typedef double RhsPacket;
108 enum { nr = EIGEN_NEON_GEBP_NR };
110 float64x2_t B_0, B_1;
113 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacket& dest)
const { dest = *b; }
115 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const {
116 dest.B_0 = vld1q_f64(b);
117 dest.B_1 = vld1q_f64(b + 2);
120 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, RhsPacket& dest)
const { loadRhs(b, dest); }
122 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
124 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, RhsPacket& dest)
const { loadRhs(b, dest); }
126 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacket& b, AccPacket& c, RhsPacket& ,
127 const FixedInt<0>&)
const {
128 c = vfmaq_n_f64(c, a, b);
134 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
135 const FixedInt<0>&)
const {
136 madd_helper<0>(a, b, c);
138 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
139 const FixedInt<1>&)
const {
140 madd_helper<1>(a, b, c);
142 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
143 const FixedInt<2>&)
const {
144 madd_helper<2>(a, b, c);
146 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
147 const FixedInt<3>&)
const {
148 madd_helper<3>(a, b, c);
152 template <
int LaneID>
153 EIGEN_STRONG_INLINE
void madd_helper(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c)
const {
154#if EIGEN_GNUC_STRICT_LESS_THAN(9, 0, 0)
159 asm(
"fmla %0.2d, %1.2d, %2.d[0]\n" :
"+w"(c) :
"w"(a),
"w"(b.B_0) :);
160 else if (LaneID == 1)
161 asm(
"fmla %0.2d, %1.2d, %2.d[1]\n" :
"+w"(c) :
"w"(a),
"w"(b.B_0) :);
162 else if (LaneID == 2)
163 asm(
"fmla %0.2d, %1.2d, %2.d[0]\n" :
"+w"(c) :
"w"(a),
"w"(b.B_1) :);
164 else if (LaneID == 3)
165 asm(
"fmla %0.2d, %1.2d, %2.d[1]\n" :
"+w"(c) :
"w"(a),
"w"(b.B_1) :);
168 c = vfmaq_laneq_f64(c, a, b.B_0, 0);
169 else if (LaneID == 1)
170 c = vfmaq_laneq_f64(c, a, b.B_0, 1);
171 else if (LaneID == 2)
172 c = vfmaq_laneq_f64(c, a, b.B_1, 0);
173 else if (LaneID == 3)
174 c = vfmaq_laneq_f64(c, a, b.B_1, 1);
183#if EIGEN_HAS_ARM64_FP16 && EIGEN_COMP_CLANG
186struct gebp_traits<half, half, false, false, Architecture::NEON>
187 : gebp_traits<half, half, false, false, Architecture::Generic> {
188 typedef half RhsPacket;
189 typedef float16x4_t RhsPacketx4;
190 typedef float16x4_t PacketHalf;
191 enum { nr = EIGEN_NEON_GEBP_NR };
193 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacket& dest)
const { dest = *b; }
195 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const { dest = vld1_f16(&b->x); }
197 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, RhsPacket& dest)
const { dest = *b; }
199 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
201 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar*, RhsPacket&)
const {
204 eigen_assert(
false &&
"Cannot loadRhsQuad for a scalar RHS.");
207 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacket& b, AccPacket& c, RhsPacket& ,
208 const FixedInt<0>&)
const {
209#if EIGEN_HAS_ARM64_FP16_VECTOR_ARITHMETIC
210 c = vfmaq_n_f16(c, a, b);
212 c = vcombine_f16(vcvt_f16_f32(vfmaq_n_f32(vcvt_f32_f16(vget_low_f16(c)), vcvt_f32_f16(vget_low_f16(a)), b)),
213 vcvt_f16_f32(vfmaq_n_f32(vcvt_f32_f16(vget_high_f16(c)), vcvt_f32_f16(vget_high_f16(a)), b)));
216 EIGEN_STRONG_INLINE
void madd(
const PacketHalf& a,
const RhsPacket& b, PacketHalf& c, RhsPacket& ,
217 const FixedInt<0>&)
const {
218#if EIGEN_HAS_ARM64_FP16_VECTOR_ARITHMETIC
219 c = vfma_n_f16(c, a, b);
221 c = vcvt_f16_f32(vfmaq_n_f32(vcvt_f32_f16(c), vcvt_f32_f16(a), b));
227 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
228 const FixedInt<0>&)
const {
229 madd_helper<0>(a, b, c);
231 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
232 const FixedInt<1>&)
const {
233 madd_helper<1>(a, b, c);
235 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
236 const FixedInt<2>&)
const {
237 madd_helper<2>(a, b, c);
239 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c, RhsPacket& ,
240 const FixedInt<3>&)
const {
241 madd_helper<3>(a, b, c);
245 template <
int LaneID>
246 EIGEN_STRONG_INLINE
void madd_helper(
const LhsPacket& a,
const RhsPacketx4& b, AccPacket& c)
const {
247#if EIGEN_HAS_ARM64_FP16_VECTOR_ARITHMETIC
248 c = vfmaq_lane_f16(c, a, b, LaneID);
250 float32x2_t bf = get_half<LaneID>(vcvt_f32_f16(b));
252 vcvt_f16_f32(vfmaq_lane_f32(vcvt_f32_f16(vget_low_f16(c)), vcvt_f32_f16(vget_low_f16(a)), bf, LaneID % 2)),
253 vcvt_f16_f32(vfmaq_lane_f32(vcvt_f32_f16(vget_high_f16(c)), vcvt_f32_f16(vget_high_f16(a)), bf, LaneID % 2)));
257#if !EIGEN_HAS_ARM64_FP16_VECTOR_ARITHMETIC
258 template <
int LaneID>
259 EIGEN_STRONG_INLINE
static std::enable_if_t<(LaneID <= 1), float32x2_t> get_half(float32x4_t vec) {
260 return vget_low_f32(vec);
263 template <
int LaneID>
264 EIGEN_STRONG_INLINE
static std::enable_if_t<(LaneID > 1), float32x2_t> get_half(float32x4_t vec) {
265 return vget_high_f32(vec);