如何使用Huggingface预训练模型获取其训练数据集对应的输出
解决方案
你需要先通过Hugging Face Datasets库加载XSUM/CNN DailyMail数据集的原始文本,再批量输入到已加载的预训练模型中生成对应摘要即可,完整实现流程如下:
依赖安装
先确保你已经安装了需要的工具库:
pip install datasets transformers torch pandas
完整实现代码
以下以CNN DailyMail数据集 + BART预训练checkpoint为例:
from datasets import load_dataset from transformers import AutoTokenizer, AutoModelForSeq2SeqLM import torch import pandas as pd # 自动适配GPU/CPU运行 device = "cuda" if torch.cuda.is_available() else "cpu" # 1. 加载目标数据集,如需XSUM替换下方参数为 load_dataset("xsum", split="test") # split可指定为train/validation/test,按需选择要生成结果的数据集拆分 dataset = load_dataset("cnn_dailymail", "3.0.0", split="test") # 2. 加载已训练好的摘要模型和分词器 model_id = "mwesner/pretrained-bart-CNN-Dailymail-summ" tokenizer = AutoTokenizer.from_pretrained(model_id) model = AutoModelForSeq2SeqLM.from_pretrained(model_id).to(device) # 3. 配置生成参数,尽量对齐模型训练时的生成设置,避免结果偏差 # XSUM场景可将max_length调整为64 generate_config = { "num_beams": 4, "max_length": 128, "early_stopping": True, "no_repeat_ngram_size": 3, "truncation": True } # 4. 批量生成摘要 batch_size = 8 # 可根据显存大小调整 result = [] for idx in range(0, len(dataset), batch_size): batch_data = dataset[idx: idx + batch_size] # 对原始文章做分词处理 tokenized_input = tokenizer( batch_data["article"], max_length=1024, truncation=True, padding="max_length", return_tensors="pt" ).to(device) # 生成摘要id序列 summary_ids = model.generate(**tokenized_input, **generate_config) # 解码为自然语言文本 generated_summaries = [ tokenizer.decode(s, skip_special_tokens=True, clean_up_tokenization_spaces=False) for s in summary_ids ] # 拼接原始文章、标准摘要、生成摘要存入结果 for article, gold_sum, gen_sum in zip(batch_data["article"], batch_data["highlights"], generated_summaries): # XSUM场景标准摘要的key为summary,替换batch_data["highlights"]为batch_data["summary"]即可 result.append({ "original_article": article, "gold_summary": gold_sum, "model_generated_summary": gen_sum }) # 5. 保存结果到本地 pd.DataFrame(result).to_csv("bart_cnn_dailymail_summaries.csv", index=False)
适配其他模型/数据集的修改说明
- 切换为Pegasus、T5的对应预训练checkpoint时,仅需替换
model_id为对应模型的ID即可,AutoTokenizer和AutoModelForSeq2SeqLM会自动适配模型结构,无需单独引入对应模型的专属类。 - 如需生成训练集、验证集的结果,修改
load_dataset的split参数为train或validation即可,训练集数据量较大,建议拆分后分批次处理避免内存溢出。
内容的提问来源于stack exchange,提问作者Kiera.K
相关产品推荐
相关产品推荐

