upstream: add GEMM kernel references from 4 repos for BI-V100 porting

Sources (all CUDA 10.2 compatible, no CUTLASS/Triton dependency):
- leimao/CUDA-GEMM-Optimization: v00-v07, fp16 WMMA variant, double buffered
- siboehm/SGEMM_CUDA: kernel 1-12, warp tiling + double buffering
- wangzyon/NVIDIA_SGEMM_PRACTICE: kernel 1-7
- edtallison/sgemm-cuda: kernel 1-12 (reimplementation with notes)

Key porting issue: ALL kernels hardcode WARPSIZE=32.
BI-V100 has warp_size=64. Need to:
1. Replace all 32U / WARPSIZE constants with 64
2. Adjust warp subtile decomposition (WMITER, WNITER, WSUBM, WSUBN)
3. Adjust shared memory bank conflict avoidance (may have different bank count)
4. Test __shfl_down_sync with mask=0xFFFFFFFFFFFFFFFF (64-bit)
This commit is contained in:
Claude
2026-08-14 15:11:57 +00:00
parent 29ecc2e602
commit 9ca33cf4d5
59 changed files with 8802 additions and 0 deletions

View File

@@ -0,0 +1,65 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v00.
// Non-coalesced read and write from global memory.
template <typename T>
__global__ void gemm_v00(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
// Compute the row and column of C that this thread is responsible for.
size_t const C_row_idx{blockIdx.x * blockDim.x + threadIdx.x};
size_t const C_col_idx{blockIdx.y * blockDim.y + threadIdx.y};
// Each thread compute
// C[C_row_idx, C_col_idx] = alpha * A[C_row_idx, :] * B[:, C_col_idx] +
// beta * C[C_row_idx, C_col_idx].
if (C_row_idx < m && C_col_idx < n)
{
T sum{static_cast<T>(0)};
for (size_t k_idx{0U}; k_idx < k; ++k_idx)
{
sum += A[C_row_idx * lda + k_idx] * B[k_idx * ldb + C_col_idx];
}
C[C_row_idx * ldc + C_col_idx] =
alpha * sum + beta * C[C_row_idx * ldc + C_col_idx];
}
}
template <typename T>
void launch_gemm_kernel_v00(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
dim3 const block_dim{32U, 32U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(m) + block_dim.x - 1U) / block_dim.x,
(static_cast<unsigned int>(n) + block_dim.y - 1U) / block_dim.y, 1U};
gemm_v00<T><<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B,
ldb, *beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v00<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v00<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v00<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,65 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v01.
// Coalesced read and write from global memory.
template <typename T>
__global__ void gemm_v01(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
// Compute the row and column of C that this thread is responsible for.
size_t const C_col_idx{blockIdx.x * blockDim.x + threadIdx.x};
size_t const C_row_idx{blockIdx.y * blockDim.y + threadIdx.y};
// Each thread compute
// C[C_row_idx, C_col_idx] = alpha * A[C_row_idx, :] * B[:, C_col_idx] +
// beta * C[C_row_idx, C_col_idx].
if (C_row_idx < m && C_col_idx < n)
{
T sum{static_cast<T>(0)};
for (size_t k_idx{0U}; k_idx < k; ++k_idx)
{
sum += A[C_row_idx * lda + k_idx] * B[k_idx * ldb + C_col_idx];
}
C[C_row_idx * ldc + C_col_idx] =
alpha * sum + beta * C[C_row_idx * ldc + C_col_idx];
}
}
template <typename T>
void launch_gemm_kernel_v01(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
dim3 const block_dim{32U, 32U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + block_dim.x - 1U) / block_dim.x,
(static_cast<unsigned int>(m) + block_dim.y - 1U) / block_dim.y, 1U};
gemm_v01<T><<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B,
ldb, *beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v01<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v01<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v01<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,112 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v02.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K>
__global__ void gemm_v02(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Compute the row and column of C that this thread is responsible for.
size_t const C_col_idx{blockIdx.x * blockDim.x + threadIdx.x};
size_t const C_row_idx{blockIdx.y * blockDim.y + threadIdx.y};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile[BLOCK_TILE_SIZE_Y][BLOCK_TILE_SIZE_K];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
T sum{static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
// Doing this results in 2 TOPS.
// Suppose blockDim.x = blockDim.y = 32.
// Effectively, for a warp, in one iteration, we read the value from
// A_thread_block_tile at the same location on the shared memory
// resulting in a broadcast, we also read 32 values that have no
// bank conflicts from B_thread_block_tile. Even with that, all the
// values have to be read from the shared memory and consequence is
// the shared memory instruction runs very intensively just to
// compute a small number of values using simple arithmetic
// instructions, which is not efficient.
sum += A_thread_block_tile[threadIdx.y][k_i] *
B_thread_block_tile[k_i][threadIdx.x];
}
__syncthreads();
}
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * sum + beta * C[C_row_idx * ldc + C_col_idx];
}
}
template <typename T>
void launch_gemm_kernel_v02(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{32U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{32U};
constexpr unsigned int BLOCK_TILE_SIZE_K{32U};
constexpr unsigned int NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS == 0U);
static_assert(BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS == 0U);
dim3 const block_dim{BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + block_dim.x - 1U) / block_dim.x,
(static_cast<unsigned int>(m) + block_dim.y - 1U) / block_dim.y, 1U};
gemm_v02<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v02<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v02<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v02<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,108 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v02.
// Coalesced read and write from global memory.
// We guarantee that matrix A, B, and C are 32 byte aligned.
// This implementation is slower because we waste a lot of threads.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K>
__global__ void gemm_v02_vectorized(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Compute the row and column of C that this thread is responsible for.
size_t const C_col_idx{blockIdx.x * blockDim.x + threadIdx.x};
size_t const C_row_idx{blockIdx.y * blockDim.y + threadIdx.y};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile[BLOCK_TILE_SIZE_Y][BLOCK_TILE_SIZE_K];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
T sum{static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
// Doing this results in 2 TOPS.
// Suppose blockDim.x = blockDim.y = 32.
// Effectively, for a warp, in one iteration, we read the value from
// A_thread_block_tile at the same location on the shared memory
// resulting in a broadcast, we also read 32 values that have no
// bank conflicts from B_thread_block_tile. Even with that, all the
// values have to be read from the shared memory and consequence is
// the shared memory instruction runs very intensively just to
// compute a small number of values using simple arithmetic
// instructions, which is not efficient.
sum += A_thread_block_tile[threadIdx.y][k_i] *
B_thread_block_tile[k_i][threadIdx.x];
}
__syncthreads();
}
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * sum + beta * C[C_row_idx * ldc + C_col_idx];
}
}
template <typename T>
void launch_gemm_kernel_v02_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{32U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{32U};
constexpr unsigned int BLOCK_TILE_SIZE_K{32U};
constexpr unsigned int NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS == 0U);
static_assert(BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS == 0U);
dim3 const block_dim{BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + block_dim.x - 1U) / block_dim.x,
(static_cast<unsigned int>(m) + block_dim.y - 1U) / block_dim.y, 1U};
gemm_v02_vectorized<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y,
BLOCK_TILE_SIZE_K><<<grid_dim, block_dim, 0U, stream>>>(
m, n, k, *alpha, A, lda, B, ldb, *beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v02_vectorized<float>(
size_t m, size_t n, size_t k, float const* alpha, float const* A,
size_t lda, float const* B, size_t ldb, float const* beta, float* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v02_vectorized<double>(
size_t m, size_t n, size_t k, double const* alpha, double const* A,
size_t lda, double const* B, size_t ldb, double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v02_vectorized<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,144 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v03.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t THREAD_TILE_SIZE_Y>
__global__ void gemm_v03(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
THREAD_TILE_SIZE_Y};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile[BLOCK_TILE_SIZE_Y][BLOCK_TILE_SIZE_K];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes BLOCK_TILE_SIZE_Y output values.
// Specifically, these values corresponds to
// C[blockIdx.y * BLOCK_TILE_SIZE_Y + threadIdx.x / BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_Y : blockIdx.y * BLOCK_TILE_SIZE_Y + (threadIdx.x /
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_Y][blockIdx.x *
// BLOCK_TILE_SIZE_X + threadIdx.x % BLOCK_TILE_SIZE_X]
T C_thread_results[THREAD_TILE_SIZE_Y] = {static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
size_t const B_thread_block_tile_row_idx{k_i};
// B_val is cached in the register to alleviate the pressure on the
// shared memory access.
T const B_val{
B_thread_block_tile[B_thread_block_tile_row_idx]
[thread_linear_idx % BLOCK_TILE_SIZE_X]};
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
size_t const A_thread_block_tile_row_idx{
thread_linear_idx / BLOCK_TILE_SIZE_X * THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const A_thread_block_tile_col_idx{k_i};
T const A_val{A_thread_block_tile[A_thread_block_tile_row_idx]
[A_thread_block_tile_col_idx]};
C_thread_results[thread_tile_row_idx] += A_val * B_val;
}
}
__syncthreads();
}
// Write the results to DRAM.
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y; ++thread_tile_row_idx)
{
size_t const C_row_idx{blockIdx.y * BLOCK_TILE_SIZE_Y +
thread_linear_idx / BLOCK_TILE_SIZE_X *
THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const C_col_idx{blockIdx.x * BLOCK_TILE_SIZE_X +
thread_linear_idx % BLOCK_TILE_SIZE_X};
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * C_thread_results[thread_tile_row_idx] +
beta * C[C_row_idx * ldc + C_col_idx];
}
}
}
template <typename T>
void launch_gemm_kernel_v03(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{64U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{64U};
constexpr unsigned int BLOCK_TILE_SIZE_K{8U};
// Each thread computes THREAD_TILE_SIZE_Y values of C.
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y / THREAD_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_Y % THREAD_TILE_SIZE_Y == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_K == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_X == 0U);
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v03<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
THREAD_TILE_SIZE_Y><<<grid_dim, block_dim, 0U, stream>>>(
m, n, k, *alpha, A, lda, B, ldb, *beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v03<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v03<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v03<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,141 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v03.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t THREAD_TILE_SIZE_Y>
__global__ void gemm_v03_vectorized(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
THREAD_TILE_SIZE_Y};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile[BLOCK_TILE_SIZE_Y][BLOCK_TILE_SIZE_K];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes BLOCK_TILE_SIZE_Y output values.
// Specifically, these values corresponds to
// C[blockIdx.y * BLOCK_TILE_SIZE_Y + threadIdx.x / BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_Y : blockIdx.y * BLOCK_TILE_SIZE_Y + (threadIdx.x /
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_Y][blockIdx.x *
// BLOCK_TILE_SIZE_X + threadIdx.x % BLOCK_TILE_SIZE_X]
T C_thread_results[THREAD_TILE_SIZE_Y] = {static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
size_t const B_thread_block_tile_row_idx{k_i};
// B_val is cached in the register to alleviate the pressure on the
// shared memory access.
T const B_val{
B_thread_block_tile[B_thread_block_tile_row_idx]
[thread_linear_idx % BLOCK_TILE_SIZE_X]};
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
size_t const A_thread_block_tile_row_idx{
thread_linear_idx / BLOCK_TILE_SIZE_X * THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const A_thread_block_tile_col_idx{k_i};
T const A_val{A_thread_block_tile[A_thread_block_tile_row_idx]
[A_thread_block_tile_col_idx]};
C_thread_results[thread_tile_row_idx] += A_val * B_val;
}
}
__syncthreads();
}
// Write the results to DRAM.
// Cannot vectorized the write to DRAM because we are writting to a column
// instead of a row in C.
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y; ++thread_tile_row_idx)
{
size_t const C_row_idx{blockIdx.y * BLOCK_TILE_SIZE_Y +
thread_linear_idx / BLOCK_TILE_SIZE_X *
THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const C_col_idx{blockIdx.x * BLOCK_TILE_SIZE_X +
thread_linear_idx % BLOCK_TILE_SIZE_X};
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * C_thread_results[thread_tile_row_idx] +
beta * C[C_row_idx * ldc + C_col_idx];
}
}
}
template <typename T>
void launch_gemm_kernel_v03_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{64U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{64U};
constexpr unsigned int BLOCK_TILE_SIZE_K{8U};
// Each thread computes THREAD_TILE_SIZE_Y values of C.
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y / THREAD_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_Y % THREAD_TILE_SIZE_Y == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_K == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_X == 0U);
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v03_vectorized<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y,
BLOCK_TILE_SIZE_K, THREAD_TILE_SIZE_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v03_vectorized<float>(
size_t m, size_t n, size_t k, float const* alpha, float const* A,
size_t lda, float const* B, size_t ldb, float const* beta, float* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v03_vectorized<double>(
size_t m, size_t n, size_t k, double const* alpha, double const* A,
size_t lda, double const* B, size_t ldb, double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v03_vectorized<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,199 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v04.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__global__ void gemm_v04(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile[BLOCK_TILE_SIZE_Y][BLOCK_TILE_SIZE_K];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes BLOCK_TILE_SIZE_Y output values.
// Specifically, these values corresponds to
// C[blockIdx.y * BLOCK_TILE_SIZE_Y + threadIdx.x / BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_Y : blockIdx.y * BLOCK_TILE_SIZE_Y + (threadIdx.x /
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_Y][blockIdx.x *
// BLOCK_TILE_SIZE_X + threadIdx.x % BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_X : blockIdx.x * BLOCK_TILE_SIZE_X + (threadIdx.x %
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_X]
T C_thread_results[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
// A_vals is cached in the register.
T A_vals[THREAD_TILE_SIZE_Y] = {static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[THREAD_TILE_SIZE_X] = {static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
size_t const A_thread_block_tile_row_idx{
thread_linear_idx / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y};
size_t const A_thread_block_tile_col_idx{k_i};
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
// There will be shared memory bank conflicts accessing the
// values from A_thread_block_tile. We can do it better by
// transposing the A_thread_block_tile when we load the data
// from DRAM.
A_vals[thread_tile_row_idx] =
A_thread_block_tile[A_thread_block_tile_row_idx +
thread_tile_row_idx]
[A_thread_block_tile_col_idx];
}
size_t const B_thread_block_tile_row_idx{k_i};
size_t const B_thread_block_tile_col_idx{
thread_linear_idx % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X};
#pragma unroll
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X;
++thread_tile_col_idx)
{
B_vals[thread_tile_col_idx] =
B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx +
thread_tile_col_idx];
}
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X;
++thread_tile_col_idx)
{
C_thread_results[thread_tile_row_idx]
[thread_tile_col_idx] +=
A_vals[thread_tile_row_idx] *
B_vals[thread_tile_col_idx];
}
}
}
__syncthreads();
}
// Write the results to DRAM.
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y; ++thread_tile_row_idx)
{
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X; ++thread_tile_col_idx)
{
size_t const C_row_idx{
blockIdx.y * BLOCK_TILE_SIZE_Y +
threadIdx.x / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const C_col_idx{
blockIdx.x * BLOCK_TILE_SIZE_X +
threadIdx.x % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X +
thread_tile_col_idx};
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * C_thread_results[thread_tile_row_idx]
[thread_tile_col_idx] +
beta * C[C_row_idx * ldc + C_col_idx];
}
}
}
}
template <typename T>
void launch_gemm_kernel_v04(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
// Each thread computes THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y values of C.
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
static_assert(BLOCK_TILE_SIZE_X % THREAD_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % THREAD_TILE_SIZE_Y == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_K == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_X == 0U);
static_assert(
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS_PER_BLOCK == 0U);
static_assert(
BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS_PER_BLOCK == 0U);
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v04<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v04<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v04<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v04<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,225 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v04.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__global__ void gemm_v04_vectorized(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile[BLOCK_TILE_SIZE_Y][BLOCK_TILE_SIZE_K];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes BLOCK_TILE_SIZE_Y output values.
// Specifically, these values corresponds to
// C[blockIdx.y * BLOCK_TILE_SIZE_Y + threadIdx.x / BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_Y : blockIdx.y * BLOCK_TILE_SIZE_Y + (threadIdx.x /
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_Y][blockIdx.x *
// BLOCK_TILE_SIZE_X + threadIdx.x % BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_X : blockIdx.x * BLOCK_TILE_SIZE_X + (threadIdx.x %
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_X]
T C_thread_results[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
// A_vals is cached in the register.
T A_vals[THREAD_TILE_SIZE_Y] = {static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[THREAD_TILE_SIZE_X] = {static_cast<T>(0)};
constexpr size_t NUM_VECTOR_UNITS{sizeof(int4) / sizeof(T)};
static_assert(sizeof(int4) % sizeof(T) == 0U);
static_assert(BLOCK_TILE_SIZE_K % NUM_VECTOR_UNITS == 0U);
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_THREAD_TILE_SIZE_X{THREAD_TILE_SIZE_X /
NUM_VECTOR_UNITS};
static_assert(THREAD_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
size_t const A_thread_block_tile_row_idx{
thread_linear_idx / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y};
size_t const A_thread_block_tile_col_idx{k_i};
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
// There will be shared memory bank conflicts accessing the
// values from A_thread_block_tile. We can do it better by
// transposing the A_thread_block_tile when we load the data
// from DRAM.
A_vals[thread_tile_row_idx] =
A_thread_block_tile[A_thread_block_tile_row_idx +
thread_tile_row_idx]
[A_thread_block_tile_col_idx];
}
size_t const B_thread_block_tile_row_idx{k_i};
size_t const B_thread_block_tile_col_idx{
thread_linear_idx % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X};
// Although the read from A_thread_block_tile cannot be vectorized, the read
// from B_thread_block_tile can be vectorized.
#pragma unroll
for (size_t thread_tile_col_vector_idx{0U};
thread_tile_col_vector_idx < VECTORIZED_THREAD_TILE_SIZE_X;
++thread_tile_col_vector_idx)
{
*reinterpret_cast<int4*>(
&B_vals[thread_tile_col_vector_idx * NUM_VECTOR_UNITS]) =
*reinterpret_cast<int4 const*>(
&B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx +
thread_tile_col_vector_idx *
NUM_VECTOR_UNITS]);
}
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X;
++thread_tile_col_idx)
{
C_thread_results[thread_tile_row_idx]
[thread_tile_col_idx] +=
A_vals[thread_tile_row_idx] *
B_vals[thread_tile_col_idx];
}
}
}
__syncthreads();
}
// Vectorized writing the results to DRAM.
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y; ++thread_tile_row_idx)
{
for (size_t thread_tile_col_vector_idx{0U};
thread_tile_col_vector_idx < VECTORIZED_THREAD_TILE_SIZE_X;
++thread_tile_col_vector_idx)
{
size_t const C_row_idx{
blockIdx.y * BLOCK_TILE_SIZE_Y +
thread_linear_idx / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const C_col_idx{
blockIdx.x * BLOCK_TILE_SIZE_X +
thread_linear_idx % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X +
thread_tile_col_vector_idx * NUM_VECTOR_UNITS};
// Vectorized read from C.
int4 C_row_vector_vals{*reinterpret_cast<int4 const*>(
&C[C_row_idx * ldc + C_col_idx])};
// Vectorized read from C_thread_results.
int4 const C_thread_results_row_vector_vals{
*reinterpret_cast<int4 const*>(
&C_thread_results[thread_tile_row_idx]
[thread_tile_col_vector_idx *
NUM_VECTOR_UNITS])};
// Update the values in C_row_vector_vals
for (size_t i{0U}; i < NUM_VECTOR_UNITS; ++i)
{
reinterpret_cast<T*>(&C_row_vector_vals)[i] =
alpha * reinterpret_cast<T const*>(
&C_thread_results_row_vector_vals)[i] +
beta * reinterpret_cast<T const*>(&C_row_vector_vals)[i];
}
// Vectorized write to C.
if (C_row_idx < m && C_col_idx < n)
{
// No need to mask out the out-of-bound invalid elements,
// because the row of C matrix is 32-byte aligned.
*reinterpret_cast<int4*>(&C[C_row_idx * ldc + C_col_idx]) =
C_row_vector_vals;
}
}
}
}
template <typename T>
void launch_gemm_kernel_v04_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
// Each thread computes THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y values of C.
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
static_assert(BLOCK_TILE_SIZE_X % THREAD_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % THREAD_TILE_SIZE_Y == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_K == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_X == 0U);
static_assert(
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS_PER_BLOCK == 0U);
static_assert(
BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS_PER_BLOCK == 0U);
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v04_vectorized<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y,
BLOCK_TILE_SIZE_K, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v04_vectorized<float>(
size_t m, size_t n, size_t k, float const* alpha, float const* A,
size_t lda, float const* B, size_t ldb, float const* beta, float* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v04_vectorized<double>(
size_t m, size_t n, size_t k, double const* alpha, double const* A,
size_t lda, double const* B, size_t ldb, double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v04_vectorized<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,196 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v05.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__global__ void gemm_v05(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T
A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_Y];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes BLOCK_TILE_SIZE_Y output values.
// Specifically, these values corresponds to
// C[blockIdx.y * BLOCK_TILE_SIZE_Y + threadIdx.x / BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_Y : blockIdx.y * BLOCK_TILE_SIZE_Y + (threadIdx.x /
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_Y][blockIdx.x *
// BLOCK_TILE_SIZE_X + threadIdx.x % BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_X : blockIdx.x * BLOCK_TILE_SIZE_X + (threadIdx.x %
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_X]
T C_thread_results[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
// A_vals is cached in the register.
T A_vals[THREAD_TILE_SIZE_Y] = {static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[THREAD_TILE_SIZE_X] = {static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_transposed<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile_transposed,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
size_t const A_thread_block_tile_row_idx{
thread_linear_idx / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y};
size_t const A_thread_block_tile_col_idx{k_i};
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
A_vals[thread_tile_row_idx] =
A_thread_block_tile_transposed[A_thread_block_tile_col_idx]
[A_thread_block_tile_row_idx +
thread_tile_row_idx];
}
size_t const B_thread_block_tile_row_idx{k_i};
size_t const B_thread_block_tile_col_idx{
thread_linear_idx % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X};
#pragma unroll
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X;
++thread_tile_col_idx)
{
B_vals[thread_tile_col_idx] =
B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx +
thread_tile_col_idx];
}
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X;
++thread_tile_col_idx)
{
C_thread_results[thread_tile_row_idx]
[thread_tile_col_idx] +=
A_vals[thread_tile_row_idx] *
B_vals[thread_tile_col_idx];
}
}
}
__syncthreads();
}
// Write the results to DRAM.
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y; ++thread_tile_row_idx)
{
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X; ++thread_tile_col_idx)
{
size_t const C_row_idx{
blockIdx.y * BLOCK_TILE_SIZE_Y +
threadIdx.x / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const C_col_idx{
blockIdx.x * BLOCK_TILE_SIZE_X +
threadIdx.x % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X +
thread_tile_col_idx};
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * C_thread_results[thread_tile_row_idx]
[thread_tile_col_idx] +
beta * C[C_row_idx * ldc + C_col_idx];
}
}
}
}
template <typename T>
void launch_gemm_kernel_v05(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
// Each thread computes THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y values of C.
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
static_assert(BLOCK_TILE_SIZE_X % THREAD_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % THREAD_TILE_SIZE_Y == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_K == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_X == 0U);
static_assert(
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS_PER_BLOCK == 0U);
static_assert(
BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS_PER_BLOCK == 0U);
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v05<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v05<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v05<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v05<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,222 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// GEMM kernel v05.
// Coalesced read and write from global memory.
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__global__ void gemm_v05_vectorized(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T
A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_Y];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes BLOCK_TILE_SIZE_Y output values.
// Specifically, these values corresponds to
// C[blockIdx.y * BLOCK_TILE_SIZE_Y + threadIdx.x / BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_Y : blockIdx.y * BLOCK_TILE_SIZE_Y + (threadIdx.x /
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_Y][blockIdx.x *
// BLOCK_TILE_SIZE_X + threadIdx.x % BLOCK_TILE_SIZE_X *
// THREAD_TILE_SIZE_X : blockIdx.x * BLOCK_TILE_SIZE_X + (threadIdx.x %
// BLOCK_TILE_SIZE_X + 1) * THREAD_TILE_SIZE_X]
T C_thread_results[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
// A_vals is cached in the register.
T A_vals[THREAD_TILE_SIZE_Y] = {static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[THREAD_TILE_SIZE_X] = {static_cast<T>(0)};
constexpr size_t NUM_VECTOR_UNITS{sizeof(int4) / sizeof(T)};
static_assert(sizeof(int4) % sizeof(T) == 0U);
static_assert(BLOCK_TILE_SIZE_K % NUM_VECTOR_UNITS == 0U);
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_THREAD_TILE_SIZE_X{THREAD_TILE_SIZE_X /
NUM_VECTOR_UNITS};
static_assert(THREAD_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile_transposed,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
size_t const A_thread_block_tile_row_idx{
thread_linear_idx / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y};
size_t const A_thread_block_tile_col_idx{k_i};
#pragma unroll
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
A_vals[thread_tile_row_idx] =
A_thread_block_tile_transposed[A_thread_block_tile_col_idx]
[A_thread_block_tile_row_idx +
thread_tile_row_idx];
}
size_t const B_thread_block_tile_row_idx{k_i};
size_t const B_thread_block_tile_col_idx{
thread_linear_idx % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X};
// Although the read from A_thread_block_tile cannot be vectorized, the read
// from B_thread_block_tile can be vectorized.
#pragma unroll
for (size_t thread_tile_col_vector_idx{0U};
thread_tile_col_vector_idx < VECTORIZED_THREAD_TILE_SIZE_X;
++thread_tile_col_vector_idx)
{
*reinterpret_cast<int4*>(
&B_vals[thread_tile_col_vector_idx * NUM_VECTOR_UNITS]) =
*reinterpret_cast<int4 const*>(
&B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx +
thread_tile_col_vector_idx *
NUM_VECTOR_UNITS]);
}
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y;
++thread_tile_row_idx)
{
for (size_t thread_tile_col_idx{0U};
thread_tile_col_idx < THREAD_TILE_SIZE_X;
++thread_tile_col_idx)
{
C_thread_results[thread_tile_row_idx]
[thread_tile_col_idx] +=
A_vals[thread_tile_row_idx] *
B_vals[thread_tile_col_idx];
}
}
}
__syncthreads();
}
// Vectorized writing the results to DRAM.
for (size_t thread_tile_row_idx{0U};
thread_tile_row_idx < THREAD_TILE_SIZE_Y; ++thread_tile_row_idx)
{
for (size_t thread_tile_col_vector_idx{0U};
thread_tile_col_vector_idx < VECTORIZED_THREAD_TILE_SIZE_X;
++thread_tile_col_vector_idx)
{
size_t const C_row_idx{
blockIdx.y * BLOCK_TILE_SIZE_Y +
thread_linear_idx / (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_Y +
thread_tile_row_idx};
size_t const C_col_idx{
blockIdx.x * BLOCK_TILE_SIZE_X +
thread_linear_idx % (BLOCK_TILE_SIZE_X / THREAD_TILE_SIZE_X) *
THREAD_TILE_SIZE_X +
thread_tile_col_vector_idx * NUM_VECTOR_UNITS};
// Vectorized read from C.
int4 C_row_vector_vals{*reinterpret_cast<int4 const*>(
&C[C_row_idx * ldc + C_col_idx])};
// Vectorized read from C_thread_results.
int4 const C_thread_results_row_vector_vals{
*reinterpret_cast<int4 const*>(
&C_thread_results[thread_tile_row_idx]
[thread_tile_col_vector_idx *
NUM_VECTOR_UNITS])};
// Update the values in C_row_vector_vals
for (size_t i{0U}; i < NUM_VECTOR_UNITS; ++i)
{
reinterpret_cast<T*>(&C_row_vector_vals)[i] =
alpha * reinterpret_cast<T const*>(
&C_thread_results_row_vector_vals)[i] +
beta * reinterpret_cast<T const*>(&C_row_vector_vals)[i];
}
// Vectorized write to C.
if (C_row_idx < m && C_col_idx < n)
{
// No need to mask out the out-of-bound invalid elements,
// because the row of C matrix is 32-byte aligned.
*reinterpret_cast<int4*>(&C[C_row_idx * ldc + C_col_idx]) =
C_row_vector_vals;
}
}
}
}
template <typename T>
void launch_gemm_kernel_v05_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
// Each thread computes THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y values of C.
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_Y /
(THREAD_TILE_SIZE_X * THREAD_TILE_SIZE_Y)};
static_assert(BLOCK_TILE_SIZE_X % THREAD_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % THREAD_TILE_SIZE_Y == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_K == 0U);
static_assert(NUM_THREADS_PER_BLOCK % BLOCK_TILE_SIZE_X == 0U);
static_assert(
BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS_PER_BLOCK == 0U);
static_assert(
BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS_PER_BLOCK == 0U);
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v05_vectorized<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y,
BLOCK_TILE_SIZE_K, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v05_vectorized<float>(
size_t m, size_t n, size_t k, float const* alpha, float const* A,
size_t lda, float const* B, size_t ldb, float const* beta, float* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v05_vectorized<double>(
size_t m, size_t n, size_t k, double const* alpha, double const* A,
size_t lda, double const* B, size_t ldb, double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v05_vectorized<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,334 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
template <typename T, size_t BLOCK_TILE_SIZE, size_t WARP_TILE_SIZE,
size_t NUM_THREAD_TILES_PER_WARP, size_t THREAD_TILE_SIZE>
__device__ void load_data_from_shared_memory_to_register_file(
T const thread_block_tile[BLOCK_TILE_SIZE],
T register_values[NUM_THREAD_TILES_PER_WARP][THREAD_TILE_SIZE],
size_t warp_idx, size_t thread_idx)
{
static_assert(BLOCK_TILE_SIZE % THREAD_TILE_SIZE == 0U);
#pragma unroll
for (size_t thread_tile_repeat_idx{0U};
thread_tile_repeat_idx < NUM_THREAD_TILES_PER_WARP;
++thread_tile_repeat_idx)
{
size_t const thread_block_tile_idx{
warp_idx * WARP_TILE_SIZE +
thread_tile_repeat_idx *
(WARP_TILE_SIZE / NUM_THREAD_TILES_PER_WARP) +
thread_idx * THREAD_TILE_SIZE};
#pragma unroll
for (size_t thread_tile_idx{0U}; thread_tile_idx < THREAD_TILE_SIZE;
++thread_tile_idx)
{
register_values[thread_tile_repeat_idx][thread_tile_idx] =
thread_block_tile[thread_block_tile_idx + thread_tile_idx];
}
}
}
template <typename T, size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__device__ void compute_thread_tile_results(
T const A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y],
T const B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X],
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X])
{
// Compute NUM_THREAD_TILES_PER_WARP_Y * NUM_THREAD_TILES_PER_WARP_X outer
// products.
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP_Y;
++thread_tile_repeat_row_idx)
{
#pragma unroll
for (size_t thread_tile_repeat_col_idx{0U};
thread_tile_repeat_col_idx < NUM_THREAD_TILES_PER_WARP_X;
++thread_tile_repeat_col_idx)
{
#pragma unroll
for (size_t thread_tile_y_idx{0U};
thread_tile_y_idx < THREAD_TILE_SIZE_Y; ++thread_tile_y_idx)
{
#pragma unroll
for (size_t thread_tile_x_idx{0U};
thread_tile_x_idx < THREAD_TILE_SIZE_X;
++thread_tile_x_idx)
{
C_thread_results[thread_tile_repeat_row_idx]
[thread_tile_repeat_col_idx]
[thread_tile_y_idx][thread_tile_x_idx] +=
A_vals[thread_tile_repeat_row_idx][thread_tile_y_idx] *
B_vals[thread_tile_repeat_col_idx][thread_tile_x_idx];
}
}
}
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t WARP_TILE_SIZE_X, size_t WARP_TILE_SIZE_Y,
size_t THREAD_TILE_SIZE_X, size_t THREAD_TILE_SIZE_Y,
size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y>
__device__ void write_results_from_register_file_to_global_memory(
T const C_thread_results[NUM_THREAD_TILES_PER_WARP_Y]
[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_Y]
[THREAD_TILE_SIZE_X],
T alpha, T beta, T* C, size_t ldc, size_t m, size_t n, size_t block_row_idx,
size_t block_col_idx, size_t warp_row_idx, size_t warp_col_idx,
size_t thread_row_idx_in_warp, size_t thread_col_idx_in_warp)
{
// Write the results to DRAM.
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP_Y;
++thread_tile_repeat_row_idx)
{
#pragma unroll
for (size_t thread_tile_repeat_col_idx{0U};
thread_tile_repeat_col_idx < NUM_THREAD_TILES_PER_WARP_X;
++thread_tile_repeat_col_idx)
{
#pragma unroll
for (size_t thread_tile_y_idx{0U};
thread_tile_y_idx < THREAD_TILE_SIZE_Y; ++thread_tile_y_idx)
{
#pragma unroll
for (size_t thread_tile_x_idx{0U};
thread_tile_x_idx < THREAD_TILE_SIZE_X;
++thread_tile_x_idx)
{
size_t const C_row_idx{
block_row_idx * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
thread_tile_repeat_row_idx *
(WARP_TILE_SIZE_Y / NUM_THREAD_TILES_PER_WARP_Y) +
thread_row_idx_in_warp * THREAD_TILE_SIZE_Y +
thread_tile_y_idx};
size_t const C_col_idx{
block_col_idx * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
thread_tile_repeat_col_idx *
(WARP_TILE_SIZE_X / NUM_THREAD_TILES_PER_WARP_X) +
thread_col_idx_in_warp * THREAD_TILE_SIZE_X +
thread_tile_x_idx};
if (C_row_idx < m && C_col_idx < n)
{
C[C_row_idx * ldc + C_col_idx] =
alpha * C_thread_results[thread_tile_repeat_row_idx]
[thread_tile_repeat_col_idx]
[thread_tile_y_idx]
[thread_tile_x_idx] +
beta * C[C_row_idx * ldc + C_col_idx];
}
}
}
}
}
}
// GEMM kernel v06.
// Each thread in the block processes THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values. Number of threads BLOCK_TILE_SIZE_Y *
// BLOCK_TILE_SIZE_X / (THREAD_TILE_SIZE_Y * THREAD_TILE_SIZE_X)
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y, size_t NUM_THREADS_PER_WARP_X,
size_t NUM_THREADS_PER_WARP_Y>
__global__ void gemm_v06(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
static_assert(NUM_THREADS_PER_WARP_X * NUM_THREADS_PER_WARP_Y == 32U);
constexpr size_t NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
constexpr size_t NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int NUM_THREAD_TILES_PER_WARP_X{
WARP_TILE_SIZE_X / (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X)};
constexpr unsigned int NUM_THREAD_TILES_PER_WARP_Y{
WARP_TILE_SIZE_Y / (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y)};
static_assert(
WARP_TILE_SIZE_X % (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X) == 0U);
static_assert(
WARP_TILE_SIZE_Y % (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y) == 0U);
constexpr unsigned int NUM_THREADS_X{NUM_WARPS_X * NUM_THREADS_PER_WARP_X};
constexpr unsigned int NUM_THREADS_Y{NUM_WARPS_Y * NUM_THREADS_PER_WARP_Y};
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{NUM_THREADS_X * NUM_THREADS_Y};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T
A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_Y];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
// A_vals is cached in the register.
T A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y] = {
static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
size_t const warp_linear_idx{thread_linear_idx / 32U};
size_t const warp_row_idx{warp_linear_idx / NUM_WARPS_X};
size_t const warp_col_idx{warp_linear_idx % NUM_WARPS_X};
size_t const thread_linear_idx_in_warp{thread_linear_idx % 32U};
size_t const thread_linear_row_idx_in_warp{thread_linear_idx_in_warp /
NUM_THREADS_PER_WARP_X};
size_t const thread_linear_col_idx_in_warp{thread_linear_idx_in_warp %
NUM_THREADS_PER_WARP_X};
// Number of outer loops to perform the sum of inner products.
// C_thread_block_tile =
// \sigma_{thread_block_tile_idx=0}^{num_thread_block_tiles-1} A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :]
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes NUM_THREAD_TILES_PER_WARP_Y *
// NUM_THREAD_TILES_PER_WARP_X * THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values.
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_transposed<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile_transposed,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
// Perform A[:, thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] where A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] and
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] are cached in the
// shared memory as A_thread_block_tile and B_thread_block_tile,
// respectively. This inner product is further decomposed to
// BLOCK_TILE_SIZE_K outer products. A_thread_block_tile *
// B_thread_block_tile = \sigma_{k_i=0}^{BLOCK_TILE_SIZE_K-1}
// A_thread_block_tile[:, k_i] @ B_thread_block_tile[k_i, :] Note that
// both A_thread_block_tile and B_thread_block_tile can be cached in the
// register.
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
// Load data from shared memory to register file for A.
load_data_from_shared_memory_to_register_file<
T, BLOCK_TILE_SIZE_Y, WARP_TILE_SIZE_Y, NUM_THREADS_PER_WARP_Y,
THREAD_TILE_SIZE_Y>(A_thread_block_tile_transposed[k_i], A_vals,
warp_row_idx,
thread_linear_row_idx_in_warp);
// Load data from shared memory to register file for B.
load_data_from_shared_memory_to_register_file<
T, BLOCK_TILE_SIZE_X, WARP_TILE_SIZE_X, NUM_THREADS_PER_WARP_X,
THREAD_TILE_SIZE_X>(B_thread_block_tile[k_i], B_vals,
warp_col_idx,
thread_linear_col_idx_in_warp);
// Compute NUM_THREAD_TILES_PER_WARP_Y * NUM_THREAD_TILES_PER_WARP_X
// outer products.
compute_thread_tile_results<T, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y,
THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y>(
A_vals, B_vals, C_thread_results);
}
__syncthreads();
}
// Write the results to DRAM.
write_results_from_register_file_to_global_memory<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, WARP_TILE_SIZE_X,
WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y,
NUM_THREAD_TILES_PER_WARP_X, NUM_THREAD_TILES_PER_WARP_Y>(
C_thread_results, alpha, beta, C, ldc, m, n, blockIdx.y, blockIdx.x,
warp_row_idx, warp_col_idx, thread_linear_row_idx_in_warp,
thread_linear_col_idx_in_warp);
}
template <typename T>
void launch_gemm_kernel_v06(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
constexpr unsigned int WARP_TILE_SIZE_X{32U};
constexpr unsigned int WARP_TILE_SIZE_Y{64U};
constexpr unsigned int NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr unsigned int NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_WARP_X{4U};
constexpr unsigned int NUM_THREADS_PER_WARP_Y{8U};
static_assert(NUM_THREADS_PER_WARP_X * NUM_THREADS_PER_WARP_Y == 32U);
static_assert(
WARP_TILE_SIZE_X % (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X) == 0U);
static_assert(
WARP_TILE_SIZE_Y % (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y) == 0U);
constexpr unsigned int NUM_THREADS_X{NUM_WARPS_X * NUM_THREADS_PER_WARP_X};
constexpr unsigned int NUM_THREADS_Y{NUM_WARPS_Y * NUM_THREADS_PER_WARP_Y};
constexpr unsigned int NUM_THREADS_PER_BLOCK{NUM_THREADS_X * NUM_THREADS_Y};
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v06<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y, NUM_THREADS_PER_WARP_X, NUM_THREADS_PER_WARP_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v06<float>(size_t m, size_t n, size_t k,
float const* alpha, float const* A,
size_t lda, float const* B,
size_t ldb, float const* beta,
float* C, size_t ldc,
cudaStream_t stream);
template void launch_gemm_kernel_v06<double>(size_t m, size_t n, size_t k,
double const* alpha,
double const* A, size_t lda,
double const* B, size_t ldb,
double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v06<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,361 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
template <typename T, size_t BLOCK_TILE_SIZE, size_t WARP_TILE_SIZE,
size_t NUM_THREAD_TILES_PER_WARP, size_t THREAD_TILE_SIZE>
__device__ void load_data_from_shared_memory_to_register_file_vectorized(
T const thread_block_tile[BLOCK_TILE_SIZE],
T register_values[NUM_THREAD_TILES_PER_WARP][THREAD_TILE_SIZE],
size_t warp_idx, size_t thread_idx)
{
static_assert(BLOCK_TILE_SIZE % THREAD_TILE_SIZE == 0U);
constexpr size_t NUM_VECTOR_UNITS{sizeof(int4) / sizeof(T)};
static_assert(sizeof(int4) % sizeof(T) == 0U);
constexpr size_t VECTORIZED_THREAD_TILE_SIZE{THREAD_TILE_SIZE /
NUM_VECTOR_UNITS};
static_assert(THREAD_TILE_SIZE % NUM_VECTOR_UNITS == 0U);
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP;
++thread_tile_repeat_row_idx)
{
size_t const thread_block_tile_row_idx{
warp_idx * WARP_TILE_SIZE +
thread_tile_repeat_row_idx *
(WARP_TILE_SIZE / NUM_THREAD_TILES_PER_WARP) +
thread_idx * THREAD_TILE_SIZE};
#pragma unroll
for (size_t thread_tile_vector_idx{0U};
thread_tile_vector_idx < VECTORIZED_THREAD_TILE_SIZE;
++thread_tile_vector_idx)
{
*reinterpret_cast<int4*>(
&register_values[thread_tile_repeat_row_idx]
[thread_tile_vector_idx * NUM_VECTOR_UNITS]) =
*reinterpret_cast<int4 const*>(
&thread_block_tile[thread_block_tile_row_idx +
thread_tile_vector_idx *
NUM_VECTOR_UNITS]);
}
}
}
template <typename T, size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__device__ void compute_thread_tile_results(
T const A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y],
T const B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X],
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X])
{
// Compute NUM_THREAD_TILES_PER_WARP_Y * NUM_THREAD_TILES_PER_WARP_X outer
// products.
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP_Y;
++thread_tile_repeat_row_idx)
{
#pragma unroll
for (size_t thread_tile_repeat_col_idx{0U};
thread_tile_repeat_col_idx < NUM_THREAD_TILES_PER_WARP_X;
++thread_tile_repeat_col_idx)
{
#pragma unroll
for (size_t thread_tile_y_idx{0U};
thread_tile_y_idx < THREAD_TILE_SIZE_Y; ++thread_tile_y_idx)
{
#pragma unroll
for (size_t thread_tile_x_idx{0U};
thread_tile_x_idx < THREAD_TILE_SIZE_X;
++thread_tile_x_idx)
{
C_thread_results[thread_tile_repeat_row_idx]
[thread_tile_repeat_col_idx]
[thread_tile_y_idx][thread_tile_x_idx] +=
A_vals[thread_tile_repeat_row_idx][thread_tile_y_idx] *
B_vals[thread_tile_repeat_col_idx][thread_tile_x_idx];
}
}
}
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t WARP_TILE_SIZE_X, size_t WARP_TILE_SIZE_Y,
size_t THREAD_TILE_SIZE_X, size_t THREAD_TILE_SIZE_Y,
size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y>
__device__ void write_results_from_register_file_to_global_memory_vectorized(
T const C_thread_results[NUM_THREAD_TILES_PER_WARP_Y]
[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_Y]
[THREAD_TILE_SIZE_X],
T alpha, T beta, T* C, size_t ldc, size_t m, size_t n, size_t block_row_idx,
size_t block_col_idx, size_t warp_row_idx, size_t warp_col_idx,
size_t thread_row_idx_in_warp, size_t thread_col_idx_in_warp)
{
constexpr size_t NUM_VECTOR_UNITS{sizeof(int4) / sizeof(T)};
static_assert(sizeof(int4) % sizeof(T) == 0U);
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_THREAD_TILE_SIZE_X{THREAD_TILE_SIZE_X /
NUM_VECTOR_UNITS};
static_assert(THREAD_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
// Write the results to DRAM.
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP_Y;
++thread_tile_repeat_row_idx)
{
#pragma unroll
for (size_t thread_tile_repeat_col_idx{0U};
thread_tile_repeat_col_idx < NUM_THREAD_TILES_PER_WARP_X;
++thread_tile_repeat_col_idx)
{
#pragma unroll
for (size_t thread_tile_y_idx{0U};
thread_tile_y_idx < THREAD_TILE_SIZE_Y; ++thread_tile_y_idx)
{
#pragma unroll
for (size_t thread_tile_x_vector_idx{0U};
thread_tile_x_vector_idx < VECTORIZED_THREAD_TILE_SIZE_X;
++thread_tile_x_vector_idx)
{
size_t const C_row_idx{
blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
thread_tile_repeat_row_idx *
(WARP_TILE_SIZE_Y / NUM_THREAD_TILES_PER_WARP_Y) +
thread_row_idx_in_warp * THREAD_TILE_SIZE_Y +
thread_tile_y_idx};
size_t const C_col_idx{
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
thread_tile_repeat_col_idx *
(WARP_TILE_SIZE_X / NUM_THREAD_TILES_PER_WARP_X) +
thread_col_idx_in_warp * THREAD_TILE_SIZE_X +
thread_tile_x_vector_idx * NUM_VECTOR_UNITS};
if (C_row_idx < m && C_col_idx < n)
{
int4 C_vals{*reinterpret_cast<int4 const*>(
&C[C_row_idx * ldc + C_col_idx])};
#pragma unroll
for (size_t i{0U}; i < NUM_VECTOR_UNITS; ++i)
{
reinterpret_cast<T*>(&C_vals)[i] =
alpha *
C_thread_results[thread_tile_repeat_row_idx]
[thread_tile_repeat_col_idx]
[thread_tile_y_idx]
[thread_tile_x_vector_idx *
NUM_VECTOR_UNITS +
i] +
beta * reinterpret_cast<T const*>(&C_vals)[i];
}
*reinterpret_cast<int4*>(
&C[C_row_idx * ldc + C_col_idx]) = C_vals;
}
}
}
}
}
}
// GEMM kernel v06.
// Each thread in the block processes THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values. Number of threads BLOCK_TILE_SIZE_Y *
// BLOCK_TILE_SIZE_X / (THREAD_TILE_SIZE_Y * THREAD_TILE_SIZE_X)
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y, size_t NUM_THREADS_PER_WARP_X,
size_t NUM_THREADS_PER_WARP_Y>
__global__ void gemm_v06_vectorized(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
static_assert(NUM_THREADS_PER_WARP_X * NUM_THREADS_PER_WARP_Y == 32U);
constexpr size_t NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
constexpr size_t NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int NUM_THREAD_TILES_PER_WARP_X{
WARP_TILE_SIZE_X / (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X)};
constexpr unsigned int NUM_THREAD_TILES_PER_WARP_Y{
WARP_TILE_SIZE_Y / (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y)};
static_assert(
WARP_TILE_SIZE_X % (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X) == 0U);
static_assert(
WARP_TILE_SIZE_Y % (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y) == 0U);
constexpr unsigned int NUM_THREADS_X{NUM_WARPS_X * NUM_THREADS_PER_WARP_X};
constexpr unsigned int NUM_THREADS_Y{NUM_WARPS_Y * NUM_THREADS_PER_WARP_Y};
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{NUM_THREADS_X * NUM_THREADS_Y};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T
A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_Y];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X];
// A_vals is cached in the register.
T A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y] = {
static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
size_t const warp_linear_idx{thread_linear_idx / 32U};
size_t const warp_row_idx{warp_linear_idx / NUM_WARPS_X};
size_t const warp_col_idx{warp_linear_idx % NUM_WARPS_X};
size_t const thread_linear_idx_in_warp{thread_linear_idx % 32U};
size_t const thread_linear_row_idx_in_warp{thread_linear_idx_in_warp /
NUM_THREADS_PER_WARP_X};
size_t const thread_linear_col_idx_in_warp{thread_linear_idx_in_warp %
NUM_THREADS_PER_WARP_X};
// Number of outer loops to perform the sum of inner products.
// C_thread_block_tile =
// \sigma_{thread_block_tile_idx=0}^{num_thread_block_tiles-1} A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :]
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes NUM_THREAD_TILES_PER_WARP_Y *
// NUM_THREAD_TILES_PER_WARP_X * THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values.
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS>(A, lda, B, ldb, A_thread_block_tile_transposed,
B_thread_block_tile, thread_block_tile_idx,
thread_linear_idx, m, n, k);
__syncthreads();
// Perform A[:, thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] where A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] and
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] are cached in the
// shared memory as A_thread_block_tile and B_thread_block_tile,
// respectively. This inner product is further decomposed to
// BLOCK_TILE_SIZE_K outer products. A_thread_block_tile *
// B_thread_block_tile = \sigma_{k_i=0}^{BLOCK_TILE_SIZE_K-1}
// A_thread_block_tile[:, k_i] @ B_thread_block_tile[k_i, :] Note that
// both A_thread_block_tile and B_thread_block_tile can be cached in the
// register.
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
// Load data from shared memory to register file for A.
load_data_from_shared_memory_to_register_file_vectorized<
T, BLOCK_TILE_SIZE_Y, WARP_TILE_SIZE_Y, NUM_THREADS_PER_WARP_Y,
THREAD_TILE_SIZE_Y>(A_thread_block_tile_transposed[k_i], A_vals,
warp_row_idx,
thread_linear_row_idx_in_warp);
// Load data from shared memory to register file for B.
load_data_from_shared_memory_to_register_file_vectorized<
T, BLOCK_TILE_SIZE_X, WARP_TILE_SIZE_X, NUM_THREADS_PER_WARP_X,
THREAD_TILE_SIZE_X>(B_thread_block_tile[k_i], B_vals,
warp_col_idx,
thread_linear_col_idx_in_warp);
// Compute NUM_THREAD_TILES_PER_WARP_Y * NUM_THREAD_TILES_PER_WARP_X
// outer products.
compute_thread_tile_results<T, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y,
THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y>(
A_vals, B_vals, C_thread_results);
}
__syncthreads();
}
// Write the results to DRAM.
write_results_from_register_file_to_global_memory_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, WARP_TILE_SIZE_X,
WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y,
NUM_THREAD_TILES_PER_WARP_X, NUM_THREAD_TILES_PER_WARP_Y>(
C_thread_results, alpha, beta, C, ldc, m, n, blockIdx.y, blockIdx.x,
warp_row_idx, warp_col_idx, thread_linear_row_idx_in_warp,
thread_linear_col_idx_in_warp);
}
template <typename T>
void launch_gemm_kernel_v06_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
constexpr unsigned int WARP_TILE_SIZE_X{32U};
constexpr unsigned int WARP_TILE_SIZE_Y{64U};
constexpr unsigned int NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr unsigned int NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_WARP_X{4U};
constexpr unsigned int NUM_THREADS_PER_WARP_Y{8U};
static_assert(NUM_THREADS_PER_WARP_X * NUM_THREADS_PER_WARP_Y == 32U);
static_assert(
WARP_TILE_SIZE_X % (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X) == 0U);
static_assert(
WARP_TILE_SIZE_Y % (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y) == 0U);
constexpr unsigned int NUM_THREADS_X{NUM_WARPS_X * NUM_THREADS_PER_WARP_X};
constexpr unsigned int NUM_THREADS_Y{NUM_WARPS_Y * NUM_THREADS_PER_WARP_Y};
constexpr unsigned int NUM_THREADS_PER_BLOCK{NUM_THREADS_X * NUM_THREADS_Y};
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v06_vectorized<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y,
BLOCK_TILE_SIZE_K, WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y,
THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y,
NUM_THREADS_PER_WARP_X, NUM_THREADS_PER_WARP_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v06_vectorized<float>(
size_t m, size_t n, size_t k, float const* alpha, float const* A,
size_t lda, float const* B, size_t ldb, float const* beta, float* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v06_vectorized<double>(
size_t m, size_t n, size_t k, double const* alpha, double const* A,
size_t lda, double const* B, size_t ldb, double const* beta, double* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v06_vectorized<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,481 @@
#include <cuda_fp16.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
template <typename T, size_t BLOCK_TILE_SIZE, size_t WARP_TILE_SIZE,
size_t NUM_THREAD_TILES_PER_WARP, size_t THREAD_TILE_SIZE>
__device__ void load_data_from_shared_memory_to_register_file_vectorized(
T const thread_block_tile[BLOCK_TILE_SIZE],
T register_values[NUM_THREAD_TILES_PER_WARP][THREAD_TILE_SIZE],
size_t warp_idx, size_t thread_idx)
{
static_assert(BLOCK_TILE_SIZE % THREAD_TILE_SIZE == 0U);
constexpr size_t NUM_VECTOR_UNITS{sizeof(int4) / sizeof(T)};
static_assert(sizeof(int4) % sizeof(T) == 0U);
constexpr size_t VECTORIZED_THREAD_TILE_SIZE{THREAD_TILE_SIZE /
NUM_VECTOR_UNITS};
static_assert(THREAD_TILE_SIZE % NUM_VECTOR_UNITS == 0U);
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP;
++thread_tile_repeat_row_idx)
{
size_t const thread_block_tile_row_idx{
warp_idx * WARP_TILE_SIZE +
thread_tile_repeat_row_idx *
(WARP_TILE_SIZE / NUM_THREAD_TILES_PER_WARP) +
thread_idx * THREAD_TILE_SIZE};
#pragma unroll
for (size_t thread_tile_vector_idx{0U};
thread_tile_vector_idx < VECTORIZED_THREAD_TILE_SIZE;
++thread_tile_vector_idx)
{
*reinterpret_cast<int4*>(
&register_values[thread_tile_repeat_row_idx]
[thread_tile_vector_idx * NUM_VECTOR_UNITS]) =
*reinterpret_cast<int4 const*>(
&thread_block_tile[thread_block_tile_row_idx +
thread_tile_vector_idx *
NUM_VECTOR_UNITS]);
}
}
}
template <typename T, size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y>
__device__ void compute_thread_tile_results(
T const A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y],
T const B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X],
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X])
{
// Compute NUM_THREAD_TILES_PER_WARP_Y * NUM_THREAD_TILES_PER_WARP_X outer
// products.
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP_Y;
++thread_tile_repeat_row_idx)
{
#pragma unroll
for (size_t thread_tile_repeat_col_idx{0U};
thread_tile_repeat_col_idx < NUM_THREAD_TILES_PER_WARP_X;
++thread_tile_repeat_col_idx)
{
#pragma unroll
for (size_t thread_tile_y_idx{0U};
thread_tile_y_idx < THREAD_TILE_SIZE_Y; ++thread_tile_y_idx)
{
#pragma unroll
for (size_t thread_tile_x_idx{0U};
thread_tile_x_idx < THREAD_TILE_SIZE_X;
++thread_tile_x_idx)
{
C_thread_results[thread_tile_repeat_row_idx]
[thread_tile_repeat_col_idx]
[thread_tile_y_idx][thread_tile_x_idx] +=
A_vals[thread_tile_repeat_row_idx][thread_tile_y_idx] *
B_vals[thread_tile_repeat_col_idx][thread_tile_x_idx];
}
}
}
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y, size_t NUM_THREADS_PER_WARP_X,
size_t NUM_THREADS_PER_WARP_Y, size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y>
__device__ void process_data_from_shared_memory_using_register_file_vectorized(
T A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y],
T B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X],
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X],
T const A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_Y],
T const B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X],
size_t warp_row_idx, size_t warp_col_idx, size_t thread_row_idx_in_warp,
size_t thread_col_idx_in_warp)
{
#pragma unroll
for (size_t k_i{0U}; k_i < BLOCK_TILE_SIZE_K; ++k_i)
{
// Load data from shared memory to register file for A.
load_data_from_shared_memory_to_register_file_vectorized<
T, BLOCK_TILE_SIZE_Y, WARP_TILE_SIZE_Y, NUM_THREADS_PER_WARP_Y,
THREAD_TILE_SIZE_Y>(A_thread_block_tile_transposed[k_i], A_vals,
warp_row_idx, thread_row_idx_in_warp);
// Load data from shared memory to register file for B.
load_data_from_shared_memory_to_register_file_vectorized<
T, BLOCK_TILE_SIZE_X, WARP_TILE_SIZE_X, NUM_THREADS_PER_WARP_X,
THREAD_TILE_SIZE_X>(B_thread_block_tile[k_i], B_vals, warp_col_idx,
thread_col_idx_in_warp);
// Compute NUM_THREAD_TILES_PER_WARP_Y *
// NUM_THREAD_TILES_PER_WARP_X outer products.
compute_thread_tile_results<T, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y,
THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y>(
A_vals, B_vals, C_thread_results);
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t WARP_TILE_SIZE_X, size_t WARP_TILE_SIZE_Y,
size_t THREAD_TILE_SIZE_X, size_t THREAD_TILE_SIZE_Y,
size_t NUM_THREAD_TILES_PER_WARP_X,
size_t NUM_THREAD_TILES_PER_WARP_Y>
__device__ void write_results_from_register_file_to_global_memory_vectorized(
T const C_thread_results[NUM_THREAD_TILES_PER_WARP_Y]
[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_Y]
[THREAD_TILE_SIZE_X],
T alpha, T beta, T* C, size_t ldc, size_t m, size_t n, size_t block_row_idx,
size_t block_col_idx, size_t warp_row_idx, size_t warp_col_idx,
size_t thread_row_idx_in_warp, size_t thread_col_idx_in_warp)
{
constexpr size_t NUM_VECTOR_UNITS{sizeof(int4) / sizeof(T)};
static_assert(sizeof(int4) % sizeof(T) == 0U);
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_THREAD_TILE_SIZE_X{THREAD_TILE_SIZE_X /
NUM_VECTOR_UNITS};
static_assert(THREAD_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
// Write the results to DRAM.
#pragma unroll
for (size_t thread_tile_repeat_row_idx{0U};
thread_tile_repeat_row_idx < NUM_THREAD_TILES_PER_WARP_Y;
++thread_tile_repeat_row_idx)
{
#pragma unroll
for (size_t thread_tile_repeat_col_idx{0U};
thread_tile_repeat_col_idx < NUM_THREAD_TILES_PER_WARP_X;
++thread_tile_repeat_col_idx)
{
#pragma unroll
for (size_t thread_tile_y_idx{0U};
thread_tile_y_idx < THREAD_TILE_SIZE_Y; ++thread_tile_y_idx)
{
#pragma unroll
for (size_t thread_tile_x_vector_idx{0U};
thread_tile_x_vector_idx < VECTORIZED_THREAD_TILE_SIZE_X;
++thread_tile_x_vector_idx)
{
size_t const C_row_idx{
blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
thread_tile_repeat_row_idx *
(WARP_TILE_SIZE_Y / NUM_THREAD_TILES_PER_WARP_Y) +
thread_row_idx_in_warp * THREAD_TILE_SIZE_Y +
thread_tile_y_idx};
size_t const C_col_idx{
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
thread_tile_repeat_col_idx *
(WARP_TILE_SIZE_X / NUM_THREAD_TILES_PER_WARP_X) +
thread_col_idx_in_warp * THREAD_TILE_SIZE_X +
thread_tile_x_vector_idx * NUM_VECTOR_UNITS};
if (C_row_idx < m && C_col_idx < n)
{
int4 C_vals{*reinterpret_cast<int4 const*>(
&C[C_row_idx * ldc + C_col_idx])};
#pragma unroll
for (size_t i{0U}; i < NUM_VECTOR_UNITS; ++i)
{
reinterpret_cast<T*>(&C_vals)[i] =
alpha *
C_thread_results[thread_tile_repeat_row_idx]
[thread_tile_repeat_col_idx]
[thread_tile_y_idx]
[thread_tile_x_vector_idx *
NUM_VECTOR_UNITS +
i] +
beta * reinterpret_cast<T const*>(&C_vals)[i];
}
*reinterpret_cast<int4*>(
&C[C_row_idx * ldc + C_col_idx]) = C_vals;
}
}
}
}
}
}
// GEMM kernel v06.
// Each thread in the block processes THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values. Number of threads BLOCK_TILE_SIZE_Y *
// BLOCK_TILE_SIZE_X / (THREAD_TILE_SIZE_Y * THREAD_TILE_SIZE_X)
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t THREAD_TILE_SIZE_X,
size_t THREAD_TILE_SIZE_Y, size_t NUM_THREADS_PER_WARP_X,
size_t NUM_THREADS_PER_WARP_Y>
__global__ void
gemm_v06_vectorized_double_buffered(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
static_assert(NUM_THREADS_PER_WARP_X * NUM_THREADS_PER_WARP_Y == 32U);
constexpr size_t NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
constexpr size_t NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int NUM_THREAD_TILES_PER_WARP_X{
WARP_TILE_SIZE_X / (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X)};
constexpr unsigned int NUM_THREAD_TILES_PER_WARP_Y{
WARP_TILE_SIZE_Y / (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y)};
static_assert(
WARP_TILE_SIZE_X % (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X) == 0U);
static_assert(
WARP_TILE_SIZE_Y % (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y) == 0U);
constexpr unsigned int NUM_THREADS_X{NUM_WARPS_X * NUM_THREADS_PER_WARP_X};
constexpr unsigned int NUM_THREADS_Y{NUM_WARPS_Y * NUM_THREADS_PER_WARP_Y};
// Avoid using blockDim.x * blockDim.y as the number of threads per block.
// Because it is a runtime constant and the compiler cannot optimize the
// loop unrolling based on that.
// Use a compile time constant instead.
constexpr size_t NUM_THREADS{NUM_THREADS_X * NUM_THREADS_Y};
constexpr size_t NUM_PIPELINES{2U};
// Only double buffer is supported in the implementation.
// But even more number of pipelines can be supported if the implementation
// is modified.
static_assert(NUM_PIPELINES == 2U);
static_assert((NUM_WARPS_X * NUM_WARPS_Y) % NUM_PIPELINES == 0U);
static_assert(NUM_THREADS % NUM_PIPELINES == 0U);
constexpr size_t NUM_THREADS_PER_PIPELINE{NUM_THREADS / NUM_PIPELINES};
constexpr size_t NUM_WARPS_PER_PIPELINE{(NUM_WARPS_X * NUM_WARPS_Y) /
NUM_PIPELINES};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T
A_thread_block_tile_transposed[NUM_PIPELINES][BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_Y];
__shared__ T B_thread_block_tile[NUM_PIPELINES][BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X];
// A_vals is cached in the register.
T A_vals[NUM_THREAD_TILES_PER_WARP_Y][THREAD_TILE_SIZE_Y] = {
static_cast<T>(0)};
// B_vals is cached in the register.
T B_vals[NUM_THREAD_TILES_PER_WARP_X][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
size_t const warp_linear_idx{thread_linear_idx / 32U};
size_t const warp_row_idx{warp_linear_idx / NUM_WARPS_X};
size_t const warp_col_idx{warp_linear_idx % NUM_WARPS_X};
size_t const thread_linear_idx_in_warp{thread_linear_idx % 32U};
size_t const thread_linear_row_idx_in_warp{thread_linear_idx_in_warp /
NUM_THREADS_PER_WARP_X};
size_t const thread_linear_col_idx_in_warp{thread_linear_idx_in_warp %
NUM_THREADS_PER_WARP_X};
// Separate the warps to different pipelines.
size_t const pipeline_index{warp_linear_idx / NUM_WARPS_PER_PIPELINE};
// Number of outer loops to perform the sum of inner products.
// C_thread_block_tile =
// \sigma_{thread_block_tile_idx=0}^{num_thread_block_tiles-1} A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :]
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
// Each thread in the block processes NUM_THREAD_TILES_PER_WARP_Y *
// NUM_THREAD_TILES_PER_WARP_X * THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values.
T C_thread_results[NUM_THREAD_TILES_PER_WARP_Y][NUM_THREAD_TILES_PER_WARP_X]
[THREAD_TILE_SIZE_Y][THREAD_TILE_SIZE_X] = {
static_cast<T>(0)};
if (pipeline_index == 0U)
{
// Pipeline 0 warps load buffer 0.
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS_PER_PIPELINE>(
A, lda, B, ldb, A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index], 0U,
thread_linear_idx - pipeline_index * NUM_THREADS_PER_PIPELINE, m, n,
k);
}
__syncthreads();
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
thread_block_tile_idx += NUM_PIPELINES)
{
if (pipeline_index == 0U)
{
// Pipeline 0 warps process buffer 0.
process_data_from_shared_memory_using_register_file_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y, NUM_THREADS_PER_WARP_X,
NUM_THREADS_PER_WARP_Y, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y>(
A_vals, B_vals, C_thread_results,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index], warp_row_idx, warp_col_idx,
thread_linear_row_idx_in_warp, thread_linear_col_idx_in_warp);
__syncthreads();
// Pipeline 0 warps process buffer 1.
if (thread_block_tile_idx + 1U < num_thread_block_tiles)
{
process_data_from_shared_memory_using_register_file_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y, NUM_THREADS_PER_WARP_X,
NUM_THREADS_PER_WARP_Y, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y>(
A_vals, B_vals, C_thread_results,
A_thread_block_tile_transposed[pipeline_index + 1],
B_thread_block_tile[pipeline_index + 1], warp_row_idx,
warp_col_idx, thread_linear_row_idx_in_warp,
thread_linear_col_idx_in_warp);
}
__syncthreads();
// Pipeline 0 warps load buffer 0.
if (thread_block_tile_idx + 2U < num_thread_block_tiles)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS_PER_PIPELINE>(
A, lda, B, ldb,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index],
thread_block_tile_idx + 2,
thread_linear_idx -
pipeline_index * NUM_THREADS_PER_PIPELINE,
m, n, k);
}
__syncthreads();
}
else
{
// Pipeline 1 warps load buffer 1.
if (thread_block_tile_idx + 1U < num_thread_block_tiles)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS_PER_PIPELINE>(
A, lda, B, ldb,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index],
thread_block_tile_idx + 1,
thread_linear_idx -
pipeline_index * NUM_THREADS_PER_PIPELINE,
m, n, k);
}
__syncthreads();
// Pipeline 1 warps process buffer 0.
process_data_from_shared_memory_using_register_file_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y, NUM_THREADS_PER_WARP_X,
NUM_THREADS_PER_WARP_Y, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y>(
A_vals, B_vals, C_thread_results,
A_thread_block_tile_transposed[pipeline_index - 1],
B_thread_block_tile[pipeline_index - 1], warp_row_idx,
warp_col_idx, thread_linear_row_idx_in_warp,
thread_linear_col_idx_in_warp);
__syncthreads();
// Pipeline 1 warps process buffer 1.
if (thread_block_tile_idx + 1U < num_thread_block_tiles)
{
process_data_from_shared_memory_using_register_file_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y, NUM_THREADS_PER_WARP_X,
NUM_THREADS_PER_WARP_Y, NUM_THREAD_TILES_PER_WARP_X,
NUM_THREAD_TILES_PER_WARP_Y>(
A_vals, B_vals, C_thread_results,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index], warp_row_idx,
warp_col_idx, thread_linear_row_idx_in_warp,
thread_linear_col_idx_in_warp);
}
__syncthreads();
}
}
// Write the results to DRAM.
write_results_from_register_file_to_global_memory_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, WARP_TILE_SIZE_X,
WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X, THREAD_TILE_SIZE_Y,
NUM_THREAD_TILES_PER_WARP_X, NUM_THREAD_TILES_PER_WARP_Y>(
C_thread_results, alpha, beta, C, ldc, m, n, blockIdx.y, blockIdx.x,
warp_row_idx, warp_col_idx, thread_linear_row_idx_in_warp,
thread_linear_col_idx_in_warp);
}
template <typename T>
void launch_gemm_kernel_v06_vectorized_double_buffered(
size_t m, size_t n, size_t k, T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
constexpr unsigned int WARP_TILE_SIZE_X{32U};
constexpr unsigned int WARP_TILE_SIZE_Y{64U};
constexpr unsigned int NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr unsigned int NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int THREAD_TILE_SIZE_X{8U};
constexpr unsigned int THREAD_TILE_SIZE_Y{8U};
constexpr unsigned int NUM_THREADS_PER_WARP_X{4U};
constexpr unsigned int NUM_THREADS_PER_WARP_Y{8U};
static_assert(NUM_THREADS_PER_WARP_X * NUM_THREADS_PER_WARP_Y == 32U);
static_assert(
WARP_TILE_SIZE_X % (THREAD_TILE_SIZE_X * NUM_THREADS_PER_WARP_X) == 0U);
static_assert(
WARP_TILE_SIZE_Y % (THREAD_TILE_SIZE_Y * NUM_THREADS_PER_WARP_Y) == 0U);
constexpr unsigned int NUM_THREADS_X{NUM_WARPS_X * NUM_THREADS_PER_WARP_X};
constexpr unsigned int NUM_THREADS_Y{NUM_WARPS_Y * NUM_THREADS_PER_WARP_Y};
constexpr unsigned int NUM_THREADS_PER_BLOCK{NUM_THREADS_X * NUM_THREADS_Y};
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v06_vectorized_double_buffered<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, THREAD_TILE_SIZE_X,
THREAD_TILE_SIZE_Y, NUM_THREADS_PER_WARP_X, NUM_THREADS_PER_WARP_Y>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v06_vectorized_double_buffered<float>(
size_t m, size_t n, size_t k, float const* alpha, float const* A,
size_t lda, float const* B, size_t ldb, float const* beta, float* C,
size_t ldc, cudaStream_t stream);
template void launch_gemm_kernel_v06_vectorized_double_buffered<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,247 @@
#include <cuda_fp16.h>
#include <mma.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// https://developer.nvidia.com/blog/cutlass-linear-algebra-cuda/
// https://github.com/NVIDIA/cutlass/blob/b7508e337938137a699e486d8997646980acfc58/media/docs/programming_guidelines.md
// GEMM kernel v07.
// Each thread in the block processes THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values. Number of threads BLOCK_TILE_SIZE_Y *
// BLOCK_TILE_SIZE_X / (THREAD_TILE_SIZE_Y * THREAD_TILE_SIZE_X)
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t BLOCK_TILE_SKEW_SIZE_X,
size_t BLOCK_TILE_SKEW_SIZE_Y, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_X,
size_t WMMA_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_K, size_t NUM_THREADS>
__global__ void gemm_v07(size_t m, size_t n, size_t k, T alpha, T const* A,
size_t lda, T const* B, size_t ldb, T beta, T* C,
size_t ldc)
{
constexpr size_t NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_Y +
BLOCK_TILE_SKEW_SIZE_Y];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X +
BLOCK_TILE_SKEW_SIZE_X];
constexpr size_t NUM_WMMA_TILES_X{WARP_TILE_SIZE_X / WMMA_TILE_SIZE_X};
static_assert(WARP_TILE_SIZE_X % WMMA_TILE_SIZE_X == 0U);
constexpr size_t NUM_WMMA_TILES_Y{WARP_TILE_SIZE_Y / WMMA_TILE_SIZE_Y};
static_assert(WARP_TILE_SIZE_Y % WMMA_TILE_SIZE_Y == 0U);
constexpr size_t NUM_WMMA_TILES_K{BLOCK_TILE_SIZE_K / WMMA_TILE_SIZE_K};
static_assert(BLOCK_TILE_SIZE_K % WMMA_TILE_SIZE_K == 0U);
// Declare the fragments.
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::col_major>
a_frags[NUM_WMMA_TILES_Y];
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::row_major>
b_frags[NUM_WMMA_TILES_X];
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
acc_frags[NUM_WMMA_TILES_Y][NUM_WMMA_TILES_X];
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
c_frag;
// Make sure the accumulator starts from 0.
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
nvcuda::wmma::fill_fragment(
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx],
static_cast<T>(0));
}
}
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
size_t const warp_linear_idx{thread_linear_idx / 32U};
size_t const warp_row_idx{warp_linear_idx / NUM_WARPS_X};
size_t const warp_col_idx{warp_linear_idx % NUM_WARPS_X};
// Number of outer loops to perform the sum of inner products.
// C_thread_block_tile =
// \sigma_{thread_block_tile_idx=0}^{num_thread_block_tiles-1} A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :]
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_transposed<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS, BLOCK_TILE_SKEW_SIZE_X, BLOCK_TILE_SKEW_SIZE_Y>(
A, lda, B, ldb, A_thread_block_tile_transposed, B_thread_block_tile,
thread_block_tile_idx, thread_linear_idx, m, n, k);
__syncthreads();
// Perform A[:, thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] where A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] and
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] are cached in the
// shared memory as A_thread_block_tile and B_thread_block_tile,
// respectively. This inner product is further decomposed to
// BLOCK_TILE_SIZE_K outer products. A_thread_block_tile *
// B_thread_block_tile = \sigma_{k_i=0}^{BLOCK_TILE_SIZE_K-1}
// A_thread_block_tile[:, k_i] @ B_thread_block_tile[k_i, :] Note that
// both A_thread_block_tile and B_thread_block_tile can be cached in the
// register.
#pragma unroll
for (size_t k_i{0U}; k_i < NUM_WMMA_TILES_K; ++k_i)
{
#pragma unroll
for (size_t wmma_tile_row_idx{0U};
wmma_tile_row_idx < NUM_WMMA_TILES_Y; ++wmma_tile_row_idx)
{
nvcuda::wmma::load_matrix_sync(
a_frags[wmma_tile_row_idx],
&A_thread_block_tile_transposed[k_i * WMMA_TILE_SIZE_K]
[warp_row_idx *
WARP_TILE_SIZE_Y +
wmma_tile_row_idx *
WMMA_TILE_SIZE_Y],
BLOCK_TILE_SIZE_Y + BLOCK_TILE_SKEW_SIZE_Y);
}
#pragma unroll
for (size_t wmma_tile_col_idx{0U};
wmma_tile_col_idx < NUM_WMMA_TILES_X; ++wmma_tile_col_idx)
{
nvcuda::wmma::load_matrix_sync(
b_frags[wmma_tile_col_idx],
&B_thread_block_tile[k_i * WMMA_TILE_SIZE_K]
[warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X);
}
#pragma unroll
for (size_t wmma_tile_row_idx{0U};
wmma_tile_row_idx < NUM_WMMA_TILES_Y; ++wmma_tile_row_idx)
{
#pragma unroll
for (size_t wmma_tile_col_idx{0U};
wmma_tile_col_idx < NUM_WMMA_TILES_X; ++wmma_tile_col_idx)
{
// Perform the matrix multiplication.
nvcuda::wmma::mma_sync(
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx],
a_frags[wmma_tile_row_idx], b_frags[wmma_tile_col_idx],
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx]);
}
}
}
__syncthreads();
}
// Write the results to DRAM.
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
#pragma unroll
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
// Load the fragment from global memory.
nvcuda::wmma::load_matrix_sync(
c_frag,
&C[(blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
wmma_tile_row_idx * WMMA_TILE_SIZE_Y) *
n +
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
n, nvcuda::wmma::mem_row_major);
// Perform scaling and addition.
for (size_t i{0}; i < c_frag.num_elements; ++i)
{
c_frag.x[i] =
alpha *
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx].x[i] +
beta * c_frag.x[i];
}
// Store the fragment back to global memory.
nvcuda::wmma::store_matrix_sync(
&C[(blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
wmma_tile_row_idx * WMMA_TILE_SIZE_Y) *
n +
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
c_frag, n, nvcuda::wmma::mem_row_major);
}
}
}
template <typename T>
void launch_gemm_kernel_v07(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
constexpr unsigned int WARP_TILE_SIZE_X{32U};
constexpr unsigned int WARP_TILE_SIZE_Y{64U};
constexpr unsigned int NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr unsigned int NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
// The skew size is used to avoid bank conflicts in shared memory.
constexpr size_t BLOCK_TILE_SKEW_SIZE_X{16U};
constexpr size_t BLOCK_TILE_SKEW_SIZE_Y{16U};
constexpr unsigned int WMMA_TILE_SIZE_X{16U};
constexpr unsigned int WMMA_TILE_SIZE_Y{16U};
constexpr unsigned int WMMA_TILE_SIZE_K{16U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{NUM_WARPS_X * NUM_WARPS_Y *
32U};
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v07<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
BLOCK_TILE_SKEW_SIZE_X, BLOCK_TILE_SKEW_SIZE_Y, WARP_TILE_SIZE_X,
WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_K, NUM_THREADS_PER_BLOCK>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v07<__half>(size_t m, size_t n, size_t k,
__half const* alpha,
__half const* A, size_t lda,
__half const* B, size_t ldb,
__half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,246 @@
#include <cuda_fp16.h>
#include <mma.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// https://developer.nvidia.com/blog/cutlass-linear-algebra-cuda/
// https://github.com/NVIDIA/cutlass/blob/b7508e337938137a699e486d8997646980acfc58/media/docs/programming_guidelines.md
// GEMM kernel v07.
// Each thread in the block processes THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values. Number of threads BLOCK_TILE_SIZE_Y *
// BLOCK_TILE_SIZE_X / (THREAD_TILE_SIZE_Y * THREAD_TILE_SIZE_X)
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t BLOCK_TILE_SKEW_SIZE_X,
size_t BLOCK_TILE_SKEW_SIZE_Y, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_X,
size_t WMMA_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_K, size_t NUM_THREADS>
__global__ void gemm_v07_vectorized(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
constexpr size_t NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_Y +
BLOCK_TILE_SKEW_SIZE_Y];
__shared__ T B_thread_block_tile[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_X +
BLOCK_TILE_SKEW_SIZE_X];
constexpr size_t NUM_WMMA_TILES_X{WARP_TILE_SIZE_X / WMMA_TILE_SIZE_X};
static_assert(WARP_TILE_SIZE_X % WMMA_TILE_SIZE_X == 0U);
constexpr size_t NUM_WMMA_TILES_Y{WARP_TILE_SIZE_Y / WMMA_TILE_SIZE_Y};
static_assert(WARP_TILE_SIZE_Y % WMMA_TILE_SIZE_Y == 0U);
constexpr size_t NUM_WMMA_TILES_K{BLOCK_TILE_SIZE_K / WMMA_TILE_SIZE_K};
static_assert(BLOCK_TILE_SIZE_K % WMMA_TILE_SIZE_K == 0U);
// Declare the fragments.
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::col_major>
a_frags[NUM_WMMA_TILES_Y];
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::row_major>
b_frags[NUM_WMMA_TILES_X];
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
acc_frags[NUM_WMMA_TILES_Y][NUM_WMMA_TILES_X];
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
c_frag;
// Make sure the accumulator starts from 0.
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
nvcuda::wmma::fill_fragment(
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx],
static_cast<T>(0));
}
}
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
size_t const warp_linear_idx{thread_linear_idx / 32U};
size_t const warp_row_idx{warp_linear_idx / NUM_WARPS_X};
size_t const warp_col_idx{warp_linear_idx % NUM_WARPS_X};
// Number of outer loops to perform the sum of inner products.
// C_thread_block_tile =
// \sigma_{thread_block_tile_idx=0}^{num_thread_block_tiles-1} A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :]
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
++thread_block_tile_idx)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS, BLOCK_TILE_SKEW_SIZE_X, BLOCK_TILE_SKEW_SIZE_Y>(
A, lda, B, ldb, A_thread_block_tile_transposed, B_thread_block_tile,
thread_block_tile_idx, thread_linear_idx, m, n, k);
__syncthreads();
// Perform A[:, thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] where A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] and
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :] are cached in the
// shared memory as A_thread_block_tile and B_thread_block_tile,
// respectively. This inner product is further decomposed to
// BLOCK_TILE_SIZE_K outer products. A_thread_block_tile *
// B_thread_block_tile = \sigma_{k_i=0}^{BLOCK_TILE_SIZE_K-1}
// A_thread_block_tile[:, k_i] @ B_thread_block_tile[k_i, :] Note that
// both A_thread_block_tile and B_thread_block_tile can be cached in the
// register.
#pragma unroll
for (size_t k_i{0U}; k_i < NUM_WMMA_TILES_K; ++k_i)
{
#pragma unroll
for (size_t wmma_tile_row_idx{0U};
wmma_tile_row_idx < NUM_WMMA_TILES_Y; ++wmma_tile_row_idx)
{
nvcuda::wmma::load_matrix_sync(
a_frags[wmma_tile_row_idx],
&A_thread_block_tile_transposed[k_i * WMMA_TILE_SIZE_K]
[warp_row_idx *
WARP_TILE_SIZE_Y +
wmma_tile_row_idx *
WMMA_TILE_SIZE_Y],
BLOCK_TILE_SIZE_Y + BLOCK_TILE_SKEW_SIZE_Y);
}
#pragma unroll
for (size_t wmma_tile_col_idx{0U};
wmma_tile_col_idx < NUM_WMMA_TILES_X; ++wmma_tile_col_idx)
{
nvcuda::wmma::load_matrix_sync(
b_frags[wmma_tile_col_idx],
&B_thread_block_tile[k_i * WMMA_TILE_SIZE_K]
[warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X);
}
#pragma unroll
for (size_t wmma_tile_row_idx{0U};
wmma_tile_row_idx < NUM_WMMA_TILES_Y; ++wmma_tile_row_idx)
{
#pragma unroll
for (size_t wmma_tile_col_idx{0U};
wmma_tile_col_idx < NUM_WMMA_TILES_X; ++wmma_tile_col_idx)
{
// Perform the matrix multiplication.
nvcuda::wmma::mma_sync(
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx],
a_frags[wmma_tile_row_idx], b_frags[wmma_tile_col_idx],
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx]);
}
}
}
__syncthreads();
}
// Write the results to DRAM.
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
#pragma unroll
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
// Load the fragment from global memory.
nvcuda::wmma::load_matrix_sync(
c_frag,
&C[(blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
wmma_tile_row_idx * WMMA_TILE_SIZE_Y) *
n +
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
n, nvcuda::wmma::mem_row_major);
// Perform scaling and addition.
for (size_t i{0}; i < c_frag.num_elements; ++i)
{
c_frag.x[i] =
alpha *
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx].x[i] +
beta * c_frag.x[i];
}
// Store the fragment back to global memory.
nvcuda::wmma::store_matrix_sync(
&C[(blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
wmma_tile_row_idx * WMMA_TILE_SIZE_Y) *
n +
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
c_frag, n, nvcuda::wmma::mem_row_major);
}
}
}
template <typename T>
void launch_gemm_kernel_v07_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
// The skew size is used to avoid bank conflicts in shared memory.
constexpr size_t BLOCK_TILE_SKEW_SIZE_X{16U};
constexpr size_t BLOCK_TILE_SKEW_SIZE_Y{16U};
constexpr unsigned int WARP_TILE_SIZE_X{32U};
constexpr unsigned int WARP_TILE_SIZE_Y{64U};
constexpr unsigned int NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr unsigned int NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int WMMA_TILE_SIZE_X{16U};
constexpr unsigned int WMMA_TILE_SIZE_Y{16U};
constexpr unsigned int WMMA_TILE_SIZE_K{16U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{NUM_WARPS_X * NUM_WARPS_Y *
32U};
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v07_vectorized<T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y,
BLOCK_TILE_SIZE_K, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y, WARP_TILE_SIZE_X,
WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_K, NUM_THREADS_PER_BLOCK>
<<<grid_dim, block_dim, 0U, stream>>>(m, n, k, *alpha, A, lda, B, ldb,
*beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v07_vectorized<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,380 @@
#include <cuda_fp16.h>
#include <mma.h>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include "cuda_gemm_utils.hpp"
// https://developer.nvidia.com/blog/cutlass-linear-algebra-cuda/
// https://github.com/NVIDIA/cutlass/blob/b7508e337938137a699e486d8997646980acfc58/media/docs/programming_guidelines.md
template <
typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t WARP_TILE_SIZE_X, size_t WARP_TILE_SIZE_Y,
size_t WMMA_TILE_SIZE_X, size_t WMMA_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_K,
size_t NUM_WMMA_TILES_X, size_t NUM_WMMA_TILES_Y, size_t NUM_WMMA_TILES_K,
size_t BLOCK_TILE_SKEW_SIZE_X, size_t BLOCK_TILE_SKEW_SIZE_Y>
__device__ void process_data_from_shared_memory_using_wmma(
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::col_major>
a_frags[NUM_WMMA_TILES_Y],
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::row_major>
b_frags[NUM_WMMA_TILES_X],
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
acc_frags[NUM_WMMA_TILES_Y][NUM_WMMA_TILES_X],
T const A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_Y +
BLOCK_TILE_SKEW_SIZE_Y],
T const B_thread_block_tile[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X],
size_t warp_row_idx, size_t warp_col_idx)
{
#pragma unroll
for (size_t k_i{0U}; k_i < NUM_WMMA_TILES_K; ++k_i)
{
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
nvcuda::wmma::load_matrix_sync(
a_frags[wmma_tile_row_idx],
&A_thread_block_tile_transposed[k_i * WMMA_TILE_SIZE_K]
[warp_row_idx *
WARP_TILE_SIZE_Y +
wmma_tile_row_idx *
WMMA_TILE_SIZE_Y],
BLOCK_TILE_SIZE_Y + BLOCK_TILE_SKEW_SIZE_Y);
}
#pragma unroll
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
nvcuda::wmma::load_matrix_sync(
b_frags[wmma_tile_col_idx],
&B_thread_block_tile[k_i * WMMA_TILE_SIZE_K]
[warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X);
}
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
#pragma unroll
for (size_t wmma_tile_col_idx{0U};
wmma_tile_col_idx < NUM_WMMA_TILES_X; ++wmma_tile_col_idx)
{
// Perform the matrix multiplication.
nvcuda::wmma::mma_sync(
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx],
a_frags[wmma_tile_row_idx], b_frags[wmma_tile_col_idx],
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx]);
}
}
}
}
// GEMM kernel v07.
// Each thread in the block processes THREAD_TILE_SIZE_Y *
// THREAD_TILE_SIZE_X output values. Number of threads BLOCK_TILE_SIZE_Y *
// BLOCK_TILE_SIZE_X / (THREAD_TILE_SIZE_Y * THREAD_TILE_SIZE_X)
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t BLOCK_TILE_SKEW_SIZE_X,
size_t BLOCK_TILE_SKEW_SIZE_Y, size_t WARP_TILE_SIZE_X,
size_t WARP_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_X,
size_t WMMA_TILE_SIZE_Y, size_t WMMA_TILE_SIZE_K, size_t NUM_THREADS>
__global__ void
gemm_v07_vectorized_double_buffered(size_t m, size_t n, size_t k, T alpha,
T const* A, size_t lda, T const* B,
size_t ldb, T beta, T* C, size_t ldc)
{
constexpr size_t NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr size_t NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr size_t NUM_WMMA_TILES_X{WARP_TILE_SIZE_X / WMMA_TILE_SIZE_X};
static_assert(WARP_TILE_SIZE_X % WMMA_TILE_SIZE_X == 0U);
constexpr size_t NUM_WMMA_TILES_Y{WARP_TILE_SIZE_Y / WMMA_TILE_SIZE_Y};
static_assert(WARP_TILE_SIZE_Y % WMMA_TILE_SIZE_Y == 0U);
constexpr size_t NUM_WMMA_TILES_K{BLOCK_TILE_SIZE_K / WMMA_TILE_SIZE_K};
static_assert(BLOCK_TILE_SIZE_K % WMMA_TILE_SIZE_K == 0U);
constexpr size_t NUM_PIPELINES{2U};
// Only double buffer is supported in the implementation.
// But even more number of pipelines can be supported if the implementation
// is modified.
static_assert(NUM_PIPELINES == 2U);
static_assert((NUM_WARPS_X * NUM_WARPS_Y) % NUM_PIPELINES == 0U);
static_assert(NUM_THREADS % NUM_PIPELINES == 0U);
constexpr size_t NUM_THREADS_PER_PIPELINE{NUM_THREADS / NUM_PIPELINES};
constexpr size_t NUM_WARPS_PER_PIPELINE{(NUM_WARPS_X * NUM_WARPS_Y) /
NUM_PIPELINES};
// Cache a tile of A and B in shared memory for data reuse.
__shared__ T
A_thread_block_tile_transposed[NUM_PIPELINES][BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_Y +
BLOCK_TILE_SKEW_SIZE_Y];
__shared__ T
B_thread_block_tile[NUM_PIPELINES][BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X];
// Declare the fragments.
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::col_major>
a_frags[NUM_WMMA_TILES_Y];
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T,
nvcuda::wmma::row_major>
b_frags[NUM_WMMA_TILES_X];
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
acc_frags[NUM_WMMA_TILES_Y][NUM_WMMA_TILES_X];
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, WMMA_TILE_SIZE_Y,
WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_K, T>
c_frag;
// Make sure the accumulator starts from 0.
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
nvcuda::wmma::fill_fragment(
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx],
static_cast<T>(0));
}
}
size_t const thread_linear_idx{threadIdx.y * blockDim.x + threadIdx.x};
size_t const warp_linear_idx{thread_linear_idx / 32U};
size_t const warp_row_idx{warp_linear_idx / NUM_WARPS_X};
size_t const warp_col_idx{warp_linear_idx % NUM_WARPS_X};
// Separate the warps to different pipelines.
size_t const pipeline_index{warp_linear_idx / NUM_WARPS_PER_PIPELINE};
// Number of outer loops to perform the sum of inner products.
// C_thread_block_tile =
// \sigma_{thread_block_tile_idx=0}^{num_thread_block_tiles-1} A[:,
// thread_block_tile_idx:BLOCK_TILE_SIZE_K] *
// B[thread_block_tile_idx:BLOCK_TILE_SIZE_K, :]
size_t const num_thread_block_tiles{(k + BLOCK_TILE_SIZE_K - 1) /
BLOCK_TILE_SIZE_K};
if (pipeline_index == 0U)
{
// Pipeline 0 warps load buffer 0.
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS_PER_PIPELINE, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
A, lda, B, ldb, A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index], 0U,
thread_linear_idx - pipeline_index * NUM_THREADS_PER_PIPELINE, m, n,
k);
}
__syncthreads();
for (size_t thread_block_tile_idx{0U};
thread_block_tile_idx < num_thread_block_tiles;
thread_block_tile_idx += NUM_PIPELINES)
{
if (pipeline_index == 0U)
{
// Pipeline 0 warps process buffer 0.
process_data_from_shared_memory_using_wmma<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X,
WMMA_TILE_SIZE_Y, WMMA_TILE_SIZE_K, NUM_WMMA_TILES_X,
NUM_WMMA_TILES_Y, NUM_WMMA_TILES_K, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
a_frags, b_frags, acc_frags,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index], warp_row_idx,
warp_col_idx);
__syncthreads();
// Pipeline 0 warps process buffer 1.
if (thread_block_tile_idx + 1U < num_thread_block_tiles)
{
process_data_from_shared_memory_using_wmma<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X,
WMMA_TILE_SIZE_Y, WMMA_TILE_SIZE_K, NUM_WMMA_TILES_X,
NUM_WMMA_TILES_Y, NUM_WMMA_TILES_K, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
a_frags, b_frags, acc_frags,
A_thread_block_tile_transposed[pipeline_index + 1],
B_thread_block_tile[pipeline_index + 1], warp_row_idx,
warp_col_idx);
}
__syncthreads();
// Pipeline 0 warps load buffer 0.
if (thread_block_tile_idx + 2U < num_thread_block_tiles)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS_PER_PIPELINE, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
A, lda, B, ldb,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index],
thread_block_tile_idx + 2,
thread_linear_idx -
pipeline_index * NUM_THREADS_PER_PIPELINE,
m, n, k);
}
__syncthreads();
}
else
{
// Pipeline 1 warps load buffer 1.
if (thread_block_tile_idx + 1U < num_thread_block_tiles)
{
load_data_from_global_memory_to_shared_memory_transposed_vectorized<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
NUM_THREADS_PER_PIPELINE, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
A, lda, B, ldb,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index],
thread_block_tile_idx + 1,
thread_linear_idx -
pipeline_index * NUM_THREADS_PER_PIPELINE,
m, n, k);
}
__syncthreads();
// Pipeline 1 warps process buffer 0.
process_data_from_shared_memory_using_wmma<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X,
WMMA_TILE_SIZE_Y, WMMA_TILE_SIZE_K, NUM_WMMA_TILES_X,
NUM_WMMA_TILES_Y, NUM_WMMA_TILES_K, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
a_frags, b_frags, acc_frags,
A_thread_block_tile_transposed[pipeline_index - 1],
B_thread_block_tile[pipeline_index - 1], warp_row_idx,
warp_col_idx);
__syncthreads();
// Pipeline 1 warps process buffer 1.
if (thread_block_tile_idx + 1U < num_thread_block_tiles)
{
process_data_from_shared_memory_using_wmma<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
WARP_TILE_SIZE_X, WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X,
WMMA_TILE_SIZE_Y, WMMA_TILE_SIZE_K, NUM_WMMA_TILES_X,
NUM_WMMA_TILES_Y, NUM_WMMA_TILES_K, BLOCK_TILE_SKEW_SIZE_X,
BLOCK_TILE_SKEW_SIZE_Y>(
a_frags, b_frags, acc_frags,
A_thread_block_tile_transposed[pipeline_index],
B_thread_block_tile[pipeline_index], warp_row_idx,
warp_col_idx);
}
__syncthreads();
}
}
// Write the results to DRAM.
#pragma unroll
for (size_t wmma_tile_row_idx{0U}; wmma_tile_row_idx < NUM_WMMA_TILES_Y;
++wmma_tile_row_idx)
{
#pragma unroll
for (size_t wmma_tile_col_idx{0U}; wmma_tile_col_idx < NUM_WMMA_TILES_X;
++wmma_tile_col_idx)
{
// Load the fragment from global memory.
nvcuda::wmma::load_matrix_sync(
c_frag,
&C[(blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
wmma_tile_row_idx * WMMA_TILE_SIZE_Y) *
n +
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
n, nvcuda::wmma::mem_row_major);
// Perform scaling and addition.
for (size_t i{0}; i < c_frag.num_elements; ++i)
{
c_frag.x[i] =
alpha *
acc_frags[wmma_tile_row_idx][wmma_tile_col_idx].x[i] +
beta * c_frag.x[i];
}
// Store the fragment back to global memory.
nvcuda::wmma::store_matrix_sync(
&C[(blockIdx.y * BLOCK_TILE_SIZE_Y +
warp_row_idx * WARP_TILE_SIZE_Y +
wmma_tile_row_idx * WMMA_TILE_SIZE_Y) *
n +
blockIdx.x * BLOCK_TILE_SIZE_X +
warp_col_idx * WARP_TILE_SIZE_X +
wmma_tile_col_idx * WMMA_TILE_SIZE_X],
c_frag, n, nvcuda::wmma::mem_row_major);
}
}
}
template <typename T>
void launch_gemm_kernel_v07_vectorized_double_buffered(
size_t m, size_t n, size_t k, T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta, T* C, size_t ldc,
cudaStream_t stream)
{
// Feel free to play with the block tile sizes.
// The algorithm correctness should always be guaranteed.
constexpr unsigned int BLOCK_TILE_SIZE_X{128U};
constexpr unsigned int BLOCK_TILE_SIZE_Y{128U};
constexpr unsigned int BLOCK_TILE_SIZE_K{16U};
// The skew size is used to avoid bank conflicts in shared memory.
constexpr size_t BLOCK_TILE_SKEW_SIZE_X{16U};
constexpr size_t BLOCK_TILE_SKEW_SIZE_Y{16U};
constexpr unsigned int WARP_TILE_SIZE_X{32U};
constexpr unsigned int WARP_TILE_SIZE_Y{64U};
constexpr unsigned int NUM_WARPS_X{BLOCK_TILE_SIZE_X / WARP_TILE_SIZE_X};
constexpr unsigned int NUM_WARPS_Y{BLOCK_TILE_SIZE_Y / WARP_TILE_SIZE_Y};
static_assert(BLOCK_TILE_SIZE_X % WARP_TILE_SIZE_X == 0U);
static_assert(BLOCK_TILE_SIZE_Y % WARP_TILE_SIZE_Y == 0U);
constexpr unsigned int WMMA_TILE_SIZE_X{16U};
constexpr unsigned int WMMA_TILE_SIZE_Y{16U};
constexpr unsigned int WMMA_TILE_SIZE_K{16U};
constexpr unsigned int NUM_THREADS_PER_BLOCK{NUM_WARPS_X * NUM_WARPS_Y *
32U};
dim3 const block_dim{NUM_THREADS_PER_BLOCK, 1U, 1U};
dim3 const grid_dim{
(static_cast<unsigned int>(n) + BLOCK_TILE_SIZE_X - 1U) /
BLOCK_TILE_SIZE_X,
(static_cast<unsigned int>(m) + BLOCK_TILE_SIZE_Y - 1U) /
BLOCK_TILE_SIZE_Y,
1U};
gemm_v07_vectorized_double_buffered<
T, BLOCK_TILE_SIZE_X, BLOCK_TILE_SIZE_Y, BLOCK_TILE_SIZE_K,
BLOCK_TILE_SKEW_SIZE_X, BLOCK_TILE_SKEW_SIZE_Y, WARP_TILE_SIZE_X,
WARP_TILE_SIZE_Y, WMMA_TILE_SIZE_X, WMMA_TILE_SIZE_Y, WMMA_TILE_SIZE_K,
NUM_THREADS_PER_BLOCK><<<grid_dim, block_dim, 0U, stream>>>(
m, n, k, *alpha, A, lda, B, ldb, *beta, C, ldc);
CHECK_LAST_CUDA_ERROR();
}
// Explicit instantiation.
template void launch_gemm_kernel_v07_vectorized_double_buffered<__half>(
size_t m, size_t n, size_t k, __half const* alpha, __half const* A,
size_t lda, __half const* B, size_t ldb, __half const* beta, __half* C,
size_t ldc, cudaStream_t stream);

