Eigen  5.0.1
 
Loading...
Searching...
No Matches
VectorKernels.h
1// SPDX-FileCopyrightText: The Eigen Authors
2// SPDX-License-Identifier: MPL-2.0
3#ifndef EIGEN_SME_VECTOR_KERNELS_H
4#define EIGEN_SME_VECTOR_KERNELS_H
5
6// IWYU pragma: private
7#include "../../InternalHeaderCheck.h"
8
9namespace Eigen {
10namespace internal {
11
12#if EIGEN_COMP_CLANG
13#define EIGEN_SME_VECTOR_UNROLL4 _Pragma("unroll")
14#elif EIGEN_COMP_GNUC
15#define EIGEN_SME_VECTOR_UNROLL4 _Pragma("GCC unroll 4")
16#else
17#define EIGEN_SME_VECTOR_UNROLL4
18#endif
19
20// FPCR.RMode == 0b10 (roundTowardNegative), the one mode where (+0) + (-0) is -0.
21static EIGEN_ALWAYS_INLINE bool sme_rounds_toward_negative() {
22 std::uint64_t fpcr;
23 asm volatile("mrs %0, fpcr" : "=r"(fpcr));
24 return ((fpcr >> 22) & 3) == 2;
25}
26
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);
35 Index i = 0;
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)));
39 }
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);
46 }
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));
49 }
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));
55 }
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)));
59 }
60}
61
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();
69 // -0 is the additive identity except under roundTowardNegative, so a zero sum gets the IEEE sign.
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));
74 Index i = 0;
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));
79 }
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);
85 }
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))));
91 }
92 return predux(pg, sum);
93}
94
95// y += alpha * ZA group k, updated in ZA like the accumulation.
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));
103}
104
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);
113 // Match the generic GEMV's scaled column batches to bound the unscaled sums.
114 for (Index first = 0; first < cols;) {
115 const Index end = first + (cols - first < block_cols ? cols - first : block_cols);
116 // Four predicated chunks per row block keep four independent FMLA chains, also in the final partial block.
117 // An inactive chunk addresses the block start and accesses no memory.
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;
125 svzero_za();
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);
133 }
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);
138 // Bound the final increment to avoid signed overflow with a 32-bit Index near its maximum.
139 i += remaining < 16 * lanes ? remaining : 16 * lanes;
140 }
141 first = end;
142 }
143}
144
145#undef EIGEN_SME_VECTOR_UNROLL4
146
147template <typename Scalar>
148struct sme_vector_scalar : false_type {};
149template <>
150struct sme_vector_scalar<float> : true_type {};
151#ifdef EIGEN_VECTORIZE_SME_F64F64
152template <>
153struct sme_vector_scalar<double> : true_type {};
154#endif
155
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)> {};
159
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> {};
163
164// Fixed crossovers, fitted on Apple M4 (SVL 512): below them the NEON/SME memory handoff outweighs the
165// streaming kernels. Core's cache estimates, including setCpuCacheSizes overrides, supply the others.
166static constexpr std::size_t kSmeDotMinBytes = 16384; // per operand
167static constexpr std::size_t kSmeAxpyMinBytes = 4096; // per operand
168static constexpr Index kSmeGemvMinRows = 128;
169static constexpr Index kSmeGemvMinCols = 4;
170// A GEMV whose RHS reaches kSmeGemvWideRhsBytes needs only kSmeGemvWideMinBytes of matrix, not Core's L1 estimate.
171static constexpr std::size_t kSmeGemvWideRhsBytes = 128;
172static constexpr std::size_t kSmeGemvWideMinBytes = 32768;
173
174// Keep cache queries out of callers so the small vector fallback can inline.
175// Per operand, the L1 estimate (divided by L1Divisor) is the lower crossover and the L2 estimate the upper one.
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);
183}
184
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());
192 // DOT has no streaming stores for a following NEON consumer: half L1 per operand suffices.
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());
197 // Under roundTowardNegative the -0 seeds turn +0 sums into -0.
198 if (result != Scalar(0) || !sme_rounds_toward_negative()) return result;
199 }
200 return Base::run(lhs, rhs);
201 }
202};
203
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)
210 return false;
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;
214 // Exact aliasing is coefficient-wise; partial overlap must retain the default traversal.
215 if (distance != 0 && distance / sizeof(Scalar) < static_cast<std::uintptr_t>(dst.size())) return false;
216 // Align streaming stores to 64 bytes without changing Eigen's allocation alignment.
217 sme_fpsr_guard status;
218 sme_axpy(dst.size(), src.data(), dst.data(), alpha, first_aligned<64>(dst.data(), dst.size()));
219 return true;
220}
221
222template <typename Dst, typename Scalar, typename Lhs, typename Rhs>
223struct Assignment<
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) {
232 // A direct operand restricts extraction to one scale, preserving nested products' evaluation order.
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);
235 }
236};
237
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)
244 return true;
245 // Thin products amortize the NEON/SME handoff once the matrix reaches Core's L1 estimate.
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);
249}
250
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) \
265 ? Index(16) \
266 : Index(4)); \
267 /* BLAS contract: alpha == 0 leaves the result unchanged. */ \
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); \
271 } else { \
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); \
275 } \
276 } \
277 };
278EIGEN_SME_GEMV_SPECIALIZATION(float)
279#ifdef EIGEN_VECTORIZE_SME_F64F64
280EIGEN_SME_GEMV_SPECIALIZATION(double)
281#endif
282#undef EIGEN_SME_GEMV_SPECIALIZATION
283#endif
284
285} // namespace internal
286} // namespace Eigen
287
288#endif