§ 参考资料 共 7 条
Softmax Attention
输入 $\boldsymbol Q_{M\times d}$、$\boldsymbol K_{N\times d}$ 和 $\boldsymbol V_{N\times d}$,计算 Attention:
$$ \mathrm{Attention}(\boldsymbol Q,\boldsymbol K,\boldsymbol V)=\mathrm{softmax}\left(\frac{\boldsymbol Q\boldsymbol K^T}{\sqrt d}\right)\boldsymbol V $$其中,Softmax 为逐行计算
例子:
Input:
Q (2x4): [[1.0, 0.0, 0.0, 0.0],
[0.0, 1.0, 0.0, 0.0]]
K (3x4): [[1.0, 0.0, 0.0, 0.0],
[0.0, 1.0, 0.0, 0.0],]
[0.0, 0.0, 1.0, 0.0]]
V (3x4): [[1.0, 2.0, 3.0, 4.0],
[5.0, 6.0, 7.0, 8.0],]
[9.0, 10.0, 11.0, 12.0]]
Output:
O (2x4): [[4.29, 5.29, 6.29, 7.29],
[5.0, 6.0, 7.0, 8.0]]
基于矩阵乘法的做法
最简单的做法就是三个 Kernel:
- 转置 + 矩阵乘法,计算 $\boldsymbol Q\boldsymbol K^T/\sqrt d$
- Softmax,每一行作为一个 Block,用不到 Global Memory
- 矩阵乘法
代码:
#include <cuda_runtime.h>
__device__ __forceinline__ constexpr int compute_offset(int row, int col,
int cols) {
return row * cols + col;
}
// A: M x K
// B: K x N
// => output: AB, M x N
// A: M x K
// B: K x N
// => output: AB^T, M x N
template <const int block_size, const bool transpose, const bool norm>
__global__ void multiple_kernel(const float* A, const float* B, float* output,
int M, int N, int K) {
__shared__ float tile_A[block_size][block_size];
__shared__ float tile_B[block_size][block_size];
const auto col = threadIdx.x + blockIdx.x * blockDim.x;
const auto row = threadIdx.y + blockIdx.y * blockDim.y;
const auto local_row = threadIdx.y;
const auto local_col = threadIdx.x;
float sum = 0.0f;
#pragma unroll
for (int i = 0; i < K; i += block_size) {
if (row < M && local_col + i < K) {
tile_A[local_row][local_col] = A[compute_offset(row, local_col + i, K)];
} else {
tile_A[local_row][local_col] = 0.0f;
}
if constexpr (transpose) {
if (col < N && local_row + i < K) {
tile_B[local_row][local_col] = B[compute_offset(col, local_row + i, K)];
} else {
tile_B[local_row][local_col] = 0.0f;
}
} else {
if (local_row + i < K && col < N) {
tile_B[local_row][local_col] = B[compute_offset(local_row + i, col, N)];
} else {
tile_B[local_row][local_col] = 0.0f;
}
}
__syncthreads();
#pragma unroll
for (int j = 0; j < block_size; j++) {
sum = fmaf(tile_A[local_row][j], tile_B[j][local_col], sum);
}
__syncthreads();
}
if (row < M && col < N) {
if constexpr (norm) {
output[compute_offset(row, col, N)] = sum * rsqrtf(static_cast<float>(K));
} else {
output[compute_offset(row, col, N)] = sum;
}
}
}
struct SoftmaxState {
float max;
float sum;
};
__device__ __forceinline__ SoftmaxState init_state() {
return {-INFINITY, 0.0f};
}
__device__ __forceinline__ SoftmaxState mono_state(float value) {
return {value, 1.0f};
}
__device__ __forceinline__ SoftmaxState combine(SoftmaxState a,
SoftmaxState b) {
if (a.max == -INFINITY) {
return b;
}
if (b.max == -INFINITY) {
return a;
}
SoftmaxState out;
out.max = fmaxf(a.max, b.max);
out.sum = a.sum * __expf(a.max - out.max) + b.sum * __expf(b.max - out.max);
return out;
}
__device__ __forceinline__ SoftmaxState warp_reduce(SoftmaxState state) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
SoftmaxState other;
other.max = __shfl_down_sync(0xffffffff, state.max, offset);
other.sum = __shfl_down_sync(0xffffffff, state.sum, offset);
state = combine(state, other);
}
return state;
}
template <const int block_size>
__device__ __forceinline__ SoftmaxState block_reduce(SoftmaxState state) {
__shared__ SoftmaxState warp_results[block_size];
const auto lane = threadIdx.x & 31;
const auto warp = threadIdx.x >> 5;
const auto num_warps = (blockDim.x + 31) >> 5;
state = warp_reduce(state);
if (lane == 0) {
warp_results[warp] = state;
}
__syncthreads();
SoftmaxState block_state = init_state();
if (warp == 0) {
if (lane < num_warps) {
block_state = warp_results[lane];
}
block_state = warp_reduce(block_state);
}
return block_state;
}
// softmax(input)
template <const int block_size>
__global__ void softmax_kernel(float* input, int M, int N) {
const auto row = blockIdx.x;
if (row >= M) {
return;
}
auto state = init_state();
#pragma unroll
for (int col = threadIdx.x; col < N; col += blockDim.x) {
const int idx = compute_offset(row, col, N);
state = combine(state, mono_state(input[idx]));
}
state = block_reduce<block_size>(state);
__shared__ SoftmaxState final_state;
if (threadIdx.x == 0) {
final_state = state;
}
__syncthreads();
#pragma unroll
for (int col = threadIdx.x; col < N; col += blockDim.x) {
const int idx = compute_offset(row, col, N);
input[idx] = __expf(input[idx] - final_state.max) / final_state.sum;
}
}
// Q, K, V, output are device pointers
// Q: M x d
// K: N x d
// QK^T: M x N
// V: N x d
// output: M x d
extern "C" void solve(const float* Q, const float* K, const float* V,
float* output, int M, int N, int d) {
auto total_items = M * N;
float* attention;
cudaMalloc(&attention, total_items * sizeof(float));
cudaMemset(attention, 0, total_items * sizeof(float));
// QK^T / sqrt(d)
{
constexpr auto block_size = 32;
const auto grid_dim = dim3((N + block_size - 1) / block_size,
(M + block_size - 1) / block_size);
const auto block_dim = dim3(block_size, block_size);
multiple_kernel<block_size, true, true>
<<<grid_dim, block_dim>>>(Q, K, attention, M, N, d);
}
// row-independent softmax(attention)
{
constexpr auto block_size = 256;
const auto grid_dim = dim3(M);
const auto block_dim = dim3(block_size);
softmax_kernel<block_size><<<grid_dim, block_dim>>>(attention, M, N);
}
// softmax(attention) * V
{
constexpr auto block_size = 32;
const auto grid_dim = dim3((d + 31) / 32, (M + 31) / 32);
const auto block_dim = dim3(32, 32);
multiple_kernel<block_size, false, false>
<<<grid_dim, block_dim>>>(attention, V, output, M, d, N);
}
cudaDeviceSynchronize();
}
然后 multiple_kernel 其实有点问题,这里对于 transpose=false 情况下 tile_B 的读取不合并(但是没有 Bank Conflict),不能直接用 col,改一下:
template <const int block_size, const bool transpose, const bool norm>
__global__ void multiple_kernel(const float* A, const float* B, float* output,
int M, int N, int K) {
__shared__ float tile_A[block_size][block_size];
__shared__ float tile_B[block_size][block_size];
const auto col = threadIdx.x + blockIdx.x * blockDim.x;
const auto row = threadIdx.y + blockIdx.y * blockDim.y;
const auto local_row = threadIdx.y;
const auto local_col = threadIdx.x;
float sum = 0.0f;
#pragma unroll
for (int i = 0; i < K; i += block_size) {
if (row < M && local_col + i < K) {
tile_A[local_row][local_col] = A[compute_offset(row, local_col + i, K)];
} else {
tile_A[local_row][local_col] = 0.0f;
}
if constexpr (transpose) {
// 左上角起点为 row=blockIdx.x * blockDim.x, col=i
// threadIdx.x 是最快的,所以用它来作为 col 的自变量,相对的,用
// threadIdx.y 作为 row 的自变量 最终得到的矩阵应该进行转置,来适配下面的
// local 矩阵乘法
const auto key = threadIdx.y + blockIdx.x * blockDim.x;
if (key < N && local_col + i < K) {
tile_B[local_col][local_row] = B[compute_offset(key, local_col + i, K)];
} else {
tile_B[local_col][local_row] = 0.0f;
}
} else {
if (local_row + i < K && col < N) {
tile_B[local_row][local_col] = B[compute_offset(local_row + i, col, N)];
} else {
tile_B[local_row][local_col] = 0.0f;
}
}
__syncthreads();
#pragma unroll
for (int j = 0; j < block_size; j++) {
sum = fmaf(tile_A[local_row][j], tile_B[j][local_col], sum);
}
__syncthreads();
}
if (row < M && col < N) {
if constexpr (norm) {
output[compute_offset(row, col, N)] = sum * rsqrtf(static_cast<float>(K));
} else {
output[compute_offset(row, col, N)] = sum;
}
}
}
更高效的矩阵乘法 Kernel:
- 对于读取矩阵部分,可以单个线程读多个元素,可以增加并行性
- 对于计算部分,也可以一个线程负责多个元素的计算,这样也可以从 Shared Memory 中先把数据缓存到寄存器中
所以把原先的 block_size 参数拆开,先把整个矩阵分成若干 BM x BN 个子矩阵,每个 Block 负责一个子矩阵。每个子矩阵再分成若干 TM x TN 的小子矩阵,每个线程负责一个小子矩阵。这样增加了并行度和数据缓存度
然后每个 Block 单次缓存 BK 长度的子矩阵
- 在读取的时候,每个线程读取若干个元素到 Shared Memory 中,共同得到
tile_A和tile_B - 在计算的时候,优先枚举
k,这样可以把当前A矩阵的第k列和B矩阵的第k行先放到寄存器中,然后再处理当前线程TM x TN小子矩阵计算时的所有k的部分,累加到一个累加器acc[TM][TN]中
template <const int BM, const int BN, const int BK, const int TM, const int TN>
__global__ void multiple_kernel(const float* A, const float* B, float* output,
int M, int N, int K) {
__shared__ float tile_A[BM][BK + 1];
__shared__ float tile_B[BK][BN + 1];
const int tx = threadIdx.x;
const int ty = threadIdx.y;
const int tid = threadIdx.y * blockDim.x + threadIdx.x;
const int num_threads = blockDim.x * blockDim.y;
// Block 负责的子矩阵的左上角坐标
const auto block_row = blockIdx.y * BM;
const auto block_col = blockIdx.x * BN;
// 当前线程负责的子矩阵的左上角坐标(相对于 Block 负责的矩阵左上角)
const int thread_row = ty * TM;
const int thread_col = tx * TN;
// 累加器
float acc[TM][TN] = {};
#pragma unroll
for (int k0 = 0; k0 < K; k0 += BK) {
// 加载 tile_A 的数据
for (int i = tid; i < BM * BK; i += num_threads) {
const int local_row = i / BK;
const int local_col = i % BK;
const int global_row = block_row + local_row;
const int global_col = k0 + local_col;
if (global_row < M && global_col < K) {
tile_A[local_row][local_col] = A[global_row * K + global_col];
} else {
tile_A[local_row][local_col] = 0.0f;
}
}
// 加载 tile_B 的数据
for (int i = tid; i < BK * BN; i += num_threads) {
const int local_row = i / BN;
const int local_col = i % BN;
const int global_row = k0 + local_row;
const int global_col = block_col + local_col;
if (global_row < K && global_col < N) {
tile_B[local_row][local_col] = B[global_row * N + global_col];
} else {
tile_B[local_row][local_col] = 0.0f;
}
}
__syncthreads();
// Register-tiled 计算,每个线程计算 TM x TN 块
#pragma unroll
for (int k = 0; k < BK; ++k) {
float a_reg[TM];
float b_reg[TN];
#pragma unroll
for (int i = 0; i < TM; ++i) {
a_reg[i] = tile_A[thread_row + i][k];
}
#pragma unroll
for (int j = 0; j < TN; ++j) {
b_reg[j] = tile_B[k][thread_col + j];
}
#pragma unroll
for (int i = 0; i < TM; ++i) {
#pragma unroll
for (int j = 0; j < TN; ++j) {
acc[i][j] = fmaf(a_reg[i], b_reg[j], acc[i][j]);
}
}
}
__syncthreads();
}
#pragma unroll
for (int i = 0; i < TM; ++i) {
const int row = block_row + thread_row + i;
#pragma unroll
for (int j = 0; j < TN; ++j) {
const int col = block_col + thread_col + j;
if (row < M && col < N) {
output[row * N + col] = acc[i][j];
}
}
}
}
{
constexpr int BM = 64;
constexpr int BN = 64;
constexpr int BK = 32;
constexpr int TM = 4;
constexpr int TN = 4;
dim3 block_dim(BN / TN, BM / TM);
dim3 grid_dim((N + BN - 1) / BN, (M + BM - 1) / BM);
multiple_kernel<BM, BN, BK, TM, TN>
<<<grid_dim, block_dim>>>(A, B, C, M, N, K);
}
简化版的 Flash Attention
前面的基于矩阵乘法的做法申请了一个 $M\times N$ 的 Attention 矩阵,并且她需要写一次然后再读一次,很不好
一个更好的思路是:不显式保存完整的 Attention 矩阵,而是在计算出一个 Score 后立即消费它。具体来说,可以将 $\boldsymbol Q\boldsymbol K^T$、Softmax 和 $\boldsymbol P\boldsymbol V$ 融合到同一个 Kernel 中,通过 Online Softmax 一边计算 Score,一边更新 Softmax 状态并累加对应的 $\boldsymbol V$
具体来说:
- 每个 Block 负责一个 Query,即计算一个
Q[i]对应的输出O[i] - 一个 Block 中包含若干个 Warp,不同 Warp 分摊 KV 维度上的工作。例如有 8 个 Warp 时,第 $w$ 个 Warp 依次处理
K[w]、K[w+8]、K[w+16]等 - 对于某个
K[j],一个 Warp 中的 32 个线程共同计算 $s_{ij}=\frac{\boldsymbol Q_i\cdot \boldsymbol K_j}{\sqrt d}$,每个线程负责若干个 Channel,先计算局部点积,再通过 Warp Reduce 得到完整的 Score - 得到
score后不将其写入 Attention 矩阵,而是立即使用 Online Softmax 更新当前 Warp 的 Softmax 状态,同时更新对 $\boldsymbol V_j$ 的加权累加结果 - 每个 Warp 独立处理自己负责的一部分 Key,因此最后还需要将多个 Warp 的 Online Softmax 状态合并,得到整行 Attention 的最终结果
这种实现的关键收益是:整个计算过程中都不需要生成 $M\times N$ 的 Attention 矩阵。中间状态只保存在寄存器和少量 Shared Memory 中,从而减少了显存占用以及对 Global Memory 的读写。
#include <cuda_runtime.h>
constexpr int k_warps_per_block = 8;
constexpr int k_threads_per_block = k_warps_per_block * 32;
constexpr int k_max_d = 128;
constexpr int k_cols_per_thread = k_max_d / 32;
__device__ __forceinline__ float warp_reduce(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_xor_sync(0xffffffffu, val, offset);
}
return val;
}
__global__ void softmax_atten_kernel(const float* __restrict__ Q,
const float* __restrict__ K,
const float* __restrict__ V,
float* __restrict__ O, int M, int N, int d,
float scale) {
const int i = blockIdx.x;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
extern __shared__ float s_mem[];
float* s_out = s_mem;
float* s_m = s_mem + d;
float* s_l = s_m + k_warps_per_block;
for (int k = threadIdx.x; k < d; k += blockDim.x) {
s_out[k] = 0.0f;
}
__syncthreads();
// cache Q[i] in registers
float q[k_cols_per_thread];
#pragma unroll
for (int c = 0; c < k_cols_per_thread; ++c) {
const int k = lane + c * 32;
q[c] = (k < d) ? Q[i * d + k] : 0.0f;
}
float m = -INFINITY;
float l = 0.0f;
float acc[k_cols_per_thread] = {};
for (int j = warp; j < N; j += k_warps_per_block) {
const float* Kj = K + j * d;
const float* Vj = V + j * d;
// score = Q[i] dot K[j]
float partial = 0.0f;
#pragma unroll
for (int c = 0; c < k_cols_per_thread; ++c) {
const int k = lane + c * 32;
if (k < d) {
partial = fmaf(q[c], Kj[k], partial);
}
}
const float score = warp_reduce(partial) * scale;
// online softmax
const float old_m = m;
m = fmaxf(m, score);
const float alpha = __expf(old_m - m);
const float beta = __expf(score - m);
l = l * alpha + beta;
// online P*V accumulation
#pragma unroll
for (int c = 0; c < k_cols_per_thread; ++c) {
const int k = lane + c * 32;
if (k < d) {
acc[c] = acc[c] * alpha + beta * Vj[k];
}
}
}
if (lane == 0) {
s_m[warp] = m;
s_l[warp] = l;
}
__syncthreads();
float g_m = -INFINITY;
float g_l = 0.0f;
#pragma unroll
for (int w = 0; w < k_warps_per_block; ++w) {
g_m = fmaxf(g_m, s_m[w]);
}
#pragma unroll
for (int w = 0; w < k_warps_per_block; ++w) {
g_l += s_l[w] * __expf(s_m[w] - g_m);
}
const float warp_scale = __expf(m - g_m);
#pragma unroll
for (int c = 0; c < k_cols_per_thread; ++c) {
const int k = lane + c * 32;
if (k < d) {
atomicAdd(&s_out[k], acc[c] * warp_scale);
}
}
__syncthreads();
const float inv_l = 1.0f / g_l;
for (int k = threadIdx.x; k < d; k += blockDim.x) {
O[i * d + k] = s_out[k] * inv_l;
}
}
// Q, K, V, output are device pointers
extern "C" void solve(const float* Q, const float* K, const float* V,
float* output, int M, int N, int d) {
constexpr int block_dim = k_threads_per_block;
const int grid_dim = M;
const int s_mem_size = (d + 2 * k_warps_per_block) * sizeof(float);
softmax_atten_kernel<<<grid_dim, block_dim, s_mem_size>>>(Q, K, V, output, M,
N, d, rsqrtf(d));
}
在这里,Warp Reduce 用了 __shfl_xor_sync 而非 __shfl_down_sync,这样的话每个 Lane 都能得到最终的结果,而无需最后再 __shfl_sync 一下了,这叫做 Butterfly Reduction
3D Convolution
输入三维体 $\boldsymbol V_{D\times H\times W}$ 和卷积核 $\boldsymbol K_{K_D\times K_H\times K_W}$,计算无填充的三维卷积 $\boldsymbol O$:
$$ \boldsymbol O_{i,j,k}=\sum_{d=0}^{K_D-1}\sum_{r=0}^{K_H-1}\sum_{c=0}^{K_W-1}\boldsymbol V_{i+d,j+r,k+c}\boldsymbol K_{d,r,c} $$例子:
Input:
V (2x2x2): [[[1.0, 2.0],
[3.0, 4.0]],
[[5.0, 6.0],
[7.0, 8.0]]]
K (2x2x2): [[[1.0, 1.0],
[1.0, 1.0]],
[[1.0, 1.0],
[1.0, 1.0]]]
Output:
O (1x1x1): [[[36.0]]]
其中:
- $1\le D,H,W\le 256$
- $1\le K_D,K_H,K_W\le 5$
- $K_D\le D$、$K_H\le H$、$K_W\le W$
由于卷积核很小,简单的思路就是按照结果矩阵划分 Block,每个 Block 负责一个子矩阵。每个子矩阵需要知道所有的 $\boldsymbol V$ 矩阵中的元素个数比较少,直接分 Tile 把原矩阵对应的元素读出来,放到 Shared Memory 中即可。把卷积核这个矩阵可以放到常量区
代码:
#include <cuda_runtime.h>
#define OFFSET(d, r, c, rows, cols) \
((d) * ((rows) * (cols)) + (r) * (cols) + (c))
__constant__ float c_kernel[256];
template <const int BD, const int BR, const int BC, const int MAX_K>
__global__ void conv3d_kernel(const float* __restrict__ input, float* output,
int input_depth, int input_rows, int input_cols,
int kernel_depth, int kernel_rows,
int kernel_cols) {
__shared__ float tile[BD + MAX_K - 1][BR + MAX_K - 1][BC + MAX_K - 1];
for (int d = threadIdx.z; d < kernel_depth + BD - 1; d += BD) {
for (int r = threadIdx.y; r < kernel_rows + BR - 1; r += BR) {
for (int c = threadIdx.x; c < kernel_cols + BC - 1; c += BC) {
const auto depth = d + blockIdx.z * BD;
const auto row = r + blockIdx.y * BR;
const auto col = c + blockIdx.x * BC;
if (depth < input_depth && row < input_rows && col < input_cols) {
tile[d][r][c] =
input[OFFSET(depth, row, col, input_rows, input_cols)];
} else {
tile[d][r][c] = 0.0f;
}
}
}
}
__syncthreads();
const auto output_depth = input_depth - kernel_depth + 1;
const auto output_rows = input_rows - kernel_rows + 1;
const auto output_cols = input_cols - kernel_cols + 1;
const auto depth = threadIdx.z + blockIdx.z * BD;
const auto row = threadIdx.y + blockIdx.y * BR;
const auto col = threadIdx.x + blockIdx.x * BC;
if (depth < output_depth && row < output_rows && col < output_cols) {
float sum = 0.0f;
for (int kd = 0; kd < kernel_depth; ++kd) {
for (int kr = 0; kr < kernel_rows; ++kr) {
for (int kc = 0; kc < kernel_cols; ++kc) {
sum += tile[threadIdx.z + kd][threadIdx.y + kr][threadIdx.x + kc]
* c_kernel[OFFSET(kd, kr, kc, kernel_rows, kernel_cols)];
}
}
}
output[OFFSET(depth, row, col, output_rows, output_cols)] = sum;
}
}
// input, kernel, output are device pointers
extern "C" void solve(const float* input, const float* kernel, float* output,
int input_depth, int input_rows, int input_cols,
int kernel_depth, int kernel_rows, int kernel_cols) {
const auto kernel_size = kernel_depth * kernel_rows * kernel_cols;
cudaMemcpyToSymbol(c_kernel, kernel, kernel_size * sizeof(float), 0,
cudaMemcpyDeviceToDevice);
constexpr auto BD = 8, BR = 8, BC = 8;
const int output_depth = input_depth - kernel_depth + 1;
const int output_rows = input_rows - kernel_rows + 1;
const int output_cols = input_cols - kernel_cols + 1;
const auto grid_dim =
dim3((output_cols + BC - 1) / BC, (output_rows + BR - 1) / BR,
(output_depth + BD - 1) / BD);
const auto block_dim = dim3(BC, BR, BD);
conv3d_kernel<BD, BR, BC, 5><<<grid_dim, block_dim>>>(
input, output, input_depth, input_rows, input_cols, kernel_depth,
kernel_rows, kernel_cols);
}
LeetGPU 的 Solution 中最快的代码使用了 REG_COL 来提高带宽,每个线程连续计算 4 列,并且尽可能地合并读取操作,代码:
#define OFFSET(d, r, c, rows, cols) \
((d) * ((rows) * (cols)) + (r) * (cols) + (c))
__constant__ float c_kernel[256];
template <const int BD, const int BR, const int BX, const int KD, const int KR,
const int KC, const int REG_COL>
__global__ void conv3d_kernel(const float* __restrict__ input, float* output,
int input_depth, int input_rows, int input_cols) {
// 一个对应 REG_COL 列,所以一个 block 总共 BX * REG_COL 列
constexpr auto TILE_D = BD + KD - 1;
constexpr auto TILE_R = BR + KR - 1;
constexpr auto TILE_C = BX * REG_COL + KC - 1;
// +1 为 降低 Bank Conflict
__shared__ float tile[TILE_D][TILE_R][TILE_C + 1];
const auto tid = threadIdx.x + threadIdx.y * BX + threadIdx.z * BX * BR;
const auto num_threads = BX * BR * BD;
const auto num_elements = TILE_D * TILE_R * TILE_C;
// 线性协同加载可以提高访存合并效果
for (int i = tid; i < num_elements; i += num_threads) {
const auto d = i / (TILE_R * TILE_C);
const auto r = (i % (TILE_R * TILE_C)) / TILE_C;
const auto c = (i % (TILE_R * TILE_C)) % TILE_C;
const auto depth = blockIdx.z * BD + d;
const auto row = blockIdx.y * BR + r;
const auto col = blockIdx.x * BX * REG_COL + c;
if (depth < input_depth && row < input_rows && col < input_cols) {
tile[d][r][c] = input[OFFSET(depth, row, col, input_rows, input_cols)];
} else {
tile[d][r][c] = 0.0f;
}
}
__syncthreads();
const auto output_depth = input_depth - KD + 1;
const auto output_rows = input_rows - KR + 1;
const auto output_cols = input_cols - KC + 1;
float acc[REG_COL] = {};
#pragma unroll
for (int kd = 0; kd < KD; ++kd) {
#pragma unroll
for (int kr = 0; kr < KR; ++kr) {
// 当前线程需要的当前行的所有元素
float input_reg[KC + REG_COL - 1];
#pragma unroll
for (int v = 0; v < KC + REG_COL - 1; ++v) {
input_reg[v] =
tile[threadIdx.z + kd][threadIdx.y + kr][threadIdx.x * REG_COL + v];
}
#pragma unroll
for (int kc = 0; kc < KC; ++kc) {
const auto kval = c_kernel[OFFSET(kd, kr, kc, KR, KC)];
#pragma unroll
for (int rc = 0; rc < REG_COL; ++rc) {
acc[rc] += input_reg[rc + kc] * kval;
}
}
}
}
#pragma unroll
for (int rc = 0; rc < REG_COL; ++rc) {
const auto depth = threadIdx.z + blockIdx.z * BD;
const auto row = threadIdx.y + blockIdx.y * BR;
const auto col = (threadIdx.x + blockIdx.x * BX) * REG_COL + rc;
if (depth < output_depth && row < output_rows && col < output_cols) {
output[OFFSET(depth, row, col, output_rows, output_cols)] = acc[rc];
}
}
}
这其实和前面的矩阵乘法思想相同,可以叫做 Register Tiling,核心思想就是不要让一个线程只计算一个输出,而是让它计算一小组相关输出:
Block 负责一个输出 Tile
Thread 负责 Tile 内的一个微块
Register acc[...] 保存该微块的多个输出
然后利用这些输出之间共享的输入数据:
Global Memory → Shared Memory → Registers → 多次计算
每向更快一级的存储搬一次数据,就尽可能多使用几次
通常可以归纳成五步:
- 选择每线程负责的输出微块,例如
TM × TN或连续REG_COL个输出 - 找出这些输出所需输入的并集
- 让整个 Block 合并地把输入从全局内存加载到 Shared Memory
- 每个线程把自己反复使用的 Shared Memory 数据取到寄存器
- 在寄存器累加器中完成多个输出,最后统一写回
Prefix Sum
输入 float 数组,输出一个每个数的前缀和
例子:
Input: [5.0, -2.0, 3.0, 1.0, -4.0]
Output: [5.0, 3.0, 6.0, 7.0, 3.0]
求前缀和这个操作被叫做 Scan,有两个普遍的算法
先说一下 Step Complexity(步数复杂度)和 Work Complexity(工作量复杂度)的概念:
- 我们可以用“盖房子”来做个生动的比喻,假设要盖一座房子,总共需要砌 10000 块砖
- Work Complexity:就是砌完这 10000 块砖的总工作量
- Step Complexity:就是如果雇佣无限多的工人同时干活,受限于工序先后顺序(比如必须先筑基、再砌墙、最后盖屋顶),最快需要多少个时间步才能盖完
一个是 Hillis–Steele / Kogge–Stone 风格的 Scan 算法,它是 Inclusive(第 $i$ 个结果包含它自己)的
思路大概是每轮让元素与距离为 $2^k$ 的前驱合并,Step Complexity 是 $O(\log n)$,Work Complexity 是 $O(n\log n)$。图示如下:

