gsm8k.py — 这个文件实现了 GSM8K 数据集的奖励评分逻辑¶
模块路径: verl.utils.reward_score.gsm8k
文件概述¶
这个文件实现了 GSM8K 数据集的奖励评分逻辑。GSM8K(Grade School Math 8K)是一个小学数学应用题数据集,标准答案格式为 #### 数字。评分的核心思路是:从模型输出中提取数字答案,与标准答案比较。
关键代码讲解¶
1. 答案提取函数 extract_solution¶
_SOLUTION_CLIP_CHARS = 300
def extract_solution(solution_str, method="strict"):
# 优化:只在最后300个字符中搜索(数学答案通常在末尾)
if len(solution_str) > _SOLUTION_CLIP_CHARS:
solution_str = solution_str[-_SOLUTION_CLIP_CHARS:]
if method == "strict":
# 严格模式:必须匹配 "#### 数字" 格式
solutions = re.findall("#### (\\-?[0-9\\.\\,]+)", solution_str)
if len(solutions) == 0:
final_answer = None
else:
final_answer = solutions[-1].replace(",", "").replace("$", "")
elif method == "flexible":
# 灵活模式:匹配最后一个有效数字
answer = re.findall("(\\-?[0-9\\.\\,]+)", solution_str)
final_answer = None
if len(answer) > 0:
for final_answer in reversed(answer):
if final_answer not in ["", "."]:
break
return final_answer
两种提取模式:
- strict(严格):要求模型输出包含 #### 答案 格式,测试模型是否学会了规定格式
- flexible(灵活):在输出中搜索最后一个数字作为答案
2. 评分函数 compute_score¶
def compute_score(solution_str, ground_truth, method="strict", format_score=0.0, score=1.0):
answer = extract_solution(solution_str=solution_str, method=method)
if answer is None:
return 0 # 没有提取到答案,得0分
else:
if answer == ground_truth:
return score # 答案正确,得满分(默认1.0)
else:
return format_score # 格式正确但答案错,得格式分(默认0.0)
评分规则简单明了:
- 没找到答案 → 0 分
- 答案正确 → score(默认 1.0)
- 格式正确但答案错误 → format_score(默认 0.0,可调整来鼓励模型至少输出正确格式)
核心类/函数列表¶
| 函数名 | 作用 |
|---|---|
extract_solution |
从模型输出中提取数字答案 |
compute_score |
比较提取的答案与标准答案,返回奖励分数 |
与其他模块的关系¶
- 被
reward_score/__init__.py中的default_compute_score调用,当data_source == "openai/gsm8k"时使用 - 参考了 ReFT 论文(Reasoning with Reinforced Fine-Tuning)的评分方式
小结¶
GSM8K 的评分是所有评分模块中最简单的一个 -- 纯字符串匹配。它的设计思路是:数学答案通常在输出末尾,用正则表达式提取后做精确匹配。format_score 参数允许在 RL 训练中给"格式正确但答案错误"的输出一点奖励,帮助模型先学会输出格式。