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的支持

template<typename Use, int m, int n, int k, typename T, typename Layout=void> class fragment;

void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm);
void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm, layout_t layout);
void store_matrix_sync(T* mptr, const fragment<...> &a, unsigned ldm, layout_t layout);
void fill_fragment(fragment<...> &a, const T& v);
void mma_sync(fragment<...> &d, const fragment<...> &a, const fragment<...> &b, const fragment<...> &c, bool satf=false);

(1)fragment

template<typename Use, int m, int n, int k, typename T, typename Layout=void> class 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

void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm);
void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm, layout_t layout);

等到所有 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

void store_matrix_sync(T* mptr, const fragment<...> &a, unsigned ldm, layout_t layout);

等待所有 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

void fill_fragment(fragment<...> &a, const T& v);

用常数值 v 填充矩阵片段。由于矩阵元素到每个片段的映射未指定,因此该函数通常由 warp 中的所有线程使用公共的 v 值来调用。

(5)mma_sync

void mma_sync(fragment<...> &d, const fragment<...> &a, const fragment<...> &b, const fragment<...> &c, bool satf=false);

等待所有 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矩阵乘法

namespace wmma = nvcuda::wmma;

enum {
GEMM_M = 4096,
GEMM_N = 4096,
GEMM_K = 4096,

WMMA_M = 16,
WMMA_N = 16,
WMMA_K = 16,

WARMUP_ITERS = 10,
BENCHMARK_ITERS = 50
};

__global__ void naive_wmma_kernel(const __half *A, const __half *B, float *C) {
int tile_row = blockIdx.y;
int tile_col = blockIdx.x;

int row = tile_row * WMMA_M;
int col = tile_col * WMMA_N;

wmma::fragment<wmma::matrix_a, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
a_frag;
wmma::fragment<wmma::matrix_b, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
b_frag;
wmma::fragment<wmma::accumulator, WMMA_M, WMMA_N, WMMA_K, float> c_frag;
wmma::fill_fragment(c_frag, 0.0f);

for (int k0 = 0; k0 < GEMM_K; k0 += WMMA_K) {
const __half *a_tile_ptr = A + row * GEMM_K + k0;
const __half *b_tile_ptr = B + k0 * GEMM_N + col;
wmma::load_matrix_sync(a_frag, a_tile_ptr, GEMM_K);
wmma::load_matrix_sync(b_frag, b_tile_ptr, GEMM_N);
wmma::mma_sync(c_frag, a_frag, b_frag, c_frag);
}
float *c_tile_ptr = C + row * GEMM_N + col;
wmma::store_matrix_sync(c_tile_ptr, c_frag, GEMM_N, wmma::mem_row_major);
}

dim3 block(32);
dim3 grid(GEMM_N / WMMA_N, GEMM_M / WMMA_M);

naive_wmma_kernel<<<grid, block>>>(d_A, d_B, d_C_wmma);

代码逐段解析

1. 参数枚举与 Tile 尺寸定义

enum {
GEMM_M = 4096, // 矩阵 A 的行数(M),也是矩阵 C 的行数
GEMM_N = 4096, // 矩阵 B 的列数(N),也是矩阵 C 的列数
GEMM_K = 4096, // 矩阵 A 的列数 = 矩阵 B 的行数(K)

WMMA_M = 16, // Tensor Core 一次 warp 级乘法处理的 m 维度 tile 大小
WMMA_N = 16, // Tensor Core 一次 warp 级乘法处理的 n 维度 tile 大小
WMMA_K = 16, // Tensor Core 一次 warp 级乘法处理的 k 维度 tile 大小

WARMUP_ITERS = 10, // GPU 预热迭代次数
BENCHMARK_ITERS = 50 // 正式性能基准迭代次数
};
  • 整体矩阵规模为 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 函数签名

__global__ void naive_wmma_kernel(const __half *A, const __half *B, float *C)
  • AB:FP16 输入矩阵(Tensor Core 要求 A、B 为低精度输入)。

  • C:FP32 输出矩阵(累加器使用更高精度以减少误差)。

3. Tile 索引计算

int tile_row = blockIdx.y;    // 当前 block 在 Y 维度的索引,对应输出矩阵的行 tile
int tile_col = blockIdx.x; // 当前 block 在 X 维度的索引,对应输出矩阵的列 tile

