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

微调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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 14:00:27