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 的奖励
小结¶
定义了奖励管理器的标准接口,所有奖励计算策略都遵循这个接口。