ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

[CUDA 优化实战] hgemm - 超越 cuBLAS:Tensor-core、cp.async、ldmatrix、mma

2026/8/30 7:27:35 拓冰建站 浏览量
[CUDA 优化实战] hgemm - 超越 cuBLAS:Tensor-core、cp.async、ldmatrix、mma [CUDA 优化实战] hgemm - 超越 cuBLASTensor-core、cp.async、ldmatrix、mma0. 序 - 半精度一统江湖本文适用于有一定 CUDA 编程基础熟悉 GEMM 优化对进阶 tensor core / 嵌入 PTX 指令 性能调优感兴趣的读者阅读完整 kernel 和测试代码可以点击vitamin-cuda: hgemm 查看没错这是超越 cuBLAS 系列之三NV 还不针对移动端显卡调优的话我们就还是超越。不过这估计是 GEMM 系列最后一篇了也可能会写 fp8 的 kernel 但应该不会是纯血 fp8 gemm 了感觉有点同质化。而且 fp8 本身就要搭配各种量化姿势使用现在 bf16 甚至 fp8 几乎统治了 llm/vlm 训练和推理传统模型可能还有在用 fp16精度略高一些但动态范围小传统模型的权重和激活值分布比较集中离群值相对少用 fp16 更佳。而 LLm、VLM 现在一般原始模型是 bf16推理可能会使用 fp8/int8/int4 等各种量化姿势。所以半精度矩阵乘法才是真正的主战场。本文以 MNK4096MxKxN, cuBLAS 最擅长的中等规模的 GEMM fp16/bf16 为例在 RTX 5060 移动版显卡上使用 cp.async、ldmatrix、mma 等 PTX 指令配合 Tensor Core 加速计算在同精度赛道上成功超越了 NVIDIA 原厂的 cuBLAS。本文将复盘这场与 cuBLAS 较量的过程。介绍 cp.async 指令、ldmatrix/mma warp 级别 PTX 指令的运用硬核的 swizzle 推导layout 分析与理解等等本文会给出 4 个 kernel 实现从基础的 grid swizzle cp.async 双 ldmatrix mma到 swizzlling 读写 smem 解决 bank conflict上 double buffer 隐藏时延最后合并 gmem 读写事务是的填了上一篇文章的 c 矩阵写回 gmem 事务合并的坑 等手段完成最终对 cuBLAS 的超越。整个过程涉及对指令要求/用法的介绍和分析swizzle 设计希望能帮助读者深入理解 Tensor Core 的使用和优化技巧。kernel 大纲如下第一个是 cuBLAS kernelhgemm_cublas bf16/fp16 版hgemm_naive bf16/fp16 版 (ldmatrix mma)hgemm_bcf bf16/fp16 版 (ldmatrix mma, As/Bs swizzle bcf, 95~99% cuBLAS’ performance)hgemm_bcf_dbf bf16/fp16 版 (ldmatrix mma, As/Bs swizzle bcf, double buffer, outperforming cuBLAS)hgemm_bcf_dbf_rw bf16/fp16 版 (ldmatrix mma, As/Bs swizzle bcf, double buffer, coalesced r/w gmem, outperforming cuBLAS)1. hgemm_naive我们之前的文章已经介绍过 cp.async, ldmatrix, mma 的基本用法这里再简单提一下只说用到指令的具体用法省略输入/输出列表详细用法请参考官方文档cp.async.cg.shared.global.L2::128B[%0],[%1],16;// cg(Cache at Global level)bypass L1拷贝16字节prefetch 128B 到 L2 cp.async.commit_group;// 提交异步拷贝任务 cp.async.wait_group0;// 表示允许当前线程后台异步的 group 数0 表示不允许后台要等待到全部完成cp.async使用姿势和 tf32 的 kernel 中的一样毫无变化mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 d, a, b, c;mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 d, a, b, c;mma指令有一点变化之前计算 shape 是 m16n8k8简称 1688现在是 m16n8k16之前用 a/b 类型使用 tf32现在换成 f16/bf16ldmatrix.sync.aligned.m8n8.x4.shared.b16{%0, %1, %2, %3},[%4];ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16{%0, %1},[%2];ldmatrix一个 warp 从 smem 协同加载一个小矩阵分块到所有线程所有线程一起 hold 着的寄存器结果叫做一个 fragment输出所以线程 hold 着的值叫 fragment c注意x2x4表示读取线程数为 1632。ldmatrix 读取数据提供地址的每个线程永远是读 16 字节读取完后会分发到各个线程组成一个 fragment。理解这一点才能理解如何解决 bank conflict由于半精度计算我们使用 mma 的 m16n8k16 shape所以要 从 As 读取 16x16 的 tile和 Bs 的 16x8 的 tileldmatrix 半精度我们也能用.trans转置了狂喜~本次我们有 tf32 的经验所以再简单说下 tiling 策略就直接上代码。128x128 的 tilingc 矩阵视角k 维度跨步为 32因为半精度2 字节我们可以激进一点 load 32 个 k 了256 的 thread block size同时做了一个 2d 的 2x4 warp tiling划分为 64x128 上下两半的 c 分块一个 warp 负责 64x32还是对应这 m16n8k164x4 轮 (k 维度另算。tiling size 选择的策略有许多考量因素具体还是参考我第一篇 SGEMM 的文章谢谢代码// a block calculate c[128][128]templateconst int BM128, const int BN128, const int BK32, typename T__global__ void hgemm_naive_kernel(T *a, T *b, T *c, int m, int n, int k){// grid swizzling int linear_idblockIdx.y * gridDim.x blockIdx.x;const int SWIZZLE_W8;// 将执行块设置为8的宽度 int bx(linear_id % SWIZZLE_W)(linear_id /(SWIZZLE_W * gridDim.y))* SWIZZLE_W;int by(linear_id / SWIZZLE_W)% gridDim.y;int tidthreadIdx.x;//0~255 int warp_idtid / WARP_SIZE;int lane_idtid % WARP_SIZE;// 搬运映射 int load_a_rowtid /4;//0~63 int load_a_col(tid %4)*8;//0,8,16,24 int load_b_rowtid /16;//0~15(K 维度 int load_b_col(tid %16)*8;//0,8,16...120(N 维度 // A/B 都行优先 __shared__ T As[BM][BK];__shared__ T Bs[BK][BN];// warp tiling // 每个 warp 负责64x32的 C 矩阵块 int warp_id_mwarp_id /4;//0,1int warp_id_nwarp_id %4;//0,1,2,3// 寄存器总量M 维4块 * N 维4块 * 每块4个寄存器64float sum[4][4][4]{0.f};// 主循环for(int bk0;bkk;bkBK){//1. cp.async load A uint32_t smem_a0static_castuint32_t(__cvta_generic_to_shared(As[load_a_row][load_a_col]));uint32_t smem_a1static_castuint32_t(__cvta_generic_to_shared(As[load_a_row 64][load_a_col]));T *global_a0a[(by * BM load_a_row)* k bk load_a_col];T *global_a1a[(by * BM load_a_row 64)* k bk load_a_col];CP_ASYNC_CG(smem_a0, global_a0);CP_ASYNC_CG(smem_a1, global_a1);//2. cp.async load B uint32_t smem_b0static_castuint32_t(__cvta_generic_to_shared(Bs[load_b_row][load_b_col]));uint32_t smem_b1static_castuint32_t(__cvta_generic_to_shared(Bs[load_b_row 16][load_b_col]));T *global_b0b[(bk load_b_row)* n bx * BN load_b_col];T *global_b1b[(bk load_b_row 16)* n bx * BN load_b_col];CP_ASYNC_CG(smem_b0, global_b0);CP_ASYNC_CG(smem_b1, global_b1);CP_ASYNC_COMMIT_GROUP();CP_ASYNC_WAIT_GROUP_0();__syncthreads();//3. Tensor Core 计算阶段(K 分2步一次16个 k)#pragma unrollfor(int k_step0;k_step2;k_step){int k_offsetk_step *16;uint32_t reg_a[4][4];uint32_t reg_b[4][2];//4次 ldmatrix A(4*1664行#pragma unrollfor(int m_idx0;m_idx4;m_idx){// ldmatrix x4 读 16x16 int a_rowwarp_id_m *64 m_idx *16(lane_id %16);int a_colk_offset (lane_id /16)*8;uint32_t smem_addrstatic_castuint32_t(__cvta_generic_to_shared(As[a_row][a_col]));LDMATRIX_X4(reg_a[m_idx][0], reg_a[m_idx][1], reg_a[m_idx][2], reg_a[m_idx][3], smem_addr);}//4次 ldmatrix B(4*832列#pragma unrollfor(int n_idx0;n_idx4;n_idx){// Lane0~15 的线程恰好覆盖了16行 两块 8x8 的首地址 int b_rowk_offset (lane_id %16);int b_colwarp_id_n *32 n_idx *8;uint32_t smem_addrstatic_castuint32_t(__cvta_generic_to_shared(Bs[b_row][b_col]));LDMATRIX_X2_TRANS(reg_b[n_idx][0], reg_b[n_idx][1], smem_addr);}// MMA 核心运算4x4 次 m16n8k16#pragma unrollfor(int m_idx0;m_idx4;m_idx){#pragma unrollfor(int n_idx0;n_idx4;n_idx){ifconstexpr(std::is_same_vT, __half){M16N8K16_F16(sum[m_idx][n_idx][0], sum[m_idx][n_idx][1], sum[m_idx][n_idx][2], sum[m_idx][n_idx][3], reg_a[m_idx][0], reg_a[m_idx][1], reg_a[m_idx][2], reg_a[m_idx][3], reg_b[n_idx][0], reg_b[n_idx][1]);}else{M16N8K16_BF16(sum[m_idx][n_idx][0], sum[m_idx][n_idx][1], sum[m_idx][n_idx][2], sum[m_idx][n_idx][3], reg_a[m_idx][0], reg_a[m_idx][1], reg_a[m_idx][2], reg_a[m_idx][3], reg_b[n_idx][0], reg_b[n_idx][1]);}}}}__syncthreads();}// ---------------- 写回 C 矩阵 ---------------- int t_rowlane_id /4;//0~7 int t_col(lane_id %4)*2;//0,2,4,6#pragma unrollfor(int m_idx0;m_idx4;m_idx){#pragma unrollfor(int n_idx0;n_idx4;n_idx){int c_base_rowby * BM warp_id_m *64 m_idx *16;int c_base_colbx * BN warp_id_n *32 n_idx *8;int idx_0(c_base_row t_row)* n c_base_col t_col;int idx_2(c_base_row t_row 8)* n c_base_col t_col;ifconstexpr(std::is_same_vT, __half){HALF2(c[idx_0])__float22half2_rn(FLOAT2(sum[m_idx][n_idx][0]));HALF2(c[idx_2])__float22half2_rn(FLOAT2(sum[m_idx][n_idx][2]));}else{BFLOAT2(c[idx_0])__float22bfloat162_rn(FLOAT2(sum[m_idx][n_idx][0]));BFLOAT2(c[idx_2])__float22bfloat162_rn(FLOAT2(sum[m_idx][n_idx][2]));}}}}2. hgemm_bcf很显然 naive 版的代码As/Bs 的读取没有经过 swizzle肯定是有 bank 冲突的。现在我们来优化共享内存访问。之前 tf32 的文章中我们已经有了两个 swizzle 宏函数#define SWIZZLE_A(row, col) ((col) ^ (((row 1) 0x3) 2))#define SWIZZLE_B(row, col) ((col) ^ (((row) 0x7) 3))那么这俩 swizzle 函数能不能直接拿到半精度的 kernel 里用显然不能因为我们现在处理的是半精度数据As 和 Bs 的 layout 是 128x32 和 32x128那么 As 的 col 会出现 0~31Bs 的 row 也是。但把这作为起点从 tf32 的 swizzle 延伸到半精度的推导却非常简单。因为仔细考虑一下半精度占两个字节而原来 tf32 是 4 个字节。那么就会导致原来 16 字节对齐的 col 下标现在会翻倍。 也就是说原来我们只需要保证 col 的低 bit01 不被扰动即可现在需要保证 bit02 不被扰动那么显然 SWIZZLE_B 本来就满足要求所以可以直接使用但是 SWIZZLE_A 就需要修改了改动很简单乘个 2 即可。由此我们得到 As/Bs 读写无冲突的两个 swizzle 函数#define SWIZZLE_A(row, col) ((col) ^ (((row 1) 0x3) 3))#define SWIZZLE_B(row, col) ((col) ^ (((row) 0x7) 3))代码核心修改// cp.async load A uint32_t smem_a0static_castuint32_t(__cvta_generic_to_shared(As[load_a_row][SWIZZLE_A(load_a_row, load_a_col)]));uint32_t smem_a1static_castuint32_t(__cvta_generic_to_shared(As[load_a_row 64][SWIZZLE_A(load_a_row 64, load_a_col)]));... // cp.async load B uint32_t smem_b0static_castuint32_t(__cvta_generic_to_shared(Bs[load_b_row][SWIZZLE_B(load_b_row, load_b_col)]));uint32_t smem_b1static_castuint32_t(__cvta_generic_to_shared(Bs[load_b_row 16][SWIZZLE_B(load_b_row 16, load_b_col)]));// 读取... uint32_t smem_addrstatic_castuint32_t(__cvta_generic_to_shared(As[a_row][SWIZZLE_A(a_row, a_col)]));... uint32_t smem_addrstatic_castuint32_t(__cvta_generic_to_shared(Bs[b_row][SWIZZLE_B(b_row, b_col)]));3. hgemm_bcf_dbf上一版代码在保障 16 字节对齐的情况下解决了 bank conflict。那么直接加上双 buffer 流水线具体流程参考我第一篇文章 代码也不贴了下一小节填坑重点说一下。4. hgemm_bcf_dbf_rw我们之前 tf32 的 kernel 中就提了还有 c 矩阵写回 global memory 没有做到事务合并。但其实是可以做的只是当时偷懒没写。现在正好前置步骤都已经理顺了上面的过程也不费脑子我们干脆把这一步也补上。核心思想很简单利用 As/Bs 的共享内存作为中转 Buffer让所有线程把 fragment c 的结果先写入 smem 排整齐最后再用 float4 向量化指令读取并一把梭写回 gmem。这里有一个极其巧妙的巧合As 和 Bs 占用的共享内存大小是 2(12832 32128)2B 1281282这正正好能放得下一个 Block 计算出的 128x128 的 C 矩阵块如果是存不下还要分批中转那我可能就真不想写了笑。在动手之前我们要重点理解一下 fragment C 的寄存器状态。mma.sync.m16n8k16 指令计算后输出的是一个 16x8 的小矩阵我们需要搞清楚每个线程究竟 hold 住了这个矩阵里的哪些值。NV 官方文档有个图mma m16n8k16 的计算 shape 下fragment c 为为了更直观我也画了一张图总结下来规律很清晰在一个 warp 内每 4 个线程负责同一行的 8 个元素fp16/bf16跨过 8 行之后这个排布再重复一次。 具体到线程级别每个线程手里攥着 4 个值。比如T0 拿着 c[0][0], c[0][1] 和跨 8 行的 c[8][0], c[8][1]T1 拿着 c[0][2], c[0][3], c[8][2], c[8][3]… 以此类推当然别忘了我们外层还有一个 4(m)x4(n) 的循环这意味着在 C 矩阵的全局坐标系下每次迭代其实是在跨越 16 行或 8 列。好理理解了寄存器分布现在的任务就是把 T0~T31 手里的数据同行同列地拼凑起来写进我们那个 128x128 的 smem 里。但是前方 bank conflict 预警一个 warp 内每 4 个线程写一行总共覆盖 8 行。而在我们这个 128 宽度的 smem 里显然不同行的同一列在物理地址上是绝对对齐的。如果直接写这 8 行的线程会瞬间撞在同一个 Bank 上引发极其惨烈的 8-way bank conflict不过要错开 Bank 也很简单我们直接把 B 矩阵的 Swizzle 宏拿过来“复用”即可。 为什么以 warp 0m/n offset 都等于 0 为例row:0,0,0,0,1,1,1,1,2,2,2,2,3,3,3,3,4,4,4,4,5,5,5,5,6,6,6,6,7,7,7,7col:0,2,4,6,0,2,4,6,0,2,4,6,0,2,4,6,0,2,4,6,0,2,4,6,0,2,4,6,0,2,4,6仔细观察二进制row 是 07即 00xxx有效变量位是 bit02然后我们还要保障 16 字节对齐所以直接取了低 3bits左移三位正好和 col 的低 3bit 错开异或后即可遍历所有 32 个 bank。哎等等有人说这里左移三位那不是越到了 bit5超过 32 了感觉不对呀。同志请注意由于我们用的是 fp16/bf16 元素的下标半精度只占两字节所以 bank id 的计算公式是 (row*128 col)*2 / 4 % 32 (col/2) % 32这里本来就要看 6 个 bits 的或者说 bit5 还是会被右移回 bit4嘿嘿~其他 warpoffset 情况同理即可。由此我们就用极其优雅的位运算把 c 矩阵零冲突地写进了共享内存里了。核心修改如下#define SWIZZLE_C(row, col) ((col) ^ (((row) 0x7) 3))... // A/B 都行优先用 union 复用同一块内存写法优雅 __shared__ __align__(128)union{// 前半段计算用的 A 和 B struct{T As[2][BM][BK];T Bs[2][BK][BN];};// 后半段写回用的 C T Cs[BM][BN];}smem;... // 复用 As/Bs 中转 __syncthreads();int t_rowlane_id /4;//0~7 int t_col(lane_id %4)*2;//0,2,4,6每四个线程负责一行的8列 共16字节 // register to Cs smem#pragma unrollfor(int m_idx0;m_idx4;m_idx){#pragma unrollfor(int n_idx0;n_idx4;n_idx){int c_base_rowwarp_id_m *64 m_idx *16;// m 跨16行 int c_base_colwarp_id_n *32 n_idx *8;// n 跨8列 //16行我们分成两次8行写入 int c_row_0c_base_row t_row;int c_row_2c_base_row t_row 8;int c_colc_base_col t_col;ifconstexpr(std::is_same_vT, __half){HALF2(smem.Cs[c_row_0][SWIZZLE_C(c_row_0, c_col)])__float22half2_rn(FLOAT2(sum[m_idx][n_idx][0]));HALF2(smem.Cs[c_row_2][SWIZZLE_C(c_row_2, c_col)])__float22half2_rn(FLOAT2(sum[m_idx][n_idx][2]));}else{BFLOAT2(smem.Cs[c_row_0][SWIZZLE_C(c_row_0, c_col)])__float22bfloat162_rn(FLOAT2(sum[m_idx][n_idx][0]));BFLOAT2(smem.Cs[c_row_2][SWIZZLE_C(c_row_2, c_col)])__float22bfloat162_rn(FLOAT2(sum[m_idx][n_idx][2]));}}}__syncthreads();// smem to gmem // 每个线程负责搬运64个元素(fp16/bf16)即8个 float4256 个线程一次写256*4*44096字节 T *c_blockc[by * BM * n bx * BN];#pragma unrollfor(int step0;step8;step){// 保证同一个 warp 的32个线程此时读取的 elem_idx 是绝对连续的 int elem_idx(step *256 tid)*8;int rowelem_idx /128;int colelem_idx %128;int s_colSWIZZLE_C(row, col);FLOAT4(c_block[row * n col])FLOAT4(smem.Cs[row][s_col]);}备注这里之所以 FLOAT4 强转写回得益于我们前面推导的SWIZZLE_C保留了 16 字节对齐特性。5.benchmark、ncu report 和分析话不多说直接上 benchmark 结果和 ncu reportn:4096, m:4096, k:4096torch mean time:4.097551ms,33.54tflops hgemm_cublas mean time:4.210246ms, speedup:0.97, tflops:32.64hgemm_naive mean time:5.191345ms, speedup:0.79, tflops:26.47hgemm_bcf mean time:4.336920ms, speedup:0.94, tflops:31.69hgemm_bcf_dbf mean time:4.096174ms, speedup:1.00, tflops:33.55hgemm_bcf_dbf_rw mean time:4.075860ms, speedup:1.01, tflops:33.72一些讨论老规矩看一下 cuBLAS 的 kernelvoid cutlass::Kernel2cutlass_80_tensorop_bf16_s16816gemm_relu_bf16_256x64_32x4_nn_align8(T1::Params)256x64 的 tilingm256n6432 的 k 维度切分4 级流水线。我的 ncu summary 还提示Uncoalesced Shared Accesses, 具体原因是 cp.async swizzle 之后源/目标地址的连续性/单调性不满足其要求导致 cp.async 需要 replay wavefronts 重复写入我最初其实是很想把它优化掉的期望写一个完美的 kernel。刚开始不理解为什么会 Uncoalesced于是我先尝试改造swizzle映射无果然后又尝试将 swizzle 改到读取 global memory平铺写入 smem再 swizzle 读出。可依然解决不了这个问题。经查阅资料和拷打 Gemini 得知一个 warp 的 32 个线程执行 cp.async 搬运512B数据时理想情况下LSU 会将其打包为 4 次完美的 128B 内存事务每次 8 个线程恰好填满一行 Smem。但由于我做了 swizzle原本这 8 个线程连续的写入地址被打散。这些散乱的地址既不连续也不单调递增crossbar 无法在一个时钟周期内把数据路由到对应的 Bank 里只能将事务拆解从而触发了多次 wavefronts 写入。备注这块底层的微架构行为我并不确定是否描述正确且严谨欢迎懂的大佬在评论区指正补充谢谢为什么 cuBLAS 没有这个问题我不知道 cuBLAS 具体实现是什么样的但它肯定有它的代价点开它的 shared memory 统计表可以发现高达 50 多万次 bank conflict而我0 冲突实际约有 0.14%冲突这是不同 warp 间冲突导致可忽略这其实是极致性能调优中的 Trade-off我选择了绝对无冲突的读牺牲了 cp.async 写的合并性而 cuBLAS 保底了异步拷贝写入的极速容忍了计算读取时产生的部分冲突。cuBLAS 和我选择了不同的妥协方向孰优孰劣欢迎朋友们评价一下6. 结束总结一下我们使用了技巧列表首先是cp.async、ldmatrix、mma 等 PTX 指令配合 Tensor Core 加速计算grid swizzle 拉满L2swizzle 解决 bank conflictdouble buffer 异步流水线隐藏时延利用As/Bs中转 完成c 矩阵写回事务合并。成功超越 cuBLAS。本文应该是纯血 gemm 系列最后一篇了从 fp32 到 tf32再到半精度从 naive 到超越 cuBLAS整个过程相信大家对 GPU 架构、CUDA 编程、性能优化有了更深刻的理解。如有错误请大家指正。完整 kernel 和测试代码可以点击vitamin-cuda: hgemm 查看本文同步发布于个人博客欢迎关注https://www.baizeway.com/article/5a219c62549f9573以上。