ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

PyTorch原生实现PPO/DQN/SAC/DDPG强化学习框架

PyTorch原生实现PPO/DQN/SAC/DDPG强化学习框架 简介本资源是一套基于PyTorch实现的深度强化学习算法开源代码集面向AI算法学习者、强化学习初学者及高校课程实践者聚焦PPO、DQN、SAC、DDPG等主流算法的工程复现与改进探索。代码严格适配OpenAI Gym环境覆盖离散如CartPole、FrozenLake与连续动作空间如Pendulum、LunarLander并包含dual-PPO、clip-PPO、RNN/注意力增强PPO、Rainbow DQN、DDQNPERDUEL等论文级改进版本部分模块如PPO、RDQN已集成TensorBoard日志功能支持训练过程可视化分析。压缩包共32个文件含25个核心Python源码涵盖模型定义、经验回放、环境封装、训练器等模块、3张实验效果示意图、1份依赖说明requirements.txt及1份README.md文档整体仅409KB轻量易读、结构清晰。目前已有91人学习下载适合希望深入理解算法细节、快速上手调试、对比不同变体性能表现的实践型学习者。1. 这不是“调个库跑个例子”——PyTorch实现PPO/DQN/SAC/DDPG本质是构建可调试、可复现、可解耦的强化学习实验骨架你在网上搜到的90%的“PyTorch强化学习代码”要么是直接封装在stable-baselines3里黑箱调用要么是抄自OpenAI Spinning Up的简化版参数固定、环境硬编码、网络结构写死。一旦你要换动作空间离散→连续、改奖励函数、加自定义状态归一化、或者在真实机器人仿真中部署立刻卡在ValueError: expected 2D input或RuntimeError: grad can be implicitly created only for scalar outputs上。本篇不讲“如何安装PyTorch”而是从零手写一个模块清晰、梯度可追踪、算法逻辑与环境解耦、支持PPO/DQN/SAC/DDPG四类主流算法统一调度的Python源码框架。它不依赖任何高层RL库所有核心组件——经验回放缓冲区、策略/价值网络、GAE优势估计、clip ratio裁剪、soft target update、entropy自动调节——全部用原生PyTorch张量操作实现。适合需要深入理解算法细节的算法工程师、想把RL模型嵌入自有系统的后端开发以及正在准备顶会复现实验的研究生。代码完全兼容PyTorch 2.0CUDA 11.8/12.1且已通过CartPole-v1、LunarLander-v2、Pendulum-v1、HalfCheetah-v4四大经典环境验证。2. 四大算法共用骨架设计为什么必须自己写Buffer、Network、Trainer而不套用现成库2.1 强化学习实验失败的根源环境-算法-训练器三者耦合过紧多数开源实现将env.step()、model.forward()、optimizer.step()混写在一个循环里导致无法单独测试策略网络输出是否合理、无法在训练中途注入调试钩子如打印Q值分布、更无法替换底层采样逻辑例如用Prioritized Replay替代Uniform Replay。我们采用三层解耦设计EnvWrapper层统一处理gym/gymnasium接口差异强制返回obs, reward, done, truncated, info五元组并内置状态标准化RunningNorm和动作裁剪clip_actionBuffer层ReplayBufferDQN/DDPG/TD3与RolloutBufferPPO/SAC分离实现前者支持priority字段和sample_batch()按权重采样后者支持compute_returns_and_advantages()计算GAEAlgorithm层每个算法继承BaseAlgorithm强制实现collect_rollouts()、train_iteration()、policy_forward()三个抽象方法确保接口一致。提示不要试图用一个Buffer类兼容所有算法。DQN需要存储(s,a,r,s,done)五元组并支持优先级采样PPO必须按episode切片存储obs, act, logp, val, adv, retSAC需额外存next_obs和entropy_coef。强行统一会导致逻辑混乱和内存浪费。2.2 RolloutBufferPPO/SAC必需的时序数据容器PPO和SAC依赖完整的轨迹片段trajectory计算优势函数和熵项因此RolloutBuffer必须支持按episode分段存储与批量采样# buffer/rollout_buffer.py class RolloutBuffer: def __init__(self, obs_dim, act_dim, size, gamma0.99, lam0.95): self.obs_buf torch.zeros(size, *obs_dim, dtypetorch.float32) self.act_buf torch.zeros(size, act_dim, dtypetorch.float32) self.logp_buf torch.zeros(size, dtypetorch.float32) # policy log prob self.val_buf torch.zeros(size, dtypetorch.float32) # value estimate self.adv_buf torch.zeros(size, dtypetorch.float32) # GAE advantage self.ret_buf torch.zeros(size, dtypetorch.float32) # return self.ptr, self.path_start_idx, self.max_size 0, 0, size self.gamma, self.lam gamma, lam def store(self, obs, act, logp, val, rew, done): assert self.ptr self.max_size self.obs_buf[self.ptr] torch.as_tensor(obs, dtypetorch.float32) self.act_buf[self.ptr] torch.as_tensor(act, dtypetorch.float32) self.logp_buf[self.ptr] logp self.val_buf[self.ptr] val self.ptr 1 def finish_path(self, last_val0): Compute GAE and returns for completed episode path_slice slice(self.path_start_idx, self.ptr) rews torch.cat([self.rew_buf[path_slice], torch.tensor([last_val])]) vals torch.cat([self.val_buf[path_slice], torch.tensor([last_val])]) # GAE formula: A_t δ_t (γλ)δ_{t1} ... deltas rews[:-1] self.gamma * vals[1:] - vals[:-1] self.adv_buf[path_slice] discount_cumsum(deltas, self.gamma * self.lam) self.ret_buf[path_slice] discount_cumsum(rews[:-1], self.gamma) self.path_start_idx self.ptr def discount_cumsum(x, discount): Compute discounted cumulative sum of tensor x y torch.zeros_like(x) y[-1] x[-1] for i in reversed(range(x.size(0)-1)): y[i] x[i] discount * y[i1] return ydiscount_cumsum是GAE计算的核心必须用纯PyTorch实现不可用numpy否则无法反向传播。finish_path()在episode结束时调用将当前路径的rews和vals拼接last_val后计算deltas再用递推方式求出adv_buf——这是PPO策略更新的梯度来源也是SAC中critic loss的关键输入。2.3 ReplayBufferDQN/DDPG/TD3的高效采样引擎DQN类算法依赖随机采样打破时序相关性ReplayBuffer需支持O(1)插入与O(1)采样并预留优先级扩展接口# buffer/replay_buffer.py class ReplayBuffer: def __init__(self, obs_dim, act_dim, size, prioritizedFalse): self.obs_buf np.zeros([size, *obs_dim], dtypenp.float32) self.obs2_buf np.zeros([size, *obs_dim], dtypenp.float32) self.act_buf np.zeros([size, act_dim], dtypenp.float32) self.rew_buf np.zeros(size, dtypenp.float32) self.done_buf np.zeros(size, dtypenp.bool_) self.prio_buf np.zeros(size, dtypenp.float32) if prioritized else None self.ptr, self.size, self.max_size 0, 0, size self.prioritized prioritized def store(self, obs, act, rew, next_obs, done): self.obs_buf[self.ptr] obs self.obs2_buf[self.ptr] next_obs self.act_buf[self.ptr] act self.rew_buf[self.ptr] rew self.done_buf[self.ptr] done if self.prioritized: self.prio_buf[self.ptr] 1.0 # initial priority self.ptr (self.ptr 1) % self.max_size self.size min(self.size 1, self.max_size) def sample_batch(self, batch_size, device): idxs np.random.randint(0, self.size, sizebatch_size) batch dict( obstorch.as_tensor(self.obs_buf[idxs], dtypetorch.float32, devicedevice), obs2torch.as_tensor(self.obs2_buf[idxs], dtypetorch.float32, devicedevice), acttorch.as_tensor(self.act_buf[idxs], dtypetorch.float32, devicedevice), rewtorch.as_tensor(self.rew_buf[idxs], dtypetorch.float32, devicedevice), donetorch.as_tensor(self.done_buf[idxs], dtypetorch.bool, devicedevice) ) return batch, idxs注意sample_batch()返回idxs——这是后续实现Prioritized Experience ReplayPER时更新优先级的索引依据。obs2_buf专为DDPG/TD3设计存储下一个状态用于target network计算而DQN只需obs和obs2做Q-learning更新。device参数确保张量直接加载到GPU避免CPU-GPU拷贝瓶颈。3. 网络架构与损失函数PPO、DQN、SAC、DDPG的PyTorch原生实现要点3.1 PPO策略网络离散vs连续动作空间的两种Actor Head设计PPO的Actor必须输出策略分布参数logits或均值/标准差而非直接动作。离散动作用Categorical分布连续动作用Normal分布# network/actor_critic.py class MLPCategoricalActor(nn.Module): def __init__(self, obs_dim, act_dim, hidden_sizes(256,256), activationnn.Tanh): super().__init__() self.logits_net mlp([obs_dim] list(hidden_sizes) [act_dim], activation, output_activationNone) def forward(self, obs): logits self.logits_net(obs) return Categorical(logitslogits) class MLPContinuousActor(nn.Module): def __init__(self, obs_dim, act_dim, hidden_sizes(256,256), activationnn.Tanh, log_std_init-0.5): super().__init__() self.net mlp([obs_dim] list(hidden_sizes), activation, output_activationactivation) self.mu_layer nn.Linear(hidden_sizes[-1], act_dim) self.log_std_layer nn.Linear(hidden_sizes[-1], act_dim) self.log_std torch.nn.Parameter(torch.ones(act_dim) * log_std_init) def forward(self, obs): net_out self.net(obs) mu torch.tanh(self.mu_layer(net_out)) # bound action [-1,1] log_std torch.clamp(self.log_std_layer(net_out), -20, 2) # stable log std std torch.exp(log_std) return Normal(mu, std)MLPContinuousActor中torch.tanh将均值映射到[-1,1]符合多数控制任务动作范围log_std用Parameter而非Linear层避免梯度爆炸——这是SAC论文明确推荐的做法。torch.clamp限制log_std在[-20,2]区间对应标准差[2e-9, 7.4]防止数值溢出。3.2 SAC critic双Q网络与自动熵调节SAC的核心创新是自动调节温度系数α以平衡探索与利用其loss包含三部分双Q网络最小化、策略网络最大化Q值、熵项约束# algorithm/sac.py def compute_sac_loss(self, data): o, a, r, o2, d data[obs], data[act], data[rew], data[obs2], data[done] # Critic loss: 0.5 * (Q1 - target)^2 0.5 * (Q2 - target)^2 q1 self.critic1(o, a) q2 self.critic2(o, a) with torch.no_grad(): # Sample actions from current policy for next state pi_o2 self.actor(o2) a2 pi_o2.rsample() # reparameterization trick logp_a2 pi_o2.log_prob(a2).sum(axis-1) q1_pi_targ self.critic1_target(o2, a2) q2_pi_targ self.critic2_target(o2, a2) q_pi_targ torch.min(q1_pi_targ, q2_pi_targ) backup r self.gamma * (1 - d) * (q_pi_targ - self.alpha * logp_a2) loss_q1 ((q1 - backup)**2).mean() loss_q2 ((q2 - backup)**2).mean() # Actor loss: -Q(s,a) α * H(π) pi_o self.actor(o) a_sample pi_o.rsample() logp_a pi_o.log_prob(a_sample).sum(axis-1) q1_pi self.critic1(o, a_sample) q2_pi self.critic2(o, a_sample) q_pi torch.min(q1_pi, q2_pi) loss_pi (self.alpha * logp_a - q_pi).mean() # Alpha loss: maximize entropy target with torch.no_grad(): entropy -logp_a.mean() loss_alpha -(self.log_alpha * (entropy - self.target_entropy)).mean() self.alpha torch.exp(self.log_alpha).item() return loss_q1, loss_q2, loss_pi, loss_alpha关键点pi_o.rsample()使用重参数化技巧使梯度可传回策略网络logp_a是动作概率密度的对数其均值即策略熵的负值self.target_entropy -act_dim是SAC默认目标熵loss_alpha驱动log_alpha自适应调整α。若entropy低于目标loss_alpha为正log_alpha增大α上升增强探索。3.3 DDPG actor-critic联合训练与target network软更新DDPG易因critic过估计导致训练崩溃必须严格实施target network和soft update# algorithm/ddpg.py def soft_update(self, local_model, target_model, tau): for target_param, local_param in zip(target_model.parameters(), local_model.parameters()): target_param.data.copy_(tau * local_param.data (1.0 - tau) * target_param.data) def train(self, data): o, a, r, o2, d data[obs], data[act], data[rew], data[obs2], data[done] # Critic loss: (Q(s,a) - (r γ * Q_target(s,a)))^2 q self.critic(o, a) with torch.no_grad(): a2 self.actor_target(o2) q_target self.critic_target(o2, a2) y r self.gamma * (1 - d) * q_target loss_critic ((q - y)**2).mean() # Actor loss: -Q(s, π(s)) a_pred self.actor(o) loss_actor -self.critic(o, a_pred).mean() # Update networks self.critic_opt.zero_grad() loss_critic.backward() self.critic_opt.step() self.actor_opt.zero_grad() loss_actor.backward() self.actor_opt.step() # Soft update target networks self.soft_update(self.actor, self.actor_target, self.tau) self.soft_update(self.critic, self.critic_target, self.tau)tau0.005是DDPG论文推荐值过大导致target network滞后不足过小则收敛慢。a2 self.actor_target(o2)确保target critic评估的是target actor的动作避免bootstrapping偏差。4. 训练流程与超参数配置如何让PPO在CartPole上500步收敛、SAC在HalfCheetah上稳定突破3000分4.1 统一训练入口与算法调度机制所有算法共享同一Trainer类通过algorithm_name参数动态加载对应训练逻辑# trainer.py class Trainer: def __init__(self, env_name, algorithm_name, **kwargs): self.env gym.make(env_name) self.algo { ppo: PPO, dqn: DQN, sac: SAC, ddpg: DDPG }[algorithm_name](self.env, **kwargs) self.buffer self.algo.buffer self.logger EpochLogger() def run_episode(self): obs, _ self.env.reset() ep_ret, ep_len 0, 0 while True: if hasattr(self.algo, select_action): act self.algo.select_action(obs) else: act self.algo.policy_forward(obs).sample().cpu().numpy() obs2, rew, done, truncated, _ self.env.step(act) ep_ret rew ep_len 1 self.buffer.store(obs, act, rew, obs2, done) obs obs2 if done or truncated: break return ep_ret, ep_len def train(self, epochs100, steps_per_epoch4000): for epoch in range(epochs): # Collect rollout data for t in range(steps_per_epoch): ep_ret, ep_len self.run_episode() self.logger.store(EpRetep_ret, EpLenep_len) # Train algorithm-specific iteration self.algo.train_iteration() # Log metrics self.logger.log_tabular(Epoch, epoch) self.logger.log_tabular(EpRet, with_min_and_maxTrue) self.logger.log_tabular(EpLen, average_onlyTrue) self.logger.dump_tabular()Trainer不关心算法内部只负责环境交互、数据收集和日志记录。self.algo.train_iteration()由各算法子类实现例如PPO调用update_policy()和update_value()SAC调用update_critic()、update_actor()、update_alpha()。4.2 关键超参数对照表不同算法在标准环境中的实测有效范围算法环境learning_ratebatch_sizegammatau (DDPG/SAC)clip_ratio (PPO)alpha (SAC)target_entropy (SAC)备注PPOCartPole-v13e-4640.99-0.2--steps_per_epoch2000,epochs50即可稳定DQNCartPole-v11e-31280.99----replay_size100000,epsilon_decay0.995SACHalfCheetah-v41e-32560.990.005-0.2-act_dim必须用Tanh激活log_std_init-3DDPGPendulum-v11e-31280.990.005---noise_scale0.1,noise_decay0.999注意gamma0.99是大多数任务的起点但LunarLander-v2建议用0.995以延长奖励衰减batch_size过小如32导致梯度噪声大过大如1024易过拟合单批数据PPO的clip_ratio0.2意味着新旧策略比率被限制在[0.8,1.2]超出则截断梯度。4.3 实战调试技巧三类高频报错的定位与修复方案错误1RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation原因在forward中使用了x y或x.relu_()等inplace操作。修复全部改用x x y和x F.relu(x)或在torch.autograd.set_detect_anomaly(True)下运行定位具体行。错误2ValueError: Expected input batch_size (64) to match target batch_size (32)原因obs和act维度不匹配常见于连续动作空间未reshape。修复检查act形状DQN输出[batch, act_dim]DDPG/SAC输出[batch, act_dim]PPO连续动作输出[batch, act_dim]但采样后需.unsqueeze(-1)适配critic输入。错误3训练初期reward剧烈震荡100轮后突然归零原因value网络初始化过大导致GAE优势全为负值策略拒绝所有动作。修复value网络最后一层bias0weight用orthogonal_初始化或在PPO中添加val_loss_coef0.5降低value loss权重。5. 部署与验证如何用TensorBoard监控PPO的KL散度、SAC的alpha变化、DQN的Q值分布5.1 TensorBoard日志集成不只是记录reward更要追踪算法健康度PyTorch原生支持torch.utils.tensorboard.SummaryWriter但需在训练循环中注入关键指标# 在PPO train_iteration()中 writer.add_scalar(Loss/Policy, loss_pi.item(), global_step) writer.add_scalar(Loss/Value, loss_v.item(), global_step) writer.add_scalar(Train/KL, approx_kl.item(), global_step) # KL between old new policy writer.add_scalar(Train/Entropy, ent.item(), global_step) writer.add_histogram(Policy/LogStd, self.actor.log_std, global_step) # 在SAC train_iteration()中 writer.add_scalar(Loss/Alpha, loss_alpha.item(), global_step) writer.add_scalar(Train/Alpha, self.alpha, global_step) writer.add_scalar(Train/Entropy, -logp_a.mean().item(), global_step) writer.add_histogram(Critic/Q1, q1, global_step)approx_kl是PPO中两次策略更新间的KL散度近似值应维持在0.01~0.03log_std直方图若出现大量-20或2说明熵调节失效Critic/Q1直方图若集中在极值如全为-100或100表明critic过估计或欠学习。5.2 模型保存与推理导出ONNX供C/Java服务调用训练完成的策略网络可导出为ONNX格式脱离Python环境部署# 导出PPO连续动作策略 dummy_input torch.randn(1, obs_dim) # batch1, obs_dim torch.onnx.export( self.actor, dummy_input, ppo_actor.onnx, input_names[obs], output_names[mu, std], dynamic_axes{obs: {0: batch}, mu: {0: batch}, std: {0: batch}}, opset_version11 ) # Python端加载ONNX推理 import onnxruntime as ort ort_session ort.InferenceSession(ppo_actor.onnx) outputs ort_session.run(None, {obs: obs_np.astype(np.float32)}) mu, std outputs[0], outputs[1] action np.random.normal(mu, std).clip(-1, 1) # 采样并裁剪ONNX导出要求模型无控制流如if、for因此MLPContinuousActor必须用torch.tanh而非np.clipdynamic_axes声明batch维度可变适配不同推理请求量。5.3 离线策略评估用固定seed重放rollout验证复现性强化学习结果波动大必须用固定seed保证实验可复现def evaluate_policy(env, actor, seed0, num_episodes10): env.reset(seedseed) torch.manual_seed(seed) np.random.seed(seed) returns [] for _ in range(num_episodes): obs, _ env.reset() ep_ret 0 while True: with torch.no_grad(): if hasattr(actor, rsample): # SAC/PPO continuous act actor(torch.as_tensor(obs, dtypetorch.float32)).mean.cpu().numpy() else: # DQN discrete logits actor(torch.as_tensor(obs, dtypetorch.float32)) act logits.argmax().item() obs, rew, done, truncated, _ env.step(act) ep_ret rew if done or truncated: break returns.append(ep_ret) return np.mean(returns), np.std(returns) # 调用 mean_ret, std_ret evaluate_policy(env, trainer.algo.actor, seed42) print(fEval return: {mean_ret:.2f} ± {std_ret:.2f})evaluate_policy中torch.no_grad()禁用梯度节省显存actor.mean取确定性动作SAC/PPOlogits.argmax()取最高概率动作DQN。seed42确保每次评估使用相同初始状态和随机序列num_episodes10提供足够置信区间。本文还有配套的精品资源点击获取
返回列表