Skip to content

CuFlash-Attn从零实现的 CUDA FlashAttention

技术白皮书 · O(N) 内存 · FP32/FP16 · 前向与反向

CuFlash-Attn
v0.5.1稳定版
FP32/16/BF16数值路径
WMMA前向 Tensor Core
C++ / CUDA最小依赖

O(N) 内存

通过 FlashAttention 分块技术避免在 HBM 中物化 O(N²) 注意力矩阵。

算法详解 →
📦

零依赖

纯 CUDA C++,无 PyTorch、无 Cutlass、无 Triton。理解每一行代码,修改每一个细节。

Kernel 逐行解读 →
🔄

完整训练支持

前向与反向传播,含梯度重计算。FP32 与 FP16,数值安全累加。

API 参考 →
🎯

可配置架构编译

源码可配置 sm_70 到 sm_90;当前公开结果只以实际归档的硬件与版本为准。

基准测试 →
📐

稳定 C ABI

稳定的 C ABI,便于与 Python、Rust 或任何支持 FFI 的语言集成。

C API 文档 →
🔬

轻量维护

文档、工作流与仓库结构保持精简,并与实际库边界持续对齐。

项目状态 →

验证状态

公开数字只在原始命令、硬件/软件、commit 与结果产物齐全时引用。历史跨 GPU 表格已隔离为不可审计快照; 当前优先级是以同一输入契约复测 PyTorch SDPA、官方 FlashAttention 与本实现,并归档 JSON 和 profiler 产物。

查看结果发布门槛 →

快速开始

5 分钟内构建并运行:

bash
git clone https://github.com/open-infra-ai/cuflash-attn.git
cd cuflash-attn

cmake --preset release
cmake --build --preset release

ctest --preset release --output-on-failure
cpp
#include "cuflash/flash_attention.h"

auto err = cuflash::flash_attention_forward(
    d_Q, d_K, d_V, d_O, d_L,
    batch_size, num_heads, seq_len, head_dim,
    scale, true, stream
);
python
import ctypes
lib = ctypes.CDLL("./build/release/libcuflash_attn.so")

lib.cuflash_attention_forward_f32(
    q_ptr, k_ptr, v_ptr, o_ptr, l_ptr,
    B, H, N, D, scale, True, None
)

核心参考文献

FlashAttention — Dao et al., NeurIPS 2022.
arXiv:2205.14135
FlashAttention-2 — Dao, ICLR 2024.
arXiv:2307.08691
Online Softmax — Milakov & Gimelshein.
arXiv:1805.02867

稳定 v0.5.0 基线 · 精简 CUDA FlashAttention 参考实现