T5 文本摘要微调实战:从 CNN/Daily Mail 数据集到一句话生成摘要
2026/9/14 1:54:33 网站建设 项目流程

T5 文本摘要微调实战:从 CNN/Daily Mail 数据集到一句话生成摘要

【免费下载链接】Transformers-TutorialsThis repository contains demos I made with the Transformers library by HuggingFace.项目地址: https://gitcode.com/GitHub_Trending/tr/Transformers-Tutorials

Transformers-Tutorials 是 HuggingFace Transformers 库的示例仓库,其中 T5/ 目录收录了长文本摘要任务的完整示例。本文以该示例为主线,走通 T5 模型微调的全流程:加载 CNN/Daily Mail 数据集、分词与任务前缀预处理、用 Accelerate 在 TPU 上训练,最后对新文章生成摘要。

为什么长文本摘要任务适合选 T5

文本到文本范式:把摘要统一为文本转换

T5(Text-to-Text Transfer Transformer)是 Google 提出的通用文本转换模型。它的核心思路是把所有 NLP 任务统一成"输入文本 → 输出文本":翻译的输入是源语文本,问答的输入是问题和上下文,摘要的输入就是整篇文章。同一套模型权重可以在不同任务间复用,切换任务时只需更换任务前缀并做少量微调。

T5 在摘要场景被频繁使用,原因可以归为三点:

  • 编码器-解码器结构:编码器先通读长文并压缩成向量表示,解码器再据此逐词产出摘要,天然适配"长变短"的任务形态。
  • 任务前缀机制:示例在输入前拼接 "Vat samen: "(荷兰语"Summarize: "),向模型声明当前任务是摘要;换任务时沿用同样机制。
  • 预训练基线强:模型经过大规模语料预训练,小规模微调就能得到可用的摘要效果。

CNN/Daily Mail 摘要数据集的加载与样例查看

安装依赖并加载训练、验证、测试三个划分

CNN/Daily Mail 是文本摘要的经典基准数据集,约含 30 万篇新闻文章及人工撰写的摘要。仓库示例使用的是荷兰语版本ml6team/cnn_dailymail_nl(ML6 对同一语料的荷兰语翻译),字段结构一致:article存原文,highlights存摘要。

pip install transformers datasets accelerate sentencepiece
from datasets import load_dataset train_ds, val_ds, test_ds = load_dataset( "ml6team/cnn_dailymail_nl", split=["train", "validation", "test"]) example = train_ds[0] print(example["article"][:200]) print(example["highlights"])

加载后打印首条样例即可确认两个字段可读。highlights通常是数条以***分隔的短句,将直接作为训练标签。

文本预处理:任务前缀、截断与标签填充

分别对文章和摘要做分词

T5 不能直接吃文本,输入必须先转成input_ids(词表中词汇的整数 ID 序列)。示例在预处理中一次完成三件事:拼接任务前缀、按固定长度截断、把填充位替换为忽略标记。

from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("flax-community/t5-base-dutch") prefix = "Vat samen: " # 荷兰语,意为 "Summarize: " def preprocess_examples(examples): inputs = tokenizer([prefix + x for x in examples["article"]], max_length=512, padding="max_length", truncation=True) labels = tokenizer(examples["highlights"], max_length=64, padding="max_length", truncation=True).input_ids inputs["labels"] = [[-100 if y == tokenizer.pad_token_id else y for y in lab] for lab in labels] return inputs

两个截断长度分别对应编码器与解码器:512 是文章侧的接收上限,64 是摘要侧的预期长度,因为highlights通常只有几个短句。

为什么填充位要用 -100 替换

💡 摘要被补齐到固定长度后,填充位置的 token ID 会留在 labels 里。若不处理,损失函数也会要求模型在这些位置"生成填充",梯度方向被带偏。标准做法是把填充位统一替换为 -100:PyTorch 的 CrossEntropyLoss 将其视为忽略索引,不计损失。预处理后可用tokenizer.decode把编码 ID 还原成文本,抽查内容没有被截断。

