尝试对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的运算里, 我们需要以下步骤:
- 通过 来计算出
- 通过 计算出
- 通过 得到下一个
讨论可行的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的数学推导&算子实现分析