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

如何获取未预训练的Hugging Face T5模型并基于自定义数据集训练?

获取未预训练的T5模型

你可以通过加载T5的配置文件,初始化一个结构完整但权重随机初始化的无预训练模型,代码如下:

from transformers import T5Config, T5ForConditionalGeneration

# 选择对应规模的T5配置(如t5-small、t5-base、t5-large)
config = T5Config.from_pretrained("t5-small")
# 初始化无预训练权重的模型
model = T5ForConditionalGeneration(config)
自定义数据集训练T5文本摘要的完整步骤

1. 数据预处理

确保你的自定义数据集是结构化格式(CSV、JSON均可),每条数据包含原文和对应摘要两个核心字段。使用T5专属Tokenizer处理数据时,需要给输入原文添加summarize: 前缀(T5模型的任务标识要求):

from transformers import T5Tokenizer
import datasets

# 加载自定义数据集(以CSV格式为例)
dataset = datasets.load_dataset("csv", data_files="your_custom_data.csv")
tokenizer = T5Tokenizer.from_pretrained("t5-small")

def preprocess_data(examples):
    # 给原文添加任务前缀
    inputs = ["summarize: " + doc for doc in examples["text"]]
    # 处理输入文本,设置最大长度、截断、填充
    model_inputs = tokenizer(inputs, max_length=512, truncation=True, padding="max_length")
    
    # 处理目标摘要
    labels = tokenizer(text_target=examples["summary"], max_length=128, truncation=True, padding="max_length")
    model_inputs["labels"] = labels["input_ids"]
    return model_inputs

# 批量处理数据集
tokenized_dataset = dataset.map(preprocess_data, batched=True)

2. 配置训练参数

用TrainingArguments定义训练核心参数,再通过Trainer绑定模型、数据集与参数:

from transformers import Trainer, TrainingArguments

training_args = TrainingArguments(
    output_dir="./t5-custom-summarizer",
    evaluation_strategy="epoch",  # 每轮训练后验证
    learning_rate=2e-5,
    per_device_train_batch_size=8,
    per_device_eval_batch_size=8,
    num_train_epochs=3,
    weight_decay=0.01,
    logging_dir="./logs",  # 日志保存路径
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized_dataset["train"],
    eval_dataset=tokenized_dataset["test"],
)

3. 启动训练

直接调用Trainer的训练方法即可:

trainer.train()

训练完成后,保存模型与Tokenizer以便后续使用:

model.save_pretrained("./t5-custom-summarizer-final")
tokenizer.save_pretrained("./t5-custom-summarizer-final")

4. 验证与推理

用训练好的模型生成摘要示例:

def generate_summary(text):
    input_text = "summarize: " + text
    inputs = tokenizer(input_text, return_tensors="pt", max_length=512, truncation=True)
    # 用beam search生成更流畅的摘要
    outputs = model.generate(**inputs, max_length=128, num_beams=4, early_stopping=True)
    return tokenizer.decode(outputs[0], skip_special_tokens=True)

# 测试生成效果
sample_text = "这里替换成你的测试原文内容"
print(generate_summary(sample_text))

内容的提问来源于stack exchange,提问作者CLox

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 00:05:41