Eigen  5.0.1
 
Loading...
Searching...
No Matches
MatrixProductMMAbfloat16.h
1// SPDX-FileCopyrightText: The Eigen Authors
2// SPDX-License-Identifier: MPL-2.0
3
4#ifndef EIGEN_MATRIX_PRODUCT_MMA_BFLOAT16_ALTIVEC_H
5#define EIGEN_MATRIX_PRODUCT_MMA_BFLOAT16_ALTIVEC_H
6
7#if EIGEN_COMP_LLVM
8#define BFLOAT16_UNROLL _Pragma("unroll 8")
9#else
10#define BFLOAT16_UNROLL _Pragma("GCC unroll(8)")
11#endif
12
13namespace Eigen {
14
15namespace internal {
16
17template <bool zero>
18EIGEN_ALWAYS_INLINE Packet8bf loadBfloat16(const bfloat16* indexA) {
19 Packet8bf lhs1 = ploadu<Packet8bf>(indexA);
20 if (zero) {
21 Packet8bf lhs2 = pset1<Packet8bf>(Eigen::bfloat16(0));
22 return vec_mergeh(lhs1.m_val, lhs2.m_val);
23 } else {
24 return lhs1;
25 }
26}
27
28template <bool zero>
29EIGEN_ALWAYS_INLINE Packet8bf loadRhsBfloat16(const bfloat16* blockB, Index strideB, Index i) {
30 return loadBfloat16<zero>(blockB + strideB * i);
31}
32
33template <Index num_acc, Index num_packets, bool zero, bool rhsExtraCols, bool lhsExtraRows, Index num_rhs,
34 Index num_lhs>
35EIGEN_ALWAYS_INLINE void KLoop(const bfloat16* indexA, const bfloat16* indexB, __vector_quad (&quad_acc)[num_acc],
36 Index strideB, Index k, Index offsetB, Index extra_cols, Index extra_rows) {
37 Packet8bf lhs[num_lhs], rhs[num_rhs];
38
39 BFLOAT16_UNROLL
40 for (Index i = 0; i < (num_rhs - (rhsExtraCols ? 1 : 0)); i++) {
41 rhs[i] = loadRhsBfloat16<zero>(indexB + k * 4, strideB, i);
42 }
43 if (rhsExtraCols) {
44 rhs[num_rhs - 1] = loadRhsBfloat16<zero>(indexB + k * extra_cols - offsetB, strideB, num_rhs - 1);
45 }
46
47 indexA += k * (lhsExtraRows ? extra_rows : num_packets);
48 if (num_lhs == 1) {
49 lhs[0] = loadBfloat16<zero>(indexA);
50 } else {
51 BFLOAT16_UNROLL
52 for (Index j = 0; j < num_lhs; j += 2) {
53 Packet8bf lhs1 = ploadu<Packet8bf>(indexA + (j + 0) * (zero ? 4 : 8));
54 if (zero) {
55 Packet8bf lhs2 = pset1<Packet8bf>(Eigen::bfloat16(0));
56 lhs[j + 0] = vec_mergeh(lhs1.m_val, lhs2.m_val);
57 lhs[j + 1] = vec_mergel(lhs1.m_val, lhs2.m_val);
58 } else {
59 lhs[j + 0] = lhs1;
60 lhs[j + 1] = ploadu<Packet8bf>(indexA + (j + 1) * 8);
61 }
62 }
63 }
64
65 BFLOAT16_UNROLL
66 for (Index i = 0, x = 0; i < num_rhs; i++) {
67 BFLOAT16_UNROLL
68 for (Index j = 0; j < num_lhs; j++, x++) {
69 __builtin_mma_xvbf16ger2pp(&(quad_acc[x]), reinterpret_cast<Packet16uc>(rhs[i].m_val),
70 reinterpret_cast<Packet16uc>(lhs[j].m_val));
71 }
72 }
73}
74
75template <Index num_acc>
76EIGEN_ALWAYS_INLINE void zeroAccumulators(__vector_quad (&quad_acc)[num_acc]) {
77 BFLOAT16_UNROLL
78 for (Index k = 0; k < num_acc; k++) __builtin_mma_xxsetaccz(&(quad_acc[k]));
79}
80
81template <Index num_acc>
82EIGEN_ALWAYS_INLINE void disassembleAccumulators(__vector_quad (&quad_acc)[num_acc], Packet4f (&acc)[num_acc][4]) {
83 BFLOAT16_UNROLL
84 for (Index k = 0; k < num_acc; k++) __builtin_mma_disassemble_acc((void*)acc[k], &(quad_acc[k]));
85}
86
87template <Index num_acc, bool rhsExtraCols, bool lhsExtraRows, Index num_rhs, Index num_lhs>
88EIGEN_ALWAYS_INLINE void outputResults(Packet4f (&acc)[num_acc][4], Index rows, const Packet4f& pAlpha, float* result,
89 const Index extra_cols, Index extra_rows) {
90 BFLOAT16_UNROLL
91 for (Index i = 0, k = 0; i < num_rhs - (rhsExtraCols ? 1 : 0); i++, result += 4 * rows) {
92 BFLOAT16_UNROLL
93 for (Index j = 0; j < num_lhs; j++, k++) {
94 storeResults<false, lhsExtraRows>(acc[k], rows, pAlpha, result + j * 4, extra_cols, extra_rows);
95 }
96 }
97 if (rhsExtraCols) {
98 storeResults<rhsExtraCols, lhsExtraRows>(acc[num_acc - 1], rows, pAlpha, result, extra_cols, extra_rows);
99 }
100}
101
102template <const Index num_acc, const Index num_packets, bool rhsExtraCols, bool lhsExtraRows, bool multiIter = false>
103EIGEN_ALWAYS_INLINE void colLoopBodyIter(Index depth, Index rows, const Packet4f& pAlpha, const bfloat16* indexA,
104 const bfloat16* indexB, Index strideB, Index offsetB, float* result,
105 const Index extra_cols, const Index extra_rows) {
106 constexpr Index num_lhs = multiIter ? (num_packets / 4) : 1;
107 constexpr Index num_rhs = (num_acc + num_lhs - 1) / num_lhs;
108
109 for (Index offset_row = 0; offset_row < num_packets; offset_row += 4, indexA += (multiIter ? 0 : 8),
110 indexB += (multiIter ? (num_rhs * strideB) : 0), result += (multiIter ? (4 * rows * num_rhs) : 4)) {
111 Packet4f acc[num_acc][4];
112 __vector_quad quad_acc[num_acc];
113
114 zeroAccumulators<num_acc>(quad_acc);
115
116 Index k;
117 for (k = 0; k + 2 <= depth; k += 2) {
118 KLoop<num_acc, num_packets, false, rhsExtraCols, lhsExtraRows, num_rhs, num_lhs>(
119 indexA, indexB, quad_acc, strideB, k, offsetB, extra_cols, extra_rows);
120 }
121 if (depth & 1) {
122 KLoop<num_acc, num_packets, true, rhsExtraCols, lhsExtraRows, num_rhs, num_lhs>(
123 indexA - (multiIter ? 0 : offset_row), indexB, quad_acc, strideB, k, offsetB, extra_cols, extra_rows);
124 }
125
126 disassembleAccumulators<num_acc>(quad_acc, acc);
127
128 outputResults<num_acc, rhsExtraCols, lhsExtraRows, num_rhs, num_lhs>(acc, rows, pAlpha, result, extra_cols,
129 extra_rows);
130 }
131}
132
133#define MAX_BFLOAT16_ACC 8
134
135template <const Index num_acc, const Index num_packets, bool rhsExtraCols, bool lhsExtraRows>
136void colLoopBody(Index& col, Index depth, Index cols, Index rows, const Packet4f& pAlpha, const bfloat16* indexA,
137 const bfloat16* indexB, Index strideB, Index offsetB, float* result) {
138 constexpr Index step = (num_acc * 4); // each accumulator has 4 elements
139 const Index extra_cols = (rhsExtraCols) ? (cols & 3) : 0;
140 const Index extra_rows = (lhsExtraRows) ? (rows & 3) : 0;
141 constexpr bool multiIters = !rhsExtraCols && (num_acc == MAX_BFLOAT16_ACC);
142 constexpr bool normIters = multiIters && ((num_acc % (num_packets / 4)) == 0);
143
144 do {
145 colLoopBodyIter<num_acc, num_packets, rhsExtraCols, lhsExtraRows, normIters>(
146 depth, rows, pAlpha, indexA, indexB, strideB, offsetB, result, extra_cols, extra_rows);
147
148 indexB += strideB * num_acc;
149 result += rows * step;
150 } while (multiIters && (step <= cols - (col += step)));
151}
152
153template <const Index num_acc, const Index num_packets, bool rhsExtraCols, bool lhsExtraRows>
154EIGEN_ALWAYS_INLINE void colLoopBodyExtraN(Index col, Index depth, Index cols, Index rows, const Packet4f& pAlpha,
155 const bfloat16* indexA, const bfloat16* blockB, Index strideB, Index offsetB,
156 float* result) {
157 if (MAX_BFLOAT16_ACC > num_acc) {
158 colLoopBody<num_acc + (rhsExtraCols ? 1 : 0), num_packets, rhsExtraCols, lhsExtraRows>(
159 col, depth, cols, rows, pAlpha, indexA, blockB, strideB, offsetB, result);
160 }
161}
162
163template <const Index num_packets, bool rhsExtraCols, bool lhsExtraRows>
164void colLoopBodyExtra(Index col, Index depth, Index cols, Index rows, const Packet4f& pAlpha, const bfloat16* indexA,
165 const bfloat16* blockB, Index strideB, Index offsetB, float* result) {
166 switch ((cols - col) >> 2) {
167 case 7:
168 colLoopBodyExtraN<7, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
169 strideB, offsetB, result);
170 break;
171 case 6:
172 colLoopBodyExtraN<6, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
173 strideB, offsetB, result);
174 break;
175 case 5:
176 colLoopBodyExtraN<5, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
177 strideB, offsetB, result);
178 break;
179 case 4:
180 colLoopBodyExtraN<4, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
181 strideB, offsetB, result);
182 break;
183 case 3:
184 colLoopBodyExtraN<3, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
185 strideB, offsetB, result);
186 break;
187 case 2:
188 colLoopBodyExtraN<2, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
189 strideB, offsetB, result);
190 break;
191 case 1:
192 colLoopBodyExtraN<1, num_packets, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
193 strideB, offsetB, result);
194 break;
195 default:
196 if (rhsExtraCols) {
197 colLoopBody<1, num_packets, true, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
198 offsetB, result);
199 }
200 break;
201 }
202}
203
204template <const Index num_packets, bool lhsExtraRows = false>
205EIGEN_ALWAYS_INLINE void colLoops(Index depth, Index cols, Index rows, const Packet4f& pAlpha, const bfloat16* indexA,
206 const bfloat16* blockB, Index strideB, Index offsetB, float* result) {
207 Index col = 0;
208 if (cols >= (MAX_BFLOAT16_ACC * 4)) {
209 colLoopBody<MAX_BFLOAT16_ACC, num_packets, false, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB,
210 strideB, 0, result);
211 blockB += (strideB >> 2) * col;
212 result += rows * col;
213 }
214 if (cols & 3) {
215 colLoopBodyExtra<num_packets, true, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB, offsetB,
216 result);
217 } else {
218 colLoopBodyExtra<num_packets, false, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB, 0,
219 result);
220 }
221}
222
223EIGEN_ALWAYS_INLINE Packet8bf convertF32toBF16(const float* res) {
224 Packet16uc fp16[2];
225 __vector_pair fp16_vp = *reinterpret_cast<__vector_pair*>(const_cast<float*>(res));
226 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(fp16), &fp16_vp);
227 fp16[0] = __builtin_vsx_xvcvspbf16(fp16[0]);
228 fp16[1] = __builtin_vsx_xvcvspbf16(fp16[1]);
229 return vec_pack(reinterpret_cast<Packet4ui>(fp16[0]), reinterpret_cast<Packet4ui>(fp16[1]));
230}
231
232template <typename DataMapper, const Index size>
233EIGEN_ALWAYS_INLINE void convertArrayF32toBF16Col(float* result, Index col, Index rows, const DataMapper& res) {
234 const DataMapper res2 = res.getSubMapper(0, col);
235 Index row;
236 float* result2 = result + col * rows;
237 for (row = 0; row + 8 <= rows; row += 8, result2 += 8) {
238 // get and save block
239 PacketBlock<Packet8bf, size> block;
240 BFLOAT16_UNROLL
241 for (Index j = 0; j < size; j++) {
242 block.packet[j] = convertF32toBF16(result2 + j * rows);
243 }
244 res2.template storePacketBlock<Packet8bf, size>(row, 0, block);
245 }
246 // extra rows
247 if (row < rows) {
248 BFLOAT16_UNROLL
249 for (Index j = 0; j < size; j++) {
250 Packet8bf fp16 = convertF32toBF16(result2 + j * rows);
251 res2.template storePacketPartial<Packet8bf>(row, j, fp16, rows & 7);
252 }
253 }
254}
255
256template <const Index size, bool non_unit_stride = false>
257EIGEN_ALWAYS_INLINE void convertPointerF32toBF16(Index& i, float* result, Index rows, bfloat16*& dst,
258 Index resInc = 1) {
259 constexpr Index extra = ((size < 8) ? 8 : size);
260 while (i + size <= rows) {
261 PacketBlock<Packet8bf, (size + 7) / 8> r32;
262 r32.packet[0] = convertF32toBF16(result + i + 0);
263 if (size >= 16) {
264 r32.packet[1] = convertF32toBF16(result + i + 8);
265 }
266 if (size >= 32) {
267 r32.packet[2] = convertF32toBF16(result + i + 16);
268 r32.packet[3] = convertF32toBF16(result + i + 24);
269 }
270 storeBF16fromResult<size, non_unit_stride, 0>(dst, r32.packet[0], resInc, rows & 7);
271 if (size >= 16) {
272 storeBF16fromResult<size, non_unit_stride, 8>(dst, r32.packet[1], resInc);
273 }
274 if (size >= 32) {
275 storeBF16fromResult<size, non_unit_stride, 16>(dst, r32.packet[2], resInc);
276 storeBF16fromResult<size, non_unit_stride, 24>(dst, r32.packet[3], resInc);
277 }
278 i += extra;
279 dst += extra * resInc;
280 if (size != 32) break;
281 }
282}
283
284template <bool non_unit_stride = false>
285EIGEN_ALWAYS_INLINE void convertArrayPointerF32toBF16(float* result, Index rows, bfloat16* dst, Index resInc = 1) {
286 Index i = 0;
287 convertPointerF32toBF16<32, non_unit_stride>(i, result, rows, dst, resInc);
288 convertPointerF32toBF16<16, non_unit_stride>(i, result, rows, dst, resInc);
289 convertPointerF32toBF16<8, non_unit_stride>(i, result, rows, dst, resInc);
290 convertPointerF32toBF16<1, non_unit_stride>(i, result, rows, dst, resInc);
291}
292
293template <typename DataMapper>
294EIGEN_ALWAYS_INLINE void convertArrayF32toBF16(float* result, Index cols, Index rows, const DataMapper& res) {
295 Index col;
296 for (col = 0; col + 4 <= cols; col += 4) {
297 convertArrayF32toBF16Col<DataMapper, 4>(result, col, rows, res);
298 }
299 // extra cols
300 switch (cols - col) {
301 case 1:
302 convertArrayF32toBF16Col<DataMapper, 1>(result, col, rows, res);
303 break;
304 case 2:
305 convertArrayF32toBF16Col<DataMapper, 2>(result, col, rows, res);
306 break;
307 case 3:
308 convertArrayF32toBF16Col<DataMapper, 3>(result, col, rows, res);
309 break;
310 }
311}
312
313template <Index size>
314EIGEN_ALWAYS_INLINE void calcColLoops(const bfloat16*& indexA, Index& row, Index depth, Index cols, Index rows,
315 const Packet4f& pAlpha, const bfloat16* indexB, Index strideB, Index offsetA,
316 Index offsetB, Index bigSuffix, float* result) {
317 if ((size == 16) || (rows & size)) {
318 indexA += size * offsetA;
319 colLoops<size>(depth, cols, rows, pAlpha, indexA, indexB, strideB, offsetB, result + row);
320 row += size;
321 indexA += bigSuffix * size / 16;
322 }
323}
324
325template <typename DataMapper>
326void gemmMMAbfloat16(const DataMapper& res, const bfloat16* indexA, const bfloat16* indexB, Index rows, Index depth,
327 Index cols, bfloat16 alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
328 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
329 const Packet4f pAlpha = pset1<Packet4f>(falpha);
330 ei_declare_aligned_stack_constructed_variable(float, result, cols* rows, 0);
331
332 convertArrayBF16toF32<DataMapper>(result, cols, rows, res);
333
334 if (strideA == -1) strideA = depth;
335 if (strideB == -1) strideB = depth;
336 // Packing is done in blocks.
337 // There are 4 possible sizes of blocks
338 // Blocks of 8 columns with 16 elements (8x16)
339 // Blocks of 8 columns with 8 elements (8x8). This happens when there's 16 > rows >= 8
340 // Blocks of 8 columns with 4 elements (8x4). This happens when there's 8 > rows >= 4
341 // Blocks of 8 columns with < 4 elements. This happens when there's less than 4 remaining rows
342
343 // Loop for LHS standard block (8x16)
344 Index bigSuffix = (2 * 8) * (strideA - offsetA);
345 indexB += 4 * offsetB;
346 strideB *= 4;
347 offsetB *= 3;
348
349 Index row = 0;
350 while (row + 16 <= rows) {
351 calcColLoops<16>(indexA, row, depth, cols, rows, pAlpha, indexB, strideB, offsetA, offsetB, bigSuffix, result);
352 }
353 // LHS (8x8) block
354 calcColLoops<8>(indexA, row, depth, cols, rows, pAlpha, indexB, strideB, offsetA, offsetB, bigSuffix, result);
355 // LHS (8x4) block
356 calcColLoops<4>(indexA, row, depth, cols, rows, pAlpha, indexB, strideB, offsetA, offsetB, bigSuffix, result);
357 // extra rows
358 if (rows & 3) {
359 // This index is the beginning of remaining block.
360 colLoops<4, true>(depth, cols, rows, pAlpha, indexA, indexB, strideB, offsetB, result + row);
361 }
362
363 // Convert back to bfloat16
364 convertArrayF32toBF16<DataMapper>(result, cols, rows, res);
365}
366
367#undef MAX_BFLOAT16_ACC
368
369#if !EIGEN_ALTIVEC_DISABLE_MMA
370template <Index num_acc, typename LhsMapper, bool zero>
371EIGEN_ALWAYS_INLINE void loadVecLoop(Index k, LhsMapper& lhs, Packet8bf (&a0)[num_acc], Packet8bf b1) {
372 a0[k + 0] = lhs.template loadPacket<Packet8bf>(k * 4, 0);
373 if (!zero) {
374 b1 = lhs.template loadPacket<Packet8bf>(k * 4, 1);
375 }
376 if (num_acc > (k + 1)) {
377 a0[k + 1] = vec_mergel(a0[k + 0].m_val, b1.m_val);
378 }
379 a0[k + 0] = vec_mergeh(a0[k + 0].m_val, b1.m_val);
380}
381
382template <Index num_acc>
383EIGEN_ALWAYS_INLINE void multVec(__vector_quad (&quad_acc)[num_acc], const Packet8bf (&a0)[num_acc],
384 const Packet8bf& b0) {
385 BFLOAT16_UNROLL
386 for (Index k = 0; k < num_acc; k++) {
387 __builtin_mma_xvbf16ger2pp(&(quad_acc[k]), reinterpret_cast<Packet16uc>(b0.m_val),
388 reinterpret_cast<Packet16uc>(a0[k].m_val));
389 }
390}
391
392template <Index num_acc, typename LhsMapper, typename RhsMapper, bool zero, bool linear>
393EIGEN_ALWAYS_INLINE void vecColLoop(Index j, LhsMapper& lhs, RhsMapper& rhs, __vector_quad (&quad_acc)[num_acc]) {
394 Packet8bf a0[num_acc];
395 Packet8bf b1 = pset1<Packet8bf>(Eigen::bfloat16(0));
396 Packet8bf b0 = loadColData<RhsMapper, linear>(rhs, j);
397
398 if (zero) {
399 b0 = vec_mergeh(b0.m_val, b1.m_val);
400 }
401
402 using LhsSubMapper = typename LhsMapper::SubMapper;
403
404 LhsSubMapper lhs2 = lhs.getSubMapper(0, j);
405 BFLOAT16_UNROLL
406 for (Index k = 0; k < num_acc; k += 2) {
407 loadVecLoop<num_acc, LhsSubMapper, zero>(k, lhs2, a0, b1);
408 }
409
410 multVec<num_acc>(quad_acc, a0, b0);
411}
412
413#define MAX_BFLOAT16_VEC_ACC 8
414
415template <const Index num_acc, typename LhsMapper, typename RhsMapper, bool extraRows, bool linear>
416void colVecColLoopBody(Index& row, Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
417 float* result) {
418 constexpr Index step = (num_acc * 4);
419 const Index extra_rows = (extraRows) ? (rows & 3) : 0;
420 constexpr bool multiIters = !extraRows && (num_acc == MAX_BFLOAT16_VEC_ACC);
421
422 do {
423 Packet4f acc[num_acc][4];
424 __vector_quad quad_acc[num_acc];
425
426 zeroAccumulators<num_acc>(quad_acc);
427
428 using LhsSubMapper = typename LhsMapper::SubMapper;
429
430 LhsSubMapper lhs2 = lhs.getSubMapper(row, 0);
431 for (Index j = 0; j + 2 <= cend; j += 2) {
432 vecColLoop<num_acc, LhsSubMapper, RhsMapper, false, linear>(j, lhs2, rhs, quad_acc);
433 }
434 if (cend & 1) {
435 vecColLoop<num_acc, LhsSubMapper, RhsMapper, true, linear>(cend - 1, lhs2, rhs, quad_acc);
436 }
437
438 disassembleAccumulators<num_acc>(quad_acc, acc);
439
440 outputVecColResults<num_acc, extraRows>(acc, result, pAlpha, extra_rows);
441
442 result += step;
443 } while (multiIters && (step <= rows - (row += step)));
444}
445
446template <const Index num_acc, typename LhsMapper, typename RhsMapper, bool extraRows, bool linear>
447EIGEN_ALWAYS_INLINE void colVecColLoopBodyExtraN(Index& row, Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
448 const Packet4f& pAlpha, float* result) {
449 if (MAX_BFLOAT16_VEC_ACC > num_acc) {
450 colVecColLoopBody<num_acc + (extraRows ? 1 : 0), LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs,
451 pAlpha, result);
452 }
453}
454
455template <typename LhsMapper, typename RhsMapper, bool extraRows, bool linear>
456EIGEN_ALWAYS_INLINE void colVecColLoopBodyExtra(Index& row, Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
457 const Packet4f& pAlpha, float* result) {
458 switch ((rows - row) >> 2) {
459 case 7:
460 colVecColLoopBodyExtraN<7, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
461 break;
462 case 6:
463 colVecColLoopBodyExtraN<6, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
464 break;
465 case 5:
466 colVecColLoopBodyExtraN<5, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
467 break;
468 case 4:
469 colVecColLoopBodyExtraN<4, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
470 break;
471 case 3:
472 colVecColLoopBodyExtraN<3, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
473 break;
474 case 2:
475 colVecColLoopBodyExtraN<2, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
476 break;
477 case 1:
478 colVecColLoopBodyExtraN<1, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
479 break;
480 default:
481 if (extraRows) {
482 colVecColLoopBody<1, LhsMapper, RhsMapper, true, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
483 }
484 break;
485 }
486}
487
488template <typename LhsMapper, typename RhsMapper, bool linear>
489EIGEN_ALWAYS_INLINE void calcVecColLoops(Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
490 float* result) {
491 Index row = 0;
492 if (rows >= (MAX_BFLOAT16_VEC_ACC * 4)) {
493 colVecColLoopBody<MAX_BFLOAT16_VEC_ACC, LhsMapper, RhsMapper, false, linear>(row, cend, rows, lhs, rhs, pAlpha,
494 result);
495 result += row;
496 }
497 if (rows & 3) {
498 colVecColLoopBodyExtra<LhsMapper, RhsMapper, true, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
499 } else {
500 colVecColLoopBodyExtra<LhsMapper, RhsMapper, false, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
501 }
502}
503
504template <typename RhsMapper, typename LhsMapper, typename = void>
505struct UseMMAStride : std::false_type {
506 static EIGEN_ALWAYS_INLINE void run(Index j2, Index jend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
507 const Packet4f& pAlpha, float* result) {
508 using RhsSubMapper = typename RhsMapper::SubMapper;
509
510 RhsSubMapper rhs2 = rhs.getSubMapper(j2, 0);
511 calcVecColLoops<LhsMapper, RhsSubMapper, false>(jend - j2, rows, lhs, rhs2, pAlpha, result);
512 }
513};
514
515template <typename RhsMapper, typename LhsMapper>
516struct UseMMAStride<RhsMapper, LhsMapper,
517 std::enable_if_t<std::is_member_function_pointer<decltype(&RhsMapper::stride)>::value>>
518 : std::true_type {
519 static EIGEN_ALWAYS_INLINE void run(Index j2, Index jend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
520 const Packet4f& pAlpha, float* result) {
521 using RhsSubMapper = typename RhsMapper::SubMapper;
522
523 RhsSubMapper rhs2 = rhs.getSubMapper(j2, 0);
524 if (rhs.stride() == 1) {
525 calcVecColLoops<LhsMapper, RhsSubMapper, true>(jend - j2, rows, lhs, rhs2, pAlpha, result);
526 } else {
527 calcVecColLoops<LhsMapper, RhsSubMapper, false>(jend - j2, rows, lhs, rhs2, pAlpha, result);
528 }
529 }
530};
531
532template <typename LhsMapper, typename RhsMapper>
533void gemvMMA_bfloat16_col(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs, bfloat16* res,
534 Index resIncr, bfloat16 alpha) {
535 EIGEN_UNUSED_VARIABLE(resIncr);
536 eigen_internal_assert(resIncr == 1);
537
538 // The following copy tells the compiler that lhs's attributes are not modified outside this function
539 // This helps GCC to generate proper code.
540 LhsMapper lhs(alhs);
541 RhsMapper rhs2(rhs);
542
543 const Index lhsStride = lhs.stride();
544
545 // TODO: improve the following heuristic:
546 const Index block_cols = cols < 128 ? cols : (lhsStride * sizeof(bfloat16) < 16000 ? 16 : 8);
547 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
548 Packet4f pAlpha = pset1<Packet4f>(falpha);
549
550 ei_declare_aligned_stack_constructed_variable(float, result, rows, 0);
551
552 convertArrayPointerBF16toF32(result, 1, rows, res);
553
554 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
555 Index jend = numext::mini(j2 + block_cols, cols);
556
557 using LhsSubMapper = typename LhsMapper::SubMapper;
558
559 LhsSubMapper lhs2 = lhs.getSubMapper(0, j2);
560 UseMMAStride<RhsMapper, LhsSubMapper>::run(j2, jend, rows, lhs2, rhs2, pAlpha, result);
561 }
562
563 convertArrayPointerF32toBF16(result, rows, res);
564}
565
566static Packet16uc p16uc_ELEMENT_VEC3 = {0x0c, 0x0d, 0x0e, 0x0f, 0x1c, 0x1d, 0x1e, 0x1f,
567 0x0c, 0x0d, 0x0e, 0x0f, 0x1c, 0x1d, 0x1e, 0x1f};
568
569template <Index num_acc>
570EIGEN_ALWAYS_INLINE void preduxVecResults2(Packet4f (&acc)[num_acc][4], Index k) {
571 if (num_acc > (k + 1)) {
572 acc[k][0] = vec_mergeh(acc[k][0], acc[k + 1][0]);
573 acc[k][1] = vec_mergeo(acc[k][1], acc[k + 1][1]);
574 acc[k][2] = vec_mergel(acc[k][2], acc[k + 1][2]);
575 acc[k][3] = vec_perm(acc[k][3], acc[k + 1][3], p16uc_ELEMENT_VEC3);
576
577 acc[k][0] = (acc[k][0] + acc[k][2]) + (acc[k][1] + acc[k][3]);
578 } else {
579 acc[k][0] = vec_mergeh(acc[k][0], acc[k][1]);
580 acc[k][0] += vec_mergel(acc[k][2], acc[k][3]);
581#ifdef _BIG_ENDIAN
582 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 12);
583#else
584 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 4);
585#endif
586 }
587}
588
589template <Index num_acc>
590EIGEN_ALWAYS_INLINE void preduxVecResults(Packet4f (&acc)[num_acc][4]) {
591 BFLOAT16_UNROLL
592 for (Index k = 0; k < num_acc; k += 4) {
593 preduxVecResults2<num_acc>(acc, k + 0);
594 if (num_acc > (k + 2)) {
595 preduxVecResults2<num_acc>(acc, k + 2);
596 acc[k + 0][0] = reinterpret_cast<Packet4f>(
597 vec_mergeh(reinterpret_cast<Packet2ul>(acc[k + 0][0]), reinterpret_cast<Packet2ul>(acc[k + 2][0])));
598 }
599 }
600}
601
602template <Index num_acc, typename LhsMapper, typename RhsMapper, bool extra>
603EIGEN_ALWAYS_INLINE void multVecLoop(__vector_quad (&quad_acc)[num_acc], const LhsMapper& lhs, RhsMapper& rhs, Index j,
604 Index extra_cols) {
605 Packet8bf a0[num_acc], b0;
606
607 if (extra) {
608 b0 = rhs.template loadPacketPartial<Packet8bf>(j, extra_cols);
609 } else {
610 b0 = rhs.template loadPacket<Packet8bf>(j);
611 }
612
613 const LhsMapper lhs2 = lhs.getSubMapper(0, j);
614 BFLOAT16_UNROLL
615 for (Index k = 0; k < num_acc; k++) {
616 if (extra) {
617 a0[k] = lhs2.template loadPacketPartial<Packet8bf>(k, 0, extra_cols);
618 } else {
619 a0[k] = lhs2.template loadPacket<Packet8bf>(k, 0);
620 }
621 }
622
623 multVec<num_acc>(quad_acc, a0, b0);
624}
625
626template <Index num_acc, typename LhsMapper, typename RhsMapper>
627EIGEN_ALWAYS_INLINE void vecLoop(Index cols, const LhsMapper& lhs, RhsMapper& rhs, __vector_quad (&quad_acc)[num_acc],
628 Index extra_cols) {
629 Index j = 0;
630 for (; j + 8 <= cols; j += 8) {
631 multVecLoop<num_acc, LhsMapper, RhsMapper, false>(quad_acc, lhs, rhs, j, extra_cols);
632 }
633
634 if (extra_cols) {
635 multVecLoop<num_acc, LhsMapper, RhsMapper, true>(quad_acc, lhs, rhs, j, extra_cols);
636 }
637}
638
639template <const Index num_acc, typename LhsMapper, typename RhsMapper>
640void colVecLoopBody(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
641 float* result) {
642 constexpr bool multiIters = (num_acc == MAX_BFLOAT16_VEC_ACC);
643 const Index extra_cols = (cols & 7);
644
645 do {
646 Packet4f acc[num_acc][4];
647 __vector_quad quad_acc[num_acc];
648
649 zeroAccumulators<num_acc>(quad_acc);
650
651 const LhsMapper lhs2 = lhs.getSubMapper(row, 0);
652 vecLoop<num_acc, LhsMapper, RhsMapper>(cols, lhs2, rhs, quad_acc, extra_cols);
653
654 disassembleAccumulators<num_acc>(quad_acc, acc);
655
656 preduxVecResults<num_acc>(acc);
657
658 outputVecResults<num_acc>(acc, result, pAlpha);
659
660 result += num_acc;
661 } while (multiIters && (num_acc <= rows - (row += num_acc)));
662}
663
664template <const Index num_acc, typename LhsMapper, typename RhsMapper>
665EIGEN_ALWAYS_INLINE void colVecLoopBodyExtraN(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs,
666 const Packet4f& pAlpha, float* result) {
667 if (MAX_BFLOAT16_VEC_ACC > num_acc) {
668 colVecLoopBody<num_acc, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
669 }
670}
671
672template <typename LhsMapper, typename RhsMapper>
673EIGEN_ALWAYS_INLINE void colVecLoopBodyExtra(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs,
674 const Packet4f& pAlpha, float* result) {
675 switch (rows - row) {
676 case 7:
677 colVecLoopBodyExtraN<7, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
678 break;
679 case 6:
680 colVecLoopBodyExtraN<6, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
681 break;
682 case 5:
683 colVecLoopBodyExtraN<5, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
684 break;
685 case 4:
686 colVecLoopBodyExtraN<4, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
687 break;
688 case 3:
689 colVecLoopBodyExtraN<3, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
690 break;
691 case 2:
692 colVecLoopBodyExtraN<2, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
693 break;
694 case 1:
695 colVecLoopBodyExtraN<1, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
696 break;
697 }
698}
699
700template <typename LhsMapper, typename RhsMapper>
701EIGEN_ALWAYS_INLINE void calcVecLoops(Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
702 float* result) {
703 Index row = 0;
704 if (rows >= MAX_BFLOAT16_VEC_ACC) {
705 colVecLoopBody<MAX_BFLOAT16_VEC_ACC, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
706 result += row;
707 }
708 colVecLoopBodyExtra<LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
709}
710
711template <typename LhsMapper, typename RhsMapper>
712EIGEN_STRONG_INLINE void gemvMMA_bfloat16_row(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs,
713 bfloat16* res, Index resIncr, bfloat16 alpha) {
714 typedef typename RhsMapper::LinearMapper LinearMapper;
715
716 // The following copy tells the compiler that lhs's attributes are not modified outside this function
717 // This helps GCC to generate proper code.
718 LhsMapper lhs(alhs);
719 LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
720
721 eigen_internal_assert(rhs.stride() == 1);
722
723 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
724 const Packet4f pAlpha = pset1<Packet4f>(falpha);
725
726 ei_declare_aligned_stack_constructed_variable(float, result, rows, 0);
727 if (resIncr == 1) {
728 convertArrayPointerBF16toF32(result, 1, rows, res);
729 } else {
730 convertArrayPointerBF16toF32<true>(result, 1, rows, res, resIncr);
731 }
732 calcVecLoops<LhsMapper, LinearMapper>(cols, rows, lhs, rhs2, pAlpha, result);
733 if (resIncr == 1) {
734 convertArrayPointerF32toBF16(result, rows, res);
735 } else {
736 convertArrayPointerF32toBF16<true>(result, rows, res, resIncr);
737 }
738}
739#endif
740
741#undef MAX_BFLOAT16_VEC_ACC
742#undef BFLOAT16_UNROLL
743
744} // namespace internal
745} // namespace Eigen
746#endif // EIGEN_MATRIX_PRODUCT_MMA_BFLOAT16_ALTIVEC_H