CUDA 算子优化:GEMM
GEMM(General Matrix Multiplication)是深度学习中最核心的计算模式之一。 Transformer 中的 Linear、Attention 投影以及 MLP 层,本质上都可以归结为矩阵乘法。 因此,GPU 对 GEMM 的优化能力直接决定了 AI 训练和推理性能。
本文将从一个最简单的 GEMM kernel 出发,逐步介绍 GPU GEMM 优化中的几个核心思想:
- 如何将矩阵计算映射到 GPU 线程;
- 如何利用 GPU memory hierarchy 减少数据访问;
- 如何提高单个线程的计算密度。
矩阵乘法
对于矩阵乘法
其中:
- A 的形状为
- B 的形状为
- C 的形状为
输出矩阵 C 中的每个元素通过下面的公式计算得到:
也就是说,每个输出元素都需要计算 A 的一行和 B 的一列之间的点积。

对于整个矩阵:
总计算量约为:
而输入矩阵大小为:
因此 GEMM 的核心特点是:
计算量大,并且存在大量数据复用。
如何充分利用这种数据复用,是 GPU GEMM 优化的关键。
Naive GEMM:让 GPU 并行计算
基于 GEMM 的计算过程,我们可以很容易写出下面的 kernel 代码,每一个线程去处理一个输出元素,读取对应的行和列,进行点积运算,然后将结果写回 memory 中。
例如:
C[row][col]对应一个 thread。
该 thread 需要读取:
- A 的第 row 行
- B 的第 col 列
然后完成点积运算。
__global__ void gemm_naive (const float* A, const float* B, float* C, int M, int N, int K){ int col = blockDim.x * blockIdx.x + threadIdx.x; int row = blockDim.y * blockIdx.y + threadIdx.y; if (row >= M || col >= N) return;
float accum = 0.0f; for (int i = 0; i < K; i++) { accum += A[row * K + i] * B[i * N + col]; } C[row * N + col] = accum;}
这种实现能够利用 GPU 的大量线程并行计算,但是存在一个严重的问题:线程之间没有共享数据。
例如计算 ,三个线程都需要读取
也就是说,同一行 A 数据会被多个线程重复从 Global Memory 中读取。而 Global Memory 是 GPU 中访问速度最慢的存储层级之一。因此,Naive GEMM 的瓶颈并不是计算能力不足,而是:
没有利用线程之间的数据复用。
这引出了 GEMM 优化的第一个方向:利用 Shared Memory 缓存公共数据。
Shared Memory Tiling:减少 Global Memory 访问
GPU 中存在多级 memory hierarchy:
Register |Shared Memory |Global Memory其中:
- Register 速度最快,但是空间有限;
- Shared Memory 位于线程块内部,可以被 block 中所有线程共享;
- Global Memory 容量最大,但是访问延迟较高。
因此,高性能 CUDA kernel 通常会尽量:
将频繁访问的数据从 Global Memory 搬运到 Shared Memory。
在 naive GEMM 中,一个 thread block 已经覆盖了 C 矩阵中的一个区域,但是 block 内线程之间彼此独立,每个线程都会从 Global Memory 中读取自己需要的数据。为了减少重复访问,我们让 block 内线程协同加载 A 和 B 的 tile 到 Shared Memory,并利用这些共享数据完成计算。

整体的计算流程大致如下:

