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

自定义T5模型自定义指标计算报错,无法获取准确率求助

解决T5模型compute_metrics中的TypeError及准确率计算问题

问题根源

你遇到的TypeError本质是对predictions的结构处理错误:T5的Trainer返回的predictions是包含logits的元组(形如(logits_array,)),且logits的维度为[batch_size, 生成序列长度, 词表大小]。直接对整个元组或错误维度做argmax,会导致把list/多维数组当成整数传入后续操作(比如tokenizer解码),触发报错。

分步解决代码

1. 正确的compute_metrics函数实现

import numpy as np

def compute_metrics(eval_pred):
    # 提取logits和标签:T5返回的predictions是(logits,)元组
    logits, labels = eval_pred
    # 对logits在词表维度取argmax,得到预测的token id
    pred_ids = np.argmax(logits, axis=-1)
    
    # 处理标签中的-100(Trainer会把padding标签设为-100,需替换为padding token id)
    labels = np.where(labels != -100, labels, tokenizer.pad_token_id)
    
    # 解码预测和标签的token id为文本
    pred_texts = tokenizer.batch_decode(pred_ids, skip_special_tokens=True)
    label_texts = tokenizer.batch_decode(labels, skip_special_tokens=True)
    
    # 计算完全匹配准确率:预测文本与标签完全一致则计数
    accuracy = sum(1 for p, l in zip(pred_texts, label_texts) if p.strip() == l.strip()) / len(pred_texts)
    
    return {"accuracy": accuracy}

2. 关键细节说明

  • 提取logits:务必从元组中取出第一个元素,不能直接用整个元组做argmax操作。
  • 维度处理:np.argmax(logits, axis=-1)是对每个token位置的词表概率取最大值,得到维度为[batch_size, 生成序列长度]的预测token id数组,符合tokenizer解码要求。
  • 标签预处理:T5训练时会把无需计算loss的padding标签设为-100,解码前必须替换为tokenizer的pad_token_id,否则会触发解码错误。
  • 准确率定义:这里用的是完全匹配准确率,如果你的任务允许部分匹配(比如仅判别式计算正确就算对),可以修改判断逻辑,比如提取预测文本中的判别式部分与标签对比。

3. Trainer调用示例

确保初始化Trainer时传入自定义的compute_metrics函数:

from transformers import Trainer, TrainingArguments

training_args = TrainingArguments(
    output_dir="./t5_quadratic_solver",
    evaluation_strategy="epoch",  # 每个epoch评估一次
    per_device_train_batch_size=8,
    per_device_eval_batch_size=8,
    num_train_epochs=3,
    logging_dir="./logs",
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    compute_metrics=compute_metrics,  # 传入自定义函数
)

常见错误排查

  • 如果仍报类似错误,检查pred_ids的形状:确保是二维数组(batch_size × seq_len),若为三维则说明argmax的axis参数设置错误。
  • 解码时出现乱码,确认tokenizer与训练时使用的是同一个实例,且skip_special_tokens=True。
  • 标签文本若有特殊格式(比如示例中的空格分隔),需确保训练时的tokenizer处理逻辑与预测解码逻辑一致。

内容的提问来源于stack exchange,提问作者ALiCe P.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 07:52:41