3#ifndef EIGEN_SME_VECTOR_KERNELS_H
4#define EIGEN_SME_VECTOR_KERNELS_H
7#include "../../InternalHeaderCheck.h"
13#define EIGEN_SME_VECTOR_UNROLL4 _Pragma("unroll")
15#define EIGEN_SME_VECTOR_UNROLL4 _Pragma("GCC unroll 4")
17#define EIGEN_SME_VECTOR_UNROLL4
21static EIGEN_ALWAYS_INLINE
bool sme_rounds_toward_negative() {
23 asm volatile(
"mrs %0, fpcr" :
"=r"(fpcr));
24 return ((fpcr >> 22) & 3) == 2;
27template <
typename Scalar,
typename Index>
28__arm_new(
"za") __arm_locally_streaming
29 EIGEN_DONT_INLINE
void sme_axpy(Index n, const Scalar* x, Scalar* y, Scalar alpha, Index prefix) {
30 using Traits = sme_packet_traits<Scalar>;
31 using Vec =
typename Traits::type;
32 const Index lanes = Traits::size();
33 const auto pn = Traits::ptrue_c();
34 const auto a = pset1<Vec>(alpha);
36 for (; i < prefix; i += prefix - i < lanes ? prefix - i : lanes) {
37 auto active = Traits::whilelt(i, prefix);
38 pstoreu(active, y + i, pmadd(active, ploadu(active, x + i), a, ploadu(active, y + i)));
40 for (; i <= n - 16 * lanes; i += 16 * lanes) {
41 EIGEN_SME_VECTOR_UNROLL4
42 for (
int k = 0; k < 4; ++k) {
43 auto xv = ploadu_x4(pn, x + i + k * 4 * lanes);
44 sme_write_za_vg1x4(k, ploadu_x4(pn, y + i + k * 4 * lanes));
45 sme_madd_za_vg1x4(k, xv, a);
47 EIGEN_SME_VECTOR_UNROLL4
48 for (
int k = 0; k < 4; ++k) pstoreu_x4(pn, y + i + k * 4 * lanes, sme_read_za_vg1x4<Scalar>(k));
50 for (; i <= n - 4 * lanes; i += 4 * lanes) {
51 auto xv = ploadu_x4(pn, x + i);
52 sme_write_za_vg1x4(0, ploadu_x4(pn, y + i));
53 sme_madd_za_vg1x4(0, xv, a);
54 pstoreu_x4(pn, y + i, sme_read_za_vg1x4<Scalar>(0));
56 for (; i < n; i += n - i < lanes ? n - i : lanes) {
57 auto tail = Traits::whilelt(i, n);
58 pstoreu(tail, y + i, pmadd(tail, ploadu(tail, x + i), a, ploadu(tail, y + i)));
62template <
typename Scalar,
typename Index>
63__arm_new(
"za") __arm_locally_streaming EIGEN_DONT_INLINE Scalar sme_dot(Index n, const Scalar* x, const Scalar* y) {
64 using Traits = sme_packet_traits<Scalar>;
65 using Vec =
typename Traits::type;
66 const Index lanes = Traits::size();
67 const auto pg = Traits::ptrue();
68 const auto pn = Traits::ptrue_c();
70 const auto negative_zero = pset1<Vec>(Scalar(-0.0));
71 EIGEN_SME_VECTOR_UNROLL4
72 for (
int k = 0; k < 4; ++k)
73 sme_write_za_vg1x4(k, pcreate(negative_zero, negative_zero, negative_zero, negative_zero));
75 for (; i <= n - 16 * lanes; i += 16 * lanes) {
76 EIGEN_SME_VECTOR_UNROLL4
77 for (
int k = 0; k < 4; ++k)
78 sme_madd_za_vg1x4(k, ploadu_x4(pn, x + i + k * 4 * lanes), ploadu_x4(pn, y + i + k * 4 * lanes));
80 for (; i <= n - 4 * lanes; i += 4 * lanes) sme_madd_za_vg1x4(0, ploadu_x4(pn, x + i), ploadu_x4(pn, y + i));
81 auto tail_accumulator = negative_zero;
82 for (; i < n; i += n - i < lanes ? n - i : lanes) {
83 auto tail = Traits::whilelt(i, n);
84 tail_accumulator = pmadd_m(tail, ploadu(tail, x + i), ploadu(tail, y + i), tail_accumulator);
86 auto sum = tail_accumulator;
87 EIGEN_SME_VECTOR_UNROLL4
88 for (
int k = 0; k < 4; ++k) {
89 auto v = sme_read_za_vg1x4<Scalar>(k);
90 sum = padd(pg, sum, padd(pg, padd(pg, pget<0>(v), pget<1>(v)), padd(pg, pget<2>(v), pget<3>(v))));
92 return predux(pg, sum);
96template <
typename Scalar,
typename Vec>
97static EIGEN_ALWAYS_INLINE
void sme_gemv_update(
unsigned int k, svcount_t active, Scalar* y,
98 Vec alpha) __arm_streaming __arm_inout(
"za") {
99 const auto sum = sme_read_za_vg1x4<Scalar>(k);
100 sme_write_za_vg1x4(k, ploadu_x4(active, y));
101 sme_madd_za_vg1x4(k, sum, alpha);
102 pstoreu_x4(active, y, sme_read_za_vg1x4<Scalar>(k));
105template <
typename Scalar,
typename Index>
106__arm_new(
"za") __arm_locally_streaming
107 EIGEN_DONT_INLINE
void sme_gemv(Index rows, Index cols, const Scalar* a, Index stride, const Scalar* x, Scalar* y,
108 Scalar alpha, Index block_cols) {
109 using Traits = sme_packet_traits<Scalar>;
110 using Vec =
typename Traits::type;
111 const Index lanes = Traits::size();
112 const auto scale = pset1<Vec>(alpha);
114 for (Index first = 0; first < cols;) {
115 const Index end = first + (cols - first < block_cols ? cols - first : block_cols);
118 for (Index i = 0; i < rows;) {
119 const Index remaining = rows - i;
120 const auto p0 = Traits::whilelt_c4(i, rows), p1 = Traits::whilelt_c4(i + std::int64_t(4) * lanes, rows),
121 p2 = Traits::whilelt_c4(i + std::int64_t(8) * lanes, rows),
122 p3 = Traits::whilelt_c4(i + std::int64_t(12) * lanes, rows);
123 const Index o1 = remaining > 4 * lanes ? 4 * lanes : 0, o2 = remaining > 8 * lanes ? 8 * lanes : 0,
124 o3 = remaining > 12 * lanes ? 12 * lanes : 0;
126 for (Index j = first; j < end; ++j) {
127 const auto b = pset1<Vec>(x[j]);
128 const Scalar* column = a + i + j * stride;
129 sme_madd_za_vg1x4(0, ploadu_x4(p0, column), b);
130 sme_madd_za_vg1x4(1, ploadu_x4(p1, column + o1), b);
131 sme_madd_za_vg1x4(2, ploadu_x4(p2, column + o2), b);
132 sme_madd_za_vg1x4(3, ploadu_x4(p3, column + o3), b);
134 sme_gemv_update(0, p0, y + i, scale);
135 sme_gemv_update(1, p1, y + i + o1, scale);
136 sme_gemv_update(2, p2, y + i + o2, scale);
137 sme_gemv_update(3, p3, y + i + o3, scale);
139 i += remaining < 16 * lanes ? remaining : 16 * lanes;
145#undef EIGEN_SME_VECTOR_UNROLL4
147template <
typename Scalar>
148struct sme_vector_scalar : false_type {};
150struct sme_vector_scalar<float> : true_type {};
151#ifdef EIGEN_VECTORIZE_SME_F64F64
153struct sme_vector_scalar<double> : true_type {};
156template <
typename Xpr>
157struct sme_vector_access : bool_constant<(sme_vector_scalar<typename traits<Xpr>::Scalar>::value &&
158 bool(Xpr::IsVectorAtCompileTime) && has_direct_access<Xpr>::value)> {};
160template <
typename Lhs,
typename Rhs>
161struct sme_dot_supported : bool_constant<sme_vector_access<Lhs>::value && sme_vector_access<Rhs>::value &&
162 is_same<typename traits<Lhs>::Scalar, typename traits<Rhs>::Scalar>::value> {};
166static constexpr std::size_t kSmeDotMinBytes = 16384;
167static constexpr std::size_t kSmeAxpyMinBytes = 4096;
168static constexpr Index kSmeGemvMinRows = 128;
169static constexpr Index kSmeGemvMinCols = 4;
171static constexpr std::size_t kSmeGemvWideRhsBytes = 128;
172static constexpr std::size_t kSmeGemvWideMinBytes = 32768;
176template <
typename Scalar,
int L1Divisor = 1>
177EIGEN_DONT_INLINE
bool sme_vector_size_suitable(Index size) {
178 std::ptrdiff_t l1, l2, l3;
179 manage_caching_sizes(GetAction, &l1, &l2, &l3);
180 return l1 > 0 && l2 > 0 && size >= 0 &&
181 static_cast<std::size_t
>(size) >= (
static_cast<std::size_t
>(l1) - 1) / (L1Divisor *
sizeof(Scalar)) + 1 &&
182 static_cast<std::size_t
>(size) <=
static_cast<std::size_t
>(l2) /
sizeof(Scalar);
185template <
typename Lhs,
typename Rhs>
186struct default_inner_product_impl<Lhs, Rhs, true, std::enable_if_t<sme_dot_supported<Lhs, Rhs>::value>>
187 : default_inner_product_impl<Lhs, Rhs, true, false_type> {
188 using Base = default_inner_product_impl<Lhs, Rhs, true, false_type>;
189 using Scalar =
typename traits<Lhs>::Scalar;
190 static EIGEN_STRONG_INLINE Scalar run(
const MatrixBase<Lhs>& lhs,
const MatrixBase<Rhs>& rhs) {
191 inner_product_assert<Lhs, Rhs>::run(lhs.derived(), rhs.derived());
193 if (lhs.size() >= Index(kSmeDotMinBytes /
sizeof(Scalar)) && sme_vector_size_suitable<Scalar, 2>(lhs.size()) &&
194 lhs.innerStride() == 1 && rhs.innerStride() == 1) {
195 sme_fpsr_guard status;
196 const Scalar result = sme_dot(lhs.size(), lhs.derived().data(), rhs.derived().data());
198 if (result != Scalar(0) || !sme_rounds_toward_negative())
return result;
200 return Base::run(lhs, rhs);
204template <
typename Dst,
typename Src>
205EIGEN_STRONG_INLINE
bool sme_try_axpy(Dst& dst,
const Src& src,
typename traits<Src>::Scalar alpha) {
206 using Scalar =
typename traits<Src>::Scalar;
207 eigen_assert(dst.rows() == src.rows() && dst.cols() == src.cols());
208 if (dst.size() < Index(kSmeAxpyMinBytes /
sizeof(Scalar)) || !sme_vector_size_suitable<Scalar>(dst.size()) ||
209 dst.innerStride() != 1 || src.innerStride() != 1)
211 const std::uintptr_t dst_address =
reinterpret_cast<std::uintptr_t
>(dst.data());
212 const std::uintptr_t src_address =
reinterpret_cast<std::uintptr_t
>(src.data());
213 const std::uintptr_t distance = dst_address > src_address ? dst_address - src_address : src_address - dst_address;
215 if (distance != 0 && distance /
sizeof(Scalar) <
static_cast<std::uintptr_t
>(dst.size()))
return false;
217 sme_fpsr_guard status;
218 sme_axpy(dst.size(), src.data(), dst.data(), alpha, first_aligned<64>(dst.data(), dst.size()));
222template <
typename Dst,
typename Scalar,
typename Lhs,
typename Rhs>
224 Dst, CwiseBinaryOp<scalar_product_op<Scalar, Scalar>, Lhs, Rhs>, add_assign_op<Scalar, Scalar>, Dense2Dense,
225 std::enable_if_t<sme_vector_access<Dst>::value &&
226 (sme_vector_access<Lhs>::value || sme_vector_access<Rhs>::value) &&
227 blas_traits<CwiseBinaryOp<scalar_product_op<Scalar, Scalar>, Lhs, Rhs>>::HasScalarFactor &&
228 is_same<Scalar, typename traits<Dst>::Scalar>::value>> {
229 using Source = CwiseBinaryOp<scalar_product_op<Scalar, Scalar>, Lhs, Rhs>;
230 using BlasTraits = blas_traits<Source>;
231 static EIGEN_STRONG_INLINE
void run(Dst& dst,
const Source& src,
const add_assign_op<Scalar, Scalar>& func) {
233 if (!sme_try_axpy(dst, BlasTraits::extract(src), BlasTraits::extractScalarFactor(src)))
234 Assignment<Dst, Source, add_assign_op<Scalar, Scalar>, Dense2Dense, false_type>::run(dst, src, func);
238#ifndef EIGEN_USE_BLAS
239template <
typename Scalar,
typename Index>
240EIGEN_STRONG_INLINE
bool sme_gemv_size_suitable(Index rows, Index cols) {
241 if (rows < Index(kSmeGemvMinRows) || cols < Index(kSmeGemvMinCols))
return false;
242 if (cols >= Index(kSmeGemvWideRhsBytes /
sizeof(Scalar)) &&
243 cols > Index(kSmeGemvWideMinBytes /
sizeof(Scalar) - 1) / rows)
246 const std::ptrdiff_t l1 = l1CacheSize();
247 return l1 > 0 &&
static_cast<std::size_t
>(cols) >
248 (
static_cast<std::size_t
>(l1) - 1) /
sizeof(Scalar) /
static_cast<std::size_t
>(rows);
251#define EIGEN_SME_GEMV_SPECIALIZATION(Scalar) \
252 template <typename Index, bool ConjugateLhs, bool ConjugateRhs> \
253 struct general_matrix_vector_product<Index, Scalar, const_blas_data_mapper<Scalar, Index, ColMajor>, ColMajor, \
254 ConjugateLhs, Scalar, const_blas_data_mapper<Scalar, Index, RowMajor>, \
255 ConjugateRhs, Specialized> { \
256 static EIGEN_STRONG_INLINE void run(Index rows, Index cols, \
257 const const_blas_data_mapper<Scalar, Index, ColMajor>& lhs, \
258 const const_blas_data_mapper<Scalar, Index, RowMajor>& rhs, Scalar* res, \
259 Index resIncr, Scalar alpha) { \
260 if (sme_gemv_size_suitable<Scalar>(rows, cols) && rhs.stride() == 1 && resIncr == 1) { \
261 const std::ptrdiff_t l1 = l1CacheSize(); \
262 const Index block_cols = cols < 128 ? cols \
263 : (l1 > 0 && static_cast<std::size_t>(lhs.stride()) <= \
264 (static_cast<std::size_t>(l1) - 1) / sizeof(Scalar) \
268 if (alpha == Scalar(0)) return; \
269 sme_fpsr_guard status; \
270 sme_gemv(rows, cols, lhs.data(), lhs.stride(), rhs.data(), res, alpha, block_cols); \
272 general_matrix_vector_product<Index, Scalar, const_blas_data_mapper<Scalar, Index, ColMajor>, ColMajor, \
273 ConjugateLhs, Scalar, const_blas_data_mapper<Scalar, Index, RowMajor>, \
274 ConjugateRhs, BuiltIn>::run(rows, cols, lhs, rhs, res, resIncr, alpha); \
278EIGEN_SME_GEMV_SPECIALIZATION(
float)
279#ifdef EIGEN_VECTORIZE_SME_F64F64
280EIGEN_SME_GEMV_SPECIALIZATION(
double)
282#undef EIGEN_SME_GEMV_SPECIALIZATION