另一种是 Belloch 算法,是一种 Exclusive Scan 算法。这个比较复杂,包括 Reduce(Up-Sweep)与 Down-Sweep 两个部分,Up-Sweep 计算子树和,Down-Sweep 计算 Exclusive 前缀和
Down-Sweep 阶段从根遍历到叶子节点,节点的值代表当前子树的 Exclusive 前缀和,即子树中最靠前的元素的 Exclusive 前缀和
开始时,设置根节点的值为 $0$(根节点的 Exclusive 前缀和为 $0$),向下遍历时:
- 父节点的值表示子树的 Exclusive 前缀和,左右子节点的值仍然表示子树的和
- 右子节点的值应该是左子树的和加上当前子树的 Exclusive 前缀和,即父节点的值加上原始左子节点的值
- 左子节点的值应该是当前子树的 Exclusive 前缀和,即父节点的值
这样就得到了 Exclusive 前缀和,Step Complexity 是 $O(2\log n)$,Work Complexity 是 $O(2n)$,图示:

可以看到,Hillis–Steele 算法在 Step Complexity 占优,Blelloch 算法在 Work Complexity 占优
类似 Reduction 的步骤,先写一个 Warp Scan,用 Hillis–Steele 算法最好,然后再写 Block Scan,然后处理所有 Block 的前缀和,最后计算 Offset 加到所有的元素上
代码如下(假设 $N<250,000$):
#include <cooperative_groups.h>
#include <cuda_runtime.h>
__device__ __forceinline__ float warp_scan(float value) {
// clang-format off
// Hillis–Steele / Kogge–Stone 风格的 scan 算法,假设 warp 大小是 8
// 0 1 2 3 4 5 6 7 8
// round 1 (offset = 1): 0 0+1 1+2 2+3 3+4 4+5 5+6 6+7 7+8
// round 2 (offset = 2): 0 0+1 0+1+2 0+1+2+3 1+2+3+4 2+3+4+5 3+4+5+6 4+5+6+7 5+6+7+8
// round 3 (offset = 4): 0 0+1 0+1+2 0+1+2+3 0+1+2+3+4 0+1+2+3+4+5 0+1+2+3+4+5+6 0+1+2+3+4+5+6+7 1+2+3+4+5+6+7+8
// round 4 (offset = 8): 0 0+1 0+1+2 0+1+2+3 0+1+2+3+4 0+1+2+3+4+5 0+1+2+3+4+5+6 0+1+2+3+4+5+6+7 0+1+2+3+4+5+6+7+8
// clang-format on
const auto lane = threadIdx.x & 31;
float other_value = 0.0f;
#pragma unroll
for (int offset = 1; offset < 32; offset <<= 1) {
other_value = __shfl_up_sync(0xffffffff, value, offset);
if (lane >= offset) {
value += other_value;
}
}
return value;
}
__device__ __forceinline__ float block_scan(float value) {
const auto warp = threadIdx.x >> 5;
const auto lane = threadIdx.x & 31;
const auto num_warps = blockDim.x >> 5;
__shared__ float warp_prefix_sums[32];
value = warp_scan(value);
if (lane == 31) {
warp_prefix_sums[warp] = value;
}
__syncthreads();
if (warp == 0) {
float warp_sum = lane < num_warps ? warp_prefix_sums[lane] : 0.0f;
warp_sum = warp_scan(warp_sum);
if (lane < num_warps) {
warp_prefix_sums[lane] = warp_sum;
}
}
__syncthreads();
if (warp == 0) {
return value;
} else {
return value + warp_prefix_sums[warp - 1];
}
}
__device__ float g_block_prefix_sums[128 + 1];
template <const int items_per_thread>
__global__ void prefix_sum_kernel(const float* __restrict__ input,
float* __restrict__ output, int N) {
cooperative_groups::grid_group grid = cooperative_groups::this_grid();
const auto tid = threadIdx.x + blockIdx.x * blockDim.x;
float thread_prefix_sums[items_per_thread];
float thread_sum = 0.0f;
#pragma unroll
for (int i = 0; i < items_per_thread; ++i) {
const auto idx = tid * items_per_thread + i;
if (idx < N) {
thread_sum += input[idx];
}
thread_prefix_sums[i] = thread_sum;
}
const auto thread_prefix_sum = block_scan(thread_sum);
const auto thread_exclusive_prefix_sum = thread_prefix_sum - thread_sum;
if (threadIdx.x == blockDim.x - 1) {
g_block_prefix_sums[blockIdx.x] = thread_prefix_sum;
}
grid.sync();
if (blockIdx.x == 0) {
float value =
threadIdx.x < gridDim.x ? g_block_prefix_sums[threadIdx.x] : 0.0f;
value = block_scan(value);
if (threadIdx.x < gridDim.x) {
g_block_prefix_sums[threadIdx.x] = value;
}
}
grid.sync();
const auto block_exclusive_prefix_sum =
blockIdx.x == 0 ? 0.0f : g_block_prefix_sums[blockIdx.x - 1];
#pragma unroll
for (int i = 0; i < items_per_thread; ++i) {
const auto idx = tid * items_per_thread + i;
if (idx < N) {
output[idx] = thread_prefix_sums[i] + block_exclusive_prefix_sum
+ thread_exclusive_prefix_sum;
}
}
}
// input, output are device pointers
extern "C" void solve(const float* input, float* output, int N) {
constexpr auto items_per_thread = 8;
// 256 / 32 = 8
constexpr auto block_dim = 256;
// 250,000 / (256 * 8) = 124
// 124 < 160(最大 block 驻留数量)
// 124 < 256
const auto grid_dim =
(N + (block_dim * items_per_thread) - 1) / (block_dim * items_per_thread);
void* args[] = {(void*)&input, (void*)&output, (void*)&N};
void* kernel = (void*)prefix_sum_kernel<items_per_thread>;
cudaLaunchCooperativeKernel(kernel, grid_dim, block_dim, args, 0, 0);
}
Decoupled Look-back Scan
参考 CUDA 算子笔记 3
General Matrix Multiplication(GEMM)
输入矩阵 $\boldsymbol A_{M\times K}$、$\boldsymbol B_{K\times N}$ 和 $\boldsymbol C_{M\times K}$ 和 $\alpha$、$\beta$,计算:
$$ \boldsymbol C=\alpha\boldsymbol A\boldsymbol B+\beta\boldsymbol C $$输入的矩阵全部为 FP16,即 half 类型,$\alpha$ 和 $\beta$ 是 FP32 float 类型
Tensor Core 与 WMMA
Tensor Core 是 NVIDIA Volta 架构及其后续架构(如 Ampere、Hopper、Ada Lovelace 架构)中引入的一种特殊计算单元。它们专门用于深度学习任务中的张量计算,如矩阵乘法和卷积运算(卷积运算也可以转换为矩阵乘法),最核心的就是加速矩阵乘法
对于矩阵乘法来说,CUDA Core 可以看成是一堆 FMA(Fused Multiply-Add)单元,而 Tensor Core 在硬件层面上执行 $4\times 4\times 4$ 的矩阵乘加(MMA,Matrix Multiply and Accumulate),即 $\boldsymbol D=\boldsymbol A\boldsymbol B+\boldsymbol C$
相比于 CUDA Core,Tensor Core 吞吐量提高了很多,Volta 一个 SM 中有 64 个 FP32 CUDA Core 和 8 个 Tensor Core
- 在一个周期内,Tensor Core 可以执行 $4\times 4\times 4=64$ 次 FMA,SM 吞吐量就是 $64\times 8=512$ 次 FMA
- 而 CUDA Core 在一个周期内只能执行一次 FMA,所以整个 SM 总共 $64$ 次 FMA
大约性能提升了八倍
此外,Tensor Core 还减少了中间数据在寄存器文件、执行单元和线程之间来回搬运的成本。假设某个线程用 CUDA Core 算 $c_{ij}=\sum_k a_{ik}b_{kj}$,典型的流程是:Shared/Global Memory -> Register File -> FP32 FMA Unit -> Register File。先读入数据到寄存器,每做一次 FMA,都要从 Register File 中读出 $a$、$b$ 和 $acc$,计算完成后还要再写回 Register File
而矩阵乘法是一个数据复用度极高的操作,Tensor Core 中,很多数据一旦进入,就可以直接在专用 datapath 中复用。做到一次读取,广播给内部的多个乘法器,中间结果不需要像 CUDA Core 的 FMA 那样频繁地写回 Register File
所以英伟达宣传性能提升了 12 倍
NVIDIA 对外提供的 Tensor Core 最主要的接口是 WMMA(Warp Matrix Multiply and Accumulate):
template<typename Use, int m, int n, int k, typename T, typename Layout=void> class fragment;
// 读取矩阵数据到 fragment,然后 Warp 内部同步
void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm);
void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm, layout_t layout);
// 写入 fragment 到矩阵,然后 Warp 内部同步
void store_matrix_sync(T* mptr, const fragment<...> &a, unsigned ldm, layout_t layout);
// 初始化 fragment
void fill_fragment(fragment<...> &a, const T& v);
// MMA 运算 + Warp 内部同步
void mma_sync(fragment<...> &d, const fragment<...> &a, const fragment<...> &b, const fragment<...> &c, bool satf=false);
其中:
fragment:Tensor Core 数据存储类,支持matrix_a、matrix_b和accumulatorload_matrix_sync:Tensor Core 数据加载 API,支持将矩阵数据从 Shared/Global Memory 加载到fragmentstore_matrix_sync:Tensor Core 结果存储 API,支持将计算结果从fragment存储到 Shared/Global Memoryfill_fragment:fragment填充 API,支持常数值填充mma_sync:Tensor Core 矩阵乘计算 API,支持 $\boldsymbol D = \boldsymbol A\boldsymbol B + \boldsymbol C$ 或者 $\boldsymbol C = \boldsymbol A\boldsymbol B + \boldsymbol C$
基于 WMMA 的 GEMM
引入 WMMA 之后,整体的矩阵计算层级就变成了:

