ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

LSTM改造PolicyGradient:用记忆通道解决部分可观测强化学习问题

LSTM改造PolicyGradient:用记忆通道解决部分可观测强化学习问题 我们做强化学习的平时用MLP当策略网络感觉一切正常。直到某一天你发现环境的观测其实藏了一部分不可见的信息或者你需要根据一连串的历史动作和反馈才能判断当前状态时问题就来了。这也是我写这份笔记的初衷当经典的PolicyGradient算法遇到部分可观测环境或者需要长程依赖记忆时该怎么给它装上LSTM这条记忆线。这篇笔记主要解决一个核心场景用LSTM改造PolicyGradientREINFORCE策略网络让它具备时序建模能力。内容包含LSTM在强化学习里的角色定位、网络结构设计、训练时序列批处理的手法以及我调试时踩过的若干坑。适合已经写过基础REINFORCE、但想在序列决策任务上进一步深入的同学参考。1. 为什么给PolicyGradient加LSTM1.1 经典PG算法的盲区从马尔可夫假设说起先回顾一下传统PolicyGradient的策略模型。我们用策略网络 π_θ(a|s) 来参数化动作分布训练目标是最小化L(θ) -E[ R_t * log π_θ(a_t|s_t) ]其中 R_t 可以是Monte Carlo回报也可以是带baseline的advantage估计。这套框架的前提是状态 s_t 具备了做决策需要的全部信息这就是马尔可夫性质。可现实任务里状态往往不是完全可观测的。机械臂抓取时视觉传感器有遮挡自动驾驶时雷达数据有延迟和噪声游戏里地图有迷雾这些都是典型的POMDP部分可观测马尔可夫决策过程环境。在POMDP条件下单帧观测 o_t 并不能唯一确定真实状态 s_t。比如一个极端例子环境里有两个长得一模一样但行为完全相反的怪物你只看当前画面无法区分它们必须记住上一帧甚至上几帧的移动轨迹。这时MLP策略网络等于拿着残缺信息做决策表现自然会崩。解决POMDP的经典思路是维护一个信念状态belief state也就是对真实状态的后验分布。但在工程实现上我们很少直接建模这个分布更常用的做法是把历史观测序列喂给一个RNN让RNN的隐状态当作“学习出来的信念状态”。LSTM作为RNN家族里最稳定的变体自然就是首选。1.2 LSTM在强化学习里到底充当什么角色LSTM插进PG算法后策略网络的输入从“当前帧观测”变成了“当前帧观测 隐状态”输出动作分布的同时还输出新的隐状态。LSTM在这里干的活简单说就是两件事记忆编码器把一段历史的观测、动作、奖励压缩成一个固定维度的向量隐状态 h_t。策略决策器基座解码出当下最优动作分布时参考隐状态里携带的时序上下文。这样改造后策略网络就不再是“看到什么就做什么”的即时映射而是“结合我记住的东西判断现在是什么局面”。这在有些资料里也被称为 RNN-based policy 或者 memory-based policy。还有一点值得注意加入LSTM后策略的梯度反向传播路径会穿越时间步。也就是说损失函数不仅回传“当前动作有多好”的信号还会通过BPTTBackpropagation Through Time把信号传到过去的隐状态上。这意味着梯度更新能同时优化“当前决策”和“记忆方式”——网络自己会学怎么记、记什么、什么时候忘。这是LSTMPG相比普通PG最质变的地方。1.3 什么样的任务适合上LSTMPG不是所有强化学习任务都需要LSTM。加LSTM会带来训练开销增大、不稳定因素增加等代价。根据我自己的实测这几类任务比较值得用观测存在遮挡或噪声单帧信息不完整动作效果有延迟当前动作的影响要好几步之后才显现同一观测在不同上下文下应该做出不同反应环境中存在周期性规律或明显的时序模式比如红绿灯切换、博弈对手的周期性策略。反之如果任务本身就是完全可观测的比如经典CartPole、Pendulum这类benchmark用MLP就够了。硬加LSTM反而容易过拟合历史信息或者训练变慢、梯度更不稳定。2. 核心细节拆解LSTMPG的关键设计点2.1 Episode序列的组织与截断策略LSTM天然按时间步消费数据所以训练样本的组织方式跟普通PG完全不同。普通PG把每条transition当成独立样本收集而LSTMPG必须按episode为单位组织样本。一个自然的做法是完整跑完一个episode拿到整条轨迹 (o_1, a_1, r_1, ..., o_T, a_T, r_T)然后把这整条轨迹当成一个序列喂给LSTM做前向计算。这样做的好处是隐状态从头到尾是连贯的BPTT也能利用完整的历史依赖。但实践中有一个现实问题episode可能很长。比如处理机械臂连续控制任务一条轨迹可能几千步在游戏环境里甚至几万步。这种超长序列直接做完整BPTT显存和计算量都顶不住。我的做法是折中按固定窗口长度截断BPTT但保留跨窗口的隐状态传递也就是truncated BPTT。比如设定窗口长度 K200第1到200步正常前向后通过时间反向传播更新参数第201到400步继续沿用第200步的隐状态接着算只是反向传播时不再回传到第1步。这种截断方式保留了长期记忆同时把单次更新的计算量限制在可控范围内。2.2 hidden state的传递与清零时机LSTM的隐状态管理是一个极其容易被忽略、但对结果影响极大的环节。隐状态从哪里初始化Episode之间要不要清零Batch里不同序列的隐状态怎么隔离这些我全部踩过坑。先说Episode内部的传递同一个episode内的所有时间步隐状态必须连续传递这是LSTM记忆发挥作用的前提。一旦中间清零网络就等于失去上下文之前的记忆全部作废和MLP没有本质区别。再说Episode之间的处理不同episode之间隐状态必须清零。因为每个episode的策略采样是独立的前一个episode结束后的隐状态带着上一局的“残留记忆”直接传给新episode会导致策略被旧上下文污染。我在CartPole的POMDP变体上做过对比实验episode之间不reset隐状态训练曲线抖动幅度会明显加大最终收敛值也会低一截。还有一个细节容易被忽略清零时不仅要把 h 置为零向量连 cell state c 也要一起置零。只清 h 不清晰 cLSTM的记忆仍然残留在 c 里等于没清干净。PyTorch的lstm.detach()常用在截断BPTT时切断梯度但只在需要截断时用别在episode内部乱调。2.3 输入输出维度的设计LSTM改造策略网络时输入输出维度是很多人第一次写就卡壳的地方。直接说结论我常用如下配置策略网络输入维度obs_dim action_dim 1可选其中1是上一时刻的奖励用于感知奖励变化趋势LSTM隐层维度hidden_dim一般取值范围64~256根据任务复杂度定不需要盲目加大输出维度动作空间维度离散动作就是类别数连续动作就是动作均值维度。有一点需要注意LSTM的输入是三维张量[seq_len, batch_size, input_dim]。这是PyTorch的默认格式和很多人的直觉[batch, seq_len]不一样。写代码时如果维度不对会在forward里直接报错或算出错误结果。还有一点关于奖励的额外输入如果把r_{t-1}作为额外特征拼进去就要求在收集轨迹时保存每步的奖励值。这在REINFORCE里是很自然的反正你也要算回报所以顺手拼进去没有额外成本。3. 实操实现PyTorch写一个LSTM-PolicyGradient3.1 整体网络结构我直接给一份核心代码骨架用PyTorch实现。注意这不是完整可跑项目而是聚焦在策略网络和训练循环的核心部分。import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim class LSTMDiscretePolicy(nn.Module): def __init__(self, input_dim, hidden_dim, num_actions, num_layers1): super().__init__() self.input_dim input_dim self.hidden_dim hidden_dim self.num_actions num_actions self.num_layers num_layers self.lstm nn.LSTM( input_sizeinput_dim, hidden_sizehidden_dim, num_layersnum_layers, batch_firstFalse, ) self.fc nn.Linear(hidden_dim, num_actions) def forward(self, obs_seq, hiddenNone): # obs_seq: [seq_len, batch_size, input_dim] lstm_out, hidden self.lstm(obs_seq, hidden) # lstm_out: [seq_len, batch_size, hidden_dim] logits self.fc(lstm_out) # [seq_len, batch_size, num_actions] return logits, hidden def init_hidden(self, batch_size, device): h0 torch.zeros(self.num_layers, batch_size, self.hidden_dim).to(device) c0 torch.zeros(self.num_layers, batch_size, self.hidden_dim).to(device) return (h0, c0)这里有个关键设计batch_firstFalse保持PyTorch LSTM默认格式。如果谁习惯batch_firstTrue那就要把所有输入输出的维度顺序调过来千万别搞混。基本上我推荐直接用默认的batch_firstFalse跟LSTM内部数学定义对齐不容易出错。3.2 训练循环与Batch组织REINFORCE训练循环的主要工作是收集一条episode轨迹计算每个时间步的折扣回报然后构造序列损失更新网络。这里我给出一个简化版的单episode训练流程def train_one_episode(env, policy, optimizer, gamma0.99, hidden_dim64): obs env.reset() done False obs_list, action_list, reward_list [], [], [] hidden policy.init_hidden(batch_size1, devicedevice) state_list [] # 保存每个时间步的memory state用于构造序列 while not done: obs_tensor torch.FloatTensor(obs).unsqueeze(0).unsqueeze(0).to(device) # obs_tensor: [1, 1, input_dim] logits, hidden policy.forward(obs_tensor, hidden) dist torch.distributions.Categorical(logitslogits.squeeze()) action dist.sample() next_obs, reward, done, _ env.step(action.item()) obs_list.append(obs) action_list.append(action) reward_list.append(reward) obs next_obs if done: break # 计算折扣回报 T len(reward_list) returns torch.zeros(T, 1) running_ret 0 for t in reversed(range(T)): running_ret reward_list[t] gamma * running_ret returns[t] running_ret # 用整条episode序列重新前向计算损失 obs_seq torch.FloatTensor(obs_list).unsqueeze(1).to(device) # [T, 1, input_dim] logits_seq, _ policy.forward(obs_seq, policy.init_hidden(1, device)) probs torch.softmax(logits_seq, dim-1) # [T, 1, num_actions] dist torch.distributions.Categorical(logitslogits_seq) log_probs dist.log_prob(torch.tensor(action_list).view(-1, 1).to(device)) loss -(log_probs * returns.to(device)).mean() optimizer.zero_grad() loss.backward() optimizer.step() return loss.item(), T注意一个细节训练时我们重新用整条episode序列做了一次前向隐状态从零开始。这和采样时的隐状态传递是两条路径。也就是说采样时的隐状态只用于决策训练前向的隐状态只用于梯度计算。两者不共用避免梯度在采样路径上传播导致不必要的计算图堆积。3.3 关键参数说明参数选择上我踩了几次坑之后形成了这套相对稳的配置直接抄作业也行参数推荐值说明gamma0.99~0.995折扣因子环境步长很长就取大一点hidden_dim64~256LSTM隐层维度从128起步比较稳num_layers1~2超过2层训练不稳定收益也不明显学习率3e-4~1e-3比MLP策略网络调低一点更稳截断窗口100~300太长显存吃不消太短记忆经常被切断这里特别说一下学习率LSTM的反向传播路径比MLP长得多梯度数值天然偏大。用MLP时代比较顺手的学习率比如1e-2放到LSTM策略上训练几个episode就会爆掉。我自己的习惯是先跑一遍检查loss和熵如果熵突然掉到0或者loss变成NaN先把学习率降一个数量级再试。奖励归一化也是REINFORCE里很有用的trick算出returns_mean和returns_std把returns减去均值除以标准差。这不等同于真正减少方差但对数值稳定很有帮助。我在batch训练时几乎必做。4. 常见问题与排查技巧实录4.1 训练不收敛loss波动很大这是最常被问到的。先检查几个点按优先级排列reward归一化REINFORCE对reward scale极其敏感reward数值本身变大比如100到200梯度幅度会随之震荡。用advantage normalization减均值除标准差是经典解法。entropy正则LSTM策略因为参数多容易过早收敛到单一动作熵直接归零。这时候加entropy loss很有效损失项变成entropy_loss -entropy_coef * dist.entropy().mean()entropy_coef从0.01起步就行。这个系数别太大否则策略一直随机探索不学东西。调低学习率前面说过LSTM的反传路径长梯度范数大学习率需要缩。4.2 训练时隐状态污染导致效果突然崩坏这个坑非常隐蔽症状是训练曲线在某个阶段突然出现一段剧烈的性能下降然后后面爬回来。排查时我一度以为是学习率问题直到仔细检查代码才发现是hidden state没在episode之间清零。如果你用batch训练还会遇到另一个问题不同episode长度不同batch里短的episode结束之后它的隐状态应该清零但PyTorch的LSTM不能对batch里单个样本清零隐状态。我的处理办法是用变长序列的pack_padded_sequence机制或者简单粗暴地把同长度的episode分到同一个batch里训练。前一种方案兼容性好就是代码复杂一些后一种方案虽然有点浪费数据但实现简单适合小规模实验。4.3 训练显存持续增长或爆显存LSTMPG训练最容易犯的错误是在采样循环里不断把隐状态和观测追加进同一个计算图导致计算图越积越长显存持续暴涨。正确做法是采样阶段用torch.no_grad()包裹前向计算并且每一步只保留隐状态和动作/奖励不保留完整的中间激活值。到训练阶段再重新用整条序列前向一遍专门构造计算图。with torch.no_grad(): logits, hidden policy.forward(obs_tensor, hidden)这个改动看起来不起眼但能避免绝大多数的显存爆炸问题。4.4 序列长度不一致的batch训练当多个episode的轨迹长度不同把它们堆成一个batch时会出现维度不齐的问题。除了之前说的按长度分batch之外还有一种更通用的做法是padding mask# 将不同长度seq padding到同一长度max_len # mask: [max_len, batch_size, 1]表示该时间步是否有效 log_probs log_probs * mask loss -(log_probs * returns * mask).sum() / mask.sum()注意loss的denominator也要用mask.sum()而不是简单用mean()否则填充的0步会稀释真实样本的梯度。这个细节我在初版代码里没处理损失数值一直偏小策略网络学习效率低下。排查了好久才意识到是padding位置贡献了无效的log_prob。4.5 LSTM初始化方式LSTM权重初始化对训练的稳定性影响比很多人想象的大。PyTorch默认初始化其实比较稳但我遇到过几次长期不收敛的情况手动设置一下正交初始化之后明显改善def init_lstm_weights(module): if isinstance(module, nn.LSTM): for name, param in module.named_parameters(): if weight_ih in name: nn.init.xavier_uniform_(param.data) elif weight_hh in name: nn.init.orthogonal_(param.data) elif bias in name: param.data.fill_(0)这个初始化方式不是必须的但如果你的模型一直不收敛、loss纹丝不动值得试一下。它等价于给了LSTM一个更好的起点让它在训练初期更容易捕捉时序依赖。5. 踩坑后的几点个人体会LSTMPolicyGradient组合本质上是在策略网络里引入一条“记忆通道”让智能体能在POMDP环境下做推理。但它的代价也很明显训练复杂度高、调参空间大、对数据组织的规范性要求高。写代码的时候我最常提醒自己的是三件事隐状态管理、序列组织方式、梯度截断时机。这三个点做对了LSTM基本就成功了一半。从我实际测试的几个任务来看LSTMPG在部分可观测、状态信息不完整的环境下带来的提升是实打实的有时候能让原本根本学不出来的任务变得可收敛。但如果是完全可观测的简单任务直接上LSTM反而会带来不必要的麻烦。最后分享一个自己常用的调试技巧如果你怀疑LSTM部分出了问题可以先把LSTM改成恒等映射把输入直接接到输出对比一下MLP策略的表现。如果LSTM版本明显差于MLP版本那就是记忆模块本身在拖后腿如果差不多说明记忆模块至少没有引入太大偏差可以继续放心调其他超参数。这个思路在定位问题时能帮上不少忙。
返回列表