如何区分Stable-Baselines3中Actor与Critic对共享特征提取器的调用?
区分Actor/Critic对共享特征提取器的调用并打印日志
要实现仅在Actor调用共享特征提取器时打印日志,这里提供两种可靠的方案:
方案一:自定义策略类,传递调用标识
通过继承Stable-Baselines3的内置策略类,在Actor和Critic调用特征提取器时传入不同的标识参数,让特征提取器明确调用来源。
1. 修改特征提取器,新增调用来源参数
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 from stable_baselines3.common.policies import CnnPolicy 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) n_input_channels = observation_space.shape[0] self.cnn = nn.Sequential( nn.Conv2d(n_input_channels, 32, kernel_size=8, stride=4, padding=0), nn.ReLU(), nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=0), 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, caller: str = "") -> th.Tensor: # 仅在Actor调用时打印日志 if caller == "actor": print(f"Actor调用特征提取器,输入形状: {observations.shape}") return self.linear(self.cnn(observations))
2. 自定义策略类,重写调用逻辑
继承CnnPolicy,在Actor和Critic的特征提取调用中传入对应的标识:
class CustomCnnPolicy(CnnPolicy): def _get_actor_forward(self, obs: th.Tensor) -> th.Tensor: # Actor调用时传递"actor"标识 features = self.extract_features(obs, caller="actor") return self.mlp_extractor.forward_actor(features) def _get_critic_forward(self, obs: th.Tensor) -> th.Tensor: # Critic调用时传递"critic"标识(或留空) features = self.extract_features(obs, caller="critic") return self.mlp_extractor.forward_critic(features) # 重写extract_features以支持传递caller参数 def extract_features(self, obs: th.Tensor, caller: str = "") -> th.Tensor: return self.features_extractor(obs, caller=caller)
3. 初始化PPO模型
policy_kwargs = dict( features_extractor_class=CustomCNN, features_extractor_kwargs=dict(features_dim=128), ) # 使用自定义策略类 model = PPO(CustomCnnPolicy, "BreakoutNoFrameskip-v4", policy_kwargs=policy_kwargs, verbose=1) model.learn(1000)
方案二:通过调用栈判断来源
如果不想修改策略类,可以通过Python的inspect模块追踪调用栈,识别调用来源。这种方法无需修改策略逻辑,但依赖Stable-Baselines3的内部函数命名,版本更新后可能需要调整:
import inspect class CustomCNN(BaseFeaturesExtractor): # __init__部分保持与原代码一致... def forward(self, observations: th.Tensor) -> th.Tensor: # 遍历调用栈,查找Actor相关的调用函数 for frame_info in inspect.stack(): # 匹配Actor相关的函数名(如_get_actor_forward、forward_actor) if "actor" in frame_info.function.lower(): print(f"Actor调用特征提取器,输入形状: {observations.shape}") break return self.linear(self.cnn(observations))
方案对比
- 方案一:稳定性高,不依赖内部实现细节,适合长期使用;
- 方案二:代码改动小,但如果Stable-Baselines3更新内部函数名,需要同步调整判断逻辑;
- 注意:训练时频繁打印日志会影响性能,建议添加频率控制(比如每N步打印一次)。
内容的提问来源于stack exchange,提问作者JayJona
相关产品推荐
相关产品推荐

