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

如何为Donut Transformer配置trainer的compute_metric函数?

Donut Transformer微调的compute_metrics正确实现方式

你的问题出在Donut是序列到序列(Seq2Seq)生成任务,不是普通分类任务,原来的代码用了分类任务的指标计算逻辑,完全不适用。以下是适配Donut的正确实现思路和代码:

核心问题分析

  1. Donut的pred.predictions是模型输出的token logits(形状为[batch_size, max_gen_len, vocab_size]),不是分类任务的类别概率;
  2. 标签pred.label_ids中包含-100(用来忽略pad部分的损失计算),不能直接拿来做分类指标;
  3. Donut的任务目标是生成符合格式的文本(比如结构化的JSON、文本描述),需要用生成任务的指标(如ROUGE、BLEU)或任务特定的结构化指标。

实现步骤与代码示例

1. 基础生成指标(ROUGE/BLEU)

以ROUGE为例,这是文本生成任务常用的指标,先安装依赖库:

pip install rouge_score

然后编写compute_metrics函数:

from datasets import load_metric
from transformers import DonutProcessor

# 加载processor和指标
processor = DonutProcessor.from_pretrained("naver-clova-ix/donut-base")
rouge = load_metric("rouge")

def compute_metrics(pred):
    # 处理预测结果:从logits取token id,解码为文本
    pred_token_ids = pred.predictions.argmax(-1)
    pred_texts = processor.batch_decode(pred_token_ids, skip_special_tokens=True)
    
    # 处理标签:将-100替换为pad token id,再解码为文本
    label_ids = pred.label_ids
    label_ids[label_ids == -100] = processor.tokenizer.pad_token_id
    label_texts = processor.batch_decode(label_ids, skip_special_tokens=True)
    
    # 计算ROUGE指标
    result = rouge.compute(predictions=pred_texts, references=label_texts, use_stemmer=True)
    
    # 简化输出,保留关键指标
    return {
        "rouge1": round(result["rouge1"].mid.fmeasure, 4),
        "rouge2": round(result["rouge2"].mid.fmeasure, 4),
        "rougeL": round(result["rougeL"].mid.fmeasure, 4),
    }

2. 结构化任务指标(如表单/文档信息抽取)

如果你的任务是抽取结构化信息(比如提取表单中的键值对),可以自定义精确匹配或字段级准确率:

import json
from transformers import DonutProcessor

def compute_metrics(pred):
    processor = DonutProcessor.from_pretrained("naver-clova-ix/donut-base")
    
    # 处理预测和标签文本
    pred_token_ids = pred.predictions.argmax(-1)
    pred_texts = processor.batch_decode(pred_token_ids, skip_special_tokens=True)
    
    label_ids = pred.label_ids
    label_ids[label_ids == -100] = processor.tokenizer.pad_token_id
    label_texts = processor.batch_decode(label_ids, skip_special_tokens=True)
    
    # 计算精确匹配数和字段准确率
    exact_match = 0
    total_fields = 0
    correct_fields = 0
    
    for pred_txt, label_txt in zip(pred_texts, label_texts):
        try:
            # 假设输出是JSON格式
            pred_json = json.loads(pred_txt)
            label_json = json.loads(label_txt)
            
            # 精确匹配
            if pred_json == label_json:
                exact_match += 1
            
            # 字段级准确率
            for key in label_json.keys():
                total_fields += 1
                if pred_json.get(key) == label_json[key]:
                    correct_fields += 1
        except json.JSONDecodeError:
            # 解码失败的样本跳过统计
            continue
    
    return {
        "exact_match": round(exact_match / len(pred_texts), 4),
        "field_accuracy": round(correct_fields / max(total_fields, 1), 4)
    }

额外注意事项

  • 确保你的TrainingArguments中设置了evaluation_strategy(如"epoch"或"steps"),否则Trainer不会执行指标计算;
  • 如果使用自定义指标,要注意处理解码失败的异常情况(比如生成的文本不是合法JSON);
  • 可以根据任务需求组合多种指标,同时输出ROUGE和结构化指标。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 21:10:36