当前位置:
AIGC文章详情

同一个 KDA 内核,三种实现跑出 2 倍差距:Kimi K3 FlashKDA 背后的性能攻防

源自97位全网作者

17:46

Kimi K3 开源满月,模型层面的讨论基本尘埃落定。arXiv但有个角落,最近两周才真正进入深水区:K3 在 93 层注意力里有 69 层用的是线性注意力 KDA(Kimi Delta Attention),而月之暗面把配套的 KDA 内核 FlashKDA 也开源了。知乎GitHub

有意思的不是开源本身,而是开源之后发生的事:社区迅速围绕同一个算法做出了三种实现,数字放在一起看,相当有戏剧性。

先把三组数字摆出来

  • 官方 FlashKDA(CUTLASS C++ 实现):在 H20 上比 FLA 的 Triton 实现快 1.72–2.22 倍,这是 Moonshot 官方口径,第三方对照仓库自测后确认"区间略有出入但基本一致"。知乎

  • Triton-TLE 复现:换到 H800 上,用基于 Triton 的扩展把官方 FlashKDA 反超了——12 组配置几何平均加速 1.38 倍,两个重点长文本配置约 1.4 倍。知乎

  • CAKE 自动生成的内核:以 PR #4262 提交给 FlashInfer,在 B200 上对官方 FlashKDA 几何平均快 2.05 倍,混合变长序列场景 2.29 倍,最好的一组 2.51 倍。GitHub

同一套 KDA 算法、同样的递推语义,三种实现跨三种卡,差距拉到 2 倍以上。而且排序完全不是很多人直觉里的"手写 CUTLASS 最快"——在 H800 和 B200 上,官方版本都是被反超的那个。

这不是打脸爽文。往细了看,每一组数字背后都站着一个真实的技术判断,串起来正好是一张"2026 年线性注意力内核工程"的地图。

KDA 内核到底难在哪

KDA 是 Kimi Linear 那条线上的产物,可以理解为对 Gated DeltaNet 的细化:不再给整个状态一个统一的标量衰减,而是按通道给不同的衰减系数。它把历史压缩进一份固定大小的状态里递推更新,避免了全注意力的二次方膨胀——K3 用它扛 1M 上下文,靠的就是这个。知乎

但落到 GPU 上,KDA 有个天然的结构性矛盾:

  • 递推是串行的:后一个 token 的状态,必须等前一个算完;

  • 纯串行又没法喂饱 Tensor Core。

FlashKDA 的解法是把序列按 16 个 token 切成 chunk:chunk 内部整体改写成小矩阵运算,可以大规模并行;chunk 之间只保留一条状态递推链。在此基础上,官方版本做了一个关键决定——拆成两个内核:

  • K1 负责并行准备:网格覆盖所有 (序列, head, chunk),256 线程一个 CTA,用 `launch_bounds(256, 8)` 压到每线程 32 个寄存器,换 8 blocks/SM 的高占用率,追求"量大管饱";

  • K2 负责串行递推:网格只覆盖 (序列, head),每个 CTA 192 线程——4 个 MMA warp 加 1 个 TMA load warp、1 个 TMA store warp,标准 warp specialization,输入 3 级流水、输出 2 级流水。

按 FlashKDA 仓库里的 deep-dive 文档口径,这次拆分带来至少 15% 的端到端收益。拆开不是为了优雅,是因为两个阶段的资源诉求完全相反:K1 要并发度,K2 要把片上资源集中喂给一条状态链。GitHub

同一个 KDA 内核,三种实现跑出 2 倍差距:Kimi K3 FlashKDA 背后的性能攻防

源码里还有一批值得抄的细节:16×16 的严格下三角矩阵 L 满足 L¹⁶=0,求逆用 Neumann 展开,通过平方倍增把 15 项压成 3 次 MMA;两块生命周期不重叠的 buffer 用 union 叠放,一处就省出约 14KB shared memory;状态用 BF16 存、FP32 做 FMA 更新(128×128 的状态正好 32KB),内部测试认为无可测精度损失;指数运算下沉到 `ex2.approx` 这类 PTX 近似指令,再用 lower_bound=-5 卡住衰减的数值范围,保证所有因果块都能走稠密 Tensor Core 矩阵乘,绕开了逐位置对角的慢路径。知乎