View File

@@ -0,0 +1,98 @@
#ifndef CUDA_GEMM_HPP
#define CUDA_GEMM_HPP
#include <cuda_runtime.h>
template <typename T>
void launch_gemm_kernel_v00(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v01(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v02(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v02_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v03(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v03_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v04(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v04_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v05(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v05_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v06(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v06_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v06_vectorized_double_buffered(
size_t m, size_t n, size_t k, T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v07(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc,
cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v07_vectorized(size_t m, size_t n, size_t k,
T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta,
T* C, size_t ldc, cudaStream_t stream);
template <typename T>
void launch_gemm_kernel_v07_vectorized_double_buffered(
size_t m, size_t n, size_t k, T const* alpha, T const* A, size_t lda,
T const* B, size_t ldb, T const* beta, T* C, size_t ldc,
cudaStream_t stream);
#endif

View File

@@ -0,0 +1,29 @@
#include <iostream>
#include <cuda_runtime.h>
#include "cuda_gemm_utils.hpp"
void check_cuda(cudaError_t err, const char* const func, const char* const file,
const int line)
{
if (err != cudaSuccess)
{
std::cerr << "CUDA Runtime Error at: " << file << ":" << line
<< std::endl;
std::cerr << cudaGetErrorString(err) << " " << func << std::endl;
std::exit(EXIT_FAILURE);
}
}
void check_cuda_last(const char* const file, const int line)
{
cudaError_t const err{cudaGetLastError()};
if (err != cudaSuccess)
{
std::cerr << "CUDA Runtime Error at: " << file << ":" << line
<< std::endl;
std::cerr << cudaGetErrorString(err) << std::endl;
std::exit(EXIT_FAILURE);
}
}

View File

@@ -0,0 +1,486 @@
#ifndef CUDA_GEMM_UTILS_CUH
#define CUDA_GEMM_UTILS_CUH
#include <cuda_runtime.h>
#include "cuda_gemm_utils.hpp"
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t NUM_THREADS,
size_t BLOCK_TILE_SKEW_SIZE_X = 0U,
size_t BLOCK_TILE_SKEW_SIZE_K = 0U>
__device__ void load_data_from_global_memory_to_shared_memory(
T const* A, size_t lda, T const* B, size_t ldb,
T A_thread_block_tile[BLOCK_TILE_SIZE_Y]
[BLOCK_TILE_SIZE_K + BLOCK_TILE_SKEW_SIZE_K],
T B_thread_block_tile[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X],
size_t thread_block_tile_idx, size_t thread_linear_idx, size_t m, size_t n,
size_t k)
{
// Load data from A on DRAM to A_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx < (BLOCK_TILE_SIZE_Y * BLOCK_TILE_SIZE_K + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const A_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) / BLOCK_TILE_SIZE_K};
size_t const A_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) % BLOCK_TILE_SIZE_K};
size_t const A_row_idx{blockIdx.y * BLOCK_TILE_SIZE_Y +
A_thread_block_tile_row_idx};
size_t const A_col_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
A_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
T val{static_cast<T>(0)};
if (A_row_idx < m && A_col_idx < k)
{
val = A[A_row_idx * lda + A_col_idx];
}
// This if will slow down the kernel.
// Add static asserts from the host code to guarantee this if is
// always true.
static_assert(BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS ==
0U);
// if (A_thread_block_tile_row_idx < BLOCK_TILE_SIZE_Y &&
// A_thread_block_tile_col_idx < BLOCK_TILE_SIZE_K)
// {
// A_thread_block_tile[A_thread_block_tile_row_idx]
// [A_thread_block_tile_col_idx] = val;
// }
A_thread_block_tile[A_thread_block_tile_row_idx]
[A_thread_block_tile_col_idx] = val;
}
// Load data from B on DRAM to B_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx < (BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_X + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const B_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) / BLOCK_TILE_SIZE_X};
size_t const B_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) % BLOCK_TILE_SIZE_X};
size_t const B_row_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
B_thread_block_tile_row_idx};
size_t const B_col_idx{blockIdx.x * BLOCK_TILE_SIZE_X +
B_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
T val{static_cast<T>(0)};
if (B_row_idx < k && B_col_idx < n)
{
val = B[B_row_idx * ldb + B_col_idx];
}
// This if will slow down the kernel.
// Add static asserts from the host code to guarantee this if is
// always true.
static_assert(BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS ==
0U);
// if (B_thread_block_tile_row_idx < BLOCK_TILE_SIZE_K &&
// B_thread_block_tile_col_idx < BLOCK_TILE_SIZE_X)
// {
// B_thread_block_tile[B_thread_block_tile_row_idx]
// [B_thread_block_tile_col_idx] = val;
// }
B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx] = val;
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t NUM_THREADS,
size_t BLOCK_TILE_SKEW_SIZE_X = 0U,
size_t BLOCK_TILE_SKEW_SIZE_Y = 0U>
__device__ void load_data_from_global_memory_to_shared_memory_transposed(
T const* A, size_t lda, T const* B, size_t ldb,
T A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_Y +
BLOCK_TILE_SKEW_SIZE_Y],
T B_thread_block_tile[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X],
size_t thread_block_tile_idx, size_t thread_linear_idx, size_t m, size_t n,
size_t k)
{
// Load data from A on DRAM to A_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx < (BLOCK_TILE_SIZE_Y * BLOCK_TILE_SIZE_K + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const A_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) / BLOCK_TILE_SIZE_K};
size_t const A_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) % BLOCK_TILE_SIZE_K};
size_t const A_row_idx{blockIdx.y * BLOCK_TILE_SIZE_Y +
A_thread_block_tile_row_idx};
size_t const A_col_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
A_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
T val{static_cast<T>(0)};
if (A_row_idx < m && A_col_idx < k)
{
val = A[A_row_idx * lda + A_col_idx];
}
// Removing the if will give another ~2 FLOPs performance on RTX
// 3090. But it will make the kernel incorrect for some GEMM
// configurations. T val{A[A_row_idx * lda + A_col_idx]}; This if
// will slow down the kernel. Add static asserts from the host code
// to guarantee this if is always true.
static_assert(BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y % NUM_THREADS ==
0U);
// if (A_thread_block_tile_row_idx < BLOCK_TILE_SIZE_Y &&
// A_thread_block_tile_col_idx < BLOCK_TILE_SIZE_K)
// {
// A_thread_block_tile[A_thread_block_tile_row_idx]
// [A_thread_block_tile_col_idx] = val;
// }
A_thread_block_tile_transposed[A_thread_block_tile_col_idx]
[A_thread_block_tile_row_idx] = val;
}
// Load data from B on DRAM to B_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx < (BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_X + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const B_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) / BLOCK_TILE_SIZE_X};
size_t const B_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) % BLOCK_TILE_SIZE_X};
size_t const B_row_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
B_thread_block_tile_row_idx};
size_t const B_col_idx{blockIdx.x * BLOCK_TILE_SIZE_X +
B_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
T val{static_cast<T>(0)};
if (B_row_idx < k && B_col_idx < n)
{
val = B[B_row_idx * ldb + B_col_idx];
}
// Removing the if will give another ~2 FLOPs performance on RTX
// 3090. But it will make the kernel incorrect for some GEMM
// configurations. T val{B[B_row_idx * ldb + B_col_idx]}; This if
// will slow down the kernel. Add static asserts from the host code
// to guarantee this if is always true.
static_assert(BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K % NUM_THREADS ==
0U);
// if (B_thread_block_tile_row_idx < BLOCK_TILE_SIZE_K &&
// B_thread_block_tile_col_idx < BLOCK_TILE_SIZE_X)
// {
// B_thread_block_tile[B_thread_block_tile_row_idx]
// [B_thread_block_tile_col_idx] = val;
// }
B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx] = val;
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t NUM_THREADS,
size_t BLOCK_TILE_SKEW_SIZE_X = 0U,
size_t BLOCK_TILE_SKEW_SIZE_K = 0U, typename VECTOR_TYPE = int4>
__device__ void load_data_from_global_memory_to_shared_memory_vectorized(
T const* A, size_t lda, T const* B, size_t ldb,
T A_thread_block_tile[BLOCK_TILE_SIZE_Y]
[BLOCK_TILE_SIZE_K + BLOCK_TILE_SKEW_SIZE_K],
T B_thread_block_tile[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X],
size_t thread_block_tile_idx, size_t thread_linear_idx, size_t m, size_t n,
size_t k)
{
constexpr size_t NUM_VECTOR_UNITS{sizeof(VECTOR_TYPE) / sizeof(T)};
static_assert(sizeof(VECTOR_TYPE) % sizeof(T) == 0U);
static_assert(BLOCK_TILE_SIZE_K % NUM_VECTOR_UNITS == 0U);
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_BLOCK_TILE_SIZE_K{BLOCK_TILE_SIZE_K /
NUM_VECTOR_UNITS};
static_assert(BLOCK_TILE_SIZE_K % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_BLOCK_TILE_SIZE_X{BLOCK_TILE_SIZE_X /
NUM_VECTOR_UNITS};
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
// The skew size could affect the data alignment in shared memory when we
// use vectorized load. We need to make sure the data alignment is correct.
static_assert((BLOCK_TILE_SIZE_K) * sizeof(T) % sizeof(VECTOR_TYPE) == 0U);
static_assert((BLOCK_TILE_SIZE_X) * sizeof(T) % sizeof(VECTOR_TYPE) == 0U);
static_assert((BLOCK_TILE_SIZE_K + BLOCK_TILE_SKEW_SIZE_K) * sizeof(T) %
sizeof(VECTOR_TYPE) ==
0U);
static_assert((BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X) * sizeof(T) %
sizeof(VECTOR_TYPE) ==
0U);
// Load data from A on DRAM to A_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx <
(BLOCK_TILE_SIZE_Y * VECTORIZED_BLOCK_TILE_SIZE_K + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const A_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) /
VECTORIZED_BLOCK_TILE_SIZE_K};
size_t const A_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) %
VECTORIZED_BLOCK_TILE_SIZE_K * NUM_VECTOR_UNITS};
size_t const A_row_idx{blockIdx.y * BLOCK_TILE_SIZE_Y +
A_thread_block_tile_row_idx};
size_t const A_col_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
A_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
VECTOR_TYPE A_row_vector_vals{0, 0, 0, 0};
if (A_row_idx < m && A_col_idx < k)
{
A_row_vector_vals = *reinterpret_cast<VECTOR_TYPE const*>(
&A[A_row_idx * lda + A_col_idx]);
}
if (A_col_idx + NUM_VECTOR_UNITS > k)
{
// Number of invalid elements in the last vector.
size_t const num_invalid_elements{A_col_idx + NUM_VECTOR_UNITS - k};
// Mask out the invalid elements.
T* const A_row_vector_vals_ptr{
reinterpret_cast<T*>(&A_row_vector_vals)};
for (size_t i{0U}; i < num_invalid_elements; ++i)
{
A_row_vector_vals_ptr[NUM_VECTOR_UNITS - 1U - i] =
static_cast<T>(0);
}
}
// If this is true, the following if can be removed.
// static_assert(VECTORIZED_BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y %
// NUM_THREADS == 0U);
if (A_thread_block_tile_row_idx < BLOCK_TILE_SIZE_Y &&
A_thread_block_tile_col_idx < BLOCK_TILE_SIZE_K)
{
*reinterpret_cast<int4*>(
&A_thread_block_tile[A_thread_block_tile_row_idx]
[A_thread_block_tile_col_idx]) =
A_row_vector_vals;
}
}
// Load data from B on DRAM to B_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx <
(BLOCK_TILE_SIZE_K * VECTORIZED_BLOCK_TILE_SIZE_X + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const B_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) /
VECTORIZED_BLOCK_TILE_SIZE_X};
size_t const B_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) %
VECTORIZED_BLOCK_TILE_SIZE_X * NUM_VECTOR_UNITS};
size_t const B_row_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
B_thread_block_tile_row_idx};
size_t const B_col_idx{blockIdx.x * BLOCK_TILE_SIZE_X +
B_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
VECTOR_TYPE B_row_vector_vals{0, 0, 0, 0};
if (B_row_idx < k && B_col_idx < n)
{
B_row_vector_vals = *reinterpret_cast<VECTOR_TYPE const*>(
&B[B_row_idx * ldb + B_col_idx]);
}
if (B_col_idx + NUM_VECTOR_UNITS > n)
{
// Number of invalid elements in the last vector.
size_t const num_invalid_elements{B_col_idx + NUM_VECTOR_UNITS - n};
// Mask out the invalid elements.
T* const B_row_vector_vals_ptr{
reinterpret_cast<T*>(&B_row_vector_vals)};
for (size_t i{0U}; i < num_invalid_elements; ++i)
{
B_row_vector_vals_ptr[NUM_VECTOR_UNITS - 1U - i] =
static_cast<T>(0);
}
}
// If this is true, the following if can be removed.
// static_assert(VECTORIZED_BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K %
// NUM_THREADS ==
// 0U);
if (B_thread_block_tile_row_idx < BLOCK_TILE_SIZE_K &&
B_thread_block_tile_col_idx < BLOCK_TILE_SIZE_X)
{
*reinterpret_cast<int4*>(
&B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx]) =
B_row_vector_vals;
}
}
}
template <typename T, size_t BLOCK_TILE_SIZE_X, size_t BLOCK_TILE_SIZE_Y,
size_t BLOCK_TILE_SIZE_K, size_t NUM_THREADS,
size_t BLOCK_TILE_SKEW_SIZE_X = 0U,
size_t BLOCK_TILE_SKEW_SIZE_Y = 0U, typename VECTOR_TYPE = int4>
__device__ void
load_data_from_global_memory_to_shared_memory_transposed_vectorized(
T const* A, size_t lda, T const* B, size_t ldb,
T A_thread_block_tile_transposed[BLOCK_TILE_SIZE_K][BLOCK_TILE_SIZE_Y +
BLOCK_TILE_SKEW_SIZE_Y],
T B_thread_block_tile[BLOCK_TILE_SIZE_K]
[BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X],
size_t thread_block_tile_idx, size_t thread_linear_idx, size_t m, size_t n,
size_t k)
{
constexpr size_t NUM_VECTOR_UNITS{sizeof(VECTOR_TYPE) / sizeof(T)};
static_assert(sizeof(VECTOR_TYPE) % sizeof(T) == 0U);
static_assert(BLOCK_TILE_SIZE_K % NUM_VECTOR_UNITS == 0U);
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_BLOCK_TILE_SIZE_K{BLOCK_TILE_SIZE_K /
NUM_VECTOR_UNITS};
static_assert(BLOCK_TILE_SIZE_K % NUM_VECTOR_UNITS == 0U);
constexpr size_t VECTORIZED_BLOCK_TILE_SIZE_X{BLOCK_TILE_SIZE_X /
NUM_VECTOR_UNITS};
static_assert(BLOCK_TILE_SIZE_X % NUM_VECTOR_UNITS == 0U);
// The skew size could affect the data alignment in shared memory when we
// use vectorized load. We need to make sure the data alignment is correct.
static_assert((BLOCK_TILE_SIZE_Y) * sizeof(T) % sizeof(VECTOR_TYPE) == 0U);
static_assert((BLOCK_TILE_SIZE_X) * sizeof(T) % sizeof(VECTOR_TYPE) == 0U);
static_assert((BLOCK_TILE_SIZE_Y + BLOCK_TILE_SKEW_SIZE_Y) * sizeof(T) %
sizeof(VECTOR_TYPE) ==
0U);
static_assert((BLOCK_TILE_SIZE_X + BLOCK_TILE_SKEW_SIZE_X) * sizeof(T) %
sizeof(VECTOR_TYPE) ==
0U);
// Load data from A on DRAM to A_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx <
(BLOCK_TILE_SIZE_Y * VECTORIZED_BLOCK_TILE_SIZE_K + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const A_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) /
VECTORIZED_BLOCK_TILE_SIZE_K};
size_t const A_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) %
VECTORIZED_BLOCK_TILE_SIZE_K * NUM_VECTOR_UNITS};
size_t const A_row_idx{blockIdx.y * BLOCK_TILE_SIZE_Y +
A_thread_block_tile_row_idx};
size_t const A_col_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
A_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
int4 A_row_vector_vals{0, 0, 0, 0};
if (A_row_idx < m && A_col_idx < k)
{
A_row_vector_vals =
*reinterpret_cast<int4 const*>(&A[A_row_idx * lda + A_col_idx]);
}
if (A_col_idx + NUM_VECTOR_UNITS > k)
{
// Number of invalid elements in the last vector.
size_t const num_invalid_elements{A_col_idx + NUM_VECTOR_UNITS - k};
// Mask out the invalid elements.
T* const A_row_vector_vals_ptr{
reinterpret_cast<T*>(&A_row_vector_vals)};
for (size_t i{0U}; i < num_invalid_elements; ++i)
{
A_row_vector_vals_ptr[NUM_VECTOR_UNITS - 1U - i] =
static_cast<T>(0);
}
}
// If this is true, the following if can be removed.
// static_assert(VECTORIZED_BLOCK_TILE_SIZE_K * BLOCK_TILE_SIZE_Y %
// NUM_THREADS ==
// 0U);
if (A_thread_block_tile_row_idx < BLOCK_TILE_SIZE_Y &&
A_thread_block_tile_col_idx < BLOCK_TILE_SIZE_K)
{
for (size_t i{0U}; i < NUM_VECTOR_UNITS; ++i)
{
A_thread_block_tile_transposed[A_thread_block_tile_col_idx +
i][A_thread_block_tile_row_idx] =
reinterpret_cast<T const*>(&A_row_vector_vals)[i];
}
}
}
// Load data from B on DRAM to B_thread_block_tile on shared memory.
#pragma unroll
for (size_t load_idx{0U};
load_idx <
(BLOCK_TILE_SIZE_K * VECTORIZED_BLOCK_TILE_SIZE_X + NUM_THREADS - 1U) /
NUM_THREADS;
++load_idx)
{
size_t const B_thread_block_tile_row_idx{
(thread_linear_idx + load_idx * NUM_THREADS) /
VECTORIZED_BLOCK_TILE_SIZE_X};
size_t const B_thread_block_tile_col_idx{
(thread_linear_idx + load_idx * NUM_THREADS) %
VECTORIZED_BLOCK_TILE_SIZE_X * NUM_VECTOR_UNITS};
size_t const B_row_idx{thread_block_tile_idx * BLOCK_TILE_SIZE_K +
B_thread_block_tile_row_idx};
size_t const B_col_idx{blockIdx.x * BLOCK_TILE_SIZE_X +
B_thread_block_tile_col_idx};
// These boundary checks might slow down the kernel to some extent.
// But they guarantee the correctness of the kernel for all
// different GEMM configurations.
int4 B_row_vector_vals{0, 0, 0, 0};
if (B_row_idx < k && B_col_idx < n)
{
B_row_vector_vals =
*reinterpret_cast<int4 const*>(&B[B_row_idx * ldb + B_col_idx]);
}
if (B_col_idx + NUM_VECTOR_UNITS > n)
{
// Number of invalid elements in the last vector.
size_t const num_invalid_elements{B_col_idx + NUM_VECTOR_UNITS - n};
// Mask out the invalid elements.
T* const B_row_vector_vals_ptr{
reinterpret_cast<T*>(&B_row_vector_vals)};
for (size_t i{0U}; i < num_invalid_elements; ++i)
{
B_row_vector_vals_ptr[NUM_VECTOR_UNITS - 1U - i] =
static_cast<T>(0);
}
}
// If this is true, the following if can be removed.
// static_assert(VECTORIZED_BLOCK_TILE_SIZE_X * BLOCK_TILE_SIZE_K %
// NUM_THREADS ==
// 0U);
if (B_thread_block_tile_row_idx < BLOCK_TILE_SIZE_K &&
B_thread_block_tile_col_idx < BLOCK_TILE_SIZE_X)
{
*reinterpret_cast<int4*>(
&B_thread_block_tile[B_thread_block_tile_row_idx]
[B_thread_block_tile_col_idx]) =
B_row_vector_vals;
}
}
}
#endif // CUDA_GEMM_UTILS_CUH

