如何在Hugging Face Trainer评估阶段记录带元数据的自定义指标?
将元数据传入Hugging Face Trainer的compute_metrics实现分组评估指标
问题背景
我正在用Hugging Face的Trainer执行句子回归任务,每个样本包含:
- input_ids:分词后的句子
- labels:数值标量目标(用于回归)
- metadata:分类字符串字段(如
project_name或task_type)
核心需求是:评估阶段按元数据类别分组计算损失(比如每个项目/任务的独立损失),梯度计算仍使用全局损失,但需要将元数据传递到compute_metrics函数中完成分组统计。
已尝试操作及问题
- 修改数据集,每个样本返回包含元数据的字典:
{ "input_ids": input_ids, "labels": torch.tensor([numerical_score]), # 标量目标 "metadata": project_name # 元数据字段 }
- 更新DataCollator,确保包含元数据的所有元素正确传入模型
- 配置Trainer时开启
include_inputs_for_metrics=True:
Trainer(..., compute_metrics=compute_metrics, args=TrainingArguments(include_inputs_for_metrics=True))
遇到的问题:元数据在评估循环中丢失,compute_metrics仅能获取predictions和labels,无法访问额外的元数据字段。希望避免重写Trainer类或整个评估循环,减少维护成本,同时保证方案兼容并行运行场景。
期望实现目标
- 按元数据类别分组记录评估指标(如各项目/任务的损失)
- 最小化对Hugging Face Trainer原生逻辑的修改
- 兼容默认Trainer评估循环,保证可扩展性和并行运行稳定性
解决方案:编码元数据为整数ID传递
这是最简洁高效的方案,核心思路是将字符串类型的元数据转换为整数ID,避免gather函数因形状不匹配导致的异常,同时通过include_inputs_for_metrics=True传递到compute_metrics中。
步骤1:编码元数据字符串为整数ID
使用标签编码器将元数据字符串映射为整数,确保能被PyTorch正确处理:
from sklearn.preprocessing import LabelEncoder import torch # 收集评估数据集所有元数据类别 all_metadata = [sample["metadata"] for sample in eval_dataset] le = LabelEncoder() le.fit(all_metadata) # 给数据集添加metadata_id字段 def add_metadata_id(sample): sample["metadata_id"] = le.transform([sample["metadata"]])[0] return sample eval_dataset = eval_dataset.map(add_metadata_id)
步骤2:自定义DataCollator,保留metadata_id
确保collator将metadata_id字段加入batch,传递到后续流程:
from transformers import DataCollatorWithPadding def custom_collator(tokenizer): base_collator = DataCollatorWithPadding(tokenizer) def collate_fn(batch): # 用默认collator处理input_ids和labels processed_batch = base_collator([ {"input_ids": x["input_ids"], "labels": x["labels"]} for x in batch ]) # 添加metadata_id字段,转为tensor processed_batch["metadata_id"] = torch.tensor([x["metadata_id"] for x in batch]) return processed_batch return collate_fn # 实例化自定义collator data_collator = custom_collator(tokenizer)
步骤3:修改compute_metrics实现分组统计
开启include_inputs_for_metrics=True后,eval_pred对象会包含inputs字段,从中提取metadata_id并映射回原字符串,计算分组损失:
import numpy as np from sklearn.metrics import mean_squared_error def compute_metrics(eval_pred): predictions, labels, inputs = eval_pred.predictions, eval_pred.label_ids, eval_pred.inputs # 处理回归任务的预测结果形状(预测结果可能为二维,需压缩为一维) preds = predictions.squeeze() labels = labels.squeeze() # 计算全局MSE损失 global_mse = mean_squared_error(labels, preds) # 将metadata_id映射回原字符串类别 metadata_ids = inputs["metadata_id"].numpy() metadata_labels = le.inverse_transform(metadata_ids) # 按元数据分组计算损失 grouped_metrics = {} unique_metadata = np.unique(metadata_labels) for meta in unique_metadata: mask = metadata_labels == meta group_preds = preds[mask] group_labels = labels[mask] if len(group_preds) > 0: group_mse = mean_squared_error(group_labels, group_preds) grouped_metrics[f"mse_{meta}"] = group_mse # 合并全局与分组指标 metrics = {"global_mse": global_mse} metrics.update(grouped_metrics) return metrics # 初始化Trainer trainer = Trainer( model=model, args=TrainingArguments( output_dir="./results", include_inputs_for_metrics=True, # 其他训练/评估参数 ), train_dataset=train_dataset, eval_dataset=eval_dataset, compute_metrics=compute_metrics, data_collator=data_collator )
方案优势
- 整数类型的
metadata_id能被PyTorch的gather函数正确处理,不会出现形状不匹配导致的挂起问题 - 无需重写Trainer类或评估循环,完全兼容原生流程
- 并行运行时,
include_inputs_for_metrics=True会自动合并所有进程的元数据,保证分组统计的正确性
替代方案:使用EvalCallback收集数据
如果不想修改数据集和collator,可以自定义EvalCallback,在评估阶段收集元数据与预测结果,最后计算分组指标:
from transformers import TrainerCallback import numpy as np from sklearn.metrics import mean_squared_error class MetadataEvalCallback(TrainerCallback): def __init__(self, eval_dataset, label_encoder): self.eval_dataset = eval_dataset self.le = label_encoder self.reset_cache() def reset_cache(self): self.all_preds = [] self.all_labels = [] self.all_metadata = [] def on_evaluate(self, args, state, control, metrics=None, **kwargs): # 收集预测结果与标签 preds = kwargs["predictions"].squeeze() labels = kwargs["label_ids"].squeeze() self.all_preds.extend(preds) self.all_labels.extend(labels) # 收集元数据 self.all_metadata.extend([sample["metadata"] for sample in self.eval_dataset]) # 计算分组指标 preds_np = np.array(self.all_preds) labels_np = np.array(self.all_labels) metadata_np = np.array(self.all_metadata) grouped_metrics = {} unique_meta = np.unique(metadata_np) for meta in unique_meta: mask = metadata_np == meta group_mse = mean_squared_error(labels_np[mask], preds_np[mask]) grouped_metrics[f"mse_{meta}"] = group_mse # 更新评估指标字典 metrics.update(grouped_metrics) # 重置缓存,避免下一次评估累积数据 self.reset_cache() return metrics # 初始化编码器与回调 le = LabelEncoder() le.fit([sample["metadata"] for sample in eval_dataset]) metadata_callback = MetadataEvalCallback(eval_dataset, le) # 初始化Trainer trainer = Trainer( model=model, args=TrainingArguments(output_dir="./results"), train_dataset=train_dataset, eval_dataset=eval_dataset, compute_metrics=lambda eval_pred: { "global_mse": mean_squared_error(eval_pred.label_ids.squeeze(), eval_pred.predictions.squeeze()) }, callbacks=[metadata_callback] )
这个方案无需修改数据集结构,但需要注意每次评估后重置缓存,避免数据累积。
内容的提问来源于stack exchange,提问作者enter_thevoid
相关产品推荐
相关产品推荐

