跳转至

__init__.py — 这是 PRIME 数学评分子包的入口文件

模块路径: verl.utils.reward_score.prime_math

文件概述

这是 PRIME 数学评分子包的入口文件,实现了基于 sympy 符号计算的高精度数学答案验证。它结合了字符串标准化(math_normalize)和符号等价判断(grader.math_equal),能够处理复杂的数学表达式等价判断。

来源于 OpenAI 的 prm800k 项目。

关键代码讲解

1. LaTeX 解析

def _parse_latex(expr: str) -> str:
    """将 LaTeX 表达式转换为 sympy 可读的格式"""
    expr = expr.replace("\\tfrac", "\\frac").replace("\\dfrac", "\\frac")
    expr = latex2text.LatexNodes2Text().latex_to_text(expr)
    expr = expr.replace("√", "sqrt").replace("π", "pi")
    return expr.strip()

使用 pylatexenc 库将 LaTeX 转为纯文本数学表达式。

2. 表达式标准化 _normalize

def _normalize(expr: str) -> str:
    # 去掉 \text{}, $, %, 单位等
    expr = expr.replace("million", "*10^6")
    for unit in ["degree", "cm", "meter", ...]:
        expr = re.sub(f"{unit}(es)?(s)? *(\\^[0-9]+)?", "", expr)
    # 尝试解析 LaTeX
    if "\\" in expr:
        expr = _parse_latex(expr)
    # 混合数处理: "7 3/4" -> "7+3/4"
    expr = _inject_implicit_mixed_number(expr)
    expr = expr.lower()
    return expr

3. sympy 等价判断

@timeout_limit(seconds=10)
def are_equal_under_sympy(ground_truth_normalized, given_normalized):
    expr = f"({ground_truth_normalized})-({given_normalized})"
    if should_allow_eval(expr):
        sympy_diff = _sympy_parse(expr)
        simplified = sympy.simplify(sympy_diff)
        if simplified == 0:
            return True
    return False

核心思路:计算两个表达式的差,如果 sympy 化简后为 0,则等价。有 10 秒超时保护防止 sympy 挂起。

4. 主评分函数

def compute_score(model_output: str, ground_truth: str) -> bool:
    is_matched, extracted_model_output = match_answer(model_output)
    if grade_answer(extracted_model_output, ground_truth):
        return True, True, extracted_model_output
    # 处理 \pi 的情况
    if "\\pi" in extracted_model_output or "\\pi" in ground_truth:
        equivs = [math_equal(extracted_model_output, ground_truth, pi=pi)
                  for pi in [math.pi, 3.14]]
        is_correct = any(equivs)
    else:
        is_correct = math_equal(extracted_model_output, ground_truth, timeout=True)
    return is_correct, format_correctness, extracted_model_output

多层验证策略:先做字符串标准化比较,再用 sympy 符号计算验证。

核心类/函数列表

函数名 作用
compute_score 主评分函数
grade_answer 多层等价判断
match_answer 从模型输出中提取答案
_normalize 表达式标准化
are_equal_under_sympy sympy 符号等价判断
_parse_latex LaTeX 转纯文本
split_tuple 处理元组/区间格式

与其他模块的关系

  • 被 __init__.py 调用,处理 Numina 系列数据集
  • 内部使用 math_normalize 和 grader 子模块
  • 使用 verl.utils.py_functional.timeout_limit 做超时保护

小结

PRIME 数学评分是最复杂的数学评分模块。它通过多层验证策略(字符串比较 → sympy 化简 → 数值近似)来判断答案等价性,能处理各种 LaTeX 变体和数学等价表达。timeout_limit 装饰器防止 sympy 在复杂表达式上挂起。