View File

@@ -0,0 +1,13 @@
#ifndef CUDA_GEMM_UTILS_HPP
#define CUDA_GEMM_UTILS_HPP
#include <cuda_runtime.h>
#define CHECK_CUDA_ERROR(val) check_cuda((val), #val, __FILE__, __LINE__)
void check_cuda(cudaError_t err, const char* const func, const char* const file,
const int line);
#define CHECK_LAST_CUDA_ERROR() check_cuda_last(__FILE__, __LINE__)
void check_cuda_last(const char* const file, const int line);
#endif // CUDA_GEMM_UTILS_HPP

View File

@@ -0,0 +1,109 @@
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include "cuda_gemm.hpp"
#include "profile_utils.cuh"
int main()
{
print_device_info();
constexpr size_t num_repeats{1U};
constexpr size_t num_warmups{1U};
__half const fp16_abs_tol{__float2half(5.0e-2f)};
double const fp16_rel_tol{1.0e-1f};
__half const fp16_tensor_core_abs_tol{__float2half(5.0e-2f)};
double const fp16_tensor_core_rel_tol{1.0e-2f};
constexpr size_t m{4096U};
constexpr size_t k{4096U};
constexpr size_t n{4096U};
constexpr size_t lda{(k + 16U - 1U) / 16U * 16U};
constexpr size_t ldb{(n + 16U - 1U) / 16U * 16U};
constexpr size_t ldc{(n + 16U - 1U) / 16U * 16U};
static_assert(lda >= k);
static_assert(ldb >= n);
static_assert(ldc >= n);
std::cout << "Matrix Size: " << "M = " << m << " N = " << n << " K = " << k
<< std::endl;
std::cout << "Matrix A: " << m << " x " << k
<< " Leading Dimension Size = " << lda << std::endl;
std::cout << "Matrix B: " << k << " x " << n
<< " Leading Dimension Size = " << ldb << std::endl;
std::cout << "Matrix C: " << m << " x " << n
<< " Leading Dimension Size = " << ldc << std::endl;
std::cout << std::endl;
// Define all the GEMM kernel launch functions to be profiled.
std::vector<std::pair<
std::string,
std::function<void(size_t, size_t, size_t, __half const*, __half const*,
size_t, __half const*, size_t, __half const*,
__half*, size_t, cudaStream_t)>>> const
gemm_fp16_kernel_launch_functions{
{"Custom GEMM Kernel V00", launch_gemm_kernel_v00<__half>},
{"Custom GEMM Kernel V01", launch_gemm_kernel_v01<__half>},
{"Custom GEMM Kernel V02", launch_gemm_kernel_v02<__half>},
{"Custom GEMM Kernel V02 Vectorized",
launch_gemm_kernel_v02_vectorized<__half>},
{"Custom GEMM Kernel V03", launch_gemm_kernel_v03<__half>},
{"Custom GEMM Kernel V03 Vectorized",
launch_gemm_kernel_v03_vectorized<__half>},
{"Custom GEMM Kernel V04", launch_gemm_kernel_v04<__half>},
{"Custom GEMM Kernel V04 Vectorized",
launch_gemm_kernel_v04_vectorized<__half>},
{"Custom GEMM Kernel V05", launch_gemm_kernel_v05<__half>},
{"Custom GEMM Kernel V05 Vectorized",
launch_gemm_kernel_v05_vectorized<__half>},
{"Custom GEMM Kernel V06", launch_gemm_kernel_v06<__half>},
{"Custom GEMM Kernel V06 Vectorized",
launch_gemm_kernel_v06_vectorized<__half>},
{"Custom GEMM Kernel V06 Vectorized Double Buffered",
launch_gemm_kernel_v06_vectorized_double_buffered<__half>},
};
for (auto const& gemm_fp16_kernel_launch_function :
gemm_fp16_kernel_launch_functions)
{
std::cout << gemm_fp16_kernel_launch_function.first << std::endl;
std::pair<__half, __half> const gemm_kernel_profile_result{
profile_gemm<__half>(
m, n, k, lda, ldb, ldc, gemm_fp16_kernel_launch_function.second,
fp16_abs_tol, fp16_rel_tol, num_repeats, num_warmups)};
std::cout << std::endl;
}
std::vector<std::pair<
std::string,
std::function<void(size_t, size_t, size_t, __half const*, __half const*,
size_t, __half const*, size_t, __half const*,
__half*, size_t, cudaStream_t)>>> const
gemm_fp16_tensor_core_kernel_launch_functions{
{"Custom GEMM Kernel V07", launch_gemm_kernel_v07<__half>},
{"Custom GEMM Kernel V07 Vectorized",
launch_gemm_kernel_v07_vectorized<__half>},
{"Custom GEMM Kernel V07 Vectorized Double Buffered",
launch_gemm_kernel_v07_vectorized_double_buffered<__half>},
};
for (auto const& gemm_fp16_tensor_core_kernel_launch_function :
gemm_fp16_tensor_core_kernel_launch_functions)
{
std::cout << gemm_fp16_tensor_core_kernel_launch_function.first
<< std::endl;
std::pair<__half, __half> const gemm_kernel_profile_result{
profile_gemm<__half>(
m, n, k, lda, ldb, ldc,
gemm_fp16_tensor_core_kernel_launch_function.second,
fp16_tensor_core_abs_tol, fp16_tensor_core_rel_tol, num_repeats,
num_warmups)};
std::cout << std::endl;
}
return 0;
}

