T5 文本摘要微调实战:CNN/Daily Mail 数据集全流程复现
2026/9/14 4:06:46 网站建设 项目流程

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 仓库里的 T5 示例,在 CNN/Daily Mail 新闻数据集上做一次微调,长文本摘要就能交给模型稳定输出。下面从装依赖到跑推理完整走一遍,读完你可以独立复现,并把数据换成自己的报告或论文语料。

先看效果:把千字新闻压成两句话

训练完成后的调用方式非常朴素:丢进一段新闻原文,模型吐出两三句话的摘要。最终你会得到类似这样的对照输出:

  • 原文:300 词以上的完整报道
  • 参考摘要:编辑人工撰写的高亮句(highlights)
  • 模型摘要:一到两句话,信息密度向参考摘要靠拢

T5 是 Google 提出的 seq2seq(序列到序列)模型——编码器把整段长文读进去、解码器一个字一个字把摘要吐出来,结构和机器翻译很像。它的特色是把摘要、翻译、问答全部统一成"文本进、文本出",所以摘要任务只需要让模型记住"什么样的原文对应什么样的摘要"。接下来我们把整个过程倒过来重放一遍。

依赖安装与 CNN/Daily Mail 数据集加载

环境一次性装齐,其中 Accelerate 是管理多卡/TPU 分布式训练的库,本教程单卡跑也用不到它的全部能力,先装上备用:

pip install transformers datasets accelerate torch sentencepiece

数据集直接用 Datasets 库从仓库拉取,指定 3.0.0 版本——这个版本的每篇文章平均约 770 词,是真正适合"长文本摘要"的版本:

from datasets import load_dataset dataset = load_dataset("cnn_dailymail", "3.0.0") print(dataset) print(dataset["train"][0]["article"][:200])

加载后包含 train / validation / test 三个子集,每条样本两个字段:article(原文)和highlights(人工摘要,句间用" "连接)。

summarize 前缀与分词截断配置

T5 做摘要时,输入前需要加一个任务前缀summarize:(prefix 即任务指令前缀,相当于告诉模型"你现在干的是摘要这活"),预训练阶段它就是按"前缀 + 输入"的格式见过的:

from transformers import T5Tokenizer tokenizer = T5Tokenizer.from_pretrained("t5-base") def preprocess_function(examples): inputs = tokenizer(["summarize: " + doc for doc in examples["article"]], max_length=512, truncation=True) labels = tokenizer(examples["highlights"], max_length=150, truncation=True) inputs["labels"] = labels["input_ids"] return inputs tokenized = dataset.map(preprocess_function, batched=True, remove_columns=dataset["train"].column_names)

这里对原文和摘要分别分词:原文截到 512 token,摘要截到 150 token。⚠️ 截断是"从尾部砍"的静默操作,512/150 只是常见起点,如果你的语料偏长或摘要老是被切掉,需要按自己的数据统计调大上限。另外你可能在别的教程里见到as_target_tokenizer()的写法——那是给 BERT 类分词器区分[CLS]/[SEP]用的;T5 输入输出共用同一套词表、不加特殊标记,直接分词得到的input_ids就能当 labels,新版 transformers 里该 API 已弃用,照搬反而会收到警告。

训练参数配置与 Trainer 启动

加载 t5-base(约 2.2 亿参数,单卡可训)并配置训练参数。fp16=True是半精度训练:权重用 16 位浮点存储,省显存还提速,代价是极小的数值误差,摘要任务上不敏感。其余像学习率衰减、warmup 步数等参数用默认值即可,这里只留核心项:

from transformers import (T5ForConditionalGeneration, Seq2SeqTrainer, Seq2SeqTrainingArguments) model = T5ForConditionalGeneration.from_pretrained("t5-base") args = Seq2SeqTrainingArguments( output_dir="./t5-cnn-dm", eval_strategy="epoch", learning_rate=2e-5, per_device_train_batch_size=16, num_train_epochs=2, weight_decay=0.01, fp16=True, predict_with_generate=True, save_total_limit=2, )

predict_with_generate=True很关键:它让评估阶段用解码生成来算 ROUGE 分(摘要领域的标准指标,统计生成文本与参考摘要的 n-gram 重合率),而不是只看逐 token 的交叉熵。

trainer = Seq2SeqTrainer( model=model, args=args, train_dataset=tokenized["train"], eval_dataset=tokenized["validation"], ) trainer.train()

Trainer 会自动管批量、梯度累积、保存和评估的循环。训练结束后在测试集上打分:

metrics = trainer.evaluate(tokenized["test"]) print(metrics)

输出里rouge2rougeL越高,说明摘要与人工写的重合度越好。

生成调用:模型摘要 vs 参考摘要

推理时把summarize:前缀加回去,用 beam search(束搜索)解码——每个位置同时保留 top-k 条候选句往下扩,最后取整体概率最高的一条,比逐字贪心更稳。num_beams=4是质量与速度的常用折中:设为 1 最快但容易重复啰嗦,调大则更慢。

def summarize(text, max_len=150, num_beams=4): ids = tokenizer("summarize: " + text, return_tensors="pt", max_length=512, truncation=True) out = model.generate(**ids, max_length=max_len, num_beams=num_beams, early_stopping=True) return tokenizer.decode(out[0], skip_special_tokens=True) sample = dataset["test"][0] print("模型摘要:", summarize(sample["article"])) print("参考摘要:", sample["highlights"])

把两行输出放在一起看,你会发现模型摘要在覆盖关键事实上已经接近参考摘要,只是措辞不完全一致——对摘要任务来说,这正是我们想要的。

仓库的 T5/ 目录 下还有配套 notebook,比如 TPU 上微调 T5 的完整示例.ipynb),里面的超参配置和分词管线与本文同源,可以对照阅读。

这套"前缀 + 截断 + Seq2SeqTrainer"的组合换个数据集就能迁移到技术报告、论文摘要场景,前缀改成summarize:之外的任务指令甚至可以复用做问答。想省显存或快速试验,可以进一步看 LoRA 轻量微调或把模型包成 API 服务。你对自己的语料有类似需求,欢迎在评论区聊聊你的场景。

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

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

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

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

立即咨询