PyTorch强化学习实战——软演员-评论家SAC算法详解与实现0. 前言1. SAC 算法2. 算法实现3. 运行结果相关链接0. 前言在本节中我们将介绍软演员-评论家 (Soft Actor-Critic,SAC) 方法测试 HalfCheetah 环境该方法于2018年由Haarnoja等人发布的论文《Soft Actor-Critic: Off-policy Maximum Entropy Deep Reinforcement Learning》中提出。目前SAC被视为连续控制问题的最佳方法之一并得到广泛应用。其核心思想更接近深度确定性策略梯度 (Deep Deterministic Policy Gradient, DDPG) 方法而非优势演员-评论家 (Advantage Actor-Critic, A2C) 方法策略梯度。我们将直接将其与长期被视为连续控制问题标准方案的近端策略优化 (Proximal Policy Optimization, PPO) 性能进行对比。1. SAC 算法软演员-评论家 (Soft Actor-Critic,SAC) 方法的核心思想是熵正则化在每个时间步添加与策略熵成正比的额外奖励。用数学符号表示我们寻找的策略如下π ∗ a r g m a x a E τ ∼ π ∑ t 0 ∞ γ t ( R ( s t , a t , s t 1 ) α H ( π ( ⋅ ∣ s t ) ) ) \pi^*\underset {a} {argmax}\mathbb E_{\tau\sim\pi}\sum_{t0}^\infty\gamma^t(R(s_t,a_t,s_{t1})\alpha H(\pi(\cdot|s_t)))π∗aargmaxEτ∼πt0∑∞γt(R(st,at,st1)αH(π(⋅∣st)))其中H ( P ) x ∼ P [ − l o g P ( x ) ] H(P)_{x∼P}[-logP(x)]H(P)Ex∼P[−logP(x)]是分布P PP的熵。换言之当智能体处于熵最大化的状态时给予额外奖励。此外SAC采用了裁剪双Q技巧除了价值函数外我们训练两个预测Q值的网络并取两者最小值进行贝尔曼近似。研究人员表示这有助于解决训练过程中的Q值高估问题。总共需要训练四个网络策略网络π ( s ) π(s)π(s)、价值网络V ( s , a ) V(s,a)V(s,a)以及两个Q网络Q 1 ( s , a ) Q_1(s,a)Q1(s,a)和Q 2 ( s , a ) Q_2(s,a)Q2(s,a)。价值网络V ( s , a ) V(s,a)V(s,a)使用目标网络。SAC训练流程如下Q 网络使用均方误差 (Mean Squared Error,MSE) 目标进行训练通过使用目标值网络进行贝尔曼近似y q ( r , s ′ ) r γ V t g t ( s ′ ) y_q ( r,s ′ ) r γV_{tgt} (s′)yq(r,s′)rγVtgt(s′)(对于非终止步骤)V 网络使用MSE目标进行训练目标为y v ( s ) m i n i 1 , 2 Q i ( s , a ~ ) − α l o g π ( a ~ ∣ s ) y_v ( s ) \underset {i 1 , 2}{min} Q i ( s, \tilde a ) − α log π_ (\tilde a | s )yv(s)i1,2minQi(s,a~)−αlogπ(a~∣s)其中a ~ \tilde aa~是从策略π ( ⋅ ∣ s ) π_ ( ⋅| s )π(⋅∣s)中采样的动作策略网络π θ π_θπθ采用DDPG风格训练通过最大化以下目标Q 1 ( s , a ~ ( s ) ) − α l o g π ( a ~ ( s ) ∣ s ) Q_1(s, \tilde a_ (s)) − αlog π_ (\tilde a_ ( s ) | s )Q1(s,a~(s))−αlogπ(a~(s)∣s)其中a ~ ( s ) \tilde a_(s)a~(s)是从π ( ⋅ ∣ s ) π_ ( ⋅| s )π(⋅∣s)中采样的动作2. 算法实现SAC方法的实现位于 train_sac.py 中。模型由以下网络组成定义在model.py中ModelActor与置信域策略优化 (Trust Region Policy Optimization, TRPO) 一节使用的策略网络相同。由于策略方差并非由状态参数化(logstd字段不是网络而只是张量)训练目标并未完全符合SAC规范。一方面这可能影响收敛性和性能因为SAC方法的核心思想——熵正则化——需要参数化方差才能实现另一方面这减少了模型参数量。我们可以扩展该示例实现策略的参数化方差从而构建完整的SAC方法ModelCritic与 TRPO 一节中的价值网络相同ModelSACTwinQ这两个网络以状态和动作为输入预测Q值(1)首个实现该方法的函数是unpack_batch_sac()定义于common.py中。其目标是获取轨迹步长的批次数据并计算V网络和双Q网络的目标值torch.no_grad()defunpack_batch_sac(batch:tt.List[lib.experience.ExperienceFirstLast],val_net:model.ModelCritic,twinq_net:model.ModelSACTwinQ,policy_net:model.ModelActor,gamma:float,ent_alpha:float,device:torch.device):states_v,actions_v,ref_q_vunpack_batch_a2c(batch,val_net,gamma,device)# references for the critic networkmu_vpolicy_net(states_v)act_distdistr.Normal(mu_v,torch.exp(policy_net.logstd))acts_vact_dist.sample()q1_v,q2_vtwinq_net(states_v,acts_v)# element-wise minimumref_vals_vtorch.min(q1_v,q2_v).squeeze()-\ ent_alpha*act_dist.log_prob(acts_v).sum(dim1)returnstates_v,actions_v,ref_vals_v,ref_q_v函数的第一步使用已定义的unpack_batch_a2c()方法该方法解包批次数据将状态和动作转换为张量并通过贝尔曼近似计算Q网络的参考值。完成这一步后我们需要根据双Q值的最小值减去缩放后的熵系数来计算V网络的参考值。熵值通过当前策略网络计算得出。如前所述我们的策略具有参数化的均值但方差是全局的且不依赖于状态。(2)在主训练循环中我们使用先前定义的函数执行三个不同的优化步骤分别针对V网络、Q网络和策略网络。首先解包批次数据以获取张量及Q网络和V网络的目标值batchbuffer.sample(BATCH_SIZE)states_v,actions_v,ref_vals_v,ref_q_vcommon.unpack_batch_sac(batch,tgt_crt_net.target_model,twinq_net,act_net,GAMMA,SAC_ENTROPY_ALPHA,device)双Q网络使用相同的目标值进行优化twinq_opt.zero_grad()q1_v,q2_vtwinq_net(states_v,actions_v)q1_loss_vF.mse_loss(q1_v.squeeze(),ref_q_v.detach())q2_loss_vF.mse_loss(q2_v.squeeze(),ref_q_v.detach())q_loss_vq1_loss_vq2_loss_v q_loss_v.backward()twinq_opt.step()评论家网络也使用已计算的目标值通过简单的MSE目标进行优化crt_opt.zero_grad()val_vcrt_net(states_v)v_loss_vF.mse_loss(val_v.squeeze(),ref_vals_v.detach())v_loss_v.backward()crt_opt.step()最后对演员网络进行优化act_opt.zero_grad()acts_vact_net(states_v)q_out_v,_twinq_net(states_v,acts_v)act_loss-q_out_v.mean()act_loss.backward()act_opt.step()与之前给出的公式相比代码中缺失了熵正则化项这更接近 DDPG 的训练方式。由于我们的方差不依赖于状态因此可以从优化目标中省略该项。3. 运行结果在HalfCheetah和Ant环境中运行了SAC训练耗时9-13小时处理了500万次观测。结果存在一定矛盾性一方面SAC的样本效率和奖励增长动态优于近端策略优化 (Proximal Policy Optimization, PPO) 方法。例如在HalfCheetah环境中SAC仅用50万次观测就获得900分奖励而PPO需要超过100万次观测才能达到相同策略水平。在MuJoCo环境SAC更是取得了7063分的策略表现展现了先进性能。但另一方面由于SAC的异策略特性其训练速度明显更慢——相比同策略方法需要执行更多计算。HalfCheetah环境的500万帧训练耗时10小时。作为对比A2C在相同时间内可处理5000万次观测。这展示了同策略与异策略方法之间的权衡如果环境运行速度快且观测获取成本低像PPO这样的同策略方法可能是最佳选择但如果观测获取困难离策略方法表现更优不过需要更多的计算量。下图显示了HalfCheetah环境的奖励动态在Ant环境中的结果则差强人意——根据得分显示学习到的策略几乎无法稳定站立。PyBullet环境的结果如下图所示MuJoCo环境的结果如下图所示可以使用play.py工具对保存的模型进行基准测试并录制学习策略的运行视频。相关链接PyTorch强化学习实战1——强化学习Reinforcement LearningRL详解PyTorch强化学习实战2——强化学习环境库GymnasiumPyTorch强化学习实战3——Gymnasium API扩展功能PyTorch强化学习实战4——PyTorch基础PyTorch强化学习实战5——PyTorch Ignite 事件驱动机制与实践PyTorch强化学习实战6——交叉熵方法详解与实现PyTorch强化学习实战7——表格学习与贝尔曼方程PyTorch强化学习实战8——Q学习详解与实现PyTorch强化学习实战9——深度Q学习PyTorch强化学习实战10——强化学习高级组件PyTorch强化学习实战11——N步DQNN-step DQNPyTorch强化学习实战12——Double DQNDDQNPyTorch强化学习实战13——噪声网络NoisyNet-DQNPyTorch强化学习实战14——优先经验回放机制PyTorch强化学习实战15——Dueling DQNPyTorch强化学习实战16——Categorical DQNPyTorch强化学习实战17——强化学习训练加速PyTorch强化学习实战18——基于DQN处理股票交易问题PyTorch强化学习实战19——策略梯度法PyTorch强化学习实战20——优势演员-评论家Advantage Actor-Critic, A2CPyTorch强化学习实战21——异步优势演员-评论家Asynchronous Advantage Actor-Critic, A3CPyTorch强化学习实战22——将强化学习应用于TextWorld互动小说游戏PyTorch强化学习实战23——强化学习在网页导航中的应用PyTorch强化学习实战24——连续动作空间中的强化学习PyTorch强化学习实战25——深度确定性策略梯度DDPGPyTorch强化学习实战26——提升随机策略梯度稳定性