View File

@@ -0,0 +1,78 @@
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include "cuda_gemm.hpp"
#include "profile_utils.cuh"
int main()
{
print_device_info();
constexpr size_t num_repeats{1U};
constexpr size_t num_warmups{1U};
float const fp32_abs_tol{1.0e-3f};
double const fp32_rel_tol{0.0e-4f};
constexpr size_t m{4096U};
constexpr size_t k{4096U};
constexpr size_t n{4096U};
constexpr size_t lda{(k + 16U - 1U) / 16U * 16U};
constexpr size_t ldb{(n + 16U - 1U) / 16U * 16U};
constexpr size_t ldc{(n + 16U - 1U) / 16U * 16U};
static_assert(lda >= k);
static_assert(ldb >= n);
static_assert(ldc >= n);
std::cout << "Matrix Size: " << "M = " << m << " N = " << n << " K = " << k
<< std::endl;
std::cout << "Matrix A: " << m << " x " << k
<< " Leading Dimension Size = " << lda << std::endl;
std::cout << "Matrix B: " << k << " x " << n
<< " Leading Dimension Size = " << ldb << std::endl;
std::cout << "Matrix C: " << m << " x " << n
<< " Leading Dimension Size = " << ldc << std::endl;
std::cout << std::endl;
// Define all the GEMM kernel launch functions to be profiled.
std::vector<std::pair<
std::string,
std::function<void(size_t, size_t, size_t, float const*, float const*,
size_t, float const*, size_t, float const*, float*,
size_t, cudaStream_t)>>> const
gemm_kernel_launch_functions{
{"Custom GEMM Kernel V00", launch_gemm_kernel_v00<float>},
{"Custom GEMM Kernel V01", launch_gemm_kernel_v01<float>},
{"Custom GEMM Kernel V02", launch_gemm_kernel_v02<float>},
{"Custom GEMM Kernel V02 Vectorized",
launch_gemm_kernel_v02_vectorized<float>},
{"Custom GEMM Kernel V03", launch_gemm_kernel_v03<float>},
{"Custom GEMM Kernel V03 Vectorized",
launch_gemm_kernel_v03_vectorized<float>},
{"Custom GEMM Kernel V04", launch_gemm_kernel_v04<float>},
{"Custom GEMM Kernel V04 Vectorized",
launch_gemm_kernel_v04_vectorized<float>},
{"Custom GEMM Kernel V05", launch_gemm_kernel_v05<float>},
{"Custom GEMM Kernel V05 Vectorized",
launch_gemm_kernel_v05_vectorized<float>},
{"Custom GEMM Kernel V06", launch_gemm_kernel_v06<float>},
{"Custom GEMM Kernel V06 Vectorized",
launch_gemm_kernel_v06_vectorized<float>},
{"Custom GEMM Kernel V06 Vectorized Double Buffered",
launch_gemm_kernel_v06_vectorized_double_buffered<float>},
};
for (auto const& gemm_kernel_launch_function : gemm_kernel_launch_functions)
{
std::cout << gemm_kernel_launch_function.first << std::endl;
std::pair<float, float> const gemm_kernel_profile_result{
profile_gemm<float>(
m, n, k, lda, ldb, ldc, gemm_kernel_launch_function.second,
fp32_abs_tol, fp32_rel_tol, num_repeats, num_warmups)};
std::cout << std::endl;
}
return 0;
}