代码:
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <mma.h>
#define OFFSET(row, col, cols) ((row) * (cols) + (col))
#define CEIL(x, y) (((x) + (y) - 1) / (y))
// A 分成 k_tile_m x k_tile_k 的块,B 分成 k_tile_k x k_tile_n 的块,C 分成
// k_tile_m x k_tile_n 的块
constexpr int k_tile_m = 16;
constexpr int k_tile_n = 16;
constexpr int k_tile_k = 16;
// 4 个 warp,分配到两个维度上:
// 0 1 n
// ┌─────
// 0 │ 0 1
// 1 │ 2 3
// m
constexpr int k_threads_per_warp = 32;
constexpr int k_warps_per_block_dim_m = 2;
constexpr int k_warps_per_block_dim_n = 2;
constexpr int k_warps_per_block =
k_warps_per_block_dim_m * k_warps_per_block_dim_n;
constexpr int k_rows_per_warp_group = k_warps_per_block_dim_m * k_tile_m;
constexpr int k_cols_per_warp_group = k_warps_per_block_dim_n * k_tile_n;
constexpr int k_threads_per_block = k_warps_per_block * k_threads_per_warp;
// 每个 warp 处理 4x4 个 tile,一个 block 处理 8x8 个 tile:
// B0 B1 B2 B3 B4 B5 B6 B7
// ┌────────────────────────
// A0 │ 0 1 0 1 0 1 0 1
// A1 │ 2 3 2 3 2 3 2 3
// A2 │ 0 1 0 1 0 1 0 1
// A3 │ 2 3 2 3 2 3 2 3
// A4 │ 0 1 0 1 0 1 0 1
// A5 │ 2 3 2 3 2 3 2 3
// A6 │ 0 1 0 1 0 1 0 1
// A7 │ 2 3 2 3 2 3 2 3
constexpr int k_tiles_per_warp_dim_m = 4;
constexpr int k_tiles_per_warp_dim_n = 4;
constexpr int k_tiles_per_warp =
k_tiles_per_warp_dim_m * k_tiles_per_warp_dim_n;
constexpr int k_tiles_per_block_dim_m =
k_warps_per_block_dim_m * k_tiles_per_warp_dim_m;
constexpr int k_tiles_per_block_dim_n =
k_warps_per_block_dim_n * k_tiles_per_warp_dim_n;
constexpr int k_rows_per_block = k_tiles_per_block_dim_m * k_tile_m;
constexpr int k_cols_per_block = k_tiles_per_block_dim_n * k_tile_n;
// 向量化加载
constexpr int k_vec_size = 8;
#define FLOAT4(value) (reinterpret_cast<float4*>(&(value))[0])
#define CFLOAT4(value) (reinterpret_cast<const float4*>(&(value))[0])
// WMMA
namespace wmma = nvcuda::wmma;
using a_fragment = wmma::fragment<wmma::matrix_a, k_tile_m, k_tile_n, k_tile_k,
half, wmma::row_major>;
using b_fragment = wmma::fragment<wmma::matrix_b, k_tile_m, k_tile_n, k_tile_k,
half, wmma::row_major>;
using acc_fragment =
wmma::fragment<wmma::accumulator, k_tile_m, k_tile_n, k_tile_k, float>;
using c_fragment =
wmma::fragment<wmma::accumulator, k_tile_m, k_tile_n, k_tile_k, half>;
__global__ void gemm_kernel(const half* __restrict__ A,
const half* __restrict__ B, half* __restrict__ C,
int M, int N, int K, float alpha, float beta) {
__shared__ half tile_a[2][k_rows_per_block][k_tile_k + 8];
__shared__ half tile_b[2][k_tile_k][k_cols_per_block + 8];
int stage = 0;
const int block_row = blockIdx.y * k_rows_per_block;
const int block_col = blockIdx.x * k_cols_per_block;
const int warp = threadIdx.x / k_threads_per_warp;
const int warp_id_dim_m = warp / k_warps_per_block_dim_n;
const int warp_id_dim_n = warp % k_warps_per_block_dim_n;
// warp 0 -> row offset 0, col offset 0
// warp 1 -> row offset 0, col offset 16
// warp 2 -> row offset 16, col offset 0
// warp 3 -> row offset 16, col offset 16
const int warp_tile_row = warp_id_dim_m * k_tile_m;
const int warp_tile_col = warp_id_dim_n * k_tile_n;
// 先把 K-stage 的第一部分读进来,后面会边读下一阶段的元素边计算
// 针对矩阵 A 把 block 内的所有 thread 分配到两个维度上
// 0 1 x
// ┌────────
// 0 │ 0 1
// 1 │ 2 3
// 2 │ 4 5
// ... │ ...
// 64 │ 126 127
// y
constexpr int a_dim_x = k_tile_k / k_vec_size,
a_dim_y = k_threads_per_block / a_dim_x;
const int a_thread_id_x = threadIdx.x % a_dim_x;
const int a_thread_id_y = threadIdx.x / a_dim_x;
// 针对矩阵 B 把 block 内的所有 thread 分配到两个维度上
// 0 1 ... 15 x
// ┌────────────────
// 0 │ 0 1 ... 15
// 1 │ 16 17 ... 31
// ... │ ...
// 8 │ 112 113 ... 127
// y
constexpr int b_dim_x = k_cols_per_block / k_vec_size,
b_dim_y = k_threads_per_block / b_dim_x;
const int b_thread_id_x = threadIdx.x % b_dim_x;
const int b_thread_id_y = threadIdx.x / b_dim_x;
#pragma unroll
for (int i = 0; i < k_rows_per_block; i += a_dim_y) {
const int local_row = i + a_thread_id_y;
const int row = block_row + local_row;
const int col = a_thread_id_x * k_vec_size;
FLOAT4(tile_a[stage][local_row][col]) = CFLOAT4(A[OFFSET(row, col, K)]);
}
#pragma unroll
for (int i = 0; i < k_tile_k; i += b_dim_y) {
const int row = i + b_thread_id_y;
const int local_col = b_thread_id_x * k_vec_size;
const int col = block_col + local_col;
FLOAT4(tile_b[stage][row][local_col]) = CFLOAT4(B[OFFSET(row, col, N)]);
}
// B0 B1 B2 B3
// ┌────┬────┬────┬────┐
// A0 │acc0│acc1│acc2│acc3│
// ├────┼────┼────┼────┤
// A1 │acc4│acc5│acc6│acc7│
// ├────┼────┼────┼────┤
// A2 │acc8│acc9│ ...│ ...│
// ├────┼────┼────┼────┤
// A3 │ ...│ ...│ ...│accF│
// └────┴────┴────┴────┘
a_fragment a_frag[k_tiles_per_warp_dim_m];
b_fragment b_frag[k_tiles_per_warp_dim_n];
acc_fragment acc_frag[k_tiles_per_warp];
c_fragment c_frag[k_tiles_per_warp];
#pragma unroll
for (int i = 0; i < k_tiles_per_warp; ++i) {
wmma::fill_fragment(acc_frag[i], 0.0f);
}
// 加载 C 矩阵的 tile 到 c_frag 中
#pragma unroll
for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
const auto row = block_row + (i * k_rows_per_warp_group) + warp_tile_row;
#pragma unroll
for (int j = 0; j < k_tiles_per_warp_dim_n; ++j) {
const auto col = block_col + (j * k_cols_per_warp_group) + warp_tile_col;
wmma::load_matrix_sync(c_frag[OFFSET(i, j, k_tiles_per_warp_dim_n)],
&C[OFFSET(row, col, N)], N, wmma::mem_row_major);
}
}
__syncthreads();
// K-stage 循环,边计算边 prefetch 下一部分的 A 和 B
// 用于存储下一部分 A 和 B 的 tile,这里只存储当前 thread 的对应的 tile 数据
half stage_a[k_rows_per_block / a_dim_y * k_vec_size];
half stage_b[k_tile_k / b_dim_y * k_vec_size];
for (int k = 0; k < K; k += k_tile_k) {
// prefetch 到 stage_a 和 stage_b 中
if (k + k_tile_k < K) {
#pragma unroll
for (int i = 0; i < k_rows_per_block; i += a_dim_y) {
const int row = block_row + i + a_thread_id_y;
const int col = k + k_tile_k + a_thread_id_x * k_vec_size;
const int idx = i / a_dim_y * k_vec_size;
FLOAT4(stage_a[idx]) = CFLOAT4(A[OFFSET(row, col, K)]);
}
#pragma unroll
for (int i = 0; i < k_tile_k; i += b_dim_y) {
const int row = i + k + k_tile_k + b_thread_id_y;
const int col = block_col + b_thread_id_x * k_vec_size;
const int idx = i / b_dim_y * k_vec_size;
FLOAT4(stage_b[idx]) = CFLOAT4(B[OFFSET(row, col, N)]);
}
}
// 从 shared memory 加载 A 和 B 到 fragment
#pragma unroll
for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
const int row = (i * k_rows_per_warp_group) + warp_tile_row;
wmma::load_matrix_sync(a_frag[i], &tile_a[stage][row][0], k_tile_k + 8);
}
#pragma unroll
for (int i = 0; i < k_tiles_per_warp_dim_n; ++i) {
const int col = (i * k_cols_per_warp_group) + warp_tile_col;
wmma::load_matrix_sync(b_frag[i], &tile_b[stage][0][col],
k_cols_per_block + 8);
}
// 执行 MMA 运算
#pragma unroll
for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
#pragma unroll
for (int j = 0; j < k_tiles_per_warp_dim_n; ++j) {
const int acc_idx = OFFSET(i, j, k_tiles_per_warp_dim_n);
wmma::mma_sync(acc_frag[acc_idx], a_frag[i], b_frag[j],
acc_frag[acc_idx]);
}
}
// 将 prefetch 的东西写入 shared memory
if (k + k_tile_k < K) {
#pragma unroll
for (int i = 0; i < k_rows_per_block; i += a_dim_y) {
const int idx = i / a_dim_y * k_vec_size;
FLOAT4(
tile_a[stage ^ 1][i + a_thread_id_y][a_thread_id_x * k_vec_size]) =
FLOAT4(stage_a[idx]);
}
#pragma unroll
for (int i = 0; i < k_tile_k; i += b_dim_y) {
const int idx = i / b_dim_y * k_vec_size;
FLOAT4(
tile_b[stage ^ 1][i + b_thread_id_y][b_thread_id_x * k_vec_size]) =
FLOAT4(stage_b[idx]);
}
stage ^= 1;
__syncthreads();
}
}
// 将结果写回全局内存
#pragma unroll
for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
#pragma unroll
for (int j = 0; j < k_tiles_per_warp_dim_n; ++j) {
const int acc_idx = OFFSET(i, j, k_tiles_per_warp_dim_n);
// fragment 是一个 warp-level 分布式对象,这里当前线程只处理自己负责的元素
for (int t = 0; t < acc_frag[acc_idx].num_elements; ++t) {
// 这里依赖了一个 UB,不一定 c_frag.x[t] 和 acc_frag.x[t] 一一对应
const float ab = acc_frag[acc_idx].x[t];
half& c = c_frag[acc_idx].x[t];
c = __float2half(alpha * ab + beta * __half2float(c));
}
const int row =
block_row + (i * k_warps_per_block_dim_m * k_tile_m) + warp_tile_row;
const int col =
block_col + (j * k_warps_per_block_dim_n * k_tile_n) + warp_tile_col;
wmma::store_matrix_sync(&C[OFFSET(row, col, N)], c_frag[acc_idx], N,
wmma::mem_row_major);
}
}
}
// A, B, and C are device pointers
extern "C" void solve(const half* A, const half* B, half* C, int M, int N,
int K, float alpha, float beta) {
const auto Mp = CEIL(M, k_rows_per_block) * k_rows_per_block;
const auto Np = CEIL(N, k_cols_per_block) * k_cols_per_block;
const auto Kp = CEIL(K, k_tile_k) * k_tile_k;
constexpr auto block_dim = k_threads_per_block;
const auto grid_dim =
dim3(CEIL(Np, k_cols_per_block), CEIL(Mp, k_rows_per_block));
if (Mp == M && Np == N && Kp == K) {
// 如果矩阵大小匹配,直接使用原矩阵
gemm_kernel<<<grid_dim, block_dim>>>(A, B, C, M, N, K, alpha, beta);
} else {
// 如果矩阵大小不匹配,那么申请新的内存,创建 padded 矩阵
half *Ap = nullptr, *Bp = nullptr, *Cp = nullptr;
cudaMalloc(&Ap, Mp * Kp * sizeof(half));
cudaMalloc(&Bp, Kp * Np * sizeof(half));
cudaMalloc(&Cp, Mp * Np * sizeof(half));
cudaMemset(Ap, 0, Mp * Kp * sizeof(half));
cudaMemset(Bp, 0, Kp * Np * sizeof(half));
cudaMemset(Cp, 0, Mp * Np * sizeof(half));
// cudaMemcpy2D(
// dst, // 目标起始地址
// dst_pitch, // 目标每一行的跨度(字节)
// src, // 源起始地址
// src_pitch, // 源每一行的跨度(字节)
// width, // 每行实际复制多少字节
// height, // 一共复制多少行
// kind // 拷贝方向
// );
cudaMemcpy2D(Ap, Kp * sizeof(half), A, K * sizeof(half), K * sizeof(half),
M, cudaMemcpyDeviceToDevice);
cudaMemcpy2D(Bp, Np * sizeof(half), B, N * sizeof(half), N * sizeof(half),
K, cudaMemcpyDeviceToDevice);
cudaMemcpy2D(Cp, Np * sizeof(half), C, N * sizeof(half), N * sizeof(half),
M, cudaMemcpyDeviceToDevice);
gemm_kernel<<<grid_dim, block_dim>>>(Ap, Bp, Cp, Mp, Np, Kp, alpha, beta);
// 拷贝结果回原矩阵
cudaMemcpy2D(C, N * sizeof(half), Cp, Np * sizeof(half), N * sizeof(half),
M, cudaMemcpyDeviceToDevice);
cudaFree(Ap);
cudaFree(Bp);
cudaFree(Cp);
}
}
这里用到了向量化加载,首先将矩阵 Pad 一下,然后再调用 Kernel,结束的时候把结果 Copy 回原始矩阵中,在加载的时候,每次加载 8 个 half 元素
此外,代码里面先读入 K-Stage 的第一阶段要用的数据,然后在 K-Stage 过程中边 Prefetch 下一部分的数据边计算
不 Prefetch 时,搬数据和计算基本是串行的:
K0: |------ global load ------|-- MMA --|
K1: |------ global load ------|-- MMA --|
K2: |------ global load ------|-- MMA --|
当前代码中,使用 Ping-Pong Buffer 将“从 Global Memory 中加载下一阶段的数据”和“MMA 计算”并行执行,流水线变成这样:
K0: |========== MMA K0 ==========|
K1 load: |------ global memory K1 ------|
↓
SMEM buffer 1
K1: |========== MMA K1 ==========|
K2 load: |------ global memory K2 ------|
↓
SMEM buffer 0
这样 Latency 变成了:
$$ T_{\text{GEMM}}:T_{\text{memory}}+T_{\text{compute}}\longrightarrow\max(T_{\text{memory}},T_{\text{compute}}) $$基于 Buffered Prefetch 的普通矩阵乘法
#include <cuda_runtime.h>
#define CEIL(x, y) (((x) + (y) - 1) / (y))
#define FLOAT4(x) (*(reinterpret_cast<float4*>(&(x))))
#define CFLOAT4(x) (*(reinterpret_cast<const float4*>(&(x))))
template <int BX, int BY, int TM, int TN, int KS>
__global__ void mm_kernel(const float* __restrict__ A,
const float* __restrict__ B, float* __restrict__ C,
int M, int N, int K, float scale = 0.0f) {
constexpr int k_vec_size = 4;
constexpr int num_threads = BX * BY;
constexpr int BM = BY * TM;
constexpr int BN = BX * TN;
__shared__ float tile_a[2][BM][KS + 4];
__shared__ float tile_b[2][KS][BN + 4];
int stage = 0;
const int tid = threadIdx.x + threadIdx.y * blockDim.x;
const int block_row = blockIdx.y * BM;
const int block_col = blockIdx.x * BN;
const int thread_row = threadIdx.y * TM;
const int thread_col = threadIdx.x * TN;
constexpr int a_block_dim_x = KS / k_vec_size,
a_block_dim_y = num_threads / a_block_dim_x;
const int a_x = tid % a_block_dim_x;
const int a_y_start = tid / a_block_dim_x;
#pragma unroll
for (int i = 0; i < BM / a_block_dim_y; ++i) {
const int a_y = a_y_start + i * a_block_dim_y;
const int row = block_row + a_y;
const int col = a_x * k_vec_size;
FLOAT4(tile_a[stage][a_y][a_x * k_vec_size]) = CFLOAT4(A[row * K + col]);
}
constexpr int b_block_dim_x = BN / k_vec_size,
b_block_dim_y = num_threads / b_block_dim_x;
const int b_x = tid % b_block_dim_x;
const int b_y_start = tid / b_block_dim_x;
#pragma unroll
for (int i = 0; i < KS / b_block_dim_y; ++i) {
const int b_y = b_y_start + i * b_block_dim_y;
const int row = b_y;
const int col = block_col + b_x * k_vec_size;
FLOAT4(tile_b[stage][b_y][b_x * k_vec_size]) = CFLOAT4(B[row * N + col]);
}
__syncthreads();
float prefetched_tile_a[(BM / a_block_dim_y) * k_vec_size];
float prefetched_tile_b[(KS / b_block_dim_y) * k_vec_size];
float acc[TM][TN] = {0.0f};
for (int k = 0; k < K; k += KS) {
if (k + KS < K) {
#pragma unroll
for (int i = 0; i < BM / a_block_dim_y; ++i) {
const int a_y = a_y_start + i * a_block_dim_y;
const int row = block_row + a_y;
const int col = k + KS + a_x * k_vec_size;
FLOAT4(prefetched_tile_a[i * k_vec_size]) = CFLOAT4(A[row * K + col]);
}
#pragma unroll
for (int i = 0; i < KS / b_block_dim_y; ++i) {
const int b_y = b_y_start + i * b_block_dim_y;
const int row = k + KS + b_y;
const int col = block_col + b_x * k_vec_size;
FLOAT4(prefetched_tile_b[i * k_vec_size]) = CFLOAT4(B[row * N + col]);
}
}
{
#pragma unroll
for (int t = 0; t < KS; ++t) {
float reg_a[TM];
float reg_b[TN];
#pragma unroll
for (int i = 0; i < TM; ++i) {
reg_a[i] = tile_a[stage][thread_row + i][t];
}
#pragma unroll
for (int i = 0; i < TN / k_vec_size; ++i) {
FLOAT4(reg_b[i * k_vec_size]) =
CFLOAT4(tile_b[stage][t][thread_col + i * k_vec_size]);
}
#pragma unroll
for (int i = 0; i < TM; ++i) {
#pragma unroll
for (int j = 0; j < TN; ++j) {
acc[i][j] = fmaf(reg_a[i], reg_b[j], acc[i][j]);
}
}
}
}
if (k + KS < K) {
#pragma unroll
for (int i = 0; i < BM / a_block_dim_y; ++i) {
const int a_y = a_y_start + i * a_block_dim_y;
FLOAT4(tile_a[stage ^ 1][a_y][a_x * k_vec_size]) =
CFLOAT4(prefetched_tile_a[i * k_vec_size]);
}
#pragma unroll
for (int i = 0; i < KS / b_block_dim_y; ++i) {
const int b_y = b_y_start + i * b_block_dim_y;
FLOAT4(tile_b[stage ^ 1][b_y][b_x * k_vec_size]) =
CFLOAT4(prefetched_tile_b[i * k_vec_size]);
}
}
stage ^= 1;
__syncthreads();
}
#pragma unroll
for (int i = 0; i < TM; ++i) {
const int row = block_row + thread_row + i;
#pragma unroll
for (int j = 0; j < TN / k_vec_size; ++j) {
const int col = block_col + thread_col + j * k_vec_size;
FLOAT4(C[row * N + col]) = CFLOAT4(acc[i][j * k_vec_size]);
}
}
}
// A, B, C are device pointers
extern "C" void solve(const float* A, const float* B, float* C, int M, int K,
int N) {
constexpr int BX = 16;
constexpr int BY = 16;
constexpr int TM = 8;
constexpr int TN = 8;
constexpr int BM = BY * TM;
constexpr int BN = BX * TN;
constexpr int KS = 16;
const int Mp = CEIL(M, BM) * BM;
const int Np = CEIL(N, BN) * BN;
const int Kp = CEIL(K, KS) * KS;
constexpr dim3 block_dim(BX, BY);
const dim3 grid_dim(Np / BN, Mp / BM);
if (Mp == M && Np == N && Kp == K) {
mm_kernel<BX, BY, TM, TN, KS><<<grid_dim, block_dim>>>(A, B, C, M, N, K);
} else {
float *Ap = nullptr, *Bp = nullptr, *Cp = nullptr;
cudaMalloc(&Ap, Mp * Kp * sizeof(float));
cudaMalloc(&Bp, Kp * Np * sizeof(float));
cudaMalloc(&Cp, Mp * Np * sizeof(float));
cudaMemsetAsync(Ap, 0, Mp * Kp * sizeof(float));
cudaMemsetAsync(Bp, 0, Kp * Np * sizeof(float));
// A: [M, K] -> [Mp, Kp]
cudaMemcpy2DAsync(Ap, Kp * sizeof(float), A, K * sizeof(float),
K * sizeof(float), M, cudaMemcpyDeviceToDevice);
// B: [K, N] -> [Kp, Np]
cudaMemcpy2DAsync(Bp, Np * sizeof(float), B, N * sizeof(float),
N * sizeof(float), K, cudaMemcpyDeviceToDevice);
mm_kernel<BX, BY, TM, TN, KS>
<<<grid_dim, block_dim>>>(Ap, Bp, Cp, Mp, Np, Kp);
// C: [Mp, Np] -> [M, N]
cudaMemcpy2DAsync(C, N * sizeof(float), Cp, Np * sizeof(float),
N * sizeof(float), M, cudaMemcpyDeviceToDevice);
cudaFree(Cp);
cudaFree(Bp);
cudaFree(Ap);
}
}
使用 Swizzle 来解决 Bank Conflict 问题
上面的代码仍然存在 Bank Conflict
从 tila_a 读取到 reg_a 时,也就是这一句话:
reg_a[i] = tile_a[stage][thread_row + i][t];
固定 i 和 t,所有的线程都读取的第 t 列,且行数随着 threadIdx.y 变化。如果 blockDim.x 为 16,会出现:
lane 0~15:threadIdx.y = 2w, threadIdx.x = 0~15, pos = [2w*TM+i][t], address = (2w*TM+i)*(KS+4)+t
lane 16~31:threadIdx.y = 2w+1, threadIdx.x = 0~15, pos = [(2w+1)*TM+i][t], address = ((2w+1)*TM+i)*(KS+4)+t
地址相差 TM*(KS+4),由于要用 float4,这个值很容易是 32 的倍数。所以前半个 Warp 和后半个 Warp 分别是广播,两个 Warp 处于同一个 Bank Conflict,是一个 2-way Bank Conflict
这几个 Shared Memory 访问的状态:
| Shared 访问 | Bank Conflict |
|---|---|
A:global/register → shared,float4 写入 | 2-way Conflict |
| A:shared → register,标量读取 | 2-way Conflict |
B:global/register → shared,float4 写入 | 无冲突 |
B:shared → register,float4 读取 | 2-way Conflict |
很难用 Padding 去解决这里的 Bank Conflict 问题,甚至我这里 Padding 为 4 还增加了一个 Bank Conflict,就是 A 的写入。所以引入 Swizzle 来解决
Swizzle 用来改变二维数组在共享内存中的映射规则,一种常见且高效的映射策略是使用按位异或(XOR):
def mapping(y, x):
return y, x ^ y
它只映射了列,在映射的时候连着行一起考虑了进来,效果:

