ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Mojo MAX 内核排查实录:gfx950 上 TileTensor 化 MHA 的精度回归定位与根因分析

Mojo MAX 内核排查实录:gfx950 上 TileTensor 化 MHA 的精度回归定位与根因分析 Mojo MAX 内核排查实录gfx950 上 TileTensor 化 MHA 的精度回归定位与根因分析【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo导读本文完整还原了 MAX 内核仓库中一次典型的 GPU 内核回归排查过程在 AMD Instinct MI355Xgfx950上两个 MHAMulti-Head AttentionPR 落地后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.01.1 测试环境调查发生在 2026 年 4 月 9-10 日硬件与模型配置如下项目配置硬件8 × AMD Instinct MI355X OAMgfx950TP8模型Llama-3.1-405BRedHatAI/Meta-Llama-3.1-405B-Instruct-FP8-dynamic调查对象 PRPR #821196936b243e61Structured prefill kernelPR #82541e25b574fad6AMD CDNA attention 内部逻辑的 TileTensor 转换1.2 问题现象两个 MHA 相关 PR 落地后405B smoke test 的精度从 ~1.0 掉到 0.0模型输出变为乱码例如Hello. the. has in..。这是典型的内核级测试全绿、端到端推理崩坏场景——精度回归往往被层层调用栈掩盖需要系统化的排查。2. 第一步Commit 二分定位回归源调查首先通过 commit 二分法锁定引入回归的提交。关键测试手段是MHA_NO_STRUCTUREDTrue环境变量用于在同一个 commit 上开关结构化 prefill 路径CheckpointCommitAccuracy状态Pre-#82119 基线0aaa35b19c41.0PASSPost-#82119structured ON6936b243e611.0PASSPost-#82119structured OFF6936b243e61MHA_NO_STRUCTUREDTrue1.0PASSPost-#82541TileTensore25b574fad60.0FAILHEAD of mainvarious0.0FAIL结论PR #82119structured prefill是干净的回归完全来自 PR #82541TileTensor 转换。这是整个调查的分水岭——排查范围从两个 PR 都可能出错收敛到只看 TileTensor 转换引入的差异。3. 测试覆盖缺口group16 从未在 AMD 上被测试405B 在 TP8 下的每 GPU 形状为num_heads16, kv_num_heads1, depth128, group16。即每个 KV head 服务 16 个 Q head。而现有的test_mha_causal_mask_amd.mojo测试最高只覆盖到group8。为什么 group16 是分水岭这与 AMD GPU 的 MFMAMatrix Fused Multiply-Add硬件结构直接相关当 group ≤ 8 时AMD buffer 资源的 OOBout-of-boundsclamping 会静默屏蔽非活跃 MMA 行上的错误——错误的寄存器/SMEM 布局恰好落在被 clamp 掉的行上测试碰巧通过当 group16 时全部 16 个 MFMA 行都是活跃的任何布局错误都会真实参与计算并污染 attention score。从当前仓库源码看这一覆盖缺口已被修复test_mha_causal_mask_amd.mojo 中已经加入了group16的 prefill 与 decode 用例如num_heads16, group16的128x128、1024x1024、seq_len1, 1024/5000等这正是文档第 4 条建议的落地。4. 连续内存 kernel 测试decode group16 最初失败后被 UInt-to-Int 重构修复4.1 原始 commite25b574fad6上的结果使用连续 Q/K/V 内存的flash_attention测试Prefillseq_len 1 group16全部形状通过Decodeseq_len1 group16 depth128失败——2048 个值中有 2 个超出 2% 相对容差且是BF16 特有FP32 零错误group4 与 group8 均通过group16 恰好在边界上失败。4.2 UInt-to-Int 重构后的结果main 分支上合入的 UInt-to-Int 类型重构commitsa89010e1e19、54d3612089a、2ad7fdb062a之后所有连续内存测试全部通过包括 group16 decode depth128。类型语义从无符号到有符号的变化直接影响了distribute映射计算 SMEM 偏移的方式从而修复了精度问题。4.3 关键洞察连续内存测试其实分派到了错误的 kernel调查发现一个容易误导人的陷阱连续内存的flash_attention在seq_len1时实际分派到的是prefill kernelmha[...]BM128而不是decode kernelmha_decoding[...]BM16。decode kernel 只有在 KV cache 的flash_attention重载overload且is_token_generationTrue时才会被触发。也就是说连续内存测试从头到尾都没有真正执行过mha_decoding代码路径。这为后续paged 测试 wrong-vs-wrong的盲区埋下了伏笔。5. Paged KV Cache 测试连续 vs Paged 精度差是先存问题5.1 混合 CE带 paged KV 的 prefilltest_mha_mixed_ce_tg.mojogroup16 bf16通过。注意该测试覆盖的是带缓存上下文的 prefill不是decode。5.2 Paged decodeTG paged KVtest_batch_kv_cache_flash_attention_causal_mask_ragged_paged.mojo 增加了 group16 形状此前被has_nvidia_gpu_accelerator()门控即只在 NVIDIA 上跑单分区cache ≤ 256通过Split-kcache 256通过bs1 多种 cache 大小通过压力测试20 seedsbs40 失败全量套件多种形状、多种 seeds出现罕见的、依赖数据的失败约 0.01 的绝对差约 0.5 个值接近 bf16 的 1 ULP。5.3 关键结论continuous-vs-paged 精度差距是先存的同样的连续 vs Paged 不匹配batch2diff0.01171875在不含任何 TileTensor 改动的 main上也出现。因此这不是 TileTensor 引入的回归而是 continuous 与 paged KV cache 路径在 group16 下的既有精度差异——只是此前从未在 AMD 上测过 group16。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_layoutwarp_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 testaccuracy0.0持续失败尽管所有 kernel 级测试都通过连续内存 group16 decodePASSPaged group16 decodePASS存在罕见的先存精度差Paged group16 prefillPASS所有既有测试形状PASS7.1 Smoke test 与 kernel 测试的六个差异点通过bazel./bazelw run smoke-test编译nn.mojopkg以编译后包形式构建而非从源码即时编译使用graph compiler对模型图做 JIT 编译运行126 层 transformer微小错误逐层累积放大TP8跨 8 块 GPU使用FP8 量化模型权重走完整 serving pipeline而非直接 kernel 调用。7.2 待解之谜root cause 候选方向bazel 编译包与mojo源码编译方式的差异与 graph compiler JIT 的交互完整 pipeline 激活的、kernel 测试未覆盖的另一条代码路径每层极小的精度差异在 126 层中复合放大。8. Op Shapes 参考405B TP8 每 GPU 的 kernel 几何形状8.1 Prefill kernelmha[...]BM128, BN64, BK32, WM32, WN64Grid(num_heads16, ceil(seq/BM), batch)8.2 Decode kernelmha_decoding[...]BM16, BN128, BK32, WM16, WN32num_threads2561 × 4 × 64Grid(num_partitions, num_heads//group1, batch)cache_length 256 时启用 split-k8.3 MMA 形状gfx950token_gendepth12816×16×32 MFMAfragment_layoutrow_major(1, 4)warp_layoutcol_major(16, 4)这些几何参数是理解第 16 节根因的必备上下文BM16 的 decode kernel 中16 个 MFMA 行恰好对应 group16 的 16 个活跃行——任何一行布局错误都会立刻暴露。9. 哈希级验证连续路径是 Bitwise Identical利用 PR #83067 引入的基于哈希的可复现性测试Anand 的方案确认当前 main 上的 TileTensor 代码与 TileTensor 之前的代码在所有测试形状上产生 bitwise 一致的结果TestTileTensor HashMain HashMatchgroup4 prefill 128x1281557827035002903837315578270350029038373YESgroup16 prefill 128x1281828273730453380381318282737304533803813YESgroup16 decode 1x1203149279285342356610934927928534235661093YESgroup8 decode 1x1203184392267807336988538439226780733698853YES结论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[raggedTrue]()flash_attention[raggedTrue]的别名这是 KV cache 重载mha.mojo 中根据 cache 状态计算is_token_generation并在中央分派点mha.mojo据此选择 decode 路径。这条链路与 paged KV cache 测试使用的重载完全相同——而我们的测试是通过的。这进一步加深了smoke test 为何失败的谜团直到第 15 节揭示了测试本身的盲区。11. THE GAPpaged 测试是在用错误对比错误这是本次调查最具方法论价值的一步paged KV cache 测试test_batch_kv_cache_flash_attention_causal_mask_ragged_paged比较的是 continuous-batching 与 paged 两条路径的结果。但两条路径分派的是同一个mha_decodingkernel、跑同一份TileTensor 代码——如果mha_decoding存在 group16 bug两条路径会产生同样的错误答案比较依然通过而连续内存哈希测试用的是 dense Q/K/V 的flash_attention它分派到mhaprefill kernel不是mha_decoding——所以哈希对比只能证明 prefill 正确对 decode 没有任何说服力。没有任何测试把 group16 的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_length1触发is_token_generationTrue调用mha_decoding用mha_gpu_naive在容差内做正确性校验并用已知好值的 MI355 哈希钉死 bitwise 可复现性compute_hash实现了 FNV-1a 风格的 64 位哈希。12. 决定性复现通过 KV cache 路径触发 mha_decoding用哈希测试比较 KV cache 重载触发mha_decoding在 TileTensor 分支与 main 上的输出ConfigMain HashTileTensor HashMatchgroup4, kv_heads838428280763526766453842828076352676645YESgroup8, kv_heads172292376569697656697229237656969765669YESgroup16, kv_heads1, small cache338874720835656989315015457017215230757NOgroup16, kv_heads1, large cache715360979146641488512939245640425186085NO单分区无 split-k与多分区都失败——bug 在mha_decoding内核核心而非 split-k 归约。13. 根因TileTensor 与 LayoutTensor 的 SMEM 写/读布局不匹配13.1 Smoking gun在 kv_buffer.mojo 的KVBufferImpl中token_genTrue的load_from_shared存在一个显式注释与 LayoutTensor fallbackelse: # 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 PATHTileTensor 约定 DRAM → 寄存器RegTileLoadercol-major 寄存器 寄存器 → SMEMtt_copy_local_to_sharedcol-major 读 → SMEM 内容位于 TileTensor 排序的位置 READ PATHLayoutTensor 约定workaround SMEM → MMA 寄存器mma_op.load_b期望 LayoutTensor 的 SMEM 顺序 → 从错误的 SMEM 位置读取对 group ≤ 8valid_rows 16意味着 OOB-clamp 的 MMA 行屏蔽了错误对 group16所有行都有效错误数据直接产生错误的 attention score。13.3 实验验证移除load_from_sharedworkaround所有路径都用 TileTensor load_bgroup4 也坏了——证实TiledMmaOp.load_b对这些 tile 形状确实产生与mma_op.load_b不同的结果让copy_to_shared改用 LayoutTensor与load_from_shared对齐group16 哈希仍不同——因为load_from_dramRegTileLoader仍以 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_localRegTileLoader.load用worker_idx lane_id()warp scope或thread_idx.xblock scope以col-major顺序存 dstcopy_dram_to_local以row-major存 dstLayoutTensor 原生顺序两者使用相同的线程分布公式。col-major 与 row-major 的寄存器存储在各约定内部自洽RegTileLoader↔tt_copy_local_to_sharedcopy_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[widthsimd]而 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 Atoken_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_bdistributeswizzle对齐 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 → 对 group16 比较两者结果。可能的剩余解释仅在完整nn.mojopkg所有 op 一起编译中显现的编译期模块交互差异graph compiler 调用 kernel 的运行时差异直接测试无法覆盖structured_kernels/amd_tile_io.mojo对相邻代码编译的影响。最终建议从当前仓库看已部分落实kernel 级 MHA 代码是正确的——group16 下连续内存与 paged KV cache 测试均通过group16 下 continuous-vs-paged 精度差是先存问题应单独调查不阻塞 TileTensor 重新合入smoke test 失败需要继续调查bazel 编译 / graph compiler / 完整 pipeline 的交互永久加入 group16 测试用例防止回归——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 二分缩小范围 → 检查测试覆盖缺口group16→ 用哈希测试钉死 bitwise 可复现性逐步逼近根因警惕测试盲区比较两条都走同一内核的路径无法发现内核自身 bug分派到 prefill 的连续测试无法覆盖 decode 路径布局语义是 GPU 内核正确性的根基TileTensor 与 LayoutTensor 的寄存器/SMEM 映射、row-major 与 col-major 存储约定、硬件 LDS 转置 intrinsicds_read_tr16_b64与通用distribute之间的语义差异会在特定几何形状group16下从无害差异变成静默错误端到端与 kernel 级测试的鸿沟bazel 编译包、graph compiler JIT、126 层误差复合、FP8 量化都可能让 kernel 级全绿的代码在完整 pipeline 中失败两者必须分别对待、分别排查。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表