ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Unity ML-Agents 中的 PyTorch 与 TensorBoard:从深度强化学习训练到 ONNX 模型导出

Unity ML-Agents 中的 PyTorch 与 TensorBoard:从深度强化学习训练到 ONNX 模型导出 人工智能强化学习深度学习机器学习游戏开发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点击查看免费下载本篇技术指南围绕 Unity ML-Agents Toolkit 的核心计算支柱展开系统讲解开源深度学习框架 PyTorch 在其中的角色、训练产出的 ONNX 模型如何被 Unity Agent 使用以及 TensorBoard 如何辅助你调优超参数。读完本文你将理解 ML-Agents 训练管线的 PyTorch 底层实现结构、模型保存与导出机制并掌握用 TensorBoard 观察训练曲线、迭代超参数的具体方法。PyTorch 在 ML-Agents 中的地位正如 机器学习背景文档 所述ML-Agents Toolkit 提供的许多强化学习算法PPO、SAC、POCA 等都依赖某种形式的深度学习而我们的具体实现正是构建在开源库 PyTorch 之上的。PyTorch 是一个使用数据流图data flow graphs执行计算的开放源码库数据流图是深度学习模型的底层表示形式。它支持在桌面、服务器或移动设备上的 CPU 与 GPU 上进行训练training和推理inference。在 ML-Agents Toolkit 中当你训练 Agent 的行为时最终输出是一个模型.onnx文件你可以将该文件与 Unity 中的 Agent 关联起来。除非你要实现全新的算法否则 PyTorch 的使用大多被抽象封装在幕后你几乎感知不到它的存在——这保证了用户只需要关心训练配置与结果而不必直接编写深度学习代码。PyTorch 封装层torch_utils 与统一设备管理从源码结构看ML-Agents 刻意将 PyTorch 的导入集中收敛在一个地方以统一管理设备与线程配置。torch.py 是训练侧唯一直接import torch的模块文件内注释明确说明 This should be the only place that we import torch directlyflak8 配置也会拦截其他位置的直接导入它承担了以下职责版本校验启动时检查已安装的 torch 版本是否不低于 1.6.0否则提示前往 PyTorch 官网安装线程数配置通过torch.set_num_threads(cpu_utils.get_num_threads_to_use())与KMP_BLOCKTIME0设置 CPU 线程策略设备选择set_torch_config()支持通过 TorchSettings 指定device未指定时自动回退为cuda if torch.cuda.is_available() else cpu并在cuda/xpu/mps设备上调用torch.set_default_device()同时将默认 dtype 设为torch.float32其余代码通过default_device()获取当前设备。所有训练相关的网络组件都通过from mlagents.torch_utils import torch, nn引用该统一入口例如 torch_policy.py 中的TorchPolicy会调用self.actor.to(default_device())将 Actor 网络搬到计算设备上并在evaluate()中使用torch.no_grad()关闭梯度计算以加速推理。网络构建torch_entities 模块训练侧的神经网络结构定义集中在 torch_entities 目录下主要包括networks.py定义ObservationEncoder、Actor 网络主体等核心模块。ObservationEncoder使用ModelUtils.create_input_processors()为每个观测创建处理器向量输入用VectorInput、视觉输入用 CNN、变长实体观测用EntityEmbedding并对变长观测使用残差自注意力Residual Self-AttentionRSA进行编码encoders.pyVectorInput、CNN等编码器实现支持观测归一化decoders.pyValueHeads等价值头与策略头layers.pyLinearEncoder、LSTM等基础层LSTM 支撑循环神经网络RNN策略distributions.py连续/离散动作的分布如高斯分布、Categorical用于采样动作与计算对数概率、熵。这些模块全部继承torch.nn.Module。你可以参考测试用例 test_networks.py 了解其标准训练循环先torch.manual_seed(0)固定随机种子然后实例化网络、用torch.optim.Adam(networkbody.parameters(), lr3e-3)构造优化器在前向传播后计算torch.nn.functional.mse_loss之类的损失再反向传播更新参数——这正是 PyTorch 典型的「定义网络 → 优化器 → 前向 → 损失 → 反向」流程。从训练到产出.pt 检查点与 .onnx 导出训练产物与模型保存训练过程中PyTorch 主要负责更新网络参数而 ML-Agents 在保存模型时区分了两种产物见 torch_model_saver.py检查点文件.pt每次save_checkpoint()都会用torch.save(state_dict, ...)保存所有已注册模块Policy、Optimizer 等的state_dict生成{behavior_name}-{step}.pt并同步保存一份默认的checkpoint.pt用于断点续训可运行模型.onnx同一时刻调用export()导出 ONNX 模型供 Unity 运行时推理使用。initialize_or_load()支持两条加载路径通过init_path指定初始模型或通过--resume/--load从checkpoint.pt恢复训练。加载时以strictFalse调用load_state_dict()对缺失键missing keys、多余键unexpected keys以及 KeyError/ValueError/RuntimeError 异常做容错处理并以 warning 日志提示然后决定是否把全局步数重置为 0。ONNX 导出机制ONNX 导出由 model_serialization.py 中的ModelSerializer完成其核心是调用torch.onnx.export()要点包括构造一批 dummy 输入全零观测、全一掩码、全零记忆形状按 ONNX 要求的NCHW通道优先格式处理Sentis 导入同样遵循该格式通过input_names如obs_0、action_masks、recurrent_in和output_names如version_number、memory_size、continuous_actions、discrete_actions、deterministic_continuous_actions等命名张量通过dynamic_axes将 batch 维度声明为动态轴保证不同批量大小的推理使用SerializationSettings.onnx_opset指定的 opset 版本整个导出过程运行在线程安全的exporting_to_onnx()上下文内is_exporting()可用于判断当前是否处于导出阶段以便在导出时切换为确定性输出。copy_final_model()会把最终生成的.onnx复制为{model_path}.onnx这就是你在 Unity 的 Behavior Parameters 组件Model 属性中要关联的模型文件。Unity 侧的模型推理由 Runtime/Inference 下的SentisPolicy、ModelRunner等实现完成通过输入名/输出名与训练端导出的张量一一对应。TensorBoard观察训练、调优超参数训练 PyTorch 模型的一个关键环节是为模型的某些属性称为超参数hyperparameters设定合适的值。找到合适的超参数值通常需要多轮迭代因此 ML-Agents 借助可视化工具 TensorBoard 来辅助这一过程。TensorBoard 允许可视化 Agent 在训练过程中的某些属性例如奖励 reward这有助于你建立对不同超参数的直觉并为你的 Unity 环境设定最优值。训练数据如何进入 TensorBoard在代码层面训练统计的写入由 stats.py 完成。它从torch.utils.tensorboard导入SummaryWriter按统计类别category为每个 writer 创建独立的SummaryWriter(filewriter_dir)从而将各类指标如累积奖励、价值损失、策略损失、熵等写入 event 文件。训练完成后在命令行启动tensorboard --logdir 训练结果目录即可在浏览器中查看曲线。从奖励曲线到超参数决策以仓库自带的 3DBall.yaml 为例PPO 训练涉及的超参数包括learning_rate、batch_size、buffer_size、beta熵正则系数、epsilonclip 范围、lambdGAE 系数、num_epoch、learning_rate_schedule等。通过 TensorBoard 观察奖励曲线可以快速定位问题奖励长期不上升可能需要调大学习率或增加num_epoch训练后期震荡剧烈可以考虑降低学习率或调整熵正则。关于各类超参数的具体含义与设置建议详见 Training ML-Agents如果你对 TensorBoard 本身还不熟悉推荐阅读 使用 TensorBoard 与 ML-Agents 指南。环境要求与开始使用在开始训练之前请确保安装满足要求的 PyTorch 与 TensorBoard。根据 setup.py 的依赖声明install_requires当前仓库要求torch2.1.1,2.8.0tensorboard2.14Python 版本3.10.1,3.10.12。安装完成并配置好 Unity 环境后通过mlagents-learn 配置文件.yaml命令启动训练对应入口为 learn.py 的mlagents-learnconsole script。训练过程中或结束后运行tensorboard --logdir results查看指标曲线。最终得到的.onnx模型拖入 Unity 场景中 Agent 的 Model 字段即可实现推理。如果torch.cuda.is_available()为真训练会自动在 GPU 上进行你也可以通过--torch-device参数显式指定计算设备。小结PyTorch 是 ML-Agents 深度强化学习训练的底层引擎torch_utils提供统一封装torch_entities定义网络结构训练产出.pt检查点与.onnx模型后者被 Unity 端 Sentis 运行时加载推理TensorBoard 则贯穿训练全程帮你通过奖励曲线等指标迭代超参数。理解了这条从 PyTorch 训练到 ONNX 部署的完整链路你就能更高效地调试训练配置、定位训练问题并为自己的 Unity 环境找到最优超参数组合。赞分享人工智能强化学习深度学习机器学习游戏开发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点击查看免费下载相关推荐从零创建 Unity 强化学习环境ML-Agents RollerBall 训练教程从零创建 Unity 强化学习环境ML Agents RollerBall 训练教程 导读 本文基于 Unity ML Agents Toolkit 官方教程人工智能强化学习深度学习机器学习游戏开发AI 应用ML-Agents 理论全景从强化学习基础到 Unity 智能体训练方法论ML Agents 理论全景从强化学习基础到 Unity 智能体训练方法论 本文是 Unity ML Agents Toolkit 的理论总览指南系统梳理该人工智能强化学习深度学习机器学习游戏开发AI 应用pytorch-image-models中的模型导出ONNX与Azure MLpytorch image models中的模型导出ONNX与Azure ML 概述 在计算机视觉领域将训练好的模型导出为标准化格式并部署到生产环境是至关重人工智能计算机视觉深度学习预训练上一篇终极ChatGPT-Midjourney完整使用指南从个人创作到商业应用的10个实战场景下一篇Rufus重磅更新zstd压缩格式让启动盘制作提速300%创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表