- 人工智能
- 大模型
- 算子库
【免费下载链接】flash-attention
Fast and memory-efficient exact attention
导读
本文围绕 flash-attention 仓库中 csrc/fused_dense_lib 这一独立的 CUDA 扩展模块展开,它实现了融合的 matmul + bias(前向与反向)以及 matmul + bias + GELU/ReLU(前向与反向),并额外支持 bfloat16 精度,是训练 GPT 等 Transformer 模型时替代朴素nn.Linear+ 激活函数组合的加速组件。读完本文,你将掌握该扩展的安装与编译细节、三个核心 C++/CUDA 入口的调用方式与形状约束、cuBLASLt 融合 epilogue 的底层原理,以及它在 Tensor Parallel / sequence parallel 场景下如何与高层 Python 模块协作。
一、模块定位:一个"麻雀虽小、五脏俱全"的融合算子库
csrc/fused_dense_lib是 flash-attention 仓库中相对独立的一个子模块,与注意力内核解耦,专门处理 Transformer 里占计算量很大一部分的 MLP / Dense 层。其 README(csrc/fused_dense_lib/README.md)给出了最核心的定位:
- 实现融合的 matmul + bias(前向与反向),以及融合的 matmul + bias + gelu(前向与反向);
- 代码改编自 Apex 的 FusedDense,但关键差异是让它支持 bfloat16;
- 为获得最佳性能,建议使用 CUDA >= 11.8(更早版本的 cuBLAS 对 bfloat16 的 matmul + bias + gelu 融合性能不佳);
- 目前只在 A100 上做过测试。
整个模块只有 4 个文件,构成了一条完整的"PyTorch 扩展"链路:
| 文件 | 职责 |
|---|---|
| csrc/fused_dense_lib/setup.py | 基于torch.utils.cpp_extension的构建脚本,定义编译参数 |
| csrc/fused_dense_lib/fused_dense.cpp | PyTorch 绑定层:参数检查、张量分配、dispatch、pybind11 导出 |
| csrc/fused_dense_lib/fused_dense_cuda.cu | CUDA 实现层:基于 cuBLAS / cuBLASLt 的 GEMM 封装与融合 epilogue |
高层封装flash_attn/ops/fused_dense.py | 提供FusedDense、FusedMLP、ColumnParallelLinear等nn.Module |
安装方法
README 给出的安装命令非常简单,且 flash-attention 的 training/README.md 在训练环境准备步骤里也引用了同样的命令:
cd csrc/fused_dense_lib && pip install .setup.py中,扩展通过CUDAExtension编译fused_dense.cpp与fused_dense_cuda.cu两个源文件,C++ 与 nvcc 都使用-O3优化,并且会调用append_nvcc_threads根据本机 CUDA 版本自动追加--threads(CUDA >= 11.2 时默认 4 线程)来加速编译。模块名为fused_dense_lib,安装后 Python 侧通过import fused_dense_lib使用。
二、三个核心 CUDA 入口:前向、权重梯度、反向融合
fused_dense.cpp通过PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)导出三个函数(csrc/fused_dense_lib/fused_dense.cpp#L209-L213):
| 导出函数 | 对应 CUDA 内核 | 作用 |
|---|---|---|
linear_act_forward | gemm_bias_act_lt | 融合的线性 + 激活(GELU/ReLU)前向 |
linear_bias_wgrad | gemm_bgradb_lt | 权重梯度 + bias 梯度(bias 梯度由 cuBLASLt 的 BGRADB epilogue 直接产出) |
bias_act_linear_dgrad_bgrad | gemm_dact_bgradb_lt | 融合的激活反向 + 输入梯度 + bias 梯度 |
2.1 前向:linear_act_forward(input, weight, bias, is_gelu, save_pre_act, heuristic)
前向本质是output = linear(input, weight, bias)之后立即施加 GELU 或 ReLU,关键点在于:
is_gelu:决定使用CUBLASLT_EPILOGUE_GELU系列还是CUBLASLT_EPILOGUE_RELU系列 epilogue;save_pre_act:是否保存激活前的pre_act张量,供反向复用,避免重算:- GELU 时
pre_act保存为与输入同 dtype 的原始值,形状为[batch, out_features]; - ReLU 时 cuBLASLt 只保存1 比特/元素 的位掩码(bit-mask),形状为
[batch, out_features / 8],dtype 为uint8,内存占用可忽略——这一点在 csrc/fused_dense_lib/fused_dense.cpp#L123-L125 有明确注释;
- GELU 时
heuristic:在 cuBLASLt 启发式返回的前 5 个算法候选中挑选第几个用于实际 matmul(heuristicResult[heuristic].algo,见 fused_dense_cuda.cu#L200-L227)。
值得注意的实现细节:代码里保留了注释// TD [2022-04-29] Somehow algo 0 and 2 are a lot slower than other algos,即开发者实测发现某些算法编号明显更慢,因此把"选哪个启发式算法"作为可配置参数暴露出来,而不是直接固定取第一个结果。
2.2 权重梯度:linear_bias_wgrad(input, d_output, has_d_bias)
反向阶段计算d_weight = d_output^T @ input与d_bias:
- CUDA >= 11.6 时走 cuBLASLt 路径
gemm_bgradb_lt,使用CUBLASLT_EPILOGUE_BGRADB在一次 GEMM 中同时产出d_weight与d_bias; - 若 cuBLASLt 路径失败(
status != 0),降级为普通cublasGemmEx计算d_weight,而d_bias在 CUDA < 11.6 时退化为d_output.view({-1, out_features}).sum(0)的 PyTorch 求和(fused_dense.cpp#L63-L67); has_d_bias为false时跳过d_bias分配,传入空指针,BGRADB epilogue 因此不启用。
2.3 激活反向:bias_act_linear_dgrad_bgrad(weight, d_output, pre_act, is_gelu, heuristic)
这一入口做的是"先过激活函数导数、再过第二层线性层"融合路径:d_input = (d_output @ weight^T) ⊙ act'(pre_act)并同时求出d_bias。它依赖前向保存的pre_act(GELU 存原始值,ReLU 存位掩码),使用的 epilogue 是CUBLASLT_EPILOGUE_DGELU_BGRAD或CUBLASLT_EPILOGUE_DRELU_BGRAD(fused_dense_cuda.cu#L462)。注释特别说明:cuBLASLt 的这个 epilogue 必须同时计算激活梯度与 bias 梯度,无法只算激活梯度,因此d_bias总是会被产出,只是调用方(如FusedMLPFunc.backward)在不需要时会丢弃它。
三、两种底层路径:cublasGemmEx 与 cuBLASLt epilogue 的取舍
fused_dense_cuda.cu的实现体现了"版本感知"的分层设计,核心逻辑受CUBLAS_VERSION宏控制:
gemm_bias(cublasGemmEx 路径,任意版本可用):为 fp16(CUDA_R_16F)和 bf16(CUDA_R_16BF)分别做了模板重载,computeType固定为CUDA_R_32F(FP32 累加),使用CUBLAS_GEMM_DEFAULT_TENSOR_OP。它只能做纯 matmul,无法把 bias / 激活融合进去,作为兼容性兜底。cuBLASLt 路径(
CUBLAS_VERSION >= 11600即 CUDA 11.6+ 启用):通过cublasLtMatmulDescInit+cublasLtMatmulDescSetAttribute配置完整的操作描述符,把 bias、pre_act、epilogue 类型全部作为属性注入,再以cublasLtMatmulAlgoGetHeuristic拿启发式算法并执行。由于 epilogue 在 GEMM 内部完成,避免了"GEMM 写出中间结果 → 读回 → 加 bias → 激活 → 再写回"的多轮显存读写。
这解释了 README 中"CUDA >= 11.8 才有最佳性能"的论断:虽然 11.6 起就有 cuBLASLt 融合 epilogue,但 bf16 的 matmul + bias + gelu 融合路径在 cuBLAS 11.8 中才达到成熟且高效的状态。若编译时 CUDA 低于 11.6,#if会直接裁剪掉三个融合内核,linear_act_forward_cuda与bias_act_linear_dgrad_bgrad_cuda直接返回失败码,由上层回退到未融合实现。
工作区内存(workspace)的分配策略
三个入口都遵循同一个工作区策略(fused_dense.cpp#L69-L73 等):
// 参考 PyTorch issue 73328,Apex 用 4M,TransformerEngine 在 Hopper 上用 32M、其他 GPU 用 4M size_t workspaceSize = 1024 * 1024 * (at::cuda::getCurrentDeviceProperties()->major >= 9 ? 32 : 4); auto lt_workspace = at::empty({static_cast<int64_t>(workspaceSize)}, opts.dtype(torch::kUInt8));即:计算能力 major >= 9(Hopper/H100 等)分配 32 MiB,其余 GPU(如 Ampere A100,major = 8)分配 4 MiB 的uint8工作区,通过CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES传给 cuBLASLt 作为算法选择的上限。这一分配策略在三个入口中完全一致,且与 PyTorch issue #73328 的结论对齐。
四、Python 侧封装:从裸算子到 nn.Module 与 Tensor Parallel
4.1 低层函数调用链
安装fused_dense_lib后,flash_attn/ops/fused_dense.py通过import fused_dense_lib as fused_dense_cuda引入原生算子(flash_attn/ops/fused_dense.py#L9),再包一层torch.autograd.Function(FusedDenseFunc、FusedMLPFunc)实现自动微分。以FusedDenseFunc为例,前向先用F.linear完成主 GEMM,反向则:
grad_weight, grad_bias = fused_dense_cuda.linear_bias_wgrad( total_x.reshape(batch_dim, total_x.shape[-1]), grad_output, ctx.needs_input_grad[2] )FusedMLPFunc则把两条 GEMM + 激活完整串起来(flash_attn/ops/fused_dense.py#L330-L335):
output1, *rest = fused_dense_cuda.linear_act_forward( total_x.reshape(batch_dim, n), weight1, bias1, is_gelu, save_pre_act, heuristic )反向时用bias_act_linear_dgrad_bgrad一步完成"激活导数 + 第二层权重梯度 + bias 梯度"(flash_attn/ops/fused_dense.py#L418-L420)。
fused_dense_func/fused_mlp_func是纯函数入口,内部做dtype 与设备资格检查(x.dtype in [torch.float16, torch.bfloat16],或 fp32 且开启了 autocast),不满足条件时自动回退到未融合的F.linear组合,保证功能正确性优先。
4.2 高层模块:FusedMLP 与 ParallelFusedMLP
flash_attn/ops/fused_dense.py在算子之上提供了 4 个可用的nn.Module:
| 类 | 用途 |
|---|---|
FusedDense | 直接替换nn.Linear,支持return_residual以便融合残差反向 |
FusedMLP | 单卡 MLP:fc1(Linear)+ GELU +fc2(Linear)全融合 |
ColumnParallelLinear/RowParallelLinear | Tensor Parallel 的列切 / 行切线性层 |
ParallelFusedMLP | 结合ColumnParallelLinear+RowParallelLinear的并行 MLP |
FusedMLP构造参数中,heuristic是理解性能的关键(flash_attn/ops/fused_dense.py#L555-L562 的 docstring 总结):
-1:不融合 GEMM + 激活,退化为独立内核,用torch.jit.fuser("fuser2")融合激活;0..4:在融合的 GEMM + 激活中使用该编号的 cuBLASLt 启发式算法;'auto'(默认):自动决策——- CUDA >= 11.8:fp16 与 bf16 均取
heuristic = 0(最佳性能); - CUDA <= 11.7:fp16 取
1,bf16 取-1(不融合,因为旧 cuBLAS 的 bf16 融合路径性能差); - H100(计算能力 9.0):fp16 与 bf16 均取
-1,实测融合 cuBLASLt 实现比未融合版本更慢。
- CUDA >= 11.8:fp16 与 bf16 均取
此外还提供checkpoint_lvl(0/1/2)三档反向重计算策略:0 不重算、1 反向重算gelu_out、2 重算pre_act与gelu_out,以"更慢的反向换取更少的内存驻留",便于在大模型训练中调节显存占用。注意 ReLU 的pre_act只是位掩码,所以即使checkpoint_lvl=1也会直接保存它而不重算(flash_attn/ops/fused_dense.py#L337-L339)。
4.3 在模型与训练脚本中的实际接线
flash_attn/modules/mlp.py在 import 时尝试引入FusedMLP/ParallelFusedMLP/ColumnParallelLinear/RowParallelLinear,未安装fused_dense_lib时置为None,并在使用处抛出ImportError("fused_dense is not installed")(flash_attn/modules/mlp.py#L70-L71)。
flash_attn/models/gpt.py的模型工厂会读取配置选择 MLP 实现(flash_attn/models/gpt.py#L219-L246):
if fused_mlp: if FusedMLP is None: raise ImportError("fused_dense is not installed") activation = ("gelu_approx" if config.activation_function in ["gelu_new", "gelu_fast", "gelu_approx", "gelu_pytorch_tanh"] else config.activation_function) mlp_cls = FusedMLP if process_group is None else ParallelFusedMLP即:单卡训练用FusedMLP,开启process_group的张量并行时自动切换为ParallelFusedMLP(其内部是ColumnParallelLinear+ 激活 +RowParallelLinear,并配套sequence_parallel的 all_gather / reduce_scatter 通信)。FusedMLP还被用于flash_attn/models/bert.py和flash_attn/models/vit.py。因此,fused_dense_lib虽小,却是整个 flash-attention 训练栈中 Dense/MLP 加速的关键依赖。
五、约束、兼容性与注意事项
综合 README 与源码,使用本扩展时需注意以下边界条件:
- dtype 只支持 fp16 与 bf16:
fused_dense.cpp中的DISPATCH_HALF_AND_BF16宏只分派这两个类型,其他 dtype 直接AT_ERROR;fp32 输入仅在开启 autocast(AMP)时会被提升后进入融合路径。 - 张量必须 CUDA 且连续(contiguous):三个入口都对
is_cuda、is_contiguous做了TORCH_CHECK,形状也有CHECK_SHAPE严格校验(如d_output必须为[batch, out_features])。 - 矩阵维度上限:Python 侧对
min(batch_dim, n, *weight.shape) > 65535 * 32抛错,即仅支持维度不超过约 2M 的矩阵("fused_dense only supports matrix dims <= 2M")。 - ReLU 的维度对齐要求:保存 pre_act 位掩码时,
dim_eligible要求最后一维能被 128 整除(ReLU)/ 8 整除(GELU),否则自动走未融合回退路径(flash_attn/ops/fused_dense.py#L494)。 - 多设备保护:三个入口都使用
at::cuda::CUDAGuard锁定输入所在设备,避免内核被错误地发射到cuda:0。 - 测试范围:README 明确说明仅在有 A100 的机器上验证过;代码中保留的算法速度注释(algo 0/2 较慢)也提示不同 GPU/驱动组合下启发式算法表现可能不同,
heuristic参数正是为此提供的调优旋钮。
六、小结
csrc/fused_dense_lib用不到 1000 行的 C++/CUDA 代码,把 Transformer 训练中最频繁的 Dense 计算路径(matmul + bias + GELU 的前向与反向)通过 cuBLASLt 的融合 epilogue 压进一次 GEMM,并率先补上了 Apex FusedDense 缺失的 bf16 支持。其价值不仅在于算子本身,更在于它支撑起flash_attn/ops/fused_dense.py中的FusedMLP、ColumnParallelLinear/RowParallelLinear等高层模块,成为 flash-attention 仓库训练 GPT/BERT/ViT 模型时 MLP 层加速与张量并行的基础设施。理解它的安装条件(CUDA >= 11.8 以获得 bf16 最佳性能)、三个原生入口的分工以及heuristic/checkpoint_lvl等调优参数,即可在自己的训练栈中安全、高效地复用它。
- 人工智能
- 大模型
- 算子库
【免费下载链接】flash-attention
Fast and memory-efficient exact attention
相关推荐
CANN ops-nn `aclnnFusedMatmulGelu` 融合算子接口详解:MatMul + Bias + GELU 的 NPU 融合计算与两阶段 aclnn 调用
CANN ops nn aclnnFusedMatmulGelu 融合算子接口详解:MatMul + Bias + GELU 的 NPU 融合计算与两阶段 ac
人工智能算子库深度学习CANNAscendCANN ops-nn 算子库 FusedMatmulGelu 融合算子:MatMul + 偏置 + GELU 的 NPU 加速实现与 aclnn 调用指南
CANN ops nn 算子库 FusedMatmulGelu 融合算子:MatMul + 偏置 + GELU 的 NPU 加速实现与 aclnn 调用指南 导
人工智能算子库深度学习CANNAscendFlash Attention 简易CUDA实现指南
Flash Attention 简易CUDA实现指南 项目介绍 Flash Attention in CUDA 是一个精简版的实现,旨在展示如何在大约100行C
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考