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:
187
upstream_ref/sgemm_cuda/10_kernel_warptiling.cuh
Normal file
187
upstream_ref/sgemm_cuda/10_kernel_warptiling.cuh
Normal file
@@ -0,0 +1,187 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
const int WARPSIZE = 32; // warpSize is not constexpr
|
||||
|
||||
namespace wt {
|
||||
template <const int BM, const int BN, const int BK, const int rowStrideA,
|
||||
const int rowStrideB>
|
||||
__device__ void loadFromGmem(int N, int K, const float *A, const float *B,
|
||||
float *As, float *Bs, int innerRowA, int innerColA,
|
||||
int innerRowB, int innerColB) {
|
||||
for (uint offset = 0; offset + rowStrideA <= BM; offset += rowStrideA) {
|
||||
const float4 tmp = reinterpret_cast<const float4 *>(
|
||||
&A[(innerRowA + offset) * K + innerColA * 4])[0];
|
||||
// float4 tmp;
|
||||
// asm("ld.global.nc.v4.f32 {%0, %1, %2, %3}, [%4];"
|
||||
// : "=f"(tmp.x), "=f"(tmp.y), "=f"(tmp.z), "=f"(tmp.w)
|
||||
// : "l"(&A[(innerRowA + offset) * K + innerColA * 4]));
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA + offset] = tmp.x;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA + offset] = tmp.y;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA + offset] = tmp.z;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA + offset] = tmp.w;
|
||||
}
|
||||
|
||||
for (uint offset = 0; offset + rowStrideB <= BK; offset += rowStrideB) {
|
||||
reinterpret_cast<float4 *>(
|
||||
&Bs[(innerRowB + offset) * BN + innerColB * 4])[0] =
|
||||
reinterpret_cast<const float4 *>(
|
||||
&B[(innerRowB + offset) * N + innerColB * 4])[0];
|
||||
// asm("ld.global.v4.f32 {%0, %1, %2, %3}, [%4];"
|
||||
// : "=f"(Bs[(innerRowB + offset) * BN + innerColB * 4 + 0]),
|
||||
// "=f"(Bs[(innerRowB + offset) * BN + innerColB * 4 + 1]),
|
||||
// "=f"(Bs[(innerRowB + offset) * BN + innerColB * 4 + 2]),
|
||||
// "=f"(Bs[(innerRowB + offset) * BN + innerColB * 4 + 3])
|
||||
// : "l"(&B[(innerRowB + offset) * N + innerColB * 4]));
|
||||
}
|
||||
}
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int WM, const int WN,
|
||||
const int WMITER, const int WNITER, const int WSUBM, const int WSUBN,
|
||||
const int TM, const int TN>
|
||||
__device__ void
|
||||
processFromSmem(float *regM, float *regN, float *threadResults, const float *As,
|
||||
const float *Bs, const uint warpRow, const uint warpCol,
|
||||
const uint threadRowInWarp, const uint threadColInWarp) {
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// populate registers for whole warptile
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[wSubRowIdx * TM + i] =
|
||||
As[(dotIdx * BM) + warpRow * WM + wSubRowIdx * WSUBM +
|
||||
threadRowInWarp * TM + i];
|
||||
}
|
||||
}
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[wSubColIdx * TN + i] =
|
||||
Bs[(dotIdx * BN) + warpCol * WN + wSubColIdx * WSUBN +
|
||||
threadColInWarp * TN + i];
|
||||
}
|
||||
}
|
||||
|
||||
// execute warptile matmul
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
// calculate per-thread results
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[(wSubRowIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
(wSubColIdx * TN) + resIdxN] +=
|
||||
regM[wSubRowIdx * TM + resIdxM] *
|
||||
regN[wSubColIdx * TN + resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace wt
|
||||
|
||||
/*
|
||||
* @tparam BM The threadblock size for M dimension SMEM caching.
|
||||
* @tparam BN The threadblock size for N dimension SMEM caching.
|
||||
* @tparam BK The threadblock size for K dimension SMEM caching.
|
||||
* @tparam WM M dim of continuous tile computed by each warp
|
||||
* @tparam WN N dim of continuous tile computed by each warp
|
||||
* @tparam WMITER The number of subwarp tiling steps in M dimension.
|
||||
* @tparam WNITER The number of subwarp tiling steps in N dimension.
|
||||
* @tparam TM The per-thread tile size for M dimension.
|
||||
* @tparam TN The per-thread tile size for N dimension.
|
||||
*/
|
||||
template <const int BM, const int BN, const int BK, const int WM, const int WN,
|
||||
const int WNITER, const int TM, const int TN, const int NUM_THREADS>
|
||||
__global__ void __launch_bounds__(NUM_THREADS)
|
||||
sgemmWarptiling(int M, int N, int K, float alpha, float *A, float *B,
|
||||
float beta, float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// Placement of the warp in the threadblock tile
|
||||
const uint warpIdx = threadIdx.x / WARPSIZE; // the warp this thread is in
|
||||
const uint warpCol = warpIdx % (BN / WN);
|
||||
const uint warpRow = warpIdx / (BN / WN);
|
||||
|
||||
// size of the warp subtile
|
||||
constexpr uint WMITER = (WM * WN) / (WARPSIZE * TM * TN * WNITER);
|
||||
constexpr uint WSUBM = WM / WMITER; // 64/2=32
|
||||
constexpr uint WSUBN = WN / WNITER; // 32/2=16
|
||||
|
||||
// Placement of the thread in the warp subtile
|
||||
const uint threadIdxInWarp = threadIdx.x % WARPSIZE; // [0, 31]
|
||||
const uint threadColInWarp = threadIdxInWarp % (WSUBN / TN); // i%(16/4)
|
||||
const uint threadRowInWarp = threadIdxInWarp / (WSUBN / TN); // i/4
|
||||
|
||||
// allocate space for the current blocktile in SMEM
|
||||
__shared__ float As[BM * BK];
|
||||
__shared__ float Bs[BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
// Move C_ptr to warp's output tile
|
||||
C += (cRow * BM + warpRow * WM) * N + cCol * BN + warpCol * WN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// we'll load 128bit / 32bit = 4 elements per thread at each step
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
constexpr uint rowStrideA = (NUM_THREADS * 4) / BK;
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
constexpr uint rowStrideB = NUM_THREADS / (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[WMITER * TM * WNITER * TN] = {0.0};
|
||||
// we cache into registers on the warptile level
|
||||
float regM[WMITER * TM] = {0.0};
|
||||
float regN[WNITER * TN] = {0.0};
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
wt::loadFromGmem<BM, BN, BK, rowStrideA, rowStrideB>(
|
||||
N, K, A, B, As, Bs, innerRowA, innerColA, innerRowB, innerColB);
|
||||
__syncthreads();
|
||||
wt::processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN, TM,
|
||||
TN>(regM, regN, threadResults, As, Bs, warpRow, warpCol,
|
||||
threadRowInWarp, threadColInWarp);
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
// move C pointer to current warp subtile
|
||||
float *C_interim = C + (wSubRowIdx * WSUBM) * N + wSubColIdx * WSUBN;
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRowInWarp * TM + resIdxM) * N +
|
||||
threadColInWarp * TN + resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
const int i = (wSubRowIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
wSubColIdx * TN + resIdxN;
|
||||
tmp.x = alpha * threadResults[i + 0] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[i + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[i + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[i + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRowInWarp * TM + resIdxM) * N +
|
||||
threadColInWarp * TN + resIdxN])[0] = tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
220
upstream_ref/sgemm_cuda/11_kernel_double_buffering.cuh
Normal file
220
upstream_ref/sgemm_cuda/11_kernel_double_buffering.cuh
Normal file
@@ -0,0 +1,220 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
namespace db {
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int rowStrideA,
|
||||
const int rowStrideB>
|
||||
__device__ void loadFromGmem(const int N, const int K, float *A, float *B,
|
||||
float *As, float *Bs, const int innerRowA,
|
||||
const int innerColA, const int innerRowB,
|
||||
const int innerColB) {
|
||||
for (uint offset = 0; offset + rowStrideA <= BM; offset += rowStrideA) {
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&A[(innerRowA + offset) * K + innerColA * 4])[0];
|
||||
// transpose A while storing it
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA + offset] = tmp.x;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA + offset] = tmp.y;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA + offset] = tmp.z;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA + offset] = tmp.w;
|
||||
}
|
||||
|
||||
for (uint offset = 0; offset + rowStrideB <= BK; offset += rowStrideB) {
|
||||
reinterpret_cast<float4 *>(
|
||||
&Bs[(innerRowB + offset) * BN + innerColB * 4])[0] =
|
||||
reinterpret_cast<float4 *>(
|
||||
&B[(innerRowB + offset) * N + innerColB * 4])[0];
|
||||
}
|
||||
}
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int WM, const int WN,
|
||||
const int WMITER, const int WNITER, const int WSUBM, const int WSUBN,
|
||||
const int TM, const int TN>
|
||||
__device__ void
|
||||
processFromSmem(float *regM, float *regN, float *threadResults, const float *As,
|
||||
const float *Bs, const uint warpRow, const uint warpCol,
|
||||
const uint threadRowInWarp, const uint threadColInWarp) {
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// populate registers for whole warptile
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[wSubRowIdx * TM + i] =
|
||||
As[(dotIdx * BM) + warpRow * WM + wSubRowIdx * WSUBM +
|
||||
threadRowInWarp * TM + i];
|
||||
}
|
||||
}
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[wSubColIdx * TN + i] =
|
||||
Bs[(dotIdx * BN) + warpCol * WN + wSubColIdx * WSUBN +
|
||||
threadColInWarp * TN + i];
|
||||
}
|
||||
}
|
||||
|
||||
// execute warptile matmul
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
// calculate per-thread results
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[(wSubRowIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
(wSubColIdx * TN) + resIdxN] +=
|
||||
regM[wSubRowIdx * TM + resIdxM] *
|
||||
regN[wSubColIdx * TN + resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace db
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int WM, const int WN,
|
||||
const int WNITER, const int TM, const int TN, const int NUM_THREADS>
|
||||
__global__ void __launch_bounds__(NUM_THREADS)
|
||||
sgemmDoubleBuffering(const int M, const int N, const int K,
|
||||
const float alpha, float *A, float *B, float beta,
|
||||
float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// Placement of the warp in the threadblock tile
|
||||
const uint warpIdx = threadIdx.x / WARPSIZE; // the warp this thread is in
|
||||
const uint warpCol = warpIdx % (BN / WN);
|
||||
const uint warpRow = warpIdx / (BN / WN);
|
||||
|
||||
// size of the warp subtile
|
||||
constexpr uint WMITER = (WM * WN) / (WARPSIZE * TM * TN * WNITER);
|
||||
constexpr uint WSUBM = WM / WMITER; // 64/2=32
|
||||
constexpr uint WSUBN = WN / WNITER; // 32/2=16
|
||||
|
||||
// Placement of the thread in the warp subtile
|
||||
const uint threadIdxInWarp = threadIdx.x % WARPSIZE; // [0, 31]
|
||||
const uint threadColInWarp = threadIdxInWarp % (WSUBN / TN); // i%(16/4)
|
||||
const uint threadRowInWarp = threadIdxInWarp / (WSUBN / TN); // i/4
|
||||
|
||||
// allocate space for the current blocktile in SMEM
|
||||
__shared__ float As[2 * BM * BK];
|
||||
__shared__ float Bs[2 * BK * BN];
|
||||
|
||||
// setup double buffering split
|
||||
bool doubleBufferIdx = threadIdx.x >= (NUM_THREADS / 2);
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
// Move C_ptr to warp's output tile
|
||||
C += (cRow * BM + warpRow * WM) * N + cCol * BN + warpCol * WN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// for the loading, we're pretending like there's half as many threads
|
||||
// as there actually are
|
||||
const uint innerRowA = (threadIdx.x % (NUM_THREADS / 2)) / (BK / 4);
|
||||
const uint innerColA = (threadIdx.x % (NUM_THREADS / 2)) % (BK / 4);
|
||||
constexpr uint rowStrideA = ((NUM_THREADS / 2) * 4) / BK;
|
||||
const uint innerRowB = (threadIdx.x % (NUM_THREADS / 2)) / (BN / 4);
|
||||
const uint innerColB = (threadIdx.x % (NUM_THREADS / 2)) % (BN / 4);
|
||||
constexpr uint rowStrideB = (NUM_THREADS / 2) / (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[WMITER * TM * WNITER * TN] = {0.0};
|
||||
// we cache into registers on the warptile level
|
||||
float regM[WMITER * TM] = {0.0};
|
||||
float regN[WNITER * TN] = {0.0};
|
||||
|
||||
if (doubleBufferIdx == 0) {
|
||||
// load first (B0)
|
||||
db::loadFromGmem<BM, BN, BK, rowStrideA, rowStrideB>(
|
||||
N, K, A, B, As, Bs, innerRowA, innerColA, innerRowB, innerColB);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += 2 * BK) {
|
||||
if (doubleBufferIdx == 0) {
|
||||
// process current (B0)
|
||||
db::processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN, TM,
|
||||
TN>(regM, regN, threadResults, As, Bs, warpRow,
|
||||
warpCol, threadRowInWarp, threadColInWarp);
|
||||
__syncthreads();
|
||||
|
||||
// process current+1 (B1)
|
||||
if (bkIdx + BK < K) {
|
||||
db::processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN,
|
||||
TM, TN>(regM, regN, threadResults, As + (BM * BK),
|
||||
Bs + (BK * BN), warpRow, warpCol,
|
||||
threadRowInWarp, threadColInWarp);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// load current + 2 (B0)
|
||||
if (bkIdx + 2 * BK < K) {
|
||||
db::loadFromGmem<BM, BN, BK, rowStrideA, rowStrideB>(
|
||||
N, K, A + 2 * BK, B + 2 * BK * N, As, Bs, innerRowA, innerColA,
|
||||
innerRowB, innerColB);
|
||||
}
|
||||
} else {
|
||||
// load current + 1 (B1)
|
||||
if (bkIdx + BK < K) {
|
||||
db::loadFromGmem<BM, BN, BK, rowStrideA, rowStrideB>(
|
||||
N, K, A + BK, B + BK * N, As + (BM * BK), Bs + (BK * BN), innerRowA,
|
||||
innerColA, innerRowB, innerColB);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// process current (B0)
|
||||
db::processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN, TM,
|
||||
TN>(regM, regN, threadResults, As, Bs, warpRow,
|
||||
warpCol, threadRowInWarp, threadColInWarp);
|
||||
__syncthreads();
|
||||
|
||||
// process current+1 (B1)
|
||||
if (bkIdx + BK < K) {
|
||||
db::processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN,
|
||||
TM, TN>(regM, regN, threadResults, As + (BM * BK),
|
||||
Bs + (BK * BN), warpRow, warpCol,
|
||||
threadRowInWarp, threadColInWarp);
|
||||
}
|
||||
}
|
||||
|
||||
A += 2 * BK; // move BK columns to right
|
||||
B += 2 * BK * N; // move BK rows down
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
// move C pointer to current warp subtile
|
||||
float *C_interim = C + (wSubRowIdx * WSUBM) * N + wSubColIdx * WSUBN;
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRowInWarp * TM + resIdxM) * N +
|
||||
threadColInWarp * TN + resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
const int i = (wSubRowIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
wSubColIdx * TN + resIdxN;
|
||||
tmp.x = alpha * threadResults[i + 0] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[i + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[i + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[i + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRowInWarp * TM + resIdxM) * N +
|
||||
threadColInWarp * TN + resIdxN])[0] = tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
229
upstream_ref/sgemm_cuda/12_kernel_double_buffering.cuh
Normal file
229
upstream_ref/sgemm_cuda/12_kernel_double_buffering.cuh
Normal file
@@ -0,0 +1,229 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cooperative_groups.h>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda/barrier>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
namespace {
|
||||
template <const int BM, const int BN, const int BK, const int rowStrideA,
|
||||
const int rowStrideB, typename T>
|
||||
__device__ void loadFromGmem(int N, int K, float *A, float *B, float *As,
|
||||
float *Bs, int innerRowA, int innerColA,
|
||||
int innerRowB, int innerColB, T &barrier) {
|
||||
|
||||
for (uint offset = 0; offset + rowStrideA <= BM; offset += rowStrideA) {
|
||||
cuda::memcpy_async(&As[(innerColA * 4 + 0) * BM + innerRowA + offset],
|
||||
&A[(innerRowA + offset) * K + innerColA * 4],
|
||||
cuda::aligned_size_t<sizeof(float)>(sizeof(float)),
|
||||
barrier);
|
||||
cuda::memcpy_async(&As[(innerColA * 4 + 1) * BM + innerRowA + offset],
|
||||
&A[(innerRowA + offset) * K + innerColA * 4 + 1],
|
||||
cuda::aligned_size_t<sizeof(float)>(sizeof(float)),
|
||||
barrier);
|
||||
cuda::memcpy_async(&As[(innerColA * 4 + 2) * BM + innerRowA + offset],
|
||||
&A[(innerRowA + offset) * K + innerColA * 4 + 2],
|
||||
cuda::aligned_size_t<sizeof(float)>(sizeof(float)),
|
||||
barrier);
|
||||
cuda::memcpy_async(&As[(innerColA * 4 + 3) * BM + innerRowA + offset],
|
||||
&A[(innerRowA + offset) * K + innerColA * 4 + 3],
|
||||
cuda::aligned_size_t<sizeof(float)>(sizeof(float)),
|
||||
barrier);
|
||||
}
|
||||
|
||||
for (uint offset = 0; offset + rowStrideB <= BK; offset += rowStrideB) {
|
||||
cuda::memcpy_async(&Bs[(innerRowB + offset) * BN + innerColB * 4],
|
||||
&B[(innerRowB + offset) * N + innerColB * 4],
|
||||
cuda::aligned_size_t<sizeof(float4)>(sizeof(float4)),
|
||||
barrier);
|
||||
}
|
||||
}
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int WM, const int WN,
|
||||
const int WMITER, const int WNITER, const int WSUBM, const int WSUBN,
|
||||
const int TM, const int TN>
|
||||
__device__ void
|
||||
processFromSmem(float *regM, float *regN, float *threadResults, const float *As,
|
||||
const float *Bs, const uint warpRow, const uint warpCol,
|
||||
const uint threadRowInWarp, const uint threadColInWarp) {
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// populate registers for whole warptile
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[wSubRowIdx * TM + i] =
|
||||
As[(dotIdx * BM) + warpRow * WM + wSubRowIdx * WSUBM +
|
||||
threadRowInWarp * TM + i];
|
||||
}
|
||||
}
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[wSubColIdx * TN + i] =
|
||||
Bs[(dotIdx * BN) + warpCol * WN + wSubColIdx * WSUBN +
|
||||
threadColInWarp * TN + i];
|
||||
}
|
||||
}
|
||||
|
||||
// execute warptile matmul
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
// calculate per-thread results
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[(wSubRowIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
(wSubColIdx * TN) + resIdxN] +=
|
||||
regM[wSubRowIdx * TM + resIdxM] *
|
||||
regN[wSubColIdx * TN + resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
/*
|
||||
* @tparam BM The threadblock size for M dimension SMEM caching.
|
||||
* @tparam BN The threadblock size for N dimension SMEM caching.
|
||||
* @tparam BK The threadblock size for K dimension SMEM caching.
|
||||
* @tparam WM M dim of continuous tile computed by each warp
|
||||
* @tparam WN N dim of continuous tile computed by each warp
|
||||
* @tparam WMITER The number of subwarp tiling steps in M dimension.
|
||||
* @tparam WNITER The number of subwarp tiling steps in N dimension.
|
||||
* @tparam TM The per-thread tile size for M dimension.
|
||||
* @tparam TN The per-thread tile size for N dimension.
|
||||
*/
|
||||
template <const int BM, const int BN, const int BK, const int WM, const int WN,
|
||||
const int WNITER, const int TM, const int TN, const int NUM_THREADS>
|
||||
__global__ void __launch_bounds__(NUM_THREADS)
|
||||
runSgemmDoubleBuffering2(int M, int N, int K, float alpha, float *A,
|
||||
float *B, float beta, float *C) {
|
||||
auto block = cooperative_groups::this_thread_block();
|
||||
__shared__ cuda::barrier<cuda::thread_scope::thread_scope_block> frontBarrier;
|
||||
__shared__ cuda::barrier<cuda::thread_scope::thread_scope_block> backBarrier;
|
||||
auto frontBarrierPtr = &frontBarrier;
|
||||
auto backBarrierPtr = &backBarrier;
|
||||
if (block.thread_rank() == 0) {
|
||||
init(&frontBarrier, block.size());
|
||||
init(&backBarrier, block.size());
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// Placement of the warp in the threadblock tile
|
||||
const uint warpIdx = threadIdx.x / WARPSIZE; // the warp this thread is in
|
||||
const uint warpCol = warpIdx % (BN / WN);
|
||||
const uint warpRow = warpIdx / (BN / WN);
|
||||
|
||||
// size of the warp subtile
|
||||
constexpr uint WMITER = (WM * WN) / (WARPSIZE * TM * TN * WNITER);
|
||||
constexpr uint WSUBM = WM / WMITER; // 64/2=32
|
||||
constexpr uint WSUBN = WN / WNITER; // 32/2=16
|
||||
|
||||
// Placement of the thread in the warp subtile
|
||||
const uint threadIdxInWarp = threadIdx.x % WARPSIZE; // [0, 31]
|
||||
const uint threadColInWarp = threadIdxInWarp % (WSUBN / TN); // i%(16/4)
|
||||
const uint threadRowInWarp = threadIdxInWarp / (WSUBN / TN); // i/4
|
||||
|
||||
// allocate space for the current blocktile in SMEM
|
||||
__shared__ float As[2 * BM * BK];
|
||||
__shared__ float Bs[2 * BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
// Move C_ptr to warp's output tile
|
||||
C += (cRow * BM + warpRow * WM) * N + cCol * BN + warpCol * WN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// we'll load 128bit / 32bit = 4 elements per thread at each step
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
constexpr uint rowStrideA = (NUM_THREADS * 4) / BK;
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
constexpr uint rowStrideB = NUM_THREADS / (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[WMITER * TM * WNITER * TN] = {0.0};
|
||||
// we cache into registers on the warptile level
|
||||
float regM[WMITER * TM] = {0.0};
|
||||
float regN[WNITER * TN] = {0.0};
|
||||
|
||||
int As_offset = 0;
|
||||
int Bs_offset = 0;
|
||||
|
||||
// double-buffering: load first blocktile into SMEM
|
||||
loadFromGmem<BM, BN, BK, rowStrideA, rowStrideB>(
|
||||
N, K, A, B, As + As_offset * BM * BK, Bs + Bs_offset * BK * BN, innerRowA,
|
||||
innerColA, innerRowB, innerColB, (*frontBarrierPtr));
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K - BK; bkIdx += BK) {
|
||||
// double-buffering: load next blocktile into SMEM
|
||||
loadFromGmem<BM, BN, BK, rowStrideA, rowStrideB>(
|
||||
N, K, A + BK, B + BK * N, As + (1 - As_offset) * BM * BK,
|
||||
Bs + (1 - Bs_offset) * BK * BN, innerRowA, innerColA, innerRowB,
|
||||
innerColB, (*backBarrierPtr));
|
||||
|
||||
// compute the current blocktile
|
||||
(*frontBarrierPtr).arrive_and_wait();
|
||||
processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN, TM, TN>(
|
||||
regM, regN, threadResults, As + As_offset * BM * BK,
|
||||
Bs + Bs_offset * BK * BN, warpRow, warpCol, threadRowInWarp,
|
||||
threadColInWarp);
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
|
||||
As_offset = 1 - As_offset;
|
||||
Bs_offset = 1 - Bs_offset;
|
||||
// swap the front and back barriers
|
||||
auto tmp = frontBarrierPtr;
|
||||
frontBarrierPtr = backBarrierPtr;
|
||||
backBarrierPtr = tmp;
|
||||
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// compute the last blocktile
|
||||
(*frontBarrierPtr).arrive_and_wait();
|
||||
processFromSmem<BM, BN, BK, WM, WN, WMITER, WNITER, WSUBM, WSUBN, TM, TN>(
|
||||
regM, regN, threadResults, As + As_offset * BM * BK,
|
||||
Bs + Bs_offset * BK * BN, warpRow, warpCol, threadRowInWarp,
|
||||
threadColInWarp);
|
||||
|
||||
// write out the results
|
||||
for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) {
|
||||
for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) {
|
||||
// move C pointer to current warp subtile
|
||||
float *C_interim = C + (wSubRowIdx * WSUBM) * N + wSubColIdx * WSUBN;
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRowInWarp * TM + resIdxM) * N +
|
||||
threadColInWarp * TN + resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
const int i = (wSubRowIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
wSubColIdx * TN + resIdxN;
|
||||
tmp.x = alpha * threadResults[i + 0] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[i + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[i + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[i + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRowInWarp * TM + resIdxM) * N +
|
||||
threadColInWarp * TN + resIdxN])[0] = tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
29
upstream_ref/sgemm_cuda/1_naive.cuh
Normal file
29
upstream_ref/sgemm_cuda/1_naive.cuh
Normal file
@@ -0,0 +1,29 @@
|
||||
#pragma once
|
||||
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
/*
|
||||
|
||||
Matrix sizes:
|
||||
MxK * KxN = MxN
|
||||
|
||||
*/
|
||||
|
||||
__global__ void sgemm_naive(int M, int N, int K, float alpha, const float *A,
|
||||
const float *B, float beta, float *C) {
|
||||
const uint x = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
const uint y = blockIdx.y * blockDim.y + threadIdx.y;
|
||||
|
||||
// if statement is necessary to make things work under tile quantization
|
||||
if (x < M && y < N) {
|
||||
float tmp = 0.0;
|
||||
for (int i = 0; i < K; ++i) {
|
||||
tmp += A[x * K + i] * B[i * N + y];
|
||||
}
|
||||
// C = α*(A@B)+β*C
|
||||
C[x * N + y] = alpha * tmp + beta * C[x * N + y];
|
||||
}
|
||||
}
|
||||
24
upstream_ref/sgemm_cuda/2_kernel_global_mem_coalesce.cuh
Normal file
24
upstream_ref/sgemm_cuda/2_kernel_global_mem_coalesce.cuh
Normal file
@@ -0,0 +1,24 @@
|
||||
#pragma once
|
||||
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
template <const uint BLOCKSIZE>
|
||||
__global__ void sgemm_global_mem_coalesce(int M, int N, int K, float alpha,
|
||||
const float *A, const float *B,
|
||||
float beta, float *C) {
|
||||
const int cRow = blockIdx.x * BLOCKSIZE + (threadIdx.x / BLOCKSIZE);
|
||||
const int cCol = blockIdx.y * BLOCKSIZE + (threadIdx.x % BLOCKSIZE);
|
||||
|
||||
// if statement is necessary to make things work under tile quantization
|
||||
if (cRow < M && cCol < N) {
|
||||
float tmp = 0.0;
|
||||
for (int i = 0; i < K; ++i) {
|
||||
tmp += A[cRow * K + i] * B[i * N + cCol];
|
||||
}
|
||||
C[cRow * N + cCol] = alpha * tmp + beta * C[cRow * N + cCol];
|
||||
}
|
||||
}
|
||||
57
upstream_ref/sgemm_cuda/3_kernel_shared_mem_blocking.cuh
Normal file
57
upstream_ref/sgemm_cuda/3_kernel_shared_mem_blocking.cuh
Normal file
@@ -0,0 +1,57 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
template <const int BLOCKSIZE>
|
||||
__global__ void sgemm_shared_mem_block(int M, int N, int K, float alpha,
|
||||
const float *A, const float *B,
|
||||
float beta, float *C) {
|
||||
// the output block that we want to compute in this threadblock
|
||||
const uint cRow = blockIdx.x;
|
||||
const uint cCol = blockIdx.y;
|
||||
|
||||
// allocate buffer for current block in fast shared mem
|
||||
// shared mem is shared between all threads in a block
|
||||
__shared__ float As[BLOCKSIZE * BLOCKSIZE];
|
||||
__shared__ float Bs[BLOCKSIZE * BLOCKSIZE];
|
||||
|
||||
// the inner row & col that we're accessing in this thread
|
||||
const uint threadCol = threadIdx.x % BLOCKSIZE;
|
||||
const uint threadRow = threadIdx.x / BLOCKSIZE;
|
||||
|
||||
// advance pointers to the starting positions
|
||||
A += cRow * BLOCKSIZE * K; // row=cRow, col=0
|
||||
B += cCol * BLOCKSIZE; // row=0, col=cCol
|
||||
C += cRow * BLOCKSIZE * N + cCol * BLOCKSIZE; // row=cRow, col=cCol
|
||||
|
||||
float tmp = 0.0;
|
||||
for (int bkIdx = 0; bkIdx < K; bkIdx += BLOCKSIZE) {
|
||||
// Have each thread load one of the elements in A & B
|
||||
// Make the threadCol (=threadIdx.x) the consecutive index
|
||||
// to allow global memory access coalescing
|
||||
As[threadRow * BLOCKSIZE + threadCol] = A[threadRow * K + threadCol];
|
||||
Bs[threadRow * BLOCKSIZE + threadCol] = B[threadRow * N + threadCol];
|
||||
|
||||
// block threads in this block until cache is fully populated
|
||||
__syncthreads();
|
||||
A += BLOCKSIZE;
|
||||
B += BLOCKSIZE * N;
|
||||
|
||||
// execute the dotproduct on the currently cached block
|
||||
for (int dotIdx = 0; dotIdx < BLOCKSIZE; ++dotIdx) {
|
||||
tmp += As[threadRow * BLOCKSIZE + dotIdx] *
|
||||
Bs[dotIdx * BLOCKSIZE + threadCol];
|
||||
}
|
||||
// need to sync again at the end, to avoid faster threads
|
||||
// fetching the next block into the cache before slower threads are done
|
||||
__syncthreads();
|
||||
}
|
||||
C[threadRow * N + threadCol] =
|
||||
alpha * tmp + beta * C[threadRow * N + threadCol];
|
||||
}
|
||||
80
upstream_ref/sgemm_cuda/4_kernel_1D_blocktiling.cuh
Normal file
80
upstream_ref/sgemm_cuda/4_kernel_1D_blocktiling.cuh
Normal file
@@ -0,0 +1,80 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int TM>
|
||||
__global__ void sgemm1DBlocktiling(int M, int N, int K, float alpha,
|
||||
const float *A, const float *B, float beta,
|
||||
float *C) {
|
||||
// If we flip x and y here we get ~30% less performance for large matrices.
|
||||
// The current, 30% faster configuration ensures that blocks with sequential
|
||||
// blockIDs access columns of B sequentially, while sharing the same row of A.
|
||||
// The slower configuration would share columns of A, but access into B would
|
||||
// be non-sequential. So the faster configuration has better spatial locality
|
||||
// and hence a greater L2 hit rate.
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// each warp will calculate 32*TM elements, with 32 being the columnar dim.
|
||||
const int threadCol = threadIdx.x % BN;
|
||||
const int threadRow = threadIdx.x / BN;
|
||||
|
||||
// allocate space for the current blocktile in SMEM
|
||||
__shared__ float As[BM * BK];
|
||||
__shared__ float Bs[BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// todo: adjust this to each thread to load multiple entries and
|
||||
// better exploit the cache sizes
|
||||
assert(BM * BK == blockDim.x);
|
||||
assert(BN * BK == blockDim.x);
|
||||
const uint innerColA = threadIdx.x % BK; // warp-level GMEM coalescing
|
||||
const uint innerRowA = threadIdx.x / BK;
|
||||
const uint innerColB = threadIdx.x % BN; // warp-level GMEM coalescing
|
||||
const uint innerRowB = threadIdx.x / BN;
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[TM] = {0.0};
|
||||
|
||||
// outer loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
As[innerRowA * BK + innerColA] = A[innerRowA * K + innerColA];
|
||||
Bs[innerRowB * BN + innerColB] = B[innerRowB * N + innerColB];
|
||||
__syncthreads();
|
||||
|
||||
// advance blocktile
|
||||
A += BK;
|
||||
B += BK * N;
|
||||
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// we make the dotproduct loop the outside loop, which facilitates
|
||||
// reuse of the Bs entry, which we can cache in a tmp var.
|
||||
float tmpB = Bs[dotIdx * BN + threadCol];
|
||||
for (uint resIdx = 0; resIdx < TM; ++resIdx) {
|
||||
threadResults[resIdx] +=
|
||||
As[(threadRow * TM + resIdx) * BK + dotIdx] * tmpB;
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint resIdx = 0; resIdx < TM; ++resIdx) {
|
||||
C[(threadRow * TM + resIdx) * N + threadCol] =
|
||||
alpha * threadResults[resIdx] +
|
||||
beta * C[(threadRow * TM + resIdx) * N + threadCol];
|
||||
}
|
||||
}
|
||||
102
upstream_ref/sgemm_cuda/5_kernel_2D_blocktiling.cuh
Normal file
102
upstream_ref/sgemm_cuda/5_kernel_2D_blocktiling.cuh
Normal file
@@ -0,0 +1,102 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void __launch_bounds__((BM * BN) / (TM * TN), 1)
|
||||
sgemm2DBlocktiling(int M, int N, int K, float alpha, const float *A,
|
||||
const float *B, float beta, float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
const uint totalResultsBlocktile = BM * BN;
|
||||
// A thread is responsible for calculating TM*TN elements in the blocktile
|
||||
const uint numThreadsBlocktile = totalResultsBlocktile / (TM * TN);
|
||||
|
||||
// ResultsPerBlock / ResultsPerThread == ThreadsPerBlock
|
||||
assert(numThreadsBlocktile == blockDim.x);
|
||||
|
||||
// BN/TN are the number of threads to span a column
|
||||
const int threadCol = threadIdx.x % (BN / TN);
|
||||
const int threadRow = threadIdx.x / (BN / TN);
|
||||
|
||||
// allocate space for the current blocktile in smem
|
||||
__shared__ float As[BM * BK];
|
||||
__shared__ float Bs[BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
const uint innerRowA = threadIdx.x / BK;
|
||||
const uint innerColA = threadIdx.x % BK;
|
||||
// calculates the number of rows of As that are being loaded in a single step
|
||||
// by a single block
|
||||
const uint strideA = numThreadsBlocktile / BK;
|
||||
const uint innerRowB = threadIdx.x / BN;
|
||||
const uint innerColB = threadIdx.x % BN;
|
||||
// for both As and Bs we want each load to span the full column-width, for
|
||||
// better GMEM coalescing (as opposed to spanning full row-width and iterating
|
||||
// across columns)
|
||||
const uint strideB = numThreadsBlocktile / BN;
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[TM * TN] = {0.0};
|
||||
// register caches for As and Bs
|
||||
float regM[TM] = {0.0};
|
||||
float regN[TN] = {0.0};
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
for (uint loadOffset = 0; loadOffset < BM; loadOffset += strideA) {
|
||||
As[(innerRowA + loadOffset) * BK + innerColA] =
|
||||
A[(innerRowA + loadOffset) * K + innerColA];
|
||||
}
|
||||
for (uint loadOffset = 0; loadOffset < BK; loadOffset += strideB) {
|
||||
Bs[(innerRowB + loadOffset) * BN + innerColB] =
|
||||
B[(innerRowB + loadOffset) * N + innerColB];
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// advance blocktile
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// block into registers
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
||||
}
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
||||
}
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[resIdxM * TN + resIdxN] +=
|
||||
regM[resIdxM] * regN[resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN] =
|
||||
alpha * threadResults[resIdxM * TN + resIdxN] +
|
||||
beta * C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
98
upstream_ref/sgemm_cuda/6_kernel_vectorize.cuh
Normal file
98
upstream_ref/sgemm_cuda/6_kernel_vectorize.cuh
Normal file
@@ -0,0 +1,98 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void sgemmVectorize(int M, int N, int K, float alpha, float *A,
|
||||
float *B, float beta, float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// BN/TN are the number of threads to span a column
|
||||
const int threadCol = threadIdx.x % (BN / TN);
|
||||
const int threadRow = threadIdx.x / (BN / TN);
|
||||
|
||||
// allocate space for the current blocktile in smem
|
||||
__shared__ float As[BM * BK];
|
||||
__shared__ float Bs[BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// we'll load 128bit / 32bit = 4 elements per thread at each step
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[TM * TN] = {0.0};
|
||||
float regM[TM] = {0.0};
|
||||
float regN[TN] = {0.0};
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
// transpose A while loading it
|
||||
float4 tmp =
|
||||
reinterpret_cast<float4 *>(&A[innerRowA * K + innerColA * 4])[0];
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA] = tmp.x;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA] = tmp.y;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA] = tmp.z;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA] = tmp.w;
|
||||
|
||||
reinterpret_cast<float4 *>(&Bs[innerRowB * BN + innerColB * 4])[0] =
|
||||
reinterpret_cast<float4 *>(&B[innerRowB * N + innerColB * 4])[0];
|
||||
__syncthreads();
|
||||
|
||||
// advance blocktile
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// block into registers
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[i] = As[dotIdx * BM + threadRow * TM + i];
|
||||
}
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
||||
}
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[resIdxM * TN + resIdxN] +=
|
||||
regM[resIdxM] * regN[resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
tmp.x = alpha * threadResults[resIdxM * TN + resIdxN] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[resIdxM * TN + resIdxN + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[resIdxM * TN + resIdxN + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[resIdxM * TN + resIdxN + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(
|
||||
&C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN])[0] =
|
||||
tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
103
upstream_ref/sgemm_cuda/7_kernel_resolve_bank_conflicts.cuh
Normal file
103
upstream_ref/sgemm_cuda/7_kernel_resolve_bank_conflicts.cuh
Normal file
@@ -0,0 +1,103 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void sgemmResolveBankConflicts(int M, int N, int K, float alpha,
|
||||
float *A, float *B, float beta,
|
||||
float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// BN/TN are the number of threads to span a column
|
||||
const int threadCol = threadIdx.x % (BN / TN);
|
||||
const int threadRow = threadIdx.x / (BN / TN);
|
||||
|
||||
// allocate space for the current blocktile in smem
|
||||
__shared__ float As[BM * BK];
|
||||
__shared__ float Bs[BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// we'll load 128bit / 32bit = 4 elements per thread at each step
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[TM * TN] = {0.0};
|
||||
float regM[TM] = {0.0};
|
||||
float regN[TN] = {0.0};
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
// transpose A while loading it
|
||||
float4 tmp =
|
||||
reinterpret_cast<float4 *>(&A[innerRowA * K + innerColA * 4])[0];
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA] = tmp.x;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA] = tmp.y;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA] = tmp.z;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA] = tmp.w;
|
||||
|
||||
// "linearize" Bs while storing it
|
||||
tmp = reinterpret_cast<float4 *>(&B[innerRowB * N + innerColB * 4])[0];
|
||||
Bs[((innerColB % 2) * 4 + innerRowB * 8 + 0) * 16 + innerColB / 2] = tmp.x;
|
||||
Bs[((innerColB % 2) * 4 + innerRowB * 8 + 1) * 16 + innerColB / 2] = tmp.y;
|
||||
Bs[((innerColB % 2) * 4 + innerRowB * 8 + 2) * 16 + innerColB / 2] = tmp.z;
|
||||
Bs[((innerColB % 2) * 4 + innerRowB * 8 + 3) * 16 + innerColB / 2] = tmp.w;
|
||||
__syncthreads();
|
||||
|
||||
// advance blocktile
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// block into registers
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[i] = As[dotIdx * BM + threadRow * TM + i];
|
||||
}
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[i] = Bs[(dotIdx * 8 + i) * 16 + threadCol];
|
||||
}
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[resIdxM * TN + resIdxN] +=
|
||||
regM[resIdxM] * regN[resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
tmp.x = alpha * threadResults[resIdxM * TN + resIdxN] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[resIdxM * TN + resIdxN + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[resIdxM * TN + resIdxN + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[resIdxM * TN + resIdxN + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(
|
||||
&C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN])[0] =
|
||||
tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
103
upstream_ref/sgemm_cuda/8_kernel_bank_extra_col.cuh
Normal file
103
upstream_ref/sgemm_cuda/8_kernel_bank_extra_col.cuh
Normal file
@@ -0,0 +1,103 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void sgemmResolveBankExtraCol(int M, int N, int K, float alpha,
|
||||
float *A, float *B, float beta,
|
||||
float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// BN/TN are the number of threads to span a column
|
||||
const int threadCol = threadIdx.x % (BN / TN);
|
||||
const int threadRow = threadIdx.x / (BN / TN);
|
||||
|
||||
// allocate space for the current blocktile in smem
|
||||
__shared__ float As[BM * BK];
|
||||
const int extraCols = 5;
|
||||
__shared__ float Bs[BK * (BN + extraCols)];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// we'll load 128bit / 32bit = 4 elements per thread at each step
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[TM * TN] = {0.0};
|
||||
float regM[TM] = {0.0};
|
||||
float regN[TN] = {0.0};
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
// transpose A while loading it
|
||||
float4 tmp =
|
||||
reinterpret_cast<float4 *>(&A[innerRowA * K + innerColA * 4])[0];
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA] = tmp.x;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA] = tmp.y;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA] = tmp.z;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA] = tmp.w;
|
||||
|
||||
tmp = reinterpret_cast<float4 *>(&B[innerRowB * N + innerColB * 4])[0];
|
||||
Bs[innerRowB * (BN + extraCols) + innerColB * 4 + 0] = tmp.x;
|
||||
Bs[innerRowB * (BN + extraCols) + innerColB * 4 + 1] = tmp.y;
|
||||
Bs[innerRowB * (BN + extraCols) + innerColB * 4 + 2] = tmp.z;
|
||||
Bs[innerRowB * (BN + extraCols) + innerColB * 4 + 3] = tmp.w;
|
||||
__syncthreads();
|
||||
|
||||
// advance blocktile
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// block into registers
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[i] = As[dotIdx * BM + threadRow * TM + i];
|
||||
}
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[i] = Bs[dotIdx * (BN + extraCols) + threadCol * TN + i];
|
||||
}
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[resIdxM * TN + resIdxN] +=
|
||||
regM[resIdxM] * regN[resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
tmp.x = alpha * threadResults[resIdxM * TN + resIdxN] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[resIdxM * TN + resIdxN + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[resIdxM * TN + resIdxN + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[resIdxM * TN + resIdxN + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(
|
||||
&C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN])[0] =
|
||||
tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
127
upstream_ref/sgemm_cuda/9_kernel_autotuned.cuh
Normal file
127
upstream_ref/sgemm_cuda/9_kernel_autotuned.cuh
Normal file
@@ -0,0 +1,127 @@
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
const int K9_NUM_THREADS = 256;
|
||||
|
||||
template <const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void __launch_bounds__(K9_NUM_THREADS)
|
||||
sgemmAutotuned(int M, int N, int K, float alpha, float *A, float *B,
|
||||
float beta, float *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
// size of warptile
|
||||
constexpr int WM = TM * 16;
|
||||
constexpr int WN = TN * 16;
|
||||
// iterations of warptile
|
||||
constexpr int WMITER = CEIL_DIV(BM, WM);
|
||||
constexpr int WNITER = CEIL_DIV(BN, WN);
|
||||
|
||||
// Placement of the thread in the warptile
|
||||
const int threadCol = threadIdx.x % (WN / TN);
|
||||
const int threadRow = threadIdx.x / (WN / TN);
|
||||
|
||||
// allocate space for the current blocktile in smem
|
||||
__shared__ float As[BM * BK];
|
||||
__shared__ float Bs[BK * BN];
|
||||
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// we'll load 128bit / 32bit = 4 elements per thread at each step
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
constexpr uint rowStrideA = (K9_NUM_THREADS * 4) / BK;
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
constexpr uint rowStrideB = K9_NUM_THREADS / (BN / 4);
|
||||
|
||||
// allocate thread-local cache for results in registerfile
|
||||
float threadResults[WMITER * WNITER * TM * TN] = {0.0};
|
||||
float regM[TM] = {0.0};
|
||||
float regN[TN] = {0.0};
|
||||
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
for (uint offset = 0; offset + rowStrideA <= BM; offset += rowStrideA) {
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&A[(innerRowA + offset) * K + innerColA * 4])[0];
|
||||
// transpose A while storing it
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA + offset] = tmp.x;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA + offset] = tmp.y;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA + offset] = tmp.z;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA + offset] = tmp.w;
|
||||
}
|
||||
|
||||
for (uint offset = 0; offset + rowStrideB <= BK; offset += rowStrideB) {
|
||||
reinterpret_cast<float4 *>(
|
||||
&Bs[(innerRowB + offset) * BN + innerColB * 4])[0] =
|
||||
reinterpret_cast<float4 *>(
|
||||
&B[(innerRowB + offset) * N + innerColB * 4])[0];
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (uint wmIdx = 0; wmIdx < WMITER; ++wmIdx) {
|
||||
for (uint wnIdx = 0; wnIdx < WNITER; ++wnIdx) {
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// block into registers
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[i] = As[dotIdx * BM + (wmIdx * WM) + threadRow * TM + i];
|
||||
}
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[i] = Bs[dotIdx * BN + (wnIdx * WN) + threadCol * TN + i];
|
||||
}
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[(wmIdx * TM + resIdxM) * (WNITER * TN) +
|
||||
wnIdx * TN + resIdxN] +=
|
||||
regM[resIdxM] * regN[resIdxN];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
// advance blocktile
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
}
|
||||
|
||||
// write out the results
|
||||
for (uint wmIdx = 0; wmIdx < WMITER; ++wmIdx) {
|
||||
for (uint wnIdx = 0; wnIdx < WNITER; ++wnIdx) {
|
||||
float *C_interim = C + (wmIdx * WM * N) + (wnIdx * WN);
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 4) {
|
||||
// load C vector into registers
|
||||
float4 tmp = reinterpret_cast<float4 *>(
|
||||
&C_interim[(threadRow * TM + resIdxM) * N + threadCol * TN +
|
||||
resIdxN])[0];
|
||||
// perform GEMM update in reg
|
||||
const int i =
|
||||
(wmIdx * TM + resIdxM) * (WNITER * TN) + wnIdx * TN + resIdxN;
|
||||
tmp.x = alpha * threadResults[i + 0] + beta * tmp.x;
|
||||
tmp.y = alpha * threadResults[i + 1] + beta * tmp.y;
|
||||
tmp.z = alpha * threadResults[i + 2] + beta * tmp.z;
|
||||
tmp.w = alpha * threadResults[i + 3] + beta * tmp.w;
|
||||
// write back
|
||||
reinterpret_cast<float4 *>(&C_interim[(threadRow * TM + resIdxM) * N +
|
||||
threadCol * TN + resIdxN])[0] =
|
||||
tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user