这套东西是教科书级的 Hopper 内核工程:TMA、warp specialization、mbarrier、swizzle,一个不少。

同一个 KDA 内核,三种实现跑出 2 倍差距:Kimi K3 FlashKDA 背后的性能攻防

反转一:H800 上,Triton 系复现反超 1.38 倍

7 月底,有人用 Triton-TLE(一个在 Triton 之上把共享内存对象、异步搬运、流水线、warp specialization 组织成一等抽象的扩展)重写了整套 KDA 内核,在 H800 上测出对官方 FlashKDA 几何平均 1.38 倍的加速。知乎

同一个 KDA 内核,三种实现跑出 2 倍差距:Kimi K3 FlashKDA 背后的性能攻防

这篇文章最有价值的不是数字,是它诊断出的两个问题:

第一,K1 在 H800 上占用率接近满载,但指令发射率没跟上。占用率高不等于 GPU 在干活——warp 都驻留了,却在等、在空转。TLE 的改法是用显式共享内存张量加异步加载组织数据复用,配合自动调优换更紧凑的 CTA 配置,把发射率顶上去。知乎

第二,K2 的启动网格只覆盖 (序列, head),当 H 不够大时,网格比 SM 数量还小,这时候继续压单个 CTA 的资源已经换不来任何并行度。TLE 干脆把反复使用的主状态留在寄存器里,减少它在 SMEM、寄存器、HBM 之间的反复搬运;整条 K2 里只有 kgᵀ@v_new 那一步状态更新上了 WGMMA。知乎

同一个 KDA 内核,三种实现跑出 2 倍差距:Kimi K3 FlashKDA 背后的性能攻防

注意口径:官方 1.72–2.22 倍是 H20 上、以 FLA Triton 为基线;反超 1.38 倍是 H800 上、以官方 FlashKDA 为基线。卡不同、基线不同、序列形状不同,两组数字谁也不否定谁——它们共同说明的是:内核的最优形态对硬件和工作负载高度敏感,官方调优针对的是自家的主力部署场景,换张卡,平衡点就挪了。

反转二:B200 上,AI 生成的内核快 2.05 倍

8 月初流出的这组数字更狠。FlashInfer 的 PR #4262 收了一个由 CAKE 框架生成的 BF16 递归 KDA prefill 后端,面向 B200/SM100a。PR 里附的 CUPTI 基准(冷 L2 冲刷、BF16 状态、六种定长/变长形状)显示,它对 MoonshotAI/FlashKDA 的加速比在 1.69–2.51 倍之间,几何平均 2.0512 倍。GitHub

为什么 B200 上能拉开这么大?因为路线变了:它没有沿用 K1/K2 两内核拆分,而是把整个 KDA prefill 数据流塞进一个融合内核——chunk 局部中间量全程驻留寄存器、SMEM、TMEM,不落 HBM。

支撑这个路线的是 Blackwell 一代的硬件特性:tcgen05 指令加 TMEM(Tensor Memory)让单个内核能持有的片上状态远超 Hopper,M128 schedule 下一次逻辑 GEMM 是 M=128、N=160、K=32,底层拆成两条 K=16 的 tcgen05.mma 共用一个累加器。chunk 也从 16 放大到 32——直接物化 32 token 的正向/逆向门控因子会超出 BF16 的指数范围,CAKE 的解法是引入一个锚点标量做指数居中,锚点在关键项里对消,数值安全性保住了,块更大、效率更高。知乎

同一个 KDA 内核,三种实现跑出 2 倍差距:Kimi K3 FlashKDA 背后的性能攻防

