尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

从CartPole到Pendulum:DQN与DDPG的PyTorch复现指南

从CartPole到Pendulum:DQN与DDPG的PyTorch复现指南 简介这是《动手学强化学习》系列的PyTorch实战资源面向希望从零上手强化学习的初学者也适合正在做课程实验或毕业设计的学生。资料围绕强化学习的核心思想展开强化学习不同于监督学习依赖标注数据而是通过智能体与环境交互获得奖励信号来不断优化策略并在探索与利用之间寻求平衡本包精选了DQN与DDPG两种经典算法分别应用于CartPole-v0和Pendulum-v0环境覆盖了从环境交互、经验回放到网络更新的完整流程能够帮助读者在实践中理解马尔可夫决策过程、值函数估计与策略优化。压缩包共9个文件以4个Python脚本为主另有2个视频演示、2个pyc缓存及1个Markdown说明文档整体仅2.6MB轻量易翻阅。该资源目前已有157人学习浏览。通过源码可分析神经网络结构与训练参数视频直观展示智能体学习效果README辅助梳理依赖与运行步骤很适合作为入门强化学习的实操参考也可作为相关课程作业的代码模板。1. 从 CartPole 到 PendulumDDPG 与 DQN 的 PyTorch 复现图谱拿到的这份《动手学强化学习》源码包拆开就是两个完整的 PyTorch 训练脚本与配套 READMEDQN CartPole-v0和DDPG Pendulum-v0。前者是离散动作空间最经典的入门题后者是连续控制任务里绕不开的基准环境。两个脚本都不带框架级的抽象网络定义、经验回放、目标网络更新、噪声注入全部平铺在眼前适合刚把 PyTorch 基础框架跑熟、想亲手看一眼强化学习算法每一步数值流向的读者。如果你需要在推荐系统、机器人控制这类场景里落地强化学习或准备复现论文对比实验这份代码也够当脚手架。它能帮你建立两个关键认知DQN 把 Q-learning 从查表换成函数逼近后为什么必须搭配经验回放与目标网络DDPG 又是如何用确定性策略与 Actor-Critic 结构绕开连续动作空间的 argmax 问题。下面从算法原理、代码路径、环境适配三个维度拆开讲。2. DQN 离散控制经验回放、目标网络与 epsilon 衰减的协作机制2.1 为什么 CartPole 是 DQN 的最佳验证场CartPole 的状态空间是 4 维连续向量小车位置、速度、杆角度、角速度动作空间是 2 个离散动作左推、右推。环境本身是马尔可夫决策过程的标准形态每一步返回一个奖励信号目标是让杆子保持竖直的时间尽可能长。这个任务的关键在于状态和动作之间不是线性映射——同样的角度在不同速度下需要不同的修正力度所以线性策略很快会失效必须靠非线性函数逼近来拟合 Q 值。DQN 在这里做的事情是用一个参数化的 Q 网络Q(s, a; θ)替代 Q-learning 中的 Q 表输入状态输出每个动作的 Q 值估计。更新目标不再是查表得到的下个状态最大值而是用目标网络计算的r γ * max_a Q_target(s, a)这就是把强化学习问题转化成监督学习问题的核心闭包。因为目标值里含有待优化的网络自身输出直接自举会导致训练震荡目标网络每隔 N 步同步一次参数让 TD 误差的计算有一个相对稳定的参照系。2.2 DQN 核心代码路径与超参设置源码中 DQN 的结构采用双隐藏层全连接网络这是处理 CartPole 这种低维状态的标准配置网络太深反而容易在少量样本上过拟合环境噪声。class DQN(nn.Module): def __init__(self, state_dim, action_dim, hidden_dim128): super(DQN, self).__init__() self.fc1 nn.Linear(state_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.fc3 nn.Linear(hidden_dim, action_dim) def forward(self, x): x torch.relu(self.fc1(x)) x torch.relu(self.fc2(x)) return self.fc3(x)训练循环里最关键的部件有三个经验回放缓冲区、目标网络软硬同步、epsilon 贪心探索。每次环境交互得到的(s, a, r, s, done)四元组先存入回放缓冲区训练时从中随机采样一个小批量打破样本之间的时序相关性——这是策略梯度类算法没有、而 DQN 必须有的机制。# 经验回放采样与 TD 目标计算 batch_state, batch_action, batch_reward, batch_next_state, batch_done memory.sample(batch_size) q_values q_net(batch_state).gather(1, batch_action.unsqueeze(1)).squeeze(1) with torch.no_grad(): max_next_q target_net(batch_next_state).max(1)[0] q_target batch_reward gamma * max_next_q * (1 - batch_done) loss nn.MSELoss()(q_values, q_target) optimizer.zero_grad() loss.backward() optimizer.step()代码里gather(1, action)是取当前状态下实际执行动作对应的 Q 值max(1)[0]是取目标网络在下个状态所有动作中的最大值。(1 - batch_done)这个掩码处理了终止状态——对话结束后的虚拟 Q 值应为 0否则会把终结后的伪价值继续往回传递。超参数上这个实现用的是入门基准但如果想跑出稳定收敛建议按下面表格调整参数默认值调整建议learning_rate1e-3若 loss 早期剧烈震荡降到 3e-4gamma0.99任务步数较短可降到 0.95batch_size64样本量大时可升到 128注意内存占用target_update100每 100 步把 eval 网络权重拷给 target 网络epsilon_decay0.995衰减太慢探索期过长太快则前期策略空洞replay_size10000至少覆盖 20 个回合的交互数据epsilon 从 1.0 开始线性或指数衰减到 0.01 附近保持。这个过程让智能体前期大量探索环境后期逐渐转为利用已有 Q 值。需要留意的是epsilon 衰减率和target_update间隔是一对互相牵制的参数探索不足时目标网络的更新频率必须调低否则 Q 值会朝一个尚未成熟的估计方向漂移。3. DDPG 连续控制Actor-Critic 架构、软更新与 OU 噪声的配合3.1 连续动作空间为什么逼着我们把策略参数化Pendulum-v0 的任务目标同样简单直观给一个摆施加力矩让它摆到竖直向上的平衡位置并保持。难点在于动作是一个 [-2, 2] 区间内的连续扭矩值而不是几个离散选项。DQN 在这类问题上直接失效——max_a Q(s, a)需要对连续动作空间求最大值遍历或网格采样都不可行。DDPG 的思路是把策略本身建模成一个确定性映射a μ(s; θμ)让 Actor 网络直接输出动作再用 Critic 网络评估这个动作的价值。这样求最大值的问题被转换成了对 Actor 参数做梯度上升的问题沿着Q(s, μ(s))对动作的梯度方向调整 Actor使输出动作的 Q 值逐步提高。3.2 Actor-Critic 网络设计与软更新实现源码中的 Actor 网络输出层用了 tanh 激活函数把原始输出压到 [-1, 1]再乘上动作边界action_bound2。这里有个隐藏细节tanh 的输出在接近 ±1 时梯度几乎为 0所以网络初始权重必须设得足够小否则早训练阶段 Actor 直接饱和在边界上Critic 拿到的永远是极端动作样本。def forward(self, state): x torch.relu(self.fc1(state)) x torch.relu(self.fc2(x)) return torch.tanh(self.fc3(x)) * self.action_boundCritic 网络的输入是状态和动作的拼接向量源码里通过torch.cat([state, action], dim1)完成。拼接维度要对齐state 是 3 维cosθ, sinθ, θ˙action 是 1 维拼接后 4 维输入。Critic 输出的是标量 Q 值用来评估(s, a)组合的好坏。DDPG 与 DQN 的目标网络更新方式不同DQN 每 N 步硬拷贝一次DDPG 每个训练步都做软更新。软更新的公式是θ_target ← τ·θ_eval (1-τ)·θ_targetτ 通常取 0.005。这个操作让目标网络以极慢的速度跟踪在线网络借用一句形象的说法目标网络是在线网络的滑动平均它的缓慢移动保证了 TD 目标的稳定性同时不引入 DQN 那种周期性跳变。def soft_update(target, source, tau): for target_param, param in zip(target.parameters(), source.parameters()): target_param.data.copy_(tau * param.data (1.0 - tau) * target_param.data)训练循环中 Critic 的更新与 DQN 相似只是 TD 目标里的max_a Q(s, a)换成了Q_target(s, μ_target(s))——即用目标 Actor 输出目标动作、目标 Critic 评估价值。Actor 的 loss 则是-Q(s, μ(s))的均值取负号是因为 PyTorch 只做梯度下降这样 Critic 给分越高的动作方向Actor 更新幅度越大Critic 的收敛质量直接决定 Actor 的学习方向是否正确。3.3 OU 噪声给确定性策略注入随机性DDPG 是确定性策略如果在训练时始终输出同一个动作环境探索量会严重不足。源码中在动作上叠加了 Ornstein-Uhlenbeck 噪声它和普通高斯噪声的区别在于拥有均值回归特性噪声项会随时间逐渐向 0 回归保证训练前期探索幅度大后期策略趋于稳定时扰动自然减小。class OUNoise: def __init__(self, action_dim, mu0.0, theta0.15, sigma0.2): self.mu mu self.theta theta self.sigma sigma self.reset() def sample(self): dx self.theta * (self.mu - self.state) self.sigma * np.random.randn(len(self.state)) self.state dx return self.statetheta控制噪声向均值回归的速度越大噪声越快地趋向 0sigma控制随机扰动的幅度。真实训练时如果发现智能体老在同一个动作附近打转把sigma提到 0.3如果训练后期策略抖动把theta降到 0.1。噪声过大的副作用是动作频繁触碰边界Pendulum 环境会因此累积较大的角速度需要更长回合才能刹住。DDPG 的超参经验值和 DQN 有明显差异。Actor 和 Critic 的学习率往往需要分开设置Critic 用 1e-3、Actor 用 1e-4 是常见组合原因是 Critic 的 loss 直接由 TD 误差驱动收敛相对快Actor 需要跟着 Critic 的估计走步子大了容易踩空。replay buffer 在 DDPG 中通常设到 100000 以上因为连续控制任务的样本利用率低需要更大规模的离线数据池来稳定训练。4. gym 环境适配与 PyTorch 版本兼容的三个坑4.1 Pendulum-v0 在 gym 新版本中已移除源码里写的是gym.make(Pendulum-v0)但 gym 0.26 及之后的版本已经移除了所有带-v0后缀的经典控制环境。直接运行会报EnvNotFound错误。解决办法是把环境名改成Pendulum-v1API 接口完全一致抽样空间和奖励函数没有变化改一行字符串即可。import gym env gym.make(Pendulum-v1) # 旧代码是 Pendulum-v0如果因为历史依赖必须留在旧版 gym可以用pip install gym0.21.0锁定版本。但更推荐的做法是直接迁移到 v1因为旧版 gym 在 Python 3.10 上 import 时就会报distutils相关的 ModuleNotFoundError那个错误和你的算法代码毫无关系纯粹是环境依赖失衡。4.2 np.float 属性在 numpy 1.24 中被移除源码里训练循环通常用np.float做数据类型转换比如state np.float32(state)没问题但如果是np.float不带 32/64numpy 1.20 开始 deprecation 警告、1.24 直接移除。运行时会看到AttributeError: module numpy has no attribute float。# 检查当前 numpy 版本 python -c import numpy; print(numpy.__version__)对依赖 Python 3.10 的环境numpy 会自动装到 1.24所以这个错误基本必现。修复方式把代码中所有np.float替换为np.float64或者降级pip install numpy1.24。前者是正道后者是权宜之计。4.3 Box space 边界检查与状态归一化Pendulum 的 observation space 是Box(shape(3,), low[-1, -1, -8], high[1, 1, 8])前两维是角度的 cos/sin 值天然落在 [-1, 1] 区间第三维是角速度范围 ±8。这个尺度差异虽然不比 CartPole 严重但训练时我习惯在进入网络前做一次标准化把观测向量除以空间的high值让所有维度落到大致相同的量纲。obs env.reset() obs obs / env.observation_space.high # 分量级归一化归一化对 DDPG 的影响比 DQN 更显著因为 Critic 同时拼接状态和动作作为输入如果状态各分量尺度差一个数量级Critic 的初始 loss 会被大数值维度主导收敛方向从一开始就偏了。DQN 对尺度不敏感但归一化能显著提升复现稳定性属于做了不亏的改动。环境兼容性问题排查时一个实用技巧是在启动训练前打印环境信息print(env.observation_space) print(env.action_space) print(env.action_space.high, env.action_space.low)这样能在算法报错之前就确认环境本身的状态空间维度、动作边界和类型。三个经典问题按出现频率排序分别是环境名失效gym 版本迁移、np.float被移除numpy 升级、env.step返回 4 元组还是 5 元组差异gym 0.26 后done被拆成terminated和truncated。最后一个问题会导致解包报错处理方式是obs, reward, terminated, truncated, info env.step(action)然后done terminated or truncated。5. 训练验证技巧奖励曲线、超参敏感性与复现性控制5.1 滑动平均奖励曲线怎么画才有意义CartPole 的单回合奖励天然是递增的趋势原始曲线抖动很大画在一张图上看不出收敛趋势。我一般会维护一个rewards_history列表每 10 个 episode 计算一次滑动平均再叠加原始曲线的淡色透明层这样既能看到波动幅度也能看到整体走向。import pandas as pd import matplotlib.pyplot as plt rewards pd.Series(episode_rewards) rolling_mean rewards.rolling(window50, min_periods10).mean() plt.figure(figsize(10, 5)) plt.plot(rewards, alpha0.3, labelraw) plt.plot(rolling_mean, colorred, label50-episode rolling mean) plt.xlabel(Episode) plt.ylabel(Total Reward) plt.legend() plt.show()窗口大小选 50 比较折中100 窗口太钝看不出中期的阶段性波动20 窗口太灵敏容易让人误判局部回退。看曲线时盯住两个特征滚动均值是否持续抬升以及单回合 reward 的最大值是否逼近环境理论值。CartPole 到了 500 分说明已到环境 step 上限Pendulum 的最优 reward 在 -140 附近低于 -300 说明还没逃离初始摆动的死区。5.2 超参敏感性的快速诊断方法如果 DQN 训练不收敛不要盲目调学习率。先看 TD loss 曲线loss 持续下降但 reward 不涨说明 Q 值在进步但 epsilon 衰减过快策略执行跟不上价值估计loss 直接发散说明学习率过高或者目标网络更新太频繁。DDPG 的常见病是 Critic loss 震荡而 Actor loss 一直不降这时优先把 Critic 学习率降一半或者把tau从 0.005 改成 0.001让目标网络更慢地跟踪在线网络。快速定位问题的一个做法是把超参写成字典集中管理跑多个 seed 对比config { lr_actor: 1e-4, lr_critic: 1e-3, gamma: 0.99, tau: 0.005, batch_size: 128, buffer_capacity: 100000, }5.3 seed 固定三件套与训练探针复现实验必须同时固定随机种子PyTorch、numpy、Python random 三个模块缺一不可。环境本身的初始状态随机性也需要固定env.seed(seed)在 gym 新版中改成了env.reset(seedseed)这一步经常被漏掉导致同一个 seed 下实验结果仍然不可复现。import random import numpy as np import torch random.seed(0) np.random.seed(0) torch.manual_seed(0) env.reset(seed0)除了画奖励曲线我还会加一个简单的探针每 500 步打印当前 batch 的平均 Q 值和平均 reward。平均 Q 值持续上涨但平均 reward 持平说明网络在“过度乐观估计”两者一起停滞大概率是探索不足或网络容量不够。这个探针能提前发现训练异常不用等整个实验跑完才在最终曲线上看到问题。数据流验证也有一个高频坑DDPG 往 replay buffer 存 transition 时动作必须乘以action_bound之后的真实环境动作而不是 Actor 网络 tanh 输出的 [-1, 1] 值。存错了的话Critic 输入的动作尺度和实际执行的动作尺度不一致训练初期 loss 会异常小等 buffer 里积累足够多的错误数据后突然发散。最开始 debug 时可以只存 1000 条 transition 就手动打印一遍 buffer 内的 reward 均值和动作范围确认流向正确再放开跑长时间训练。本文还有配套的精品资源点击获取
返回列表