微调T5 Transformer生成重复输出问题排查与解决
T5微调后生成重复食材内容的问题分析与解决方案
核心原因及对应解决方案
1. 训练数据质量不足
如果训练数据中存在大量输出仅重复输入食材的样本,或样本数量少、多样性差,模型会优先学习到“复制食材”的模式,而非生成菜谱逻辑。
解决方案:
- 清洗数据集:移除所有输出仅为食材罗列的样本,确保每个样本的输出是完整的菜谱步骤(含做法、烹饪流程)。
- 扩充数据集:增加不同食材组合的高质量菜谱样本,覆盖更多烹饪场景。
代码示例(数据过滤):
from datasets import load_dataset # 加载自定义菜谱数据集(JSON格式,含input和output字段) dataset = load_dataset("json", data_files="recipes_dataset.json") # 过滤输出仅为食材的样本(假设有效菜谱至少包含2个步骤) def filter_low_quality_samples(example): return len(example["output"].split("\n")) >= 2 dataset = dataset.filter(filter_low_quality_samples)
2. 解码策略选择不当
T5默认使用贪心解码(Greedy Search),这种策略会每次选择概率最高的token,极易陷入循环重复,尤其是生成较长文本时。
解决方案:
切换到更适合开放文本生成的解码策略,比如Beam Search结合重复n-gram限制,或随机采样(Top-K/Top-P):
方案A:Beam Search + 禁止重复n-gram
from transformers import T5ForConditionalGeneration, T5Tokenizer model = T5ForConditionalGeneration.from_pretrained("./fine-tuned-t5") tokenizer = T5Tokenizer.from_pretrained("t5-small") input_text = "generate a recipe: whole chicken, rice, eggs" inputs = tokenizer(input_text, return_tensors="pt", truncation=True) # 生成时禁止重复2元组,避免连续重复 outputs = model.generate( **inputs, max_length=250, num_beams=5, no_repeat_ngram_size=2, early_stopping=True, skip_special_tokens=True ) print(tokenizer.decode(outputs[0]))
方案B:随机采样(增加生成多样性)
outputs = model.generate( **inputs, max_length=250, do_sample=True, top_k=50, top_p=0.95, temperature=0.7, # 控制随机性,值越低越保守 skip_special_tokens=True )
3. 训练参数不合理
- 欠拟合:训练轮数不足、学习率过低,模型未充分学习菜谱生成逻辑;
- 过拟合:学习率过高、未加正则化,模型过度拟合训练数据中的重复模式。
解决方案:
- 调整训练轮数,结合验证集监控模型性能(如Perplexity、生成样本质量);
- 降低学习率(推荐
2e-5~5e-5),添加权重衰减、Dropout等正则化手段; - 启用
load_best_model_at_end,保存验证集表现最优的模型。
代码示例(训练参数配置):
from transformers import Seq2SeqTrainingArguments, Seq2SeqTrainer training_args = Seq2SeqTrainingArguments( output_dir="./t5-recipe-generator", per_device_train_batch_size=8, per_device_eval_batch_size=8, learning_rate=3e-5, num_train_epochs=12, weight_decay=0.01, # 权重衰减抑制过拟合 logging_steps=100, evaluation_strategy="epoch", # 每轮验证一次 save_strategy="epoch", load_best_model_at_end=True, # 加载验证集最优模型 metric_for_best_model="perplexity" ) # 加载模型时启用Dropout model = T5ForConditionalGeneration.from_pretrained("t5-small", dropout_rate=0.15)
4. Prompt格式不一致
T5是Prompt驱动的模型,训练时的输入Prompt必须与推理时完全一致,否则模型无法正确理解任务指令。此外,若训练数据的输出包含重复的Prompt前缀(如a recipe:),模型会延续这种重复模式。
解决方案:
- 统一训练与推理的Prompt格式:比如训练时输入固定为
generate a recipe: {ingredients},推理时也使用完全相同的格式; - 确保训练数据的输出直接从菜谱内容开始,不要包含Prompt前缀(如避免输出开头是
a recipe:)。
示例训练数据格式:
{ "input": "generate a recipe: whole chicken, rice, eggs", "output": "1. Preheat your oven to 375°F. Pat the whole chicken dry with paper towels and season with salt, pepper, and garlic powder... 2. Cook the rice according to package instructions... 3. Scramble the eggs in a pan with a bit of oil until fluffy..." }
额外排查点
- 检查模型生成的
max_length设置:若设置过小,模型可能无法生成完整菜谱,只能重复食材; - 确认tokenizer的
pad_token和eos_token配置正确,避免生成过程无法终止; - 若使用自定义数据集,确保输入输出的tokenization过程一致(如是否开启
truncation、padding)。
内容的提问来源于stack exchange,提问作者onStackOverflowListener
相关产品推荐
相关产品推荐