int row = tile_row * WMMA_M; // 当前 block 负责的输出 tile 的起始行
int col = tile_col * WMMA_N; // 当前 block 负责的输出 tile 的起始列
  • 使用 2D grid 划分:blockIdx.y 表示行方向上的 tile 索引,blockIdx.x 表示列方向上的 tile 索引。

  • 每个 block 负责一个完整的 16 × 16 输出 tile。

4. 声明 WMMA Fragment

wmma::fragment<wmma::matrix_a, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
a_frag;
wmma::fragment<wmma::matrix_b, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
b_frag;
wmma::fragment<wmma::accumulator, WMMA_M, WMMA_N, WMMA_K, float> c_frag;
wmma::fill_fragment(c_frag, 0.0f);
  • 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 维度循环:分块累加

for (int k0 = 0; k0 < GEMM_K; k0 += WMMA_K) {
const __half *a_tile_ptr = A + row * GEMM_K + k0;
const __half *b_tile_ptr = B + k0 * GEMM_N + col;

wmma::load_matrix_sync(a_frag, a_tile_ptr, GEMM_K);
wmma::load_matrix_sync(b_frag, b_tile_ptr, GEMM_N);
wmma::mma_sync(c_frag, a_frag, b_frag, c_frag);
}

这是朴素 WMMA GEMM 的核心循环逻辑:

  1. 计算 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 子块。

  2. 加载矩阵片段

    • load_matrix_sync(a_frag, a_tile_ptr, GEMM_K):从全局内存加载 A 的一个 16 × 16 tile 到 a_fragldm = GEMM_K = 4096 告诉硬件矩阵连续行之间的跨度(leading dimension)。

    • load_matrix_sync(b_frag, b_tile_ptr, GEMM_N):从全局内存加载 B 的一个 16 × 16 tile 到 b_fragldm = GEMM_N = 4096

  3. 执行矩阵乘累加

    • mma_sync(c_frag, a_frag, b_frag, c_frag):执行 C += A × B

    • 这是就地运算:输出累加器 c_frag 同时作为源累加器 C 和目标累加器 D

    • 此调用是 Warp 同步的,所有 32 个线程必须同时到达此处。

  4. 循环迭代k0 从 0 步进到 4080(GEMM_K - WMMA_K),共 256 步。每一步将结果累加到 c_frag 中。

6. 存储结果

float *c_tile_ptr = C + row * GEMM_N + col;
wmma::store_matrix_sync(c_tile_ptr, c_frag, GEMM_N, wmma::mem_row_major);
  • 将累加完成的 16 × 16 输出矩阵从 c_frag 写回全局内存。

  • store_matrix_sync 同样是 Warp 同步操作。

  • ldm = GEMM_N = 4096:矩阵 C 的 leading dimension。

  • 布局指定为 mem_row_major,即按行主序存储。

7. Grid 与 Block 配置

dim3 block(32);                                // 每个 block 1 个 Warp = 32 线程
dim3 grid(GEMM_N / WMMA_N, GEMM_M / WMMA_M); // grid = (256, 256)

naive_wmma_kernel<<<grid, block>>>(d_A, d_B, d_C_wmma);
  • 每个 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,但存在明显的性能问题:

  1. 全局内存访问未合并:每个 Warp 的 32 个线程从全局内存加载矩阵 A 和 B 的 tile 时,访问模式是非连续(stride 为 GEMM_KGEMM_N)的跨步访问,导致内存带宽利用率低下。

  2. 无共享内存缓存:每个 Warp 每次循环迭代都直接从全局内存加载 tile,没有利用共享内存进行数据复用。同一个矩阵元素可能被多个 Warp 多次读取(例如,矩阵 A 的一行需要被多个输出 tile 共享),但这里完全没有复用。

  3. 寄存器压力敏感:每个 Warp 持有多个 fragment(a_frag、b_frag、c_frag),这些都占用寄存器。32 线程 × 每个 fragment 所需的寄存器 → 可能成为 occupancy 瓶颈。

  4. 缺乏双缓冲(Double Buffering):加载和计算串行执行,没有通过 pipeline 掩盖内存延迟。

为了解决这些问题,更高效的实现通常会引入共享内存分块(Shared Memory Tiling)和向量化加载。

WMMA共享内存

namespace wmma = nvcuda::wmma;

