Article
第 1 章:Evaluate 库简介与安装
1.1 什么是 Evaluate 库
| 概念名称 | 说明 | 注意事项 |
|---|---|---|
| Evaluate 库 | Hugging Face 提供的统一评估库,用于衡量机器学习模型在各种任务上的性能。 | 不要与 transformers.Trainer 中的 compute_metrics 混淆,Evaluate 是其底层支持工具。 |
| 开源项目 | 属于 Hugging Face 生态系统的一部分,代码托管于 GitHub,支持社区贡献新指标。 | 可访问 huggingface/evaluate 查看源码和文档。 |
| 指标即服务 | 提供标准化接口加载和使用评估指标,支持本地运行和在线加载。 | 所有指标均可通过 evaluate.load() 统一调用,简化使用流程。 |
| 支持任务类型 | 覆盖分类、回归、文本生成、语义相似度、问答、命名实体识别等多种 NLP 任务。 | 也逐步支持语音、多模态等任务的评估。 |
| 与 Datasets 集成 | 可直接应用于 datasets.Dataset 对象,实现高效批量评估。 | 需确保输入数据格式符合指标要求(如 list 而非单个值)。 |
1.2 Evaluate 的核心功能概述
| 功能名称 | 说明 | 注意事项 |
|---|---|---|
| 统一加载接口 | 使用 evaluate.load() 加载任意内置或自定义指标,支持远程和本地加载。 | 需指定正确的指标名称;若为社区贡献指标,需确保名称拼写准确。 |
| 多种指标支持 | 内置数十种常用评估指标,涵盖 accuracy、F1、BLEU、ROUGE、BERTScore 等。 | 不同指标适用于不同任务,需根据任务类型选择合适指标。 |
| 批量计算能力 | 支持对多个样本进行向量化评估,提升大规模数据集上的计算效率。 | 输入应为 Python 列表(list),避免逐个传入单一样本以提高性能。 |
| 可扩展性 | 允许用户定义并注册自己的评估指标,支持本地开发与共享发布。 | 自定义指标需遵循 EvaluationModule 接口规范。 |
| 缓存机制 | 自动缓存已下载的指标脚本和配置,避免重复下载,提升后续加载速度。 | 缓存路径可通过 HF_HOME 或 EVALUATE_CACHE 环境变量自定义。 |
| 与 Transformers 集成 | 可直接在 Trainer 中作为 compute_metrics 函数使用,无缝集成训练流程。 | 需将 compute 方法封装为无状态函数传递给 TrainingArguments。 |
| 指标信息查询 | 提供 evaluate.list() 和 .inspect() 方法查看可用指标及其文档。 | list() 可筛选任务类型、指标类型等,便于发现合适指标。 |
1.3 安装与环境配置
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 安装 evaluate | 在终端执行:pip install evaluate | 建议在虚拟环境中安装,避免依赖冲突。 |
| 升级到最新版 | 使用命令:pip install --upgrade evaluate | 旧版本可能存在 API 不兼容问题,建议保持更新。 |
| 验证安装 | 在 Python 中运行:import evaluateprint(evaluate.version) | 若无报错且能输出版本号,则安装成功。 |
| 设置缓存路径 | 设置环境变量:export EVALUATE_CACHE=/path/to/cache或 export HF_HOME=/path/to/hf_home | 自定义路径可节省主目录空间,适合多用户或服务器环境。 |
| 安装额外依赖 | 某些指标需额外包,如 bertscore 需:pip install bert-scorerouge 需:pip install rouge-score | 使用特定指标前应查看其文档,安装对应依赖,否则 load 会报错。 |
| 使用 conda 安装 | conda 官方无包,但可通过 pip 在 conda 环境内安装:conda activate myenvpip install evaluate | 推荐使用 pip 安装,conda 用户也适用。 |
| 离线使用 | 提前下载所需指标,离线时通过本地路径加载:evaluate.load("./local_metric/") | 需保证本地目录包含完整的指标脚本(.py)和配置文件。 |
第 2 章:基本使用入门
2.1 加载内置指标(load 函数)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
evaluate.load | evaluate.load(path, config_name=None, cache_dir=None, keep_in_memory=False, revision=None, trust_remote_code=False, **kwargs) | 从 Hugging Face Hub 或本地加载指定的评估指标模块 | import evaluateaccuracy = evaluate.load("accuracy")# 加载 F1 分数f1 = evaluate.load("f1", config_name="multiclass")# 加载本地自定义指标custom_metric = evaluate.load("./my_metrics/precision_at_k/") | path 必须为字符串,表示指标名称(如 "f1"、"bleu");若为本地路径则需指向包含 .py 文件的目录;config_name 可用于区分同一指标的不同变体;trust_remote_code=True 才能加载含自定义逻辑的指标;本地路径必须包含 __init__.py 或 .py 主文件 |
2.2 计算单个指标(compute 方法)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
compute | metric.compute(references, predictions, **kwargs) | 根据参考值和预测值计算评估分数 | import evaluateacc = evaluate.load("accuracy")result = acc.compute(references=[1, 0, 1], predictions=[1, 1, 1])print(result) # {'accuracy': 0.6666666666666666}# 计算二分类 F1f1 = evaluate.load("f1")score = f1.compute(references=[0, 1, 0, 1], predictions=[0, 1, 1, 1], average="binary") | references 和 predictions 必须为 list 类型,即使只传入单个样本也需用列表包裹;不同指标支持的 **kwargs 不同,需查阅文档;例如 F1 支持 average 参数("micro", "macro", "weighted" 等);输入长度必须一致 |
2.3 批量计算多个样本(批量输入处理)
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 输入格式统一 | 将所有真实标签组织为一个列表(references),所有预测结果组织为另一个列表(predictions),两个列表长度相等 | 避免逐个样本调用 compute,应一次性传入完整列表以提升性能 |
| 批量传入 compute | 使用 metric.compute(references=refs, predictions=preds) 一次性处理多个样本 | 列表元素数量应与数据集样本数一致;不支持 generator,需转换为 list |
| 处理文本生成任务 | 对于 BLEU、ROUGE 等,references 可为嵌套列表(每个样本对应多个参考翻译) | 如 references=[[ref1a, ref1b], [ref2a]] 表示第一个样本有两个参考答案 |
| 性能优化建议 | 利用向量化操作,避免 for 循环;若数据过大可分块处理 | 超大列表可能导致内存溢出,建议分批计算后合并结果 |
| 示例:批量准确率计算 | refs = [0, 1, 1, 0, 1]preds = [0, 1, 0, 0, 1]acc = evaluate.load("accuracy")result = acc.compute(references=refs, predictions=preds) | 输出为字典形式:{'accuracy': 0.8},便于后续分析 |
2.4 查看指标信息(list 和 inspect 方法)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
evaluate.list | evaluate.list(evaluation_modules=None, with_details=False) | 列出当前可用的所有评估指标 | import evaluateall_metrics = evaluate.list()print(all_metrics[:5]) # 查看前5个# 获取所有分类相关指标cls_metrics = evaluate.list("classification")# 获取详细信息detailed = evaluate.list(with_details=True) | 默认返回指标名称列表;with_details=True 时返回包含描述、citation 等的完整信息字典;支持模糊匹配关键字(如 "text-generation");返回结果可能因网络状态略有延迟 |
inspect | evaluate.inspect(evaluation_module) | 打印指定指标的源码、文档字符串和使用示例 | evaluate.inspect("accuracy")evaluate.inspect("rouge")# 查看 BLEU 指标详情evaluate.inspect("bleu") | 需确保已成功加载该指标;可用于学习指标内部实现逻辑;输出包含 __init__.py 内容、compute 函数定义和测试样例,适合调试与理解参数含义 |
第 3 章:常用指标详解
3.1 分类任务指标(accuracy, precision, recall, f1)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
accuracy | evaluate.load("accuracy").compute(references, predictions) | 计算预测正确的样本占比 | acc = evaluate.load("accuracy")result = acc.compute(references=[0,1,1], predictions=[0,1,0])print(result) # {'accuracy': 0.6666666666666666} | 适用于多类与二分类;不适用于类别严重不平衡场景 |
precision | evaluate.load("precision").compute(references, predictions, average) | 计算精确率:预测为正类中实际为正类的比例 | prec = evaluate.load("precision")score = prec.compute(references=[0,1,1,0], predictions=[1,1,1,0], average="binary") | average 可选 "binary"(二分类), "micro", "macro", "weighted";多分类需指定 |
recall | evaluate.load("recall").compute(references, predictions, average) | 计算召回率:实际正类中被正确预测的比例 | rec = evaluate.load("recall")score = rec.compute(references=[0,1,1,0], predictions=[1,1,1,0], average="macro") | 同 precision,average 控制聚合方式;"macro" 对各类平等加权 |
f1 | evaluate.load("f1").compute(references, predictions, average) | 计算 F1 分数:精确率与召回率的调和平均 | f1 = evaluate.load("f1")score = f1.compute(references=[0,1,1,0], predictions=[1,1,1,0], average="weighted") | F1 综合反映模型性能;"weighted" 按类别频次加权,适合不平衡数据 |
3.2 回归任务指标(mse, mae, r_squared)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
mse | evaluate.load("mse").compute(references, predictions, squared) | 计算均方误差(MSE)或均方根误差(RMSE) | mse = evaluate.load("mse")score = mse.compute(references=[1.2, 2.3], predictions=[1.1, 2.5], squared=True)rmse = mse.compute(references=[1.2,2.3], predictions=[1.1,2.5], squared=False) | squared=True 返回 MSE,False 返回 RMSE;值越小越好 |
mae | evaluate.load("mae").compute(references, predictions) | 计算平均绝对误差(MAE) | mae = evaluate.load("mae")result = mae.compute(references=[1.0,2.0,3.0], predictions=[1.1,1.9,3.2])print(result) # {'mae': 0.13333333333333333} | 对异常值比 MSE 更鲁棒;反映预测偏差的平均大小 |
r_squared | evaluate.load("r_squared").compute(references, predictions) | 计算决定系数(R²),反映模型解释方差的比例 | r2 = evaluate.load("r_squared")score = r2.compute(references=[1,2,3], predictions=[1.1,1.9,3.1]) | 取值范围 (-∞, 1],越接近 1 拟合越好;0 表示不优于均值预测 |
3.3 文本生成指标(bleu, rouge, meteor, ter)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
bleu | evaluate.load("bleu").compute(predictions, references) | 基于 n-gram 精确率评估机器翻译或文本生成质量 | bleu = evaluate.load("bleu")refs = [["the cat is on the mat"], ["the dog sat"]]preds = ["the cat is on mat"]result = bleu.compute(predictions=preds, references=refs) | references 为嵌套列表,每个样本可有多个参考文本;需安装 sentencepiece |
rouge | evaluate.load("rouge").compute(predictions, references, rouge_types, use_stemmer) | 评估摘要任务,支持 ROUGE-N, ROUGE-L, ROUGE-W 等 | rouge = evaluate.load("rouge")result = rouge.compute(predictions=["hello world"], references=["hello there"], rouge_types=["rouge1", "rougeL"]) | rouge_types 指定计算哪些变体;use_stemmer=True 可启用词干提取 |
meteor | evaluate.load("meteor").compute(predictions, references) | 基于同义词、词干和句法匹配的文本相似度指标 | meteor = evaluate.load("meteor")score = meteor.compute(predictions=["the cat is there"], references=["there is a cat"]) | 比 BLEU 更贴近人工评价;需安装 nltk 和 java(部分实现) |
ter | evaluate.load("ter").compute(predictions, references, case_sensitive) | 翻译编辑率(TER):将预测转为参考所需最少编辑次数 | ter = evaluate.load("ter")score = ter.compute(predictions=["the cat is here"], references=["the cat is there"], case_sensitive=False) | 值越小越好;case_sensitive 控制是否区分大小写;需安装 sacrebleu 支持 |
3.4 嵌入相似度指标(sacrebleu, bertscore)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
sacrebleu | evaluate.load("sacrebleu").compute(predictions, references, lowercase) | 使用 sacrebleu 库计算标准化 BLEU 分数 | sacrebleu = evaluate.load("sacrebleu")result = sacrebleu.compute(predictions=["the cat is on"], references=[["the cat is on the mat"]], lowercase=True) | 自动处理标记化和标准化;lowercase 可统一转小写;结果与 sacrebleu 命令行工具一致 |
bertscore | evaluate.load("bertscore").compute(predictions, references, model_type) | 基于 BERT 嵌入计算预测与参考之间的语义相似度 | bertscore = evaluate.load("bertscore")result = bertscore.compute(predictions=["a cat"], references=["a kitten"], model_type="bert-base-uncased") | 需安装 bert-score:pip install bert-score;model_type 指定预训练模型;计算较慢但更符合语义 |
第 4 章:高级功能与配置
4.1 自定义指标配置(config_name 参数)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
evaluate.load | evaluate.load(path, config_name=None, **kwargs) | 通过 config_name 加载同一指标的不同配置变体 | f1_multi = evaluate.load("f1", config_name="multiclass")f1_binary = evaluate.load("f1", config_name="binary")# 加载 BLEU 的不同配置bleu_detok = evaluate.load("bleu", config_name="detok")bleu_default = evaluate.load("bleu", config_name="default") | config_name 用于区分同一指标在不同任务下的预设配置;必须是该指标支持的配置名称;若 config_name 不存在会抛出 FileNotFoundError;可通过 inspect 查看可用配置 |
| 预设配置示例 | binary:二分类设置multiclass:多分类默认multilabel:多标签分类detok:带去标记化的 BLEU | 为常见使用场景提供默认参数封装 | # ROUGE 不同变体rouge_l = evaluate.load("rouge", config_name="rougeL")rouge_2 = evaluate.load("rouge", config_name="rouge2") | 使用 config_name 可避免重复传递 kwargs,提升代码可读性 |
4.2 多指标组合计算(combine_metrics)
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
evaluate.combine | evaluate.combine(metric_configs) | 将多个指标合并为一个可统一调用的评估模块 | metrics = evaluate.combine(["accuracy", "f1", "precision", "recall"])results = metrics.compute(references=[0,1], predictions=[0,0], average="binary")# 使用字典指定配置configs = [{"name": "f1", "config_name": "binary"}, "accuracy"]metrics = evaluate.combine(configs) | metric_configs 可为字符串列表或字典列表;返回一个支持 compute 的组合对象 |
compute(组合后) | combined_metric.compute(**kwargs) | 在同一输入上同时计算多个指标 | results = metrics.compute(references=refs, predictions=preds, average="binary")print(results) # {'accuracy': ..., 'f1': ..., 'precision': ..., 'recall': ...} | 所有指标必须接受相同的输入参数;若某指标不支持某 kwargs 会报错 |
4.3 指标缓存机制与性能优化
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 缓存路径设置 | 默认缓存至 ~/.cache/huggingface/evaluate;可通过 EVALUATE_CACHE 或 HF_HOME 环境变量修改 | 多用户系统中建议设置统一缓存路径避免重复下载 |
| 首次加载行为 | 第一次 load 某指标时从 Hugging Face Hub 下载脚本并缓存 | 需联网;下载后后续调用直接使用本地缓存,速度显著提升 |
| 缓存清理 | 手动删除缓存目录中的对应指标文件夹即可 | 清理后下次加载会重新下载;可用于解决脚本损坏问题 |
keep_in_memory | load 时设置 keep_in_memory=True 可将指标常驻内存 | 适合频繁调用场景;增加内存占用 |
| 批量处理优化 | 始终使用 list 一次性传入所有样本,避免循环调用 compute | 向量化操作显著提升性能;尤其对基于模型的指标(如 bertscore)更明显 |
| 离线使用 | 确保指标已缓存后断网使用;或通过本地路径加载 | 本地路径应包含完整的指标 .py 文件和 __init__.py |
4.4 使用附加参数(compute 的 kwargs 参数详解)
| 方法参数 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
average | average="binary" / "micro" / "macro" / "weighted" / None | 指定多分类/多标签下指标的聚合方式 | f1 = evaluate.load("f1")f1.compute(references=[0,1,2], predictions=[0,1,1], average="macro") | "binary" 仅用于二分类;None 返回每个类别的分数(数组) |
squared | squared=True / False | 控制 MSE 是否返回平方形式 | mse = evaluate.load("mse")mse.compute(references=[1,2], predictions=[1.1,2.2], squared=False) | squared=False 即 RMSE;默认 True |
rouge_types | rouge_types=["rouge1", "rouge2", "rougeL", "rougeLsum"] | 指定 ROUGE 计算哪些 n-gram 和长度变体 | rouge = evaluate.load("rouge")rouge.compute(predictions=["..."], references=["..."], rouge_types=["rouge1", "rougeL"]) | "rougeLsum" 适用于摘要(处理换行符) |
use_stemmer | use_stemmer=True / False | 是否启用词干提取器(Porter stemmer) | rouge.compute(..., use_stemmer=True) | 提升语义匹配合理性;轻微性能开销 |
lowercase | lowercase=True / False | 是否统一转为小写再计算 | sacrebleu.compute(..., lowercase=True) | 减少大小写差异影响;某些任务(如命名实体)应设为 False |
tokenizer | tokenizer=custom_fn | 自定义分词函数(如用于 BLEU) | bleu.compute(..., tokenizer=lambda x: x.split()) | 覆盖默认分词逻辑;需返回 token 列表 |
model_type | model_type="bert-base-uncased" 等 | 指定 BERTScore 使用的预训练模型 | bertscore = evaluate.load("bertscore")bertscore.compute(..., model_type="roberta-large") | 模型越大越准但越慢;需提前下载或能联网访问 |
num_layers | num_layers: int | 指定 BERTScore 使用模型的哪一层计算表示 | bertscore.compute(..., num_layers=8) | 默认通常为最后一层;可调整以优化效果 |
第 5 章:自定义指标开发
5.1 使用 evaluate.EvaluationModule 创建自定义指标
| 概念名称 | 说明 | 注意事项 |
|---|---|---|
EvaluationModule 类 | Hugging Face Evaluate 中表示评估模块的基类,可通过子类化创建自定义指标 | 在新版 evaluate 库中,推荐使用函数式接口或直接实现 compute 函数,而非继承该类;此方式主要用于兼容旧代码 |
| 自定义指标结构 | 需创建一个 Python 文件(如 my_metric.py),包含 _info() 和 compute() 两个核心方法 | 文件应放置于独立目录中以便加载 |
_info 方法 | 返回 MetricInfo 对象,描述指标名称、描述、citation、输入类型等元数据 | 必须实现,用于提供指标文档信息 |
MetricInfo 字段 | name, description, citation, features, inputs_description 等 | features 定义输入数据格式(基于 datasets.Dataset) |
| 创建步骤概览 | 1. 定义 .py 文件2. 实现 _info()3. 实现 compute()4. 保存到本地目录 | 不需要显式继承 EvaluationModule 即可被 load 加载 |
5.2 定义 compute 方法逻辑
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
compute | def compute(references, predictions, **kwargs):# 自定义逻辑return {"metric_name": value} | 执行实际的评估计算,返回字典形式的结果 | def compute(references, predictions):correct = sum([1 for r,p in zip(references, predictions) if r == p])accuracy = correct / len(references)return {"custom_accuracy": accuracy} | 必须接受 references 和 predictions 参数;返回值必须为 dict |
| 输入处理 | references: list, predictions: list | 接收真实标签和模型预测的列表 | # 示例输入refs = [0, 1, 0, 1]preds = [0, 1, 1, 1] | 确保长度一致;支持任意 Python 可序列化类型(str, int, float 等) |
支持 **kwargs | **kwargs 可接收额外参数用于配置 | 提供灵活性,如设置阈值、权重等 | def compute(references, predictions, threshold=0.5):preds_bin = [1 if p > threshold else 0 for p in predictions]# 继续计算... | 建议为所有可选参数设置默认值 |
| 返回格式 | return { "metric_name": scalar 或 list } | 输出应为包含一个或多个指标键值对的字典 | return {"f1_score": 0.85, "precision": 0.82} | 支持返回多个指标;scalar 值通常为 float/int |
5.3 添加输入验证与文档说明
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
实现 _info() 方法 | 返回 datasets.MetricInfo 或 dict 形式的元数据 | 必须包含 name, description, citation 字段;features 描述输入结构 |
features 配置 | 使用 datasets.Features 定义输入 schema | 如 features=datasets.Features({"predictions": datasets.Value("float32"), "references": datasets.ClassLabel(names=["neg", "pos"])}) |
| 输入验证逻辑 | 在 compute 开头添加类型和长度检查 | if len(references) != len(predictions):raise ValueError("References and predictions must have same length") |
| 文档字符串 | 在 .py 文件顶部添加 docstring,并在 _info().description 中提供详细说明 | """Custom Accuracy - 计算分类任务准确率""" |
| 异常处理 | 使用 try-except 包裹关键计算,提供清晰错误信息 | try:result = some_calc(...)except Exception as e:raise RuntimeError(f"Failed to compute metric: {e}") |
5.4 注册并本地加载自定义指标
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 创建指标文件 | 新建目录如 ./my_metrics/accuracy_plus/内含 accuracy_plus.py 文件 | 目录名即为指标名;.py 文件中需包含 _info 和 compute |
编写 accuracy_plus.py | 包含 _info() 和 compute() 函数定义 | 可参考 Hugging Face 官方指标源码结构 |
| 本地加载 | 使用 evaluate.load("./my_metrics/accuracy_plus/") | 路径指向包含 .py 文件的目录;路径末尾斜杠可选 |
| 测试加载 | import evaluatemetric = evaluate.load("./my_metrics/custom_acc/")result = metric.compute(references=[1,0], predictions=[1,1]) | 确保无导入错误;compute 能正常返回结果 |
| 发布到 Hub(可选) | 将目录推送到 Hugging Face Model Hub 的 Dataset 空间 | 之后可通过 evaluate.load("username/metric-name") 全局加载 |
| 缓存行为 | 首次加载后会被缓存,后续修改需清缓存或改名测试 | 缓存路径通常为 ~/.cache/huggingface/evaluate/ |
第 6 章:与 Hugging Face 生态集成
6.1 在 Transformers 中集成 Evaluate
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
作为 compute_metrics 函数 | 将 evaluate 指标封装为函数,传入 Trainer 的 compute_metrics 参数 | 必须定义一个接收 eval_pred 参数的函数,该参数是 EvalPrediction 对象 |
EvalPrediction 结构 | eval_pred.predictions 包含模型输出 logits 或序列,eval_pred.label_ids 包含真实标签 | 通常需对 predictions 做 argmax 或解码处理才能与 label_ids 比较 |
| 封装示例:准确率 | def compute_metrics(eval_pred):predictions, labels = eval_predpredictions = predictions.argmax(axis=-1)accuracy = evaluate.load("accuracy")return accuracy.compute(references=labels, predictions=predictions) | 每次调用都重新 load 会降低性能,建议在外部加载后复用实例 |
| 复用指标实例优化 | accuracy_metric = evaluate.load("accuracy")def compute_metrics(eval_pred):preds, labels = eval_predpreds = preds.argmax(-1)return accuracy_metric.compute(references=labels, predictions=preds) | 避免重复初始化,提升评估效率 |
| 多分类 F1 集成 | f1_metric = evaluate.load("f1")def compute_metrics(eval_pred):preds, labels = eval_predpreds = preds.argmax(-1)return f1_metric.compute(references=labels, predictions=preds, average="weighted") | 注意传递 average 等必要 kwargs |
6.2 在 Datasets 中使用指标
| 方法名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
| map 方法集成 | dataset.map(function, batched=True) | 对 dataset 的每个样本或批次应用评估逻辑 | def eval_batch(batch):refs = batch["labels"]preds = model_predict(batch["inputs"])acc = accuracy.compute(references=refs, predictions=preds)return {"accuracy": [acc["accuracy"]] * len(refs)}results = dataset.map(eval_batch, batched=True) | batched=True 提升处理速度;注意返回值需为列表以匹配 batch 大小 |
| 与 Dataset 一起传递 | 将指标与数据集字段关联进行批量评估 | accuracy = evaluate.load("accuracy")results = accuracy.compute(references=dataset["label"],predictions=pred_list) | dataset["field"] 返回 list,可直接作为输入 | |
| 使用 DatasetDict 评估多集 | 分别对 train/val/test 计算指标 | splits = ["train", "test"]for split in splits:preds = get_preds(dset[split])score = metric.compute(references=dset[split]["label"],predictions=preds)print(f"{split}: {score}") | 便于比较不同数据集上的模型表现 | |
| 性能建议 | 避免在 map 中重复加载指标 | 在 map 外部初始化 metric 实例 | 内部加载会导致每个 batch 都重新读取脚本,严重降低性能 |
6.3 在 TrainingArguments 中启用评估
| 参数名称 | 语法 | 用途 | 代码示例 | 注意事项 |
|---|---|---|---|---|
evaluation_strategy | evaluation_strategy="no" / "steps" / "epoch" | 控制评估触发时机 | from transformers import TrainingArgumentsargs = TrainingArguments(output_dir="./results",evaluation_strategy="epoch",save_strategy="epoch") | "steps" 需配合 eval_steps 使用;"epoch" 每轮结束评估 |
eval_steps | eval_steps=500 | 当 evaluation_strategy="steps" 时,每隔多少步评估一次 | args = TrainingArguments(evaluation_strategy="steps",eval_steps=100,...) | 仅当 evaluation_strategy="steps" 时生效 |
compute_metrics | compute_metrics=function | 指定评估时调用的指标计算函数 | trainer = Trainer(model=model,args=args,train_dataset=train_ds,eval_dataset=eval_ds,compute_metrics=compute_metrics) | 必须提供该函数才能输出评估结果 |
eval_accumulation_steps | eval_accumulation_steps=10 | 控制评估时累计多少步的 logits 后再计算指标 | args = TrainingArguments(..., eval_accumulation_steps=16) | 减少 GPU 显存占用,适合大模型或大数据集 |
| 输出内容 | 训练日志自动打印 metrics | {'eval_loss': ..., 'eval_accuracy': ..., 'eval_f1': ..., 'eval_runtime': ..., 'eval_samples_per_second': ...} | 所有 compute_metrics 返回的键都会加上 eval_ 前缀输出 |
6.4 在 AutoTrain 或 Hub 上发布指标
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 准备指标目录 | 创建包含 .py 文件、README.md、__init__.py 的文件夹 | 结构如:my_metric/├── my_metric.py├── README.md└── __init__.py |
| 编写 README.md | 包含指标名称、描述、用法示例、引用文献等 | 使用 Markdown 格式;建议包含 load 和 compute 的代码示例 |
| 登录 Hugging Face | huggingface-cli login | 输入 token 完成认证 |
| 推送到 Dataset 空间 | 使用 huggingface_hub 库或 Web 上传 | api = HfApi()api.create_repo("your-username/my-metric", repo_type="dataset")api.upload_folder(folder_path="./my_metric",repo_id="your-username/my-metric",repo_type="dataset") |
| 全局加载方式 | 发布后可通过用户名/仓库名加载 | metric = evaluate.load("your-username/my-metric") |
| AutoTrain 兼容性 | AutoTrain 目前主要支持内置指标 | 自定义指标需手动集成到训练脚本中 |
| 版本控制 | 支持 Git 提交历史和 revision 加载 | evaluate.load("user/metric", revision="v1.0.0") |
第 7 章:实战案例分析
7.1 在文本分类任务中评估模型性能
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 加载分类指标 | 使用 evaluate 加载 accuracy、f1、precision、recall | 建议组合使用,全面评估模型表现 |
| 数据准备 | 提取真实标签和模型预测结果为列表格式 | predictions 需经过 argmax 或 label decoding 处理 |
| 计算多指标 | 通过 combine 同时计算多个指标 | metrics = evaluate.combine(["accuracy", "f1", "precision", "recall"])results = metrics.compute(references=labels,predictions=preds,average="weighted") |
| 单独计算 F1(多分类) | 分别计算各类别的 F1 分数 | f1_per_class = evaluate.load("f1")result = f1_per_class.compute(references=labels, predictions=preds,average=None, labels=[0,1,2])print(result["f1"]) # [0.8, 0.75, 0.82] |
| 输出结果 | 打印或记录评估结果字典 | {'accuracy': 0.88, 'f1': 0.87, 'precision': 0.86, 'recall': 0.88} |
7.2 在机器翻译任务中使用 BLEU 与 METEOR
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 准备翻译结果 | 获取模型生成的翻译文本列表(predictions)和参考翻译列表(references) | references 应为嵌套列表,每个样本可有多个参考答案 |
| 加载 BLEU 指标 | 使用 evaluate.load("bleu") | 需安装 sentencepiece 支持;自动处理标记化 |
| 计算 BLEU 分数 | 调用 compute 方法传入 predictions 和 references | bleu = evaluate.load("bleu")result = bleu.compute(predictions=preds_list,references=refs_nested_list) |
| 加载 METEOR 指标 | 使用 evaluate.load("meteor") | 需安装 nltk:pip install nltk |
| 计算 METEOR 分数 | meteor = evaluate.load("meteor")score = meteor.compute(predictions=preds,references=refs) | METEOR 更关注语义匹配,通常与人工评价更一致 |
| 对比分析 | 比较 BLEU 和 METEOR 结果 | 若 BLEU 高但 METEOR 低,可能生成文本 n-gram 匹配好但语义偏差大 |
7.3 在摘要任务中综合使用 ROUGE 指标
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 加载 ROUGE 指标 | evaluate.load("rouge") | 支持 rouge1, rouge2, rougeL, rougeLsum |
| 准备摘要数据 | 获取生成摘要(predictions)和参考摘要(references)列表 | 注意处理换行符;长文本建议使用 rougeLsum |
指定 rouge_types | 明确计算哪些 ROUGE 变体 | rouge = evaluate.load("rouge")results = rouge.compute(predictions=gen_sums,references=ref_sums,rouge_types=["rouge1", "rouge2", "rougeL", "rougeLsum"]) |
| 启用词干提取 | use_stemmer=True 提升语义匹配 | results = rouge.compute(..., use_stemmer=True) |
| 解析输出结果 | 返回字典包含各指标的 precision, recall, fmeasure | {'rouge1': 0.45, 'rouge2': 0.23, 'rougeL': 0.40, 'rougeLsum': 0.41} |
| 结果解读 | rouge1 > rouge2 通常正常;若 rougeL 远低于 rouge1 可能连贯性差 | 结合人工检查判断摘要质量 |
7.4 构建端到端评估流水线
| 操作步骤 | 操作细节 | 注意事项 |
|---|---|---|
| 定义评估函数 | 封装加载、计算、输出全过程 | def evaluate_model(model, tokenizer, dataset, metrics_config):preds = []for example in dataset:input_text = example["text"]pred = generate(model, tokenizer, input_text)preds.append(pred)results = {}for name in metrics_config:metric = evaluate.load(name)results.update(metric.compute(...))return results |
| 集成多任务评估 | 根据任务类型动态选择指标 | if task == "classification":metrics = ["accuracy", "f1"]elif task == "summarization":metrics = ["rouge"] |
| 批量处理优化 | 使用 dataset.map 或 DataLoader 批量生成预测 | from torch.utils.data import DataLoaderdataloader = DataLoader(dataset, batch_size=8, collate_fn=collate_fn) |
| 缓存中间结果 | 保存 predictions 到文件,避免重复生成 | import jsonwith open("predictions.json", "w") as f:json.dump(preds, f) |
| 输出结构化报告 | 将结果保存为 JSON 或 CSV | import pandas as pddf = pd.DataFrame([results])df.to_csv("eval_results.csv") |
| 自动化脚本示例 | 创建 eval_pipeline.py 脚本,支持命令行参数 | python eval_pipeline.py --model bert-base-uncased --data test.json --task classification |
第 8 章:常见问题与最佳实践
8.1 常见错误与调试技巧
| 问题现象 | 可能原因 | 解决方案 | 调试技巧 |
|---|---|---|---|
ModuleNotFoundError 或 FileNotFoundError when loading metric | 指标名称拼写错误、网络问题导致下载失败、本地路径不存在 | 1. 检查指标名称是否正确(如 "bleu" 而非 "BLEU")2. 确保网络畅通或使用已缓存的指标 3. 验证本地路径包含有效的 .py 文件 | 使用 evaluate.list() 查看可用指标列表;通过 evaluate.inspect("metric_name") 验证是否存在 |
ValueError: References and predictions must have the same shape | references 和 predictions 列表长度不一致 | 检查数据预处理流程,确保两者样本数相同 | 打印 len(references) 和 len(predictions) 进行对比;使用断言 assert len(refs) == len(preds) |
RuntimeError: Can't initialize tokenizer(如用于 BLEU) | 缺少依赖库(如 sentencepiece, sacrebleu) | 安装缺失依赖:pip install sentencepiecepip install sacrebleupip install bert-score(用于 BERTScore) | 根据错误信息判断所需包;建议在虚拟环境中统一管理依赖 |
compute 返回 NaN 或异常值 | 输入包含非法值(None, NaN)、标签越界、average 参数不匹配 | 清洗输入数据,移除空值;验证标签范围;检查 average 是否适用于当前任务(如 binary 不用于多类) | 在 compute 前添加数据验证逻辑,例如:assert all(r is not None for r in references) |
| 自定义指标加载失败 | 目录结构错误、缺少 _info() 函数、compute 接口不匹配 | 确保目录下有 .py 文件且导出函数正确;参考官方指标结构 | 使用 evaluate.inspect("./path/to/metric") 查看加载详情和源码 |
8.2 指标选择指南(按任务类型)
| 任务类型 | 推荐指标 | 说明 | 备选/补充指标 |
|---|---|---|---|
| 二分类 | accuracy, f1, precision, recall | accuracy 衡量整体正确率;f1 综合 precision 和 recall,适合不平衡数据 | 若关注假阳性用 precision,关注漏检用 recall;AUC-ROC 需额外计算 |
| 多分类 | accuracy, f1 (weighted/macro) | accuracy 仍有效;f1 使用 weighted(按类别频次加权)或 macro(各类平等) | 可结合混淆矩阵分析类别级表现 |
| 多标签分类 | f1 (micro/macro), precision, recall | micro 更关注总体性能,macro 关注各类平均 | 注意 labels 应为 multi-hot 编码 |
| 回归任务 | mae, mse/rmse, r_squared | MAE 易解释;MSE 对异常值敏感;R² 衡量拟合优度 | 根据业务需求选择,如金融预测常用 RMSE |
| 机器翻译 | sacrebleu, meteor, ter | sacrebleu 是标准 BLEU 实现;METEOR 更贴近人工评价;TER 越低越好 | 建议组合使用,避免单一指标偏差 |
| 文本摘要 | rouge1, rouge2, rougeL, rougeLsum | rouge1 关键词覆盖;rougeL 句子结构连贯性;长摘要用 rougeLsum | 可辅以 BERTScore 提升语义评估质量 |
| 语义相似度 | bertscore, cosine_similarity(嵌入) | BERTScore 基于上下文嵌入,优于 n-gram 匹配 | 需注意模型选择(如 'roberta-large')影响结果 |
| 生成任务(通用) | bleu, rouge, meteor, bertscore | 浅层匹配用 BLEU/ROUGE,深层语义用 METEOR/BERTScore | 结合人工评估更可靠 |
8.3 性能瓶颈分析与优化建议
| 瓶颈来源 | 分析方法 | 优化策略 |
|---|---|---|
| 频繁加载指标 | 多次调用 evaluate.load() 而非复用实例 | ✅ 复用指标对象:在循环外初始化 metric = evaluate.load("f1"),内部仅调用 metric.compute() |
| 逐样本调用 compute | 在 for 循环中对每个样本单独计算 | ✅ 批量处理:收集所有 predictions 和 references 后一次性传入 compute |
| 大模型嵌入计算慢(如 BERTScore) | BERTScore 使用大型语言模型计算相似度 | ✅ 降低 batch size ✅ 使用较小模型(如 distilbert-base-uncased)✅ 启用 GPU 加速 |
| I/O 阻塞 | 频繁读写磁盘或网络请求 | ✅ 启用缓存(默认已开启) ✅ 离线模式使用本地指标 ✅ 预加载数据到内存 |
| 高维输出解析开销 | 返回大量中间结果或日志 | ✅ 仅返回必要指标 ✅ 使用 dict.update() 合并结果减少调用次数 |
| CPU 密集型计算(如 stemmer) | ROUGE 启用 use_stemmer 导致变慢 | ✅ 若非必需可关闭 stemmer ✅ 使用多进程并行处理多个样本集 |
8.4 版本兼容性与 API 变更记录
| 版本范围 | 主要变更 | 影响 | 迁移建议 |
|---|---|---|---|
evaluate < 0.2 → >= 0.2 | 引入 combine_metrics 功能;重构部分指标接口 | 旧脚本可能无法直接运行 combine | 更新代码使用新语法 evaluate.combine(["acc", "f1"]) |
evaluate < 0.3 → >= 0.3 | EvaluationModule 类逐步弃用;推荐函数式接口 | 继承 EvaluationModule 的自定义指标可能警告 | 改为直接实现 .py 文件中的 compute 和 _info |
evaluate < 0.4 → >= 0.4 | 改进缓存机制,默认使用 HF_HOME | 缓存路径变化可能导致重复下载 | 设置 HF_HOME 或清理旧缓存 ~/.cache/huggingface/evaluate |
evaluate >= 0.5 | 支持 config_name 更灵活配置(如 "binary", "multiclass") | 需更新调用方式以利用预设配置 | 使用 evaluate.load("f1", config_name="binary") 替代手动传参 |
evaluate >= 0.6 | 增强与 transformers.Trainer 集成;支持更多 kwargs 透传 | 训练脚本中 compute_metrics 更稳定 | 检查 compute_metrics 函数是否正确接收 eval_pred |
| 当前最新版 (2025) | 统一指标命名规范;增强文档与 inspect 功能 | inspect 输出更详细,便于调试 | 善用 evaluate.inspect("metric_name") 学习用法 |
| 通用建议 | - | - | ✅ 使用 pip show evaluate 查看当前版本✅ 参考 Hugging Face Evaluate 文档 ✅ 在生产环境固定版本号(如 evaluate==0.6.0)避免意外 break |