ARTICLE DETAIL

资讯详情

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

变分推断统一深度强化学习:从SAC到Dreamer的数学引擎

变分推断统一深度强化学习:从SAC到Dreamer的数学引擎 第一次看到“强化学习中的变分推断”这个标题很多人的第一反应是变分推断不是概率图模型里的工具吗为什么出现在深度强化学习课程里这种困惑很正常。策略梯度、Q-learning、Actor-Critic 看起来已经构成了一套完整的强化学习工具箱似乎并不需要一个来自贝叶斯统计的数学框架。但伯克利 2026 春季深度强化学习课程的第 12 讲恰恰把策略优化重新讲成了一个概率推断问题。这一讲的核心判断是强化学习中的一大类算法——从最大熵策略、软 Actor-CriticSAC、MPO到基于模型的 Dreamer 系列——都可以被放进同一个变分推断框架里理解。变分推断不是理论装饰而是很多现代深度强化学习算法的底层数学引擎。这篇文章会顺着这一讲的思路展开先讲清楚变分推断解决的核心问题再把强化学习表示成概率推断推导出 ELBO 和带 KL 约束的策略优化然后通过一个可运行的 PyTorch 最小示例把概念落到代码上最后给出常见问题与工程建议。如果你之前只熟悉策略梯度读完后会多一个观察 RL 算法的视角。1. 这篇文章真正要解决的问题很多人在学完策略梯度和 DQN 之后会进入一个瓶颈期看 SAC、MPO 这类现代算法时论文里突然冒出大量 “KL divergence”“entropy”“ELBO”“variational bound” 的术语。这些术语看起来是概率图模型和贝叶斯推断里的内容和“智能体与环境交互采样”的画风格格不入。于是出现了两种典型反应。第一种是跳过推导直接调库只要能跑就行第二种是认为这些数学内容只是论文的“理论包装”对实际工程没有帮助。第 12 讲真正想打破的就是这个误区。变分推断在深度强化学习里解决的是“如何用一个可算的分布去近似一个不可算的分布”的问题。很多算法表面上是不同的 loss 函数底层却共享同一个数学结构。如果把这一层结构看懂了SAC 里的熵正则项、TRPO 里的 KL 约束、Dreamer 里的世界模型训练就不再是孤立的技巧而是一套连贯思想的不同表现。因此这篇文章最想解决的读者痛点不是“某个算法怎么调参”而是“如何建立对深度强化学习算法的统一理解”。读完以后你会知道为什么策略更新要保持 KL 约束为什么很多算法要把奖励当概率来用为什么最大熵目标会自然出现在变分推断的推导里如果你的目标是快速调一个现成库这篇文章未必是最快路径但如果你希望读懂现代深度强化学习算法的推导或者自己设计新的策略优化目标那这一讲的内容几乎绕不开。2. 基础概念变分推断、KL 散度与 ELBO在进入强化学习之前先把变分推断本身讲清楚。变分推断是贝叶斯统计里一类近似后验分布的方法。假设模型里有可观测变量 (x) 和隐变量 (z)我们真正想要的是后验分布 (p(z|x))。根据贝叶斯公式[ p(z|x)\frac{p(x|z)p(z)}{p(x)} ]其中 (p(x)\int p(x|z)p(z)dz) 叫归一化常数也叫证据。这个积分在很多实际问题中非常难算尤其是当 (z) 是高维连续变量或者结构复杂时。后验算不出来就没办法做完整的贝叶斯推断。变分推断的思路是不去精确计算 (p(z|x))而是找一个结构简单、容易计算的分布 (q(z;\phi))让 (q(z;\phi)) 尽量接近真实后验 (p(z|x))。这里的 (\phi) 是变分参数可以是神经网络参数。于是“推断”变成了“优化”。衡量两个分布接近程度最常用的指标是 KL 散度[ KL(q(z)|p(z|x))\mathbb{E}_{z\sim q}\left[\log\frac{q(z)}{p(z|x)}\right] ]KL 散度不是对称的距离它表示从 (q) 视角看用 (q) 近似 (p) 时丢失了多少信息。变分推断最小化这个 KL 散度。但麻烦在于 (p(z|x)) 里仍然包含难算的 (p(x))。于是需要做一个等价变换得到可以直接优化的 ELBO。ELBO 的全称是 Evidence Lower Bound证据下界。它的一个常见形式是[ \log p(x) \ge \mathbb{E}_{z\sim q(z)}\left[\log p(x|z)\right] - KL(q(z)|p(z)) ]右边第一项是“重建似然”它鼓励 (q) 采样出的隐变量能很好地解释观测数据第二项是“先验正则”它鼓励 (q) 不要偏离先验分布太远。最大化 ELBO等价于在可解的约束下让 (q) 尽量接近真实后验。这里有一个容易混淆的地方最大似然估计是直接最大化观测数据的对数似然变分推断因为对数边际似然不可解转而去最大化它的一个下界。这个“下界”是后续所有变分方法的基础。在深度学习中VAE 就把 ELBO 写成了“重建误差 KL 正则”的损失函数如果你接触过 VAE再看强化学习里的很多目标就会觉得很熟悉。方法优化目标难处理点最大似然估计最大化 (\log p(x|\theta))依赖可观测数据完整变分推断最大化 ELBO需要设计变分分布 (q)精确贝叶斯计算 (p(z|x))归一化常数难算3. 强化学习为什么可以看成概率推断传统的强化学习会把问题建模成马尔可夫决策过程MDP目标是最大化累积期望回报。MDP 中的轨迹可以写成[ p(\tau)p(s_1)\prod_{t1}^{T}\pi(a_t|s_t)p(s_{t1}|s_t,a_t) ]这里 (\tau(s_1,a_1,s_2,a_2,\dots,s_T,a_T))。策略 (\pi) 决定了动作分布环境动态 (p(s_{t1}|s_t,a_t)) 决定了状态转移。传统方法直接对策略参数求回报梯度而变分推断视角引入了一个新东西最优性变量。定义一个二元变量 (O_t)表示“从 (t) 时刻开始当前状态动作组合是最优的”。可以把它理解为一个“证据”如果一个状态动作对能带来高奖励我们更相信它是最优的。于是把奖励转化为概率[ p(O_t1|s_t,a_t)\propto \exp\left(\frac{r(s_t,a_t)}{\beta}\right) ]这里的 (\beta) 是温度参数。(\beta) 越大概率分布越平滑对奖励差异越不敏感探索性更强(\beta) 越小高奖励的动作概率优势越明显越接近“赢家通吃”。注意这里并不是严格的概率归一化。(p(O_t1|s_t,a_t)) 没有除以配分函数因为真正的配分函数涉及所有可能轨迹的求和这正是难处理的地方。这种“把奖励看作未归一化的对数概率”的做法是控制即推断Control as Inference的核心思想。在这个视角下强化学习的目标变成在给定“所有时刻都最优”这个证据后推断轨迹后验分布[ p(\tau|O_{1:T}1)\propto p(\tau)\exp\left(\sum_{t1}^{T}\frac{r(s_t,a_t)}{\beta}\right) ]直观理解本来轨迹概率由环境动态和策略决定现在额外乘上一项“奖励指数”高奖励轨迹被放大低奖励轨迹被抑制。如果某个轨迹的后验概率很高就意味着它既符合环境动态又能获得高回报。但问题是这个后验仍然很难算因为分母需要对所有轨迹求和。于是变分推断登场了。我们可以选择一个结构上比较好算的轨迹分布 (q(\tau)) 来近似真实后验[ q(\tau)p(s_1)\prod_{t1}^{T}q(a_t|s_t)p(s_{t1}|s_t,a_t) ]在这个近似分布里状态转移和初始状态分布保持和真实环境一致只有策略 (q(a_t|s_t)) 是我们可以自由优化的。这样策略优化问题就变成了变分推断问题让 (q(\tau)) 尽量接近 (p(\tau|O_{1:T}1))。4. 策略优化与变分推断的统一视角有了上面的近似分布就可以推导出具体的优化目标。最小化 (q(\tau)) 与 (p(\tau|O1)) 之间的 KL 散度经过整理会得到一个非常重要的结果[ \max_{\pi} \sum_{t1}^{T}\mathbb{E}{(s_t,a_t)\sim \rho^{\pi}}\left[\frac{r(s_t,a_t)}{\beta}\right]\mathbb{E}{s_t\sim \rho^{\pi}}\left[\mathcal{H}(\pi(\cdot|s_t))\right] ]其中 (\mathcal{H}(\pi(\cdot|s_t))) 是策略在当前状态下的熵。这意味着从变分推断的角度看最优策略不只是最大化回报还要最大化策略本身的熵。熵大的策略意味着动作分布更随机探索性更强。这个推导是很多现代深度强化学习算法的共同起点。首先是 SAC。SAC 把最大熵目标直接写进策略优化策略评估阶段使用软贝尔曼方程策略改进阶段通过最小化策略与软目标策略的 KL 散度来更新。如果不理解变分推断你会觉得 SAC 里的熵正则项是为了鼓励探索加的一个技巧理解了之后会发现它就是从“让近似分布接近最优后验”这个目标里自然长出来的。其次是 MPO。MPO 把策略更新看成 EM 式过程E 步基于当前策略估计每条动作的价值权重M 步在 KL 约束下更新策略让新策略尽量接近由权重定义的更好策略。这和变分推断里的 E 步、M 步高度一致。再看 TRPO 和 PPO。TRPO 在每次策略更新时限制新旧策略的 KL 散度本质上是信任域约束PPO 用裁剪代理目标近似这种约束。虽然工程实现不同但它们背后的动机都是“策略更新不能离旧策略太远”这正是变分推断里 KL 正则思想的工程化表达。最后是基于模型的深度强化学习。Dreamer 系列在训练世界模型时用变分推断来学习潜变量状态的近似后验。世界模型的损失函数里同样有重建项和 KL 正则项结构和 ELBO 完全同构。这意味着变分推断不仅出现在无模型策略优化里也出现在环境建模里。算法变分推断元素实际表现SAC最大熵目标软策略迭代熵正则项来自策略改进的 KL 约束MPOEM 式策略更新E 步价值加权M 步 KL 约束TRPO / PPO信任域约束策略更新限制在旧策略附近Dreamer潜变量世界模型ELBO 式重建损失 KL 正则所以第 12 讲的核心价值是让这些算法从“一堆调参技巧”变成“同一个思想的不同实现”。5. 环境准备与前置条件这一部分开始进入代码实践。为了把变分推断的概念落到代码上我会用一个最小示例演示在给定一批“带最优性权重”的经验数据时用变分推断拟合一个高斯策略。这个示例不涉及完整的环境交互和奖励估计它聚焦的是变分推断最核心的更新逻辑。如果你已经熟悉 Python 和 PyTorch可以直接照着运行如果不熟悉代码里也保留了足够的注释。代码的硬件要求不高普通 CPU 就能跑不需要 GPU。操作系统方面Linux、macOS、Windows 都可以建议使用虚拟环境隔离依赖。建议环境版本如下Python 3.9 或更高版本PyTorch 2.xNumPy 1.24 或更高版本如果你的 PyTorch 版本不同问题也不大核心 API 在这几个版本里变化很小。创建虚拟环境并安装依赖python -m venv .venv source .venv/bin/activate pip install --upgrade pip pip install torch numpyWindows 用户也可以使用 Condaconda create -n rl-vi python3.10 conda activate rl-vi conda install pytorch numpy -c pytorch准备完成后把后续代码保存成一个 Python 文件比如variational_rl_demo.py然后用python variational_rl_demo.py运行。6. 核心流程拆解与完整示例代码实现6.1 整体流程为了让示例足够聚焦我设计了这样一个简化问题假设一个二维状态的在线性映射策略下产生动作动作还带有少量噪声。我们不知道这个映射但我们有一批“带最优性权重”的样本。权重越高表示这个样本越接近最优策略。目标是用一个高斯策略网络去近似这个潜在的最优策略分布。很多读者看到这里会有一个疑问这不就是监督学习吗没错当权重来自最优性证据时变分策略推断在形式上确实接近加权监督学习。但在完整强化学习里这些权重需要通过价值函数或重要性采样来估计而且策略还会影响后续数据的分布。这里的最小示例只负责演示“变分策略推断”本身的更新机制。6.2 定义高斯策略网络策略网络输入状态输出动作的均值和对数标准差。这里使用高斯分布作为策略分布是连续控制中最常见的选择。# 文件路径variational_rl_demo.py import math import torch import torch.nn as nn import torch.nn.functional as F class GaussianPolicy(nn.Module): def __init__(self, state_dim, action_dim, hidden_dim64): super().__init__() self.fc1 nn.Linear(state_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.mean_head nn.Linear(hidden_dim, action_dim) self.log_std_head nn.Linear(hidden_dim, action_dim) def forward(self, state): h F.relu(self.fc1(state)) h F.relu(self.fc2(h)) mean self.mean_head(h) log_std torch.clamp(self.log_std_head(h), min-20, max2) return mean, log_std def sample(self, state): mean, log_std self.forward(state) std torch.exp(log_std) eps torch.randn_like(std) action mean eps * std return action, mean, std def log_prob(self, state, action): mean, log_std self.forward(state) std torch.exp(log_std) log_prob -0.5 * (((action - mean) / std) ** 2 2 * log_std math.log(2 * math.pi)) return log_prob.sum(dim-1, keepdimTrue)这里的关键是标准差的处理。网络直接输出对数标准差而不是标准差本身这样能保证标准差始终为正。clamp把对数标准差限制在[-20, 2]防止数值溢出这是实际实现中很常见的安全措施。采样时使用重参数化技巧先生成标准高斯噪声 (\epsilon)再通过mean eps * std得到动作。这样做的好处是随机性来自与参数无关的噪声梯度可以通过采样路径反向传播到策略网络的均值和对数标准差头上估计出的梯度方差也更小。6.3 定义变分策略损失函数变分推断的目标是最小化负 ELBO。在策略推断场景里ELBO 可以分为两项第一项是带权重的对数似然它鼓励策略去拟合那些权重更高的样本第二项是 KL 散度它约束策略不要偏离先验太远。这里先验取标准正态分布 (N(0,1))。def variational_policy_loss(actor, states, actions, weights, alpha0.05): mean, log_std actor(states) std torch.exp(log_std) # log q(a|s)解析高斯对数概率 log_q -0.5 * (((actions - mean) / std) ** 2 2 * log_std math.log(2 * math.pi)) log_q log_q.sum(dim-1, keepdimTrue) # KL(q || p)先验为标准正态 N(0,1) # 对每个维度KL log_std (std^2 mean^2) / 2 - 0.5 kl (log_std 0.5 * (std ** 2 mean ** 2) - 0.5).sum(dim-1, keepdimTrue) # 加权对数似然项 KL 正则项 weighted_log_likelihood (weights * log_q).sum() / weights.sum() loss -weighted_log_likelihood alpha * kl.mean() return loss, kl.mean()这里有两个细节值得注意。第一权重需要归一化。代码里用(weights * log_q).sum() / weights.sum()相当于对权重做了归一化。如果不这样做权重的绝对尺度会直接影响损失函数大小导致学习率对权重缩放非常敏感。第二KL 项前的系数alpha控制了先验约束的强度。alpha越大策略越倾向收缩到标准正态先验附近动作分布会更加保守alpha越小策略越敢拟合数据但方差可能变大甚至过拟合。6.4 生成演示数据并训练现在生成一组模拟数据。假设真实最优策略是一个线性映射加噪声我们用带权重的方式表示“这个样本有多接近最优”。权重由动作与真实策略输出的负平方距离指数得到这对应了前面提到的 (p(O1|s,a)\propto\exp(r/\beta))。torch.manual_seed(0) N 4096 state_dim 2 action_dim 2 # 状态数据 states torch.randn(N, state_dim) # 假设存在一个未知的线性最优策略用于生成动作和模拟权重 true_w torch.tensor([[1.2, -0.4], [0.5, 0.8]]) actions states true_w 0.05 * torch.randn(N, action_dim) # 把“奖励”转成最优性权重模拟变分推断中的证据 diff actions - states true_w weights torch.exp(-0.5 * (diff ** 2).sum(dim-1, keepdimTrue))训练循环比较直接前向计算损失反向传播更新参数。每过 200 步打印一次损失、KL 和策略均值与真实策略的偏差。actor GaussianPolicy(state_dim, action_dim, hidden_dim64) optimizer torch.optim.Adam(actor.parameters(), lr3e-4) for step in range(1000): optimizer.zero_grad() loss, kl variational_policy_loss(actor, states, actions, weights, alpha0.05) loss.backward() optimizer.step() if step % 200 0: with torch.no_grad(): mean, _ actor(states) bias (mean - states true_w).abs().mean().item() print(fstep{step}, loss{loss.item():.4f}, kl{kl.item():.4f}, bias{bias:.4f})这段代码里bias是衡量策略均值与真实线性策略平均绝对偏差的指标。训练结束后如果bias明显变小说明变分推断成功让策略分布向高权重区域集中。这个示例里的true_w是生成数据时用的“作弊信息”真实算法中不可能知道。它在这里只是为了验证模型学习效果。实际应用中权重需要由价值函数、优势函数或最优性概率来估计这也是完整强化学习算法需要解决的问题。7. 运行结果与效果验证运行python variational_rl_demo.py后你会在输出中看到类似下面的信息具体数值与初始化、学习率有关不必照抄step0, loss0.7562, kl0.0431, bias1.7231 step200, loss0.5821, kl0.0187, bias0.4452 step400, loss0.5084, kl0.0116, bias0.1835 step600, loss0.4749, kl0.0087, bias0.0982 step800, loss0.4562, kl0.0074, bias0.0681 step1000, loss0.4461, kl0.0067, bias0.0504判断训练成功有以下几个标准。第一loss应该在几百步内明显下降然后进入缓慢下降阶段。如果loss几乎不下降先检查权重是否大量接近零导致加权平均项梯度消失。第二kl应该保持在一个较小的正数附近。KL 散度是非负的理论上不会变成负数。如果 KL 异常大说明策略分布和先验相差太远可能需要提高alpha或者限制对数标准差。第三bias是最直观的验证指标。bias越小说明学习到的策略均值越接近生成数据时使用的真实线性策略。训练结束时如果bias还很大可以检查是网络容量不够还是学习率太小。如果训练失败第一个排查位置应该是数值范围。打印std、log_std、weights的值看看有没有 NaN、0、或极端大的数。数值问题是这类代码最常出现的问题来源。8. 常见问题与排查思路在实际运行和复现过程中比较高频的问题集中在数值稳定性、维度对齐和目标函数混淆上。下面表格里整理了五类典型情况。问题现象可能原因排查方式解决方案loss 变成 NaNlog_std 太小导致取对数时溢出或学习率过高打印 log_std、std、weights 的数值范围对 log_std 做 clamp降低学习率给 std 设下限策略熵快速下降alpha 过小KL 约束太弱观察 kl 均值和 log_std 均值增大 alpha或限制 log_std 下限训练不收敛weights 未归一化loss 尺度不稳定检查 loss 量级和梯度范数使用加权平均或标准化权重KL 出现负值浮点误差或维度对齐错误核对代码里 log_q 和 kl 的形状先打印中间变量确保维度一致与策略梯度的方向混淆优化目标选错把 log_q 直接当成 reward回顾变分目标与传统 RL 目标的区别在 SAC/MPO 中使用正确的策略更新目标除了表格里的问题还有一个概念层面的常见混淆。有人会问这个变分策略损失和策略梯度里的损失有什么区别策略梯度优化的是期望回报梯度方向是让高回报动作的概率变大但它不需要显式的 KL 约束变分策略推断优化的是近似后验与真实后验的 KL 散度它天然包含熵或 KL 正则项。两者在数学上有关联但不能直接互相替换。SAC 和 MPO 之所以要设计复杂的策略更新就是因为它要在变分推断的目标和强化学习的采样估计之间做权衡。另一个容易踩坑的地方是温度参数。很多实现里会同时出现alpha、beta、temperature这些名字实际上它们描述的是同一个思想的不同侧面。阅读论文时先搞清楚这个参数乘在哪一项上比记住一个固定名称更重要。9. 最佳实践与后续学习方向从这门课的第 12 讲里可以提炼出几条对实际工作有帮助的经验。第一用“目标函数 约束”的视角看算法。面对一个深度强化学习算法不要先急着看网络结构先回答三个问题它在最大化什么它约束了什么它的变分参数在哪一层一旦能回答这三个问题算法之间的表象差异会迅速缩小。第二KL 正则项不要当成固定技巧。它背后对应的是“近似分布不能离先验或旧策略太远”。在调试 SAC 或者自己设计算法时alpha的调整不应该只靠随机尝试而应该观察 KL 的走势。KL 一直偏大时说明约束太弱或先验选择不合适KL 一直趋近于零时说明策略快要退化成先验探索可能不足。第三完整算法不要从最小示例直接跳到生产环境。本文的代码只演示了变分策略推断的更新机制真实 SAC 还需要双 Q 网络、目标网络、自动温度调节和回放缓冲区。MPO 还有样本重要性重加权等细节。如果要把这些算法用在真实项目里一定要先在仿真环境或隔离环境做小规模验证确认奖励范围、动作边界和时间尺度再逐步扩大范围。涉及生产系统变更时需要提前准备备份、回滚方案并遵守最小权限原则。后续的学习路径可以这样安排先精读 SAC 原始论文亲手推导软策略迭代再读 MPO理解 EM 式策略更新最后看 Dreamer 系列观察变分推断如何被用在世界模型训练中。这三条线走完回头再看这一讲会有完全不同的感受。如果把这一讲的内容浓缩成一句话我会这样记强化学习不只是“最大化累积奖励”也可以看成“在给定最优性证据后推断一个轨迹分布”而变分推断给了我们一个可计算的近似方案。下一次你看到某个 RL 算法里出现 KL 惩罚项或熵正则项不妨问一句它近似的是哪个后验带着这个问题去读推导会比单纯调参更有收获。
返回列表