如果是 $2^k$ 乘 $2^k$ 的方阵,那么这样完全可以,但是对于不是这样的矩阵,就得用复杂一点的方法了,而且需要具体问题具体分析
比如这里的 Bank Conflict,可以这样,先去掉 Padding,这样只有两个 Bank Conflict 了,一个是 A 的读取,一个是 B 的读取
对于 A 的读取,如果 blockDim.x 为 16,只需要保证上下两个半 Warp 处于不同的列就可以了,对于 blockDim.x 为 8 的情况,需要保证四组 Thread Group 处于不同的列,以此类推,可以这么搞:
template <int TM, int KS>
__device__ __forceinline__ int swizzle_a_col(int row, int col) {
// 每行有多少个 float4,由于仍然需要向量化读取,所以这里把 float4
// 作为一个整体,KS 是 tile_a 的列数
constexpr int Q = KS / k_vec_size;
// 由于一个 warp 只有 32 个线程,也就是 8 个 float 4
// KS <= 32 时,每个 8-lane 写入分组覆盖若干完整的行,映射只置换每行内部的
// float4 块,覆盖的 bank 集合不变
// KS > 32 时,每个分组写入同一行的连续 8 个 float4,映射只置换这 8 个向量块
constexpr int L = 32 / k_vec_size;
constexpr int G = Q < L ? Q : L;
// 处于哪一个 thread group
// row = threadIdx.y * TM + i,所以 row / TM == threadIdx.y
const int group = (row / TM) & (G - 1);
// group=0: col^0, eg. col=8 -> col^0=8 col=12 -> col^0=12
// group=1: col^4, eg. col=8 -> col^4=12 col=12 -> col^4=8
// 也就是对于 group=1 的行,两两交换相邻的 float4 的实际位置
return col ^ (group * k_vec_size);
}
对于 B,同理
代码:
#include <cuda_runtime.h>
constexpr int k_vec_size = 4;
__host__ __device__ constexpr int ceil(int x, int y) {
return (x + y - 1) / y;
}
template <typename T>
__device__ constexpr float4& to_float4(T& x) {
return *(reinterpret_cast<float4*>(&(x)));
}
template <typename T>
const __device__ constexpr float4& to_cfloat4(const T& x) {
return *(reinterpret_cast<const float4*>(&(x)));
}
template <int TM, int KS>
__device__ __forceinline__ int swizzle_a_col(int row, int col) {
// 每行有多少个 float4,由于仍然需要向量化读取,所以这里把 float4
// 作为一个整体,KS 是 tile_a 的列数
constexpr int Q = KS / k_vec_size;
// 由于一个 warp 只有 32 个线程,也就是 8 个 float 4
// KS <= 32 时,每个 8-lane 写入分组覆盖若干完整的行,映射只置换每行内部的
// float4 块,覆盖的 bank 集合不变
// KS > 32 时,每个分组写入同一行的连续 8 个 float4,映射只置换这 8 个向量块
constexpr int L = 32 / k_vec_size;
constexpr int G = Q < L ? Q : L;
// 处于哪一个 thread group
// row = threadIdx.y * TM + i,所以 row / TM == threadIdx.y
const int group = (row / TM) & (G - 1);
// group=0: col^0, eg. col=8 -> col^0=8 col=12 -> col^0=12
// group=1: col^4, eg. col=8 -> col^4=12 col=12 -> col^4=8
// 也就是对于 group=1 的行,两两交换相邻的 float4 的实际位置
return col ^ (group * k_vec_size);
}
__host__ __device__ constexpr int ilog2_pow2(int x) {
int result = 0;
while (x > 1) {
x >>= 1;
++result;
}
return result;
}
template <int TN>
__device__ __forceinline__ int swizzle_b_col(int col) {
constexpr int d = ilog2_pow2(TN / 4);
constexpr int shift = d > 3 ? d : 3;
constexpr int bits = d < 3 ? d : 3;
constexpr int mask = (1 << bits) - 1;
const int q = col >> 2;
const int p = q ^ ((q >> shift) & mask);
return (p << 2) | (col & 3);
}
template <int BX, int BY, int TM, int TN, int KS>
__global__ void mm_kernel(const float* __restrict__ A,
const float* __restrict__ B, float* __restrict__ C,
int M, int N, int K, float scale = 0.0f) {
constexpr int num_threads = BX * BY;
constexpr int BM = BY * TM;
constexpr int BN = BX * TN;
__shared__ float tile_a[2][BM][KS];
__shared__ float tile_b[2][KS][BN];
int stage = 0;
const int tid = threadIdx.x + threadIdx.y * blockDim.x;
const int block_row = blockIdx.y * BM;
const int block_col = blockIdx.x * BN;
const int thread_row = threadIdx.y * TM;
const int thread_col = threadIdx.x * TN;
constexpr int a_block_dim_x = KS / k_vec_size,
a_block_dim_y = num_threads / a_block_dim_x;
const int a_x = tid % a_block_dim_x;
const int a_y_start = tid / a_block_dim_x;
#pragma unroll
for (int i = 0; i < BM / a_block_dim_y; ++i) {
const int a_y = a_y_start + i * a_block_dim_y;
const int row = block_row + a_y;
const int col = a_x * k_vec_size;
to_float4(
tile_a[stage][a_y][swizzle_a_col<TM, KS>(a_y, a_x * k_vec_size)]) =
to_cfloat4(A[row * K + col]);
}
constexpr int b_block_dim_x = BN / k_vec_size,
b_block_dim_y = num_threads / b_block_dim_x;
const int b_x = tid % b_block_dim_x;
const int b_y_start = tid / b_block_dim_x;
#pragma unroll
for (int i = 0; i < KS / b_block_dim_y; ++i) {
const int b_y = b_y_start + i * b_block_dim_y;
const int row = b_y;
const int col = block_col + b_x * k_vec_size;
to_float4(tile_b[stage][b_y][swizzle_b_col<TN>(b_x * k_vec_size)]) =
to_cfloat4(B[row * N + col]);
}
__syncthreads();
float prefetched_tile_a[(BM / a_block_dim_y) * k_vec_size];
float prefetched_tile_b[(KS / b_block_dim_y) * k_vec_size];
float acc[TM][TN] = {0.0f};
for (int k = 0; k < K; k += KS) {
if (k + KS < K) {
#pragma unroll
for (int i = 0; i < BM / a_block_dim_y; ++i) {
const int a_y = a_y_start + i * a_block_dim_y;
const int row = block_row + a_y;
const int col = k + KS + a_x * k_vec_size;
to_float4(prefetched_tile_a[i * k_vec_size]) =
to_cfloat4(A[row * K + col]);
}
#pragma unroll
for (int i = 0; i < KS / b_block_dim_y; ++i) {
const int b_y = b_y_start + i * b_block_dim_y;
const int row = k + KS + b_y;
const int col = block_col + b_x * k_vec_size;
to_float4(prefetched_tile_b[i * k_vec_size]) =
to_cfloat4(B[row * N + col]);
}
}
{
#pragma unroll
for (int t = 0; t < KS; ++t) {
float reg_a[TM];
float reg_b[TN];
#pragma unroll
for (int i = 0; i < TM; ++i) {
const int r = thread_row + i;
reg_a[i] = tile_a[stage][r][swizzle_a_col<TM, KS>(r, t)];
}
#pragma unroll
for (int i = 0; i < TN / k_vec_size; ++i) {
const int c = thread_col + i * 4;
to_float4(reg_b[i * k_vec_size]) =
to_cfloat4(tile_b[stage][t][swizzle_b_col<TN>(c)]);
}
#pragma unroll
for (int i = 0; i < TM; ++i) {
#pragma unroll
for (int j = 0; j < TN; ++j) {
acc[i][j] = fmaf(reg_a[i], reg_b[j], acc[i][j]);
}
}
}
}
if (k + KS < K) {
#pragma unroll
for (int i = 0; i < BM / a_block_dim_y; ++i) {
const int a_y = a_y_start + i * a_block_dim_y;
to_float4(tile_a[stage ^ 1][a_y]
[swizzle_a_col<TM, KS>(a_y, a_x * k_vec_size)]) =
to_cfloat4(prefetched_tile_a[i * k_vec_size]);
}
#pragma unroll
for (int i = 0; i < KS / b_block_dim_y; ++i) {
const int b_y = b_y_start + i * b_block_dim_y;
to_float4(tile_b[stage ^ 1][b_y][swizzle_b_col<TN>(b_x * k_vec_size)]) =
to_cfloat4(prefetched_tile_b[i * k_vec_size]);
}
}
stage ^= 1;
__syncthreads();
}
#pragma unroll
for (int i = 0; i < TM; ++i) {
const int row = block_row + thread_row + i;
#pragma unroll
for (int j = 0; j < TN / k_vec_size; ++j) {
const int col = block_col + thread_col + j * k_vec_size;
to_float4(C[row * N + col]) = to_cfloat4(acc[i][j * k_vec_size]);
}
}
}
// A, B, C are device pointers
extern "C" void solve(const float* A, const float* B, float* C, int M, int K,
int N) {
constexpr int BX = 16;
constexpr int BY = 16;
constexpr int TM = 8;
constexpr int TN = 8;
constexpr int BM = BY * TM;
constexpr int BN = BX * TN;
constexpr int KS = 16;
const int Mp = ceil(M, BM) * BM;
const int Np = ceil(N, BN) * BN;
const int Kp = ceil(K, KS) * KS;
constexpr dim3 block_dim(BX, BY);
const dim3 grid_dim(Np / BN, Mp / BM);
if (Mp == M && Np == N && Kp == K) {
mm_kernel<BX, BY, TM, TN, KS><<<grid_dim, block_dim>>>(A, B, C, M, N, K);
} else {
float *Ap = nullptr, *Bp = nullptr, *Cp = nullptr;
cudaMalloc(&Ap, Mp * Kp * sizeof(float));
cudaMalloc(&Bp, Kp * Np * sizeof(float));
cudaMalloc(&Cp, Mp * Np * sizeof(float));
cudaMemsetAsync(Ap, 0, Mp * Kp * sizeof(float));
cudaMemsetAsync(Bp, 0, Kp * Np * sizeof(float));
// A: [M, K] -> [Mp, Kp]
cudaMemcpy2DAsync(Ap, Kp * sizeof(float), A, K * sizeof(float),
K * sizeof(float), M, cudaMemcpyDeviceToDevice);
// B: [K, N] -> [Kp, Np]
cudaMemcpy2DAsync(Bp, Np * sizeof(float), B, N * sizeof(float),
N * sizeof(float), K, cudaMemcpyDeviceToDevice);
mm_kernel<BX, BY, TM, TN, KS>
<<<grid_dim, block_dim>>>(Ap, Bp, Cp, Mp, Np, Kp);
// C: [Mp, Np] -> [M, N]
cudaMemcpy2DAsync(C, N * sizeof(float), Cp, Np * sizeof(float),
N * sizeof(float), M, cudaMemcpyDeviceToDevice);
cudaFree(Cp);
cudaFree(Bp);
cudaFree(Ap);
}
}