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

如何将Stable-Baselines3训练的A2C模型导出为PyTorch模型?

问题:如何将Stable-Baselines3的A2C模型导出为PyTorch模型?

我用Stable-Baselines3训练了一个基于MlpPolicy的A2C模型(刚入门强化学习,觉得这个工具很适合新手)。现在想通过可解释强化学习(XRL)的DeepSHAP方法来理解模型——我熟悉SHAP框架,而且DeepSHAP的实现也很完善。因为DeepSHAP基于PyTorch(Stable-Baselines3的底层框架),所以目标是提取底层的PyTorch模型,但过程中遇到了问题。

我找过相关讨论线程,但里面的模型架构和A2C不一样,只得到了部分帮助,没彻底解决问题。

我知道Stable-Baselines3支持用model.policy.state_dict()导出模型,但没法成功导入导出的内容。

打印A2C_model.policy后,看到的PyTorch模型结构是:

ActorCriticPolicy(
  (features_extractor): FlattenExtractor(
    (flatten): Flatten(start_dim=1, end_dim=-1)
  )
  (pi_features_extractor): FlattenExtractor(
    (flatten): Flatten(start_dim=1, end_dim=-1)
  )
  (vf_features_extractor): FlattenExtractor(
    (flatten): Flatten(start_dim=1, end_dim=-1)
  )
  (mlp_extractor): MlpExtractor(
    (policy_net): Sequential(
      (0): Linear(in_features=49, out_features=64, bias=True)
      (1): Tanh()
      (2): Linear(in_features=64, out_features=64, bias=True)
      (3): Tanh()
    )
    (value_net): Sequential(
      (0): Linear(in_features=49, out_features=64, bias=True)
      (1): Tanh()
      (2): Linear(in_features=64, out_features=64, bias=True)
      (3): Tanh()
    )
  )
  (action_net): Linear(in_features=64, out_features=5, bias=True)
  (value_net): Linear(in_features=64, out_features=1, bias=True)
)

我尝试根据这个结构自己在PyTorch里重建模型,但因为对PyTorch不熟悉没成功,写的代码如下:

import torch as th
import torch.nn as nn

class PyTorchMlp(nn.Module):  
        def __init__(self):
                nn.Module.__init__(self)

                n_inputs = 49
                n_actions = 5
        
                self.features_extractor = nn.Flatten(start_dim = 1, end_dim = -1)
        
                self.pi_features_extractor = nn.Flatten(start_dim = 1, end_dim = -1)
        
                self.vf_features_extractor = nn.Flatten(start_dim = 1, end_dim = -1)
        
                self.mlp_extractor = nn.Sequentail(
                    self.policy_net = nn.Sequential(
                        nn.Linear(in_features = n_inputs, out_features = 64),
                        nn.Tanh(),
                        nn.Linear(in_features = 64, out_features = 64),
                        nn.Tanh()
                    ),
        
                    self.value_net = nn.Sequential(
                        nn.Linear(in_features = n_inputs, out_features = 64),
                        nn.Tanh(),
                        nn.Linear(in_features = 64, out_features = 64),
                        nn.Tanh()
                    )
                )
        
                self.action_net = nn.Linear(in_features = 64, out_features = 5)
        
                self.value_net = nn.Linear(in_features = 64, out_features = 1)
        
    
            def forward(self, x):
                pass

现在想问:如何将Stable-Baselines3模型导出为PyTorch模型?


解决方法

方法1:直接使用原模型的policy(无需重建)

Stable-Baselines3的A2C_model.policy本身就是一个PyTorch的nn.Module,你完全可以直接拿它来用,不需要额外导出或重建。比如要用于DeepSHAP,只需要把这个policy作为模型传入即可,示例代码:

import shap
import torch as th

# 假设A2C_model是你训练好的模型
policy = A2C_model.policy
policy.eval()  # 切换到评估模式

# 准备输入样本(比如环境观测数据)
background = th.randn(100, 49)  # 100个随机背景样本,维度和你的输入一致
explainer = shap.DeepExplainer(policy, background)

# 对单个样本计算SHAP值
sample = th.randn(1, 49)
shap_values = explainer.shap_values(sample)

方法2:正确导出并加载模型参数

如果你确实需要导出模型到独立的PyTorch文件再加载,步骤如下:

  1. 导出模型参数
# 保存policy的状态字典
th.save(A2C_model.policy.state_dict(), "a2c_policy.pt")
  1. 正确重建模型结构
    你之前的代码有几个错误:
  • nn.Sequentail拼写错误,应该是nn.Sequential
  • mlp_extractor不能用Sequential包裹policy_net和value_net,因为Stable-Baselines3的MlpExtractor是一个自定义模块,不是Sequential
  • forward方法没有实现前向逻辑

正确的模型重建代码:

import torch as th
import torch.nn as nn

class MlpExtractor(nn.Module):
    def __init__(self, n_inputs):
        super().__init__()
        self.policy_net = nn.Sequential(
            nn.Linear(n_inputs, 64),
            nn.Tanh(),
            nn.Linear(64, 64),
            nn.Tanh()
        )
        self.value_net = nn.Sequential(
            nn.Linear(n_inputs, 64),
            nn.Tanh(),
            nn.Linear(64, 64),
            nn.Tanh()
        )
    
    def forward(self, features):
        return self.policy_net(features), self.value_net(features)

class ActorCriticPolicy(nn.Module):
    def __init__(self, n_inputs=49, n_actions=5):
        super().__init__()
        # 特征提取器(这里都是Flatten)
        self.features_extractor = nn.Flatten(start_dim=1, end_dim=-1)
        self.pi_features_extractor = nn.Flatten(start_dim=1, end_dim=-1)
        self.vf_features_extractor = nn.Flatten(start_dim=1, end_dim=-1)
        
        self.mlp_extractor = MlpExtractor(n_inputs)
        
        self.action_net = nn.Linear(64, n_actions)
        self.value_net = nn.Linear(64, 1)
    
    def forward(self, x):
        # 前向传播逻辑,和原policy一致
        features = self.features_extractor(x)
        pi_features = self.pi_features_extractor(x)
        vf_features = self.vf_features_extractor(x)
        
        policy_features, value_features = self.mlp_extractor(features)
        
        action_logits = self.action_net(policy_features)
        value = self.value_net(value_features)
        
        return action_logits, value

# 初始化模型并加载参数
reconstructed_policy = ActorCriticPolicy()
reconstructed_policy.load_state_dict(th.load("a2c_policy.pt"))
reconstructed_policy.eval()

这样就能正确加载导出的模型参数,得到和原模型一致的PyTorch模型。


内容的提问来源于stack exchange,提问作者kaiwy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 14:07:03