IQ.Pilot Release Commit @ 0798119

This commit is contained in:
IQ.Lvbs history cleanup
2026-08-22 23:42:42 -05:00
commit b42569dbca
4529 changed files with 1132125 additions and 0 deletions

View File

@@ -0,0 +1,17 @@
/**
* @file
* @brief An aggregate header for all group-scope MMA operations.
*/
// All compilation targets can use the warp-scope MMA operations.
#include "warp/warp.cuh"
// Hopper has its own warpgroup-scope MMA operations.
#ifdef KITTENS_HOPPER
#include "warpgroup/warpgroup.cuh"
#endif
// Blackwell has its own tensor-scope MMA operations.
#ifdef KITTENS_BLACKWELL
#include "tensor/tensor.cuh"
#endif

View File

@@ -0,0 +1,172 @@
/**
* @file Group-level tcgen05 MMA operations.
*/
template<int trans_a, int n_trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1, int ncta=1>
__device__ static inline void mma(D &d, const A &a, const B &b, semaphore &sem) {
if(laneid() == 0) ::kittens::mma<trans_a, n_trans_b, D, A, B, acc, ncta>(d, a, b, sem);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1>
__device__ static inline void mma2(D &d, const A &a, const B &b, semaphore &sem) {
mma<trans_a, trans_b, D, A, B, acc, 2>(d, a, b, sem);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm(D &d, const A &a, const B &b, semaphore &sem) {
mma<trans_a, trans_b, D, A, B, 0>(d, a, b, sem);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2(D &d, const A &a, const B &b, semaphore &sem) {
mma2<trans_a, trans_b, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::T, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::T, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::T, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::T, D, A, B, 0>(d, a, b, sem);
}
// no sem versions
template<int trans_a, int n_trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1, int ncta=1>
__device__ static inline void mma(D &d, const A &a, const B &b) {
if(laneid() == 0) ::kittens::mma<trans_a, n_trans_b, D, A, B, acc, ncta>(d, a, b);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1>
__device__ static inline void mma2(D &d, const A &a, const B &b) {
mma<trans_a, trans_b, D, A, B, acc, 2>(d, a, b);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm(D &d, const A &a, const B &b) {
mma<trans_a, trans_b, D, A, B, 0>(d, a, b);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2(D &d, const A &a, const B &b) {
mma2<trans_a, trans_b, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AB(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AB(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_ABt(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_ABt(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtB(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtB(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtBt(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtBt(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AB(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AB(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_ABt(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::T, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_ABt(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::T, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtB(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtB(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtBt(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::T, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtBt(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::T, D, A, B, 0>(d, a, b);
}

View File

@@ -0,0 +1,947 @@
/**
* @file
* @brief Matrix multiply-accumulate operations for tiles stored in registers.
*/
/**
* @brief Perform the HMMA.16816 operation.
*
* This function performs the half-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32` instruction.
*
* @param[out] d0 The first half of the output float2 accumulator.
* @param[out] d1 The second half of the output float2 accumulator.
* @param[in] a0 The first half of the first input bf16_2 matrix.
* @param[in] a1 The second half of the first input bf16_2 matrix.
* @param[in] a2 The first half of the second input bf16_2 matrix.
* @param[in] a3 The second half of the second input bf16_2 matrix.
* @param[in] b0 The first half of the bf16_2 matrix B.
* @param[in] b1 The second half of the bf16_2 matrix B.
* @param[in] c0 The first half of the float2 accumulator matrix C.
* @param[in] c1 The second half of the float2 accumulator matrix C.
*/
__device__ static inline void hmma16816( float2 &d0, float2 &d1,
const bf16_2 &a0, const bf16_2 &a1, const bf16_2 &a2, const bf16_2 &a3,
const bf16_2 &b0, const bf16_2 &b1,
const float2 &c0, const float2 &c1 ) {
asm volatile(
// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#multiply-and-accumulate-instruction-mma
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 " \
"{%0, %1, %2, %3}, " \
"{%4, %5, %6, %7}, " \
"{%8, %9}, " \
"{%10, %11, %12, %13};"
// D matrix
: "+f"(d0.x), "+f"(d0.y),
"+f"(d1.x), "+f"(d1.y)
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"f"(c0.x), "f"(c0.y),
"f"(c1.x), "f"(c1.y)
);
}
/**
* @brief Perform the HMMA.16816 operation with inputs as fp16 and fp32 accumulators
*
* This function performs the half-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32` instruction.
*
* @param[out] d0 The first half of the output float2 accumulator.
* @param[out] d1 The second half of the output float2 accumulator.
* @param[in] a0 The first half of the first input half_2 matrix.
* @param[in] a1 The second half of the first input half_2 matrix.
* @param[in] a2 The first half of the second input half_2 matrix.
* @param[in] a3 The second half of the second input half_2 matrix.
* @param[in] b0 The first half of the half_2 matrix B.
* @param[in] b1 The second half of the half_2 matrix B.
* @param[in] c0 The first half of the float2 accumulator matrix C.
* @param[in] c1 The second half of the float2 accumulator matrix C.
*/
__device__ static inline void hmma16816( float2 &d0, float2 &d1,
const half_2 &a0, const half_2 &a1, const half_2 &a2, const half_2 &a3,
const half_2 &b0, const half_2 &b1,
const float2 &c0, const float2 &c1 ) {
asm volatile(
// https://docs.nvidia.com/cuda/parallel-thread-execution/#multiply-and-accumulate-instruction-mma
"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 " \
"{%0, %1, %2, %3}, " \
"{%4, %5, %6, %7}, " \
"{%8, %9}, " \
"{%10, %11, %12, %13};"
// D matrix
: "+f"(d0.x), "+f"(d0.y),
"+f"(d1.x), "+f"(d1.y)
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"f"(c0.x), "f"(c0.y),
"f"(c1.x), "f"(c1.y)
);
}
/**
* @brief Perform the HMMA.16816 operation.
*
* This function performs the half-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16` instruction.
*
* @param[out] d0 The first half of the output half_2 accumulator.
* @param[out] d1 The second half of the output half_2 accumulator.
* @param[in] a0 The first half of the first input half_2 matrix.
* @param[in] a1 The second half of the first input half_2 matrix.
* @param[in] a2 The first half of the second input half_2 matrix.
* @param[in] a3 The second half of the second input half_2 matrix.
* @param[in] b0 The first half of the half_2 matrix B.
* @param[in] b1 The second half of the half_2 matrix B.
* @param[in] c0 The first half of the half_2 accumulator matrix C.
* @param[in] c1 The second half of the half_2 accumulator matrix C.
*/
__device__ static inline void hmma16816( half_2 &d0, half_2 &d1,
const half_2 &a0, const half_2 &a1, const half_2 &a2, const half_2 &a3,
const half_2 &b0, const half_2 &b1,
const half_2 &c0, const half_2 &c1 ) {
asm volatile(
// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#multiply-and-accumulate-instruction-mma
"mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 " \
"{%0, %1}, " \
"{%2, %3, %4, %5}, " \
"{%6, %7}, " \
"{%8, %9};"
// D matrix
: "=r"(*(uint32_t*)(&d0)), "=r"(*(uint32_t*)(&d1))
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"r"(*(uint32_t*)(&c0)), "r"(*(uint32_t*)(&c1))
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Perform the HMMA.16816 operation for FP8 using fp8e4m3_2.
*
* Using mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 instruction
* but with fp8e4m3_2 (2 FP8 values) instead of fp8e4m3_4
*/
/**
* @brief Perform the HMMA.16816 operation for FP8.
*
* This function performs the fp8-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32` instruction.
*
* @param[out] d0 The first half of the output float2 accumulator.
* @param[out] d1 The second half of the output float2 accumulator.
* @param[in] a0,a1,a2,a3 Input FP8 matrix A values
* @param[in] b0,b1 Input FP8 matrix B values
* @param[in] c0,c1 Input float2 accumulator matrix C values
*/
__device__ static inline void hmma16816( float2 &d0, float2 &d1,
const fp8e4m3_4 &a0, const fp8e4m3_4 &a1,
const fp8e4m3_4 &a2, const fp8e4m3_4 &a3,
const fp8e4m3_4 &b0, const fp8e4m3_4 &b1,
const float2 &c0, const float2 &c1) {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0, %1, %2, %3}, "
"{%4, %5, %6, %7}, "
"{%8, %9}, "
"{%10, %11, %12, %13};"
// D matrix (output)
: "+f"(d0.x), "+f"(d0.y),
"+f"(d1.x), "+f"(d1.y)
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"f"(c0.x), "f"(c0.y),
"f"(c1.x), "f"(c1.y)
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, row_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::row> &a,
const rt_base<bf16, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base matrix multiply-accumulate operation for row layout
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, row_layout> matrix.
* @param[in] b The second input rt_base<half_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::row> &a,
const rt_base<half, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base matrix multiply-accumulate operation for row layout.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3, row_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::row> &a,
const rt_base<fp8e4m3, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<half_2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, row_layout> matrix.
* @param[in] b The second input rt_base<half_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<half_2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<half, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::row> &a,
const rt_base<half, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<half, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base dot product operation for row layout.
*
* This function performs the base dot product operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, row_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_ABt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::row> &a,
const rt_base<bf16, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2], // for some reason this one seems to need to be backwards
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3], // for some reason this one seems to need to be backwards
c.data[2], c.data[3]
);
}
/**
* @brief Base dot product operation for row layout
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base dot product operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, row_layout> matrix.
* @param[in] b The second input rt_base<half_2, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_ABt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::row> &a,
const rt_base<half, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2], // for some reason this one seems to need to be backwards
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3], // for some reason this one seems to need to be backwards
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base dot product operation for row layout.
*
* This function performs the base dot product operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3x4, row_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3x4, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_ABt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::row> &a,
const rt_base<fp8e4m3, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2], // for some reason this one seems to need to be backwards
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3], // for some reason this one seems to need to be backwards
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, col_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::col> &a,
const rt_base<bf16, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, col_layout> matrix.
* @param[in] b The second input rt_base<half_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::col> &a,
const rt_base<half, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3x4, col_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3x4, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::col> &a,
const rt_base<fp8e4m3, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A and B.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, col_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtBt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::col> &a,
const rt_base<bf16, ducks::rt_layout::row> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A and B
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, col_layout> matrix.
* @param[in] b The second input rt_base<half_2, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtBt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::col> &a,
const rt_base<half, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A and B.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3x4, col_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3x4, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtBt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::col> &a,
const rt_base<fp8e4m3, ducks::rt_layout::row> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Matrix multiply-accumulate operation.
*
* This function performs the matrix multiply-accumulate operation
* using the `hmma16816` function.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_hf<N, M, row_layout> accumulator.
* @param[in] a The first input rt_hf<N, K, row_layout> matrix.
* @param[in] b The second input rt_hf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_hf<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::row_layout A, ducks::rt::col_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_AB(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::rows && D::cols == B::cols); // Check D matches A, B
static_assert(A::cols == B::rows); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_AB_base(
d.tiles[n][m],
a.tiles[n][0],
b.tiles[0][m],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::width; k++) {
mma_AB_base(
d.tiles[n][m],
a.tiles[n][k],
b.tiles[k][m],
d.tiles[n][m]
);
}
}
}
}
/**
* @brief Dot product operation for row layout.
*
* This function performs the dot product operation
* using the `hmma16816` function.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_bf<N, K, row_layout> matrix.
* @param[in] b The second input rt_bf<M, K, row_layout> matrix in row-major mode.
* @param[in] c The input rt_fl<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::row_layout A, ducks::rt::row_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_ABt(D &d,
const A &a,
const B &b, // notice row and (M, K) instead of col and (K, M)
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::rows && D::cols == B::rows); // Check D matches A, B
static_assert(A::cols == B::cols); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_ABt_base(
d.tiles[n][m],
a.tiles[n][0],
b.tiles[m][0],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::width; k++) {
mma_ABt_base(
d.tiles[n][m],
a.tiles[n][k],
b.tiles[m][k],
d.tiles[n][m]
);
}
}
}
}
/**
* @brief Matrix multiply-accumulate operation with transposed A.
*
* This function performs the matrix multiply-accumulate operation
* using the `hmma16816` instruction.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_bf<K, N, row_layout> matrix.
* @param[in] b The second input rt_bf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_fl<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::col_layout A, ducks::rt::col_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_AtB(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::cols && D::cols == B::cols); // Check D matches A, B
static_assert(A::rows == B::rows); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_AtB_base(
d.tiles[n][m],
a.tiles[0][n],
b.tiles[0][m],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::height; k++) {
mma_AtB_base(
d.tiles[n][m],
a.tiles[k][n],
b.tiles[k][m],
d.tiles[n][m]
);
}
}
}
}
/**
* @brief Matrix multiply-accumulate operation with transposed A and B.
*
* This function performs the matrix multiply-accumulate operation
* using the `hmma16816` instruction.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_bf<K, N, col_layout> matrix.
* @param[in] b The second input rt_bf<M, K, row_layout> matrix in column-major mode.
* @param[in] c The input rt_fl<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::col_layout A, ducks::rt::row_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_AtBt(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::cols && D::cols == B::rows); // Check D matches A, B
static_assert(A::rows == B::cols); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_AtBt_base(
d.tiles[n][m],
a.tiles[0][n],
b.tiles[m][0],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::height; k++) {
mma_AtBt_base(
d.tiles[n][m],
a.tiles[k][n],
b.tiles[m][k],
d.tiles[n][m]
);
}
}
}
}
template<int trans_A, int trans_B, ducks::rt::all D, ducks::rt::all A, ducks::rt::all B, ducks::rt::all C>
__device__ static inline void mma(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
if constexpr(trans_A == transpose::T) {
if constexpr(trans_B == transpose::T) {
mma_AtBt(d, a, b, c);
} else {
mma_AtB(d, a, b, c);
}
} else {
if constexpr(trans_B == transpose::T) {
mma_ABt(d, a, b, c);
} else {
mma_AB(d, a, b, c);
}
}
}
template<int trans_A, int trans_B, ducks::rt::all A, ducks::rt::all B, ducks::rt::all C>
__device__ static inline C mma(const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
C d;
if constexpr(trans_A == transpose::T) {
if constexpr(trans_B == transpose::T) {
mma_AtBt(d, a, b, c);
} else {
mma_AtB(d, a, b, c);
}
} else {
if constexpr(trans_B == transpose::T) {
mma_ABt(d, a, b, c);
} else {
mma_AB(d, a, b, c);
}
}
return d;
}
// --------------------------------------------------------------------------------------------------------------------
// --------------------------------------------------------------------------------------------------------------------
// -------------------------------------------------- COMPLEX INPUTS --------------------------------------------------
// --------------------------------------------------------------------------------------------------------------------
// --------------------------------------------------------------------------------------------------------------------
/**
* @brief Matrix multiply-accumulate operation for complex tiles
*
* This function calls mma_AB with hf arguments
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_cmplx_hf<N, M, row_layout> accumulator.
* @param[in] a The first input rt_cmplx_hf<N, K, row_layout> matrix.
* @param[in] b The second input rt_cmplx_hf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_cmplx_hf<N, M, row_layout> accumulator matrix.
*/
template<int N, int K, int M>
__device__ static inline void mma_AB(crt_hf<N, M, ducks::rt_layout::row> &d,
const crt_hf<N, K, ducks::rt_layout::row> &a,
const crt_hf<K, M, ducks::rt_layout::col> &b,
const crt_hf<N, M, ducks::rt_layout::row> &c) {
KITTENS_CHECK_WARP
// Copy data from input accumulate register into output
::kittens::group<1>::copy(d.real, c.real);
::kittens::group<1>::copy(d.imag, c.imag);
// Negative on B matrix so we can use single accum register
rt_hf<N, K, ducks::rt_layout::row> tmp;
// Hex value for -1 in float16
constexpr half factor = std::bit_cast<__half>(uint16_t(0xFB80));
::kittens::group<1>::mul(tmp, a.imag, factor);
mma_AB(d.real, a.real, b.real, d.real);
mma_AB(d.real, tmp, b.imag, d.real);
mma_AB(d.imag, a.real, b.imag, d.imag);
mma_AB(d.imag, a.imag, b.real, d.imag);
}
/**
* @brief Matrix multiply-accumulate operation for complex tiles
*
* This function calls mma_AB with bf16 arguments
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_cmplx_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_cmplx_bf<N, K, row_layout> matrix.
* @param[in] b The second input rt_cmplx_bf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_cmplx_fl<N, M, row_layout> accumulator matrix.
*/
template<int N, int K, int M>
__device__ static inline void mma_AB(crt_fl<N, M, ducks::rt_layout::row> &d,
const crt_bf<N, K, ducks::rt_layout::row> &a,
const crt_bf<K, M, ducks::rt_layout::col> &b,
const crt_fl<N, M, ducks::rt_layout::row> &c) {
KITTENS_CHECK_WARP
// Copy data from input accumulate register into output
::kittens::group<1>::copy(d.real, c.real);
::kittens::group<1>::copy(d.imag, c.imag);
// Negative on B matrix so we can use single accum register
kittens::rt_bf<N, K, ducks::rt_layout::row> tmp;
// Hex value for -1 in bf16
constexpr bf16 factor = std::bit_cast<__nv_bfloat16>(uint16_t(0xBF80));
::kittens::group<1>::mul(tmp, a.imag, factor);
mma_AB(d.real, a.real, b.real, d.real);
mma_AB(d.real, tmp, b.imag, d.real);
mma_AB(d.imag, a.real, b.imag, d.imag);
mma_AB(d.imag, a.imag, b.real, d.imag);
}

View File

@@ -0,0 +1,334 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 112, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 112, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %61, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"{%56, %57, %58, %59}, " \
"%60, " \
"p, 1, %63, %62;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %61, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"{%56, %57, %58, %59}, " \
"%60, " \
"p, 1, %63, %62;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %33, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27}, " \
"{%28, %29, %30, %31}, " \
"%32, " \
"p, 1, %35, %34;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 112, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %58, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"%56, " \
"%57, " \
"p, 1, %61, %59, %60;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %58, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"%56, " \
"%57, " \
"p, 1, %61, %59, %60;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %30, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27}, " \
"%28, " \
"%29, " \
"p, 1, %33, %31, %32;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,813 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 128, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 128, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %71, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %71, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %39, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 128, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %69, %67, %68;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %69, %67, %68;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %37, %35, %36;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %67;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %67;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,382 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 144, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 144, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %77, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"{%72, %73, %74, %75}, " \
"%76, " \
"p, 1, %79, %78;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %77, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"{%72, %73, %74, %75}, " \
"%76, " \
"p, 1, %79, %78;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %41, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35}, " \
"{%36, %37, %38, %39}, " \
"%40, " \
"p, 1, %43, %42;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 144, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %74, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"%72, " \
"%73, " \
"p, 1, %77, %75, %76;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %74, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"%72, " \
"%73, " \
"p, 1, %77, %75, %76;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %38, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35}, " \
"%36, " \
"%37, " \
"p, 1, %41, %39, %40;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,190 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 16, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 16, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %15, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %15, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %9, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f16.f16.f16 " \
"{%0, %1, %2, %3}, " \
"{%4, %5, %6, %7}, " \
"%8, " \
"p, 1, %11, %10;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 16, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %13, %11, %12;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %13, %11, %12;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %6, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f16.f16.f16 " \
"{%0, %1, %2, %3}, " \
"%4, " \
"%5, " \
"p, 1, %9, %7, %8;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,666 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 160, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 160, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %87, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %87, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %45, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"{%40, %41, %42, %43}, " \
"%44, " \
"p, 1, %47, %46;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 160, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %85, %83, %84;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %85, %83, %84;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %42, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"%40, " \
"%41, " \
"p, 1, %45, %43, %44;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %83;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %83;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,430 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 176, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 176, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %93, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"{%88, %89, %90, %91}, " \
"%92, " \
"p, 1, %95, %94;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %93, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"{%88, %89, %90, %91}, " \
"%92, " \
"p, 1, %95, %94;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %49, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43}, " \
"{%44, %45, %46, %47}, " \
"%48, " \
"p, 1, %51, %50;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 176, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %90, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"%88, " \
"%89, " \
"p, 1, %93, %91, %92;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %90, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"%88, " \
"%89, " \
"p, 1, %93, %91, %92;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %46, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43}, " \
"%44, " \
"%45, " \
"p, 1, %49, %47, %48;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,674 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 192, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 192, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %103, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %103, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %55, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 192, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %98, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"%96, " \
"%97, " \
"p, 1, %101, %99, %100;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %98, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"%96, " \
"%97, " \
"p, 1, %101, %99, %100;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %53, %51, %52;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %98, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"%96, " \
"%97, " \
"p, 1, %99;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,478 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 208, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 208, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %109, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"{%104, %105, %106, %107}, " \
"%108, " \
"p, 1, %111, %110;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %109, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"{%104, %105, %106, %107}, " \
"%108, " \
"p, 1, %111, %110;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %57, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51}, " \
"{%52, %53, %54, %55}, " \
"%56, " \
"p, 1, %59, %58;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 208, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %106, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"%104, " \
"%105, " \
"p, 1, %109, %107, %108;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %106, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"%104, " \
"%105, " \
"p, 1, %109, %107, %108;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %54, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51}, " \
"%52, " \
"%53, " \
"p, 1, %57, %55, %56;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,826 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 224, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 224, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %119, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %119, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %61, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"{%56, %57, %58, %59}, " \
"%60, " \
"p, 1, %63, %62;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 224, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %58, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"%56, " \
"%57, " \
"p, 1, %61, %59, %60;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,526 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 240, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 240, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %125, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"{%120, %121, %122, %123}, " \
"%124, " \
"p, 1, %127, %126;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %125, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"{%120, %121, %122, %123}, " \
"%124, " \
"p, 1, %127, %126;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %65, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59}, " \
"{%60, %61, %62, %63}, " \
"%64, " \
"p, 1, %67, %66;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 240, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %122, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"%120, " \
"%121, " \
"p, 1, %125, %123, %124;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %122, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"%120, " \
"%121, " \
"p, 1, %125, %123, %124;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %62, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59}, " \
"%60, " \
"%61, " \
"p, 1, %65, %63, %64;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,446 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 32, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 32, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %23, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %23, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %15, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 32, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %21, %19, %20;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %21, %19, %20;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %13, %11, %12;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %11;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %11;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,238 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 48, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 48, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %31, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %31, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %17, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11}, " \
"{%12, %13, %14, %15}, " \
"%16, " \
"p, 1, %19, %18;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 48, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %29, %27, %28;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %29, %27, %28;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %14, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11}, " \
"%12, " \
"%13, " \
"p, 1, %17, %15, %16;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,587 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 64, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 64, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %39, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %39, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %23, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 64, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %37, %35, %36;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %37, %35, %36;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %21, %19, %20;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b), // transpose is not supported for FP8
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b), // transpose is not supported for FP8
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,286 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 80, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 80, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %45, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"{%40, %41, %42, %43}, " \
"%44, " \
"p, 1, %47, %46;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %45, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"{%40, %41, %42, %43}, " \
"%44, " \
"p, 1, %47, %46;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %25, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19}, " \
"{%20, %21, %22, %23}, " \
"%24, " \
"p, 1, %27, %26;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 80, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %42, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"%40, " \
"%41, " \
"p, 1, %45, %43, %44;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %42, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"%40, " \
"%41, " \
"p, 1, %45, %43, %44;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %22, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19}, " \
"%20, " \
"%21, " \
"p, 1, %25, %23, %24;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,703 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 96, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 96, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %55, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %55, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %31, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 96, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %53, %51, %52;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %53, %51, %52;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %29, %27, %28;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %51;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %51;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %27;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %27;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};

View File

@@ -0,0 +1,47 @@
#pragma once
#include "../../../../../common/common.cuh"
#include "../../../../../types/types.cuh"
namespace kittens {
namespace detail {
namespace wgmma {
// templated wrapper for PTX
template<typename T_D, typename T_AB, int cols, int trans_a, int trans_b, int inv=1>
struct base {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, cols, ducks::rt_layout::row> &dst,
const rt<T_AB, 16, cols, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
);
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, cols, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
);
};
// all the ptx's
#include "64x16.impl"
#include "64x32.impl"
#include "64x48.impl"
#include "64x64.impl"
#include "64x80.impl"
#include "64x96.impl"
#include "64x112.impl"
#include "64x128.impl"
#include "64x144.impl"
#include "64x160.impl"
#include "64x176.impl"
#include "64x192.impl"
#include "64x208.impl"
#include "64x224.impl"
#include "64x240.impl"
#include "64x256.impl"
} // namespace wgmma
} // namespace detail
} // namespace kittens

File diff suppressed because it is too large Load Diff