别再手动写摘要了:3步完成T5微调,长文本摘要自动生成实战
【免费下载链接】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 数据集上做一次 T5 微调,目标是实现长文本摘要自动生成。全程只需三步:选型、数据适配、训练与推理,单张消费级显卡就能跑完。
先看效果预期:数据集里每篇article通常有五六百词,而highlights参考摘要只有一三句话。微调到位后,模型输出就是这种"完整主谓宾、不是句子拼贴"的人话——这正是生成式摘要相对抽取式的核心优势。
怎么选T5摘要模型:T5 vs BART vs Pegasus 📌
摘要任务常见两条路线:抽取式直接从原文挑句子,快且稳,但句子生硬;生成式让模型自己写,T5、BART、Pegasus 都走这条路。
为什么最终选 T5?三个理由:
t5-base参数量约 2.2 亿,消费级显卡装得下;- 它把摘要统一成"文本到文本"格式,输入前面加一句
summarize:就定义了任务,几乎零额外代码; - HuggingFace 生态兼容性最好,训练和推理各几行 API,完整可运行的示例就放在仓库的 T5/ 目录里。
加载CNN/Daily Mail数据集并预处理:3个易踩的坑 🧰
CNN/Daily Mail 是摘要任务的基准数据集,约 30 万篇文章配人工摘要,训个小模型绰绰有余:
from datasets import load_dataset dataset = load_dataset("cnn_dailymail", "3.0.0") # 3.0.0:生成式版本 print(dataset["train"][0]["article"][:200]) print(dataset["train"][0]["highlights"])两个坑先说清楚:默认版本是抽取式,"摘要"只是原文句子直接拼贴,做生成式摘要必须显式指定3.0.0;另外每条样本只有article和highlights两个字段,够用但不多。
T5 只认input_ids不认文本,把分词逻辑封装成函数,用map批量处理:
from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("t5-base") prefix = "summarize: " def preprocess(examples): inputs = tokenizer([prefix + t for t in examples["article"]], max_length=512, truncation=True) labels = tokenizer(text_target=examples["highlights"], max_length=150, truncation=True) inputs["labels"] = labels["input_ids"] return inputs encoded = dataset.map(preprocess, batched=True)经验提示:
max_length=512意味着超长文本只取前 512 个 token。业务上处理长篇报告的话可以提到 1024,但显存占用会同步上涨,改完记得看一遍 loss 曲线。
怎么微调T5模型的训练参数(Seq2Seq训练参数配置) ⚡
训练用 Seq2SeqTrainer,Seq2Seq 训练参数配置的核心就下面几行:
model = T5ForConditionalGeneration.from_pretrained("t5-base") args = Seq2SeqTrainingArguments( output_dir="./t5-summarization", learning_rate=2e-5, per_device_train_batch_size=16, num_train_epochs=4, evaluation_strategy="epoch", predict_with_generate=True, fp16=True, ) trainer = Seq2SeqTrainer( model, args, train_dataset=encoded["train"], eval_dataset=encoded["validation"], ) trainer.train()learning_rate设多少?T5 微调取2e-5是安全起点,loss 震荡就降到1e-5。显存够不够?fp16混合精度下,t5-base配 batch 16 在 24G 卡上比较从容。轮数为什么是 4?30 万样本一个 epoch 就是一万多步,跑 4 轮基本收敛,再加save_total_limit=3防止 checkpoint 把磁盘写爆。
想上 TPU 呢?Transformers-Tutorials 里配套了一个荷兰语版本示例(Fine_tuning_Dutch_T5_base_on_CNN_Daily_Mail_for_summarization),做法是用 HuggingFace Accelerate 的Accelerator()包住训练函数,模型和数据自动切分到 TPU 多核,你只写纯 PyTorch 代码,剩下的交给框架。更多模型示例可以在 README.md 里按目录找。
摘要模型推理生成:3个最影响输出的generate参数 🔍
训完之后,生成一条摘要只需要一次generate()调用:
def summarize(text): inputs = tokenizer(prefix + text, return_tensors="pt", max_length=512, truncation=True) out = model.generate(**inputs, max_length=150, num_beams=4, early_stopping=True) return tokenizer.decode(out[0], skip_special_tokens=True) print(summarize(dataset["test"][0]["article"]))三个参数值得逐个试:
max_length=150限定摘要长度,太长会"车轱辘话",太短会掐掉关键信息,按业务场景调;num_beams=4是束搜索宽度,越大越稳也越慢,多数场景 4 就够;early_stopping=True让所有束都输出结束符后立刻截断,省下一半推理时间。
避坑要点:推理时输入漏掉
summarize:前缀,模型会把你的文章当成"陌生任务",生成一堆答非所问的内容——这是 T5 微调后第一大 bug 来源。
下一步:微调之后的3个延伸方向
- 用 LoRA 做参数高效微调,显存能压到零头,小卡也能训;
- 在同一数据集上跑一版 BART,A/B 对比生成质量,心里有底;
- 用 FastAPI 或 Gradio 把模型包成 API 服务,接到自己的报告系统里。
挑一个贴合手头工作的方向,接着往下挖就行。
【免费下载链接】Transformers-TutorialsThis repository contains demos I made with the Transformers library by HuggingFace.项目地址: https://gitcode.com/GitHub_Trending/tr/Transformers-Tutorials
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考