跳转至

math_normalize.py — 这个文件实现了 LaTeX 数学表达式的标准化处理

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

文件概述

这个文件实现了 LaTeX 数学表达式的标准化处理,来源于 Hendrycks' MATH 数据集的评分代码(通过 OpenAI prm800k)。它与 math_reward.py 中的 strip_string 功能类似,但作为独立模块被 prime_math 使用。

关键代码讲解

主函数 normalize_answer

def normalize_answer(answer: Optional[str]) -> Optional[str]:
    if answer is None:
        return None
    answer = answer.strip()
    # 去掉 \text{} 包裹
    m = re.search(r"^\\text\{(?P<text>.+?)\}$", answer)
    if m is not None:
        answer = m.group("text").strip()
    return _strip_string(answer)

内部标准化函数 _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 = _remove_right_units(string)       # 去单位
    string = _fix_sqrt(string)                 # 标准化根号
    string = string.replace(" ", "")           # 去空格
    string = _fix_fracs(string)                # 标准化分数
    if string == "0.5":
        string = "\\frac{1}{2}"                # 特殊处理
    string = _fix_a_slash_b(string)            # a/b → \frac{a}{b}
    return string

处理的 LaTeX 变体包括: - \tfrac / \dfrac → \frac - \sqrt3 → \sqrt{3} - \frac12 → \frac{1}{2} - 1/2 → \frac{1}{2} - 0.5 → \frac{1}{2}

核心类/函数列表

函数名 作用
normalize_answer 主标准化入口
_strip_string 综合标准化处理
_fix_fracs 修复分数表示
_fix_sqrt 修复根号表示
_fix_a_slash_b 斜杠分数转 LaTeX 分数
_remove_right_units 移除单位

与其他模块的关系

  • 被 prime_math/__init__.py 的 grade_answer 调用
  • 功能与 math_reward.py 中的 strip_string 高度重叠,但来自不同的代码源

小结

这是一个纯字符串处理的工具模块,将各种 LaTeX 写法统一为标准格式。它是 prime_math 多层验证策略中的第一层(字符串标准化比较)的基础。