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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 11:40:07