Prhub

#1995 fix(multi-agent): preserve rollout logprobs

原始 PR 作者 Jiang020609 合并时间 2026-06-01 16:49 文件变更 2 提交数 1 评论 0 代码增减 +12 / -1

执行摘要

保留多智能体 rollout logprobs

Issue #1976 报告多智能体示例因缺失 rollout_log_probs 导致 GRPO 优势计算时报错 'NoneType' object is not iterable。根本原因是在 examples/multi_agent/agent_system.py 中请求 return_logprob=True 但只提取 token id,丢弃了对应的 logprob。

可直接合并。修复了 issue 中报告的问题,代码简洁且与主框架的 rollout 实现一致。

讨论亮点

无 review 讨论。

实现拆解

  1. 提取 logprob:在 examples/multi_agent/agent_system.pygenerate_response 函数中,将 output_token_logprobs 中的 logprob(item[0])提取到 new_response_log_probs 列表。
  2. 写入 Sample:初始化 sample.rollout_log_probs(若为 None 则设为空列表),然后拼接 new_response_log_probs,并添加断言确保长度与响应 tokens 一致。
  3. 启用训练参数:在 examples/multi_agent/run-qwen3-30B-A3B-multi-agent.sh 的 GRPO 参数中添加 --use-rollout-logprobs,使训练能够使用保存的 rollout logprobs。
文件 模块 状态 重要度
examples/multi_agent/agent_system.py 多智能体示例 modified 5.95
examples/multi_agent/run-qwen3-30B-A3B-multi-agent.sh 多智能体示例 modified 2.0

关键符号

generate_response

关键源码片段

examples/multi_agent/agent_system.py core-logic

核心修复文件,在 generate_response 中添加了 logprob 提取和写入 Sample 的逻辑。

async def generate_response(args, prompt, key):
    try:
        # ... 前置代码 ...
        payload = {"input_ids": prompt_token_ids, "sampling_params": current_sampling_params, "return_logprob": True}
        output = await post(url, payload)
​
        # 提取新响应 tokens 和对应的 logprobs
        if "output_token_logprobs" in output["meta_info"]:
            output_token_logprobs = output["meta_info"]["output_token_logprobs"]
            # item[1] 是 token id, item[0] 是 logprob
            new_response_tokens = [item[1] for item in output_token_logprobs]
            new_response_log_probs = [item[0] for item in output_token_logprobs]
        else:
            new_response_tokens = []
            new_response_log_probs = []
​
        # 更新 sample tokens
        sample.tokens = sample.tokens + new_response_tokens
        sample.response_length += len(new_response_tokens)
        # 初始化 rollout_log_probs 并追加
        if sample.rollout_log_probs is None:
            sample.rollout_log_probs = []
        sample.rollout_log_probs += new_response_log_probs
        # 断言长度一致,避免静默错误
        assert len(sample.rollout_log_probs) == sample.response_length, (
            f"rollout logprob length mismatch: {len(sample.rollout_log_probs)} logprobs "
            f"vs {sample.response_length} response tokens"
        )
        sample.response = output["text"]
        # ... 后续代码 ...
    except Exception as e:
        print(f"Error generating response: {e}")
        return None

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险较低。变更范围小,仅影响多智能体示例代码,不涉及核心框架。断言可确保 logprobs 长度一致,避免静默错误。

仅影响多智能体示例的用户。修复后 GRPO 训练可正常使用 rollout logprobs,避免崩溃。对现有其他功能无影响。

缺少测试覆盖

关联 Issue

#1976 [Bug] examples/multi_agent can crash with NoneType in GRPO advantage computation because old logprobs are missing

完整报告

参与讨论