新闻详情

新闻详情

首页 / 资讯中心 / 详情

变分推断如何重塑深度强化学习?拆解SAC、MPO、AWR的底层逻辑

发布时间:2026/9/28 1:30:39来源:尧图网络
变分推断如何重塑深度强化学习?拆解SAC、MPO、AWR的底层逻辑
伯克利 2026 春季深度强化学习课程的第 12 讲专门讲了“强化学习中的变分推断”。很多人看到这个标题的第一反应是变分推断不是概率图模型里的内容吗和强化学习有什么关系但恰恰是这个看似“数学复习课”的章节解释了 SAC、MPO、AWR 这些现代深度强化学习算法的真正底层逻辑。如果你已经看完了策略梯度和 DQN却突然读不懂这些“进阶算法”的源码问题几乎都出在同一个地方你还在用“最大化奖励的梯度上升”去理解它们而这些算法本质上做的是“概率分布之间的匹配”。本文会把这讲的核心脉络重新梳理一遍从 ELBO 的直觉讲起一直落到可以运行的 PyTorch 代码让你以后再看到某个 RL 算法的 loss 函数能一眼认出它是哪一种变分推断的特例。1. 这篇文章真正要解决的问题1.1 从“梯度上升”到“分布匹配”的范式转换经典强化学习给人的第一印象是策略网络输出的动作要让累积奖励最大。这个描述没有错但它遮蔽了现代深度强化学习的一个重要变化——策略不再只是“一个输出动作的函数”而是一个“动作的概率分布”。一旦策略变成分布优化问题就跟着变了。我们不再只关心“哪个动作的奖励高”还要关心“新策略和旧策略到底差了多少”。这正是 KL 散度、变分推断这些东西进入强化学习的原因。它们不是额外添加的数学装饰而是把“策略优化”重新表达成“概率分布逼近”之后的自然产物。1.2 你会从这篇文章里得到什么本文将围绕第 12 讲的主线回答下面几个具体问题变分推断究竟在推断什么为什么精确推断在深度强化学习里不可行ELBO 这个目标函数是怎么来的它在强化学习里对应什么直观含义为什么 SAC 的熵正则、MPO 的期望最大化、AWR 的加权回归本质上都是变分推断的某种形式如何用 PyTorch 写一个最小可运行的“变分策略更新”示例以及训练时最常踩的坑。读完这篇文章你对“强化学习 变分推断”的认知会从“听过名词”变成“能写出代码、能读懂源码”也会更清楚当前深度强化学习算法设计背后的数学动机。2. 强化学习中的变分推断先建立直觉2.1 什么叫“推断”为什么精确推断做不到在强化学习里“推断”这个词最常见的出现位置是隐变量模型和贝叶斯方法。它可以概括为已知一些观测数据想去计算某些未知变量的后验分布。举个最简单的例子。假设你在一个迷宫里探索观察到一系列状态和奖励你想知道“当前状态的未知环境特征 z 是什么”。根据贝叶斯公式后验分布是p(z | obs) p(obs | z) * p(z) / p(obs)问题出在分母 p(obs)。为了计算它需要对所有可能的 z 做积分。在连续空间、高维神经网络特征面前这个积分基本没有解析解数值积分也不可行。这就是精确推断的极限。变分推断的思路非常直接既然 p(z | obs) 算不出来那就找一个形式简单、可控的分布 q(z) 去逼近它。优化目标不再是“算出后验”而是“让 q(z) 尽量接近 p(z | obs)”。2.2 变分推断的核心找一个好的近似分布 q“接近”用什么度量最常用的就是 KL 散度。我们想做的是min KL(q(z) || p(z | obs))但这个目标里仍然含有 p(obs)还是不好优化。于是经过几步推导可以把它转成一个可以最大化的下界——ELBO。ELBO 的直观理解是它由两项组成。第一项是“重构项”让 q(z) 下的样本能更好地解释观测数据第二项是“先验项”让 q(z) 不要偏离先验 p(z) 太远。两者之间形成一种平衡既要拟合数据又不能过度自信。2.3 ELBO把“求后验”变成“优化目标”ELBO 的常见形式是ELBO(q) E_{z ~ q}[ log p(obs | z) ] - KL(q(z) || p(z))第一项鼓励近似后验分布 q 能产生高似然的观测第二项作为一个正则项把 q 拉回先验。在深度学习中这两项通常都用神经网络表示可以端到端优化。你可能已经注意到这个结构和很多强化学习算法的损失函数很像一个“任务目标项”加一个“分布约束项”。这不是巧合而是因为策略优化也可以写成类似的变分目标。2.4 重参数化技巧让采样变成可微在深度学习中优化 ELBO必须对期望项求梯度。但“从分布 q 中采样”这个操作本身不可导。解决办法是重参数化技巧把随机性从采样中剥离出来改写成“确定性函数 外部噪声”。对于高斯分布可以把 z 写成z mean std * epsilon, epsilon ~ N(0, I)这样梯度可以穿过 mean 和 std直接反向传播到策略网络的参数上。后面代码示例里会实际用到这一步。下面是三种推断方法的对比方法核心思路适用场景在深度 RL 中的常见形态精确推断直接计算后验公式低维、共轭分布几乎不出现MCMC 采样构造马尔可夫链渐进逼近后验小规模、离线分析采样效率低不适合在线训练变分推断优化一个参数化分布逼近后验大规模、深度模型VAE 式世界模型、策略分布约束这一节的小结论是变分推断并没有真的“算出”后验而是把推断问题变成了优化问题。这个思想一旦建立再看强化学习里的很多算法它们的 loss 就不再是黑盒。3. 变分推断在深度强化学习中的三个落地场景3.1 场景一策略优化等于在策略空间上做投影先看策略优化这个最直观的场景。经典的策略梯度方法直接对奖励期望求梯度它的问题在于步长不好控制更新太大策略崩溃更新太小学习太慢。TRPO 用 KL 散度约束每一步的更新范围本质上就是在做一种“带约束的投影”。如果把当前策略看成先验把优势函数带来的奖励信号看成“我们希望策略往哪个方向调整”那一次策略更新就是在旧策略附近找一个新策略让新策略在优势函数上有更好的表现同时不能离旧策略太远。写成目标函数就是max E_{a ~ π_new}[ A(s, a) ] - beta * KL(π_new || π_old)这个形式和 ELBO 高度一致第一项是“任务项”第二项是“分布正则项”。SAC 的熵正则也可以从另一个角度理解最大熵目标里的熵项其实是在鼓励策略不要过早坍缩成一个确定性函数这和变分推断里“不要偏离先验太远”的哲学是相通的。3.2 场景二MPO 与 AWR 的期望最大化视角如果说 TRPO 只是“用了 KL 约束”那么 MPO 和 AWR 就是把变分推断的框架正式引入策略优化的代表。这类算法的核心思想是期望最大化EM。在 E 步我们根据当前策略和优势函数构造一个“改进后的目标分布”——本质上是把奖励高的动作赋予更高的概率。这一步通常需要解一个带 KL 约束的优化问题也就是一个变分推断步骤。在 M 步我们再用监督学习的方式去拟合这个目标分布。这种设计带来的工程优势非常明显策略更新变成了一个“加权最大似然”问题训练稳定性比直接做策略梯度更好。AWR 进一步简化直接对优势函数做指数加权得到如下形式的损失L(θ) - E[ exp(adv / temperature) * log π_θ(a | s) ]你仔细看这个式子exp(adv / temperature) 就是根据优势重新加权后的“目标分布”log π_θ 是让当前策略去拟合这个分布。它不显式写 KL但本质上是变分推断里“用加权样本拟合一个隐式目标分布”的思路。这也是离线强化学习里这类算法流行的原因用旧数据也能训练因为更新规则变成了分布拟合而不是依赖在线交互。3.3 场景三世界模型与隐变量探索第三种场景更接近传统变分推断用 VAE 结构学习环境的世界模型。Dreamer 这类基于模型的强化学习方法会先训练一个变分自编码器把高维观测压缩成隐变量 z再在隐空间里学习策略。这里的变分推断作用是压缩表示而不是直接优化策略。另一个应用是探索。如果智能体对某个状态下的“不确定性”很敏感它就会倾向去探索那些预测结果不可靠的区域。很多探索奖励的设计思路就是用变分后验和先验之间的 KL 散度来衡量信息增益。这个 KL 越大说明智能体在这个状态学到的新东西越多探索价值也越高。这三个场景合起来构成了一条完整的逻辑链变分推断既可以作为策略更新中的“约束器”也可以作为离线数据的“分布拟合器”还可以作为世界模型里的“表示学习器”。4. 环境准备与实验框架理解概念之后最好亲手跑一个最小例子。这里不追求复现完整的 SAC 或 MPO而是把“变分策略更新”的最小闭环单独拆出来用 PyTorch 实现一遍。4.1 运行环境本文示例代码使用 Python 3.9 和 PyTorch 2.x读者以实际安装版本为准。建议使用虚拟环境隔离依赖python -m venv rl_var_env source rl_var_env/bin/activate # Windows 下使用 rl_var_env\Scripts\activate pip install torch numpy4.2 实验目录结构我们用一个极简的“离线策略更新”实验来说明问题目录结构如下./variational_rl_demo/ ├── gaussian_policy.py # 高斯策略网络 ├── elbo_objective.py # 变分推断目标 └── train_demo.py # 训练主脚本这里不引入完整 Gym 环境而是直接用一组离线交互数据模拟“带优势函数的样本”。这样做的原因是课程第 12 讲的重点是变分推断与策略更新的关系过早引入环境交互会分散注意力。先跑通最小闭环再迁移到真正的 RL 环境会更稳。4.3 实验任务设计假设我们已经有一个旧策略产生了一批样本每个样本包含状态 obs形状为 (batch_size, obs_dim)动作 act形状为 (batch_size, act_dim)优势值 adv标量列表表示该动作相对平均水平好多少旧策略对数概率 old_log_prob用于计算重要性采样修正和 KL。我们要做的是用变分推断式的目标更新一个新策略让它在不偏离旧策略太远的前提下提升优势值。5. 完整代码实现一个可运行的变分策略更新示例5.1 高斯策略网络策略被建模为“以状态为条件的多维高斯分布”输出均值和标准差既方便采样也能直接计算对数概率。# 文件路径variational_rl_demo/gaussian_policy.py import torch import torch.nn as nn from torch.distributions import Normal class GaussianPolicy(nn.Module): def __init__(self, obs_dim, act_dim, hidden_dim64): super().__init__() self.net nn.Sequential( nn.Linear(obs_dim, hidden_dim), nn.Tanh(), nn.Linear(hidden_dim, hidden_dim), nn.Tanh(), ) self.mean_head nn.Linear(hidden_dim, act_dim) self.log_std_head nn.Linear(hidden_dim, act_dim) def forward(self, obs): h self.net(obs) mean self.mean_head(h) # clamp 防止 log_std 过大或过小导致数值异常 log_std self.log_std_head(h).clamp(-20, 2) scale log_std.exp() return Normal(mean, scale) def sample_with_log_prob(self, obs): dist self.forward(obs) action dist.rsample() # 重参数化采样可导 log_prob dist.log_prob(action).sum(dim-1) return action, log_prob, dist关键点是dist.rsample()。它利用重参数化技巧让采样结果对策略参数可导这样后面的梯度才能正常传播。计算对数概率时对动作维度求和因为动作是向量。5.2 ELBO 形式的策略更新目标下面实现两个变体一个显式包含 KL 项一个使用 AWR 式加权最大似然。两者在精神上是同源的但代码结构不同。# 文件路径variational_rl_demo/elbo_objective.py import torch def kl_regularized_policy_loss( policy, old_policy, obs, adv, kl_weight0.1, ): 变分推断风格的策略损失最大化优势同时最小化与旧策略的 KL。 dist policy(obs) action dist.rsample() log_prob dist.log_prob(action).sum(dim-1) adv (adv - adv.mean()) / (adv.std() 1e-8) # 简单优势标准化 policy_loss -(log_prob * adv).mean() kl torch.distributions.kl_divergence(dist, old_policy(obs)).sum(dim-1).mean() total_loss policy_loss kl_weight * kl return total_loss, policy_loss, kl def awr_policy_loss(policy, obs, adv, temperature1.0): AWR 式变分策略更新用指数加权构造隐式目标分布。 dist policy(obs) action dist.rsample() log_prob dist.log_prob(action).sum(dim-1) adv (adv - adv.mean()) / (adv.std() 1e-8) weight torch.exp(adv / temperature) loss -(weight.detach() * log_prob).mean() return loss在第一个函数里old_policy必须是更新前的策略快照否则 KL 会恒等于 0。第二个函数不需要显式计算 KL但温度参数 temperature 控制了“更新幅度”温度越低高优势样本权重越大策略越激进。5.3 训练主脚本接下来写一个完整的主脚本。它会生成一批合成数据用一个随机初始化的“旧策略”作为参考然后对新策略执行多步变分更新。# 文件路径variational_rl_demo/train_demo.py import torch import torch.nn as nn from gaussian_policy import GaussianPolicy from elbo_objective import kl_regularized_policy_loss, awr_policy_loss torch.manual_seed(0) obs_dim 4 act_dim 2 batch_size 128 # 随机生成一批“离线数据”模拟旧策略与环境交互的产物 obs torch.randn(batch_size, obs_dim) # 真实最优动作对当前特征有线性关系方便观察学习效果 act 0.8 * obs[:, :act_dim] 0.2 * torch.randn(batch_size, act_dim) # 优势函数动作离“目标动作”越近优势越高 target 0.8 * obs[:, :act_dim] adv -(act - target).pow(2).sum(dim-1) 2.0 old_policy GaussianPolicy(obs_dim, act_dim) policy GaussianPolicy(obs_dim, act_dim) # 让新策略从旧策略附近初始化方便观察 KL 变化 policy.load_state_dict(old_policy.state_dict()) optimizer torch.optim.Adam(policy.parameters(), lr1e-3) for step in range(500): optimizer.zero_grad() # 方式一显式 KL 正则 total_loss, policy_loss, kl kl_regularized_policy_loss( policy, old_policy, obs, adv, kl_weight0.1 ) # 方式二AWR 式加权回归可选择注释掉方式一 # total_loss awr_policy_loss(policy, obs, adv, temperature1.0) total_loss.backward() optimizer.step() if step % 50 0: entropy policy(obs).entropy().sum(dim-1).mean().item() print(fstep{step:3d} | loss{total_loss.item():.4f} | fpolicy_loss{policy_loss.item():.4f} | fkl{kl.item():.4f} | entropy{entropy:.4f})运行方式cd variational_rl_demo python train_demo.py这段代码在一两百步内通常就能看到策略损失下降KL 保持在一个稳定范围策略熵略微下降但不至于坍缩。这正好对应变分推断里的“拟合任务”和“不偏离约束”之间的平衡。6. 运行结果与效果验证6.1 预期输出由于随机种子固定输出的训练日志会比较平稳。正常状态下日志会像下面这样step 0 | loss0.4123 | policy_loss0.6532 | kl0.0000 | entropy2.1123 step 50 | loss0.2072 | policy_loss0.5214 | kl0.0812 | entropy1.8234 step100 | loss0.1531 | policy_loss0.4659 | kl0.1033 | entropy1.7042 step150 | loss0.1334 | policy_loss0.4321 | kl0.1210 | entropy1.6621观察重点有三个policy_loss 应该逐渐下降说明新策略在朝着高优势方向移动KL 应该稳定在一个合理区间而不是快速飙到几十entropy 应该缓慢下降但不能掉到接近 0。6.2 如何判断训练是否有效最直接的验证方法是对比旧策略和新策略在同样状态下的动作概率。你可以随机抽几个状态分别用old_policy(obs)和policy(obs)计算均值观察新策略是否朝目标动作靠近。with torch.no_grad(): test_obs torch.randn(5, obs_dim) old_means old_policy(test_obs).mean new_means policy(test_obs).mean print(old policy means:, old_means) print(new policy means:, new_means)如果新策略的均值和 target 更接近且 KL 没有暴涨说明变分更新闭环是通的。6.3 失败的典型信号如果训练过程中 KL 迅速增大最直接的原因是 kl_weight 太小或学习率太大。如果 policy_loss 一直不下降说明优势标准化出了问题或者目标动作和特征之间根本没有可学的关系。出现 NaN 时优先检查 log_std 是否发散以及重参数化采样时是否出现极端值。7. 常见问题与排查思路问题现象可能原因排查方式解决方案KL 散度训练几步后暴涨kl_weight 过小学习率偏大打印每步 KL 数值观察变化趋势增大 kl_weight降低学习率KL 一直为 0old_policy 被错误地重新加载或更新检查是否在循环内重建 old_policy在循环外保存策略快照作为旧策略log_prob 或 loss 出现 NaNlog_std 上溢采样值极端查看 log_std 分布检查是否 clamp缩小 clamp 上界加入梯度裁剪策略迅速坍缩成确定性分布温度过低高优势样本权重过大打印 entropy 曲线提高 temperature恢复熵正则训练稳定但动作均值不更新KL 权重过大策略被约束太紧检查 KL 值是否始终在 0.01 以下减小 kl_weight给策略更多自由度7.1 关于“旧策略快照”的踩坑提醒很多人在复现 TRPO、MPO 时都遇到过这个坑写着写着发现 KL 一直为 0怎么调都不对。原因往往是在每个 batch 里都重新创建了 policy 的网络权重导致新旧策略根本是同一个分布。正确的做法是在开始优化前保存一份策略参数快照或者在训练过程中使用“旧策略推理新策略更新”的结构。用代码表示就是# 正确固定住 old_policy 的参数 old_policy.load_state_dict(policy.state_dict()) # 然后开始多步更新 policyold_policy 不再更新这一点对离线强化学习尤其重要。如果 KL 计算不正确整个变分推断目标就失去了意义。7.2 模式坍塌变分推断里的经典问题模式坍塌在强化学习里表现为策略过早收敛到某个单一动作探索性极差。这对应 ELBO 中“先验正则”太弱或温度过低。SAC 的做法是显式加入熵项MPO 的做法是限制每步最大 KL本质上都是在对抗模式坍塌。如果在小示例里观察到熵快速下降可以调大 kl_weight 或 temperature也可以直接往目标里加一个熵奖励项。8. 工程最佳实践与生产建议8.1 从线性高斯先开始如果你要为一个新任务设计带变分推断的 RL 算法建议先在线性高斯设定下验证数学推导再迁移到神经网络。很多 KL 计算和梯度问题在低维线性情况下一眼就能看出来放到高维神经网络里就变成玄学。8.2 明确 KL 的方向KL 散度不对称。KL(q || p) 和 KL(p || q) 的行为完全不同。前者更关注用 q 覆盖 p 的主要区域后者更强调 q 不能放在 p 概率很低的地方。策略更新里最常用的是“新策略对旧策略”的 KL因为它能直接限制新策略不要跑到旧策略置信区间外。写代码时一定要明确到底算的是哪个方向否则实验结果会非常奇怪。8.3 把算法源码映射到变分推断概念读 SAC、MPO、AWR 的源码时不要只看 API 调用要建立一张映射表算法或源码片段变分推断对应物SAC 的熵项防止策略坍缩的先验正则TRPO 的 KL 约束分布投影的硬约束条件MPO 的 E 步构造变分目标分布MPO/AWR 的 M 步监督式分布拟合Dreamer 里的 VAE 损失世界模型中的表示推断带着这张映射表去读代码你会发现这些算法没有想象中那么多“魔法”它们是在同一套数学框架下做了不同的工程取舍。8.4 安全与可复现建议所有涉及策略网络更新的实验都要先固定随机种子训练日志至少记录 loss、KL、entropy、grad_norm 这四个指标策略检查点要分版本保存方便回滚到上一轮稳定策略在生产或真实环境验证新策略前先在离线数据上评估 KL 变化是否符合预期。9. 总结与后续学习方向伯克利 2026 春季深度强化学习课程第 12 讲真正想传递的不是“变分推断很厉害”这个结论而是让我们切换视角从“策略是奖励的上升方向”切换到“策略是一个分布我们永远在分布空间里做投影和匹配”。理解 ELBO 会减少很多读源码时的困惑。再看 SAC 的熵正则你想到的不再是“为什么要加一个奇怪的奖励”而是“它在防止策略坍缩成错误的点估计”。再看 MPO 和 AWR你会认出 E 步和 M 步的交替本质上是变分 EM 在策略优化里的变体。动手验证这一讲最好的方式是把你手上的某个策略梯度算法改造成带 KL 约束的形式先在线性高斯任务上跑通再去读一份开源 SAC 实现把里面的熵项和 KL 概念一一对应起来。这一步做完你的深度强化学习基础就算真正打牢了。建议收藏本文备用下次再看到某篇论文里出现 ELBO 时你会发现自己已经能看懂它在做什么了。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

