如何使用Stable-Baselines3通过模仿学习预训练模型
在Stable-Baselines3中实现预训练(替代旧版SB的pretrain功能)
核心思路
旧版Stable-Baselines的pretrain本质是行为克隆(Behavior Cloning, BC),用专家数据监督训练策略网络。SB3中没有直接的ExpertDataset和pretrain方法,但可以通过以下两种方式实现相同功能:
方法一:使用SB3-Contrib中的BC算法(推荐)
SB3的扩展库sb3-contrib提供了现成的BC实现,能直接加载专家数据进行预训练。
步骤1:安装sb3-contrib
pip install sb3-contrib
步骤2:准备专家数据集
确保你的.npz文件包含以下键(和旧版格式匹配):
observations: 专家观测数据,形状为(n_samples, obs_dim)actions: 专家动作数据,形状为(n_samples, act_dim)
步骤3:加载数据并训练
import numpy as np from sb3_contrib import BC from stable_baselines3 import PPO from stable_baselines3.common.env_util import make_vec_env # 加载专家数据 expert_data = np.load('expert_cartpole.npz') observations = expert_data['observations'] actions = expert_data['actions'] # 初始化环境和PPO模型(和旧版一致) env = make_vec_env('CartPole-v1', n_envs=1) model = PPO('MlpPolicy', env, verbose=1) # 用BC预训练模型策略网络 bc = BC( policy=model.policy, observations=observations, actions=actions, verbose=1 ) bc.learn(total_timesteps=1000 * len(observations)) # 对应旧版的n_epochs=1000 # 预训练完成后,继续用PPO训练或直接使用模型 model.set_policy(bc.policy) # 后续可以继续PPO训练 # model.learn(total_timesteps=10000)
方法二:手动实现预训练逻辑(无需额外库)
如果不想依赖扩展库,可以手动写行为克隆的训练循环,直接优化模型的策略网络。
import numpy as np import torch from stable_baselines3 import PPO from stable_baselines3.common.env_util import make_vec_env # 加载专家数据 expert_data = np.load('expert_cartpole.npz') observations = torch.tensor(expert_data['observations'], dtype=torch.float32) actions = torch.tensor(expert_data['actions'], dtype=torch.float32) # 初始化PPO模型 env = make_vec_env('CartPole-v1', n_envs=1) model = PPO('MlpPolicy', env, verbose=1) # 设置优化器 optimizer = torch.optim.Adam(model.policy.parameters(), lr=3e-4) n_epochs = 1000 batch_size = 128 # 手动预训练循环 for epoch in range(n_epochs): # 随机打乱数据 permutation = torch.randperm(len(observations)) obs_batch = observations[permutation] acts_batch = actions[permutation] for i in range(0, len(observations), batch_size): batch_obs = obs_batch[i:i+batch_size] batch_acts = acts_batch[i:i+batch_size] # 计算策略网络的动作概率损失 dist = model.policy.get_distribution(batch_obs) loss = -dist.log_prob(batch_acts).mean() # 反向传播优化 optimizer.zero_grad() loss.backward() optimizer.step() if (epoch + 1) % 100 == 0: print(f"Epoch {epoch+1}/{n_epochs}, Loss: {loss.item():.4f}") # 预训练完成后可直接使用模型
注意事项
- 确保专家数据的格式和模型的观测/动作空间匹配,比如CartPole的动作是离散的,数据要对应整数类型(如果用离散策略)。
- 旧版的
traj_limitation参数可以通过截取数据的前N条轨迹实现,比如从expert_data中提取前N个轨迹的观测和动作。
内容的提问来源于stack exchange,提问作者Lord-Goku
相关产品推荐
相关产品推荐

