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 → 计算瓶颈负载
GEMM 分块策略
分块不会减少总计算量,但可以把局部数据保存在共享内存和寄存器中,让同一批 A/B 数据被多次复用。这样可以提高计算访存比,避免计算单元长时间等待全局内存数据。
分块的核心思想:
- 将大矩阵分解为小块
- 每块数据加载到高速存储器(共享内存)
- 在高速存储器中重复使用数据
分块参数选择:
| 参数 | 含义 | 典型值 | 限制因素 |
|---|---|---|---|
TILE_M | C 矩阵 M 维度分块 | 64-256 | 共享内存容量 |
TILE_N | C 矩阵 N 维度分块 | 64-256 | 共享内存容量 |
TILE_K | K 维度分块 | 8-32 | 共享内存容量 |
THREAD_M | 每线程计算 M 维度 | 8-16 | 寄存器数量 |
THREAD_N | 每线程计算 N 维度 | 8-16 | 寄存器数量 |
经典分块实现
// 分块 GEMM 实现:每个线程计算一个 C 元素
template<int TILE>
__global__ void gemmTiled(float* A, float* B, float* C, int M, int N, int K) {
__shared__ float As[TILE][TILE];
__shared__ float Bs[TILE][TILE];
int tx = threadIdx.x;
int ty = threadIdx.y;
int row = blockIdx.y * TILE + ty;
int col = blockIdx.x * TILE + tx;
float accum = 0.0f;
// 分块迭代
for (int t = 0; t < (K + TILE - 1) / TILE; t++) {
// 协作加载 A 块到共享内存(合并访问)
int a_col = t * TILE + tx;
if (row < M && a_col < K) {
As[ty][tx] = A[row * K + a_col];
} else {
As[ty][tx] = 0.0f;
}
// 协作加载 B 块到共享内存(合并访问)
int b_row = t * TILE + ty;
if (b_row < K && col < N) {
Bs[ty][tx] = B[b_row * N + col];
} else {
Bs[ty][tx] = 0.0f;
}
__syncthreads();
// 在共享内存中计算
#pragma unroll
for (int k = 0; k < TILE; k++) {
accum += As[ty][k] * Bs[k][tx];
}
__syncthreads();
}
// 写回结果
if (row < M && col < N) {
C[row * N + col] = accum;
}
}
// 启动配置
dim3 blockDim(16, 16); // 256 线程/Block
dim3 gridDim((N + 15) / 16, (M + 15) / 16);
gemmTiled<16><<<gridDim, blockDim>>>(A, B, C, M, N, K);
优化要点
1. 合并访存:
// A 矩阵加载:行优先,threadIdx.x 映射到连续列
As[ty][tx] = A[c_row * K + (t * TILE_K + tx)];
// B 矩阵加载:列优先,threadIdx.y 映射到连续行
Bs[ty][tx] = B[(t * TILE_K + ty) * N + c_col];
2. 寄存器优化:
// 每线程计算 8×8 子块,64 个累加器存储在寄存器
float accum[8][8];
// 优势:
// - 减少共享内存访问
// - 提高计算/访存比
// - 充分利用寄存器带宽
3. 循环展开:
// #pragma unroll 自动展开内层循环
// 增加指令级并行
// 隐藏访存延迟
张量核心(Tensor Core)优化
张量核心(Tensor Core)简介
张量核心(Tensor Core)是专为矩阵乘法加速的硬件单元:
- 单指令完成
- A、B:16 位浮点(FP16/BF16)
- C、D:16 位或 32 位浮点
- MT GPU S5000 及后续架构支持
MT GPU 张量核心特性:
- 支持 FP16、BF16、INT8 精度
- 每指令 256 FLOPs(4×4×4 矩阵乘)
- 需对齐到 16 字节边界
张量核心(Tensor Core)GEMM 实现
推荐:使用 MUTLASS 库
对于生产环境,建议使用 MUTLASS(MUSA Tensor Core Library,MUSA 张量核心库)库来实现张量核心(Tensor Core)加速的 GEMM。MUTLASS 提供了经过深度优化的张量核心(Tensor Core)实现,支持 FP16/BF16/TF32/FP8 等多种精度。
详见:MUTLASS 库文档
MUTLASS 主循环(Mainloop)示例
下面两个标签页展示 MUTLASS 两阶段主循环的核心片段。第一个片段强调基础拷贝流程,第二个片段 展示边界谓词拷贝。两者都省略了张量视图、迭代器、谓词和类型定义等上下文,不能作为独立 kernel 直接编译。
- 基础拷贝
- 谓词拷贝
该片段用于说明 MUTLASS_PRAGMA_NO_UNROLL 如何配合静态循环使用:
#include <mutlass/mutlass.h>
#include <mute/algorithm/gemm.hpp>
// MUTLASS mainloop 两阶段实现
template <class... Args>
__global__ void gemm_mainloop_kernel(Args... args) {
// ... 初始化代码省略 ...
// K 维度上的外积(outer product)数量
auto K_BLOCK_MAX = size<2>(tCrA);
// 使用 MUTLASS_PRAGMA_NO_UNROLL 禁用循环展开
// 静态 for 循环需要手动控制展开行为
MUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > -1)
{
// 使用 for_each 进行静态循环流水线化
for_each(make_int_sequence<K_BLOCK_MAX>{}, [&] (auto k_block)
{
if (k_block == K_BLOCK_MAX - 1)
{
__syncthreads();
// 将寄存器数据写回共享内存(双缓冲)
copy(tArA, tAsA);
copy(tBrB, tBsB);
__syncthreads();
}
// 预加载下一个 K 块的数据(共享内存 smem → 寄存器内存 rmem)
int k_block_next = (k_block + Int<1>{}) % K_BLOCK_MAX; // 静态取模
copy(tCsA(_,_,k_block_next), tCrA_copy_view(_,_,k_block_next));
copy(tCsB(_,_,k_block_next), tCrB_copy_view(_,_,k_block_next));
if (k_block == 0)
{
// 从全局内存(gmem)加载下一个 K 块(全局内存 gmem → 寄存器内存 rmem)
copy(gmem_tiled_copy_a, tAgA(_,_,_,*k_tile_iter), tArA);
copy(gmem_tiled_copy_b, tBgB(_,_,_,*k_tile_iter), tBrB);
++k_tile_iter;
--k_tile_count;
}
// 计算前变换数据
mute::transform(tCrA(_,_,k_block), TransformA{});
mute::transform(tCrB(_,_,k_block), TransformB{});
// 线程级寄存器矩阵乘累加
mute::gemm(tiled_mma, accum, tCrA(_,_,k_block), tCrB(_,_,k_block), src_accum);
});
}
}
该片段用于说明边界条件下如何使用 copy_if 控制全局内存加载:
// MUTLASS mainloop 两阶段实现
template <class... Args>
__global__ void gemm_mainloop_kernel(Args... args) {
// ... 初始化代码省略 ...
auto K_BLOCK_MAX = size<2>(tCrA);
// 使用 MUTLASS_PRAGMA_NO_UNROLL 禁用循环展开
MUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > -1)
{
// 使用 for_each 进行静态循环流水线化
for_each(make_int_sequence<K_BLOCK_MAX>{}, [&] (auto k_block)
{
if (k_block == K_BLOCK_MAX - 1)
{
__syncthreads();
// 将寄存器数据写回共享内存(双缓冲)
copy(tArA, tAsA);
copy(tBrB, tBsB);
__syncthreads();
}
// 预加载下一个 K 块的数据(共享内存 smem → 寄存器内存 rmem)
int k_block_next = (k_block + Int<1>{}) % K_BLOCK_MAX; // 静态取模
copy(tCsA(_,_,k_block_next), tCrA_copy_view(_,_,k_block_next));
copy(tCsB(_,_,k_block_next), tCrB_copy_view(_,_,k_block_next));
if (k_block == 0)
{
// 从全局内存(gmem)加载下一个 K 块(全局内存 gmem → 寄存器内存 rmem)
copy_if(gmem_tiled_copy_a, tApA, tAgA(_,_,_,*k_tile_iter), tArA);
copy_if(gmem_tiled_copy_b, tBpB, tBgB(_,_,_,*k_tile_iter), tBrB);
++k_tile_iter;
--k_tile_count;
}
// 计算前变换数据
mute::transform(tCrA(_,_,k_block), TransformA{});
mute::transform(tCrB(_,_,k_block), TransformB{});
// 线程级寄存器矩阵乘累加
mute::gemm(tiled_mma, accum, tCrA(_,_,k_block), tCrB(_,_,k_block), src_accum);
});
}
}
关键实现要点
| 要点 | 说明 |
|---|---|
| 双缓冲(Double Buffering) | 使用 K_BLOCK_MAX 块进行数据流水线,避免等待 |
| 静态循环 | 使用 for_each(make_int_sequence<K_BLOCK_MAX>{}) 替代普通 for 循环 |
| MUTLASS_PRAGMA_NO_UNROLL | 禁用 while 循环展开,保持静态循环的确定性 |
| 寄存器复用 | 数据在共享内存(smem)→ 寄存器内存(rmem)→ 寄存器中流动,最大化复用 |
张量核心(Tensor Core)优化要点
| 优化项 | 要求 | 理由 |
|---|---|---|
| 精度选择 | FP16/BF16 | 张量核心(Tensor Core)原生支持 |
| 对齐要求 | 16 字节边界 | 避免非对齐访问开销 |
| 分块大小 | 16 的倍数 | 匹配张量核心(Tensor Core)指令形状 |
| 累加器精度 | FP32 | 保持数值稳定性 |
性能收益趋势
GEMV 优化收益趋势
| 优化阶段 | 主要收益 | 说明 |
|---|---|---|
| 基础并行 | 建立行级并行 | 一个线程束或多个线程束协作计算输出元素 |
| + 合并访存/向量化 | 提高全局内存带宽利用率 | 连续线程访问连续地址,向量化减少访存指令数量 |
| + 共享内存缓存 | 降低向量重复读取压力 | 将向量分段缓存到共享内存,在多个线程束间复用 |
GEMM 优化收益趋势
| 优化阶段 | 主要收益 | 说明 |
|---|---|---|
| 基础实现(Naive) | 实现简单 | 计算访存比较低,难以充分利用计算单元 |
| 共享内存分块 | 提升数据复用 | 将 A/B 分块缓存到共享内存,降低全局内存访问压力 |
| + 寄存器优化 | 提升计算密度 | 每线程计算多个输出元素,累加器存放在寄存器中 |
| + 张量核心(Tensor Core) | 提升矩阵乘吞吐 | 适用于 FP16/BF16/INT8 等张量核心支持的数据类型 |
实际性能收益与矩阵形状、数据类型、硬件架构、访存对齐和 kernel 实现有关,应以目标环境实测为准。