ML-Agents On/Off-Policy Trainer 架构解析:PPO 与 SAC 训练器源码级指南
2026/9/20 14:23:22 网站建设 项目流程
  • 人工智能
  • 强化学习
  • 深度学习
  • 机器学习
  • 游戏开发
  • AI 应用

【免费下载链接】ml-agents

The Unity Machine Learning Agents Toolkit (ML-Agents) is an open-source project that enables games and simulations to serve as environments for training intelligent agents using deep reinforcement learning and imitation learning.

项目地址:https://gitcode.com/gh_mirrors/ml/ml-agents
点击查看免费下载

本文以 Unity ML-Agents Toolkit 仓库中 Python-On-Off-Policy-Trainer-Documentation.md 为骨架,结合ml-agents/mlagents/trainers下训练器实际源码,系统讲解 on-policy(PPO)与 off-policy(SAC)训练器的类层级、核心 API、更新流程与配置参数。读完本文,你将理解 Trainer 体系的继承关系、训练循环中"收集轨迹—更新策略—发布新策略"的完整链路,并能看懂并编写对应训练器的 YAML 配置文件。

一、训练器体系总览:从 Trainer 到 PPO/SAC 的三级继承

从文档与源码结构看,ML-Agents 的训练器(Trainer)采用三级抽象继承:

Trainer (abc.ABC) # ml-agents/mlagents/trainers/trainer/trainer.py └── RLTrainer (Trainer) # 使用 Reward Signals 的训练器基类 ├── OnPolicyTrainer (RLTrainer) # PPO 算法实现 └── OffPolicyTrainer (RLTrainer) # SAC 算法实现
  • trainer.py 定义了所有训练器共有的抽象基类Trainer
  • rl_trainer.py 中RLTrainer是"使用奖励信号(Reward Signals)的训练器"的基类,文档原文明确"RLTrainer(Trainer):This class is the base class for trainers that use Reward Signals";
  • on_policy_trainer.py 中OnPolicyTrainer是 PPO 算法的实现(源码头注释引用 arxiv.org/abs/1707.06347);
  • off_policy_trainer.py 中OffPolicyTrainer是 SAC 算法的实现(源码头注释引用 arxiv.org/abs/1801.01290),并支持离散动作与循环网络。

