首页
/ Transformers 训练示例脚本实战指南:以摘要任务为例,覆盖安装、分布式训练、TPU 与 Accelerate

Transformers 训练示例脚本实战指南:以摘要任务为例,覆盖安装、分布式训练、TPU 与 Accelerate

2026-09-04 14:33:28作者:江焘钦

本文基于 Transformers 官方文档中“使用脚本训练”(run scripts)章节编写,围绕 examples/pytorch/summarization/run_summarization.py 这一参考实现,完整讲解示例脚本的定位、从源码安装环境的步骤、单卡/多卡/TPU/🤗 Accelerate 四种运行方式、自定义 CSV/JSONL 数据集的接法、小样本快速调试与断点续训、以及训练后一键推送模型到 Hub 的方法。读完本文,你可以直接复制本文中的命令在自己机器上复现 T5-small 在 CNN/DailyMail 上的微调流程,并知道每个命令行参数在源码中的对应位置。

一、示例脚本在 Transformers 中的定位

Transformers 提供两种训练入口:交互式的 notebooks 与面向真实训练的示例脚本(run scripts)。示例脚本展示了如何用 PyTorch 完成“下载数据 → 预处理 → 用 Trainer 微调 → 评估”的完整闭环,是官方推荐给需要脱离 notebook 场景的用户的写法。

需要区分三类脚本的来源与维护状态(文档原文明确说明):

  1. 主仓库 examples/ 下的脚本:与当前开发版本同步维护,是本文的主角;
  2. examples/research_projects/ 下的脚本:对应研究项目,见 研究项目说明,同样主要由社区贡献;
  3. 旧版本脚本(legacy):文档提示这些脚本“未获活跃支持,可能要求特定版本的 Transformers,且很可能与最新库不兼容”。

文档同时给出两个重要预期管理:

  • 示例脚本不保证开箱即用地解决你的所有问题,你需要按自己的任务修改脚本;绝大多数脚本刻意保留了完整的“数据预处理”代码段,方便你就地编辑;
  • 如果你想在脚本中新增功能,官方建议先在社区论坛或 issue 中讨论,以牺牲可读性为代价换取更多功能的 PR 通常不会被合并,但修 bug 的 PR 受欢迎。

从当前仓库结构看,examples/ 目录保留了 PyTorch 示例量化示例modular-transformers研究项目调度器示例;文档中同时提到的 TensorFlow 与 JAX/Flax 示例目录已不在当前主干中,若需要它们,应按下文“切换旧版本”一节 checkout 到对应 tag(这也是文档把旧版本清单挂在同一页面的原因)。

二、环境准备:从源码安装 Transformers

文档明确要求:为了让最新版示例脚本跑通,必须从源码安装 Transformers,并建议放在一个干净的虚拟环境中。当前仓库主干即对应这一方式:

git clone https://gitcode.com/GitHub_Trending/tra/transformers
cd transformers
pip install .

安装完成后,进入你要运行的示例目录,安装该目录自带的依赖。以摘要任务为例,其依赖清单为 examples/pytorch/summarization/requirements.txt

accelerate >= 0.12.0
datasets >= 1.8.0
sentencepiece != 0.1.92
protobuf
rouge-score
nltk
py7zr
torch >= 1.3
evaluate

可以看到摘要任务额外需要 rouge-scorenltk(脚本内部用 rouge 指标评估、用 nltk.sent_tokenize 做后处理,见下文源码分析)。安装命令为:

pip install -r requirements.txt

切换旧版本示例脚本

文档列出了从 v1.0.0 到 v4.5.1 的旧版本示例(v4.5.1、v4.4.2、v4.3.3、v4.2.2、v4.1.1、v4.0.1、v3.5.1、v3.4.0、v3.3.1、v3.2.0、v3.1.0、v3.0.2、v2.11.0 … v1.0.0 等 tag)。使用方式是先按上面的方式安装,再把仓库切到目标 tag:

git checkout tags/v3.5.1

