尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

FlashKDA 16×16 fp16 Neumann求逆逐行讲解:寄存器级矩阵乘法终极指南

FlashKDA 16×16 fp16 Neumann求逆逐行讲解:寄存器级矩阵乘法终极指南 FlashKDA 16×16 fp16 Neumann求逆逐行讲解寄存器级矩阵乘法终极指南【免费下载链接】FlashKDAFlashKDA: high-performance Kimi Delta Attention kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashKDAFlashKDA 是基于 CUTLASS 构建的 Kimi Delta AttentionKDA高性能 CUDA kernel 库。本文将逐行剖析其核心函数 neumann_inv_fused_1warp一个 warp32 线程如何仅靠 Neumann 级数展开 寄存器级mma.m16n8k16矩阵乘法指令完成 16×16 fp16 矩阵求逆——全程不碰共享内存、不依赖循环分解是理解 FlashKDA 数值精妙之处的最佳入口。为什么 Kimi Delta Attention 需要矩阵求逆 KDA 属于线性注意力家族。为了在长序列上并行计算FlashKDA 把序列切成CHUNK 16的块chunked recurrence。块内 delta rule 递推可以写成块级闭式解其中就需要对形如(I - L)的矩阵求逆L是一个严格下三角的 16×16 矩阵由衰减后的 key 内积构造因为L严格下三角所以L^16 0幂零性——这是后面一切技巧的数学根基。为什么偏偏选 16官方 deep-dive 文档给了三个理由数值范围适配 bf16、16×16 求逆远比 64×64 便宜且可免分解、全部运算可映射到 SM80 MMA 指令。详见 docs/20260420-flashkda-v1-deep-dive.md。Neumann 级数把求逆变成 6 次矩阵乘法传统 LU 分解求逆需要 O(n³) 复杂度的通用算法。但利用L^16 0Neumann 级数在此有限收敛(I - L)⁻¹ I L L² ... L¹⁵ L¹⁶ 0求和到此为止FlashKDA 的巧思在于不逐项累加 16 项而是保持恒等式INV (I - L) · (I L²)(I L⁴)(I L⁸) I - L L² - L³ ... L¹⁴ - L¹⁵于是求逆退化为3 次自乘L²→L⁴→L⁸ 3 次累乘INV × L²/⁴/⁸共 6 次 16×16 fp16 矩阵乘法每线程只需 12 条硬件 MMA 指令。寄存器级实现neumann_inv_fused_1warp 逐行拆解函数位于 csrc/smxx/utils.cuh由 K1 kernel 中compute_tid 256的全部线程调用但函数内只让前 32 个线程一个 warp干活auto mma make_tiled_mma( SM80_16x8x16_F16F16F16F16_TN{}, // fp16 输入、fp16 累加器 LayoutShape_1,_1{}, Tile_16,_16,_16{} // 16×16 输出 32 线程 ); if (tid int(size(mma))) return; // 线程 32~255 直接返回寄存器角色分配每个线程持有 4 个uint32_t寄存器 2 个 fp16 数组成的 fragment整个矩阵被摊在 warp 的寄存器堆里寄存器角色L_a(4×u32)L的 A 操作数行主序 fragmentINV_a/INV_c初始I - L/ 累加结果Lpow_c/Lpow_b当前幂L^{2^j}及其转置B 操作数tmp_a/mm_c复用暂存 / 单次乘积结果输入矩阵先通过SM75_U32x4_LDSM_Nldmatrix从共享内存批量搬入寄存器auto smem_copy_A make_tiled_copy_A(Copy_AtomSM75_U32x4_LDSM_N, FP16{}, mma); // L 与 INV 各加载一次之后全程留在寄存器核心积木两条指令拼出 16×16 MMA硬件指令是m16n8k16所以 16×16 需要沿 N 方向执行两次auto mma_16x16 [](uint32_t* d, uint32_t const* a, uint32_t const* b, uint32_t const* c) { SM80_16x8x16_F16F16F16F16_TN::fma(d[0], d[1], a[0], a[1], a[2], a[3], b[0], b[1], c[0], c[1]); SM80_16x8x16_F16F16F16F16_TN::fma(d[2], d[3], a[0], a[1], a[2], a[3], b[2], b[3], c[2], c[3]); };由于 MMA 要求 B 操作数为列主序TN 布局转置不能用共享内存往返而是用SM75_U32x1_MOVM_Tmovmatrix在寄存器内原地完成——这是 K2 中被反复强调的 register-file transpose 技巧的同一招数。逐行走一遍计算链第 1 步L² L × Ltranspose_u32x4(L_a, Lpow_b); // 寄存器内转置 L → B 操作数 clear_u32x4(Lpow_c); // 累加器清零MMA 的 C 必须为 0 mma_16x16(Lpow_c, L_a, Lpow_b, Lpow_c); // L²第 2 步INV INV × L²transpose_u32x4(Lpow_c, Lpow_b); // 顺手转置 L²供第 3 步复用 copy_u32x4(INV_a, INV_c); // 把 I - L 拷入累加器 clear_u32x4(mm_c); mma_16x16(mm_c, INV_a, Lpow_b, mm_c); // mm (I - L) · L² add_fp16x2_u32x4(INV_c, mm_c); // fp16x2 SIMD 加法第 3~6 步L⁴、L⁸ 自乘 两次累加完全相同的模式滚动两次copy_u32x4(Lpow_c, tmp_a); // A L² clear_u32x4(Lpow_c); mma_16x16(Lpow_c, tmp_a, Lpow_b, Lpow_c); // L⁴ L² × L²复用 L² 的转置 // ... 同样地算出 L⁸并两次执行 INV INV × L^{2^j}注意两个工程细节Lpow_b中上一轮的转置被下一轮自乘直接复用省掉重复转置tmp_a作为 A 操作数的暂存位避免污染正在写入的Lpow_c。收尾fp16 → bf16 落回共享内存cute::transform(tCrC, tCrC_bf16, [] __device__ (FP16 x) - BF16 { return BF16(x); }); auto smem_tiled_store make_tiled_copy_C(Copy_AtomSM90_U32x4_STSM_N, BF16{}, mma); // stmatrix 写回 copy(smem_tiled_store, tCrC_st_view, tCsC_st);结果以 bf16 写入共享内存INV槽位随后 K1 通过 TMA 把INV存到 workspace 供 K2 使用csrc/smxx/fwd_kernel1.cuh。为什么选 fp16 而不是 bf16 求逆一个容易被忽略的细节整个求逆用fp16完成。deep-dive 文档的解释是逆矩阵元素有界于[-1, 1]fp16 的窄动态范围完全够用而它比 bf16 多出 3 位尾数10 bit vs 7 bit给 Neumann 级数的 16 项累加留出了额外精度余量同时省去了 bf16 MMA 所需的fp32 → bf16转换docs/20260420-flashkda-v1-deep-dive.md。内部测试验证了 FlashKDA 的数值精度与fla_chunk_kda参考实现高度一致调用链全景它发生在 K1 的哪一步结合 csrc/smxx/fwd_kernel1.cuhK1token 并行grid N×H×num_chunks的完整流水是g 激活 → L2 归一化 → 衰减得到k_decayed/k_inv构造 L 与 Mqk两个 warp 分别用 MMA 算出L k_decayed k_invᵀ与Mqk掩码 初始化256 线程并行把上三角置零、乘sigmoid(β)同时写出INV I - LNeumann 求逆neumann_inv_fused_1warp(L_fp16, INV_fp16, INV, compute_tid)一发入魂TMA 存出INV、Mqk等 workspace交给 K2 做 chunk 间递推。小结✅ 幂零性L^16 0让 Neumann 级数有限收敛把求逆降级为6 次 16×16 fp16 MMA✅INV INV × L^{2^j}的滚动更新只维护一个累加器寄存器占用约 28 个 u32/线程✅ldmatrix进、movmatrix转、stmatrix出共享内存只做首尾搬运计算全在寄存器堆完成✅ fp16 用 3 位尾数换来级数累加的精度余量是精度与性能双赢的取舍。想继续深入 K2 的状态递推与MOVM_T优化可以接着读 csrc/smxx/fwd_kernel2.cuh或从 tests/torch_ref.py 对照 PyTorch 参考实现理解每个张量的含义。【免费下载链接】FlashKDAFlashKDA: high-performance Kimi Delta Attention kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashKDA创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表