跳转至

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 学会检查和修正自己的答案。