forked from IQ.Lvbs/IQ.Pilot
IQ.Pilot Release Commit @ 0798119
This commit is contained in:
@@ -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
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -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
Reference in New Issue
Block a user