enum {
GEMM_M = 4096,
GEMM_N = 4096,
GEMM_K = 4096,

WMMA_M = 16,
WMMA_N = 16,
WMMA_K = 16,

BLOCK_M = 64,
BLOCK_N = 64,

WARP_TILE_M = 16,
WARP_TILE_N = 32,

WARPS_M = BLOCK_M / WARP_TILE_M,
WARPS_N = BLOCK_N / WARP_TILE_N,
WARPS_PER_BLOCK = WARPS_M * WARPS_N,
WARP_SIZE = 32,
BLOCK_THREADS = WARPS_PER_BLOCK * WARP_SIZE,

C_TILES_PER_WARP = WARP_TILE_N / WMMA_N,

WARMUP_ITERS = 10,
BENCHMARK_ITERS = 50
};

__global__ void shared_wmma_kernel(const __half *A, const __half *B, float *C) {
__shared__ __half shared_A[BLOCK_M * WMMA_K];

__shared__ __half shared_B[WMMA_K * BLOCK_N];

int tid = threadIdx.x;
int warp_id = tid / WARP_SIZE;

int warp_m = warp_id / WARPS_N;
int warp_n = warp_id % WARPS_N;

int block_row = blockIdx.y * BLOCK_M;
int block_col = blockIdx.x * BLOCK_N;

int warp_row = warp_m * WARP_TILE_M;
int warp_col = warp_n * WARP_TILE_N;

wmma::fragment<wmma::matrix_a, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
a_frag;

wmma::fragment<wmma::matrix_b, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
b_frag[C_TILES_PER_WARP];

wmma::fragment<wmma::accumulator, WMMA_M, WMMA_N, WMMA_K, float>
c_frag[C_TILES_PER_WARP];

for (int n = 0; n < C_TILES_PER_WARP; ++n) {
wmma::fill_fragment(c_frag[n], 0.0f);
}

for (int k0 = 0; k0 < GEMM_K; k0 += WMMA_K) {
/*
* Cooperatively load A[BLOCK_M, WMMA_K] from global memory.
*/
for (int index = tid; index < BLOCK_M * WMMA_K; index += BLOCK_THREADS) {
int local_row = index / WMMA_K;
int local_k = index % WMMA_K;

shared_A[index] = A[(block_row + local_row) * GEMM_K + (k0 + local_k)];
}

/*
* Cooperatively load B[WMMA_K, BLOCK_N] from global memory.
*/
for (int index = tid; index < WMMA_K * BLOCK_N; index += BLOCK_THREADS) {
int local_k = index / BLOCK_N;
int local_col = index % BLOCK_N;

shared_B[index] = B[(k0 + local_k) * GEMM_N + (block_col + local_col)];
}

/*
* All global-to-shared stores must finish before any warp reads
* the current K-stage.
*/
__syncthreads();

const __half *a_tile_ptr = shared_A + warp_row * WMMA_K;

wmma::load_matrix_sync(a_frag, a_tile_ptr, WMMA_K);

for (int n = 0; n < C_TILES_PER_WARP; ++n) {
const __half *b_tile_ptr = shared_B + warp_col + n * WMMA_N;

wmma::load_matrix_sync(b_frag[n], b_tile_ptr, BLOCK_N);

wmma::mma_sync(c_frag[n], a_frag, b_frag[n], c_frag[n]);
}

/*
* No warp may overwrite shared_A/shared_B with the next K-stage
* until all warps have finished reading the current stage.
*/
__syncthreads();
}

int output_row = block_row + warp_row;
int output_col = block_col + warp_col;

for (int n = 0; n < C_TILES_PER_WARP; ++n) {
float *c_tile_ptr = C + output_row * GEMM_N + output_col + n * WMMA_N;

wmma::store_matrix_sync(c_tile_ptr, c_frag[n], GEMM_N, wmma::mem_row_major);
}
}

dim3 block(BLOCK_THREADS, 1, 1);
dim3 grid(GEMM_N / BLOCK_N, GEMM_M / BLOCK_M, 1);
shared_wmma_kernel<<<grid, block>>>(d_A, d_B, d_C_wmma);

代码逐段解析

1. 参数枚举与 Tile 尺寸定义

