§ 参考资料 共 6 条
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);
}
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);
}
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}}) $$