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

为SFTTrainer添加准确率、召回率与F1指标的技术问询

解决SFTTrainer添加分类评估指标的可行方案

一、验证集预处理修正

首先要确保验证集格式适配指标计算需求:

  • 原验证集为prompt + data + response,需将response单独提取为标签列,保留prompt + data作为输入列,而非直接移除response。
  • 预处理示例代码:
def preprocess_validation_dataset(ds):
    ds = ds.map(lambda x: {
        "input_text": x["prompt"] + x["data"],
        "label_text": x["response"]
    })
    return ds

二、适配SFTTrainer的preprocess_logits_for_metrics实现

SFTTrainer输出的logits是生成任务格式,需先转换为可匹配标签的文本:

def preprocess_logits_for_metrics(logits, labels):
    # 解码预测的token序列为文本
    preds = tokenizer.batch_decode(torch.argmax(logits, dim=-1), skip_special_tokens=True)
    # 处理标签中的padding标记(-100),解码为文本
    labels = torch.where(labels != -100, labels, tokenizer.pad_token_id)
    labels = tokenizer.batch_decode(labels, skip_special_tokens=True)
    return preds, labels

三、compute_metrics函数的正确实现

使用evaluate库计算准确率、F1等指标,需先统一文本格式消除干扰:

import evaluate

acc_metric = evaluate.load("accuracy")
f1_metric = evaluate.load("f1")

def compute_metrics(eval_pred):
    preds, labels = eval_pred
    # 统一文本格式:去除首尾空格、转为小写
    preds = [p.strip().lower() for p in preds]
    labels = [l.strip().lower() for l in labels]
    
    # 计算指标
    accuracy = acc_metric.compute(predictions=preds, references=labels)["accuracy"]
    f1 = f1_metric.compute(predictions=preds, references=labels, average="weighted")["f1"]
    
    return {
        "accuracy": accuracy,
        "f1": f1
    }

四、SFTTrainer初始化的正确配置

确保训练器传入预处理后的验证集及正确的指标计算参数:

from trl import SFTTrainer

trainer = SFTTrainer(
    model=model,
    train_dataset=train_dataset,
    eval_dataset=preprocess_validation_dataset(valid_dataset),
    tokenizer=tokenizer,
    args=training_args,
    compute_metrics=compute_metrics,
    preprocess_logits_for_metrics=preprocess_logits_for_metrics,
    # 其他SFT参数,如max_seq_length、dataset_text_field="input_text"等
)

关键注意事项

  • 若为多分类任务,建议将标签文本映射为整数ID,避免文本匹配误差:
    # 构建标签-ID映射
    label_list = list(set(valid_dataset["label_text"]))
    label2id = {label: i for i, label in enumerate(label_list)}
    
    # 修改预处理逻辑,将文本转为ID
    preds = [label2id.get(p.strip().lower(), -1) for p in preds]
    labels = [label2id.get(l.strip().lower(), -1) for l in labels]
    
  • 确保tokenizer解码时跳过特殊令牌(如<s>、</s>),防止干扰文本匹配。
  • 若预测文本存在冗余,可截取与标签长度一致的片段,或使用Jaccard相似度等模糊匹配方式计算指标。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 07:44:58