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

为SB3-contrib的RecurrentPPO实现行为克隆的技术求助

解决方案:适配RecurrentPPO的行为克隆预训练代码修改

核心修改点

RecurrentPPO的LSTM策略依赖**episode起始标记(episode_starts)和LSTM隐藏状态(lstm_states)**处理序列数据,原代码未覆盖这两个关键逻辑,需从数据加载、状态初始化、模型前向传播三个层面调整:

1. 确保专家数据集包含episode_starts字段

训练/测试用的专家数据集需返回三元组:(观测数据, 专家动作, episode_starts),其中episode_starts是布尔数组,标记每个样本是否为对应episode的第一步。

2. 调整DataLoader迭代输出结构

无需修改DataLoader初始化代码,但要确保数据集的__getitem__方法返回上述三元组,迭代时能同时获取观测、动作和episode起始标记。

3. 重写train和test函数(适配LSTM状态)

以下是修改后的核心函数,添加LSTM状态初始化与重置逻辑,适配RecurrentPPO的策略调用:

def train(model, device, train_loader, optimizer):
    model.train()
    for batch_idx, (data, target, episode_starts) in enumerate(train_loader):
        data, target, episode_starts = data.to(device), target.to(device), episode_starts.to(device)
        optimizer.zero_grad()

        # 初始化LSTM隐藏状态:(细胞状态, 隐藏状态),形状为(n_layers, batch_size, hidden_size)
        lstm_states = model._get_lstm_states(data.shape[0])
        lstm_states = (lstm_states[0].to(device), lstm_states[1].to(device))

        # 根据episode_starts重置对应样本的LSTM状态,避免跨episode干扰
        for i, start in enumerate(episode_starts):
            if start:
                lstm_states[0][:, i, :] = 0.0
                lstm_states[1][:, i, :] = 0.0

        # RecurrentPPO策略前向传播:必须传入lstm_states和episode_starts
        if isinstance(env.action_space, gym.spaces.Box):
            action, _, _, lstm_states = model(data, lstm_states=lstm_states, episode_starts=episode_starts)
            action_prediction = action.double()
        else:
            dist, _, lstm_states = model(data, lstm_states=lstm_states, episode_starts=episode_starts)
            action_prediction = dist.distribution.logits
            target = target.long()

        loss = criterion(action_prediction, target)
        loss.backward()
        optimizer.step()

        if batch_idx % log_interval == 0:
            print(
                "Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}".format(
                    epoch,
                    batch_idx * len(data),
                    len(train_loader.dataset),
                    100.0 * batch_idx / len(train_loader),
                    loss.item(),
                )
            )

def test(model, device, test_loader):
    model.eval()
    test_loss = 0.0
    with th.no_grad():
        for data, target, episode_starts in test_loader:
            data, target, episode_starts = data.to(device), target.to(device), episode_starts.to(device)

            # 初始化LSTM隐藏状态
            lstm_states = model._get_lstm_states(data.shape[0])
            lstm_states = (lstm_states[0].to(device), lstm_states[1].to(device))

            # 根据episode_starts重置状态
            for i, start in enumerate(episode_starts):
                if start:
                    lstm_states[0][:, i, :] = 0.0
                    lstm_states[1][:, i, :] = 0.0

            # RecurrentPPO策略前向传播
            if isinstance(env.action_space, gym.spaces.Box):
                action, _, _, lstm_states = model(data, lstm_states=lstm_states, episode_starts=episode_starts)
                action_prediction = action.double()
            else:
                dist, _, lstm_states = model(data, lstm_states=lstm_states, episode_starts=episode_starts)
                action_prediction = dist.distribution.logits
                target = target.long()

            test_loss += criterion(action_prediction, target).item() * len(data)

    test_loss /= len(test_loader.dataset)
    print(f"Test set: Average loss: {test_loss:.4f}")

4. 修正模型植入代码

原代码末尾的a2c_student为笔误,替换为传入的student:

# 将训练后的策略植入RecurrentPPO智能体
student.policy = model

关键说明

  • model._get_lstm_states(batch_size)是RecurrentPPO策略内置方法,用于生成初始的LSTM细胞状态和隐藏状态。
  • 必须根据episode_starts重置对应样本的LSTM状态,否则不同episode的序列信息会互相干扰,导致训练失效。
  • 若对数据集做shuffle,需保证每个episode内的样本连续性,否则LSTM无法学习到有效的序列依赖。

内容的提问来源于stack exchange,提问作者GerardL

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 00:48:07