12#ifndef EIGEN_TRIANGULAR_SOLVER_MATRIX_H
13#define EIGEN_TRIANGULAR_SOLVER_MATRIX_H
16#include "../InternalHeaderCheck.h"
22template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStride,
33 static void kernel(Index size, Index otherSize,
const Scalar* _tri, Index triStride, Scalar* _other, Index otherIncr,
37template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStride,
42 static void kernel(Index size, Index otherSize,
const Scalar* _tri, Index triStride, Scalar* _other, Index otherIncr,
49template <
typename Scalar>
50struct triangular_solve_packet_traits {
51 static constexpr bool Enabled = packet_traits<Scalar>::Vectorizable &&
52 (std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value) &&
53 std::numeric_limits<Scalar>::is_iec559 && std::numeric_limits<Scalar>::radix == 2;
54 static constexpr int PacketSize = packet_traits<Scalar>::size;
55 static constexpr int RegisterRows = 4;
56 static constexpr int NumberOfRegisters = gebp_traits<Scalar, Scalar>::NumberOfRegisters;
57 static constexpr int RhsPackets = plain_enum_max(
58 1, plain_enum_min((NumberOfRegisters - 2) / (RegisterRows + 1), NumberOfRegisters / (2 * RegisterRows)));
61 static constexpr int tile_rows(
int packets) {
return packets * 2 <= RhsPackets ? 2 * RegisterRows : RegisterRows; }
63 static constexpr int WorkspaceRows = 128;
64#if defined(EIGEN_VECTORIZE_AVX512) && EIGEN_USE_AVX512_TRSM_L_KERNELS
65 static constexpr bool UseUnblocked =
false;
67 static constexpr bool UseUnblocked = Enabled;
70 template <
typename Index>
71 static EIGEN_STRONG_INLINE
bool use_unblocked(Index size, Index cols, std::ptrdiff_t l1) {
72 if (!UseUnblocked || size < RegisterRows || size > WorkspaceRows || cols < PacketSize)
return false;
73 const int packets = cols >= RhsPackets * PacketSize ? RhsPackets : 1;
76 const std::ptrdiff_t rowBytes = (packets * PacketSize + RegisterRows + 1) *
sizeof(Scalar);
77 const std::ptrdiff_t transposeBytes = PacketSize * PacketSize *
sizeof(Scalar);
78 return std::ptrdiff_t(size) * rowBytes + transposeBytes <= l1 / 2;
82template <
typename Scalar,
typename Index,
int Mode,
int TriStorageOrder>
83struct triangular_solve_packet_kernel {
84 using Traits = triangular_solve_packet_traits<Scalar>;
85 using Packet =
typename packet_traits<Scalar>::type;
86 using TriMapper = const_blas_data_mapper<Scalar, Index, TriStorageOrder>;
87 static constexpr int PacketSize = Traits::PacketSize;
88 static constexpr bool IsLower = (Mode &
Lower) != 0;
90 template <
int RhsPackets>
91 static EIGEN_STRONG_INLINE
bool all_finite(
const PacketBlock<Packet, RhsPackets>& x, bool_constant<true>) {
93 using FloatTraits = binary_floating_point_traits<Scalar>;
94 bool nonfinite =
false;
95 for (
int p = 0; p < RhsPackets; ++p) {
96 Scalar values[PacketSize];
97 pstoreu(values, x.packet[p]);
98 EIGEN_FAST_MATH_CONSTANT_BARRIER(values);
99 for (
int c = 0; c < PacketSize; ++c)
100 nonfinite |= (FloatTraits::bits(values[c]) & FloatTraits::kExponentMask) == FloatTraits::kExponentMask;
106 template <
int RhsPackets>
107 static EIGEN_STRONG_INLINE
bool all_finite(
const PacketBlock<Packet, RhsPackets>&, bool_constant<false>) {
111 template <
int RhsPackets>
112 static EIGEN_STRONG_INLINE
void update(PacketBlock<Packet, RhsPackets>& x,
const PacketBlock<Packet, RhsPackets>& y,
114 const Packet pa = pset1<Packet>(a);
115 for (
int p = 0; p < RhsPackets; ++p) x.packet[p] = pnmadd(pa, y.packet[p], x.packet[p]);
118 template <
int RhsPackets>
119 static EIGEN_STRONG_INLINE
void scale(PacketBlock<Packet, RhsPackets>& x, Scalar a) {
120 EIGEN_IF_CONSTEXPR (!(Mode &
UnitDiag)) {
121 const Packet pa = pset1<Packet>(a);
122 for (
int p = 0; p < RhsPackets; ++p) x.packet[p] = pmul(x.packet[p], pa);
126 template <std::size_t Row,
int RhsPackets, std::size_t... Next>
127 static EIGEN_STRONG_INLINE
void solve_row(PacketBlock<Packet, RhsPackets>* x,
const TriMapper& a,
128 const Scalar* inverse, Index r0, Index step, std::index_sequence<Next...>) {
129 const Index row = r0 + Index(Row) * step;
130 scale(x[Row], inverse[row]);
131 int unroll[] = {0, (update(x[Row + 1 + Next], x[Row], a(r0 + Index(Row + 1 + Next) * step, row)), 0)...};
132 EIGEN_UNUSED_VARIABLE(unroll);
136 template <
int RhsPackets, std::size_t... Rows>
137 static EIGEN_STRONG_INLINE
void solve_block(Index i,
const TriMapper& a,
const Scalar* inverse,
138 PacketBlock<Packet, RhsPackets>* work, Index r0, Index step,
139 std::index_sequence<Rows...>) {
140 PacketBlock<Packet, RhsPackets> x[] = {work[r0 + Index(Rows) * step]...};
141 for (Index k = 0; k < i; ++k) {
142 const Index c = IsLower ? k : r0 + i - k;
143 const PacketBlock<Packet, RhsPackets> y = work[c];
144 int unroll[] = {0, (update(x[Rows], y, a(r0 + Index(Rows) * step, c)), 0)...};
145 EIGEN_UNUSED_VARIABLE(unroll);
149 0, (solve_row<Rows>(x, a, inverse, r0, step, std::make_index_sequence<
sizeof...(Rows) - Rows - 1>{}), 0)...};
150 EIGEN_UNUSED_VARIABLE(solve_rows);
151 int store_rows[] = {0, (work[r0 + Index(Rows) * step] = x[Rows], 0)...};
152 EIGEN_UNUSED_VARIABLE(store_rows);
155 template <
int RhsPackets>
156 static EIGEN_STRONG_INLINE
void solve(Index size,
const TriMapper& a,
const Scalar* inverse, Scalar* other,
158 PacketBlock<Packet, RhsPackets> work[Traits::WorkspaceRows];
159 for (
int p = 0; p < RhsPackets; ++p) {
161 for (; row + PacketSize <= size; row += PacketSize) {
162 PacketBlock<Packet, PacketSize> block;
163 for (
int c = 0; c < PacketSize; ++c)
164 block.packet[c] = ploadu<Packet>(other + row + (p * PacketSize + c) * otherStride);
166 for (
int r = 0; r < PacketSize; ++r) work[row + r].packet[p] = block.packet[r];
168 for (; row < size; ++row)
169 work[row].packet[p] = pgather<Scalar, Packet>(other + row + p * PacketSize * otherStride, otherStride);
172 const Index step = IsLower ? 1 : -1;
173 constexpr int TileRows = Traits::tile_rows(RhsPackets);
174 EIGEN_IF_CONSTEXPR (TileRows != Traits::RegisterRows) {
175 for (; i + TileRows <= size; i += TileRows)
176 solve_block(i, a, inverse, work, IsLower ? i : size - i - 1, step, std::make_index_sequence<TileRows>{});
178 for (; i + Traits::RegisterRows <= size; i += Traits::RegisterRows) {
179 const Index r0 = IsLower ? i : size - i - 1;
181 EIGEN_IF_CONSTEXPR (Traits::RegisterRows == 4) {
182 const Index r1 = r0 + step, r2 = r1 + step, r3 = r2 + step;
183 PacketBlock<Packet, RhsPackets> x0 = work[r0], x1 = work[r1], x2 = work[r2], x3 = work[r3];
184 for (Index k = 0; k < i; ++k) {
185 const Index c = IsLower ? k : size - k - 1;
186 const PacketBlock<Packet, RhsPackets> y = work[c];
187 update(x0, y, a(r0, c));
188 update(x1, y, a(r1, c));
189 update(x2, y, a(r2, c));
190 update(x3, y, a(r3, c));
192 scale(x0, inverse[r0]);
193 update(x1, x0, a(r1, r0));
194 update(x2, x0, a(r2, r0));
195 update(x3, x0, a(r3, r0));
196 scale(x1, inverse[r1]);
197 update(x2, x1, a(r2, r1));
198 update(x3, x1, a(r3, r1));
199 scale(x2, inverse[r2]);
200 update(x3, x2, a(r3, r2));
201 scale(x3, inverse[r3]);
207 solve_block(i, a, inverse, work, r0, step, std::make_index_sequence<Traits::RegisterRows>{});
210 for (; i < size; ++i) {
211 const Index r = IsLower ? i : size - i - 1;
212 PacketBlock<Packet, RhsPackets> x = work[r];
213 for (Index k = 0; k < i; ++k) {
214 const Index c = IsLower ? k : size - k - 1;
215 update(x, work[c], a(r, c));
217 scale(x, inverse[r]);
220 EIGEN_IF_CONSTEXPR (TriStorageOrder ==
RowMajor) {
224 if (!all_finite(work[IsLower ? size - 1 : 0], bool_constant<Traits::Enabled>{})) {
225 const Index origin = IsLower ? 0 : size - 1;
226 trsmKernelL<Scalar, Index, Mode, false, TriStorageOrder, 1, false>::kernel(
227 size, Index(RhsPackets * PacketSize), &a(origin, origin), a.stride(), other + origin, Index(1),
232 for (
int p = 0; p < RhsPackets; ++p) {
234 for (; row + PacketSize <= size; row += PacketSize) {
235 PacketBlock<Packet, PacketSize> block;
236 for (
int r = 0; r < PacketSize; ++r) block.packet[r] = work[row + r].packet[p];
238 for (
int c = 0; c < PacketSize; ++c) pstoreu(other + row + (p * PacketSize + c) * otherStride, block.packet[c]);
240 for (; row < size; ++row)
241 pscatter<Scalar, Packet>(other + row + p * PacketSize * otherStride, work[row].packet[p], otherStride);
249 static EIGEN_DONT_INLINE
void solve_padded(Index size,
const TriMapper& a,
const Scalar* inverse, Scalar* other,
250 Index otherStride, Index cols) {
251 Map<Matrix<Scalar, Dynamic, Dynamic, ColMajor>,
Unaligned, OuterStride<>> rest(other, size, cols,
252 OuterStride<>(otherStride));
253 Matrix<Scalar, Dynamic, Dynamic, ColMajor, Traits::WorkspaceRows, PacketSize> padded(size, Index(PacketSize));
254 padded.leftCols(cols) = rest;
255 padded.rightCols(PacketSize - cols).setZero();
256 solve<1>(size, a, inverse, padded.data(), size);
257 rest = padded.leftCols(cols);
260 static EIGEN_DONT_INLINE
void kernel(Index size, Index cols,
const Scalar* tri, Index triStride, Scalar* other,
262 eigen_internal_assert(size <= Traits::WorkspaceRows && cols >= PacketSize &&
263 (TriStorageOrder ==
RowMajor || cols % PacketSize == 0));
264 EIGEN_IF_CONSTEXPR (!IsLower) {
265 tri -= (size - 1) * (triStride + 1);
268 TriMapper a(tri, triStride);
269 Scalar inverse[Traits::WorkspaceRows];
270 Map<Vector<Scalar, Dynamic>> mapped(inverse, size);
271 EIGEN_IF_CONSTEXPR (Mode &
UnitDiag) {
274 const Map<const Vector<Scalar, Dynamic>,
Unaligned, InnerStride<Dynamic>> diagonal(
275 tri, size, InnerStride<Dynamic>(triStride + 1));
276 mapped = diagonal.cwiseInverse();
279 EIGEN_IF_CONSTEXPR (Traits::RhsPackets > 1) {
280 for (; j + Traits::RhsPackets * PacketSize <= cols; j += Traits::RhsPackets * PacketSize)
281 solve<Traits::RhsPackets>(size, a, inverse, other + j * otherStride, otherStride);
283 EIGEN_IF_CONSTEXPR (Traits::RhsPackets > 2) {
284 for (; j + 2 * PacketSize <= cols; j += 2 * PacketSize)
285 solve<2>(size, a, inverse, other + j * otherStride, otherStride);
287 for (; j + PacketSize <= cols; j += PacketSize) solve<1>(size, a, inverse, other + j * otherStride, otherStride);
288 EIGEN_IF_CONSTEXPR (TriStorageOrder ==
RowMajor) {
289 if (j < cols) solve_padded(size, a, inverse, other + j * otherStride, otherStride, cols - j);
294template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStride,
296EIGEN_STRONG_INLINE
void trsmKernelL<Scalar, Index, Mode, Conjugate, TriStorageOrder, OtherInnerStride,
297 Specialized>::kernel(Index size, Index otherSize,
const Scalar* _tri,
298 Index triStride, Scalar* _other, Index otherIncr,
300 EIGEN_IF_CONSTEXPR ((Specialized && OtherInnerStride == 1 && triangular_solve_packet_traits<Scalar>::Enabled)) {
301 if (size >= triangular_solve_packet_traits<Scalar>::RegisterRows &&
302 size <= triangular_solve_packet_traits<Scalar>::WorkspaceRows && otherSize >= packet_traits<Scalar>::size) {
304 const Index packetCols =
305 TriStorageOrder ==
RowMajor ? otherSize : numext::round_down(otherSize, Index(packet_traits<Scalar>::size));
306 triangular_solve_packet_kernel<Scalar, Index, Mode, TriStorageOrder>::kernel(size, packetCols, _tri, triStride,
307 _other, otherStride);
308 if (packetCols == otherSize)
return;
309 otherSize -= packetCols;
310 _other += packetCols * otherStride;
313 using TriMapper = const_blas_data_mapper<Scalar, Index, TriStorageOrder>;
314 using OtherMapper = blas_data_mapper<Scalar, Index, ColMajor, Unaligned, OtherInnerStride>;
315 TriMapper tri(_tri, triStride);
316 OtherMapper other(_other, otherStride, otherIncr);
319 conj_if<Conjugate> conj;
322 for (Index k = 0; k < size; ++k) {
324 Index i = IsLower ? k : -k;
325 Index rs = size - k - 1;
326 Index s = TriStorageOrder ==
RowMajor ? (IsLower ? 0 : i + 1) : IsLower ? i + 1 : i - rs;
328 Scalar a = (Mode &
UnitDiag) ? Scalar(1) : Scalar(Scalar(1) / conj(tri(i, i)));
329 for (Index j = 0; j < otherSize; ++j) {
330 EIGEN_IF_CONSTEXPR (TriStorageOrder ==
RowMajor) {
332 const Scalar* l = &tri(i, s);
333 typename OtherMapper::LinearMapper r = other.getLinearMapper(s, j);
334 for (Index i3 = 0; i3 < k; ++i3) b += conj(l[i3]) * r(i3);
336 other(i, j) = (other(i, j) - b) * a;
338 Scalar& otherij = other(i, j);
341 typename OtherMapper::LinearMapper r = other.getLinearMapper(s, j);
342 typename TriMapper::LinearMapper l = tri.getLinearMapper(s, i);
343 for (Index i3 = 0; i3 < rs; ++i3) r(i3) -= b * conj(l(i3));
349template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStride,
351EIGEN_STRONG_INLINE
void trsmKernelR<Scalar, Index, Mode, Conjugate, TriStorageOrder, OtherInnerStride,
352 Specialized>::kernel(Index size, Index otherSize,
const Scalar* _tri,
353 Index triStride, Scalar* _other, Index otherIncr,
355 using RealScalar =
typename NumTraits<Scalar>::Real;
356 using LhsMapper = blas_data_mapper<Scalar, Index, ColMajor, Unaligned, OtherInnerStride>;
357 using RhsMapper = const_blas_data_mapper<Scalar, Index, TriStorageOrder>;
358 LhsMapper lhs(_other, otherStride, otherIncr);
359 RhsMapper rhs(_tri, triStride);
362 conj_if<Conjugate> conj;
364 for (Index k = 0; k < size; ++k) {
365 Index j = IsLower ? size - k - 1 : k;
367 typename LhsMapper::LinearMapper r = lhs.getLinearMapper(0, j);
368 EIGEN_IF_CONSTEXPR (OtherInnerStride == 1 && packet_traits<Scalar>::Vectorizable) {
369 using Packet =
typename packet_traits<Scalar>::type;
370 constexpr Index PS = unpacket_traits<Packet>::size;
373 for (; k3 + 3 < k; k3 += 4) {
374 Index col0 = IsLower ? j + 1 + k3 : k3;
375 Scalar b0 = conj(rhs(col0, j));
376 Scalar b1 = conj(rhs(col0 + 1, j));
377 Scalar b2 = conj(rhs(col0 + 2, j));
378 Scalar b3 = conj(rhs(col0 + 3, j));
379 Packet neg_pb0 = pset1<Packet>(-b0);
380 Packet neg_pb1 = pset1<Packet>(-b1);
381 Packet neg_pb2 = pset1<Packet>(-b2);
382 Packet neg_pb3 = pset1<Packet>(-b3);
383 typename LhsMapper::LinearMapper a0 = lhs.getLinearMapper(0, col0);
384 typename LhsMapper::LinearMapper a1 = lhs.getLinearMapper(0, col0 + 1);
385 typename LhsMapper::LinearMapper a2 = lhs.getLinearMapper(0, col0 + 2);
386 typename LhsMapper::LinearMapper a3 = lhs.getLinearMapper(0, col0 + 3);
388 for (; i + PS <= otherSize; i += PS) {
389 Packet pr = r.template loadPacket<Packet>(i);
390 pr = pmadd(a0.template loadPacket<Packet>(i), neg_pb0, pr);
391 pr = pmadd(a1.template loadPacket<Packet>(i), neg_pb1, pr);
392 pr = pmadd(a2.template loadPacket<Packet>(i), neg_pb2, pr);
393 pr = pmadd(a3.template loadPacket<Packet>(i), neg_pb3, pr);
394 r.template storePacket<Packet>(i, pr);
396 for (; i < otherSize; ++i) {
397 r(i) -= a0(i) * b0 + a1(i) * b1 + a2(i) * b2 + a3(i) * b3;
401 for (; k3 < k; ++k3) {
402 Scalar b = conj(rhs(IsLower ? j + 1 + k3 : k3, j));
403 typename LhsMapper::LinearMapper a = lhs.getLinearMapper(0, IsLower ? j + 1 + k3 : k3);
404 Packet neg_pb = pset1<Packet>(-b);
406 for (; i + PS <= otherSize; i += PS) {
407 Packet pr = r.template loadPacket<Packet>(i);
408 pr = pmadd(a.template loadPacket<Packet>(i), neg_pb, pr);
409 r.template storePacket<Packet>(i, pr);
411 for (; i < otherSize; ++i) r(i) -= a(i) * b;
414 EIGEN_IF_CONSTEXPR ((Mode &
UnitDiag) == 0) {
415 Scalar inv_rjj = RealScalar(1) / conj(rhs(j, j));
416 Packet pinv = pset1<Packet>(inv_rjj);
418 for (; i + PS <= otherSize; i += PS) {
419 r.template storePacket<Packet>(i, pmul(r.template loadPacket<Packet>(i), pinv));
421 for (; i < otherSize; ++i) r(i) *= inv_rjj;
424 for (Index k3 = 0; k3 < k; ++k3) {
425 Scalar b = conj(rhs(IsLower ? j + 1 + k3 : k3, j));
426 typename LhsMapper::LinearMapper a = lhs.getLinearMapper(0, IsLower ? j + 1 + k3 : k3);
427 for (Index i = 0; i < otherSize; ++i) r(i) -= a(i) * b;
429 EIGEN_IF_CONSTEXPR ((Mode &
UnitDiag) == 0) {
430 Scalar inv_rjj = RealScalar(1) / conj(rhs(j, j));
431 for (Index i = 0; i < otherSize; ++i) r(i) *= inv_rjj;
438template <
typename Scalar,
typename Index,
int Side,
int Mode,
bool Conjugate,
int TriStorageOrder,
439 int OtherInnerStride>
440struct triangular_solve_matrix<Scalar, Index, Side, Mode, Conjugate, TriStorageOrder,
RowMajor, OtherInnerStride> {
441 static void run(Index size, Index cols,
const Scalar* tri, Index triStride, Scalar* _other, Index otherIncr,
442 Index otherStride, level3_blocking<Scalar, Scalar>& blocking) {
443 triangular_solve_matrix<
446 OtherInnerStride>::run(size, cols, tri, triStride, _other, otherIncr, otherStride, blocking);
453template <
typename Scalar>
454std::ptrdiff_t triangular_solve_budget(std::ptrdiff_t l2, std::ptrdiff_t l3) {
455 return (numext::maxi)(l3 / 4, l2) / std::ptrdiff_t(
sizeof(Scalar));
462template <
typename Index>
463Index triangular_solve_panel_columns(Index size, Index cols, std::ptrdiff_t budget, Index nr) {
464 eigen_internal_assert(size > 0);
465 const std::ptrdiff_t width = numext::round_down<std::ptrdiff_t>(budget / size, nr);
466 return width >= 512 && width < cols ? Index(width) : cols;
479template <
typename Scalar,
typename Index>
480Index triangular_solve_kc(Index size, Index otherSize, Index extent, std::ptrdiff_t budget,
bool slabRuns,
481 level3_blocking<Scalar, Scalar>& blocking) {
482 EIGEN_UNUSED_VARIABLE(extent);
483 const bool deep = std::ptrdiff_t(size) * otherSize > budget || (slabRuns && std::ptrdiff_t(size) * size / 2 > budget);
484 if (!deep || blocking.blockA() !=
nullptr)
return blocking.kc();
485 Index kc = size, mc = size, nc = otherSize;
486 computeProductBlockingSizes<Scalar, Scalar>(kc, mc, nc);
487 kc = (numext::mini)(kc, numext::round_down((numext::mini)(size / 8, Index(160)), Index(8)));
488#if defined(EIGEN_ALLOCA) && !defined(EIGEN_NO_ALLOCA)
489 const std::ptrdiff_t stackKc =
490 std::ptrdiff_t(EIGEN_STACK_ALLOCATION_LIMIT) / (std::ptrdiff_t(
sizeof(Scalar)) * extent);
491 if (blocking.kc() <= stackKc) kc = Index((numext::mini)(std::ptrdiff_t(kc), stackKc));
493 return (numext::maxi)(kc, blocking.kc());
498template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStr
ide>
499struct triangular_solve_matrix<Scalar, Index,
OnTheLeft, Mode, Conjugate, TriStorageOrder,
ColMajor, OtherInnerStride> {
500 static EIGEN_DONT_INLINE
void run(Index size, Index otherSize,
const Scalar* _tri, Index triStride, Scalar* _other,
501 Index otherIncr, Index otherStride, level3_blocking<Scalar, Scalar>& blocking);
504template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStr
ide>
505EIGEN_DONT_INLINE
void triangular_solve_matrix<Scalar, Index,
OnTheLeft, Mode, Conjugate, TriStorageOrder,
ColMajor,
506 OtherInnerStride>::run(Index size, Index otherSize,
const Scalar* _tri,
507 Index triStride, Scalar* _other, Index otherIncr,
509 level3_blocking<Scalar, Scalar>& blocking) {
510 std::ptrdiff_t l1, l2, l3;
511 manage_caching_sizes(GetAction, &l1, &l2, &l3);
512 EIGEN_IF_CONSTEXPR ((OtherInnerStride == 1 && triangular_solve_packet_traits<Scalar>::Enabled)) {
513 using PacketTraits = triangular_solve_packet_traits<Scalar>;
514 if (PacketTraits::use_unblocked(size, otherSize, l1)) {
515 const Index origin = (Mode &
Lower) ? 0 : size - 1;
516 trsmKernelL<Scalar, Index, Mode, Conjugate, TriStorageOrder, OtherInnerStride, true>::kernel(
517 size, otherSize, _tri + origin * (triStride + 1), triStride, _other + origin, otherIncr, otherStride);
521#if defined(EIGEN_VECTORIZE_AVX512) && defined(EIGEN_USE_AVX512_TRSM_L_KERNELS) && EIGEN_USE_AVX512_TRSM_L_KERNELS && \
522 EIGEN_ENABLE_AVX512_NOCOPY_TRSM_L_CUTOFFS
523 EIGEN_IF_CONSTEXPR ((OtherInnerStride == 1 &&
524 (std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value))) {
529 if (size < avx512_trsm_cutoff<Scalar>(l2, otherSize, L2Cap)) {
530 trsmKernelL<Scalar, Index, Mode, Conjugate, TriStorageOrder, 1,
true>::kernel(
531 size, otherSize, _tri, triStride, _other, 1, otherStride);
537 using TriMapper = const_blas_data_mapper<Scalar, Index, TriStorageOrder>;
538 using OtherMapper = blas_data_mapper<Scalar, Index, ColMajor, Unaligned, OtherInnerStride>;
539 TriMapper tri(_tri, triStride);
541 using Traits = gebp_traits<Scalar, Scalar>;
543 enum { SmallPanelWidth = plain_enum_max(Traits::mr, Traits::nr), IsLower = (Mode &
Lower) ==
Lower };
549 const std::ptrdiff_t budget = triangular_solve_budget<Scalar>(l2, l3);
550 const Index nc = triangular_solve_panel_columns(size, otherSize, budget, Index(Traits::nr));
551 const Index mc = (numext::mini)(size, blocking.mc());
554 const Index blockARows = (numext::maxi)(mc, Index(SmallPanelWidth));
556 const Index kc = triangular_solve_kc<Scalar>(size, otherSize, (numext::maxi)(blockARows, nc), budget,
559 std::size_t sizeA = kc * blockARows;
560 std::size_t sizeB = kc * nc;
562 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
563 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
565 gebp_kernel<Scalar, Scalar, Index, OtherMapper, Traits::mr, Traits::nr, Conjugate, false> gebp_kernel;
566 gemm_pack_lhs<Scalar, Index, TriMapper, Traits::mr, Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
569 gemm_pack_rhs<Scalar, Index, OtherMapper, Traits::nr, ColMajor, false, true> pack_rhs;
573 Index subcols = otherSize > 0 ? l2 / (4 *
sizeof(Scalar) * numext::maxi<Index>(otherStride, size)) : 0;
574 Index colStep = Traits::nr;
576 Index panelWidth = SmallPanelWidth;
577 EIGEN_IF_CONSTEXPR ((OtherInnerStride == 1 && triangular_solve_packet_traits<Scalar>::UseUnblocked)) {
583 using PacketTraits = triangular_solve_packet_traits<Scalar>;
584 const Index rows = Index(PacketTraits::RegisterRows), packet = Index(PacketTraits::PacketSize);
585 if (otherSize >= Index(plain_enum_min(2, PacketTraits::RhsPackets)) * packet) {
586 const Index panels = numext::div_ceil(kc, Index(PacketTraits::WorkspaceRows));
587 panelWidth = (numext::mini)(blockARows, rows * numext::div_ceil(numext::div_ceil(kc, panels), rows));
588 colStep = colStep % packet == 0 ? colStep : packet % colStep == 0 ? packet : colStep * packet;
589 subcols = Index(l2 / (2 * std::ptrdiff_t(
sizeof(Scalar)) * kc));
592 subcols = numext::maxi<Index>(numext::round_down(subcols, colStep), colStep);
594 for (Index j0 = 0; j0 < otherSize; j0 += nc) {
595 const Index cols = (numext::mini)(otherSize - j0, nc);
596 Scalar* _panel = _other + j0 * otherStride;
597 OtherMapper other(_panel, otherStride, otherIncr);
599 for (Index k2 = IsLower ? 0 : size; IsLower ? k2 < size : k2 > 0; IsLower ? k2 += kc : k2 -= kc) {
600 const Index actual_kc = (numext::mini)(IsLower ? size - k2 : k2, kc);
615 for (Index j2 = 0; j2 < cols; j2 += subcols) {
616 Index actual_cols = (numext::mini)(cols - j2, subcols);
618 for (Index k1 = 0; k1 < actual_kc; k1 += panelWidth) {
619 Index actualPanelWidth = numext::mini<Index>(actual_kc - k1, panelWidth);
622 Index i = IsLower ? k2 + k1 : k2 - k1 - 1;
623#if defined(EIGEN_VECTORIZE_AVX512) && defined(EIGEN_USE_AVX512_TRSM_L_KERNELS) && EIGEN_USE_AVX512_TRSM_L_KERNELS
624 EIGEN_IF_CONSTEXPR ((OtherInnerStride == 1 &&
625 (std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value))) {
626 i = IsLower ? k2 + k1 : k2 - k1 - actualPanelWidth;
629 using DiagonalKernel =
630 trsmKernelL<Scalar, Index, Mode, Conjugate, TriStorageOrder, OtherInnerStride,
true>;
631 const Scalar* diagonal = _tri + i + i * triStride;
632 Scalar* rhs = _panel + i * otherIncr + j2 * otherStride;
635 if (panelWidth == Index(SmallPanelWidth))
636 DiagonalKernel::kernel(numext::mini<Index>(actualPanelWidth, SmallPanelWidth), actual_cols, diagonal,
637 triStride, rhs, otherIncr, otherStride);
639 DiagonalKernel::kernel(actualPanelWidth, actual_cols, diagonal, triStride, rhs, otherIncr, otherStride);
642 Index lengthTarget = actual_kc - k1 - actualPanelWidth;
643 Index startBlock = IsLower ? k2 + k1 : k2 - k1 - actualPanelWidth;
644 Index blockBOffset = IsLower ? k1 : lengthTarget;
647 pack_rhs(blockB + actual_kc * j2, other.getSubMapper(startBlock, j2), actualPanelWidth, actual_cols,
648 actual_kc, blockBOffset);
651 if (lengthTarget > 0) {
652 Index startTarget = IsLower ? k2 + k1 + actualPanelWidth : k2 - actual_kc;
654 pack_lhs(blockA, tri.getSubMapper(startTarget, startBlock), actualPanelWidth, lengthTarget);
656 gebp_kernel(other.getSubMapper(startTarget, j2), blockA, blockB + actual_kc * j2, lengthTarget,
657 actualPanelWidth, actual_cols, Scalar(-1), actualPanelWidth, actual_kc, 0, blockBOffset);
664 Index start = IsLower ? k2 + kc : 0;
665 Index end = IsLower ? size : k2 - kc;
666 for (Index i2 = start; i2 < end; i2 += mc) {
667 const Index actual_mc = (numext::mini)(mc, end - i2);
669 pack_lhs(blockA, tri.getSubMapper(i2, IsLower ? k2 : k2 - kc), actual_kc, actual_mc);
671 gebp_kernel(other.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, cols, Scalar(-1), -1, -1, 0,
682template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStr
ide>
683struct triangular_solve_matrix<Scalar, Index,
OnTheRight, Mode, Conjugate, TriStorageOrder,
ColMajor,
685 static EIGEN_DONT_INLINE
void run(Index size, Index otherSize,
const Scalar* _tri, Index triStride, Scalar* _other,
686 Index otherIncr, Index otherStride, level3_blocking<Scalar, Scalar>& blocking);
689template <
typename Scalar,
typename Index,
int Mode,
bool Conjugate,
int TriStorageOrder,
int OtherInnerStr
ide>
690EIGEN_DONT_INLINE
void triangular_solve_matrix<Scalar, Index,
OnTheRight, Mode, Conjugate, TriStorageOrder,
ColMajor,
691 OtherInnerStride>::run(Index size, Index otherSize,
const Scalar* _tri,
692 Index triStride, Scalar* _other, Index otherIncr,
694 level3_blocking<Scalar, Scalar>& blocking) {
695 Index rows = otherSize;
697 std::ptrdiff_t l1, l2, l3;
698 manage_caching_sizes(GetAction, &l1, &l2, &l3);
700#if defined(EIGEN_VECTORIZE_AVX512) && defined(EIGEN_USE_AVX512_TRSM_R_KERNELS) && EIGEN_USE_AVX512_TRSM_R_KERNELS && \
701 EIGEN_ENABLE_AVX512_NOCOPY_TRSM_R_CUTOFFS
702 EIGEN_IF_CONSTEXPR ((OtherInnerStride == 1 &&
703 (std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value))) {
706 if (size < avx512_trsm_cutoff<Scalar>(l2, rows, L2Cap)) {
707 trsmKernelR<Scalar, Index, Mode, Conjugate, TriStorageOrder, OtherInnerStride,
true>::kernel(
708 size, rows, _tri, triStride, _other, 1, otherStride);
714 using LhsMapper = blas_data_mapper<Scalar, Index, ColMajor, Unaligned, OtherInnerStride>;
715 using RhsMapper = const_blas_data_mapper<Scalar, Index, TriStorageOrder>;
716 LhsMapper lhs(_other, otherStride, otherIncr);
717 RhsMapper rhs(_tri, triStride);
719 using Traits = gebp_traits<Scalar, Scalar>;
721 RhsStorageOrder = TriStorageOrder,
722 SmallPanelWidth = plain_enum_max(Traits::mr, Traits::nr),
729 const std::ptrdiff_t budget = triangular_solve_budget<Scalar>(l2, l3);
730 Index mc = (numext::mini)(rows, blocking.mc());
732 const Index kc = triangular_solve_kc<Scalar>(size, rows, (numext::maxi)(mc, size), budget,
733 TriStorageOrder ==
ColMajor, blocking);
737 if (kc > blocking.kc()) {
738 const std::ptrdiff_t maxA = (numext::maxi)(std::ptrdiff_t(blocking.kc()) * mc, budget / 2);
739 mc = (numext::mini)(mc, (numext::maxi)(Index(Traits::mr), numext::round_down(Index(maxA / kc), Index(Traits::mr))));
742 std::size_t sizeA = kc * mc;
743 std::size_t sizeB = kc * size;
745 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
746 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
748 gebp_kernel<Scalar, Scalar, Index, LhsMapper, Traits::mr, Traits::nr, false, Conjugate> gebp_kernel;
749 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
750 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder, false, true> pack_rhs_panel;
751 gemm_pack_lhs<Scalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
ColMajor,
755 for (Index k2 = IsLower ? size : 0; IsLower ? k2 > 0 : k2 < size; IsLower ? k2 -= kc : k2 += kc) {
756 const Index actual_kc = (numext::mini)(IsLower ? k2 : size - k2, kc);
757 Index actual_k2 = IsLower ? k2 - actual_kc : k2;
759 Index startPanel = IsLower ? 0 : k2 + actual_kc;
760 Index rs = IsLower ? actual_k2 : size - actual_k2 - actual_kc;
761 Scalar* geb = blockB + actual_kc * actual_kc;
763 if (rs > 0) pack_rhs(geb, rhs.getSubMapper(actual_k2, startPanel), actual_kc, rs);
768 for (Index j2 = 0; j2 < actual_kc; j2 += SmallPanelWidth) {
769 Index actualPanelWidth = numext::mini<Index>(actual_kc - j2, SmallPanelWidth);
770 Index actual_j2 = actual_k2 + j2;
771 Index panelOffset = IsLower ? j2 + actualPanelWidth : 0;
772 Index panelLength = IsLower ? actual_kc - j2 - actualPanelWidth : j2;
775 pack_rhs_panel(blockB + j2 * actual_kc, rhs.getSubMapper(actual_k2 + panelOffset, actual_j2), panelLength,
776 actualPanelWidth, actual_kc, panelOffset);
780 for (Index i2 = 0; i2 < rows; i2 += mc) {
781 const Index actual_mc = (numext::mini)(mc, rows - i2);
786 for (Index j2 = IsLower ? (actual_kc - ((actual_kc % SmallPanelWidth) ? Index(actual_kc % SmallPanelWidth)
787 : Index(SmallPanelWidth)))
789 IsLower ? j2 >= 0 : j2 < actual_kc; IsLower ? j2 -= SmallPanelWidth : j2 += SmallPanelWidth) {
790 Index actualPanelWidth = numext::mini<Index>(actual_kc - j2, SmallPanelWidth);
791 Index absolute_j2 = actual_k2 + j2;
792 Index panelOffset = IsLower ? j2 + actualPanelWidth : 0;
793 Index panelLength = IsLower ? actual_kc - j2 - actualPanelWidth : j2;
796 if (panelLength > 0) {
797 gebp_kernel(lhs.getSubMapper(i2, absolute_j2), blockA, blockB + j2 * actual_kc, actual_mc, panelLength,
798 actualPanelWidth, Scalar(-1), actual_kc, actual_kc,
799 panelOffset, panelOffset);
804 trsmKernelR<Scalar, Index, Mode, Conjugate, TriStorageOrder, OtherInnerStride,
805 true>::kernel(actualPanelWidth, actual_mc,
806 _tri + absolute_j2 + absolute_j2 * triStride, triStride,
807 _other + i2 * otherIncr + absolute_j2 * otherStride, otherIncr,
811 pack_lhs_panel(blockA, lhs.getSubMapper(i2, absolute_j2), actualPanelWidth, actual_mc, actual_kc, j2);
816 gebp_kernel(lhs.getSubMapper(i2, startPanel), blockA, geb, actual_mc, actual_kc, rs, Scalar(-1), -1, -1, 0, 0);
@ UnitDiag
Definition Constants.h:216
@ Lower
Definition Constants.h:212
@ Upper
Definition Constants.h:214
@ Unaligned
Definition Constants.h:236
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
@ OnTheLeft
Definition Constants.h:332
@ OnTheRight
Definition Constants.h:334