☰
FlexGen 仓库中的 Performer 高效注意力微调实战:基于 Flax/JAX 的掩码语言模型训练指南
2026/9/25 1:47:42 网站建设 项目流程
  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

Running large language models on a single GPU for throughput-oriented scenarios.

项目地址:https://gitcode.com/gh_mirrors/fl/FlexGen
点击查看免费下载

导读

本文围绕 FlexGen 仓库内置的 Hugging Face Transformers research_projects 示例,完整讲解如何基于Performer(FAVOR+ 快速注意力机制)在 Flax/JAX 生态下对 BERT 进行掩码语言建模(MLM)微调。你将掌握两个开箱即用的训练脚本(简单百科与完整英文维基百科)、全部核心命令行参数(--performer、--reinitialize、--wandb_user_name等),并通过阅读源码理解 Performer 用随机特征映射把注意力复杂度从二次降为线性的底层原理,从而能够在资源受限环境下训练超长序列模型。

一、项目背景:这份示例在 FlexGen 仓库中的位置

该 Performer 微调示例位于 benchmark/third_party/transformers/examples/research_projects/performer 目录下,是 FlexGen 仓库 benchmark 体系中以third_party方式内置的 Hugging Face Transformers 研究项目。它展示了与 FlexGen 主项目"在单 GPU 上以吞吐为导向运行大语言模型"主题高度互补的另一个维度:在训练/微调阶段通过高效注意力降低显存与算力消耗。

目录内共包含 6 个文件,构成一套完整的最小可运行研究项目:

文件作用
README.md项目说明、依赖、示例与关键参数
sanity_script.sh快速验证脚本:bert-base-cased + Simple Wikipedia
full_script.sh完整实验脚本:bert-large-cased + English Wikipedia
run_mlm_performer.py训练入口:参数解析、数据加载、MLM 数据整理器、训练/评估循环
modeling_flax_performer.pyFlax 版 Performer 模型定义(注意力替换为 FAVOR+)
modeling_flax_performer_utils.py快速注意力核心算法(随机特征映射、低秩分解)

二、环境依赖与安装前提

原文档明确列出运行时依赖为datasets、flax和jax,并说明wandb集成是内置的(可选启用)。结合 run_mlm_performer.py 的导入语句,实际依赖清单如下:

  • datasets:从 Hugging Face Hub 加载维基百科数据集(load_dataset);
  • flax与jax:模型定义、jax.pmap多设备并行、jax.lax底层算子全部基于 JAX 生态;
  • transformers:提供BertConfig、FlaxBertForMaskedLM、AutoTokenizer、HfArgumentParser、TrainingArguments等基础组件;
  • numpy、tqdm:数据处理与进度显示;
  • 可选tensorboard:脚本在启动时会调用is_tensorboard_available()检测,未安装则打印提示并降级为不记录指标;
  • 可选wandb:只有传入--wandb_user_name时才在运行时import wandb并初始化。

一个值得注意的细节是:脚本在模型加载阶段使用的是jnp.float32dtype,并没有默认开启混合精度;如果希望加速训练,可以关注TrainingArguments中的 fp16 相关开关。这保证了示例代码在任何 JAX 可用的 CPU/GPU/TPU 环境都能直接运行。

三、两个开箱即用的微调脚本

原文档提供了两个脚本,分别用于快速验证和完整实验,其命令均可直接复制运行。

3.1 快速验证脚本 sanity_script.sh

sanity_script.sh 从bert-base-cased检查点出发,在Simple Wikipedia 数据集(datasets 提供的一个小规模、用简单英语编写的维基百科子集)上微调:

TOKENIZERS_PARALLELISM=true python run_mlm_performer.py --output_dir experiments \ --dataset_name wikipedia --dataset_config_name 20200501.simple \ --model_name_or_path bert-base-cased --tokenizer_name bert-base-cased \ --do_train --overwrite_output_dir \ --per_device_train_batch_size 4 --learning_rate 5e-4 \ --warmup_steps 100 --num_train_epochs 3 --performer

3.2 完整实验脚本 full_script.sh

full_script.sh 使用更大的bert-large-cased检查点,在English Wikipedia 完整数据集上微调,适合真正验证 Performer 在长序列、大数据量下的效果:

