10#ifndef EIGEN_PACKET_MATH_SME_H
11#define EIGEN_PACKET_MATH_SME_H
14#include "../../InternalHeaderCheck.h"
47template <
typename RealScalar>
50struct sme_tile_count<float> {
51 static constexpr int value = 4;
54struct sme_tile_count<double> {
55 static constexpr int value = 8;
65#if !__is_identifier(__arm_agnostic)
66#define EIGEN_SME_ZA_AGNOSTIC __arm_agnostic("sme_za_state")
69#ifndef EIGEN_SME_ZA_AGNOSTIC
70#define EIGEN_SME_ZA_AGNOSTIC
78template <
typename Scalar>
79struct sme_packet_traits {};
82struct sme_packet_traits<float> {
83 using type = svfloat32_t;
84 using type_x2 = svfloat32x2_t;
85 using type_x4 = svfloat32x4_t;
86 static EIGEN_ALWAYS_INLINE
int size() __arm_streaming_compatible EIGEN_SME_ZA_AGNOSTIC {
87 return static_cast<int>(svcntsw());
89 static EIGEN_ALWAYS_INLINE svbool_t ptrue() __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
return svptrue_b32(); }
90 static EIGEN_ALWAYS_INLINE svcount_t ptrue_c() __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
return svptrue_c32(); }
91 static EIGEN_ALWAYS_INLINE svbool_t whilelt(int64_t begin, int64_t end) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
92 return svwhilelt_b32(begin, end);
94 static EIGEN_ALWAYS_INLINE svcount_t whilelt_c4(int64_t begin, int64_t end) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
95 return svwhilelt_c32_s64(begin, end, 4);
99#ifdef EIGEN_VECTORIZE_SME_F64F64
101struct sme_packet_traits<double> {
102 using type = svfloat64_t;
103 using type_x2 = svfloat64x2_t;
104 using type_x4 = svfloat64x4_t;
105 static EIGEN_ALWAYS_INLINE
int size() __arm_streaming_compatible EIGEN_SME_ZA_AGNOSTIC {
106 return static_cast<int>(svcntsd());
108 static EIGEN_ALWAYS_INLINE svbool_t ptrue() __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
return svptrue_b64(); }
109 static EIGEN_ALWAYS_INLINE svcount_t ptrue_c() __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
return svptrue_c64(); }
110 static EIGEN_ALWAYS_INLINE svbool_t whilelt(int64_t begin, int64_t end) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
111 return svwhilelt_b64(begin, end);
113 static EIGEN_ALWAYS_INLINE svcount_t whilelt_c4(int64_t begin, int64_t end) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
114 return svwhilelt_c64_s64(begin, end, 4);
121template <
typename Packet>
122struct sme_unpacket_traits {};
124struct sme_unpacket_traits<svfloat32_t> {
126 static EIGEN_ALWAYS_INLINE svfloat32_t dup(
float from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
127 return svdup_f32(from);
130#ifdef EIGEN_VECTORIZE_SME_F64F64
132struct sme_unpacket_traits<svfloat64_t> {
134 static EIGEN_ALWAYS_INLINE svfloat64_t dup(
double from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
135 return svdup_f64(from);
146struct unpacket_traits<svfloat32_t> {};
147#ifdef EIGEN_VECTORIZE_SME_F64F64
149struct unpacket_traits<svfloat64_t> {};
154template <
typename Packet>
155EIGEN_ALWAYS_INLINE Packet
156pset1(
typename sme_unpacket_traits<Packet>::type from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
157 return sme_unpacket_traits<Packet>::dup(from);
167template <
typename Scalar>
168EIGEN_ALWAYS_INLINE
typename sme_packet_traits<Scalar>::type ploadu(
169 svbool_t pg,
const Scalar* from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
170 return svld1(pg, from);
172template <
typename Scalar>
173EIGEN_ALWAYS_INLINE
typename sme_packet_traits<Scalar>::type_x2 ploadu_x2(
174 svcount_t pn,
const Scalar* from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
175 return svld1_x2(pn, from);
177template <
typename Scalar>
178EIGEN_ALWAYS_INLINE
typename sme_packet_traits<Scalar>::type_x4 ploadu_x4(
179 svcount_t pn,
const Scalar* from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
180 return svld1_x4(pn, from);
182template <
typename Scalar>
183EIGEN_ALWAYS_INLINE
void pstoreu(svbool_t pg, Scalar* to,
184 typename sme_packet_traits<Scalar>::type from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
187template <
typename Scalar>
188EIGEN_ALWAYS_INLINE
void pstoreu_x2(
189 svcount_t pn, Scalar* to,
typename sme_packet_traits<Scalar>::type_x2 from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
192template <
typename Scalar>
193EIGEN_ALWAYS_INLINE
void pstoreu_x4(
194 svcount_t pn, Scalar* to,
typename sme_packet_traits<Scalar>::type_x4 from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
199template <
typename Scalar>
200EIGEN_ALWAYS_INLINE
typename sme_packet_traits<Scalar>::type_x2 pld2(
201 svbool_t pg,
const Scalar* from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
202 return svld2(pg, from);
204template <
typename Scalar>
205EIGEN_ALWAYS_INLINE
void pst2(svbool_t pg, Scalar* to,
206 typename sme_packet_traits<Scalar>::type_x2 from) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
213template <
typename Packet>
214EIGEN_ALWAYS_INLINE Packet padd(svbool_t pg, Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
215 return svadd_x(pg, a, b);
217template <
typename Packet>
218EIGEN_ALWAYS_INLINE Packet pmul(svbool_t pg, Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
219 return svmul_x(pg, a, b);
221template <
typename Packet>
222EIGEN_ALWAYS_INLINE Packet pnegate(svbool_t pg, Packet a) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
223 return svneg_x(pg, a);
225template <
typename Packet>
226EIGEN_ALWAYS_INLINE Packet pmadd(svbool_t pg, Packet a, Packet b, Packet c) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
227 return svmla_x(pg, c, a, b);
229template <
typename Packet>
230EIGEN_ALWAYS_INLINE Packet pnmadd(svbool_t pg, Packet a, Packet b, Packet c) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
231 return svmls_x(pg, c, a, b);
234template <
typename Packet>
235EIGEN_ALWAYS_INLINE Packet pmadd_m(svbool_t pg, Packet a, Packet b, Packet c) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
236 return svmla_m(pg, c, a, b);
239template <
typename Packet>
240EIGEN_ALWAYS_INLINE
typename sme_unpacket_traits<Packet>::type predux(svbool_t pg,
241 Packet a) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
242 return svaddv(pg, a);
247EIGEN_ALWAYS_INLINE svfloat32_t pget(svfloat32x2_t v) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
248 return svget2_f32(v, Lane);
251EIGEN_ALWAYS_INLINE svfloat32_t pget(svfloat32x4_t v) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
252 return svget4_f32(v, Lane);
254#ifdef EIGEN_VECTORIZE_SME_F64F64
256EIGEN_ALWAYS_INLINE svfloat64_t pget(svfloat64x2_t v) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
257 return svget2_f64(v, Lane);
260EIGEN_ALWAYS_INLINE svfloat64_t pget(svfloat64x4_t v) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
261 return svget4_f64(v, Lane);
264template <
typename Packet>
265EIGEN_ALWAYS_INLINE
auto pcreate(Packet a, Packet b) EIGEN_SME_ZA_AGNOSTIC __arm_streaming
266 ->
decltype(svcreate2(a, b)) {
267 return svcreate2(a, b);
269template <
typename Packet>
270EIGEN_ALWAYS_INLINE
auto pcreate(Packet a, Packet b, Packet c, Packet d) EIGEN_SME_ZA_AGNOSTIC __arm_streaming
271 ->
decltype(svcreate4(a, b, c, d)) {
272 return svcreate4(a, b, c, d);
278template <
typename Packet>
279EIGEN_ALWAYS_INLINE Packet puzp1(Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
282template <
typename Packet>
283EIGEN_ALWAYS_INLINE Packet puzp2(Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
286template <
typename Packet>
287EIGEN_ALWAYS_INLINE Packet pzip1(Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
290template <
typename Packet>
291EIGEN_ALWAYS_INLINE Packet pzip2(Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
294template <
typename Packet>
295EIGEN_ALWAYS_INLINE Packet psplice(svbool_t pg, Packet a, Packet b) __arm_streaming EIGEN_SME_ZA_AGNOSTIC {
296 return svsplice(pg, a, b);
302struct sme_fpsr_guard {
303 EIGEN_ALWAYS_INLINE sme_fpsr_guard() {
asm volatile(
"mrs %0, fpsr" :
"=r"(value) : :
"memory"); }
304 EIGEN_ALWAYS_INLINE ~sme_fpsr_guard() {
asm volatile(
"msr fpsr, %0" : :
"r"(value) :
"memory"); }
305 sme_fpsr_guard(
const sme_fpsr_guard&) =
delete;
306 sme_fpsr_guard& operator=(
const sme_fpsr_guard&) =
delete;
314EIGEN_ALWAYS_INLINE T sme_min(T a, T b) __arm_streaming_compatible EIGEN_SME_ZA_AGNOSTIC {
315 return a < b ? a : b;
323EIGEN_ALWAYS_INLINE T* sme_offset(T* p, Index n) __arm_streaming_compatible EIGEN_SME_ZA_AGNOSTIC {
324 return reinterpret_cast<T*
>(uintptr_t(p) + ptrdiff_t(n) *
sizeof(T));
331EIGEN_ALWAYS_INLINE
void sme_ld1_hor_za(uint32_t slice, svbool_t pg,
const float* p) __arm_streaming __arm_inout(
"za") {
332 svld1_hor_za32(Tile, slice, pg, p);
335EIGEN_ALWAYS_INLINE svfloat32_t sme_read_hor_za(svfloat32_t zero, svbool_t pg,
336 uint32_t slice) __arm_streaming __arm_in(
"za") {
337 return svread_hor_za32_f32_m(zero, pg, Tile, slice);
340EIGEN_ALWAYS_INLINE svfloat32_t sme_read_ver_za(svfloat32_t zero, svbool_t pg,
341 uint32_t slice) __arm_streaming __arm_in(
"za") {
342 return svread_ver_za32_f32_m(zero, pg, Tile, slice);
346EIGEN_ALWAYS_INLINE
void sme_write_hor_za_vg4(uint32_t slice, svfloat32_t a, svfloat32_t b, svfloat32_t c,
347 svfloat32_t d) __arm_streaming __arm_inout(
"za") {
348 svwrite_hor_za32_f32_vg4(Tile, slice, svcreate4_f32(a, b, c, d));
351EIGEN_ALWAYS_INLINE
void sme_write_ver_za_vg4(uint32_t slice, svfloat32x4_t v) __arm_streaming __arm_inout(
"za") {
352 svwrite_ver_za32_f32_vg4(Tile, slice, v);
356EIGEN_ALWAYS_INLINE
void sme_mopa(svbool_t pm, svbool_t pn, svfloat32_t a,
357 svfloat32_t b) __arm_streaming __arm_inout(
"za") {
358 svmopa_za32_f32_m(Tile, pm, pn, a, b);
361EIGEN_ALWAYS_INLINE
void sme_mops(svbool_t pm, svbool_t pn, svfloat32_t a,
362 svfloat32_t b) __arm_streaming __arm_inout(
"za") {
363 svmops_za32_f32_m(Tile, pm, pn, a, b);
367EIGEN_ALWAYS_INLINE
void sme_madd_za_vg1x4(uint32_t slice, svfloat32x4_t x,
368 svfloat32_t y) __arm_streaming __arm_inout(
"za") {
369 svmla_single_za32_f32_vg1x4(slice, x, y);
371EIGEN_ALWAYS_INLINE
void sme_madd_za_vg1x4(uint32_t slice, svfloat32x4_t x,
372 svfloat32x4_t y) __arm_streaming __arm_inout(
"za") {
373 svmla_za32_f32_vg1x4(slice, x, y);
375EIGEN_ALWAYS_INLINE
void sme_write_za_vg1x4(uint32_t slice, svfloat32x4_t x) __arm_streaming __arm_inout(
"za") {
376 svwrite_za32_f32_vg1x4(slice, x);
379#ifdef EIGEN_VECTORIZE_SME_F64F64
381EIGEN_ALWAYS_INLINE
void sme_ld1_hor_za(uint32_t slice, svbool_t pg,
382 const double* p) __arm_streaming __arm_inout(
"za") {
383 svld1_hor_za64(Tile, slice, pg, p);
386EIGEN_ALWAYS_INLINE svfloat64_t sme_read_hor_za(svfloat64_t zero, svbool_t pg,
387 uint32_t slice) __arm_streaming __arm_in(
"za") {
388 return svread_hor_za64_f64_m(zero, pg, Tile, slice);
391EIGEN_ALWAYS_INLINE svfloat64_t sme_read_ver_za(svfloat64_t zero, svbool_t pg,
392 uint32_t slice) __arm_streaming __arm_in(
"za") {
393 return svread_ver_za64_f64_m(zero, pg, Tile, slice);
396EIGEN_ALWAYS_INLINE
void sme_write_hor_za_vg4(uint32_t slice, svfloat64_t a, svfloat64_t b, svfloat64_t c,
397 svfloat64_t d) __arm_streaming __arm_inout(
"za") {
398 svwrite_hor_za64_f64_vg4(Tile, slice, svcreate4_f64(a, b, c, d));
401EIGEN_ALWAYS_INLINE
void sme_write_ver_za_vg4(uint32_t slice, svfloat64x4_t v) __arm_streaming __arm_inout(
"za") {
402 svwrite_ver_za64_f64_vg4(Tile, slice, v);
405EIGEN_ALWAYS_INLINE
void sme_mopa(svbool_t pm, svbool_t pn, svfloat64_t a,
406 svfloat64_t b) __arm_streaming __arm_inout(
"za") {
407 svmopa_za64_f64_m(Tile, pm, pn, a, b);
410EIGEN_ALWAYS_INLINE
void sme_mops(svbool_t pm, svbool_t pn, svfloat64_t a,
411 svfloat64_t b) __arm_streaming __arm_inout(
"za") {
412 svmops_za64_f64_m(Tile, pm, pn, a, b);
414EIGEN_ALWAYS_INLINE
void sme_madd_za_vg1x4(uint32_t slice, svfloat64x4_t x,
415 svfloat64_t y) __arm_streaming __arm_inout(
"za") {
416 svmla_single_za64_f64_vg1x4(slice, x, y);
418EIGEN_ALWAYS_INLINE
void sme_madd_za_vg1x4(uint32_t slice, svfloat64x4_t x,
419 svfloat64x4_t y) __arm_streaming __arm_inout(
"za") {
420 svmla_za64_f64_vg1x4(slice, x, y);
422EIGEN_ALWAYS_INLINE
void sme_write_za_vg1x4(uint32_t slice, svfloat64x4_t x) __arm_streaming __arm_inout(
"za") {
423 svwrite_za64_f64_vg1x4(slice, x);
429template <
typename Scalar>
432struct sme_za_read<float> {
434 static EIGEN_ALWAYS_INLINE svfloat32x4_t ver_vg4(uint32_t slice) __arm_streaming __arm_in(
"za") {
435 return svread_ver_za32_f32_vg4(Tile, slice);
437 static EIGEN_ALWAYS_INLINE svfloat32x4_t vg1x4(uint32_t slice) __arm_streaming __arm_in(
"za") {
438 return svread_za32_f32_vg1x4(slice);
441#ifdef EIGEN_VECTORIZE_SME_F64F64
443struct sme_za_read<double> {
445 static EIGEN_ALWAYS_INLINE svfloat64x4_t ver_vg4(uint32_t slice) __arm_streaming __arm_in(
"za") {
446 return svread_ver_za64_f64_vg4(Tile, slice);
448 static EIGEN_ALWAYS_INLINE svfloat64x4_t vg1x4(uint32_t slice) __arm_streaming __arm_in(
"za") {
449 return svread_za64_f64_vg1x4(slice);
453template <
int Tile,
typename Scalar>
454EIGEN_ALWAYS_INLINE
typename sme_packet_traits<Scalar>::type_x4 sme_read_ver_za_vg4(
455 uint32_t slice) __arm_streaming __arm_in(
"za") {
456 return sme_za_read<Scalar>::template ver_vg4<Tile>(slice);
458template <
typename Scalar>
459EIGEN_ALWAYS_INLINE
typename sme_packet_traits<Scalar>::type_x4 sme_read_za_vg1x4(
460 uint32_t slice) __arm_streaming __arm_in(
"za") {
461 return sme_za_read<Scalar>::vg1x4(slice);
466template <
int Tile,
bool Subtract,
typename Packet>
467EIGEN_ALWAYS_INLINE
void sme_mopa_signed(svbool_t pm, svbool_t pn, Packet a,
468 Packet b) __arm_streaming __arm_inout(
"za") {
469 EIGEN_IF_CONSTEXPR (Subtract) {
470 sme_mops<Tile>(pm, pn, a, b);
472 sme_mopa<Tile>(pm, pn, a, b);