如何为Donut Transformer配置trainer的compute_metric函数?
Donut Transformer微调的compute_metrics正确实现方式
你的问题出在Donut是序列到序列(Seq2Seq)生成任务,不是普通分类任务,原来的代码用了分类任务的指标计算逻辑,完全不适用。以下是适配Donut的正确实现思路和代码:
核心问题分析
- Donut的
pred.predictions是模型输出的token logits(形状为[batch_size, max_gen_len, vocab_size]),不是分类任务的类别概率; - 标签
pred.label_ids中包含-100(用来忽略pad部分的损失计算),不能直接拿来做分类指标; - 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
相关产品推荐
相关产品推荐

