You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.30 21:27:04