ML-Agents 自定义训练器插件(Custom Trainer Plugin)开发实战:基于 A2C / DQN 示例扩展你的强化学习算法

发布时间:2026/9/21 1:09:05
ML-Agents 自定义训练器插件(Custom Trainer Plugin)开发实战:基于 A2C / DQN 示例扩展你的强化学习算法 ML-Agents 自定义训练器插件Custom Trainer Plugin开发实战基于 A2C / DQN 示例扩展你的强化学习算法【免费下载链接】ml-agentsThe 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 开发者的自定义训练器插件开发指南。ML-Agents 在训练层引入了一套基于 Pythonentry_points的可扩展插件系统允许你将全新的强化学习算法如 A2C、DQN以独立 pip 包的形式注册进mlagents-learn训练 CLI并让自定义算法的超参数直接出现在 YAML 配置文件中。读完本文你将掌握插件包的目录结构、setup.py注册方式、Trainer / Optimizer / Hyperparameter 三层 API 的继承与实现要点并能基于仓库自带的mlagents_trainer_plugin示例A2C 与 DQN 完整实现快速搭建属于你自己的训练算法插件。一、为什么需要自定义训练器插件ML-Agents 官方内置了 PPO、SAC、POCA 等训练算法但不同研究与应用场景往往需要更多样的强化学习算法。为此ML-Agents 在ml-agents训练包中引入了可扩展的插件系统使用户可以基于高层训练器 API 定义新的训练器Trainer从而将mlagents-learnCLI 的训练流程路由到自定义训练器用自定义算法专属的超参数扩展 YAML 配置文件通过pip install -e将算法包注册为 entry point开箱即用。文档明确指出这一插件机制以setuptools的插件体系为基础当前仍处于beta阶段接口在未来版本中可能调整。关于 Python 插件系统的通用工作原理可参考 Training-Plugins.md。On-policy 与 Off-policy 算法概览无模型强化学习算法大体分为两类这也是插件 API 划分 Trainer 基类的基本依据On-policy在策略算法基于当前策略采集的数据执行更新。典型代表是 PPO 与本文示例中的 A2C。Off-policy离策略算法从历史数据缓冲区中学习 Q 函数再以 Q 函数指导决策。典型代表是 DQN 与 SAC。在 ML-Agents 的语境下off-policy 算法有三个关键优势其一由于可以从缓冲区反复抽取并复用数据通常比 on-policy 算法更节省样本其二可以将玩家演示player demonstrations以在线方式插入缓冲区与强化学习数据混合使用从而开启基于玩家数据流的模仿学习新范式。插件系统为这两类算法分别提供了OnPolicyTrainer与OffPolicyTrainer基类位于 ml-agents/mlagents/trainers/ 下你实现的自定义训练器需要根据算法策略类型选择继承其中一个。二、插件系统的工作原理entry_points 注册机制插件的核心是 Python 包元数据中的entry_points声明。ML-Agents 在 ml-agents/mlagents/plugins/init.py 中定义了插件接口名称常量ML_AGENTS_STATS_WRITER mlagents.stats_writer ML_AGENTS_TRAINER_TYPE mlagents.trainer_type其中ML_AGENTS_TRAINER_TYPE即字符串mlagents.trainer_type就是训练器插件接口。训练启动时ml-agents/mlagents/plugins/trainer_type.py 中的register_trainer_plugins()会遍历importlib.metadata.entry_points()中所有注册到该接口的 entry point逐个调用其注册函数并将返回的{训练器名: Trainer类}与{训练器名: Settings类}合并到全局注册表中for entry_point in entry_points: plugin_func entry_point.load() plugin_trainer_types, plugin_trainer_settings plugin_func() mla_plugins.all_trainer_types.update(plugin_trainer_types) mla_plugins.all_trainer_settings.update(plugin_trainer_settings)从源码可以看到trainer_type.py注册过程对每个 entry point 使用try/except BaseException包裹即使某个插件初始化失败也只会记录日志并跳过该插件而不会让整个训练流程崩溃——这保证了用户自定义代码的健壮性。官方默认训练器 PPO、SAC、POCA 同样通过get_default_trainer_types()trainer_type.py注册到mlagents.trainer_type接口下与你的自定义插件在同一个注册表中平等共存。三、参考实现mlagents_trainer_plugin 包结构为了让用户直观了解插件包的组织方式仓库提供了官方示例包 ml-agents-trainer-plugin其中完整实现了A2C在策略与DQN离策略两个算法。其目录结构如下├── mlagents_trainer_plugin │ ├── __init__.py │ ├── a2c │ │ ├── __init__.py │ │ ├── a2c_3DBall.yaml │ │ ├── a2c_optimizer.py │ │ └── a2c_trainer.py │ └── dqn │ ├── __init__.py │ ├── dqn_basic.yaml │ ├── dqn_optimizer.py │ └── dqn_trainer.py └── setup.py每个算法子包遵循一致的四文件约定YAML 配置文件定义训练超参数*_trainer.py实现 Trainer 类*_optimizer.py实现 Optimizer 与 Settings 超参数类。包根目录的setup.py负责向插件系统声明注册from setuptools import setup from mlagents.plugins import ML_AGENTS_TRAINER_TYPE setup( namemlagents_trainer_plugin, version0.0.1, entry_points{ ML_AGENTS_TRAINER_TYPE: [ a2cmlagents_trainer_plugin.a2c.a2c_trainer:get_type_and_setting, dqnmlagents_trainer_plugin.dqn.dqn_trainer:get_type_and_setting, ] }, )对照 ml-agents-trainer-plugin/setup.py 可见entry point 的写法为{entry point name}{plugin module}:{plugin function}a2c/dqn是训练器类型名也就是 YAML 配置中trainer_type字段的取值mlagents_trainer_plugin.a2c.a2c_trainer是模块路径get_type_and_setting是注册函数它返回{训练器名: Trainer类}与{训练器名: Settings类}两个字典。四、安装与执行4.1 前置安装如果你还没有安装基础包请先按 安装指南 完成ml-agents-envs与ml-agents的安装。在仓库根目录下可以按以下顺序安装pip3 install -e ./ml-agents-envs pip3 install -e ./ml-agents4.2 安装插件包从仓库根目录安装ml-agents-trainer-plugin替换为你自己的插件包名即可pip3 install -e ./ml-agents-trainer-plugin-eeditable模式会把包以开发模式安装进当前 Python 虚拟环境安装后你的插件即被注册为mlagents.trainer_type接口下的 entry point。4.3 运行自定义训练器完成安装后就可以使用包含新训练器的配置文件启动mlagents-learnmlagents-learn ml-agents-trainer-plugin/mlagents_trainer_plugin/a2c/a2c_3DBall.yaml --run-id run-id-name --env env-executable其中a2c_3DBall.yaml中trainer_type: a2c指定了使用 A2C 训练器--run-id为本次运行的标识--env指向 Unity 环境可执行文件或省略以连接 Unity Editor 中的训练场景。mlagents-learn会通过 entry point 自动解析a2c对应的 Trainer 与 Settings 类。五、逐步实现自定义训练器四步教程下面结合 Tutorial-Custom-Trainer-Plugin.md 的四步流程说明如何从零实现一个名为YourCustomTrainer的训练器。5.1 准备虚拟环境建议使用venv或conda创建并激活独立的 Python 虚拟环境conda create -n trainer-env python3.10.12 conda activate trainer-envStep 1编写自定义 Trainer 类插件用户负责按 API 标准实现训练器类根据算法策略选择继承OnPolicyTrainer或OffPolicyTrainer。自定义训练器负责采集经验并训练模型其角色相当于 Policy 与 Optimizer 之间的协调者。首先在create_policy方法中创建 Policy 对象def create_policy( self, parsed_behavior_id: BehaviorIdentifiers, behavior_spec: BehaviorSpec ) - TorchPolicy: actor_cls: Union[Type[SimpleActor], Type[SharedActorCritic]] SimpleActor actor_kwargs: Dict[str, Any] { conditional_sigma: False, tanh_squash: False, } if self.shared_critic: reward_signal_configs self.trainer_settings.reward_signals reward_signal_names [ key.value for key, _ in reward_signal_configs.items() ] actor_cls SharedActorCritic actor_kwargs.update({stream_names: reward_signal_names}) policy TorchPolicy( self.seed, behavior_spec, self.trainer_settings.network_settings, actor_cls, actor_kwargs, ) return policymlagents.trainers.torch_entities.networks提供了SimpleActorActor 与 Critic 分离和SharedActorCritic共享网络两种网络结构供选择。接着在create_optimizer方法中创建 Optimizer 并连接到 Policydef create_optimizer(self) - TorchOptimizer: return TorchPPOOptimizer( # type: ignore cast(TorchPolicy, self.policy), self.trainer_settings # type: ignore ) # type: ignoreRLTrainer基类有两个抽象方法必须实现_process_trajectory与_update_policy。_process_trajectory接收一个Trajectory负责计算价值估计与优势目标将处理后的数据写入更新缓冲区。其典型流程不完整示例为将 trajectory 转为 agent buffer调用self.optimizer.get_trajectory_value_estimates获取各 reward signal 的价值估计计算奖励最后调用self._append_to_update_buffer(agent_buffer_trajectory)追加到更新缓冲区供 Optimizer 更新模型使用def _process_trajectory(self, trajectory: Trajectory) - None: super()._process_trajectory(trajectory) agent_id trajectory.agent_id # All the agents should have the same ID agent_buffer_trajectory trajectory.to_agentbuffer() # Get all value estimates ( value_estimates, value_next, value_memories, ) self.optimizer.get_trajectory_value_estimates( agent_buffer_trajectory, trajectory.next_obs, trajectory.done_reached and not trajectory.interrupted, ) for name, v in value_estimates.items(): agent_buffer_trajectory[RewardSignalUtil.value_estimates_key(name)].extend( v ) self._stats_reporter.add_stat( fPolicy/{self.optimizer.reward_signals[name].name.capitalize()} Value Estimate, np.mean(v), ) # Evaluate all reward functions self.collected_rewards[environment][agent_id] np.sum( agent_buffer_trajectory[BufferKey.ENVIRONMENT_REWARDS] ) for name, reward_signal in self.optimizer.reward_signals.items(): evaluate_result ( reward_signal.evaluate(agent_buffer_trajectory) * reward_signal.strength ) agent_buffer_trajectory[RewardSignalUtil.rewards_key(name)].extend( evaluate_result ) # Report the reward signals self.collected_rewards[name][agent_id] np.sum(evaluate_result) self._append_to_update_buffer(agent_buffer_trajectory)Trajectory 本质上是字符串到任意值的字典列表对 Policy 调用forward时参数包含来自最后一步的experience字典字段包括 observation、action、reward、done 状态、group_reward、LSTM 记忆状态等forward会生成动作并产出下一个 experience 字典。Step 2实现自定义 Optimizer以官方 PPO 优化器TorchPPOOptimizer(TorchOptimizer)为参照Optimizer 接收 Policy 与训练参数在update方法中完成价值估计与损失计算。先为自定义优化器定义超参数类继承OnPolicyHyperparamSettingson-policy 算法或OffPolicyHyperparamSettingsoff-policy 算法class PPOSettings(OnPolicyHyperparamSettings): beta: float 5.0e-3 epsilon: float 0.2 lambd: float 0.95 num_epoch: int 3 shared_critic: bool False learning_rate_schedule: ScheduleType ScheduleType.LINEAR beta_schedule: ScheduleType ScheduleType.LINEAR epsilon_schedule: ScheduleType ScheduleType.LINEAR然后按如下接口实现update方法其中从 Trainer 生成的AgentBuffer计算各类损失与指标def update(self, batch: AgentBuffer, num_sequences: int) - Dict[str, float]:PPO 的典型计算模式不完整示例为通过self.policy.actor.get_stats获取 log 概率与熵通过self.critic.critic_pass获取价值估计再用ModelUtils.trust_region_policy_loss计算信任区域策略损失run_out self.policy.actor.get_stats( current_obs, actions, masksact_masks, memoriesmemories, sequence_lengthself.policy.sequence_length, ) log_probs run_out[log_probs] entropy run_out[entropy] values, _ self.critic.critic_pass( current_obs, memoriesvalue_memories, sequence_lengthself.policy.sequence_length, ) policy_loss ModelUtils.trust_region_policy_loss( ModelUtils.list_to_tensor(batch[BufferKey.ADVANTAGES]), log_probs, old_log_probs, loss_masks, decay_eps, ) loss ( policy_loss 0.5 * value_loss - decay_bet * ModelUtils.masked_mean(entropy, loss_masks) )最后更新模型并返回包含损失与衰减后学习率的统计字典ModelUtils.update_learning_rate(self.optimizer, decay_lr) self.optimizer.zero_grad() loss.backward() self.optimizer.step() update_stats { Losses/Policy Loss: torch.abs(policy_loss).item(), Losses/Value Loss: value_loss.item(), Policy/Learning Rate: decay_lr, Policy/Epsilon: decay_eps, Policy/Beta: decay_bet, }Step 3将自定义训练器集成进插件系统在setup.py的setup()调用中向entry_points字典添加ML_AGENTS_TRAINER_TYPE条目entry_points{ ML_AGENTS_TRAINER_TYPE: [ your_trainer_typeyour_package.your_custom_trainer:get_type_and_setting ] },关键元素含义元素含义ML_AGENTS_TRAINER_TYPE训练器类型插件接口的字符串常量mlagents.trainer_typeyour_trainer_type自定义训练器类型名用于配置文件中的trainer_type字段your_package包含自定义训练器实现的可 pip 安装包在YourCustomTrainer类中定义get_type_and_setting注册函数def get_type_and_setting(): return {YourCustomTrainer.get_trainer_name(): YourCustomTrainer}, { YourCustomTrainer.get_trainer_name(): YourCustomSetting }最后在配置文件behaviors段中指定训练器类型behaviors: 3DBall: trainer_type: your_trainer_type ...Step 4安装自定义训练器并启动训练确保ml-agents-envs与ml-agents已安装然后安装你的插件包若包可 pip 安装pip3 install your_custom_package或直接安装官方示例包pip3 install -e ./ml-agents-trainer-plugin安装完成后用包含新训练器类型的配置文件启动训练mlagents-learn ml-agents-trainer-plugin/mlagents_trainer_plugin/a2c/a2c_3DBall.yaml --run-id run-id-name --env env-executable验证你的实现在干净环境Python 3.10.12中按前述步骤安装所有依赖后可以通过以下方式验证pip3 install -e ./ml-agents-envs pip3 install -e ./ml-agents pip3 install -e ./ml-agents-trainer-plugin确认配置文件中已指定trainer_type: a2c然后试运行mlagents-learn ml-agents-trainer-plugin/mlagents_trainer_plugin/a2c/a2c_3DBall.yaml --run-id test-trainer你还可以在 Python REPL 中列出当前注册的所有训练器 import pkg_resources for entry in pkg_resources.iter_entry_points(mlagents.trainer_type): ... print(entry) ... default mlagents.plugins.trainer_type:get_default_trainer_types a2c mlagents_trainer_plugin.a2c.a2c_trainer:get_type_and_setting dqn mlagents_trainer_plugin.dqn.dqn_trainer:get_type_and_setting若安装正确会看到 Unity 标志与训练启动提示[INFO] Listening on port 5004. Start training by pressing the Play button in the Unity Editor.若出现以下报错说明训练器类型拼写错误或该训练器未安装mlagents.trainers.exception.TrainerConfigError: Invalid trainer type a2c was found六、源码级解读A2C 与 DQN 示例实现6.1 A2C在策略算法的参考实现A2C 训练器继承自OnPolicyTrainer完整实现位于 ml-agents-trainer-plugin/mlagents_trainer_plugin/a2c/a2c_trainer.py。其核心要点在__init__中通过cast(A2CSettings, self.trainer_settings.hyperparameters)将超参数断言为自定义的A2CSettings类型并读取shared_critic开关a2c_trainer.py_process_trajectory中调用get_gae计算 GAE 优势与回报写入BufferKey.ADVANTAGES与BufferKey.DISCOUNTED_RETURNSa2c_trainer.pycreate_policy根据shared_critic选择SimpleActor或SharedActorCritica2c_trainer.py文件末尾的get_type_and_setting返回训练器与超参数注册表a2c_trainer.py。对应的优化器 a2c_optimizer.py 中定义了A2CSettings(OnPolicyHyperparamSettings)并带有一个强约束校验器——A2C 只做单次遍历num_epoch必须为 1否则抛出TrainerConfigErrorattr.s(auto_attribsTrue) class A2CSettings(OnPolicyHyperparamSettings): beta: float 5.0e-3 lambd: float 0.95 num_epoch: int attr.ib(default1) # A2C does just one pass shared_critic: bool False num_epoch.validator def _check_num_epoch_one(self, attribute, value): if value ! 1: raise TrainerConfigError(A2C requires num_epoch 1) learning_rate_schedule: ScheduleType ScheduleType.LINEAR beta_schedule: ScheduleType ScheduleType.LINEAR见 a2c_optimizer.pyA2COptimizer.update中损失函数为经典的policy_loss 0.5 * value_loss - decay_bet * entropy组合使用 Adam 优化器并输出Losses/Policy Loss、Losses/Value Loss、Policy/Learning Rate、Policy/Beta等统计量a2c_optimizer.py。6.2 DQN离策略算法的参考实现DQN 训练器继承自OffPolicyTrainer实现位于 ml-agents-trainer-plugin/mlagents_trainer_plugin/dqn/dqn_trainer.py。其特点_process_trajectory将轨迹写入回放缓冲区replay buffer并在last_step.interrupted时对最后一步做 bootstrap 处理复制观测并清除 done 标志dqn_trainer.pycreate_policy将exploration_initial_eps传入QNetwork作为 actor 网络并调用self.maybe_load_replay_buffer()支持回放缓冲区持久化dqn_trainer.py。DQNSettings(OffPolicyHyperparamSettings)定义了探索率调度与目标网络更新等 DQN 专属参数attr.s(auto_attribsTrue) class DQNSettings(OffPolicyHyperparamSettings): gamma: float 0.99 exploration_schedule: ScheduleType ScheduleType.LINEAR exploration_initial_eps: float 0.1 exploration_final_eps: float 0.05 target_update_interval: int 10000 tau: float 0.005 steps_per_update: float 1 save_replay_buffer: bool False reward_signal_steps_per_update: float attr.ib()见 dqn_optimizer.pyDQNOptimizer.update中探索率通过decay_exploration_rate随训练步数衰减并写回self.policy.actor.exploration_ratedqn_optimizer.py损失使用smooth_l1_lossHuber 损失计算当前 Q 与目标 Q 之间的误差并在每步后用soft_updatetau软更新将在线网络参数融合进目标网络dqn_optimizer.py。QNetwork同时继承nn.Module、Actor、Critic实现get_action_and_stats时以exploration_rate决定随机探索动作或贪心动作dqn_optimizer.py。七、配置文件与超参数详解7.1 A2C 配置a2c_3DBall.yaml官方 A2C 示例配置位于 ml-agents-trainer-plugin/mlagents_trainer_plugin/a2c/a2c_3DBall.yamlbehaviors: 3DBall: trainer_type: a2c hyperparameters: batch_size: 1000 buffer_size: 1000 learning_rate: 0.0003 beta: 0.001 lambd: 0.99 num_epoch: 1 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: 1000各超参数与A2CSettings字段一一对应batch_size与buffer_size控制每次更新的样本量与缓冲区容量learning_rate配合learning_rate_schedule: linear做线性衰减源码中ModelUtils.DecayedValue从初始学习率衰减至1e-10见 a2c_optimizer.pybeta为熵正则系数lambd为 GAE 衰减因子num_epoch: 1必须为 1A2C 单次遍历特性配置不合法会触发校验错误。7.2 DQN 配置dqn_basic.yaml官方 DQN 示例配置位于 ml-agents-trainer-plugin/mlagents_trainer_plugin/dqn/dqn_basic.yamlbehaviors: Basic: trainer_type: dqn hyperparameters: learning_rate: 0.0003 learning_rate_schedule: constant batch_size: 64 buffer_size: 50000 tau: 0.005 steps_per_update: 10.0 save_replay_buffer: false exploration_schedule: linear exploration_initial_eps: 0.8 exploration_final_eps: 0.05 network_settings: normalize: false hidden_units: 20 num_layers: 2 vis_encode_type: simple reward_signals: extrinsic: gamma: 0.99 strength: 1.0 keep_checkpoints: 5 max_steps: 500000 time_horizon: 10 summary_freq: 1000DQN 专属参数中tau为软更新系数steps_per_update为每次更新间隔的环境步数exploration_initial_eps/exploration_final_eps与exploration_schedule: linear共同控制 ε-greedy 探索率的线性衰减衰减窗口为 20000 步见 dqn_optimizer.pysave_replay_buffer决定是否持久化回放缓冲区buffer_size: 50000对应离策略算法的大容量经验缓冲区。八、更多插件接口StatsWriter训练器插件并非唯一的扩展点。ML-Agents 的插件系统还提供了StatsWriter接口接口名mlagents.stats_writer用于在训练过程中接收各类统计信息如每个 summary 周期的平均 Agent 奖励。默认实现将统计信息输出到控制台与 TensorBoard。StatsWriter的注册函数接收RunOptions参数并返回StatsWriter列表。官方参考实现位于 ml-agents-plugin-examples/mlagents_plugin_examples/example_stats_writer.py其setup.py中的注册方式为entry_points{ ML_AGENTS_STATS_WRITER: [ examplemlagents_plugin_examples.example_stats_writer:get_example_stats_writer ] }register_stats_writer_plugins()的实现见 ml-agents/mlagents/plugins/stats_writer.py与训练器插件同样采用遍历 entry point 并逐项容错加载的模式。默认的get_default_stats_writers()会返回 TensorBoardWriter、GaugeWriter、ConsoleWriter 三个写入器stats_writer.py。九、常见问题与排查TrainerConfigError: Invalid trainer type xxx was found配置文件中的trainer_type拼写错误或对应插件包未安装。请检查 entry point 名称与配置值是否一致并用 5.4 节的 REPL 命令列出所有已注册训练器。插件初始化失败导致训练中断register_trainer_plugins会捕获插件加载阶段的异常并跳过该插件同时打印异常日志。排查时优先查看mlagents-learn启动日志中的Error initializing Trainer plugins信息。超参数校验失败例如 A2C 的num_epoch必须为 1。自定义 Settings 类可以通过attr的 validator 对超参数施加约束配置不合法时会在解析阶段直接报错。关于 Trainer 与 Optimizer 的 API 详细参考可继续阅读 Python-On-Off-Policy-Trainer-Documentation.md 与 Python-Optimizer-Documentation.md插件机制的通用说明见 Training-Plugins.md。动手实践时以仓库中的 ml-agents-trainer-plugin 为起点将 A2C 或 DQN 替换为你自己的算法实现即可快速完成首个自定义训练器。【免费下载链接】ml-agentsThe 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),仅供参考