如何在Stable Baselines3中基于预训练权重训练自定义特征提取器?
问题
我为StableBaselines3模型使用了以下自定义特征提取器:
import torch.nn as nn from stable_baselines3 import PPO class Encoder(nn.Module): def __init__(self, input_dim, embedding_dim, hidden_dim, output_dim=2): super(Encoder, self).__init__() self.encoder = nn.Sequential( nn.Linear(input_dim, embedding_dim), nn.ReLU() ) self.regressor = nn.Sequential( nn.Linear(embedding_dim, hidden_dim), nn.ReLU(), ) def forward(self, x): x = self.encoder(x) x = self.regressor(x) return x model = Encoder(input_dim, embedding_dim, hidden_dim) model.load_state_dict(torch.load('trained_model.pth')) # Freeze all layers for param in model.parameters(): param.requires_grad = False class CustomFeatureExtractor(BaseFeaturesExtractor): def __init__(self, observation_space, features_dim): super(CustomFeatureExtractor, self).__init__(observation_space, features_dim) self.model = model # Use the pre-trained model as the feature extractor self._features_dim = features_dim def forward(self, observations): features = self.model(observations) return features policy_kwargs = { "features_extractor_class": CustomFeatureExtractor, "features_extractor_kwargs": {"features_dim": 64} } model = PPO("MlpPolicy", env=envs, policy_kwargs=policy_kwargs)
目前模型训练无问题且效果良好,现在我希望不再冻结权重,尝试从初始预训练权重开始训练该特征提取器。对于这种以嵌套类形式定义的自定义特征提取器,我该如何实现?由于我的特征提取器与官方文档中的定义不同,我不确定它是否会被训练,或者解冻层后是否会自动开始训练?
解决方案
要解冻预训练的Encoder模型并让它随PPO一起训练,只需调整两个关键环节:
1. 解除权重冻结
直接移除原有的冻结代码,或者将参数的requires_grad设为True——PyTorch中模型参数默认就是可训练的,只要不手动冻结,就会参与训练:
# 解冻所有层 for param in model.parameters(): param.requires_grad = True
2. 确认特征提取器被纳入优化流程
你的CustomFeatureExtractor将预训练的Encoder作为子模块,StableBaselines3初始化PPO时会递归遍历整个policy的所有子模块,自动把所有requires_grad=True的参数加入优化器。所以只要Encoder的参数处于可训练状态,就会被正常更新,和官方定义的特征提取器逻辑完全一致,不用纠结嵌套结构的问题。
完整修改后的代码示例
import torch.nn as nn from stable_baselines3 import PPO from stable_baselines3.common.torch_layers import BaseFeaturesExtractor # 补充缺失的导入 class Encoder(nn.Module): def __init__(self, input_dim, embedding_dim, hidden_dim, output_dim=2): super(Encoder, self).__init__() self.encoder = nn.Sequential( nn.Linear(input_dim, embedding_dim), nn.ReLU() ) self.regressor = nn.Sequential( nn.Linear(embedding_dim, hidden_dim), nn.ReLU(), ) def forward(self, x): x = self.encoder(x) x = self.regressor(x) return x # 加载预训练模型 encoder_model = Encoder(input_dim, embedding_dim, hidden_dim) encoder_model.load_state_dict(torch.load('trained_model.pth')) # 解冻所有层(或直接省略此步骤,默认参数可训练) for param in encoder_model.parameters(): param.requires_grad = True class CustomFeatureExtractor(BaseFeaturesExtractor): def __init__(self, observation_space, features_dim): super(CustomFeatureExtractor, self).__init__(observation_space, features_dim) self.model = encoder_model # 传入解冻后的预训练模型 self._features_dim = features_dim def forward(self, observations): features = self.model(observations) return features policy_kwargs = { "features_extractor_class": CustomFeatureExtractor, "features_extractor_kwargs": {"features_dim": 64} } # 初始化PPO模型,特征提取器参数自动纳入优化 model = PPO("MlpPolicy", env=envs, policy_kwargs=policy_kwargs) model.learn(total_timesteps=100000)
验证训练状态的方法
训练前可以打印特征提取器的参数状态,确认是否可训练:
# 检查参数是否开启梯度更新 for name, param in model.policy.features_extractor.named_parameters(): print(f"参数 {name}: requires_grad={param.requires_grad}")
如果输出均为requires_grad=True,说明这些参数会在训练中被更新。
内容的提问来源于stack exchange,提问作者Sayyor Y
相关产品推荐
相关产品推荐

