grader.py — 这个文件是 PRIME 数学评分的核心评判器¶
模块路径: verl.utils.reward_score.prime_math.grader
文件概述¶
这个文件是 PRIME 数学评分的核心评判器,实现了 math_equal 函数,它是最全面的数学等价判断函数。代码整合自多个开源项目(Hendrycks MATH、ToRA、CRITIC、OpenAI prm800k),支持数值比较、字符串比较、符号计算比较,以及区间、矩阵、元组等复杂数学对象的比较。
关键代码讲解¶
1. 数字检测与标准化¶
def is_digit(s):
"""判断字符串是否为数字,处理千分位逗号"""
try:
if "{,}" in str(s):
num = float(str(s).replace("{,}", ""))
return True, num
num = float(str(s).replace(",", ""))
return True, num
except ValueError:
return False, None
def normalize(answer, pi):
"""处理美元符号、百分号、进制、pi 等"""
if isinstance(answer, str) and bool(re.match(r"\$\d+(\.\d+)?", answer)):
return answer[1:] # 去掉 $
answer = handle_base(answer) # 处理 _base 记法
answer = handle_pi(answer, pi) # 将 \pi 替换为数值
return answer
2. 核心等价判断 math_equal¶
def math_equal(prediction, reference, include_percentage=True, tolerance=1e-4, timeout=10.0, pi=math.pi):
prediction = normalize(prediction, pi)
reference = normalize(reference, pi)
# 第0层: 字符串比较
if prediction.strip().lower() == reference.strip().lower():
return True
# 第1层: 数值比较
if is_digit(prediction)[0] and is_digit(reference)[0]:
gt_result = [reference/100, reference, reference*100] if include_percentage else [reference]
for item in gt_result:
if isclose(item, prediction, rel_tol=tolerance):
return True
return False
# 第2层: 符号比较
prediction = format_intervals(prediction) # 处理 Interval() 格式
# 处理 [a,b] vs [c,d] 的元素级比较
# 处理 Point() 和 Matrix 的比较
...
return symbolic_equal(prediction, reference, tolerance, timeout)
三层渐进式比较,从简单到复杂。
3. 符号等价判断¶
def symbolic_equal(a, b, tolerance, timeout=10.0):
def _parse(s):
for f in [parse_expr, parse_latex]:
try:
with timeout_limit(seconds=timeout):
return f(s)
except TimeoutError:
continue
return s
a = _parse(a)
b = _parse(b)
# 尝试化简 a - b == 0
try:
with timeout_limit(seconds=timeout):
if simplify(a - b) == 0:
return True
except:
pass
# 尝试数值近似比较
try:
with timeout_limit(seconds=timeout):
if isclose(N(a), N(b), rel_tol=tolerance):
return True
except:
pass
return False
先用 sympy 的 simplify 化简差值,失败则用 N() 数值近似比较。
4. 区间格式处理¶
def format_intervals(prediction):
# 将 Interval(a, b) → [a, b]
# Interval.Ropen(a, b) → [a, b)
# Interval.Lopen(a, b) → (a, b]
# Interval.open(a, b) → (a, b)
...
核心类/函数列表¶
| 函数名 | 作用 |
|---|---|
math_equal |
核心数学等价判断函数 |
symbolic_equal |
符号计算等价判断 |
normalize |
答案标准化(处理 $, %, pi, 进制) |
is_digit |
数字检测 |
format_intervals |
区间格式转换 |
handle_pi |
处理 pi 的数值替换 |
与其他模块的关系¶
- 被
prime_math/__init__.py导入并使用 - 使用
verl.utils.py_functional.timeout_limit做超时保护 - 使用 sympy 的
parse_expr、parse_latex、simplify等功能
小结¶
这是整个 reward_score 体系中最复杂的文件,实现了涵盖数值、字符串、符号计算三个层次的数学等价判断。多个超时保护确保 sympy 不会在复杂表达式上无限运行。百分号的特殊处理(答案 50 等价于 50%)是一个实用的设计考量。