enum {
GEMM_M = 4096, // 矩阵 A 的行数
GEMM_N = 4096, // 矩阵 B 的列数
GEMM_K = 4096, // 矩阵 A 的列数 = 矩阵 B 的行数

WMMA_M = 16, // Tensor Core 每次 warp 级乘法处理的 m 维度
WMMA_N = 16, // Tensor Core 每次 warp 级乘法处理的 n 维度
WMMA_K = 16, // Tensor Core 每次 warp 级乘法处理的 k 维度

BLOCK_M = 64, // 每个 block 负责的输出行数
BLOCK_N = 64, // 每个 block 负责的输出列数

WARP_TILE_M = 16, // 每个 warp 在行方向上的 tile 高度
WARP_TILE_N = 32, // 每个 warp 在列方向上的 tile 宽度

WARPS_M = BLOCK_M / WARP_TILE_M, // block 行方向上 warp 数量 = 4
WARPS_N = BLOCK_N / WARP_TILE_N, // block 列方向上 warp 数量 = 2
WARPS_PER_BLOCK = WARPS_M * WARPS_N, // 每个 block 总 warp 数 = 8
WARP_SIZE = 32, // CUDA warp 线程数
BLOCK_THREADS = WARPS_PER_BLOCK * WARP_SIZE, // 每个 block 线程数 = 256

C_TILES_PER_WARP = WARP_TILE_N / WMMA_N, // 每个 warp 输出的 WMMA tile 数 = 2

WARMUP_ITERS = 10,
BENCHMARK_ITERS = 50
};

与朴素版本的关键区别:

  • 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__ __half shared_A[BLOCK_M * WMMA_K];   // 64 × 16
__shared__ __half shared_B[WMMA_K * BLOCK_N]; // 16 × 64

这是优化的核心——引入共享内存作为全局内存与寄存器之间的缓存层

  • shared_A:缓存矩阵 A 的一个 64 × 16 子块(BLOCK_M × WMMA_K)。

  • shared_B:缓存矩阵 B 的一个 16 × 64 子块(WMMA_K × BLOCK_N)。

  • 每次 K 循环迭代,block 内的所有线程协作加载这两个 tile 到共享内存,然后 8 个 warp 从共享内存中读取各自所需的子 tile 进行计算。

为什么共享内存能提升性能?

  1. 合并全局内存访问:协作加载时,连续的线程访问连续的内存地址,实现全局内存的合并访问(coalesced access),最大化内存带宽利用率。
  2. 数据复用:同一个 64×16 的 A tile 被 4 个 warp 共享(每行方向上的 warp 读取 A 的相同行);同一个 16×64 的 B tile 被 2 个 warp 共享(每列方向上的 warp 读取 B 的相同列)。
  3. 减少全局内存访问次数:全局内存访问量从每个 warp 每次 K 迭代 2 次减少到每个 block 每次 K 迭代 2 次,降低了 8 倍。

3. 线程与 Warp 索引计算

int tid = threadIdx.x;
int warp_id = tid / WARP_SIZE; // 当前线程所属的 warp(0-7)

int warp_m = warp_id / WARPS_N; // warp 在 block 中的行索引(0-3)
int warp_n = warp_id % WARPS_N; // warp 在 block 中的列索引(0-1)

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
int block_row = blockIdx.y * BLOCK_M;   // 当前 block 在全局矩阵中的起始行
int block_col = blockIdx.x * BLOCK_N; // 当前 block 在全局矩阵中的起始列

int warp_row = warp_m * WARP_TILE_M; // 当前 warp 在 block 内的起始行偏移
int warp_col = warp_n * WARP_TILE_N; // 当前 warp 在 block 内的起始列偏移
  • block_row / block_col:全局坐标,标识当前 block 负责输出矩阵的 64 × 64 区域。

  • warp_row / warp_col:局部坐标,标识当前 warp 在该 block 内的 16 × 32 子区域。

4. 声明多 Tile Fragment 数组

wmma::fragment<wmma::matrix_a, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
a_frag;

wmma::fragment<wmma::matrix_b, WMMA_M, WMMA_N, WMMA_K, __half,
wmma::row_major>
b_frag[C_TILES_PER_WARP]; // 2 个 b fragment

wmma::fragment<wmma::accumulator, WMMA_M, WMMA_N, WMMA_K, float>
c_frag[C_TILES_PER_WARP]; // 2 个累加器 fragment

一个 warp 负责 16 × 32 的输出区域,即 2 个相邻的 WMMA tile(每个 16 × 16,沿列方向排列):

