ARTICLE DETAIL

资讯详情

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

LSTM加持策略梯度:让强化学习智能体在部分可观测环境中学会记忆

LSTM加持策略梯度:让强化学习智能体在部分可观测环境中学会记忆 到了策略梯度系列第九篇总算是时候让智能体带着记忆去决策了。前八篇我们写过的PolicyGradient、REINFORCE、带baseline的各种变体核心都在优化 (\nabla_\theta J(\theta)\mathbb{E}[\nabla_\theta \log \pi_\theta(a|s)R]) 这个式子公式本身很干净但它默认了一个前提每一步拿到的状态 (s) 已经包含了做决策所需的全部信息。现实里这个前提经常不成立比如智能体只能看到部分环境状态或者单帧观测根本推断不出速度、方向、历史上下文这些关键信息。这时候就需要把LSTM塞进策略网络让策略从 (\pi_\theta(a|s)) 变成 (\pi_\theta(a|h_t))其中 (h_t) 是循环网络对历史信息的压缩摘要。这篇笔记不搞复杂的推导直接从代码落地的角度讲清楚LSTM和PolicyGradient怎么结合适合已经跑通基础REINFORCE、又正在被“部分可观测环境”折磨的读者参考。1. 为什么策略梯度里要加LSTM从马尔可夫假设到记忆1.1 标准PolicyGradient的隐含前提几乎所有基础强化学习教程开篇都会强调MDP也就是马尔可夫决策过程。MDP的核心要求是当前状态 (s_t) 必须包含足够信息让智能体只根据 (s_t) 就能做出最优决策不需要依赖更早的历史。这个假设在数学上非常漂亮因为它保证了策略 (\pi(a|s)) 的形式是完备的也正是REINFORCE这类策略梯度算法的理论根基。但到了真实工程里完全满足马尔可夫性质的环境并不多。举几个常见的例子用单帧RGB图像玩Atari游戏时只看一帧画面根本判断不了物体的运动方向机械臂抓取时如果传感器的观测有遮挡或者只能返回部分关节角度单步观测同样不足以确定当前系统状态自动驾驶里摄像头画面的单帧更是无法体现前车的加速度。这些场景统称为部分可观测马尔可夫决策过程也就是POMDP。在POMDP环境下继续用 (\pi(a|s)) 这种无记忆的策略本质上是在拿一个信息不完整的“快照”做决策。从数据角度看一个确定性环境里同一个观测可能对应多个完全不同的真实状态策略网络会尝试拟合一个多对一的映射结果通常是学出一团浆糊或者只能靠随机猜测的运气维持表现。1.2 LSTM在策略网络里到底扮演什么角色LSTM全称长短期记忆网络是RNN家族里最常用的变体。它的引入解决了一个工程师很头疼的问题如何让策略网络自己决定“哪些历史信息需要记住记多久”。LSTM内部有输入门、遗忘门、输出门和细胞状态 (c_t)这些机制可以看作一个小型的读写控制器(c_t) 就是随身携带的笔记门控决定每步写入多少新信息、丢弃多少旧信息、输出多少给决策层。在强化学习里使用LSTM最核心的转变是这个 [ \pi_\theta(a|h_t), \quad h_t \text{LSTM}(h_{t-1}, o_t) ] 这里的 (o_t) 是环境返回的观测(h_t) 是LSTM处理完当前观测后输出的隐藏状态。策略网络不再直接吃原始观测而是吃LSTM归纳出的历史当前信息摘要。可以理解成智能体带了一个小笔记本每走一步就把关键信息写进去再翻开笔记本做决策。有一个地方需要说清楚虽然标题写的是LSTM加持PolicyGradient但实际上任何基于Actor的算法比如Actor-Critic、PPO、TRPO都可以用同样的方式接入LSTM。只需要把价值网络和策略网络改成接收 (h_t) 而不是 (s_t) 即可。这篇笔记聚焦PolicyGradient但思想是通用的。1.3 LSTM引入之后训练目标有什么变化从公式上看REINFORCE的更新式几乎不变仍然是“增大好轨迹里动作的log概率减小差轨迹里动作的log概率”只是现在log概率的形式变成了 (\log \pi_\theta(a_t|h_t))。损失函数真正影响的路径变了除了常规的策略层参数梯度还要通过LSTM的时序依赖反向传播回去这就是BPTT沿时间反向传播。这个变化对训练提出了几个新要求首先采样时不能打乱样本一个episode内部的时序关系必须保留其次每次梯度更新需要考虑多步历史计算图会比普通全连接网络长得多最后LSTM内部的梯度传播容易消失或爆炸需要额外的技巧对冲比如梯度裁剪、正交初始化、合理的学习率调节。这些细节后面会展开。2. LSTM接入策略网络的几种姿势与选型2.1 网络结构设计从观测到动作的完整通路LSTM接入策略网络最直接的就是拿LSTM替换掉原来策略网络里的第一层全连接层。典型的离散动作空间结构如下观测 o_t - LSTM - h_t - Linear - logits - softmax/sample - Linear - V(s) (如果带baseline)这里有个决策要点LSTM输入是原始观测 (o_t)还是先经过一个特征提取网络的特征向量 (f(o_t))如果任务本身就是低维向量输入比如CartPole的四维状态通常直接进LSTM即可。如果是图像等高维输入一般先用CNN卷积层把单帧观测压缩成一个低维特征向量再送入LSTM。这个小细节决定了网络训练的速度和最终性能核心原则是不要让LSTM去承担过于底层的特征提取工作。另一个常被问到的点是要不要把上一时刻的动作 (a_{t-1}) 也拼进LSTM输入这个做法在机器人控制任务里很常见因为策略需要知道自己上一步做了什么才能平滑地调整当前动作。对于本身状态比较完善的任务拼不拼影响不大对于部分可观测性较强的连续控制任务把 (a_{t-1}) 拼进去通常是有益的。我一般建议默认不拼遇到收敛不良可以先试这个改动。2.2 双向LSTM能用吗不行在监督学习的时序任务里双向LSTM效果很好因为它能同时看到过去和未来的上下文。但在强化学习在线决策场景智能体在时刻 (t) 做动作时根本不可能拿到未来的观测所以策略网络里的LSTM必须是单向的。这一条看起来简单但我在不少入门者的代码里见过把双向LSTM直接迁移过来的情况训练时用完整轨迹倒是能跑一旦部署到在线交互就会彻底崩掉因为未来的信息不存在了。如果你是在做离线强化学习手头有完整轨迹数据理论上训练时用双向编码也不是完全不行但推理阶段依然要换成单向网络这会导致训练和部署不一致所以不值得为了几个点的收益去冒这个险。2.3 一条episode就是一个序列别打乱监督学习里我习惯把样本shuffle成mini-batch防止模型记住数据顺序。但强化学习场景完全不同一个episode天然就是一条时间序列LSTM依赖这条序列来传递隐藏状态样本一旦被打乱(h_{t-1} \to h_t) 的递推关系就断了模型根本学不到任何时序依赖。训练LSTM策略时最常用的组织方式是一次rollout收集一整条episode训练时把这条episode的所有观测按顺序整理成形状为 (1, T, input_dim) 的张量一次forward得到所有时间步的logits和value再统一算loss、backward。如果环境步数太长也可以把一条episode拆成多个片段做截断BPTT但片段之间的顺序依然不能乱。3. 核心实现PyTorch搭建带LSTM的策略网络与REINFORCE训练3.1 网络定义与初始化细节先说网络结构。我实现一个同时输出动作概率和价值估计的网络这样训练REINFORCE时可以顺手加baseline方差能小很多。import torch import torch.nn as nn import torch.nn.functional as F from torch.distributions import Categorical class LSTMPolicyNetwork(nn.Module): def __init__(self, input_dim, hidden_dim, num_actions, num_layers1, drop_prob0.0): super().__init__() self.hidden_dim hidden_dim self.num_layers num_layers self.lstm nn.LSTM( input_sizeinput_dim, hidden_sizehidden_dim, num_layersnum_layers, batch_firstTrue, dropoutdrop_prob ) self.actor_head nn.Linear(hidden_dim, num_actions) self.critic_head nn.Linear(hidden_dim, 1) # 关键初始化正交初始化对RNN训练帮助很大 for name, param in self.lstm.named_parameters(): if weight_ih in name: nn.init.xavier_uniform_(param) elif weight_hh in name: nn.init.orthogonal_(param) elif bias in name: nn.init.zeros_(param) def forward(self, obs_seq, hidden_stateNone): lstm_out, hidden_state self.lstm(obs_seq, hidden_state) logits self.actor_head(lstm_out) value self.critic_head(lstm_out).squeeze(-1) return logits, value, hidden_state有几个值得细说的点。第一LSTM的初始隐藏状态可以直接给NonePyTorch默认会初始化成全零但要记住它需要在每个episode开始时重新置空。第二weight_hh用正交初始化这是因为RNN在时间维度的连乘结构对初始权重非常敏感正交矩阵能尽量延缓梯度消失或爆炸。第三dropout在LSTM中要慎重PyTorch的LSTM只在层间加dropout不会作用在时间步上层数少于2时这个参数没有效果不用为了凑超参数硬加。3.2 采样收集轨迹隐藏状态怎么传训练REINFORCE的第一步是采样一条episode。和普通策略网络不同这里每一步都要把上一步的隐藏状态传回去让LSTM维持记忆。def collect_episode(env, policy, max_steps300): obs_list, action_list, reward_list [], [], [] log_prob_list [] obs, _ env.reset() hidden_state None done False step 0 while not done and step max_steps: obs_tensor ( torch.FloatTensor(obs) .unsqueeze(0) .unsqueeze(0) ) # 注意这里要保留梯度因为后面要直接 backward logits, value, hidden_state policy(obs_tensor, hidden_state) dist Categorical(logitslogits.squeeze(1)) action dist.sample() log_prob dist.log_prob(action) obs_list.append(obs) action_list.append(action.item()) log_prob_list.append(log_prob) obs, reward, terminated, truncated, _ env.step(action.item()) reward_list.append(reward) done terminated or truncated step 1 obs_seq torch.FloatTensor(obs_list).unsqueeze(0) # (1, T, input_dim) rewards torch.FloatTensor(reward_list) log_probs torch.stack(log_prob_list) return obs_seq, action_list, rewards, log_probs, hidden_state这里有一个很重要的工程取舍。我在代码里让log_prob带着整个计算图好处是训练时可以直接对这个episode的loss做backward逻辑很顺。缺点是如果episode特别长计算图会非常庞大显存和内存都会吃紧。后面避坑部分我会给截断BPTT的替代方案。如果采样时用的是torch.no_grad()那log_prob不会保留梯度训练时就需要用保存的obs重新forward一遍来计算log_prob。这种做法内存省但会引入“rollout策略和当前策略不一致”的隐患。REINFORCE本来就是每条episode只做一次更新这个隐患实际影响不大但理解清楚更好。3.3 折扣回报计算与baseline设计REINFORCE下一步是把稀疏的奖励折算成每个时间步的回报。折扣因子 (\gamma) 的意义是越早的决策影响越长远所以早期时间步的回报应该包含未来所有折扣后的奖励。def compute_discounted_returns(rewards, gamma0.99): returns [] G 0.0 for r in reversed(rewards): G r gamma * G returns.insert(0, G) return torch.FloatTensor(returns)直接使用原始回报做损失函数梯度方差非常大因为不同episode之间的总回报可能相差悬殊。两个常用手段一是引入价值函数作为baseline用优势 (A_t G_t - V(s_t)) 替换 (G_t)二是对回报做标准化让它变成零均值单位方差。两者可以同时使用实战效果通常也不错。我的做法是让价值网络直接吃LSTM的隐藏状态 (h_t)这比让它单独吃一个全连接层更好因为 (h_t) 本身已经包含了历史信息对当前状态的价值估计会更准确。训练损失由两部分组成def compute_loss(log_probs, values, rewards, gamma0.99): returns compute_discounted_returns(rewards, gamma) returns (returns - returns.mean()) / (returns.std() 1e-9) advantages (returns - values.detach()).detach() policy_loss -(log_probs * advantages).mean() value_loss F.mse_loss(values, returns) return policy_loss 0.5 * value_lossvalues.detach()是为了让价值函数的梯度不会干扰策略梯度。这里我特意把advantage也detach掉避免出现二阶梯度的问题。初学的时候容易把returns - values直接拿去做乘法这会使得策略部分的梯度穿过价值网络导致两个head互相干扰训练很不稳定。3.4 完整训练循环与梯度裁剪训练循环本身不复杂关键是要在正确的地方重置LSTM隐藏状态并在backward之前做梯度裁剪。def train_lstm_reinforce(env, policy, optimizer, episodes1000, gamma0.99, max_steps300, clip_grad0.5): for episode_idx in range(episodes): obs_seq, actions, rewards, log_probs, _ collect_episode( env, policy, max_steps ) # 重新forward一次得到每个时间步的value with torch.no_grad(): logits, values, _ policy(obs_seq, None) loss compute_loss(log_probs, values, rewards, gamma) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(policy.parameters(), clip_grad) optimizer.step()这里我用了with torch.no_grad()重新获得values因为采样阶段得到的value虽然也有梯度但它只是用于估计baseline完全不需要回传到LSTM。梯度裁剪的max_norm0.5是我在这个任务里常用的值如果发现训练不稳定可以先尝试把它调到0.25或1.0。LSTM的时间维度连乘很容易让梯度的范数暴涨裁剪几乎是必需品而不是可选项。4. 实验对比从CartPole到“记忆型”环境4.1 在CartPole上验证LSTM策略的稳定性第一个实验用CartPole-v1这个环境状态维数低、完全可观测按道理不需要LSTM。我拿它先验证代码有没有写错顺便观察LSTM引入后对标准任务的负面影响有多大。超参数hidden_dim128lr1e-3gamma0.992000个episode。训练结果符合预期带LSTM的网络最终也能收敛到500分满分的水平但速度比两层全连接网络要慢一些大约多花20%到30%的episode。原因也很直白LSTM引入了额外的时序参数和更长的反向传播路径在本身不需要记忆的任务上属于“过度设计”。不过这个实验的真正价值在于排除bug如果LSTM策略在CartPole这种简单任务上都学不好那基本可以判断是代码实现问题而不是任务太难。我建议所有想在自己的任务里引入LSTM的人都先拿一个标准MDP环境跑通验证再切换到部分可观测任务。4.2 构造一个必须“记住第一步”的小环境为了直观展示LSTM的价值我构造了一个尽可能简单的记忆任务环境。核心逻辑第一步给智能体一个one-hot向量表示一个从0到N-1的随机数字之后的每一步都返回全零向量整个episode长度为T最后一步智能体必须输出一个动作如果动作等于第一步的那个数字奖励为1否则为0。class MemEnv: def __init__(self, n_symbols5, seq_len10): self.n n_symbols self.T seq_len self.action_space_n n_symbols def reset(self): self.target int(torch.randint(self.n, (1,)).item()) self.t 0 return self._obs() def _obs(self): if self.t 0: return torch.eye(self.n)[self.target].numpy() return torch.zeros(self.n).numpy() def step(self, action): self.t 1 done self.t self.T reward 1.0 if done and int(action) self.target else 0.0 return self._obs(), reward, done, {}这个环境的含义非常清晰普通全连接策略在非第一步的所有时间步收到的观测都是全零它根本无从判断第一步看到的是哪个数字最优表现只能是1/5的随机猜测。而LSTM可以在第一步把one-hot信息写进细胞状态之后一路保持最后根据记忆输出动作。这样的对比能让读者一眼看出记忆模块的威力。4.3 实验结果全连接策略的挣扎与LSTM的稳定收敛我在n_symbols5、seq_len10的设置下跑了1000个episode两组实验分别使用两层全连接网络和一层LSTM网络其余超参数保持一致。结果非常典型。模型平均回报最后100个episode达到80%准确率的episode数两层全连接0.19左右接近随机1/5未达到LSTM策略0.85到0.92约400个episode全连接策略在这个任务上基本就是随机猜测因为观测里没有任何关于目标数字的线索。LSTM策略经过几百个episode后能稳定输出正确动作说明它确实学会了把第一步的信息长期保存在隐藏状态中。这个任务虽然简单但它把“记忆能力”这个抽象概念变成了一个可以被直接观测和统计的实验指标很适合用来验证带记忆策略网络的实现是否正确。在真实部分可观测任务里效果不如这个玩具环境这么泾渭分明但只要环境需要短期记忆LSTM带来的提升一般都在可感知的范围内。我还在一个只给部分关节角度的机械臂仿真任务里测过隐藏状态维数128的LSTM策略相比全连接策略成功率提升了大概15个百分点代价是训练时间翻倍这个trade-off需要根据自己的任务权衡。5. 避坑指南训练LSTM策略时我踩过的几个坑5.1 隐藏状态重置时机一个让人抓狂的bug很多人第一次写LSTM策略最容易犯的错误是忘记在episode结束时重置hidden state。如果连续多个episode共用同一个隐藏状态序列LSTM会“记忆串场”把上一个episode的信息带到当前episode里导致训练曲线忽高忽低甚至完全无法收敛。正确的做法是每个episode开始都把hidden state设为None也就是重置为全零。需要注意的一点在gymnasium环境里terminated和truncated都应该视为一个episode的结束。terminated代表到达终止状态truncated代表超时或环境强制截断这两种情况下隐藏状态都必须清空否则下一个episode的开头会带着上一个episode的尾巴。5.2 计算图太长长episode下的显存与内存爆炸REINFORCE里如果直接把整条episode的log_prob求和再backward计算图长度等于整个episode的步数。CartPole这种几百步的还好如果任务动辄上千步或者观测是图像内存和显存会迅速见底。解决方案是截断BPTT思路是隐藏状态继续向后传递但梯度只回传最近N步。最简单粗暴的实现是每过N步把hidden_state中需要梯度的那部分detach掉。这样当前片段只保留最近N步的计算图既保留了LSTM对长期信息的依赖又不会让内存无限制膨胀。N一般取64到128太小会失去长期记忆的效果太大又回到内存爆炸的老问题。# 示意在 collect 过程中每隔 trunc_len 步 # 将 hidden_state 两个分量都 detach 一次 if (step 1) % trunc_len 0: hidden_state ( hidden_state[0].detach(), hidden_state[1].detach(), )这样处理之后每个片段结束位置适合单独做一次loss回传逻辑会比单episode统一backward复杂一点但工程上是值得的。如果任务episode普遍不超过200步可以先不搞截断保持代码简单。5.3 梯度消失与梯度爆炸LSTM的老朋友LSTM的门控机制已经大幅缓解了梯度消失问题但在长序列任务里梯度爆炸依然时有发生。除了梯度裁剪之外权重初始化也很关键。nn.LSTM默认的初始化方式是均匀分布很多情况下效果一般我习惯手动把输入权重设为xavier把循环权重设为正交初始化实测对收敛速度有明显改善。学习率方面LSTM策略网络通常需要比全连接网络更低的学习率。全连接REINFORCE用1e-3很常见但到了LSTM这里1e-3可能训练几十个episode就出现loss爆炸。我一般先用5e-4起步如果训练曲线太保守再逐步上调。记住LSTM的时间反向传播路径长高学习率的风险是成倍放大的。5.4 观测标准化与奖励缩放容易被忽视的前置操作LSTM对输入尺度比全连接网络更敏感。如果一个观测维度的数值范围是[0, 10000]另一个是[-1, 1]LSTM的门控计算会把数值范围大的维度当成主导信号直接影响遗忘和写入行为。所以原始观测一定要先标准化到[-1, 1]或者零均值单位方差再做归一化后送入网络。奖励的尺度同样值得警惕。REINFORCE的梯度大小和回报的大小线性相关如果奖励非常大比如1000那么即使梯度裁剪策略更新的步长也会极不稳定。我习惯在计算loss之前对returns做标准化这是成本最低的降方差手段。如果任务允许还可以对每一步奖励做一个缩放比如统一乘0.01这比改学习率更直观。5.5 用hidden state的范数作为诊断指标训练普通策略网络时我主要盯着loss曲线和平均回报但训练LSTM策略时我强烈建议额外记录每一层的hidden state的L2范数输出到TensorBoard或者本地日志里。hidden state范数的变化能反映很多问题如果范数急剧放大大概率是梯度爆炸的前兆如果范数普遍趋近于零说明记忆模块退化了模型没有在有效地保存历史信息如果范数波动异常则往往提示观测输入分布不稳定或者奖励信号有问题。这个指标是我在实际调试中用过的最有效的早期预警之一。比起事后分析loss曲线才意识到问题看hidden state范数可以提前好几个episode发现异常。做了多年实验之后我的习惯是每当引入新的RNN结构就把这个状态统计量当成标准配置而不是临时想起来才看一眼。写在实际调试之后我个人在实际操作中最大的体会是LSTM不是万能药它解决的是“状态信息不完整”这一类具体问题。如果环境本身满足马尔可夫性质强行加LSTM只会增加训练难度不会带来收益但如果任务确实需要记忆比如单帧画面测不到速度、传感器存在遮挡、策略需要根据几步前的信息做决策那LSTM几乎是性价比最高的方案之一。它会带来训练时间翻倍、调参难度上升、代码复杂度增加这些代价不过这些代价都是可以管理和接受的。最后再分享一个小技巧如果项目用到了带记忆的策略网络先从最简单的隐藏维度64起步用我前面提到的MemEnv这类玩具任务验证代码正确性确认LSTM确实学到了记忆能力再迁移到真实任务。这套流程帮我避免过很多次“在复杂环境里debug到怀疑人生”的糟糕体验。这个系列的下一篇我大概率会沿着这个方向写一写带记忆的PPO实现里面涉及的技巧会更复杂一些但核心思想是一致的。
返回列表