尝试对Kimi KDA的数学推导&算子实现分析

prologue

kimi KDA是kimi k3所使用的kernel, 但远早于k3发布(2026.2.17), 囿于笔者今年3月份才开始探究ai & ai infra, 且近来事务繁忙, 一直都没有认真阅读其论文和公式. Kimi 发布K3 report的时候, 也是对着K3的公式一脸茫然. 读剑林老师的博客的时候, 对模型结构的数学设计感到无比优雅, 心向往之. 恰好笔者近来有较多空闲, 于是便去研读了一遍KDA 的论文和一些网上的资料, 对一些数学推导和算子实现便成为如下笔记, 希望大家原谅笔者的数学水平和难以评价的表达.

mainloop

直接进入正题, kimi k3的关键核心式如下(笔者懒惰懒得花功夫排版了, 大家只需要知道就好:

即:

发现方阵且可独立计算:

对于每一个chunk, 我们给予一个作为

则此时对于每一个, 我们尝试仅根据计算

为了方便展开, 我们令

则:

认为现在已经很明朗了, 尝试继续推进:

对Infra工作而言,

单个计算是无依赖的, 但是考虑到矩阵乘没有交换性, 因此该实现并不优雅, 且若考虑到显存更是毫无linear attention本该具有的优势.

于是考虑另外的方向:

我们重新出发:

发现此时的累乘非常好处理, 于是进行如下定义:

这里我们再令:

则:

然后于是就发生了这样的代换:

代入原式:

进一步化简:

之后似乎陷入了一个困境, 在括号里很难拿出来, 由此两个也无法进行合并.

但是注意到:

这两个之间似乎满足某种神奇的对偶关系.且我们更能发现为对角阵

于是不妨这样定义:

于是, 我们的递推式很好地变成了!:

现在不妨这样定义:

那我们现在的递推式就是:

那现在我们便很方便地写出一个chunk内的状态:

代回:

诶嘿, 这步代换堪称神来之笔, 我们得到了我们想要的标量

于是很自然地, 我们定义:

然后看上去目前天然便是下三角矩阵:

于是, 我们便得到了一个很优美的方程:

即为:

嗯很好地, 我们现在只需要解出即可, 鉴于是下三角矩阵, 有引理:

显然的逆矩阵是有解析解而且比较可求的, 记为:

则:

之前我们有这样的式子:

现在将整个chunk展开:

定义:

以及:

则:

注意我们现在写的公式是对于token C处的展开, 这样我们就得到了下一个chunk的S

现在的核心问题转而变成了我们如何得到chunk内每一个token的

从第(30)式出发固然是一个很好的想法, 但是显然会产生一个prefixsum的东西, 然而大多数情况下不需要我们显式地得到, 而且如果把所有的均算出来之后, 其为一个的一个张量, linear attention的优势何在!?

于是, 我们选择直接将相乘:

所谓仅仅是将乘上了一个scale标量, 无伤大雅. 但是此时式子里有一个重要的东西:

这个东西是一个标量!

因此所有token:

这样来看, 在KDA的运算里, 我们需要以下步骤:

  1. 通过 来计算出
  2. 通过 计算出
  3. 通过 得到下一个

讨论可行的Kernel实现思路:

我们还是先看看kimi官方怎么做的吧, 看上去把该kernel拆成了两个kernel.把chunk的大小固定为16了(难道不太小了吗?)

一个chunk开始的时候, 我们默认已经存在:

官方实现里把chunk-local工作放给了K1, 把依赖上一个chunk的工作放到了K2.

K1的launch grid为

K2的launch grid为. 为了便于理解, 可以直接把H省略掉只剩下单头注意力.

回顾我们上文的推导, 原始的递推方程是这样的:

我们定义了chunk内的decay:

然后进行了归一化:

我们也已经推过:

其中:

又定义:

于是:

这里我们相当于把之前的那些公式重新都写了一遍. 但是现在官方实现中引入了另一个变量把吸收掉了:

那现在的state update就是:

展开:

注意到, 所以

又有:

展开:

嗯, 我们把问题实际放到kernel里, 实际上我们是想解这个problem.

定义:

以及:

则所有的token可以一次写作:

嗯可能我们看到这里已经有一些晕了, 但是重新想起来相当于我们数学推导的时候的.

可以显然明显地看到, 的计算并不依赖, 因此, 计算逆矩阵也并不依赖其他chunk的state, 但是其他的需要.

所以FlashKDA的大致分工即为:

K1

K1负责计算T和S, 具体流程如下:

首先我们首先要构造的是一个chunk的

看上去是一个技术含量不高的scan.

之后构造:

然后做第一个外积:

对于每一个元素乘上当前row的, 然后取严格下三角:

然后计算逆矩阵.

至于这个逆矩阵如何计算, 官方采用了纽曼级数的方法, 这也是我们选的C这么小的原因.

理论上K1的职责到此结束, 但是发现在的计算中还是有一个可以并行计算的项:

好了, 如上可见, 每一个chunk可以正好并行.

K2

K2的主要思路是串行进行所有chunk, 因为S具有依赖性(但是这非常不符合算子人的想法, 不知道对这个有无太多思路).

我们首先要计算:

所以首先计算:

然后V减去它是常见gemm操作.

之后乘和K1算出的:

然后我们便可以去计算O了:

我们已经算过所以很容易地就可以算出

剩下的最后一个问题是

由之前的公式可以直接计算:

不妨直接设:

于是可以计算:

epilogue

好了, 我们完成了对Kimi KDA attention的推导和对flash KDA v1 kernel的分析, 这是笔者第一次接触到linear attention kernel, 对于kda的一些设计也不能太过理解. 但是按照kimi k3的模型质量来看, 应该是没有训炸(x).

最近在对算子很感兴趣, 看flash kda v1的实现, 感觉并不算很成熟, 看flash infer里的一些pr也可以看出很容易地就可以实现2x的加速比.按照笔者的理解, 在K2里面只有head可以并行化, S串行化太过硬伤, 但是Chunksize又不能选大(选大之后neumann invmethod又会出现一些问题). 今天读了一下tian qi新发的论文CAKE: Compiler-Agent Co-Design for Frontier Kernel Evolution, 里面也提到了cake ir对kda的优化, 哎不知道对kda范式的革新是否还像fla一样由人类作出了…

尝试对Kimi KDA的数学推导&算子实现分析

https://hjcheng0602.github.io/blog/kdd/

AuthorHan Jincheng
Posted on08-13-2026
Updated on08-13-2026
kernel optimization KDA linear attention