如何获取未预训练的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
相关产品推荐
相关产品推荐

