跳转至

gsm8k_tool.py — verl/tools/gsm8k_tool.py

文件路径

verl/tools/gsm8k_tool.py

文件概述

Gsm8kTool 是一个面向 GSM8K 数学题数据集的演示工具,用于在 RL 训练中计算 LLM 解答数学题的奖励分数。GSM8K 是一个包含 8000+ 小学数学题的基准数据集,广泛用于评估 LLM 的数学推理能力。

核心功能:LLM 提交一个答案,工具与标准答案比对,返回奖励分数。如果答案没有改进,还会给一个 -0.05 的惩罚,鼓励 LLM 每次提交都要有进步。

关键代码讲解

1. 创建实例

async def create(self, instance_id=None, ground_truth=None, **kwargs):
    if instance_id is None:
        instance_id = str(uuid4())
    if ground_truth is None:
        ground_truth = kwargs.get("create_kwargs", {}).get("ground_truth", None)
    self._instance_dict[instance_id] = {
        "response": "",
        "ground_truth": ground_truth,
        "reward": 0.0,
    }
    return instance_id, ToolResponse()

创建时需要传入 ground_truth(标准答案)。答案可以通过参数直接传入,也可以通过 create_kwargs 间接传入。

2. 执行工具(核心逻辑)

@rollout_trace_op
async def execute(self, instance_id, parameters, **kwargs):
    answer = parameters.get("answer", "")
    if not isinstance(answer, str):
        answer = str(answer)

    if answer.startswith("#### "):
        self._instance_dict[instance_id]["response"] = answer
    else:
        self._instance_dict[instance_id]["response"] = "#### " + answer

    reward = await self.calc_reward(instance_id)
    # penalty for non improved answer submission
    tool_reward = 0.0 if reward > self._instance_dict[instance_id]["reward"] else -0.05
    # update the reward
    self._instance_dict[instance_id]["reward"] = reward

    return ToolResponse(text=f"Current parsed {answer=} {reward=}"), tool_reward, {}

执行流程: 1. 从参数获取 answer。 2. 给答案加上 "#### " 前缀(GSM8K 的标准格式)。 3. 计算与标准答案的比对奖励。 4. 惩罚机制:如果新答案的奖励不比之前高,返回 -0.05 的惩罚分。这鼓励 LLM 每次调用工具时都提交更好的答案。 5. 返回当前答案和奖励信息给 LLM。

3. 计算奖励

async def calc_reward(self, instance_id, **kwargs):
    return gsm8k.compute_score(
        self._instance_dict[instance_id]["response"],
        self._instance_dict[instance_id]["ground_truth"],
        method="flexible",
        format_score=0.0,
        score=1.0,
    )

使用 gsm8k.compute_score 函数比对答案: - method="flexible":灵活匹配模式。 - format_score=0.0:格式不正确时得 0 分。 - score=1.0:答案正确时得 1 分。

核心类/函数列表

类/方法 作用
Gsm8kTool GSM8K 数学题奖励计算工具
create 创建实例并设置标准答案
execute 接收 LLM 的答案,计算奖励
calc_reward 调用 gsm8k 评分函数比对答案
release 释放实例

与其他模块的关系

  • 继承自 BaseTool。
  • 依赖 verl.utils.reward_score.gsm8k 模块计算分数。
  • 类似于 Geo3kTool(几何题工具),结构几乎一致。
  • 与 Gsm8kInteraction(interactions 模块)对比:Gsm8kTool 是作为工具使用(LLM 主动调用),Gsm8kInteraction 是作为交互环境使用(环境主动评判)。

小结

Gsm8kTool 是一个简单但典型的工具实现示例,展示了如何在 RL 训练中通过工具为 LLM 提供即时反馈。它的惩罚机制是一个有趣的设计——通过扣分来防止 LLM 反复提交相同的错误答案。