perf: speed up fp8 gemm tiles and scheduling
- K tile 32->64 (new default): fewer barriers, more MMA per stage; generalize tile_at swizzle and load_operand_tile accordingly - 64x128 small-M CTA for m<=64 (2x at 64x4096x4096) - L2 rasterization for crosswise-A layouts (+6..21%) - micro-bench: NT 4096^3 +35%; linear fwd 1.24-1.76x, bwd 1.71-2.27x vs bf16 - add csrc/tests/fp8_test.cu (single MMA demo + GEMM layouts x K-tiles vs CPU reference)
This commit is contained in:
@@ -1,146 +0,0 @@
|
||||
/*
|
||||
Single-kernel BF16 -> FP8 MMA -> BF16 demo for Ada (sm_89).
|
||||
|
||||
nvcc -I csrc -arch=sm_89 -std=c++17 -O3 --use_fast_math \
|
||||
--ptxas-options=-O3,-v csrc/tests/fp8_mma_test.cu -o fp8_mma_test \
|
||||
&& ./fp8_mma_test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
|
||||
#include <cuda_fp8.h>
|
||||
|
||||
#include "../kernels/common/mma.cuh"
|
||||
|
||||
#include <algorithm>
|
||||
#include <vector>
|
||||
|
||||
constexpr int M = 16;
|
||||
constexpr int N = 8;
|
||||
constexpr int K = 32;
|
||||
|
||||
__device__ __forceinline__ unsigned pack_fp8x4(float x0, float x1, float x2,
|
||||
float x3) {
|
||||
__nv_fp8_e4m3 q0(x0);
|
||||
__nv_fp8_e4m3 q1(x1);
|
||||
__nv_fp8_e4m3 q2(x2);
|
||||
__nv_fp8_e4m3 q3(x3);
|
||||
return static_cast<unsigned>(q0.__x) |
|
||||
(static_cast<unsigned>(q1.__x) << 8) |
|
||||
(static_cast<unsigned>(q2.__x) << 16) |
|
||||
(static_cast<unsigned>(q3.__x) << 24);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ unsigned load_quantize_fp8x4(
|
||||
const bf16* src, float scale_inv) {
|
||||
return pack_fp8x4(__bfloat162float(src[0]) * scale_inv,
|
||||
__bfloat162float(src[1]) * scale_inv,
|
||||
__bfloat162float(src[2]) * scale_inv,
|
||||
__bfloat162float(src[3]) * scale_inv);
|
||||
}
|
||||
|
||||
__global__ void fused_bf16_fp8_mma_kernel(
|
||||
const bf16* __restrict__ a, const bf16* __restrict__ b,
|
||||
bf16* __restrict__ out, float scale_a, float scale_b) {
|
||||
const int lane = threadIdx.x;
|
||||
const int group = lane >> 2;
|
||||
const int thread_in_group = lane & 3;
|
||||
const int k0 = thread_in_group * 4;
|
||||
|
||||
// PTX m16n8k32 A fragment: two rows, two 16-column K partitions.
|
||||
unsigned a_frag[4];
|
||||
a_frag[0] = load_quantize_fp8x4(&a[group * K + k0], 1.0f / scale_a);
|
||||
a_frag[1] = load_quantize_fp8x4(&a[(group + 8) * K + k0], 1.0f / scale_a);
|
||||
a_frag[2] = load_quantize_fp8x4(&a[group * K + k0 + 16], 1.0f / scale_a);
|
||||
a_frag[3] = load_quantize_fp8x4(&a[(group + 8) * K + k0 + 16],
|
||||
1.0f / scale_a);
|
||||
|
||||
// B is supplied as row-major [N,K], equivalent to the col-major [K,N]
|
||||
// operand required by the MMA instruction.
|
||||
unsigned b_frag[2];
|
||||
b_frag[0] = load_quantize_fp8x4(&b[group * K + k0], 1.0f / scale_b);
|
||||
b_frag[1] = load_quantize_fp8x4(&b[group * K + k0 + 16], 1.0f / scale_b);
|
||||
|
||||
float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
|
||||
astrai::mma_sync<__nv_fp8_e4m3>(acc, a_frag, b_frag, acc);
|
||||
|
||||
const int col = thread_in_group * 2;
|
||||
const float output_scale = scale_a * scale_b;
|
||||
*reinterpret_cast<__nv_bfloat162*>(&out[group * N + col]) =
|
||||
__floats2bfloat162_rn(acc[0] * output_scale,
|
||||
acc[1] * output_scale);
|
||||
*reinterpret_cast<__nv_bfloat162*>(&out[(group + 8) * N + col]) =
|
||||
__floats2bfloat162_rn(acc[2] * output_scale,
|
||||
acc[3] * output_scale);
|
||||
}
|
||||
|
||||
static float quantize_e4m3(float value) {
|
||||
return static_cast<float>(__nv_fp8_e4m3(value));
|
||||
}
|
||||
|
||||
int main() {
|
||||
srand(0);
|
||||
std::vector<float> a(M * K), b(N * K), reference(M * N, 0.0f);
|
||||
std::vector<bf16> a_bf16(M * K), b_bf16(N * K), output(M * N);
|
||||
for (float& value : a) value = randf() * 4.0f;
|
||||
for (float& value : b) value = randf() * 4.0f;
|
||||
for (int i = 0; i < M * K; ++i) {
|
||||
a_bf16[i] = f2bf(a[i]);
|
||||
a[i] = bf2f(a_bf16[i]);
|
||||
}
|
||||
for (int i = 0; i < N * K; ++i) {
|
||||
b_bf16[i] = f2bf(b[i]);
|
||||
b[i] = bf2f(b_bf16[i]);
|
||||
}
|
||||
|
||||
const float amax = *std::max_element(
|
||||
a.begin(), a.end(), [](float x, float y) { return fabsf(x) < fabsf(y); });
|
||||
const float bmax = *std::max_element(
|
||||
b.begin(), b.end(), [](float x, float y) { return fabsf(x) < fabsf(y); });
|
||||
const float scale_a = fabsf(amax) / 448.0f;
|
||||
const float scale_b = fabsf(bmax) / 448.0f;
|
||||
|
||||
for (int row = 0; row < M; ++row) {
|
||||
for (int col = 0; col < N; ++col) {
|
||||
float sum = 0.0f;
|
||||
for (int k = 0; k < K; ++k) {
|
||||
float qa = quantize_e4m3(a[row * K + k] / scale_a);
|
||||
float qb = quantize_e4m3(b[col * K + k] / scale_b);
|
||||
sum = fmaf(qa, qb, sum);
|
||||
}
|
||||
reference[row * N + col] = sum * scale_a * scale_b;
|
||||
}
|
||||
}
|
||||
|
||||
bf16 *d_a, *d_b, *d_out;
|
||||
CUDA_CHECK(cudaMalloc(&d_a, a_bf16.size() * sizeof(bf16)));
|
||||
CUDA_CHECK(cudaMalloc(&d_b, b_bf16.size() * sizeof(bf16)));
|
||||
CUDA_CHECK(cudaMalloc(&d_out, output.size() * sizeof(bf16)));
|
||||
CUDA_CHECK(cudaMemcpy(d_a, a_bf16.data(), a_bf16.size() * sizeof(bf16),
|
||||
cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(cudaMemcpy(d_b, b_bf16.data(), b_bf16.size() * sizeof(bf16),
|
||||
cudaMemcpyHostToDevice));
|
||||
|
||||
fused_bf16_fp8_mma_kernel<<<1, 32>>>(d_a, d_b, d_out, scale_a, scale_b);
|
||||
CUDA_CHECK(cudaDeviceSynchronize());
|
||||
CUDA_CHECK(cudaMemcpy(output.data(), d_out, output.size() * sizeof(bf16),
|
||||
cudaMemcpyDeviceToHost));
|
||||
|
||||
float max_abs_error = 0.0f;
|
||||
float max_rel_error = 0.0f;
|
||||
for (int i = 0; i < M * N; ++i) {
|
||||
float error = fabsf(bf2f(output[i]) - reference[i]);
|
||||
max_abs_error = fmaxf(max_abs_error, error);
|
||||
max_rel_error = fmaxf(max_rel_error,
|
||||
error / fmaxf(fabsf(reference[i]), 1e-4f));
|
||||
}
|
||||
const bool pass = max_abs_error < 0.05f;
|
||||
print_test_header();
|
||||
print_test_row("M=16 N=8 K=32 fused BF16->E4M3 MMA", max_abs_error,
|
||||
max_rel_error, pass);
|
||||
|
||||
cudaFree(d_a);
|
||||
cudaFree(d_b);
|
||||
cudaFree(d_out);
|
||||
return pass ? 0 : 1;
|
||||
}
|
||||
@@ -0,0 +1,312 @@
|
||||
/*
|
||||
FP8 family tests: single-warp MMA demo + full GEMM correctness.
|
||||
|
||||
Part 1 exercises one bf16 -> fp8 -> mma.sync m16n8k32 instruction pair
|
||||
(sanity for astrai::mma_sync + the fragment layout contract).
|
||||
Part 2 checks launch_fp8_gemm across all four operand layouts, both K
|
||||
tiles, and ragged shapes against an fp32 CPU reference.
|
||||
|
||||
nvcc -I csrc -arch=sm_89 -std=c++17 -O3 csrc/tests/fp8_test.cu -o /tmp/fp8_test \
|
||||
&& /tmp/fp8_test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
|
||||
#include <cuda_fp8.h>
|
||||
#include <algorithm>
|
||||
#include <cmath>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cuda_runtime.h>
|
||||
#include <type_traits>
|
||||
#include <vector>
|
||||
|
||||
#include "../kernels/common/mma.cuh"
|
||||
#include "../kernels/fp8/gemm.cuh"
|
||||
|
||||
using namespace astrai::fp8;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Part 1: single-kernel BF16 -> FP8 MMA -> BF16 demo (m16n8k32)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
namespace {
|
||||
|
||||
constexpr int kMmaM = 16;
|
||||
constexpr int kMmaN = 8;
|
||||
constexpr int kMmaK = 32;
|
||||
|
||||
__device__ __forceinline__ unsigned pack_fp8x4(float x0, float x1, float x2,
|
||||
float x3) {
|
||||
__nv_fp8_e4m3 q0(x0);
|
||||
__nv_fp8_e4m3 q1(x1);
|
||||
__nv_fp8_e4m3 q2(x2);
|
||||
__nv_fp8_e4m3 q3(x3);
|
||||
return static_cast<unsigned>(q0.__x) |
|
||||
(static_cast<unsigned>(q1.__x) << 8) |
|
||||
(static_cast<unsigned>(q2.__x) << 16) |
|
||||
(static_cast<unsigned>(q3.__x) << 24);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ unsigned load_quantize_fp8x4(
|
||||
const bf16* src, float scale_inv) {
|
||||
return pack_fp8x4(__bfloat162float(src[0]) * scale_inv,
|
||||
__bfloat162float(src[1]) * scale_inv,
|
||||
__bfloat162float(src[2]) * scale_inv,
|
||||
__bfloat162float(src[3]) * scale_inv);
|
||||
}
|
||||
|
||||
__global__ void fused_bf16_fp8_mma_kernel(
|
||||
const bf16* __restrict__ a, const bf16* __restrict__ b,
|
||||
bf16* __restrict__ out, float scale_a, float scale_b) {
|
||||
const int lane = threadIdx.x;
|
||||
const int group = lane >> 2;
|
||||
const int thread_in_group = lane & 3;
|
||||
const int k0 = thread_in_group * 4;
|
||||
|
||||
// PTX m16n8k32 A fragment: two rows, two 16-column K partitions.
|
||||
unsigned a_frag[4];
|
||||
a_frag[0] = load_quantize_fp8x4(&a[group * kMmaK + k0], 1.0f / scale_a);
|
||||
a_frag[1] =
|
||||
load_quantize_fp8x4(&a[(group + 8) * kMmaK + k0], 1.0f / scale_a);
|
||||
a_frag[2] =
|
||||
load_quantize_fp8x4(&a[group * kMmaK + k0 + 16], 1.0f / scale_a);
|
||||
a_frag[3] = load_quantize_fp8x4(&a[(group + 8) * kMmaK + k0 + 16],
|
||||
1.0f / scale_a);
|
||||
|
||||
// B is supplied as row-major [N,K], equivalent to the col-major [K,N]
|
||||
// operand required by the MMA instruction.
|
||||
unsigned b_frag[2];
|
||||
b_frag[0] = load_quantize_fp8x4(&b[group * kMmaK + k0], 1.0f / scale_b);
|
||||
b_frag[1] =
|
||||
load_quantize_fp8x4(&b[group * kMmaK + k0 + 16], 1.0f / scale_b);
|
||||
|
||||
float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
|
||||
astrai::mma_sync<__nv_fp8_e4m3>(acc, a_frag, b_frag, acc);
|
||||
|
||||
const int col = thread_in_group * 2;
|
||||
const float output_scale = scale_a * scale_b;
|
||||
*reinterpret_cast<__nv_bfloat162*>(&out[group * kMmaN + col]) =
|
||||
__floats2bfloat162_rn(acc[0] * output_scale, acc[1] * output_scale);
|
||||
*reinterpret_cast<__nv_bfloat162*>(&out[(group + 8) * kMmaN + col]) =
|
||||
__floats2bfloat162_rn(acc[2] * output_scale, acc[3] * output_scale);
|
||||
}
|
||||
|
||||
static float quantize_e4m3(float value) {
|
||||
return static_cast<float>(__nv_fp8_e4m3(value));
|
||||
}
|
||||
|
||||
static bool test_single_mma() {
|
||||
srand(0);
|
||||
std::vector<float> a(kMmaM * kMmaK), b(kMmaN * kMmaK),
|
||||
reference(kMmaM * kMmaN, 0.0f);
|
||||
std::vector<bf16> a_bf16(kMmaM * kMmaK), b_bf16(kMmaN * kMmaK),
|
||||
output(kMmaM * kMmaN);
|
||||
for (float& value : a) value = randf() * 4.0f;
|
||||
for (float& value : b) value = randf() * 4.0f;
|
||||
for (int i = 0; i < kMmaM * kMmaK; ++i) {
|
||||
a_bf16[i] = f2bf(a[i]);
|
||||
a[i] = bf2f(a_bf16[i]);
|
||||
}
|
||||
for (int i = 0; i < kMmaN * kMmaK; ++i) {
|
||||
b_bf16[i] = f2bf(b[i]);
|
||||
b[i] = bf2f(b_bf16[i]);
|
||||
}
|
||||
|
||||
const float amax = *std::max_element(
|
||||
a.begin(), a.end(),
|
||||
[](float x, float y) { return fabsf(x) < fabsf(y); });
|
||||
const float bmax = *std::max_element(
|
||||
b.begin(), b.end(),
|
||||
[](float x, float y) { return fabsf(x) < fabsf(y); });
|
||||
const float scale_a = fabsf(amax) / 448.0f;
|
||||
const float scale_b = fabsf(bmax) / 448.0f;
|
||||
|
||||
for (int row = 0; row < kMmaM; ++row) {
|
||||
for (int col = 0; col < kMmaN; ++col) {
|
||||
float sum = 0.0f;
|
||||
for (int k = 0; k < kMmaK; ++k) {
|
||||
float qa = quantize_e4m3(a[row * kMmaK + k] / scale_a);
|
||||
float qb = quantize_e4m3(b[col * kMmaK + k] / scale_b);
|
||||
sum = fmaf(qa, qb, sum);
|
||||
}
|
||||
reference[row * kMmaN + col] = sum * scale_a * scale_b;
|
||||
}
|
||||
}
|
||||
|
||||
bf16 *d_a, *d_b, *d_out;
|
||||
CUDA_CHECK(cudaMalloc(&d_a, a_bf16.size() * sizeof(bf16)));
|
||||
CUDA_CHECK(cudaMalloc(&d_b, b_bf16.size() * sizeof(bf16)));
|
||||
CUDA_CHECK(cudaMalloc(&d_out, output.size() * sizeof(bf16)));
|
||||
CUDA_CHECK(cudaMemcpy(d_a, a_bf16.data(), a_bf16.size() * sizeof(bf16),
|
||||
cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(cudaMemcpy(d_b, b_bf16.data(), b_bf16.size() * sizeof(bf16),
|
||||
cudaMemcpyHostToDevice));
|
||||
|
||||
fused_bf16_fp8_mma_kernel<<<1, 32>>>(d_a, d_b, d_out, scale_a, scale_b);
|
||||
CUDA_CHECK(cudaDeviceSynchronize());
|
||||
CUDA_CHECK(cudaMemcpy(output.data(), d_out, output.size() * sizeof(bf16),
|
||||
cudaMemcpyDeviceToHost));
|
||||
|
||||
float max_abs_error = 0.0f;
|
||||
float max_rel_error = 0.0f;
|
||||
for (int i = 0; i < kMmaM * kMmaN; ++i) {
|
||||
float error = fabsf(bf2f(output[i]) - reference[i]);
|
||||
max_abs_error = fmaxf(max_abs_error, error);
|
||||
max_rel_error = fmaxf(
|
||||
max_rel_error, error / fmaxf(fabsf(reference[i]), 1e-4f));
|
||||
}
|
||||
const bool pass = max_abs_error < 0.05f;
|
||||
print_test_row("M=16 N=8 K=32 fused BF16->E4M3 MMA", max_abs_error,
|
||||
max_rel_error, pass);
|
||||
|
||||
cudaFree(d_a);
|
||||
cudaFree(d_b);
|
||||
cudaFree(d_out);
|
||||
return pass;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Part 2: GEMM correctness — layouts x K-tiles vs fp32 CPU reference
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <typename LA, typename LB, int kK, int Stages>
|
||||
static bool run_gemm_case(const float* ha, const float* hb, int m, int n,
|
||||
int k, int a_ld, int b_ld) {
|
||||
__nv_fp8_e4m3 *da, *db;
|
||||
__nv_bfloat16* dout;
|
||||
float *dsa, *dsb;
|
||||
cudaMalloc(&da, (size_t)m * k);
|
||||
cudaMalloc(&db, (size_t)n * k);
|
||||
cudaMalloc(&dout, (size_t)m * n * 2);
|
||||
cudaMalloc(&dsa, 4);
|
||||
cudaMalloc(&dsb, 4);
|
||||
float one = 1.0f;
|
||||
cudaMemcpy(dsa, &one, 4, cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dsb, &one, 4, cudaMemcpyHostToDevice);
|
||||
// quantize inputs to e4m3 on host and upload byte-by-byte
|
||||
std::vector<unsigned char> qa(m * k), qb(n * k);
|
||||
for (int i = 0; i < m * k; ++i) {
|
||||
__nv_fp8_e4m3 q(ha[i]);
|
||||
qa[i] = *(unsigned char*)&q;
|
||||
}
|
||||
for (int i = 0; i < n * k; ++i) {
|
||||
__nv_fp8_e4m3 q(hb[i]);
|
||||
qb[i] = *(unsigned char*)&q;
|
||||
}
|
||||
cudaMemcpy(da, qa.data(), qa.size(), cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(db, qb.data(), qb.size(), cudaMemcpyHostToDevice);
|
||||
|
||||
FP8Params p = {};
|
||||
p.a_ptr = da;
|
||||
p.b_ptr = db;
|
||||
p.out_ptr = dout;
|
||||
p.scale_a = dsa;
|
||||
p.scale_b = dsb;
|
||||
p.m = m;
|
||||
p.n = n;
|
||||
p.k = k;
|
||||
p.a_ld = a_ld;
|
||||
p.b_ld = b_ld;
|
||||
launch_fp8_gemm<FP8Format::E4M3, false, LA, LB, kK, Stages>(p, 0);
|
||||
cudaError_t e = cudaDeviceSynchronize();
|
||||
if (e != cudaSuccess) {
|
||||
printf(" CUDA err: %s\n", cudaGetErrorString(e));
|
||||
return false;
|
||||
}
|
||||
std::vector<unsigned short> hb16(m * n);
|
||||
cudaMemcpy(hb16.data(), dout, (size_t)m * n * 2, cudaMemcpyDeviceToHost);
|
||||
const float tol = 0.06f;
|
||||
double max_rel = 0;
|
||||
bool ok = true;
|
||||
for (int i = 0; i < m && ok; ++i) {
|
||||
for (int j = 0; j < n && ok; ++j) {
|
||||
float ref = 0;
|
||||
for (int kk = 0; kk < k; ++kk) {
|
||||
// A reference reads the actual uploaded buffer: LA ColMajor
|
||||
// means the buffer is [K][M] (ha_t), else [M][K].
|
||||
float av = std::is_same_v<LA, ColMajor>
|
||||
? (float)__nv_fp8_e4m3(ha[kk * m + i])
|
||||
: (float)__nv_fp8_e4m3(ha[i * k + kk]);
|
||||
float bv;
|
||||
if (std::is_same_v<LB, ColMajor>)
|
||||
bv = (float)__nv_fp8_e4m3(hb[j * k + kk]);
|
||||
else
|
||||
bv = (float)__nv_fp8_e4m3(hb[kk * n + j]);
|
||||
ref += av * bv;
|
||||
}
|
||||
float got =
|
||||
__bfloat162float(__ushort_as_bfloat16(hb16[i * n + j]));
|
||||
float err = fabsf(got - ref);
|
||||
float rel = err / fmaxf(fabsf(ref), 0.5f);
|
||||
if (rel > max_rel) max_rel = rel;
|
||||
if (err > tol * fmaxf(fabsf(ref), 1.0f)) ok = false;
|
||||
}
|
||||
}
|
||||
printf(" max_rel=%.4f %s\n", max_rel, ok ? "PASS" : "FAIL");
|
||||
cudaFree(da);
|
||||
cudaFree(db);
|
||||
cudaFree(dout);
|
||||
cudaFree(dsa);
|
||||
cudaFree(dsb);
|
||||
return ok;
|
||||
}
|
||||
|
||||
static bool test_gemm() {
|
||||
struct {
|
||||
int m, n, k;
|
||||
} cfgs[] = {
|
||||
{128, 128, 128}, {256, 128, 256}, {128, 256, 64},
|
||||
{100, 130, 96}, {64, 64, 160}, {300, 200, 320},
|
||||
};
|
||||
bool all = true;
|
||||
for (auto& c : cfgs) {
|
||||
float* ha = new float[c.m * c.k];
|
||||
float* hb_rowmajor = new float[c.k * c.n]; // [K][N] for B RowMajor
|
||||
float* hb_colmajor = new float[c.n * c.k]; // [N][K] for B ColMajor
|
||||
for (int i = 0; i < c.m * c.k; ++i) ha[i] = randf();
|
||||
for (int i = 0; i < c.k * c.n; ++i) hb_rowmajor[i] = randf();
|
||||
for (int i = 0; i < c.k * c.n; ++i)
|
||||
hb_colmajor[i / c.k * c.k + i % c.k] = hb_rowmajor[i];
|
||||
float* ha_t = new float[c.k * c.m]; // [K][M] for A ColMajor
|
||||
for (int i = 0; i < c.m; ++i)
|
||||
for (int p = 0; p < c.k; ++p) ha_t[p * c.m + i] = ha[i * c.k + p];
|
||||
printf("%dx%dx%d:\n", c.m, c.n, c.k);
|
||||
printf(" NT K32:");
|
||||
all &= run_gemm_case<RowMajor, ColMajor, 32, 3>(ha, hb_colmajor, c.m,
|
||||
c.n, c.k, c.k, c.k);
|
||||
printf(" NT K64:");
|
||||
all &= run_gemm_case<RowMajor, ColMajor, 64, 2>(ha, hb_colmajor, c.m,
|
||||
c.n, c.k, c.k, c.k);
|
||||
printf(" NN K32:");
|
||||
all &= run_gemm_case<RowMajor, RowMajor, 32, 3>(ha, hb_rowmajor, c.m,
|
||||
c.n, c.k, c.k, c.n);
|
||||
printf(" NN K64:");
|
||||
all &= run_gemm_case<RowMajor, RowMajor, 64, 2>(ha, hb_rowmajor, c.m,
|
||||
c.n, c.k, c.k, c.n);
|
||||
printf(" TN K32:");
|
||||
all &= run_gemm_case<ColMajor, ColMajor, 32, 3>(ha_t, hb_colmajor, c.m,
|
||||
c.n, c.k, c.m, c.k);
|
||||
printf(" TN K64:");
|
||||
all &= run_gemm_case<ColMajor, ColMajor, 64, 2>(ha_t, hb_colmajor, c.m,
|
||||
c.n, c.k, c.m, c.k);
|
||||
printf(" TT K64:");
|
||||
all &= run_gemm_case<ColMajor, RowMajor, 64, 2>(ha_t, hb_rowmajor, c.m,
|
||||
c.n, c.k, c.m, c.n);
|
||||
delete[] ha;
|
||||
delete[] hb_rowmajor;
|
||||
delete[] hb_colmajor;
|
||||
delete[] ha_t;
|
||||
}
|
||||
return all;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
int main() {
|
||||
print_test_header();
|
||||
bool ok = test_single_mma();
|
||||
ok &= test_gemm();
|
||||
printf(ok ? "All PASS\n" : "FAILURES\n");
|
||||
return ok ? 0 : 1;
|
||||
}
|
||||
Reference in New Issue
Block a user