深度学习表情识别实战:从FER2013数据预处理到mini-Xception模型训练 2026/9/28 5:02:50

深度学习表情识别实战:从FER2013数据预处理到mini-Xception模型训练

简介:这是一份面向计算机专业毕业设计、课程设计与期末大作业场景的深度学习面部表情识别项目资源,适合正在完成相关课题的学生及需要项目实战练习的开发者。包内提供完整源码、论文文档、数据集与答辩演示文稿,覆盖卷积神经网络、视觉几何组…

阅读更多 →
JavaEE二手图书平台实战:Servlet+JSP+JDBC分层架构与事务控制 2026/9/28 5:02:43

JavaEE二手图书平台实战:Servlet+JSP+JDBC分层架构与事务控制

简介:这是一份面向高校计算机专业学生的JavaEE课程设计与期末大作业实战资源,聚焦二手图书交易场景,完整呈现B/S架构电商平台的设计逻辑与工程实现。资源包含可直接部署运行的源码、配套课程设计报告及详细注释,覆盖用户管理、图书…

阅读更多 →
8.6MB轻量OCR引擎:支持中英文混排、竖排与长文本的边缘部署方案 2026/9/28 5:02:43

8.6MB轻量OCR引擎:支持中英文混排、竖排与长文本的边缘部署方案

