Mojo MAX 内核排查实录:gfx950 上 TileTensor 化 MHA 的精度回归定位与根因分析
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
导读
本文完整还原了 MAX 内核仓库中一次典型的 GPU 内核回归排查过程:在 AMD Instinct MI355X(gfx950)上,两个 MHA(Multi-Head Attention)PR 落地后,Llama-3.1-405B 的 smoke test 精度从 ~1.0 骤降到 0.0(输出变成乱码)。调查经历了 commit 二分、测试覆盖缺口分析、哈希级可复现性验证,最终定位到mha_decoding内核中 TileTensor 与 LayoutTensor 在 SMEM 读写路径上的布局不匹配这一根因。通过本文,读者可以掌握一套完整的 GPU 内核精度回归排查方法论,理解 TileTensor/LayoutTensor 两种张量抽象的寄存器-共享内存映射差异,以及"连续内存 vs Paged KV cache"测试范式之间的盲区。
1. 问题背景:405B smoke test 精度从 1.0 掉到 0.0
1.1 测试环境
调查发生在 2026 年 4 月 9-10 日,硬件与模型配置如下:
| 项目 | 配置 |
|---|---|
| 硬件 | 8 × AMD Instinct MI355X OAM(gfx950),TP=8 |
| 模型 | Llama-3.1-405B(RedHatAI/Meta-Llama-3.1-405B-Instruct-FP8-dynamic) |
| 调查对象 PR | PR #82119(6936b243e61):Structured prefill kernel |
PR #82541(e25b574fad6):AMD CDNA attention 内部逻辑的 TileTensor 转换 |
1.2 问题现象
两个 MHA 相关 PR 落地后,405B smoke test 的精度从 ~1.0 掉到 0.0,模型输出变为乱码(例如"Hello. the. has in..")。这是典型的"内核级测试全绿、端到端推理崩坏"场景——精度回归往往被层层调用栈掩盖,需要系统化的排查。
2. 第一步:Commit 二分定位回归源
调查首先通过 commit 二分法锁定引入回归的提交。关键测试手段是MHA_NO_STRUCTURED=True环境变量,用于在同一个 commit 上开关结构化 prefill 路径:
| Checkpoint | Commit | Accuracy | 状态 |
|---|---|---|---|
| Pre-#82119 基线 | 0aaa35b19c4 | 1.0 | PASS |
| Post-#82119(structured ON) | 6936b243e61 | 1.0 | PASS |
| Post-#82119(structured OFF) | 6936b243e61+MHA_NO_STRUCTURED=True | 1.0 | PASS |
| Post-#82541(TileTensor) | e25b574fad6 | 0.0 | FAIL |
| HEAD of main | various | 0.0 | FAIL |
结论:PR #82119(structured prefill)是干净的,回归完全来自 PR #82541(TileTensor 转换)。这是整个调查的分水岭——排查范围从"两个 PR 都可能出错"收敛到"只看 TileTensor 转换引入的差异"。
3. 测试覆盖缺口:group=16 从未在 AMD 上被测试
405B 在 TP=8 下的每 GPU 形状为:num_heads=16, kv_num_heads=1, depth=128, group=16。即每个 KV head 服务 16 个 Q head。
而现有的test_mha_causal_mask_amd.mojo测试最高只覆盖到group=8。为什么 group=16 是分水岭?这与 AMD GPU 的 MFMA(Matrix Fused Multiply-Add)硬件结构直接相关:
- 当 group ≤ 8 时,AMD buffer 资源的 OOB(out-of-bounds)clamping 会静默屏蔽非活跃 MMA 行上的错误——错误的寄存器/SMEM 布局恰好落在被 clamp 掉的行上,测试"碰巧"通过;
- 当 group=16 时,全部 16 个 MFMA 行都是活跃的,任何布局错误都会真实参与计算并污染 attention score。
从当前仓库源码看,这一覆盖缺口已被修复:test_mha_causal_mask_amd.mojo 中已经加入了group=16的 prefill 与 decode 用例(如num_heads=16, group=16的128x128、1024x1024、seq_len=1, 1024/5000等),这正是文档第 4 条建议的落地。
4. 连续内存 kernel 测试:decode group=16 最初失败,后被 UInt-to-Int 重构修复
4.1 原始 commite25b574fad6上的结果
使用连续 Q/K/V 内存的flash_attention测试:
- Prefill(seq_len > 1)+ group=16:全部形状通过;
- Decode(seq_len=1)+ group=16 + depth=128:失败——2048 个值中有 2 个超出 2% 相对容差,且是BF16 特有(FP32 零错误);group=4 与 group=8 均通过,group=16 恰好在边界上失败。
4.2 UInt-to-Int 重构后的结果
main 分支上合入的 UInt-to-Int 类型重构(commitsa89010e1e19、54d3612089a、2ad7fdb062a)之后,所有连续内存测试全部通过,包括 group=16 decode depth=128。类型语义从无符号到有符号的变化,直接影响了distribute映射计算 SMEM 偏移的方式,从而修复了精度问题。
4.3 关键洞察:连续内存测试其实分派到了错误的 kernel
调查发现一个容易误导人的陷阱:连续内存的flash_attention在seq_len=1时,实际分派到的是prefill kernel(mha[...],BM=128),而不是decode kernel(mha_decoding[...],BM=16)。decode kernel 只有在 KV cache 的flash_attention重载(overload)且is_token_generation=True时才会被触发。
也就是说:连续内存测试从头到尾都没有真正执行过
mha_decoding代码路径。这为后续"paged 测试 wrong-vs-wrong"的盲区埋下了伏笔。
5. Paged KV Cache 测试:连续 vs Paged 精度差是"先存问题"
5.1 混合 CE(带 paged KV 的 prefill)
test_mha_mixed_ce_tg.mojo,group=16 + bf16:通过。注意该测试覆盖的是带缓存上下文的 prefill,不是decode。
5.2 Paged decode(TG + paged KV)
test_batch_kv_cache_flash_attention_causal_mask_ragged_paged.mojo 增加了 group=16 形状(此前被has_nvidia_gpu_accelerator()门控,即只在 NVIDIA 上跑):
- 单分区(cache ≤ 256):通过;
- Split-k(cache > 256):通过;
- bs=1 + 多种 cache 大小:通过;
- 压力测试(20 seeds,bs=4):0 失败;
- 全量套件(多种形状、多种 seeds):出现罕见的、依赖数据的失败(约 0.01 的绝对差,约 0.5 个值,接近 bf16 的 1 ULP)。
5.3 关键结论:continuous-vs-paged 精度差距是"先存"的
同样的连续 vs Paged 不匹配(batch=2,diff=0.01171875)在不含任何 TileTensor 改动的 main上也出现。因此这不是 TileTensor 引入的回归,而是 continuous 与 paged KV cache 路径在 group=16 下的既有精度差异——只是此前从未在 AMD 上测过 group=16。
6. PRegisterBuffer.copy_to_shared 分析:嫌疑排除
6.1 原始 commit 上的失败
TileTensor 的distribute[tt_col_major[warp_m, warp_n, 1]()]方法产生了错误的 SMEM lane 分配。回退到旧的copy_local_to_shared[thread_layout=warp_layout]方式可以修复连续内存 kernel 测试。
6.2 UInt-to-Int 重构后
copy_to_shared的 bug 被 UInt-to-Int 重构顺带修复——连续内存测试无需改动copy_to_shared即通过。签名与无符号语义的变化很可能改变了distribute映射计算 SMEM 偏移的方式。
6.3 关键结论
对 smoke test 单独应用copy_to_shared修复(回退到copy_local_to_shared)并不能修复 405B smoke test——精度仍是 0.0。这证明copy_to_shared不是 smoke test 失败的首要原因,排查视线必须继续上移。
7. Smoke Test 状态:所有 kernel 级测试通过,端到端仍失败
将 TileTensor 改动 cherry-pick 到当前 main 后,405B smoke test(accuracy=0.0)持续失败,尽管所有 kernel 级测试都通过:
- 连续内存 group=16 decode:PASS
- Paged group=16 decode:PASS(存在罕见的先存精度差)
- Paged group=16 prefill:PASS
- 所有既有测试形状:PASS
7.1 Smoke test 与 kernel 测试的六个差异点
- 通过bazel(
./bazelw run smoke-test)编译,nn.mojopkg以编译后包形式构建,而非从源码即时编译; - 使用graph compiler对模型图做 JIT 编译;
- 运行126 层 transformer,微小错误逐层累积放大;
- TP=8,跨 8 块 GPU;
- 使用FP8 量化模型权重;
- 走完整 serving pipeline,而非直接 kernel 调用。
7.2 待解之谜
root cause 候选方向:
- bazel 编译包与
mojo源码编译方式的差异; - 与 graph compiler JIT 的交互;
- 完整 pipeline 激活的、kernel 测试未覆盖的另一条代码路径;
- 每层极小的精度差异在 126 层中复合放大。
8. Op Shapes 参考:405B TP=8 每 GPU 的 kernel 几何形状
8.1 Prefill kernel(mha[...])
- BM=128, BN=64, BK=32, WM=32, WN=64
- Grid:(num_heads=16, ceil(seq/BM), batch)
8.2 Decode kernel(mha_decoding[...])
- BM=16, BN=128, BK=32, WM=16, WN=32
- num_threads=256(1 × 4 × 64)
- Grid:(num_partitions, num_heads//group=1, batch)
- cache_length > 256 时启用 split-k
8.3 MMA 形状(gfx950,token_gen,depth=128)
- 16×16×32 MFMA
- fragment_layout:row_major(1, 4)
- warp_layout:col_major(16, 4)
这些几何参数是理解第 16 节根因的必备上下文:BM=16 的 decode kernel 中,16 个 MFMA 行恰好对应 group=16 的 16 个活跃行——任何一行布局错误都会立刻暴露。
9. 哈希级验证:连续路径是 Bitwise Identical
利用 PR #83067 引入的基于哈希的可复现性测试(Anand 的方案),确认当前 main 上的 TileTensor 代码与 TileTensor 之前的代码在所有测试形状上产生 bitwise 一致的结果:
| Test | TileTensor Hash | Main Hash | Match |
|---|---|---|---|
| group=4 prefill 128x128 | 15578270350029038373 | 15578270350029038373 | YES |
| group=16 prefill 128x128 | 18282737304533803813 | 18282737304533803813 | YES |
| group=16 decode 1x12031 | 4927928534235661093 | 4927928534235661093 | YES |
| group=8 decode 1x12031 | 8439226780733698853 | 8439226780733698853 | YES |
结论:main 上的 UInt-to-Int 重构(a89010e1e19、54d3612089a、2ad7fdb062a)似乎解决了原始e25b574fad6commit 中存在的所有数值差异。
10. Pipeline Dispatch 追踪:serving 走的是 KV cache 重载
完整 serving pipeline 的分派链路:
- Python 层:
flash_attention_ragged()→ops.inplace_custom("mo.mha.ragged.paged") - MOGG 层:
_execute_mha_ragged_paged_scalar_args()→generic_flash_attention_kv_cache_ragged() - Mojo 层:
_flash_attention_dispatch()→gpu_flash_attention[ragged=True]()(flash_attention[ragged=True]的别名) - 这是 KV cache 重载(mha.mojo 中根据 cache 状态计算
is_token_generation),并在中央分派点(mha.mojo)据此选择 decode 路径。
这条链路与 paged KV cache 测试使用的重载完全相同——而我们的测试是通过的。这进一步加深了"smoke test 为何失败"的谜团,直到第 15 节揭示了测试本身的盲区。
11. THE GAP:paged 测试是在"用错误对比错误"
这是本次调查最具方法论价值的一步:
paged KV cache 测试(test_batch_kv_cache_flash_attention_causal_mask_ragged_paged)比较的是 continuous-batching 与 paged 两条路径的结果。但两条路径分派的是同一个mha_decodingkernel、跑同一份TileTensor 代码——如果mha_decoding存在 group=16 bug,两条路径会产生同样的错误答案,比较依然通过!
而连续内存哈希测试用的是 dense Q/K/V 的flash_attention,它分派到mha(prefill kernel),不是mha_decoding——所以哈希对比只能证明 prefill 正确,对 decode 没有任何说服力。
没有任何测试把 group=16 的mha_decoding输出与已知正确参照做对比。这就是缺失的测试。
这一洞察最终沉淀为新测试 test_mha_decoding_vs_naive.mojo(约 39 秒运行时间)。其文件头注释直白地写明了动机:"test_batch_kv_cache_flash_attention_*比较的是 continuous-vs-paged(两者走同一个mha_decodingkernel,所以检测不到mha_decoding自身的 bug)"。测试通过 KV cache 重载(max_prompt_length=1触发is_token_generation=True)调用mha_decoding,用mha_gpu_naive在容差内做正确性校验,并用已知好值的 MI355 哈希钉死 bitwise 可复现性(compute_hash实现了 FNV-1a 风格的 64 位哈希)。
12. 决定性复现:通过 KV cache 路径触发 mha_decoding
用哈希测试比较 KV cache 重载(触发mha_decoding)在 TileTensor 分支与 main 上的输出:
| Config | Main Hash | TileTensor Hash | Match |
|---|---|---|---|
| group=4, kv_heads=8 | 3842828076352676645 | 3842828076352676645 | YES |
| group=8, kv_heads=1 | 7229237656969765669 | 7229237656969765669 | YES |
| group=16, kv_heads=1, small cache | 3388747208356569893 | 15015457017215230757 | NO |
| group=16, kv_heads=1, large cache | 7153609791466414885 | 12939245640425186085 | NO |
单分区(无 split-k)与多分区都失败——bug 在mha_decoding内核核心,而非 split-k 归约。
13. 根因:TileTensor 与 LayoutTensor 的 SMEM 写/读布局不匹配
13.1 Smoking gun
在 kv_buffer.mojo 的KVBufferImpl中,token_gen=True的load_from_shared存在一个显式注释与 LayoutTensor fallback:
else: # Token-gen: use LayoutTensor path (TileTensor distribute # produces different offsets for single-row token-gen tiles).即 SMEM读路径使用 LayoutTensor 分布(mma_op.load_b,经_load_matrix_frag/ds_read_tr16_b64),而 SMEM写路径使用 TileTensor 分布(load_from_dram的RegTileLoader.load()、copy_to_shared的tt_copy_local_to_shared)。
13.2 为什么会坏
RegTileLoader.load以col-major索引存寄存器(dst_idx = i + j * M),而旧的copy_dram_to_local以row-major存储;随后tt_copy_local_to_shared以 col-major 读寄存器(与 RegTileLoader 匹配),copy_local_to_shared以 row-major 读(与旧 DMA 匹配)。在各自约定内(旧的全程 row-major、新的全程 col-major),寄存器 ↔ SMEM 映射是自洽的;但load_from_shared的 workaround 从 TileTensor 切回 LayoutTensor,打破了自洽性:
WRITE PATH(TileTensor 约定): DRAM → 寄存器(RegTileLoader,col-major 寄存器) 寄存器 → SMEM(tt_copy_local_to_shared,col-major 读) → SMEM 内容位于 TileTensor 排序的位置 READ PATH(LayoutTensor 约定,workaround): SMEM → MMA 寄存器(mma_op.load_b,期望 LayoutTensor 的 SMEM 顺序) → 从错误的 SMEM 位置读取对 group ≤ 8,valid_rows < 16意味着 OOB-clamp 的 MMA 行屏蔽了错误;对 group=16,所有行都有效,错误数据直接产生错误的 attention score。
13.3 实验验证
- 移除
load_from_sharedworkaround(所有路径都用 TileTensor load_b):group=4 也坏了——证实TiledMmaOp.load_b对这些 tile 形状确实产生与mma_op.load_b不同的结果; - 让
copy_to_shared改用 LayoutTensor(与load_from_shared对齐):group=16 哈希仍不同——因为load_from_dram(RegTileLoader)仍以 col-major 写寄存器,而 LayoutTensor 的copy_to_shared以 row-major 读; load_from_dram与copy_to_shared都用 LayoutTensor:编译错误——copy_dram_to_local期望特定 element_layout 的 LayoutTensor src,而TileTensor.to_layout_tensor()产生标量元素。
13.4 TileTensor distribute 分析
深入对比 tile_tensor.mojo 的distribute/distribute_with_offset与 LayoutTensordistribute:对平面 2D row-major 布局,两者产生相同的偏移公式:
thread_coord_i = (thread_id // thread_stride[i]) % thread_shape[i] offset = sum(thread_coord_i * data_stride[i])真正的差异不在distribute本身,而在于:
TiledMmaOp.load_b用distribute[col_major[...]]+ 在 MMA 子 tile 上 vectorize——通用做法;mma_op.load_b(旧 LayoutTensor)用_load_matrix_frag,其内部调用ds_read_tr16_b64——硬件特定的 LDS 转置读 intrinsic,按硬件定义的 pattern 读元素。
两种方式产生不同的寄存器级 MMA operand 布局:硬件 intrinsic 在 LDS 读取期间执行了一次物理转置,而通用distribute无法复刻这次转置。
13.5 RegTileLoader.load vs copy_dram_to_local
RegTileLoader.load:用worker_idx = lane_id()(warp scope)或thread_idx.x(block scope),以col-major顺序存 dst;copy_dram_to_local:以row-major存 dst(LayoutTensor 原生顺序);- 两者使用相同的线程分布公式。
col-major 与 row-major 的寄存器存储在各约定内部自洽(RegTileLoader↔tt_copy_local_to_shared,copy_dram_to_local↔copy_local_to_shared),一旦 workaround 混用约定(TileTensor 写 + LayoutTensor 读),就产生跨约定错位。
从当前仓库的 kv_buffer.mojo 源码看,这一问题区域已演进为精细的逐分支处理:对非转置 V 路径的 strided SMEM tile,代码明确注释了 TileTensorvectorize[simd, 1]只跟踪标量element_size、会丢失 element-layout 步长、从而在 strided tile 上发出一条读错字节的连续load[width=simd],而 LayoutTensor 的vectorize通过zipped_divide保留 element layout、逐元素迭代——这正是 LayoutTensor 路径可行、而 TileTensor 路径需要显式 BLOCK 分布标量 strided 读的原因。该处实现同时给出了 MFMA B 非转置寄存器布局的关键语义:lane (tr, tc)(tr = lane // MMA_N、tc = lane % MMA_N)持有input_frag_size个连续K 行,而非 CYCLIC 分布。
14. 修复选项评估
Option A:token_gen 的三步全部改用 LayoutTensor
修复load_from_dram与copy_to_shared使其在 token_gen 时使用 LayoutTensor,与现有load_from_sharedworkaround 对齐。受阻:copy_dram_to_local与 TileTensor 派生的 LayoutTensor 存在类型不匹配(element_layout / SIMD 宽度)。
Option B:让 TiledMmaOp.load_b 匹配 mma_op.load_b
让TiledMmaOp.load_b使用与旧mma_op.load_b相同的_load_matrix_frag/ds_read_tr16_b64硬件 intrinsic,然后移除load_from_sharedworkaround,三步统一用 TileTensor。这是工程上最直接的对齐方案。
Option C:用 distribute 让 TiledMmaOp.load_b 产生正确布局
精确理解ds_read_tr16_b64产生的寄存器布局,并用 TileTensordistribute配合正确的 thread layout 与向量化 pattern 复刻。这是最干净的 TileTensor 原生修复——从当前 kv_buffer.mojo 的实现看,BF16 K 转置路径已通过TiledMmaOp.load_b+distribute+swizzle对齐 TensorCore.load_b 的 vector-granularity 语义,而非转置 V 路径则显式发射 BLOCK 分布的 strided 标量读,两条路径都已摆脱了对 LayoutTensor workaround 的依赖,与 Option C 的方向一致。
15. 下一步与建议
调查尾声提出的下一个测试方案:构造带数据的 KV cache → 用 KV cache 调用flash_attention(触发mha_decoding)→ 用mha_gpu_naive计算同样的 attention → 对 group=16 比较两者结果。可能的剩余解释:
- 仅在完整
nn.mojopkg(所有 op 一起编译)中显现的编译期模块交互差异; - graph compiler 调用 kernel 的运行时差异,直接测试无法覆盖;
structured_kernels/amd_tile_io.mojo对相邻代码编译的影响。
最终建议(从当前仓库看已部分落实):
- kernel 级 MHA 代码是正确的——group=16 下连续内存与 paged KV cache 测试均通过;
- group=16 下 continuous-vs-paged 精度差是先存问题,应单独调查,不阻塞 TileTensor 重新合入;
- smoke test 失败需要继续调查bazel 编译 / graph compiler / 完整 pipeline 的交互;
- 永久加入 group=16 测试用例防止回归——
test_mha_causal_mask_amd.mojo与test_batch_kv_cache_flash_attention_causal_mask_ragged_paged.mojo均需覆盖;此外,新增的test_mha_decoding_vs_naive.mojo以mha_gpu_naive+ 已知好值哈希双重校验,直接填补了"没有参照物对比mha_decoding"的盲区。
总结:方法论要点
- 先二分、再覆盖、后哈希:commit 二分缩小范围 → 检查测试覆盖缺口(group=16)→ 用哈希测试钉死 bitwise 可复现性,逐步逼近根因;
- 警惕测试盲区:"比较两条都走同一内核的路径"无法发现内核自身 bug;"分派到 prefill 的连续测试"无法覆盖 decode 路径;
- 布局语义是 GPU 内核正确性的根基:TileTensor 与 LayoutTensor 的寄存器/SMEM 映射、row-major 与 col-major 存储约定、硬件 LDS 转置 intrinsic(
ds_read_tr16_b64)与通用distribute之间的语义差异,会在特定几何形状(group=16)下从"无害差异"变成"静默错误"; - 端到端与 kernel 级测试的鸿沟:bazel 编译包、graph compiler JIT、126 层误差复合、FP8 量化都可能让 kernel 级全绿的代码在完整 pipeline 中失败,两者必须分别对待、分别排查。
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考