对应代码:
#define BLOCK_SIZE 32__global__ void gemm_v1 (const float* A, const float* B, float* C, int M, int N, int K){ int col = blockDim.x * blockIdx.x + threadIdx.x; int row = blockDim.y * blockIdx.y + threadIdx.y; if (row >= M || col >= N) return;
int bx = blockIdx.x; int by = blockIdx.y; int tx = threadIdx.x; int ty = threadIdx.y;
const int BM = BLOCK_SIZE; const int BN = BLOCK_SIZE; const int BK = BLOCK_SIZE;
__shared__ float As[BM * BK]; __shared__ float Bs[BK * BN];
A = &A[(by * BM) * K]; B = &B[bx * BN]; C = &C[(by * BM) * N + bx * BN];
float accum = 0.0f; for (int k = 0; k < K; k += BK) { // global ==> shared As[ty * BK + tx] = A[ty * K + tx]; Bs[ty * BN + tx] = B[ty * N + tx];
__syncthreads ();
A = A + BK; B = B + BK * N; for (int i = 0; i < BK; i++) { accum += As[ty * BK + i] * Bs[i * BN + tx]; } __syncthreads (); } C[ty * N + tx] = accum;}在 kernel 中:
__shared__ float As[BM * BK];__shared__ float Bs[BK * BN];定义两个 shared memory tile:
As保留 A 的局部块;Bs保留 B 的局部块。
每轮循环:
for (int k = 0; k < K; k += BK)完成一次:
Global Memory | vShared Memory | vMatrix Multiply因此,相比 naive GEMM:
- Global Memory 访问次数降低;
- 数据复用增加;
- 算术强度(Arithmetic Intensity)提高。
经过 Shared Memory 优化后,Global Memory 访问已经大幅减少,但是计算过程仍然存在进一步优化空间。
在 GPU 中,数据访问速度存在明显差异:
Register ↑Shared Memory ↑Global Memory其中 register 是线程私有的高速存储空间。CUDA 中的局部变量通常会被编译器放入寄存器中,因此如果一个线程能够同时计算多个输出元素,这些中间结果可以保存在寄存器中,减少重复的数据访问。
Thread Tile:提高线程计算密度
Shared Memory Tiling 解决了线程块之间的数据复用问题,但是此时每个线程仍然只负责计算一个输出元素:
accum += As[...] * Bs[...];一个线程完成一次矩阵乘累加后立即结束,导致线程内部计算量较低。
为了进一步提高计算资源利用率,可以让一个线程负责多个输出元素。这样,一个线程加载的数据可以参与更多计算,同时多个输出的累加结果可以保存在 register 中。
例如,一个 thread 不再计算:
C[0][0]而是计算:
C[0][0] C[0][1]C[1][0] C[1][1]这就是 thread tile。
不要说 Thread Tile 是为了利用寄存器,这种说法会稍微本末倒置。 更准确的因果关系应该是:
想提高线程计算量 ↓一个线程计算多个输出 ↓需要保存多个累加结果 ↓这些结果放在寄存器 ↓获得更高计算密度也就是说 Thread Tile 是优化目标,Register 是实现手段。

代码如下:
template <const int BM, const int BN, const int BK, const int TM, const int TN>__global__ void sgemm(float* A, float* B, float* C, int M, int N, int K){ int bx = blockIdx.x; int by = blockIdx.y;
int block_row_thread = BN / TN; // block中一行的thread数量 int block_col_thread = BM / TM; // block中一列的thread数量 int thread_num = block_row_thread * block_col_thread; // block中thread总量
int tx = (threadIdx.x % block_row_thread) * TN; // threadtile左上角x坐标 int ty = (threadIdx.x / block_row_thread) * TM; // threadtile左上角y坐标
__shared__ float As[BM * BK]; __shared__ float Bs[BK * BN];
A = &A[by * BM * K]; B = &B[bx * BN]; C = &C[by * BM * N + bx * BN];
int a_tile_row = threadIdx.x / BK; int a_tile_col = threadIdx.x % BK; int a_tile_stride = thread_num / BK; // BM/(BM/(thread_num/BK)) = thread_num/BK = stride
int b_tile_row = threadIdx.x / BN; int b_tile_col = threadIdx.x % BN; int b_tile_stride = thread_num / BN;
float accum[TM][TN] = { 0.0f }; for (int k = 0; k < K; k += BK) { for (int i = 0; i < BM; i += a_tile_stride) { As[(a_tile_row + i) * BK + a_tile_col] = A[(a_tile_row + i) * K + a_tile_col]; } for (int i = 0; i < BK; i += b_tile_stride) { Bs[(b_tile_row + i) * BN + b_tile_col] = B[(b_tile_row + i) * N + b_tile_col]; } __syncthreads();
A += BK; B += BK * N;
for (int row = 0; row < TM; row++) { for (int col = 0; col < TN; col++) { for (int i = 0; i < BK; i++) { accum[row][col] += As[(ty + row) * BK + i] * Bs[i * BN + (tx + col)]; } } } __syncthreads (); } for (int row = 0; row < TM; row++) { for (int col = 0; col < TN; col++) { C[(ty + row) * N + (tx + col)] = accum[row][col]; } }}支持与分享
如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!