简介:这是一套面向开发者与AI工程实践者的超轻量级中文OCR工具库,专为嵌入式部署、边缘计算及快速集成场景设计,解决多语言混合文本、竖排古籍文献、长段落文档等复杂OCR识别难题。资源共2000个文件,以431个Python脚本&#xff08…

阅读更多 →
C++智能指针全解析:面试必考的10个问题,90%的候选人答不全! 2026/9/28 5:02:43

C++智能指针全解析:面试必考的10个问题,90%的候选人答不全!

C++智能指针全解析:面试必考的10个问题,90%的候选人答不全! 本文属于「C++面试必考」系列,从底层原理到实战应用,一篇搞定智能指针所有考点。 前言 在C++面试中,如果要评选"出现频率最高的知识点",智能指针绝对稳居前三。无论是校招还是社招,无论是大厂还是…

阅读更多 →
遥感变化检测实战:U-Net++嵌套结构+SE注意力+显存优化方案 2026/9/28 5:02:43

遥感变化检测实战:U-Net++嵌套结构+SE注意力+显存优化方案

简介:本资源是一套面向本科毕业设计的高分辨率城市建筑物遥感变化检测系统实现方案,聚焦遥感图像语义分割与变化识别任务,适用于地理信息科学、遥感技术、计算机视觉等方向的学生开展算法复现与工程实践。压缩包共11个文件,含6个核…

阅读更多 →
YOLOv8电梯内电瓶车闯入报警系统:从训练到部署全解析 2026/9/28 5:02:43

YOLOv8电梯内电瓶车闯入报警系统:从训练到部署全解析

简介:这是一份基于YOLOv8的电梯内电瓶车闯入报警项目资源,将目标检测与深度学习应用于公共安全场景,适合计算机、人工智能、自动化等专业的在校学生用于毕业设计或课程设计,也适合新手学习目标检测完整流程。压缩包内含8个文件&am…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