📌 项目地址:MoonshotAI/FlashKDA | ⭐ 948 颗星 | 🔧 Cuda | 📜 未标注
核心价值:将Kimi Delta Attention 的 token 级延迟压到极致
Kimi Delta Attention(KDA)是月之暗面在长上下文场景中使用的 attention 变体,它通过门控机制和可学习的衰减向量来兼顾记忆容量与计算效率。FlashKDA 不是一个新的 attention 结构,而是一组基于 CUTLASS 的手写 CUDA 核(kernels),专门替换 flash-linear-attention (FLA) 中原本用 Triton 实现的 chunk_kda 算子。实测在 H20 上,FlashKDA 的前向计算速度是 Triton 版本的 2.5~5 倍(详见仓库内 BENCHMARK_H20.md),并且保证数值精度与 PyTorch bf16 参考实现完全匹配。
这个项目的痛点很具体:当你用 FLA 库跑 KDA 时,默认的 Triton kernel 在大批次或长序列场景下会成为瓶颈。FlashKDA 直接接入 FLA 的自动调度,只需一行环境变量或装好库,就能无感替换,不需要改模型代码。
安装与配置
硬件和驱动要求很严格:必须使用 NVIDIA SM90 架构(H100/H200/B100/B200)或更新的 GPU,CUDA 版本 ≥ 12.9,PyTorch ≥ 2.4。安装步骤:
git clone https://github.com/MoonshotAI/FlashKDA.git flash-kda
cd flash-kda
git submodule update --init --recursive
pip install -v --no-build-isolation .
默认只编译当前显卡的架构。如果要打包成 wheel 或为 CI 编译多个架构,可以显式指定:
FLASH_KDA_CUDA_ARCHS=all pip install -v --no-build-isolation .
支持的值有 auto(默认)、all,或逗号分隔的架构列表,如 90a,100a。
作为 FLA 后端的用法
安装 FlashKDA 后,它会被 flash-linear-attention(版本 ≥0.5.0)的 chunk_kda 函数自动调度。你的模型代码只需要确保在 torch.inference_mode() 上下文中调用:
import torch
from fla.ops.kda import chunk_kda
with torch.inference_mode():
out, final_state = chunk_kda(
q=q, k=k, v=v, g=g, beta=beta,
scale=scale,
initial_state=h0,
output_final_state=True,
use_gate_in_kernel=True,
use_qk_l2norm_in_kernel=True,
use_beta_sigmoid_in_kernel=True,
safe_gate=True,
A_log=A_log, dt_bias=dt_bias,
lower_bound=lower_bound,
transpose_state_layout=True,
cu_seqlens=cu_seqlens,
)
取消 FlashKDA:设置环境变量 FLA_FLASH_KDA=0 即可回退到 Triton 路径。如果想确认是否命中了 FlashKDA,可以加入 logging:
import logging
logging.basicConfig(level=logging.INFO)
运行时会打印 [FLA Backend] kda.chunk_kda -> flashkda 表示命中,否则打印 ... rejected: ... 并说明原因。
性能与注意事项
- 性能差异:FlashKDA 在短序列(≤2K)上优势不大,但在 8K+ 长序列和批量场景下提升明显。具体数据参见仓库内
BENCHMARK_H20.md。 - 数值一致性:测试命令
bash tests/test.sh已包含前向对齐测试,FlashKDA 输出与 PyTorch bf16 参考实现的数值误差在机器 epsilon 量级。 - 局限性:
- 仅支持 bf16 输入,输出也是 bf16。
A_log和dt_bias参数必须是 fp32。- 不支持训练时的反向传播(当前只有前向 kernel)。
- 所有输入张量必须在同一 GPU 上且显存连续。
- 许可证:仓库根目录的
LICENSE文件定义了使用条款,请在使用前确认。 - 相关资源:详细的设计思路可阅读仓库内
docs/20260420-flashkda-v1-deep-dive.md技术博客。