CUDA 算子优化:GEMM

1955 字
10 分钟
CUDA 算子优化:GEMM
2026-07-23

GEMM(General Matrix Multiplication)是深度学习中最核心的计算模式之一。 Transformer 中的 Linear、Attention 投影以及 MLP 层,本质上都可以归结为矩阵乘法。 因此,GPU 对 GEMM 的优化能力直接决定了 AI 训练和推理性能。

本文将从一个最简单的 GEMM kernel 出发,逐步介绍 GPU GEMM 优化中的几个核心思想:

  1. 如何将矩阵计算映射到 GPU 线程;
  2. 如何利用 GPU memory hierarchy 减少数据访问;
  3. 如何提高单个线程的计算密度。

矩阵乘法#

对于矩阵乘法

C=A×BC = A \times B

其中:

  • A 的形状为 M×KM \times K
  • B 的形状为 K×NK \times N
  • C 的形状为 M×NM \times N

输出矩阵 C 中的每个元素通过下面的公式计算得到:

Cij=k=0K1AikBkjC_{ij} = \sum_{k = 0}^{K - 1}A_{ik}B_{kj}

也就是说,每个输出元素都需要计算 A 的一行和 B 的一列之间的点积。

矩阵乘法计算过程
矩阵乘法计算过程

对于整个矩阵:

CRM×NC\in \mathbb{R}^{M\times N}

总计算量约为:

2MNK2MNK

而输入矩阵大小为:

MK+KNMK+KN

因此 GEMM 的核心特点是:

Important

计算量大,并且存在大量数据复用。

如何充分利用这种数据复用,是 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;
}

Naive GEMM 实现
Naive GEMM 实现

这种实现能够利用 GPU 的大量线程并行计算,但是存在一个严重的问题:线程之间没有共享数据

例如计算 C0,0,C0,1,C0,2C_{0,0}, C_{0,1}, C_{0,2},三个线程都需要读取

A0,0,A0,1,...,A0,K1A_{0,0}, A_{0,1}, ..., A_{0,K-1}

也就是说,同一行 A 数据会被多个线程重复从 Global Memory 中读取。而 Global Memory 是 GPU 中访问速度最慢的存储层级之一。因此,Naive GEMM 的瓶颈并不是计算能力不足,而是:

Important

没有利用线程之间的数据复用。

这引出了 GEMM 优化的第一个方向:利用 Shared Memory 缓存公共数据

Shared Memory Tiling:减少 Global Memory 访问#

GPU 中存在多级 memory hierarchy:

Register
|
Shared Memory
|
Global Memory

其中:

  • Register 速度最快,但是空间有限;
  • Shared Memory 位于线程块内部,可以被 block 中所有线程共享;
  • Global Memory 容量最大,但是访问延迟较高。

因此,高性能 CUDA kernel 通常会尽量:

Important

将频繁访问的数据从 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
|
v
Shared Memory
|
v
Matrix 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。

Note

不要说 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];
}
}
}

支持与分享

如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!

赞助
CUDA 算子优化:GEMM
https://llm-tech.com.cn/posts/cuda-gemm/
作者
Ming
发布于
2026-07-23
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author
Ming
你是来找 Ming 学习的吗
🎉 欢迎来到 Ming 的博客
这里是我的个人博客,分享 AI Infra、LLM 等技术内容。欢迎关注交流!
分类
标签
站点统计
文章
19
分类
8
标签
16
总字数
55,114
运行时长
0
最后活动
0 天前

目录