之后再 pip install .,即可获得与该 tag 匹配的库版本与脚本版本。

脚本自身的版本自检

run_summarization.py 开头就做了两道版本关卡,解释了为什么必须从源码安装:

# Will error if the minimal version of Transformers is not installed. Remove at your own risks.
check_min_version("4.57.0.dev0")

require_version("datasets>=1.8.0", "To fix: pip install -r examples/pytorch/summarization/requirements.txt")

也就是说,如果你的 pip 安装的是过旧的正式版 Transformers,脚本会在启动时直接报错并提示如何修复,这为“环境与脚本不匹配”的问题提供了最早的诊断点。

三、运行摘要训练脚本:T5-small + CNN/DailyMail

文档以文本摘要任务为贯穿全文的例子。该脚本的工作流程是(原文档逐条列出):

  1. 从 🤗 Datasets 库下载并预处理数据集;
  2. Trainer 把一个支持“序列到序列”架构的预训练模型微调为摘要模型;
  3. 官方示例把 google-t5/t5-small 微调在 cnn_dailymail 数据集上;
  4. 由于 T5 的训练方式,需要额外提供 source_prefix 参数,告诉 T5“这是一篇摘要任务”,官方取值是 "summarize: "
python examples/pytorch/summarization/run_summarization.py \
    --model_name_or_path google-t5/t5-small \
    --do_train \
    --do_eval \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --predict_with_generate

参数在源码中的落点

脚本用 HfArgumentParser 把参数分成三组数据类(run_summarization.py#L331-L337):ModelArguments(模型相关)、DataTrainingArguments(数据相关)、Seq2SeqTrainingArguments(训练相关,定义于 training_args.py)。上表命令中涉及的关键参数及其源码语义如下:

参数 定义位置/默认值 作用
model_name_or_path ModelArguments,必填 预训练模型名或本地路径
dataset_name / dataset_config DataTrainingArguments 从 Hub 加载数据集的名称与配置版本
source_prefix L273-L275,默认 None 拼接到每条源文本前缀,T5 系列必给 "summarize: "
max_source_length L192-L200,默认 1024 输入侧分词后最大长度,超长截断
max_target_length L201-L209,默认 128 目标侧(摘要)最大长度
num_beams 默认 1(即贪心解码),传入 model.generate 评估时束搜索宽度
predict_with_generate 来自 Seq2SeqTrainingArguments 评估/预测时走生成而非 teacher forcing
per_device_train_batch_size TrainingArguments 每设备批大小

脚本还会主动校验 T5 的使用方式:若模型是 google-t5/t5-smallt5-baset5-larget5-3bt5-11b 之一却未传 source_prefix,会打印警告(L364-L374)——这是文档强调“T5 需要 source_prefix”这一条的代码证据。

数据集的输入/输出列也无需手动指定:脚本内置了常用摘要数据集的列名映射 summarization_name_mappingL310-L323),例如 cnn_dailymail 自动对应 article/highlightsxsum 对应 document/summary;未命中映射时默认取数据集第 0 列做文本、第 1 列做摘要(L517-L534)。

评估指标由脚本内置:通过 evaluate.load("rouge") 计算 ROUGE(L627-L657),--predict_with_generate 时才启用该 compute_metrics。训练结束后还会自动创建 model card(含 finetuned_fromtasks: summarization 等元信息,L740-L755)。

支持的模型架构

摘要任务 READMErun_summarization.py 支持这些架构:BartForConditionalGenerationMBartForConditionalGenerationMarianMTModelPegasusForConditionalGenerationT5ForConditionalGenerationMT5ForConditionalGeneration(以及仅翻译的 FSMTForConditionalGeneration)。脚本内部统一通过 AutoModelForSeq2SeqLM 加载(L437-L445)。

四、分布式训练与混合精度

文档指出 Trainer 原生支持分布式与混合精度,脚本可直接继承这两个能力,只需两处改动:

  • 追加 --fp16 开启混合精度;
  • torchrun--nproc_per_node 指定 GPU 数量。
