news 2026/9/13 1:12:15

基于PyTorch的强化学习入门:从环境搭建到DQN实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于PyTorch的强化学习入门:从环境搭建到DQN实现

简介:这是一份基于PyTorch的强化学习动手实践系列资源,面向希望从代码层面理解RL算法的初学者与进阶者。内容聚焦DQN、DDPG两类经典算法,并结合OpenAI Gym中的CartPole-v0、Pendulum-v0等标准环境展示落地实现,涵盖从马尔可夫决策过程建模、经验回放到目标网络更新等关键环节,能够帮助读者摆脱“只懂理论、难以下手”的困境。资源共9个文件,以Python脚本为核心,包含可直接运行的训练与测试代码,适合对照逐步调试;mp4演示视频用于直观查看智能体从随机探索到策略收敛的学习过程;md文档则简明梳理原理与用法,便于快速索引。压缩包整体约2.6MB,轻量易下载,目录结构清楚,适合在本地环境快速复现实验。目前已有157人学习浏览,整体以“代码+视频+文档”的配套形式呈现,读者既能获得一手可运行的算法工程,也能在调试参数和对比结果中加深对强化学习核心思想的理解,是动手入门强化学习的一条高效路径。

1. 为什么「动手学强化学习」系列都基于 PyTorch 展开

强化学习和普通监督学习最大的区别,不是模型结构,而是“数据从哪来”。监督学习一个 DataLoader 喂到底,强化学习每一步都要和环境交互、拿状态、拿奖励,再决定下一个动作。这个闭环决定了框架必须支持频繁修改结构、快速打印中间张量、随时打断看梯度,PyTorch 的动态图在这种试错式迭代里比静态图顺手得多。

另一个更现实的原因是生态。Stable-Baselines3、CleanRL、ElegantRL 这些被引用最多的强化学习库,后端全是 PyTorch。手写 DQN 遇到瓶颈想换个现成实现对照效果,装了 PyTorch 就能直接跑,不需要为对比单独维护一套环境配置。所以这套教程基于 PyTorch,不是情怀,是当前强化学习领域事实上的默认路径。

下面从环境准备开始,把「照着能跑、跑完能理解」的流程完整过一遍。内容会覆盖 PyTorch 安装、gymnasium 接口、一个最小的 DQN 实现,以及训练不收敛时最值得先动的几个参数。

2. 动手前的环境准备:PyTorch + Gymnasium + Stable-Baselines3

2.1 用 Anaconda 建独立环境,Python 版本选 3.10

动手学强化学习系列最劝退的通常不是算法本身,而是装环境这一步。我的习惯是先建空环境再装 PyTorch,避免把系统 Python 搞乱:

conda create -n rl python=3.10 -y conda activate rl

Python 选 3.10 而不是最新版,是为了跟 PyTorch 2.x 的预编译包和常见扩展库的 wheel 兼容性对齐。新版本的 Python 往往要等 PyTorch 和一堆科学计算库适配完才值得升。如果你后面要装 examples 依赖或者 mujoco 这类物理引擎,3.10 也是踩坑最少的版本。

2.2 按显卡实际情况装 torch,别默认拉最大安装包

PyTorch 安装有个最常踩的坑:直接pip install torch会下载默认带 CUDA 的版本,没有 NVIDIA GPU 的机器纯白等几个 GB。正确做法是先看机器情况,再决定装哪个分支:

# CPU 机器 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu # 有 NVIDIA 显卡,先执行 nvidia-smi 看驱动支持的 CUDA 版本 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121

cu121 对应 CUDA 12.1,cu118 对应 CUDA 11.8,这是 PyTorch 官方 index 里最常用的两个组合。驱动版本比包里的 CUDA runtime 新或者相等都可以,不用严格相等。装完验证一下:

python -c "import torch; print(torch.__version__, torch.cuda.is_available())"

输出里torch.cuda.is_available()是 True 才说明 GPU 分支真的可用。这个验证我一般放在所有环境配置的最后一步,因为后面 gymnasium 和训练脚本都要依赖它。CPU 机器跑 CartPole 这类小型环境完全没压力,真正卡瓶颈的是渲染和大量并行采样,不是网络前向。

