1#ifndef __MATH_DMMA_TMA_KERNEL_H__
2#define __MATH_DMMA_TMA_KERNEL_H__
95#include <cuda_runtime.h>
107#define DMMA_NG_OPGRAD 9
111#define DMMA_NG_CONV1 13
119#if defined(CUDART_VERSION) && (CUDART_VERSION >= 12000)
120#define NEKO_TMA_TOOLKIT 1
122#define NEKO_TMA_TOOLKIT 0
135#ifndef NEKO_TMA_ARCH_COMPILED
136#if defined(__CUDA_ARCH_LIST__)
198template< const
int LX >
213template< const
int LX >
228template< const
int LX >
241template< const
int LX >
247template< const
int LX >
258template< const
int LX >
295#define NEKO_DMMA_TMA_BATCH_SMEM ((int) sizeof(dmma_tma_batch_smem))
300 "dmma tma batch block exceeds the opt-in shared memory ceiling");
303 "dmma tma batch bulk copy targets must stay 128 byte aligned");
372#define NEKO_OPGRAD_TMA_SMEM ((int) sizeof(opgrad_tma_smem))
401#define NEKO_CONV1_TMA_SMEM ((int) sizeof(conv1_tma_smem))
405 "tma block exceeds the opt-in shared memory ceiling");
410 "tma bulk copy targets must stay 128 byte aligned");
434template< const
int LX >
456 const void *
h1,
const void *
g11,
457 const void *
g22,
const void *
g33,
458 const void *
g12,
const void *
g13,
470 const void *
aw,
const void *
u,
471 const void *
v,
const void *
w,
472 const void *
h1,
const void *
g11,
473 const void *
g22,
const void *
g33,
474 const void *
g12,
const void *
g13,
484 const void *
dr,
const void *
ds,
528 const void *
dr,
const void *
ds,
538 const void *
vx,
const void *
vy,
557 const char *
v =
getenv(
"NEKO_DMMA_TMA_NW");
567#define NEKO_TUNE_LOG_DMMA_TMA_BATCH(LX, T5) \
569 for (int c = 0; c < NEKO_DMMA_CANDIDATES; c++) { \
570 if ((T5)[c] >= NEKO_TUNE_INIT) { continue; } \
572 sprintf(lbl_, "TMAB %dw", NEKO_DMMA_NW(c)); \
573 sprintf(neko_log_buf, "%-13s: %9.2f us/call", lbl_, \
574 NEKO_TUNE_US((T5)[c], iters)); \
575 log_message(neko_log_buf); \
580#define NEKO_TUNE_LOG_DMMA_TMA(LX, T4) \
582 for (int c = 0; c < NEKO_DMMA_CANDIDATES; c++) { \
583 if ((T4)[c] >= NEKO_TUNE_INIT) { continue; } \
585 sprintf(lbl_, "TMA %dw", NEKO_DMMA_NW(c)); \
586 sprintf(neko_log_buf, "%-13s: %9.2f us/call", lbl_, \
587 NEKO_TUNE_US((T4)[c], iters)); \
588 log_message(neko_log_buf); \
592#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) && \
593 (__CUDA_ARCH__ < 1000) && NEKO_TMA_TOOLKIT
628 asm volatile(
"mbarrier.init.shared::cta.b64 [%0], %1;"
639 asm volatile(
"mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;"
646 unsigned long long *
bar)
648 asm volatile(
"cp.async.bulk.shared::cluster.global"
649 ".mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
670 "mbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
671 "selp.b32 %0, 1, 0, p;\n"
685 asm volatile(
"fence.proxy.async.shared::cta;" :::
"memory");
694 asm volatile(
"cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;"
705 asm volatile(
"cp.async.bulk.commit_group;" :::
"memory");
706 asm volatile(
"cp.async.bulk.wait_group 0;" :::
"memory");
__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)
__global__ void T *__restrict__ 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 T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dtdy
__global__ void T *__restrict__ 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 T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dtdx
__global__ void T *__restrict__ 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 T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dtdz
__global__ void T *__restrict__ T *__restrict__ aw
__global__ void T *__restrict__ T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ w
__global__ void T *__restrict__ 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 T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ jacinv
__global__ void T *__restrict__ 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 T *__restrict__ const T *__restrict__ const T *__restrict__ dsdz
__global__ void T *__restrict__ T *__restrict__ const T *__restrict__ u
__global__ void T *__restrict__ av
__global__ void T *__restrict__ 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__ drdz
__global__ void T *__restrict__ 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__ drdx
__global__ void T *__restrict__ 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 T *__restrict__ dsdx
__global__ void T *__restrict__ T *__restrict__ const T *__restrict__ const T *__restrict__ v
__global__ void T *__restrict__ 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 T *__restrict__ const T *__restrict__ dsdy
__global__ void T *__restrict__ T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ h1
__global__ void T *__restrict__ 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__ drdy
__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 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 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
__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 const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ g11
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ ds
__global__ void const T *__restrict__ x
__global__ void const T *__restrict__ const T *__restrict__ dr
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ dt
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ const T *__restrict__ vz
__global__ void const T *__restrict__ const T *__restrict__ vx
__global__ void const T *__restrict__ const T *__restrict__ const T *__restrict__ vy
#define NEKO_DMMA_CANDIDATES
static bool dmma_tma_aligned(const void *w, const void *u, const void *h1, const void *g11, const void *g22, const void *g33, const void *g12, const void *g13, const void *g23)
static bool dmma_tma_vector_aligned(const void *au, const void *av, const void *aw, const void *u, const void *v, const void *w, const void *h1, const void *g11, const void *g22, const void *g33, const void *g12, const void *g13, const void *g23)
static bool cuda_have_tma_conv1()
static bool dmma_tma_cdtp_aligned(const void *dtx, const void *x, const void *dr, const void *ds, const void *dt)
static bool cuda_have_tma_batch()
static bool dmma_tma_batch_lx_supported()
#define NEKO_OPGRAD_TMA_SMEM
static bool dmma_tma_vector_lx_supported()
static bool tma_arch_compiled()
#define NEKO_CONV1_TMA_SMEM
static bool dmma_tma_cdtp_lx_supported()
static int cuda_tma_smem_optin()
static bool dmma_tma_dudxyz_lx_supported()
static bool cuda_have_tma_opgrad()
static bool dmma_tma_conv1_lx_supported()
static bool dmma_tma_lx_supported()
static int neko_dmma_tma_env()
static bool dmma_tma_metrics_aligned(const void *drdx, const void *dsdx, const void *dtdx, const void *drdy, const void *dsdy, const void *dtdy, const void *drdz, const void *dsdz, const void *dtdz)
#define NEKO_DMMA_TMA_BATCH_SMEM
static bool dmma_tma_opgrad_aligned(const void *u, const void *drdx, const void *dsdx, const void *dtdx, const void *drdy, const void *dsdy, const void *dtdy, const void *drdz, const void *dsdz, const void *dtdz)
static bool dmma_tma_opgrad_lx_supported()
static bool dmma_tma_conv1_aligned(const void *du, const void *u, const void *vx, const void *vy, const void *vz, const void *jacinv, const void *drdx, const void *dsdx, const void *dtdx, const void *drdy, const void *dsdy, const void *dtdy, const void *drdz, const void *dsdz, const void *dtdz)
static bool dmma_tma_dudxyz_aligned(const void *du, const void *u, const void *dr, const void *ds, const void *dt, const void *jacinv)
static bool cuda_have_tma()
static bool dmma_tma_ptr_aligned(const void *p)