新闻详情

新闻详情

首页 / 资讯中心 / 详情

PaddleNLP FlashMask 灵活注意力掩码完全指南:列式稀疏表示与长序列训练加速实战

发布时间:2026/9/25 6:01:02来源:尧图网络
PaddleNLP FlashMask 灵活注意力掩码完全指南:列式稀疏表示与长序列训练加速实战
人工智能大模型预训练微调LoRARLHF强化学习分布式训练【免费下载链接】PaddleNLPEasy-to-use and powerful LLM and SLM library with awesome model zoo.项目地址https://gitcode.com/gh_mirrors/pa/PaddleNLP点击查看免费下载本文聚焦 PaddleNLP 中集成的高性能注意力掩码技术 FlashMask系统讲解其解决长序列大模型训练中掩码冗余计算与 O(N²) 存储瓶颈的核心原理——列式稀疏掩码表示、FlashAttention-2 扩展与块级跳过计算并结合仓库中 Llama 系列 SFT、LoRA、DPO、RM 的完整配置与运行命令给出可直接复现的端到端训练加速实战方案。读完本文你将掌握 FlashMask 的算法设计、性能收益边界以及在 PaddleNLP 中开启并调优该特性的完整路径。1. 背景稠密注意力掩码是长序列训练的瓶颈在 Transformer 类大模型训练中注意力Attention机制需要确定哪些 Query-Key token 之间执行有效计算业界通常借助二维稠密注意力掩码Attention Mask完成这一约束。然而稠密掩码带来了两个层面的代价冗余计算大量被 mask 的无效 token 间注意力仍然会被实际计算存储压力掩码的空间复杂度为 O(N²)N 为序列长度在长序列训练场景下会形成巨大的显存占用成为高效训练的阻碍。业界已有 Memory Efficient AttentionMEA与 FlashAttention 等加速方案但它们支持的掩码类型较为有限。如图 1 所示FlashAttention 通常只能支持纯因果掩码Causal、滑动窗口掩码Sliding Window、因果文档掩码Causal Document Mask和文档掩码Document Mask等固定形式而真实训练任务SFT、DPO、RM、多模态等中的掩码形态丰富多变现有技术难以满足灵活性的要求。2. FlashMask 核心创新列式稀疏掩码表示2.1 关键洞察掩码区域的列向连续性FlashMask 的核心洞察在于大模型常见注意力掩码模式中Query-Key token 的掩码模式具有连续性——对于每一个 Key token被置为无效计算的 Query token 是相邻排列的。换句话说二维掩码矩阵中每一列上被 mask 的灰色区域沿列方向连续分布。基于这一观察FlashMask 将二维稠密掩码矩阵压缩为一维的行索引区间从而以紧凑形式表达掩码并显著降低存储需求可公式化表示为$$M_{j} [start_j, end_j), \quad \forall j \in {1, \ldots, N}$$其中 N 为 Key 的序列长度$M_j$ 为二维稠密掩码矩阵的第 j 列$[start_j, end_j)$ 为连续的行索引区间表示 $start_j$ 到 $end_j - 1$ 的连续 Query token 被 mask 掉置为无效 Attention 计算。2.2 列式稀疏表示四个一维向量为高效处理因果与双向注意力场景中的复杂掩码FlashMask 以对角线为区分使用四个一维向量表示掩码LTSLower Triangular Start下三角起始行索引LTELower Triangular End下三角结束行索引UTSUpper Triangular Start上三角起始行索引UTEUpper Triangular End上三角结束行索引。下三角被 mask 的行索引区间记为 $[LTS, LTE)$上三角记为 $[UTS, UTE)$。以 16 个 Query token 与 16 个 Key token 的因果注意力掩码为例图 2灰色单元格为 mask 区域仅用 LTS、LTE 两个向量即可完整表达col_idx0123456789101112131415$LTS$135556699912121216161616$LTE$15141415121211111616161616161616以第 1 列col_idx0为例开始 mask 的行号为 13结束 mask 的行号为 15开区间表示位置 13、14 的 Query token 不与位置 0 的 Key token 做有效 Attention 计算。对于图 1 中的各类掩码FlashMask 均可用该列式稀疏表示完整表达。其中$-$的空缺表示在不同场景下有不同默认值LTS 与 UTS 的默认值为 0mask 区域默认从第 0 行开始LTE 与 UTE 的默认值为 Query 序列长度mask 区域默认结束于最后一行。这种表示将存储复杂度从 O(N²) 降低到O(N)是长序列高效训练的基础。3. 扩展 FlashAttention块级分类与实时跳过计算FlashMask 将列式掩码表示集成进FlashAttention-2算法其高性能 Kernel 实现分为两个阶段预处理与实时块跳过计算。在 FlashAttention 的 Kernel 中得分矩阵score matrix按 Tile Block 分块计算如图 4 的简化示意整个得分矩阵被分为 4×4 块每个块含 4 个 Query token × 4 个 Key token 的交互。FlashMask 的原始输入是 token 级别的逐列表示经预处理转化为块级别表示供实时阶段快速判断每个块的类型。3.1 预处理阶段生成 8 个分块极值向量预处理阶段先将列式稀疏掩码向量 LTS、LTE、UTS、UTE 加载到高带宽存储HBM再依据 FlashAttention 的分块列大小将列式向量分块并计算每个分块内所有列的向量最大值与最小值生成 8 个中间向量$LTStart^{min}$、$LTStart^{max}$$LTEnd^{min}$、$LTEnd^{max}$$UTStart^{min}$、$UTStart^{max}$$UTEnd^{min}$、$UTEnd^{max}$以图 4 最左侧 4 个分块为例分块含 4 列$LTS[13,5,5,5]$、$LTE[15,14,14,15]$因此 $LTStart^{min}min(LTS)5$、$LTStart^{max}max(LTS)13$、$LTEnd^{min}min(LTE)14$、$LTEnd^{max}max(LTE)15$其余分块的极值计算结果如图 5 所示。3.2 实时块跳过计算阶段三种块类型实时计算阶段利用预处理得到的极值向量对得分矩阵的每个分块分类以提升计算效率。分类依据为以下三种类型完全掩码块若 $BlockRow_{min} \geq Start^{max}$ 且 $BlockRow_{max} \leq End^{min}$则该块所有元素均被掩码计算可直接跳过部分掩码块若 $BlockRow_{min} End^{max}$ 且 $BlockRow_{max} Start^{min}$则该块部分元素被掩码需要逐元素掩码计算未掩码块其余情况块内所有元素均未被掩码可简化计算不做额外掩码操作。图 4 展示了因果掩码场景下使用 LTS、LTE 进行 Kernel 计算的完整过程三类块的判定实例如下完全掩码块图 4 中 [3,2] 位置的块最小行号 12 ≥ $LTStart^{max}12$最大行号 15 ≤ $LTEnd^{max}16$块内全部被掩码计算直接跳过部分掩码块图 4 中 [1,1] 位置的块最小行号 4 $LTEnd^{max}12$最大行号 7 $LTStart^{min}6$需逐元素掩码计算未掩码块图 4 中 [3,1] 位置的块最小行号 12 ≥ $LTEnd^{max}12$块内全部有效无需额外掩码操作。算法 1原论文 [3] 中的伪代码给出了 FlashMask 扩展 FlashAttention-2 的完整前向计算流程其中浅蓝色阴影部分为 FlashMask 新增的计算步骤。3.3 效率提升与精度保证FlashMask 充分利用掩码稀疏性通过跳过完全掩码块减少计算开销同时不改变算法精度与使用稠密掩码矩阵的注意力计算保持比特级别的数值等效性确保精度无损。3.4 仓库中的 Kernel 实现佐证从源码结构看FlashMask 在仓库中以自定义算子的形式集成反向算子实现在 ops/csrc/paddle_bwd_ops/flashmask_attn_bwd.cc通过PD_BUILD_OP(flashmask_attn_bwd)注册输入为q/k/v、startend_row_indices、out、softmax_lse、seed_offset、out_grad输出 q/k/v 三者的梯度并带dropout与causal两个属性。其中startend_row_indices正是列式稀疏掩码的区间表示在算子接口层面的落地形态——它在数据侧被构建为每行一个区间终点的紧凑向量详见下文第 5 节Kernel 内部据此完成块级分类与跳过计算。4. FlashMask 优势速度与存储的双重提升4.1 端到端训练吞吐量提升在 Llama-2 7B、13B、70B 三种规模下针对 SFT、LoRA、DPO、RM 四种下游训练场景与不同序列长度的实验表明相比基于稠密掩码矩阵的方法FlashMask 实现1.65 倍至 3.22 倍的端到端吞吐量提升图 6并显著降低训练峰值显存消耗图 7、支持更长的序列长度。在 Llama2 7B 上 FlashMask 对比 FlashAttentionCausalTrue的显存消耗数据见表 2单位 GB。4.2 端到端训练收敛验证在 Llama 3.1 模型上的实验验证 FlashMask对收敛精度无影响。作为一种精确算法通过控制计算过程的随机性如 FlashAttention 反向 Query 梯度计算使用 atomicAdd 操作FlashMask 可以与使用稠密掩码的 FlashAttention 在比特级别精确对齐。图 8 展示了 Llama3.1 8B 在 SFT、LoRA、DPO、RM 四个下游任务上端到端训练 Loss 与稠密掩码基线完全一致。4.3 稀疏度与 Kernel 时延的线性关系FlashMask 利用注意力掩码的块稀疏性跳过完全掩码块的计算将计算复杂度降低到 $O((1 - ρ)T_rT_c)$其中 ρ 表示块稀疏度。针对因果文档掩码、共享问题掩码、文档掩码三种掩码类型、多组稀疏度的实验图 9表明Kernel 执行时延与稀疏度呈线性关系稀疏度越高FlashMask 计算速度越快。4.4 Kernel 性能对比优于 FlexAttention与 PyTorch 使用编译器技术支持注意力掩码的FlexAttention相比在各类常见掩码模式下 FlashMask 展现了更高的计算效率在 TFLOPs/s 指标上FlashMask 比 FlexAttention高出 12.1% 至 60.7%并在 A100 GPU 上实现37.8% 至 62.3%的理论峰值计算性能图 10FlexAttention 使用 PyTorch 2.6.0.dev20240920cu124测试平台为 A100-SXM 80G。5. FlashMask 应用场景赋能大语言模型5.1 大语言模型下游训练加速FlashMask 可广泛应用于大语言模型的下游训练如 SFT、LoRA、DPO、RM 等。尤其在DPO 与 RM训练中数据由问题和回答对组成多个答案可共享同一个问题从而大幅减少对问题 token 的冗余计算稀疏收益尤为显著。5.2 支持单向/双向混合注意力掩码模式FlashMask 同时支持因果掩码单向注意力与文档掩码双向注意力因此能灵活覆盖混合注意力场景例如全局 滑动窗口掩码兼顾全局上下文与局部细节FlashMask 可高效处理该混合掩码前缀语言模型前缀部分关注全部 token其余部分使用因果掩码如 T5 预训练FlashMask 可同时支持两种注意力模式提升训练与推理效率。5.3 支持多模态图文混合多分辨率训练针对多模态数据的不同分辨率与掩码策略FlashMask 可通过不同注意力模式灵活适配其长序列处理能力有助于模型学习不同模态数据之间的关联例如在图文匹配任务中更有效地对齐图像与文本的关键信息。FlashMask 开源代码已在 PaddlePaddle 与 PaddleNLP 平台发布支持超过千亿参数规模的模型以及超过128K tokens的上下文长度。6. 快速开始在 PaddleNLP 中开启 FlashMask6.1 环境依赖Python 3.8PaddlePaddle 3.0.0b0若未安装请参照飞桨官网安装指引完成安装通过以下命令安装最新 develop 分支代码pip install --pre --upgrade paddlenlp -f https://www.paddlepaddle.org.cn/whl/paddlenlp.html6.2 数据准备SFT LoRAPaddleNLP 精调数据为每行一个字典的 json 文件字典字段如下srcstr, List(str)模型的输入指令instruction、提示prompt即模型应执行的任务tgtstr, List(str)模型的输出。样例数据单轮多轮对话均可src与tgt按对话轮次一一对应{ src: [Show me the most compelling argument for the existence of God from a theists perspective and then contrast that with the most compelling argument from an atheists perspective. 1 / 1, The most compelling argument for the existence of God ..., Please cite your sources for these.1 / 1, Sure! Here are the sources for the arguments I presented: ...], tgt: [The most compelling argument for the existence of God from a theists perspective is the cosmological argument, ..., Please cite your sources for these.1 / 1, Sure! Here are the sources for the arguments I presented: ..., Why are these arguments considered the most compelling?1 / 1] }为便于测试也可直接使用 tulu-v2-sft-mixture 数据集mkdir data wget https://paddlenlp.bj.bcebos.com/datasets/examples/tulu.jsonl mv tulu.jsonl data/train.json6.3 SFT 与 LoRA 启动SFT 启动命令python -u -m paddle.distributed.launch --gpus 0,1,2,3,4,5,6,7 run_finetune.py ./config/llama/flashmask/sft.jsonLoRA 启动命令python -u -m paddle.distributed.launch --gpus 0,1,2,3,4,5,6,7 run_finetune.py ./config/llama/flashmask/lora.json6.4 配置参数详解以仓库真实配置为准仓库 llm/config/llama/flashmask/ 目录下提供了开箱即用的 SFT、LoRA、DPO 配置文件。以 SFT 配置 sft.json 为例核心参数如下参数值说明model_name_or_pathmeta-llama/Meta-Llama-3.1-8B-Instruct基座模型dataset_name_or_path./data训练数据目录output_dir./checkpoints/llama_sft_flashmask输出目录per_device_train_batch_size1每卡 batch sizegradient_accumulation_steps4梯度累积步数max_steps12000最大训练步数learning_rate2e-05学习率warmup_ratio0.03warmup 比例src_length/max_length8000 / 8192输入长度 / 最大序列长度bf16truebfloat16 混合精度fp16_opt_levelO2混合精度优化级别tensor_parallel_degree2张量并行度pipeline_parallel_degree1流水线并行度shardingstage2参数分片策略LoRA 配置使用 stage1use_flash_attentiontrue开启 FlashAttentionflash_masktrue开启 FlashMask 灵活掩码核心开关zero_padding/greedy_zero_paddingtrue / true零填充与 FlashMask 配合recomputefalse关闭重计算以释放显存用于更长序列lazytrue惰性初始化LoRA 配置 lora.json 与 SFT 结构一致差异在于learning_rate提升至 2e-04、sharding使用 stage1、增加lora: true与benchmark: truebenchmark 模式便于吞吐测试。从源码看flash_mask参数在 paddlenlp/trl/model_config.py 中定义flash_mask: bool field(defaultFalse, ...)默认关闭数据侧在 paddlenlp/trl/trl_data.py 中根据该开关决定构造列式稀疏掩码向量attn_mask_startend_row_indices每个 token 记录一个区间终点O(N) 存储还是构造 O(N²) 的稠密attention_masknp.tri三角矩阵直观印证了前文的存储收益。6.5 数据准备DPO RMDPO/RM 精调数据同样为每行一个字典的 json 文件字段如下srcstr, List(str)用户对话内容tgtstr, List(str)系统回复内容responsestr, List(str)包含 chosen 与 rejected 回复sortList(int)用于区分 chosen 与 rejectedsort 值小的是 rejectedsort 值大的是 chosen。样例数据{ src: [In this task, you are given a second sentence. Your task is to generate the first sentence ...], tgt: [], response: [Could you provide some context ..., As an AI assistant, its essential to generate the first sentence ...], sort: [1, 0] }为便于测试可直接使用 ultrafeedback_binarized 数据集mkdir dpo_data wget https://paddlenlp.bj.bcebos.com/datasets/examples/ultrafeedback.jsonl mv ultrafeedback.jsonl dpo_data/6.6 DPO 与 RM 启动DPO 启动命令python -u -m paddle.distributed.launch --gpus 0,1,2,3,4,5,6,7 ./alignment/dpo/run_dpo.py ./config/llama/flashmask/dpo.jsonRM 启动命令python -u -m paddle.distributed.launch --gpus 0,1,2,3,4,5,6,7 ./alignment/rm/flashmask/run_reward.py ./config/llama/flashmask/rm.jsonDPO 配置 dpo.json 的关键差异参数train_dataset_path指向dpo_data/ultrafeedback.jsonl、learning_rate为 5e-07、max_seq_len/max_prompt_len为 8192/8000、tensor_parallel_degree为 4。需要说明仓库当前 llm/config/llama/flashmask/ 目录下已提供 sft.json、lora.json、dpo.json 三个配置RM 训练入口脚本为 llm/alignment/rm/run_reward.py其源码run_reward.py会校验flash_mask必须与 zero padding 和 flash attention 配合使用提示实际运行时需按该约束组织配置。6.7 FlashMask 与 ZeroPadding 的配合机制从数据流水线看FlashMask 与零填充ZeroPadding深度绑定在 paddlenlp/datasets/zero_padding_dataset.py 中ZeroPadding类将attn_mask_startend_row_indices列为受支持输入键拼接 batch 时会对每行记录的区间终点累加序列偏移attn_mask_startend_row_indices [i sequence_sum for i in record[...]]把多个样本的列式掩码无缝拼接成跨样本的全局区间表示同时若该键存在则不再构造稠密 attention_mask对应分支会生成 O(N²) 的块对角矩阵。这与 DPO 数据构造trl_data.py 中 prompt 区间、双 response 区间的attn_mask_startend_row_indices生成逻辑共同构成了DPO/RM 共享问题、掩码多区间这一典型场景的完整闭环也正是第 5.1 节所述稀疏收益的工程基础。7. 小结FlashMask 以列式稀疏 块级分类跳过的设计在不损失数值精度比特级等效的前提下将注意力掩码的存储复杂度从 O(N²) 降到 O(N)为长序列128K 上下文大模型训练提供了显著加速端到端吞吐提升 1.653.22 倍、Kernel 峰值算力利用率达 37.8%62.3%、较 FlexAttention 快 12.1%60.7%。在 PaddleNLP 中只需在配置中同时开启flash_mask、use_flash_attention与zero_padding配合仓库自带的 sft.json、lora.json、dpo.json 即可复现全流程加速。8. 参考文献[1] Self-attention Does Not Need O(n^2) MemoryMemory Efficient Attention 原始论文[2] FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning[3] FlashMask: Efficient and Rich Mask Extension of FlashAttentionFlashMask 原始论文[4] FlexAttention: The Flexibility of PyTorch with the Performance of FlashAttentionPyTorch 官方博客进一步阅读本文对应的英文版本位于 docs/en/llm/docs/flashmask.mdFlashMask 的 RL/对齐训练细节可参考 llm/docs/rlhf.md 与 llm/docs/alignment_tutorial.mdRM 训练入口见 llm/alignment/rm/run_reward.py。赞分享人工智能大模型预训练微调LoRARLHF强化学习分布式训练【免费下载链接】PaddleNLPEasy-to-use and powerful LLM and SLM library with awesome model zoo.项目地址https://gitcode.com/gh_mirrors/pa/PaddleNLP点击查看免费下载相关推荐Label Studio OCR 发票标注3 个控件搭出 BIO 预标注流水线一次跑通Label Studio OCR 发票标注3 个控件搭出 BIO 预标注流水线一次跑通 OCR 发票标注这条链路上OCR 模型跑完并不能直接喂给 NER 训人工智能大模型预训练微调LoRARLHF强化学习分布式训练模型推理服务推理引擎模型量化模型压缩本地部署NLPPaddleNLP FlashMask 实战与原理用列式稀疏掩码把大模型长序列训练提速 1.65x~3.22xPaddleNLP FlashMask 实战与原理用列式稀疏掩码把大模型长序列训练提速 1.65x~3.22x 导读 FlashMask 是 PaddlePa人工智能大模型预训练微调LoRARLHF强化学习分布式训练模型推理服务推理引擎模型量化模型压缩本地部署NLPPaddleNLP FlashMask深度解析列式稀疏掩码如何让长序列大模型训练提速3.2倍PaddleNLP FlashMask深度解析列式稀疏掩码如何让长序列大模型训练提速3.2倍 PaddleNLP 内置的 FlashMask 灵活注意力掩码人工智能大模型预训练微调LoRARLHF强化学习分布式训练模型推理服务推理引擎模型量化模型压缩本地部署NLP上一篇CompileFlow终极指南阿里巴巴高性能流程编排引擎如何提升业务效率10倍下一篇微信聊天记录导出终极指南WeChatMsg 免费开源把对话永久留痕创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