2.3 用 gymnasium 跑一次随机策略,验证环境接口

环境库现在统一用 gymnasium,而不是老版的 gym。老 gym 已经停止维护,gymnasium 是社区活跃维护的后继分支,接口变化主要在两点:reset 返回两个值,step 返回五个值:

import gymnasium as gym env = gym.make("CartPole-v1", render_mode="rgb_array") obs, _ = env.reset(seed=42) for t in range(200): action = env.action_space.sample() obs, reward, terminated, truncated, info = env.step(action) if terminated or truncated: break env.close()

这里reset(seed=42)是 gymnasium 推荐的做法,在构造 env 之前就固定随机性。terminated表示回合是否真正结束,truncated表示是否因为时步上限被截断,它们在计算 DQN 的 target 值时必须区别对待——后面第 3 节会单独讲。这一步跑通,说明环境链路已经通了,接下来可以直接写算法。

3. 基于 PyTorch 的 DQN 最小实现:回放、目标网络与训练循环

3.1 经验回放用 deque,为什么不直接用 list

DQN 相比传统 Q-learning 的三个核心机制是经验回放、目标网络、epsilon-greedy 探索。经验回放负责打破连续样本之间的相关性,实现很简单:

from collections import deque import random class ReplayBuffer: def __init__(self, capacity: int): self.buffer = deque(maxlen=capacity) def push(self, s, a, r, s_next, done): self.buffer.append((s, a, r, s_next, done)) def sample(self, batch_size: int): return random.sample(self.buffer, batch_size) def __len__(self): return len(self.buffer)

deque(maxlen=capacity)而不是 list,是因为容量满了之后会自动把头部旧数据弹出,省去手动检查长度的代码。random.sample是不重复抽样,确保一个 batch 内不出现同一条经验,梯度更新更稳。这里我刻意没有用 priority replay,动手学阶段先搞懂均匀采样,后续再考虑 PER 的差别。

3.2 Q 网络和目标网络:何时更新、怎么更新

Q 网络是一个普通的三层全连接,输入是观测维度,输出是每个动作的 Q 值:

import torch import torch.nn as nn class QNet(nn.Module): def __init__(self, obs_dim: int, act_dim: int): super().__init__() self.net = nn.Sequential( nn.Linear(obs_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, act_dim), ) def forward(self, x): return self.net(x)

目标网络初始化时直接复制一份参数,之后每若干步同步一次。同步有两种方式:硬更新是每 C 步直接拷贝参数,软更新是每步都做一次插值。软更新在 DDPG、TD3、SAC 这些连续控制算法里是标配,DQN 用软更新也完全可以:

def soft_update(target: nn.Module, source: nn.Module, tau: float = 0.005): for target_param, source_param in zip(target.parameters(), source.parameters()): target_param.data.copy_( target_param.data * (1 - tau) + source_param.data * tau )

tau 取 0.005 表示目标网络每步向当前网络挪动 0.5%,相当于约 200 步做一次完整替换的平滑版本。软更新的好处是目标值变化平滑,训练更稳定,代价是前期收敛稍慢。如果你在复现论文中的结果,论文如果写的是硬更新,就按论文来,而不是无脑套软更新。

3.3 训练循环里的关键一步:target 怎么算

一个 train_step 里最容易被新手写错的就是 target 的计算。常见做法是:

def train_step(batch, q_net, target_net, optimizer, gamma=0.99): batch_state = torch.FloatTensor(np.array([t[0] for t in batch])).to(device) batch_action = torch.LongTensor(np.array([t[1] for t in batch])).unsqueeze(1).to(device) batch_reward = torch.FloatTensor(np.array([t[2] for t in batch])).unsqueeze(1).to(device) batch_next_state = torch.FloatTensor(np.array([t[3] for t in batch])).to(device) batch_done = torch.FloatTensor(np.array([t[4] for t in batch])).unsqueeze(1).to(device) q_values = q_net(batch_state).gather(1, batch_action) with torch.no_grad(): max_next_q = target_net(batch_next_state).max(dim=1, keepdim=True)[0] target = batch_reward + gamma * (1 - batch_done) * max_next_q loss = nn.functional.mse_loss(q_values, target) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()

