跳转至

math_reward.py — 这个文件实现了 MATH 数据集(Hendrycks' MATH)的奖励评分逻辑

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

文件概述

这个文件实现了 MATH 数据集(Hendrycks' MATH)的奖励评分逻辑。MATH 数据集的标准答案使用 LaTeX 的 \boxed{} 格式包裹。评分的核心是:从模型输出中提取 \boxed{} 中的内容,经过 LaTeX 标准化处理后,与标准答案比较。

该模块改编自 EleutherAI 的 lm-evaluation-harness 项目。

关键代码讲解

1. 提取 \boxed{} 中的内容

def last_boxed_only_string(string):
    """找到字符串中最后一个 \\boxed{...} 表达式"""
    idx = string.rfind("\\boxed")
    if idx < 0:
        idx = string.rfind("\\fbox")
        if idx < 0:
            return None
    # 通过计数大括号来找到匹配的右括号
    i = idx
    num_left_braces_open = 0
    while i < len(string):
        if string[i] == "{":
            num_left_braces_open += 1
        if string[i] == "}":
            num_left_braces_open -= 1
            if num_left_braces_open == 0:
                right_brace_idx = i
                break
        i += 1
    return string[idx : right_brace_idx + 1]

通过大括号计数法定位 \boxed{} 的完整范围,处理嵌套大括号的情况。

2. LaTeX 字符串标准化 strip_string

def strip_string(string):
    string = string.replace("\n", "")           # 去换行
    string = string.replace("tfrac", "frac")    # 统一分数表示
    string = string.replace("dfrac", "frac")
    string = string.replace("\\left", "")       # 去括号修饰符
    string = string.replace("\\right", "")
    string = string.replace("^{\\circ}", "")    # 去角度符号
    string = string.replace(" ", "")            # 去空格
    string = fix_sqrt(string)                   # 标准化根号
    string = fix_fracs(string)                  # 标准化分数
    string = fix_a_slash_b(string)              # 将 a/b 转为 \frac{a}{b}
    if string == "0.5":
        string = "\\frac{1}{2}"                 # 特殊处理
    return string

这里做了大量 LaTeX 表达式的标准化处理,确保 \frac12 和 \frac{1}{2} 被视为相同。

3. 等价性判断与评分

def is_equiv(str1, str2, verbose=False):
    """判断两个数学表达式在标准化后是否相等"""
    ss1 = strip_string(str1)
    ss2 = strip_string(str2)
    return ss1 == ss2

def compute_score(solution_str, ground_truth) -> float:
    string_in_last_boxed = last_boxed_only_string(solution_str)
    if string_in_last_boxed is not None:
        answer = remove_boxed(string_in_last_boxed)
        if is_equiv(answer, ground_truth):
            return 1.0
    return 0.0

评分流程:提取 \boxed{} 内容 → 去掉 \boxed{} 包裹 → 标准化 → 与标准答案比较。

核心类/函数列表

函数名 作用
compute_score 主评分函数
last_boxed_only_string 提取最后一个 \boxed{} 表达式
remove_boxed 去除 \boxed{} 包裹
is_equiv 判断两个表达式标准化后是否等价
strip_string LaTeX 字符串标准化
fix_fracs 修复分数格式
fix_sqrt 修复根号格式
fix_a_slash_b 将 a/b 转为 \frac{a}{b}

与其他模块的关系

  • 被 __init__.py 中 default_compute_score 调用,用于 lighteval/MATH 等数据集
  • 与 math_verify.py 是同功能的替代关系(后者使用第三方 math-verify 库,精度更高)
  • math_batch.py 导入本模块的 compute_score 实现批量评分

小结

这个模块的核心挑战在于 LaTeX 数学表达式的标准化。同一个数学答案可以有很多不同的 LaTeX 写法(如 \frac12 vs \frac{1}{2} vs 0.5),strip_string 函数将这些变体统一为标准格式后再比较。这是一种基于字符串的"浅层"等价判断,更深层的符号计算判断由 prime_math 模块实现。