资讯详情

资讯详情

ML-Agents 训练插件机制实战:用 setuptools entry_points 自定义 StatsWriter 与 Trainer

人工智能强化学习深度学习机器学习游戏开发AI 应用【免费下载链接】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点击查看免费下载ML-AgentsUnity Machine Learning Agents Toolkit的训练器mlagents-learn支持通过 setuptools 的 entry points入口点机制在训练流程中注入用户自己用 Python 实现的具体接口。本文以官方文档 Training-Plugins.md 为主线结合仓库内的插件源码ml-agents/mlagents/plugins、ml-agents-plugin-examples、ml-agents-trainer-plugin深入讲解插件的编写、注册、安装与加载原理。读完本文你将能够编写并注册自定义的StatsWriter统计信息写出器与TrainerType自定义训练器类型插件把训练统计输出到自建监控系统或为 ML-Agents 接入 PPO/SAC/POCA 之外的新强化学习算法。注意插件接口目前仍处于 beta 阶段官方文档原文为Plugin interfaces should currently be considered in beta接口数量有限且可能在后续版本中调整升级 ml-agents 版本后需要注意兼容性。插件机制概述基于 setuptools entry points 的扩展架构ML-Agents 的插件系统建立在 Python 包管理工具 setuptools 的 entry points入口点机制之上这也是社区中许多插件化 Python 项目如 pytest、flake8采用的经典方案插件作者在自己的 Python 包中声明入口点而宿主程序这里是mlagents-learn在启动时通过importlib.metadata扫描所有已安装包中注册的入口点并调用其中的注册函数来获得插件实例。从仓库源码可以看到当前 ML-Agents 定义了两个插件接口常量位于 ml-agents/mlagents/plugins/init.pyML_AGENTS_STATS_WRITER mlagents.stats_writer ML_AGENTS_TRAINER_TYPE mlagents.trainer_typeML_AGENTS_STATS_WRITER字符串常量值为mlagents.stats_writer统计信息写出器接口ML_AGENTS_TRAINER_TYPE字符串常量值为mlagents.trainer_type自定义训练器类型接口。即使是 ML-Agents 自身的默认行为也是通过这一机制注册的。在 ml-agents/setup.py 中可以看到entry_points{ console_scripts: [ mlagents-learnmlagents.trainers.learn:main, mlagents-run-experimentmlagents.trainers.run_experiment:main, mlagents-push-to-hfmlagents.utils.push_to_hf:main, mlagents-load-from-hfmlagents.utils.load_from_hf:main, ], # Plugins - each plugin type should have an entry here for the default behavior ML_AGENTS_STATS_WRITER: [ defaultmlagents.plugins.stats_writer:get_default_stats_writers ], ML_AGENTS_TRAINER_TYPE: [ defaultmlagents.plugins.trainer_type:get_default_trainer_types ], },也就是说默认的统计写出器与默认的 PPO/SAC/POCA 训练器本质上也是以插件形式注册的默认入口点名称为default。这保证了插件加载逻辑的完全统一。如何编写你自己的插件编写一个 ML-Agents 训练插件核心工作是两部分实现插件接口对应的 Python 类/注册函数以及在setup.py中声明 entry_points。仓库中的 ml-agents-plugin-examples 目录提供了每个插件接口的参考实现是绝佳的起步模板。第一步准备 setup.py 并声明 entry_points如果你还没有为你的 Python 代码准备setup.py需要先补上。参考实现 ml-agents-plugin-examples/setup.py 是一个极简示例from setuptools import setup from mlagents.plugins import ML_AGENTS_STATS_WRITER setup( namemlagents_plugin_examples, version0.0.1, # Example of how to add your own registration functions that will be called # by mlagents-learn. # # Here, the get_example_stats_writer() function in mlagents_plugin_examples/example_stats_writer.py # will get registered with the ML_AGENTS_STATS_WRITER plugin interface. entry_points{ ML_AGENTS_STATS_WRITER: [ examplemlagents_plugin_examples.example_stats_writer:get_example_stats_writer ] }, )在setup()调用中需要为每个你实现的插件接口在entry_points字典中增加对应条目。其形式为{插件接口名}: [ {插件实现名}{插件模块}:{插件注册函数} ]各组成部分的含义组成部分示例说明插件接口名ML_AGENTS_STATS_WRITER即字符串常量mlagents.stats_writer必须是官方提供的接口之一见下文插件接口。可以直接从mlagents.plugins导入常量避免手写字符串拼错插件实现名example该插件实现的名称可以随意取仅用于标识插件模块mlagents_plugin_examples.example_stats_writer指向注册函数所在的 Python 模块插件注册函数get_example_stats_writer运行mlagents-learn时会被调用的函数。不同插件接口对它的参数与返回值要求不同见下文需要特别注意的是插件注册函数与插件实现类是两回事。注册函数是被mlagents-learn启动时调用、用于生产插件实例的工厂函数而真正的功能逻辑写在实现了接口抽象基类的类中。第二步本地安装插件在setup.py中定义好entry_points之后需要在你安装了mlagents的同一个 Python 虚拟环境中执行可编辑安装pip install -e [path to your plugin code]例如仓库中两个参考插件包的安装方式为pip install -e ml-agents-plugin-examples pip install -e ml-agents-trainer-plugin-eeditable模式会把包以开发模式安装源码修改即时生效无需反复重装同时 setuptools 会把entry_points元数据写入当前环境中之后importlib.metadata.entry_points()才能扫描到它们。环境必须与mlagents一致——插件运行期需要导入mlagents.trainers.settings、mlagents.trainers.stats等模块跨环境安装会导致扫描不到入口点或导入失败。另外从 ml-agents/setup.py 可以看到当前mlagents要求 Python 版本3.10.1,3.10.12插件也应在此范围内开发测试。插件接口一StatsWriter统计信息写出器StatsWriter负责接收训练过程中产生的各类统计信息例如每个 summary 周期内 Agent 的平均奖励并决定如何输出。默认情况下ML-Agents 会把这类信息打印到控制台并写入 TensorBoard。如果你希望把统计信息转发到自定义的可视化平台、数据库或消息队列实现一个自定义StatsWriter是最直接的方式。接口定义StatsWriter抽象基类定义在 ml-agents/mlagents/trainers/stats.py。任何派生类必须实现write_stats()方法此外还可以按需覆写on_add_stat()与add_property()。class StatsWriter(abc.ABC): def on_add_stat( self, category: str, key: str, value: float, aggregation: StatsAggregationMethod StatsAggregationMethod.AVERAGE, ) - None: 每条统计值上报到 StatsReporter.add_stat / set_stat 时的回调。 pass abc.abstractmethod def write_stats( self, category: str, values: Dict[str, StatsSummary], step: int ) - None: 记录训练信息的回调。 pass def add_property( self, category: str, property_type: StatsPropertyType, value: Any ) - None: 向 StatsWriter 添加通用属性如超参、最大步数、训练器类型等。 pass参数含义与官方文档一致并可从源码注释确认category统计信息的类别通常是被训练 Agent 的 behavior 名称values字符串键到StatsSummary值组成的字典Dict[str, StatsSummary]StatsSummary内含mean、std、num等汇总字段即统计周期内的聚合结果step当前的训练步数on_add_stat()每条统计值上报时触发的回调可以在这里为每个统计发射注册自定义处理其aggregation参数默认使用StatsAggregationMethod.AVERAGE平均值聚合。默认实现TensorboardWriter / GaugeWriter / ConsoleWriter默认的统计写出器集合定义在 ml-agents/mlagents/plugins/stats_writer.py 的get_default_stats_writers()中mlagents-learn始终使用这三者def get_default_stats_writers(run_options: RunOptions) - List[StatsWriter]: checkpoint_settings run_options.checkpoint_settings return [ TensorboardWriter( checkpoint_settings.write_path, clear_past_datanot checkpoint_settings.resume, hidden_keys[Is Training, Step], ), GaugeWriter(), ConsoleWriter(), ]TensorboardWriter把统计写入 TensorBoard写路径来自checkpoint_settings.write_path非 resume 续训时清空历史数据隐藏Is Training、Step两个键GaugeWriter把统计写入计时器仪表盘见 stats.py 的实现会将 category 与键名中的/与空格做净化处理后写入 gaugeConsoleWriter输出到标准输出 stdout。这三者的实现类同样位于 ml-agents/mlagents/trainers/stats.py可以作为编写自定义写出器时参考的范本。StatsWriter 插件的注册StatsWriter的注册函数签名是接收一个RunOptions参数返回一个StatsWriter列表。仓库参考实现 ml-agents-plugin-examples/mlagents_plugin_examples/example_stats_writer.py 完整代码如下from typing import Dict, List from mlagents.trainers.settings import RunOptions from mlagents.trainers.stats import StatsWriter, StatsSummary class ExampleStatsWriter(StatsWriter): Example implementation of the StatsWriter abstract class. This doesnt do anything interesting, just prints the stats that it gets. def write_stats( self, category: str, values: Dict[str, StatsSummary], step: int ) - None: print(fExampleStatsWriter category: {category} values: {values}) def get_example_stats_writer(run_options: RunOptions) - List[StatsWriter]: Registration function. This is referenced in setup.py and will be called by mlagents-learn when it starts to determine the list of StatsWriters to use. It must return a list of StatsWriters. print(Creating a new stats writer! This is so exciting!) return [ExampleStatsWriter()]两点值得注意注册函数中可以先对run_options做判断例如根据配置决定启用哪些写出器再返回一个或多个StatsWriter实例返回列表run_options类型为mlagents.trainers.settings.RunOptions它携带了本次训练的全部运行选项含checkpoint_settings你可以借此读取 write path、resume 标志等实现与默认 TensorboardWriter 类似的路径逻辑。加载原理register_stats_writer_pluginsmlagents-learn启动时调用 ml-agents/mlagents/plugins/stats_writer.py 中的register_stats_writer_plugins(run_options)来加载所有统计写出器插件def register_stats_writer_plugins(run_options: RunOptions) - List[StatsWriter]: all_stats_writers: List[StatsWriter] [] if ML_AGENTS_STATS_WRITER not in importlib_metadata.entry_points(): logger.warning( fUnable to find any entry points for {ML_AGENTS_STATS_WRITER}, even the default ones. Uninstalling and reinstalling ml-agents via pip should resolve. Using default plugins for now. ) return get_default_stats_writers(run_options) entry_points importlib_metadata.entry_points()[ML_AGENTS_STATS_WRITER] for entry_point in entry_points: try: plugin_func entry_point.load() plugin_stats_writers plugin_func(run_options) all_stats_writers plugin_stats_writers except BaseException: logger.exception( fError initializing StatsWriter plugins for {entry_point.name}. This plugin will not be used. ) return all_stats_writers其关键逻辑可以概括为通过importlib.metadata.entry_points()扫描环境内所有注册到mlagents.stats_writer接口的入口点逐个调用每个入口点的注册函数entry_point.load()后调用把返回的StatsWriter列表累加进all_stats_writers任何插件初始化抛出的异常BaseException都会被捕获并记录日志该插件被跳过但不会中断整个训练——这是刻意设计避免用户代码出错导致训练崩溃源码注释Catch all exceptions from setting up the plugin, so that bad user code doesnt break things.如果环境中完全找不到该接口的任何入口点连默认的都没有会发出警告并回退到get_default_stats_writers(run_options)同时提示卸载并重装 ml-agents通常是修复办法。这意味着你的自定义 StatsWriter 插件会与默认的 TensorboardWriter、GaugeWriter、ConsoleWriter 一起生效而不是替换它们因为所有入口点返回的写出器都会被合并。插件接口二TrainerType自定义训练器类型ML_AGENTS_TRAINER_TYPEmlagents.trainer_type接口允许你为mlagents-learn注册全新的训练算法。仓库中的 ml-agents-trainer-plugin 目录提供了完整的参考实现它注册了a2c与dqn两种训练器每个训练器均配套定义了训练器类、优化器与超参设置。默认训练器同样由插件注册默认的三种训练器 PPO、SAC、POCA 也是通过默认入口点注册的见 ml-agents/mlagents/plugins/trainer_type.pydef get_default_trainer_types() - Tuple[Dict[str, Any], Dict[str, Any]]: mla_plugins.all_trainer_types.update( { PPOTrainer.get_trainer_name(): PPOTrainer, SACTrainer.get_trainer_name(): SACTrainer, POCATrainer.get_trainer_name(): POCATrainer, } ) mla_plugins.all_trainer_settings.update( { PPOTrainer.get_trainer_name(): PPOSettings, SACTrainer.get_trainer_name(): SACSettings, POCATrainer.get_trainer_name(): POCASettings, } ) return mla_plugins.all_trainer_types, mla_plugins.all_trainer_settings这里可以看到 TrainerType 插件的注册函数契约不接收参数返回一个二元组(all_trainer_types, all_trainer_settings)其中all_trainer_typesDict[str, Trainer]训练器名称如ppo、sac、poca由Trainer.get_trainer_name()返回到训练器类的映射all_trainer_settingsDict[str, HyperparamSettings]训练器名称到其超参设置类如PPOSettings、SACSettings、POCASettings的映射。这两个全局字典定义于 ml-agents/mlagents/plugins/init.py# TODO: the real type is Dict[str, HyperparamSettings] all_trainer_types: Dict[str, Any] {} all_trainer_settings: Dict[str, Any] {}register_trainer_plugins()同一文件 trainer_type.py负责加载所有 TrainerType 入口点并把各插件返回的类型/设置映射update进全局字典因此自定义训练器与内置训练器按名称共存。参考实现dqn 训练器插件ml-agents-trainer-plugin/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, ] }, )其中dqn训练器的实现位于 ml-agents-trainer-plugin/mlagents_trainer_plugin/dqn/dqn_trainer.py核心结构为定义训练器名称常量TRAINER_NAME dqnDQNTrainer继承自mlagents.trainers.trainer.off_policy_trainer.OffPolicyTrainerDQN 属于 off-policy 算法在构造函数中接收behavior_name、reward_buff_cap、trainer_settingsTrainerSettings、training、load、seed、artifact_path等参数配套DQNOptimizer、DQNSettings定义在dqn_optimizer.py以及QNetwork网络结构注册函数get_type_and_setting()返回({TRAINER_NAME: DQNTrainer}, {TRAINER_NAME: DQNSettings})二元组。插件注册后就可以在 YAML 训练配置中直接以trainer: dqn或a2c引用该算法。仓库内各算法的参考配置位于 config/ppo、config/sac、config/poca 等目录例如 config/ppo/3DBall.yaml其中trainer字段对应的就是 TrainerType 插件注册的算法名称。自定义训练器的完整开发流程还可参考官方教程 Tutorial-Custom-Trainer-Plugin.md。插件加载与运行链路小结把以上内容串起来mlagents-learn启动时与插件相关的完整调用链为入口命令mlagents-learnconsole_scripts指向 ml-agents/mlagents/trainers/learn.py 的main解析训练配置YAML为RunOptions调用register_stats_writer_plugins(run_options)扫描mlagents.stats_writer入口点 → 逐个调用注册函数 → 合并所有StatsWriter含默认三件套调用register_trainer_plugins()扫描mlagents.trainer_type入口点 → 调用各注册函数 → 把训练器类型与超参设置合并进all_trainer_types/all_trainer_settings全局字典训练循环中每个 summary 周期把统计聚合为StatSummary后逐一向所有已注册的StatsWriter分发write_stats(category, values, step)。从代码结构可以推断插件注册在训练启动阶段一次性完成之后统计写出与训练器选择均基于注册结果这种设计让第三方算法如 DQN、A2C无需改动 ml-agents 主仓库代码即可接入实现了训练生态的开放扩展。常见问题与排错启动时出现 Unable to find any entry points ... even the default ones 警告说明环境中完全找不到mlagents.stats_writer或mlagents.trainer_type的任何入口点通常是因为mlagents包安装异常。官方提示的修复方式是卸载并重装 ml-agentspip uninstall mlagents后重新pip install mlagents此时会回退到默认插件继续运行。插件未生效确认插件包与mlagents安装在同一个虚拟环境中且使用pip install -e 插件路径完成安装修改setup.py后需要重新执行安装以刷新 entry_points 元数据。插件初始化报错但训练未中断这是有意设计——register_stats_writer_plugins与register_trainer_plugins都捕获了BaseException并记录日志后跳过该插件。排查时请查看日志中的Error initializing ... plugin异常栈。接口处于 beta 阶段插件接口在未来版本中可能变化升级 ml-agents 版本后请留意 CHANGELOG.md 中与 plugins 相关的变更。参考与进一步阅读官方插件文档Training-Plugins.md插件接口常量与全局注册表ml-agents/mlagents/plugins/init.pyStatsWriter 插件加载实现ml-agents/mlagents/plugins/stats_writer.pyTrainerType 插件加载实现ml-agents/mlagents/plugins/trainer_type.pyStatsWriter 抽象基类与默认实现ml-agents/mlagents/trainers/stats.pyStatsWriter 参考插件ml-agents-plugin-examplesTrainerType 参考插件a2c / dqnml-agents-trainer-plugin自定义训练器完整教程Tutorial-Custom-Trainer-Plugin.md各算法训练配置示例config赞分享人工智能强化学习深度学习机器学习游戏开发AI 应用【免费下载链接】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点击查看免费下载相关推荐Detectron2 训练实战指南自定义训练循环、Trainer 抽象与 Hook 机制全解析Detectron2 训练实战指南自定义训练循环、Trainer 抽象与 Hook 机制全解析 导读 在完成自定义模型与数据加载器的搭建之后如何高效地把它们人工智能计算机视觉深度学习机器学习YOLO26多任务学习框架一站式解决计算机视觉需求YOLO26多任务学习框架一站式解决计算机视觉需求 YOLO26是Ultralytics推出的最新一代计算机视觉框架它集成了 目标检测 、 实例分割 、 图ML-Agents 自定义 Side Channel 实战用 C 与 Python 构建 Unity 与训练脚本之间的自定义通信通道ML Agents 自定义 Side Channel 实战用 C 与 Python 构建 Unity 与训练脚本之间的自定义通信通道 Side Channel人工智能强化学习深度学习机器学习游戏开发AI 应用创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →