测试平台:RTX 5080
前言:v2.0 已将 CUDA Core 路线的优化基本做完,下一步引入 Tensor Core。但 mma.sync 这类 warp 级集体指令的输入输出布局由硬件严格规定,布局错误既不会触发编译报错,也不会引发运行时崩溃,只会表现为数值悄然偏离。因此在动手修改 v2.1 之前,先用一个独立的最小实验将布局规则一次性验证清楚。本篇记录三件事:为什么选择裸写 PTX 指令、mma 与 fragment 的底层机制、实验的设计方法。
1 为什么先做实验,而不是直接改 kernel
v2.0 复盘的结论很明确:三条证据链(指令数、wavefront、时长)都指向同一个瓶颈——每 warp 每轮 4468 条 FFMA 挤在发射端口上,标量数据通路已经到头。剩下的唯一路径是把 S=QK^T 和 O=PV 两处 GEMM 迁移到 Tensor Core。
而 mma.sync 有一个很恶劣的性质:布局错误是静默的。
- 地址算错只是访问错位,并非非法访问,编译照样通过
- 运行时不崩溃,读回来的都是合法的 half 数值
- flash 里还套着 softmax,若数值错得不夸张,diff 只是变大一些,无法定位错误源头
在四百行 kernel 中混入布局错误,排查成本极高。正确的做法是把布局相关的风险隔离到一个小实验里单独验证:单 warp、手工构造数据、期望值可人工验算、出错一格即可直接定位到「哪个 lane 的哪个寄存器」。这五十行是保险费,不能省。
2 为什么不用 wmma
第一反应是用 wmma——NVIDIA 官方的 C++ 封装,wmma::fragment + load_matrix_sync + mma_sync 三件套即可完成调用。但要用在 FlashAttention 上,有三个坎过不去:
第一,fragment 是黑盒。 wmma 规范明确说明 fragment 内部布局是 implementation-defined,调用者无法得知 fragment<a, 16, 16, 16> 中哪个线程持有哪个元素。而 FA 最关键的一步是 S→P 的转换:S 算完是 C fragment → 逐元素做 exp → 转手喂给下一个 mma 当 A operand。如果布局不透明,这条路只能走 shared memory 中转:S 写回 shared、同步、按 A 布局重新读取——一来一回多出 2×Br×Bc 的 shared 流量外加两次同步。裸 mma 的布局在 PTX ISA 里是明确定义的,可以推导证明 S 的 C fragment 换一个视角就是 P 的 A fragment,数据在寄存器中原地可用,零搬运。
第二,粒度不合适。 wmma 的 fp16 只有 16×16×16 一种形状,而 FA 里 S 的列数是 Bc、O 的列数是 D,用 16 宽的块去凑 8 对齐的边界(causal mask 逐元素判断、softmax 的列归约)处处别扭。mma.m16n8k16 的 n=8 与这些边界对得更齐。
第三,官方实现就是裸写的。 翻 FA2 源码,MMA_Atom<SM80_16x8x16_F32F16F16F32_TN> 底层就是这条 mma.sync,SM75_U32x4_LDSM_N 就是 ldmatrix.x4。想读懂官方实现、想借鉴它的 softmax 细节,必须先掌握这套裸指令的语言。
3 mma.sync 是一条什么指令
完整指令名如下:
mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32
逐段拆解:
- m16n8k16:D=A×B+C,其中 D 是 16×8,A 是 16×16,B 是 16×8(k=16 是收缩维)
- row.col:A 按行主序理解、B 按列主序理解(数学上的约定,实际传入的是 fragment)
- f32.f16.f16.f32:D 是 f32,A 是 f16,B 是 f16,C 是 f32——重点是累加器为 fp32。flash 的 O 要跨 Tc 个块累加,half 累加精度会严重劣化,所以 fp32 累加是底线
- .sync.aligned:warp 集体指令,32 个线程必须全部到齐一起执行
它的本质是一条指令让 32 个线程协作完成一个 16×8×16 的张量乘。对比标量路线:同样 2048 FLOP,CUDA Core 要 32 个线程各发 32 条 FFMA;mma 是每人发 1 条,乘加全部进 Tensor Core。
关键在于「协作」如何达成——这就要讲数据交接协议。一条 mma 指令的每个线程需要备好:
| operand | 每线程提供 | 总量对账 |
|---|---|---|
| A(16×16 half) | 4 个 b32(8 个 half) | 32 人 × 8 half = 256 = 16×16 ✓ |
| B(16×8 half) | 2 个 b32(4 个 half) | 32 人 × 4 half = 128 = 16×8 ✓ |
| C/D(16×8 f32) | 4 个 fp32 | 32 人 × 4 = 128 = 16×8 ✓ |
也就是说,矩阵并非整块存放在某个线程手中,而是以碎片形式分布于 32 个线程。mma 执行时硬件从全员处收齐碎片、拼装成完整矩阵、完成计算、再按固定规则将结果拆分发回。这个「拆分规则」就是 fragment 布局。
4 fragment:数据的分布规则
4.1 C/D 的归属(最直观,先记这个)
设 g = lane/4(行组 0-7),t = lane%4(列组 0-3),则该线程持有 D 的四个元素:
c0 = (g, 2t) c1 = (g, 2t+1)
c2 = (g+8, 2t) c3 = (g+8, 2t+1)
推导不难:D 每行 8 列,4 个 t 每人占 2 个连续列正好盖满;8 个 g 盖上半 8 行,加 8 偏移盖下半。32 人 × 4 元素 = 128 元素,严丝合缝。
这条规则是 v2.1 里一半操作的地基:
- causal mask 逐元素判断:本线程的 c0 位于哪一行哪一列,直接反算
(q_start + warp*32 + mi*16 + g, kv_start + ni*8 + 2t) - softmax 的行归约:第 g 行的 16 个元素分布于 lane 4g..4g+3 四个线程,因此 max/sum 需要跨这 4 个线程做 shuffle 归约(xor 2 → xor 1 两步)
4.2 A 的归属
与 C 同款逻辑,只是 k 维长度为 16,沿 k 方向重复一段:
a0 = {A[g][2t], A[g][2t+1]}
a1 = {A[g+8][2t], A[g+8][2t+1]}
a2 = {A[g][2t+8], A[g][2t+9]}
a3 = {A[g+8][2t+8], A[g+8][2t+9]}
注意 a0 中两个 half 是同一行相邻两列——沿 n 方向打包成 half2。
4.3 B 的归属(最容易记错的地方)
b0 = {B[2t][g], B[2t+1][g]}
b1 = {B[2t+8][g], B[2t+9][g]}
A 的 half2 沿 n 打包,B 的 half2 沿 k 打包。B fragment 的视角里矩阵按 [k][n] 排列——4 个 t 沿 k 盖 8 行,8 个 g 沿 n 盖 8 列,本质是列主序。
这个不对称有其道理:A 和 B 的收缩维都是 k,TC 内部成对消费相邻元素,所以 fragment 把相邻性都安排在收缩方向上。但记忆时特别容易把 B 也当成沿行打包——错一处的症状就是乘出来的数如同「做了转置」一样整体错位。
5 ldmatrix:硬件替你完成拆分
fragment 如此细碎,用普通 LDS 拼凑会非常痛苦:每个线程 4 次独立寻址、每次只有 4B(fragment 中相邻的两个 half 不保证地址连续),又慢又容易错。
ldmatrix 是官方提供的「拆分专用」指令:
ldmatrix.sync.aligned.m8n8.x4.shared.b16
协议是反向的:不是每个线程去取自己需要的数据,而是 32 个线程每人上报一个 16B 地址(指向 shared 里某行的 8 个 half),硬件把这些行收上来,按 fragment 布局自动分发到 32 人的寄存器里。调用者只需要保证「上报的行集合」是正确的。
x4 表示一次搬 4 个 8×8 tile,每人回来 4 个 b32。地址约定(关键中的关键):
lane 0-7 → tile0 的第 0-7 行
lane 8-15 → tile1 的第 0-7 行
lane 16-23 → tile2 的第 0-7 行
lane 24-31 → tile3 的第 0-7 行
寄存器 r0-r3 分别对应 4 个 tile 里「分给该 lane 的那 2 个 half」。因此地址计算的算式是纯几何问题:4 个 tile 如何铺在矩阵上、该 lane 负责哪个 tile 的哪一行。
.trans 后缀:同一批地址,分发时做转置——线程 T 拿到的不再是「T/4 行的 2(T%4)、2(T%4)+1 列」,而是「2(T%4) 行、2(T%4)+1 行的 T/4 列」。正好把「沿 n 打包」换成「沿 k 打包」——B fragment 需要的形状。V 矩阵就用这个方式读取。
至此,mma + ldmatrix + fragment 三件套的机制说明完毕。剩下的问题是:这套理解是否正确?纸面推导一百遍,不如运行一次实验。
6 实验设计
6.1 数据设计:期望值无需计算
实验里最花心思的是数据。设计如下:
A[i][j] = i(行号) B[k][j] = j(列号)
→ D[i][j] = Σ_k i·j = 16·i·j
任何格子的期望值可以直接心算:D[3][5] = 240,D[15][7] = 1680。不需要对拍不需要 golden reference,肉眼可验。
更妙的是出错可定位:D 的元素归属是 (g, 2t),如果 D[3][5] 错了,那一定是 lane 13(g=3, t=1)的 c1 寄存器出了问题——错一格就直接锁定到「哪个 lane 的哪个寄存器」,排查成本从「大海捞针」降到「看一眼」。
6.2 代码走读
实验 kernel 只有 32 个线程(单 warp),流程是:手工构造数据进 shared → ldmatrix 加载 A 和 B → 一条 mma → 逐格校验。
加载 A(直读)。mma 的 A fragment 需要 4 个 8×8 tile:
a0 = A[0:8, 0:8] a1 = A[8:16, 0:8]
a2 = A[0:8, 8:16] a3 = A[8:16, 8:16]
对照 ldmatrix.x4 的地址约定(lane 0-7 给 tile0 的行……)翻译成地址算式:
int tl = l / 8, r = l % 8;
int arow = (tl & 1) ? r + 8 : r; // tile1/3 → 下半 8 行
int acol = (tl & 2) ? 8 : 0; // tile2/3 → 右半 8 列
uint32_t aaddr = smem_u32addr(&sA[arow * 16 + acol]);
tile 编号的二进制位直接当作行/列的偏移开关:tl&1 是下半,tl&2 是右半。这个「tile 号即坐标」的技巧后面在 v2.1 里反复使用。
加载 B(.trans 读)。B 是 16×8,只需要 2 个 tile(上下各一个 8×8)。x4 强制要求 4 个地址,解决办法是给重复地址凑数,寄存器只取 b0/b1:
int btl = l / 8, br = l % 8;
int brow = (btl & 1) ? br + 8 : br; // tile1/3 重复也无害
uint32_t baddr = smem_u32addr(&sB[brow * 8]);
这里能凑数的前提是 sB 只有 8 列,tile0 和 tile2 给同样地址,分发出来 b2/b3 与 b0/b1 内容相同,不使用即可。
执行 + 校验:
uint32_t a0, a1, a2, a3, b0, b1, b2, b3;
ldmatrix_x4(a0, a1, a2, a3, aaddr); // A:normal 直读
ldmatrix_x4_trans(b0, b1, b2, b3, baddr); // B:转置分发
float d0 = 0.f, d1 = 0.f, d2 = 0.f, d3 = 0.f;
mma_m16n8k16(d0, d1, d2, d3, a0, a1, a2, a3, b0, b1);
// 本线程 4 个格子:c0=(g,2t) c1=(g,2t+1) c2=(g+8,2t) c3=(g+8,2t+1)
int row0 = g, col0 = 2 * t;
float e0 = 16.f * row0 * col0; // 期望值直接心算
...
if (fabsf(d0 - e0) > 0.5f || ...) {
printf("[FAIL] lane %2d 持有 ...: got(...) expect(...)\n", ...);
}
校验分两层:每线程先自检自己的 4 格(出错即打印 lane 号和坐标),lane 0 再全量复核 128 格。容差 0.5 是给 half 输入留的舍入余量(期望值全是 16 的倍数,实际误差远小于 1)。
6.3 编译的坑
实验代码注释是中文的,第一次编译直接雪崩——错误行号和文件内容对不上。排查后发现是无 BOM 的 UTF-8 被 EDG 前端按 936 代码页误读,中文注释的某个字节序列被当成了转义符,解析从这里开始全线跑偏。解法很朴素但有效:所有带中文注释的源文件统一加 UTF-8 BOM。后来 bench 接 v2.1 时又踩了一次(新旧文件 BOM 混用,连老文件都跟着崩),最后全项目统一带 BOM 才恢复稳定。
7 结果:一次通过
[PASS] D[i][j] = 16*i*j 全部 128 格吻合
这次 PASS 验证了三件事:
- C fragment 坐标公式(c0=(g,2t)...)——v2.1 里 causal mask 的坐标反算、softmax 的 quad shuffle 归约全部建立在它上面
- ldmatrix normal 的分发协议——Q(A operand)直读的方式正确。v2.1 里 K 也是 normal 直读(Ks 行主序 = K^T 的转置视角,分发出来的相邻 half 恰好就是 B fragment 需要的沿 k 相邻),这个推理同样建立在 normal 协议之上,v2.1 跑通反过来验证了它
- ldmatrix .trans 的分发协议——V 的加载方式原样复用
更重要的是,这个实验把 v2.1 最大的风险(布局链路)从「四百行 kernel 里排雷」变成「五十行实验里定点清除」。后面写 v2.1 时,布局相关的代码基本是照着实验的结论直接编写,一次跑通,正确性 diff 在 half 精度预期之内。
回头看,这个实验五十行,写起来不到一小时,但它省下的排查时间是以天计的。裸写 PTX 指令的正确打开方式就是:先构造一个可人工验算的玩具,把协议钉死,再上真实 kernel。
下一篇:v2.1 正式实现——Q/K/V 三个加载姿势、S→P 的寄存器直传布局咬合、fragment 上的 online softmax。
