☰
flash-attention 仓库 fused_dense_lib 深度指南:融合 Matmul+Bias+GELU 的 CUDA 扩展实现与使用
2026/9/30 2:27:47 网站建设 项目流程
  • 人工智能
  • 大模型
  • 算子库

【免费下载链接】flash-attention

Fast and memory-efficient exact attention

项目地址:https://gitcode.com/GitHub_Trending/fl/flash-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.cppPyTorch 绑定层:参数检查、张量分配、dispatch、pybind11 导出
csrc/fused_dense_lib/fused_dense_cuda.cuCUDA 实现层:基于 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_forwardgemm_bias_act_lt融合的线性 + 激活(GELU/ReLU)前向
linear_bias_wgradgemm_bgradb_lt权重梯度 + bias 梯度(bias 梯度由 cuBLASLt 的 BGRADB epilogue 直接产出)
bias_act_linear_dgrad_bgradgemm_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 有明确注释;
  • 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宏控制:

  1. gemm_bias(cublasGemmEx 路径,任意版本可用):为 fp16(CUDA_R_16F)和 bf16(CUDA_R_16BF)分别做了模板重载,computeType固定为CUDA_R_32F(FP32 累加),使用CUBLAS_GEMM_DEFAULT_TENSOR_OP。它只能做纯 matmul,无法把 bias / 激活融合进去,作为兼容性兜底。

  2. 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/RowParallelLinearTensor 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 实现比未融合版本更慢。

此外还提供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 与源码,使用本扩展时需注意以下边界条件:

  1. dtype 只支持 fp16 与 bf16:fused_dense.cpp中的DISPATCH_HALF_AND_BF16宏只分派这两个类型,其他 dtype 直接AT_ERROR;fp32 输入仅在开启 autocast(AMP)时会被提升后进入融合路径。
  2. 张量必须 CUDA 且连续(contiguous):三个入口都对is_cuda、is_contiguous做了TORCH_CHECK,形状也有CHECK_SHAPE严格校验(如d_output必须为[batch, out_features])。
  3. 矩阵维度上限:Python 侧对min(batch_dim, n, *weight.shape) > 65535 * 32抛错,即仅支持维度不超过约 2M 的矩阵("fused_dense only supports matrix dims <= 2M")。
  4. ReLU 的维度对齐要求:保存 pre_act 位掩码时,dim_eligible要求最后一维能被 128 整除(ReLU)/ 8 整除(GELU),否则自动走未融合回退路径(flash_attn/ops/fused_dense.py#L494)。
  5. 多设备保护:三个入口都使用at::cuda::CUDAGuard锁定输入所在设备,避免内核被错误地发射到cuda:0。
  6. 测试范围: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

项目地址:https://gitcode.com/GitHub_Trending/fl/flash-attention
点击查看免费下载
上一篇:CNN 可视化工具怎么用?10 分钟上手交互式卷积神经网络学习神器
下一篇:单人档玩腻了?《骑马与砍杀2》多人联机 BannerlordCoop 快速开黑指南

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询