TOKENIZERS_PARALLELISM=true python run_mlm_performer.py --output_dir experiments \ --dataset_name wikipedia --dataset_config_name 20200501.en \ --model_name_or_path bert-large-cased --tokenizer_name bert-large-cased \ --do_train --overwrite_output_dir \ --per_device_train_batch_size 4 --learning_rate 5e-4 \ --warmup_steps 100 --num_train_epochs 3 --performer

两个脚本的差异仅在--dataset_config_name(20200501.simplevs20200501.en)与--model_name_or_path/--tokenizer_name(base vs large)。命令开头的TOKENIZERS_PARALLELISM=true用于关闭 tokenizer 并行警告,属于 Transformers 生态的常见惯例。

四、核心命令行参数详解

原文档列出了五个关键参数,这里结合 run_mlm_performer.py 中的ModelArguments、DataTrainingArguments、WandbArguments三个 dataclass 逐一展开:

4.1--performer:启用 FAVOR+ 注意力(核心开关)

原文档说明"移除--performer参数即可使用标准 Bert 模型"。在源码中这个开关直接决定模型类别的选择:

lm_class = FlaxPerformerForMaskedLM if model_args.performer else FlaxBertForMaskedLM

即传--performer时加载 modeling_flax_performer.py 中的FlaxPerformerForMaskedLM,其注意力层会调用make_fast_softmax_attention构造的快速注意力函数;不传则回退为标准的FlaxBertForMaskedLM。这使你可以用同一套数据与训练逻辑,直接对比标准注意力与 Performer 的精度和吞吐差异——这正是研究型实验最需要的 A/B 能力。

4.2--reinitialize:从空白模型开始训练

原文档说明"添加--reinitialize将从空白模型(blank model)而非 Bert 检查点开始"。源码对应逻辑为:

if model_args.reinitialize: model = lm_class(config=BertConfig.from_pretrained(model_args.model_name_or_path)) else: model = lm_class.from_pretrained(model_args.model_name_or_path, ...)

传--reinitialize时,仅用model_name_or_path读取BertConfig的架构参数(层数、头数、隐层维度等),权重全部随机初始化;不传则加载预训练权重。这用于回答"Performer 在从头训练时的收敛性如何"这类研究问题。

4.3--model_name_or_path:更换 BERT 规模

原文档指出可将该参数换成 Hugging Face Hub 上任意预训练检查点来改变 BERT 规模。源码中该参数有双重用途:

  • 不传--reinitialize时,作为权重初始化的来源;
  • 同时通过BertConfig.from_pretrained(model_args.model_name_or_path)决定模型架构。

配套参数--tokenizer_name允许指定与模型不同的 tokenizer;若不传则回退使用--model_name_or_path对应的 tokenizer。

4.4--wandb_user_name:触发 Weights & Biases 日志

原文档说明"传入你的用户名将触发 wandb 日志记录"。源码在WandbArguments中定义了:

wandb_user_name: Optional[str] = field( default=None, metadata={"help": "The WandB user name for potential logging. If left None, no logging"}, ) wandb_project_name: Optional[str] = field( default="performer-experiments", metadata={"help": "The WandB project name for potential logging"}, )

训练循环中,每完成一个 batch 会记录Training loss,每完成一个 epoch 记录Eval loss;项目名默认为performer-experiments,可用--wandb_project_name覆盖。只要不传--wandb_user_name,wandb 完全不会被激活,因此默认运行无需登录 wandb 账号。

4.5--dataset_name与--dataset_config:选择数据集

原文档说明可通过这两个参数选择数据集,并建议使用 Hub 的数据集查看器辅助定位。源码的DataTrainingArguments中,dataset_name是数据集名称,dataset_config_name是配置名(如维基百科的语言/时间快照版本)。加载逻辑为:

datasets = load_dataset(data_args.dataset_name, data_args.dataset_config_name) if "validation" not in datasets.keys(): datasets["validation"] = load_dataset(..., split=f"train[:{data_args.validation_split_percentage}%]") datasets["train"] = load_dataset(..., split=f"train[{data_args.validation_split_percentage}%:]")

当数据集本身没有 validation 划分时,脚本会自动按--validation_split_percentage(默认 5)从训练集切分出验证集。此外,脚本同样支持本地数据文件:传入--train_file/--validation_file(支持 csv、json、txt 三种格式)即可脱离 Hub 训练,脚本会取名为text的列或第一列作为语料。

