跳转至

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%)是一个实用的设计考量。