在运行层面,mlagents-learn根据配置文件中每个 behavior 的trainer_type字段创建对应训练器:ppo映射到 ppo/trainer.py(继承 OnPolicyTrainer),sac映射到 sac/trainer.py(继承 OffPolicyTrainer),多智能体场景还有基于 on-policy 的 poca/trainer.py` 类签名为:

class Trainer(abc.ABC)

__init__签名与参数含义:

def __init__( self, brain_name: str, # Brain(behavior)名称,即要训练的行为名 trainer_settings: TrainerSettings, # 训练器参数(TrainerSettings 对象) training: bool, # 是否处于训练模式 load: bool, # 是否加载已保存的模型 artifact_path: str, # 训练产物(模型、检查点)的存储目录 reward_buff_cap: int = 1, # 奖励缓冲最大容量(默认 1) ):

文档与源码(trainer.py)共同定义了它"负责收集经验并训练神经网络模型"的职责。其核心属性与方法如下:

成员类型说明
stats_reporterproperty返回与该训练器关联的 StatsReporter,用于向 TensorBoard 输出统计
parametersproperty返回TrainerSettings类型的训练器参数
get_max_stepsproperty返回最大训练步数,用于判断何时停止训练
get_stepproperty返回训练器已执行的步数
threadedproperty是否在线程中运行训练器。True允许训练器在环境采样的同时更新策略;False则强制严格的 on-policy 更新(即采样时不更新策略)
should_still_trainproperty是否应继续训练。源码实现为self.is_training and self.get_step <= self.get_max_steps,即未训练或达到 max_steps 时返回 False
reward_bufferproperty返回Deque[float]类型的奖励缓冲,保存最近若干个已完成 episode 的累计奖励
save_model()abstractmethod保存与该训练器关联的一个或多个策略的模型文件
end_episode()abstractmethodepisode 结束信号,必须重置缓冲,仅在 Academy 重置时调用
create_policy()abstractmethod创建 Policy 对象
add_policy()abstractmethod将策略添加到训练器
get_policy(name_behavior_id)方法按完整 behavior 名获取关联策略
advance()abstractmethod推进训练器:典型实现是从所有已订阅的轨迹队列(trajectory_queues)取出轨迹并用其中的 step 更新策略,必要时把新策略推送到策略队列(policy_queues
publish_policy_queue()方法注册一个策略队列,训练器更新策略后向该队列发布新策略
subscribe_trajectory_queue()方法注册一个轨迹队列,训练器从该队列摄取 Trajectory

threaded属性是区分严格 on-policy 训练的关键开关:当为False时,训练器与环境采样同步推进,策略更新发生在采样间隙,保证用于更新的数据全部来自当前策略;当为True时,训练器在独立线程中边采样边更新,吞吐更高但策略更新时可能有部分轨迹来自旧策略。

三、RLTrainer:奖励信号与训练循环的中枢

rl_trainer.py 中RLTrainer(Trainer)是所有使用奖励信号训练器的基类。文档列出的核心方法:

def end_episode(self) -> None: ... @abc.abstractmethod def create_optimizer(self) -> TorchOptimizer: ... def save_model(self) -> None: ... def advance(self) -> None: ...

各方法在源码中的实际语义:

  • end_episode:将collected_rewards中所有 reward signal 的累计奖励清零。collected_rewards是一个"奖励信号名 → agent_id → 累计奖励"的字典,其中"environment"条目始终保留(环境奖励必须上报 TensorBoard,无论配置了哪些奖励信号)。
  • create_optimizer:抽象方法,返回TorchOptimizer对象,由子类(PPO/SAC)各自实现。
  • save_model:保存与训练器关联的策略。实现会先调用_checkpoint()生成检查点,再通过TorchModelSaver.copy_final_model()复制出最终的.onnx模型文件,并调用ModelCheckpointManager.track_final_checkpoint()登记最终检查点。
  • advance:训练循环的主入口(rl_trainer.py)。它先遍历所有trajectory_queues,消费队列中的轨迹(_process_trajectory);随后在should_still_train_is_ready_update()为真时调用_update_policy(),若更新成功则把更新后的策略放入所有policy_queues,供环境侧 AgentProcessor 取用。注释特别说明:每次最多抓取队列最大长度数量的轨迹,确保队列中的轨迹是 on-policy 的。

此外,RLTrainer 还管理检查点与摘要的写入节奏:_maybe_save_model()checkpoint_interval间隔保存检查点,_maybe_write_summary()summary_freq间隔写 TensorBoard 摘要,二者都确保"在更新步写入而不是采样过程中写入"。_increment_step()同步推进训练器步数与策略步数。

四、OnPolicyTrainer:PPO 算法的落地实现

4.1 类定义与构造参数

class OnPolicyTrainer(RLTrainer)

文档明确指出"The PPOTrainer is an implementation of the PPO algorithm",即 OnPolicyTrainer 就是 PPO 训练器。其__init__签名:

def __init__( self, behavior_name: str, # 与训练器配置关联的 behavior 名 reward_buff_cap: int, # 奖励缓冲中追踪的最大奖励历史 trainer_settings: TrainerSettings, # 训练器参数 training: bool, # 是否处于训练模式 load: bool, # 是否加载模型 seed: int, # 模型初始化使用的随机种子 artifact_path: str, # 训练产物存储目录 ):

源码中(on_policy_trainer.py)构造器将trainer_settings.hyperparameters强制转换为OnPolicyHyperparamSettings,并保存seed、预留policyoptimizer字段。其职责是"收集经验并训练一个 on-policy 模型"。

4.2 add_policy:注册策略并初始化优化器

def add_policy(self, parsed_behavior_id: BehaviorIdentifiers, policy: Policy) -> None:

实现要点:

  1. 若已有策略存在,则输出警告:"你的环境包含多个 team,但该训练器不支持对抗游戏,如需训练对抗游戏请启用 self-play";
  2. 保存策略到self.policyself.policies[behavior_id]
  3. 调用create_optimizer()创建优化器,并为每个 reward signal 初始化累计奖励计数器;
  4. 通过model_saver.register()注册策略与优化器,然后initialize_or_load()(按load标志决定初始化新模型或加载已有模型);
  5. policy.get_current_step()恢复训练步数,保证断点续训时步数连续。

4.3 _update_policy:PPO 的 minibatch 更新流程

这是 on-policy 训练的核心(on_policy_trainer.py):

  1. 就绪判断_is_ready_update()检查update_buffer.num_experiences > hyperparameters.buffer_size,即缓冲中经验数超过buffer_size才触发更新。
  2. batch 对齐序列长度batch_size = batch_size - batch_size % policy.sequence_length,且保证至少一个序列(max(batch_size, sequence_length)),因为训练时要重塑为batch_size × sequence_length张量;n_sequencesbatch_size / sequence_length
  3. 优势标准化:读取缓冲中的 ADVANTAGES,做 z-score 标准化((advantages - mean) / (std + 1e-10))。
  4. 多 epoch 小批量循环:按num_epoch轮次,每轮对update_buffer.shuffle(sequence_length=...)打乱,然后切出max_num_batch个 minibatch,逐个调用optimizer.update(minibatch, n_sequences)optimizer.update_reward_signals(minibatch),把统计量汇总到batch_update_stats
  5. 统计上报:对每个统计项取均值写入 stats reporter。
  6. 行为克隆:若配置了bc_module(行为克隆),额外调用bc_module.update()并上报其统计。
  7. 清空缓冲_clear_update_buffer()重置 update buffer,进入下一轮数据收集。

4.4 与其他 on-policy 变体的关系

从源码结构看,PPO 与 POCA 都是 on-policy 家族:OnPolicyTrainer作为通用 on-policy 基类被 ppo/trainer.py 与 poca/trainer.py 复用,区别在于各自创建不同的优化器与策略。这正是"on-policy"训练的统一特征:策略更新所消耗的经验必须来自当前策略的采样,因此每次更新后缓冲会被清空重建。

五、OffPolicyTrainer:SAC 算法的落地实现

5.1 类定义与构造参数

class OffPolicyTrainer(RLTrainer)

文档指出"The SACTrainer is an implementation of the SAC algorithm, with support for discrete actions and recurrent networks"——即 off-policy 训练器就是 SAC 实现,且同时支持离散动作与循环(recurrent)网络。其__init__签名与 OnPolicyTrainer 完全一致(behavior_name、reward_buff_cap、trainer_settings、training、load、seed、artifact_path)。

构造器中(off_policy_trainer.py)额外初始化了 off-policy 特有的节奏参数:

self.update_steps = 1 # 策略更新次数计数 self.reward_signal_update_steps = 1 # 奖励信号更新次数计数 self.steps_per_update = hyperparameters.steps_per_update self.reward_signal_steps_per_update = hyperparameters.reward_signal_steps_per_update self.checkpoint_replay_buffer = hyperparameters.save_replay_buffer

5.2 经验回放缓冲:save/load replay buffer

off-policy 与 on-policy 最本质的区别在于经验可以重复利用(存入回放缓冲反复采样)。为此 OffPolicyTrainer 重写了模型保存相关方法:

def save_model(self) -> None: ... # 保存最终模型,并顺带保存回放缓冲 def save_replay_buffer(self) -> None: ... # 将更新缓冲保存为 pickle 文件 def load_replay_buffer(self) -> None: ... # 从文件加载最近一次保存的回放缓冲
  • save_model():"Saves the final training model to memory. Overrides the default to save the replay buffer."——在调用基类save_model()后,若checkpoint_replay_buffer为真,追加调用save_replay_buffer()
  • save_replay_buffer():把update_buffer序列化到os.path.join(artifact_path, "last_replay_buffer.hdf5"),并打印保存的文件大小日志。
  • load_replay_buffer():从同一路径读取缓冲,load_from_file后记录加载的经验数量。

maybe_load_replay_buffer()会在load标志且checkpoint_replay_buffer为真时尝试加载,若文件缺失(FileNotFoundError)或格式异常(AttributeError)则警告"从零开始"。

5.3 就绪判断与更新节奏

def _is_ready_update(self) -> bool: return ( self.update_buffer.num_experiences >= self.hyperparameters.batch_size and self._step >= self.hyperparameters.buffer_init_steps )

即:缓冲中经验数达到batch_size,且已走完buffer_init_steps步预热(buffer warm-up)。这与 on-policy 的"缓冲超过 buffer_size"判据完全不同。

_update_policy()(off_policy_trainer.py)按步数比例循环更新策略:只要(self._step - buffer_init_steps) / self.update_steps > self.steps_per_update就继续采样 minibatch 并调用optimizer.update(),每次更新update_steps += 1。随后_update_reward_signals()用独立的reward_signal_steps_per_update节奏单独更新奖励信号(模拟 arxiv.org/abs/1809.02925 等论文中"策略更新 N 次后再更新奖励信号 N 次"的做法)。最后,若缓冲超过buffer_size,按BUFFER_TRUNCATE_PERCENT = 0.8的比例截断(truncate(int(buffer_size * 0.8), sequence_length)),避免每次更新都触发大缓冲截断。

5.4 add_policy:恢复训练节奏

与 OnPolicyTrainer 类似,add_policy()注册策略、创建优化器、初始化模型保存器;不同之处在于它还根据当前步数恢复更新节奏计数:

self._step = policy.get_current_step() self.update_steps = int(max(1, self._step / self.steps_per_update)) self.reward_signal_update_steps = int(max(1, self._step / self.reward_signal_steps_per_update))

这样断点续训后,策略更新与奖励信号更新频率能正确衔接历史节奏。

六、配置参数解析:On/Off-Policy 超参数与 TrainerSettings

训练器参数在 settings.py 中以attrs类定义,并在TrainerSettings.structure()中根据trainer_type动态选择超参数类完成 YAML 反序列化(strict_to_cls(d_copy[key], all_trainer_settings[trainer_type])),同时调用check_hyperparam_schedules()校验学习率调度。

6.1 通用超参数 HyperparamSettings

batch_size: int = 1024 # 每次更新使用的经验批量大小 buffer_size: int = 10240 # 经验缓冲容量 learning_rate: float = 3.0e-4 # 学习率 learning_rate_schedule: ScheduleType = ScheduleType.CONSTANT # constant 或 linear

6.2 OnPolicyHyperparamSettings

class OnPolicyHyperparamSettings(HyperparamSettings): num_epoch: int = 3 # 每轮更新中遍历缓冲的 epoch 数

on-policy 训练器额外只有num_epoch(PPO 的多轮 minibatch 迭代次数,见 4.3 节)。ScheduleType枚举目前仅支持CONSTANTLINEAR两种(源码注释留有 lesson 调度的 TODO)。

6.3 OffPolicyHyperparamSettings

class OffPolicyHyperparamSettings(HyperparamSettings): batch_size: int = 128 # 每次采样更新的 minibatch 大小 buffer_size: int = 50000 # 回放缓冲容量 buffer_init_steps: int = 0 # 开始更新前需收集的步数(预热) steps_per_update: float = 1 # 每步平均触发的策略更新次数(1 表示每步更新一次) save_replay_buffer: bool = False # 是否保存回放缓冲(.hdf5) reward_signal_steps_per_update: float = 4 # 奖励信号更新的节奏

注意 off-policy 的默认batch_sizebuffer_size显著不同于 on-policy:SAC 用小批量(128)从大缓冲(50000)中反复随机采样(sample_mini_batch),这正是经验复用的体现;而 PPO 需要缓冲攒够buffer_size(如 12000)后一次性全部用于多轮更新。

6.4 TrainerSettings:训练器顶层配置

settings.py 中TrainerSettings的核心字段:

字段默认值说明
trainer_type"ppo"训练器类型(ppo / sac / poca 等)
hyperparameters按 trainer_type 自动选择超参数对象
network_settingsNetworkSettings()网络结构(hidden_units=128、num_layers=2、memory、normalize 等)
reward_signals{extrinsic: RewardSignalSettings()}奖励信号字典,默认只含外部奖励
checkpoint_interval500000检查点保存间隔
keep_checkpoints5保留的检查点数量
even_checkpointsFalse为 True 时按max_steps / keep_checkpoints均分检查点间隔
max_steps500000最大训练步数
time_horizon64时间视野,轨迹截断步数
summary_freq50000TensorBoard 摘要写入间隔
threadedFalse是否线程化(见 2 节)
init_pathNone初始化模型路径
self_playNone自博弈设置
behavioral_cloningNone行为克隆设置(demo_path、steps、strength 等)

其中reward_signals支持extrinsicgailcuriosityrnd四种类型(见RewardSignalType枚举,settings.py),每个信号由RewardSignalSettings定义gamma(默认 0.99)、strength(默认 1.0)与network_settings;若配置中出现已废弃的encoding_size,系统会警告并将其映射为network_settings.hidden_units

6.5 完整的 YAML 配置示例

仓库 config/ppo/3DBall.yaml 给出了 on-policy 训练器的完整配置:

behaviors: 3DBall: trainer_type: ppo hyperparameters: batch_size: 64 buffer_size: 12000 learning_rate: 0.0003 beta: 0.001 epsilon: 0.2 lambd: 0.99 num_epoch: 3 learning_rate_schedule: linear network_settings: normalize: true hidden_units: 128 num_layers: 2 vis_encode_type: simple reward_signals: extrinsic: gamma: 0.99 strength: 1.0 keep_checkpoints: 5 max_steps: 500000 time_horizon: 1000 summary_freq: 12000

其中betaepsilonlambd是 PPO 特有的优化参数(KL 惩罚系数、裁剪阈值、GAE 系数),属于 ppo/trainer.py 定义的 PPO 优化器专属字段。SAC 的对应示例见 config/sac/3DBall.yaml,其超参数会使用buffer_init_stepssteps_per_updatesave_replay_buffer等 off-policy 字段。更多字段说明可查阅 Training-Configuration-File.md。

七、参数随机化、课程学习与命令行装配

文档后半部分集中描述了与训练器配置紧密相关的辅助设置类,它们是"环境参数 → 训练过程"的桥梁:

7.1 参数随机化采样器

ParameterRandomizationSettings抽象类定义了通过EnvironmentParametersChannel向环境下发采样器设置的apply(key, env_channel)抽象方法,并借助structure/unstructure静态方法与 cattrs 注册钩子完成 YAML 与对象的互转。文档列出的四个具体采样器:

  • ConstantSettings:常量采样器,value字段,apply调用env_channel.set_float_parameter(key, value);YAML 中直接写一个 float 即可触发(structure会把 float/int 转成ConstantSettings)。
  • UniformSettings:均匀采样器,min_value(默认 0.0)、max_value(默认 1.0),校验min_value <= max_valueapply调用set_uniform_sampler_parameters(key, min, max, seed)
  • GaussianSettings:高斯采样器,mean(默认 1.0)、st_dev(默认 1.0),apply调用set_gaussian_sampler_parameters(key, mean, st_dev, seed)
  • MultiRangeUniformSettings:多区间均匀采样器,intervals[min, max]列表,校验每个区间恰好两个值且 min ≤ max,apply调用set_multirangeuniform_sampler_parameters(key, intervals, seed)

对应枚举ParameterRandomizationType提供uniformgaussianmultirangeuniformconstant四类,配置中必须同时给出sampler_typesampler_parameters,否则抛出TrainerConfigError

7.2 课程学习:CompletionCriteria 与 Lesson

  • CompletionCriteriaSettings:判断下一课(lesson)是否开始的依据,字段包括behavior(参照的 behavior 名)、measureprogressreward,默认 reward)、min_lesson_length(最小 episode 数,默认 0)、signal_smoothing(奖励平滑,默认 True)、threshold(阈值,measure 为 progress 时必须介于 0~1)、require_reset。核心方法need_increment(progress, reward_buffer, smoothing)返回(是否进入下一课, 新的平滑值):按 reward 度量时,若启用平滑则measure = 0.25 * smoothing + 0.75 * measure再与阈值比较。
  • Lesson:单个课程的数据结构,包含环境参数名name、采样器valuecompletion_criteria;若 completion_criteria 为 None 则是课程链中的最后一课。
  • EnvironmentParameterSettings:一个环境参数按顺序排列的课程列表(curriculum: List[Lesson])。structure会校验课程链:非末课必须有completion_criteria,末课不得携带completion_criteria(带了只告警忽略);若配置不是带curriculum的映射,则视为单课课程,直接用采样器配置包装成唯一的 Lesson。

7.3 检查点与运行选项

  • CheckpointSettings:封装命令行级检查点选项(run_idinitialize_fromload_modelresumeforcetrain_modelinferenceresults_dir),提供write_pathmaybe_init_pathrun_logs_dir路径属性。prioritize_resume_init()解决冲突:若命令行同时给了resumeinitialize_from,优先 resume 并告警;仅 YAML 同时设置时同样优先 resume。
  • RunOptionsmlagents-learn运行时所有选项的汇总(behaviors、env_settings、engine_settings、environment_parameters、checkpoint_settings、torch_settings、debug)。静态方法from_argparse(args)读取parse_command_line产生的 argparse.Namespace,加载 YAML 配置文件,再用命令行非默认参数覆盖 YAML 值,最终构造RunOptionsTrainerSettings.structure中的deep_update_dict支持对嵌套 dict 做递归合并,这正是"default_settings + behaviors 覆盖"机制的底层实现。

八、总结:如何选择 On-Policy 与 Off-Policy 训练器

  • OnPolicyTrainer(PPO):策略更新的经验必须来自当前策略,缓冲攒满buffer_size后整体用于num_epoch轮 minibatch 更新再清空。适合大多数单智能体、环境奖励稀疏度适中的任务;threaded=False可保证严格 on-policy 语义。
  • OffPolicyTrainer(SAC):经验存入大容量回放缓冲反复随机采样,通过steps_per_update控制更新频率,支持离散动作与循环网络,并可将回放缓冲保存为last_replay_buffer.hdf5供断点续训复用;适合样本效率要求高、奖励稠密或动作空间较大的任务。
  • 无论是哪一类,其训练循环都统一遵循 Trainer 基类定义的"轨迹队列消费 →_update_policy()→ 策略队列发布"模式,差异仅在于"何时算就绪、如何从缓冲取数据、如何更新优化器与奖励信号"。

进一步阅读:训练配置完整参考见 Training-Configuration-File.md,训练器源码入口在 trainer/,所有官方示例配置位于 config/ 下的ppo/sac/poca/imitation/目录。

  • 人工智能
  • 强化学习
  • 深度学习
  • 机器学习
  • 游戏开发
  • AI 应用

【免费下载链接】ml-agents

The Unity Machine Learning Agents Toolkit (ML-Agents) is an open-source project that enables games and simulations to serve as environments for training intelligent agents using deep reinforcement learning and imitation learning.

项目地址:https://gitcode.com/gh_mirrors/ml/ml-agents
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询