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

资讯详情

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

SAC强化学习实战:最大熵、双Q与重参数化连续控制解析

SAC强化学习实战:最大熵、双Q与重参数化连续控制解析 1. SAC到底在优化什么把熵写进回报之后SACSoft Actor-Critic是我在强化学习里反复实现过最多的算法之一第一遍照着论文抄第二遍自己重写网络第三遍才开始琢磨它为什么这么设计。这篇笔记不讲某某论文提出了某某方法那种套话只讲我踩过的坑和真正理解之后觉得精妙的地方。如果你已经看过DQN、DDPG、TD3想找一个在连续控制任务里稳定、样本效率又不错的算法SAC基本是绕不过去的一站。它也常出现在python强化学习实战、机械臂强化学习实战这类场景里因为连续动作空间的机械臂控制恰好是SAC的主场。先说清楚SAC解决的问题域连续动作空间下的model-free强化学习。这里的state value function、Q函数、策略网络三者的关系是理解SAC的入口。很多人一上来就去抄代码结果Q值震荡、alpha爆炸最后调不出来就放弃了根子在于没搞懂它的目标函数和别人不一样。所以这一节先把目标函数讲透后面的网络结构、损失函数才有落脚点。1.1 标准回报和最大熵目标的差别在哪普通强化学习的目标是最大化期望折扣累积回报J(π) E[ Σ γ^t r(s_t, a_t) ]而SAC在这条公式里硬塞进了一项——熵J(π) E[ Σ γ^t ( r(s_t, a_t) α · H(π(·|s_t)) ) ]其中H(π(·|s_t))是策略在状态s_t下的熵α是温度系数用来平衡拿奖励和保持随机这两件事的比例。翻译成人话SAC不满足于只拿到高回报它还要求策略在每个状态下尽量别太自信只要动作不会明显拉低回报就保留一定的随机性。我第一次看这个公式时的反应是凭什么加熵加了熵回报不就变低了吗实际跑下来发现正是这一项让SAC在早期的探索阶段比DDPG稳得多。DDPG的策略是确定性的前期一旦陷入局部最优几乎跳不出来SAC因为始终保留一份随机性探索是内建的不需要额外往动作上加高斯噪声。1.2 熵为什么能让探索和利用不再互相打架传统做法里探索和利用是分开的训练时往动作上加噪声比如DDPG的OU噪声、TD3的动作噪声评估时把噪声关掉。噪声大小是个玄学超参加多了不收敛加少了不探索。SAC把这件事变成了同一个目标里的两项。策略越随机熵越大α·H那项奖励越高但越随机回报项通常越低。于是策略网络在梯度下降时会自动去找那个回报尽量高、又不至于过早收缩成确定性的平衡点。这就像一个人做事既要把事情办成回报又不想把自己逼到走投无路熵留一点余地。手上保留几个备选方案长期看反而更抗风险。注意SAC的随机性不是加在动作输出上的噪声而是策略本身的分布。这是它和DDPG最本质的区别评估时SAC可以取分布的均值也可以继续采样两种做法都可以但一般取均值更稳定。1.3 把SAC、DDPG、TD3摆在一起看我整理过一张对比表每次复习都翻出来看维度DDPGTD3SAC策略类型确定性确定性随机高斯Q网络数量12取min2取min探索方式外部加噪声外部加噪声内建熵温度系数无无有可自动调目标策略平滑无有无靠熵代替训练稳定性一般较好好TD3用双Q取min和目标策略平滑来解决Q值高估SAC同样用了双Q取min但它不需要目标策略平滑因为它天然就是随机的策略平滑那一手被熵项覆盖了。这就是为什么我说这几个算法是递进关系理解TD3之后再学SAC会顺很多如果直接跳到SAC很容易把双Q、目标网络、重参数化这几个概念搅在一起。2. 三张网络加一个温度系数SAC的骨架拆解SAC的网络结构看起来比TD3复杂但拆开就四块东西两个Q网络、一个策略网络Actor、一个温度系数α。如果按代码里的optimizer数量数通常是三个优化器一个更新两个Q可以合成一个一个更新策略一个更新α。搞清楚每一块负责什么、梯度从哪儿来就不会在写loss的时候手忙脚乱。2.1 双Q网络和clipped double Q的取舍SAC里有两个结构相同、初始化不同的Q网络记作Q1和Q2。计算目标值时取两者最小值y r γ · ( min(Q1(s, a), Q2(s, a)) - α · log π(a|s) )这里Q1、Q2是目标网络。取min的原因是抑制Q值高估——单个Q网络在bootstrapping时容易把估计值越推越高因为max操作会系统性地偏向被高估的动作。这个问题DQN时代就有Double DQN靠解耦动作选择和价值评估来缓解SAC用的是更粗暴也更稳的办法两个网络取小的那个当目标。训练时两个Q网络各自算一份MSE损失但都用同一个目标值y实现上是with torch.no_grad(): next_action, next_logp actor.sample(next_state) q1_next critic1_target(next_state, next_action) q2_next critic2_target(next_state, next_action) q_next torch.min(q1_next, q2_next) - alpha * next_logp target_q reward gamma * (1 - done) * q_next q1_loss F.mse_loss(critic1(state, action), target_q) q2_loss F.mse_loss(critic2(state, action), target_q)提示(1 - done)这一项别漏。episode结束时如果没有正确置零bootstrapping会把终止状态的Q值错误地传回来训练后期会出现莫名其妙的Q值膨胀。2.2 策略网络为什么要输出均值和标准差SAC的策略网络不是直接输出一个动作而是输出一个高斯分布的均值μ和对数标准差log σ然后从这个分布里采样动作。这么设计的原因是要计算log π(a|s)而熵项-α·log π(a|s)必须有明确的概率密度才写得出来。确定性策略没有密度自然也就没有熵。网络结构上通常共享一个两层MLP隐藏层256维ReLU激活然后分两个头一个输出μ一个输出log σ。log σ一般会做截断比如夹到[-20, 2]之间防止标准差过大或过小导致数值问题。采样时用重参数化ε ~ N(0, I) a_raw μ σ · ε a tanh(a_raw)先采样标准正态再线性变换最后用tanh压到动作范围[-1, 1]。这三步每一步都有它的道理第3节展开讲。2.3 目标网络和软更新里的tau和DQN、DDPG一样SAC也有目标网络但更新方式是软更新Polyak平均θ ← τ · θ (1 - τ) · θτ通常取0.005。意思是每次更新目标网络只朝当前网络挪0.5%靠大量步数慢慢逼近。这个值看起来小得离谱但实测就是这么用的——大了目标值不稳定小了学习太慢。我试过把τ调到0.05前几百步看着还行一旦Q值开始增长目标网络跟得太快直接震荡。软更新的写法有两种一种是手动做参数插值另一种是循环遍历state_dictfor param, target_param in zip(critic.parameters(), critic_target.parameters()): target_param.data.copy_(tau * param.data (1 - tau) * target_param.data)注意目标网络的参数更新必须放在torch.no_grad()上下文里或者用.data否则会把目标网络也拉进计算图显存和时间都会炸。3. 重参数化技巧策略网络的梯度到底从哪来这是SAC里我卡最久的地方。策略网络的损失里含log π(a|s)而a又是从π里采样出来的采样这个操作本身不可导——你没法对随机采样求导。SAC用重参数化reparameterization绕过了这个问题把随机性从采样操作里剥离出去。3.1 采样不可导的问题怎么绕假设策略输出随机动作a ~ N(μ_θ, σ_θ)我们想通过a反向传播更新θ。直接在采样节点上求导是不行的因为sampling是随机操作没有梯度。重参数化的思路是把随机性挪到一个外部的、与参数无关的噪声ε上让动作变成θ的确定性函数a μ_θ σ_θ · ε, ε ~ N(0, I)这样一来a对θ的依赖变成了显式的加法和乘法梯度可以通过μ_θ和σ_θ顺利回传。梯度路径上ε是常数不影响求导。这就是重参数化名字的由来——重新参数化了随机变量的生成方式。3.2 策略损失逐项拆解SAC的策略损失是action, log_prob actor.sample(state) q1 critic1(state, action) q2 critic2(state, action) q torch.min(q1, q2) actor_loss (alpha * log_prob - q).mean()拆开看每一项q策略选出的动作能拿到的价值我们希望它越大越好所以最小化-q。alpha * log_prob熵项的另一种写法。log_prob越小概率密度越低、越不自信-log_prob越大熵越大。这里 alpha * log_prob是把熵奖励折算进损失等价于最大化q - alpha * log_prob。提示不同实现里正负号写得不一样有的写alpha * log_prob - q有的写-q alpha * log_prob本质相同。抄代码时一定要自己对一遍符号符号写反了训练也能跑但策略会往收缩成确定性、甚至反向的方向走。3.3 tanh压缩带来的log_prob修正动作范围通常限制在[-1, 1]所以要对高斯样本做tanh。但tanh是个非线性变换会改变概率密度。如果直接拿变换前高斯分布的log_prob去算熵就错了。正确的做法是按变量替换公式修正log π(a|s) log μ(a_raw|s) - Σ log(1 - tanh²(a_raw))代码里通常这么写log_prob dist.log_prob(action_raw) log_prob - torch.log(1 - action.pow(2) 1e-6) log_prob log_prob.sum(dim-1, keepdimTrue)那个1e-6是防止tanh接近 ±1 时1 - a²变成0导致log(0)爆炸。这个细节在很多教程里被忽略了结果就是训练到后期偶尔出现inf或nan。我第一次复现SAC时就是卡在这里log_prob时不时冒nan查了两天才定位到是tanh饱和区的数值问题。4. 温度系数自动调节alpha不是超参是被学出来的早期的SAC论文里α是个固定超参需要针对不同任务手动调。后来作者发现这个值对结果影响太大改成了自动调节。现在大家说的SAC基本都是带自动温度调节的版本。4.1 固定alpha的痛点α大策略偏随机探索充分但回报上不去α小策略偏确定初期可能收敛到局部最优。而且同一个α在HalfCheetah上合适换到机械臂抓取任务上可能完全不对。因为不同任务的奖励尺度差好几个数量级奖励大的任务熵项相对就小奖励小的任务熵项就压过回报。手动调α有点像同时拧温度和水量的花洒调好一个另一个又不对。我最早跑一个自定义的机械臂环境奖励设计成距离的负值量级很小固定α时策略几乎全程在乱晃因为熵奖励压过了回报。改成自动调节之后才正常。4.2 自动熵调节的机制自动调节的核心思路是设定一个目标熵H_target让算法自己调整α使策略的实际熵尽量贴近目标。目标熵一般取动作维度的负数H_target -dim(A)。比如动作是6维的机械臂目标熵就是-6。α的损失函数是alpha_loss -(log_alpha * (log_prob target_entropy).detach()).mean()实现上log_alpha才是可学习参数保证α恒正优化器更新的是log_alpha用的时候取指数alpha log_alpha.exp()如果实际熵比目标熵大-log_prob target_entropyloss会推动log_alpha变小减小熵的权重反之则变大。自动调节让α跟着任务走省掉了大量调参。4.3 我怎么观测alpha的变化训练时我习惯把alpha、平均熵、平均Q值、policy loss都打出来。正常收敛的曲线大概是alpha从小逐渐下降并稳定熵从大到小并稳定在目标附近Q值平滑上升。如果alpha一路飙到几百通常是奖励尺度太大或者Q值发散如果alpha掉到接近0说明策略已经收缩成确定性了探索不足——这时候得回头检查熵项符号是不是写反了。观测项正常范围异常表现与可能原因alpha0.01 ~ 1 逐步稳定持续上升奖励尺度过大持续归零熵项符号反了平均熵逐渐逼近 -dim(A)远大于目标策略太散远小于策略太死Q值平滑上升后趋稳持续下降学习率过大或目标网络同步过快policy loss小幅波动发散梯度裁剪没做或α过大5. 从零实现一个可以跑起来的SAC光看公式容易浮在表面我把关键模块的代码骨架列出来。环境用的是标准Gym风格接口PyTorch 1.x以上都能跑。动手写一遍之后前面几节的概念会清晰很多。5.1 网络结构和初始化策略网络和Q网络都是两层MLP隐藏层256维class GaussianPolicy(nn.Module): def __init__(self, state_dim, action_dim, hidden256): super().__init__() self.net nn.Sequential( nn.Linear(state_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), ) self.mu nn.Linear(hidden, action_dim) self.log_std nn.Linear(hidden, action_dim) self.log_std_min, self.log_std_max -20, 2 def forward(self, state): h self.net(state) mu self.mu(h) log_std self.log_std(h).clamp(self.log_std_min, self.log_std_max) return mu, log_stdQ网络接受(state, action)拼接后的输入输出一个标量class QNet(nn.Module): def __init__(self, state_dim, action_dim, hidden256): super().__init__() self.net nn.Sequential( nn.Linear(state_dim action_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, 1), ) def forward(self, state, action): return self.net(torch.cat([state, action], dim-1))提示策略网络最后一层self.mu和self.log_std的初始化权重和偏置建议调小比如用较小的方差初始化否则训练一开始mu偏大tanh直接饱和梯度几乎传不回去。5.2 回放缓冲区和采样回放缓冲区就是个大数组存(s, a, r, s, done)容量一般1e6。采样用均匀随机采样一次256条class ReplayBuffer: def __init__(self, capacity, state_dim, action_dim): self.s np.zeros((capacity, state_dim), dtypenp.float32) self.a np.zeros((capacity, action_dim), dtypenp.float32) self.r np.zeros((capacity, 1), dtypenp.float32) self.s_ np.zeros((capacity, state_dim), dtypenp.float32) self.d np.zeros((capacity, 1), dtypenp.float32) self.ptr, self.size, self.capacity 0, 0, capacity存储用numpy的预分配数组而不是list因为list会随着训练不断增长到后期内存占用量级完全不同。这个细节在跑长任务时很关键我用list写过一版跑了几十万步之后内存直接吃满。5.3 训练循环的关键顺序每一步的顺序不能乱我按经验总结成这个流程用当前策略选动作加不加探索噪声都行熵已经内建了存进buffer。从buffer采一个batch。算目标Q值更新两个Q网络。用当前Q值更新策略网络。更新温度系数α。软更新目标网络。核心超参我列一张表都是实践里验证过比较稳的超参取值说明学习率3e-4Q、策略、α可以用同一个batch size256显存不够可降到128γ0.99长周期任务可用0.995τ0.005软更新系数buffer1e6至少是batch的几百倍隐藏层256×2复杂任务可加到512初始alpha0.2自动调节后会自动变目标熵-dim(A)动作维度取负5.4 训练不收敛时的排查顺序我遇到过大概七八种SAC训崩的情况总结下来按这个顺序查最省时间先看Q值曲线。如果Q值爆炸式增长八成是done没处理对或者γ设太大。再看alpha。如果alpha持续上涨到几百检查奖励尺度SAC对奖励量级敏感最好把奖励归一化到[-1, 1]附近。然后看熵项符号。actor_loss alpha * log_prob - q如果写成了alpha * (-log_prob) - q策略会被推向最大化log_prob即收缩成确定性alpha会一路下跌。最后看目标网络有没有更新。漏更新目标网络的话Q值会停滞甚至缓慢衰减看起来像学习率太小。
返回列表