ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

Transformer结合强化学习:策略网络实现长时序决策实战

Transformer结合强化学习:策略网络实现长时序决策实战 最近不少同学在选题时都在纠结同一个问题长时序决策到底该用什么模型过去大家普遍在 RNN、LSTM 与时间卷积网络之间做选择但近几年 Transformer 异军突起直接渗透到了强化学习的各个环节——从策略网络、价值函数到环境世界模型、离线强化学习到处都能看到它的身影。如果你正在为毕业设计或科研方向发愁Transformer 结合强化学习绝对是值得投入的方向之一。这篇文章不会只停留在概念层面我会从“为什么 Transformer 适合长时序决策”开始手把手带大家搭建一个可运行的 Transformer 策略网络并用 REINFORCE 算法训练它完成一个部分可观测的 CartPole 任务。代码会拆分成独立文件方便你直接修改和扩展也会重点讲清楚训练过程中的关键坑点与工程化建议。文章适合三类读者刚接触强化学习、想找论文方向的学生已经会用 PPO 等算法、希望把序列建模能力引入策略网络的研究者以及需要在真实项目中做长时序控制决策的工程师。读完后你应该能理解 Transformer 与强化学习结合的三种主流路线并且能动手实现一个最小可运行的 Transformer Policy。1. 为什么 Transformer 能和强化学习走到一起1.1 长时序决策到底难在哪里强化学习处理的是“智能体通过与环境交互来学习策略”的问题。传统强化学习算法通常基于马尔可夫假设也就是认为当前状态已经包含了做决策所需的全部信息过去的历史不再重要。但在很多真实任务中这个假设是不成立的。比如自动驾驶需要根据过去几秒的车辆轨迹判断意图机器人操作需要结合之前的传感器读数判断当前阶段金融交易需要参考历史行情走势。这些场景都具备明显的“部分可观测性”智能体看到的每一个时刻信息都不完整必须依赖一段历史窗口才能做出正确决策。这类问题在学术上被称为 POMDP部分可观测马尔可夫决策过程也是长时序决策的核心难点。除了观测不完整长时序决策还有另一个难点信用分配。当智能体做了一个错误决策但错误影响在几十步之后才体现出来时算法很难判断究竟是哪一步出了问题。RNN 和 LSTM 虽然天生适合处理序列但在长序列上存在梯度消失、训练并行度低等问题。这也给 Transformer 进入强化学习留下了空间。1.2 Transformer 的注意力机制带来了什么Transformer 最初是为机器翻译设计的核心是自注意力机制。自注意力允许序列中的每个位置直接与所有历史位置计算关联因此能够捕捉远距离依赖。相比于 RNN 逐时间步推进Transformer 在并行计算上也有天然优势。在强化学习中使用 Transformer最直观的方式是把一段连续观测当作序列输入利用自注意力建模观测之间的时间关联。比如当前第 10 步的观测可以直接与第 2 步的关键信息建立权重连接而不需要像 RNN 那样一步步传递隐状态。这种能力对部分可观测环境、长时记忆任务、多智能体协同决策都非常有价值。当然Transformer 也不是万能的。自注意力的计算复杂度是 O(T²)序列越长训练成本越高同时它缺少 RNN 那样的时序归纳偏置必须额外添加位置编码否则模型无法知道每个观测的先后顺序。这些细节在后续代码中都会体现。1.3 热门结合方向一览Transformer 与强化学习的结合并非只有一种模式目前主流路线大致分为三类。第一类是“策略网络替换”也就是把原本的 MLP 或 RNN 策略网络替换成 Transformer。智能体将最近 T 步观测组成序列输入 Transformer 编码器输出当前动作的概率分布。本文实战部分会实现这一种。第二类是“序列决策建模”代表工作是 Decision Transformer。它把强化学习问题转化为条件序列生成问题给定目标回报和过往轨迹让模型生成未来的动作序列。这类方法在离线强化学习、长时序决策中表现非常亮眼。第三类是“世界模型”也就是用 Transformer 充当环境模型预测未来状态和奖励。智能体先在“想象”中训练策略再迁移到真实环境能显著提升样本效率。像基于模型的强化学习、元强化学习等方向也经常使用 Transformer 作为骨干网络。这三条路线各有优劣但对于入门来说我建议先理解策略网络替换这条最简单也最通用的路线。2. 环境准备与版本说明2.1 依赖安装本文示例代码基于 Python 和 PyTorch环境版本建议如下Python 3.9 及以上PyTorch 2.0 及以上本文用到的torch.nn.TransformerEncoder在 1.8 之后都可用Gymnasium 0.29 及以上用于提供 CartPole 环境NumPy用于数据处理不建议继续使用旧版gym新项目推荐直接用gymnasium接口更清晰且维护更积极。先创建虚拟环境python -m venv venv source venv/bin/activate # Windows 下使用 venv\Scripts\activate然后安装依赖pip install torch gymnasium numpy如果你的机器有 NVIDIA 显卡建议安装对应 CUDA 版本的 PyTorch否则用 CPU 跑本文的小实验也足够。2.2 项目结构为了不让代码堆在一个文件里我们把项目拆成三个模块rl_transformer_demo/ ├── envs.py # 部分可观测 CartPole 环境封装 ├── policy.py # Transformer 策略网络 ├── train.py # REINFORCE 训练脚本 └── requirements.txt # 依赖清单这样写的好处是每一层职责清晰环境负责处理观测序列策略网络负责序列建模与动作输出训练脚本负责采样、计算回报和反向传播。后续你想把 REINFORCE 换成 PPO也只需要修改train.py不需要动模型结构。3. 核心原理拆解从 Transformer 到强化学习策略网络3.1 Transformer 基础回顾Transformer 的核心是自注意力机制它通过 Query、Key、Value 三个矩阵计算每个 token 与其他 token 的相关性。简单理解就是让序列中的每个位置都能“看到”整个序列并从中聚合信息。在 PyTorch 中我们一般直接使用nn.TransformerEncoderLayer和nn.TransformerEncoder组合成编码器。输入形状通常是(batch, sequence_length, feature_dim)其中batch_firstTrue可以让维度顺序更符合直觉。但要注意Transformer 本身不感知输入顺序。因此我们必须为序列添加位置编码否则对于模型来说交换两个时间步的输入会得到完全一样的结果这在时序决策里是不可接受的。位置编码可以是固定正弦函数也可以是可学习参数本文选择实现一个简单的可学习位置编码。3.2 强化学习基础回顾强化学习中有几个核心概念策略、回报、价值函数。策略π(a|s)表示在状态s下选择动作a的概率。强化学习的目标是最大化累计折扣回报。策略梯度方法通过直接对策略求导来更新参数而 REINFORCE 是最经典的蒙特卡洛策略梯度算法。REINFORCE 的更新规则可以表示为gradient E[ -log π(a|s) * G ]其中G是从当前时刻到 episode 结束的折扣回报。直观理解是如果某个动作带来了较高的回报就增大它的选择概率如果带来了较低回报就减小它的选择概率。REINFORCE 的缺点是方差较大但它的实现简单非常适合用来理解策略梯度思想。本文实战先用 REINFORCE 跑通流程进阶部分我会说明如何切换到 PPO。3.3 两者结合为什么能提升长时序决策把 Transformer 用作策略网络时输入不再是单步状态而是最近 T 步的历史状态序列。每个时间步的观测经过嵌入层后再叠加上位置编码进入 Transformer Encoder最后取最后一个时间步的输出来预测动作。这种设计的优势非常明显。Transformer 可以在多个时间步之间自动建立依赖关系相当于让策略网络在“看”历史轨迹后做判断。相比直接输入单步状态历史窗口提供了更多上下文相比 RNNTransformer 的并行度和长程依赖建模能力更强。不过也要提醒一点Transformer 不是记忆网络它不会主动保留无限长期的状态。我们需要在外部维护一个固定长度的滑动窗口并且这个窗口长度决定了模型最多能回溯多远的依赖。窗口太短信息不足窗口太长计算成本高需要根据具体任务调节。4. 完整实战用 TransformerPolicy REINFORCE 训练部分可观测 CartPole4.1 构造部分可观测环境经典 CartPole 环境默认提供 4 个观测值小车位置、速度、杆子角度、角速度。如果我们只保留位置和角度两个观测值那就丢失了速度信息环境就变成了部分可观测智能体必须结合连续多步的观测序列才能推断出运动趋势。这里我用 Gymnasium 的Wrapper做一个封装内部维护一个deque作为历史窗口每次返回的观测形状为(history_len, obs_dim)。# envs.py from collections import deque import gymnasium as gym import numpy as np class PartialObservableCartPole(gym.Wrapper): 将 CartPole 包装成部分可观测环境。 只保留 obs_indices 对应的观测维度并维护最近 history_len 步的观测序列。 def __init__(self, history_len8, obs_indices(0, 2)): env gym.make(CartPole-v1) super().__init__(env) self.history_len history_len self.obs_indices obs_indices self.history deque(maxlenhistory_len) obs_dim len(obs_indices) self.observation_space gym.spaces.Box( low-np.inf, highnp.inf, shape(history_len, obs_dim), dtypenp.float32, ) def _pack_obs(self, obs): partial np.asarray(obs, dtypenp.float32)[list(self.obs_indices)] self.history.append(partial) # 历史不足时用当前观测重复填充保证序列长度一致 while len(self.history) self.history_len: self.history.append(partial.copy()) return np.vstack(self.history).astype(np.float32) def reset(self, **kwargs): obs, info self.env.reset(**kwargs) self.history.clear() return self._pack_obs(obs), info def step(self, action): obs, reward, terminated, truncated, info self.env.step(action) return self._pack_obs(obs), reward, terminated, truncated, info这段代码有几个细节需要注意。deque(maxlenhistory_len)会自动丢弃最旧的数据所以窗口始终保留最近的历史。初始化时历史为空用当前观测重复填充这样可以避免训练刚开始时出现 NaN 或者维度不齐。4.2 实现 Transformer 策略网络策略网络接收形状为(B, T, obs_dim)的序列输入输出每个动作的 logits。模型结构如下输入线性层将obs_dim维特征映射到d_model维。可学习位置编码为每个时间步添加位置信息。TransformerEncoder包含多层自注意力。取最后一个时间步的输出经过线性层输出动作 logits。代码实现如下# policy.py import torch import torch.nn as nn from torch.distributions import Categorical class PositionalEncoding(nn.Module): 简单的可学习位置编码。 def __init__(self, d_model, max_len64): super().__init__() self.embed nn.Embedding(max_len, d_model) def forward(self, x): # x: (B, T, D) pos torch.arange(x.size(1), devicex.device) return x self.embed(pos) class TransformerPolicy(nn.Module): 基于 TransformerEncoder 的策略网络。 输入: (B, T, obs_dim) 的历史观测序列 输出: (B, action_dim) 的动作 logits def __init__( self, obs_dim2, action_dim2, d_model64, nhead4, num_layers2, dim_feedforward128, max_len64, ): super().__init__() self.input_proj nn.Linear(obs_dim, d_model) self.position_embedding PositionalEncoding(d_model, max_len) encoder_layer nn.TransformerEncoderLayer( d_modeld_model, nheadnhead, dim_feedforwarddim_feedforward, batch_firstTrue, ) self.encoder nn.TransformerEncoder(encoder_layer, num_layersnum_layers) self.action_head nn.Linear(d_model, action_dim) def forward(self, obs): # obs: (B, T, obs_dim) h self.input_proj(obs) h self.position_embedding(h) h self.encoder(h) # 取最后一个时间步对应的输出 h h[:, -1, :] logits self.action_head(h) return logits torch.no_grad() def select_action(self, obs): 推理阶段使用不计算梯度。 logits self(obs.unsqueeze(0)) dist Categorical(logitslogits) return dist.sample().item()为什么取最后一个时间步的输出而不是所有位置的平均因为这一步动作只与“当前状态及历史”有关最后一个时间步对应的输出通过自注意力已经聚合了前面所有时刻的信息逻辑上最自然。你也可以尝试h.mean(dim1)做全局平均池化两种方式在某些任务上效果会有差异但后者更偏向全局特征。4.3 实现 REINFORCE 训练循环训练部分用 REINFORCE。整体流程是重置环境得到初始观测序列。在一个 episode 内重复采样动作、执行动作、记录 log_prob 和 reward。episode 结束后计算折扣回报。利用-log_prob * return作为损失反向传播更新策略网络。# train.py import torch from torch.distributions import Categorical from envs import PartialObservableCartPole from policy import TransformerPolicy GAMMA 0.99 LR 3e-4 EPISODES 800 HISTORY_LEN 8 LOG_INTERVAL 50 def compute_returns(rewards, gammaGAMMA): 计算每个时间步的折扣回报 G_t r_t gamma * r_{t1} ... returns [] G 0.0 for r in reversed(rewards): G r gamma * G returns.append(G) returns.reverse() return torch.tensor(returns, dtypetorch.float32) def main(): env PartialObservableCartPole(history_lenHISTORY_LEN) policy TransformerPolicy( obs_dim2, action_dimenv.action_space.n, d_model64, nhead4, num_layers2, dim_feedforward128, max_lenHISTORY_LEN 1, ) optimizer torch.optim.Adam(policy.parameters(), lrLR) episode_returns [] for ep in range(1, EPISODES 1): obs, _ env.reset() log_probs [] rewards [] done False while not done: # 训练阶段需要保留梯度不能使用 select_action logits policy(obs.unsqueeze(0)) dist Categorical(logitslogits) action dist.sample() log_probs.append(dist.log_prob(action)) obs, reward, terminated, truncated, info env.step(action.item()) rewards.append(reward) done terminated or truncated returns compute_returns(rewards) # 对回报做标准化可以显著降低 REINFORCE 的方差 returns (returns - returns.mean()) / (returns.std() 1e-8) policy_loss [] for log_prob, R in zip(log_probs, returns): policy_loss.append(-log_prob * R) loss torch.stack(policy_loss).sum() optimizer.zero_grad() loss.backward() optimizer.step() total_reward sum(rewards) episode_returns.append(total_reward) if ep % LOG_INTERVAL 0: avg_reward sum(episode_returns[-LOG_INTERVAL:]) / len(episode_returns[-LOG_INTERVAL:]) print(fEpisode {ep}, avg_reward_last_{LOG_INTERVAL}{avg_reward:.2f}) env.close() print(Training finished.) if __name__ __main__: main()这里有一个容易踩坑的地方select_action加上了torch.no_grad()因此训练循环中不能调用它否则梯度会被截断模型永远不会更新。训练时应该直接调用policy(obs)得到 logits再构建Categorical分布采样。4.4 运行与验证在项目目录下执行python train.py预期输出类似Episode 50, avg_reward_last_5022.36 Episode 100, avg_reward_last_5078.54 ... Episode 800, avg_reward_last_50153.20 Training finished.注意CartPole-v1 的最大 episode 长度是 500所以平均奖励理论上限是 500。由于 Transformer 参数比普通 MLP 更多REINFORCE 的方差又比较大实际训练时曲线会有明显震荡这属于正常现象。想更稳定地看到效果可以把EPISODES提高到 2000或者引入并行环境采样。需要说明的是这个实战示例的核心目的是展示“如何把 Transformer 接入强化学习训练流程”而不是追求刷高 CartPole 分数。为了做严格对比你可以额外实现一个直接用obs[-1]单步观测的 MLP 策略网络再与 Transformer 历史窗口策略对比就能明显感受到部分可观测环境中序列建模带来的差异。4.5 进阶将 REINFORCE 换成 PPOREINFORCE 虽然简单但在更复杂的任务中方差过大更新效率低。实际项目和研究中最常用的是 PPO。PPO 与 REINFORCE 的核心差别在于PPO 通过重要性采样可以重复使用一段采样数据多次更新。PPO 引入 clipped surrogate objective防止每次更新幅度过大。PPO 使用 GAE 估计优势函数替代原始折扣回报。PPO 的 actor loss 可以用下面的伪代码理解ratio (new_log_prob - old_log_prob).exp() surr1 ratio * advantage surr2 torch.clamp(ratio, 1 - clip_eps, 1 clip_eps) * advantage actor_loss -torch.min(surr1, surr2).mean()如果把 TransformerPolicy 接入 PPO只需要在模型里增加一个 value head类似class TransformerActorCritic(nn.Module): def __init__(self, ...): super().__init__() # 共用 Transformer Encoder 部分 ... self.action_head nn.Linear(d_model, action_dim) self.value_head nn.Linear(d_model, 1) def forward(self, obs): h self.input_proj(obs) h self.position_embedding(h) h self.encoder(h) h h[:, -1, :] logits self.action_head(h) value self.value_head(h) return logits, valuePPO 已经是工业界和学术界最常用的 on-policy 算法如果你要拿这个方向做实验建议不要自己从零实现可以直接基于 Stable-Baselines3 扩展自定义策略网络或者参考 CleanRL 的 PPO 实现。5. 常见问题与排查Transformer 结合强化学习的代码写起来并不复杂但训练过程中往往会出现各种问题。下面整理了几个高频问题与排查思路。问题现象常见原因解决思路训练不收敛奖励曲线长时间接近 0学习率过大或过小调低学习率到 1e-4 到 3e-4 区间观察 loss 是否下降奖励曲线剧烈震荡REINFORCE 方差大加入 reward 标准化、增加 GAE、换 PPOTransformer 收敛速度明显慢于 MLPTransformer 参数多数据利用率低减小 d_model 和层数增加并行环境数量使用 PPO序列窗口增加后训练时间暴涨自注意力复杂度为 O(T²)缩短序列长度或者改用 FlashAttention、线性注意力模型完全不学习梯度被 torch.no_grad() 截断检查训练循环中是否误用了只用于推理的采样函数位置编码导致报错序列长度超过 max_len确保 max_len 不小于实际历史窗口长度部分可观测环境中策略依然表现差历史窗口信息不足增大 history_len或改为 RNN 隐状态编码5.1 关于收敛速度的补充说明很多同学第一次跑 Transformer 强化学习会拿它和普通 MLP 策略对比发现 MLP 在 CartPole 上收敛得更快。这是很正常的现象。CartPole 本身状态维度低、任务简单MLP 参数少自然容易训练。Transformer 的优势要等到任务需要真正的历史依赖时才体现出来。因此在做实验时不要只盯着 CartPole 的绝对分数而应该对比“单步观测 MLP”与“历史窗口 Transformer”在部分可观测设置下的差距。这种对比才是 Transformer 策略网络价值的证据。5.2 关于注意力顺序的讨论本文使用的是TransformerEncoder默认没有因果掩码也就是说每个时间步都可以看到未来信息。这在策略网络中会造成一定的问题吗严格来说在决策时刻模型不应该使用未来观测。但由于我们在部署时只输入当前时刻之前的历史序列未来信息并不存在因此训练和推理之间是一致的。如果你希望序列建模更严谨可以尝试自己做 causal mask用TransformerDecoder或自定义 mask 的TransformerEncoder。这个细节在很多论文中会被专门讨论也是你可以做改进的切入点之一。6. 最佳实践与工程建议6.1 网络设计与输入特征不要盲目堆叠 Transformer 的层数和维度。强化学习的训练数据来自与环境交互成本远高于监督学习。模型参数越多需要的数据就越多。建议从d_model64、num_layers2开始逐步增加复杂度。观测特征一定要做归一化。CartPole 的观测范围还算温和但真实任务中不同特征的量纲可能差异巨大比如位置是米、速度是米每秒、角度是弧度。如果不做归一化Transformer 的注意力权重很容易被某个特征主导训练极不稳定。输入序列的设计也很关键。你可以只输入最近 N 步原始观测也可以额外拼接动作、奖励等历史信息构造更丰富的轨迹上下文。这些都是可调整的工程细节。6.2 训练流程与稳定技巧在强化学习中稳定训练比模型结构更重要。我建议至少做到以下几点。第一使用向量化环境。Python 单环境采样的速度很慢且单条轨迹的随机性很大。可以一次性开多个并行环境采样用 batch 数据更新模型能显著提升稳定性和训练速度。第二使用 GAE 替代原始折扣回报。GAE 通过平衡偏差与方差能让优势函数的估计更加平滑。REINFORCE 适合演示原理但工程实验尽量用 PPO GAE。第三加一个熵正则项。策略网络很容易过早收敛到确定性策略导致探索不足。在 loss 中加入-beta * entropy可以鼓励模型保持一定的随机性。第四做好实验记录。每次实验固定随机种子记录模型结构、超参数、reward 曲线。否则等到第二天你根本想不起来同一个实验结果是用哪组参数跑出来的。6.3 毕设与科研方向建议如果你准备把这个方向作为毕业设计我有几个具体建议。第一个方向是复现并改进 Decision Transformer。它把轨迹整理成(return-to-go, state, action)形式的序列然后用 Transformer 直接生成动作。你可以在 Atari、MuJoCo 或自定义环境中复现它然后尝试修改 return-to-go 的 embedding 方式或者把离线强化学习算法 IQL 的 value 函数融合进来。第二个方向是探索 Transformer 在组合优化问题中的强化学习求解。比如用 Transformer 作为策略网络解决 TSP 问题这是近几年比较热门的研究点。这类问题天然是序列决策非常适合发挥 Transformer 的序列建模能力。第三个方向是研究长时序决策中的位置编码问题。固定正弦编码、可学习编码、RoPE 等不同位置编码策略对强化学习训练稳定性的影响目前还没有非常系统的结论比较容易找到创新点。不论选择哪个方向我的建议都是先跑通一个最小实现再在 baseline 上做一个小改进最后通过多组对比实验验证改进的有效性。不要一上来就试图设计一个全新的模型那样大概率会在环境搭建阶段就耗尽时间。7. 总结与学习路线本文从长时序决策的难点出发解释了 Transformer 与强化学习结合的原因然后通过一个完整的实战项目实现了基于 Transformer 策略网络的 REINFORCE 算法并讨论了 PPO 的扩展方式以及常见训练问题。读完这篇文章你应该已经掌握了一条清晰的动手路线先理解策略网络与强化学习目标再实现历史窗口与 Transformer 模型最后用策略梯度算法训练。下一步建议按以下顺序继续深入读 PPO 原始论文并用 CleanRL 或 Stable-Baselines3 复现 PPO。阅读 Decision Transformer 论文和开源代码理解序列决策建模思路。选择一个带时间依赖的场景例如机器人控制、交通信号控制、组合优化把你的策略网络迁移上去。做好实验对比记录 Transformer 与 MLP、LSTM 在长时序任务上的差异。如果你正在纠结毕业设计题目完全没有必要追求特别宏大的模型。把一个任务吃透跑通 Transformer 与强化学习的完整训练链路并且能清楚地解释每一步的原理就已经超过大多数只停留在概念层面的同学了。代码有任何问题欢迎在评论区一起讨论。
返回列表