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:
@@ -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);
|
||||
@@ -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);
|
||||
112
upstream_ref/cuda_gemm_optimization/02_2d_block_tiling.cu
Normal file
112
upstream_ref/cuda_gemm_optimization/02_2d_block_tiling.cu
Normal 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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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*>(
|
||||
®ister_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);
|
||||
@@ -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*>(
|
||||
®ister_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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
98
upstream_ref/cuda_gemm_optimization/cuda_gemm.hpp
Normal file
98
upstream_ref/cuda_gemm_optimization/cuda_gemm.hpp
Normal 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
|
||||
29
upstream_ref/cuda_gemm_optimization/cuda_gemm_utils.cu
Normal file
29
upstream_ref/cuda_gemm_optimization/cuda_gemm_utils.cu
Normal 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);
|
||||
}
|
||||
}
|
||||
486
upstream_ref/cuda_gemm_optimization/cuda_gemm_utils.cuh
Normal file
486
upstream_ref/cuda_gemm_optimization/cuda_gemm_utils.cuh
Normal 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
|
||||
13
upstream_ref/cuda_gemm_optimization/cuda_gemm_utils.hpp
Normal file
13
upstream_ref/cuda_gemm_optimization/cuda_gemm_utils.hpp
Normal 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
|
||||
109
upstream_ref/cuda_gemm_optimization/profile_cuda_gemm_fp16.cu
Normal file
109
upstream_ref/cuda_gemm_optimization/profile_cuda_gemm_fp16.cu
Normal 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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
400
upstream_ref/cuda_gemm_optimization/profile_utils.cuh
Normal file
400
upstream_ref/cuda_gemm_optimization/profile_utils.cuh
Normal 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
|
||||
Reference in New Issue
Block a user