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

如何在Huggingface Trainer中使用SQuAD指标输出F1与准确率?

在训练阶段获取抽取式QA的SQuAD指标并集成到Trainer

要在训练过程中实时查看SQuAD的F1和Exact Match(精确匹配)指标,核心是保留原始答案信息并正确关联模型预测与真实答案,以下是具体步骤:

1. 预处理时保留关键信息

分词后的数据集默认会过滤掉原始answers字段,所以在预处理函数里必须手动保留原始答案、offset_mapping(用于将token位置还原为原始文本跨度):

def preprocess_function(examples):
    questions = [q.strip() for q in examples["question"]]
    inputs = tokenizer(
        questions,
        examples["context"],
        max_length=384,
        truncation="only_second",
        return_offsets_mapping=True,  # 必须开启,用于还原文本跨度
        padding="max_length",
    )

    # 保留原始answers字段和样本ID,确保后续评估能对应
    inputs["answers"] = examples["answers"]
    inputs["id"] = examples["id"]

    # 常规处理token级start/end位置的逻辑
    start_positions = []
    end_positions = []
    for i, offset in enumerate(inputs["offset_mapping"]):
        answer = examples["answers"][i]
        start_char = answer["answer_start"][0]
        end_char = start_char + len(answer["text"][0])
        sequence_ids = inputs.sequence_ids(i)

        # 定位context对应的token区间
        idx = 0
        while sequence_ids[idx] != 1:
            idx += 1
        context_start = idx
        while sequence_ids[idx] == 1:
            idx += 1
        context_end = idx - 1

        # 匹配字符位置到token位置
        start_token = None
        end_token = None
        for j in range(context_start, context_end + 1):
            if offset[j][0] <= start_char and offset[j][1] >= start_char:
                start_token = j
            if offset[j][0] <= end_char and offset[j][1] >= end_char:
                end_token = j
        start_positions.append(start_token if start_token is not None else 0)
        end_positions.append(end_token if end_token is not None else 0)

    inputs["start_positions"] = start_positions
    inputs["end_positions"] = end_positions
    return inputs

# 应用预处理,保留所有需要的字段
tokenized_datasets = raw_datasets.map(preprocess_function, batched=True)

2. 自定义SQuAD指标计算函数

使用evaluate库加载SQuAD评估模块,将模型的token级预测转换为原始文本,再与真实答案比对:

import evaluate
import numpy as np

squad_metric = evaluate.load("squad")

def compute_metrics(eval_pred):
    start_logits, end_logits = eval_pred.predictions
    eval_dataset = tokenized_datasets["validation"]
    raw_eval = raw_datasets["validation"]

    predictions = []
    references = []

    for i in range(len(start_logits)):
        # 还原预测的文本跨度
        offset_mapping = eval_dataset[i]["offset_mapping"]
        context = raw_eval[i]["context"]
        start_pred = np.argmax(start_logits[i])
        end_pred = np.argmax(end_logits[i])

        # 处理无效的预测跨度(start > end)
        if start_pred > end_pred:
            pred_text = ""
        else:
            start_char = offset_mapping[start_pred][0]
            end_char = offset_mapping[end_pred][1]
            pred_text = context[start_char:end_char]

        # 整理成SQuAD评估要求的格式
        predictions.append({
            "id": eval_dataset[i]["id"],
            "prediction_text": pred_text
        })
        references.append({
            "id": eval_dataset[i]["id"],
            "answers": eval_dataset[i]["answers"]
        })

    # 计算F1和Exact Match指标
    return squad_metric.compute(predictions=predictions, references=references)

3. 配置Trainer实现训练时评估

在TrainingArguments中设置评估策略,将自定义的compute_metrics传入Trainer,这样每个epoch结束后会自动计算并输出指标:

from transformers import Trainer, TrainingArguments

training_args = TrainingArguments(
    output_dir="./qa_trainer",
    evaluation_strategy="epoch",  # 每个epoch后执行一次评估
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    num_train_epochs=3,
    logging_dir="./logs",
    logging_strategy="epoch",  # 同步记录日志和指标
    save_strategy="epoch",
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized_datasets["train"],
    eval_dataset=tokenized_datasets["validation"],
    compute_metrics=compute_metrics,  # 传入指标计算函数
)

# 启动训练,训练过程中会在每个epoch后输出F1和Exact Match值
trainer.train()

关键注意点

  • 必须开启return_offsets_mapping=True,否则无法将token位置转换为原始文本的字符位置。
  • 处理预测时要判断start_pred > end_pred的情况,避免出现无效的文本跨度。
  • 确保样本ID在预处理前后保持一致,否则会出现预测与真实答案不匹配的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 09:37:14