flash-attention 仓库 fused_dense_lib 深度指南:融合 Matmul+Bias+GELU 的 CUDA 扩展实现与使用
发布时间:2026/9/30 2:27:49来源:尧图网络
人工智能大模型算子库【免费下载链接】flash-attentionFast 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 层。其 READMEcsrc/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 版本自动追加--threadsCUDA 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系列 epiloguesave_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 个算法候选中挑选第几个用于实际 matmulheuristicResult[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_biasCUDA 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-L67has_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_actGELU 存原始值ReLU 存位掩码使用的 epilogue 是CUBLASLT_EPILOGUE_DGELU_BGRAD或CUBLASLT_EPILOGUE_DRELU_BGRADfused_dense_cuda.cu#L462。注释特别说明cuBLASLt 的这个 epilogue 必须同时计算激活梯度与 bias 梯度无法只算激活梯度因此d_bias总是会被产出只是调用方如FusedMLPFunc.backward在不需要时会丢弃它。三、两种底层路径cublasGemmEx 与 cuBLASLt epilogue 的取舍fused_dense_cuda.cu的实现体现了版本感知的分层设计核心逻辑受CUBLAS_VERSION宏控制gemm_biascublasGemmEx 路径任意版本可用为 fp16CUDA_R_16F和 bf16CUDA_R_16BF分别做了模板重载computeType固定为CUDA_R_32FFP32 累加使用CUBLAS_GEMM_DEFAULT_TENSOR_OP。它只能做纯 matmul无法把 bias / 激活融合进去作为兼容性兜底。cuBLASLt 路径CUBLAS_VERSION 11600即 CUDA 11.6 启用通过cublasLtMatmulDescInitcublasLtMatmulDescSetAttribute配置完整的操作描述符把 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 73328Apex 用 4MTransformerEngine 在 Hopper 上用 32M、其他 GPU 用 4M size_t workspaceSize 1024 * 1024 * (at::cuda::getCurrentDeviceProperties()-major 9 ? 32 : 4); auto lt_workspace at::empty({static_castint64_t(workspaceSize)}, opts.dtype(torch::kUInt8));即计算能力 major 9Hopper/H100 等分配 32 MiB其余 GPU如 Ampere A100major 8分配 4 MiB 的uint8工作区通过CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES传给 cuBLASLt 作为算法选择的上限。这一分配策略在三个入口中完全一致且与 PyTorch issue #73328 的结论对齐。四、Python 侧封装从裸算子到 nn.Module 与 Tensor Parallel4.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.FunctionFusedDenseFunc、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-L335output1, *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 与 ParallelFusedMLPflash_attn/ops/fused_dense.py在算子之上提供了 4 个可用的nn.Module类用途FusedDense直接替换nn.Linear支持return_residual以便融合残差反向FusedMLP单卡 MLPfc1Linear GELU fc2Linear全融合ColumnParallelLinear/RowParallelLinearTensor Parallel 的列切 / 行切线性层ParallelFusedMLP结合ColumnParallelLinearRowParallelLinear的并行 MLPFusedMLP构造参数中heuristic是理解性能的关键flash_attn/ops/fused_dense.py#L555-L562 的 docstring 总结-1不融合 GEMM 激活退化为独立内核用torch.jit.fuser(fuser2)融合激活0..4在融合的 GEMM 激活中使用该编号的 cuBLASLt 启发式算法auto默认自动决策——CUDA 11.8fp16 与 bf16 均取heuristic 0最佳性能CUDA 11.7fp16 取1bf16 取-1不融合因为旧 cuBLAS 的 bf16 融合路径性能差H100计算能力 9.0fp16 与 bf16 均取-1实测融合 cuBLASLt 实现比未融合版本更慢。此外还提供checkpoint_lvl0/1/2三档反向重计算策略0 不重算、1 反向重算gelu_out、2 重算pre_act与gelu_out以更慢的反向换取更少的内存驻留便于在大模型训练中调节显存占用。注意 ReLU 的pre_act只是位掩码所以即使checkpoint_lvl1也会直接保存它而不重算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-L246if 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 与 bf16fused_dense.cpp中的DISPATCH_HALF_AND_BF16宏只分派这两个类型其他 dtype 直接AT_ERRORfp32 输入仅在开启 autocastAMP时会被提升后进入融合路径。张量必须 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-attentionFast and memory-efficient exact attention项目地址https://gitcode.com/GitHub_Trending/fl/flash-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上一篇CNN 可视化工具怎么用10 分钟上手交互式卷积神经网络学习神器下一篇单人档玩腻了《骑马与砍杀2》多人联机 BannerlordCoop 快速开黑指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网