CUDA算子优化(5):WMMA矩阵乘法
WMMA(Warp-level Matrix Multiply-Accumulate)是 NVIDIA CUDA 编程模型中提供的一组底层 API,专门用于调用 GPU 的 Tensor Core 进行Warp 级别的矩阵乘加运算。通过该接口,开发者可直接在 Warp 层面高效执行半精度(FP16)或混合精度的矩阵计算,从而显著加速深度学习模型的训练与推理。使用时需注意内存对齐与寄存器分配,通常需配合 Volta 架构及以上的 GPU 硬件。
WMMA API解释
Volta架构引入的新特性Tensor Cores是可编程乘法和加法单元,能够显著增加浮点计算吞吐量。每个Tensor Core执行D = A * B + C的运算。在Volta架构中,矩阵相乘的输入A和B是FP16的矩阵,而C和D可能是FP16或FP32的矩阵;而在图灵架构中,添加了对IN8和IN4的支持。
|
(1)fragment
|
fragment:包含分布在Warp中所有线程上的矩阵片段。
Use:可用值:matrix_a、matrix_b、accumulator。matrix_a表示当前fragment用作第一个被乘数,A;matrix_b表示当前fragment用作第二个被乘数,B;accumulator表示当前fragment用作源或目标累加器(C或者D)。
m、n、k:描述参与乘法累加运算的Warp宽度矩阵块的形状,matrix_a块的尺寸为m * k;matrix_b的尺寸为k * n;accumulator块的尺寸为m * n。
数据T表示数据类型,对于A和B数据类型T可以是double、float、__half、__nv_bfloat16、char或者unsigned char;对于C和D数据类型T可以是double、float、int或者__half。
必须为matrix_a和matrix_b指定Layout参数,row_major 或 col_major 分别指示矩阵行或列中的元素在内存中是连续的。 accumulator 矩阵的 Layout 参数应保留默认值 void 。仅当如下所述加载或存储累加器时才指定行或列布局。
(2)load_matrix_sync
|
等到所有 Warp 通道都到达 load_matrix_sync 时,才从内存加载矩阵fragment a。
mptr 必须是指向内存中矩阵第一个元素的 mptr 位对齐指针 ldm 描述连续行(对于行主布局)或列(对于列主布局)之间元素的步幅,并且对于 __half 元素类型必须是 8 的倍数,对于 float 元素类型必须是 4 的倍数。(即,两种情况下均为 16 字节的倍数)。
如果片段是 accumulator ,则必须将 layout 参数指定为 mem_row_major 或 mem_col_major 。对于 matrix_a 和 matrix_b 片段,布局是从片段的 layout 参数推断出来的。
(3)store_matrix_sync
|
等待所有 Warp 通道都到达 store_matrix_sync 后,将矩阵片段 a 存储到内存中。
mptr 必须是指向内存中矩阵首元素的 mptr 位对齐指针 ldm 描述连续行(对于行主布局)或列(对于列主布局)之间元素的步幅,对于 __half 元素类型,必须是 8 的倍数;对于 float 元素类型,必须是 4 的倍数。(即,两种情况下均为 16 字节的倍数)。
输出矩阵的布局必须指定为 mem_row_major 或 mem_col_major 。
所有 Warp 线程的 mptr 、 ldm 、 layout 和 a 的所有模板参数的值必须相同。
(4)fill_fragment
|
用常数值 v 填充矩阵片段。由于矩阵元素到每个片段的映射未指定,因此该函数通常由 warp 中的所有线程使用公共的 v 值来调用。
(5)mma_sync
|
等待所有 Warp 通道到达 mma_sync 后,执行 Warp 同步矩阵乘法累加运算 D=A*B+C 。同时支持就地运算 C=A*B+C 。Warp 中所有线程的每个矩阵片段的 satf 值和模板参数值必须相同。此外,片段 A 、 B 、 C 和 D 的模板参数 m 、 n 和 k 必须匹配。Warp 中的所有线程都必须调用此函数,否则结果未定义。
如果 satf (饱和为有限值)模式为 true ,则以下附加数值属性适用于目标累加器:
如果元素结果为 +Infinity,则相应的累加器将包含 +MAX_NORM
如果元素结果为-Infinity,则相应的累加器将包含 -MAX_NORM
如果元素结果为 NaN,则相应的累加器将包含 +0
朴素WMMA矩阵乘法
|
代码逐段解析
1. 参数枚举与 Tile 尺寸定义
|
-
整体矩阵规模为 4096 × 4096。
-
每个 Warp(32 个线程)负责计算一个 16 × 16 的输出 tile(
WMMA_M × WMMA_N)。 -
每个
mma_sync调用完成 16 × 16 × 16 的矩阵乘累加,即一次处理 16 个 k 维度元素。 -
由于
GEMM_K = 4096,内层循环需要执行4096 / 16 = 256次迭代才能完成整个 K 维度的累加。
2. Kernel 函数签名
|
-
A、B:FP16 输入矩阵(Tensor Core 要求 A、B 为低精度输入)。 -
C:FP32 输出矩阵(累加器使用更高精度以减少误差)。
3. Tile 索引计算
|
-
使用 2D grid 划分:
blockIdx.y表示行方向上的 tile 索引,blockIdx.x表示列方向上的 tile 索引。 -
每个 block 负责一个完整的 16 × 16 输出 tile。
4. 声明 WMMA Fragment
|
-
a_frag:类型为matrix_a,形状为 16 × 16(WMMA_M × WMMA_K),row_major布局,FP16 数据。 -
b_frag:类型为matrix_b,形状为 16 × 16(WMMA_K × WMMA_N),row_major布局,FP16 数据。 -
c_frag:类型为accumulator,形状为 16 × 16(WMMA_M × WMMA_N),FP32 数据(累加器不需要布局参数,保留void)。 -
fill_fragment(c_frag, 0.0f):将累加器初始化为 0。
重要理解:Fragment 并不是一个普通的 C++ 数组,而是 分布在同一个 Warp 中 32 个线程上的寄存器集合。每个线程持有 fragment 的一部分元素,开发者无法直接索引访问单个元素(映射关系未指定),只能通过
load_matrix_sync/store_matrix_sync进行批量加载和存储。
5. K 维度循环:分块累加
|
这是朴素 WMMA GEMM 的核心循环逻辑:
-
计算 tile 指针:
-
a_tile_ptr = A + row * GEMM_K + k0:定位矩阵 A 中,第row行、第k0列开始的 16×16 子块。 -
b_tile_ptr = B + k0 * GEMM_N + col:定位矩阵 B 中,第k0行、第col列开始的 16×16 子块。
-
-
加载矩阵片段:
-
load_matrix_sync(a_frag, a_tile_ptr, GEMM_K):从全局内存加载 A 的一个 16 × 16 tile 到a_frag。ldm = GEMM_K = 4096告诉硬件矩阵连续行之间的跨度(leading dimension)。 -
load_matrix_sync(b_frag, b_tile_ptr, GEMM_N):从全局内存加载 B 的一个 16 × 16 tile 到b_frag。ldm = GEMM_N = 4096。
-
-
执行矩阵乘累加:
-
mma_sync(c_frag, a_frag, b_frag, c_frag):执行C += A × B。 -
这是就地运算:输出累加器
c_frag同时作为源累加器C和目标累加器D。 -
此调用是 Warp 同步的,所有 32 个线程必须同时到达此处。
-
-
循环迭代:
k0从 0 步进到 4080(GEMM_K - WMMA_K),共 256 步。每一步将结果累加到c_frag中。
6. 存储结果
|
-
将累加完成的 16 × 16 输出矩阵从
c_frag写回全局内存。 -
store_matrix_sync同样是 Warp 同步操作。 -
ldm = GEMM_N = 4096:矩阵 C 的 leading dimension。 -
布局指定为
mem_row_major,即按行主序存储。
7. Grid 与 Block 配置
|
-
每个 block 只包含 32 个线程(一个 Warp)。这是因为 WMMA API 本身就是 Warp 级别的操作,不需要跨 Warp 协作。
-
Grid 尺寸为 256 × 256 = 65536 个 block。每个 block 输出 16 × 16 = 256 个元素,总共输出 65536 × 256 = 16,777,216 个元素,正好是 4096 × 4096 的矩阵大小。
-
理论上 GPU 可以并行调度所有 block,但实际受限于 SM 数量和寄存器资源。
朴素实现的局限性
这段代码虽然正确实现了基于 WMMA 的 GEMM,但存在明显的性能问题:
-
全局内存访问未合并:每个 Warp 的 32 个线程从全局内存加载矩阵 A 和 B 的 tile 时,访问模式是非连续(stride 为
GEMM_K或GEMM_N)的跨步访问,导致内存带宽利用率低下。 -
无共享内存缓存:每个 Warp 每次循环迭代都直接从全局内存加载 tile,没有利用共享内存进行数据复用。同一个矩阵元素可能被多个 Warp 多次读取(例如,矩阵 A 的一行需要被多个输出 tile 共享),但这里完全没有复用。
-
寄存器压力敏感:每个 Warp 持有多个 fragment(a_frag、b_frag、c_frag),这些都占用寄存器。32 线程 × 每个 fragment 所需的寄存器 → 可能成为 occupancy 瓶颈。
-
缺乏双缓冲(Double Buffering):加载和计算串行执行,没有通过 pipeline 掩盖内存延迟。
为了解决这些问题,更高效的实现通常会引入共享内存分块(Shared Memory Tiling)和向量化加载。
WMMA共享内存
|
代码逐段解析
1. 参数枚举与 Tile 尺寸定义
|
与朴素版本的关键区别:
-
BLOCK_M = 64, BLOCK_N = 64:每个 block 现在负责计算 64 × 64 的输出矩阵,而不是之前的 16 × 16。 -
WARP_TILE_M = 16, WARP_TILE_N = 32:每个 warp 负责 16 × 32 的输出区域,即在列方向上覆盖 2 个 WMMA_N tile。 -
WARPS_M = 4, WARPS_N = 2:每个 block 包含 4 × 2 = 8 个 warp,共 256 个线程。 -
C_TILES_PER_WARP = 2:每个 warp 需要管理 2 个输出 fragment(因为列方向为 32 ÷ 16 = 2)。
这种分层设计是高性能 GEMM 的典型模式:Grid → Block → Warp → WMMA Tile,每一层解决不同粒度的数据复用问题。
2. 共享内存声明
|
这是优化的核心——引入共享内存作为全局内存与寄存器之间的缓存层:
-
shared_A:缓存矩阵 A 的一个 64 × 16 子块(BLOCK_M × WMMA_K)。 -
shared_B:缓存矩阵 B 的一个 16 × 64 子块(WMMA_K × BLOCK_N)。 -
每次 K 循环迭代,block 内的所有线程协作加载这两个 tile 到共享内存,然后 8 个 warp 从共享内存中读取各自所需的子 tile 进行计算。
为什么共享内存能提升性能?
- 合并全局内存访问:协作加载时,连续的线程访问连续的内存地址,实现全局内存的合并访问(coalesced access),最大化内存带宽利用率。
- 数据复用:同一个 64×16 的 A tile 被 4 个 warp 共享(每行方向上的 warp 读取 A 的相同行);同一个 16×64 的 B tile 被 2 个 warp 共享(每列方向上的 warp 读取 B 的相同列)。
- 减少全局内存访问次数:全局内存访问量从每个 warp 每次 K 迭代 2 次减少到每个 block 每次 K 迭代 2 次,降低了 8 倍。
3. 线程与 Warp 索引计算
|
Block 内的 8 个 warp 按 4 行 × 2 列 的网格排列:
| warp_id | warp_m | warp_n |
|---|---|---|
| 0 | 0 | 0 |
| 1 | 0 | 1 |
| 2 | 1 | 0 |
| 3 | 1 | 1 |
| 4 | 2 | 0 |
| 5 | 2 | 1 |
| 6 | 3 | 0 |
| 7 | 3 | 1 |
|
-
block_row / block_col:全局坐标,标识当前 block 负责输出矩阵的 64 × 64 区域。 -
warp_row / warp_col:局部坐标,标识当前 warp 在该 block 内的 16 × 32 子区域。
4. 声明多 Tile Fragment 数组
|
一个 warp 负责 16 × 32 的输出区域,即 2 个相邻的 WMMA tile(每个 16 × 16,沿列方向排列):
|
-
a_frag:只有 1 个,因为两个输出 tile 共享同一个 A tile(相同行,相同的 K 维度子块)。 -
b_frag[2]:需要 2 个不同的 B tile(同一行,不同列范围)。 -
c_frag[2]:2 个累加器,分别累加两个输出 tile 的结果。
|
5. K 维度循环:协作加载 + WMMA 计算
这是共享内存版本的核心,分为三个阶段:协作加载 → 同步 → WMMA 计算 → 同步。
Stage 1:协作加载 A Tile 到共享内存
|
-
循环模式:256 个线程以
tid为起始索引,步长 256,覆盖64 × 16 = 1024个元素。 -
local_row = index / 16:共享内存中的行索引(0-63)。 -
local_k = index % 16:共享内存中的列索引(0-15)。 -
全局内存地址:
(block_row + local_row) * GEMM_K + (k0 + local_k)。- 相邻线程的
local_k相差 1,因此访问全局内存时相邻线程访问相邻列 → 完美合并访问。
- 相邻线程的
Stage 2:协作加载 B Tile 到共享内存
|
-
覆盖
16 × 64 = 1024个元素。 -
local_k = index / 64:K 维度索引(0-15)。 -
local_col = index % 64:列索引(0-63)。 -
全局内存地址:
(k0 + local_k) * GEMM_N + (block_col + local_col)。- 相邻线程的
local_col相差 1 → 也是合并访问。
- 相邻线程的
Stage 3:线程屏障 —— 确保数据就绪
|
-
所有 256 个线程必须全部完成共享内存的写入后,任何 warp 才能开始读取。
-
如果没有这个屏障,某个 warp 可能读到其他 warp 尚未写入完成的脏数据。
Stage 4:Warp 从共享内存加载 Fragment 并计算
|
-
A tile 指针:
shared_A + warp_row * WMMA_K,定位到当前 warp 所需的 16 行在共享内存中的起始位置。-
例如,
warp_m = 2时,warp_row = 32,读取shared_A[32 * 16]开始的 16 行。 -
重要观察:虽然
load_matrix_sync的ldm参数为WMMA_K = 16(远小于朴素版本的 4096),但这只是共享内存内部的 stride,不存在全局内存未被合并的问题。
-
-
B tile 指针:
shared_B + warp_col + n * WMMA_N,定位到当前 warp 所需的列范围(0-31 或 32-63)。-
n = 0:b_tile_ptr = shared_B + warp_col(列偏移 0-31)。 -
n = 1:b_tile_ptr = shared_B + warp_col + 16(列偏移 16-47,即 32-63)。
-
-
ldm = BLOCK_N = 64:B 矩阵在共享内存中每行有 64 个元素,leading dimension 为 64。 -
mma_sync:两个 WMMA tile 共享同一个a_frag,分别用不同的b_frag[n]计算,结果累加到对应的c_frag[n]。
Stage 5:线程屏障 —— 防止数据覆写
|
-
在所有 warp 完成对当前 K 阶段共享内存的读取之前,不允许任何线程写入下一轮的数据。
-
这个屏障保证了
shared_A和shared_B中的数据的正确生命周期。
6. 存储结果
|
-
将两个累加器 fragment 分别写回全局内存。
-
输出位置:当前 warp 负责的 16 × 32 区域中的第
n个 16 × 16 tile。 -
ldm = GEMM_N = 4096,这是输出矩阵 C 的 leading dimension。
7. Grid 与 Block 配置
|
-
每个 block = 256 线程 = 8 个 Warp,负责输出 64 × 64 的矩阵块。
-
Grid 尺寸 = 64 × 64 = 4096 个 block,覆盖整个 4096 × 4096 输出矩阵。
-
与朴素版本的 65536 个 block 相比,block 数量减少了 16 倍,每个 block 的线程数增加了 8 倍。
共享内存版本的优势
| 对比维度 | 朴素版本 | 共享内存版本 |
|---|---|---|
| 每个 block 线程数 | 32(1 warp) | 256(8 warps) |
| 每个 block 输出大小 | 16 × 16 | 64 × 64 |
| 全局内存访问 | 每次 K 迭代 2 次/warp | 每次 K 迭代 2 次/block(复用 8 倍) |
| 全局内存合并 | ❌ 跨步访问(stride = 4096) | ✅ 连续访问(coalesced) |
| 数据复用 | ❌ 无 | ✅ block 内 warp 间共享 |
| 共享内存占用 | 0 | 64×16 + 16×64 = 2048 half ≈ 4 KB |
朴素实现的 grid 包含 65536 个线程块,每块 256 个线程,每个线程分别从 A 和 B 各读取一次全局内存,所以总访问量为 65536 × 256 × 2 = 33,554,432。 通过共享内存引入,访问量降低到 4096 × 256 × 2 = 2,097,152 次,减少了 16 倍。同时,合并访问模式使得每个内存事务都能充分利用 128 字节的缓存行,进一步提升了有效带宽。