4.6 更多可调参数(从源码补充)

  • --mlm_probability:掩码概率,默认 0.15;
  • --max_seq_length:最大序列长度,默认取模型最大输入长度。这是 Performer 相对标准 BERT 优势最明显的场景——标准注意力在序列长度增长时显存呈平方增长,而 Performer 可支撑更长的序列;
  • --pad_to_max_length:是否将所有样本填充到max_seq_length,默认 False(动态按 batch 内最大长度填充);
  • --overwrite_cache/--preprocessing_num_workers:控制数据集预处理缓存与并行进程数;
  • --use_fast_tokenizer:是否使用 fast tokenizer,默认 True;
  • 其余训练参数(batch size、学习率、warmup、epochs、seed 等)全部继承自 Transformers 的TrainingArguments。

五、源码级拆解:训练脚本如何工作

run_mlm_performer.py 完整实现了一个 Flax 版 MLM 训练器,主流程如下:

  1. 参数解析:用HfArgumentParser((ModelArguments, DataTrainingArguments, TrainingArguments, WandbArguments))解析四组参数;若命令行只传一个以.json结尾的参数,则按 JSON 配置文件解析;
  2. 数据集准备:如上节所述,从 Hub 或本地文件加载并切分训练/验证集,然后 tokenize(return_special_tokens_mask=True、截断、可选填充);
  3. 模型加载:按--performer与--reinitialize选择模型类与初始化方式;
  4. 数据整理器(FlaxDataCollatorForLanguageModeling):负责动态 padding 与掩码生成,掩码策略遵循经典 MLM 约定——80% 替换为[MASK]、10% 替换为随机词、10% 保持原词,且只对被掩码的 token 计算 loss(未掩码 token 的标签设为 -100);
  5. 优化器与学习率调度:使用 Flax 的Adam优化器,并实现了一个因子可组合的学习率调度器create_learning_rate_scheduler,支持constant、linear_warmup、rsqrt_decay、rsqrt_normalized_decay、decay_every、cosine_decay六种因子,默认组合为"constant * linear_warmup * rsqrt_decay"(即学习率先线性热身,再按步数平方根倒数衰减,warmup 步数取max(warmup_steps, 1));
  6. 并行训练:jax.pmap将training_step/eval_step映射到所有本地设备,jax_utils.replicate复制优化器参数,训练时通过jax.lax.pmean做跨设备梯度平均——这意味着脚本天然支持多 GPU/TPU 数据并行;
  7. 指标记录:每个 epoch 计算 loss 与 accuracy,保存 TensorBoard 标量(若可用),可选同步到 wandb。

六、Performer 模型实现:注意力如何从 O(n²) 降到 O(n)

6.1 模型结构:BERT 骨架 + 快速注意力插槽

modeling_flax_performer.py 完整复刻了 BERT 的组件层次:FlaxPerformerEmbeddings(词/位置/类型三路 embedding 求和 + LayerNorm)→FlaxPerformerEncoder(N 层FlaxPerformerLayer)→ MLM 头。与标准 BERT 的唯一结构性差异在注意力层:

class FlaxPerformerAttention(nn.Module): num_heads: int head_size: int @nn.compact def __call__(self, hidden_state, attention_mask): single_head_dim = self.head_size // self.num_heads fast_softmax_attention = make_fast_softmax_attention(qkv_dim=single_head_dim) self_att = nn.attention.SelfAttention( num_heads=self.num_heads, qkv_features=self.head_size, name="self", attention_fn=fast_softmax_attention )(hidden_state, attention_mask) layer_norm = FlaxPerformerLayerNorm(name="layer_norm")(self_att + hidden_state) return layer_norm

关键在nn.attention.SelfAttention的attention_fn参数——它把 Flax 原生自注意力的点积注意力函数替换成了make_fast_softmax_attention返回的快速注意力函数,从而让"换注意力机制"变成了一次函数注入,其余 BERT 结构(FFN、残差、LayerNorm、pooler)完全复用。

6.2 权重迁移:PyTorch 检查点如何进入 Flax

脚本允许从 PyTorch 版 BERT 检查点初始化 Flax 权重,modeling_flax_performer.py 中的convert_from_pytorch静态方法负责这一映射,处理了四大类差异:

  • 全连接层:PyTorch 的dense.weight→ Flax 的dense.kernel;
  • 注意力头分解:query/key/value的权重按num_attention_heads重塑并转置,以匹配 FlaxSelfAttention的头分解存储;
  • 层归一化:LayerNorm.weight/bias→layer_norm.gamma/beta;
  • 参数转置:intermediate.dense.kernel、output.dense.kernel、pooler.dense.kernel等需要转置,attention.output.dense与attention.output.LayerNorm则要消除一层嵌套。

