如何将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文件再加载,步骤如下:
- 导出模型参数
# 保存policy的状态字典 th.save(A2C_model.policy.state_dict(), "a2c_policy.pt")
- 正确重建模型结构
你之前的代码有几个错误:
nn.Sequentail拼写错误,应该是nn.Sequentialmlp_extractor不能用Sequential包裹policy_net和value_net,因为Stable-Baselines3的MlpExtractor是一个自定义模块,不是Sequentialforward方法没有实现前向逻辑
正确的模型重建代码:
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
相关产品推荐
相关产品推荐

