跳转至

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 训练中给"格式正确但答案错误"的输出一点奖励,帮助模型先学会输出格式。