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// Copyright (C) 2024 Kseniya Zaytseva <kseniya.zaytseva@syntacore.com>
5// Copyright (C) 2025 Chip Kerchner <ckerchner@tenstorrent.com>
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_RVV10_GENERAL_BLOCK_KERNEL_H
13#define EIGEN_RVV10_GENERAL_BLOCK_KERNEL_H
14
15// IWYU pragma: private
16#include "../../InternalHeaderCheck.h"
17
18namespace Eigen {
19namespace internal {
20
21/********************************* real ************************************/
22
23template <>
24struct gebp_traits<float, float, false, false, Architecture::RVV10, GEBPPacketFull>
25 : gebp_traits<float, float, false, false, Architecture::Generic, GEBPPacketFull> {
26 typedef float RhsPacket;
27 typedef QuadPacket<float> RhsPacketx4;
28 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacket& dest) const { dest = pset1<RhsPacket>(*b); }
29 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
30 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
31 }
32
33 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, RhsPacket& dest) const { loadRhs(b, dest); }
34
35 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
36
37 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, RhsPacket& dest) const { dest = ploadquad<RhsPacket>(b); }
38
39 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacket& b, AccPacket& c, RhsPacket& /*tmp*/,
40 const FixedInt<0>&) const {
41#if EIGEN_RISCV64_DEFAULT_LMUL == 1
42 c = __riscv_vfmadd_vf_f32m1(a, b, c, unpacket_traits<AccPacket>::size);
43#elif EIGEN_RISCV64_DEFAULT_LMUL == 2
44 c = __riscv_vfmadd_vf_f32m2(a, b, c, unpacket_traits<AccPacket>::size);
45#elif EIGEN_RISCV64_DEFAULT_LMUL == 4
46 c = __riscv_vfmadd_vf_f32m4(a, b, c, unpacket_traits<AccPacket>::size);
47#endif
48 }
49
50#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
51 EIGEN_STRONG_INLINE void madd(const Packet1Xf& a, const RhsPacket& b, Packet1Xf& c, RhsPacket& /*tmp*/,
52 const FixedInt<0>&) const {
53 c = __riscv_vfmadd_vf_f32m1(a, b, c, unpacket_traits<Packet1Xf>::size);
54 }
55#endif
56#if EIGEN_RISCV64_DEFAULT_LMUL == 4
57 EIGEN_STRONG_INLINE void madd(const Packet2Xf& a, const RhsPacket& b, Packet2Xf& c, RhsPacket& /*tmp*/,
58 const FixedInt<0>&) const {
59 c = __riscv_vfmadd_vf_f32m2(a, b, c, unpacket_traits<Packet2Xf>::size);
60 }
61#endif
62
63 template <typename LaneIdType>
64 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacketx4& b, AccPacket& c, RhsPacket& /*tmp*/,
65 const LaneIdType& lane) const {
66#if EIGEN_RISCV64_DEFAULT_LMUL == 1
67 c = __riscv_vfmadd_vf_f32m1(a, b.get(lane), c, unpacket_traits<AccPacket>::size);
68#elif EIGEN_RISCV64_DEFAULT_LMUL == 2
69 c = __riscv_vfmadd_vf_f32m2(a, b.get(lane), c, unpacket_traits<AccPacket>::size);
70#elif EIGEN_RISCV64_DEFAULT_LMUL == 4
71 c = __riscv_vfmadd_vf_f32m4(a, b.get(lane), c, unpacket_traits<AccPacket>::size);
72#endif
73 }
74};
75
76template <>
77struct gebp_traits<double, double, false, false, Architecture::RVV10, GEBPPacketFull>
78 : gebp_traits<double, double, false, false, Architecture::Generic, GEBPPacketFull> {
79 typedef double RhsPacket;
80 typedef QuadPacket<double> RhsPacketx4;
81 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacket& dest) const { dest = pset1<RhsPacket>(*b); }
82 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
83 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
84 }
85
86 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, RhsPacket& dest) const { loadRhs(b, dest); }
87
88 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
89
90 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, RhsPacket& dest) const { dest = ploadquad<RhsPacket>(b); }
91
92 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacket& b, AccPacket& c, RhsPacket& /*tmp*/,
93 const FixedInt<0>&) const {
94#if EIGEN_RISCV64_DEFAULT_LMUL == 1
95 c = __riscv_vfmadd_vf_f64m1(a, b, c, unpacket_traits<AccPacket>::size);
96#elif EIGEN_RISCV64_DEFAULT_LMUL == 2
97 c = __riscv_vfmadd_vf_f64m2(a, b, c, unpacket_traits<AccPacket>::size);
98#elif EIGEN_RISCV64_DEFAULT_LMUL == 4
99 c = __riscv_vfmadd_vf_f64m4(a, b, c, unpacket_traits<AccPacket>::size);
100#endif
101 }
102
103#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
104 EIGEN_STRONG_INLINE void madd(const Packet1Xd& a, const RhsPacket& b, Packet1Xd& c, RhsPacket& /*tmp*/,
105 const FixedInt<0>&) const {
106 c = __riscv_vfmadd_vf_f64m1(a, b, c, unpacket_traits<Packet1Xd>::size);
107 }
108#endif
109#if EIGEN_RISCV64_DEFAULT_LMUL == 4
110 EIGEN_STRONG_INLINE void madd(const Packet2Xd& a, const RhsPacket& b, Packet2Xd& c, RhsPacket& /*tmp*/,
111 const FixedInt<0>&) const {
112 c = __riscv_vfmadd_vf_f64m2(a, b, c, unpacket_traits<Packet2Xd>::size);
113 }
114#endif
115
116 template <typename LaneIdType>
117 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacketx4& b, AccPacket& c, RhsPacket& /*tmp*/,
118 const LaneIdType& lane) const {
119#if EIGEN_RISCV64_DEFAULT_LMUL == 1
120 c = __riscv_vfmadd_vf_f64m1(a, b.get(lane), c, unpacket_traits<AccPacket>::size);
121#elif EIGEN_RISCV64_DEFAULT_LMUL == 2
122 c = __riscv_vfmadd_vf_f64m2(a, b.get(lane), c, unpacket_traits<AccPacket>::size);
123#elif EIGEN_RISCV64_DEFAULT_LMUL == 4
124 c = __riscv_vfmadd_vf_f64m4(a, b.get(lane), c, unpacket_traits<AccPacket>::size);
125#endif
126 }
127};
128
129#if defined(EIGEN_VECTORIZE_RVV10FP16)
130
131template <>
132struct gebp_traits<half, half, false, false, Architecture::RVV10>
133 : gebp_traits<half, half, false, false, Architecture::Generic> {
134 typedef half RhsPacket;
135 typedef PacketXh LhsPacket;
136 typedef PacketXh AccPacket;
137 typedef QuadPacket<half> RhsPacketx4;
138
139 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacket& dest) const { dest = pset1<RhsPacket>(*b); }
140 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
141 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
142 }
143
144 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, RhsPacket& dest) const { loadRhs(b, dest); }
145
146 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
147
148 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, RhsPacket& dest) const { dest = pload<RhsPacket>(b); }
149
150 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacket& b, AccPacket& c, RhsPacket& /*tmp*/,
151 const FixedInt<0>&) const {
152#if EIGEN_RISCV64_DEFAULT_LMUL == 1
153 c = __riscv_vfmadd_vf_f16m1(a, numext::bit_cast<_Float16>(b), c, unpacket_traits<AccPacket>::size);
154#else
155 c = __riscv_vfmadd_vf_f16m2(a, numext::bit_cast<_Float16>(b), c, unpacket_traits<AccPacket>::size);
156#endif
157 }
158
159#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
160 EIGEN_STRONG_INLINE void madd(const Packet1Xh& a, const RhsPacket& b, Packet1Xh& c, RhsPacket& /*tmp*/,
161 const FixedInt<0>&) const {
162 c = __riscv_vfmadd_vf_f16m1(a, numext::bit_cast<_Float16>(b), c, unpacket_traits<Packet1Xh>::size);
163 }
164#endif
165
166 template <typename LaneIdType>
167 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacketx4& b, AccPacket& c, RhsPacket& /*tmp*/,
168 const LaneIdType& lane) const {
169#if EIGEN_RISCV64_DEFAULT_LMUL == 1
170 c = __riscv_vfmadd_vf_f16m1(a, numext::bit_cast<_Float16>(b.get(lane)), c, unpacket_traits<AccPacket>::size);
171#else
172 c = __riscv_vfmadd_vf_f16m2(a, numext::bit_cast<_Float16>(b.get(lane)), c, unpacket_traits<AccPacket>::size);
173#endif
174 }
175};
176
177#endif
178
179#if defined(EIGEN_VECTORIZE_RVV10BF16)
180
181template <>
182struct gebp_traits<bfloat16, bfloat16, false, false, Architecture::RVV10>
183 : gebp_traits<bfloat16, bfloat16, false, false, Architecture::Generic> {
184 typedef bfloat16 RhsPacket;
185 typedef PacketXbf LhsPacket;
186 typedef PacketXbf AccPacket;
187 typedef QuadPacket<bfloat16> RhsPacketx4;
188
189 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacket& dest) const { dest = pset1<RhsPacket>(*b); }
190 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
191 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
192 }
193
194 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, RhsPacket& dest) const { loadRhs(b, dest); }
195
196 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
197
198 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, RhsPacket& dest) const { dest = pload<RhsPacket>(b); }
199
200 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacket& b, AccPacket& c, RhsPacket& /*tmp*/,
201 const FixedInt<0>&) const {
202#if EIGEN_RISCV64_DEFAULT_LMUL == 1
203 c = F32ToBf16(
204 __riscv_vfwmaccbf16_vf_f32m2(Bf16ToF32(c), numext::bit_cast<__bf16>(b), a, unpacket_traits<AccPacket>::size));
205#else
206 c = F32ToBf16(
207 __riscv_vfwmaccbf16_vf_f32m4(Bf16ToF32(c), numext::bit_cast<__bf16>(b), a, unpacket_traits<AccPacket>::size));
208#endif
209 }
210
211#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
212 EIGEN_STRONG_INLINE void madd(const Packet1Xbf& a, const RhsPacket& b, Packet1Xbf& c, RhsPacket& /*tmp*/,
213 const FixedInt<0>&) const {
214 c = F32ToBf16(
215 __riscv_vfwmaccbf16_vf_f32m2(Bf16ToF32(c), numext::bit_cast<__bf16>(b), a, unpacket_traits<Packet1Xbf>::size));
216 }
217#endif
218
219 template <typename LaneIdType>
220 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacketx4& b, AccPacket& c, RhsPacket& /*tmp*/,
221 const LaneIdType& lane) const {
222#if EIGEN_RISCV64_DEFAULT_LMUL == 1
223 c = F32ToBf16(__riscv_vfwmaccbf16_vf_f32m2(Bf16ToF32(c), numext::bit_cast<__bf16>(b.get(lane)), a,
224 unpacket_traits<AccPacket>::size));
225#else
226 c = F32ToBf16(__riscv_vfwmaccbf16_vf_f32m4(Bf16ToF32(c), numext::bit_cast<__bf16>(b.get(lane)), a,
227 unpacket_traits<AccPacket>::size));
228#endif
229 }
230};
231
232#endif
233
234#if EIGEN_RISCV64_DEFAULT_LMUL == 1
235/********************************* complex ************************************/
236
237#define PACKET_DECL_COND_POSTFIX(postfix, name, packet_size) \
238 typedef typename packet_conditional< \
239 packet_size, typename packet_traits<name##Scalar>::type, typename packet_traits<name##Scalar>::half, \
240 typename unpacket_traits<typename packet_traits<name##Scalar>::half>::half>::type name##Packet##postfix
241
242#define RISCV_COMPLEX_PACKET_DECL_COND_SCALAR(packet_size) \
243 typedef typename packet_conditional< \
244 packet_size, typename packet_traits<Scalar>::type, typename packet_traits<Scalar>::half, \
245 typename unpacket_traits<typename packet_traits<Scalar>::half>::half>::type ScalarPacket
246
247template <typename RealScalar, bool ConjLhs_, bool ConjRhs_, int PacketSize_>
248struct gebp_traits<std::complex<RealScalar>, std::complex<RealScalar>, ConjLhs_, ConjRhs_, Architecture::RVV10,
249 PacketSize_> : gebp_traits<std::complex<RealScalar>, std::complex<RealScalar>, ConjLhs_, ConjRhs_,
250 Architecture::Generic, PacketSize_> {
251 typedef std::complex<RealScalar> Scalar;
252 typedef std::complex<RealScalar> LhsScalar;
253 typedef std::complex<RealScalar> RhsScalar;
254 typedef std::complex<RealScalar> ResScalar;
255 typedef typename packet_traits<std::complex<RealScalar>>::type RealPacket;
256
257 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
258 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
259 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
260 RISCV_COMPLEX_PACKET_DECL_COND_SCALAR(PacketSize_);
261#undef RISCV_COMPLEX_PACKET_DECL_COND_SCALAR
262
263 enum {
264 ConjLhs = ConjLhs_,
265 ConjRhs = ConjRhs_,
266 Vectorizable = unpacket_traits<RealPacket>::vectorizable && unpacket_traits<ScalarPacket>::vectorizable,
267 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
268 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
269 RhsPacketSize = Vectorizable ? unpacket_traits<RhsScalar>::size : 1,
270 RealPacketSize = Vectorizable ? unpacket_traits<RealPacket>::size : 1,
271
272 nr = 4,
273 mr = ResPacketSize,
274
275 LhsProgress = ResPacketSize,
276 RhsProgress = 1
277 };
278
279 typedef DoublePacket<RealPacket> DoublePacketType;
280
281 typedef std::conditional_t<Vectorizable, ScalarPacket, Scalar> LhsPacket4Packing;
282 typedef std::conditional_t<Vectorizable, RealPacket, Scalar> LhsPacket;
283 typedef std::conditional_t<Vectorizable, DoublePacket<RealScalar>, Scalar> RhsPacket;
284 typedef std::conditional_t<Vectorizable, ScalarPacket, Scalar> ResPacket;
285 typedef std::conditional_t<Vectorizable, DoublePacketType, Scalar> AccPacket;
286
287 typedef QuadPacket<RhsPacket> RhsPacketx4;
288
289 EIGEN_STRONG_INLINE void initAcc(Scalar& p) { p = Scalar(0); }
290
291 EIGEN_STRONG_INLINE void initAcc(DoublePacketType& p) {
292 p.first = pset1<RealPacket>(RealScalar(0));
293 p.second = pset1<RealPacket>(RealScalar(0));
294 }
295
296 // Scalar path
297 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, ScalarPacket& dest) const { dest = pset1<ScalarPacket>(*b); }
298
299 // Vectorized path
300 template <typename RealPacketType>
301 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, DoublePacket<RealPacketType>& dest) const {
302 dest.first = pset1<RealPacketType>(numext::real(*b));
303 dest.second = pset1<RealPacketType>(numext::imag(*b));
304 }
305
306 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
307 loadRhs(b, dest.B_0);
308 loadRhs(b + 1, dest.B1);
309 loadRhs(b + 2, dest.B2);
310 loadRhs(b + 3, dest.B3);
311 }
312
313 // Scalar path
314 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, ScalarPacket& dest) const { loadRhs(b, dest); }
315
316 // Vectorized path
317 template <typename RealPacketType>
318 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, DoublePacket<RealPacketType>& dest) const {
319 loadRhs(b, dest);
320 }
321
322 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
323
324 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, ResPacket& dest) const { loadRhs(b, dest); }
325 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, DoublePacket<RealScalar>& dest) const {
326 loadQuadToDoublePacket(b, dest);
327 }
328
329 // nothing special here
330 EIGEN_STRONG_INLINE void loadLhs(const LhsScalar* a, LhsPacket& dest) const {
331 dest = pload<LhsPacket>((const typename unpacket_traits<LhsPacket>::type*)(a));
332 }
333
334 template <typename LhsPacketType>
335 EIGEN_STRONG_INLINE void loadLhsUnaligned(const LhsScalar* a, LhsPacketType& dest) const {
336 dest = ploadu<LhsPacketType>((const typename unpacket_traits<LhsPacketType>::type*)(a));
337 }
338
339 EIGEN_STRONG_INLINE Packet1Xcf pmadd_scalar(const Packet1Xcf& a, float b, const Packet1Xcf& c) const {
340 return Packet1Xcf(__riscv_vfmadd_vf_f32m1(a.v, b, c.v, unpacket_traits<Packet1Xf>::size));
341 }
342
343 EIGEN_STRONG_INLINE Packet1Xcd pmadd_scalar(const Packet1Xcd& a, double b, const Packet1Xcd& c) const {
344 return Packet1Xcd(__riscv_vfmadd_vf_f64m1(a.v, b, c.v, unpacket_traits<Packet1Xd>::size));
345 }
346
347#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
348 EIGEN_STRONG_INLINE Packet2Xcf pmadd_scalar(const Packet2Xcf& a, float b, const Packet2Xcf& c) const {
349 return Packet2Xcf(__riscv_vfmadd_vf_f32m2(a.v, b, c.v, unpacket_traits<Packet2Xf>::size));
350 }
351
352 EIGEN_STRONG_INLINE Packet2Xcd pmadd_scalar(const Packet2Xcd& a, double b, const Packet2Xcd& c) const {
353 return Packet2Xcd(__riscv_vfmadd_vf_f64m2(a.v, b, c.v, unpacket_traits<Packet2Xd>::size));
354 }
355#endif
356
357#if EIGEN_RISCV64_DEFAULT_LMUL == 4
358 EIGEN_STRONG_INLINE Packet4Xcf pmadd_scalar(const Packet4Xcf& a, float b, const Packet4Xcf& c) const {
359 return Packet4Xcf(__riscv_vfmadd_vf_f32m4(a.v, b, c.v, unpacket_traits<Packet4Xf>::size));
360 }
361
362 EIGEN_STRONG_INLINE Packet4Xcd pmadd_scalar(const Packet4Xcd& a, double b, const Packet4Xcd& c) const {
363 return Packet4Xcd(__riscv_vfmadd_vf_f64m4(a.v, b, c.v, unpacket_traits<Packet4Xd>::size));
364 }
365#endif
366
367 template <typename LhsPacketType, typename RhsPacketType, typename ResPacketType, typename TmpType,
368 typename LaneIdType>
369 EIGEN_STRONG_INLINE std::enable_if_t<!is_same<RhsPacketType, RhsPacketx4>::value> madd(const LhsPacketType& a,
370 const RhsPacketType& b,
371 DoublePacket<ResPacketType>& c,
372 TmpType& /*tmp*/,
373 const LaneIdType&) const {
374 c.first = pmadd_scalar(a, b.first, c.first);
375 c.second = pmadd_scalar(a, b.second, c.second);
376 }
377
378 template <typename LhsPacketType, typename AccPacketType, typename LaneIdType>
379 EIGEN_STRONG_INLINE void madd(const LhsPacketType& a, const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
380 const LaneIdType& lane) const {
381 madd(a, b.get(lane), c, tmp, lane);
382 }
383
384 template <typename LaneIdType>
385 EIGEN_STRONG_INLINE void madd(const LhsPacket& a, const RhsPacket& b, ResPacket& c, RhsPacket& /*tmp*/,
386 const LaneIdType&) const {
387 c = cj.pmadd(a, b, c);
388 }
389
390 template <typename RealPacketType, typename ResPacketType>
391 EIGEN_STRONG_INLINE void acc(const DoublePacket<RealPacketType>& c, const ResPacketType& alpha,
392 ResPacketType& r) const {
393 // assemble c
394 ResPacketType tmp;
395 EIGEN_IF_CONSTEXPR ((!ConjLhs) && (!ConjRhs)) {
396 tmp = pcplxflip(pconj(ResPacketType(c.second)));
397 tmp = padd(ResPacketType(c.first), tmp);
398 } else EIGEN_IF_CONSTEXPR ((!ConjLhs) && (ConjRhs)) {
399 tmp = pconj(pcplxflip(ResPacketType(c.second)));
400 tmp = padd(ResPacketType(c.first), tmp);
401 } else EIGEN_IF_CONSTEXPR ((ConjLhs) && (!ConjRhs)) {
402 tmp = pcplxflip(ResPacketType(c.second));
403 tmp = padd(pconj(ResPacketType(c.first)), tmp);
404 } else EIGEN_IF_CONSTEXPR ((ConjLhs) && (ConjRhs)) {
405 tmp = pcplxflip(ResPacketType(c.second));
406 tmp = psub(pconj(ResPacketType(c.first)), tmp);
407 }
408
409 r = pmadd(tmp, alpha, r);
410 }
411
412 // madd() already applied the operand conjugations; neither c nor alpha is conjugated here.
413 EIGEN_STRONG_INLINE void acc(const Scalar& c, const Scalar& alpha, Scalar& r) const { r += alpha * c; }
414
415 protected:
416 conj_helper<LhsScalar, RhsScalar, ConjLhs, ConjRhs> cj;
417};
418
419#define PACKET_DECL_COND_SCALAR_POSTFIX(postfix, packet_size) \
420 typedef typename packet_conditional< \
421 packet_size, typename packet_traits<Scalar>::type, typename packet_traits<Scalar>::half, \
422 typename unpacket_traits<typename packet_traits<Scalar>::half>::half>::type ScalarPacket##postfix
423
424template <typename RealScalar, bool ConjRhs_, int PacketSize_>
425class gebp_traits<RealScalar, std::complex<RealScalar>, false, ConjRhs_, Architecture::RVV10, PacketSize_>
426 : public gebp_traits<RealScalar, std::complex<RealScalar>, false, ConjRhs_, Architecture::Generic, PacketSize_> {
427 public:
428 typedef std::complex<RealScalar> Scalar;
429 typedef RealScalar LhsScalar;
430 typedef Scalar RhsScalar;
431 typedef Scalar ResScalar;
432 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
433 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
434 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
435 PACKET_DECL_COND_POSTFIX(_, Real, PacketSize_);
436 PACKET_DECL_COND_SCALAR_POSTFIX(_, PacketSize_);
437#undef PACKET_DECL_COND_SCALAR_POSTFIX
438
439 enum {
440 ConjLhs = false,
441 ConjRhs = ConjRhs_,
442 Vectorizable = unpacket_traits<RealPacket_>::vectorizable && unpacket_traits<ScalarPacket_>::vectorizable,
443 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
444 RhsPacketSize = Vectorizable ? unpacket_traits<RhsPacket_>::size : 1,
445 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
446
447 NumberOfRegisters = EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS,
448 // FIXME: should depend on NumberOfRegisters
449 nr = 4,
450 mr = (plain_enum_min(16, NumberOfRegisters) / 2 / nr) * ResPacketSize,
451
452 LhsProgress = ResPacketSize,
453 RhsProgress = 1
454 };
455
456 typedef std::conditional_t<Vectorizable, LhsPacket_, LhsScalar> LhsPacket;
457 typedef RhsScalar RhsPacket;
458 typedef std::conditional_t<Vectorizable, ResPacket_, ResScalar> ResPacket;
459 typedef LhsPacket LhsPacket4Packing;
460 typedef QuadPacket<RhsPacket> RhsPacketx4;
461 typedef ResPacket AccPacket;
462
463 EIGEN_STRONG_INLINE void initAcc(AccPacket& p) const { p = pset1<ResPacket>(ResScalar(0)); }
464
465 template <typename RhsPacketType>
466 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketType& dest) const {
467 dest = pset1<RhsPacketType>(*b);
468 }
469
470 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
471 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
472 }
473
474 template <typename RhsPacketType>
475 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, RhsPacketType& dest) const {
476 loadRhs(b, dest);
477 }
478
479 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
480
481 EIGEN_STRONG_INLINE void loadLhs(const LhsScalar* a, LhsPacket& dest) const { dest = ploaddup<LhsPacket>(a); }
482
483 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, RhsPacket& dest) const { dest = ploadquad<RhsPacket>(b); }
484
485 template <typename LhsPacketType>
486 EIGEN_STRONG_INLINE void loadLhsUnaligned(const LhsScalar* a, LhsPacketType& dest) const {
487 dest = ploaddup<LhsPacketType>(a);
488 }
489
490 template <typename LhsPacketType, typename RhsPacketType, typename AccPacketType, typename LaneIdType>
491 EIGEN_STRONG_INLINE void madd(const LhsPacketType& a, const RhsPacketType& b, AccPacketType& c, RhsPacketType& tmp,
492 const LaneIdType&) const {
493 madd_impl(a, b, c, tmp, std::conditional_t<Vectorizable, std::true_type, std::false_type>());
494 }
495
496 EIGEN_STRONG_INLINE Packet1Xcf pmadd_scalar(const Packet1Xf& a, std::complex<float> b, const Packet1Xcf& c) const {
497 Packet1Xcf bp = pset1<Packet1Xcf>(b);
498 return Packet1Xcf(__riscv_vfmadd_vv_f32m1(a, bp.v, c.v, unpacket_traits<Packet1Xf>::size));
499 }
500
501 EIGEN_STRONG_INLINE Packet1Xcd pmadd_scalar(const Packet1Xd& a, std::complex<double> b, const Packet1Xcd& c) const {
502 Packet1Xcd bp = pset1<Packet1Xcd>(b);
503 return Packet1Xcd(__riscv_vfmadd_vv_f64m1(a, bp.v, c.v, unpacket_traits<Packet1Xd>::size));
504 }
505
506#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
507 EIGEN_STRONG_INLINE Packet2Xcf pmadd_scalar(const Packet2Xf& a, std::complex<float> b, const Packet2Xcf& c) const {
508 Packet2Xcf bp = pset1<Packet2Xcf>(b);
509 return Packet2Xcf(__riscv_vfmadd_vv_f32m2(a, bp.v, c.v, unpacket_traits<Packet2Xf>::size));
510 }
511
512 EIGEN_STRONG_INLINE Packet2Xcd pmadd_scalar(const Packet2Xd& a, std::complex<double> b, const Packet2Xcd& c) const {
513 Packet2Xcd bp = pset1<Packet2Xcd>(b);
514 return Packet2Xcd(__riscv_vfmadd_vv_f64m2(a, bp.v, c.v, unpacket_traits<Packet2Xd>::size));
515 }
516#endif
517
518#if EIGEN_RISCV64_DEFAULT_LMUL == 4
519 EIGEN_STRONG_INLINE Packet4Xcf pmadd_scalar(const Packet4Xf& a, std::complex<float> b, const Packet4Xcf& c) const {
520 Packet4Xcf bp = pset1<Packet4Xcf>(b);
521 return Packet4Xcf(__riscv_vfmadd_vv_f32m4(a, bp.v, c.v, unpacket_traits<Packet4Xf>::size));
522 }
523
524 EIGEN_STRONG_INLINE Packet4Xcd pmadd_scalar(const Packet4Xd& a, std::complex<double> b, const Packet4Xcd& c) const {
525 Packet4Xcd bp = pset1<Packet4Xcd>(b);
526 return Packet4Xcd(__riscv_vfmadd_vv_f64m4(a, bp.v, c.v, unpacket_traits<Packet4Xd>::size));
527 }
528#endif
529
530 template <typename LhsPacketType, typename RhsPacketType, typename AccPacketType>
531 EIGEN_STRONG_INLINE void madd_impl(const LhsPacketType& a, const RhsPacketType& b, AccPacketType& c,
532 RhsPacketType& tmp, const std::true_type&) const {
533 EIGEN_UNUSED_VARIABLE(tmp);
534 c = pmadd_scalar(a, b, c);
535 }
536
537 EIGEN_STRONG_INLINE void madd_impl(const LhsScalar& a, const RhsScalar& b, ResScalar& c, RhsScalar& /*tmp*/,
538 const std::false_type&) const {
539 c += a * b;
540 }
541
542 template <typename LhsPacketType, typename AccPacketType, typename LaneIdType>
543 EIGEN_STRONG_INLINE void madd(const LhsPacketType& a, const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
544 const LaneIdType& lane) const {
545 madd(a, b.get(lane), c, tmp, lane);
546 }
547
548 template <typename ResPacketType, typename AccPacketType>
549 EIGEN_STRONG_INLINE void acc(const AccPacketType& c, const ResPacketType& alpha, ResPacketType& r) const {
550 conj_helper<ResPacketType, ResPacketType, false, ConjRhs> cj;
551 r = cj.pmadd(alpha, c, r);
552 }
553};
554
555template <typename RealScalar, bool ConjLhs_, int PacketSize_>
556class gebp_traits<std::complex<RealScalar>, RealScalar, ConjLhs_, false, Architecture::RVV10, PacketSize_>
557 : public gebp_traits<RealScalar, std::complex<RealScalar>, ConjLhs_, false, Architecture::Generic, PacketSize_> {
558 public:
559 typedef std::complex<RealScalar> LhsScalar;
560 typedef RealScalar RhsScalar;
561 typedef typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType ResScalar;
562
563 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
564 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
565 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
566#undef PACKET_DECL_COND_POSTFIX
567
568 enum {
569 ConjLhs = ConjLhs_,
570 ConjRhs = false,
571 Vectorizable = unpacket_traits<LhsPacket_>::vectorizable && unpacket_traits<RhsPacket_>::vectorizable,
572 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
573 RhsPacketSize = Vectorizable ? unpacket_traits<RhsPacket_>::size : 1,
574 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
575
576 nr = 4,
577 mr = 3 * LhsPacketSize,
578
579 LhsProgress = LhsPacketSize,
580 RhsProgress = 1
581 };
582
583 typedef std::conditional_t<Vectorizable, LhsPacket_, LhsScalar> LhsPacket;
584 typedef RhsScalar RhsPacket;
585 typedef std::conditional_t<Vectorizable, ResPacket_, ResScalar> ResPacket;
586 typedef LhsPacket LhsPacket4Packing;
587
588 typedef QuadPacket<RhsPacket> RhsPacketx4;
589
590 typedef ResPacket AccPacket;
591
592 EIGEN_STRONG_INLINE void initAcc(AccPacket& p) { p = pset1<ResPacket>(ResScalar(0)); }
593
594 template <typename RhsPacketType>
595 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketType& dest) const {
596 dest = pset1<RhsPacketType>(*b);
597 }
598
599 EIGEN_STRONG_INLINE void loadRhs(const RhsScalar* b, RhsPacketx4& dest) const {
600 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
601 }
602
603 template <typename RhsPacketType>
604 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar* b, RhsPacketType& dest) const {
605 loadRhs(b, dest);
606 }
607
608 EIGEN_STRONG_INLINE void updateRhs(const RhsScalar*, RhsPacketx4&) const {}
609
610 EIGEN_STRONG_INLINE void loadRhsQuad(const RhsScalar* b, RhsPacket& dest) const {
611 loadRhsQuad_impl(b, dest, std::conditional_t<RhsPacketSize == 16, std::true_type, std::false_type>());
612 }
613
614 EIGEN_STRONG_INLINE void loadRhsQuad_impl(const RhsScalar* b, RhsPacket& dest, const std::true_type&) const {
615 // FIXME we can do better!
616 // what we want here is a ploadheight
617 RhsScalar tmp[4] = {b[0], b[0], b[1], b[1]};
618 dest = ploadquad<RhsPacket>(tmp);
619 }
620
621 EIGEN_STRONG_INLINE void loadRhsQuad_impl(const RhsScalar* b, RhsPacket& dest, const std::false_type&) const {
622 eigen_internal_assert(RhsPacketSize <= 8);
623 dest = pset1<RhsPacket>(*b);
624 }
625
626 EIGEN_STRONG_INLINE void loadLhs(const LhsScalar* a, LhsPacket& dest) const { dest = pload<LhsPacket>(a); }
627
628 template <typename LhsPacketType>
629 EIGEN_STRONG_INLINE void loadLhsUnaligned(const LhsScalar* a, LhsPacketType& dest) const {
630 dest = ploadu<LhsPacketType>(a);
631 }
632
633 template <typename LhsPacketType, typename RhsPacketType, typename AccPacketType, typename LaneIdType>
634 EIGEN_STRONG_INLINE void madd(const LhsPacketType& a, const RhsPacketType& b, AccPacketType& c, RhsPacketType& tmp,
635 const LaneIdType&) const {
636 madd_impl(a, b, c, tmp, std::conditional_t<Vectorizable, std::true_type, std::false_type>());
637 }
638
639 EIGEN_STRONG_INLINE Packet1Xcf pmadd_scalar(const Packet1Xcf& a, float b, const Packet1Xcf& c) const {
640 return Packet1Xcf(__riscv_vfmadd_vf_f32m1(a.v, b, c.v, unpacket_traits<Packet1Xf>::size));
641 }
642
643 EIGEN_STRONG_INLINE Packet1Xcd pmadd_scalar(const Packet1Xcd& a, double b, const Packet1Xcd& c) const {
644 return Packet1Xcd(__riscv_vfmadd_vf_f64m1(a.v, b, c.v, unpacket_traits<Packet1Xd>::size));
645 }
646
647#if EIGEN_RISCV64_DEFAULT_LMUL >= 2
648 EIGEN_STRONG_INLINE Packet2Xcf pmadd_scalar(const Packet2Xcf& a, float b, const Packet2Xcf& c) const {
649 return Packet2Xcf(__riscv_vfmadd_vf_f32m2(a.v, b, c.v, unpacket_traits<Packet2Xf>::size));
650 }
651
652 EIGEN_STRONG_INLINE Packet2Xcd pmadd_scalar(const Packet2Xcd& a, double b, const Packet2Xcd& c) const {
653 return Packet2Xcd(__riscv_vfmadd_vf_f64m2(a.v, b, c.v, unpacket_traits<Packet2Xd>::size));
654 }
655#endif
656
657#if EIGEN_RISCV64_DEFAULT_LMUL == 4
658 EIGEN_STRONG_INLINE Packet4Xcf pmadd_scalar(const Packet4Xcf& a, float b, const Packet4Xcf& c) const {
659 return Packet4Xcf(__riscv_vfmadd_vf_f32m4(a.v, b, c.v, unpacket_traits<Packet4Xf>::size));
660 }
661
662 EIGEN_STRONG_INLINE Packet4Xcd pmadd_scalar(const Packet4Xcd& a, double b, const Packet4Xcd& c) const {
663 return Packet4Xcd(__riscv_vfmadd_vf_f64m4(a.v, b, c.v, unpacket_traits<Packet4Xd>::size));
664 }
665#endif
666
667 template <typename LhsPacketType, typename RhsPacketType, typename AccPacketType>
668 EIGEN_STRONG_INLINE void madd_impl(const LhsPacketType& a, const RhsPacketType& b, AccPacketType& c,
669 RhsPacketType& tmp, const std::true_type&) const {
670 EIGEN_UNUSED_VARIABLE(tmp);
671 c = pmadd_scalar(a, b, c);
672 }
673
674 EIGEN_STRONG_INLINE void madd_impl(const LhsScalar& a, const RhsScalar& b, ResScalar& c, RhsScalar& /*tmp*/,
675 const std::false_type&) const {
676 c += a * b;
677 }
678
679 template <typename LhsPacketType, typename AccPacketType, typename LaneIdType>
680 EIGEN_STRONG_INLINE void madd(const LhsPacketType& a, const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
681 const LaneIdType& lane) const {
682 madd(a, b.get(lane), c, tmp, lane);
683 }
684
685 template <typename ResPacketType, typename AccPacketType>
686 EIGEN_STRONG_INLINE void acc(const AccPacketType& c, const ResPacketType& alpha, ResPacketType& r) const {
687 conj_helper<ResPacketType, ResPacketType, ConjLhs, false> cj;
688 r = cj.pmadd(c, alpha, r);
689 }
690};
691#endif
692
693} // namespace internal
694} // namespace Eigen
695
696#endif // EIGEN_RVV10_GENERAL_BLOCK_KERNEL_H