- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
本文以 Transformers Flax 摘要微调示例为骨架,完整讲解如何用run_summarization_flax.py在 GPU/TPU 上对 BART、T5、Pegasus 等序列到序列(Seq2Seq)模型进行摘要任务的端到端微调。你将掌握完整的训练/评估/预测命令行用法、三类参数(模型、数据、训练)的默认值与语义,并从源码层面理解 JAX/Flax 的函数式训练循环、分布式pmap并行、ROUGE 指标计算与 Model Hub 推送机制,可直接复用到自己的摘要乃至其他 Seq2Seq 任务中。
该文档位于 benchmark/third_party/transformers/examples/flax/summarization/README.md,配套的完整训练脚本为 run_summarization_flax.py。该目录属于当前仓库 benchmark/third_party 下随附的 Transformers 框架源码,作为开源参考实现随仓库分发。
为什么用 JAX/Flax 做摘要微调
摘要(Summarization)是典型的序列到序列任务:输入长文档、输出短摘要。JAX/Flax 技术栈在这一场景下的核心优势体现在:
- JAX暴露 NumPy 风格的 API,并具备强大的转换能力:
jit可以把纯函数 trace 后编译为 GPU/TPU 上高效、融合的加速代码;grad求任意梯度、pmap多设备并行、remat梯度检查点、vmap自动向量化、pjit自动分片的模型并行,且这些转换可以任意组合。 - Flax在 JAX 之上提供基于 dataclass 的模块抽象,代码简洁显式;其"lifted"转换(如
vmap、remat)允许任意嵌套。 - 函数式与不可变性:JAX/Flax 模型不可变,以纯函数方式更新,天然适合
pmap级别的简单高效模型并行。
需要特别说明的是,这一时期的 Flax 示例没有 Trainer 抽象,所有训练循环都是显式写在脚本中的(参见 examples/flax/README.md),因此阅读本脚本的源码也是学习 JAX/Flax 训练范式的最佳入口之一。
环境准备与依赖安装
脚本运行依赖的最小集合记录在 requirements.txt:
datasets >= 1.1.3 jax>=0.2.8 jaxlib>=0.1.59 flax>=0.3.5 optax>=0.0.8 evaluate>=0.2.0若还要运行仓库附带的示例测试(test_flax_examples.py),需要额外安装 _tests_requirements.txt 中列出的pytest、nltk、rouge-score、seqeval、tensorboard、conllu等包。
安装 JAX 本身需按运行环境区分:
- GPU:JAX 的 pip 安装与 CUDA/CuDNN 版本强相关,需根据本机 CUDA 版本选择对应的
jaxlib安装方式。 - TPU:JAX/Flax 官方示例均以在 Cloud TPU 上高效运行为设计目标,多设备并行开箱即用。
另外脚本在启动时会自动检测并下载 NLTK 的punkt分词数据(用于 ROUGE 计算前的句子切分),离线环境(设置了TRANSFORMERS_OFFLINE)下会直接抛出提示,要求先联网完成下载。
支持哪些模型架构
脚本通过FlaxAutoModelForSeq2SeqLM自动加载模型。从源码映射表 modeling_flax_auto.py(FLAX_MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING_NAMES)可以确认本脚本支持以下 Seq2Seq 架构:
| model_type | 对应 Flax 模型类 |
|---|---|
bart | FlaxBartForConditionalGeneration |
blenderbot | FlaxBlenderbotForConditionalGeneration |
blenderbot-small | FlaxBlenderbotSmallForConditionalGeneration |
encoder-decoder | FlaxEncoderDecoderModel |
longt5 | FlaxLongT5ForConditionalGeneration |
marian | FlaxMarianMTModel |
mbart | FlaxMBartForConditionalGeneration |
mt5 | FlaxMT5ForConditionalGeneration |
pegasus | FlaxPegasusForConditionalGeneration |
t5 | FlaxT5ForConditionalGeneration |
也就是说,README 命令示例中的 BART 只是其中之一,你也可以直接替换为t5-small、google/pegasus-xsum等 checkpoint。模型加载流程为AutoConfig.from_pretrained→AutoTokenizer.from_pretrained→FlaxAutoModelForSeq2SeqLM.from_pretrained(或from_config从头训练),权重精度由--dtype控制(float32/float16/bfloat16,默认float32)。
训练命令与预期结果(README 核心示例)
README 给出的完整训练命令如下,可以直接复制运行:
python run_summarization_flax.py \ --output_dir ./bart-base-xsum \ --model_name_or_path facebook/bart-base \ --tokenizer_name facebook/bart-base \ --dataset_name="xsum" \ --do_train --do_eval --do_predict --predict_with_generate \ --num_train_epochs 6 \ --learning_rate 5e-5 --warmup_steps 0 \ --per_device_train_batch_size 64 \ --per_device_eval_batch_size 64 \ --overwrite_output_dir \ --max_source_length 512 --max_target_length 64 \ --push_to_hub按 README 记载,该配置在 6 个 epoch 后约37 分钟完成训练,得到验证集 loss1.7785、ROUGE217.01(训练统计可在 TensorBoard.dev 上查看)。需要说明的是:
- 这里使用的是
generate的默认生成参数;README 特别提醒,若针对xsum数据集的特点(如合适的 beam 数、最小生成长度等)显式配置生成参数,ROUGE 分数可以进一步提升。 - 上述耗时与指标与当时的硬件环境、依赖版本相关,在不同机器上复现时会有差异,不应视为固定基准。
三类命令行参数详解(对应源码 dataclass)
脚本使用HfArgumentParser解析三类参数(源码中分别定义为ModelArguments、DataTrainingArguments、TrainingArguments三个 dataclass)。除命令行外,还支持把唯一参数写成 JSON 配置文件路径,脚本会自动调用parse_json_file读取。
模型参数(ModelArguments)
| 参数 | 默认值 | 说明 |
|---|---|---|
--model_name_or_path | None | 预训练 checkpoint 名称或路径;不设置则从头训练 |
--model_type | None | 从头训练时指定模型类型(上述 10 种之一) |
--config_name | None | 与模型名不同的 config 名称/路径 |
--tokenizer_name | None | 与模型名不同的分词器名称/路径 |
--cache_dir | None | 预训练模型下载缓存目录 |
--use_fast_tokenizer | True | 是否使用 tokenizers 库的快速分词器 |
--dtype | float32 | 权重初始化与训练的浮点格式:float32/float16/bfloat16 |
--use_auth_token | False | 加载私有模型时使用huggingface-cli login生成的 token |
数据参数(DataTrainingArguments)
| 参数 | 默认值 | 说明 |
|---|---|---|
--dataset_name | None | 使用 🤗 Datasets Hub 上的数据集名称 |
--dataset_config_name | None | 数据集的 configuration 名称(如多语言数据集的语言子集) |
--text_column/--summary_column | None | 自定义数据集中"正文/摘要"列名;不指定则自动推断 |
--train_file/--validation_file/--test_file | None | 本地训练/验证/测试文件,仅支持json或csv |
--max_source_length | 1024 | 源文本最大 token 数,超长截断、不足填充 |
--max_target_length | 128 | 目标摘要最大 token 数 |
--val_max_target_length | None | 验证/预测时的目标长度,默认取max_target_length;同时会覆盖model.generate的max_length |
--max_train_samples/--max_eval_samples/--max_predict_samples | None | 调试用:截断样本数以加速 |
--preprocessing_num_workers | None | 预处理进程数 |
--source_prefix | None | 加在每条源文本前的前缀(T5 类模型常用) |
--predict_with_generate | False | 是否用generate计算生成式指标(ROUGE/BLEU) |
--num_beams | None | 评估时 beam 数,会传给model.generate,默认取模型 config |
--overwrite_cache | False | 覆盖预处理缓存 |
训练参数(TrainingArguments)
| 参数 | 默认值 | 说明 |
|---|---|---|
--output_dir | (必填) | 模型预测与 checkpoint 输出目录 |
--overwrite_output_dir | False | 输出目录已存在且非空时是否覆盖;也可用于从 checkpoint 目录续训 |
--do_train/--do_eval/--do_predict | False | 是否执行训练/评估/预测 |
--per_device_train_batch_size | 8 | 每个 GPU/TPU core/CPU 上的训练 batch |
--per_device_eval_batch_size | 8 | 每个设备上的评估 batch |
--learning_rate | 5e-5 | AdamW 初始学习率 |
--weight_decay | 0.0 | AdamW 权重衰减 |
--adam_beta1/--adam_beta2/--adam_epsilon | 0.9 / 0.999 / 1e-8 | AdamW 优化器超参 |
--label_smoothing_factor | 0.0 | 标签平滑系数(0 表示不启用) |
--adafactor | False | 是否用 Adafactor 替代 AdamW |
--num_train_epochs | 3.0 | 总训练 epoch 数 |
--warmup_steps | 0 | 线性预热步数 |
--logging_steps | 500 | 每 N 步输出一次日志 |
--save_steps | 500 | 每 N 步保存 checkpoint |
--eval_steps | None | 每 N 步做一次评估 |
--seed | 42 | 随机种子 |
--push_to_hub | False | 训练后是否上传模型到 Model Hub |
--hub_model_id | None | Hub 仓库全名(含用户名/组织名),如yourname/bart-base-xsum |
--hub_token | None | 推送 Hub 用的 token |
--gradient_checkpointing | False | 梯度检查点:以更慢的反向传播换取显存节省 |
注意:总 batch size 是"每设备 batch × 设备数"(源码中train_batch_size = per_device_train_batch_size * jax.device_count()),脚本会自动使用检测到的全部 GPU/TPU core,分布式训练开箱即用。
数据加载与预处理:从 Hub 或本地 jsonlines/csv
脚本支持两条数据路径:
- Hub 数据集:指定
--dataset_name(如xsum),通过load_dataset自动下载。 - 本地文件:指定
--train_file/--validation_file/--test_file,扩展名必须是csv或json(源码中有断言校验),load_dataset按扩展名自动解析。
列名自动推断:源码维护了一张summarization_name_mapping字典(run_summarization_flax.py),为常见数据集预置了正文/摘要列名:
| 数据集 | 正文列 | 摘要列 |
|---|---|---|
cnn_dailymail | article | highlights |
xsum | document | summary |
samsum | dialogue | summary |
big_patent | description | abstract |
xglue | news_body | news_title |
orange_sum/pn_summary/psc/thaisum/wiki_summary/amazon_reviews_multi | 见源码映射 | 见源码映射 |
若未命中映射或使用自定义文件,默认取数据集第一列为正文、第二列为摘要;也可用--text_column/--summary_column显式指定。
预处理细节(源码preprocess_function):由于 jitted 函数需要固定长度输入,预处理对源/目标统一使用padding="max_length"补齐到max_source_length/max_target_length并截断。针对 Flax 模型不接受labels的特性,脚本通过模型模块的shift_tokens_right函数把目标序列右移一位生成decoder_input_ids,同时保留decoder_attention_mask用于在损失中屏蔽 pad token。
训练循环内部实现:JAX/Flax 微调原理拆解
阅读 run_summarization_flax.py 的训练部分,可以完整还原 JAX/Flax 微调的典型范式:
- 训练状态:自定义
TrainState(继承 Flaxtrain_state.TrainState)额外携带dropout_rng,replicate()将参数复制到所有设备并把 dropout 随机键按设备分片(shard_prng_key)。 - 学习率调度:
create_learning_rate_fn用optax.linear_schedule构造"线性预热 + 线性衰减",再通过optax.join_schedules拼接,总步数 =数据集大小 // 总batch × epoch数。 - 权重衰减掩码:
decay_mask_fn遍历参数树,识别名称中含layernorm/layer_norm/ln的 LayerNorm 参数及所有bias,对它们不施加权重衰减——这是 AdamW 的标准最佳实践。 - 优化器:
optax.adamw配合上述学习率调度与衰减掩码。 - 损失函数:
loss_fn实现了带标签平滑的交叉熵(one-hot 软化标签,confidence = 1 - label_smoothing_factor),并用decoder_attention_mask屏蔽 padding token,最后对非 pad 位置求均值归一化。 - 梯度更新:
train_step内用jax.value_and_grad(compute_loss, has_aux=True)求梯度,用jax.lax.psum跨设备规约损失与样本数,再归一化后apply_gradients更新状态。 - 多设备并行:
jax.pmap(..., "batch")把 train/eval/generate 步编译成 SPMD 并行程序;训练数据经shard分发到各设备,评估与生成使用pad_shard_unpad(自动补齐不完整 batch、卸载结果)。这就是 JAX"纯函数 + 任意组合转换"特性的直接体现。 - 数据加载:
data_loader用jax.random.permutation生成随机 batch 索引,drop_last=True时跳过不完整的末尾 batch(评估时则保留)。
评估与生成:ROUGE 指标计算流程
当指定--predict_with_generate时,评估/预测循环除了计算 loss,还会执行生成并计算 ROUGE:
- 生成参数:
gen_kwargs由val_max_target_length(或模型 config 的max_length)与num_beams(或模型 config 的num_beams)构成,通过model.generate(batch["input_ids"], attention_mask=..., **gen_kwargs)批量解码。 - 指标计算:
compute_metrics先用 tokenizer 解码预测与标签(skip_special_tokens=True),再用nltk.sent_tokenize把每段文本按句切分换行(ROUGE-LSum 要求句子间换行),最后evaluate.load("rouge").compute(..., use_stemmer=True)得到各 ROUGE 分数并乘以 100,同时记录平均生成长度gen_len。 - 输出:预测阶段结束后,主进程把指标写入
output_dir/test_results.json(键名如test_rouge1、test_rouge2、test_rougeL、test_rougeLsum)。
日志、Checkpoint 与 Model Hub 推送
- TensorBoard:主进程(
jax.process_index() == 0)用 FlaxSummaryWriter写入训练/评估标量(train_loss、eval_loss、各 ROUGE 值等),write_metric负责落盘。 - Checkpoint:每个 epoch 结束后,主进程从
state.params取出首设备副本,调用model.save_pretrained(output_dir, params=params)与tokenizer.save_pretrained(output_dir)保存模型与分词器。 - Hub 推送:
--push_to_hub会基于output_dir目录名(或--hub_model_id指定的全名)创建/克隆远端仓库,每个 epoch 以Saving weights and logs of epoch N为提交信息异步推送。使用前需要本地登录(huggingface-cli login)或通过--hub_token传入认证 token。
仓库如何验证该示例可运行
test_flax_examples.py 中的test_run_summarization用t5-small配合tests/fixtures/tests_samples/xsum/sample.json小样本做冒烟测试:num_train_epochs=3、warmup_steps=8、learning_rate=2e-4、per_device_train_batch_size=2、per_device_eval_batch_size=1,并断言test_rouge1 >= 10、test_rouge2 >= 2、test_rougeL >= 7、test_rougeLsum >= 7。这既验证了脚本端到端可运行,也给出了"小数据快速验证"的参数参考——先用极小样本跑通流程,再切换到完整数据集。
自定义数据集:jsonlines/csv 快速上手
如果你有自己的摘要数据,可按以下方式组织:
- jsonlines:每行一个 JSON 对象,包含正文与摘要两个字段,如
{"document": "...", "summary": "..."}。 - csv:第一列为正文、第二列为摘要(或通过
--text_column/--summary_column指定列名)。
启动示例:
python run_summarization_flax.py \ --model_name_or_path facebook/bart-base \ --train_file ./data/train.json \ --validation_file ./data/val.json \ --test_file ./data/test.json \ --text_column document --summary_column summary \ --do_train --do_eval --do_predict --predict_with_generate \ --num_train_epochs 3 \ --per_device_train_batch_size 16 \ --output_dir ./my-bart-summarizer \ --overwrite_output_dir运行注意事项小结
output_dir已存在且非空、同时启用了--do_train且未加--overwrite_output_dir时会直接报错终止,这是防止误覆盖已有 checkpoint 的保护机制。- 模型 config 必须正确设置
decoder_start_token_id,否则脚本会在启动时抛错。 max_source_length/max_target_length决定显存占用与训练速度;摘要任务常见配置为源 512~1024、目标 64~128(README 示例即 512/64)。- 若目标是把模型发布到 Hub,建议显式指定
--hub_model_id,保证仓库名符合用户名/模型名规范。 - 文中涉及的文件均为仓库只读参考实现,运行与配置均在本地完成,无需修改仓库内容。
- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
相关推荐
Transformers 中的 PEGASUS-X:面向长文本摘要的序列到序列模型实战指南
Transformers 中的 PEGASUS X:面向长文本摘要的序列到序列模型实战指南 PEGASUS X 是 Google 在 2022 年发布、并已完整
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 摘要生成实战指南:基于 T5 微调 BillSum 法律文本摘要模型
Transformers 摘要生成实战指南:基于 T5 微调 BillSum 法律文本摘要模型 摘要生成(Summarization)是 🤗 Transfor
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 中的越南语大规模序列到序列模型:BARTpho 架构、分词原理与文本摘要实战
Transformers 中的越南语大规模序列到序列模型:BARTpho 架构、分词原理与文本摘要实战 BARTpho 是面向越南语(Vietnamese)的大
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考