Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorContractionGpu.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014-2015 Benoit Steiner <benoit.steiner.goog@gmail.com>
5// Copyright (C) 2015 Navdeep Jaitly <ndjaitly@google.com>
6// Copyright (C) 2014 Eric Martin <eric@ericmart.in>
7//
8// This Source Code Form is subject to the terms of the Mozilla
9// Public License v. 2.0. If a copy of the MPL was not distributed
10// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
11// SPDX-License-Identifier: MPL-2.0
12
13#ifndef EIGEN_TENSOR_TENSOR_CONTRACTION_GPU_H
14#define EIGEN_TENSOR_TENSOR_CONTRACTION_GPU_H
15
16#if defined(EIGEN_USE_GPU) && defined(EIGEN_GPUCC)
17
18// IWYU pragma: private
19#include "./InternalHeaderCheck.h"
20
21namespace Eigen {
22
23template <typename Scalar, typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper,
24 bool needs_edge_check>
25__device__ EIGEN_STRONG_INLINE void EigenContractionKernelInternal(const LhsMapper lhs, const RhsMapper rhs,
26 const OutputMapper output, Scalar* lhs_shmem,
27 Scalar* rhs_shmem, const Index m_size,
28 const Index n_size, const Index k_size) {
29 const Index m_block_idx = blockIdx.x;
30 const Index n_block_idx = blockIdx.y;
31
32 const Index base_m = 64 * m_block_idx;
33 const Index base_n = 64 * n_block_idx;
34 const Index thread_x = threadIdx.x;
35 const Index thread_y = threadIdx.y;
36 const Index thread_z = threadIdx.z;
37
38 // declare and initialize 64 registers for output 8x8 block
39
40 // prefetch registers
41 Scalar lhs_pf0;
42 Scalar lhs_pf1;
43 Scalar lhs_pf2;
44 Scalar lhs_pf3;
45 Scalar lhs_pf4;
46 Scalar lhs_pf5;
47 Scalar lhs_pf6;
48 Scalar lhs_pf7;
49
50 Scalar rhs_pf0;
51 Scalar rhs_pf1;
52 Scalar rhs_pf2;
53 Scalar rhs_pf3;
54 Scalar rhs_pf4;
55 Scalar rhs_pf5;
56 Scalar rhs_pf6;
57 Scalar rhs_pf7;
58
59 // shared memory is formatted
60 // (contract idx in block, nocontract idx in block, block idx)
61 // where block idx is column major. This transposition limits the number of
62 // bank conflicts when reading the LHS. The core idea is that since the contracting
63 // index is shared by both sides, then the contracting index should be in threadIdx.x.
64
65 // On the LHS, we pad each row inside of each block with an extra element. This makes
66 // each block 8 rows of 9 elements, which is 72 elements. This gives no bank conflicts
67 // on writes and very few 2-way conflicts on reads. There is an 8x8 grid of these blocks.
68
69 // On the RHS we just add 8 padding elements to the end of each block. This gives no bank
70 // conflicts on writes and also none on reads.
71
72 // storage indices
73 const Index lhs_store_idx_base = thread_y * 72 + thread_x * 9 + thread_z;
74 const Index rhs_store_idx_base = thread_y * 72 + thread_z * 8 + thread_x;
75
76 const Index lhs_store_idx_0 = lhs_store_idx_base + 576 * 0;
77 const Index lhs_store_idx_1 = lhs_store_idx_base + 576 * 1;
78 const Index lhs_store_idx_2 = lhs_store_idx_base + 576 * 2;
79 const Index lhs_store_idx_3 = lhs_store_idx_base + 576 * 3;
80 const Index lhs_store_idx_4 = lhs_store_idx_base + 576 * 4;
81 const Index lhs_store_idx_5 = lhs_store_idx_base + 576 * 5;
82 const Index lhs_store_idx_6 = lhs_store_idx_base + 576 * 6;
83 const Index lhs_store_idx_7 = lhs_store_idx_base + 576 * 7;
84
85 const Index rhs_store_idx_0 = rhs_store_idx_base + 576 * 0;
86 const Index rhs_store_idx_1 = rhs_store_idx_base + 576 * 1;
87 const Index rhs_store_idx_2 = rhs_store_idx_base + 576 * 2;
88 const Index rhs_store_idx_3 = rhs_store_idx_base + 576 * 3;
89 const Index rhs_store_idx_4 = rhs_store_idx_base + 576 * 4;
90 const Index rhs_store_idx_5 = rhs_store_idx_base + 576 * 5;
91 const Index rhs_store_idx_6 = rhs_store_idx_base + 576 * 6;
92 const Index rhs_store_idx_7 = rhs_store_idx_base + 576 * 7;
93
94 // in the loading code, the following variables are important:
95 // thread_x: the vertical position in an 8x8 block
96 // thread_y: the vertical index of the 8x8 block in the grid
97 // thread_z: the horizontal position in an 8x8 block
98 // k: the horizontal index of the 8x8 block in the grid
99 //
100 // The k parameter is implicit (it was the loop counter for a loop that went
101 // from 0 to <8, but now that loop is unrolled in the below code).
102
103 const Index load_idx_vert = thread_x + 8 * thread_y;
104 const Index lhs_vert = base_m + load_idx_vert;
105
106#define prefetchIntoRegisters(base_k) \
107 { \
108 lhs_pf0 = conv(0); \
109 lhs_pf1 = conv(0); \
110 lhs_pf2 = conv(0); \
111 lhs_pf3 = conv(0); \
112 lhs_pf4 = conv(0); \
113 lhs_pf5 = conv(0); \
114 lhs_pf6 = conv(0); \
115 lhs_pf7 = conv(0); \
116 \
117 rhs_pf0 = conv(0); \
118 rhs_pf1 = conv(0); \
119 rhs_pf2 = conv(0); \
120 rhs_pf3 = conv(0); \
121 rhs_pf4 = conv(0); \
122 rhs_pf5 = conv(0); \
123 rhs_pf6 = conv(0); \
124 rhs_pf7 = conv(0); \
125 \
126 if (!needs_edge_check || lhs_vert < m_size) { \
127 const Index lhs_horiz_0 = base_k + thread_z + 0 * 8; \
128 const Index lhs_horiz_1 = base_k + thread_z + 1 * 8; \
129 const Index lhs_horiz_2 = base_k + thread_z + 2 * 8; \
130 const Index lhs_horiz_3 = base_k + thread_z + 3 * 8; \
131 const Index lhs_horiz_4 = base_k + thread_z + 4 * 8; \
132 const Index lhs_horiz_5 = base_k + thread_z + 5 * 8; \
133 const Index lhs_horiz_6 = base_k + thread_z + 6 * 8; \
134 const Index lhs_horiz_7 = base_k + thread_z + 7 * 8; \
135 \
136 if (!needs_edge_check || lhs_horiz_7 < k_size) { \
137 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
138 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
139 lhs_pf2 = lhs(lhs_vert, lhs_horiz_2); \
140 lhs_pf3 = lhs(lhs_vert, lhs_horiz_3); \
141 lhs_pf4 = lhs(lhs_vert, lhs_horiz_4); \
142 lhs_pf5 = lhs(lhs_vert, lhs_horiz_5); \
143 lhs_pf6 = lhs(lhs_vert, lhs_horiz_6); \
144 lhs_pf7 = lhs(lhs_vert, lhs_horiz_7); \
145 } else if (lhs_horiz_6 < k_size) { \
146 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
147 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
148 lhs_pf2 = lhs(lhs_vert, lhs_horiz_2); \
149 lhs_pf3 = lhs(lhs_vert, lhs_horiz_3); \
150 lhs_pf4 = lhs(lhs_vert, lhs_horiz_4); \
151 lhs_pf5 = lhs(lhs_vert, lhs_horiz_5); \
152 lhs_pf6 = lhs(lhs_vert, lhs_horiz_6); \
153 } else if (lhs_horiz_5 < k_size) { \
154 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
155 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
156 lhs_pf2 = lhs(lhs_vert, lhs_horiz_2); \
157 lhs_pf3 = lhs(lhs_vert, lhs_horiz_3); \
158 lhs_pf4 = lhs(lhs_vert, lhs_horiz_4); \
159 lhs_pf5 = lhs(lhs_vert, lhs_horiz_5); \
160 } else if (lhs_horiz_4 < k_size) { \
161 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
162 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
163 lhs_pf2 = lhs(lhs_vert, lhs_horiz_2); \
164 lhs_pf3 = lhs(lhs_vert, lhs_horiz_3); \
165 lhs_pf4 = lhs(lhs_vert, lhs_horiz_4); \
166 } else if (lhs_horiz_3 < k_size) { \
167 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
168 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
169 lhs_pf2 = lhs(lhs_vert, lhs_horiz_2); \
170 lhs_pf3 = lhs(lhs_vert, lhs_horiz_3); \
171 } else if (lhs_horiz_2 < k_size) { \
172 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
173 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
174 lhs_pf2 = lhs(lhs_vert, lhs_horiz_2); \
175 } else if (lhs_horiz_1 < k_size) { \
176 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
177 lhs_pf1 = lhs(lhs_vert, lhs_horiz_1); \
178 } else if (lhs_horiz_0 < k_size) { \
179 lhs_pf0 = lhs(lhs_vert, lhs_horiz_0); \
180 } \
181 } \
182 \
183 const Index rhs_vert = base_k + load_idx_vert; \
184 if (!needs_edge_check || rhs_vert < k_size) { \
185 const Index rhs_horiz_0 = base_n + thread_z + 0 * 8; \
186 const Index rhs_horiz_1 = base_n + thread_z + 1 * 8; \
187 const Index rhs_horiz_2 = base_n + thread_z + 2 * 8; \
188 const Index rhs_horiz_3 = base_n + thread_z + 3 * 8; \
189 const Index rhs_horiz_4 = base_n + thread_z + 4 * 8; \
190 const Index rhs_horiz_5 = base_n + thread_z + 5 * 8; \
191 const Index rhs_horiz_6 = base_n + thread_z + 6 * 8; \
192 const Index rhs_horiz_7 = base_n + thread_z + 7 * 8; \
193 \
194 if (rhs_horiz_7 < n_size) { \
195 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
196 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
197 rhs_pf2 = rhs(rhs_vert, rhs_horiz_2); \
198 rhs_pf3 = rhs(rhs_vert, rhs_horiz_3); \
199 rhs_pf4 = rhs(rhs_vert, rhs_horiz_4); \
200 rhs_pf5 = rhs(rhs_vert, rhs_horiz_5); \
201 rhs_pf6 = rhs(rhs_vert, rhs_horiz_6); \
202 rhs_pf7 = rhs(rhs_vert, rhs_horiz_7); \
203 } else if (rhs_horiz_6 < n_size) { \
204 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
205 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
206 rhs_pf2 = rhs(rhs_vert, rhs_horiz_2); \
207 rhs_pf3 = rhs(rhs_vert, rhs_horiz_3); \
208 rhs_pf4 = rhs(rhs_vert, rhs_horiz_4); \
209 rhs_pf5 = rhs(rhs_vert, rhs_horiz_5); \
210 rhs_pf6 = rhs(rhs_vert, rhs_horiz_6); \
211 } else if (rhs_horiz_5 < n_size) { \
212 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
213 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
214 rhs_pf2 = rhs(rhs_vert, rhs_horiz_2); \
215 rhs_pf3 = rhs(rhs_vert, rhs_horiz_3); \
216 rhs_pf4 = rhs(rhs_vert, rhs_horiz_4); \
217 rhs_pf5 = rhs(rhs_vert, rhs_horiz_5); \
218 } else if (rhs_horiz_4 < n_size) { \
219 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
220 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
221 rhs_pf2 = rhs(rhs_vert, rhs_horiz_2); \
222 rhs_pf3 = rhs(rhs_vert, rhs_horiz_3); \
223 rhs_pf4 = rhs(rhs_vert, rhs_horiz_4); \
224 } else if (rhs_horiz_3 < n_size) { \
225 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
226 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
227 rhs_pf2 = rhs(rhs_vert, rhs_horiz_2); \
228 rhs_pf3 = rhs(rhs_vert, rhs_horiz_3); \
229 } else if (rhs_horiz_2 < n_size) { \
230 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
231 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
232 rhs_pf2 = rhs(rhs_vert, rhs_horiz_2); \
233 } else if (rhs_horiz_1 < n_size) { \
234 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
235 rhs_pf1 = rhs(rhs_vert, rhs_horiz_1); \
236 } else if (rhs_horiz_0 < n_size) { \
237 rhs_pf0 = rhs(rhs_vert, rhs_horiz_0); \
238 } \
239 } \
240 }
241
242#define writeRegToShmem() \
243 lhs_shmem[lhs_store_idx_0] = lhs_pf0; \
244 rhs_shmem[rhs_store_idx_0] = rhs_pf0; \
245 \
246 lhs_shmem[lhs_store_idx_1] = lhs_pf1; \
247 rhs_shmem[rhs_store_idx_1] = rhs_pf1; \
248 \
249 lhs_shmem[lhs_store_idx_2] = lhs_pf2; \
250 rhs_shmem[rhs_store_idx_2] = rhs_pf2; \
251 \
252 lhs_shmem[lhs_store_idx_3] = lhs_pf3; \
253 rhs_shmem[rhs_store_idx_3] = rhs_pf3; \
254 \
255 lhs_shmem[lhs_store_idx_4] = lhs_pf4; \
256 rhs_shmem[rhs_store_idx_4] = rhs_pf4; \
257 \
258 lhs_shmem[lhs_store_idx_5] = lhs_pf5; \
259 rhs_shmem[rhs_store_idx_5] = rhs_pf5; \
260 \
261 lhs_shmem[lhs_store_idx_6] = lhs_pf6; \
262 rhs_shmem[rhs_store_idx_6] = rhs_pf6; \
263 \
264 lhs_shmem[lhs_store_idx_7] = lhs_pf7; \
265 rhs_shmem[rhs_store_idx_7] = rhs_pf7;
266
267 // declare and initialize result array
268#define res(i, j) _res_##i##j
269#define initResultRow(i) \
270 Scalar res(i, 0) = conv(0); \
271 Scalar res(i, 1) = conv(0); \
272 Scalar res(i, 2) = conv(0); \
273 Scalar res(i, 3) = conv(0); \
274 Scalar res(i, 4) = conv(0); \
275 Scalar res(i, 5) = conv(0); \
276 Scalar res(i, 6) = conv(0); \
277 Scalar res(i, 7) = conv(0);
278
279 internal::scalar_cast_op<int, Scalar> conv;
280 initResultRow(0);
281 initResultRow(1);
282 initResultRow(2);
283 initResultRow(3);
284 initResultRow(4);
285 initResultRow(5);
286 initResultRow(6);
287 initResultRow(7);
288#undef initResultRow
289
290 for (Index base_k = 0; base_k < k_size; base_k += 64) {
291 // wait for previous iteration to finish with shmem. Despite common sense,
292 // the code is a bit faster with this here than at bottom of loop
293 __syncthreads();
294
295 prefetchIntoRegisters(base_k);
296 writeRegToShmem();
297
298#undef prefetchIntoRegisters
299#undef writeRegToShmem
300
301 // wait for shared mem packing to be done before starting computation
302 __syncthreads();
303
304 // compute 8x8 matrix product by outer product. This involves packing one column
305 // of LHS and one row of RHS into registers (takes 16 registers).
306
307#define lcol(i) _lcol##i
308 Scalar lcol(0);
309 Scalar lcol(1);
310 Scalar lcol(2);
311 Scalar lcol(3);
312 Scalar lcol(4);
313 Scalar lcol(5);
314 Scalar lcol(6);
315 Scalar lcol(7);
316
317#define rrow(j) _rrow##j
318 Scalar rrow(0);
319 Scalar rrow(1);
320 Scalar rrow(2);
321 Scalar rrow(3);
322 Scalar rrow(4);
323 Scalar rrow(5);
324 Scalar rrow(6);
325 Scalar rrow(7);
326
327 // Now x corresponds to k, y to m, and z to n
328 const Scalar* lhs_block = &lhs_shmem[thread_x + 9 * thread_y];
329 const Scalar* rhs_block = &rhs_shmem[thread_x + 8 * thread_z];
330
331#define lhs_element(i, j) lhs_block[72 * ((i) + 8 * (j))]
332#define rhs_element(i, j) rhs_block[72 * ((i) + 8 * (j))]
333
334#define loadData(i, j) \
335 lcol(0) = lhs_element(0, j); \
336 rrow(0) = rhs_element(i, 0); \
337 lcol(1) = lhs_element(1, j); \
338 rrow(1) = rhs_element(i, 1); \
339 lcol(2) = lhs_element(2, j); \
340 rrow(2) = rhs_element(i, 2); \
341 lcol(3) = lhs_element(3, j); \
342 rrow(3) = rhs_element(i, 3); \
343 lcol(4) = lhs_element(4, j); \
344 rrow(4) = rhs_element(i, 4); \
345 lcol(5) = lhs_element(5, j); \
346 rrow(5) = rhs_element(i, 5); \
347 lcol(6) = lhs_element(6, j); \
348 rrow(6) = rhs_element(i, 6); \
349 lcol(7) = lhs_element(7, j); \
350 rrow(7) = rhs_element(i, 7);
351
352#define computeCol(j) \
353 res(0, j) += lcol(0) * rrow(j); \
354 res(1, j) += lcol(1) * rrow(j); \
355 res(2, j) += lcol(2) * rrow(j); \
356 res(3, j) += lcol(3) * rrow(j); \
357 res(4, j) += lcol(4) * rrow(j); \
358 res(5, j) += lcol(5) * rrow(j); \
359 res(6, j) += lcol(6) * rrow(j); \
360 res(7, j) += lcol(7) * rrow(j);
361
362#define computePass(i) \
363 loadData(i, i); \
364 \
365 computeCol(0); \
366 computeCol(1); \
367 computeCol(2); \
368 computeCol(3); \
369 computeCol(4); \
370 computeCol(5); \
371 computeCol(6); \
372 computeCol(7);
373
374 computePass(0);
375 computePass(1);
376 computePass(2);
377 computePass(3);
378 computePass(4);
379 computePass(5);
380 computePass(6);
381 computePass(7);
382
383#undef lcol
384#undef rrow
385#undef lhs_element
386#undef rhs_element
387#undef loadData
388#undef computeCol
389#undef computePass
390 } // end loop over k
391
392 // we've now iterated over all of the large (ie width 64) k blocks and
393 // accumulated results in registers. At this point thread (x, y, z) contains
394 // the sum across all big k blocks of the product of little k block of index (x, y)
395 // with block of index (y, z). To compute the final output, we need to reduce
396 // the 8 threads over y by summation.
397 // HIP uses non-sync warp shuffles; CUDA requires the _sync variants.
398#if defined(EIGEN_HIPCC)
399#define shuffleInc(i, j, mask) res(i, j) += __shfl_xor(res(i, j), mask)
400#else
401#define shuffleInc(i, j, mask) res(i, j) += __shfl_xor_sync(0xFFFFFFFF, res(i, j), mask)
402#endif
403
404#define reduceRow(i, mask) \
405 shuffleInc(i, 0, mask); \
406 shuffleInc(i, 1, mask); \
407 shuffleInc(i, 2, mask); \
408 shuffleInc(i, 3, mask); \
409 shuffleInc(i, 4, mask); \
410 shuffleInc(i, 5, mask); \
411 shuffleInc(i, 6, mask); \
412 shuffleInc(i, 7, mask);
413
414#define reduceMatrix(mask) \
415 reduceRow(0, mask); \
416 reduceRow(1, mask); \
417 reduceRow(2, mask); \
418 reduceRow(3, mask); \
419 reduceRow(4, mask); \
420 reduceRow(5, mask); \
421 reduceRow(6, mask); \
422 reduceRow(7, mask);
423
424 // actually perform the reduction, now each thread of index (_, y, z)
425 // contains the correct values in its registers that belong in the output
426 // block
427 reduceMatrix(1);
428 reduceMatrix(2);
429 reduceMatrix(4);
430
431#undef shuffleInc
432#undef reduceRow
433#undef reduceMatrix
434
435 // now we need to copy the 64 values into main memory. We can't split work
436 // among threads because all variables are in registers. There's 2 ways
437 // to do this:
438 // (1) have 1 thread do 64 writes from registers into global memory
439 // (2) have 1 thread do 64 writes into shared memory, and then 8 threads
440 // each do 8 writes into global memory. We can just overwrite the shared
441 // memory from the problem we just solved.
442 // (2) is slightly faster than (1) due to less branching and more ILP
443
444 // TODO: won't yield much gain, but could just use currently unused shared mem
445 // and then we won't have to sync
446 // wait for shared mem to be out of use
447 __syncthreads();
448
449#define writeResultShmem(i, j) lhs_shmem[i + 8 * thread_y + 64 * thread_z + 512 * j] = res(i, j);
450
451#define writeRow(i) \
452 writeResultShmem(i, 0); \
453 writeResultShmem(i, 1); \
454 writeResultShmem(i, 2); \
455 writeResultShmem(i, 3); \
456 writeResultShmem(i, 4); \
457 writeResultShmem(i, 5); \
458 writeResultShmem(i, 6); \
459 writeResultShmem(i, 7);
460
461 if (thread_x == 0) {
462 writeRow(0);
463 writeRow(1);
464 writeRow(2);
465 writeRow(3);
466 writeRow(4);
467 writeRow(5);
468 writeRow(6);
469 writeRow(7);
470 }
471#undef writeResultShmem
472#undef writeRow
473
474 const int max_i_write = numext::mini((int)((m_size - base_m - thread_y + 7) / 8), 8);
475 const int max_j_write = numext::mini((int)((n_size - base_n - thread_z + 7) / 8), 8);
476
477 if (thread_x < max_i_write) {
478 if (max_j_write == 8) {
479 // TODO: Can we trade bank conflicts for coalesced writes?
480 Scalar val0 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 0];
481 Scalar val1 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 1];
482 Scalar val2 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 2];
483 Scalar val3 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 3];
484 Scalar val4 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 4];
485 Scalar val5 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 5];
486 Scalar val6 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 6];
487 Scalar val7 = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * 7];
488
489 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 0) = val0;
490 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 1) = val1;
491 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 2) = val2;
492 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 3) = val3;
493 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 4) = val4;
494 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 5) = val5;
495 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 6) = val6;
496 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * 7) = val7;
497 } else {
498#pragma unroll 7
499 for (int j = 0; j < max_j_write; j++) {
500 Scalar val = lhs_shmem[thread_x + 8 * thread_y + 64 * thread_z + 512 * j];
501 output(base_m + thread_y + 8 * thread_x, base_n + thread_z + 8 * j) = val;
502 }
503 }
504 }
505#undef res
506}
507
508template <typename Scalar, typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper>
509__global__ void
510#if defined(EIGEN_HIPCC)
511__launch_bounds__(512, 1)
512#else
513__launch_bounds__(512)
514#endif
515 EigenContractionKernel(const LhsMapper lhs, const RhsMapper rhs, const OutputMapper output, const Index m_size,
516 const Index n_size, const Index k_size) {
517 __shared__ Scalar lhs_shmem[72 * 64];
518 __shared__ Scalar rhs_shmem[72 * 64];
519
520 const Index m_block_idx = blockIdx.x;
521 const Index n_block_idx = blockIdx.y;
522
523 const Index base_m = 64 * m_block_idx;
524 const Index base_n = 64 * n_block_idx;
525
526 if (base_m + 63 < m_size && base_n + 63 < n_size) {
527 EigenContractionKernelInternal<Scalar, Index, LhsMapper, RhsMapper, OutputMapper, false>(
528 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size);
529 } else {
530 EigenContractionKernelInternal<Scalar, Index, LhsMapper, RhsMapper, OutputMapper, true>(
531 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size);
532 }
533}
534
535template <typename Scalar, typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper>
536__global__ void
537#if defined(EIGEN_HIPCC)
538__launch_bounds__(256, 1)
539#else
540__launch_bounds__(256)
541#endif
542 EigenContractionKernelNaive(const LhsMapper lhs, const RhsMapper rhs, const OutputMapper output, const Index m_size,
543 const Index n_size, const Index k_size) {
544 const Index row = static_cast<Index>(blockIdx.x) * static_cast<Index>(blockDim.x) + static_cast<Index>(threadIdx.x);
545 const Index col = static_cast<Index>(blockIdx.y) * static_cast<Index>(blockDim.y) + static_cast<Index>(threadIdx.y);
546
547 if (row >= m_size || col >= n_size) {
548 return;
549 }
550
551 internal::scalar_cast_op<int, Scalar> conv;
552 Scalar result = conv(0);
553 for (Index k = 0; k < k_size; ++k) {
554 result += lhs(row, k) * rhs(k, col);
555 }
556 output(row, col) = result;
557}
558
559template <typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper, bool CHECK_LHS_BOUNDARY,
560 bool CHECK_RHS_BOUNDARY>
561__device__ __forceinline__ void EigenFloatContractionKernelInternal16x16(const LhsMapper lhs, const RhsMapper rhs,
562 const OutputMapper output,
563 float2 lhs_shmem2[][16],
564 float2 rhs_shmem2[][8], const Index m_size,
565 const Index n_size, const Index k_size,
566 const Index base_m, const Index base_n) {
567 // prefetch registers
568 float4 lhs_pf0, rhs_pf0;
569
570 float4 results[4];
571 const Index thread_x = threadIdx.x;
572 const Index thread_y = threadIdx.y;
573 for (int i = 0; i < 4; i++) {
574 results[i].x = results[i].y = results[i].z = results[i].w = 0;
575 }
576
577#define prefetch_lhs(reg, row, col) \
578 if (!CHECK_LHS_BOUNDARY) { \
579 if (col < k_size) { \
580 reg = lhs.template loadPacket<float4, Unaligned>(row, col); \
581 } \
582 } else { \
583 if (col < k_size) { \
584 if (row + 3 < m_size) { \
585 reg = lhs.template loadPacket<float4, Unaligned>(row, col); \
586 } else if (row + 2 < m_size) { \
587 reg.x = lhs(row + 0, col); \
588 reg.y = lhs(row + 1, col); \
589 reg.z = lhs(row + 2, col); \
590 } else if (row + 1 < m_size) { \
591 reg.x = lhs(row + 0, col); \
592 reg.y = lhs(row + 1, col); \
593 } else if (row < m_size) { \
594 reg.x = lhs(row + 0, col); \
595 } \
596 } \
597 }
598
599 Index lhs_vert = base_m + thread_x * 4;
600
601 for (Index k = 0; k < k_size; k += 16) {
602 lhs_pf0 = internal::pset1<float4>(0);
603 rhs_pf0 = internal::pset1<float4>(0);
604
605 Index lhs_horiz = thread_y + k;
606 prefetch_lhs(lhs_pf0, lhs_vert, lhs_horiz)
607
608 Index rhs_vert = k + (thread_x % 4) * 4;
609 Index rhs_horiz0 = (thread_x >> 2) + thread_y * 4 + base_n;
610
611 if (!CHECK_RHS_BOUNDARY) {
612 if ((rhs_vert + 3) < k_size) {
613 // just CHECK_RHS_BOUNDARY
614 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
615 } else if (rhs_vert + 2 < k_size) {
616 // just CHECK_RHS_BOUNDARY
617 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
618 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
619 rhs_pf0.z = rhs(rhs_vert + 2, rhs_horiz0);
620 } else if (rhs_vert + 1 < k_size) {
621 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
622 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
623 } else if (rhs_vert < k_size) {
624 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
625 }
626 } else {
627 if (rhs_horiz0 < n_size) {
628 if ((rhs_vert + 3) < k_size) {
629 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
630 } else if ((rhs_vert + 2) < k_size) {
631 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
632 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
633 rhs_pf0.z = rhs(rhs_vert + 2, rhs_horiz0);
634 } else if ((rhs_vert + 1) < k_size) {
635 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
636 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
637 } else if (rhs_vert < k_size) {
638 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
639 }
640 }
641 }
642 float x1, x2;
643 // TODO: The following can be a bitwise operation.
644 if ((thread_x % 8) < 4) {
645 x1 = rhs_pf0.y;
646 x2 = rhs_pf0.w;
647 } else {
648 x1 = rhs_pf0.x;
649 x2 = rhs_pf0.z;
650 }
651#if defined(EIGEN_HIPCC)
652 x1 = __shfl_xor(x1, 4);
653 x2 = __shfl_xor(x2, 4);
654#else
655 x1 = __shfl_xor_sync(0xFFFFFFFF, x1, 4);
656 x2 = __shfl_xor_sync(0xFFFFFFFF, x2, 4);
657#endif
658 if ((thread_x % 8) < 4) {
659 rhs_pf0.y = x1;
660 rhs_pf0.w = x2;
661 } else {
662 rhs_pf0.x = x1;
663 rhs_pf0.z = x2;
664 }
665
666 // We have 64 features.
667 // Row 0 -> times (0, 4, 8, 12, 1, 5, 9, 13) for features 0, 1.
668 // Row 1 -> times (0, 4, 8, 12, 1, 5, 9, 13) for features 2, 3.
669 // ...
670 // Row 31 -> times (0, 4, 8, 12, 1, 5, 9, 13) for features 62, 63
671 // Row 32 -> times (2, 6, 10, 14, 3, 7, 11, 15) for features 0, 1
672 // ...
673 rhs_shmem2[(thread_x >> 3) + thread_y * 2][thread_x % 8] = make_float2(rhs_pf0.x, rhs_pf0.y);
674 rhs_shmem2[(thread_x >> 3) + thread_y * 2 + 32][thread_x % 8] = make_float2(rhs_pf0.z, rhs_pf0.w);
675
676 // Row 0 (time 0) -> features (0, 1), (4, 5), .. (28, 29), (32, 33), .. (60, 61)
677 // Row 1 (time 1) -> features (0, 1), (4, 5), .. (28, 29), (32, 33), .. (60, 61)
678 // ...
679 // Row 15 (time 15) -> features (0, 1), (4, 5), .. (28, 29), (32, 33), .. (60, 61)
680 // Row 16 (time 0) -> features (2, 3), (6, 7), .. (30, 31), (34, 35), .. (62, 63)
681 // ...
682
683 lhs_shmem2[thread_y][thread_x] = make_float2(lhs_pf0.x, lhs_pf0.y);
684 lhs_shmem2[thread_y + 16][thread_x] = make_float2(lhs_pf0.z, lhs_pf0.w);
685
686#define add_vals(fl1, fl2, fr1, fr2) \
687 results[0].x += fl1.x * fr1.x; \
688 results[0].y += fl1.y * fr1.x; \
689 results[0].z += fl2.x * fr1.x; \
690 results[0].w += fl2.y * fr1.x; \
691 \
692 results[1].x += fl1.x * fr1.y; \
693 results[1].y += fl1.y * fr1.y; \
694 results[1].z += fl2.x * fr1.y; \
695 results[1].w += fl2.y * fr1.y; \
696 \
697 results[2].x += fl1.x * fr2.x; \
698 results[2].y += fl1.y * fr2.x; \
699 results[2].z += fl2.x * fr2.x; \
700 results[2].w += fl2.y * fr2.x; \
701 \
702 results[3].x += fl1.x * fr2.y; \
703 results[3].y += fl1.y * fr2.y; \
704 results[3].z += fl2.x * fr2.y; \
705 results[3].w += fl2.y * fr2.y;
706
707 __syncthreads();
708
709// Do the multiplies.
710#pragma unroll
711 for (int koff = 0; koff < 16; koff++) {
712 // 32 x threads.
713 float2 fl1 = lhs_shmem2[koff][thread_x];
714 float2 fl2 = lhs_shmem2[koff + 16][thread_x];
715
716 int start_feature = thread_y * 4;
717 float2 fr1 = rhs_shmem2[(start_feature >> 1) + 32 * ((koff % 4) / 2)][koff / 4 + (koff % 2) * 4];
718 float2 fr2 = rhs_shmem2[(start_feature >> 1) + 1 + 32 * ((koff % 4) / 2)][koff / 4 + (koff % 2) * 4];
719
720 add_vals(fl1, fl2, fr1, fr2)
721 }
722 __syncthreads();
723 }
724
725#undef prefetch_lhs
726#undef add_vals
727
728 Index horiz_base = thread_y * 4 + base_n;
729 if (!CHECK_LHS_BOUNDARY && !CHECK_RHS_BOUNDARY) {
730 for (int i = 0; i < 4; i++) {
731 output(lhs_vert, horiz_base + i) = results[i].x;
732 output(lhs_vert + 1, horiz_base + i) = results[i].y;
733 output(lhs_vert + 2, horiz_base + i) = results[i].z;
734 output(lhs_vert + 3, horiz_base + i) = results[i].w;
735 }
736 } else if (!CHECK_RHS_BOUNDARY) {
737 // CHECK LHS
738 if (lhs_vert + 3 < m_size) {
739 for (int i = 0; i < 4; i++) {
740 output(lhs_vert, horiz_base + i) = results[i].x;
741 output(lhs_vert + 1, horiz_base + i) = results[i].y;
742 output(lhs_vert + 2, horiz_base + i) = results[i].z;
743 output(lhs_vert + 3, horiz_base + i) = results[i].w;
744 }
745 } else if (lhs_vert + 2 < m_size) {
746 for (int i = 0; i < 4; i++) {
747 output(lhs_vert, horiz_base + i) = results[i].x;
748 output(lhs_vert + 1, horiz_base + i) = results[i].y;
749 output(lhs_vert + 2, horiz_base + i) = results[i].z;
750 }
751 } else if (lhs_vert + 1 < m_size) {
752 for (int i = 0; i < 4; i++) {
753 output(lhs_vert, horiz_base + i) = results[i].x;
754 output(lhs_vert + 1, horiz_base + i) = results[i].y;
755 }
756 } else if (lhs_vert < m_size) {
757 for (int i = 0; i < 4; i++) {
758 output(lhs_vert, horiz_base + i) = results[i].x;
759 }
760 }
761 } else if (!CHECK_LHS_BOUNDARY) {
762 // CHECK RHS
763 for (int i = 0; i < 4; i++) {
764 if (horiz_base + i < n_size) {
765 output(lhs_vert, horiz_base + i) = results[i].x;
766 output(lhs_vert + 1, horiz_base + i) = results[i].y;
767 output(lhs_vert + 2, horiz_base + i) = results[i].z;
768 output(lhs_vert + 3, horiz_base + i) = results[i].w;
769 }
770 }
771 } else {
772 // CHECK both boundaries.
773 for (int i = 0; i < 4; i++) {
774 if (horiz_base + i < n_size) {
775 if (lhs_vert < m_size) output(lhs_vert, horiz_base + i) = results[i].x;
776 if (lhs_vert + 1 < m_size) output(lhs_vert + 1, horiz_base + i) = results[i].y;
777 if (lhs_vert + 2 < m_size) output(lhs_vert + 2, horiz_base + i) = results[i].z;
778 if (lhs_vert + 3 < m_size) output(lhs_vert + 3, horiz_base + i) = results[i].w;
779 }
780 }
781 }
782}
783
784template <typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper, bool CHECK_LHS_BOUNDARY,
785 bool CHECK_RHS_BOUNDARY>
786__device__ __forceinline__ void EigenFloatContractionKernelInternal(const LhsMapper lhs, const RhsMapper rhs,
787 const OutputMapper output, float2 lhs_shmem2[][32],
788 float2 rhs_shmem2[][8], const Index m_size,
789 const Index n_size, const Index k_size,
790 const Index base_m, const Index base_n) {
791 // prefetch registers
792 float4 lhs_pf0, lhs_pf1, lhs_pf2, lhs_pf3;
793 float4 rhs_pf0, rhs_pf1;
794
795 float4 results[8];
796 const Index thread_x = threadIdx.x;
797 const Index thread_y = threadIdx.y;
798
799 for (int i = 0; i < 8; i++) {
800 results[i].x = results[i].y = results[i].z = results[i].w = 0;
801 }
802
803 Index lhs_vert = base_m + thread_x * 4 + (thread_y % 4) * 32;
804
805 for (Index k = 0; k < k_size; k += 32) {
806 lhs_pf0 = internal::pset1<float4>(0);
807 lhs_pf1 = internal::pset1<float4>(0);
808 lhs_pf2 = internal::pset1<float4>(0);
809 lhs_pf3 = internal::pset1<float4>(0);
810
811 rhs_pf0 = internal::pset1<float4>(0);
812 rhs_pf1 = internal::pset1<float4>(0);
813
814 if (!CHECK_LHS_BOUNDARY) {
815 if ((thread_y / 4 + k + 24) < k_size) {
816 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
817 lhs_pf1 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 8));
818 lhs_pf2 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 16));
819 lhs_pf3 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 24));
820 } else if ((thread_y / 4 + k + 16) < k_size) {
821 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
822 lhs_pf1 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 8));
823 lhs_pf2 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 16));
824 } else if ((thread_y / 4 + k + 8) < k_size) {
825 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
826 lhs_pf1 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 8));
827 } else if ((thread_y / 4 + k) < k_size) {
828 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
829 }
830 } else {
831 // just CHECK_LHS_BOUNDARY
832 if (lhs_vert + 3 < m_size) {
833 if ((thread_y / 4 + k + 24) < k_size) {
834 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
835 lhs_pf1 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 8));
836 lhs_pf2 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 16));
837 lhs_pf3 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 24));
838 } else if ((thread_y / 4 + k + 16) < k_size) {
839 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
840 lhs_pf1 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 8));
841 lhs_pf2 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 16));
842 } else if ((thread_y / 4 + k + 8) < k_size) {
843 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
844 lhs_pf1 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k + 8));
845 } else if ((thread_y / 4 + k) < k_size) {
846 lhs_pf0 = lhs.template loadPacket<float4, Unaligned>(lhs_vert, (thread_y / 4 + k));
847 }
848 } else if (lhs_vert + 2 < m_size) {
849 if ((thread_y / 4 + k + 24) < k_size) {
850 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
851 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
852 lhs_pf0.z = lhs(lhs_vert + 2, (thread_y / 4 + k));
853 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
854 lhs_pf1.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 8));
855 lhs_pf1.z = lhs(lhs_vert + 2, (thread_y / 4 + k + 8));
856 lhs_pf2.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 16));
857 lhs_pf2.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 16));
858 lhs_pf2.z = lhs(lhs_vert + 2, (thread_y / 4 + k + 16));
859 lhs_pf3.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 24));
860 lhs_pf3.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 24));
861 lhs_pf3.z = lhs(lhs_vert + 2, (thread_y / 4 + k + 24));
862 } else if ((thread_y / 4 + k + 16) < k_size) {
863 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
864 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
865 lhs_pf0.z = lhs(lhs_vert + 2, (thread_y / 4 + k));
866 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
867 lhs_pf1.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 8));
868 lhs_pf1.z = lhs(lhs_vert + 2, (thread_y / 4 + k + 8));
869 lhs_pf2.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 16));
870 lhs_pf2.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 16));
871 lhs_pf2.z = lhs(lhs_vert + 2, (thread_y / 4 + k + 16));
872 } else if ((thread_y / 4 + k + 8) < k_size) {
873 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
874 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
875 lhs_pf0.z = lhs(lhs_vert + 2, (thread_y / 4 + k));
876 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
877 lhs_pf1.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 8));
878 lhs_pf1.z = lhs(lhs_vert + 2, (thread_y / 4 + k + 8));
879 } else if ((thread_y / 4 + k) < k_size) {
880 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
881 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
882 lhs_pf0.z = lhs(lhs_vert + 2, (thread_y / 4 + k));
883 }
884 } else if (lhs_vert + 1 < m_size) {
885 if ((thread_y / 4 + k + 24) < k_size) {
886 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
887 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
888 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
889 lhs_pf1.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 8));
890 lhs_pf2.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 16));
891 lhs_pf2.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 16));
892 lhs_pf3.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 24));
893 lhs_pf3.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 24));
894 } else if ((thread_y / 4 + k + 16) < k_size) {
895 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
896 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
897 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
898 lhs_pf1.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 8));
899 lhs_pf2.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 16));
900 lhs_pf2.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 16));
901 } else if ((thread_y / 4 + k + 8) < k_size) {
902 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
903 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
904 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
905 lhs_pf1.y = lhs(lhs_vert + 1, (thread_y / 4 + k + 8));
906 } else if ((thread_y / 4 + k) < k_size) {
907 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
908 lhs_pf0.y = lhs(lhs_vert + 1, (thread_y / 4 + k));
909 }
910 } else if (lhs_vert < m_size) {
911 if ((thread_y / 4 + k + 24) < k_size) {
912 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
913 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
914 lhs_pf2.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 16));
915 lhs_pf3.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 24));
916 } else if ((thread_y / 4 + k + 16) < k_size) {
917 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
918 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
919 lhs_pf2.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 16));
920 } else if ((thread_y / 4 + k + 8) < k_size) {
921 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
922 lhs_pf1.x = lhs(lhs_vert + 0, (thread_y / 4 + k + 8));
923 } else if ((thread_y / 4 + k) < k_size) {
924 lhs_pf0.x = lhs(lhs_vert + 0, (thread_y / 4 + k));
925 }
926 }
927 }
928 __syncthreads();
929 Index rhs_vert = k + thread_x * 4;
930 Index rhs_horiz0 = thread_y * 2 + base_n;
931 Index rhs_horiz1 = thread_y * 2 + 1 + base_n;
932 if (!CHECK_RHS_BOUNDARY) {
933 if ((rhs_vert + 3) < k_size) {
934 // just CHECK_RHS_BOUNDARY
935 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
936 rhs_pf1 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz1);
937 } else if (rhs_vert + 2 < k_size) {
938 // just CHECK_RHS_BOUNDARY
939 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
940 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
941 rhs_pf0.z = rhs(rhs_vert + 2, rhs_horiz0);
942 rhs_pf1.x = rhs(rhs_vert, rhs_horiz1);
943 rhs_pf1.y = rhs(rhs_vert + 1, rhs_horiz1);
944 rhs_pf1.z = rhs(rhs_vert + 2, rhs_horiz1);
945 } else if (rhs_vert + 1 < k_size) {
946 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
947 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
948 rhs_pf1.x = rhs(rhs_vert, rhs_horiz1);
949 rhs_pf1.y = rhs(rhs_vert + 1, rhs_horiz1);
950 } else if (rhs_vert < k_size) {
951 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
952 rhs_pf1.x = rhs(rhs_vert, rhs_horiz1);
953 }
954 } else {
955 if (rhs_horiz1 < n_size) {
956 if ((rhs_vert + 3) < k_size) {
957 // just CHECK_RHS_BOUNDARY
958 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
959 rhs_pf1 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz1);
960 } else if (rhs_vert + 2 < k_size) {
961 // just CHECK_RHS_BOUNDARY
962 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
963 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
964 rhs_pf0.z = rhs(rhs_vert + 2, rhs_horiz0);
965 rhs_pf1.x = rhs(rhs_vert, rhs_horiz1);
966 rhs_pf1.y = rhs(rhs_vert + 1, rhs_horiz1);
967 rhs_pf1.z = rhs(rhs_vert + 2, rhs_horiz1);
968 } else if (k + thread_x * 4 + 1 < k_size) {
969 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
970 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
971 rhs_pf1.x = rhs(rhs_vert, rhs_horiz1);
972 rhs_pf1.y = rhs(rhs_vert + 1, rhs_horiz1);
973 } else if (k + thread_x * 4 < k_size) {
974 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
975 rhs_pf1.x = rhs(rhs_vert, rhs_horiz1);
976 }
977 } else if (rhs_horiz0 < n_size) {
978 if ((rhs_vert + 3) < k_size) {
979 // just CHECK_RHS_BOUNDARY
980 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
981 } else if ((rhs_vert + 2) < k_size) {
982 // just CHECK_RHS_BOUNDARY
983 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
984 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
985 rhs_pf0.z = rhs(rhs_vert + 2, rhs_horiz0);
986 } else if ((rhs_vert + 1) < k_size) {
987 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
988 rhs_pf0.y = rhs(rhs_vert + 1, rhs_horiz0);
989 } else if (rhs_vert < k_size) {
990 rhs_pf0.x = rhs(rhs_vert, rhs_horiz0);
991 }
992 }
993 }
994 __syncthreads();
995 // Loaded. Do computation
996 // Row 0 -> times (0, 4, 8, .. 28) for features 0, 1.
997 // Row 1 -> times (0, 4, 8, .. 28) for features 2, 3.
998 // ..
999 // Row 31 -> times (0, 4, 8, .. 28) for features 62, 63
1000 rhs_shmem2[thread_y][thread_x] = make_float2(rhs_pf0.x, rhs_pf1.x);
1001 // Row 32 -> times (1, 5, 9, .. 29) for features 0, 1.
1002 // Row 33 -> times (1, 5, 9, .. 29) for features 2, 3.
1003 // ..
1004 rhs_shmem2[thread_y + 32][thread_x] = make_float2(rhs_pf0.y, rhs_pf1.y);
1005 // Row 64 -> times (2, 6, 10, .. 30) for features 0, 1.
1006 // Row 65 -> times (2, 6, 10, .. 30) for features 2, 3.
1007 rhs_shmem2[thread_y + 64][thread_x] = make_float2(rhs_pf0.z, rhs_pf1.z);
1008 // Row 96 -> times (3, 7, 11, .. 31) for features 0, 1.
1009 // Row 97 -> times (3, 7, 11, .. 31) for features 2, 3.
1010 rhs_shmem2[thread_y + 96][thread_x] = make_float2(rhs_pf0.w, rhs_pf1.w);
1011
1012 // LHS.
1013 // Row 0 (time 0) -> features (0, 1), (4, 5), .. (28, 29), (32, 33), .. (60, 61) .. (124, 125)
1014 // Row 1 (time 1) -> features (0, 1), (4, 5), .. (28, 29), (32, 33), .. (60, 61) .. (124, 125)
1015 // ...
1016 // Row 8 (time 0) -> features (2, 3), (6, 7), .. (30, 31), (34, 35), .. (62, 63) .. (126, 127)
1017 // Row 15 (time 7) -> features (2, 3), (6, 7), .. (30, 31), (34, 35), .. (62, 63) .. (126, 127)
1018
1019#define add_vals(a_feat1, a_feat2, f1, f2, f3, f4) \
1020 results[0].x += a_feat1.x * f1.x; \
1021 results[1].x += a_feat1.x * f1.y; \
1022 results[2].x += a_feat1.x * f2.x; \
1023 results[3].x += a_feat1.x * f2.y; \
1024 results[4].x += a_feat1.x * f3.x; \
1025 results[5].x += a_feat1.x * f3.y; \
1026 results[6].x += a_feat1.x * f4.x; \
1027 results[7].x += a_feat1.x * f4.y; \
1028 \
1029 results[0].y += a_feat1.y * f1.x; \
1030 results[1].y += a_feat1.y * f1.y; \
1031 results[2].y += a_feat1.y * f2.x; \
1032 results[3].y += a_feat1.y * f2.y; \
1033 results[4].y += a_feat1.y * f3.x; \
1034 results[5].y += a_feat1.y * f3.y; \
1035 results[6].y += a_feat1.y * f4.x; \
1036 results[7].y += a_feat1.y * f4.y; \
1037 \
1038 results[0].z += a_feat2.x * f1.x; \
1039 results[1].z += a_feat2.x * f1.y; \
1040 results[2].z += a_feat2.x * f2.x; \
1041 results[3].z += a_feat2.x * f2.y; \
1042 results[4].z += a_feat2.x * f3.x; \
1043 results[5].z += a_feat2.x * f3.y; \
1044 results[6].z += a_feat2.x * f4.x; \
1045 results[7].z += a_feat2.x * f4.y; \
1046 \
1047 results[0].w += a_feat2.y * f1.x; \
1048 results[1].w += a_feat2.y * f1.y; \
1049 results[2].w += a_feat2.y * f2.x; \
1050 results[3].w += a_feat2.y * f2.y; \
1051 results[4].w += a_feat2.y * f3.x; \
1052 results[5].w += a_feat2.y * f3.y; \
1053 results[6].w += a_feat2.y * f4.x; \
1054 results[7].w += a_feat2.y * f4.y;
1055
1056 lhs_shmem2[thread_y / 4][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf0.x, lhs_pf0.y);
1057 lhs_shmem2[thread_y / 4 + 8][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf1.x, lhs_pf1.y);
1058 lhs_shmem2[thread_y / 4 + 16][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf2.x, lhs_pf2.y);
1059 lhs_shmem2[thread_y / 4 + 24][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf3.x, lhs_pf3.y);
1060
1061 lhs_shmem2[thread_y / 4 + 32][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf0.z, lhs_pf0.w);
1062 lhs_shmem2[thread_y / 4 + 40][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf1.z, lhs_pf1.w);
1063 lhs_shmem2[thread_y / 4 + 48][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf2.z, lhs_pf2.w);
1064 lhs_shmem2[thread_y / 4 + 56][thread_x + (thread_y % 4) * 8] = make_float2(lhs_pf3.z, lhs_pf3.w);
1065
1066 __syncthreads();
1067
1068// Do the multiplies.
1069#pragma unroll
1070 for (int koff = 0; koff < 32; koff++) {
1071 float2 a3 = lhs_shmem2[koff][thread_x + (thread_y % 4) * 8];
1072 float2 a4 = lhs_shmem2[koff + 32][thread_x + (thread_y % 4) * 8];
1073
1074 // first feature is at (thread_y/4) * 8 last is at start + 8.
1075 int start_feature = (thread_y / 4) * 8;
1076
1077 float2 br1 = rhs_shmem2[start_feature / 2 + (koff % 4) * 32][koff / 4];
1078 float2 br2 = rhs_shmem2[start_feature / 2 + 1 + (koff % 4) * 32][koff / 4];
1079 float2 br3 = rhs_shmem2[start_feature / 2 + 2 + (koff % 4) * 32][koff / 4];
1080 float2 br4 = rhs_shmem2[start_feature / 2 + 3 + (koff % 4) * 32][koff / 4];
1081
1082 add_vals(a3, a4, br1, br2, br3, br4)
1083 }
1084 __syncthreads();
1085 } // end loop over k
1086
1087#undef add_vals
1088
1089 __syncthreads();
1090 Index horiz_base = (thread_y / 4) * 8 + base_n;
1091 if (!CHECK_LHS_BOUNDARY && !CHECK_RHS_BOUNDARY) {
1092 for (int i = 0; i < 8; i++) {
1093 output(lhs_vert, horiz_base + i) = results[i].x;
1094 output(lhs_vert + 1, horiz_base + i) = results[i].y;
1095 output(lhs_vert + 2, horiz_base + i) = results[i].z;
1096 output(lhs_vert + 3, horiz_base + i) = results[i].w;
1097 }
1098 } else if (!CHECK_RHS_BOUNDARY) {
1099 if (lhs_vert + 3 < m_size) {
1100 for (int i = 0; i < 8; i++) {
1101 output(lhs_vert, horiz_base + i) = results[i].x;
1102 output(lhs_vert + 1, horiz_base + i) = results[i].y;
1103 output(lhs_vert + 2, horiz_base + i) = results[i].z;
1104 output(lhs_vert + 3, horiz_base + i) = results[i].w;
1105 }
1106 } else if (lhs_vert + 2 < m_size) {
1107 for (int i = 0; i < 8; i++) {
1108 output(lhs_vert, horiz_base + i) = results[i].x;
1109 output(lhs_vert + 1, horiz_base + i) = results[i].y;
1110 output(lhs_vert + 2, horiz_base + i) = results[i].z;
1111 }
1112 } else if (lhs_vert + 1 < m_size) {
1113 for (int i = 0; i < 8; i++) {
1114 output(lhs_vert, horiz_base + i) = results[i].x;
1115 output(lhs_vert + 1, horiz_base + i) = results[i].y;
1116 }
1117 } else if (lhs_vert < m_size) {
1118 for (int i = 0; i < 8; i++) {
1119 output(lhs_vert, horiz_base + i) = results[i].x;
1120 }
1121 }
1122 } else if (!CHECK_LHS_BOUNDARY) {
1123 // CHECK BOUNDARY_B
1124 for (int i = 0; i < 8; i++) {
1125 if (horiz_base + i < n_size) {
1126 output(lhs_vert, horiz_base + i) = results[i].x;
1127 output(lhs_vert + 1, horiz_base + i) = results[i].y;
1128 output(lhs_vert + 2, horiz_base + i) = results[i].z;
1129 output(lhs_vert + 3, horiz_base + i) = results[i].w;
1130 }
1131 }
1132 } else {
1133 // CHECK both boundaries.
1134 for (int i = 0; i < 8; i++) {
1135 if (horiz_base + i < n_size) {
1136 if (lhs_vert < m_size) output(lhs_vert, horiz_base + i) = results[i].x;
1137 if (lhs_vert + 1 < m_size) output(lhs_vert + 1, horiz_base + i) = results[i].y;
1138 if (lhs_vert + 2 < m_size) output(lhs_vert + 2, horiz_base + i) = results[i].z;
1139 if (lhs_vert + 3 < m_size) output(lhs_vert + 3, horiz_base + i) = results[i].w;
1140 }
1141 }
1142 }
1143}
1144
1145template <typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper>
1146__global__ void
1147#if defined(EIGEN_HIPCC)
1148__launch_bounds__(256, 1)
1149#else
1150__launch_bounds__(256)
1151#endif
1152 EigenFloatContractionKernel(const LhsMapper lhs, const RhsMapper rhs, const OutputMapper output, const Index m_size,
1153 const Index n_size, const Index k_size) {
1154 __shared__ float2 lhs_shmem[64 * 32];
1155 __shared__ float2 rhs_shmem[128 * 8];
1156
1157 typedef float2 LHS_MEM[64][32];
1158 typedef float2 RHS_MEM[128][8];
1159
1160 const Index m_block_idx = blockIdx.x;
1161 const Index n_block_idx = blockIdx.y;
1162
1163 const Index base_m = 128 * m_block_idx;
1164 const Index base_n = 64 * n_block_idx;
1165
1166 bool check_rhs = (base_n + 63) >= n_size;
1167 bool check_lhs128 = (base_m + 127) >= m_size;
1168
1169 if (!check_rhs) {
1170 if (!check_lhs128) {
1171 // >= 128 rows left
1172 EigenFloatContractionKernelInternal<Index, LhsMapper, RhsMapper, OutputMapper, false, false>(
1173 lhs, rhs, output, *((LHS_MEM*)lhs_shmem), *((RHS_MEM*)rhs_shmem), m_size, n_size, k_size, base_m, base_n);
1174 } else {
1175 EigenFloatContractionKernelInternal<Index, LhsMapper, RhsMapper, OutputMapper, true, false>(
1176 lhs, rhs, output, *((LHS_MEM*)lhs_shmem), *((RHS_MEM*)rhs_shmem), m_size, n_size, k_size, base_m, base_n);
1177 }
1178 } else {
1179 if (!check_lhs128) {
1180 // >= 128 rows left
1181 EigenFloatContractionKernelInternal<Index, LhsMapper, RhsMapper, OutputMapper, false, true>(
1182 lhs, rhs, output, *((LHS_MEM*)lhs_shmem), *((RHS_MEM*)rhs_shmem), m_size, n_size, k_size, base_m, base_n);
1183 } else {
1184 EigenFloatContractionKernelInternal<Index, LhsMapper, RhsMapper, OutputMapper, true, true>(
1185 lhs, rhs, output, *((LHS_MEM*)lhs_shmem), *((RHS_MEM*)rhs_shmem), m_size, n_size, k_size, base_m, base_n);
1186 }
1187 }
1188}
1189
1190template <typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper>
1191__global__ void
1192#if defined(EIGEN_HIPCC)
1193__launch_bounds__(256, 1)
1194#else
1195__launch_bounds__(256)
1196#endif
1197 EigenFloatContractionKernel16x16(const LhsMapper lhs, const RhsMapper rhs, const OutputMapper output,
1198 const Index m_size, const Index n_size, const Index k_size) {
1199 __shared__ float2 lhs_shmem[32][16];
1200 __shared__ float2 rhs_shmem[64][8];
1201
1202 const Index m_block_idx = blockIdx.x;
1203 const Index n_block_idx = blockIdx.y;
1204
1205 const Index base_m = 64 * m_block_idx;
1206 const Index base_n = 64 * n_block_idx;
1207
1208 if (base_m + 63 < m_size) {
1209 if (base_n + 63 < n_size) {
1210 EigenFloatContractionKernelInternal16x16<Index, LhsMapper, RhsMapper, OutputMapper, false, false>(
1211 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size, base_m, base_n);
1212 } else {
1213 EigenFloatContractionKernelInternal16x16<Index, LhsMapper, RhsMapper, OutputMapper, false, true>(
1214 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size, base_m, base_n);
1215 }
1216 } else {
1217 if (base_n + 63 < n_size) {
1218 EigenFloatContractionKernelInternal16x16<Index, LhsMapper, RhsMapper, OutputMapper, true, false>(
1219 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size, base_m, base_n);
1220 } else {
1221 EigenFloatContractionKernelInternal16x16<Index, LhsMapper, RhsMapper, OutputMapper, true, true>(
1222 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size, base_m, base_n);
1223 }
1224 }
1225}
1226
1227template <typename Indices, typename LeftArgType, typename RightArgType, typename OutputKernelType>
1228struct TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, GpuDevice>
1229 : public TensorContractionEvaluatorBase<TensorEvaluator<
1230 const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, GpuDevice> > {
1231 typedef GpuDevice Device;
1232
1233 typedef TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device> Self;
1234 typedef TensorContractionEvaluatorBase<Self> Base;
1235
1236 typedef TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType> XprType;
1237 typedef std::remove_const_t<typename XprType::Scalar> Scalar;
1238 typedef typename XprType::Index Index;
1239 typedef typename XprType::CoeffReturnType CoeffReturnType;
1240 typedef typename PacketType<CoeffReturnType, GpuDevice>::type PacketReturnType;
1241
1242 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
1243
1244 // Most of the code is assuming that both input tensors are ColMajor. If the
1245 // inputs are RowMajor, we will "cheat" by swapping the LHS and RHS:
1246 // If we want to compute A * B = C, where A is LHS and B is RHS, the code
1247 // will pretend B is LHS and A is RHS.
1248 typedef std::conditional_t<Layout == static_cast<int>(ColMajor), LeftArgType, RightArgType> EvalLeftArgType;
1249 typedef std::conditional_t<Layout == static_cast<int>(ColMajor), RightArgType, LeftArgType> EvalRightArgType;
1250
1251 static constexpr int LDims =
1252 internal::array_size<typename TensorEvaluator<EvalLeftArgType, Device>::Dimensions>::value;
1253 static constexpr int RDims =
1254 internal::array_size<typename TensorEvaluator<EvalRightArgType, Device>::Dimensions>::value;
1255 static constexpr int ContractDims = internal::array_size<Indices>::value;
1256
1257 typedef array<Index, LDims> left_dim_mapper_t;
1258 typedef array<Index, RDims> right_dim_mapper_t;
1259
1260 typedef array<Index, ContractDims> contract_t;
1261 typedef array<Index, LDims - ContractDims> left_nocontract_t;
1262 typedef array<Index, RDims - ContractDims> right_nocontract_t;
1263
1264 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
1265
1266 typedef DSizes<Index, NumDims> Dimensions;
1267
1268 // typedefs needed in evalTo
1269 typedef std::remove_const_t<typename EvalLeftArgType::Scalar> LhsScalar;
1270 typedef std::remove_const_t<typename EvalRightArgType::Scalar> RhsScalar;
1271
1272 typedef TensorEvaluator<EvalLeftArgType, Device> LeftEvaluator;
1273 typedef TensorEvaluator<EvalRightArgType, Device> RightEvaluator;
1274
1275 typedef typename LeftEvaluator::Dimensions LeftDimensions;
1276 typedef typename RightEvaluator::Dimensions RightDimensions;
1277
1278 TensorEvaluator(const XprType& op, const Device& device) : Base(op, device) {
1279 EIGEN_STATIC_ASSERT((std::is_same<OutputKernelType, const NoOpOutputKernel>::value),
1280 GPU_TENSOR_CONTRACTION_DOES_NOT_SUPPORT_OUTPUT_KERNELS);
1281 }
1282
1283 // We need to redefine this method to make nvcc happy
1284 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(Scalar* data) {
1285 this->m_leftImpl.evalSubExprsIfNeeded(nullptr);
1286 this->m_rightImpl.evalSubExprsIfNeeded(nullptr);
1287 if (data) {
1288 evalTo(data);
1289 return false;
1290 } else {
1291 this->m_result = static_cast<Scalar*>(this->m_device.allocate(this->dimensions().TotalSize() * sizeof(Scalar)));
1292 evalTo(this->m_result);
1293 return true;
1294 }
1295 }
1296
1297 void evalTo(Scalar* buffer) const {
1298 if (this->m_lhs_inner_dim_contiguous) {
1299 if (this->m_rhs_inner_dim_contiguous) {
1300 if (this->m_rhs_inner_dim_reordered) {
1301 evalTyped<true, true, true, Unaligned>(buffer);
1302 } else {
1303 evalTyped<true, true, false, Unaligned>(buffer);
1304 }
1305 } else {
1306 if (this->m_rhs_inner_dim_reordered) {
1307 evalTyped<true, false, true, Unaligned>(buffer);
1308 } else {
1309 evalTyped<true, false, false, Unaligned>(buffer);
1310 }
1311 }
1312 } else {
1313 if (this->m_rhs_inner_dim_contiguous) {
1314 if (this->m_rhs_inner_dim_reordered) {
1315 evalTyped<false, true, true, Unaligned>(buffer);
1316 } else {
1317 evalTyped<false, true, false, Unaligned>(buffer);
1318 }
1319 } else {
1320 if (this->m_rhs_inner_dim_reordered) {
1321 evalTyped<false, false, true, Unaligned>(buffer);
1322 } else {
1323 evalTyped<false, false, false, Unaligned>(buffer);
1324 }
1325 }
1326 }
1327 }
1328
1329 template <int Alignment>
1330 void evalProduct(Scalar* buffer) const {
1331 evalTo(buffer);
1332 }
1333
1334 template <typename LhsScalar, typename RhsScalar, typename Index, typename LhsMapper, typename RhsMapper,
1335 typename OutputMapper, bool UseNaiveKernel>
1336 struct LaunchKernelsImpl;
1337
1338 template <typename LhsScalar, typename RhsScalar, typename Index, typename LhsMapper, typename RhsMapper,
1339 typename OutputMapper>
1340 struct LaunchKernelsImpl<LhsScalar, RhsScalar, Index, LhsMapper, RhsMapper, OutputMapper, false> {
1341 static void Run(const LhsMapper& lhs, const RhsMapper& rhs, const OutputMapper& output, Index m, Index n, Index k,
1342 const GpuDevice& device) {
1343 const Index m_blocks = (m + 63) / 64;
1344 const Index n_blocks = (n + 63) / 64;
1345 const dim3 num_blocks(m_blocks, n_blocks, 1);
1346 const dim3 block_size(8, 8, 8);
1347 LAUNCH_GPU_KERNEL((EigenContractionKernel<Scalar, Index, LhsMapper, RhsMapper, OutputMapper>), num_blocks,
1348 block_size, 0, device, lhs, rhs, output, m, n, k);
1349 }
1350 };
1351
1352 template <typename LhsScalar, typename RhsScalar, typename Index, typename LhsMapper, typename RhsMapper,
1353 typename OutputMapper>
1354 struct LaunchKernelsImpl<LhsScalar, RhsScalar, Index, LhsMapper, RhsMapper, OutputMapper, true> {
1355 static void Run(const LhsMapper& lhs, const RhsMapper& rhs, const OutputMapper& output, Index m, Index n, Index k,
1356 const GpuDevice& device) {
1357 const dim3 block_size(16, 16, 1);
1358 const dim3 num_blocks((m + 15) / 16, (n + 15) / 16, 1);
1359 LAUNCH_GPU_KERNEL((EigenContractionKernelNaive<Scalar, Index, LhsMapper, RhsMapper, OutputMapper>), num_blocks,
1360 block_size, 0, device, lhs, rhs, output, m, n, k);
1361 }
1362 };
1363
1364 template <typename LhsScalar, typename RhsScalar, typename Index, typename LhsMapper, typename RhsMapper,
1365 typename OutputMapper>
1366 // The optimized generic kernel reserves two 72x64 shared-memory tiles. With 8-byte scalars that exceeds
1367 // the 48KB static shared-memory limit of common CUDA targets, so use a slower no-shared-memory fallback.
1368 struct LaunchKernels
1369 : LaunchKernelsImpl<LhsScalar, RhsScalar, Index, LhsMapper, RhsMapper, OutputMapper, (sizeof(Scalar) > 4)> {};
1370
1371 template <typename Index, typename LhsMapper, typename RhsMapper, typename OutputMapper>
1372 struct LaunchKernels<float, float, Index, LhsMapper, RhsMapper, OutputMapper> {
1373 static void Run(const LhsMapper& lhs, const RhsMapper& rhs, const OutputMapper& output, Index m, Index n, Index k,
1374 const GpuDevice& device) {
1375 if (m < 768 || n < 768) {
1376 const Index m_blocks = (m + 63) / 64;
1377 const Index n_blocks = (n + 63) / 64;
1378 const dim3 num_blocks(m_blocks, n_blocks, 1);
1379 const dim3 block_size(16, 16, 1);
1380 LAUNCH_GPU_KERNEL((EigenFloatContractionKernel16x16<Index, LhsMapper, RhsMapper, OutputMapper>), num_blocks,
1381 block_size, 0, device, lhs, rhs, output, m, n, k);
1382 } else {
1383 const Index m_blocks = (m + 127) / 128;
1384 const Index n_blocks = (n + 63) / 64;
1385 const dim3 num_blocks(m_blocks, n_blocks, 1);
1386 const dim3 block_size(8, 32, 1);
1387 LAUNCH_GPU_KERNEL((EigenFloatContractionKernel<Index, LhsMapper, RhsMapper, OutputMapper>), num_blocks,
1388 block_size, 0, device, lhs, rhs, output, m, n, k);
1389 }
1390 }
1391 };
1392
1393 template <bool lhs_inner_dim_contiguous, bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered, int Alignment>
1394 void evalTyped(Scalar* buffer) const {
1395 // columns in left side, rows in right side
1396 const Index k = this->m_k_size;
1397 // rows in left side
1398 const Index m = this->m_i_size;
1399
1400 // columns in right side
1401 const Index n = this->m_j_size;
1402
1403 if (m == 0 || n == 0) return;
1404
1405 // zero out the result buffer (which must be of size at least m * n * sizeof(Scalar))
1406 this->m_device.fill(buffer, buffer + m * n, Scalar(0));
1407
1408 if (k == 0) return;
1409
1410 typedef internal::TensorContractionInputMapper<LhsScalar, Index, internal::Lhs, LeftEvaluator, left_nocontract_t,
1411 contract_t, 4, lhs_inner_dim_contiguous, false, Unaligned>
1412 LhsMapper;
1413
1414 typedef internal::TensorContractionInputMapper<RhsScalar, Index, internal::Rhs, RightEvaluator, right_nocontract_t,
1415 contract_t, 4, rhs_inner_dim_contiguous, rhs_inner_dim_reordered,
1416 Unaligned>
1417 RhsMapper;
1418
1419 typedef internal::blas_data_mapper<Scalar, Index, ColMajor> OutputMapper;
1420
1421 // initialize data mappers
1422 LhsMapper lhs(this->m_leftImpl, this->m_left_nocontract_strides, this->m_i_strides,
1423 this->m_left_contracting_strides, this->m_k_strides);
1424
1425 RhsMapper rhs(this->m_rightImpl, this->m_right_nocontract_strides, this->m_j_strides,
1426 this->m_right_contracting_strides, this->m_k_strides);
1427
1428 OutputMapper output(buffer, m);
1429 LaunchKernels<LhsScalar, RhsScalar, Index, LhsMapper, RhsMapper, OutputMapper>::Run(lhs, rhs, output, m, n, k,
1430 this->m_device);
1431 }
1432};
1433
1434} // end namespace Eigen
1435
1436#endif // EIGEN_USE_GPU and EIGEN_GPUCC
1437#endif // EIGEN_TENSOR_TENSOR_CONTRACTION_GPU_H
Definition TensorContraction.h:335
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47