Stable Baselines3 模型导出实战:从 PyTorch 策略到 ONNX、C++、TensorFlow.js 与 Coral 的多框架部署指南

发布时间:2026/9/15 4:12:25
Stable Baselines3 模型导出实战:从 PyTorch 策略到 ONNX、C++、TensorFlow.js 与 Coral 的多框架部署指南 Stable Baselines3 模型导出实战从 PyTorch 策略到 ONNX、C、TensorFlow.js 与 Coral 的多框架部署指南【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3训练完成一个强化学习智能体之后真正的挑战往往才开始如何把它部署到另一个语言、另一个推理框架甚至是浏览器、移动端或边缘设备上本指南以 Stable Baselines3SB3官方导出文档为核心完整讲解策略Policy作为控制器的导出原理并给出从 ONNX、PyTorch JITC、ONNX Runtime Web、TensorFlow.js 到 TFLite/CoralEdge TPU的全链路可运行代码以及基于state_dict/get_parameters的手动导出和 SBXSB3 JAX经 PyTorch 中转导出的完整示例。读完本文你将能够在脱离 Gym 与 SB3 运行时的前提下在任何目标框架中完成智能体的推理。背景策略Policy就是控制器在 Stable Baselines3 中真正决定看到什么状态、输出什么动作的控制器存放在**策略policy**对象内部。每个学习算法DQN、A2C、PPO、SAC、TD3、DDPG 等都持有一个 policy 对象代表当前学到的行为可通过model.policy访问。策略中保存了推理predict action所需的全部信息——网络结构、权重参数、观察预处理逻辑——因此导出策略就等价于导出控制器。这一点在源码中非常清晰以 policy 基类 为例ActorCriticPolicy.forward()stable_baselines3/common/policies.py#L636对 PPO 这类 actor-critic 算法返回(actions, values, log_prob)三元组即动作、状态值和动作对数概率而所有推理的糖衣逻辑观察转换、图像归一化、动作后处理都在predict()方法中stable_baselines3/common/policies.py#L331-L386。导出时我们绕开predict()直接对 policy 模块本身做 trace 或转换为中间格式就能在其他框架复现同样的前向计算。提示可结合 examples.md 中的训练示例理解策略在完整训练流程中的角色。导出前必须弄清的两处隐式处理直接导出 policy 时有两处藏在内部的处理如果不手动补齐导出的模型在目标框架中会产生错误输出。CNN 观察的归一化内置于 policy 内部当使用 CNN 策略处理图像观察时观察会在预处理阶段被归一化除以 255把像素值缩放到 [0, 1]。这个预处理发生在 policy 内部preprocess_obs见 stable_baselines3/common/preprocessing.py#L91-L139因此导出后的模型会自动包含该归一化逻辑外部无需再除一次 255。这也意味着如果你在导出前手动把图像除以 255再送入导出的模型就会得到双重归一化。建议用原始 0–255 的像素值作为模型输入。# 预处理核心逻辑摘录自 stable_baselines3/common/preprocessing.py if normalize_images and is_image_space(observation_space): return obs.float() / 255.0另外要注意PyTorch 使用channel-first通道在前的布局Gymnasium 的图像观察通常是HWC高度、宽度、通道。若需要从外部传入图像可能需要先转置为 channel-firstmaybe_transpose在 stable_baselines3/common/policies.py#L236-L277 的obs_to_tensor中处理了这一逻辑。连续动作的后处理clip 或 unscale第二个隐藏步骤是连续动作的后处理。predict()在得到网络原始输出后还会对动作做一步修正stable_baselines3/common/policies.py#L372-L379若策略启用了squash_outputSAC 的 actor 默认用 tanh 把动作压到 [-1, 1]则调用unscale_action把动作从 [-1, 1] 重新映射回action_space的实际范围如 Pendulum-v1 的 [-2, 2]否则直接把动作np.clip到action_space.low/action_space.high避免高斯采样越界。这两个后处理步骤都不会出现在你直接导出的 ONNX/JIT 模型里导出的模型返回的是归一化/未缩放的原始网络输出。因此在目标框架推理后必须由你自己补上 clip 或 unscale 逻辑下文各小节会给出对应代码。导出到 ONNXPyTorch 2.0 / Opset 14如果使用 PyTorch 2.0 与 ONNX Opset 14可以非常轻松地把 SB3 策略导出为 ONNX。核心思路是把model.policy包进一个th.nn.Module包装类在forward中直接调用 policy 并固定deterministicTrue若要导出随机策略则改为deterministicFalse。再次强调以下导出的模型不包含连续动作的后处理步骤clip 或 unscale 到正确的动作空间。PPO 完整导出示例以 PPO 为例Pendulum-v1MlpPolicy完整代码如下——训练、保存、加载、包装、导出、用 onnxruntime 校验一气呵成import torch as th from typing import Tuple from stable_baselines3 import PPO from stable_baselines3.common.policies import BasePolicy class OnnxableSB3Policy(th.nn.Module): def __init__(self, policy: BasePolicy): super().__init__() self.policy policy def forward(self, observation: th.Tensor) - Tuple[th.Tensor, th.Tensor, th.Tensor]: # NOTE: Preprocessing is included, but postprocessing # (clipping/inscaling actions) is not, # If needed, you also need to transpose the images so that they are channel first # use deterministicFalse if you want to export the stochastic policy # policy() returns actions, values, log_prob for PPO return self.policy(observation, deterministicTrue) # Example: model PPO(MlpPolicy, Pendulum-v1) PPO(MlpPolicy, Pendulum-v1).save(PathToTrainedModel) model PPO.load(PathToTrainedModel.zip, devicecpu) onnx_policy OnnxableSB3Policy(model.policy) observation_size model.observation_space.shape dummy_input th.randn(1, *observation_size) th.onnx.export( onnx_policy, dummy_input, my_ppo_model.onnx, opset_version17, input_names[input], ) ##### Load and test with onnx import onnx import onnxruntime as ort import numpy as np onnx_path my_ppo_model.onnx onnx_model onnx.load(onnx_path) onnx.checker.check_model(onnx_model) observation np.zeros((1, *observation_size)).astype(np.float32) ort_sess ort.InferenceSession(onnx_path) actions, values, log_prob ort_sess.run(None, {input: observation}) print(actions, values, log_prob) # Check that the predictions are the same with th.no_grad(): print(model.policy(th.as_tensor(observation), deterministicTrue))几个值得注意的细节dummy input 的 batch 维度th.randn(1, *observation_size)显式带上 batch 维1这是 ONNX 导出所必需的固定形状。导出后如果用 batch size 为 1 推理可直接使用需要动态 batch 时可另行配置dynamic_axes。加载到 CPUPPO.load(..., devicecpu)避免 GPU 环境导出后在其他机器上因缺少 CUDA 而无法加载。校验一致性用onnx.checker.check_model做结构合法性检查再用相同输入对比 ONNX Runtime 与原始 PyTorch policy 的输出确保转换无误。由于 PPO 是 actor-critic 结构forward返回三个输出actions、values、log_probONNX 模型的输出也是三个。对于MultiInputPolicy字典观察如 goal-conditioned 任务导出方式类似社区在相关 issue 中有更详细的讨论可参考 GH#1873 中的方案。SAC只导出 actor 网络对 SAC过程类似但示例中只导出 actor 网络——因为 rollout 时 actor 足以决定动作不需要 critic。SAC 的 actor 输出是 tanh 压到 [-1, 1] 的缩放动作因此后处理unscale必不可少import torch as th from stable_baselines3 import SAC class OnnxablePolicy(th.nn.Module): def __init__(self, actor: th.nn.Module): super().__init__() self.actor actor def forward(self, observation: th.Tensor) - th.Tensor: # NOTE: You may have to postprocess (unnormalize) actions # to the correct bounds (see commented code below) return self.actor(observation, deterministicTrue) # Example: model SAC(MlpPolicy, Pendulum-v1) SAC(MlpPolicy, Pendulum-v1).save(PathToTrainedModel.zip) model SAC.load(PathToTrainedModel.zip, devicecpu) onnxable_model OnnxablePolicy(model.policy.actor) observation_size model.observation_space.shape dummy_input th.randn(1, *observation_size) th.onnx.export( onnxable_model, dummy_input, my_sac_actor.onnx, opset_version17, input_names[input], ) ##### Load and test with onnx import onnxruntime as ort import numpy as np onnx_path my_sac_actor.onnx observation np.zeros((1, *observation_size)).astype(np.float32) ort_sess ort.InferenceSession(onnx_path) scaled_action ort_sess.run(None, {input: observation})[0] print(scaled_action) # Post-process: rescale to correct space # Rescale the action from [-1, 1] to [low, high] # low, high model.action_space.low, model.action_space.high # post_processed_action low (0.5 * (scaled_action 1.0) * (high - low)) # Check that the predictions are the same with th.no_grad(): print(model.actor(th.as_tensor(observation), deterministicTrue))这里model.policy.actor就是 SAC 策略中的 actor 模块见 stable_baselines3/sac/policies.py 中SACPolicy通过self.make_actor()构建 actor 的实现stable_baselines3/sac/policies.py#L281。onnxable_model的输入是观察、输出是 [-1, 1] 区间的缩放动作拿到scaled_action后需按如下公式还原到动作空间真实范围low, high model.action_space.low, model.action_space.high post_processed_action low (0.5 * (scaled_action 1.0) * (high - low))这一公式与 SB3 内部的unscale_action完全一致stable_baselines3/common/policies.py#L402-L413。导出到 CPyTorch JIT Trace如果目标是把模型嵌入 C 推理代码可以用 PyTorch JIT 对模型进行trace、freeze 并优化保存为 TorchScript 文件然后在任意支持 TorchScript 的环境包括 libtorch C中加载推理# See ONNX export for imports and OnnxablePolicy jit_path sac_traced.pt # Trace and optimize the module traced_module th.jit.trace(onnxable_model.eval(), dummy_input) frozen_module th.jit.freeze(traced_module) frozen_module th.jit.optimize_for_inference(frozen_module) th.jit.save(frozen_module, jit_path) ##### Load and test with torch import torch as th dummy_input th.randn(1, *observation_size) loaded_module th.jit.load(jit_path) action_jit loaded_module(dummy_input)要点说明th.jit.trace(module.eval(), dummy_input)基于一组 dummy 输入记录执行路径生成计算图。务必先.eval()把 BatchNorm/Dropout 切换到推理模式保证 trace 出的图稳定。th.jit.freeze将权重固化进图中、移除梯度信息进一步减小体积。th.jit.optimize_for_inference针对推理做算子融合等优化适合部署场景。th.jit.load后再喂入同样的dummy_input即可在纯 TorchScript 环境下得到与 ONNX 一致的动作输出同样不含后处理。社区RL Zoo 项目中已有 C 导出相关的草案实现可供参考思路与本小节完全一致trace → freeze → 在 C 侧加载.pt。导出到 ONNX-JS / ONNX Runtime Web浏览器推理把模型跑在浏览器里是常见的部署需求官方推荐的做法是先用前面的方法导出 ONNX再用onnxruntime-web在浏览器中加载推理。社区有完整的端到端示例一个躲避汽车的驾驶环境流程为创建/训练一个 PPO 模型将模型导出为 ONNX同时把归一化统计量normalization stats存成 JSON在浏览器中用onnxruntime-web加载 ONNX 模型并在前端做同样的归一化达到与本地训练近似的效果。下面的最小示例演示转 ONNX → 浏览器无后处理推理以 SAC 为例导出部分与上文相同import torch as th from stable_baselines3 import SAC class OnnxablePolicy(th.nn.Module): def __init__(self, actor: th.nn.Module): super().__init__() self.actor actor def forward(self, observation: th.Tensor) - th.Tensor: # NOTE: You may have to postprocess (unnormalize or renormalize) return self.actor(observation, deterministicTrue) # Example: model SAC(MlpPolicy, Pendulum-v1) SAC(MlpPolicy, Pendulum-v1).save(PathToTrainedModel.zip) model SAC.load(PathToTrainedModel.zip, devicecpu) onnxable_model OnnxablePolicy(model.policy.actor) observation_size model.observation_space.shape dummy_input th.randn(1, *observation_size) th.onnx.export( onnxable_model, dummy_input, my_sac_actor.onnx, opset_version17, input_names[input], )前端 JavaScript 推理依赖通过npm install onnxruntime-web安装测试版本为 1.19也可用 CDN 引入// Install using npm install onnxruntime-web (tested with version 1.19) or using cdn import * as ort from onnxruntime-web; async function runInference() { const session await ort.InferenceSession.create(my_sac_actor.onnx); // The observation_size 3 (for Pendulum-v1) const inputData Float32Array.from([0.1, -0.2, 0.3]); const inputTensor new ort.Tensor(float32, inputData, [1, 3]); const results await session.run({ input: inputTensor }); const outputName session.outputNames[0]; const action results[outputName].data; console.log(Predicted action, action); } runInference();注意这里input名称必须与th.onnx.export时指定的input_names[input]一致Pendulum-v1 的观察维度为 3因此输入张量形状是[1, 3]。同样的输出动作是缩放/未归一化的原始值业务侧如需落到真实动作空间仍需自行后处理。导出到 TensorFlow.js把模型跑在 TensorFlow.js 中需要走一条较长的转换链SB3Torch⇒ ONNX ⇒ TensorFlow ⇒ TensorFlow.js中间涉及多个工具链的版本兼容问题因此本节给出经过验证的完整方案。注意截至 2025 年 11 月的信息onnx2tf尚不支持 TensorFlow.js因此必须改用tfjs-converter。但tfjs-converter目前维护不活跃要求使用较旧的 opset 与 TensorFlow 版本。关键约束是ONNX 的 opset 版本必须改为 14上文 ONNX 导出示例中opset_version17是为了更高版本下的更稳定用法而此处需要opset_version14。第一步仍是 SB3 ⇒ ONNX与上文一致把opset_version改为 14。随后在全新环境中安装指定版本的依赖已在 Python 3.10 下测试通过pip install --use-deprecatedlegacy-resolver tensorflow2.13.0 keras2.13.1 onnx1.16.0 onnx-tf1.9.0 tensorflow-probability0.21.0 tensorflowjs4.15.0 jax0.4.26 jaxlib0.4.26然后执行 ONNX ⇒ TensorFlow 的转换import onnx import onnx_tf.backend import tensorflow as tf ONNX_FILE_PATH my_sac_actor.onnx MODEL_PATH tf_model onnx_model onnx.load(ONNX_FILE_PATH) onnx.checker.check_model(onnx_model) print(onnx.helper.printable_graph(onnx_model.graph)) print(Converting ONNX to TF...) tf_rep onnx_tf.backend.prepare(onnx_model) tf_rep.export_graph(MODEL_PATH) # After this do not forget to use tensorflowjs_converter若目录结构正确且无报错再执行命令行转换把 TensorFlow SavedModel 转成 tfjs_graph_modeltensorflowjs_converter --input_formattf_saved_model --output_formattfjs_graph_model tf_model tfjs_model如果tensorflowjs_converter报错先升级 TensorFlow 相关包pip install --upgrade tensorflow tensorflow-decision-forests tensorflowjs再重试通常即可成功无需重跑上一步转换代码。前端加载tfjs_model目录中的model.json进行推理import * as tf from https://cdn.jsdelivr.net/npm/tensorflow/tfjs4.15.0/esm; // Post processing not included async function runInference() { const MODEL_URL ./tfjs_model/model.json; const model await tf.loadGraphModel(MODEL_URL); // Observation_size is 3 for Pendulum-v1 const inputData [1.0, 0.0, 0.0]; const inputTensor tf.tensor2d([inputData], [1, 3]); const resultTensor model.execute(inputTensor); const action await resultTensor.data(); console.log(Predicted action, action); inputTensor.dispose(); resultTensor.dispose(); } runInference();tf.loadGraphModel加载的是图模型graph model用model.execute(inputTensor)执行前向并读取结果。注意代码同样省略了动作后处理且观察输入[1, 3]对应 Pendulum-v1 的 3 维状态。导出到 TFLite / CoralEdge TPUGoogle 为边缘端 AI 部署推出了Coral芯片它有多种形态包括 USB 版本。把 SB3 训练的模型跑在树莓派 Coral USB 上正是社区相关示例的最初动机。Coral 芯片推理快、功耗极低但设备端训练能力有限所以需要先把网络量化到与 Coral 能力匹配的形式。从 SB3 到 Coral 的完整链路为SB3 (Torch) ⇒ ONNX ⇒ TensorFlow ⇒ TFLite ⇒ Coral社区有完整的最小示例覆盖了整条导出链路并演示了多数导出变体的前向推理同时专门处理了以下容易踩坑的问题让 Gym 的观察能够正确适配 ONNX形状、dtype、归一化合理量化 TFLite 模型既对齐 Gym 的动作语义又充分利用 Coral 的加速能力复用前文介绍的OnnxablePolicy包装类完成第一步导出。整体思路与 TensorFlow.js 一节类似先在opset_version14下完成 SB3 ⇒ ONNX再用onnx-tf转成 TensorFlow最后用 TFLite 转换器生成量化模型并部署到 Coral 设备。手动导出state_dict 与 get_parameters如果不依赖任何转换工具也可以手动导出需要的参数权重在目标框架中自行重建网络。SB3 提供了两种取参方式model.get_parameters()返回 agent 全部网络的 state-dict 映射按对象名组织的字典见 stable_baselines3/common/base_class.py#L804-L817。它基于_get_torch_save_params()收集所有需要保存的模块策略、critic 等并逐一导出state_dict()如果还需要访问优化器的状态字典就必须用这个接口。model.policy.state_dict()policy 本身也是 PyTorchnn.Module可以直接调用标准的state_dict()拿到网络参数。关于架构信息每个算法的网络结构请查看各自目录下的policies.py例如 stable_baselines3/ppo/policies.py、stable_baselines3/sac/policies.py、stable_baselines3/dqn/policies.py 等里面定义了 MLP 层数、激活函数、特征提取器如NatureCNN、MlpExtractor等结构细节据此即可在目标框架中用相同层序重建前向计算。建议大多数情况下优先使用 PyTorch 标准的state_dict()/load_state_dict()只涉及网络参数只有当确实需要优化器状态时才改用get_parameters()。SBXSB3 JAX导出到 ONNX作为手动导出的典型示例Stable Baselines JaxSBX的策略可以借助中间 PyTorch 表示导出到 ONNX先把 JAX 的 Flax 参数字典映射成 PyTorchstate_dict重建等价网络再走标准th.onnx.export。这样不仅演示了跨框架取参也示范了 SBX 中 actor 网络的内部结构。import numpy as np import sbx import torch as th class TorchPolicy(th.nn.Module): def __init__(self, obs_dim: int, hidden_dim: int, act_dim: int): super().__init__() self.net th.nn.Sequential( th.nn.Linear(obs_dim, hidden_dim), th.nn.Tanh(), th.nn.Linear(hidden_dim, hidden_dim), th.nn.Tanh(), th.nn.Linear(hidden_dim, act_dim), ) def forward(self, x: th.Tensor) - th.Tensor: return self.net(x) model sbx.PPO(MlpPolicy, Pendulum-v1) # Also possible: load a trained model # model sbx.PPO.load(PathToTrainedModel.zip) params model.policy.actor_state.params[params] # For debug: print( SBX params ) for key, value in params.items(): if isinstance(value, dict): for name, val in value.items(): print(f{key}.{name}: {val.shape}, end ) else: print(f{key}: {value.shape}, end ) print(\n * 20 \n) obs_dim model.observation_space.shape act_dim model.action_space.shape # Number of units in the hidden layers (assume a network architecture like [64, 64]) hidden_dim params[Dense_0][kernel].shape[1] # map params to torch state_dict keys num_layers len([k for k in params.keys() if k.startswith(Dense_)]) state_dict {} for i in range(num_layers): layer_name fDense_{i} state_dict[fnet.{i * 2}.bias] th.from_numpy(np.array(params[layer_name][bias])) state_dict[fnet.{i * 2}.weight] th.from_numpy(np.array(params[layer_name][kernel].T)) torch_policy TorchPolicy(obs_dim[0], hidden_dim, act_dim[0]) print( Torch params ) print( .join(f{key}:{tuple(value.shape)} for key, value in torch_policy.named_parameters())) print( * 20 \n) torch_policy.load_state_dict(state_dict) torch_policy.eval() dummy_input th.zeros((1, *obs_dim)) # Use normal Torch export th.onnx.export( torch_policy, (dummy_input,), my_ppo_actor.onnx, opset_version18, input_names[input], output_names[action], ) ##### Load and test with onnx import onnxruntime as ort onnx_path my_ppo_actor.onnx ort_sess ort.InferenceSession(onnx_path) observation np.random.random((1, *obs_dim)).astype(np.float32) action ort_sess.run(None, {input: observation})[0] print(action) sbx_action, _ model.predict(observation, deterministicTrue) with th.no_grad(): torch_action torch_policy(th.as_tensor(observation)) # Check that the predictions are the same assert np.allclose(sbx_action, action) assert np.allclose(sbx_action, torch_action.numpy())这段代码有三个关键点值得展开参数来源SBX 的 actor 权重位于model.policy.actor_state.params[params]按 Flax 惯例以Dense_0、Dense_1等命名层每个 Dense 层包含kernel与bias两个数组。打印各参数形状便于对照调试。映射细节Flax 的kernel形状是(in_features, out_features)而 PyTorchnn.Linear的weight形状是(out_features, in_features)因此必须做转置kernel.Tbias 直接复制。重建的网络结构两层 64 隐层 Tanh 激活 输出层要严格对应 SBX PPO 默认的[64, 64]架构hidden_dim从params[Dense_0][kernel].shape[1]动态读取。一致性校验load_state_dict后分别跑 SBX 原生predict、PyTorch 重建网络、ONNX Runtime 三个前向用np.allclose断言三者输出一致确保跨框架取参 → 重建 → 转 ONNX整条链路没有偏差。小结Stable Baselines3 的模型导出本质上围绕一个事实展开策略对象持有推理所需的全部信息。从导出目标反推选择对应的技术路径即可部署目标导出路径关键注意点通用跨框架ONNXpolicy/actor ⇒ ONNXOpset 14PyTorch 2.0输出为原始/缩放动作需自行补 clip 或 unscaleClibtorchJIT trace ⇒ freeze ⇒ optimize_for_inference ⇒.pt导出前.eval()用固定 dummy input浏览器onnxruntime-webONNX ⇒ JS 前端加载input_names需与前端键名一致浏览器TensorFlow.jsONNXOpset 14⇒ TensorFlow ⇒ tfjs依赖旧版 TensorFlow/tfjs-converter 工具链边缘设备CoralONNX ⇒ TensorFlow ⇒ TFLite ⇒ 量化 ⇒ Coral观察适配与量化是主要坑点任意框架手动state_dict()/get_parameters() 重建网络网络架构参考各算法目录的policies.pySBXJAXFlax params ⇒ PyTorch 重建 ⇒ ONNXkernel 需转置注意 Flax 与 PyTorch 形状约定差异无论走哪条路径请牢记两点CNN 策略的图像归一化除以 255已在 policy 内部完成外部不要再重复归一化连续动作的 clip/unscale 后处理不会随模型导出需要在目标框架的推理侧自行实现。掌握这些原则后SB3 训练出的智能体即可无缝迁移到 C、浏览器、移动端与边缘设备真正打通训练—部署的最后一步。【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考