生产-消费的组织也很讲究:5 个固定的 SMEM 槽位构成环形缓冲,5 组 producer warp 提前准备后续 chunk,consumer 按序消费,mbarrier 负责就绪通知和背压;槽位按张量生命周期做别名复用,这才让 5 级前瞻塞得进 B200 每 CTA 的 SMEM 预算。知乎

圈里有人看了布局之后评论:这套 SMEM 设计"看上去像是 AGENT 辅助搜出来的,人类脑容量终究比不过 AI"。知乎这话半开玩笑,但指向一个真实的趋势——内核搜索空间大到一定程度后,自动生成的方案开始能和顶级手写内核掰手腕,甚至赢。

顺带一提,第三方实现不止这一家,FF-KDA(flash-flash-kda)等也在同一时期冒出来,KDA 内核已经成了线性注意力圈子的公共擂台。知乎

同一套算法,为什么差这么多

把三组数字放在一起,能提炼出三条不那么显而易见的规律:

1. 硬件代际会改变最优拆分。 Hopper 上 TMA+WGMMA+warp specialization 的组合,让"两内核分工"成为合理取舍;到了 Blackwell,tcgen05 和 TMEM 大幅扩张片上容量,融合内核反而赢。GEMM 圈爱说"瓶颈灵活性"——不同形状瓶颈不同;KDA 内核还要再加一句:不同卡,瓶颈也不同。

2. 资源取舍永远是两头的。 占用率和发射率是两头,并行度和状态驻留是两头。K1 宽浅、K2 窄深,本来是自洽的设计,但只要硬件参数一换(H20 换 H800、Hopper 换 Blackwell),原来的平衡点就不成立。TLE 那篇文章给出的教训值得记住:占用率满了先别高兴,去看发射率。

3. 基准数字要带着条件读。 H20/H800/B200 是三种卡,定长/变长/均匀短序列是三类形状,FLA 和 FlashKDA 是两个不同的基线,几何平均还会再抹平差异。每篇文章的数字在自己的口径里都是真的,但跨口径不能直接比大小。谁要是拿这三组数字给你排出一个"全球最快 KDA 内核",直接拉黑即可。

这些对你有什么用

如果你在做 K3 或 KDA 系模型的推理部署:看你手里是什么卡。H20 上官方 FlashKDA 依然是经过验证的基线;H800 上值得把 Triton-TLE 这条路线拉进来测一轮自己的真实形状;B200 用户直接盯住 FlashInfer——#4262 已经合入 main,接下来值得看的是它进正式发布版本的节奏,以及社区在真实变长流量下的复测。GitHub另外别忘了这些都是内核级微基准,端到端的 TTFT 和吞吐收益还要看 KDA 层在你流量里的实际占比。

如果你在写自己的注意力或递归类内核:FlashKDA 是一套现成的工程模板——两阶段拆分、union 复用 SMEM、BF16 存状态+FP32 更新、近似指令下沉,这些套路跟 KDA 本身无关,任何"并行准备+串行递推"结构的算子都能借。

如果你只是想看个热闹:接下来值得盯的信号有三个——FlashInfer 合入 CAKE 内核后、真实部署场景下的复测数据,FF-KDA 等第三方实现的演进,以及 CUTLASS 本身的节奏。4.7.0 刚上了 CuTe DSL 的 Primitives API 和 Task Scheduling 框架,8 月底 PyPI 上的 nvidia-cutlass-dsl 又更新了 4.7.1。GitHubPyPI内核工具链这一年的迭代速度,比大多数人意识到的快得多。

模型权重迟早会被复刻,架构论文也会被读穿。真正每天在产 token、省钱、拉开部署效率差距的,是这层看不见的水位差。Kimi K3 开源的最值钱的部分,可能恰恰是这些代码里写不出论文的东西。

内容由AI生成
0
扫一下,分享更方便,购买更轻松
0评论

当前文章无评论,是时候发表评论了
提示信息

取消
确认
评论举报

最新文章 热门文章