Eigen  5.0.1
 
Loading...
Searching...
No Matches
GeneralBlockPanelKernel.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// This Source Code Form is subject to the terms of the Mozilla
5// Public License v. 2.0. If a copy of the MPL was not distributed
6// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
7// SPDX-FileCopyrightText: The Eigen Authors
8// SPDX-License-Identifier: MPL-2.0
9
10#ifndef EIGEN_SME_GENERALBLOCKPANELKERNEL_H
11#define EIGEN_SME_GENERALBLOCKPANELKERNEL_H
12
13// IWYU pragma: private
14#include "../../InternalHeaderCheck.h"
15
16#include <arm_sme.h>
17
18namespace Eigen {
19namespace internal {
20
21// ---------------------------------------------------------------------------
22// Streaming vector length and tile geometry.
23//
24// The micro-kernel is organised around a logical mr x nr output block, packed
25// depth-major (mr contiguous scalars per depth step). Those dimensions are
26// compile-time constants: they feed gebp_traits (cache blocking) and the
27// packers, none of which can depend on a runtime value.
28//
29// The *physical* tiling of that block onto ZA tiles, on the other hand, is
30// driven by the runtime streaming vector length. A ZA tile of Scalar is
31// svl x svl, where svl is the number of Scalars in a streaming vector
32// (svcntsw() for fp32, svcntsd() for fp64). The block is covered by up to a
33// 2x2 grid of svl x svl tiles, iterated in sub-block passes when the grid is
34// smaller than the block (and predicated down to it when larger).
35//
36// fp32 uses the 4 ZA.S tiles, so the 2x2 grid is all of ZA. fp64 uses ZA.D, of
37// which there are 8, and deliberately leaves tiles 4-7 idle: a 2x2 grid loads 2
38// packed vectors per side per depth step to feed 4 FMOPAs, i.e. 64 bytes of
39// packed panel per FMOPA at either element width, and FMOPA issues at the same
40// rate for both. A 2x4 grid over all eight needs a quarter less panel traffic
41// per FMOPA and still measures 0.92-1.00x of the 2x2 on Apple M4, so the wider
42// block is not worth its L1 footprint.
43//
44// A complex accumulator is a *pair* of tiles holding its real and imaginary
45// halves, since FMOPA only takes real operands, so complex<float> gets a 1x2
46// grid of pairs out of ZA.S and complex<double> a 2x2 grid out of ZA.D. The
47// packed panels are correspondingly split -- one depth step holds the real
48// parts of the panel width then the imaginary ones -- so the four real outer
49// products a complex one expands to reuse both operands and the panel traffic
50// per FMOPA halves relative to the real kernels.
51//
52// This translation unit must be built without -msve-vector-bits (scalable/VLA
53// mode); see the guard in ConfigureVectorization.h for the rationale.
54// Everything below derives lane counts/predicates from the runtime svl; when a
55// block matches the tile grid exactly, the micro-kernel additionally switches
56// to a hand-scheduled multi-vector-load loop (see sme_process).
57// ---------------------------------------------------------------------------
58
59// The streaming vector width the block sizes are chosen at: 512 bits, where a
60// vector holds 64 / sizeof(element) elements. Other SVLs tile the block at
61// runtime. If a future SVL ever justifies a larger block, this is the only
62// knob -- but don't grow it speculatively, a doubled block measures slower at
63// SVL=512.
64static constexpr int kSmeDesignVectorBytes = 64;
65
66// Logical micro-kernel block (LHS/RHS panel widths), as a grid of ZA tiles each
67// svl x svl elements of the ZA element width -- Scalar itself for a real
68// scalar, its real part for a complex one, whose halves accumulate into
69// separate tiles -- a full 2x2 grid for real scalars.
70// A complex accumulator takes a tile pair, so complex<float> gets half the grid
71// cells of float; the grid stays two cells wide, which measures 1.2-1.9x a
72// two-cell-tall one below 128 on Apple M4 and matches it above.
73template <typename Scalar>
74struct sme_block {
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));
79};
80
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));
87};
88
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;
98#endif
99
100// ---------------------------------------------------------------------------
101// Packed panel layout and the primitives that produce it.
102//
103// A packed panel is depth-major. For a real element type one depth step holds
104// `width` contiguous scalars; for a complex one it holds the `width` real parts
105// followed by the `width` imaginary parts, so the micro-kernel feeds a ZA tile
106// pair from two contiguous vector loads and never deinterleaves inside the
107// depth loop. Either way a width x depth panel occupies width * depth Scalars,
108// which is what the GEMM driver allocates.
109//
110// `Conjugate` negates the imaginary half. Conjugation is the identity on real
111// scalars, so the real overloads ignore it -- and Conjugate=true instantiations
112// do reach them, from the SYMM above-diagonal transposed pack.
113// ---------------------------------------------------------------------------
114
115// Copy `width` contiguous source columns per depth step into a depth-major
116// packed panel of width `width`, for the depth sub-range [k0, k1). Both dst and
117// src are indexed by the absolute depth index k (dst[k*width+off],
118// src[k*src_stride+off]); the caller offsets `src` to the region's column base
119// and `dst` to the panel base. Generalised over the runtime svl: the panel is
120// covered in svl-wide column chunks, each streamed over the depth sub-range.
121// The chunk loop is outermost so each chunk's predicate is computed once instead
122// of per depth step (the runtime chunk count keeps the compiler from hoisting it
123// on its own). The symm packers reuse this for the diagonal-split direct/
124// transposed regions (a contiguous depth sub-range at a depth offset).
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) {
132 // Full panel: one two-vector load and store per depth step, four steps in flight.
133 const svcount_t pn = Traits::ptrue_c();
134 Index k = k0;
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);
144 }
145 for (; k < k1; ++k) pstoreu_x2(pn, &dst[k * width], ploadu_x2(pn, &src[k * src_stride]));
146 return;
147 }
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]));
152 }
153 }
154}
155
156// Four full 2*svl-wide panels from one 8*svl-wide source strip (panel p at dst + p * panel_stride): two four-vector
157// loads per depth step, and the core prefetches the strip into L2 four steps ahead.
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;
168 if (k + 4 < depth)
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)));
176 }
177}
178
179// Packs the leading groups of four full MR-wide panels with sve_copy_panel_quad; returns the rows it covered.
180// Complex panels keep the per-panel path (their packed layout splits real and imaginary halves).
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,
184 Index dst_stride,
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;
188 Index i = 0;
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);
192 }
193 return i;
194}
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 {
199 return 0;
200}
201
202// Complex overload: UZP1/UZP2 split the source's interleaved pairs into the two
203// halves of the packed depth step. A chunk of w complex elements spans 2*w
204// interleaved reals, hence a pair of source predicates; the upper one is empty
205// whenever 2*w fits in one vector.
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);
226 // pg_hi is all-false when 2*w fits in one vector, and an inactive lane
227 // makes no memory access -- but p + svl may still be past the source, so
228 // the address is formed through sme_offset.
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);
233 }
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);
236 }
237 }
238}
239
240// Copy the full depth [0, depth): thin wrapper used by the (non-symm) gemm
241// packers, which always pack a whole panel.
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);
247}
248
249// Full 2*svl-wide panel through all four tiles, four rows or depth steps per ZA move, source prefetched into L2
250// two steps ahead; returns the first depth index of the tail (< 2*svl) it leaves to the caller.
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();
258 Index k = k0;
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;
264 if (pf)
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);
268 }
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));
277 }
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)));
293 }
294 }
295 return k;
296}
297
298// A row of a partial panel, or zeros past its last row (those slices are never stored).
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);
304}
305
306// Partial panel (width < 2*svl) through the tiles as sme_transpose_pack_pair does, rows past the width as zeros and
307// predicated stores of `width` scalars per depth step; returns the first depth index of the tail it leaves.
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,
311 Index k0, Index k1,
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);
321 Index k = k0;
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));
331 }
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));
340 }
341 if (2 * width <= svl) {
342 // Two depth steps per store: half as many stores, and full lines at width svl/2.
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)));
352 }
353 continue;
354 }
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));
368 if (rows_hi > 0) {
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));
379 }
380 }
381 }
382 return k;
383}
384
385// Transpose-pack `width` source rows into depth-major packed output using ZA's
386// 2D store as a free transpose, for the depth sub-range [k0, k1): a svl x svl
387// block of source (svl rows x svl depth) is loaded as horizontal ZA slices,
388// then read back as vertical slices, which emits it depth-major. Row-groups of
389// svl rows are processed two at a time through ZA tiles 0 and 1: ZA is not
390// renamed, so a single tile would stall every load pass on the previous read
391// pass (write-after-read); two tiles in flight keep the phases independent.
392// Trailing row-groups (when width is not a multiple of 2*svl) use tile 0 with
393// predicated rows -- which is also what a panel narrower than 2*svl gets in
394// full: complex<float> has mr = svl at SVL=512, so its LHS panels take the
395// single-tile path and do not get the write-after-read overlap described above.
396// Widening the gate would need the pairing to run over depth instead of rows.
397// Both dst and src are indexed by the absolute depth index k:
398// dst[k*width + r] = src[r*src_stride + k], k in [k0,k1), r in [0,width).
399// The symm packers reuse this for the diagonal-split transposed/direct regions
400// (a depth sub-range at a depth offset, with a tail-panel width < mr).
401//
402// NegateOddRows negates every odd output depth row. That is how the complex
403// overload below conjugates: in the real view of a complex panel those rows are
404// exactly the imaginary halves.
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,
408 Index k0, Index k1,
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();
415
416 Index k = k0;
417 EIGEN_IF_CONSTEXPR (!NegateOddRows) {
418 // Short ranges keep the single-tile path: the four-slice moves only pay off over several fills, and need a tile
419 // of at least four slices.
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);
424 }
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);
428 int r0 = 0;
429 // Pairs of full row-groups: tiles 0 and 1 in flight.
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]);
434 }
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);
442 }
443 }
444 pstoreu(pg_all, &dst[(k + c) * width + r0], v0);
445 pstoreu(pg_all, &dst[(k + c) * width + r0 + svl], v1);
446 }
447 }
448 // Trailing row-groups (at most two svl-wide passes remain, since the pair
449 // loop consumed all multiples of 2*svl): predicate down to the remaining
450 // rows. A single `if` would drop rows when a tail width lands in
451 // (svl, 2*svl); a loop handles any leftover.
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]);
457 }
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);
462 }
463 pstoreu(pg_r, &dst[(k + c) * width + r0], v0);
464 }
465 }
466 }
467}
468
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);
474}
475
476// Complex overload. A ColMajor (RowMajor) complex operand is a ColMajor
477// (RowMajor) real one of twice the depth and twice the stride, and transposing
478// that real view already emits the split layout: real-view depth 2k lands at
479// packed offset k*(2*width) and depth 2k+1 at k*(2*width) + width, the real and
480// imaginary halves of packed depth step k.
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);
488}
489
490// Transpose-pack a whole `width`-wide panel over the full depth [0, depth):
491// thin wrapper used by the (non-symm) gemm packers.
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);
497}
498
499// Transposing copy for a panel narrower than the pack width:
500// dst_panel[k*tail + i] = src[i*src_stride + k].
501// Kept outside the caller's __arm_locally_streaming region: it needs neither SVE
502// nor ZA, and streaming mode runs scalar floating-point ~40x slower on Apple M4.
503// Outside it the source rows are contiguous in k, so PacketSize of them
504// transpose in register as in sme_pack_rhs_fallback; a product with cols < nr is
505// packed entirely here. NegateOddRows is as in sme_transpose_pack_real.
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));
513
514 Index i = 0;
515 for (; i < peeled_tail; i += Index(PacketSize)) {
516 Index k = 0;
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);
521 }
522 ptranspose(block);
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);
527 }
528 pstoreu(dst_panel + (k + Index(p)) * tail + i, row);
529 }
530 }
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;
536 }
537 }
538 }
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;
544 }
545 }
546}
547
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);
552}
553
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,
557 Index tail) {
558 tail_transpose_pack_real<Conjugate>(reinterpret_cast<RealScalar*>(dst_panel),
559 reinterpret_cast<const RealScalar*>(src), Index(2) * src_stride, Index(2) * depth,
560 tail);
561}
562
563// Tail panel of a deep block through the ZA transposer (sme_transpose_pack_partial); shallow ones and those of at
564// most 4 columns keep the NEON tail_transpose_pack, which is faster there.
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);
569}
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,
572 Index tail) {
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));
576 return;
577 }
578 }
579 tail_transpose_pack<Conjugate>(dst_panel, src, src_stride, depth, tail);
580}
581
582// De-interleaving load and interleaving store of PS complex values.
583static EIGEN_ALWAYS_INLINE void sme_neon_ld2(const float* p, Packet4f& re, Packet4f& im) {
584 const float32x4x2_t v = vld2q_f32(p);
585 re = v.val[0];
586 im = v.val[1];
587}
588static EIGEN_ALWAYS_INLINE void sme_neon_st2(float* p, Packet4f re, Packet4f im) {
589 float32x4x2_t v;
590 v.val[0] = re;
591 v.val[1] = im;
592 vst2q_f32(p, v);
593}
594static EIGEN_ALWAYS_INLINE void sme_neon_ld2(const double* p, Packet2d& re, Packet2d& im) {
595 const float64x2x2_t v = vld2q_f64(p);
596 re = v.val[0];
597 im = v.val[1];
598}
599static EIGEN_ALWAYS_INLINE void sme_neon_st2(double* p, Packet2d re, Packet2d im) {
600 float64x2x2_t v;
601 v.val[0] = re;
602 v.val[1] = im;
603 vst2q_f64(p, v);
604}
605
606// NEON copy packers for the direct-access, unit-stride case: dst[k*w + r] =
607// src[r + k*src_stride], w contiguous scalars per depth step; complex panels
608// split into the depth step's real then imaginary halves, conjugated if asked.
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,
611 Index w) {
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;
619 Index r = 0;
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);
623 pstoreu(d + r, p0);
624 pstoreu(d + r + PS, p1);
625 pstoreu(d + r + 2 * PS, p2);
626 pstoreu(d + r + 3 * PS, p3);
627 }
628 for (; r < peeled; r += PS) pstoreu(d + r, ploadu<Packet>(s + r));
629 for (; r < w; ++r) d[r] = s[r];
630 }
631}
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,
635 Index w) {
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;
644 Index r = 0;
645 for (; r < peeled; r += PS) {
646 Packet re, im;
647 sme_neon_ld2(s + 2 * r, re, im);
648 pstoreu(d + r, re);
649 pstoreu(d + w + r, Conjugate ? pnegate(im) : im);
650 }
651 for (; r < w; ++r) {
652 d[r] = s[2 * r];
653 d[w + r] = Conjugate ? -s[2 * r + 1] : s[2 * r + 1];
654 }
655 }
656}
657
658// ---------------------------------------------------------------------------
659// Generic (mapper-based) packing fallback.
660//
661// The streaming pack_lhs_*/pack_rhs_* helpers take &lhs(0,0) once and walk it by
662// raw pointer + lhs.stride(). That breaks for two DataMapper families:
663// - TensorContractionSubMapper::operator() returns by value, so &lhs(0,0) is
664// address-of-rvalue (a compile error, not just wrong results);
665// - blas_data_mapper with Incr != 1 (inner-strided Maps, e.g. from
666// TriangularSolverMatrix) can't be walked by stride() alone.
667// These fall back to the mapper's packet/element interface, emitting the
668// identical depth-major panel layout so gebp_kernel can't tell the paths apart.
669// ---------------------------------------------------------------------------
670
671// True iff DataMapper exposes .incr() (the blas_data_mapper family); others are
672// unit-inner-stride by construction.
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 {};
677
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());
681}
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&) {
684 return Index(1);
685}
686
687// Whether operator()(i,j) returns an lvalue reference into caller storage (so
688// &m(0,0) + stride walking is valid). False for by-value mappers (Tensor's).
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;
693};
694
695// Store one element of a packed depth step of width w. Real scalars land at
696// dst_step[r]; complex ones split into the step's real and imaginary halves,
697// which is the same base pointer reinterpreted, since a complex depth step of
698// width w spans 2*w reals.
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);
703}
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);
711}
712
713// LHS fallback: pack via the mapper's packet interface, shared by both
714// gemm_pack_lhs specializations. Taken by mappers without direct lvalue access
715// (TensorContractionSubMapper returns by value) or with a non-unit inner
716// stride. Vectorised with NEON packets, exactly like the generic packers drive
717// these same mappers. Tensor sub-mappers (the hot path -- tensor contractions
718// pack through this on both sides) have contiguous packet loads, but their
719// ordinary operator()/loadPacket functions cannot be called from a streaming
720// context. Inner-strided ColMajor blas mappers instead require gathers;
721// streaming-mode gathers need FEAT_SME_FA64 (absent on e.g. Apple M4), while
722// NEON's pgather uses scalar source loads and a contiguous packet store. The
723// packet path assumes the mapper's packets advance the first index; that holds
724// for ColMajor tensor and blas mappers, but not for RowMajor mappers, whose
725// packets run along the storage-inner second index. RowMajor dispatches pass
726// vectorise = false and take the scalar element loop. Complex scalars always
727// take it too: a complex packet store would emit the interleaved layout, not the
728// split one the kernel reads.
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;
735
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;
742 Index r = 0;
743 for (; r < peeled_w; r += PacketSize) {
744 pstoreu(dst_step + r, lhs.template loadPacket<Packet>(i + r, k));
745 }
746 for (; r < w; ++r) {
747 sme_pack_store<Conjugate>(dst_step, w, r, lhs(i + r, k));
748 }
749 }
750 }
751}
752
753// The PacketSize column sub-mappers one packed column group loads from.
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))...}};
758}
759
760// RHS fallback, mirroring sme_pack_lhs_fallback (including the vectorise
761// contract: LinearMapper packets must advance the first (depth) index). The
762// packed layout wants consecutive columns contiguous while the mapper's
763// packets run along the depth k, so PacketSize columns are loaded as packets
764// along k and transposed in-register (the same LinearMapper + ptranspose
765// scheme as the generic gemm_pack_rhs).
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);
774
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);
779 Index c = 0;
780 for (; c < peeled_w; c += Index(PacketSize)) {
781 // Loop-invariant in k, but not hoisted out of the k loop by the compiler
782 // for a mapper that returns its sub-mappers by value -- which is the hot
783 // path here: tensor contractions pack through TensorContractionSubMapper.
784 const std::array<LinearMapper, PacketSize> dm =
785 sme_column_mappers(rhs, j + c, std::make_index_sequence<PacketSize>{});
786 Index k = 0;
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);
791 }
792 ptranspose(block);
793 for (int p = 0; p < PacketSize; ++p) {
794 pstoreu(dst_panel + (k + Index(p)) * w + c, block.packet[p]);
795 }
796 }
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));
800 }
801 }
802 }
803 for (; c < w; ++c) {
804 for (Index k = 0; k < depth; ++k) {
805 sme_pack_store<Conjugate>(dst_panel + k * w, w, c, rhs(k, j + c));
806 }
807 }
808 }
809}
810
811// NEON dispatch invariant, constants fitted on Apple M4 (SVL 512). Packers: panel depth (the stride in
812// panel mode) <= sme_neon_max_depth && width <= sme_neon_max_panel. Kernel: both sides pass the packer
813// test with strideA / strideB as the panel depths && (max(rows, cols) <= w(depth) || min <= thin_dim).
814// So a NEON kernel reads NEON-packed panels, and ZA over NEON-packed panels (~40 ns/KB) is bounded by
815// sme_neon_max_panel; the exception is SYRK, which packs B at the full size but runs 32x32 diagonal blocks.
816#ifndef EIGEN_SME_NEON_MAX_DEPTH
817template <typename Scalar>
818struct sme_neon_max_depth : std::integral_constant<int, 16> {};
819template <>
820struct sme_neon_max_depth<float> : std::integral_constant<int, 24> {};
821template <>
822struct sme_neon_max_depth<std::complex<double>> : std::integral_constant<int, 8> {};
823#else
824template <typename Scalar>
825struct sme_neon_max_depth : std::integral_constant<int, EIGEN_SME_NEON_MAX_DEPTH> {};
826#endif
827#ifndef EIGEN_SME_NEON_MAX_WIDTH
828template <typename Scalar>
829struct sme_neon_max_width : std::integral_constant<int, 48> {};
830template <>
831struct sme_neon_max_width<double> : std::integral_constant<int, 40> {};
832template <>
833struct sme_neon_max_width<std::complex<float>> : std::integral_constant<int, 16> {};
834template <>
835struct sme_neon_max_width<std::complex<double>> : std::integral_constant<int, 24> {};
836#else
837template <typename Scalar>
838struct sme_neon_max_width : std::integral_constant<int, EIGEN_SME_NEON_MAX_WIDTH> {};
839#endif
840// Up to sme_neon_shallow_depth the NEON range widens to sme_neon_shallow_width:
841// the blocked decompositions update sub-blocks of that depth in place, at
842// offsets the ZA slice stores handle poorly.
843#ifndef EIGEN_SME_NEON_SHALLOW_DEPTH
844template <typename Scalar>
845struct sme_neon_shallow_depth : std::integral_constant<int, 8> {};
846#else
847template <typename Scalar>
848struct sme_neon_shallow_depth : std::integral_constant<int, EIGEN_SME_NEON_SHALLOW_DEPTH> {};
849#endif
850#ifndef EIGEN_SME_NEON_SHALLOW_WIDTH
851template <typename Scalar>
852struct sme_neon_shallow_width : std::integral_constant<int, 96> {};
853template <>
854struct sme_neon_shallow_width<double> : std::integral_constant<int, 64> {};
855template <>
856struct sme_neon_shallow_width<std::complex<float>> : std::integral_constant<int, 32> {};
857template <>
858struct sme_neon_shallow_width<std::complex<double>> : std::integral_constant<int, 24> {};
859#else
860template <typename Scalar>
861struct sme_neon_shallow_width : std::integral_constant<int, EIGEN_SME_NEON_SHALLOW_WIDTH> {};
862#endif
863// No panel wider than this is packed with NEON at any depth: the triangular solvers slice a tall
864// panel into depth-8 pieces for ZA. The packers cannot mirror the kernel's width rule, or a thin
865// block's wide side would be streaming-packed and read by NEON, the costly direction; the price is
866// a shallow panel in (w, this] read by ZA, 1.4x the streaming-packed time at 128x128x16 float.
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> {};
872#else
873template <typename Scalar>
874struct sme_neon_max_panel : std::integral_constant<int, EIGEN_SME_NEON_MAX_PANEL> {};
875#endif
876// A block this narrow on one side runs on NEON whatever its other side: the
877// blocked decompositions update long strips of that width in place.
878#ifndef EIGEN_SME_NEON_THIN_DIM
879template <typename Scalar>
880struct sme_neon_thin_dim : std::integral_constant<int, 32> {};
881template <>
882struct sme_neon_thin_dim<std::complex<float>> : std::integral_constant<int, 8> {};
883template <>
884struct sme_neon_thin_dim<std::complex<double>> : std::integral_constant<int, 4> {};
885#else
886template <typename Scalar>
887struct sme_neon_thin_dim : std::integral_constant<int, EIGEN_SME_NEON_THIN_DIM> {};
888#endif
889
890// `width` is the panel's rows (LHS) or cols (RHS).
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);
896 return false;
897#elif defined(EIGEN_SME_FORCE_NEON_SMALL_BLOCKS)
898 EIGEN_UNUSED_VARIABLE(depth);
899 EIGEN_UNUSED_VARIABLE(width);
900 return true;
901#else
902 return depth <= Index(sme_neon_max_depth<Scalar>::value) && width <= Index(sme_neon_max_panel<Scalar>::value);
903#endif
904}
905
906// strideA / strideB are the packed panels' depths, which exceed `depth` when
907// the caller packed them in panel mode and runs the kernel on a slice.
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);
913#else
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));
919#endif
920}
921
922// Shared dispatch for the four gemm_pack specializations: raw-pointer walk
923// when the mapper grants direct unit-inner-stride access, otherwise the
924// packet/element fallback. Tag-dispatched so &m(0,0) is only compiled for
925// lvalue mappers. UsePacketPath records whether the mapper's packets advance
926// the index the fallback needs, independently of its direct-access category.
927// In panel mode the call packs one depth slice of a panel `stride` deep that the
928// kernel consumes whole (the triangular solvers and TRMM), so the mode decision
929// uses the panel's depth, not the slice's.
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,
934 std::true_type /* direct access */) {
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);
939 } else {
940 sme_fpsr_guard fpsr;
941 direct(block, src, m.stride(), depth, n, stride, offset);
942 }
943 } else {
944 fallback(block, m, depth, n, stride, offset, UsePacketPath);
945 }
946}
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,
951 std::false_type /* no direct access */) {
952 fallback(block, m, depth, n, stride, offset, UsePacketPath);
953}
954
955/*****************************************************************************
956 * gebp_traits specializations for SME (float x float, double x double)
957 *
958 * Override mr and nr so that:
959 * - gemm_pack_lhs receives Pack1 = mr, creating uniform LHS panels
960 * - gemm_pack_rhs receives nr, creating uniform RHS panels
961 * - mc is rounded to a multiple of mr, nc to a multiple of nr
962 * - Cache blocking (kc, mc, nc) is recomputed accordingly
963 *
964 * We provide custom gemm_pack_lhs/gemm_pack_rhs specializations for both
965 * scalars, so both ColMajor and RowMajor source matrices produce an identical,
966 * simple packed format that the SME kernel consumes.
967 *
968 * Mixed-scalar products (e.g. MatrixXf * MatrixXcf) also instantiate
969 * gemm_pack_lhs<float, ...>, but with Pack1/nr from the generic
970 * gebp_traits<float, complex<float>> (mr=6, nr=4) and are consumed by the
971 * generic gebp_kernel, not the SME one. So the specializations below pin
972 * Pack1/nr_ to the SME block sizes: only the instantiation that feeds the SME
973 * gebp_kernel matches; mixed-scalar ones fall through to the generic template.
974 * This is load-bearing: it relies on no other consumer of the same scalar
975 * instantiating the packer with mr == the SME block size (holds today --
976 * generic float traits give mr <= 12). The kernel side is self-checking (the
977 * SME gebp_kernel static_asserts mr/nr against the block sizes, so a traits
978 * change breaks the build instead of silently mispairing packer and kernel);
979 * the packer side is enforced by the static_asserts below for the in-tree
980 * mixed-scalar traits (downstream code instantiating the packers with
981 * hand-picked mr/nr remains uncovered).
982 *****************************************************************************/
983
984template <>
985class gebp_traits<float, float, false, false, Architecture::Target, GEBPPacketFull>
986 : public gebp_traits<float, float, false, false, Architecture::Target, GEBPPacketHalf> {
987 public:
988 // The base class provides all the standard typedefs (LhsPacket, etc.)
989 // We only override the register-block sizes.
990 enum {
991 mr = kSmeMr, // LHS panel width
992 nr = kSmeNr // RHS panel width
993 };
994};
995
996// The packers do not know the opposite scalar type, so the SME block sizes are
997// effectively SME-format tags. Ensure the in-tree mixed-scalar traits cannot
998// select an SME packer whose output would be consumed by the generic kernel.
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");
1005
1006#ifdef EIGEN_VECTORIZE_SME_F64F64
1007template <>
1008class gebp_traits<double, double, false, false, Architecture::Target, GEBPPacketFull>
1009 : public gebp_traits<double, double, false, false, Architecture::Target, GEBPPacketHalf> {
1010 public:
1011 // As above, only the register-block sizes are overridden.
1012 static constexpr int mr = kSmeMrD;
1013 static constexpr int nr = kSmeNrD;
1014};
1015
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");
1022#endif
1023
1024// Complex block sizes, as above, but left open over the conjugation flags. A
1025// complex operand really does reach the kernel conjugated -- from an adjoint or
1026// conjugate product -- and the generic complex traits keep mr/nr independent of
1027// that, so pinning <false, false> here would hand a conjugated instantiation the
1028// generic block sizes while its kernel expects the SME ones.
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,
1032 GEBPPacketHalf> {
1033 public:
1034 static constexpr int mr = kSmeMrC;
1035 static constexpr int nr = kSmeNrC;
1036};
1037
1038// The mixed-scalar guard, with the roles of the two operands swapped relative
1039// to the real case: gemm_pack_lhs is instantiated with the LHS scalar and
1040// Traits::mr, gemm_pack_rhs with the RHS scalar and Traits::nr.
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");
1047
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,
1052 GEBPPacketHalf> {
1053 public:
1054 static constexpr int mr = kSmeMrCD;
1055 static constexpr int nr = kSmeNrCD;
1056};
1057
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");
1064#endif
1065
1066/*****************************************************************************
1067 * gemm_pack_lhs for SME (ColMajor source)
1068 *
1069 * Packs the LHS matrix into uniform panels of width mr.
1070 * Each depth step k writes exactly MR contiguous scalars.
1071 *****************************************************************************/
1072
1073template <typename Scalar, int MR, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1074struct sme_pack_lhs_colmajor {
1075 // Non-streaming NEON copy of every panel, for the shallow blocks the NEON
1076 // kernel consumes.
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);
1083 }
1084 }
1085
1086 // EIGEN_DONT_INLINE: GCC (14.1 through trunk) may inline a __arm_locally_streaming function into a
1087 // non-streaming caller and drop its mode switch. The other streaming entry points are __arm_new("za"),
1088 // which GCC does not inline.
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);
1095
1096 // Full panels of width MR, streamed in svl-wide predicated chunks.
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);
1100 }
1101
1102 // Tail panel: rows < MR, use predicated SVE.
1103 if (peeled_rows < rows) {
1104 const Index tail = rows - peeled_rows;
1105 Scalar* dst_panel =
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));
1108 }
1109 }
1110
1111 EIGEN_DONT_INLINE void operator()(Scalar* blockA, const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
1112 Index offset = 0) {
1113 if (PanelMode) {
1114 eigen_assert(stride >= depth && offset <= stride);
1115 }
1116 // Inner-strided ColMajor blas mappers' packets advance the row index, so
1117 // the fallback may use them.
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>{});
1121 }
1122};
1123
1124// RowMajor LHS packer -- SME in-ZA transpose.
1125//
1126// The packed output wants depth-major layout (MR rows contiguous per depth
1127// step) but the RowMajor source has rows contiguous (strided by depth per
1128// row). A natural SVE gather would be slow; instead we use ZA's 2D store
1129// as a free transpose: load svl rows as horizontal slices of a ZA tile,
1130// then read vertical slices to produce depth-major output (see
1131// sme_transpose_pack).
1132template <typename Scalar, int MR, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1133struct sme_pack_lhs_rowmajor {
1134 // Non-streaming NEON copy of every panel, for the shallow blocks the NEON
1135 // kernel consumes.
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);
1142 }
1143 }
1144
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);
1152 }
1153 }
1154
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);
1158
1159 if (peeled_rows > 0) {
1160 pack_full_panels(dst_base, src, src_stride, depth, peeled_rows, dst_stride, dst_offset);
1161 }
1162
1163 // Row tail (rows - peeled_rows in [1, MR-1]), at most once per call: see tail_pack.
1164 if (peeled_rows < rows) {
1165 const Index tail = rows - peeled_rows;
1166 Scalar* dst_panel =
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);
1169 }
1170 }
1171
1172 EIGEN_DONT_INLINE void operator()(Scalar* blockA, const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
1173 Index offset = 0) {
1174 if (PanelMode) {
1175 eigen_assert(stride >= depth && offset <= stride);
1176 }
1177 // Inner-strided RowMajor blas mappers' packets advance the depth index, not
1178 // the row index, so the fallback must stay scalar (see
1179 // sme_pack_lhs_fallback).
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>{});
1183 }
1184};
1185
1186/*****************************************************************************
1187 * gemm_pack_rhs for SME (ColMajor source) -- SME in-ZA transpose, mirroring
1188 * the RowMajor LHS packer.
1189 *
1190 * Packs the RHS matrix into panels of width nr. ColMajor source has
1191 * columns contiguous; we load NR columns as horizontal ZA slices and then
1192 * read verticals to produce depth-major packed output.
1193 *****************************************************************************/
1194
1195template <typename Scalar, int NR, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1196struct sme_pack_rhs_colmajor {
1197 // Non-streaming NEON copy of every panel, for the shallow blocks the NEON
1198 // kernel consumes.
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);
1205 }
1206 }
1207
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);
1215 }
1216 }
1217
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);
1221
1222 if (peeled_cols > 0) {
1223 pack_full_panels(dst_base, src, src_stride, depth, peeled_cols, dst_stride, dst_offset);
1224 }
1225
1226 // Col tail (cols - peeled_cols in [1, NR-1]), at most once per call: see tail_pack.
1227 if (peeled_cols < cols) {
1228 const Index tail = cols - peeled_cols;
1229 Scalar* dst_panel =
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);
1232 }
1233 }
1234
1235 EIGEN_DONT_INLINE void operator()(Scalar* blockB, const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
1236 Index offset = 0) {
1237 if (PanelMode) {
1238 eigen_assert(stride >= depth && offset <= stride);
1239 }
1240 // Inner-strided ColMajor blas mappers' LinearMapper packets advance the
1241 // depth index, which is what the fallback transposes.
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>{});
1245 }
1246};
1247
1248// RowMajor RHS packer -- streaming SVE copy (mirrors the ColMajor LHS packer).
1249// Rows are contiguous in the source, so each depth-step is NR contiguous scalars.
1250template <typename Scalar, int NR, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1251struct sme_pack_rhs_rowmajor {
1252 // Non-streaming NEON copy of every panel, for the shallow blocks the NEON
1253 // kernel consumes.
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);
1260 }
1261 }
1262
1263 // EIGEN_DONT_INLINE: as in sme_pack_lhs_colmajor::pack_direct.
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;
1268
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);
1272 }
1273
1274 if (peeled_cols < cols) {
1275 const Index tail = cols - peeled_cols;
1276 Scalar* dst_panel =
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));
1279 }
1280 }
1281
1282 EIGEN_DONT_INLINE void operator()(Scalar* blockB, const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
1283 Index offset = 0) {
1284 if (PanelMode) {
1285 eigen_assert(stride >= depth && offset <= stride);
1286 }
1287 // Inner-strided RowMajor blas mappers' LinearMapper packets advance the
1288 // column index, not depth, so the fallback must stay scalar (see
1289 // sme_pack_rhs_fallback).
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>{});
1293 }
1294};
1295
1296// Pack1/nr_ are pinned to the SME block sizes (rather than left open) so these
1297// specializations only match consumers that actually feed the SME gebp_kernel
1298// -- see "Mixed-scalar products" in the gebp_traits doc comment above.
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> {}; \
1303 \
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> {}; \
1307 \
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> {}; \
1311 \
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> {};
1315
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)
1321#endif
1322
1323#undef EIGEN_SME_DECLARE_GEMM_PACKERS
1324
1325/*****************************************************************************
1326 * sme_store_za_tile -- Store one ZA tile back to C with alpha scaling.
1327 *
1328 * `pw` is the row-predicate width for this tile, `cw` the col-predicate width
1329 * (both <= the runtime svl).
1330 *****************************************************************************/
1331
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);
1340 // FMLA and FADD have equal latency/throughput on ARMv9 cores, and
1341 // multiplying by alpha=1.0 is exact in IEEE-754 so the FMLA form is
1342 // bit-identical to FADD in that case. A single unconditional FMLA
1343 // keeps the store compact and measures no worse (and a few percent
1344 // better on small matrices, where the branch would otherwise disrupt
1345 // instruction scheduling).
1346 const Vec vzero = pset1<Vec>(Scalar(0));
1347 const Vec valpha = pset1<Vec>(alpha);
1348
1349 // Two C slices are loaded before either is stored: a C line the caller wrote
1350 // from non-streaming code just before the kernel does not forward across the
1351 // mode switch on Apple M4, and a serial load/store pays that latency per slice.
1352 // C = A*B meets the condition on every call, since evalTo zeroes the
1353 // destination first. SVE vectors are sizeless, hence the spelled-out pair.
1354 if (C_stride_row == 1) {
1355 // Column-major C: extract vertical slices (columns of the ZA tile)
1356 int ci = 0;
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));
1364 }
1365 if (ci < cw) {
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));
1369 }
1370 } else if (C_stride_col == 1) {
1371 // Row-major C: extract horizontal slices (rows of the ZA tile)
1372 int ri = 0;
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));
1380 }
1381 if (ri < pw) {
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));
1385 }
1386 } else {
1387 // General stride: extract rows to temp buffer, scatter to C. scratch
1388 // holds one ZA row; every caller passes cw <= min(svl, nr) (a tile
1389 // never spans more than the logical block), so nr is a static
1390 // bound independent of the runtime svl.
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];
1398 }
1399 }
1400 }
1401}
1402
1403/*****************************************************************************
1404 * sme_store_2x2_grid -- store the (up to) 2x2 grid of svl x svl ZA tiles.
1405 *
1406 * Tile layout: 0 = (row-lo, col-lo) 1 = (row-lo, col-hi)
1407 * 2 = (row-hi, col-lo) 3 = (row-hi, col-hi)
1408 * The col-hi tiles (1, 3) are stored only when chi > 0 and the row-hi tiles
1409 * (2, 3) only when rhi > 0, so a single tile, a 1x2/2x1 pair, or the full grid
1410 * all route through here. Runs once per sub-block pass, after a depth loop
1411 * that dwarfs it, so the branches cost nothing and predict perfectly (the
1412 * pattern repeats across blocks).
1413 *****************************************************************************/
1414
1415// Full 2*svl x 2*svl block into column-major C: four columns per vertical four-slice read of the two tiles that
1416// share them, each column loaded and stored whole with two-vector accesses, all four loaded before any is stored.
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);
1430 pstoreu_x2(pn, p,
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))));
1438 }
1439}
1440
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);
1450 return;
1451 }
1452 sme_store_za_tile<Scalar, 0>(C, C_stride_row, C_stride_col, alpha, row_start, rlo, col_start, clo);
1453 if (chi > 0) {
1454 sme_store_za_tile<Scalar, 1>(C, C_stride_row, C_stride_col, alpha, row_start, rlo, col_start + svl, chi);
1455 }
1456 if (rhi > 0) {
1457 sme_store_za_tile<Scalar, 2>(C, C_stride_row, C_stride_col, alpha, row_start + svl, rhi, col_start, clo);
1458 if (chi > 0) {
1459 sme_store_za_tile<Scalar, 3>(C, C_stride_row, C_stride_col, alpha, row_start + svl, rhi, col_start + svl, chi);
1460 }
1461 }
1462}
1463
1464// One depth step's worth of the exact-match grid: the four FMOPAs that take the
1465// lo/hi halves of a packed A column and a packed B column and accumulate the
1466// 2x2 ZA-tile outer product. `all` is the all-true predicate because this is
1467// only used on the exact-match path, where the block fills the grid, so
1468// factoring it out is identical to the inline form.
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);
1479}
1480
1481/*****************************************************************************
1482 * Complex accumulator: a pair of ZA tiles holding the real and imaginary
1483 * halves of one grid cell.
1484 *
1485 * With a = ar + i*sa*ai and b = br + i*sb*bi, where sa is -1 when the LHS is
1486 * conjugated and sb likewise for the RHS,
1487 *
1488 * re(a*b) = ar*br - (sa*sb) * ai*bi, im(a*b) = sb * ar*bi + sa * ai*br,
1489 *
1490 * so all four real outer products differ only in whether they accumulate
1491 * (FMOPA) or subtract (FMOPS) -- a compile-time choice, with no work in the
1492 * depth loop and no separate conjugating packer.
1493 *
1494 * Slices come back out through sme_read_slice, which applies the complex
1495 * alpha and interleaves the halves with ZIP1/ZIP2 into the two vectors that
1496 * cover one slice's worth of contiguous complex results. The kernel that
1497 * folds a real alpha keeps them deinterleaved instead (see
1498 * sme_accumulate_pair_real_alpha).
1499 *
1500 * A ZA tile number is an instruction immediate, so cells outside the tile grid
1501 * (complex<float> has two tile pairs, hence a single grid row) are dropped by
1502 * the InGrid specialization rather than by a runtime guard, which would still
1503 * have to name an in-range tile.
1504 *****************************************************************************/
1505
1506// One slice of a tile pair, re-interleaved into the two vectors that cover its
1507// complex results: `lo` the first half, `hi` the second. ScaleByAlpha applies
1508// the complex alpha, four predicated FP ops the caller skips when alpha is 1
1509// (see sme_store_za_pair) -- streaming-mode FP is de-rated enough on Apple M4
1510// that those four cost about as much as the rest of the slice.
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;
1518 Vec re, im;
1519 EIGEN_IF_CONSTEXPR (Vertical) {
1520 re = sme_read_ver_za<TileRe>(vzero, pg, slice);
1521 im = sme_read_ver_za<TileIm>(vzero, pg, slice);
1522 } else {
1523 re = sme_read_hor_za<TileRe>(vzero, pg, slice);
1524 im = sme_read_hor_za<TileIm>(vzero, pg, slice);
1525 }
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);
1531 } else {
1532 lo = pzip1(re, im);
1533 hi = pzip2(re, im);
1534 }
1535}
1536
1537// Accumulate a tile pair's `slices` slices into C along its contiguous axis,
1538// `step` reals apart. `lanes` is twice a slice's complex count; when it fits
1539// one vector the high half's predicate is empty, so its load and store are
1540// no-ops even though the destination has nothing at p + svl to point at.
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) {
1552 Vec lo, hi;
1553 sme_read_slice<RealScalar, TileRe, TileIm, Vertical, ScaleByAlpha>(pg, valpha_re, valpha_im, vzero, uint32_t(s), lo,
1554 hi);
1555 pstoreu(pl0, p, padd(pl0, ploadu(pl0, p), lo));
1556 // pl1 is all-false when one vector covers the slice, and an inactive lane
1557 // neither reads nor writes -- so this needs no `lanes > svl` guard, only an
1558 // address the destination is allowed to form.
1559 RealScalar* EIGEN_RESTRICT phi = sme_offset(p, Index(svl));
1560 pstoreu(pl1, phi, padd(pl1, ploadu(pl1, phi), hi));
1561 }
1562}
1563
1564// A real alpha on slices spanning two vectors, in the FoldRealAlpha kernel:
1565// LD2/ST2 keep C deinterleaved with one predicate lane per complex result, so
1566// each half is one FMA, c + alpha*acc, as in the real-scalar store. A single
1567// FMA rounds once, so it overflows only when the exact result does. This
1568// covers C.noalias() -= A*B, whose alpha is -1.
1569//
1570// `limit`, the slices left in the C block from the first one, bounds the
1571// prefetch to the block (unbounded, it measured 0.45-0.5x on 32^3 and 64^3).
1572// The prefetch runs 8 slices ahead of the C load that heads each FMA; without
1573// it the fold ran at 0.94x of that at 1024^3. The barrier keeps the prefetch
1574// address setup inside its branch.
1575//
1576// A complex alpha keeps sme_accumulate_pair_impl, which forms alpha*acc before
1577// adding C. Folding C into either of its two FMAs, e.g. (c_re + re*ar) - im*ai,
1578// saves two ops per slice but overflows when C and the first product exceed
1579// the range before the second cancels them (!3139 review: acc = (h,h),
1580// alpha = (1,1), C = (3h,0)). On Apple M4, LD2/ST2 measured no gain with
1581// unchanged arithmetic and 0.68-0.75x on one-vector slices, and FCMLA 0.70x
1582// (#3129).
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) {
1594 RealScalar* pf = p;
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);
1599 }
1600 Vec re, im;
1601 EIGEN_IF_CONSTEXPR (Vertical) {
1602 re = sme_read_ver_za<TileRe>(vzero, pg, uint32_t(s));
1603 im = sme_read_ver_za<TileIm>(vzero, pg, uint32_t(s));
1604 } else {
1605 re = sme_read_hor_za<TileRe>(vzero, pg, uint32_t(s));
1606 im = sme_read_hor_za<TileIm>(vzero, pg, uint32_t(s));
1607 }
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))));
1610 }
1611}
1612
1613// FoldRealAlpha selects the kernel instantiation. The folding kernel only
1614// sees a real alpha other than 1 (see sme_gebp_dispatch); the other kernel
1615// keeps exactly the unscaled and complex-alpha paths, so the fold's code never
1616// reaches products that cannot use it.
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,
1626 valpha_im, vzero);
1627 } else {
1628 sme_accumulate_pair_impl<RealScalar, TileRe, TileIm, Vertical, false>(p, step, slices, lanes, pg, valpha_re,
1629 valpha_im, vzero);
1630 }
1631}
1632
1633// Store one complex tile pair back to C. `pw` is the row-predicate width for
1634// this cell and `cw` the column one, both <= the runtime svl.
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));
1647 // std::complex's accessors are ordinary functions, which clang cannot inline
1648 // into a streaming context; the resulting mode switch would sit in this loop.
1649 // A complex is layout-compatible with its two-element real array, so read the
1650 // parts through that view instead.
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]);
1654 // Scaling by 1 + 0i is exact, so skipping it is bit-identical -- and it is by
1655 // far the common case, since a plain product carries alpha = 1.
1656 // With FoldRealAlpha, sme_gebp_dispatch has already established that alpha
1657 // is real and not 1.
1658 const bool scale = !(alpha_parts[0] == RealScalar(1) && alpha_parts[1] == RealScalar(0));
1659 RealScalar* EIGEN_RESTRICT rC = reinterpret_cast<RealScalar*>(C);
1660
1661 if (C_stride_row == 1) {
1662 // Column-major C: vertical slices are the tile pair's columns, and one
1663 // slice is pw contiguous complex results, i.e. 2*pw reals.
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) {
1668 // Row-major C: horizontal slices are the tile pair's rows.
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);
1672 } else {
1673 // General stride: interleave a row into a temp buffer, scatter to C. Every
1674 // caller passes cw <= min(svl, nr), so nr is a static bound on the buffer,
1675 // independent of the runtime svl. This path is scalar anyway, so it always
1676 // takes the scaling form.
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) {
1683 Vec lo, hi;
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);
1686 // scratch is nr complex, i.e. 2*nr >= 2*svl reals, so rscratch + svl is
1687 // always in bounds; pl1 is all-false when one vector already covers cw.
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];
1691 }
1692 }
1693 }
1694}
1695
1696// Grid cell (R, C) of a complex block: its tile pair, the four signed outer
1697// products that feed it, and its store.
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;
1705
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);
1712 }
1713
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);
1720 }
1721};
1722
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") {
1727 }
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") {}
1731};
1732
1733/*****************************************************************************
1734 * sme_process -- micro-kernel for one pw x cw output block.
1735 *
1736 * Tiles the block into svl x svl ZA tiles, processed in passes of up to a 2x2
1737 * tile grid: several (2*svl) x (2*svl) sub-block passes when the grid is
1738 * smaller than the block, tiles predicated down to the block width when it is
1739 * larger. blA/blB are packed depth-major with depth-strides pw and cw
1740 * respectively.
1741 *
1742 * When the block matches the tile grid exactly (pw == cw == 2 * svl), the
1743 * packed rows are also contiguous across depth steps, enabling the
1744 * hand-scheduled loop below: per 4 unrolled depth steps, 2 x4 loads per
1745 * side (each spanning 2 depth steps) feed 16 FMOPAs -- a 1:1 compute:load
1746 * ratio at the vector level. All other geometries use predicated
1747 * per-depth-step loads.
1748 *****************************************************************************/
1749
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") {
1755 // Conjugation is the identity on real scalars, so this overload ignores it.
1756 // a_step is the distance between A's depth steps: pw for a packed panel, the
1757 // column stride for a ColMajor source read in place.
1758 using Traits = sme_packet_traits<Scalar>;
1759 using Vec = typename Traits::type;
1760 const int svl = Traits::size();
1761
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; // >= 0; > 0 only when rpw > svl, in which case rlo == svl
1766 const svbool_t pg_rlo = Traits::whilelt(rt, pw);
1767 const svbool_t pg_rhi = Traits::whilelt(rt + svl, pw);
1768
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);
1775
1776 svzero_za();
1777 if (pw == 2 * svl && cw == 2 * svl) {
1778 // The block is exactly one full-grid patch (single pass, rt == ct ==
1779 // 0, rlo == rhi == clo == chi == svl), so a packed row is the
1780 // patch's slice and rows are contiguous across depth steps: x4 loads
1781 // each span 2 of them, e.g. va_01 = [d0 lo, d0 hi, d1 lo, d1 hi].
1782 const svcount_t pn = Traits::ptrue_c();
1783 const Index depth_4 = (depth / 4) * 4;
1784 Index k = 0;
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]);
1789
1790 // d0
1791 outer_product_2x2<Scalar>(pget<0>(va_01), pget<1>(va_01), pget<0>(vb_01), pget<1>(vb_01));
1792 // d1
1793 outer_product_2x2<Scalar>(pget<2>(va_01), pget<3>(va_01), pget<2>(vb_01), pget<3>(vb_01));
1794
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]);
1797
1798 // d2
1799 outer_product_2x2<Scalar>(pget<0>(va_23), pget<1>(va_23), pget<0>(vb_23), pget<1>(vb_23));
1800 // d3
1801 outer_product_2x2<Scalar>(pget<2>(va_23), pget<3>(va_23), pget<2>(vb_23), pget<3>(vb_23));
1802 }
1803 } else {
1804 // A read in place: its depth steps are a_step apart, one x2 load each.
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));
1816 }
1817 }
1818 // Depth tail: one x2 load per side per step.
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));
1823 }
1824 } else {
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]);
1828
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));
1831
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);
1837 }
1838 }
1839 }
1840
1841 // Store the (up to) 2x2 grid of tiles for this sub-block pass.
1842 sme_store_2x2_grid(C, C_stride_row, C_stride_col, alpha, row_start + rt, rlo, rhi, col_start + ct, clo, chi);
1843 }
1844 }
1845}
1846
1847/*****************************************************************************
1848 * sme_store_complex_grid -- store the (up to) 2x2 grid of tile pairs, exactly
1849 * as sme_store_2x2_grid does for single tiles. Cells outside the grid are
1850 * dropped at compile time by sme_complex_cell, so a narrower grid simply never
1851 * reaches them (its hi widths are structurally zero).
1852 *****************************************************************************/
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);
1861 if (chi > 0) {
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);
1864 }
1865 if (rhi > 0) {
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);
1868 if (chi > 0) {
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);
1871 }
1872 }
1873}
1874
1875/*****************************************************************************
1876 * sme_process, complex overload -- same sub-block structure over a grid of
1877 * complex accumulators, each a pair of ZA tiles (see sme_complex_cell).
1878 *
1879 * The packed panels are read through their real view: one depth step is `pw`
1880 * (`cw`) reals followed by as many imaginary ones, so a cell's four operands
1881 * are four contiguous predicated loads at a fixed offset apart, and the four
1882 * outer products they feed reuse all of them.
1883 *****************************************************************************/
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") {
1890 // Complex panels are always packed (split real/imaginary halves): lhs_step is pw.
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>;
1902
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);
1907
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);
1914
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);
1921
1922 svzero_za();
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);
1926 // The second-tile loads are unconditional: their predicates are empty
1927 // when the grid is a single tile wide or tall, so they touch no memory.
1928 // Grouping them lets the compiler issue them in parallel, and gating the
1929 // outer products on svptest_any keeps the depth loop unswitched without
1930 // materialising the r1/c1 counts here.
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);
1942 }
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);
1947 }
1948 }
1949 }
1950
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);
1953 }
1954 }
1955}
1956
1957// A RHS panel of at most svl columns fills only one tile column of the 2 x 2 grid, so the four tiles are stacked along
1958// M instead: a full LHS panel (tiles 0 and 1) and the next one (tiles 2 and 3, pw1 rows) against the same B vector.
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);
1973 svzero_za();
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);
1983 }
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,
1987 cw);
1988 if (pw1 > svl)
1989 sme_store_za_tile<Scalar, 3>(C, C_stride_row, C_stride_col, alpha, row_start + 3 * svl, pw1 - svl, col_start, cw);
1990}
1991
1992// Tile Dst += tile Src, four columns per step.
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))));
2001 }
2002}
2003
2004// A block that fills one or two tiles (pw or cw <= svl): FMOPAs into one tile wait on each other, so consecutive
2005// depth steps go to the spare tiles and the partial sums are folded before the store.
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);
2016 svzero_za();
2017 Index k = 0;
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);
2028 }
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) {
2035 // Tiles 0 and 1 take the even depth steps, 2 and 3 the odd ones.
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);
2044 }
2045 if (k < depth) {
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));
2049 }
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);
2053 } else {
2054 // cw <= svl < pw: tiles 0 and 2 take the even depth steps, 1 and 3 the odd ones.
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);
2063 }
2064 if (k < depth) {
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);
2068 }
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);
2072 }
2073}
2074
2075// One pw x cw block: sme_process_split when it fills at most two tiles of the 2 x 2 grid; complex blocks keep
2076// sme_process. `split_ok` (sme_split_ok) holds once per call, outside the block loops.
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);
2085 else
2086 sme_process<ConjLhs, ConjRhs>(C, rs, cs, blA, blB, depth, alpha, row_start, pw, col_start, cw, a_step);
2087}
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,
2095 rows, cols);
2096}
2097
2098// Whether the vector length gives the tile shapes the block kernels assume: the depth-split kernel covers one 2 x 2
2099// grid of blocks up to MR x NR with four-slice tile folds, and the narrow kernel stacks two LHS panels of 2*svl rows.
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;
2104}
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();
2108}
2109
2110// sme_process_narrow for real scalars; complex panels keep the 2 x 2 path.
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);
2117}
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") {}
2125
2126// Core-side prefetch into L2 of the column-major C block the next sme_process call reads and writes.
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);
2137}
2138
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;
2145
2146 // Column-outer, row-inner: keeps blB (one kc × NR panel) hot in L1 while
2147 // smaller blA tiles stream from L2. The outer GOTO loop in
2148 // GeneralMatrixMatrix.h ensures blockA fits in L2 via mc-blocking. Each
2149 // packed panel is depth-major with depth-stride equal to its width (MR/NR
2150 // for full panels, the tail width otherwise), so that width is passed as
2151 // both the logical block size and the load stride to the block kernels.
2152 // The narrow kernel stacks two LHS panels of exactly 2*svl rows.
2153 // A small C block (256 KB or less) gains nothing from the prefetch and pays its instructions on every call.
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;
2159
2160 Index i = 0;
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,
2167 alpha, i, j);
2168 }
2169 }
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;
2173 if (prefetch_c)
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);
2178 }
2179 }
2180}
2181
2182// Picks the sme_gebp_impl instantiation. A complex product with a real alpha
2183// other than 1, e.g. the -1 of C.noalias() -= A*B, gets the kernel that folds
2184// C into one FMA per component (see sme_accumulate_pair_real_alpha); every
2185// other product gets the kernel without that path. Real scalars have one.
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,
2191 Index offsetB) {
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);
2194 }
2195};
2196
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,
2202 Index offsetB) {
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);
2207 } else {
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);
2210 }
2211 }
2212};
2213
2214// gebp with the LHS read from a ColMajor source: rows i..i+pw of column k are
2215// at lhs + i + k * lda, the packed layout with a_step = lda. Real scalars only.
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;
2226 Index i = 0;
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);
2235 }
2236 }
2237}
2238
2239// In-place ColMajor LHS: deep enough for the ZA kernel, a stride that is not a 4 KB multiple (such columns map to one
2240// L1 set) and a block of at most 4 MB, past which the strided re-reads lose to the packed panel.
2241#ifndef EIGEN_SME_DIRECT_LHS_MAX_STRIDE_BYTES
2242#define EIGEN_SME_DIRECT_LHS_MAX_STRIDE_BYTES 16384
2243#endif
2244#ifndef EIGEN_SME_DIRECT_LHS_MAX_BLOCK_BYTES
2245#define EIGEN_SME_DIRECT_LHS_MAX_BLOCK_BYTES (4 << 20)
2246#endif
2247#ifndef EIGEN_SME_DIRECT_LHS_MAX_SPAN_BYTES
2248#define EIGEN_SME_DIRECT_LHS_MAX_SPAN_BYTES (16 << 20)
2249#endif
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);
2257 return false;
2258#else
2259 if (NumTraits<Scalar>::IsComplex || depth <= Index(sme_neon_max_depth<Scalar>::value)) return false;
2260 // A RHS of one panel reads each LHS element once, so packing it only adds a copy: always for the narrow kernel,
2261 // and for the 2 x 2 one while a panel spans few enough pages.
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;
2267 }
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);
2270#endif
2271}
2272
2273// NEON path for small blocks: a shallow block pays the ZA enable and a dependent FMOPA chain per
2274// tile for little work, so its panels are consumed with NEON in the SME layout (one depth step is
2275// pw contiguous scalars, complex ones as pw reals then pw imaginaries).
2276// The NCol columns of one depth step, loaded once; each column is then an
2277// immediate lane of a fused multiply-add (nmadd: acc - a * b[lane]).
2278template <typename Scalar, int NCol>
2279struct sme_neon_cols;
2280template <>
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); }
2284 template <int C>
2285 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2286 return vfmaq_laneq_f32(acc, a, b, C);
2287 }
2288 template <int C>
2289 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2290 return vfmsq_laneq_f32(acc, a, b, C);
2291 }
2292};
2293template <>
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); }
2297 template <int C>
2298 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2299 return vfmaq_lane_f32(acc, a, b, C);
2300 }
2301 template <int C>
2302 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2303 return vfmsq_lane_f32(acc, a, b, C);
2304 }
2305};
2306template <>
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); }
2310 template <int C>
2311 static EIGEN_ALWAYS_INLINE Packet4f madd(Packet4f acc, Packet4f a, Vec b) {
2312 return vfmaq_f32(acc, a, b);
2313 }
2314 template <int C>
2315 static EIGEN_ALWAYS_INLINE Packet4f nmadd(Packet4f acc, Packet4f a, Vec b) {
2316 return vfmsq_f32(acc, a, b);
2317 }
2318};
2319template <>
2320struct sme_neon_cols<double, 4> {
2321 struct Vec {
2322 float64x2_t lo, hi;
2323 };
2324 static EIGEN_ALWAYS_INLINE Vec load(const double* b) { return {vld1q_f64(b), vld1q_f64(b + 2)}; }
2325 template <int C>
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);
2328 }
2329 template <int C>
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);
2332 }
2333};
2334template <>
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); }
2338 template <int C>
2339 static EIGEN_ALWAYS_INLINE Packet2d madd(Packet2d acc, Packet2d a, Vec b) {
2340 return vfmaq_laneq_f64(acc, a, b, C);
2341 }
2342 template <int C>
2343 static EIGEN_ALWAYS_INLINE Packet2d nmadd(Packet2d acc, Packet2d a, Vec b) {
2344 return vfmsq_laneq_f64(acc, a, b, C);
2345 }
2346};
2347template <>
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); }
2351 template <int C>
2352 static EIGEN_ALWAYS_INLINE Packet2d madd(Packet2d acc, Packet2d a, Vec b) {
2353 return vfmaq_f64(acc, a, b);
2354 }
2355 template <int C>
2356 static EIGEN_ALWAYS_INLINE Packet2d nmadd(Packet2d acc, Packet2d a, Vec b) {
2357 return vfmsq_f64(acc, a, b);
2358 }
2359};
2360
2361template <>
2362struct sme_neon_cols<float, 8> {
2363 struct Vec {
2364 float32x4_t lo, hi;
2365 };
2366 static EIGEN_ALWAYS_INLINE Vec load(const float* b) { return {vld1q_f32(b), vld1q_f32(b + 4)}; }
2367 template <int C>
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);
2370 }
2371 template <int C>
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);
2374 }
2375};
2376template <>
2377struct sme_neon_cols<double, 8> {
2378 struct Vec {
2379 float64x2_t v[4];
2380 };
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)}};
2383 }
2384 template <int C>
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);
2387 }
2388 template <int C>
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);
2391 }
2392};
2393
2394// Compile-time loop over the columns of a tile (lane numbers are immediates).
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);
2401 }
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);
2414 }
2415 sme_neon_col_loop<C + 1, NCol>::template cplx<ConjLhs, ConjRhs, Cols, Packet, NPack>(acc_re, acc_im, are, aim, bre,
2416 bim);
2417 }
2418};
2419template <int NCol>
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) {}
2426};
2427
2428// One (NPack * PacketSize) x NCol micro-tile of a real block; NPack == 0 is a
2429// single scalar row. Pointers arrive offset to the tile's first row and column.
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;
2435 // Narrow tiles split the depth over KU accumulator chains to hide FMA latency.
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) {
2438 Packet av[NPack];
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));
2441 }
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));
2448 Index k = 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);
2451 }
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;
2460 if (rs == 1) {
2461 pstoreu(pc, pmadd(acc[0][p][c], valpha, ploadu<Packet>(pc)));
2462 } else {
2463 pscatter<Scalar, Packet>(pc, pmadd(acc[0][p][c], valpha, pgather<Scalar, Packet>(pc, rs)), rs);
2464 }
2465 }
2466 }
2467 }
2468};
2469
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);
2479 Index k = 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];
2485 }
2486 }
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];
2491 }
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];
2495 }
2496};
2497
2498// Complex counterpart over the real view of the panels. Conjugation is folded
2499// into the signs of the four real products.
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));
2513 }
2514 }
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);
2522 }
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));
2525 }
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);
2531 if (rs == 1) {
2532 // Unit row stride: the column is PS interleaved (re, im) pairs.
2533 Packet cre, cim;
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);
2538 } else {
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];
2546 }
2547 }
2548 }
2549 }
2550 }
2551};
2552
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;
2572 }
2573 }
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];
2578 }
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];
2584 }
2585 }
2586};
2587
2588// Row loop of one column group: 3, 2 and 1 packets of rows, then scalar rows.
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);
2593 Index r = 0;
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);
2600 for (; r < pw; ++r)
2601 sme_neon_tile<Scalar, Index, 0, NCol>::run(C + r * rs, rs, cs, blA + r, blB, depth, alpha, pw, cw);
2602}
2603
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) {
2608 // Two accumulators per packet, so the ladder stops at two packets of rows.
2609 constexpr Index PS = Index(packet_traits<RealScalar>::size);
2610 Index r = 0;
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,
2613 cw);
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,
2616 cw);
2617 for (; r < pw; ++r)
2618 sme_neon_ctile<RealScalar, Index, ConjLhs, ConjRhs, 0, NCol>::run(C + r * rs, rs, cs, rA + r, rB, depth, alpha, pw,
2619 cw);
2620}
2621
2622// One pw x cw block: columns in groups of eight, four, two, then single columns.
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) {
2626 Index c = 0;
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);
2633 for (; c < cw; ++c)
2634 sme_neon_column_group<ConjLhs, ConjRhs, 1>(C + c * cs, rs, cs, blA, blB + c, depth, alpha, pw, cw);
2635}
2636
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);
2643 Index c = 0;
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);
2649}
2650
2651// Same panel walk as sme_gebp_impl, outside any streaming region.
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);
2666 }
2667 }
2668}
2669
2670template <typename Scalar, typename Index, typename DataMapper, int mr, int nr, bool ConjugateLhs, bool ConjugateRhs>
2671struct sme_gebp_kernel {
2672 using ResScalar = Scalar;
2673
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) {
2677 // Real scalars never reach the kernel conjugated (conj_helper folds it into
2678 // the identity long before), so the real path stays free of the flags.
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");
2683
2684 if (strideA == -1) strideA = depth;
2685 if (strideB == -1) strideB = depth;
2686
2687 if (rows <= 0 || cols <= 0 || depth <= 0) return;
2688
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);
2692
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);
2696 return;
2697 }
2698
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);
2702 }
2703
2704 // The LHS block read from its ColMajor source (see sme_direct_lhs_ok); only
2705 // instantiated for real scalars.
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);
2717 }
2718};
2719
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> {};
2724
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>)
2730#endif
2731
2732#undef EIGEN_SME_DECLARE_GEBP_KERNEL
2733
2734// sme_has_gebp_kernel (products/GeneralBlockPanelKernel.h) drives the cache
2735// blocking and the GEMM loop order, and is declared before this header. A pair
2736// listed there but not specialized here would be packed and blocked for SME and
2737// then handed to the generic 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");
2745#else
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");
2750#endif
2751
2752// ---------------------------------------------------------------------------
2753// Selfadjoint (SYMM) packers.
2754//
2755// product_selfadjoint_matrix packs the selfadjoint operand (stored as one
2756// triangle) through symm_pack_lhs/symm_pack_rhs, which materialize the full
2757// matrix as they pack. The generic SYMM packers emit packet-width sub-panels
2758// for the generic gebp_kernel, whereas the SME kernel expects uniform
2759// mr/nr-wide depth-major panels. These packers perform the same
2760// triangle mirroring in the SME layout.
2761//
2762// The packer receives the operand in an orientation where row >= col is the
2763// stored triangle. It reads that half directly and mirrors the other half:
2764// full(row,col) = (row >= col) ? m(row,col) : conj(m(col,row))
2765// and a selfadjoint view defines the diagonal's imaginary part as zero, so
2766// full(k,k) is real(m(k,k)). Both reduce to the identity on real scalars.
2767//
2768// Regions wholly below or above the diagonal use the normal dense copy or
2769// transpose packers. Only the width-wide part of a panel crossed by the
2770// diagonal needs special handling: each depth row is split between the stored
2771// triangle and its mirrored half.
2772//
2773// For a panel at offset j (entries j+c, c in [0,w)) and global row k2+k, the
2774// three depth regions are:
2775// transposed k in [0, j-k2) : k2+k < j+c for all c -> m(j+c, k2+k)
2776// straddle k in [j-k2, j+w-k2) : diagonal crosses -> per-k split
2777// direct k in [j+w-k2, depth): k2+k > j+c for all c -> m(k2+k, j+c)
2778//
2779// The RHS packs full(k2+k, j+c) and uses this mapping directly, so its mirrored
2780// half is the transposed region. The LHS packs full(j+r, k), which is the
2781// conjugate of full(k, j+r), so it reuses the same mapping with k2 == 0
2782// relative to its diagonal-anchored base pointer but conjugates the opposite
2783// regions -- the direct one and the straddle band's head. IsLhs selects which.
2784// ---------------------------------------------------------------------------
2785
2786// Streaming packer shared by the LHS (k2 == 0) and RHS symm specializations.
2787// ColM selects the ColMajor selfadjoint operand.
2788// Depth-region boundaries for the panel at outer offset `j`, all clamped to
2789// [0, depth]: the diagonal splits it into a transposed head [0, t_end), a
2790// straddle band [t_end, s_end) and a direct tail [s_end, depth).
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);
2797}
2798
2799// The two dense regions of every panel, which are ordinary copies or ZA
2800// transposes of the stored triangle. ColM selects the ColMajor operand.
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);
2806 // The transposed region is the RHS's mirrored half and the direct one the
2807 // LHS's, so exactly one of the two is conjugated (see above).
2808 constexpr bool ConjTransposed = !IsLhs;
2809 constexpr bool ConjDirect = IsLhs;
2810
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; // depth-major panel of width w
2814 Index t_end, s_end;
2815 sme_symm_panel_regions(j, w, depth, k2, t_end, s_end);
2816
2817 // Transposed region: full(k2+k, j+c) = m(j+c, k2+k).
2818 if (t_end > 0) {
2819 EIGEN_IF_CONSTEXPR (ColM) {
2820 sve_copy_panel_range<ConjTransposed>(dst, base + j + k2 * stride, stride, Index(0), t_end, w);
2821 } else {
2822 sme_transpose_pack_range<ConjTransposed>(dst, base + j * stride + k2, stride, Index(0), t_end, w);
2823 }
2824 }
2825 // Direct region: full(k2+k, j+c) = m(k2+k, j+c).
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);
2829 } else {
2830 sve_copy_panel_range<ConjDirect>(dst, base + k2 * stride + j, stride, s_end, depth, w);
2831 }
2832 }
2833 }
2834}
2835
2836// The diagonal band of every panel: the diagonal crosses at c* = (k2+k) - j
2837// (in [0, w) throughout the band), so each depth step splits into a direct head
2838// (c < c*: m(k2+k, j+c)) and a mirrored tail (c >= c*: m(j+c, k2+k); at c == c*
2839// both name the diagonal element).
2840//
2841// Kept out of the streaming region above for the reason tail_transpose_pack gives,
2842// at the cost of a second pass over the panels: it is scalar floating-point,
2843// and fusing it made the float SYMM packers 2-11x slower.
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);
2849 // The band's tail is the RHS's mirrored half and its head the LHS's, exactly
2850 // as the dense regions above.
2851 constexpr bool ConjHead = IsLhs;
2852 constexpr bool ConjTail = !IsLhs;
2853
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;
2857 Index t_end, s_end;
2858 sme_symm_panel_regions(j, w, depth, k2, t_end, s_end);
2859
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; // m(row, j+c): stride-strided
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; // m(j+c, row): contiguous
2869 int c = cs;
2870 // A selfadjoint view defines the diagonal's imaginary part as zero;
2871 // for a real scalar that is already true, so the peel folds away.
2872 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsComplex) {
2873 sme_pack_store<false>(dst_row, wi, Index(cs), Scalar(numext::real(tail[cs])));
2874 ++c;
2875 }
2876 for (; c < w; ++c) sme_pack_store<ConjTail>(dst_row, wi, Index(c), tail[c]);
2877 } else {
2878 const Scalar* head = base + row * stride + j; // m(row, j+c): contiguous
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; // m(j+c, row): stride-strided
2881 int c = cs;
2882 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsComplex) {
2883 sme_pack_store<false>(dst_row, wi, Index(cs), Scalar(numext::real(*tail)));
2884 tail += stride;
2885 ++c;
2886 }
2887 for (; c < w; ++c, tail += stride) sme_pack_store<ConjTail>(dst_row, wi, Index(c), *tail);
2888 }
2889 }
2890 }
2891}
2892
2893// Packer shared by the LHS (k2 == 0) and RHS symm specializations.
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) {
2897 {
2898 sme_fpsr_guard fpsr;
2899 sme_symm_pack_dense_regions<Scalar, StorageOrder, IsLhs, Index>(block, base, stride, depth, outer, k2);
2900 }
2901 sme_symm_pack_straddle<Scalar, StorageOrder, IsLhs, Index>(block, base, stride, depth, outer, k2);
2902}
2903
2904// symm_pack_lhs/rhs SME specializations: emit the uniform mr/nr panels
2905// sme_gebp_impl reads. Pack1/nr pinned exactly as gemm_pack_lhs/rhs above.
2906template <typename Scalar, int StorageOrder, typename Index>
2907struct sme_symm_pack_lhs {
2908 // Note: generic symm_pack_lhs's "cols" is the depth extent, and the LHS
2909 // block is diagonal-anchored (base = &lhs(k2,k2)), so its depth offset is 0.
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));
2912 }
2913};
2914
2915template <typename Scalar, int StorageOrder, typename Index>
2916struct sme_symm_pack_rhs {
2917 // Note: generic symm_pack_rhs's "rows" is the depth extent (end_k = k2 + rows), not a row count.
2918 EIGEN_DONT_INLINE void operator()(Scalar* blockB, const Scalar* rhs_, Index rhsStride, Index rows, Index cols,
2919 Index k2) const {
2920 sme_symm_pack_panels<Scalar, StorageOrder, false, Index>(blockB, rhs_, rhsStride, rows, cols, k2);
2921 }
2922};
2923
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> {}; \
2928 \
2929 template <typename Index, int StorageOrder> \
2930 struct symm_pack_rhs<SCALAR, Index, NR, StorageOrder> : sme_symm_pack_rhs<SCALAR, StorageOrder, Index> {};
2931
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)
2937#endif
2938
2939#undef EIGEN_SME_DECLARE_SYMM_PACKERS
2940
2941// Tiny results (at most 2 * PS x 8): a NEON outer-product kernel with the whole result in registers, called before
2942// any blocking or packing. Both the coeff-based product (one dependent FMA chain per result) and the packed paths
2943// are several times slower there.
2944template <typename Scalar>
2945struct sme_tiny_neon;
2946template <>
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); }
2952 // The first n (< PS) scalars at p, zeros after them.
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);
2958 }
2959 // A full load permuted by byte indices; indices of 16 or more give zeros.
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));
2962 }
2963 template <int L>
2964 static EIGEN_ALWAYS_INLINE V fma_lane(V c, V a, V b) {
2965 return vfmaq_laneq_f32(c, a, b, L);
2966 }
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); }
2970 // Rows in, columns out.
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));
2980 }
2981};
2982template <>
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));
2991 }
2992 template <int L>
2993 static EIGEN_ALWAYS_INLINE V fma_lane(V c, V a, V b) {
2994 return vfmaq_laneq_f64(c, a, b, L);
2995 }
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]);
3002 r[0] = t0;
3003 }
3004};
3005
3006// One depth chunk of FMAs; the lane index must be a constant, hence the recursion over it. Depth step u of a chunk
3007// accumulates into set u % S: four FMA pipes of latency four need about 16 independent chains.
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;
3012 // ColMajor B: b[j] holds depth steps k .. k + PS - 1 of column j; lane U is step k + U.
3013 template <int U>
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>) {
3016 EIGEN_UNROLL_LOOP
3017 for (int j = 0; j < NC; ++j) {
3018 EIGEN_UNROLL_LOOP
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]);
3020 }
3021 col(acc, a, b, std::integral_constant<int, U + 1>());
3022 }
3023 static EIGEN_ALWAYS_INLINE void col(V (&)[S][NC][RV], const V (&)[PS][RV], const V (&)[NC],
3024 std::integral_constant<int, PS>) {}
3025 // RowMajor B: b[c] holds columns c * PS .. of one depth step.
3026 template <int J>
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>) {
3029 EIGEN_UNROLL_LOOP
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>());
3032 }
3033 static EIGEN_ALWAYS_INLINE void row(V (&)[NC][RV], const V (&)[RV], const V (&)[NC / PS],
3034 std::integral_constant<int, NC>) {}
3035};
3036
3037// C (m x n, m <= RV * PS, n <= NC) += alpha * A * B, reading only A's m rows and B's n columns: rows and columns
3038// past the result repeat the last valid one or load as zeros, and their sums are dropped.
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>;
3046 V acc[S][NC][RV];
3047 EIGEN_UNROLL_LOOP
3048 for (int t = 0; t < S; ++t) {
3049 EIGEN_UNROLL_LOOP
3050 for (int j = 0; j < NC; ++j) {
3051 EIGEN_UNROLL_LOOP
3052 for (int r = 0; r < RV; ++r) acc[t][j][r] = T::zero();
3053 }
3054 }
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;
3059 // A column (B row) vector r holds cnt valid scalars. Once m (n) >= PS, each is one full load ending at its last
3060 // valid scalar, shifted down by a byte table with zeros after; below that, a partial load.
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;
3068 return vld1q_u8(t);
3069 };
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]);
3074 }
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]);
3079 }
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) {
3083 V a[PS][RV];
3084 EIGEN_IF_CONSTEXPR (LhsOrder == ColMajor) {
3085 EIGEN_UNROLL_LOOP
3086 for (int u = 0; u < PS; ++u) {
3087 EIGEN_UNROLL_LOOP
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]);
3092 }
3093 } else {
3094 EIGEN_UNROLL_LOOP
3095 for (int r = 0; r < RV; ++r) {
3096 V blk[PS];
3097 EIGEN_UNROLL_LOOP
3098 for (int q = 0; q < PS; ++q) blk[q] = T::ld(a_row[r * PS + q] + k);
3099 T::transpose(blk);
3100 EIGEN_UNROLL_LOOP
3101 for (int u = 0; u < PS; ++u) a[u][r] = blk[u];
3102 }
3103 }
3104 EIGEN_IF_CONSTEXPR (RhsOrder == ColMajor) {
3105 V b[NC];
3106 EIGEN_UNROLL_LOOP
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>());
3109 } else {
3110 EIGEN_UNROLL_LOOP
3111 for (int u = 0; u < PS; ++u) {
3112 V b[NC / PS];
3113 EIGEN_UNROLL_LOOP
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>());
3120 }
3121 }
3122 }
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];
3127 V a[RV];
3128 EIGEN_UNROLL_LOOP
3129 for (int r = 0; r < RV; ++r) a[r] = T::ld(at + r * PS);
3130 EIGEN_UNROLL_LOOP
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)];
3133 EIGEN_UNROLL_LOOP
3134 for (int r = 0; r < RV; ++r) acc[0][j][r] = T::fma_n(acc[0][j][r], a[r], b);
3135 }
3136 }
3137 EIGEN_UNROLL_LOOP
3138 for (int t = 1; t < S; ++t) {
3139 EIGEN_UNROLL_LOOP
3140 for (int j = 0; j < NC; ++j) {
3141 EIGEN_UNROLL_LOOP
3142 for (int r = 0; r < RV; ++r) acc[0][j][r] = T::add(acc[0][j][r], acc[t][j][r]);
3143 }
3144 }
3145 EIGEN_ALIGN16 Scalar out[NC][MR];
3146 EIGEN_UNROLL_LOOP
3147 for (int j = 0; j < NC; ++j) {
3148 EIGEN_UNROLL_LOOP
3149 for (int r = 0; r < RV; ++r) T::st(&out[j][r * PS], acc[0][j][r]);
3150 }
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];
3153}
3154
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);
3164 else if (one_vec)
3165 sme_tiny_gemm_kernel<Scalar, 1, 8, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3166 resIncr, resStride, alpha);
3167 else if (four_cols)
3168 sme_tiny_gemm_kernel<Scalar, 2, 4, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3169 resIncr, resStride, alpha);
3170 else
3171 sme_tiny_gemm_kernel<Scalar, 2, 8, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res,
3172 resIncr, resStride, alpha);
3173 return true;
3174}
3175
3176} // namespace internal
3177} // namespace Eigen
3178
3179#endif // EIGEN_SME_GENERALBLOCKPANELKERNEL_H
@ ColMajor
Definition Constants.h:319
@ Vertical
Definition Constants.h:267