torchrun \
    --nproc_per_node 8 examples/pytorch/summarization/run_summarization.py \
    --fp16 \
    --model_name_or_path google-t5/t5-small \
    --do_train \
    --do_eval \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --predict_with_generate

源码中有两处与分布式细节相关的设计,值得留意:

  1. 数据与模型加载都包在 main_process_first / from_pretrained 的文件锁语义里,保证分布式下只有一个进程下载(脚本注释 L417-L445);
  2. fp16 开启时,DataCollatorForSeq2Seqpad_to_multiple_of=8L618-L625),这是 Tensor Core 对齐的常见处理。

文档补充:TensorFlow 版本脚本依赖 MirroredStrategy,多 GPU 默认自动生效、无需额外参数(适用于旧版本中仍含 TF 示例的 tag)。

五、在 TPU 上运行

文档说明 PyTorch 通过 XLA 支持 TPU。使用方式为运行仓库自带的 xla_spawn.py 启动器,并用 --num_cores 指定 TPU 核数(1 或 8):

python xla_spawn.py --num_cores 8 \
    summarization/run_summarization.py \
    --model_name_or_path google-t5/t5-small \
    --do_train \
    --do_eval \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --predict_with_generate

从源码看,这套机制的实现链路很短:xla_spawn.py 把训练脚本当作模块导入,重写 sys.argv 后调用 xmp.spawn(mod._mp_fn, args=(), nprocs=args.num_cores)xla_spawn.py#L66-L78);而训练脚本末尾必须提供 _mp_fn 入口(run_summarization.py#L760-L762):

def _mp_fn(index):
    # For xla_spawn (TPUs)
    main()

这正是“为什么 TPU 训练脚本要额外定义 _mp_fn”的原因——它是 xmp.spawn 约定的多进程入口。

六、使用 🤗 Accelerate 运行

文档介绍的另一条路线是 PyTorch 专用的 Accelerate 库:它以统一方式支持“纯 CPU、单/多 GPU、TPU”多种硬件,同时保留对 PyTorch 训练循环的完全掌控。

文档给出的注意事项:由于 Accelerate 迭代很快,需安装其 Git 版才能运行这些脚本;requirements.txt 中也已把 accelerate >= 0.12.0 列为依赖。

使用 Accelerate 时,脚本要换成目录内的 run_summarization_no_trainer.py——文档的判别规则是:被 Accelerate 支持的脚本会以 *_no_trainer.py 命名(暴露裸训练循环,方便你直接改优化器、DataLoader 等,但可调参数比 Trainer 版少)。三步走:

# 1. 交互式生成并保存配置
accelerate config

# 2. 自检配置是否可用
accelerate test

# 3. 启动训练
accelerate launch run_summarization_no_trainer.py \
    --model_name_or_path google-t5/t5-small \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --output_dir ~/tmp/tst-summarization

这条 accelerate launch 命令对纯 CPU、单 GPU、单机/多机多卡、TPU 通用,无需改动脚本本身。

七、使用自定义数据集

文档明确:run_summarization.py 支持自定义数据集,前提是 CSV 或 JSON Lines 格式。使用自有文件时需额外提供三个参数:

  • --train_file / --validation_file:训练与验证文件路径(脚本在 DataTrainingArguments.__post_init__ 中会断言扩展名必须是 csvjsonL288-L307);
  • --text_column:要被摘要的输入文本列名;
  • --summary_column:目标摘要列名。
python examples/pytorch/summarization/run_summarization.py \
    --model_name_or_path google-t5/t5-small \
    --do_train \
    --do_eval \
    --train_file path_to_csv_or_jsonlines_file \
    --validation_file path_to_csv_or_jsonlines_file \
    --text_column text_column_name \
    --summary_column summary_column_name \
    --source_prefix "summarize: " \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --predict_with_generate

