gsm8k_interaction.py — verl/interactions/gsm8k_interaction.py¶
文件路径¶
verl/interactions/gsm8k_interaction.py
文件概述¶
Gsm8kInteraction 是面向 GSM8K 数学题的交互环境,用于实现"多轮自我反思"训练。与 Gsm8kTool(LLM 主动提交答案)不同,Gsm8kInteraction 模拟的是一个"老师"角色——它评判 LLM 的回答是否正确,如果不正确就要求 LLM 重新思考。
训练效果:通过多轮交互,LLM 学会了在答错时进行自我反思和纠正。
关键代码讲解¶
1. 开始交互¶
async def start_interaction(self, instance_id=None, ground_truth=None, **kwargs):
if instance_id is None:
instance_id = str(uuid4())
self._instance_dict[instance_id] = {
"response": "",
"ground_truth": ground_truth,
"reward": 0.0,
}
return instance_id
与 Gsm8kTool.create 类似,需要传入标准答案 ground_truth。
2. 生成响应(核心逻辑)¶
async def generate_response(self, instance_id, messages, **kwargs):
# 从对话历史中找到 LLM 最近的回答
content = ""
for i in range(len(messages) - 1, -1, -1):
item = messages[i]
if item.get("role") == "assistant":
content = item.get("content")
break
self._instance_dict[instance_id]["response"] = content
reward = await self.calculate_score(instance_id)
if reward == 1.0:
response = "Your response is correct!"
should_terminate_sequence = True
else:
response = "Your response is incorrect! You need to reflect on your answer and try again."
should_terminate_sequence = False
return should_terminate_sequence, response, reward, {}
执行逻辑:
1. 倒序遍历消息:找到 LLM 最近一次 role="assistant" 的回答。
2. 评分:调用 calculate_score 比对答案。
3. 生成反馈:
- 答对了(reward=1.0):返回鼓励,终止对话。
- 答错了:返回"你的回答不正确,请反思后重试",继续对话。
这就是"多轮自我反思"的核心机制——LLM 会反复尝试直到答对或达到最大轮次。
3. 计算分数¶
async def calculate_score(self, instance_id, **kwargs):
return gsm8k.compute_score(
self._instance_dict[instance_id]["response"],
self._instance_dict[instance_id]["ground_truth"],
method="strict",
format_score=0.0,
score=1.0,
)
注意与 Gsm8kTool 的区别:这里使用 method="strict"(严格匹配),而 Gsm8kTool 使用 method="flexible"(灵活匹配)。
与 Gsm8kTool 的对比¶
| 特性 | Gsm8kTool | Gsm8kInteraction |
|---|---|---|
| 角色 | 被 LLM 主动调用的工具 | 主动评判 LLM 的环境 |
| 交互方式 | LLM 提交答案 | 环境评判并反馈 |
| 评分方法 | flexible(灵活) | strict(严格) |
| 终止控制 | 由 LLM 决定 | 由环境决定 |
| 惩罚机制 | 未改进扣 0.05 分 | 无额外惩罚 |
核心类/函数列表¶
| 类/方法 | 作用 |
|---|---|
Gsm8kInteraction |
GSM8K 交互环境 |
start_interaction |
初始化交互会话 |
generate_response |
评判 LLM 回答,生成反馈 |
calculate_score |
严格模式评分 |
finalize_interaction |
清理实例 |
与其他模块的关系¶
- 继承自
BaseInteraction(base.py)。 - 依赖
verl.utils.reward_score.gsm8k计算分数。 - 与
Gsm8kTool功能相关但角色不同。 - 通过
interaction_registry.py被注册和初始化。
小结¶
Gsm8kInteraction 展示了交互系统的典型用法——模拟一个"老师"角色来训练 LLM 的自我反思能力。答对了终止对话,答错了给出反馈继续,这种机制鼓励 LLM 学会检查和修正自己的答案。