用 HuggingFace Accelerate 在 TPU 上跑模型微调

训练关键超参数及其含义

示例在 Colab TPU 上训练。HuggingFace Accelerate 负责把模型和数据分发到 8 个 TPU 核心,训练代码按普通"单设备循环"编写即可,分发细节由 Accelerator 托管。示例中的关键超参数:

超参数示例取值含义
learning_rate1e-4学习率,即每次权重更新的步长
train_batch_size2每个 TPU 核心的批大小,分发后实际乘 8
num_epochs1000设得很大,实际结束由早停决定
patience3验证损失连续不提升多少轮后停止训练
seed42随机种子,保证结果可复现

📌 "轮数设很大 + 早停"是摘要任务的常见做法:不必预估训练多少轮,交给验证损失决定。

包装训练函数并用 notebook_launcher 启动

from transformers import T5ForConditionalGeneration, AdamW, set_seed from accelerate import Accelerator, notebook_launcher hyperparameters = {"learning_rate": 1e-4, "num_epochs": 1000, "train_batch_size": 2, "patience": 3} def training_function(): accelerator = Accelerator() set_seed(42) model = T5ForConditionalGeneration.from_pretrained( "flax-community/t5-base-dutch") optimizer = AdamW(model.parameters(), lr=hyperparameters["learning_rate"]) # 构建 DataLoader,经 accelerator 做 prepare,跑标准训练循环 # 每个 epoch 监控验证损失,超过 patience 即早停 notebook_launcher(training_function)

notebook_launcher以多进程方式运行training_function。训练完成后模型保存到output_dir,示例还建议把模型上传到模型仓库(Hub)供他人复用。

微调效果验证:用新文章生成摘要

加载检查点并调用 generate

训练目标就是让模型对"前缀 + 文章"的输入产出简短摘要。验证时取一篇训练未见过的新闻,直接走generate

trained_model = T5ForConditionalGeneration.from_pretrained(hyperparameters["output_dir"]) input_ids = tokenizer(text, return_tensors="pt").input_ids generated_ids = trained_model.generate( input_ids, do_sample=True, max_length=50, top_k=0, temperature=0.7) summary = tokenizer.decode(generated_ids.squeeze(), skip_special_tokens=True) print(summary)

三个生成参数的作用:

  • do_sample=True:从概率分布中采样而非总取概率最大的词,让输出有变化。
  • top_k=0:即 top-p(核)采样,只在累计概率覆盖 p 的最小词集合内采样。
  • temperature=0.7:温度低于 1 时概率分布更尖锐,输出更稳定。

若想要更保守的输出,可改用束搜索(num_beams=4);批量化生产时也可改用单卡 GPU 运行。

优化与进阶方向

调摘要长度、模型与硬件

  • 摘要长度:若发现highlights被截断,把max_target_length从 64 提到 128。
  • 数据与模型:荷兰语语料规模较小,换用英文t5-base时可尝试更大的cnn_dailymail数据集。
  • 硬件:Accelerate 与设备无关,把 TPU 运行时换成多卡 GPU 环境,训练代码基本不用改。

常见的迁移场景

"文本到文本"的形态让迁移到其他摘要场景的成本很低:技术报告换成"标题 + 正文"两个字段即可复用同一条流水线;多语言摘要可换 mT5 等多语言预训练模型;同目录还有代码文档生成示例 Fine_tune_CodeT5_for_generating_docstrings_from_Ruby_code.ipynb,微调逻辑与本文完全相通。

以上就是 T5 文本摘要模型微调的完整过程:数据加载、预处理、Accelerate 训练、一句话生成验证,每一环都是可以独立跑通的代码片段。完整示例 notebook 见 T5/,把模型名和数据集名替换掉,就能把这套流水线搬进自己的业务里。

【免费下载链接】Transformers-TutorialsThis repository contains demos I made with the Transformers library by HuggingFace.项目地址: https://gitcode.com/GitHub_Trending/tr/Transformers-Tutorials

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

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

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

立即咨询