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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 18:15:38