列名规则的细节(补充自 摘要任务 README,与源码逻辑一致):

  • 两列 CSV:默认第 1 列当 text、第 2 列当 summary,无需指定列名;
  • 多列 CSV:如表头为 id,date,text,summary,则必须显式传 --text_column text --summary_column summary
  • JSON Lines:每行一个 JSON 对象,默认取第 1 个值当文本、第 2 个值当摘要,键名任意;如想显式指定,同样用 --text_column / --summary_column 传键名。

源码层面,文件加载走 datasets.load_dataset(extension, data_files=...)L397-L413),预处理函数 preprocess_function 会先剔除任一侧为 None 的样本、拼上 source_prefix 再分词(L546-L569)。

八、先用小样本调试脚本

文档建议:正式训练前,先把数据集砍到几十条跑一遍,确认流水线无误,再投入可能耗时数小时的全量训练。对应三个参数:--max_train_samples--max_eval_samples--max_predict_samples

python examples/pytorch/summarization/run_summarization.py \
    --model_name_or_path google-t5/t5-small \
    --max_train_samples 50 \
    --max_eval_samples 50 \
    --max_predict_samples 50 \
    --do_train \
    --do_eval \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --predict_with_generate

实现上,脚本用 dataset.select(range(max_train_samples)) 截取前 N 条(L571-L616)。文档同时提醒:并非所有示例脚本都实现了 --max_predict_samples,不确定时可加 -h 查看帮助:

python examples/pytorch/summarization/run_summarization.py -h

九、从检查点恢复训练

文档介绍了第二个实用开关:当训练被中断时,从既有检查点继续而不是从头再来。方式为追加 --resume_from_checkpoint path_to_specific_checkpoint

python examples/pytorch/summarization/run_summarization.py \
    --model_name_or_path google-t5/t5-small \
    --do_train \
    --do_eval \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --resume_from_checkpoint path_to_specific_checkpoint \
    --predict_with_generate

源码对应逻辑很直白(L681-L685):

if training_args.do_train:
    checkpoint = None
    if training_args.resume_from_checkpoint is not None:
        checkpoint = training_args.resume_from_checkpoint
    train_result = trainer.train(resume_from_checkpoint=checkpoint)

resume_from_checkpoint 本身定义于 training_args.py,接受检查点目录路径或 "last-checkpoint",由 Trainer.train() 负责恢复模型、优化器与随机状态。

十、把训练好的模型发布到 Hub

文档最后一步是分享模型:所有示例脚本都支持把最终模型推送到模型中心。步骤:

  1. 先登录 Hugging Face:
hf auth login
  1. 给脚本追加 --push_to_hub 参数。它会根据你的 Hub 用户名与 --output_dir 中的文件夹名自动创建仓库;
  2. 想给仓库起指定名字,再加 --push_to_hub_model_id,例如:
python examples/pytorch/summarization/run_summarization.py \
    --model_name_or_path google-t5/t5-small \
    --do_train \
    --do_eval \
    --dataset_name cnn_dailymail \
    --dataset_config "3.0.0" \
    --source_prefix "summarize: " \
    --push_to_hub \
    --push_to_hub_model_id finetuned-t5-cnn_dailymail \
    --output_dir /tmp/tst-summarization \
    --per_device_train_batch_size=4 \
    --per_device_eval_batch_size=4 \
    --predict_with_generate

源码中该分支位于脚本收尾处(L752-L755):开启 push_to_hub 时调用 trainer.push_to_hub(**kwargs),否则仅 create_model_card(**kwargs) 生成本地 model card;kwargs 携带 finetuned_fromtasks: summarization、数据集标签等元信息,使上传的仓库自带完整的溯源信息。

十一、延伸阅读与相关路径

适用前提提示:本文所有命令以当前仓库主干的 examples/pytorch 目录结构为准;脚本会强制要求满足其内置的最低 Transformers 版本(check_min_version)。如需 TensorFlow/Flax 示例或更早版本的脚本,请通过 git checkout tags/vX.Y.Z 切换到文档列出的历史 tag,并重新安装对应版本的库。

登录后查看全文
热门项目推荐
相关项目推荐