FlashKDA里的L矩阵和Mqk矩阵是什么?详解chunk内注意力的构造

发布时间:2026/9/21 1:37:10
FlashKDA里的L矩阵和Mqk矩阵是什么?详解chunk内注意力的构造 FlashKDA里的L矩阵和Mqk矩阵是什么详解chunk内注意力的构造【免费下载链接】FlashKDAFlashKDA: high-performance Kimi Delta Attention kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashKDAFlashKDA 是基于 CUTLASS 构建的高性能 Kimi Delta AttentionKDACUDA 内核它将 KDA 的计算拆成两个内核K1 负责按 token 并行的准备工作构造L 矩阵与Mqk 矩阵、求逆K2 负责按头并行的循环递推与输出。本文用尽量少的代码讲清楚这两个 16×16 小矩阵在 chunk 内注意力中是怎么来的、又是干什么用的。30秒速览两个矩阵各管一件事KDA 的delta rule要求每个 token 先读取旧状态、再写入修正量。同一段 chunk 里后面 token 写入的内容会污染前面 token 的读取-写入结果必须做级联修正。L 矩阵刻画这种chunk 内互相干扰的依赖矩阵求它的逆矩阵 (I − L)⁻¹ 后一次性把所有级联修正算完Mqk 矩阵chunk 内的衰减注意力分数矩阵decay 版 q·k 点积用 U (I−L)⁻¹·u 一步算出 chunk 内的注意力输出。一句话总结 chunk 内的核心公式U (I − L)⁻¹ ( β ∘ (v − k_decayed S) ) # 修正后的 delta out q_decayed S Mqk U # chunk 内注意力输出 S e^{g_total} · S k_restoredᵀ U # 状态推进其中 S 是进入本 chunk 之前的递归状态。L 和 Mqk 正是让这三个公式可以在 GPU 上纯用矩阵乘法完成的关键。背景KDA 的 delta 规则为什么需要修正KDA 按 token 顺序的递推关系是out_t q_t S_t S_t S_{t-1} · e^{g_t} k_t ⊗ ( v_t − k_t (S_{t-1}·e^{g_t}) ) · β_t括号里的v_t − k_t S_{t-1}就是delta先读出旧状态中已存的信息只把增量写回去。问题在于如果直接把一整个 chunk 内的 token 独立并行展开第 t 个 token 读取的S_{t-1}里其实已经包含了本 chunk 内更早 token 的写入。于是每个 delta 都要被前面的写入逐级修正真正的 U u L·u L²·u L³·u … u 是未修正的原始 delta这正是 L 矩阵存在的意义。L 矩阵的构造严格下三角 beta 缩放先定义 chunk 内的门控前缀和G_i Σ_{j≤i} g_j第 i 个 token 的累积衰减然后k_decayed_i k_i · e^{G_i} # 前向衰减 k_inv_j k_j · e^{−G_j} # 反向衰减 L[i, j] β_i · (k_decayed_i · k_inv_j) β_i · e^{G_i−G_j} · (k_i·k_j)仅当 i j也就是说L 的第 (i, j) 个元素是第 i 个 token 的 beta 门控 × 两者之间衰减差下的 key 相似度。它在源码中的构造是一次 16×16 的 MMA GEMML 的构造与下三角掩码、beta 缩放、(I−L) 组装fwd_kernel1.cuh对应的 PyTorch 参考实现可直接对照理解每个符号torch_ref.py两个容易看错的地方L 只保留严格下三角i j对角线为零。因为 token 对自身的修正已经包含在逆矩阵的单位阵 I里只有下三角才乘 β_i按行缩放上三角直接清零随后INV I − L。Mqk 矩阵的构造含对角线的衰减注意力分数Mqk 和 L 的构造方式完全对称只是把 k 换成了 qq_decayed_i q_i · e^{G_i} · scale Mqk[i, j] e^{G_i−G_j} · (q_decayed_i · k_inv_j)当 i ≥ jMqk 的构造fwd_kernel1.cuh与 L 并行各用一个 warp 做 16×16 MMAMqk 保留主对角线i ≥ j 全部保留上三角清零。对角线元素scale·(q_i·k_i)对应token 对自身修正 delta 的注意力不能丢。对比项L 矩阵Mqk 矩阵尺寸16×16CHUNK1616×16元素公式β_i·e^{G_i−G_j}·(k_i·k_j)scale·e^{G_i−G_j}·(q_i·k_j)保留区域严格下三角ij下三角含对角线i≥j角色delta 干扰矩阵用于求 (I−L)⁻¹chunk 内注意力分数矩阵消费者K1 内部组装 INVK2计算 out qs MqkU显存开销512 字节/块512 字节/块两者都小到可以轻松放进共享内存bf16 下各 512B见 utils.cuh这也是 FlashKDA 选择 CHUNK16 的原因之一。为什么 CHUNK16让逆矩阵免费L 是 16×16 严格下三角矩阵天然幂零L¹⁶ 0。因此 (I−L)⁻¹ 是有限级数可以精确表示(I − L)⁻¹ I − L L² − L³ … L¹⁵FlashKDA 没有直接做 LU 分解而是利用重复平方的 Neumann 级数只花 6 次 16×16 矩阵乘法就得到精确逆L² → L⁴ → L⁸ 三次平方 三次累乘实现见 neumann_inv_fused_1warp 以及调用点 fwd_kernel1.cuh。这也是官方深度解析文档中强调的设计决策16×16 的逆矩阵足够便宜可以用 Neumann 级数直接展开且 CHUNK16 让exp(cumsum(g))的数值范围落在 bf16 可表示区间内省掉大块 chunk 需要的复杂重缩放技巧。完整背景可读 20260420-flashkda-v1-deep-dive.md。K2 内核如何消费 L 和 Mqk四步出结果K1 把k_decayed、q_decayed、INV(I−L)⁻¹、Mqk、k_restored等中间结果通过 TMA 写入全局 workspaceK2 逐 chunk 读取并做 4 步计算源码fwd_kernel2.cuh读旧状态u_acc k_decayed Sout q_decayed S两个 16×128 GEMM算原始 deltau (v − u_acc) · β逐元素寄存器内完成级联修正U INV u—— L 矩阵的逆在这里发力chunk 内输出out Mqk U—— Mqk 在这里发力。最后状态推进S e^{g_total}·S k_restoredᵀ U复用同一份 U无需额外 GEMM。整个 K2 里 L 和 Mqk 各自只参与一次 16×16 的小矩阵乘法代价极低。精度表现bf16 存储的逆矩阵可靠吗INV 以 bf16 存储、中间级数用 fp16 累加逆矩阵元素有界于 [−1,1]fp16 动态范围足够状态用 bf16 存储但更新走 fp32 FMA。下面是 FlashKDA 与fla_chunk_kda在不同 g/beta 配置下的误差对比图Max Error、MAE、RMSE 与误差分布可见 FlashKDA 在绝大多数配置下误差与参考实现同量级小结L 矩阵 chunk 内 delta 规则的干扰依赖矩阵严格下三角按行乘 β其精确逆 (I−L)⁻¹ 把级联修正折叠成一次 GEMMMqk 矩阵 chunk 内衰减注意力分数矩阵下三角含对角线与修正后的 U 相乘即得 chunk 内注意力输出二者都是 16×16 的小矩阵构造、求逆、消费全部走 SM80 MMA 指令这是 FlashKDA 既快又省显存的结构性原因想动手核对每个符号对照 tests/torch_ref.py 的逐 chunk 参考实现即可。理解了 L 与 Mqk你就掌握了 FlashKDA 最核心的设计用两个小矩阵 一次精确求逆把 delta 注意力的串行依赖变成可并行的矩阵运算。【免费下载链接】FlashKDAFlashKDA: high-performance Kimi Delta Attention kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashKDA创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考