跳转至

batch.py — Batch 奖励管理器

文件概述

批量奖励计算管理器,与 Naive 版本功能类似但提供 verify 方法。

核心类

BatchRewardManager

@register("batch")
class BatchRewardManager(AbstractRewardManager):
    def verify(self, data):
        """批量验证所有样本,将结果存为 acc tensor

        与 NaiveRewardManager 的区别:
        - 提供独立的 verify 方法,可以被外部调用
        - 批量处理所有样本
        """
        scores = []
        for i in range(len(data)):
            score = self.compute_score(...)
            scores.append(score)
        data.batch["acc"] = torch.tensor(scores)
        return scores

与其他模块的关系

  • 继承 abstract.py 的 AbstractRewardManager
  • 增加了 verify 方法供训练器在需要时单独调用

小结

BatchRewardManager 在 NaiveRewardManager 基础上增加了批量验证能力。