为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
相关产品推荐
相关产品推荐