WS2812B灯环实战:Arduino流水灯与彩虹灯环完整指南 2026/9/25 6:29:11

WS2812B灯环实战:Arduino流水灯与彩虹灯环完整指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
Allegro到立创EDA专业版:PCB导入转换的实操指南 2026/9/25 6:29:11

Allegro到立创EDA专业版:PCB导入转换的实操指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
中科蓝讯LB2002定制RV32工具链深度解析 2026/9/25 6:29:11

中科蓝讯LB2002定制RV32工具链深度解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
Jetson Orin Nano Super深度评测:249美元边缘AI工作站实战解析 2026/9/25 6:29:11

Jetson Orin Nano Super深度评测:249美元边缘AI工作站实战解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
微信小程序课堂签到系统:四种签到模式与SSM后台实现 2026/9/25 6:29:11

微信小程序课堂签到系统:四种签到模式与SSM后台实现

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
计及光伏逆变器快速无功响应的分布式电源优化配置 2026/9/25 6:29:05

计及光伏逆变器快速无功响应的分布式电源优化配置

1. 为什么传统分布式电源配置方案会“浪费”光伏的无功价值?先抛个实际场景。你在做一个园区或县域配电网的分布式电源规划,手里拿着几个光伏电站的接入申请,业主要求在满足电压质量的前提下尽量多装容量。于是你按常规思路建了个优化模型&am…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