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 基础上增加了批量验证能力。