跳转至

registry.py — 奖励管理器注册表

文件概述

提供奖励管理器的注册和查找机制(约 56 行)。

核心内容

注册表

REWARD_MANAGER_REGISTRY: dict[str, type[AbstractRewardManager]] = {}

def register(name: str):
    """装饰器:注册奖励管理器

    用法:
    @register("naive")
    class NaiveRewardManager(AbstractRewardManager):
        ...
    """
    def decorator(cls):
        REWARD_MANAGER_REGISTRY[name] = cls
        return cls
    return decorator

def get_reward_manager_cls(name: str):
    """根据名称获取奖励管理器类

    已注册的名称: naive, batch, dapo, prime
    """
    return REWARD_MANAGER_REGISTRY[name]

与其他模块的关系

  • 被所有奖励管理器类使用 @register 注册
  • 被训练器根据配置调用 get_reward_manager_cls 获取类

小结

简洁的注册表模式,支持通过配置字符串选择奖励管理器。