ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Transformers 机器翻译实战:基于 T5 的英法翻译模型微调、SacreBLEU 评估与推理全流程

Transformers 机器翻译实战:基于 T5 的英法翻译模型微调、SacreBLEU 评估与推理全流程 Transformers 机器翻译实战基于 T5 的英法翻译模型微调、SacreBLEU 评估与推理全流程【免费下载链接】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 仓库的阿拉伯语任务教程 docs/source/ar/tasks/translation.md 整理并深化完整讲解如何用Seq2SeqTrainer在 OPUS Books 英法子集上微调 T5 实现英文到法文翻译覆盖数据集加载、动态填充预处理、SacreBLEU 评估、训练参数配置与推理生成等全链路并结合仓库中 Seq2SeqTrainer、DataCollatorForSeq2Seq 的源码实现帮助读者理解每个训练与评估环节背后的调用机制。1. 翻译任务与总体路线翻译Translation是将文本从一种语言转换到另一种语言的过程在 Transformers 中通常被建模为序列到序列Seq2Seq任务编码器读取源语言句子解码器自回归地生成目标语言句子。除了最典型的双语文本翻译这类模型同样可以扩展到语音翻译、语音转文本等跨模态场景。本指南对应仓库教程中的实战目标是在 OPUS Books 数据集的英法en-fr子集上微调google-t5/t5-small实现英文到法文的翻译使用微调后的模型进行推理并通过pipeline与model.generate两种方式使用它。在动手之前先确认安装了必要的依赖库pip install transformers datasets evaluate sacrebleu如果希望把训练产物上传到 Hub 并与社区共享建议先登录 Hugging Face 账号在交互式环境中执行notebook_login()即可 from huggingface_hub import notebook_login notebook_login()后续训练阶段会通过push_to_hubTrue自动推送模型这一步是前提。2. 加载 OPUS Books 英法数据集使用datasets库加载 OPUS Books 数据集的英法子集 from datasets import load_dataset books load_dataset(opus_books, en-fr)教程采用train_test_split将训练部分再拆分为训练集与验证集留出 20% 作为测试数据而不是使用数据集自带的划分字段 books books[train].train_test_split(test_size0.2)查看一条样本可以确认数据结构 books[train][0] {id: 90560, translation: {en: But this lofty plateau measured only a few fathoms, and soon we reentered Our Element., fr: Mais ce plateau élevé ne mesurait que quelques toises, et bientôt nous fûmes rentrés dans notre élément.}}其中translation字段是一个字典分别存放同一段文字的英文与法文译文这正是 Seq2Seq 训练所需的“输入-目标”配对id为样本编号训练时不需要。3. 数据预处理任务前缀、目标语言与动态填充3.1 加载 T5 分词器 from transformers import AutoTokenizer checkpoint google-t5/t5-small tokenizer AutoTokenizer.from_pretrained(checkpoint)3.2 预处理函数的三个关键职责按照教程预处理函数preprocess_function需要完成三件事添加任务前缀prefixT5 这类多任务模型需要通过输入前缀识别当前任务这里为translate English to French: 。仓库中翻译示例的 READMEexamples/pytorch/translation/README.md同样强调google-t5/t5-small等 T5 系列模型必须配合--source_prefix translate English to Romanian: 使用如果 BLEU 分数异常低首先要检查是否漏掉了 source prefix。设置目标语言text_target将法文目标句传入分词器的text_target参数以保证法文文本被正确切分。若不设置分词器会默认按英文处理目标文本破坏目标句的编码。截断过长序列通过max_length与truncationTrue将序列限制在指定长度以内。 source_lang en target_lang fr prefix translate English to French: def preprocess_function(examples): ... inputs [prefix example[source_lang] for example in examples[translation]] ... targets [example[target_lang] for example in examples[translation]] ... model_inputs tokenizer(inputs, text_targettargets, max_length128, truncationTrue) ... return model_inputs用datasets的map方法批量应用预处理batchedTrue让函数一次处理多条样本显著加快编码速度 tokenized_books books.map(preprocess_function, batchedTrue)3.3 DataCollatorForSeq2Seq为什么用动态填充教程建议使用 [DataCollatorForSeq2Seq] 组装批次并明确指出在组批时把句子动态填充dynamic padding到批次内最长句比把整个数据集统一填充到最大长度更高效。从源码看DataCollatorForSeq2Seq 的默认paddingTrue即longest策略正是“按批内最长序列填充”。它还有三个值得注意的实现细节label_pad_token_id默认值为-100用于填充labels。PyTorch 的交叉熵损失会自动忽略-100位置因此填充部分不会污染训练损失当传入model且模型实现了prepare_decoder_input_ids_from_labels时collator 会直接从labels生成decoder_input_ids见 data_collator.py#L608-L615免去解码器输入的准备开销并在启用 label smoothing 时避免损失被重复计算标签填充方向跟随tokenizer.padding_side默认右侧保证input_ids与labels的位置对齐。 from transformers import DataCollatorForSeq2Seq data_collator DataCollatorForSeq2Seq(tokenizertokenizer, modelcheckpoint)注意第二个参数名为model教程中直接传入了 checkpoint 字符串实际效果是让 Trainer 在运行期绑定到已加载的模型对象如需完全显式控制也可以传入模型实例。4. 评估SacreBLEU 指标与 compute_metrics翻译质量通常用 BLEU 分数衡量。教程通过evaluate库加载 SacreBLEU 指标 import evaluate metric evaluate.load(sacrebleu)随后定义compute_metrics函数接收 Trainer 传入的EvalPrediction预测 token id 与标签 token id完成解码、后处理与打分 import numpy as np def postprocess_text(preds, labels): ... preds [pred.strip() for pred in preds] ... labels [[label.strip()] for label in labels] ... return preds, labels def compute_metrics(eval_preds): ... preds, labels eval_preds ... if isinstance(preds, tuple): ... preds preds[0] ... decoded_preds tokenizer.batch_decode(preds, skip_special_tokensTrue) ... labels np.where(labels ! -100, labels, tokenizer.pad_token_id) ... decoded_labels tokenizer.batch_decode(labels, skip_special_tokensTrue) ... decoded_preds, decoded_labels postprocess_text(decoded_preds, decoded_labels) ... result metric.compute(predictionsdecoded_preds, referencesdecoded_labels) ... result {bleu: result[score]} ... prediction_lens [np.count_nonzero(pred ! tokenizer.pad_token_id) for pred in preds] ... result[gen_len] np.mean(prediction_lens) ... result {k: round(v, 4) for k, v in result.items()} ... return result几个关键实现点值得注意labels ! -100的判断呼应了 3.3 节-100是动态填充留下的占位符解码前先替换回pad_token_id避免把-100这类非法 id 送入batch_decodeprediction_lens统计每条生成序列中非 pad token 的数量并取平均得到gen_len用于观察模型生成是否过长或过短if isinstance(preds, tuple)是对predict_with_generateTrue场景的兼容此时 Trainer 返回的预测是(生成 token, 损失)元组需要取第一个元素。该函数在训练配置阶段传入Seq2SeqTrainer每轮评估时自动被调用。5. 训练Seq2SeqTrainingArguments 与 Seq2SeqTrainer5.1 加载模型用AutoModelForSeq2SeqLM加载 T5 的 Seq2Seq 版本带语言建模头可直接用于翻译 from transformers import AutoModelForSeq2SeqLM, Seq2SeqTrainingArguments, Seq2SeqTrainer model AutoModelForSeq2SeqLM.from_pretrained(checkpoint)5.2 训练参数配置教程给出三步骤配置训练参数、构造 Trainer、执行训练。完整参数如下并附逐项说明 training_args Seq2SeqTrainingArguments( ... output_dirmy_awesome_opus_books_model, ... eval_strategyepoch, ... learning_rate2e-5, ... per_device_train_batch_size16, ... per_device_eval_batch_size16, ... weight_decay0.01, ... save_total_limit3, ... num_train_epochs2, ... predict_with_generateTrue, ... fp16True, #change to bf16True for XPU ... push_to_hubTrue, ... )参数取值说明output_dirmy_awesome_opus_books_model唯一必填项本地模型与检查点的保存目录eval_strategyepoch每个 epoch 结束后触发一次评估与检查点保存learning_rate2e-5全参数微调 T5-small 的常用学习率量级per_device_train_batch_size/per_device_eval_batch_size16/16单卡训练/评估批大小weight_decay0.01权重衰减辅助正则化save_total_limit3最多保留 3 个检查点避免磁盘占满num_train_epochs2训练轮数predict_with_generateTrue评估时用generate生成文本再计算 BLEU而非仅看损失fp16True启用 FP16 混合精度XPU 设备请改为bf16Truepush_to_hubTrue训练结束后自动推送模型到 Hub需已登录从源码看Seq2SeqTrainingArguments 在TrainingArguments基础上新增了 Seq2Seq 专用的predict_with_generate、generation_max_length、generation_num_beams、generation_config与sortish_sampler字段。其中predict_with_generateTrue直接决定评估走“生成式指标”路径Seq2SeqTrainer.prediction_step 在self.args.predict_with_generate为真时调用self.model.generate(...)产出整句译文并把生成结果右填充到max_length后连同标签返回供compute_metrics计算 BLEUgeneration_max_length与generation_num_beams则作为评估循环中generate的默认长度与束搜索宽度优先级低于显式传入的gen_kwargs。此外 Seq2SeqTrainer.init支持通过generation_config加载完整的GenerationConfig并做严格校验覆盖模型默认生成配置。5.3 执行训练与上传 Hub trainer Seq2SeqTrainer( ... modelmodel, ... argstraining_args, ... train_datasettokenized_books[train], ... eval_datasettokenized_books[test], ... processing_classtokenizer, ... data_collatordata_collator, ... compute_metricscompute_metrics, ... ) trainer.train()参数对应关系train_dataset/eval_dataset是第 2 节拆分出的tokenized_books[train]与tokenized_books[test]processing_class传入分词器用于填充、解码等操作data_collator与compute_metrics分别绑定第 3、4 节的实现。训练结束后调用trainer.push_to_hub()把模型发布到 Hub供后续推理直接加载 trainer.push_to_hub()不熟悉 Trainer 训练流程的读者可参考仓库内的 Trainer 教程 docs/source/ar/training.md 补充基础概念。6. 推理pipeline 与手动 generate 两种方式6.1 pipeline 方式对 T5 而言输入必须带上任务前缀。翻译英文句子时写成 text translate English to French: Legumes share resources with nitrogen-fixing bacteria.最省事的用法是pipeline任务名为translation_xx_to_yyxx换为源语言代码yy换为目标语言代码如en、fr、de、es、zh等 from transformers import pipeline translator pipeline(translation_xx_to_yy, modelusername/my_awesome_opus_books_model) translator(text) [{translation_text: Legumes partagent des ressources avec des bactéries azotantes.}]6.2 手动 tokenization generate如果需要控制生成参数或复用自定义解码逻辑可以手动复现 pipeline 的内部流程。第一步把带前缀的文本编码为 PyTorch 张量 from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(username/my_awesome_opus_books_model) inputs tokenizer(text, return_tensorspt).input_ids第二步调用model.generate生成译文教程示例使用了采样式生成参数 from transformers import AutoModelForSeq2SeqLM model AutoModelForSeq2SeqLM.from_pretrained(username/my_awesome_opus_books_model) outputs model.generate(inputs, max_new_tokens40, do_sampleTrue, top_k30, top_p0.95)参数含义max_new_tokens40限制最多新生成 40 个 tokendo_sampleTrue开启随机采样top_k30与top_p0.95分别通过 k 采样与 nucleus 采样收窄候选范围提升译文流畅度。更多生成策略贪心、束搜索、各类采样组合可查阅仓库的生成模块 src/transformers/generation。第三步把 token id 解码回文本 tokenizer.decode(outputs[0], skip_special_tokensTrue) Les lignées partagent des ressources avec des bactéries enfixant lazote.7. 进阶命令行微调脚本 run_translation.py仓库在 examples/pytorch/translation 提供了与教程同源的脚本化微调方案 run_translation.py其 README 明确列出了支持的翻译架构BartForConditionalGeneration、MBartForConditionalGeneration、MarianMTModel、PegasusForConditionalGeneration、T5ForConditionalGeneration、MT5ForConditionalGeneration、FSMTForConditionalGeneration仅翻译。该脚本同样从datasets库拉取数据或读取 jsonlines/csv 文件完成下载、预处理、微调与评估。以 MarianMT 为例无需任务前缀python run_translation.py \ --model_name_or_path Helsinki-NLP/opus-mt-en-ro \ --do_train \ --do_eval \ --source_lang en \ --target_lang ro \ --dataset_name wmt16 \ --dataset_config_name ro-en \ --output_dir /tmp/tst-translation \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generate而 T5 系列模型必须额外指定--source_prefix与第 3.2 节手工教程中的前缀逻辑一致python run_translation.py \ --model_name_or_path google-t5/t5-small \ --do_train \ --do_eval \ --source_lang en \ --target_lang ro \ --source_prefix translate English to Romanian: \ --dataset_name wmt16 \ --dataset_config_name ro-en \ --output_dir /tmp/tst-translation \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generateREADME 还给出两点实用提醒若 BLEU 分数很差先确认是否遗漏--source_prefixT5 系列更换语言对时--source_lang、--target_lang与--source_prefix三处必须同步修改MBart 模型的语言代码需要国家码后缀形式如en_XX、ro_RO而非裸的en、ro。此外该目录下的 run_translation_no_trainer.py 演示了不依赖 Trainer、自行编写训练循环的写法适合作为理解 Trainer 内部机制的对照阅读。8. 小结本教程完整覆盖了 Transformers 中 Seq2Seq 翻译任务的工程闭环数据load_dataset(opus_books, en-fr)train_test_split构建英法平行语料预处理任务前缀 text_target目标语言编码 DataCollatorForSeq2Seq动态填充-100标签填充、decoder_input_ids自动生成均由 collator 源码保障评估SacreBLEU 指标配合predict_with_generateTrue在 Seq2SeqTrainer.prediction_step 中触发generate生成式评估训练Seq2SeqTrainingArguments控制学习率、精度、检查点与 Hub 推送trainer.train()一键完成微调推理pipeline(translation_xx_to_yy)快速出结果或model.generate精细控制max_new_tokens、do_sample、top_k、top_p等生成行为。掌握这条链路后替换 checkpoint、语言对与前缀即可将同一套流程迁移到其他 T5/Bart/M2M 等 Seq2Seq 架构的翻译场景中。【免费下载链接】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),仅供参考
返回列表