View File

@@ -0,0 +1,400 @@
#ifndef PROFILE_UTILS_CUH
#define PROFILE_UTILS_CUH
#include <cassert>
#include <cmath>
#include <functional>
#include <iostream>
#include <random>
#include "cuda_gemm.hpp"
#include "cuda_gemm_utils.cuh"
#include <cublas_v2.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
template <typename T>
float measure_performance(std::function<T(cudaStream_t)> bound_function,
cudaStream_t stream, size_t num_repeats = 100,
size_t num_warmups = 100)
{
cudaEvent_t start, stop;
float time;
CHECK_CUDA_ERROR(cudaEventCreate(&start));
CHECK_CUDA_ERROR(cudaEventCreate(&stop));
for (size_t i{0}; i < num_warmups; ++i)
{
bound_function(stream);
}
CHECK_CUDA_ERROR(cudaStreamSynchronize(stream));
CHECK_CUDA_ERROR(cudaEventRecord(start, stream));
for (size_t i{0}; i < num_repeats; ++i)
{
bound_function(stream);
}
CHECK_CUDA_ERROR(cudaEventRecord(stop, stream));
CHECK_CUDA_ERROR(cudaEventSynchronize(stop));
CHECK_LAST_CUDA_ERROR();
CHECK_CUDA_ERROR(cudaEventElapsedTime(&time, start, stop));
CHECK_CUDA_ERROR(cudaEventDestroy(start));
CHECK_CUDA_ERROR(cudaEventDestroy(stop));
float const latency{time / num_repeats};
return latency;
}
#define CHECK_CUBLASS_ERROR(val) check_cublass((val), #val, __FILE__, __LINE__)
void check_cublass(cublasStatus_t err, const char* const func,
const char* const file, const int line)
{
if (err != CUBLAS_STATUS_SUCCESS)
{
std::cerr << "cuBLAS Error at: " << file << ":" << line << std::endl;
std::cerr << cublasGetStatusString(err) << std::endl;
std::exit(EXIT_FAILURE);
}
}
// Determine CUDA data type from type.
template <typename T,
typename std::enable_if<std::is_same<T, float>::value ||
std::is_same<T, double>::value ||
std::is_same<T, __half>::value,
bool>::type = true>
constexpr cudaDataType_t cuda_data_type_trait()
{
if (std::is_same<T, float>::value)
{
return CUDA_R_32F;
}
else if (std::is_same<T, double>::value)
{
return CUDA_R_64F;
}
else if (std::is_same<T, __half>::value)
{
return CUDA_R_16F;
}
else
{
throw std::runtime_error("Unsupported data type.");
}
}
template <typename T,
typename std::enable_if<std::is_same<T, float>::value ||
std::is_same<T, double>::value ||
std::is_same<T, __half>::value,
bool>::type = true>
void launch_gemm_cublas(size_t m, size_t n, size_t k, T const* alpha,
T const* A, size_t lda, T const* B, size_t ldb,
T const* beta, T* C, size_t ldc, cublasHandle_t handle)
{
// Non-TensorCore algorithm?
constexpr cublasGemmAlgo_t algo{CUBLAS_GEMM_DEFAULT};
constexpr cudaDataType_t data_type{cuda_data_type_trait<T>()};
// All the matrix are in row-major order.
// https://docs.nvidia.com/cuda/cublas/#cublasgemmex
// A: m x k row-major -> A: k x m column-major non-transposed
// B: k x n row-major -> B: n x k column-major non-transposed
// C: m x n row-major -> C: n x m column-major non-transposed
// Thus, without padding, the leading dimension of the matrix in row-major
// order is the number of columns, i.e., k for A, n for B, and n for C.
// Row-major order: C = AB + C
// Column-major order: C = BA + C
// The cuBLAS API requires the leading dimension of the matrix in
// column-major order. This API call looks non-intuitive, but it is correct.
CHECK_CUBLASS_ERROR(cublasGemmEx(
handle, CUBLAS_OP_N, CUBLAS_OP_N, n, m, k, alpha, B, data_type, ldb, A,
data_type, lda, beta, C, data_type, ldc, data_type, algo));
}
template <typename T,
typename std::enable_if<std::is_same<T, float>::value ||
std::is_same<T, double>::value,
bool>::type = true>
void launch_gemm_cpu(size_t m, size_t n, size_t k, T const* alpha, T const* A,
size_t lda, T const* B, size_t ldb, T const* beta, T* C,
size_t ldc)
{
// Compute GEMM using CPU.
for (size_t i{0U}; i < m; ++i)
{
for (size_t j{0U}; j < n; ++j)
{
T sum{static_cast<T>(0)};
for (size_t l{0U}; l < k; ++l)
{
sum += A[i * lda + l] * B[l * ldb + j];
}
C[i * ldc + j] = (*alpha) * sum + (*beta) * C[i * ldc + j];
}
}
}
// Many different implementations have been tried for FP16 GEMM on CPU.
// There is always a discrepancy between the results from CPU and GPU (cuBLAS or
// custom kernel).
template <typename T, typename std::enable_if<std::is_same<T, __half>::value,
bool>::type = true>
void launch_gemm_cpu(size_t m, size_t n, size_t k, T const* alpha, T const* A,
size_t lda, T const* B, size_t ldb, T const* beta, T* C,
size_t ldc)
{
// Compute GEMM using CPU.
for (size_t i{0U}; i < m; ++i)
{
for (size_t j{0U}; j < n; ++j)
{
float sum{0.0f};
for (size_t l{0U}; l < k; ++l)
{
sum += __half2float(__hmul(A[i * lda + l], B[l * ldb + j]));
}
C[i * ldc + j] = __float2half(__half2float(*alpha) * sum +
__half2float(*beta) *
__half2float(C[i * ldc + j]));
}
}
}
template <typename T>
bool all_close(T const* C, T const* C_ref, size_t m, size_t n, size_t ldc,
T abs_tol, double rel_tol)
{
bool status{true};
for (size_t i{0U}; i < m; ++i)
{
for (size_t j{0U}; j < n; ++j)
{
double const C_val{static_cast<double>(C[i * ldc + j])};
double const C_ref_val{static_cast<double>(C_ref[i * ldc + j])};
double const diff{C_val - C_ref_val};
double const diff_val{std::abs(diff)};
if (diff_val >
std::max(static_cast<double>(abs_tol),
static_cast<double>(std::abs(C_ref_val)) * rel_tol))
{
std::cout << "C[" << i << ", " << j << "] = " << C_val
<< " C_ref[" << i << ", " << j << "] = " << C_ref_val
<< " Abs Diff: " << diff_val
<< " Abs Diff Threshold: "
<< static_cast<double>(abs_tol)
<< " Rel->Abs Diff Threshold: "
<< static_cast<double>(
static_cast<double>(std::abs(C_ref_val)) *
rel_tol)
<< std::endl;
status = false;
return status;
}
}
}
return status;
}
void print_device_info()
{
int device_id{0};
cudaGetDevice(&device_id);
cudaDeviceProp device_prop;
cudaGetDeviceProperties(&device_prop, device_id);
std::cout << "Device Name: " << device_prop.name << std::endl;
float const memory_size{static_cast<float>(device_prop.totalGlobalMem) /
(1 << 30)};
std::cout << "Memory Size: " << memory_size << " GB" << std::endl;
float const peak_bandwidth{
static_cast<float>(2.0f * device_prop.memoryClockRate *
(device_prop.memoryBusWidth / 8) / 1.0e6)};
std::cout << "Peak Bandwitdh: " << peak_bandwidth << " GB/s" << std::endl;
std::cout << std::endl;
}
template <typename T>
float compute_effective_bandwidth(size_t m, size_t n, size_t k, float latency)
{
return ((m * k + k * n + m * n) * sizeof(T)) / (latency * 1e-3) / 1e9;
}
float compute_effective_tflops(size_t m, size_t n, size_t k, float latency)
{
return (2.0 * m * k * n) / (latency * 1e-3) / 1e12;
}
template <typename T,
typename std::enable_if<std::is_same<T, float>::value ||
std::is_same<T, double>::value ||
std::is_same<T, __half>::value,
bool>::type = true>
void random_initialize_matrix(T* A, size_t m, size_t n, size_t lda,
unsigned int seed = 0U)
{
std::default_random_engine eng(seed);
// The best way to verify is to use integer values.
std::uniform_int_distribution<int> dis(0, 5);
// std::uniform_real_distribution<float> dis(-1.0f, 1.0f);
auto const rand = [&dis, &eng]() { return dis(eng); };
for (size_t i{0U}; i < m; ++i)
{
for (size_t j{0U}; j < n; ++j)
{
A[i * lda + j] = static_cast<T>(rand());
}
}
}
void print_performance_result(size_t m, size_t n, size_t k, float latency)
{
float const effective_bandwidth{
compute_effective_bandwidth<float>(m, n, k, latency)};
float const effective_tflops{compute_effective_tflops(m, n, k, latency)};
std::cout << "Latency: " << latency << " ms" << std::endl;
std::cout << "Effective Bandwidth: " << effective_bandwidth << " GB/s"
<< std::endl;
std::cout << "Effective TFLOPS: " << effective_tflops << " TFLOPS"
<< std::endl;
}
template <typename T,
typename std::enable_if<std::is_same<T, float>::value ||
std::is_same<T, double>::value ||
std::is_same<T, __half>::value,
bool>::type = true>
std::pair<float, float> profile_gemm(
size_t m, size_t n, size_t k, size_t lda, size_t ldb, size_t ldc,
std::function<void(size_t, size_t, size_t, T const*, T const*, size_t,
T const*, size_t, T const*, T*, size_t, cudaStream_t)>
gemm_kernel_launch_function,
T abs_tol, double rel_tol, size_t num_repeats = 10, size_t num_warmups = 10,
unsigned int seed = 0U)
{
T const alpha{static_cast<T>(1.0)};
T const beta{static_cast<T>(0.0)};
// Create CUDA stream.
cudaStream_t stream;
CHECK_CUDA_ERROR(cudaStreamCreate(&stream));
// Allocate memory on host.
T* A_host{nullptr};
T* B_host{nullptr};
T* C_host{nullptr};
T* C_host_ref{nullptr};
T* C_host_from_device{nullptr};
CHECK_CUDA_ERROR(cudaMallocHost(&A_host, m * lda * sizeof(T)));
CHECK_CUDA_ERROR(cudaMallocHost(&B_host, k * ldb * sizeof(T)));
CHECK_CUDA_ERROR(cudaMallocHost(&C_host, m * ldc * sizeof(T)));
CHECK_CUDA_ERROR(cudaMallocHost(&C_host_ref, m * ldc * sizeof(T)));
CHECK_CUDA_ERROR(cudaMallocHost(&C_host_from_device, m * ldc * sizeof(T)));
// Initialize matrix A and B.
random_initialize_matrix(A_host, m, k, lda);
random_initialize_matrix(B_host, k, n, ldb);
random_initialize_matrix(C_host, m, n, ldc);
// Allocate memory on device.
T* A_device{nullptr};
T* B_device{nullptr};
T* C_device{nullptr};
CHECK_CUDA_ERROR(cudaMalloc(&A_device, m * lda * sizeof(T)));
CHECK_CUDA_ERROR(cudaMalloc(&B_device, k * ldb * sizeof(T)));
CHECK_CUDA_ERROR(cudaMalloc(&C_device, m * ldc * sizeof(T)));
// Copy matrix A and B from host to device.
CHECK_CUDA_ERROR(cudaMemcpy(A_device, A_host, m * lda * sizeof(T),
cudaMemcpyHostToDevice));
CHECK_CUDA_ERROR(cudaMemcpy(B_device, B_host, k * ldb * sizeof(T),
cudaMemcpyHostToDevice));
CHECK_CUDA_ERROR(cudaMemcpy(C_device, C_host, m * ldc * sizeof(T),
cudaMemcpyHostToDevice));
CHECK_CUDA_ERROR(cudaMemcpy(C_host_ref, C_host, m * ldc * sizeof(T),
cudaMemcpyHostToHost));
// Create cuBLAS handle.
cublasHandle_t handle;
CHECK_CUBLASS_ERROR(cublasCreate(&handle));
CHECK_CUBLASS_ERROR(cublasSetStream(handle, stream));
// Compute reference output using cuBLAS.
launch_gemm_cublas<T>(m, n, k, &alpha, A_device, lda, B_device, ldb, &beta,
C_device, ldc, handle);
CHECK_CUDA_ERROR(cudaStreamSynchronize(stream));
// Copy matrix C from device to host.
CHECK_CUDA_ERROR(cudaMemcpy(C_host_ref, C_device, m * ldc * sizeof(T),
cudaMemcpyDeviceToHost));
// // Compute reference output using CPU.
// std::cout << "Computing reference output using CPU..." << std::endl;
// launch_gemm_cpu<T>(m, n, k, &alpha, A_host, lda, B_host, ldb, &beta,
// C_host_ref, ldc);
// std::cout << "Done." << std::endl;
// Launch CUDA GEMM.
CHECK_CUDA_ERROR(cudaMemcpy(C_device, C_host, m * ldc * sizeof(T),
cudaMemcpyHostToDevice));
// Verify the correctness of CUDA GEMM.
gemm_kernel_launch_function(m, n, k, &alpha, A_device, lda, B_device, ldb,
&beta, C_device, ldc, stream);
// launch_gemm_cublas<T>(m, n, k, &alpha, A_device, lda, B_device, ldb,
// &beta,
// C_device, ldc, handle);
CHECK_CUDA_ERROR(cudaStreamSynchronize(stream));
CHECK_CUDA_ERROR(cudaMemcpy(C_host_from_device, C_device,
m * ldc * sizeof(T), cudaMemcpyDeviceToHost));
assert(all_close<T>(C_host_from_device, C_host_ref, m, n, ldc, abs_tol,
rel_tol));
// Launch cuBLAS GEMM.
float const latency_cublas{measure_performance<void>(
[&](cudaStream_t stream)
{
launch_gemm_cublas<T>(m, n, k, &alpha, A_device, lda, B_device, ldb,
&beta, C_device, ldc, handle);
return;
},
stream, num_repeats, num_warmups)};
float const latency_cuda_gemm{measure_performance<void>(
[&](cudaStream_t stream)
{
gemm_kernel_launch_function(m, n, k, &alpha, A_device, lda,
B_device, ldb, &beta, C_device, ldc,
stream);
return;
},
stream, num_repeats, num_warmups)};
// Release resources.
CHECK_CUDA_ERROR(cudaFree(A_device));
CHECK_CUDA_ERROR(cudaFree(B_device));
CHECK_CUDA_ERROR(cudaFree(C_device));
CHECK_CUDA_ERROR(cudaFreeHost(A_host));
CHECK_CUDA_ERROR(cudaFreeHost(B_host));
CHECK_CUDA_ERROR(cudaFreeHost(C_host));
CHECK_CUDA_ERROR(cudaFreeHost(C_host_ref));
CHECK_CUDA_ERROR(cudaFreeHost(C_host_from_device));
CHECK_CUBLASS_ERROR(cublasDestroy(handle));
CHECK_CUDA_ERROR(cudaStreamDestroy(stream));
std::cout << "cuBLAS GEMM Kernel Performance" << std::endl;
print_performance_result(m, n, k, latency_cublas);
std::cout << "Custom GEMM Kernel Performance" << std::endl;
print_performance_result(m, n, k, latency_cuda_gemm);
std::cout << "Custom GEMM VS cuBLAS GEMM Performance: "
<< latency_cublas / latency_cuda_gemm * 100.0f << "%"
<< std::endl;
return std::pair<float, float>{latency_cublas, latency_cuda_gemm};
}
#endif // PROFILE_UTILS_CUH