MoonshotAI/FlashKDA
GitHub: MoonshotAI/FlashKDA
基于 CUTLASS 构建的高性能 KDA(Kimi Delta Attention)CUDA kernel,为 flash-linear-attention 提供加速后端。
Stars: 501 | Forks: 50
# FlashKDA
FlashKDA:Flash Kimi Delta Attention — 基于 CUTLASS 构建的高性能 KDA kernel
## 新闻
- **2026-04-22** — 深度解析博客:FlashKDA v1 背后的设计决策,请阅读[这里](docs/20260420-flashkda-v1-deep-dive.md)。
## 环境要求
- SM90 及以上
- 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 .
```
默认情况下,构建过程会检测当前的 CUDA 设备并针对该架构进行编译。对于 wheel 或 CI 构建,请显式编译所有受支持的架构:
```
FLASH_KDA_CUDA_ARCHS=all pip install -v --no-build-isolation .
```
支持的值包括 `auto`(默认)、`all`,或者以逗号分隔的架构列表,例如 `90a,100a`。
## 将 FlashKDA 作为 FLA 后端使用
安装完成后,FlashKDA 会被 `flash-linear-attention` 的 `chunk_kda` 自动调度(auto-dispatch)。有关集成细节,请参阅 [fla-org/flash-linear-attention#852](https://github.com/fla-org/flash-linear-attention/pull/852)。
**要求**
1. 安装 `flash-linear-attention >= 0.5.0`:
pip install -U flash-linear-attention
2. 在 `torch.inference_mode()` 下调用 `chunk_kda`
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,
)
**退出机制:** 设置 `FLA_FLASH_KDA=0` 以回退到 Triton 路径。
**调试调度:** 添加 `logging.basicConfig(level=logging.INFO)` 以在命中时查看 `[FLA Backend] kda.chunk_kda -> flashkda`,或在未命中时查看 `... rejected: `。
## 性能表现
请参阅 [BENCHMARK_H20.md](BENCHMARK_H20.md)。
## 测试
```
bash tests/test.sh
```
- `tests/test_fwd.py` — 正确性测试(与 torch 参考实现完全匹配;并与 `flash-linear-attention` 进行了对比)
## Kernel API
### `flash_kda.fwd`
```
flash_kda.fwd(q, k, v, g, beta, scale, out, A_log, dt_bias, lower_bound,
initial_state=None, final_state=None, cu_seqlens=None)
```
**参数:**
| 参数 | 数据类型 | 形状 | 描述 |
|---|---|---|---|
| `q` | bf16 | `[B, T, H, K]` | Query |
| `k` | bf16 | `[B, T, H, K]` | Key |
| `v` | bf16 | `[B, T, H, V]` | Value |
| `g` | bf16 | `[B, T, H, K]` | 激活前的 Gate |
| `beta` | bf16 | `[B, T, H]` | Beta logits(激活前;内部会应用 sigmoid) |
| `scale` | float | 标量 | 缩放因子 |
| `out` | bf16 | `[B, T, H, V]` | 输出 tensor |
| `A_log` | fp32 | `[H]` | Log-gate 参数 |
| `dt_bias` | fp32 | `[H, K]` | Gate 偏置 |
| `lower_bound` | float | 标量 | Gate 下限(范围从 -5.0 到 0) |
| `initial_state` | bf16/fp32/None | `[B, H, V, K]` 或 `[N, H, V, K]` | (可选)初始循环状态 |
| `final_state` | bf16/fp32/None | `[B, H, V, K]` 或 `[N, H, V, K]` | (可选,输出)最终循环状态 |
| `cu_seqlens` | int64 | `[N+1]` | (可选)用于变长批处理的累积序列长度 |
- 目前要求 `K = V = 128`。
- `initial_state` / `final_state` 接受 `None`(无状态)、bf16 或 fp32 tensor。如果同时提供这两者,它们的数据类型必须匹配。
- 当提供 `cu_seqlens` 时,`B` 必须为 1,`T` 是所有序列的总长度,且 `initial_state` / `final_state` 的形状为 `[N, H, V, K]`。
- 当 `cu_seqlens` 为 `None` 时,每个批处理元素被视为独立的序列,状态形状为 `[B, H, V, K]`。
## 开发说明
要为 CUDA/C++ 源码设置 IntelliSense (clangd),请运行:
```
bash setup_clangd.sh
```
这将生成一个包含正确仓库路径的 `.clangd` 文件,并将全局 clangd `config.yaml` 安装到 `~/.config/clangd/`。
## 引用
```
@misc{flashkda2026,
title={FlashKDA: Flash Kimi Delta Attention},
author={Yutian Chen, Zhiyuan Li, Yucheng Wang, Ming Wei},
year={2026},
publisher = {GitHub},
howpublished = {\url{https://github.com/MoonshotAI/FlashKDA}},
}
```
标签:AI, CUDA, Vectored Exception Handling, 凭据扫描, 注意力机制, 算子加速, 自动化代码审查, 逆向工具, 高性能计算