1#ifndef __MATH_AX_HELM_KERNEL_H__
2#define __MATH_AX_HELM_KERNEL_H__
45template<
typename T, const
int LX, const
int CHUNKS >
119 for (
int l = 0; l<
LX; l++){
141 for (
int l = 0; l<
LX; l++){
152template<
typename T, const
int LX, const
int EB >
178 static_assert(
sizeof(
shdx) +
185 "kstep block exceeds the LDS budget");
221 for (
int k = 0;
k <
LX; ++
k){
233 for (
int l = 0; l <
LX; l++){
241 for (
int l = 0; l <
LX; l++){
262 for (
int l = 0; l <
LX; l++){
271 for (
int k = 0;
k <
LX; ++
k){
282template<
typename T, const
int LX, const
int EB >
308 static_assert(
sizeof(
shdx) +
315 "kstep block exceeds the LDS budget");
342 for(
int k = 0;
k <
LX; ++
k){
350 for (
int k = 0;
k <
LX; ++
k){
362 for (
int l = 0; l <
LX; l++){
370 for (
int l = 0; l <
LX; l++){
391 for (
int l = 0; l <
LX; l++){
400 for (
int k = 0;
k <
LX; ++
k){
446#if defined(__gfx90a__) || defined(__gfx942__)
462template<
typename T, const
int LX, const
int NWF, const
int TILE >
496 "wavefronts per block must split evenly over the elements");
506 static_assert(
sizeof(
shdx) +
sizeof(
shdy) +
sizeof(
shdz) +
509 "mfma block exceeds the shared memory budget");
513 const int tid = wf * 64 +
lane;
516 const int eb = wf /
WPE;
539 for (
int q = 0; q <
SPT; q++) {
556 for (
int q = 0; q <
SPT; q++) {
586 for (
int q = 0; q <
NG; q++) {
599 mfma_contract_sel<T, LX, 0, false, false, WPE, TILE, PAD>::run(
shr +
sh,
601 mfma_contract_sel<T, LX, 1, false, false, WPE, TILE, PAD>::run(
shs +
sh,
603 mfma_contract_sel<T, LX, 2, false, false, WPE, TILE, PAD>::run(
sht +
sh,
610 for (
int q = 0; q <
SPT; q++) {
618 const int r =
GREG ? q : 0;
644 mfma_contract_sel<T, LX, 0, true, false, WPE, TILE, PAD>::run(
shu +
sh,
647 mfma_contract_sel<T, LX, 1, true, true, WPE, TILE, PAD>::run(
shu +
sh,
650 mfma_contract_sel<T, LX, 2, true, true, WPE, TILE, PAD>::run(
shu +
sh,
656 for (
int q = 0; q <
SPT; q++) {
673template<
typename T, const
int LX, const
int NWF, const
int TILE >
676 const T *,
const T *,
const T *,
const T *,
677 const T *,
const T *,
const T *,
const int) {}
680#if defined(__gfx90a__) || defined(__gfx942__)
683#define NEKO_AX_HELM_MFMA_DISPATCH(TYPE, LXV) \
684 template< const int NWF, const int TILE > \
685 struct ax_helm_mfma_dispatch< TYPE, LXV, NWF, TILE > { \
686 __device__ static void run(TYPE *w, const TYPE *u, \
687 const TYPE *dx, const TYPE *dy, \
688 const TYPE *dz, const TYPE *h1, \
689 const TYPE *g11, const TYPE *g22, \
690 const TYPE *g33, const TYPE *g12, \
691 const TYPE *g13, const TYPE *g23, \
693 ax_helm_mfma_elem< TYPE, LXV, NWF, TILE >(w, u, dx, dy, dz, h1, \
694 g11, g22, g33, g12, g13, g23, \
726template<
typename T, const
int LX, const
int NWF, const
int TILE >
751template<
typename T, const
int LX, const
int EB >
789 static_assert(
sizeof(
shdx) +
802 "kstep block exceeds the LDS budget");
835 for(
int k = 0;
k <
LX; ++
k){
849 for (
int k = 0;
k <
LX; ++
k){
865 for (
int l = 0; l <
LX; l++){
881 for (
int l = 0; l <
LX; l++){
937 for (
int l = 0; l <
LX; l++){
956 for (
int k = 0;
k <
LX; ++
k){
964template<
typename T, const
int LX, const
int EB >
1002 static_assert(
sizeof(
shdx) +
1015 "kstep block exceeds the LDS budget");
1050 for(
int k = 0;
k <
LX; ++
k){
1064 for (
int k = 0;
k <
LX; ++
k){
1080 for (
int l = 0; l <
LX; l++){
1096 for (
int l = 0; l <
LX; l++){
1152 for (
int l = 0; l <
LX; l++){
1171 for (
int k = 0;
k <
LX; ++
k){
1211#if defined(__gfx90a__) || defined(__gfx942__)
1215template<
typename T, const
int LX, const
int NWF, const
int TILE >
1253 "wavefronts per block must split evenly over the elements");
1263 static_assert(
sizeof(
shdx) +
sizeof(
shdy) +
sizeof(
shdz) +
1266 "mfma vector block exceeds the shared memory budget");
1270 const int tid = wf * 64 +
lane;
1273 const int eb = wf /
WPE;
1274 const int sub = wf %
WPE;
1291 for (
int q = 0; q <
SPT; q++) {
1316 for (
int q = 0; q <
NG; q++) {
1329 for (
int c = 0; c < 3; c++) {
1333 const T *
const cin = (c == 0) ?
u : (c == 1) ?
v :
w;
1334 T *
const cout = (c == 0) ?
au : (c == 1) ?
av :
aw;
1340 for (
int q = 0; q <
SPT; q++) {
1348 mfma_contract_sel<T, LX, 0, false, false, WPE, TILE, PAD>::run(
shr +
sh,
1350 mfma_contract_sel<T, LX, 1, false, false, WPE, TILE, PAD>::run(
shs +
sh,
1352 mfma_contract_sel<T, LX, 2, false, false, WPE, TILE, PAD>::run(
sht +
sh,
1361 for (
int q = 0; q <
SPT; q++) {
1369 const int r =
GREG ? q : 0;
1390 mfma_contract_sel<T, LX, 0, true, false, WPE, TILE, PAD>::run(
shc +
sh,
1393 mfma_contract_sel<T, LX, 1, true, true, WPE, TILE, PAD>::run(
shc +
sh,
1396 mfma_contract_sel<T, LX, 2, true, true, WPE, TILE, PAD>::run(
shc +
sh,
1402 for (
int q = 0; q <
SPT; q++) {
1431template<
typename T, const
int LX, const
int NWF, const
int TILE >
1434 const T *,
const T *,
const T *,
1435 const T *,
const T *,
const T *,
const T *,
1436 const T *,
const T *,
const T *,
1437 const T *,
const T *,
const T *,
const int) {}
1440#if defined(__gfx90a__) || defined(__gfx942__)
1443#define NEKO_AX_HELM_MFMA_VECTOR_DISPATCH(TYPE, LXV) \
1444 template< const int NWF, const int TILE > \
1445 struct ax_helm_mfma_vector_dispatch< TYPE, LXV, NWF, TILE > { \
1446 __device__ static void run(TYPE *au, TYPE *av, TYPE *aw, \
1447 const TYPE *u, const TYPE *v, const TYPE *w, \
1448 const TYPE *dx, const TYPE *dy, \
1449 const TYPE *dz, const TYPE *h1, \
1450 const TYPE *g11, const TYPE *g22, \
1451 const TYPE *g33, const TYPE *g12, \
1452 const TYPE *g13, const TYPE *g23, \
1454 ax_helm_mfma_vector_elem< TYPE, LXV, NWF, TILE >(au, av, aw, u, v, w, \
1457 g12, g13, g23, nelv); \
1490template<
typename T, const
int LX, const
int NWF, const
int TILE >
1516template<
typename T >
1530 for (
int i = idx;
i < n;
i +=
str) {
__global__ void ale_add_kinematics_kernel(const int n, T *__restrict__ wx, T *__restrict__ wy, T *__restrict__ wz, const T *__restrict__ x_ref, const T *__restrict__ y_ref, const T *__restrict__ z_ref, const T *__restrict__ phi, const T *__restrict__ x, const T *__restrict__ y, const T *__restrict__ z, const kinematics_params_t kin_params)
__shared__ T shus[EB *LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g23
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g22
__global__ void T *__restrict__ av
__shared__ T shv[EB *LX *LX]
__shared__ T shwr[EB *LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g13
__global__ void ax_helm_kernel_1d(T *__restrict__ w, const T *__restrict__ u, const T *__restrict__ dx, const T *__restrict__ dy, const T *__restrict__ dz, const T *__restrict__ dxt, const T *__restrict__ dyt, const T *__restrict__ dzt, const T *__restrict__ h1, const T *__restrict__ g11, const T *__restrict__ g22, const T *__restrict__ g33, const T *__restrict__ g12, const T *__restrict__ g13, const T *__restrict__ g23)
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const int nelv
__shared__ T shur[EB *LX *LX]
__shared__ T shvs[EB *LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g12
__shared__ T shw[EB *LX *LX]
__shared__ T shws[EB *LX *LX]
__shared__ T shu[EB *LX *LX]
__global__ void T *__restrict__ T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ w
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g33
__global__ void T *__restrict__ T *__restrict__ aw
__shared__ T shvr[EB *LX *LX]
__global__ void const T *__restrict__ u
__global__ void const T *__restrict__ const T *__restrict__ dx
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dz
__shared__ T shdz[LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ dy
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ h1
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g11
__global__ void T *__restrict__ T *__restrict__ const T *__restrict__ const T *__restrict__ v
__global__ void ax_helm_kernel_vector_part2(T *__restrict__ au, T *__restrict__ av, T *__restrict__ aw, const T *__restrict__ u, const T *__restrict__ v, const T *__restrict__ w, const T *__restrict__ h2, const T *__restrict__ B, const int n)
__shared__ T shdy[LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dyt
__shared__ T shdzt[LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dzt
__shared__ T shdyt[LX *LX]
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dxt
#define NEKO_EB_BOUNDS(NT)
__global__ void __launch_bounds__((LX *LX *EB)) ax_helm_kernel_kstep(T *__restrict__ w
#define NEKO_MFMA_CUBE_N(LX, SZ)
#define NEKO_MFMA_EB_N(NWF, LX, SZ)
#define NEKO_MFMA_VECTOR_GREG_N(SPT, SZ)
#define NEKO_MFMA_SPT_N(WPE, LX, SZ)
#define NEKO_MFMA_PAD_N(LX, SZ)
#define NEKO_MFMA_DMAT_N(LX, SZ)
static __device__ void run(T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const int)
static __device__ void run(T *, T *, T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const T *, const int)