Eigen  5.0.1
 
Loading...
Searching...
No Matches
PacketMathBF16.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2025 Chip Kerchner <ckerchner@tenstorrent.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_PACKET_MATH_BF16_RVV10_H
12#define EIGEN_PACKET_MATH_BF16_RVV10_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18namespace internal {
19
20typedef eigen_packet_wrapper<vbfloat16m1_t __attribute__((riscv_rvv_vector_bits(EIGEN_RISCV64_RVV_VL))), 26> Packet1Xbf;
21typedef eigen_packet_wrapper<vbfloat16m2_t __attribute__((riscv_rvv_vector_bits(EIGEN_RISCV64_RVV_VL * 2))), 27>
22 Packet2Xbf;
23
24template <>
25struct rvv_half_packet<Packet2Xbf> {
26 typedef Packet1Xbf type;
27};
28
29template <>
30struct unpacket_traits<Packet1Xbf> : rvv_default_unpacket_traits<bfloat16, Packet1Xbf, 1> {
31 typedef Packet1Xs integer_packet;
32 typedef PacketMask16 packet_mask;
33};
34
35template <>
36struct unpacket_traits<Packet2Xbf> : rvv_default_unpacket_traits<bfloat16, Packet2Xbf, 2> {
37 typedef Packet2Xs integer_packet;
38 typedef PacketMask8 packet_mask;
39};
40
41#if EIGEN_RISCV64_DEFAULT_LMUL == 1
42typedef Packet1Xbf PacketXbf;
43#else
44typedef Packet2Xbf PacketXbf;
45#endif
46
47template <>
48struct packet_traits<bfloat16> : rvv_default_packet_traits<bfloat16, PacketXbf> {
49 enum {
50 HasAdd = 1,
51 HasSub = 1,
52 HasShift = 1,
53 HasMul = 1,
54 HasNegate = 1,
55 HasAbs = 1,
56 HasArg = 0,
57 HasAbs2 = 1,
58 HasMin = 1,
59 HasMax = 1,
60 HasConj = 1,
61 HasSetLinear = 0,
62 HasBlend = 0,
63 HasReduxp = 0,
64 HasSign = 0,
65
66 HasCmp = 1,
67 HasDiv = 1,
68 HasRound = 0,
69
70 HasSin = 0,
71 HasCos = 0,
72 HasLog = 0,
73 HasExp = 0,
74 HasSqrt = 1,
75 HasTanh = 0,
76 HasErf = 0
77 };
78};
79
80/********************************* Packet1Xbf ************************************/
81
82EIGEN_STRONG_INLINE Packet1Xbf __riscv_vreinterpret_v_u32m1_bf16m1(const Packet1Xu& a) {
83 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vreinterpret_v_u32m1_u16m1(a));
84}
85
86EIGEN_STRONG_INLINE Packet2Xf Bf16ToF32(const Packet1Xbf& a) {
87 return __riscv_vfwcvtbf16_f_f_v_f32m2(a, unpacket_traits<Packet1Xbf>::size);
88}
89
90EIGEN_STRONG_INLINE Packet1Xbf F32ToBf16(const Packet2Xf& a) {
91 return __riscv_vfncvtbf16_f_f_w_bf16m1(a, unpacket_traits<Packet2Xf>::size);
92}
93
94template <>
95EIGEN_STRONG_INLINE Packet1Xbf ptrue<Packet1Xbf>(const Packet1Xbf& /*a*/) {
96 Packet1Xbf r = __riscv_vreinterpret_bf16m1(
97 __riscv_vmv_v_x_u16m1(static_cast<numext::uint16_t>(0xffffu), unpacket_traits<Packet1Xbf>::size));
98 EIGEN_FAST_MATH_CONSTANT_BARRIER(r);
99 return r;
100}
101
102template <>
103EIGEN_STRONG_INLINE Packet1Xbf pzero<Packet1Xbf>(const Packet1Xbf& /*a*/) {
104 return __riscv_vreinterpret_bf16m1(
105 __riscv_vmv_v_x_i16m1(numext::bit_cast<int16_t>(static_cast<__bf16>(0.0)), unpacket_traits<Packet1Xbf>::size));
106}
107
108template <>
109EIGEN_STRONG_INLINE Packet1Xbf pabs(const Packet1Xbf& a) {
110 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vand_vx_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a),
111 static_cast<numext::uint16_t>(0x7fffu),
112 unpacket_traits<Packet1Xs>::size));
113}
114
115template <>
116EIGEN_STRONG_INLINE Packet1Xbf pset1<Packet1Xbf>(const bfloat16& from) {
117 return __riscv_vreinterpret_bf16m1(
118 __riscv_vmv_v_x_i16m1(numext::bit_cast<int16_t>(from), unpacket_traits<Packet1Xbf>::size));
119}
120
121template <>
122EIGEN_STRONG_INLINE Packet1Xbf pset1frombits<Packet1Xbf>(numext::uint16_t from) {
123 return __riscv_vreinterpret_bf16m1(__riscv_vmv_v_x_u16m1(from, unpacket_traits<Packet1Xbf>::size));
124}
125
126template <>
127EIGEN_STRONG_INLINE Packet1Xbf plset<Packet1Xbf>(const bfloat16& a) {
128 return F32ToBf16(plset<Packet2Xf>(static_cast<float>(a)));
129}
130
131template <>
132EIGEN_STRONG_INLINE void pbroadcast4<Packet1Xbf>(const bfloat16* a, Packet1Xbf& a0, Packet1Xbf& a1, Packet1Xbf& a2,
133 Packet1Xbf& a3) {
134 vint16m1_t aa = __riscv_vle16_v_i16m1(reinterpret_cast<const int16_t*>(a), 4);
135 a0 = __riscv_vreinterpret_bf16m1(__riscv_vrgather_vx_i16m1(aa, 0, unpacket_traits<Packet1Xs>::size));
136 a1 = __riscv_vreinterpret_bf16m1(__riscv_vrgather_vx_i16m1(aa, 1, unpacket_traits<Packet1Xs>::size));
137 a2 = __riscv_vreinterpret_bf16m1(__riscv_vrgather_vx_i16m1(aa, 2, unpacket_traits<Packet1Xs>::size));
138 a3 = __riscv_vreinterpret_bf16m1(__riscv_vrgather_vx_i16m1(aa, 3, unpacket_traits<Packet1Xs>::size));
139}
140
141template <>
142EIGEN_STRONG_INLINE Packet1Xbf padd<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
143 // b + (1 * a)
144 return F32ToBf16(__riscv_vfwmaccbf16_vf_f32m2(Bf16ToF32(b),
145 numext::bit_cast<__bf16>(static_cast<numext::int16_t>(0x3f80u)), a,
146 unpacket_traits<Packet1Xbf>::size));
147}
148
149template <>
150EIGEN_STRONG_INLINE Packet1Xbf psub<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
151 // a + (-1 * b)
152 return F32ToBf16(__riscv_vfwmaccbf16_vf_f32m2(Bf16ToF32(a),
153 numext::bit_cast<__bf16>(static_cast<numext::int16_t>(0xbf80u)), b,
154 unpacket_traits<Packet1Xbf>::size));
155}
156
157template <>
158EIGEN_STRONG_INLINE Packet1Xbf pabsdiff(const Packet1Xbf& a, const Packet1Xbf& b) {
159 return pabs<Packet1Xbf>(psub<Packet1Xbf>(a, b));
160}
161
162template <>
163EIGEN_STRONG_INLINE Packet1Xbf pnegate(const Packet1Xbf& a) {
164 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vxor_vx_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a),
165 static_cast<numext::uint16_t>(0x8000u),
166 unpacket_traits<Packet1Xs>::size));
167}
168
169template <>
170EIGEN_STRONG_INLINE Packet1Xbf psignbit(const Packet1Xbf& a) {
171 return __riscv_vreinterpret_v_i16m1_bf16m1(
172 __riscv_vsra_vx_i16m1(__riscv_vreinterpret_v_bf16m1_i16m1(a), 15, unpacket_traits<Packet1Xs>::size));
173}
174
175template <>
176EIGEN_STRONG_INLINE Packet1Xbf pmul<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
177 Packet2Xf c;
178 return F32ToBf16(__riscv_vfwmaccbf16_vv_f32m2(pzero<Packet2Xf>(c), a, b, unpacket_traits<Packet1Xbf>::size));
179}
180
181template <>
182EIGEN_STRONG_INLINE Packet1Xbf pdiv<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
183 return F32ToBf16(pdiv<Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
184}
185
186template <>
187EIGEN_STRONG_INLINE Packet1Xbf pmadd(const Packet1Xbf& a, const Packet1Xbf& b, const Packet1Xbf& c) {
188 return F32ToBf16(__riscv_vfwmaccbf16_vv_f32m2(Bf16ToF32(c), a, b, unpacket_traits<Packet1Xbf>::size));
189}
190
191template <>
192EIGEN_STRONG_INLINE Packet1Xbf pmsub(const Packet1Xbf& a, const Packet1Xbf& b, const Packet1Xbf& c) {
193 return F32ToBf16(
194 __riscv_vfwmaccbf16_vv_f32m2(Bf16ToF32(pnegate<Packet1Xbf>(c)), a, b, unpacket_traits<Packet1Xbf>::size));
195}
196
197template <>
198EIGEN_STRONG_INLINE Packet1Xbf pnmadd(const Packet1Xbf& a, const Packet1Xbf& b, const Packet1Xbf& c) {
199 return F32ToBf16(
200 __riscv_vfwmaccbf16_vv_f32m2(Bf16ToF32(c), pnegate<Packet1Xbf>(a), b, unpacket_traits<Packet1Xbf>::size));
201}
202
203template <>
204EIGEN_STRONG_INLINE Packet1Xbf pnmsub(const Packet1Xbf& a, const Packet1Xbf& b, const Packet1Xbf& c) {
205 return pnegate<Packet1Xbf>(
206 F32ToBf16(__riscv_vfwmaccbf16_vv_f32m2(Bf16ToF32(c), a, b, unpacket_traits<Packet1Xbf>::size)));
207}
208
209template <>
210EIGEN_STRONG_INLINE Packet1Xbf pmin<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
211 return F32ToBf16(pmin<Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
212}
213
214template <>
215EIGEN_STRONG_INLINE Packet1Xbf pmin<PropagateNaN, Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
216 return F32ToBf16(pmin<PropagateNaN, Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
217}
218
219template <>
220EIGEN_STRONG_INLINE Packet1Xbf pmin<PropagateNumbers, Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
221 return F32ToBf16(pmin<PropagateNumbers, Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
222}
223
224template <>
225EIGEN_STRONG_INLINE Packet1Xbf pmax<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
226 return F32ToBf16(pmax<Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
227}
228
229template <>
230EIGEN_STRONG_INLINE Packet1Xbf pmax<PropagateNaN, Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
231 return F32ToBf16(pmax<PropagateNaN, Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
232}
233
234template <>
235EIGEN_STRONG_INLINE Packet1Xbf pmax<PropagateNumbers, Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
236 return F32ToBf16(pmax<PropagateNumbers, Packet2Xf>(Bf16ToF32(a), Bf16ToF32(b)));
237}
238
239// Comparisons are performed in float32 and the resulting vbool mask is expanded to all-ones/all-zeros
240// 16-bit lanes with a vmerge. Narrowing the float32 comparison result with F32ToBf16 instead would go
241// through an arithmetic conversion (vfncvtbf16) that canonicalizes the all-ones (NaN) lanes to 0x7fc0
242// and corrupts the mask.
243template <>
244EIGEN_STRONG_INLINE Packet1Xbf pcmp_le<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
245 PacketMask16 mask = __riscv_vmfle_vv_f32m2_b16(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet1Xbf>::size);
246 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vmerge_vvm_u16m1(
247 __riscv_vreinterpret_v_bf16m1_u16m1(pzero<Packet1Xbf>(a)),
248 __riscv_vreinterpret_v_bf16m1_u16m1(ptrue<Packet1Xbf>(a)), mask, unpacket_traits<Packet1Xbf>::size));
249}
250
251template <>
252EIGEN_STRONG_INLINE Packet1Xbf pcmp_lt<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
253 PacketMask16 mask = __riscv_vmflt_vv_f32m2_b16(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet1Xbf>::size);
254 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vmerge_vvm_u16m1(
255 __riscv_vreinterpret_v_bf16m1_u16m1(pzero<Packet1Xbf>(a)),
256 __riscv_vreinterpret_v_bf16m1_u16m1(ptrue<Packet1Xbf>(a)), mask, unpacket_traits<Packet1Xbf>::size));
257}
258
259template <>
260EIGEN_STRONG_INLINE Packet1Xbf pcmp_eq<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
261 PacketMask16 mask = __riscv_vmfeq_vv_f32m2_b16(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet1Xbf>::size);
262 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vmerge_vvm_u16m1(
263 __riscv_vreinterpret_v_bf16m1_u16m1(pzero<Packet1Xbf>(a)),
264 __riscv_vreinterpret_v_bf16m1_u16m1(ptrue<Packet1Xbf>(a)), mask, unpacket_traits<Packet1Xbf>::size));
265}
266
267template <>
268EIGEN_STRONG_INLINE Packet1Xbf pcmp_lt_or_nan<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
269 PacketMask16 mask = __riscv_vmfge_vv_f32m2_b16(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet1Xbf>::size);
270 return __riscv_vreinterpret_v_u16m1_bf16m1(
271 __riscv_vmerge_vxm_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(ptrue<Packet1Xbf>(a)),
272 static_cast<numext::uint16_t>(0), mask, unpacket_traits<Packet1Xbf>::size));
273}
274
275// Classify on the raw bits, as the scalar isinf/isnan/isfinite do: |a| ==, >, < 0x7f80. The generic forms compare
276// the widened floats, and -ffinite-math-only folds the a >= a of pisnan to true.
277template <>
278EIGEN_STRONG_INLINE Packet1Xbf pisinf<Packet1Xbf>(const Packet1Xbf& a) {
279 const vuint16m1_t abs_bits =
280 __riscv_vand_vx_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a), 0x7fffu, unpacket_traits<Packet1Xbf>::size);
281 const PacketMask16 mask = __riscv_vmseq_vx_u16m1_b16(abs_bits, 0x7f80u, unpacket_traits<Packet1Xbf>::size);
282 return __riscv_vreinterpret_v_i16m1_bf16m1(__riscv_vmerge_vxm_i16m1(
283 __riscv_vreinterpret_v_bf16m1_i16m1(pzero<Packet1Xbf>(a)), -1, mask, unpacket_traits<Packet1Xbf>::size));
284}
285
286template <>
287EIGEN_STRONG_INLINE Packet1Xbf pisnan<Packet1Xbf>(const Packet1Xbf& a) {
288 const vuint16m1_t abs_bits =
289 __riscv_vand_vx_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a), 0x7fffu, unpacket_traits<Packet1Xbf>::size);
290 const PacketMask16 mask = __riscv_vmsgtu_vx_u16m1_b16(abs_bits, 0x7f80u, unpacket_traits<Packet1Xbf>::size);
291 return __riscv_vreinterpret_v_i16m1_bf16m1(__riscv_vmerge_vxm_i16m1(
292 __riscv_vreinterpret_v_bf16m1_i16m1(pzero<Packet1Xbf>(a)), -1, mask, unpacket_traits<Packet1Xbf>::size));
293}
294
295template <>
296EIGEN_STRONG_INLINE Packet1Xbf pisfinite<Packet1Xbf>(const Packet1Xbf& a) {
297 const vuint16m1_t abs_bits =
298 __riscv_vand_vx_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a), 0x7fffu, unpacket_traits<Packet1Xbf>::size);
299 const PacketMask16 mask = __riscv_vmsltu_vx_u16m1_b16(abs_bits, 0x7f80u, unpacket_traits<Packet1Xbf>::size);
300 return __riscv_vreinterpret_v_i16m1_bf16m1(__riscv_vmerge_vxm_i16m1(
301 __riscv_vreinterpret_v_bf16m1_i16m1(pzero<Packet1Xbf>(a)), -1, mask, unpacket_traits<Packet1Xbf>::size));
302}
303
304EIGEN_STRONG_INLINE Packet1Xbf pselect(const PacketMask16& mask, const Packet1Xbf& a, const Packet1Xbf& b) {
305 return __riscv_vreinterpret_v_i16m1_bf16m1(__riscv_vmerge_vvm_i16m1(__riscv_vreinterpret_v_bf16m1_i16m1(b),
306 __riscv_vreinterpret_v_bf16m1_i16m1(a), mask,
307 unpacket_traits<Packet1Xbf>::size));
308}
309
310EIGEN_STRONG_INLINE Packet1Xbf pselect(const Packet1Xbf& mask, const Packet1Xbf& a, const Packet1Xbf& b) {
311 PacketMask16 mask2 =
312 __riscv_vmsne_vx_i16m1_b16(__riscv_vreinterpret_v_bf16m1_i16m1(mask), 0, unpacket_traits<Packet1Xbf>::size);
313 return __riscv_vreinterpret_v_i16m1_bf16m1(__riscv_vmerge_vvm_i16m1(__riscv_vreinterpret_v_bf16m1_i16m1(b),
314 __riscv_vreinterpret_v_bf16m1_i16m1(a), mask2,
315 unpacket_traits<Packet1Xbf>::size));
316}
317
318// Logical Operations are not supported for bfloat16, so reinterpret casts
319template <>
320EIGEN_STRONG_INLINE Packet1Xbf pand<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
321 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vand_vv_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a),
322 __riscv_vreinterpret_v_bf16m1_u16m1(b),
323 unpacket_traits<Packet1Xbf>::size));
324}
325
326template <>
327EIGEN_STRONG_INLINE Packet1Xbf por<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
328 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vor_vv_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a),
329 __riscv_vreinterpret_v_bf16m1_u16m1(b),
330 unpacket_traits<Packet1Xbf>::size));
331}
332
333template <>
334EIGEN_STRONG_INLINE Packet1Xbf pxor<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
335 return __riscv_vreinterpret_v_u16m1_bf16m1(__riscv_vxor_vv_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a),
336 __riscv_vreinterpret_v_bf16m1_u16m1(b),
337 unpacket_traits<Packet1Xbf>::size));
338}
339
340template <>
341EIGEN_STRONG_INLINE Packet1Xbf pandnot<Packet1Xbf>(const Packet1Xbf& a, const Packet1Xbf& b) {
342 return __riscv_vreinterpret_v_i16m1_bf16m1(
343 pandnot<Packet1Xs>(__riscv_vreinterpret_v_bf16m1_i16m1(a), __riscv_vreinterpret_v_bf16m1_i16m1(b)));
344}
345
346template <>
347EIGEN_STRONG_INLINE Packet1Xbf pnot<Packet1Xbf>(const Packet1Xbf& a) {
348 return __riscv_vreinterpret_v_u16m1_bf16m1(
349 __riscv_vnot_v_u16m1(__riscv_vreinterpret_v_bf16m1_u16m1(a), unpacket_traits<Packet1Xbf>::size));
350}
351
352template <>
353EIGEN_STRONG_INLINE Packet1Xbf pload<Packet1Xbf>(const bfloat16* from) {
354 EIGEN_DEBUG_ALIGNED_LOAD return __riscv_vle16_v_bf16m1(reinterpret_cast<const __bf16*>(from),
355 unpacket_traits<Packet1Xbf>::size);
356}
357
358template <>
359EIGEN_STRONG_INLINE Packet1Xbf ploadu<Packet1Xbf>(const bfloat16* from) {
360 EIGEN_DEBUG_UNALIGNED_LOAD return __riscv_vle16_v_bf16m1(reinterpret_cast<const __bf16*>(from),
361 unpacket_traits<Packet1Xbf>::size);
362}
363
364template <>
365EIGEN_STRONG_INLINE Packet1Xbf ploaddup<Packet1Xbf>(const bfloat16* from) {
366 return __riscv_vreinterpret_v_i16m1_bf16m1(ploaddup<Packet1Xs>(reinterpret_cast<const numext::int16_t*>(from)));
367}
368
369template <>
370EIGEN_STRONG_INLINE Packet1Xbf ploadquad<Packet1Xbf>(const bfloat16* from) {
371 return __riscv_vreinterpret_v_i16m1_bf16m1(ploadquad<Packet1Xs>(reinterpret_cast<const numext::int16_t*>(from)));
372}
373
374template <>
375EIGEN_STRONG_INLINE void pstore<bfloat16>(bfloat16* to, const Packet1Xbf& from) {
376 EIGEN_DEBUG_ALIGNED_STORE __riscv_vse16_v_bf16m1(reinterpret_cast<__bf16*>(to), from,
377 unpacket_traits<Packet1Xbf>::size);
378}
379
380template <>
381EIGEN_STRONG_INLINE void pstoreu<bfloat16>(bfloat16* to, const Packet1Xbf& from) {
382 EIGEN_DEBUG_UNALIGNED_STORE __riscv_vse16_v_bf16m1(reinterpret_cast<__bf16*>(to), from,
383 unpacket_traits<Packet1Xbf>::size);
384}
385
386template <>
387EIGEN_DEVICE_FUNC inline Packet1Xbf pgather<bfloat16, Packet1Xbf>(const bfloat16* from, Index stride) {
388 return __riscv_vlse16_v_bf16m1(reinterpret_cast<const __bf16*>(from), stride * sizeof(bfloat16),
389 unpacket_traits<Packet1Xbf>::size);
390}
391
392template <>
393EIGEN_DEVICE_FUNC inline void pscatter<bfloat16, Packet1Xbf>(bfloat16* to, const Packet1Xbf& from, Index stride) {
394 __riscv_vsse16(reinterpret_cast<__bf16*>(to), stride * sizeof(bfloat16), from, unpacket_traits<Packet1Xbf>::size);
395}
396
397template <>
398EIGEN_STRONG_INLINE bfloat16 pfirst<Packet1Xbf>(const Packet1Xbf& a) {
399 return numext::bit_cast<bfloat16>(__riscv_vmv_x_s_i16m1_i16(__riscv_vreinterpret_v_bf16m1_i16m1(a)));
400}
401
402template <>
403EIGEN_STRONG_INLINE Packet1Xbf psqrt(const Packet1Xbf& a) {
404 return F32ToBf16(psqrt<Packet2Xf>(Bf16ToF32(a)));
405}
406
407template <>
408EIGEN_STRONG_INLINE Packet1Xbf print<Packet1Xbf>(const Packet1Xbf& a) {
409 return F32ToBf16(print<Packet2Xf>(Bf16ToF32(a)));
410}
411
412template <>
413EIGEN_STRONG_INLINE Packet1Xbf pfloor<Packet1Xbf>(const Packet1Xbf& a) {
414 return F32ToBf16(pfloor<Packet2Xf>(Bf16ToF32(a)));
415}
416
417template <>
418EIGEN_STRONG_INLINE Packet1Xbf preverse(const Packet1Xbf& a) {
419 return __riscv_vreinterpret_v_i16m1_bf16m1(preverse<Packet1Xs>(__riscv_vreinterpret_v_bf16m1_i16m1(a)));
420}
421
422template <>
423EIGEN_STRONG_INLINE bfloat16 predux<Packet1Xbf>(const Packet1Xbf& a) {
424 return static_cast<bfloat16>(predux<Packet2Xf>(Bf16ToF32(a)));
425}
426
427template <>
428EIGEN_STRONG_INLINE bool predux_any(const Packet1Xbf& a) {
429 const PacketMask16 mask =
430 __riscv_vmsne_vx_u16m1_b16(__riscv_vreinterpret_v_bf16m1_u16m1(a), 0, unpacket_traits<Packet1Xbf>::size);
431 return __riscv_vcpop_m_b16(mask, unpacket_traits<Packet1Xbf>::size) != 0;
432}
433
434template <>
435EIGEN_STRONG_INLINE bool predux_all(const Packet1Xbf& a) {
436 const PacketMask16 mask = __riscv_vmfeq_vf_f32m2_b16(Bf16ToF32(a), 0.0f, unpacket_traits<Packet1Xbf>::size);
437 return __riscv_vcpop_m_b16(mask, unpacket_traits<Packet1Xbf>::size) == 0;
438}
439
440template <>
441EIGEN_STRONG_INLINE bfloat16 predux_mul<Packet1Xbf>(const Packet1Xbf& a) {
442 return static_cast<bfloat16>(predux_mul<Packet2Xf>(Bf16ToF32(a)));
443}
444
445template <>
446EIGEN_STRONG_INLINE bfloat16 predux_min<Packet1Xbf>(const Packet1Xbf& a) {
447 return static_cast<bfloat16>(predux_min<Packet2Xf>(Bf16ToF32(a)));
448}
449
450template <>
451EIGEN_STRONG_INLINE bfloat16 predux_max<Packet1Xbf>(const Packet1Xbf& a) {
452 return static_cast<bfloat16>(predux_max<Packet2Xf>(Bf16ToF32(a)));
453}
454
455template <int N>
456EIGEN_DEVICE_FUNC inline void ptranspose(PacketBlock<Packet1Xbf, N>& kernel) {
457 bfloat16 buffer[unpacket_traits<Packet1Xbf>::size * N];
458 int i = 0;
459
460 for (i = 0; i < N; i++) {
461 __riscv_vsse16(reinterpret_cast<__bf16*>(&buffer[i]), N * sizeof(bfloat16), kernel.packet[i],
462 unpacket_traits<Packet1Xbf>::size);
463 }
464
465 for (i = 0; i < N; i++) {
466 kernel.packet[i] = __riscv_vle16_v_bf16m1(reinterpret_cast<__bf16*>(&buffer[i * unpacket_traits<Packet1Xbf>::size]),
467 unpacket_traits<Packet1Xbf>::size);
468 }
469}
470
471/********************************* Packet2Xbf ************************************/
472
473EIGEN_STRONG_INLINE Packet2Xbf __riscv_vreinterpret_v_u32m2_bf16m2(const Packet2Xu& a) {
474 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vreinterpret_v_u32m2_u16m2(a));
475}
476
477EIGEN_STRONG_INLINE Packet4Xf Bf16ToF32(const Packet2Xbf& a) {
478 return __riscv_vfwcvtbf16_f_f_v_f32m4(a, unpacket_traits<Packet2Xbf>::size);
479}
480
481EIGEN_STRONG_INLINE Packet2Xbf F32ToBf16(const Packet4Xf& a) {
482 return __riscv_vfncvtbf16_f_f_w_bf16m2(a, unpacket_traits<Packet4Xf>::size);
483}
484
485template <>
486EIGEN_STRONG_INLINE Packet2Xbf ptrue<Packet2Xbf>(const Packet2Xbf& /*a*/) {
487 Packet2Xbf r = __riscv_vreinterpret_bf16m2(
488 __riscv_vmv_v_x_u16m2(static_cast<numext::uint16_t>(0xffffu), unpacket_traits<Packet2Xbf>::size));
489 EIGEN_FAST_MATH_CONSTANT_BARRIER(r);
490 return r;
491}
492
493template <>
494EIGEN_STRONG_INLINE Packet2Xbf pzero<Packet2Xbf>(const Packet2Xbf& /*a*/) {
495 return __riscv_vreinterpret_bf16m2(
496 __riscv_vmv_v_x_i16m2(numext::bit_cast<int16_t>(static_cast<__bf16>(0.0)), unpacket_traits<Packet2Xbf>::size));
497}
498
499template <>
500EIGEN_STRONG_INLINE Packet2Xbf pabs(const Packet2Xbf& a) {
501 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vand_vx_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a),
502 static_cast<numext::uint16_t>(0x7fffu),
503 unpacket_traits<Packet2Xs>::size));
504}
505
506template <>
507EIGEN_STRONG_INLINE Packet2Xbf pset1<Packet2Xbf>(const bfloat16& from) {
508 return __riscv_vreinterpret_bf16m2(
509 __riscv_vmv_v_x_i16m2(numext::bit_cast<int16_t>(from), unpacket_traits<Packet2Xbf>::size));
510}
511
512template <>
513EIGEN_STRONG_INLINE Packet2Xbf pset1frombits<Packet2Xbf>(numext::uint16_t from) {
514 return __riscv_vreinterpret_bf16m2(__riscv_vmv_v_x_u16m2(from, unpacket_traits<Packet2Xbf>::size));
515}
516
517template <>
518EIGEN_STRONG_INLINE Packet2Xbf plset<Packet2Xbf>(const bfloat16& a) {
519 return F32ToBf16(plset<Packet4Xf>(static_cast<float>(a)));
520}
521
522template <>
523EIGEN_STRONG_INLINE void pbroadcast4<Packet2Xbf>(const bfloat16* a, Packet2Xbf& a0, Packet2Xbf& a1, Packet2Xbf& a2,
524 Packet2Xbf& a3) {
525 vint16m2_t aa = __riscv_vle16_v_i16m2(reinterpret_cast<const int16_t*>(a), 4);
526 a0 = __riscv_vreinterpret_bf16m2(__riscv_vrgather_vx_i16m2(aa, 0, unpacket_traits<Packet2Xs>::size));
527 a1 = __riscv_vreinterpret_bf16m2(__riscv_vrgather_vx_i16m2(aa, 1, unpacket_traits<Packet2Xs>::size));
528 a2 = __riscv_vreinterpret_bf16m2(__riscv_vrgather_vx_i16m2(aa, 2, unpacket_traits<Packet2Xs>::size));
529 a3 = __riscv_vreinterpret_bf16m2(__riscv_vrgather_vx_i16m2(aa, 3, unpacket_traits<Packet2Xs>::size));
530}
531
532template <>
533EIGEN_STRONG_INLINE Packet2Xbf padd<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
534 // b + (1 * a)
535 return F32ToBf16(__riscv_vfwmaccbf16_vf_f32m4(Bf16ToF32(b),
536 numext::bit_cast<__bf16>(static_cast<numext::int16_t>(0x3f80u)), a,
537 unpacket_traits<Packet2Xbf>::size));
538}
539
540template <>
541EIGEN_STRONG_INLINE Packet2Xbf psub<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
542 // a + (-1 * b)
543 return F32ToBf16(__riscv_vfwmaccbf16_vf_f32m4(Bf16ToF32(a),
544 numext::bit_cast<__bf16>(static_cast<numext::int16_t>(0xbf80u)), b,
545 unpacket_traits<Packet2Xbf>::size));
546}
547
548template <>
549EIGEN_STRONG_INLINE Packet2Xbf pabsdiff(const Packet2Xbf& a, const Packet2Xbf& b) {
550 return pabs<Packet2Xbf>(psub<Packet2Xbf>(a, b));
551}
552
553template <>
554EIGEN_STRONG_INLINE Packet2Xbf pnegate(const Packet2Xbf& a) {
555 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vxor_vx_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a),
556 static_cast<numext::uint16_t>(0x8000u),
557 unpacket_traits<Packet2Xs>::size));
558}
559
560template <>
561EIGEN_STRONG_INLINE Packet2Xbf psignbit(const Packet2Xbf& a) {
562 return __riscv_vreinterpret_v_i16m2_bf16m2(
563 __riscv_vsra_vx_i16m2(__riscv_vreinterpret_v_bf16m2_i16m2(a), 15, unpacket_traits<Packet2Xs>::size));
564}
565
566template <>
567EIGEN_STRONG_INLINE Packet2Xbf pmul<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
568 Packet4Xf c;
569 return F32ToBf16(__riscv_vfwmaccbf16_vv_f32m4(pzero<Packet4Xf>(c), a, b, unpacket_traits<Packet2Xbf>::size));
570}
571
572template <>
573EIGEN_STRONG_INLINE Packet2Xbf pdiv<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
574 return F32ToBf16(pdiv<Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
575}
576
577template <>
578EIGEN_STRONG_INLINE Packet2Xbf pmadd(const Packet2Xbf& a, const Packet2Xbf& b, const Packet2Xbf& c) {
579 return F32ToBf16(__riscv_vfwmaccbf16_vv_f32m4(Bf16ToF32(c), a, b, unpacket_traits<Packet2Xbf>::size));
580}
581
582template <>
583EIGEN_STRONG_INLINE Packet2Xbf pmsub(const Packet2Xbf& a, const Packet2Xbf& b, const Packet2Xbf& c) {
584 return F32ToBf16(
585 __riscv_vfwmaccbf16_vv_f32m4(Bf16ToF32(pnegate<Packet2Xbf>(c)), a, b, unpacket_traits<Packet2Xbf>::size));
586}
587
588template <>
589EIGEN_STRONG_INLINE Packet2Xbf pnmadd(const Packet2Xbf& a, const Packet2Xbf& b, const Packet2Xbf& c) {
590 return F32ToBf16(
591 __riscv_vfwmaccbf16_vv_f32m4(Bf16ToF32(c), pnegate<Packet2Xbf>(a), b, unpacket_traits<Packet2Xbf>::size));
592}
593
594template <>
595EIGEN_STRONG_INLINE Packet2Xbf pnmsub(const Packet2Xbf& a, const Packet2Xbf& b, const Packet2Xbf& c) {
596 return pnegate<Packet2Xbf>(
597 F32ToBf16(__riscv_vfwmaccbf16_vv_f32m4(Bf16ToF32(c), a, b, unpacket_traits<Packet2Xbf>::size)));
598}
599
600template <>
601EIGEN_STRONG_INLINE Packet2Xbf pmin<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
602 return F32ToBf16(pmin<Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
603}
604
605template <>
606EIGEN_STRONG_INLINE Packet2Xbf pmin<PropagateNaN, Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
607 return F32ToBf16(pmin<PropagateNaN, Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
608}
609
610template <>
611EIGEN_STRONG_INLINE Packet2Xbf pmin<PropagateNumbers, Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
612 return F32ToBf16(pmin<PropagateNumbers, Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
613}
614
615template <>
616EIGEN_STRONG_INLINE Packet2Xbf pmax<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
617 return F32ToBf16(pmax<Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
618}
619
620template <>
621EIGEN_STRONG_INLINE Packet2Xbf pmax<PropagateNaN, Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
622 return F32ToBf16(pmax<PropagateNaN, Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
623}
624
625template <>
626EIGEN_STRONG_INLINE Packet2Xbf pmax<PropagateNumbers, Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
627 return F32ToBf16(pmax<PropagateNumbers, Packet4Xf>(Bf16ToF32(a), Bf16ToF32(b)));
628}
629
630// See the Packet1Xbf comparisons above: the vbool mask is expanded with a vmerge because an arithmetic
631// narrowing conversion (vfncvtbf16) would canonicalize the all-ones (NaN) lanes to 0x7fc0 and corrupt
632// the mask.
633template <>
634EIGEN_STRONG_INLINE Packet2Xbf pcmp_le<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
635 PacketMask8 mask = __riscv_vmfle_vv_f32m4_b8(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet2Xbf>::size);
636 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vmerge_vvm_u16m2(
637 __riscv_vreinterpret_v_bf16m2_u16m2(pzero<Packet2Xbf>(a)),
638 __riscv_vreinterpret_v_bf16m2_u16m2(ptrue<Packet2Xbf>(a)), mask, unpacket_traits<Packet2Xbf>::size));
639}
640
641template <>
642EIGEN_STRONG_INLINE Packet2Xbf pcmp_lt<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
643 PacketMask8 mask = __riscv_vmflt_vv_f32m4_b8(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet2Xbf>::size);
644 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vmerge_vvm_u16m2(
645 __riscv_vreinterpret_v_bf16m2_u16m2(pzero<Packet2Xbf>(a)),
646 __riscv_vreinterpret_v_bf16m2_u16m2(ptrue<Packet2Xbf>(a)), mask, unpacket_traits<Packet2Xbf>::size));
647}
648
649template <>
650EIGEN_STRONG_INLINE Packet2Xbf pcmp_eq<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
651 PacketMask8 mask = __riscv_vmfeq_vv_f32m4_b8(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet2Xbf>::size);
652 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vmerge_vvm_u16m2(
653 __riscv_vreinterpret_v_bf16m2_u16m2(pzero<Packet2Xbf>(a)),
654 __riscv_vreinterpret_v_bf16m2_u16m2(ptrue<Packet2Xbf>(a)), mask, unpacket_traits<Packet2Xbf>::size));
655}
656
657template <>
658EIGEN_STRONG_INLINE Packet2Xbf pcmp_lt_or_nan<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
659 PacketMask8 mask = __riscv_vmfge_vv_f32m4_b8(Bf16ToF32(a), Bf16ToF32(b), unpacket_traits<Packet2Xbf>::size);
660 return __riscv_vreinterpret_v_u16m2_bf16m2(
661 __riscv_vmerge_vxm_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(ptrue<Packet2Xbf>(a)),
662 static_cast<numext::uint16_t>(0), mask, unpacket_traits<Packet2Xbf>::size));
663}
664
665// See the Packet1Xbf classification above.
666template <>
667EIGEN_STRONG_INLINE Packet2Xbf pisinf<Packet2Xbf>(const Packet2Xbf& a) {
668 const vuint16m2_t abs_bits =
669 __riscv_vand_vx_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a), 0x7fffu, unpacket_traits<Packet2Xbf>::size);
670 const PacketMask8 mask = __riscv_vmseq_vx_u16m2_b8(abs_bits, 0x7f80u, unpacket_traits<Packet2Xbf>::size);
671 return __riscv_vreinterpret_v_i16m2_bf16m2(__riscv_vmerge_vxm_i16m2(
672 __riscv_vreinterpret_v_bf16m2_i16m2(pzero<Packet2Xbf>(a)), -1, mask, unpacket_traits<Packet2Xbf>::size));
673}
674
675template <>
676EIGEN_STRONG_INLINE Packet2Xbf pisnan<Packet2Xbf>(const Packet2Xbf& a) {
677 const vuint16m2_t abs_bits =
678 __riscv_vand_vx_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a), 0x7fffu, unpacket_traits<Packet2Xbf>::size);
679 const PacketMask8 mask = __riscv_vmsgtu_vx_u16m2_b8(abs_bits, 0x7f80u, unpacket_traits<Packet2Xbf>::size);
680 return __riscv_vreinterpret_v_i16m2_bf16m2(__riscv_vmerge_vxm_i16m2(
681 __riscv_vreinterpret_v_bf16m2_i16m2(pzero<Packet2Xbf>(a)), -1, mask, unpacket_traits<Packet2Xbf>::size));
682}
683
684template <>
685EIGEN_STRONG_INLINE Packet2Xbf pisfinite<Packet2Xbf>(const Packet2Xbf& a) {
686 const vuint16m2_t abs_bits =
687 __riscv_vand_vx_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a), 0x7fffu, unpacket_traits<Packet2Xbf>::size);
688 const PacketMask8 mask = __riscv_vmsltu_vx_u16m2_b8(abs_bits, 0x7f80u, unpacket_traits<Packet2Xbf>::size);
689 return __riscv_vreinterpret_v_i16m2_bf16m2(__riscv_vmerge_vxm_i16m2(
690 __riscv_vreinterpret_v_bf16m2_i16m2(pzero<Packet2Xbf>(a)), -1, mask, unpacket_traits<Packet2Xbf>::size));
691}
692
693EIGEN_STRONG_INLINE Packet2Xbf pselect(const PacketMask8& mask, const Packet2Xbf& a, const Packet2Xbf& b) {
694 return __riscv_vreinterpret_v_i16m2_bf16m2(__riscv_vmerge_vvm_i16m2(__riscv_vreinterpret_v_bf16m2_i16m2(b),
695 __riscv_vreinterpret_v_bf16m2_i16m2(a), mask,
696 unpacket_traits<Packet2Xbf>::size));
697}
698
699EIGEN_STRONG_INLINE Packet2Xbf pselect(const Packet2Xbf& mask, const Packet2Xbf& a, const Packet2Xbf& b) {
700 PacketMask8 mask2 =
701 __riscv_vmsne_vx_i16m2_b8(__riscv_vreinterpret_v_bf16m2_i16m2(mask), 0, unpacket_traits<Packet2Xbf>::size);
702 return __riscv_vreinterpret_v_i16m2_bf16m2(__riscv_vmerge_vvm_i16m2(__riscv_vreinterpret_v_bf16m2_i16m2(b),
703 __riscv_vreinterpret_v_bf16m2_i16m2(a), mask2,
704 unpacket_traits<Packet2Xbf>::size));
705}
706
707// Logical Operations are not supported for bfloat16, so reinterpret casts
708template <>
709EIGEN_STRONG_INLINE Packet2Xbf pand<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
710 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vand_vv_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a),
711 __riscv_vreinterpret_v_bf16m2_u16m2(b),
712 unpacket_traits<Packet2Xbf>::size));
713}
714
715template <>
716EIGEN_STRONG_INLINE Packet2Xbf por<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
717 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vor_vv_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a),
718 __riscv_vreinterpret_v_bf16m2_u16m2(b),
719 unpacket_traits<Packet2Xbf>::size));
720}
721
722template <>
723EIGEN_STRONG_INLINE Packet2Xbf pxor<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
724 return __riscv_vreinterpret_v_u16m2_bf16m2(__riscv_vxor_vv_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a),
725 __riscv_vreinterpret_v_bf16m2_u16m2(b),
726 unpacket_traits<Packet2Xbf>::size));
727}
728
729template <>
730EIGEN_STRONG_INLINE Packet2Xbf pandnot<Packet2Xbf>(const Packet2Xbf& a, const Packet2Xbf& b) {
731 return __riscv_vreinterpret_v_i16m2_bf16m2(
732 pandnot<Packet2Xs>(__riscv_vreinterpret_v_bf16m2_i16m2(a), __riscv_vreinterpret_v_bf16m2_i16m2(b)));
733}
734
735template <>
736EIGEN_STRONG_INLINE Packet2Xbf pnot<Packet2Xbf>(const Packet2Xbf& a) {
737 return __riscv_vreinterpret_v_u16m2_bf16m2(
738 __riscv_vnot_v_u16m2(__riscv_vreinterpret_v_bf16m2_u16m2(a), unpacket_traits<Packet2Xbf>::size));
739}
740
741template <>
742EIGEN_STRONG_INLINE Packet2Xbf pload<Packet2Xbf>(const bfloat16* from) {
743 EIGEN_DEBUG_ALIGNED_LOAD return __riscv_vle16_v_bf16m2(reinterpret_cast<const __bf16*>(from),
744 unpacket_traits<Packet2Xbf>::size);
745}
746
747template <>
748EIGEN_STRONG_INLINE Packet2Xbf ploadu<Packet2Xbf>(const bfloat16* from) {
749 EIGEN_DEBUG_UNALIGNED_LOAD return __riscv_vle16_v_bf16m2(reinterpret_cast<const __bf16*>(from),
750 unpacket_traits<Packet2Xbf>::size);
751}
752
753template <>
754EIGEN_STRONG_INLINE Packet2Xbf ploaddup<Packet2Xbf>(const bfloat16* from) {
755 return __riscv_vreinterpret_v_i16m2_bf16m2(ploaddup<Packet2Xs>(reinterpret_cast<const numext::int16_t*>(from)));
756}
757
758template <>
759EIGEN_STRONG_INLINE Packet2Xbf ploadquad<Packet2Xbf>(const bfloat16* from) {
760 return __riscv_vreinterpret_v_i16m2_bf16m2(ploadquad<Packet2Xs>(reinterpret_cast<const numext::int16_t*>(from)));
761}
762
763template <>
764EIGEN_STRONG_INLINE void pstore<bfloat16>(bfloat16* to, const Packet2Xbf& from) {
765 EIGEN_DEBUG_ALIGNED_STORE __riscv_vse16_v_bf16m2(reinterpret_cast<__bf16*>(to), from,
766 unpacket_traits<Packet2Xbf>::size);
767}
768
769template <>
770EIGEN_STRONG_INLINE void pstoreu<bfloat16>(bfloat16* to, const Packet2Xbf& from) {
771 EIGEN_DEBUG_UNALIGNED_STORE __riscv_vse16_v_bf16m2(reinterpret_cast<__bf16*>(to), from,
772 unpacket_traits<Packet2Xbf>::size);
773}
774
775template <>
776EIGEN_DEVICE_FUNC inline Packet2Xbf pgather<bfloat16, Packet2Xbf>(const bfloat16* from, Index stride) {
777 return __riscv_vlse16_v_bf16m2(reinterpret_cast<const __bf16*>(from), stride * sizeof(bfloat16),
778 unpacket_traits<Packet2Xbf>::size);
779}
780
781template <>
782EIGEN_DEVICE_FUNC inline void pscatter<bfloat16, Packet2Xbf>(bfloat16* to, const Packet2Xbf& from, Index stride) {
783 __riscv_vsse16(reinterpret_cast<__bf16*>(to), stride * sizeof(bfloat16), from, unpacket_traits<Packet2Xbf>::size);
784}
785
786template <>
787EIGEN_STRONG_INLINE bfloat16 pfirst<Packet2Xbf>(const Packet2Xbf& a) {
788 return numext::bit_cast<bfloat16>(__riscv_vmv_x_s_i16m2_i16(__riscv_vreinterpret_v_bf16m2_i16m2(a)));
789}
790
791template <>
792EIGEN_STRONG_INLINE Packet2Xbf psqrt(const Packet2Xbf& a) {
793 return F32ToBf16(psqrt<Packet4Xf>(Bf16ToF32(a)));
794}
795
796template <>
797EIGEN_STRONG_INLINE Packet2Xbf print<Packet2Xbf>(const Packet2Xbf& a) {
798 return F32ToBf16(print<Packet4Xf>(Bf16ToF32(a)));
799}
800
801template <>
802EIGEN_STRONG_INLINE Packet2Xbf pfloor<Packet2Xbf>(const Packet2Xbf& a) {
803 return F32ToBf16(pfloor<Packet4Xf>(Bf16ToF32(a)));
804}
805
806template <>
807EIGEN_STRONG_INLINE Packet2Xbf preverse(const Packet2Xbf& a) {
808 return __riscv_vreinterpret_v_i16m2_bf16m2(preverse<Packet2Xs>(__riscv_vreinterpret_v_bf16m2_i16m2(a)));
809}
810
811template <>
812EIGEN_STRONG_INLINE bfloat16 predux<Packet2Xbf>(const Packet2Xbf& a) {
813 return static_cast<bfloat16>(predux<Packet4Xf>(Bf16ToF32(a)));
814}
815
816template <>
817EIGEN_STRONG_INLINE bool predux_any(const Packet2Xbf& a) {
818 const PacketMask8 mask =
819 __riscv_vmsne_vx_u16m2_b8(__riscv_vreinterpret_v_bf16m2_u16m2(a), 0, unpacket_traits<Packet2Xbf>::size);
820 return __riscv_vcpop_m_b8(mask, unpacket_traits<Packet2Xbf>::size) != 0;
821}
822
823template <>
824EIGEN_STRONG_INLINE bool predux_all(const Packet2Xbf& a) {
825 const PacketMask8 mask = __riscv_vmfeq_vf_f32m4_b8(Bf16ToF32(a), 0.0f, unpacket_traits<Packet2Xbf>::size);
826 return __riscv_vcpop_m_b8(mask, unpacket_traits<Packet2Xbf>::size) == 0;
827}
828
829template <>
830EIGEN_STRONG_INLINE bfloat16 predux_mul<Packet2Xbf>(const Packet2Xbf& a) {
831 return static_cast<bfloat16>(predux_mul<Packet4Xf>(Bf16ToF32(a)));
832}
833
834template <>
835EIGEN_STRONG_INLINE bfloat16 predux_min<Packet2Xbf>(const Packet2Xbf& a) {
836 return static_cast<bfloat16>(predux_min<Packet4Xf>(Bf16ToF32(a)));
837}
838
839template <>
840EIGEN_STRONG_INLINE bfloat16 predux_max<Packet2Xbf>(const Packet2Xbf& a) {
841 return static_cast<bfloat16>(predux_max<Packet4Xf>(Bf16ToF32(a)));
842}
843
844template <int N>
845EIGEN_DEVICE_FUNC inline void ptranspose(PacketBlock<Packet2Xbf, N>& kernel) {
846 bfloat16 buffer[unpacket_traits<Packet2Xbf>::size * N];
847 int i = 0;
848
849 for (i = 0; i < N; i++) {
850 __riscv_vsse16(reinterpret_cast<__bf16*>(&buffer[i]), N * sizeof(bfloat16), kernel.packet[i],
851 unpacket_traits<Packet2Xbf>::size);
852 }
853
854 for (i = 0; i < N; i++) {
855 kernel.packet[i] = __riscv_vle16_v_bf16m2(reinterpret_cast<__bf16*>(&buffer[i * unpacket_traits<Packet2Xbf>::size]),
856 unpacket_traits<Packet2Xbf>::size);
857 }
858}
859
860template <typename Packet = Packet2Xbf>
861EIGEN_STRONG_INLINE std::enable_if_t<
862 std::is_same<Packet, Packet2Xbf>::value && (unpacket_traits<Packet2Xbf>::size % 8) == 0, Packet1Xbf>
863predux_half(const Packet2Xbf& a) {
864 return padd<Packet1Xbf>(__riscv_vget_v_bf16m2_bf16m1(a, 0), __riscv_vget_v_bf16m2_bf16m1(a, 1));
865}
866
867template <>
868EIGEN_STRONG_INLINE Packet1Xbf pcast<Packet1Xs, Packet1Xbf>(const Packet1Xs& a) {
869 return __riscv_vreinterpret_v_i16m1_bf16m1(a);
870}
871
872template <>
873EIGEN_STRONG_INLINE Packet2Xbf pcast<Packet2Xs, Packet2Xbf>(const Packet2Xs& a) {
874 return __riscv_vreinterpret_v_i16m2_bf16m2(a);
875}
876
877template <>
878EIGEN_STRONG_INLINE Packet1Xs pcast<Packet1Xbf, Packet1Xs>(const Packet1Xbf& a) {
879 return __riscv_vreinterpret_v_bf16m1_i16m1(a);
880}
881
882template <>
883EIGEN_STRONG_INLINE Packet2Xs pcast<Packet2Xbf, Packet2Xs>(const Packet2Xbf& a) {
884 return __riscv_vreinterpret_v_bf16m2_i16m2(a);
885}
886
887} // namespace internal
888} // namespace Eigen
889
890#endif // EIGEN_PACKET_MATH_BF16_RVV10_H