新闻详情

新闻详情

首页 / 资讯中心 / 详情

Stable Baselines3 策略网络完全指南:从默认架构到自定义 Policy 与特征提取器

发布时间:2026/9/14 20:46:39来源:尧图网络
Stable Baselines3 策略网络完全指南:从默认架构到自定义 Policy 与特征提取器
Stable Baselines3 策略网络完全指南从默认架构到自定义 Policy 与特征提取器【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3Stable Baselines3SB3将深度强化学习算法中的策略网络抽象为可组合的两段式结构特征提取器Features Extractor加全连接网络net_arch并针对图像、向量、字典观测分别提供CnnPolicy、MlpPolicy与MultiInputPolicy。本文以官方指南 docs/guide/custom_policy.md 为骨架结合仓库内 torch_layers.py、policies.py 等源码与 test_custom_policy.py 测试系统讲解 SB3 策略网络的默认结构、policy_kwargs配置、自定义网络架构、自定义特征提取器以及多输入字典观测的处理方式读完即可掌握从调参到完全自定义策略的全部实战手法。SB3 Policy 的两段式结构在 SB3 中策略policy这个称呼沿用了强化学习术语但实际上它指代的是负责训练所需的全部网络及配套优化器的类——不仅包含用于预测动作的网络即学习到的控制器还包含价值网络、目标网络等。官方文档明确指出这是一种相对于 RL 术语的用语滥用abuse of language。一个 SB3 Policy 内部被拆成两个主要部分特征提取器Features Extractor负责把高维观测转换成特征向量。例如用 CNN 从图像中提取特征。它对应features_extractor_class参数并可通过features_extractor_kwargs修改该提取器的默认参数。全连接网络把特征映射到动作或价值其结构由net_arch参数控制。值得注意的是所有观测在送入特征提取器之前都会先经过预处理见 preprocessing.py 的preprocess_obs图像会除以 255 归一化到[0, 1]离散观测会被转换成 one-hot 向量对于普通的向量观测特征提取器退化为一个Flatten层即 FlattenExtractor。策略的组成Actor/Critic 目标网络 优化器SB3 的 Policy 通常由多个网络组合而成actor策略网络与 critic价值网络必要时还包括目标网络以及各自关联的优化器。每个网络都遵循特征提取器 全连接网络的同一模式如下图所示上图中右上角的注释进一步澄清了术语在 RL 文献中policy通常只指预测动作的 actor而在 SB3 中policy是指管理训练所需全部网络的类。默认网络架构随算法与观测空间变化SB3 的默认网络架构取决于算法和观测空间的类型。官方文档给出的速查表如下观测空间PPO / A2C / DQNSACTD3 / DDPG1D 向量观测2 层全连接每层 64 单元2 层全连接每层 256 单元[400, 300]单元源自 TD3 原始论文图像观测Nature CNN 线性层Nature CNN 保持原有全连接网络Nature CNN 保持原有全连接网络字典观测混合图像走 CNN、向量走全连接CNN 输出尺寸更小同上同上这些默认值可以在源码中找到依据ActorCriticPolicy.init当net_arch is None时若特征提取器是NatureCNN则net_arch []CNN 之后只接线性层否则net_arch dict(pi[64, 64], vf[64, 64])SACPolicy.init默认net_arch [256, 256]DQNPolicy.init与 ActorCritic 相同CNN 时为空列表否则[64, 64]。图像观测Nature CNN对于图像观测空间SB3 默认使用 Nature CNN源自 DQN Nature 论文其结构定义在 torch_layers.py 的 NatureCNN三层卷积Conv2d(8,4,0) → Conv2d(4,2,0) → Conv2d(3,1,0)每层后接 ReLU 并最终Flatten最后接一个nn.Linear(n_flatten, features_dim)加 ReLU。它要求观测空间必须是spaces.Box且满足图像空间检查默认要求uint8dtype、取值范围[0, 255]若使用VecNormalize或已经归一化的通道优先图像应传入normalize_imagesFalse此时会自动把normalized_imageTrue写入特征提取器参数见 BaseModel.init。特征提取器是否共享On-policy 与 Off-policy 的分野On-policy 算法A2C/PPO默认在 actor 和 critic 之间共享特征提取器以减少计算量share_features_extractorTrueOff-policy 算法TD3/DDPG/SACactor 与 critic 各用独立的特征提取器因为实践表明这种配置效果最好源码默认share_features_extractorFalse见 SACPolicy。这一共享行为可以在 ActorCriticPolicy.init中看到共享时pi_features_extractor与vf_features_extractor指向同一个实例不共享时则分别创建。对 On-policy 算法你也可以通过policy_kwargs里的share_features_extractorFalse强制关闭共享On/Off-policy 算法都支持该参数。可视化策略结构默认架构不满足需求时可以用print(model.policy)直接打印策略网络的完整结构便于逐层核对参数与维度。自定义网络架构通过 policy_kwargs 调参最常用的定制方式是创建模型时传入policy_kwargs参数。其中net_arch既可以是列表actor 与 critic 同构也可以是字典二者异构。import gymnasium as gym import torch as th from stable_baselines3 import PPO # 自定义 actor (pi) 和 value function (vf) 网络 # 各为两层、每层 32 个单元、ReLU 激活 # 注意pi 和 vf 网络之上还会各自动追加一层线性输出层 policy_kwargs dict(activation_fnth.nn.ReLU, net_archdict(pi[32, 32], vf[32, 32])) # 创建智能体 model PPO(MlpPolicy, CartPole-v1, policy_kwargspolicy_kwargs, verbose1) # 获取环境 env model.get_env() # 训练智能体 model.learn(total_timesteps20_000) # 保存智能体 model.save(ppo_cartpole) del model # policy_kwargs 会自动随模型加载 model PPO.load(ppo_cartpole, envenv)额外的线性输出层必须记住一个重要细节SB3 会在net_arch指定的层之上再额外添加一层线性层用于输出正确维度的动作分布参数并应用合适的激活函数例如离散动作的 Softmax。以 CartPole观测维度 4、动作维度 2为例net_archdict(pi[32, 32], vf[32, 32])的最终结构为obs 4 / \ 32 32 | | 32 32 | | 2 1 action value这一行为在 ActorCriticPolicy._build 中体现action_net与value_net分别由latent_dim_pi、latent_dim_vf构造前者对接动作分布离散则为 logits、连续则结合log_std后者固定输出 1 维价值。激活函数默认值差异On-policyA2C/PPO默认activation_fnnn.TanhOff-policySAC/TD3/DDPG与 DQN默认activation_fnnn.ReLU。连续动作的边界处理有一个常被忽略但很关键的行为差异A2C 和 PPO 对连续动作会在训练与测试时进行 clip避免超出动作空间边界报错而SAC、DDPG、TD3 则用tanh()变换对动作进行 squash能更正确地处理边界。这一逻辑体现在 BasePolicy.predictsquash_outputTrue时调用unscale_action反缩放回[low, high]否则直接用np.clip截断。后者依赖高斯采样的均值如 SAC 的 Actor 输出mu与log_std前者的分布则可能产生越界样本。自定义特征提取器继承 BaseFeaturesExtractor当内置的FlattenExtractor向量或NatureCNN图像不满足需求时可以自定义特征提取器。做法是派生BaseFeaturesExtractor实现__init__与forward然后通过features_extractor_class/features_extractor_kwargs传入模型。import torch as th import torch.nn as nn from gymnasium import spaces from stable_baselines3 import PPO from stable_baselines3.common.torch_layers import BaseFeaturesExtractor class CustomCNN(BaseFeaturesExtractor): :param observation_space: (gym.Space) :param features_dim: (int) Number of features extracted. This corresponds to the number of unit for the last layer. def __init__(self, observation_space: spaces.Box, features_dim: int 256): super().__init__(observation_space, features_dim) # 假设是 CxHxW 图像通道在前 # 重排序会由预处理器或 wrapper 完成 n_input_channels observation_space.shape[0] self.cnn nn.Sequential( nn.Conv2d(n_input_channels, 32, kernel_size8, stride4, padding0), nn.ReLU(), nn.Conv2d(32, 64, kernel_size4, stride2, padding0), nn.ReLU(), nn.Flatten(), ) # 通过一次前向传播计算展平后的维度 with th.no_grad(): n_flatten self.cnn( th.as_tensor(observation_space.sample()[None]).float() ).shape[1] self.linear nn.Sequential(nn.Linear(n_flatten, features_dim), nn.ReLU()) def forward(self, observations: th.Tensor) - th.Tensor: return self.linear(self.cnn(observations)) policy_kwargs dict( features_extractor_classCustomCNN, features_extractor_kwargsdict(features_dim128), ) model PPO(CnnPolicy, BreakoutNoFrameskip-v4, policy_kwargspolicy_kwargs, verbose1) model.learn(1000)BaseFeaturesExtractor 源码要点BaseFeaturesExtractor 本质上就是带observation_space与features_dim约定的nn.Module构造函数要求features_dim 0通过features_dim属性向外暴露输出特征维度SB3 会调用make_features_extractor()见 policies.py实例化特征提取器并把features_dim传入后续网络构造。自定义提取器内用一次前向传播计算展平维度的写法与 NatureCNN 的实现 完全一致是处理卷积输出尺寸不确定时的标准做法。特征提取器如何被使用观测进入网络后会先经过 preprocess_obs 预处理归一化、one-hot 等再由extract_features交给特征提取器最后进入MlpExtractor。on-policy 默认共享特征提取器关闭共享后ActorCriticPolicy.extract_features 会分别为 actor 与 critic 各跑一次特征提取并返回元组(pi_features, vf_features)。多输入与字典观测MultiInputPolicy 与 CombinedExtractorSB3 通过 Gymnasium 的Dict空间原生支持多输入观测。此时使用MultiInputPolicy其默认特征提取器是CombinedExtractor负责把多个输入合并成一个向量再交给net_arch网络。CombinedExtractor 的处理流程即官方文档所述的三步规则图像输入通过is_image_space见 preprocessing.py自动检测若是图像则用 Nature CNN 处理输出 256 维 latent 向量非图像输入直接展平不加层拼接把所有子向量th.cat(..., dim1)拼成一个长向量送入策略网络。其构造参数cnn_output_dim256用于控制每个 CNN 子模块的输出维度默认 256避免网络过大normalized_image控制图像校验行为。自定义多输入特征提取器下面的例子假设环境的观测空间字典包含两个键image 是(1,H,W)的单通道图像通道在前vector 是(D,)维向量。我们对 image 做简单的 4×4 下采样对 vector 接一个线性层import gymnasium as gym import torch as th from torch import nn from stable_baselines3.common.torch_layers import BaseFeaturesExtractor class CustomCombinedExtractor(BaseFeaturesExtractor): def __init__(self, observation_space: gym.spaces.Dict): # 在遍历完所有子空间之前我们不知道最终特征维度 # 所以先用一个占位值PyTorch 要求先调用 nn.Module.__init__ super().__init__(observation_space, features_dim1) extractors {} total_concat_size 0 # 需要知道这个提取器的输出大小所以遍历所有子空间并计算输出特征尺寸 for key, subspace in observation_space.spaces.items(): if key image: # 对图像的单通道做 4x4 下采样后展平 extractors[key] nn.Sequential(nn.MaxPool2d(4), nn.Flatten()) total_concat_size subspace.shape[1] // 4 * subspace.shape[2] // 4 elif key vector: # 过一个简单的 MLP extractors[key] nn.Linear(subspace.shape[0], 16) total_concat_size 16 self.extractors nn.ModuleDict(extractors) # 手动更新特征维度 self._features_dim total_concat_size def forward(self, observations) - th.Tensor: encoded_tensor_list [] # self.extractors 里是完成具体处理的 nn.Module for key, extractor in self.extractors.items(): encoded_tensor_list.append(extractor(observations[key])) # 返回 (B, self._features_dim) 的 PyTorch 张量B 是批维度 return th.cat(encoded_tensor_list, dim1)注意源码中的关键细节由于nn.Module.__init__必须先于子模块注册执行这里先以features_dim1占位遍历完所有子空间后再手动更新self._features_dimtorch_layers.py 的CombinedExtractor也采用同样的先占位、后修正模式。On-policy 算法PPO / A2C的自定义网络异构 actor/criticdict(pi, vf)PPO、A2C 若需要 actor 与 critic 使用不同架构用字典形式指定policy_kwargs dict(net_archdict(pi[32, 32], vf[64, 64]))等价写法actor 与 critic 结构相同可直接用列表net_arch[128, 128]即两个 128 单元隐藏层等价于dict(pi[128, 128], vf[128, 128])。两种写法的效果对比net_arch[128, 128]actor/critic 同构obs / \ 128 128 | | 128 128 | | action valuenet_archdict(pi[32, 32], vf[64, 64])异构obs / \ 32 64 | | 32 64 | | action value这些输入形式由 MlpExtractor 解析字典形式时pi_layers_dims net_arch.get(pi, [])、vf_layers_dims net_arch.get(vf, [])缺省键按空列表处理即退化为线性网络列表形式则两者共用同一列表。测试 test_flexible_mlp 覆盖了[]、[4]、[4, 4]、dict(vf[16], pi[8])等各类组合其中旧版格式[dict(pi..., vf...)]会触发UserWarning因为自 SB3 1.8.0 起共享层已被移除policies.py 会自动把这种旧格式转换为字典。高级示例完全自定义策略网络如果需要更细粒度的控制可以直接重定义策略类。核心要点继承ActorCriticPolicy重写_build_mlp_extractor自定义网络需保存latent_dim_pi与latent_dim_vf供上层创建动作分布与价值输出层使用可借此关闭正交初始化kwargs[ortho_init] False。from typing import Callable, Dict, List, Optional, Tuple, Type, Union from gymnasium import spaces import torch as th from torch import nn from stable_baselines3 import PPO from stable_baselines3.common.policies import ActorCriticPolicy class CustomNetwork(nn.Module): 策略与价值函数的自定义网络。 输入为特征提取器提取的特征。 :param feature_dim: 特征提取器输出的特征维度如 CNN 的特征 :param last_layer_dim_pi: (int) 策略网络最后一层的单元数 :param last_layer_dim_vf: (int) 价值网络最后一层的单元数 def __init__( self, feature_dim: int, last_layer_dim_pi: int 64, last_layer_dim_vf: int 64, ): super().__init__() # 重要保存输出维度用于创建动作分布 self.latent_dim_pi last_layer_dim_pi self.latent_dim_vf last_layer_dim_vf # 策略网络 self.policy_net nn.Sequential( nn.Linear(feature_dim, last_layer_dim_pi), nn.ReLU() ) # 价值网络 self.value_net nn.Sequential( nn.Linear(feature_dim, last_layer_dim_vf), nn.ReLU() ) def forward(self, features: th.Tensor) - Tuple[th.Tensor, th.Tensor]: :return: (th.Tensor, th.Tensor) latent_policy, latent_value 如果所有层共享则 latent_policy latent_value return self.forward_actor(features), self.forward_critic(features) def forward_actor(self, features: th.Tensor) - th.Tensor: return self.policy_net(features) def forward_critic(self, features: th.Tensor) - th.Tensor: return self.value_net(features) class CustomActorCriticPolicy(ActorCriticPolicy): def __init__( self, observation_space: spaces.Space, action_space: spaces.Space, lr_schedule: Callable[[float], float], *args, **kwargs, ): # 关闭正交初始化 kwargs[ortho_init] False super().__init__( observation_space, action_space, lr_schedule, # 剩余参数传给基类 *args, **kwargs, ) def _build_mlp_extractor(self) - None: self.mlp_extractor CustomNetwork(self.features_dim) model PPO(CustomActorCriticPolicy, CartPole-v1, verbose1) model.learn(5000)官方文档指出如果需要在 actor 与 critic 之间共享层就必须走这种自定义策略网络的路线。默认的MlpExtractor对 actor/critic 的 MLP 是完全独立的policy_net与value_net各成一支见 torch_layers.py。底层细节正交初始化默认情况下 A2C/PPO 采用正交初始化ortho_initTrue且各模块使用不同的增益特征提取器与 MLP 为sqrt(2)、action_net为0.01、value_net为1见 policies.py。若自定义策略时不需要这种初始化如高级示例那样设置ortho_initFalse即可。Off-policy 算法SAC / DDPG / TD3的自定义网络Off-policy 算法使用dict(pi[...], qf[...])结构分别指定 actorpi与 criticQ 函数qf的架构from stable_baselines3 import SAC # 自定义 actor 架构两个 64 单元层 # 自定义 critic 架构400 与 300 单元的两层 policy_kwargs dict(net_archdict(pi[64, 64], qf[400, 300])) # 创建智能体 model SAC(MlpPolicy, Pendulum-v1, policy_kwargspolicy_kwargs, verbose1) model.learn(5000)若 actor 与 critic 同构直接写net_arch[256, 256]即可。解析逻辑与约束字典形式由 get_actor_critic_arch 解析要求必须同时提供pi与qf两个键否则抛出AssertionError。与 on-policy 不同off-policy 算法不允许 actor 与 critic 之间存在除特征提取器之外的共享层这是为了避免目标网络相关的梯度问题源码 docstring 对此有明确说明。Critic 的实现ContinuousCriticOff-policy 的 criticQ 函数与 on-policy 的价值网络有本质区别它以观测特征与动作拼接后的向量th.cat([features, actions], dim1)为输入输出单个 Q 值Q(s, a)。默认创建两个 critic 网络n_critics2以配合 clipped Q-learning 抑制过估计见 ContinuousCritic。SAC 的 actor 输出则经tanhsquashSquashedDiagGaussianDistributionsac/policies.py并默认将log_std限制在[-20, 2]区间。更多定制入口除net_arch外policy_kwargs还支持activation_fn、optimizer_class/optimizer_kwargs、n_critics、share_features_extractor等参数测试 test_custom_optimizer 即演示了optimizer_classth.optim.AdamW的用法。对 off-policy 算法更进阶的定制建议直接阅读对应算法的policies.py源码并需对算法本身有充分理解。实战检查清单结构认知策略 特征提取器features_extractor_class/features_extractor_kwargs 全连接网络net_arch所有观测先预处理后进入特征提取器快速调整改层数/宽度用policy_kwargs dict(net_arch[...])同构或dict(pi[...], vf[...])on-policy 异构/dict(pi[...], qf[...])off-policy 异构记得会自动追加输出层图像任务继承BaseFeaturesExtractor自定义 CNN用features_dim与一次前向传播确定展平维度传入features_extractor_class多模态任务用Dict空间 MultiInputPolicy默认CombinedExtractor自动识别图像/向量需要更复杂的分支处理时继承BaseFeaturesExtractor自定义完整重写on-policy 继承ActorCriticPolicy重写_build_mlp_extractor可关闭正交初始化边界动作A2C/PPO 用 clipSAC/DDPG/TD3 用tanhsquash二者机制不同自定义环境时需留意验证print(model.policy)打印网络结构核对维度仓库测试 test_custom_policy.py 覆盖了从net_arch组合到自定义优化器、create_mlp单元级验证的多种场景可作为自定义后回归验证的参考。【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

