13#ifndef EIGEN_TENSOR_TENSOR_CONTRACTION_GPU_H
14#define EIGEN_TENSOR_TENSOR_CONTRACTION_GPU_H
16#if defined(EIGEN_USE_GPU) && defined(EIGEN_GPUCC)
19#include "./InternalHeaderCheck.h"
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;
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;
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;
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;
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;
103 const Index load_idx_vert = thread_x + 8 * thread_y;
104 const Index lhs_vert = base_m + load_idx_vert;
106#define prefetchIntoRegisters(base_k) \
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; \
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); \
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; \
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); \
242#define writeRegToShmem() \
243 lhs_shmem[lhs_store_idx_0] = lhs_pf0; \
244 rhs_shmem[rhs_store_idx_0] = rhs_pf0; \
246 lhs_shmem[lhs_store_idx_1] = lhs_pf1; \
247 rhs_shmem[rhs_store_idx_1] = rhs_pf1; \
249 lhs_shmem[lhs_store_idx_2] = lhs_pf2; \
250 rhs_shmem[rhs_store_idx_2] = rhs_pf2; \
252 lhs_shmem[lhs_store_idx_3] = lhs_pf3; \
253 rhs_shmem[rhs_store_idx_3] = rhs_pf3; \
255 lhs_shmem[lhs_store_idx_4] = lhs_pf4; \
256 rhs_shmem[rhs_store_idx_4] = rhs_pf4; \
258 lhs_shmem[lhs_store_idx_5] = lhs_pf5; \
259 rhs_shmem[rhs_store_idx_5] = rhs_pf5; \
261 lhs_shmem[lhs_store_idx_6] = lhs_pf6; \
262 rhs_shmem[rhs_store_idx_6] = rhs_pf6; \
264 lhs_shmem[lhs_store_idx_7] = lhs_pf7; \
265 rhs_shmem[rhs_store_idx_7] = rhs_pf7;
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);
279 internal::scalar_cast_op<int, Scalar> conv;
290 for (Index base_k = 0; base_k < k_size; base_k += 64) {
295 prefetchIntoRegisters(base_k);
298#undef prefetchIntoRegisters
299#undef writeRegToShmem
307#define lcol(i) _lcol##i
317#define rrow(j) _rrow##j
328 const Scalar* lhs_block = &lhs_shmem[thread_x + 9 * thread_y];
329 const Scalar* rhs_block = &rhs_shmem[thread_x + 8 * thread_z];
331#define lhs_element(i, j) lhs_block[72 * ((i) + 8 * (j))]
332#define rhs_element(i, j) rhs_block[72 * ((i) + 8 * (j))]
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);
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);
362#define computePass(i) \
398#if defined(EIGEN_HIPCC)
399#define shuffleInc(i, j, mask) res(i, j) += __shfl_xor(res(i, j), mask)
401#define shuffleInc(i, j, mask) res(i, j) += __shfl_xor_sync(0xFFFFFFFF, res(i, j), mask)
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);
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); \
449#define writeResultShmem(i, j) lhs_shmem[i + 8 * thread_y + 64 * thread_z + 512 * j] = res(i, j);
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);
471#undef writeResultShmem
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);
477 if (thread_x < max_i_write) {
478 if (max_j_write == 8) {
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];
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;
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;
508template <
typename Scalar,
typename Index,
typename LhsMapper,
typename RhsMapper,
typename OutputMapper>
510#if defined(EIGEN_HIPCC)
511__launch_bounds__(512, 1)
513__launch_bounds__(512)
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];
520 const Index m_block_idx = blockIdx.x;
521 const Index n_block_idx = blockIdx.y;
523 const Index base_m = 64 * m_block_idx;
524 const Index base_n = 64 * n_block_idx;
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);
530 EigenContractionKernelInternal<Scalar, Index, LhsMapper, RhsMapper, OutputMapper, true>(
531 lhs, rhs, output, lhs_shmem, rhs_shmem, m_size, n_size, k_size);
535template <
typename Scalar,
typename Index,
typename LhsMapper,
typename RhsMapper,
typename OutputMapper>
537#if defined(EIGEN_HIPCC)
538__launch_bounds__(256, 1)
540__launch_bounds__(256)
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);
547 if (row >= m_size || col >= n_size) {
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);
556 output(row, col) = result;
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) {
568 float4 lhs_pf0, rhs_pf0;
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;
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); \
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); \
599 Index lhs_vert = base_m + thread_x * 4;
601 for (Index k = 0; k < k_size; k += 16) {
602 lhs_pf0 = internal::pset1<float4>(0);
603 rhs_pf0 = internal::pset1<float4>(0);
605 Index lhs_horiz = thread_y + k;
606 prefetch_lhs(lhs_pf0, lhs_vert, lhs_horiz)
608 Index rhs_vert = k + (thread_x % 4) * 4;
609 Index rhs_horiz0 = (thread_x >> 2) + thread_y * 4 + base_n;
611 if (!CHECK_RHS_BOUNDARY) {
612 if ((rhs_vert + 3) < k_size) {
614 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
615 }
else if (rhs_vert + 2 < k_size) {
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);
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);
644 if ((thread_x % 8) < 4) {
651#if defined(EIGEN_HIPCC)
652 x1 = __shfl_xor(x1, 4);
653 x2 = __shfl_xor(x2, 4);
655 x1 = __shfl_xor_sync(0xFFFFFFFF, x1, 4);
656 x2 = __shfl_xor_sync(0xFFFFFFFF, x2, 4);
658 if ((thread_x % 8) < 4) {
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);
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);
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; \
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; \
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; \
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;
711 for (
int koff = 0; koff < 16; koff++) {
713 float2 fl1 = lhs_shmem2[koff][thread_x];
714 float2 fl2 = lhs_shmem2[koff + 16][thread_x];
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];
720 add_vals(fl1, fl2, fr1, fr2)
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;
736 }
else if (!CHECK_RHS_BOUNDARY) {
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;
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;
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;
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;
761 }
else if (!CHECK_LHS_BOUNDARY) {
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;
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;
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) {
792 float4 lhs_pf0, lhs_pf1, lhs_pf2, lhs_pf3;
793 float4 rhs_pf0, rhs_pf1;
796 const Index thread_x = threadIdx.x;
797 const Index thread_y = threadIdx.y;
799 for (
int i = 0; i < 8; i++) {
800 results[i].x = results[i].y = results[i].z = results[i].w = 0;
803 Index lhs_vert = base_m + thread_x * 4 + (thread_y % 4) * 32;
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);
811 rhs_pf0 = internal::pset1<float4>(0);
812 rhs_pf1 = internal::pset1<float4>(0);
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));
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));
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));
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));
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));
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) {
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) {
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);
955 if (rhs_horiz1 < n_size) {
956 if ((rhs_vert + 3) < k_size) {
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) {
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);
977 }
else if (rhs_horiz0 < n_size) {
978 if ((rhs_vert + 3) < k_size) {
980 rhs_pf0 = rhs.template loadPacket<float4, Unaligned>(rhs_vert, rhs_horiz0);
981 }
else if ((rhs_vert + 2) < k_size) {
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);
1000 rhs_shmem2[thread_y][thread_x] = make_float2(rhs_pf0.x, rhs_pf1.x);
1004 rhs_shmem2[thread_y + 32][thread_x] = make_float2(rhs_pf0.y, rhs_pf1.y);
1007 rhs_shmem2[thread_y + 64][thread_x] = make_float2(rhs_pf0.z, rhs_pf1.z);
1010 rhs_shmem2[thread_y + 96][thread_x] = make_float2(rhs_pf0.w, rhs_pf1.w);
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; \
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; \
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; \
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;
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);
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);
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];
1075 int start_feature = (thread_y / 4) * 8;
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];
1082 add_vals(a3, a4, br1, br2, br3, br4)
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;
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;
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;
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;
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;
1122 }
else if (!CHECK_LHS_BOUNDARY) {
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;
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;
1145template <
typename Index,
typename LhsMapper,
typename RhsMapper,
typename OutputMapper>
1147#if defined(EIGEN_HIPCC)
1148__launch_bounds__(256, 1)
1150__launch_bounds__(256)
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];
1157 typedef float2 LHS_MEM[64][32];
1158 typedef float2 RHS_MEM[128][8];
1160 const Index m_block_idx = blockIdx.x;
1161 const Index n_block_idx = blockIdx.y;
1163 const Index base_m = 128 * m_block_idx;
1164 const Index base_n = 64 * n_block_idx;
1166 bool check_rhs = (base_n + 63) >= n_size;
1167 bool check_lhs128 = (base_m + 127) >= m_size;
1170 if (!check_lhs128) {
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);
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);
1179 if (!check_lhs128) {
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);
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);
1190template <
typename Index,
typename LhsMapper,
typename RhsMapper,
typename OutputMapper>
1192#if defined(EIGEN_HIPCC)
1193__launch_bounds__(256, 1)
1195__launch_bounds__(256)
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];
1202 const Index m_block_idx = blockIdx.x;
1203 const Index n_block_idx = blockIdx.y;
1205 const Index base_m = 64 * m_block_idx;
1206 const Index base_n = 64 * n_block_idx;
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);
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);
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);
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);
1227template <
typename Indices,
typename LeftArgType,
typename RightArgType,
typename OutputKernelType>
1229 :
public TensorContractionEvaluatorBase<TensorEvaluator<
1230 const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, GpuDevice> > {
1231 typedef GpuDevice Device;
1233 typedef TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device> Self;
1234 typedef TensorContractionEvaluatorBase<Self> Base;
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;
1242 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
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;
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;
1257 typedef array<Index, LDims> left_dim_mapper_t;
1258 typedef array<Index, RDims> right_dim_mapper_t;
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;
1264 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
1266 typedef DSizes<Index, NumDims> Dimensions;
1269 typedef std::remove_const_t<typename EvalLeftArgType::Scalar> LhsScalar;
1270 typedef std::remove_const_t<typename EvalRightArgType::Scalar> RhsScalar;
1272 typedef TensorEvaluator<EvalLeftArgType, Device> LeftEvaluator;
1273 typedef TensorEvaluator<EvalRightArgType, Device> RightEvaluator;
1275 typedef typename LeftEvaluator::Dimensions LeftDimensions;
1276 typedef typename RightEvaluator::Dimensions RightDimensions;
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);
1284 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(Scalar* data) {
1285 this->m_leftImpl.evalSubExprsIfNeeded(
nullptr);
1286 this->m_rightImpl.evalSubExprsIfNeeded(
nullptr);
1291 this->m_result =
static_cast<Scalar*
>(this->m_device.allocate(this->dimensions().TotalSize() *
sizeof(Scalar)));
1292 evalTo(this->m_result);
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);
1303 evalTyped<true, true, false, Unaligned>(buffer);
1306 if (this->m_rhs_inner_dim_reordered) {
1307 evalTyped<true, false, true, Unaligned>(buffer);
1309 evalTyped<true, false, false, Unaligned>(buffer);
1313 if (this->m_rhs_inner_dim_contiguous) {
1314 if (this->m_rhs_inner_dim_reordered) {
1315 evalTyped<false, true, true, Unaligned>(buffer);
1317 evalTyped<false, true, false, Unaligned>(buffer);
1320 if (this->m_rhs_inner_dim_reordered) {
1321 evalTyped<false, false, true, Unaligned>(buffer);
1323 evalTyped<false, false, false, Unaligned>(buffer);
1329 template <
int Alignment>
1330 void evalProduct(Scalar* buffer)
const {
1334 template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename LhsMapper,
typename RhsMapper,
1335 typename OutputMapper,
bool UseNaiveKernel>
1336 struct LaunchKernelsImpl;
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);
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);
1364 template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename LhsMapper,
typename RhsMapper,
1365 typename OutputMapper>
1368 struct LaunchKernels
1369 : LaunchKernelsImpl<LhsScalar, RhsScalar, Index, LhsMapper, RhsMapper, OutputMapper, (sizeof(Scalar) > 4)> {};
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);
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);
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 {
1396 const Index k = this->m_k_size;
1398 const Index m = this->m_i_size;
1401 const Index n = this->m_j_size;
1403 if (m == 0 || n == 0)
return;
1406 this->m_device.fill(buffer, buffer + m * n, Scalar(0));
1410 typedef internal::TensorContractionInputMapper<LhsScalar, Index, internal::Lhs, LeftEvaluator, left_nocontract_t,
1411 contract_t, 4, lhs_inner_dim_contiguous,
false,
Unaligned>
1414 typedef internal::TensorContractionInputMapper<RhsScalar, Index, internal::Rhs, RightEvaluator, right_nocontract_t,
1415 contract_t, 4, rhs_inner_dim_contiguous, rhs_inner_dim_reordered,
1419 typedef internal::blas_data_mapper<Scalar, Index, ColMajor> OutputMapper;
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);
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);
1428 OutputMapper output(buffer, m);
1429 LaunchKernels<LhsScalar, RhsScalar, Index, LhsMapper, RhsMapper, OutputMapper>::Run(lhs, rhs, output, m, n, k,
Definition TensorContraction.h:335
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47