跳转至

abstract.py — 奖励管理器抽象基类

文件概述

定义 AbstractRewardManager 抽象基类。

核心类

AbstractRewardManager(ABC)

class AbstractRewardManager(ABC):
    @abstractmethod
    def __call__(self, data: DataProto, return_dict=False) -> torch.Tensor | dict:
        """计算奖励

        输入: DataProto (包含 prompts, responses, attention_mask 等)
        输出: reward_tensor (shape: [batch, response_len])
             奖励值放在每个序列最后一个有效 token 位置

        如果 return_dict=True, 返回 {"reward_tensor": ...}
        """
        pass

    def _extract_reward_from_rm_scores(self, data, return_dict=False):
        """从已有的奖励模型分数提取奖励

        如果数据中已包含 rm_scores(由奖励模型预先计算),
        直接使用,无需重新计算。
        """

奖励放置方式

prompt: [t1, t2, t3]  response: [t4, t5, t6, <pad>]
reward: [ 0,  0,  0,   0,  0, 0.8,    0]
                             ↑ 奖励放在最后一个有效 token

为什么这样?因为 RL 中奖励是 episodic 的,
一个 episode 结束时才给出总奖励。

与其他模块的关系

  • 被 naive, batch, dapo, prime 继承实现
  • 被训练器调用来计算每个 batch 的奖励

小结

定义了奖励管理器的标准接口,所有奖励计算策略都遵循这个接口。