这里gather(1, batch_action)的作用是从 Q 网络输出的每个动作 Q 值里,取出本次采样用到的那个动作对应的值。max_next_q 在 target_net 上计算,并且包在no_grad()里,因为梯度只需要回流到 q_net,不能让 target 的生成过程参与反向传播。(1 - batch_done)是价值 bootstrapping 的开关——当回合终止时,后续没有奖励了,target 直接等于当前奖励;未终止时才累加下一状态的价值。

3.4 terminated 和 truncated 会直接影响 target 的值

gymnasium 把回合终止细分为两个布尔值,这在 DQN 里不是小事。truncated 只是时间预算用完了,不表示任务失败,未来状态的价值依然存在;terminated 才是 MDP 意义上的真正终止。很多公开代码为了省事直接写done = terminated or truncated,在 CartPole 这种 200 步必截断的环境里会产生系统性偏差。我的做法是:训练时把terminated当 done 用,truncated当作未终止继续回传折扣价值;只有评估时才把两者都当作回合结束。

4. 训练不收敛时按顺序排查的 5 个参数

DQN 训练不收敛,先不要怀疑网络结构,绝大多数问题出在超参数和环境交互的细节上。下面这张表是我常用的起点配置,也是排查方向的索引。

参数推荐起点常见问题信号
learning_rate1e-4 ~ 3e-4loss 变成 NaN 先查它
batch_size128太大训练慢,太小震荡明显
replay_capacity100k ~ 500k太小样本反复被清除
gamma0.99接近 1 时 Q 值方差大
epsilon_decay1.0 → 0.01,50 万步降到 0.01 太快,策略早熟

4.1 奖励绝对值过大:先做 reward scaling

DQN 的 target 是即时奖励加折扣未来价值。如果环境的奖励量纲本身很大,比如很多机械臂环境单步给 +10,Q 值会跟着膨胀,MSE loss 会出现陡峭尖峰。常见处理是把 reward 除以一个常数,例如reward = reward / 100.0,或者做简单的 clip。CartPole 这类 ±1 的环境不需要处理,但从其他环境迁移过来时,这一步务必检查。

4.2 epsilon 衰减:别前 10% 步数就把探索用光

epsilon 从 1.0 线性降到 0.01 是最直白的写法,但实际训到后面会发现,前期火力全开探索,中期策略就已经定型,后期几乎不动。我一般用指数衰减并分成两段:前 60% 步数从 1.0 降到 0.1,后 40% 从 0.1 缓慢降向 0.01。这样保证中后期还有一定随机性去跳出局部策略。衰减太快是最隐蔽的过拟合来源,模型记住了前期几条轨迹,后面再怎么训练也翻不了身。

4.3 batch_size 和学习率要一起调,别单动一个

大 batch 的梯度方向更准,学习率可以适当调大;小 batch 噪声大,学习率就得调小。常见的错误是 batch_size 调到 256 而学习率还在 3e-3,loss 会直接震荡。反过来,如果用 32 的 batch 但学习率只有 1e-4,训练会慢到以为死机了。这两个参数是联动的,不要单独调其一。换环境换网络时,我通常先固定 batch_size,把学习率按 3 倍粒度扫三档,确定稳定区间后再微调 batch_size。

4.4 回放容量:不是越大越好,但太小一定坏

replay buffer 容量的经验值是 100k 到 1M。容量太小,旧经验一直被新经验挤掉,样本多样性不够,训练会反复在最近几条轨迹上打转。容量太大的副作用主要是内存,CartPole 状态只有 4 维时无所谓,图像输入环境下 1M 条经验会吃掉好几个 GB,这时候反而要下调容量并配合更积极的探索。做消融实验时,这个参数常常被人忽略。

4.5 随机种子:三个层面都要固定,否则没法对比

RL 实验结果方差很大,排查问题时如果不固定随机性,同一个参数跑两遍结果能差出一大截。固定种子的位置有三个层面:

