Skip to content

架构总览

本页提供 CuFlash-Attn 的全面架构视图,面向需要理解系统设计的研究人员和工程师。


系统架构


数据流

前向传播

反向传播


内存布局


内核分块策略

分块维度

前向 tile 集中定义在 src/kernels/impl/tile_io.cuhForwardTilingConfig):

参数描述scalar 前向WMMA 前向
B_rQuery 分块大小64(hd128 为 32)64(hd128 为 32)
B_cKey/Value 分块大小64(hd128 为 32)32
D头维度32, 64, 12832, 64, 128
T_r每 Query 分块线程数128128

反向使用更保守的 BackwardTilingConfig(shared memory 要容纳更多梯度张量),例如 hd128 为 16×32,hd64 为 32×32

内存复杂度

SRAM=O(Br×D+Bc×D+Br×Bc)

例如 scalar 前向 (Br=64,Bc=64,D=64,FP32):

SRAM 元素=64×64(Q)+64×64(K)+64×64(V)+64×64(S)+64×64(O)+64+64=20608

80 KB(超过默认 48 KB 上限时由 launcher opt-in 动态共享内存),精确值见 ForwardTilingConfig::smem_bytes


目录结构

cuflash-attn/
├── include/cuflash/          # 公开 API 头文件
│   ├── flash_attention.h     # C++ 命名空间 API(含 C ABI 声明)
│   ├── export.h              # 可见性宏
│   └── version.h.in          # 版本头文件模板
├── src/
│   ├── api/                  # API 调度层
│   │   └── flash_attention_api.cu
│   ├── forward/              # 前向内核(统一模板)
│   │   ├── flash_attention_forward_typed.cu   # FP32/FP16/BF16 scalar 前向
│   │   └── flash_attention_forward_wmma.cu    # FP16/BF16 WMMA(Tensor Core)前向
│   ├── backward/             # 反向内核(统一模板,scalar)
│   │   └── flash_attention_backward_typed.cu
│   └── kernels/              # 共享工具
│       ├── impl/             # 内部实现细节(tiling 配置、online softmax、type adapter)
│       ├── matmul.cu / online_softmax.cu / tile_io.cu
│       └── kernel_launch_utils.cuh
└── tests/
    ├── unit/                  # 单元测试(gtest)
    ├── integration/           # API smoke + PyTorch 对比
    └── package_smoke/         # 安装包冒烟

错误处理流程


性能特征

操作内存计算带宽受限
前向O(N)O(N2)是 (低 D)
反向O(N)O(N2)是 (低 D)
重计算O(1)O(N2)

关键洞察

FlashAttention 通过永不物化完整注意力矩阵,将内存从 O(N2) 降至 O(N)。代价是在反向传播时重计算注意力分数,这是计算受限的操作,因此在现代 GPU 上效率很高。

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