输出区域(16 × 32):
┌──────────────┬──────────────┐
│ c_frag[0] │ c_frag[1] │
│ (16 × 16) │ (16 × 16) │
└──────────────┴──────────────┘
  • a_frag:只有 1 个,因为两个输出 tile 共享同一个 A tile(相同行,相同的 K 维度子块)。

  • b_frag[2]:需要 2 个不同的 B tile(同一行,不同列范围)。

  • c_frag[2]:2 个累加器,分别累加两个输出 tile 的结果。

for (int n = 0; n < C_TILES_PER_WARP; ++n) {
wmma::fill_fragment(c_frag[n], 0.0f); // 两个累加器都初始化为 0
}

5. K 维度循环:协作加载 + WMMA 计算

这是共享内存版本的核心,分为三个阶段:协作加载 → 同步 → WMMA 计算 → 同步

Stage 1:协作加载 A Tile 到共享内存
for (int index = tid; index < BLOCK_M * WMMA_K; index += BLOCK_THREADS) {
int local_row = index / WMMA_K;
int local_k = index % WMMA_K;

shared_A[index] = A[(block_row + local_row) * GEMM_K + (k0 + local_k)];
}
  • 循环模式: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 到共享内存
for (int index = tid; index < WMMA_K * BLOCK_N; index += BLOCK_THREADS) {
int local_k = index / BLOCK_N;
int local_col = index % BLOCK_N;

shared_B[index] = B[(k0 + local_k) * GEMM_N + (block_col + local_col)];
}
  • 覆盖 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:线程屏障 —— 确保数据就绪
__syncthreads();
  • 所有 256 个线程必须全部完成共享内存的写入后,任何 warp 才能开始读取。

  • 如果没有这个屏障,某个 warp 可能读到其他 warp 尚未写入完成的脏数据。

Stage 4:Warp 从共享内存加载 Fragment 并计算
const __half *a_tile_ptr = shared_A + warp_row * WMMA_K;
wmma::load_matrix_sync(a_frag, a_tile_ptr, WMMA_K);

for (int n = 0; n < C_TILES_PER_WARP; ++n) {
const __half *b_tile_ptr = shared_B + warp_col + n * WMMA_N;
wmma::load_matrix_sync(b_frag[n], b_tile_ptr, BLOCK_N);
wmma::mma_sync(c_frag[n], a_frag, b_frag[n], c_frag[n]);
}
  • A tile 指针shared_A + warp_row * WMMA_K,定位到当前 warp 所需的 16 行在共享内存中的起始位置。

    • 例如,warp_m = 2 时,warp_row = 32,读取 shared_A[32 * 16] 开始的 16 行。

    • 重要观察:虽然 load_matrix_syncldm 参数为 WMMA_K = 16(远小于朴素版本的 4096),但这只是共享内存内部的 stride,不存在全局内存未被合并的问题。

  • B tile 指针shared_B + warp_col + n * WMMA_N,定位到当前 warp 所需的列范围(0-31 或 32-63)。

    • n = 0b_tile_ptr = shared_B + warp_col(列偏移 0-31)。

    • n = 1b_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:线程屏障 —— 防止数据覆写
__syncthreads();
  • 在所有 warp 完成对当前 K 阶段共享内存的读取之前,不允许任何线程写入下一轮的数据。

  • 这个屏障保证了 shared_Ashared_B 中的数据的正确生命周期

6. 存储结果

int output_row = block_row + warp_row;
int output_col = block_col + warp_col;

for (int n = 0; n < C_TILES_PER_WARP; ++n) {
float *c_tile_ptr = C + output_row * GEMM_N + output_col + n * WMMA_N;
wmma::store_matrix_sync(c_tile_ptr, c_frag[n], GEMM_N, wmma::mem_row_major);
}
  • 将两个累加器 fragment 分别写回全局内存。

  • 输出位置:当前 warp 负责的 16 × 32 区域中的第 n16 × 16 tile。

  • ldm = GEMM_N = 4096,这是输出矩阵 C 的 leading dimension。

7. Grid 与 Block 配置

dim3 block(BLOCK_THREADS, 1, 1);                    // 256 线程 / block
dim3 grid(GEMM_N / BLOCK_N, GEMM_M / BLOCK_M, 1); // grid = (64, 64)

shared_wmma_kernel<<<grid, block>>>(d_A, d_B, d_C_wmma);
  • 每个 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 字节的缓存行,进一步提升了有效带宽。