在半导体制造的后道工序中,半导体测试机与分选机或探针台的自动化协同,是整个封测厂(OSAT)上位机系统(EAP / Tester Controller / ATE-MES)的核心 2026/9/14 21:19:42

在半导体制造的后道工序中,半导体测试机与分选机或探针台的自动化协同,是整个封测厂(OSAT)上位机系统(EAP / Tester Controller / ATE-MES)的核心

在半导体制造的后道工序(Back-End / ATE Test)中,半导体测试机(Tester,如 Advantest、Teradyne)与分选机(Handler,如 Delta, Cohu)或探针台(Prober,如 TEL, Accretech)的自动化协同,是整个封测厂(OSAT)上位机系统(EAP / Tester Controller / ATE-MES)的核心。…

阅读更多 →
直流配电网最优潮流建模与YALMIP实现 2026/9/14 21:19:42

直流配电网最优潮流建模与YALMIP实现

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
贪心算法与单调栈:数字删除问题的最优解 2026/9/14 21:19:42

贪心算法与单调栈:数字删除问题的最优解

1. 问题背景与核心需求这个算法问题看似简单,却蕴含着典型的贪心算法思想。给定一个12位整数,我们需要删除其中8个数字(保留4个),使得剩下的数字按原始顺序排列时,形成的4位数是所有可能组合中最小的那个。…

阅读更多 →
GitLab CI/CD配置文件.gitlab-ci.yml详解与实战 2026/9/14 21:19:42

GitLab CI/CD配置文件.gitlab-ci.yml详解与实战

1. .gitlab-ci.yml文件的核心作用解析.gitlab-ci.yml是GitLab CI/CD流水线的核心配置文件,它定义了自动化构建、测试和部署的完整流程。这个YAML格式的文件需要放置在项目根目录下,GitLab Runner会根据其中的指令自动执行预设任务。关键特性:…

阅读更多 →
Kubernetes镜像自动化构建与部署的Shell脚本实践 2026/9/14 21:19:42

Kubernetes镜像自动化构建与部署的Shell脚本实践

1. 项目背景与核心需求在容器化部署成为主流的今天,Kubernetes(k8s)集群中的镜像管理一直是运维工作的重点环节。每次代码更新后,开发团队需要手动构建Docker镜像、打标签、推送到私有仓库,再通知运维人员更新部署——…

阅读更多 →
免费学术查重工具评测与使用指南 2026/9/14 21:16:41

免费学术查重工具评测与使用指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