You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何区分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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 11:58:13