GEMM/GEMV 优化
- GEMV(矩阵向量乘法)是典型的内存瓶颈负载,优化重点是带宽利用率
- GEMM(矩阵矩阵乘法)是计算瓶颈负载,优化重点是计算单元利用率
- 分块(Tiling)是核心策略,平衡容量与计算密度
- 充分利用张量核心(Tensor Core)等专用硬件单元
阅读本文时,建议先判断算子属于 GEMV 还是 GEMM,再选择对应优化路径。GEMV 优先检查访存是否合并;GEMM 优先检查分块复用、寄存器累加和张量核心使用。
| 场景 | 主要瓶颈 | 优先检查项 |
|---|---|---|
| GEMV | 内存带宽 | 合并访存、向量化加载、向量缓存 |
| GEMM | 计算吞吐 | 分块复用、寄存器累加、张量核心 |
基础示例:从基础实现(Naive)到共享内存优化
下面的示例先给出直接实现,再给出共享内存分块实现。两者计算同一个矩阵乘法,差异在于是否把 A/B 的局部数据缓存到共享内存中复用。
- 基础实现
- 共享内存分块
基础实现(Naive,无共享内存)
每个线程直接从全局内存读取数据,计算 C 矩阵的一个元素:
// 矩阵存储结构
typedef struct {
int width;
int height;
float* elements;
} Matrix;
#define BLOCK_SIZE 16
__global__ void MatMulNaive(Matrix A, Matrix B, Matrix C)
{
// 每个线程计算 C 的一个元素
float Cvalue = 0;
int row = blockIdx.y * blockDim.y + threadIdx.y;
int col = blockIdx.x * blockDim.x + threadIdx.x;
if (row < C.height && col < C.width) {
for (int e = 0; e < A.width; ++e) {
Cvalue += A.elements[row * A.width + e] * B.elements[e * B.width + col];
}
C.elements[row * C.width + col] = Cvalue;
}
}
// 启动配置
dim3 dimBlock(BLOCK_SIZE, BLOCK_SIZE);
dim3 dimGrid((B.width + dimBlock.x - 1) / dimBlock.x,
(A.height + dimBlock.y - 1) / dimBlock.y);
MatMulNaive<<<dimGrid, dimBlock>>>(A, B, C);
查看问题分析
问题分析:
| 矩阵 | 全局内存读取次数 | 原因 |
|---|---|---|
| A | B.width 次 | 每个线程独立读取 A 的行元素 |
| B | A.height 次 | 每个线程独立读取 B 的列元素 |
示例:对于 1024×1024 矩阵,BLOCK_SIZE=16
- 全局内存读取:约 2 × 1024 × 1024 × 16 = 33M 次
- 带宽利用率极低
共享内存优化实现
使用分块(Tiling)策略,将数据加载到共享内存后重复使用:
// 增加 stride 字段以支持子矩阵
typedef struct {
int width;
int height;
int stride;
float* elements;
} Matrix;
#define BLOCK_SIZE 16
// 获取子矩阵
__device__ Matrix GetSubMatrix(Matrix M, int row, int col)
{
Matrix Msub;
Msub.width = BLOCK_SIZE;
Msub.height = BLOCK_SIZE;
Msub.stride = M.stride;
Msub.elements = &M.elements[M.stride * BLOCK_SIZE * row + BLOCK_SIZE * col];
return Msub;
}
__global__ void MatMulTiled(Matrix A, Matrix B, Matrix C)
{
int blockRow = blockIdx.y;
int blockCol = blockIdx.x;
// 每个线程块计算 C 的一个子矩阵
Matrix Csub = GetSubMatrix(C, blockRow, blockCol);
float Cvalue = 0;
int row = threadIdx.y;
int col = threadIdx.x;
__shared__ float As[BLOCK_SIZE][BLOCK_SIZE];
__shared__ float Bs[BLOCK_SIZE][BLOCK_SIZE];
// 分块迭代
for (int m = 0; m < (A.width + BLOCK_SIZE - 1) / BLOCK_SIZE; ++m) {
// 获取 A、B 的子矩阵
Matrix Asub = GetSubMatrix(A, blockRow, m);
Matrix Bsub = GetSubMatrix(B, m, blockCol);
// 协作加载:每个线程加载一个元素
int a_global_row = blockRow * BLOCK_SIZE + row;
int a_global_col = m * BLOCK_SIZE + col;
int b_global_row = m * BLOCK_SIZE + row;
int b_global_col = blockCol * BLOCK_SIZE + col;
As[row][col] = (a_global_row < A.height && a_global_col < A.width)
? Asub.elements[row * Asub.stride + col]
: 0.0f;
Bs[row][col] = (b_global_row < B.height && b_global_col < B.width)
? Bsub.elements[row * Bsub.stride + col]
: 0.0f;
// 同步:确保所有数据加载完成
__syncthreads();
// 在共享内存中计算
for (int e = 0; e < BLOCK_SIZE; ++e) {
Cvalue += As[row][e] * Bs[e][col];
}
__syncthreads(); // 确保计算完成后再加载下一块数据
}
int c_global_row = blockRow * BLOCK_SIZE + row;
int c_global_col = blockCol * BLOCK_SIZE + col;
if (c_global_row < C.height && c_global_col < C.width) {
Csub.elements[row * Csub.stride + col] = Cvalue;
}
}
查看优化效果和关键设计
优化效果:
| 矩阵 | 全局内存读取次数 | 降低倍数 |
|---|---|---|
| A | B.width / BLOCK_SIZE 次 | 16 倍 |
| B | A.height / BLOCK_SIZE 次 | 16 倍 |
关键设计:
- 分块策略:将大矩阵分解为 BLOCK_SIZE×BLOCK_SIZE 的小块
- 数据复用:每个块内元素被 BLOCK_SIZE 个线程重复使用
- 协作加载:每个线程只加载一个元素,避免重复
- 同步原语:
__syncthreads()确保数据一致性
访存趋势对比
对于 1024×1024 矩阵乘法:
| 实现 | 全局内存读取趋势 | 说明 |
|---|---|---|
| 基础实现 | 重复读取 A/B 元素 | 实现简单,但全局内存访问压力高 |
| 共享内存分块 | 每个块内复用 A/B 元素 | 减少重复读取,实际收益需结合矩阵形状和硬件实测 |
GEMV 优化(矩阵×向量)
GEMV 定义与特征
GEMV(矩阵向量乘法)常出现在小批量推理、归约后的线性变换、向量投影等场景。它的计算量相对有限,但需要持续读取矩阵和向量数据,因此通常更受内存带宽限制,而不是受计算单元峰值能力限制。
GEMV(General Matrix-Vector Multiplication) 操作:
其中:
- 是 矩阵
- 是 维向量
- 是 维输出向量
算术强度(以 float 为例):
特征:
- 算存比远低于 1 → 典型的内存瓶颈负载
- 性能受限于内存带宽,而非计算能力
GEMV 逐步优化策略
GEMV 优化应按顺序推进:先保证合并访存,再做向量化加载,最后根据向量复用情况考虑共享内存缓存。下面三个标签页展示同一优化路径的不同阶段。
| 阶段 | 解决的问题 | 适用判断 |
|---|---|---|
| 基础并行化 | 建立行级并行 | 需要先把单线程点积改为线程束协作 |
| 向量化加载 | 减少访存指令压力 | 已满足连续地址访问和对齐要求 |
| 缓存向量 x | 降低 x 的重复读取 | 同一线程块内多个输出行复用同一段 x |
- 基础并行化
- 向量化加载
- 缓存向量 x
步骤 1: 基础并行化(线程束级)
最直接的并行方式是让一个线程束协作计算一个输出元素。线程束内每个线程负责一部分 K 维度,最后通过线程束内归约得到该行的点积结果。
__device__ float warpReduceSum(float value) {
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_down_sync(0xffffffff, value, offset);
}
return value;
}
// 一个线程 束计算一个输出元素,每个线程负责部分 K 维度
__global__ void gemvStep1(float* A, float* x, float* y, int m, int k) {
int row = blockIdx.x;
if (row >= m) return;
int lane = threadIdx.x % 32;
float partial = 0.0f;
for (int col = lane; col < k; col += 32) {
partial += A[row * k + col] * x[col];
}
// 线程束内归约,得到该行输出
partial = warpReduceSum(partial);
if (lane == 0) {
y[row] = partial;
}
}
// 启动配置
int threadsPerBlock = 32;
int blocksPerGrid = m;
gemvStep1<<<blocksPerGrid, threadsPerBlock>>>(A, x, y, m, k);
查看该阶段的限制
问题:
- 每个线程每次只加载 32-bit 数据,访存指令数量较多
x会被不同行重复读取,缓存和外部内存带宽压力较高- 当 较大时,可能需要多个线程束协作计算一个输出元素,并通过共享内存完成跨线程束归约
步骤 2: 向 量化加载(以合并访存为前提)
在 GEMV 中,访存模式比单纯增加计算更关键。只有当同一线程束内相邻线程访问连续地址时,硬件才能合并内存事务;在此基础上再使用向量化加载,才能减少 LSU 指令压力并提高带宽利用率。
以下示例复用上文的 warpReduceSum,并假设 A 的每行起始地址和 x 的起始地址满足 16 字节对齐。
// 一个线程束计算一行;同一线程束内相邻线程加载连续的 float4 数据
__global__ void gemvStep2(float* A, float* x, float* y, int m, int k) {
int warp_id = threadIdx.x / 32;
int lane = threadIdx.x % 32;
int warps_per_block = blockDim.x / 32;
int row = blockIdx.x * warps_per_block + warp_id;
if (row >= m) return;
float partial = 0.0f;
int vec_end = (k / 4) * 4;
// 每个 lane 读取一个 float4;相邻 lane 读取连续数据段
for (int col = lane * 4; col < vec_end; col += 32 * 4) {
float4 a_val = *reinterpret_cast<float4*>(&A[row * k + col]);
float4 x_val = *reinterpret_cast<float4*>(&x[col]);
partial += a_val.x * x_val.x;
partial += a_val.y * x_val.y;
partial += a_val.z * x_val.z;
partial += a_val.w * x_val.w;
}
// 处理尾部非 4 对齐元素
for (int col = vec_end + lane; col < k; col += 32) {
partial += A[row * k + col] * x[col];
}
// 线程束内归约 partial,得到该行输出
partial = warpReduceSum(partial);
if (lane == 0) {
y[row] = partial;
}
}
查看优化要点
优化要点:
- GEMV 首先要保证矩阵 A 的读取是合并访存,否则向量化和缓存优化收益有限。
- 合并访存后,可在满足对齐要求时使用 128-bit 或更宽的向量化加载进一步减少访存指令数量;MT S5000 支持最高 1024-bit 的单条访存请求。
- 使用
float4时需要保证向量化加载地址满足对齐要求。 - 块内归约需要使用高效的线程束级或线程块级归约实现。
步骤 3: 共享内存缓存向量 x
矩阵 A 的每一行只会被对应输出元素消费,但向量 x 会被多个输出行反复读取。因此可以让线程块协作把 x 的分段缓存到共享内存中,再由块内多个线程束复用这段数据。
以下示例复用上文的 warpReduceSum,并假设 blockDim.x == TILE_V == 256。
// 使用共享内存缓存向量 x(分段式)
__global__ void gemvStep3(float* A, float* x, float* y, int m, int k) {
__shared__ float x_shared[256]; // TILE_V = 256
int warp_id = threadIdx.x / 32;
int lane = threadIdx.x % 32;
int warps_per_block = blockDim.x / 32;
int row = blockIdx.x * warps_per_block + warp_id;
bool valid_row = row < m;
float partial = 0.0f;
int TILE_V = 256;
int num_tiles = (k + TILE_V - 1) / TILE_V;
for (int tile_idx = 0; tile_idx < num_tiles; tile_idx++) {
// 协作加载向量 x 的分段到共享内存
int x_idx = tile_idx * TILE_V + threadIdx.x;
if (threadIdx.x < TILE_V && x_idx < k) {
x_shared[threadIdx.x] = x[x_idx];
} else if (threadIdx.x < TILE_V) {
x_shared[threadIdx.x] = 0.0f;
}
__syncthreads();
// 保持合并访存:同一线程束内相邻线程读取 A 的连续列
for (int j = lane; j < TILE_V; j += 32) {
int col = tile_idx * TILE_V + j;
if (valid_row && col < k) {
partial += A[row * k + col] * x_shared[j];
}
}
__syncthreads();
}
partial = warpReduceSum(partial);
if (valid_row && lane == 0) {
y[row] = partial;
}
}
查看优化要点
优化要点:
TILE_V可调节:512, 1024, 2048- 在共享内存容量与占用率间平衡
- 向量 x 被线程块内多个线程束重复使用,同时保持矩阵 A 的合并访存
GEMV 优化检查清单
-
访存是否合并?
- 一个线程束协作计算一行,线程块内可包含 多个线程束
- 同一线程束内相邻线程访问矩阵 A 的连续列,避免跨步访问
-
是否使用向量化加载?
- 在合并访存基础上使用 128-bit 或更宽的向量化加载
- 确保数据 16B 对齐
-
是否缓存向量 x?
- 使用共享内存分段缓存
- TILE_V 根据容量调整
GEMM 优化(矩阵×矩阵)
GEMM 定义与特征
GEMM(通用矩阵乘法)是深度学习、科学计算和图形计算中的核心算子。在 LLM、卷积、全连接层等场景中,GEMM 通常占据大量计算时间,因此 GEMM 优化是获得端到端性能收益的关键路径。
GEMM(General Matrix Multiplication) 操作:
其中:
- 是 矩阵
- 是 矩阵
- 是 矩阵
通常把输入矩阵记为 A/B,把输出矩阵记为 C。优化 GEMM 时,需要同时关注计算密度、片上存储容量、共享内存带宽和寄存器占用。
算术强度(以 float 为例):
当 较大时,算术强度远高于 1 → 计算瓶颈负载