10#ifndef EIGEN_SME_GENERALBLOCKPANELKERNEL_H
11#define EIGEN_SME_GENERALBLOCKPANELKERNEL_H
14#include "../../InternalHeaderCheck.h"
64static constexpr int kSmeDesignVectorBytes = 64;
73template <
typename Scalar>
75 static constexpr int kGridRows = 2;
76 static constexpr int kGridCols = 2;
77 static constexpr int mr = kGridRows * kSmeDesignVectorBytes / int(
sizeof(Scalar));
78 static constexpr int nr = kGridCols * kSmeDesignVectorBytes / int(
sizeof(Scalar));
81template <
typename RealScalar>
82struct sme_block<std::complex<RealScalar>> {
83 static constexpr int kGridCols = 2;
84 static constexpr int kGridRows = sme_tile_count<RealScalar>::value / (2 * kGridCols);
85 static constexpr int mr = kGridRows * kSmeDesignVectorBytes / int(
sizeof(RealScalar));
86 static constexpr int nr = kGridCols * kSmeDesignVectorBytes / int(
sizeof(RealScalar));
89static constexpr int kSmeMr = sme_block<float>::mr;
90static constexpr int kSmeNr = sme_block<float>::nr;
91static constexpr int kSmeMrC = sme_block<std::complex<float>>::mr;
92static constexpr int kSmeNrC = sme_block<std::complex<float>>::nr;
93#ifdef EIGEN_VECTORIZE_SME_F64F64
94static constexpr int kSmeMrD = sme_block<double>::mr;
95static constexpr int kSmeNrD = sme_block<double>::nr;
96static constexpr int kSmeMrCD = sme_block<std::complex<double>>::mr;
97static constexpr int kSmeNrCD = sme_block<std::complex<double>>::nr;
125template <
bool Conjugate,
typename Scalar,
typename Index>
126static EIGEN_ALWAYS_INLINE
void sve_copy_panel_range(Scalar* EIGEN_RESTRICT dst,
const Scalar* EIGEN_RESTRICT src,
127 Index src_stride, Index k0, Index k1,
128 int width) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
129 using Traits = sme_packet_traits<Scalar>;
130 const int svl = Traits::size();
131 if (width == 2 * svl) {
133 const svcount_t pn = Traits::ptrue_c();
135 for (; k + 4 <= k1; k += 4) {
136 const auto v0 = ploadu_x2(pn, &src[k * src_stride]);
137 const auto v1 = ploadu_x2(pn, &src[(k + 1) * src_stride]);
138 const auto v2 = ploadu_x2(pn, &src[(k + 2) * src_stride]);
139 const auto v3 = ploadu_x2(pn, &src[(k + 3) * src_stride]);
140 pstoreu_x2(pn, &dst[k * width], v0);
141 pstoreu_x2(pn, &dst[(k + 1) * width], v1);
142 pstoreu_x2(pn, &dst[(k + 2) * width], v2);
143 pstoreu_x2(pn, &dst[(k + 3) * width], v3);
145 for (; k < k1; ++k) pstoreu_x2(pn, &dst[k * width], ploadu_x2(pn, &src[k * src_stride]));
148 for (
int off = 0; off < width; off += svl) {
149 const svbool_t pred = sme_packet_traits<Scalar>::whilelt(off, width);
150 for (Index k = k0; k < k1; ++k) {
151 pstoreu(pred, &dst[k * width + off], ploadu(pred, &src[k * src_stride + off]));
158template <
typename Scalar,
typename Index>
159static EIGEN_ALWAYS_INLINE
void sve_copy_panel_quad(Scalar* EIGEN_RESTRICT dst, Index panel_stride,
160 const Scalar* EIGEN_RESTRICT src, Index src_stride,
161 Index depth) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
162 using Traits = sme_packet_traits<Scalar>;
163 const int w = 2 * Traits::size();
164 const svcount_t pn = Traits::ptrue_c();
165 const int strip_bytes = 4 * w * int(
sizeof(Scalar));
166 for (Index k = 0; k < depth; ++k) {
167 const Scalar* c = src + k * src_stride;
169 for (
int l = 0; l < strip_bytes; l += 128)
170 __builtin_prefetch(
reinterpret_cast<const char*
>(c + 4 * src_stride) + l, 0, 2);
171 const auto a = ploadu_x4(pn, c), b = ploadu_x4(pn, c + 2 * w);
172 pstoreu_x2(pn, dst + k * w, pcreate(pget<0>(a), pget<1>(a)));
173 pstoreu_x2(pn, dst + panel_stride + k * w, pcreate(pget<2>(a), pget<3>(a)));
174 pstoreu_x2(pn, dst + 2 * panel_stride + k * w, pcreate(pget<0>(b), pget<1>(b)));
175 pstoreu_x2(pn, dst + 3 * panel_stride + k * w, pcreate(pget<2>(b), pget<3>(b)));
181template <
bool PanelMode,
typename Scalar,
typename Index>
182static EIGEN_ALWAYS_INLINE Index sme_pack_quad_panels(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src,
183 Index src_stride, Index depth, Index rows, Index mr,
185 Index dst_offset) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
186 if (mr != 2 * sme_packet_traits<Scalar>::size())
return 0;
187 const Index panel_stride = PanelMode ? mr * dst_stride : mr * depth;
189 for (; i + 4 * mr <= rows; i += 4 * mr) {
190 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * mr : dst_base + i * depth;
191 sve_copy_panel_quad(dst_panel, panel_stride, src + i, src_stride, depth);
195template <
bool PanelMode,
typename RealScalar,
typename Index>
196static EIGEN_ALWAYS_INLINE Index sme_pack_quad_panels(std::complex<RealScalar>*,
const std::complex<RealScalar>*, Index,
197 Index, Index, Index, Index,
198 Index) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
206template <
bool Conjugate,
typename RealScalar,
typename Index>
207static EIGEN_ALWAYS_INLINE
void sve_copy_panel_range(std::complex<RealScalar>* EIGEN_RESTRICT dst,
208 const std::complex<RealScalar>* EIGEN_RESTRICT src,
209 Index src_stride, Index k0, Index k1,
210 int width) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
211 using Traits = sme_packet_traits<RealScalar>;
212 using Vec =
typename Traits::type;
213 RealScalar* EIGEN_RESTRICT rdst =
reinterpret_cast<RealScalar*
>(dst);
214 const RealScalar* EIGEN_RESTRICT rsrc =
reinterpret_cast<const RealScalar*
>(src);
215 const int svl = Traits::size();
216 const Index step = Index(2 * width);
217 for (
int off = 0; off < width; off += svl) {
218 const int w = width - off;
219 const int lanes = 2 * w;
220 const svbool_t pg_w = Traits::whilelt(off, width);
221 const svbool_t pg_lo = Traits::whilelt(0, lanes);
222 const svbool_t pg_hi = Traits::whilelt(svl, lanes);
223 for (Index k = k0; k < k1; ++k) {
224 const RealScalar* p = rsrc + Index(2) * (k * src_stride + Index(off));
225 const Vec v_lo = ploadu(pg_lo, p);
229 const Vec v_hi = ploadu(pg_hi, sme_offset(p, svl));
230 Vec im = puzp2(v_lo, v_hi);
231 EIGEN_IF_CONSTEXPR (Conjugate) {
232 im = pnegate(pg_w, im);
234 pstoreu(pg_w, &rdst[k * step + Index(off)], puzp1(v_lo, v_hi));
235 pstoreu(pg_w, &rdst[k * step + Index(width + off)], im);
242template <
bool Conjugate,
typename Scalar,
typename Index>
243static EIGEN_ALWAYS_INLINE
void sve_copy_panel(Scalar* EIGEN_RESTRICT dst,
const Scalar* EIGEN_RESTRICT src,
244 Index src_stride, Index depth,
245 int width) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
246 sve_copy_panel_range<Conjugate>(dst, src, src_stride, Index(0), depth, width);
251template <
typename RealScalar,
typename Index>
252static EIGEN_ALWAYS_INLINE Index sme_transpose_pack_pair(RealScalar* EIGEN_RESTRICT dst,
253 const RealScalar* EIGEN_RESTRICT src, Index src_stride,
254 Index k0, Index k1) __arm_streaming __arm_inout(
"za") {
255 using Traits = sme_packet_traits<RealScalar>;
256 const int svl = Traits::size(), w = 2 * svl;
257 const svcount_t pn = Traits::ptrue_c();
259 for (; k + w <= k1; k += w) {
260 const bool pf = k + 2 * w <= k1;
261 for (
int r = 0; r < svl; r += 4) {
262 const RealScalar* p = src + Index(r) * src_stride + k;
263 const RealScalar* q = p + Index(svl) * src_stride;
265 for (
int u = 0; u < 4; ++u) {
266 __builtin_prefetch(p + u * src_stride + 2 * w, 0, 2);
267 __builtin_prefetch(q + u * src_stride + 2 * w, 0, 2);
269 const auto a0 = ploadu_x2(pn, p), a1 = ploadu_x2(pn, p + src_stride), a2 = ploadu_x2(pn, p + 2 * src_stride),
270 a3 = ploadu_x2(pn, p + 3 * src_stride);
271 sme_write_hor_za_vg4<0>(uint32_t(r), pget<0>(a0), pget<0>(a1), pget<0>(a2), pget<0>(a3));
272 sme_write_hor_za_vg4<1>(uint32_t(r), pget<1>(a0), pget<1>(a1), pget<1>(a2), pget<1>(a3));
273 const auto b0 = ploadu_x2(pn, q), b1 = ploadu_x2(pn, q + src_stride), b2 = ploadu_x2(pn, q + 2 * src_stride),
274 b3 = ploadu_x2(pn, q + 3 * src_stride);
275 sme_write_hor_za_vg4<2>(uint32_t(r), pget<0>(b0), pget<0>(b1), pget<0>(b2), pget<0>(b3));
276 sme_write_hor_za_vg4<3>(uint32_t(r), pget<1>(b0), pget<1>(b1), pget<1>(b2), pget<1>(b3));
278 for (
int c = 0; c < svl; c += 4) {
279 const auto t0 = sme_read_ver_za_vg4<0, RealScalar>(uint32_t(c)),
280 t2 = sme_read_ver_za_vg4<2, RealScalar>(uint32_t(c));
281 const auto t1 = sme_read_ver_za_vg4<1, RealScalar>(uint32_t(c)),
282 t3 = sme_read_ver_za_vg4<3, RealScalar>(uint32_t(c));
283 RealScalar* d0 = dst + (k + c) * w;
284 RealScalar* d1 = d0 + Index(svl) * w;
285 pstoreu_x2(pn, d0, pcreate(pget<0>(t0), pget<0>(t2)));
286 pstoreu_x2(pn, d0 + w, pcreate(pget<1>(t0), pget<1>(t2)));
287 pstoreu_x2(pn, d0 + 2 * w, pcreate(pget<2>(t0), pget<2>(t2)));
288 pstoreu_x2(pn, d0 + 3 * w, pcreate(pget<3>(t0), pget<3>(t2)));
289 pstoreu_x2(pn, d1, pcreate(pget<0>(t1), pget<0>(t3)));
290 pstoreu_x2(pn, d1 + w, pcreate(pget<1>(t1), pget<1>(t3)));
291 pstoreu_x2(pn, d1 + 2 * w, pcreate(pget<2>(t1), pget<2>(t3)));
292 pstoreu_x2(pn, d1 + 3 * w, pcreate(pget<3>(t1), pget<3>(t3)));
299template <
typename RealScalar>
300static EIGEN_ALWAYS_INLINE
typename sme_packet_traits<RealScalar>::type_x2 sme_row_or_zero(
301 const RealScalar* p,
bool valid, svcount_t pn,
302 typename sme_packet_traits<RealScalar>::type zero) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
303 return valid ? ploadu_x2(pn, p) : pcreate(zero, zero);
308template <
typename RealScalar,
typename Index>
309static EIGEN_ALWAYS_INLINE Index sme_transpose_pack_partial(RealScalar* EIGEN_RESTRICT dst,
310 const RealScalar* EIGEN_RESTRICT src, Index src_stride,
312 int width) __arm_streaming __arm_inout(
"za") {
313 using Traits = sme_packet_traits<RealScalar>;
314 using Vec =
typename Traits::type;
315 const int svl = Traits::size(), w2 = 2 * svl;
316 const svcount_t pn = Traits::ptrue_c();
317 const Vec zero = pset1<Vec>(RealScalar(0));
318 const int rows_lo = sme_min(width, svl), rows_hi = width - rows_lo;
319 const svbool_t pg_lo = Traits::whilelt(0, rows_lo), pg_hi = Traits::whilelt(0, rows_hi);
320 const svbool_t pg_two = Traits::whilelt(0, 2 * width);
322 for (; k + w2 <= k1; k += w2) {
323 for (
int r = 0; r < rows_lo; r += 4) {
324 const RealScalar* q = src + Index(r) * src_stride + k;
325 const auto a0 = sme_row_or_zero<RealScalar>(q, r < rows_lo, pn, zero);
326 const auto a1 = sme_row_or_zero<RealScalar>(q + src_stride, r + 1 < rows_lo, pn, zero);
327 const auto a2 = sme_row_or_zero<RealScalar>(q + 2 * src_stride, r + 2 < rows_lo, pn, zero);
328 const auto a3 = sme_row_or_zero<RealScalar>(q + 3 * src_stride, r + 3 < rows_lo, pn, zero);
329 sme_write_hor_za_vg4<0>(uint32_t(r), pget<0>(a0), pget<0>(a1), pget<0>(a2), pget<0>(a3));
330 sme_write_hor_za_vg4<1>(uint32_t(r), pget<1>(a0), pget<1>(a1), pget<1>(a2), pget<1>(a3));
332 for (
int r = 0; r < rows_hi; r += 4) {
333 const RealScalar* q = src + Index(svl + r) * src_stride + k;
334 const auto a0 = sme_row_or_zero<RealScalar>(q, r < rows_hi, pn, zero);
335 const auto a1 = sme_row_or_zero<RealScalar>(q + src_stride, r + 1 < rows_hi, pn, zero);
336 const auto a2 = sme_row_or_zero<RealScalar>(q + 2 * src_stride, r + 2 < rows_hi, pn, zero);
337 const auto a3 = sme_row_or_zero<RealScalar>(q + 3 * src_stride, r + 3 < rows_hi, pn, zero);
338 sme_write_hor_za_vg4<2>(uint32_t(r), pget<0>(a0), pget<0>(a1), pget<0>(a2), pget<0>(a3));
339 sme_write_hor_za_vg4<3>(uint32_t(r), pget<1>(a0), pget<1>(a1), pget<1>(a2), pget<1>(a3));
341 if (2 * width <= svl) {
343 for (
int c = 0; c < svl; c += 4) {
344 const auto t0 = sme_read_ver_za_vg4<0, RealScalar>(uint32_t(c)),
345 t1 = sme_read_ver_za_vg4<1, RealScalar>(uint32_t(c));
346 RealScalar* d0 = dst + (k + c) * width;
347 RealScalar* d1 = d0 + Index(svl) * width;
348 pstoreu(pg_two, d0, psplice(pg_lo, pget<0>(t0), pget<1>(t0)));
349 pstoreu(pg_two, d0 + 2 * width, psplice(pg_lo, pget<2>(t0), pget<3>(t0)));
350 pstoreu(pg_two, d1, psplice(pg_lo, pget<0>(t1), pget<1>(t1)));
351 pstoreu(pg_two, d1 + 2 * width, psplice(pg_lo, pget<2>(t1), pget<3>(t1)));
355 for (
int c = 0; c < svl; c += 4) {
356 const auto t0 = sme_read_ver_za_vg4<0, RealScalar>(uint32_t(c)),
357 t1 = sme_read_ver_za_vg4<1, RealScalar>(uint32_t(c));
358 RealScalar* d0 = dst + (k + c) * width;
359 RealScalar* d1 = d0 + Index(svl) * width;
360 pstoreu(pg_lo, d0, pget<0>(t0));
361 pstoreu(pg_lo, d0 + width, pget<1>(t0));
362 pstoreu(pg_lo, d0 + 2 * width, pget<2>(t0));
363 pstoreu(pg_lo, d0 + 3 * width, pget<3>(t0));
364 pstoreu(pg_lo, d1, pget<0>(t1));
365 pstoreu(pg_lo, d1 + width, pget<1>(t1));
366 pstoreu(pg_lo, d1 + 2 * width, pget<2>(t1));
367 pstoreu(pg_lo, d1 + 3 * width, pget<3>(t1));
369 const auto t2 = sme_read_ver_za_vg4<2, RealScalar>(uint32_t(c)),
370 t3 = sme_read_ver_za_vg4<3, RealScalar>(uint32_t(c));
371 pstoreu(pg_hi, d0 + svl, pget<0>(t2));
372 pstoreu(pg_hi, d0 + width + svl, pget<1>(t2));
373 pstoreu(pg_hi, d0 + 2 * width + svl, pget<2>(t2));
374 pstoreu(pg_hi, d0 + 3 * width + svl, pget<3>(t2));
375 pstoreu(pg_hi, d1 + svl, pget<0>(t3));
376 pstoreu(pg_hi, d1 + width + svl, pget<1>(t3));
377 pstoreu(pg_hi, d1 + 2 * width + svl, pget<2>(t3));
378 pstoreu(pg_hi, d1 + 3 * width + svl, pget<3>(t3));
405template <
bool NegateOddRows,
typename RealScalar,
typename Index>
406static EIGEN_ALWAYS_INLINE
void sme_transpose_pack_real(RealScalar* EIGEN_RESTRICT dst,
407 const RealScalar* EIGEN_RESTRICT src, Index src_stride,
409 int width) __arm_streaming __arm_inout(
"za") {
410 using Traits = sme_packet_traits<RealScalar>;
411 using Vec =
typename Traits::type;
412 const Vec zero = pset1<Vec>(RealScalar(0));
413 const svbool_t pg_all = Traits::ptrue();
414 const int svl = Traits::size();
417 EIGEN_IF_CONSTEXPR (!NegateOddRows) {
420 if (svl >= 4 && width == 2 * svl && k1 - k0 >= Index(8 * svl))
421 k = sme_transpose_pack_pair(dst, src, src_stride, k0, k1);
422 else if (svl >= 4 && width < 2 * svl && k1 - k0 >= Index(8 * svl))
423 k = sme_transpose_pack_partial(dst, src, src_stride, k0, k1, width);
425 for (; k < k1; k += svl) {
426 const int dk =
static_cast<int>(sme_min(k1 - k, Index(svl)));
427 const svbool_t pg_d = Traits::whilelt(k, k1);
430 for (; r0 + 2 * svl <= width; r0 += 2 * svl) {
431 for (
int r = 0; r < svl; ++r) {
432 sme_ld1_hor_za<0>(uint32_t(r), pg_d, &src[(r0 + r) * src_stride + k]);
433 sme_ld1_hor_za<1>(uint32_t(r), pg_d, &src[(r0 + svl + r) * src_stride + k]);
435 for (
int c = 0; c < dk; ++c) {
436 Vec v0 = sme_read_ver_za<0>(zero, pg_all, uint32_t(c));
437 Vec v1 = sme_read_ver_za<1>(zero, pg_all, uint32_t(c));
438 EIGEN_IF_CONSTEXPR (NegateOddRows) {
439 if (((k + Index(c)) & Index(1)) != Index(0)) {
440 v0 = pnegate(pg_all, v0);
441 v1 = pnegate(pg_all, v1);
444 pstoreu(pg_all, &dst[(k + c) * width + r0], v0);
445 pstoreu(pg_all, &dst[(k + c) * width + r0 + svl], v1);
452 for (; r0 < width; r0 += svl) {
453 const int rg = sme_min(width - r0, svl);
454 const svbool_t pg_r = Traits::whilelt(r0, width);
455 for (
int r = 0; r < rg; ++r) {
456 sme_ld1_hor_za<0>(uint32_t(r), pg_d, &src[(r0 + r) * src_stride + k]);
458 for (
int c = 0; c < dk; ++c) {
459 Vec v0 = sme_read_ver_za<0>(zero, pg_r, uint32_t(c));
460 EIGEN_IF_CONSTEXPR (NegateOddRows) {
461 if (((k + Index(c)) & Index(1)) != Index(0)) v0 = pnegate(pg_r, v0);
463 pstoreu(pg_r, &dst[(k + c) * width + r0], v0);
469template <
bool Conjugate,
typename Scalar,
typename Index>
470static EIGEN_ALWAYS_INLINE
void sme_transpose_pack_range(Scalar* EIGEN_RESTRICT dst,
const Scalar* EIGEN_RESTRICT src,
471 Index src_stride, Index k0, Index k1,
472 int width) __arm_streaming __arm_inout(
"za") {
473 sme_transpose_pack_real<false>(dst, src, src_stride, k0, k1, width);
481template <
bool Conjugate,
typename RealScalar,
typename Index>
482static EIGEN_ALWAYS_INLINE
void sme_transpose_pack_range(std::complex<RealScalar>* EIGEN_RESTRICT dst,
483 const std::complex<RealScalar>* EIGEN_RESTRICT src,
484 Index src_stride, Index k0, Index k1,
485 int width) __arm_streaming __arm_inout(
"za") {
486 sme_transpose_pack_real<Conjugate>(
reinterpret_cast<RealScalar*
>(dst),
reinterpret_cast<const RealScalar*
>(src),
487 Index(2) * src_stride, Index(2) * k0, Index(2) * k1, width);
492template <
bool Conjugate,
typename Scalar,
typename Index>
493static EIGEN_ALWAYS_INLINE
void sme_transpose_pack(Scalar* EIGEN_RESTRICT dst,
const Scalar* EIGEN_RESTRICT src,
494 Index src_stride, Index depth,
495 int width) __arm_streaming __arm_inout(
"za") {
496 sme_transpose_pack_range<Conjugate>(dst, src, src_stride, Index(0), depth, width);
506template <
bool NegateOddRows,
typename RealScalar,
typename Index>
507static void tail_transpose_pack_real(RealScalar* EIGEN_RESTRICT dst_panel,
const RealScalar* EIGEN_RESTRICT src,
508 Index src_stride, Index depth, Index tail) {
509 using Packet =
typename packet_traits<RealScalar>::type;
510 constexpr int PacketSize = int(packet_traits<RealScalar>::size);
511 const Index peeled_tail = (tail / Index(PacketSize)) * Index(PacketSize);
512 const Index peeled_depth = numext::round_down(depth, Index(PacketSize));
515 for (; i < peeled_tail; i += Index(PacketSize)) {
517 for (; k < peeled_depth; k += Index(PacketSize)) {
518 PacketBlock<Packet, PacketSize> block;
519 for (
int p = 0; p < PacketSize; ++p) {
520 block.packet[p] = ploadu<Packet>(src + (i + Index(p)) * src_stride + k);
523 for (
int p = 0; p < PacketSize; ++p) {
524 Packet row = block.packet[p];
525 EIGEN_IF_CONSTEXPR (NegateOddRows) {
526 if (((k + Index(p)) & Index(1)) != Index(0)) row = pnegate(row);
528 pstoreu(dst_panel + (k + Index(p)) * tail + i, row);
531 for (; k < depth; ++k) {
532 const bool negate = NegateOddRows && ((k & Index(1)) != Index(0));
533 for (Index p = 0; p < Index(PacketSize); ++p) {
534 const RealScalar v = src[(i + p) * src_stride + k];
535 dst_panel[k * tail + i + p] = negate ? -v : v;
539 for (; i < tail; ++i) {
540 for (Index k = 0; k < depth; ++k) {
541 const bool negate = NegateOddRows && ((k & Index(1)) != Index(0));
542 const RealScalar v = src[i * src_stride + k];
543 dst_panel[k * tail + i] = negate ? -v : v;
548template <
bool Conjugate,
typename Scalar,
typename Index>
549static void tail_transpose_pack(Scalar* EIGEN_RESTRICT dst_panel,
const Scalar* EIGEN_RESTRICT src, Index src_stride,
550 Index depth, Index tail) {
551 tail_transpose_pack_real<false>(dst_panel, src, src_stride, depth, tail);
554template <
bool Conjugate,
typename RealScalar,
typename Index>
555static void tail_transpose_pack(std::complex<RealScalar>* EIGEN_RESTRICT dst_panel,
556 const std::complex<RealScalar>* EIGEN_RESTRICT src, Index src_stride, Index depth,
558 tail_transpose_pack_real<Conjugate>(
reinterpret_cast<RealScalar*
>(dst_panel),
559 reinterpret_cast<const RealScalar*
>(src), Index(2) * src_stride, Index(2) * depth,
565template <
bool Conjugate,
typename Scalar,
typename Index>
566__arm_locally_streaming __arm_new(
"za") static
void sme_tail_pack_streaming(Scalar* dst_panel, const Scalar* src,
567 Index src_stride, Index depth,
int tail) {
568 sme_transpose_pack<Conjugate>(dst_panel, src, src_stride, depth, tail);
570template <
bool Conjugate,
typename Scalar,
typename Index>
571static EIGEN_ALWAYS_INLINE
void tail_pack(Scalar* dst_panel,
const Scalar* src, Index src_stride, Index depth,
573 EIGEN_IF_CONSTEXPR (!NumTraits<Scalar>::IsComplex) {
574 if (tail > 4 && depth >= Index(4 * sme_block<Scalar>::nr)) {
575 sme_tail_pack_streaming<Conjugate>(dst_panel, src, src_stride, depth,
static_cast<int>(tail));
579 tail_transpose_pack<Conjugate>(dst_panel, src, src_stride, depth, tail);
583static EIGEN_ALWAYS_INLINE
void sme_neon_ld2(
const float* p, Packet4f& re, Packet4f& im) {
584 const float32x4x2_t v = vld2q_f32(p);
588static EIGEN_ALWAYS_INLINE
void sme_neon_st2(
float* p, Packet4f re, Packet4f im) {
594static EIGEN_ALWAYS_INLINE
void sme_neon_ld2(
const double* p, Packet2d& re, Packet2d& im) {
595 const float64x2x2_t v = vld2q_f64(p);
599static EIGEN_ALWAYS_INLINE
void sme_neon_st2(
double* p, Packet2d re, Packet2d im) {
609template <
bool Conjugate,
typename Scalar,
typename Index>
610static void neon_copy_panel(Scalar* EIGEN_RESTRICT dst,
const Scalar* EIGEN_RESTRICT src, Index src_stride, Index depth,
612 using Packet =
typename packet_traits<Scalar>::type;
613 constexpr Index PS = Index(packet_traits<Scalar>::size);
614 const Index peeled4 = numext::round_down(w, 4 * PS);
615 const Index peeled = numext::round_down(w, PS);
616 for (Index k = 0; k < depth; ++k) {
617 const Scalar* s = src + k * src_stride;
618 Scalar* d = dst + k * w;
620 for (; r < peeled4; r += 4 * PS) {
621 const Packet p0 = ploadu<Packet>(s + r), p1 = ploadu<Packet>(s + r + PS);
622 const Packet p2 = ploadu<Packet>(s + r + 2 * PS), p3 = ploadu<Packet>(s + r + 3 * PS);
624 pstoreu(d + r + PS, p1);
625 pstoreu(d + r + 2 * PS, p2);
626 pstoreu(d + r + 3 * PS, p3);
628 for (; r < peeled; r += PS) pstoreu(d + r, ploadu<Packet>(s + r));
629 for (; r < w; ++r) d[r] = s[r];
632template <
bool Conjugate,
typename RealScalar,
typename Index>
633static void neon_copy_panel(std::complex<RealScalar>* EIGEN_RESTRICT dst,
634 const std::complex<RealScalar>* EIGEN_RESTRICT src, Index src_stride, Index depth,
636 using Packet =
typename packet_traits<RealScalar>::type;
637 constexpr Index PS = Index(packet_traits<RealScalar>::size);
638 RealScalar* rd =
reinterpret_cast<RealScalar*
>(dst);
639 const RealScalar* rs =
reinterpret_cast<const RealScalar*
>(src);
640 const Index peeled = numext::round_down(w, PS);
641 for (Index k = 0; k < depth; ++k) {
642 const RealScalar* s = rs + k * 2 * src_stride;
643 RealScalar* d = rd + k * 2 * w;
645 for (; r < peeled; r += PS) {
647 sme_neon_ld2(s + 2 * r, re, im);
649 pstoreu(d + w + r, Conjugate ? pnegate(im) : im);
653 d[w + r] = Conjugate ? -s[2 * r + 1] : s[2 * r + 1];
673template <
typename DataMapper,
typename EnableIf =
void>
674struct sme_has_incr : std::false_type {};
675template <
typename DataMapper>
676struct sme_has_incr<DataMapper, void_t<decltype(std::declval<const DataMapper&>().incr())>> : std::true_type {};
678template <typename Index, typename DataMapper, std::enable_if_t<sme_has_incr<DataMapper>::value,
bool> =
true>
679EIGEN_ALWAYS_INLINE Index sme_mapper_incr(
const DataMapper& m) {
680 return static_cast<Index
>(m.incr());
682template <typename Index, typename DataMapper, std::enable_if_t<!sme_has_incr<DataMapper>::value,
bool> =
true>
683EIGEN_ALWAYS_INLINE Index sme_mapper_incr(
const DataMapper&) {
689template <
typename DataMapper,
typename Index>
690struct sme_mapper_has_direct_access {
691 static constexpr bool value = std::is_lvalue_reference<decltype(std::declval<const DataMapper&>()(
692 std::declval<Index>(), std::declval<Index>()))>::value;
699template <
bool Conjugate,
typename Scalar,
typename Index>
700EIGEN_ALWAYS_INLINE
void sme_pack_store(Scalar* dst_step, Index w, Index r,
const Scalar& v) {
701 EIGEN_UNUSED_VARIABLE(w);
702 dst_step[r] = conj_if<Conjugate>()(v);
704template <
bool Conjugate,
typename RealScalar,
typename Index>
705EIGEN_ALWAYS_INLINE
void sme_pack_store(std::complex<RealScalar>* dst_step, Index w, Index r,
706 const std::complex<RealScalar>& v) {
707 const std::complex<RealScalar> cv = conj_if<Conjugate>()(v);
708 RealScalar* p =
reinterpret_cast<RealScalar*
>(dst_step);
709 p[r] = numext::real(cv);
710 p[w + r] = numext::imag(cv);
729template <
typename Scalar,
int MR,
typename Index,
typename DataMapper,
bool Conjugate,
bool PanelMode>
730void sme_pack_lhs_fallback(Scalar* dst_base,
const DataMapper& lhs, Index depth, Index rows, Index dst_stride,
731 Index dst_offset,
bool vectorise) {
732 using Packet =
typename packet_traits<Scalar>::type;
733 constexpr Index PacketSize = Index(packet_traits<Scalar>::size);
734 constexpr bool HasPacketPath = !NumTraits<Scalar>::IsComplex;
736 for (Index i = 0; i < rows; i += MR) {
737 const Index w = numext::mini(rows - i, Index(MR));
738 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * w : dst_base + i * depth;
739 const Index peeled_w = (vectorise && HasPacketPath) ? numext::round_down(w, Index(PacketSize)) : Index(0);
740 for (Index k = 0; k < depth; ++k) {
741 Scalar* dst_step = dst_panel + k * w;
743 for (; r < peeled_w; r += PacketSize) {
744 pstoreu(dst_step + r, lhs.template loadPacket<Packet>(i + r, k));
747 sme_pack_store<Conjugate>(dst_step, w, r, lhs(i + r, k));
754template <
typename DataMapper,
typename Index, std::size_t... Is>
755EIGEN_ALWAYS_INLINE std::array<
typename DataMapper::LinearMapper,
sizeof...(Is)> sme_column_mappers(
756 const DataMapper& rhs, Index col, std::index_sequence<Is...>) {
757 return {{rhs.getLinearMapper(0, col + Index(Is))...}};
766template <
typename Scalar,
int NR,
typename Index,
typename DataMapper,
bool Conjugate,
bool PanelMode>
767void sme_pack_rhs_fallback(Scalar* dst_base,
const DataMapper& rhs, Index depth, Index cols, Index dst_stride,
768 Index dst_offset,
bool vectorise) {
769 using Packet =
typename packet_traits<Scalar>::type;
770 using LinearMapper =
typename DataMapper::LinearMapper;
771 constexpr int PacketSize = int(packet_traits<Scalar>::size);
772 constexpr bool HasPacketPath = !NumTraits<Scalar>::IsComplex;
773 const Index peeled_depth = (depth / Index(PacketSize)) * Index(PacketSize);
775 for (Index j = 0; j < cols; j += NR) {
776 const Index w = numext::mini(cols - j, Index(NR));
777 Scalar* dst_panel = PanelMode ? dst_base + j * dst_stride + dst_offset * w : dst_base + j * depth;
778 const Index peeled_w = (vectorise && HasPacketPath) ? numext::round_down(w, Index(PacketSize)) : Index(0);
780 for (; c < peeled_w; c += Index(PacketSize)) {
784 const std::array<LinearMapper, PacketSize> dm =
785 sme_column_mappers(rhs, j + c, std::make_index_sequence<PacketSize>{});
787 for (; k < peeled_depth; k += Index(PacketSize)) {
788 PacketBlock<Packet, PacketSize> block;
789 for (
int p = 0; p < PacketSize; ++p) {
790 block.packet[p] = dm[p].template loadPacket<Packet>(k);
793 for (
int p = 0; p < PacketSize; ++p) {
794 pstoreu(dst_panel + (k + Index(p)) * w + c, block.packet[p]);
797 for (; k < depth; ++k) {
798 for (Index p = 0; p < Index(PacketSize); ++p) {
799 sme_pack_store<Conjugate>(dst_panel + k * w, w, c + p, rhs(k, j + c + p));
804 for (Index k = 0; k < depth; ++k) {
805 sme_pack_store<Conjugate>(dst_panel + k * w, w, c, rhs(k, j + c));
816#ifndef EIGEN_SME_NEON_MAX_DEPTH
817template <
typename Scalar>
818struct sme_neon_max_depth : std::integral_constant<int, 16> {};
820struct sme_neon_max_depth<float> : std::integral_constant<int, 24> {};
822struct sme_neon_max_depth<std::complex<double>> : std::integral_constant<int, 8> {};
824template <
typename Scalar>
825struct sme_neon_max_depth : std::integral_constant<int, EIGEN_SME_NEON_MAX_DEPTH> {};
827#ifndef EIGEN_SME_NEON_MAX_WIDTH
828template <
typename Scalar>
829struct sme_neon_max_width : std::integral_constant<int, 48> {};
831struct sme_neon_max_width<double> : std::integral_constant<int, 40> {};
833struct sme_neon_max_width<std::complex<float>> : std::integral_constant<int, 16> {};
835struct sme_neon_max_width<std::complex<double>> : std::integral_constant<int, 24> {};
837template <
typename Scalar>
838struct sme_neon_max_width : std::integral_constant<int, EIGEN_SME_NEON_MAX_WIDTH> {};
843#ifndef EIGEN_SME_NEON_SHALLOW_DEPTH
844template <
typename Scalar>
845struct sme_neon_shallow_depth : std::integral_constant<int, 8> {};
847template <
typename Scalar>
848struct sme_neon_shallow_depth : std::integral_constant<int, EIGEN_SME_NEON_SHALLOW_DEPTH> {};
850#ifndef EIGEN_SME_NEON_SHALLOW_WIDTH
851template <
typename Scalar>
852struct sme_neon_shallow_width : std::integral_constant<int, 96> {};
854struct sme_neon_shallow_width<double> : std::integral_constant<int, 64> {};
856struct sme_neon_shallow_width<std::complex<float>> : std::integral_constant<int, 32> {};
858struct sme_neon_shallow_width<std::complex<double>> : std::integral_constant<int, 24> {};
860template <
typename Scalar>
861struct sme_neon_shallow_width : std::integral_constant<int, EIGEN_SME_NEON_SHALLOW_WIDTH> {};
867#ifndef EIGEN_SME_NEON_MAX_PANEL
868template <
typename Scalar>
869struct sme_neon_max_panel : std::integral_constant<int, 256> {};
870template <
typename RealScalar>
871struct sme_neon_max_panel<std::complex<RealScalar>> : std::integral_constant<int, 128> {};
873template <
typename Scalar>
874struct sme_neon_max_panel : std::integral_constant<int, EIGEN_SME_NEON_MAX_PANEL> {};
878#ifndef EIGEN_SME_NEON_THIN_DIM
879template <
typename Scalar>
880struct sme_neon_thin_dim : std::integral_constant<int, 32> {};
882struct sme_neon_thin_dim<std::complex<float>> : std::integral_constant<int, 8> {};
884struct sme_neon_thin_dim<std::complex<double>> : std::integral_constant<int, 4> {};
886template <
typename Scalar>
887struct sme_neon_thin_dim : std::integral_constant<int, EIGEN_SME_NEON_THIN_DIM> {};
891template <
typename Scalar,
typename Index>
892EIGEN_ALWAYS_INLINE
bool sme_pack_with_neon(Index depth, Index width) {
893#if defined(EIGEN_SME_NO_NEON_SMALL_BLOCKS)
894 EIGEN_UNUSED_VARIABLE(depth);
895 EIGEN_UNUSED_VARIABLE(width);
897#elif defined(EIGEN_SME_FORCE_NEON_SMALL_BLOCKS)
898 EIGEN_UNUSED_VARIABLE(depth);
899 EIGEN_UNUSED_VARIABLE(width);
902 return depth <= Index(sme_neon_max_depth<Scalar>::value) && width <= Index(sme_neon_max_panel<Scalar>::value);
908template <
typename Scalar,
typename Index>
909EIGEN_ALWAYS_INLINE
bool sme_kernel_with_neon(Index rows, Index cols, Index depth, Index strideA, Index strideB) {
910#if defined(EIGEN_SME_NO_NEON_SMALL_BLOCKS) || defined(EIGEN_SME_FORCE_NEON_SMALL_BLOCKS)
911 EIGEN_UNUSED_VARIABLE(depth);
912 return sme_pack_with_neon<Scalar>(strideA, rows) && sme_pack_with_neon<Scalar>(strideB, cols);
914 const Index wide = numext::maxi(rows, cols);
915 const Index w = depth <= Index(sme_neon_shallow_depth<Scalar>::value) ? Index(sme_neon_shallow_width<Scalar>::value)
916 : Index(sme_neon_max_width<Scalar>::value);
917 return sme_pack_with_neon<Scalar>(strideA, rows) && sme_pack_with_neon<Scalar>(strideB, cols) &&
918 (wide <= w || numext::mini(rows, cols) <= Index(sme_neon_thin_dim<Scalar>::value));
930template <
bool UsePacketPath,
bool PanelMode,
typename Scalar,
typename Index,
typename DataMapper,
typename DirectFn,
931 typename NeonFn,
typename FallbackFn>
932EIGEN_ALWAYS_INLINE
void sme_dispatch_pack(DirectFn direct, NeonFn neon, FallbackFn fallback, Scalar* block,
933 const DataMapper& m, Index depth, Index n, Index stride, Index offset,
935 if (sme_mapper_incr<Index>(m) == 1) {
936 const Scalar* src = (n > 0 && depth > 0) ? &m(0, 0) :
nullptr;
937 if (sme_pack_with_neon<Scalar>(PanelMode ? stride : depth, n)) {
938 neon(block, src, m.stride(), depth, n, stride, offset);
941 direct(block, src, m.stride(), depth, n, stride, offset);
944 fallback(block, m, depth, n, stride, offset, UsePacketPath);
947template <
bool UsePacketPath,
bool PanelMode,
typename Scalar,
typename Index,
typename DataMapper,
typename DirectFn,
948 typename NeonFn,
typename FallbackFn>
949EIGEN_ALWAYS_INLINE
void sme_dispatch_pack(DirectFn, NeonFn, FallbackFn fallback, Scalar* block,
const DataMapper& m,
950 Index depth, Index n, Index stride, Index offset,
952 fallback(block, m, depth, n, stride, offset, UsePacketPath);
985class gebp_traits<float, float, false, false, Architecture::Target, GEBPPacketFull>
986 :
public gebp_traits<float, float, false, false, Architecture::Target, GEBPPacketHalf> {
999static_assert(int(gebp_traits<
float, std::complex<float>>::mr) != kSmeMr,
1000 "gebp_traits<float, complex<float>>::mr collides with kSmeMr: the SME gemm_pack_lhs would silently "
1001 "emit SME panel layout for the generic gebp_kernel");
1002static_assert(int(gebp_traits<std::complex<float>,
float>::nr) != kSmeNr,
1003 "gebp_traits<complex<float>, float>::nr collides with kSmeNr: the SME gemm_pack_rhs would silently "
1004 "emit SME panel layout for the generic gebp_kernel");
1006#ifdef EIGEN_VECTORIZE_SME_F64F64
1008class gebp_traits<double, double, false, false, Architecture::Target, GEBPPacketFull>
1009 :
public gebp_traits<double, double, false, false, Architecture::Target, GEBPPacketHalf> {
1012 static constexpr int mr = kSmeMrD;
1013 static constexpr int nr = kSmeNrD;
1016static_assert(int(gebp_traits<
double, std::complex<double>>::mr) != kSmeMrD,
1017 "gebp_traits<double, complex<double>>::mr collides with kSmeMrD: the SME gemm_pack_lhs would silently "
1018 "emit SME panel layout for the generic gebp_kernel");
1019static_assert(int(gebp_traits<std::complex<double>,
double>::nr) != kSmeNrD,
1020 "gebp_traits<complex<double>, double>::nr collides with kSmeNrD: the SME gemm_pack_rhs would silently "
1021 "emit SME panel layout for the generic gebp_kernel");
1029template <
bool ConjLhs_,
bool ConjRhs_>
1030class gebp_traits<std::complex<float>, std::complex<float>, ConjLhs_, ConjRhs_, Architecture::Target, GEBPPacketFull>
1031 :
public gebp_traits<std::complex<float>, std::complex<float>, ConjLhs_, ConjRhs_, Architecture::Target,
1034 static constexpr int mr = kSmeMrC;
1035 static constexpr int nr = kSmeNrC;
1041static_assert(int(gebp_traits<std::complex<float>,
float>::mr) != kSmeMrC,
1042 "gebp_traits<complex<float>, float>::mr collides with kSmeMrC: the SME gemm_pack_lhs would silently "
1043 "emit SME panel layout for the generic gebp_kernel");
1044static_assert(int(gebp_traits<
float, std::complex<float>>::nr) != kSmeNrC,
1045 "gebp_traits<float, complex<float>>::nr collides with kSmeNrC: the SME gemm_pack_rhs would silently "
1046 "emit SME panel layout for the generic gebp_kernel");
1048#ifdef EIGEN_VECTORIZE_SME_F64F64
1049template <
bool ConjLhs_,
bool ConjRhs_>
1050class gebp_traits<std::complex<double>, std::complex<double>, ConjLhs_, ConjRhs_, Architecture::Target, GEBPPacketFull>
1051 :
public gebp_traits<std::complex<double>, std::complex<double>, ConjLhs_, ConjRhs_, Architecture::Target,
1054 static constexpr int mr = kSmeMrCD;
1055 static constexpr int nr = kSmeNrCD;
1058static_assert(int(gebp_traits<std::complex<double>,
double>::mr) != kSmeMrCD,
1059 "gebp_traits<complex<double>, double>::mr collides with kSmeMrCD: the SME gemm_pack_lhs would silently "
1060 "emit SME panel layout for the generic gebp_kernel");
1061static_assert(int(gebp_traits<
double, std::complex<double>>::nr) != kSmeNrCD,
1062 "gebp_traits<double, complex<double>>::nr collides with kSmeNrCD: the SME gemm_pack_rhs would silently "
1063 "emit SME panel layout for the generic gebp_kernel");
1073template <
typename Scalar,
int MR,
typename Index,
typename DataMapper,
bool Conjugate,
bool PanelMode>
1074struct sme_pack_lhs_colmajor {
1077 static EIGEN_ALWAYS_INLINE
void pack_neon(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src, Index src_stride,
1078 Index depth, Index rows, Index dst_stride, Index dst_offset) {
1079 for (Index i = 0; i < rows; i += MR) {
1080 const Index w = numext::mini(Index(MR), rows - i);
1081 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * w : dst_base + i * depth;
1082 neon_copy_panel<Conjugate>(dst_panel, src + i, src_stride, depth, w);
1089 __arm_locally_streaming EIGEN_DONT_INLINE
static void pack_direct(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src,
1090 Index src_stride, Index depth, Index rows,
1091 Index dst_stride, Index dst_offset) {
1092 const Index peeled_rows = (rows / MR) * MR;
1093 const Index i0 = sme_pack_quad_panels<PanelMode>(dst_base, src, src_stride, depth, peeled_rows, Index(MR),
1094 dst_stride, dst_offset);
1097 for (Index i = i0; i < peeled_rows; i += MR) {
1098 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * MR : dst_base + i * depth;
1099 sve_copy_panel<Conjugate>(dst_panel, src + i, src_stride, depth, MR);
1103 if (peeled_rows < rows) {
1104 const Index tail = rows - peeled_rows;
1106 PanelMode ? dst_base + peeled_rows * dst_stride + dst_offset * tail : dst_base + peeled_rows * depth;
1107 sve_copy_panel<Conjugate>(dst_panel, src + peeled_rows, src_stride, depth,
static_cast<int>(tail));
1111 EIGEN_DONT_INLINE
void operator()(Scalar* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
1114 eigen_assert(stride >= depth && offset <= stride);
1118 sme_dispatch_pack<true, PanelMode>(
1119 &pack_direct, &pack_neon, &sme_pack_lhs_fallback<Scalar, MR, Index, DataMapper, Conjugate, PanelMode>, blockA,
1120 lhs, depth, rows, stride, offset, bool_constant<sme_mapper_has_direct_access<DataMapper, Index>::value>{});
1132template <
typename Scalar,
int MR,
typename Index,
typename DataMapper,
bool Conjugate,
bool PanelMode>
1133struct sme_pack_lhs_rowmajor {
1136 static EIGEN_ALWAYS_INLINE
void pack_neon(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src, Index src_stride,
1137 Index depth, Index rows, Index dst_stride, Index dst_offset) {
1138 for (Index i = 0; i < rows; i += MR) {
1139 const Index w = numext::mini(Index(MR), rows - i);
1140 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * w : dst_base + i * depth;
1141 tail_transpose_pack<Conjugate>(dst_panel, src + i * src_stride, src_stride, depth, w);
1145 __arm_locally_streaming __arm_new(
"za") static
void pack_full_panels(Scalar* dst_base,
1146 const Scalar* EIGEN_RESTRICT src,
1147 Index src_stride, Index depth, Index peeled_rows,
1148 Index dst_stride, Index dst_offset) {
1149 for (Index i = 0; i < peeled_rows; i += MR) {
1150 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * MR : dst_base + i * depth;
1151 sme_transpose_pack<Conjugate>(dst_panel, src + i * src_stride, src_stride, depth, MR);
1155 static void pack_direct(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src, Index src_stride, Index depth, Index rows,
1156 Index dst_stride, Index dst_offset) {
1157 const Index peeled_rows = numext::round_down(rows, MR);
1159 if (peeled_rows > 0) {
1160 pack_full_panels(dst_base, src, src_stride, depth, peeled_rows, dst_stride, dst_offset);
1164 if (peeled_rows < rows) {
1165 const Index tail = rows - peeled_rows;
1167 PanelMode ? dst_base + peeled_rows * dst_stride + dst_offset * tail : dst_base + peeled_rows * depth;
1168 tail_pack<Conjugate>(dst_panel, src + peeled_rows * src_stride, src_stride, depth, tail);
1172 EIGEN_DONT_INLINE
void operator()(Scalar* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
1175 eigen_assert(stride >= depth && offset <= stride);
1180 sme_dispatch_pack<false, PanelMode>(
1181 &pack_direct, &pack_neon, &sme_pack_lhs_fallback<Scalar, MR, Index, DataMapper, Conjugate, PanelMode>, blockA,
1182 lhs, depth, rows, stride, offset, bool_constant<sme_mapper_has_direct_access<DataMapper, Index>::value>{});
1195template <
typename Scalar,
int NR,
typename Index,
typename DataMapper,
bool Conjugate,
bool PanelMode>
1196struct sme_pack_rhs_colmajor {
1199 static EIGEN_ALWAYS_INLINE
void pack_neon(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src, Index src_stride,
1200 Index depth, Index cols, Index dst_stride, Index dst_offset) {
1201 for (Index i = 0; i < cols; i += NR) {
1202 const Index w = numext::mini(Index(NR), cols - i);
1203 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * w : dst_base + i * depth;
1204 tail_transpose_pack<Conjugate>(dst_panel, src + i * src_stride, src_stride, depth, w);
1208 __arm_locally_streaming __arm_new(
"za") static
void pack_full_panels(Scalar* dst_base,
1209 const Scalar* EIGEN_RESTRICT src,
1210 Index src_stride, Index depth, Index peeled_cols,
1211 Index dst_stride, Index dst_offset) {
1212 for (Index j = 0; j < peeled_cols; j += NR) {
1213 Scalar* dst_panel = PanelMode ? dst_base + j * dst_stride + dst_offset * NR : dst_base + j * depth;
1214 sme_transpose_pack<Conjugate>(dst_panel, src + j * src_stride, src_stride, depth, NR);
1218 static void pack_direct(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src, Index src_stride, Index depth, Index cols,
1219 Index dst_stride, Index dst_offset) {
1220 const Index peeled_cols = numext::round_down(cols, NR);
1222 if (peeled_cols > 0) {
1223 pack_full_panels(dst_base, src, src_stride, depth, peeled_cols, dst_stride, dst_offset);
1227 if (peeled_cols < cols) {
1228 const Index tail = cols - peeled_cols;
1230 PanelMode ? dst_base + peeled_cols * dst_stride + dst_offset * tail : dst_base + peeled_cols * depth;
1231 tail_pack<Conjugate>(dst_panel, src + peeled_cols * src_stride, src_stride, depth, tail);
1235 EIGEN_DONT_INLINE
void operator()(Scalar* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
1238 eigen_assert(stride >= depth && offset <= stride);
1242 sme_dispatch_pack<true, PanelMode>(
1243 &pack_direct, &pack_neon, &sme_pack_rhs_fallback<Scalar, NR, Index, DataMapper, Conjugate, PanelMode>, blockB,
1244 rhs, depth, cols, stride, offset, bool_constant<sme_mapper_has_direct_access<DataMapper, Index>::value>{});
1250template <
typename Scalar,
int NR,
typename Index,
typename DataMapper,
bool Conjugate,
bool PanelMode>
1251struct sme_pack_rhs_rowmajor {
1254 static EIGEN_ALWAYS_INLINE
void pack_neon(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src, Index src_stride,
1255 Index depth, Index cols, Index dst_stride, Index dst_offset) {
1256 for (Index i = 0; i < cols; i += NR) {
1257 const Index w = numext::mini(Index(NR), cols - i);
1258 Scalar* dst_panel = PanelMode ? dst_base + i * dst_stride + dst_offset * w : dst_base + i * depth;
1259 neon_copy_panel<Conjugate>(dst_panel, src + i, src_stride, depth, w);
1264 __arm_locally_streaming EIGEN_DONT_INLINE
static void pack_direct(Scalar* dst_base,
const Scalar* EIGEN_RESTRICT src,
1265 Index src_stride, Index depth, Index cols,
1266 Index dst_stride, Index dst_offset) {
1267 const Index peeled_cols = (cols / NR) * NR;
1269 for (Index j = 0; j < peeled_cols; j += NR) {
1270 Scalar* dst_panel = PanelMode ? dst_base + j * dst_stride + dst_offset * NR : dst_base + j * depth;
1271 sve_copy_panel<Conjugate>(dst_panel, src + j, src_stride, depth, NR);
1274 if (peeled_cols < cols) {
1275 const Index tail = cols - peeled_cols;
1277 PanelMode ? dst_base + peeled_cols * dst_stride + dst_offset * tail : dst_base + peeled_cols * depth;
1278 sve_copy_panel<Conjugate>(dst_panel, src + peeled_cols, src_stride, depth,
static_cast<int>(tail));
1282 EIGEN_DONT_INLINE
void operator()(Scalar* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
1285 eigen_assert(stride >= depth && offset <= stride);
1290 sme_dispatch_pack<false, PanelMode>(
1291 &pack_direct, &pack_neon, &sme_pack_rhs_fallback<Scalar, NR, Index, DataMapper, Conjugate, PanelMode>, blockB,
1292 rhs, depth, cols, stride, offset, bool_constant<sme_mapper_has_direct_access<DataMapper, Index>::value>{});
1299#define EIGEN_SME_DECLARE_GEMM_PACKERS(SCALAR, MR, NR) \
1300 template <typename Index, typename DataMapper, int Pack2, typename Packet, bool Conjugate, bool PanelMode> \
1301 struct gemm_pack_lhs<SCALAR, Index, DataMapper, MR, Pack2, Packet, ColMajor, Conjugate, PanelMode> \
1302 : sme_pack_lhs_colmajor<SCALAR, MR, Index, DataMapper, Conjugate, PanelMode> {}; \
1304 template <typename Index, typename DataMapper, int Pack2, typename Packet, bool Conjugate, bool PanelMode> \
1305 struct gemm_pack_lhs<SCALAR, Index, DataMapper, MR, Pack2, Packet, RowMajor, Conjugate, PanelMode> \
1306 : sme_pack_lhs_rowmajor<SCALAR, MR, Index, DataMapper, Conjugate, PanelMode> {}; \
1308 template <typename Index, typename DataMapper, bool Conjugate, bool PanelMode> \
1309 struct gemm_pack_rhs<SCALAR, Index, DataMapper, NR, ColMajor, Conjugate, PanelMode> \
1310 : sme_pack_rhs_colmajor<SCALAR, NR, Index, DataMapper, Conjugate, PanelMode> {}; \
1312 template <typename Index, typename DataMapper, bool Conjugate, bool PanelMode> \
1313 struct gemm_pack_rhs<SCALAR, Index, DataMapper, NR, RowMajor, Conjugate, PanelMode> \
1314 : sme_pack_rhs_rowmajor<SCALAR, NR, Index, DataMapper, Conjugate, PanelMode> {};
1316EIGEN_SME_DECLARE_GEMM_PACKERS(
float, kSmeMr, kSmeNr)
1317EIGEN_SME_DECLARE_GEMM_PACKERS(std::complex<float>, kSmeMrC, kSmeNrC)
1318#ifdef EIGEN_VECTORIZE_SME_F64F64
1319EIGEN_SME_DECLARE_GEMM_PACKERS(
double, kSmeMrD, kSmeNrD)
1320EIGEN_SME_DECLARE_GEMM_PACKERS(std::complex<double>, kSmeMrCD, kSmeNrCD)
1323#undef EIGEN_SME_DECLARE_GEMM_PACKERS
1332template <
typename Scalar,
int TileId,
typename Index>
1333EIGEN_ALWAYS_INLINE
void sme_store_za_tile(Scalar* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
1334 Scalar alpha, Index row_start,
int pw, Index col_start,
1335 int cw) __arm_streaming __arm_inout(
"za") {
1336 using Traits = sme_packet_traits<Scalar>;
1337 using Vec =
typename Traits::type;
1338 const svbool_t pg_m = Traits::whilelt(0, pw);
1339 const svbool_t pg_n = Traits::whilelt(0, cw);
1346 const Vec vzero = pset1<Vec>(Scalar(0));
1347 const Vec valpha = pset1<Vec>(alpha);
1354 if (C_stride_row == 1) {
1357 for (; ci + 2 <= cw; ci += 2) {
1358 Scalar* p0 = C + row_start + (col_start + ci) * C_stride_col;
1359 Scalar* p1 = p0 + C_stride_col;
1360 Vec c0 = ploadu(pg_m, p0);
1361 Vec c1 = ploadu(pg_m, p1);
1362 pstoreu(pg_m, p0, pmadd(pg_m, sme_read_ver_za<TileId>(vzero, pg_m, (uint32_t)ci), valpha, c0));
1363 pstoreu(pg_m, p1, pmadd(pg_m, sme_read_ver_za<TileId>(vzero, pg_m, (uint32_t)(ci + 1)), valpha, c1));
1366 Scalar* pC = C + row_start + (col_start + ci) * C_stride_col;
1367 Vec vc = ploadu(pg_m, pC);
1368 pstoreu(pg_m, pC, pmadd(pg_m, sme_read_ver_za<TileId>(vzero, pg_m, (uint32_t)ci), valpha, vc));
1370 }
else if (C_stride_col == 1) {
1373 for (; ri + 2 <= pw; ri += 2) {
1374 Scalar* p0 = C + (row_start + ri) * C_stride_row + col_start;
1375 Scalar* p1 = p0 + C_stride_row;
1376 Vec c0 = ploadu(pg_n, p0);
1377 Vec c1 = ploadu(pg_n, p1);
1378 pstoreu(pg_n, p0, pmadd(pg_n, sme_read_hor_za<TileId>(vzero, pg_n, (uint32_t)ri), valpha, c0));
1379 pstoreu(pg_n, p1, pmadd(pg_n, sme_read_hor_za<TileId>(vzero, pg_n, (uint32_t)(ri + 1)), valpha, c1));
1382 Scalar* pC = C + (row_start + ri) * C_stride_row + col_start;
1383 Vec vc = ploadu(pg_n, pC);
1384 pstoreu(pg_n, pC, pmadd(pg_n, sme_read_hor_za<TileId>(vzero, pg_n, (uint32_t)ri), valpha, vc));
1391 Scalar scratch[sme_block<Scalar>::nr];
1392 for (
int ri = 0; ri < pw; ++ri) {
1393 Vec vres = sme_read_hor_za<TileId>(vzero, pg_n, (uint32_t)ri);
1394 vres = pmul(pg_n, vres, valpha);
1395 pstoreu(pg_n, scratch, vres);
1396 for (
int ci = 0; ci < cw; ++ci) {
1397 C[(row_start + ri) * C_stride_row + (col_start + ci) * C_stride_col] += scratch[ci];
1417template <
int TileLo,
int TileHi,
typename Scalar,
typename Index>
1418EIGEN_ALWAYS_INLINE
void sme_store_tile_pair_colmajor(Scalar* EIGEN_RESTRICT C, Index ldc, Scalar alpha,
1419 int svl) __arm_streaming __arm_inout(
"za") {
1420 using Traits = sme_packet_traits<Scalar>;
1421 const svcount_t pn = Traits::ptrue_c();
1422 const svbool_t pg = Traits::ptrue();
1423 const typename Traits::type valpha = pset1<typename Traits::type>(alpha);
1424 for (
int c = 0; c < svl; c += 4) {
1425 const auto lo = sme_read_ver_za_vg4<TileLo, Scalar>(uint32_t(c));
1426 const auto hi = sme_read_ver_za_vg4<TileHi, Scalar>(uint32_t(c));
1427 Scalar* p = C + Index(c) * ldc;
1428 const auto c0 = ploadu_x2(pn, p), c1 = ploadu_x2(pn, p + ldc), c2 = ploadu_x2(pn, p + 2 * ldc),
1429 c3 = ploadu_x2(pn, p + 3 * ldc);
1431 pcreate(pmadd(pg, pget<0>(lo), valpha, pget<0>(c0)), pmadd(pg, pget<0>(hi), valpha, pget<1>(c0))));
1432 pstoreu_x2(pn, p + ldc,
1433 pcreate(pmadd(pg, pget<1>(lo), valpha, pget<0>(c1)), pmadd(pg, pget<1>(hi), valpha, pget<1>(c1))));
1434 pstoreu_x2(pn, p + 2 * ldc,
1435 pcreate(pmadd(pg, pget<2>(lo), valpha, pget<0>(c2)), pmadd(pg, pget<2>(hi), valpha, pget<1>(c2))));
1436 pstoreu_x2(pn, p + 3 * ldc,
1437 pcreate(pmadd(pg, pget<3>(lo), valpha, pget<0>(c3)), pmadd(pg, pget<3>(hi), valpha, pget<1>(c3))));
1441template <
typename Scalar,
typename Index>
1442EIGEN_ALWAYS_INLINE
void sme_store_2x2_grid(Scalar* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
1443 Scalar alpha, Index row_start,
int rlo,
int rhi, Index col_start,
int clo,
1444 int chi) __arm_streaming __arm_inout(
"za") {
1445 const int svl = sme_packet_traits<Scalar>::size();
1446 if (svl >= 4 && C_stride_row == 1 && rlo == svl && rhi == svl && clo == svl && chi == svl) {
1447 Scalar* c0 = C + row_start + col_start * C_stride_col;
1448 sme_store_tile_pair_colmajor<0, 2>(c0, C_stride_col, alpha, svl);
1449 sme_store_tile_pair_colmajor<1, 3>(c0 + Index(svl) * C_stride_col, C_stride_col, alpha, svl);
1452 sme_store_za_tile<Scalar, 0>(C, C_stride_row, C_stride_col, alpha, row_start, rlo, col_start, clo);
1454 sme_store_za_tile<Scalar, 1>(C, C_stride_row, C_stride_col, alpha, row_start, rlo, col_start + svl, chi);
1457 sme_store_za_tile<Scalar, 2>(C, C_stride_row, C_stride_col, alpha, row_start + svl, rhi, col_start, clo);
1459 sme_store_za_tile<Scalar, 3>(C, C_stride_row, C_stride_col, alpha, row_start + svl, rhi, col_start + svl, chi);
1469template <
typename Scalar>
1470static EIGEN_ALWAYS_INLINE
void outer_product_2x2(
1471 typename sme_packet_traits<Scalar>::type a_lo,
typename sme_packet_traits<Scalar>::type a_hi,
1472 typename sme_packet_traits<Scalar>::type b_lo,
1473 typename sme_packet_traits<Scalar>::type b_hi) __arm_streaming __arm_inout(
"za") {
1474 const svbool_t all = sme_packet_traits<Scalar>::ptrue();
1475 sme_mopa<0>(all, all, a_lo, b_lo);
1476 sme_mopa<1>(all, all, a_lo, b_hi);
1477 sme_mopa<2>(all, all, a_hi, b_lo);
1478 sme_mopa<3>(all, all, a_hi, b_hi);
1511template <
typename RealScalar,
int TileRe,
int TileIm,
bool Vertical,
bool ScaleByAlpha>
1512EIGEN_ALWAYS_INLINE
void sme_read_slice(
1513 svbool_t pg,
typename sme_packet_traits<RealScalar>::type valpha_re,
1514 typename sme_packet_traits<RealScalar>::type valpha_im,
typename sme_packet_traits<RealScalar>::type vzero,
1515 uint32_t slice,
typename sme_packet_traits<RealScalar>::type& lo,
1516 typename sme_packet_traits<RealScalar>::type& hi) __arm_streaming __arm_inout(
"za") {
1517 using Vec =
typename sme_packet_traits<RealScalar>::type;
1520 re = sme_read_ver_za<TileRe>(vzero, pg, slice);
1521 im = sme_read_ver_za<TileIm>(vzero, pg, slice);
1523 re = sme_read_hor_za<TileRe>(vzero, pg, slice);
1524 im = sme_read_hor_za<TileIm>(vzero, pg, slice);
1526 EIGEN_IF_CONSTEXPR (ScaleByAlpha) {
1527 const Vec out_re = pnmadd(pg, im, valpha_im, pmul(pg, re, valpha_re));
1528 const Vec out_im = pmadd(pg, re, valpha_im, pmul(pg, im, valpha_re));
1529 lo = pzip1(out_re, out_im);
1530 hi = pzip2(out_re, out_im);
1541template <
typename RealScalar,
int TileRe,
int TileIm,
bool Vertical,
bool ScaleByAlpha,
typename Index>
1542EIGEN_ALWAYS_INLINE
void sme_accumulate_pair_impl(
1543 RealScalar* EIGEN_RESTRICT p, Index step,
int slices,
int lanes, svbool_t pg,
1544 typename sme_packet_traits<RealScalar>::type valpha_re,
typename sme_packet_traits<RealScalar>::type valpha_im,
1545 typename sme_packet_traits<RealScalar>::type vzero) __arm_streaming __arm_inout(
"za") {
1546 using Traits = sme_packet_traits<RealScalar>;
1547 using Vec =
typename Traits::type;
1548 const int svl = Traits::size();
1549 const svbool_t pl0 = Traits::whilelt(0, lanes);
1550 const svbool_t pl1 = Traits::whilelt(svl, lanes);
1551 for (
int s = 0; s < slices; ++s, p += step) {
1553 sme_read_slice<RealScalar, TileRe, TileIm, Vertical, ScaleByAlpha>(pg, valpha_re, valpha_im, vzero, uint32_t(s), lo,
1555 pstoreu(pl0, p, padd(pl0, ploadu(pl0, p), lo));
1559 RealScalar* EIGEN_RESTRICT phi = sme_offset(p, Index(svl));
1560 pstoreu(pl1, phi, padd(pl1, ploadu(pl1, phi), hi));
1583template <
typename RealScalar,
int TileRe,
int TileIm,
bool Vertical,
typename Index>
1584EIGEN_ALWAYS_INLINE
void sme_accumulate_pair_real_alpha(
1585 RealScalar* EIGEN_RESTRICT p, Index step,
int slices, Index limit, svbool_t pg,
1586 typename sme_packet_traits<RealScalar>::type valpha,
1587 typename sme_packet_traits<RealScalar>::type vzero) __arm_streaming __arm_inout(
"za") {
1588 using Traits = sme_packet_traits<RealScalar>;
1589 using Vec =
typename Traits::type;
1590 const int svl = Traits::size();
1591 constexpr int kPrefetchSlices = 8;
1592 for (
int s = 0; s < slices; ++s, p += step) {
1593 if (s + kPrefetchSlices < limit) {
1595 EIGEN_OPTIMIZATION_BARRIER(pf)
1596 pf = sme_offset(pf, Index(kPrefetchSlices) * step);
1597 __builtin_prefetch(pf, 1, 3);
1598 __builtin_prefetch(sme_offset(pf, Index(svl)), 1, 3);
1602 re = sme_read_ver_za<TileRe>(vzero, pg, uint32_t(s));
1603 im = sme_read_ver_za<TileIm>(vzero, pg, uint32_t(s));
1605 re = sme_read_hor_za<TileRe>(vzero, pg, uint32_t(s));
1606 im = sme_read_hor_za<TileIm>(vzero, pg, uint32_t(s));
1608 const typename Traits::type_x2 c = pld2(pg, p);
1609 pst2(pg, p, pcreate(pmadd(pg, re, valpha, pget<0>(c)), pmadd(pg, im, valpha, pget<1>(c))));
1617template <
typename RealScalar,
int TileRe,
int TileIm,
bool Vertical,
bool FoldRealAlpha,
typename Index>
1618EIGEN_ALWAYS_INLINE
void sme_accumulate_pair(
1619 bool scale_by_alpha, RealScalar* EIGEN_RESTRICT p, Index step,
int slices, Index limit,
int lanes, svbool_t pg,
1620 typename sme_packet_traits<RealScalar>::type valpha_re,
typename sme_packet_traits<RealScalar>::type valpha_im,
1621 typename sme_packet_traits<RealScalar>::type vzero) __arm_streaming __arm_inout(
"za") {
1622 if (FoldRealAlpha && lanes > sme_packet_traits<RealScalar>::size()) {
1623 sme_accumulate_pair_real_alpha<RealScalar, TileRe, TileIm, Vertical>(p, step, slices, limit, pg, valpha_re, vzero);
1624 }
else if (scale_by_alpha) {
1625 sme_accumulate_pair_impl<RealScalar, TileRe, TileIm, Vertical, true>(p, step, slices, lanes, pg, valpha_re,
1628 sme_accumulate_pair_impl<RealScalar, TileRe, TileIm, Vertical, false>(p, step, slices, lanes, pg, valpha_re,
1635template <
typename RealScalar,
int TileRe,
int TileIm,
bool FoldRealAlpha,
typename Index>
1636EIGEN_ALWAYS_INLINE
void sme_store_za_pair(std::complex<RealScalar>* EIGEN_RESTRICT C, Index C_stride_row,
1637 Index C_stride_col, std::complex<RealScalar> alpha, Index row_start,
int pw,
1638 Index col_start,
int cw, Index rows,
1639 Index cols) __arm_streaming __arm_inout(
"za") {
1640 using Scalar = std::complex<RealScalar>;
1641 using Traits = sme_packet_traits<RealScalar>;
1642 using Vec =
typename Traits::type;
1643 const int svl = Traits::size();
1644 const svbool_t pg_m = Traits::whilelt(0, pw);
1645 const svbool_t pg_n = Traits::whilelt(0, cw);
1646 const Vec vzero = pset1<Vec>(RealScalar(0));
1651 const RealScalar* alpha_parts =
reinterpret_cast<const RealScalar*
>(&alpha);
1652 const Vec valpha_re = pset1<Vec>(alpha_parts[0]);
1653 const Vec valpha_im = pset1<Vec>(alpha_parts[1]);
1658 const bool scale = !(alpha_parts[0] == RealScalar(1) && alpha_parts[1] == RealScalar(0));
1659 RealScalar* EIGEN_RESTRICT rC =
reinterpret_cast<RealScalar*
>(C);
1661 if (C_stride_row == 1) {
1664 RealScalar* p = rC + Index(2) * (row_start + col_start * C_stride_col);
1665 sme_accumulate_pair<RealScalar, TileRe, TileIm, true, FoldRealAlpha>(
1666 scale, p, Index(2) * C_stride_col, cw, cols - col_start, 2 * pw, pg_m, valpha_re, valpha_im, vzero);
1667 }
else if (C_stride_col == 1) {
1669 RealScalar* p = rC + Index(2) * (row_start * C_stride_row + col_start);
1670 sme_accumulate_pair<RealScalar, TileRe, TileIm, false, FoldRealAlpha>(
1671 scale, p, Index(2) * C_stride_row, pw, rows - row_start, 2 * cw, pg_n, valpha_re, valpha_im, vzero);
1677 Scalar scratch[sme_block<Scalar>::nr];
1678 RealScalar* rscratch =
reinterpret_cast<RealScalar*
>(scratch);
1679 const int lanes = 2 * cw;
1680 const svbool_t pl0 = Traits::whilelt(0, lanes);
1681 const svbool_t pl1 = Traits::whilelt(svl, lanes);
1682 for (
int ri = 0; ri < pw; ++ri) {
1684 sme_read_slice<RealScalar, TileRe, TileIm, false, true>(pg_n, valpha_re, valpha_im, vzero, uint32_t(ri), lo, hi);
1685 pstoreu(pl0, rscratch, lo);
1688 pstoreu(pl1, rscratch + svl, hi);
1689 for (
int ci = 0; ci < cw; ++ci) {
1690 C[(row_start + ri) * C_stride_row + (col_start + ci) * C_stride_col] += scratch[ci];
1698template <
typename Scalar,
int R,
int C,
bool ConjLhs,
bool ConjRhs,
1699 bool InGrid = (R < sme_block<Scalar>::kGridRows && C < sme_block<Scalar>::kGridCols)>
1700struct sme_complex_cell {
1701 using RealScalar =
typename NumTraits<Scalar>::Real;
1702 using Vec =
typename sme_packet_traits<RealScalar>::type;
1703 static constexpr int kTileRe = 2 * (R * sme_block<Scalar>::kGridCols + C);
1704 static constexpr int kTileIm = kTileRe + 1;
1706 static EIGEN_ALWAYS_INLINE
void accumulate(svbool_t pm, svbool_t pn, Vec a_re, Vec a_im, Vec b_re,
1707 Vec b_im) __arm_streaming __arm_inout(
"za") {
1708 sme_mopa_signed<kTileRe, false>(pm, pn, a_re, b_re);
1709 sme_mopa_signed<kTileRe, ConjLhs == ConjRhs>(pm, pn, a_im, b_im);
1710 sme_mopa_signed<kTileIm, ConjRhs>(pm, pn, a_re, b_im);
1711 sme_mopa_signed<kTileIm, ConjLhs>(pm, pn, a_im, b_re);
1714 template <
bool FoldRealAlpha,
typename Index>
1715 static EIGEN_ALWAYS_INLINE
void store(Scalar* EIGEN_RESTRICT dst, Index C_stride_row, Index C_stride_col,
1716 Scalar alpha, Index row_start,
int pw, Index col_start,
int cw, Index rows,
1717 Index cols) __arm_streaming __arm_inout(
"za") {
1718 sme_store_za_pair<RealScalar, kTileRe, kTileIm, FoldRealAlpha>(dst, C_stride_row, C_stride_col, alpha, row_start,
1719 pw, col_start, cw, rows, cols);
1723template <
typename Scalar,
int R,
int C,
bool ConjLhs,
bool ConjRhs>
1724struct sme_complex_cell<Scalar, R, C, ConjLhs, ConjRhs, false> {
1725 using Vec =
typename sme_packet_traits<typename NumTraits<Scalar>::Real>::type;
1726 static EIGEN_ALWAYS_INLINE
void accumulate(svbool_t, svbool_t, Vec, Vec, Vec, Vec) __arm_streaming __arm_inout(
"za") {
1728 template <
bool FoldRealAlpha,
typename Index>
1729 static EIGEN_ALWAYS_INLINE
void store(Scalar*, Index, Index, Scalar, Index,
int, Index,
int, Index,
1730 Index) __arm_streaming __arm_inout(
"za") {}
1750template <
bool ConjLhs,
bool ConjRhs,
typename Scalar,
typename Index>
1751EIGEN_ALWAYS_INLINE
void sme_process(Scalar* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
1752 const Scalar* EIGEN_RESTRICT blA,
const Scalar* EIGEN_RESTRICT blB, Index depth,
1753 Scalar alpha, Index row_start,
int pw, Index col_start,
int cw,
1754 Index a_step) __arm_streaming __arm_inout(
"za") {
1758 using Traits = sme_packet_traits<Scalar>;
1759 using Vec =
typename Traits::type;
1760 const int svl = Traits::size();
1762 for (
int rt = 0; rt < pw; rt += 2 * svl) {
1763 const int rpw = sme_min(pw - rt, 2 * svl);
1764 const int rlo = sme_min(rpw, svl);
1765 const int rhi = rpw - rlo;
1766 const svbool_t pg_rlo = Traits::whilelt(rt, pw);
1767 const svbool_t pg_rhi = Traits::whilelt(rt + svl, pw);
1769 for (
int ct = 0; ct < cw; ct += 2 * svl) {
1770 const int cpw = sme_min(cw - ct, 2 * svl);
1771 const int clo = sme_min(cpw, svl);
1772 const int chi = cpw - clo;
1773 const svbool_t pg_clo = Traits::whilelt(ct, cw);
1774 const svbool_t pg_chi = Traits::whilelt(ct + svl, cw);
1777 if (pw == 2 * svl && cw == 2 * svl) {
1782 const svcount_t pn = Traits::ptrue_c();
1783 const Index depth_4 = (depth / 4) * 4;
1785 if (a_step == Index(pw)) {
1786 for (; k < depth_4; k += 4) {
1787 typename Traits::type_x4 va_01 = ploadu_x4(pn, &blA[k * pw]);
1788 typename Traits::type_x4 vb_01 = ploadu_x4(pn, &blB[k * cw]);
1791 outer_product_2x2<Scalar>(pget<0>(va_01), pget<1>(va_01), pget<0>(vb_01), pget<1>(vb_01));
1793 outer_product_2x2<Scalar>(pget<2>(va_01), pget<3>(va_01), pget<2>(vb_01), pget<3>(vb_01));
1795 typename Traits::type_x4 va_23 = ploadu_x4(pn, &blA[(k + 2) * pw]);
1796 typename Traits::type_x4 vb_23 = ploadu_x4(pn, &blB[(k + 2) * cw]);
1799 outer_product_2x2<Scalar>(pget<0>(va_23), pget<1>(va_23), pget<0>(vb_23), pget<1>(vb_23));
1801 outer_product_2x2<Scalar>(pget<2>(va_23), pget<3>(va_23), pget<2>(vb_23), pget<3>(vb_23));
1805 for (; k < depth_4; k += 4) {
1806 typename Traits::type_x2 va_0 = ploadu_x2(pn, &blA[k * a_step]);
1807 typename Traits::type_x2 va_1 = ploadu_x2(pn, &blA[(k + 1) * a_step]);
1808 typename Traits::type_x4 vb_01 = ploadu_x4(pn, &blB[k * cw]);
1809 outer_product_2x2<Scalar>(pget<0>(va_0), pget<1>(va_0), pget<0>(vb_01), pget<1>(vb_01));
1810 outer_product_2x2<Scalar>(pget<0>(va_1), pget<1>(va_1), pget<2>(vb_01), pget<3>(vb_01));
1811 typename Traits::type_x2 va_2 = ploadu_x2(pn, &blA[(k + 2) * a_step]);
1812 typename Traits::type_x2 va_3 = ploadu_x2(pn, &blA[(k + 3) * a_step]);
1813 typename Traits::type_x4 vb_23 = ploadu_x4(pn, &blB[(k + 2) * cw]);
1814 outer_product_2x2<Scalar>(pget<0>(va_2), pget<1>(va_2), pget<0>(vb_23), pget<1>(vb_23));
1815 outer_product_2x2<Scalar>(pget<0>(va_3), pget<1>(va_3), pget<2>(vb_23), pget<3>(vb_23));
1819 for (; k < depth; ++k) {
1820 typename Traits::type_x2 va = ploadu_x2(pn, &blA[k * a_step]);
1821 typename Traits::type_x2 vb = ploadu_x2(pn, &blB[k * cw]);
1822 outer_product_2x2<Scalar>(pget<0>(va), pget<1>(va), pget<0>(vb), pget<1>(vb));
1825 for (Index k = 0; k < depth; ++k) {
1826 Vec a_lo = ploadu(pg_rlo, &blA[k * a_step + rt]);
1827 Vec b_lo = ploadu(pg_clo, &blB[k * cw + ct]);
1829 Vec a_hi = ploadu(pg_rhi, sme_offset(blA, k * a_step + rt + svl));
1830 Vec b_hi = ploadu(pg_chi, sme_offset(blB, k * cw + ct + svl));
1832 sme_mopa<0>(pg_rlo, pg_clo, a_lo, b_lo);
1833 if (svptest_any(pg_chi, pg_chi)) sme_mopa<1>(pg_rlo, pg_chi, a_lo, b_hi);
1834 if (svptest_any(pg_rhi, pg_rhi)) {
1835 sme_mopa<2>(pg_rhi, pg_clo, a_hi, b_lo);
1836 if (svptest_any(pg_chi, pg_chi)) sme_mopa<3>(pg_rhi, pg_chi, a_hi, b_hi);
1842 sme_store_2x2_grid(C, C_stride_row, C_stride_col, alpha, row_start + rt, rlo, rhi, col_start + ct, clo, chi);
1853template <
typename Scalar,
bool ConjLhs,
bool ConjRhs,
bool FoldRealAlpha,
typename Index>
1854EIGEN_ALWAYS_INLINE
void sme_store_complex_grid(Scalar* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
1855 Scalar alpha, Index row_start,
int rlo,
int rhi, Index col_start,
1856 int clo,
int chi, Index rows,
1857 Index cols) __arm_streaming __arm_inout(
"za") {
1858 const int svl = sme_packet_traits<typename NumTraits<Scalar>::Real>::size();
1859 sme_complex_cell<Scalar, 0, 0, ConjLhs, ConjRhs>::template store<FoldRealAlpha>(
1860 C, C_stride_row, C_stride_col, alpha, row_start, rlo, col_start, clo, rows, cols);
1862 sme_complex_cell<Scalar, 0, 1, ConjLhs, ConjRhs>::template store<FoldRealAlpha>(
1863 C, C_stride_row, C_stride_col, alpha, row_start, rlo, col_start + svl, chi, rows, cols);
1866 sme_complex_cell<Scalar, 1, 0, ConjLhs, ConjRhs>::template store<FoldRealAlpha>(
1867 C, C_stride_row, C_stride_col, alpha, row_start + svl, rhi, col_start, clo, rows, cols);
1869 sme_complex_cell<Scalar, 1, 1, ConjLhs, ConjRhs>::template store<FoldRealAlpha>(
1870 C, C_stride_row, C_stride_col, alpha, row_start + svl, rhi, col_start + svl, chi, rows, cols);
1884template <
bool ConjLhs,
bool ConjRhs,
bool FoldRealAlpha,
typename RealScalar,
typename Index>
1885EIGEN_ALWAYS_INLINE
void sme_process(std::complex<RealScalar>* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
1886 const std::complex<RealScalar>* EIGEN_RESTRICT blA,
1887 const std::complex<RealScalar>* EIGEN_RESTRICT blB, Index depth,
1888 std::complex<RealScalar> alpha, Index row_start,
int pw, Index col_start,
int cw,
1889 Index lhs_step, Index rows, Index cols) __arm_streaming __arm_inout(
"za") {
1891 EIGEN_UNUSED_VARIABLE(lhs_step);
1892 using Scalar = std::complex<RealScalar>;
1893 using Traits = sme_packet_traits<RealScalar>;
1894 using Vec =
typename Traits::type;
1895 const int svl = Traits::size();
1896 constexpr int GridRows = sme_block<Scalar>::kGridRows;
1897 constexpr int GridCols = sme_block<Scalar>::kGridCols;
1898 using Cell00 = sme_complex_cell<Scalar, 0, 0, ConjLhs, ConjRhs>;
1899 using Cell01 = sme_complex_cell<Scalar, 0, 1, ConjLhs, ConjRhs>;
1900 using Cell10 = sme_complex_cell<Scalar, 1, 0, ConjLhs, ConjRhs>;
1901 using Cell11 = sme_complex_cell<Scalar, 1, 1, ConjLhs, ConjRhs>;
1903 const RealScalar* EIGEN_RESTRICT rA =
reinterpret_cast<const RealScalar*
>(blA);
1904 const RealScalar* EIGEN_RESTRICT rB =
reinterpret_cast<const RealScalar*
>(blB);
1905 const Index a_step = Index(2 * pw);
1906 const Index b_step = Index(2 * cw);
1908 for (
int rt = 0; rt < pw; rt += GridRows * svl) {
1909 const int rpw = sme_min(pw - rt, GridRows * svl);
1910 const int r0 = sme_min(rpw, svl);
1911 const int r1 = rpw - r0;
1912 const svbool_t pg_r0 = Traits::whilelt(0, rpw);
1913 const svbool_t pg_r1 = Traits::whilelt(svl, rpw);
1915 for (
int ct = 0; ct < cw; ct += GridCols * svl) {
1916 const int cpw = sme_min(cw - ct, GridCols * svl);
1917 const int c0 = sme_min(cpw, svl);
1918 const int c1 = cpw - c0;
1919 const svbool_t pg_c0 = Traits::whilelt(0, cpw);
1920 const svbool_t pg_c1 = Traits::whilelt(svl, cpw);
1923 for (Index k = 0; k < depth; ++k) {
1924 const RealScalar* pa = rA + k * a_step + Index(rt);
1925 const RealScalar* pb = rB + k * b_step + Index(ct);
1931 const Vec a0_re = ploadu(pg_r0, pa);
1932 const Vec a0_im = ploadu(pg_r0, pa + pw);
1933 const Vec b0_re = ploadu(pg_c0, pb);
1934 const Vec b0_im = ploadu(pg_c0, pb + cw);
1935 const Vec a1_re = ploadu(pg_r1, sme_offset(pa, Index(svl)));
1936 const Vec a1_im = ploadu(pg_r1, sme_offset(pa, Index(pw) + Index(svl)));
1937 const Vec b1_re = ploadu(pg_c1, sme_offset(pb, Index(svl)));
1938 const Vec b1_im = ploadu(pg_c1, sme_offset(pb, Index(cw) + Index(svl)));
1939 Cell00::accumulate(pg_r0, pg_c0, a0_re, a0_im, b0_re, b0_im);
1940 if (svptest_any(pg_c1, pg_c1)) {
1941 Cell01::accumulate(pg_r0, pg_c1, a0_re, a0_im, b1_re, b1_im);
1943 if (svptest_any(pg_r1, pg_r1)) {
1944 Cell10::accumulate(pg_r1, pg_c0, a1_re, a1_im, b0_re, b0_im);
1945 if (svptest_any(pg_c1, pg_c1)) {
1946 Cell11::accumulate(pg_r1, pg_c1, a1_re, a1_im, b1_re, b1_im);
1951 sme_store_complex_grid<Scalar, ConjLhs, ConjRhs, FoldRealAlpha>(
1952 C, C_stride_row, C_stride_col, alpha, row_start + rt, r0, r1, col_start + ct, c0, c1, rows, cols);
1959template <
typename Scalar,
typename Index>
1960EIGEN_ALWAYS_INLINE
void sme_process_narrow(Scalar* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
1961 const Scalar* EIGEN_RESTRICT blA0, Index step0,
1962 const Scalar* EIGEN_RESTRICT blA1, Index step1,
int pw1,
1963 const Scalar* EIGEN_RESTRICT blB,
int cw, Index depth, Scalar alpha,
1964 Index row_start, Index col_start) __arm_streaming __arm_inout(
"za") {
1965 using Traits = sme_packet_traits<Scalar>;
1966 using Vec =
typename Traits::type;
1967 const int svl = Traits::size();
1968 const svcount_t pc = Traits::ptrue_c();
1969 const svbool_t all = Traits::ptrue();
1970 const svbool_t pn = Traits::whilelt(0, cw);
1971 const svbool_t p1lo = Traits::whilelt(0, pw1);
1972 const svbool_t p1hi = Traits::whilelt(svl, pw1);
1974 for (Index k = 0; k < depth; ++k) {
1975 const auto a0 = ploadu_x2(pc, blA0 + k * step0);
1976 const Vec a1lo = ploadu(p1lo, blA1 + k * step1);
1977 const Vec a1hi = ploadu(p1hi, sme_offset(blA1, k * step1 + svl));
1978 const Vec b = ploadu(pn, blB + k * cw);
1979 sme_mopa<0>(all, pn, pget<0>(a0), b);
1980 sme_mopa<1>(all, pn, pget<1>(a0), b);
1981 sme_mopa<2>(p1lo, pn, a1lo, b);
1982 sme_mopa<3>(p1hi, pn, a1hi, b);
1984 sme_store_za_tile<Scalar, 0>(C, C_stride_row, C_stride_col, alpha, row_start, svl, col_start, cw);
1985 sme_store_za_tile<Scalar, 1>(C, C_stride_row, C_stride_col, alpha, row_start + svl, svl, col_start, cw);
1986 sme_store_za_tile<Scalar, 2>(C, C_stride_row, C_stride_col, alpha, row_start + 2 * svl, sme_min(pw1, svl), col_start,
1989 sme_store_za_tile<Scalar, 3>(C, C_stride_row, C_stride_col, alpha, row_start + 3 * svl, pw1 - svl, col_start, cw);
1993template <
int Dst,
int Src,
typename Scalar>
1994EIGEN_ALWAYS_INLINE
void sme_fold_tile() __arm_streaming __arm_inout(
"za") {
1995 const svbool_t all = sme_packet_traits<Scalar>::ptrue();
1996 for (
int c = 0; c < sme_packet_traits<Scalar>::size(); c += 4) {
1997 const auto d = sme_read_ver_za_vg4<Dst, Scalar>(uint32_t(c)), s = sme_read_ver_za_vg4<Src, Scalar>(uint32_t(c));
1998 sme_write_ver_za_vg4<Dst>(uint32_t(c),
1999 pcreate(padd(all, pget<0>(d), pget<0>(s)), padd(all, pget<1>(d), pget<1>(s)),
2000 padd(all, pget<2>(d), pget<2>(s)), padd(all, pget<3>(d), pget<3>(s))));
2006template <
typename Scalar,
typename Index>
2007EIGEN_ALWAYS_INLINE
void sme_process_split(Scalar* EIGEN_RESTRICT C, Index C_stride_row, Index C_stride_col,
2008 const Scalar* EIGEN_RESTRICT blA,
const Scalar* EIGEN_RESTRICT blB,
2009 Index depth, Scalar alpha, Index row_start,
int pw, Index col_start,
int cw,
2010 Index a_step) __arm_streaming __arm_inout(
"za") {
2011 using Traits = sme_packet_traits<Scalar>;
2012 using Vec =
typename Traits::type;
2013 const int svl = Traits::size();
2014 const svbool_t pr0 = Traits::whilelt(0, pw), pr1 = Traits::whilelt(svl, pw);
2015 const svbool_t pc0 = Traits::whilelt(0, cw), pc1 = Traits::whilelt(svl, cw);
2018 if (pw <= svl && cw <= svl) {
2019 for (; k + 4 <= depth; k += 4) {
2020 const Vec a0 = ploadu(pr0, blA + k * a_step), a1 = ploadu(pr0, blA + (k + 1) * a_step);
2021 const Vec a2 = ploadu(pr0, blA + (k + 2) * a_step), a3 = ploadu(pr0, blA + (k + 3) * a_step);
2022 const Vec b0 = ploadu(pc0, blB + k * cw), b1 = ploadu(pc0, blB + (k + 1) * cw);
2023 const Vec b2 = ploadu(pc0, blB + (k + 2) * cw), b3 = ploadu(pc0, blB + (k + 3) * cw);
2024 sme_mopa<0>(pr0, pc0, a0, b0);
2025 sme_mopa<1>(pr0, pc0, a1, b1);
2026 sme_mopa<2>(pr0, pc0, a2, b2);
2027 sme_mopa<3>(pr0, pc0, a3, b3);
2029 for (; k < depth; ++k) sme_mopa<0>(pr0, pc0, ploadu(pr0, blA + k * a_step), ploadu(pc0, blB + k * cw));
2030 sme_fold_tile<0, 1, Scalar>();
2031 sme_fold_tile<2, 3, Scalar>();
2032 sme_fold_tile<0, 2, Scalar>();
2033 sme_store_za_tile<Scalar, 0>(C, C_stride_row, C_stride_col, alpha, row_start, pw, col_start, cw);
2034 }
else if (pw <= svl) {
2036 for (; k + 2 <= depth; k += 2) {
2037 const Vec a0 = ploadu(pr0, blA + k * a_step), a1 = ploadu(pr0, blA + (k + 1) * a_step);
2038 const Vec b0 = ploadu(pc0, blB + k * cw), b0h = ploadu(pc1, blB + k * cw + svl);
2039 const Vec b1 = ploadu(pc0, blB + (k + 1) * cw), b1h = ploadu(pc1, blB + (k + 1) * cw + svl);
2040 sme_mopa<0>(pr0, pc0, a0, b0);
2041 sme_mopa<1>(pr0, pc1, a0, b0h);
2042 sme_mopa<2>(pr0, pc0, a1, b1);
2043 sme_mopa<3>(pr0, pc1, a1, b1h);
2046 const Vec a0 = ploadu(pr0, blA + k * a_step);
2047 sme_mopa<0>(pr0, pc0, a0, ploadu(pc0, blB + k * cw));
2048 sme_mopa<1>(pr0, pc1, a0, ploadu(pc1, blB + k * cw + svl));
2050 sme_fold_tile<0, 2, Scalar>();
2051 sme_fold_tile<1, 3, Scalar>();
2052 sme_store_2x2_grid(C, C_stride_row, C_stride_col, alpha, row_start, pw, 0, col_start, svl, cw - svl);
2055 for (; k + 2 <= depth; k += 2) {
2056 const Vec a0 = ploadu(pr0, blA + k * a_step), a0h = ploadu(pr1, blA + k * a_step + svl);
2057 const Vec a1 = ploadu(pr0, blA + (k + 1) * a_step), a1h = ploadu(pr1, blA + (k + 1) * a_step + svl);
2058 const Vec b0 = ploadu(pc0, blB + k * cw), b1 = ploadu(pc0, blB + (k + 1) * cw);
2059 sme_mopa<0>(pr0, pc0, a0, b0);
2060 sme_mopa<2>(pr1, pc0, a0h, b0);
2061 sme_mopa<1>(pr0, pc0, a1, b1);
2062 sme_mopa<3>(pr1, pc0, a1h, b1);
2065 const Vec b0 = ploadu(pc0, blB + k * cw);
2066 sme_mopa<0>(pr0, pc0, ploadu(pr0, blA + k * a_step), b0);
2067 sme_mopa<2>(pr1, pc0, ploadu(pr1, blA + k * a_step + svl), b0);
2069 sme_fold_tile<0, 1, Scalar>();
2070 sme_fold_tile<2, 3, Scalar>();
2071 sme_store_2x2_grid(C, C_stride_row, C_stride_col, alpha, row_start, svl, pw - svl, col_start, cw, 0);
2077template <
bool ConjLhs,
bool ConjRhs,
bool FoldRealAlpha,
typename Scalar,
typename Index>
2078EIGEN_ALWAYS_INLINE
void sme_process_block(Scalar* C, Index rs, Index cs,
const Scalar* blA,
const Scalar* blB,
2079 Index depth, Scalar alpha, Index row_start,
int pw, Index col_start,
int cw,
2080 Index a_step,
bool split_ok, Index,
2081 Index) __arm_streaming __arm_inout(
"za") {
2082 const int svl = sme_packet_traits<Scalar>::size();
2083 if (split_ok && (pw <= svl || cw <= svl))
2084 sme_process_split(C, rs, cs, blA, blB, depth, alpha, row_start, pw, col_start, cw, a_step);
2086 sme_process<ConjLhs, ConjRhs>(C, rs, cs, blA, blB, depth, alpha, row_start, pw, col_start, cw, a_step);
2088template <
bool ConjLhs,
bool ConjRhs,
bool FoldRealAlpha,
typename RealScalar,
typename Index>
2089EIGEN_ALWAYS_INLINE
void sme_process_block(std::complex<RealScalar>* C, Index rs, Index cs,
2090 const std::complex<RealScalar>* blA,
const std::complex<RealScalar>* blB,
2091 Index depth, std::complex<RealScalar> alpha, Index row_start,
int pw,
2092 Index col_start,
int cw, Index a_step,
bool, Index rows,
2093 Index cols) __arm_streaming __arm_inout(
"za") {
2094 sme_process<ConjLhs, ConjRhs, FoldRealAlpha>(C, rs, cs, blA, blB, depth, alpha, row_start, pw, col_start, cw, a_step,
2100template <
typename Scalar>
2101EIGEN_ALWAYS_INLINE
bool sme_split_ok() __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
2102 const int svl = sme_packet_traits<typename NumTraits<Scalar>::Real>::size();
2103 return svl >= 4 && sme_block<Scalar>::mr <= 2 * svl && sme_block<Scalar>::nr <= 2 * svl;
2105template <
typename Scalar>
2106EIGEN_ALWAYS_INLINE
bool sme_narrow_ok() __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
2107 return sme_block<Scalar>::mr == 2 * sme_packet_traits<typename NumTraits<Scalar>::Real>::size();
2111template <
typename Scalar,
typename Index>
2112EIGEN_ALWAYS_INLINE
void sme_process_narrow_dispatch(Scalar* C, Index rs, Index cs,
const Scalar* blA0, Index step0,
2113 const Scalar* blA1, Index step1,
int pw1,
const Scalar* blB,
2114 int cw, Index depth, Scalar alpha, Index row_start,
2115 Index col_start) __arm_streaming __arm_inout(
"za") {
2116 sme_process_narrow(C, rs, cs, blA0, step0, blA1, step1, pw1, blB, cw, depth, alpha, row_start, col_start);
2118template <
typename RealScalar,
typename Index>
2119EIGEN_ALWAYS_INLINE
void sme_process_narrow_dispatch(std::complex<RealScalar>*, Index, Index,
2120 const std::complex<RealScalar>*, Index,
2121 const std::complex<RealScalar>*, Index,
int,
2122 const std::complex<RealScalar>*,
int, Index,
2123 std::complex<RealScalar>, Index,
2124 Index) __arm_streaming __arm_inout(
"za") {}
2127template <
typename Scalar,
typename Index>
2128static EIGEN_ALWAYS_INLINE
void sme_prefetch_next_c(
const Scalar* C, Index C_stride_row, Index C_stride_col, Index i,
2129 Index j, Index rows, Index cols,
int mr,
2130 int nr) __arm_streaming_compatible EIGEN_SME_ZA_AGNOSTIC {
2131 if (C_stride_row != 1 || j >= cols)
return;
2132 const Index h = sme_min(rows - i, Index(mr)), w = sme_min(cols - j, Index(nr));
2133 const Index bytes = h * Index(
sizeof(Scalar));
2134 for (Index c = 0; c < w; ++c)
2135 for (Index b = 0; b < bytes; b += 128)
2136 __builtin_prefetch(
reinterpret_cast<const char*
>(C + i + (j + c) * C_stride_col) + b, 1, 2);
2139template <
typename Scalar,
bool ConjLhs,
bool ConjRhs,
bool FoldRealAlpha = false,
typename Index>
2140EIGEN_DONT_INLINE __arm_locally_streaming __arm_new(
"za") void sme_gebp_impl(
2141 Scalar* C, Index C_stride_row, Index C_stride_col, const Scalar* blockA, const Scalar* blockB, Index rows,
2142 Index depth, Index cols, Scalar alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
2143 constexpr int MR = sme_block<Scalar>::mr;
2144 constexpr int NR = sme_block<Scalar>::nr;
2154 const bool prefetch_c = rows * cols * Index(
sizeof(Scalar)) > Index(256 * 1024);
2155 const bool split_ok = sme_split_ok<Scalar>(), narrow_ok = sme_narrow_ok<Scalar>();
2156 for (Index j = 0; j < cols; j += NR) {
2157 const int cw =
static_cast<int>(sme_min(cols - j, Index(NR)));
2158 const Scalar* blB = blockB + j * strideB + offsetB * cw;
2161 EIGEN_IF_CONSTEXPR (!NumTraits<Scalar>::IsComplex) {
2162 if (narrow_ok && cw <= sme_packet_traits<
typename NumTraits<Scalar>::Real>::size())
2163 for (; i + MR < rows; i += 2 * MR) {
2164 const int pw1 =
static_cast<int>(sme_min(rows - i - MR, Index(MR)));
2165 sme_process_narrow_dispatch(C, C_stride_row, C_stride_col, blockA + i * strideA + offsetA * MR, Index(MR),
2166 blockA + (i + MR) * strideA + offsetA * pw1, Index(pw1), pw1, blB, cw, depth,
2170 for (; i < rows; i += MR) {
2171 const int pw =
static_cast<int>(sme_min(rows - i, Index(MR)));
2172 const Scalar* blA = blockA + i * strideA + offsetA * pw;
2174 sme_prefetch_next_c(C, C_stride_row, C_stride_col, i + MR < rows ? i + MR : Index(0),
2175 i + MR < rows ? j : j + NR, rows, cols, MR, NR);
2176 sme_process_block<ConjLhs, ConjRhs, FoldRealAlpha>(C, C_stride_row, C_stride_col, blA, blB, depth, alpha, i, pw,
2177 j, cw, Index(pw), split_ok, rows, cols);
2186template <typename Scalar, bool ConjLhs, bool ConjRhs, bool IsComplex = NumTraits<Scalar>::IsComplex>
2187struct sme_gebp_dispatch {
2188 template <
typename Index>
2189 static void run(Scalar* C, Index C_stride_row, Index C_stride_col,
const Scalar* blockA,
const Scalar* blockB,
2190 Index rows, Index depth, Index cols, Scalar alpha, Index strideA, Index strideB, Index offsetA,
2192 sme_gebp_impl<Scalar, ConjLhs, ConjRhs, false>(C, C_stride_row, C_stride_col, blockA, blockB, rows, depth, cols,
2193 alpha, strideA, strideB, offsetA, offsetB);
2197template <
typename Scalar,
bool ConjLhs,
bool ConjRhs>
2198struct sme_gebp_dispatch<Scalar, ConjLhs, ConjRhs, true> {
2199 template <
typename Index>
2200 static void run(Scalar* C, Index C_stride_row, Index C_stride_col,
const Scalar* blockA,
const Scalar* blockB,
2201 Index rows, Index depth, Index cols, Scalar alpha, Index strideA, Index strideB, Index offsetA,
2203 using RealScalar =
typename NumTraits<Scalar>::Real;
2204 if (numext::imag(alpha) == RealScalar(0) && numext::real(alpha) != RealScalar(1)) {
2205 sme_gebp_impl<Scalar, ConjLhs, ConjRhs, true>(C, C_stride_row, C_stride_col, blockA, blockB, rows, depth, cols,
2206 alpha, strideA, strideB, offsetA, offsetB);
2208 sme_gebp_impl<Scalar, ConjLhs, ConjRhs, false>(C, C_stride_row, C_stride_col, blockA, blockB, rows, depth, cols,
2209 alpha, strideA, strideB, offsetA, offsetB);
2216template <
typename Scalar,
typename Index>
2217EIGEN_DONT_INLINE __arm_locally_streaming __arm_new(
"za") void sme_gebp_impl_direct_lhs(
2218 Scalar* C, Index C_stride_row, Index C_stride_col, const Scalar* lhs, Index lda, const Scalar* blockB, Index rows,
2219 Index depth, Index cols, Scalar alpha, Index strideB, Index offsetB) {
2220 constexpr int MR = sme_block<Scalar>::mr;
2221 constexpr int NR = sme_block<Scalar>::nr;
2222 const bool split_ok = sme_split_ok<Scalar>(), narrow_ok = sme_narrow_ok<Scalar>();
2223 for (Index j = 0; j < cols; j += NR) {
2224 const int cw =
static_cast<int>(sme_min(cols - j, Index(NR)));
2225 const Scalar* blB = blockB + j * strideB + offsetB * cw;
2227 if (narrow_ok && cw <= sme_packet_traits<Scalar>::size())
2228 for (; i + MR < rows; i += 2 * MR)
2229 sme_process_narrow(C, C_stride_row, C_stride_col, lhs + i, lda, lhs + i + MR, lda,
2230 static_cast<int>(sme_min(rows - i - MR, Index(MR))), blB, cw, depth, alpha, i, j);
2231 for (; i < rows; i += MR) {
2232 const int pw =
static_cast<int>(sme_min(rows - i, Index(MR)));
2233 sme_process_block<false, false, false>(C, C_stride_row, C_stride_col, lhs + i, blB, depth, alpha, i, pw, j, cw,
2234 lda, split_ok, rows, cols);
2241#ifndef EIGEN_SME_DIRECT_LHS_MAX_STRIDE_BYTES
2242#define EIGEN_SME_DIRECT_LHS_MAX_STRIDE_BYTES 16384
2244#ifndef EIGEN_SME_DIRECT_LHS_MAX_BLOCK_BYTES
2245#define EIGEN_SME_DIRECT_LHS_MAX_BLOCK_BYTES (4 << 20)
2247#ifndef EIGEN_SME_DIRECT_LHS_MAX_SPAN_BYTES
2248#define EIGEN_SME_DIRECT_LHS_MAX_SPAN_BYTES (16 << 20)
2250template <
typename Scalar,
typename Index>
2251bool sme_direct_lhs_ok(Index lhsStride, Index rows, Index depth, Index cols) {
2252#ifdef EIGEN_SME_FORCE_NEON_SMALL_BLOCKS
2253 EIGEN_UNUSED_VARIABLE(lhsStride);
2254 EIGEN_UNUSED_VARIABLE(rows);
2255 EIGEN_UNUSED_VARIABLE(depth);
2256 EIGEN_UNUSED_VARIABLE(cols);
2259 if (NumTraits<Scalar>::IsComplex || depth <= Index(sme_neon_max_depth<Scalar>::value))
return false;
2262 const std::size_t stride_bytes = std::size_t(lhsStride) *
sizeof(Scalar);
2263 if (cols <= Index(sme_block<Scalar>::nr)) {
2264 const std::size_t span_limit = std::size_t(EIGEN_SME_DIRECT_LHS_MAX_SPAN_BYTES);
2265 if (span_limit == 0)
return false;
2266 return cols <= Index(sme_block<Scalar>::nr / 2) || std::size_t(depth) * stride_bytes <= span_limit;
2268 return stride_bytes % 4096 != 0 && stride_bytes <= std::size_t(EIGEN_SME_DIRECT_LHS_MAX_STRIDE_BYTES) &&
2269 std::size_t(rows) * std::size_t(depth) *
sizeof(Scalar) <= std::size_t(EIGEN_SME_DIRECT_LHS_MAX_BLOCK_BYTES);
2278template <
typename Scalar,
int NCol>
2279struct sme_neon_cols;
2281struct sme_neon_cols<float, 4> {
2282 using Vec = float32x4_t;
2283 static EIGEN_ALWAYS_INLINE Vec load(
const float* b) {
return vld1q_f32(b); }
2285 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2286 return vfmaq_laneq_f32(acc, a, b, C);
2289 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2290 return vfmsq_laneq_f32(acc, a, b, C);
2294struct sme_neon_cols<float, 2> {
2295 using Vec = float32x2_t;
2296 static EIGEN_ALWAYS_INLINE Vec load(
const float* b) {
return vld1_f32(b); }
2298 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2299 return vfmaq_lane_f32(acc, a, b, C);
2302 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2303 return vfmsq_lane_f32(acc, a, b, C);
2307struct sme_neon_cols<float, 1> {
2308 using Vec = float32x4_t;
2309 static EIGEN_ALWAYS_INLINE Vec load(
const float* b) {
return vld1q_dup_f32(b); }
2311 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2312 return vfmaq_f32(acc, a, b);
2315 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2316 return vfmsq_f32(acc, a, b);
2320struct sme_neon_cols<double, 4> {
2324 static EIGEN_ALWAYS_INLINE Vec load(
const double* b) {
return {vld1q_f64(b), vld1q_f64(b + 2)}; }
2326 static EIGEN_ALWAYS_INLINE Packet2d madd(Packet2d acc, Packet2d a, Vec b) {
2327 return vfmaq_laneq_f64(acc, a, C < 2 ? b.lo : b.hi, C & 1);
2330 static EIGEN_ALWAYS_INLINE Packet2d nmadd(Packet2d acc, Packet2d a, Vec b) {
2331 return vfmsq_laneq_f64(acc, a, C < 2 ? b.lo : b.hi, C & 1);
2335struct sme_neon_cols<double, 2> {
2336 using Vec = float64x2_t;
2337 static EIGEN_ALWAYS_INLINE Vec load(
const double* b) {
return vld1q_f64(b); }
2339 static EIGEN_ALWAYS_INLINE Packet2d madd(Packet2d acc, Packet2d a, Vec b) {
2340 return vfmaq_laneq_f64(acc, a, b, C);
2343 static EIGEN_ALWAYS_INLINE Packet2d nmadd(Packet2d acc, Packet2d a, Vec b) {
2344 return vfmsq_laneq_f64(acc, a, b, C);
2348struct sme_neon_cols<double, 1> {
2349 using Vec = float64x2_t;
2350 static EIGEN_ALWAYS_INLINE Vec load(
const double* b) {
return vld1q_dup_f64(b); }
2352 static EIGEN_ALWAYS_INLINE Packet2d madd(Packet2d acc, Packet2d a, Vec b) {
2353 return vfmaq_f64(acc, a, b);
2356 static EIGEN_ALWAYS_INLINE Packet2d nmadd(Packet2d acc, Packet2d a, Vec b) {
2357 return vfmsq_f64(acc, a, b);
2362struct sme_neon_cols<float, 8> {
2366 static EIGEN_ALWAYS_INLINE Vec load(
const float* b) {
return {vld1q_f32(b), vld1q_f32(b + 4)}; }
2368 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2369 return vfmaq_laneq_f32(acc, a, C < 4 ? b.lo : b.hi, C & 3);
2372 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2373 return vfmsq_laneq_f32(acc, a, C < 4 ? b.lo : b.hi, C & 3);
2377struct sme_neon_cols<double, 8> {
2381 static EIGEN_ALWAYS_INLINE Vec load(
const double* b) {
2382 return {{vld1q_f64(b), vld1q_f64(b + 2), vld1q_f64(b + 4), vld1q_f64(b + 6)}};
2385 static EIGEN_ALWAYS_INLINE Packet2d madd(Packet2d acc, Packet2d a, Vec b) {
2386 return vfmaq_laneq_f64(acc, a, b.v[C >> 1], C & 1);
2389 static EIGEN_ALWAYS_INLINE Packet2d nmadd(Packet2d acc, Packet2d a, Vec b) {
2390 return vfmsq_laneq_f64(acc, a, b.v[C >> 1], C & 1);
2395template <
int C,
int NCol>
2396struct sme_neon_col_loop {
2397 template <
typename Cols,
typename Packet,
int NPack>
2398 static EIGEN_ALWAYS_INLINE
void real(Packet (&acc)[NPack][NCol],
const Packet (&av)[NPack],
typename Cols::Vec b) {
2399 for (
int p = 0; p < NPack; ++p) acc[p][C] = Cols::template madd<C>(acc[p][C], av[p], b);
2400 sme_neon_col_loop<C + 1, NCol>::template real<Cols, Packet, NPack>(acc, av, b);
2402 template <
bool ConjLhs,
bool ConjRhs,
typename Cols,
typename Packet,
int NPack>
2403 static EIGEN_ALWAYS_INLINE
void cplx(Packet (&acc_re)[NPack][NCol], Packet (&acc_im)[NPack][NCol],
2404 const Packet (&are)[NPack],
const Packet (&aim)[NPack],
typename Cols::Vec bre,
2405 typename Cols::Vec bim) {
2406 for (
int p = 0; p < NPack; ++p) {
2407 acc_re[p][C] = Cols::template madd<C>(acc_re[p][C], are[p], bre);
2408 acc_re[p][C] = (ConjLhs != ConjRhs) ? Cols::template madd<C>(acc_re[p][C], aim[p], bim)
2409 : Cols::template nmadd<C>(acc_re[p][C], aim[p], bim);
2410 acc_im[p][C] = ConjLhs ? Cols::template nmadd<C>(acc_im[p][C], aim[p], bre)
2411 : Cols::template madd<C>(acc_im[p][C], aim[p], bre);
2412 acc_im[p][C] = ConjRhs ? Cols::template nmadd<C>(acc_im[p][C], are[p], bim)
2413 : Cols::template madd<C>(acc_im[p][C], are[p], bim);
2415 sme_neon_col_loop<C + 1, NCol>::template cplx<ConjLhs, ConjRhs, Cols, Packet, NPack>(acc_re, acc_im, are, aim, bre,
2420struct sme_neon_col_loop<NCol, NCol> {
2421 template <
typename Cols,
typename Packet,
int NPack>
2422 static EIGEN_ALWAYS_INLINE
void real(Packet (&)[NPack][NCol],
const Packet (&)[NPack],
typename Cols::Vec) {}
2423 template <
bool ConjLhs,
bool ConjRhs,
typename Cols,
typename Packet,
int NPack>
2424 static EIGEN_ALWAYS_INLINE
void cplx(Packet (&)[NPack][NCol], Packet (&)[NPack][NCol],
const Packet (&)[NPack],
2425 const Packet (&)[NPack],
typename Cols::Vec,
typename Cols::Vec) {}
2430template <
typename Scalar,
typename Index,
int NPack,
int NCol>
2431struct sme_neon_tile {
2432 using Packet =
typename packet_traits<Scalar>::type;
2433 using Cols = sme_neon_cols<Scalar, NCol>;
2434 static constexpr int PS = packet_traits<Scalar>::size;
2436 static constexpr int KU = (NPack * NCol >= 4) ? 1 : 4 / (NPack * NCol);
2437 static EIGEN_ALWAYS_INLINE
void step(Packet (&acc)[NPack][NCol],
const Scalar* a,
const Scalar* b) {
2439 for (
int p = 0; p < NPack; ++p) av[p] = ploadu<Packet>(a + p * PS);
2440 sme_neon_col_loop<0, NCol>::template real<Cols, Packet, NPack>(acc, av, Cols::load(b));
2442 static EIGEN_ALWAYS_INLINE
void run(Scalar* C, Index rs, Index cs,
const Scalar* blA,
const Scalar* blB, Index depth,
2443 Scalar alpha, Index pw, Index cw) {
2444 Packet acc[KU][NPack][NCol];
2445 for (
int u = 0; u < KU; ++u)
2446 for (
int p = 0; p < NPack; ++p)
2447 for (
int c = 0; c < NCol; ++c) acc[u][p][c] = pset1<Packet>(Scalar(0));
2449 for (; k + KU <= depth; k += KU) {
2450 for (
int u = 0; u < KU; ++u) step(acc[u], blA + (k + u) * pw, blB + (k + u) * cw);
2452 for (; k < depth; ++k) step(acc[0], blA + k * pw, blB + k * cw);
2453 for (
int u = 1; u < KU; ++u)
2454 for (
int p = 0; p < NPack; ++p)
2455 for (
int c = 0; c < NCol; ++c) acc[0][p][c] = padd(acc[0][p][c], acc[u][p][c]);
2456 const Packet valpha = pset1<Packet>(alpha);
2457 for (
int c = 0; c < NCol; ++c) {
2458 for (
int p = 0; p < NPack; ++p) {
2459 Scalar* pc = C + Index(p * PS) * rs + Index(c) * cs;
2461 pstoreu(pc, pmadd(acc[0][p][c], valpha, ploadu<Packet>(pc)));
2463 pscatter<Scalar, Packet>(pc, pmadd(acc[0][p][c], valpha, pgather<Scalar, Packet>(pc, rs)), rs);
2470template <
typename Scalar,
typename Index,
int NCol>
2471struct sme_neon_tile<Scalar, Index, 0, NCol> {
2472 static EIGEN_ALWAYS_INLINE
void run(Scalar* C, Index rs, Index cs,
const Scalar* blA,
const Scalar* blB, Index depth,
2473 Scalar alpha, Index pw, Index cw) {
2474 EIGEN_UNUSED_VARIABLE(rs);
2475 constexpr int KU = (NCol >= 4) ? 1 : 4 / NCol;
2476 Scalar acc[KU][NCol];
2477 for (
int u = 0; u < KU; ++u)
2478 for (
int c = 0; c < NCol; ++c) acc[u][c] = Scalar(0);
2480 for (; k + KU <= depth; k += KU) {
2481 for (
int u = 0; u < KU; ++u) {
2482 const Scalar a = blA[(k + u) * pw];
2483 const Scalar* b = blB + (k + u) * cw;
2484 for (
int c = 0; c < NCol; ++c) acc[u][c] += a * b[c];
2487 for (; k < depth; ++k) {
2488 const Scalar a = blA[k * pw];
2489 const Scalar* b = blB + k * cw;
2490 for (
int c = 0; c < NCol; ++c) acc[0][c] += a * b[c];
2492 for (
int u = 1; u < KU; ++u)
2493 for (
int c = 0; c < NCol; ++c) acc[0][c] += acc[u][c];
2494 for (
int c = 0; c < NCol; ++c) C[Index(c) * cs] += alpha * acc[0][c];
2500template <
typename RealScalar,
typename Index,
bool ConjLhs,
bool ConjRhs,
int NPack,
int NCol>
2501struct sme_neon_ctile {
2502 using Scalar = std::complex<RealScalar>;
2503 using Packet =
typename packet_traits<RealScalar>::type;
2504 using Cols = sme_neon_cols<RealScalar, NCol>;
2505 static constexpr int PS = packet_traits<RealScalar>::size;
2506 static EIGEN_ALWAYS_INLINE
void run(Scalar* C, Index rs, Index cs,
const RealScalar* rA,
const RealScalar* rB,
2507 Index depth, Scalar alpha, Index pw, Index cw) {
2508 Packet acc_re[NPack][NCol], acc_im[NPack][NCol];
2509 for (
int p = 0; p < NPack; ++p) {
2510 for (
int c = 0; c < NCol; ++c) {
2511 acc_re[p][c] = pset1<Packet>(RealScalar(0));
2512 acc_im[p][c] = pset1<Packet>(RealScalar(0));
2515 for (Index k = 0; k < depth; ++k) {
2516 const RealScalar* a = rA + k * 2 * pw;
2517 const RealScalar* b = rB + k * 2 * cw;
2518 Packet are[NPack], aim[NPack];
2519 for (
int p = 0; p < NPack; ++p) {
2520 are[p] = ploadu<Packet>(a + p * PS);
2521 aim[p] = ploadu<Packet>(a + pw + p * PS);
2523 sme_neon_col_loop<0, NCol>::template cplx<ConjLhs, ConjRhs, Cols, Packet, NPack>(
2524 acc_re, acc_im, are, aim, Cols::load(b), Cols::load(b + cw));
2526 const RealScalar ar = numext::real(alpha), ai = numext::imag(alpha);
2527 const Packet valpha_re = pset1<Packet>(ar), valpha_im = pset1<Packet>(ai);
2528 for (
int c = 0; c < NCol; ++c) {
2529 for (
int p = 0; p < NPack; ++p) {
2530 RealScalar* pc =
reinterpret_cast<RealScalar*
>(C + Index(p * PS) * rs + Index(c) * cs);
2534 sme_neon_ld2(pc, cre, cim);
2535 cre = pnmadd(valpha_im, acc_im[p][c], pmadd(valpha_re, acc_re[p][c], cre));
2536 cim = pmadd(valpha_im, acc_re[p][c], pmadd(valpha_re, acc_im[p][c], cim));
2537 sme_neon_st2(pc, cre, cim);
2539 RealScalar re[PS], im[PS];
2540 pstoreu(re, acc_re[p][c]);
2541 pstoreu(im, acc_im[p][c]);
2542 for (
int l = 0; l < PS; ++l) {
2543 RealScalar* pl = pc + Index(l) * 2 * rs;
2544 pl[0] += ar * re[l] - ai * im[l];
2545 pl[1] += ar * im[l] + ai * re[l];
2553template <
typename RealScalar,
typename Index,
bool ConjLhs,
bool ConjRhs,
int NCol>
2554struct sme_neon_ctile<RealScalar, Index, ConjLhs, ConjRhs, 0, NCol> {
2555 using Scalar = std::complex<RealScalar>;
2556 static EIGEN_ALWAYS_INLINE
void run(Scalar* C, Index rs, Index cs,
const RealScalar* rA,
const RealScalar* rB,
2557 Index depth, Scalar alpha, Index pw, Index cw) {
2558 EIGEN_UNUSED_VARIABLE(rs);
2559 constexpr int KU = (NCol >= 4) ? 1 : 4 / NCol;
2560 RealScalar acc_re[KU][NCol], acc_im[KU][NCol];
2561 for (
int u = 0; u < KU; ++u)
2562 for (
int c = 0; c < NCol; ++c) acc_re[u][c] = acc_im[u][c] = RealScalar(0);
2563 for (Index k = 0; k < depth; ++k) {
2564 const int u = int(k % KU);
2565 const RealScalar* a = rA + k * 2 * pw;
2566 const RealScalar* b = rB + k * 2 * cw;
2567 const RealScalar are = a[0], aim = ConjLhs ? -a[pw] : a[pw];
2568 for (
int c = 0; c < NCol; ++c) {
2569 const RealScalar bre = b[c], bim = ConjRhs ? -b[cw + c] : b[cw + c];
2570 acc_re[u][c] += are * bre - aim * bim;
2571 acc_im[u][c] += are * bim + aim * bre;
2574 for (
int u = 1; u < KU; ++u)
2575 for (
int c = 0; c < NCol; ++c) {
2576 acc_re[0][c] += acc_re[u][c];
2577 acc_im[0][c] += acc_im[u][c];
2579 const RealScalar ar = numext::real(alpha), ai = numext::imag(alpha);
2580 for (
int c = 0; c < NCol; ++c) {
2581 RealScalar* pc =
reinterpret_cast<RealScalar*
>(C + Index(c) * cs);
2582 pc[0] += ar * acc_re[0][c] - ai * acc_im[0][c];
2583 pc[1] += ar * acc_im[0][c] + ai * acc_re[0][c];
2589template <
bool ConjLhs,
bool ConjRhs,
int NCol,
typename Scalar,
typename Index>
2590EIGEN_ALWAYS_INLINE
void sme_neon_column_group(Scalar* C, Index rs, Index cs,
const Scalar* blA,
const Scalar* blB,
2591 Index depth, Scalar alpha, Index pw, Index cw) {
2592 constexpr Index PS = Index(packet_traits<Scalar>::size);
2594 for (; r + 3 * PS <= pw; r += 3 * PS)
2595 sme_neon_tile<Scalar, Index, 3, NCol>::run(C + r * rs, rs, cs, blA + r, blB, depth, alpha, pw, cw);
2596 for (; r + 2 * PS <= pw; r += 2 * PS)
2597 sme_neon_tile<Scalar, Index, 2, NCol>::run(C + r * rs, rs, cs, blA + r, blB, depth, alpha, pw, cw);
2598 for (; r + PS <= pw; r += PS)
2599 sme_neon_tile<Scalar, Index, 1, NCol>::run(C + r * rs, rs, cs, blA + r, blB, depth, alpha, pw, cw);
2601 sme_neon_tile<Scalar, Index, 0, NCol>::run(C + r * rs, rs, cs, blA + r, blB, depth, alpha, pw, cw);
2604template <
bool ConjLhs,
bool ConjRhs,
int NCol,
typename RealScalar,
typename Index>
2605EIGEN_ALWAYS_INLINE
void sme_neon_column_group(std::complex<RealScalar>* C, Index rs, Index cs,
const RealScalar* rA,
2606 const RealScalar* rB, Index depth, std::complex<RealScalar> alpha,
2607 Index pw, Index cw) {
2609 constexpr Index PS = Index(packet_traits<RealScalar>::size);
2611 for (; r + 2 * PS <= pw; r += 2 * PS)
2612 sme_neon_ctile<RealScalar, Index, ConjLhs, ConjRhs, 2, NCol>::run(C + r * rs, rs, cs, rA + r, rB, depth, alpha, pw,
2614 for (; r + PS <= pw; r += PS)
2615 sme_neon_ctile<RealScalar, Index, ConjLhs, ConjRhs, 1, NCol>::run(C + r * rs, rs, cs, rA + r, rB, depth, alpha, pw,
2618 sme_neon_ctile<RealScalar, Index, ConjLhs, ConjRhs, 0, NCol>::run(C + r * rs, rs, cs, rA + r, rB, depth, alpha, pw,
2623template <
bool ConjLhs,
bool ConjRhs,
typename Scalar,
typename Index>
2624EIGEN_ALWAYS_INLINE
void sme_neon_block(Scalar* C, Index rs, Index cs,
const Scalar* blA,
const Scalar* blB,
2625 Index depth, Scalar alpha, Index pw, Index cw) {
2627 for (; c + 8 <= cw; c += 8)
2628 sme_neon_column_group<ConjLhs, ConjRhs, 8>(C + c * cs, rs, cs, blA, blB + c, depth, alpha, pw, cw);
2629 for (; c + 4 <= cw; c += 4)
2630 sme_neon_column_group<ConjLhs, ConjRhs, 4>(C + c * cs, rs, cs, blA, blB + c, depth, alpha, pw, cw);
2631 for (; c + 2 <= cw; c += 2)
2632 sme_neon_column_group<ConjLhs, ConjRhs, 2>(C + c * cs, rs, cs, blA, blB + c, depth, alpha, pw, cw);
2634 sme_neon_column_group<ConjLhs, ConjRhs, 1>(C + c * cs, rs, cs, blA, blB + c, depth, alpha, pw, cw);
2637template <
bool ConjLhs,
bool ConjRhs,
typename RealScalar,
typename Index>
2638EIGEN_ALWAYS_INLINE
void sme_neon_block(std::complex<RealScalar>* C, Index rs, Index cs,
2639 const std::complex<RealScalar>* blA,
const std::complex<RealScalar>* blB,
2640 Index depth, std::complex<RealScalar> alpha, Index pw, Index cw) {
2641 const RealScalar* rA =
reinterpret_cast<const RealScalar*
>(blA);
2642 const RealScalar* rB =
reinterpret_cast<const RealScalar*
>(blB);
2644 for (; c + 4 <= cw; c += 4)
2645 sme_neon_column_group<ConjLhs, ConjRhs, 4>(C + c * cs, rs, cs, rA, rB + c, depth, alpha, pw, cw);
2646 for (; c + 2 <= cw; c += 2)
2647 sme_neon_column_group<ConjLhs, ConjRhs, 2>(C + c * cs, rs, cs, rA, rB + c, depth, alpha, pw, cw);
2648 for (; c < cw; ++c) sme_neon_column_group<ConjLhs, ConjRhs, 1>(C + c * cs, rs, cs, rA, rB + c, depth, alpha, pw, cw);
2652template <
typename Scalar,
bool ConjLhs,
bool ConjRhs,
typename Index>
2653EIGEN_ALWAYS_INLINE
void sme_gebp_neon(Scalar* C, Index C_stride_row, Index C_stride_col,
const Scalar* blockA,
2654 const Scalar* blockB, Index rows, Index depth, Index cols, Scalar alpha,
2655 Index strideA, Index strideB, Index offsetA, Index offsetB) {
2656 constexpr Index MR = Index(sme_block<Scalar>::mr);
2657 constexpr Index NR = Index(sme_block<Scalar>::nr);
2658 for (Index j = 0; j < cols; j += NR) {
2659 const Index cw = numext::mini(cols - j, NR);
2660 const Scalar* blB = blockB + j * strideB + offsetB * cw;
2661 for (Index i = 0; i < rows; i += MR) {
2662 const Index pw = numext::mini(rows - i, MR);
2663 const Scalar* blA = blockA + i * strideA + offsetA * pw;
2664 sme_neon_block<ConjLhs, ConjRhs>(C + i * C_stride_row + j * C_stride_col, C_stride_row, C_stride_col, blA, blB,
2665 depth, alpha, pw, cw);
2670template <
typename Scalar,
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
2671struct sme_gebp_kernel {
2672 using ResScalar = Scalar;
2674 EIGEN_DONT_INLINE
void operator()(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index rows,
2675 Index depth, Index cols, ResScalar alpha, Index strideA = -1, Index strideB = -1,
2676 Index offsetA = 0, Index offsetB = 0) {
2679 static_assert(NumTraits<Scalar>::IsComplex || (!ConjugateLhs && !ConjugateRhs),
2680 "the SME kernel does not support conjugation of real scalars");
2681 static_assert(mr == sme_block<Scalar>::mr && nr == sme_block<Scalar>::nr,
2682 "the SME kernel expects packed panels of the SME block width");
2684 if (strideA == -1) strideA = depth;
2685 if (strideB == -1) strideB = depth;
2687 if (rows <= 0 || cols <= 0 || depth <= 0)
return;
2689 Scalar* C_base =
const_cast<Scalar*
>(&res(0, 0));
2690 const Index C_stride_row = &res(1, 0) - &res(0, 0);
2691 const Index C_stride_col = &res(0, 1) - &res(0, 0);
2693 if (sme_kernel_with_neon<Scalar>(rows, cols, depth, strideA, strideB)) {
2694 sme_gebp_neon<Scalar, ConjugateLhs, ConjugateRhs>(C_base, C_stride_row, C_stride_col, blockA, blockB, rows, depth,
2695 cols, alpha, strideA, strideB, offsetA, offsetB);
2699 sme_fpsr_guard fpsr;
2700 sme_gebp_dispatch<Scalar, ConjugateLhs, ConjugateRhs>::run(C_base, C_stride_row, C_stride_col, blockA, blockB, rows,
2701 depth, cols, alpha, strideA, strideB, offsetA, offsetB);
2706 EIGEN_DONT_INLINE
void run_direct_lhs(
const DataMapper& res,
const Scalar* lhs, Index lhsStride,
const Scalar* blockB,
2707 Index rows, Index depth, Index cols, ResScalar alpha, Index strideB = -1,
2708 Index offsetB = 0) {
2709 if (strideB == -1) strideB = depth;
2710 if (rows <= 0 || cols <= 0 || depth <= 0)
return;
2711 Scalar* C_base =
const_cast<Scalar*
>(&res(0, 0));
2712 const Index C_stride_row = &res(1, 0) - &res(0, 0);
2713 const Index C_stride_col = &res(0, 1) - &res(0, 0);
2714 sme_fpsr_guard fpsr;
2715 sme_gebp_impl_direct_lhs<Scalar, Index>(C_base, C_stride_row, C_stride_col, lhs, lhsStride, blockB, rows, depth,
2716 cols, alpha, strideB, offsetB);
2720#define EIGEN_SME_DECLARE_GEBP_KERNEL(SCALAR) \
2721 template <typename Index, typename DataMapper, int mr, int nr, bool ConjugateLhs, bool ConjugateRhs> \
2722 struct gebp_kernel<SCALAR, SCALAR, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> \
2723 : sme_gebp_kernel<SCALAR, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {};
2725EIGEN_SME_DECLARE_GEBP_KERNEL(
float)
2726EIGEN_SME_DECLARE_GEBP_KERNEL(std::complex<float>)
2727#ifdef EIGEN_VECTORIZE_SME_F64F64
2728EIGEN_SME_DECLARE_GEBP_KERNEL(
double)
2729EIGEN_SME_DECLARE_GEBP_KERNEL(std::complex<double>)
2732#undef EIGEN_SME_DECLARE_GEBP_KERNEL
2738static_assert(sme_has_gebp_kernel<float, float>::value,
"the SME float kernel is not advertised to the GEMM driver");
2739static_assert(sme_has_gebp_kernel<std::complex<float>, std::complex<float>>::value,
2740 "the SME complex<float> kernel is not advertised to the GEMM driver");
2741#ifdef EIGEN_VECTORIZE_SME_F64F64
2742static_assert(sme_has_gebp_kernel<double, double>::value,
"the SME double kernel is not advertised to the GEMM driver");
2743static_assert(sme_has_gebp_kernel<std::complex<double>, std::complex<double>>::value,
2744 "the SME complex<double> kernel is not advertised to the GEMM driver");
2746static_assert(!sme_has_gebp_kernel<double, double>::value,
2747 "double is advertised to the GEMM driver without FEAT_SME_F64F64 to implement it");
2748static_assert(!sme_has_gebp_kernel<std::complex<double>, std::complex<double>>::value,
2749 "complex<double> is advertised to the GEMM driver without FEAT_SME_F64F64 to implement it");
2791template <
typename Index>
2792static EIGEN_ALWAYS_INLINE
void sme_symm_panel_regions(Index j,
int w, Index depth, Index k2, Index& t_end,
2793 Index& s_end) __arm_streaming_compatible EIGEN_SME_ZA_AGNOSTIC {
2794 const Index raw_t = j - k2, raw_s = j + Index(w) - k2;
2795 t_end = raw_t <= 0 ? Index(0) : sme_min(raw_t, depth);
2796 s_end = raw_s <= 0 ? Index(0) : sme_min(raw_s, depth);
2801template <
typename Scalar,
int StorageOrder,
bool IsLhs,
typename Index>
2802EIGEN_DONT_INLINE __arm_locally_streaming __arm_new(
"za") void sme_symm_pack_dense_regions(
2803 Scalar* block, const Scalar* EIGEN_RESTRICT base, Index stride, Index depth, Index outer, Index k2) {
2804 constexpr int PACK = IsLhs ? sme_block<Scalar>::mr : sme_block<Scalar>::nr;
2805 constexpr bool ColM = (StorageOrder ==
ColMajor);
2808 constexpr bool ConjTransposed = !IsLhs;
2809 constexpr bool ConjDirect = IsLhs;
2811 for (Index j = 0; j < outer; j += PACK) {
2812 const int w =
static_cast<int>(sme_min(outer - j, Index(PACK)));
2813 Scalar* dst = block + j * depth;
2815 sme_symm_panel_regions(j, w, depth, k2, t_end, s_end);
2819 EIGEN_IF_CONSTEXPR (ColM) {
2820 sve_copy_panel_range<ConjTransposed>(dst, base + j + k2 * stride, stride, Index(0), t_end, w);
2822 sme_transpose_pack_range<ConjTransposed>(dst, base + j * stride + k2, stride, Index(0), t_end, w);
2826 if (s_end < depth) {
2827 EIGEN_IF_CONSTEXPR (ColM) {
2828 sme_transpose_pack_range<ConjDirect>(dst, base + k2 + j * stride, stride, s_end, depth, w);
2830 sve_copy_panel_range<ConjDirect>(dst, base + k2 * stride + j, stride, s_end, depth, w);
2844template <
typename Scalar,
int StorageOrder,
bool IsLhs,
typename Index>
2845EIGEN_DONT_INLINE
void sme_symm_pack_straddle(Scalar* block,
const Scalar* EIGEN_RESTRICT base, Index stride,
2846 Index depth, Index outer, Index k2) {
2847 constexpr int PACK = IsLhs ? sme_block<Scalar>::mr : sme_block<Scalar>::nr;
2848 constexpr bool ColM = (StorageOrder ==
ColMajor);
2851 constexpr bool ConjHead = IsLhs;
2852 constexpr bool ConjTail = !IsLhs;
2854 for (Index j = 0; j < outer; j += PACK) {
2855 const int w =
static_cast<int>(numext::mini(outer - j, Index(PACK)));
2856 Scalar* dst = block + j * depth;
2858 sme_symm_panel_regions(j, w, depth, k2, t_end, s_end);
2860 for (Index k = t_end; k < s_end; ++k) {
2861 const Index row = k2 + k;
2862 const int cs =
static_cast<int>(row - j);
2863 Scalar* dst_row = dst + k * w;
2864 const Index wi = Index(w);
2865 EIGEN_IF_CONSTEXPR (ColM) {
2866 const Scalar* head = base + row + j * stride;
2867 for (
int c = 0; c < cs; ++c, head += stride) sme_pack_store<ConjHead>(dst_row, wi, Index(c), *head);
2868 const Scalar* tail = base + j + row * stride;
2872 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsComplex) {
2873 sme_pack_store<false>(dst_row, wi, Index(cs), Scalar(numext::real(tail[cs])));
2876 for (; c < w; ++c) sme_pack_store<ConjTail>(dst_row, wi, Index(c), tail[c]);
2878 const Scalar* head = base + row * stride + j;
2879 for (
int c = 0; c < cs; ++c) sme_pack_store<ConjHead>(dst_row, wi, Index(c), head[c]);
2880 const Scalar* tail = base + (j + Index(cs)) * stride + row;
2882 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsComplex) {
2883 sme_pack_store<false>(dst_row, wi, Index(cs), Scalar(numext::real(*tail)));
2887 for (; c < w; ++c, tail += stride) sme_pack_store<ConjTail>(dst_row, wi, Index(c), *tail);
2894template <
typename Scalar,
int StorageOrder,
bool IsLhs,
typename Index>
2895EIGEN_DONT_INLINE
void sme_symm_pack_panels(Scalar* block,
const Scalar* EIGEN_RESTRICT base, Index stride, Index depth,
2896 Index outer, Index k2) {
2898 sme_fpsr_guard fpsr;
2899 sme_symm_pack_dense_regions<Scalar, StorageOrder, IsLhs, Index>(block, base, stride, depth, outer, k2);
2901 sme_symm_pack_straddle<Scalar, StorageOrder, IsLhs, Index>(block, base, stride, depth, outer, k2);
2906template <
typename Scalar,
int StorageOrder,
typename Index>
2907struct sme_symm_pack_lhs {
2910 EIGEN_DONT_INLINE
void operator()(Scalar* blockA,
const Scalar* lhs_, Index lhsStride, Index cols, Index rows)
const {
2911 sme_symm_pack_panels<Scalar, StorageOrder, true, Index>(blockA, lhs_, lhsStride, cols, rows, Index(0));
2915template <
typename Scalar,
int StorageOrder,
typename Index>
2916struct sme_symm_pack_rhs {
2918 EIGEN_DONT_INLINE
void operator()(Scalar* blockB,
const Scalar* rhs_, Index rhsStride, Index rows, Index cols,
2920 sme_symm_pack_panels<Scalar, StorageOrder, false, Index>(blockB, rhs_, rhsStride, rows, cols, k2);
2924#define EIGEN_SME_DECLARE_SYMM_PACKERS(SCALAR, MR, NR) \
2925 template <typename Index, int Pack2_dummy, int StorageOrder> \
2926 struct symm_pack_lhs<SCALAR, Index, MR, Pack2_dummy, StorageOrder> \
2927 : sme_symm_pack_lhs<SCALAR, StorageOrder, Index> {}; \
2929 template <typename Index, int StorageOrder> \
2930 struct symm_pack_rhs<SCALAR, Index, NR, StorageOrder> : sme_symm_pack_rhs<SCALAR, StorageOrder, Index> {};
2932EIGEN_SME_DECLARE_SYMM_PACKERS(
float, kSmeMr, kSmeNr)
2933EIGEN_SME_DECLARE_SYMM_PACKERS(std::complex<float>, kSmeMrC, kSmeNrC)
2934#ifdef EIGEN_VECTORIZE_SME_F64F64
2935EIGEN_SME_DECLARE_SYMM_PACKERS(
double, kSmeMrD, kSmeNrD)
2936EIGEN_SME_DECLARE_SYMM_PACKERS(std::complex<double>, kSmeMrCD, kSmeNrCD)
2939#undef EIGEN_SME_DECLARE_SYMM_PACKERS
2944template <
typename Scalar>
2945struct sme_tiny_neon;
2947struct sme_tiny_neon<float> {
2948 using V = float32x4_t;
2949 static constexpr int PS = 4;
2950 static EIGEN_ALWAYS_INLINE V ld(
const float* p) {
return vld1q_f32(p); }
2951 static EIGEN_ALWAYS_INLINE V zero() {
return vdupq_n_f32(0.f); }
2953 static EIGEN_ALWAYS_INLINE V ld_part(
const float* p,
int n) {
2954 if (n <= 0)
return zero();
2955 if (n == 1)
return vld1q_lane_f32(p, zero(), 0);
2956 const V v = vcombine_f32(vld1_f32(p), vdup_n_f32(0.f));
2957 return n == 2 ? v : vld1q_lane_f32(p + 2, v, 2);
2960 static EIGEN_ALWAYS_INLINE V ld_tbl(
const float* p, uint8x16_t idx) {
2961 return vreinterpretq_f32_u8(vqtbl1q_u8(vreinterpretq_u8_f32(vld1q_f32(p)), idx));
2964 static EIGEN_ALWAYS_INLINE V fma_lane(V c, V a, V b) {
2965 return vfmaq_laneq_f32(c, a, b, L);
2967 static EIGEN_ALWAYS_INLINE V fma_n(V c, V a,
float b) {
return vfmaq_n_f32(c, a, b); }
2968 static EIGEN_ALWAYS_INLINE
void st(
float* p, V v) { vst1q_f32(p, v); }
2969 static EIGEN_ALWAYS_INLINE V add(V a, V b) {
return vaddq_f32(a, b); }
2971 static EIGEN_ALWAYS_INLINE
void transpose(V (&r)[PS]) {
2972 const float64x2_t t0 = vreinterpretq_f64_f32(vtrn1q_f32(r[0], r[1]));
2973 const float64x2_t t1 = vreinterpretq_f64_f32(vtrn2q_f32(r[0], r[1]));
2974 const float64x2_t t2 = vreinterpretq_f64_f32(vtrn1q_f32(r[2], r[3]));
2975 const float64x2_t t3 = vreinterpretq_f64_f32(vtrn2q_f32(r[2], r[3]));
2976 r[0] = vreinterpretq_f32_f64(vtrn1q_f64(t0, t2));
2977 r[1] = vreinterpretq_f32_f64(vtrn1q_f64(t1, t3));
2978 r[2] = vreinterpretq_f32_f64(vtrn2q_f64(t0, t2));
2979 r[3] = vreinterpretq_f32_f64(vtrn2q_f64(t1, t3));
2983struct sme_tiny_neon<double> {
2984 using V = float64x2_t;
2985 static constexpr int PS = 2;
2986 static EIGEN_ALWAYS_INLINE V ld(
const double* p) {
return vld1q_f64(p); }
2987 static EIGEN_ALWAYS_INLINE V zero() {
return vdupq_n_f64(0.); }
2988 static EIGEN_ALWAYS_INLINE V ld_part(
const double* p,
int n) {
return n > 0 ? vld1q_lane_f64(p, zero(), 0) : zero(); }
2989 static EIGEN_ALWAYS_INLINE V ld_tbl(
const double* p, uint8x16_t idx) {
2990 return vreinterpretq_f64_u8(vqtbl1q_u8(vreinterpretq_u8_f64(vld1q_f64(p)), idx));
2993 static EIGEN_ALWAYS_INLINE V fma_lane(V c, V a, V b) {
2994 return vfmaq_laneq_f64(c, a, b, L);
2996 static EIGEN_ALWAYS_INLINE V fma_n(V c, V a,
double b) {
return vfmaq_n_f64(c, a, b); }
2997 static EIGEN_ALWAYS_INLINE
void st(
double* p, V v) { vst1q_f64(p, v); }
2998 static EIGEN_ALWAYS_INLINE V add(V a, V b) {
return vaddq_f64(a, b); }
2999 static EIGEN_ALWAYS_INLINE
void transpose(V (&r)[PS]) {
3000 const V t0 = vtrn1q_f64(r[0], r[1]);
3001 r[1] = vtrn2q_f64(r[0], r[1]);
3008template <
typename T,
int RV,
int NC,
int S>
3009struct sme_tiny_step {
3010 using V =
typename T::V;
3011 static constexpr int PS = T::PS;
3014 static EIGEN_ALWAYS_INLINE
void col(V (&acc)[S][NC][RV],
const V (&a)[PS][RV],
const V (&b)[NC],
3015 std::integral_constant<int, U>) {
3017 for (
int j = 0; j < NC; ++j) {
3019 for (
int r = 0; r < RV; ++r) acc[U % S][j][r] = T::template fma_lane<U>(acc[U % S][j][r], a[U][r], b[j]);
3021 col(acc, a, b, std::integral_constant<int, U + 1>());
3023 static EIGEN_ALWAYS_INLINE
void col(V (&)[S][NC][RV],
const V (&)[PS][RV],
const V (&)[NC],
3024 std::integral_constant<int, PS>) {}
3027 static EIGEN_ALWAYS_INLINE
void row(V (&acc)[NC][RV],
const V (&a)[RV],
const V (&b)[NC / PS],
3028 std::integral_constant<int, J>) {
3030 for (
int r = 0; r < RV; ++r) acc[J][r] = T::template fma_lane<J % PS>(acc[J][r], a[r], b[J / PS]);
3031 row(acc, a, b, std::integral_constant<int, J + 1>());
3033 static EIGEN_ALWAYS_INLINE
void row(V (&)[NC][RV],
const V (&)[RV],
const V (&)[NC / PS],
3034 std::integral_constant<int, NC>) {}
3039template <
typename Scalar,
int RV,
int NC,
int LhsOrder,
int RhsOrder,
typename Index>
3040EIGEN_DONT_INLINE
void sme_tiny_gemm_kernel(Index m, Index n, Index depth,
const Scalar* A, Index lda,
const Scalar* B,
3041 Index ldb, Scalar* C, Index incr, Index ldc, Scalar alpha) {
3042 using T = sme_tiny_neon<Scalar>;
3043 using V =
typename T::V;
3044 constexpr int PS = T::PS, MR = RV * PS, S = 16 / (RV * NC) < PS ? 16 / (RV * NC) : PS;
3045 using Step = sme_tiny_step<T, RV, NC, S>;
3048 for (
int t = 0; t < S; ++t) {
3050 for (
int j = 0; j < NC; ++j) {
3052 for (
int r = 0; r < RV; ++r) acc[t][j][r] = T::zero();
3055 const Scalar* a_row[MR];
3056 for (
int i = 0; i < MR; ++i) a_row[i] = A + numext::mini(Index(i), m - 1) * lda;
3057 const Scalar* b_col[NC];
3058 for (
int j = 0; j < NC; ++j) b_col[j] = B + numext::mini(Index(j), n - 1) * ldb;
3061 int a_cnt[RV], b_cnt[NC / PS];
3062 Index a_off[RV], b_off[NC / PS];
3063 uint8x16_t a_idx[RV], b_idx[NC / PS];
3064 const auto shift_table = [](
int cnt) {
3065 EIGEN_ALIGN16 uint8_t t[16];
3066 for (
int b = 0; b < 16; ++b)
3067 t[b] = b < cnt *
int(
sizeof(Scalar)) ? uint8_t(b + (PS - cnt) *
int(
sizeof(Scalar))) : 0xff;
3070 for (
int r = 0; r < RV; ++r) {
3071 a_cnt[r] = int(numext::mini(Index(PS), m - r * PS));
3072 a_off[r] = r * PS + a_cnt[r] - PS;
3073 a_idx[r] = shift_table(a_cnt[r]);
3075 for (
int c = 0; c < NC / PS; ++c) {
3076 b_cnt[c] = int(numext::maxi(Index(0), numext::mini(Index(PS), n - c * PS)));
3077 b_off[c] = numext::maxi(Index(0), Index(c * PS + b_cnt[c] - PS));
3078 b_idx[c] = shift_table(b_cnt[c]);
3080 const bool a_full = m >= PS, b_full = n >= PS;
3081 const Index kmain = (depth / PS) * PS;
3082 for (Index k = 0; k < kmain; k += PS) {
3084 EIGEN_IF_CONSTEXPR (LhsOrder ==
ColMajor) {
3086 for (
int u = 0; u < PS; ++u) {
3088 for (
int r = 0; r < RV; ++r)
3089 a[u][r] = a_cnt[r] == PS ? T::ld(A + (k + u) * lda + r * PS)
3090 : a_full ? T::ld_tbl(A + (k + u) * lda + a_off[r], a_idx[r])
3091 : T::ld_part(A + (k + u) * lda + r * PS, a_cnt[r]);
3095 for (
int r = 0; r < RV; ++r) {
3098 for (
int q = 0; q < PS; ++q) blk[q] = T::ld(a_row[r * PS + q] + k);
3101 for (
int u = 0; u < PS; ++u) a[u][r] = blk[u];
3104 EIGEN_IF_CONSTEXPR (RhsOrder ==
ColMajor) {
3107 for (
int j = 0; j < NC; ++j) b[j] = T::ld(b_col[j] + k);
3108 Step::col(acc, a, b, std::integral_constant<int, 0>());
3111 for (
int u = 0; u < PS; ++u) {
3114 for (
int c = 0; c < NC / PS; ++c)
3115 b[c] = b_cnt[c] == PS ? T::ld(B + (k + u) * ldb + c * PS)
3116 : b_cnt[c] == 0 ? T::zero()
3117 : b_full ? T::ld_tbl(B + (k + u) * ldb + b_off[c], b_idx[c])
3118 : T::ld_part(B + (k + u) * ldb + c * PS, b_cnt[c]);
3119 Step::row(acc[u % S], a[u], b, std::integral_constant<int, 0>());
3123 for (Index k = kmain; k < depth; ++k) {
3124 EIGEN_ALIGN16 Scalar at[MR];
3125 for (
int i = 0; i < MR; ++i)
3126 at[i] = LhsOrder ==
ColMajor ? A[k * lda + numext::mini(Index(i), m - 1)] : a_row[i][k];
3129 for (
int r = 0; r < RV; ++r) a[r] = T::ld(at + r * PS);
3131 for (
int j = 0; j < NC; ++j) {
3132 const Scalar b = RhsOrder ==
ColMajor ? b_col[j][k] : B[k * ldb + numext::mini(Index(j), n - 1)];
3134 for (
int r = 0; r < RV; ++r) acc[0][j][r] = T::fma_n(acc[0][j][r], a[r], b);
3138 for (
int t = 1; t < S; ++t) {
3140 for (
int j = 0; j < NC; ++j) {
3142 for (
int r = 0; r < RV; ++r) acc[0][j][r] = T::add(acc[0][j][r], acc[t][j][r]);
3145 EIGEN_ALIGN16 Scalar out[NC][MR];
3147 for (
int j = 0; j < NC; ++j) {
3149 for (
int r = 0; r < RV; ++r) T::st(&out[j][r * PS], acc[0][j][r]);
3151 for (Index j = 0; j < n; ++j)
3152 for (Index i = 0; i < m; ++i) C[i * incr + j * ldc] += alpha * out[j][i];
3155template <
typename Scalar,
int LhsOrder,
int RhsOrder,
typename Index>
3156bool sme_tiny_gemm(Index rows, Index cols, Index depth,
const Scalar* lhs, Index lhsStride,
const Scalar* rhs,
3157 Index rhsStride, Scalar* res, Index resIncr, Index resStride, Scalar alpha) {
3158 constexpr int PS = sme_tiny_neon<Scalar>::PS;
3159 if (!sme_tiny_gemm_wins<Scalar>(rows, cols, depth))
return false;
3160 const bool one_vec = rows <= PS, four_cols = cols <= 4;
3161 if (one_vec && four_cols)
3162 sme_tiny_gemm_kernel<Scalar, 1, 4, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3163 resIncr, resStride, alpha);
3165 sme_tiny_gemm_kernel<Scalar, 1, 8, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3166 resIncr, resStride, alpha);
3168 sme_tiny_gemm_kernel<Scalar, 2, 4, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3169 resIncr, resStride, alpha);
3171 sme_tiny_gemm_kernel<Scalar, 2, 8, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3172 resIncr, resStride, alpha);
@ ColMajor
Definition Constants.h:319
@ Vertical
Definition Constants.h:267