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

微调BART模型做问题生成时遇Trainer._maybe_log_save_evaluate IndexError

问题诊断与修复方案

核心问题分析

你的代码存在几个关键错误,直接导致验证阶段触发IndexError:

  1. 模型初始化逻辑错误
    BartForConditionalGeneration(config).from_pretrained("facebook/bart-base") 是错误调用方式,from_pretrained是类方法,无需先实例化模型再调用;且手动加.cuda()会和Accelerator的设备管理冲突,引发设备不兼容问题。

  2. 评估指标不匹配
    你使用的SQuAD metric是为问答任务(给定问题+上下文预测答案)设计的,而你的任务是问题生成(给定上下文生成问题),两者输入输出格式完全不兼容,直接调用会导致维度/格式不匹配的索引错误。同时,当predict_with_generate=True时,eval_pred.predictions是模型生成的token序列,不是模型输出的logits,调用argmax(axis=-1)完全错误。

  3. 标签处理缺失
    BART训练时需要将padding的token ID替换为-100(PyTorch会忽略该值计算损失),你直接用padding后的input_ids作为labels,会导致模型计算padding部分的无效损失,同时干扰后续评估流程。

  4. 分词器冗余包装
    accelerator.prepare(tokenizer)完全多余,分词器是数据预处理工具,不需要适配分布式设备。


修复后的完整代码

from datasets import load_dataset
from evaluate import load
from accelerate import Accelerator
from transformers import BartForConditionalGeneration, BartTokenizer
from transformers import Seq2SeqTrainingArguments, Seq2SeqTrainer 
import numpy as np

# 加载数据集和适配问题生成的评估指标(ROUGE为生成式任务标准指标)
dataset = load_dataset("squad")
metric = load("rouge")
accelerator = Accelerator()

def model_init():
    # 正确加载预训练BART模型,由Accelerator自动处理设备分配
    model = BartForConditionalGeneration.from_pretrained("facebook/bart-base")
    return accelerator.prepare(model)

# 分词器无需Accelerator包装
tokenizer = BartTokenizer.from_pretrained("facebook/bart-base")

def preprocess_function(data):
    inputs = tokenizer(
        data['context'], 
        add_special_tokens=True, 
        max_length=256, 
        padding="max_length", 
        truncation=True
    )
    targets = tokenizer(
        data['question'], 
        add_special_tokens=True, 
        max_length=32, 
        padding="max_length", 
        truncation=True
    )
    # 将padding token替换为-100,避免计算无效损失
    labels = [
        [-100 if token == tokenizer.pad_token_id else token for token in label]
        for label in targets['input_ids']
    ]
    return {
        'input_ids': inputs['input_ids'], 
        'attention_mask': inputs['attention_mask'], 
        'labels': labels
    }

# 预处理数据集
dataset = dataset.map(preprocess_function, batched=True).shuffle(seed=777)

training_args = Seq2SeqTrainingArguments(
    output_dir="./results",
    evaluation_strategy="steps",
    eval_steps=500,
    save_steps=50000,
    learning_rate=2e-5,
    per_device_train_batch_size=4,
    per_device_eval_batch_size=4,
    num_train_epochs=2,
    weight_decay=0.01,
    predict_with_generate=True,
    # 添加生成参数,适配问题生成任务
    generation_max_length=32,
    generation_num_beams=4,
)

def compute_metrics(eval_pred):
    predictions, labels = eval_pred
    # 将labels中的-100还原为pad_token_id,方便解码
    labels = np.where(labels != -100, labels, tokenizer.pad_token_id)
    
    # 解码预测和标签为可读文本
    decoded_preds = tokenizer.batch_decode(predictions, skip_special_tokens=True)
    decoded_labels = tokenizer.batch_decode(labels, skip_special_tokens=True)
    
    # 计算ROUGE指标并格式化输出
    result = metric.compute(predictions=decoded_preds, references=decoded_labels, use_stemmer=True)
    result = {key: round(value * 100, 2) for key, value in result.items()}
    return result

# 初始化Trainer
trainer = Seq2SeqTrainer(
    args=training_args,
    train_dataset=dataset["train"],
    eval_dataset=dataset["validation"],
    tokenizer=tokenizer,
    model_init=model_init,
    compute_metrics=compute_metrics,
)

trainer.train()

关键修复点说明

  • 模型初始化:直接通过from_pretrained加载预训练模型,由Accelerator自动处理设备分配,避免手动干预设备设置。
  • 评估指标替换:用ROUGE指标替代SQuAD,适配问题生成的文本输出格式,解决指标不兼容问题。
  • 标签处理:将padding token替换为-100,既避免无效损失计算,也能保证后续解码评估的正确性。
  • 生成参数配置:添加generation_max_length和generation_num_beams,确保模型生成的序列符合任务长度要求,提升生成质量。
  • 分词器优化:移除冗余的accelerator.prepare(tokenizer),回归分词器正常使用流程。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 21:07:11