尧图精选

Transformers 示例训练脚本实战:从环境搭建、单卡运行到分布式、TPU 与 Accelerate 执行

🕒 发布时间:2026/9/7 5:02:13 📁 来源:尧图网络
Transformers 示例训练脚本实战从环境搭建、单卡运行到分布式、TPU 与 Accelerate 执行【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers本篇以 Transformers 官方文档《Trainieren mit einem Skript》脚本训练指南见 docs/source/de/run_scripts.md为主体系统讲解如何基于仓库自带的示例脚本完成模型微调包括从源码安装库与示例依赖、在 CNN/DailyMail 数据集上微调 T5-small 摘要模型的完整命令与参数、分布式与混合精度训练、TPU 执行、使用 Accelerate 运行无 Trainer 脚本、接入自定义 CSV/JSONL 数据集、小样本调试、断点续训与模型上传。读完本文你可以直接复制运行仓库 examples/pytorch/summarization 下的训练脚本并能依据源码理解每个命令行参数背后的处理逻辑。示例脚本的定位与使用边界与 Notebooks 不同Transformers 提供了一批可直接执行的训练脚本用于演示如何针对特定任务微调模型。文档同时提及过研究项目脚本transformers-research-projects和 Legacy 示例它们大多由社区贡献、不再积极维护且依赖特定历史版本与新版库的兼容性没有保证。在当前仓库中示例目录已被精简主示例脚本集中在 examples/pytorch 下如 examples/pytorch/README.mdexamples/research_projects仅保留说明性文件。文档明确了三条使用预期值得在动手前了解示例脚本不保证开箱即用。它们可能无法直接适配你的问题你需要自行改造例如修改数据预处理逻辑大部分脚本会完整暴露数据预处理流程方便你按应用场景修改如果你想在某个示例脚本中新增功能应先在官方论坛或 issue 中讨论再提交 Pull Request——维护方欢迎 bug 修复但通常不会合并以增加功能、损害可读性的改动。本文以文本摘要summarization任务为主线因为它同时覆盖了文档中介绍的全部运行方式。环境搭建从源码安装库与示例依赖文档强调要运行最新示例脚本必须在新的虚拟环境中从源码安装 Transformers而不是直接pip install transformersgit clone https://gitcode.com/GitHub_Trending/tra/transformers cd transformers pip install .如果你的示例对应某个历史版本文档列出了 v2.0.0 至 v4.5.1 的旧版示例索引可以在克隆后切换到对应 tag例如git checkout tags/v3.5.1之后进入目标示例目录安装示例专属依赖cd examples/pytorch/summarization pip install -r requirements.txt摘要示例的依赖清单见 examples/pytorch/summarization/requirements.txt实际包含依赖约束用途accelerate 0.12.0分布式/混合精度训练支持datasets 1.8.0下载与预处理数据集torch 1.3PyTorch 运行框架rouge-score、evaluate-摘要质量评估指标ROUGEnltk、py7zr-ROUGE 计算所需的分句等文本处理sentencepiece! 0.1.92T5 等模型的分词器protobuf-分词/数据格式依赖从源码看run_summarization.py 在入口处做了硬性版本校验check_min_version(4.57.0.dev0) require_version(datasets1.8.0, To fix: pip install -r examples/pytorch/summarization/requirements.txt)这意味着旧版示例对应旧 tag在新源码环境下可能因该检查而报错这正是文档要求按版本配对运行的原因。运行摘要训练脚本摘要脚本从 Datasets 库加载数据集并做预处理然后使用Seq2SeqTrainerTrainer的子类对支持摘要的 seq2seq 架构进行微调。以 T5-small 在 CNN/DailyMail 上微调为例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_name 3.0.0 \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate版本提示文档原文写作--dataset_config 3.0.0旧版字段名。当前仓库脚本中对应的 dataclass 字段名为dataset_config_name见 run_summarization.py#L157-L159使用当前源码运行时应采用--dataset_config_name。参数含义结合源码注释--model_name_or_path预训练模型标识或本地路径ModelArgumentsL96-L100配套字段还有config_name、tokenizer_name、model_revision默认main、trust_remote_code默认False等--do_train/--do_eval控制执行训练/评估来自Seq2SeqTrainingArguments完整的训练参数定义见 src/transformers/training_args.py--dataset_name--dataset_config_name通过 datasets 库下载公开数据集对应load_dataset(data_args.dataset_name, data_args.dataset_config_name, ...)调用L388-L396--source_prefix summarize: T5 系列模型因预训练方式需要显式任务前缀来表明这是摘要任务。源码中甚至有硬性提醒当模型属于google-t5/t5-small、t5-base、t5-large、t5-3b、t5-11b且未提供source_prefix时会打印警告L364-L374前缀会在预处理时拼接到每条输入前L482、L555--per_device_train_batch_size/--per_device_eval_batch_size每个设备的批次大小--predict_with_generate评估/预测时使用model.generate生成摘要文本只有开启它脚本才会计算 ROUGE 指标compute_metricscompute_metrics if training_args.predict_with_generate else NoneL677。脚本还接受若干有默认值的数据参数例如max_source_length默认 1024、max_target_length默认 128、val_max_target_length默认回落为max_target_length同时覆盖model.generate的max_length、num_beams默认 1、ignore_pad_token_for_loss默认True将 pad 位置的 label 置为 -100 不参与损失见 L561-L566。支持哪些模型根据示例目录 READMEexamples/pytorch/summarization/README.mdrun_summarization.py支持的架构为BartForConditionalGeneration、MBartForConditionalGeneration、MarianMTModel、PegasusForConditionalGeneration、T5ForConditionalGeneration、MT5ForConditionalGeneration以及仅用于翻译的FSMTForConditionalGeneration。脚本通过AutoModelForSeq2SeqLM.from_pretrained(...)加载模型L437-L445因此换成上面任一种架构只需修改--model_name_or_path。训练流程在源码中如何串联脚本入口使用HfArgumentParser((ModelArguments, DataTrainingArguments, Seq2SeqTrainingArguments))解析三组参数L331也支持把整份参数写入 JSON 文件一次传入。随后流程为加载数据集Hub 数据集或本地 CSV/JSON 文件见下文自定义数据集一节加载配置、分词器与模型必要时自动resize_token_embeddings并在max_source_length超过模型位置编码长度时自动/显式resize_position_embeddingsL447-L480通过dataset.map(preprocess_function, batchedTrue, ...)完成分词与截断并按main_process_first避免分布式环境下重复预处理使用DataCollatorForSeq2Seq动态填充fp16 训练时pad_to_multiple_of8L618-L625用evaluate.load(rouge)计算指标compute_metrics中会做分句后处理并额外输出gen_len平均生成长度L630-L657初始化Seq2SeqTrainer并执行train()/evaluate()/predict()。分布式训练与混合精度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_name 3.0.0 \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate从源码结构看--fp16不只影响数值精度DataCollatorForSeq2Seq在 fp16 下会把填充对齐到 8 的倍数L624这是 AMP 训练对张量形状的要求训练开始时脚本还会打印当前进程的 rank、device 与是否处于 16 位训练L357-L361。文档同时说明TensorFlow 示例脚本使用MirroredStrategy实现分布式脚本默认在有多卡可用时自动使用多 GPU无需额外参数。在 TPU 上运行脚本TPU 需要借助 PyTorch 的 XLA 支持。仓库提供了启动器脚本 examples/pytorch/xla_spawn.py它接收--num_cores1 或 8与待启动的训练脚本及其全部参数python examples/pytorch/xla_spawn.py --num_cores 8 \ examples/pytorch/summarization/run_summarization.py \ --model_name_or_path google-t5/t5-small \ --do_train \ --do_eval \ --dataset_name cnn_dailymail \ --dataset_config_name 3.0.0 \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate其工作原理在源码中非常直白把训练脚本作为模块导入改写sys.argv再调用xmp.spawn(mod._mp_fn, args(), nprocsargs.num_cores)在 N 个 TPU core 上并行xla_spawn.py#L66-L78。这也是为什么run_summarization.py末尾定义了_mp_fn(index)钩子L760-L762——它是 xla_spawn 约定的多进程入口。使用 Accelerate 运行无 Trainer 脚本 Accelerate 是纯 PyTorch 的分布式训练库提供统一接口在 CPU、单 GPU、多 GPU单机/多机乃至 TPU 上训练同时保留原生 PyTorch 训练循环的可读性。使用它的前提是安装 Accelerate——文档建议安装 Git 源码版本Accelerate 迭代很快而当前示例 requirements.txt 要求accelerate 0.12.0。与run_summarization.py不同Accelerate 路线使用暴露裸训练循环的run_summarization_no_trainer.py约定为task_no_trainer.py命名见 examples/pytorch/summarization/run_summarization_no_trainer.py。它暴露了更底层的循环以便快速实验与自定义如直接修改优化器、DataLoader 配置但选项比 Trainer 版少。操作步骤# 1. 交互式生成配置文件 accelerate config # 2. 验证环境是否正确 accelerate test # 3. 启动训练 accelerate launch examples/pytorch/summarization/run_summarization_no_trainer.py \ --model_name_or_path google-t5/t5-small \ --dataset_name cnn_dailymail \ --dataset_config_name 3.0.0 \ --source_prefix summarize: \ --output_dir ~/tmp/tst-summarization同一条accelerate launch命令即可覆盖 CPU 单机、单 GPU、多 GPU 分布式与 TPU 四种部署形态无需为不同硬件写不同启动命令。使用自定义数据集摘要脚本支持自定义数据集前提是 CSV 或 JSON Line 文件。除--train_file与--validation_file文件路径外多列文件还需指定列名--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_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate源码中的行为边界DataTrainingArguments.post_initdataset_name、train_file、validation_file、test_file至少要提供一个否则抛ValueError文件扩展名必须是csv或json否则assert失败未指定text_column/summary_column时优先匹配内置数据集列名映射summarization_name_mapping如cnn_dailymail → (article, highlights)L310-L323未收录的数据集默认取第 1、2 列或 JSON 的第一个、第二个键指定了列名但文件中不存在时会明确报错并列出可选列名L519-L534。CSV 双列示例examples/pytorch/summarization/README.mdtext,summary Im sitting here in a boring room. ...,Im sitting in a room where Im waiting for something to happenJSONL 示例{text: I see trees so green, red roses too. ..., summary: Im a gardener and Im a big fan of flowers.}用小样本测试脚本全量数据集训练动辄数小时文档建议先用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_name 3.0.0 \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate对应实现是对数据集做前 N 条截取train_dataset.select(range(max_train_samples))等L571-L575。并非所有示例脚本都实现max_predict_samples不确定时加-h查看帮助examples/pytorch/summarization/run_summarization.py -h从检查点恢复训练训练意外中断时可用--resume_from_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_name 3.0.0 \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --resume_from_checkpoint path_to_specific_checkpoint \ --predict_with_generate源码中该参数直接透传给 Trainertrain_result trainer.train(resume_from_checkpointcheckpoint)L681-L685恢复优化器/调度器状态等细节由 src/transformers/trainer.py 中的train()实现。分享模型推送到 Hub所有脚本训练结束后都可把最终模型上传到 Model Hub。先确保已登录当前 CLI 命令hf auth login然后给脚本加--push_to_hub。该参数会创建一个以你的 Hub 用户名 output_dir文件夹名命名的仓库。如需指定仓库名使用对应参数显式命名文档示例写作--push_to_hub_model_id注意当前源码中训练参数的字段名为hub_model_id仓库初始化时读取self.args.hub_model_id见 src/transformers/trainer.py#L584 与 #L3968仓库自动挂在你自己的命名空间下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_name 3.0.0 \ --source_prefix summarize: \ --push_to_hub \ --hub_model_id finetuned-t5-cnn_dailymail \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate脚本末尾会构造 model card 元数据finetuned_from、tasks: summarization、数据集标签与语言有push_to_hub时调用trainer.push_to_hub(**kwargs)否则本地生成create_model_card(**kwargs)L740-L755。小结运行示例脚本的标准路径是源码安装 Transformers →checkout匹配版本如需旧示例→ 安装示例requirements.txt→ 按任务目录运行脚本摘要示例以 run_summarization.py 为主线覆盖 Hub 数据集与自定义 CSV/JSONL 数据、T5 等七种 seq2seq 架构、ROUGE 评估分布式与混合精度靠torchrun --nproc_per_node N--fp16TPU 靠 xla_spawn.py想保留裸训练循环则改用 run_summarization_no_trainer.py 加accelerate launch调试用max_*_samples截断数据中断恢复用--resume_from_checkpoint成果分享用--push_to_hub。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联 返回资讯列表 →