import random import numpy as np import torch def set_seed(seed: int = 42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)

注意cudnn.deterministic = Truebenchmark = False只在完全确定性的网络里有效,CartPole 没问题,全连接网络没影响,但涉及卷积和图像输入时会牺牲少量速度。调试问题时值得开,跑正式实验时我会关掉deterministic,只保留随机种子,因为 RL 的高熵探索本身就需要一定随机性。

5. 训练中途就该做的离线评估:用干净环境给模型打分

训练时不要等到全部跑完才看效果,DQN 的 epsilon 探索噪声会干扰你对模型真实水平的判断。更好的做法是每隔固定步数,在关闭探索的情况下跑若干回合评估,然后保存历史平均奖励最高的那组参数。

5.1 评估函数:关探索、开 no_grad

评估环境的逻辑和训练一致,但动作选择不再走 epsilon-greedy,而是直接 argmax:

def evaluate(env, q_net, episodes=10): q_net.eval() returns = [] for _ in range(episodes): obs, _ = env.reset() total_reward = 0.0 terminated, truncated = False, False while not (terminated or truncated): with torch.no_grad(): obs_t = torch.FloatTensor(obs).unsqueeze(0) action = q_net(obs_t).argmax(dim=1).item() obs, reward, terminated, truncated, _ = env.step(action) total_reward += reward returns.append(total_reward) return float(np.mean(returns))

评估环境不要开启render_mode="human",否则每个回合弹出渲染窗口,几百步下来非常拖速度;需要录轨迹时用rgb_array。episodes 至少跑 10 回合,因为单回合的随机性太大,一次高分不能说明问题。

5.2 保存最优模型而不是最后一步

模型保存用 state_dict 而不是整个 model,方便之后按需加载:

torch.save(q_net.state_dict(), f"dqn_best_{global_step}.pt")

配合一个全局变量记录历史最高平均奖励,只有当当前评估分数更高时才覆盖保存。这样即使训练后期过拟合导致分数回落,手里的 checkpoint 还是之前的最好状态。

训练中期可以同时看三条曲线来判断模型状态:eval_avg_reward 上升并稳定说明策略在变好;q_value 增长过快但 reward 不见涨,是高估的信号,这时可以考虑切到 Double DQN 架构;td loss 有震荡但整体呈下降趋势,属于正常现象,不用因为单次 spike 就去调学习率。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/13 1:11:47

Boss直聘数据分析实战:薪资解析与投递量预测全流程

简介:面向求职市场数据分析与期末作业参考的实战案例包,以 Boss 直聘招聘数据为对象,完整覆盖数据获取、预处理、探索性分析与机器学习建模等环节。压缩包约 12.51MB,包含 Data-Analysis-Project-master 项目文件夹,内…

作者头像 李华
网站建设 2026/9/13 0:44:50

WiFi温湿度传感器MQTT接入故障排查与实战配置指南

1. 项目概述:为什么一个WiFi温湿度传感器的配置,值得花一整篇干货来写?你手头刚拆开一个标着“WiFi温湿度传感器”的小盒子,背面贴着DHT22或SHT30的标签,说明书里印着“支持MQTT”、“兼容2.4GHz WiFi”,但…

作者头像 李华
网站建设 2026/9/13 0:27:06

8款主流AI论文平台横向实测,本硕博论文避坑全攻略

前言:AI 写论文乱象频发,实测 8 款工具理清适配边界 每到毕业季,本科生、硕博生都会集中寻找 AI 论文辅助工具,市面各类写作软件层出不穷,但普遍存在几类硬伤:虚假参考文献、无法匹配本校格式、不支持公式代…

作者头像 李华
网站建设 2026/9/13 0:24:32

天津数字沙盘制作公司哪家好?本地服务商推荐与案例参考

天津作为京津冀协同发展的核心城市,房地产市场稳步发展,滨海新区、武清、西青、津南、宝坻等区域新项目密集,对数字沙盘的需求持续增长。但天津本地专业数字沙盘公司相对较少,很多开发商选择北京公司服务。本文介绍天津数字沙盘市…

作者头像 李华