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 反复提交相同的错误答案。