为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
相关产品推荐
相关产品推荐

