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_{t=0}^\infty\gamma^t(R(s_t,a_t,s_{t+1})+\alpha H(\pi(\cdot|s_t)))π∗=aargmaxEτ∼πt=0∑∞γt(R(st,at,st+1)+α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)=i=1,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_v=unpack_batch_a2c(batch,val_net,gamma,device)# references for the critic networkmu_v=policy_net(states_v)act_dist=distr.Normal(mu_v,torch.exp(policy_net.logstd))acts_v=act_dist.sample()q1_v,q2_v=twinq_net(states_v,acts_v)# element-wise minimumref_vals_v=torch.min(q1_v,q2_v).squeeze()-\ ent_alpha*act_dist.log_prob(acts_v).sum(dim=1)returnstates_v,actions_v,ref_vals_v,ref_q_v函数的第一步使用已定义的unpack_batch_a2c()方法,该方法解包批次数据,将状态和动作转换为张量,并通过贝尔曼近似计算Q网络的参考值。完成这一步后,我们需要根据双Q值的最小值减去缩放后的熵系数来计算V网络的参考值。熵值通过当前策略网络计算得出。如前所述,我们的策略具有参数化的均值,但方差是全局的且不依赖于状态。
(2)在主训练循环中,我们使用先前定义的函数执行三个不同的优化步骤:分别针对V网络、Q网络和策略网络。
首先解包批次数据以获取张量及Q网络和V网络的目标值:
batch=buffer.sample(BATCH_SIZE)states_v,actions_v,ref_vals_v,ref_q_v=common.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_v=twinq_net(states_v,actions_v)q1_loss_v=F.mse_loss(q1_v.squeeze(),ref_q_v.detach())q2_loss_v=F.mse_loss(q2_v.squeeze(),ref_q_v.detach())q_loss_v=q1_loss_v+q2_loss_v q_loss_v.backward()twinq_opt.step()评论家网络也使用已计算的目标值通过简单的MSE目标进行优化:
crt_opt.zero_grad()val_v=crt_net(states_v)v_loss_v=F.mse_loss(val_v.squeeze(),ref_vals_v.detach())v_loss_v.backward()crt_opt.step()最后对演员网络进行优化:
act_opt.zero_grad()acts_v=act_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 Learning,RL)详解
PyTorch强化学习实战(2)——强化学习环境库Gymnasium
PyTorch强化学习实战(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步DQN(N-step DQN)
PyTorch强化学习实战(12)——Double DQN(DDQN)
PyTorch强化学习实战(13)——噪声网络(NoisyNet-DQN)
PyTorch强化学习实战(14)——优先经验回放机制
PyTorch强化学习实战(15)——Dueling DQN
PyTorch强化学习实战(16)——Categorical DQN
PyTorch强化学习实战(17)——强化学习训练加速
PyTorch强化学习实战(18)——基于DQN处理股票交易问题
PyTorch强化学习实战(19)——策略梯度法
PyTorch强化学习实战(20)——优势演员-评论家(Advantage Actor-Critic, A2C)
PyTorch强化学习实战(21)——异步优势演员-评论家(Asynchronous Advantage Actor-Critic, A3C)
PyTorch强化学习实战(22)——将强化学习应用于TextWorld互动小说游戏
PyTorch强化学习实战(23)——强化学习在网页导航中的应用
PyTorch强化学习实战(24)——连续动作空间中的强化学习
PyTorch强化学习实战(25)——深度确定性策略梯度(DDPG)
PyTorch强化学习实战(26)——提升随机策略梯度稳定性