☰
基于 JAX/Flax 微调序列到序列摘要模型:Transformers `run_summarization_flax.py` 实战指南
2026/9/25 3:54:31 网站建设 项目流程
  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

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

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

本文以 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 模型类
bartFlaxBartForConditionalGeneration
blenderbotFlaxBlenderbotForConditionalGeneration
blenderbot-smallFlaxBlenderbotSmallForConditionalGeneration
encoder-decoderFlaxEncoderDecoderModel
longt5FlaxLongT5ForConditionalGeneration
marianFlaxMarianMTModel
mbartFlaxMBartForConditionalGeneration
mt5FlaxMT5ForConditionalGeneration
pegasusFlaxPegasusForConditionalGeneration
t5FlaxT5ForConditionalGeneration

也就是说,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_pathNone预训练 checkpoint 名称或路径;不设置则从头训练
--model_typeNone从头训练时指定模型类型(上述 10 种之一)
--config_nameNone与模型名不同的 config 名称/路径
--tokenizer_nameNone与模型名不同的分词器名称/路径
--cache_dirNone预训练模型下载缓存目录
--use_fast_tokenizerTrue是否使用 tokenizers 库的快速分词器
--dtypefloat32权重初始化与训练的浮点格式:float32/float16/bfloat16
--use_auth_tokenFalse加载私有模型时使用huggingface-cli login生成的 token

数据参数(DataTrainingArguments)

参数默认值说明
--dataset_nameNone使用 🤗 Datasets Hub 上的数据集名称
--dataset_config_nameNone数据集的 configuration 名称(如多语言数据集的语言子集)
--text_column/--summary_columnNone自定义数据集中"正文/摘要"列名;不指定则自动推断
--train_file/--validation_file/--test_fileNone本地训练/验证/测试文件,仅支持json或csv
--max_source_length1024源文本最大 token 数,超长截断、不足填充
--max_target_length128目标摘要最大 token 数
--val_max_target_lengthNone验证/预测时的目标长度,默认取max_target_length;同时会覆盖model.generate的max_length
--max_train_samples/--max_eval_samples/--max_predict_samplesNone调试用:截断样本数以加速
--preprocessing_num_workersNone预处理进程数
--source_prefixNone加在每条源文本前的前缀(T5 类模型常用)
--predict_with_generateFalse是否用generate计算生成式指标(ROUGE/BLEU)
--num_beamsNone评估时 beam 数,会传给model.generate,默认取模型 config
--overwrite_cacheFalse覆盖预处理缓存

训练参数(TrainingArguments)

参数默认值说明
--output_dir(必填)模型预测与 checkpoint 输出目录
--overwrite_output_dirFalse输出目录已存在且非空时是否覆盖;也可用于从 checkpoint 目录续训
--do_train/--do_eval/--do_predictFalse是否执行训练/评估/预测
--per_device_train_batch_size8每个 GPU/TPU core/CPU 上的训练 batch
--per_device_eval_batch_size8每个设备上的评估 batch
--learning_rate5e-5AdamW 初始学习率
--weight_decay0.0AdamW 权重衰减
--adam_beta1/--adam_beta2/--adam_epsilon0.9 / 0.999 / 1e-8AdamW 优化器超参
--label_smoothing_factor0.0标签平滑系数(0 表示不启用)
--adafactorFalse是否用 Adafactor 替代 AdamW
--num_train_epochs3.0总训练 epoch 数
--warmup_steps0线性预热步数
--logging_steps500每 N 步输出一次日志
--save_steps500每 N 步保存 checkpoint
--eval_stepsNone每 N 步做一次评估
--seed42随机种子
--push_to_hubFalse训练后是否上传模型到 Model Hub
--hub_model_idNoneHub 仓库全名(含用户名/组织名),如yourname/bart-base-xsum
--hub_tokenNone推送 Hub 用的 token
--gradient_checkpointingFalse梯度检查点:以更慢的反向传播换取显存节省

注意:总 batch size 是"每设备 batch × 设备数"(源码中train_batch_size = per_device_train_batch_size * jax.device_count()),脚本会自动使用检测到的全部 GPU/TPU core,分布式训练开箱即用。

数据加载与预处理:从 Hub 或本地 jsonlines/csv

脚本支持两条数据路径:

  1. Hub 数据集:指定--dataset_name(如xsum),通过load_dataset自动下载。
  2. 本地文件:指定--train_file/--validation_file/--test_file,扩展名必须是csv或json(源码中有断言校验),load_dataset按扩展名自动解析。

列名自动推断:源码维护了一张summarization_name_mapping字典(run_summarization_flax.py),为常见数据集预置了正文/摘要列名:

数据集正文列摘要列
cnn_dailymailarticlehighlights
xsumdocumentsummary
samsumdialoguesummary
big_patentdescriptionabstract
xgluenews_bodynews_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.

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

相关推荐

上一篇:终极系统备份与恢复指南:使用Rescuezilla保护你的数据安全
下一篇:5步掌握Steam API:打造个性化游戏数据解决方案

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

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

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

立即咨询