ForcePilot/test/gaia_eval/scorer.py

98 lines
2.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""GAIA 答案评分模块
实现 GAIA 官方的 quasi exact match 评分逻辑。
参考: https://huggingface.co/spaces/gaia-benchmark/leaderboard
"""
import re
import string
class GaiaScorer:
"""GAIA 准精确匹配评分器"""
@staticmethod
def normalize_text(text: str) -> str:
"""文本标准化
- 去除首尾空白
- 统一为小写
- 移除冠词 (a, an, the)
- 移除标点
- 合并连续空白
"""
text = text.strip().lower()
# 移除冠词
text = re.sub(r"\b(a|an|the)\b", " ", text)
# 移除标点
text = text.translate(str.maketrans("", "", string.punctuation))
# 合并连续空白
text = re.sub(r"\s+", " ", text).strip()
return text
@staticmethod
def normalize_number(text: str) -> str:
"""数字标准化
- "1,000""1000"
- "3.0""3"
- "$100""100"
- "100%""100"
"""
# 去除货币符号
text = re.sub(r"[$€£¥]", "", text)
# 去除百分号
text = text.rstrip("%")
# 去除千分位逗号
text = text.replace(",", "")
# 判断是否为数字,若是则标准化
try:
num = float(text)
# 如果是整数则去掉 .0
if num == int(num):
return str(int(num))
return str(num)
except ValueError:
return text
@classmethod
def score(cls, prediction: str, gold: str) -> bool:
"""计算单条评估的准精确匹配分数
Args:
prediction: 模型预测答案
gold: 黄金标准答案
Returns:
True 表示匹配False 表示不匹配
"""
if not prediction or not gold:
return False
# 先尝试数字比较
pred_num = cls.normalize_number(prediction.strip())
gold_num = cls.normalize_number(gold.strip())
try:
if float(pred_num) == float(gold_num):
return True
except ValueError:
pass
# 检查是否为列表答案(逗号分隔)
if "," in gold:
pred_items = sorted(cls.normalize_text(item) for item in prediction.split(","))
gold_items = sorted(cls.normalize_text(item) for item in gold.split(","))
if pred_items == gold_items:
return True
# 文本标准化比较
pred_normalized = cls.normalize_text(prediction)
gold_normalized = cls.normalize_text(gold)
return pred_normalized == gold_normalized