6.3 快速注意力核心算法:随机特征映射与低秩分解

modeling_flax_performer_utils.py 的头部注释说明该文件复制自 Google Research 的fast_self_attention.py,核心思路是利用结构化随机特征映射(RFM)技术对注意力矩阵做低秩分解,从而近似快速 softmax 注意力。

make_fast_softmax_attention是构造入口,关键可调参数包括:

  • nb_features=256:随机特征数量,特征越多近似越精确、计算开销越大;
  • ortho_features=True:默认使用高斯正交随机矩阵(GaussianOrthogonalRandomMatrix)而非非结构化高斯矩阵(GaussianUnstructuredRandomMatrix),正交矩阵可降低近似方差;
  • nonnegative_features=True:默认使用非负 softmax 核特征(nonnegative_softmax_kernel_feature_creator),即用exp(投影 - 范数项) + eps构造非负特征来近似 softmax;若置 False 则改用 sin/cos 特征(sincos_softmax_kernel_feature_creator);
  • redraw_features=True:每个注意力调用根据 query 的和重新抽取投影矩阵(保证注意力是 permutation equivariant);
  • renormalize_attention=True:对结果做重归一化以匹配 softmax 的归一化特性;
  • numerical_stabilizer=0.000001:数值稳定项。

在FastAttentionviaLowRankDecomposition.dot_product_attention中,算法将序列维度上的O(L²)注意力矩阵乘法重构为两个低秩步骤:先算Z = (K')ᵀV(key 特征与 value 的缩并),再算W = Q'Z(query 特征与 Z 的缩并),配合R = Q'(K')ᵀ1计算归一化因子。这样每个位置的计算量与序列长度 L 呈线性关系而非平方关系——这正是 Performer 支持超长序列、降低显存占用的根本原因。归一化后的输出还会经过jnp.reciprocal与数值稳定化处理,保证训练稳定性。

另外,工具模块还提供了make_fast_generalized_attention,可将 softmax 注意力推广到任意核函数(如jax.nn.relu),支持ortho、iid、deterministic三种特征类型,为扩展实验留有余地。

七、运行与验证建议

  1. 先跑 sanity_script.sh:Simple Wikipedia 规模小,在单卡上几分钟内即可完成 3 个 epoch,用于验证环境(jax/flax/datasets 版本兼容、模型加载、wandb 可选集成)是否就绪;
  2. 对比基线:去掉--performer跑一遍同一命令,即可获得标准FlaxBertForMaskedLM的 loss/accuracy 曲线作为对照;
  3. 验证长序列能力:增大--max_seq_length,观察 Performer 在序列变长时的显存增长是否明显慢于标准注意力;
  4. 尝试从头训练:加--reinitialize验证随机初始化下 Performer 的收敛行为;
  5. 监控指标:不加--wandb_user_name时关注终端输出的逐 epochLoss/Acc;若机器装有 TensorBoard,指标会写入--output_dir/logs目录。

需要说明的是:本项目属于 research_projects 研究示例,追求的是算法原理验证与快速迭代,而非生产级训练框架;在 FlexGen 仓库上下文中,它适合作为评估高效注意力方案对后续推理吞吐影响的实验前置环节。

八、总结

这份 Performer 微调示例虽然目录不大,却是一条完整的"研究闭环":sanity_script.sh/full_script.sh提供可直接运行的实验入口,run_mlm_performer.py提供与标准 BERT 无缝切换的训练框架,modeling_flax_performer.py用函数注入的方式将 FAVOR+ 快速注意力嵌入 BERT 骨架,modeling_flax_performer_utils.py则落地了随机特征映射与低秩分解的核心数学。对希望在有限显存下训练长序列 Transformer 的开发者而言,这份示例既是可复用的 Flax 训练脚手架,也是理解高效注意力机制最直接的源码教材。

  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

Running large language models on a single GPU for throughput-oriented scenarios.

项目地址:https://gitcode.com/gh_mirrors/fl/FlexGen
点击查看免费下载

相关推荐

上一篇:ThinkPad X230黑苹果完美教程:轻松实现macOS体验
下一篇:BG3ModManager:5步解决博德之门3模组管理难题

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

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

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

立即咨询