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 多层验证策略中的第一层(字符串标准化比较)的基础。