跳转至

prime.py — PRIME 奖励管理器

文件概述

PRIME 项目的奖励管理器实现(约 190 行),使用多进程并行计算奖励,支持超时控制。

核心类

PrimeRewardManager

@register("prime")
class PrimeRewardManager(AbstractRewardManager):
    def verify(self, data):
        """使用多进程并行验证所有样本"""
        scores = run_reward_scoring(
            self.compute_score,
            completions=sequences_str,
            references=ground_truth,
            tasks=data_sources,
            num_processes=64,  # 64 个子进程并行
        )
        data.batch["acc"] = torch.tensor(scores)

并行计算框架

async def parallel_compute_score_async(evaluation_func, completions, references, tasks, num_processes=64):
    """异步并行计算奖励

    使用 ProcessPoolExecutor 进行多进程并行:
    - 每个样本在独立进程中计算
    - 支持 300 秒超时
    - 异常处理: 超时或错误返回 0 分
    """
    with ProcessPoolExecutor(max_workers=num_processes) as executor:
        tasks_async = [
            single_compute_score(evaluation_func, c, r, t, ei, executor, timeout=300.0)
            for c, r, t, ei in zip(completions, references, tasks, extra_info)
        ]
        results = await asyncio.gather(*tasks_async)

设计特点

  • 多进程隔离: 每个样本在独立进程中评估,防止单个异常影响整体
  • 超时控制: 300 秒超时,防止无限循环的评估函数
  • 进程清理: 使用 psutil 确保所有子进程被正确终止

与其他模块的关系

  • 继承 abstract.py 的 AbstractRewardManager
  • 来源于 PRIME-RL 项目

小结

PRIME 奖励管理器通过多进程并行实现了高吞吐量的奖励计算,适合需要复杂评估逻辑的场景。