最近在整理深度强化学习课程笔记时,最容易被卡住的一个点就是“为什么要学变分推断”。前面刚把策略梯度、Q 学习、Actor-Critic 这些看得见的算法学完,进入模型型强化学习、技能发现、高效探索之后,突然就全是“隐变量”“后验分布”“ELBO”这些抽象概念。网上的资料又往往只讲变分推断本身,很少解释它为什么会出现在强化学习里,更缺少可以直接跑通的代码。
伯克利 2026 春季深度强化学习课程第 12 讲“强化学习中的变分推断”刚好补上这条线。这一讲把变分推断从概率图模型引入强化学习,解释了隐变量动力学模型、多模态预测、技能发现和基于模型的探索背后的统一数学框架。本文以这一讲为主线,先把变分推断的核心原理拆开讲清楚,再给出一个完整的 PyTorch 隐空间动力学模型实战,最后整理常见调试问题和工程建议。
先说清楚一个边界:本文不是课程视频的逐字稿,而是把第 12 讲涉及的知识点重新组织成一份可查阅、可运行的学习笔记。代码部分使用合成数据便于演示,真实项目中需要替换成仿真器或真实环境采样的转移数据。
1. 为什么强化学习需要变分推断
1.1 变分推断解决什么问题
变分推断(Variational Inference)是贝叶斯推断中的一类近似计算方法。在常规监督学习里,我们通常直接建模观测与标签之间的关系;但很多场景中,数据背后还存在一个无法直接观测的隐变量 z。例如:
- 一段轨迹中存在没有记录在状态里的“意图”;
- 同一个状态和动作,在不同语义下会走向完全不同的执行模式;
- 环境动力学中存在无法用当前状态完全解释的外部干扰。
把观测记为 x,隐变量记为 z,完整生成模型可以写为:
p(x) = ∫ p(x|z) p(z) dz问题在于,这个积分通常没有解析解,直接采样 z 又会发现绝大多数样本落在低概率区域,导致最大似然估计根本无法执行。变分推断的思路是:不去精确计算真实后验 p(z|x),而是引入一个可学习的近似分布 q(z|x),通过优化让 q(z|x) 尽量接近 p(z|x)。这样一来,本来“算不动”的推断问题变成“优化一个神经网络输出分布”的问题,正好可以用反向传播求解。
1.2 强化学习中的典型触发场景
强化学习里变分推断的触发场景大致可以分成四类:
- 部分可观测:智能体只能看到观测 o_t,真正的状态 s_t 可能是隐藏的,需要在隐空间做状态推断;
- 数据多峰:相同状态和动作可能导向多个不同结果,单峰高斯模型表达能力不够,需要用隐变量描述不同分支;
- 技能与选项发现:希望智能体自发分化出多种行为模式,技能本身就是一种隐变量;
- 高效探索:智能体应当优先尝试信息增益最大的动作,而信息增益需要近似后验才能计算。
这四类场景会在后面的模型型强化学习、规划、机器人控制中反复出现。理解了它们,再看具体算法时就不会觉得“为什么突然冒出一个编码器”。
1.3 与前序内容的衔接
伯克利深度强化学习课程的前半部分,重心在无模型方法:策略梯度、Q 学习、Actor-Critic。到了第 12 讲,课程开始转向“模型 + 推理”的路线。这里的关键转变是:智能体不再只学一个策略,而是先学会描述环境的结构,再在结构上进行规划、探索或技能提取。变分推断就是这个阶段的理论基础设施。
2. 环境准备与版本说明
本文代码使用 Python + PyTorch,核心依赖如下:
- Python 3.9 及以上;
- PyTorch 2.x;
- numpy;
- 可选 gymnasium,用于在真实仿真环境替换合成数据。
版本需要根据你的项目实际情况调整,本文以常见稳定环境为例,重点演示配置思路,不绑定某个精确版本。建议先创建虚拟环境并安装依赖:
python -m venv rl_vi_v2 source rl_vi_v2/bin/activate # Windows 下使用 rl_vi_v2\Scripts\activate pip install torch numpy示例项目结构如下:
variational_rl/ ├── model.py # 变分动力学模型定义 ├── train.py # ELBO 损失与训练函数 ├── generate_data.py # 合成数据生成 ├── main.py # 训练入口脚本 └── README.md3. 变分推断核心原理拆解
3.1 从最大似然到无法直接计算的积分
假设我们观测到一组数据 x,希望通过隐变量 z 学习一个生成模型 p(x|z)。最大似然目标要求最大化 log p(x),但完整表达式要把隐变量积分掉:
log p(x) = log ∫ p(x|z) p(z) dz当 p(x|z) 由神经网络表达时,这个积分既没有解析解,也无法用朴素蒙特卡洛估计:直接从先验 p(z) 采样得到的 z 几乎不会被编码到高概率区域。这正是隐变量模型难训练的本质原因,也是为什么需要变分推断。
3.2 ELBO:证据下界
变分推断的关键是引入近似后验 q(z|x),然后对 log p(x) 做变换:
log p(x) = log ∫ q(z|x) * p(x|z) * p(z) / q(z|x) dz把积分看成 q(z|x) 下的期望,再利用 Jensen 不等式将 log 移入期望内部,得到:
log p(x) >= E_{q(z|x)}[log p(x|z)] - KL(q(z|x) || p(z))等式右边就是证据下界(Evidence Lower Bound,简称 ELBO)。它由两部分组成:
- 重建项:E_{q(z|x)}[log p(x|z)],衡量从隐变量 z 重建观测 x 的效果;
- KL 项:KL(q(z|x) || p(z)),衡量近似后验与先验分布之间的距离。
最大化 log p(x) 等价于最大化 ELBO。当 q(z|x) 恰好等于真实后验 p(z|x) 时,不等式取等号。
当 q(z|x) 和 p(z) 都是高斯分布时,KL 散度有解析解,不需要采样估计。若 q 的参数为 μ_q、σ_q,p 的参数为 μ_p、σ_p,则:
KL(q||p) = log(σ_p / σ_q) + (σ_q^2 + (μ_q - μ_p)^2) / (2 σ_p^2) - 1/2这个公式在代码里几乎每天都要用到,建议直接背下来。
3.3 重参数化技巧
ELBO 的重建项是一个“采样期望”。如果直接从 q(z|x) 采样 z,梯度无法通过采样点回传到编码器参数。重参数化技巧的做法是把随机性从参数中剥离:
z = μ + σ * ε,其中 ε ~ N(0, I)先采样标准正态噪声 ε,再通过确定性变换生成 z。这样 z 对 μ 和 σ 的依赖是确定性的,反向传播可以正常进行。这个技巧是现代变分自编码器(VAE)能用梯度训练的核心原因,也是后续所有变分强化学习算法的公共底座。
3.4 与 EM 算法的联系
如果忽略神经网络参数,变分推断和 EM 算法(期望最大化)的关系非常直接。EM 的 E 步计算或近似后验 q(z|x),M 步在此基础上最大化对数似然。变分推断相当于把 E 步也参数化并通过梯度下降完成,同时允许先验、似然和近似后验都是可微神经网络。在强化学习中,这个视角帮助我们理解一个反复出现的名词:变分自编码器本质上是“用神经网络实现的、可微分的 EM”。
4. 变分推断在强化学习中的应用场景
4.1 隐变量动力学模型
隐变量动力学模型(latent dynamics model)把环境转移建模成两步:
z ~ p(z | s, a) s' ~ p(s' | s, a, z)当环境转移存在多种模态时,例如机器人推进正转和反转产生不同轨迹,确定性网络学到的是所有分支的“平均结果”,这种平均会导致长时预测误差快速累积。引入隐变量之后,模型可以显式表示多模态分支,规划器也能利用预测方差评估风险。这是变分推断在模型型强化学习中最直接的应用。
4.2 多模态轨迹预测与模型型强化学习
模型型强化学习(model-based RL)中,智能体先学习环境模型,再通过规划器选择动作。变分推断在这里有两个作用:
- 用 ELBO 训练带隐变量的转移模型,增强多模态表达能力;
- 利用后验方差作为预测不确定性,避免规划被过度自信的错误预测误导。
类似的隐空间建模思路在 PlaNet、Dreamer 等基于模型的强化学习系统中也有体现。它们把原始像素或状态嵌入到隐空间,再在隐空间里做状态递推和规划。理解 ELBO 之后,再看这些系统会觉得结构清晰很多。
4.3 技能发现:DIAYN 与变分选项发现
“技能发现”问题希望智能体不借助外部奖励也能分化出多种可用行为。以 DIAYN(Diversity Is All You Need)为例,技能 z 被建模为隐变量,优化目标是技能与状态之间的互信息 I(z; s)。互信息本身难算,但可以改写成变分下界,再用一个判别器去近似。变分选项发现(VOD)则把选项作为隐变量,用 ELBO 联合训练高层控制器与低层策略。
这两类方法的共同点是:把“发现结构”转化为“优化变分下界”。先有 ELBO 的概念,再看这类论文会顺畅很多。
4.4 探索与信息增益:VIME
VIME(Variational Information Maximizing Exploration)把探索解释为:选择能最大化环境动力学不确定性下降幅度的动作。实现时维护一个贝叶斯神经网络,用变分推断近似网络参数的后验分布,并以信息增益作为内在奖励。这个方向把“好奇心”形式化为 KL 散度和熵的变化,是变分推断在探索领域最经典的例子之一。
4.5 时序变分模型:TD-VAE 的思想
TD-VAE(Temporal Difference Variational Autoencoder)把变分推断与时序差分思想结合起来。它解决的问题是:智能体不仅要推断当前信息,还要预测较远未来会到达的状态,但未来无法直接观测。TD-VAE 使用跳跃式的变分目标,把“当前能推断什么”与“预测未来能带来的信息”统一到同一个下界里。它常用于 Play 类数据和目标条件强化学习,也是理解“规划在隐空间中进行”的重要铺垫。
5. 实战:PyTorch 实现隐空间变分动力学模型
下面实现一个完整的隐空间变分动力学模型。模型结构包含三部分:
- 编码器 q(z | s, a, s'):从转移结果推断隐变量后验;
- 先验网络 p(z | s, a):从当前状态和动作预测隐变量先验;
- 解码器 p(s' | s, a, z):利用隐变量重建下一状态。
训练目标是最大化 ELBO,即最小化“重建误差 + KL 散度”。
5.1 创建项目结构
mkdir -p variational_rl && cd variational_rl5.2 生成模拟数据
# 文件路径:variational_rl/generate_data.py import torch def generate_synthetic_data(num_samples=8000, state_dim=3, action_dim=2, seed=2026): """生成带隐变量结构的人工转移数据。 这里人为构造了一个离散隐因子 z_factor: 当 a 的第一维为正时,状态沿 a 的方向更新; 当 a 的第一维为负时,状态沿 -a 的方向更新。 真实场景中这类多峰结构可能来自不同环境模式或未观测因素, 这里用合成数据方便快速跑通训练流程。 """ torch.manual_seed(seed) s = torch.randn(num_samples, state_dim) a = torch.randn(num_samples, action_dim) z_factor = torch.sign(a[:, :1]).clamp(-1.0, 1.0) update = z_factor * torch.tanh(a[:, :state_dim]) s_prime = s + 0.3 * update + 0.05 * torch.randn(num_samples, state_dim) return s, a, s_prime5.3 定义模型
# 文件路径:variational_rl/model.py import torch import torch.nn as nn class VariationalDynamicsModel(nn.Module): """基于变分推断的隐空间动力学模型。 输入:当前状态 s、动作 a、下一时刻状态 s'。 训练时使用编码器 q(z | s, a, s') 采样隐变量; 推理时不依赖 s',直接使用先验 p(z | s, a) 采样。 """ def __init__(self, state_dim, action_dim, latent_dim=16, hidden_dim=128): super().__init__() # 编码器部分 self.encoder = nn.Sequential( nn.Linear(state_dim + action_dim + state_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), ) self.enc_mean = nn.Linear(hidden_dim, latent_dim) self.enc_logvar = nn.Linear(hidden_dim, latent_dim) # 先验网络部分 self.prior = nn.Sequential( nn.Linear(state_dim + action_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), ) self.prior_mean = nn.Linear(hidden_dim, latent_dim) self.prior_logvar = nn.Linear(hidden_dim, latent_dim) # 解码器部分 self.decoder = nn.Sequential( nn.Linear(state_dim + action_dim + latent_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, state_dim), ) @staticmethod def reparameterize(mean, logvar): """重参数化采样:z = mu + sigma * eps""" eps = torch.randn_like(mean) std = torch.exp(0.5 * logvar) return mean + eps * std @staticmethod def kl_gaussian(enc_mean, enc_logvar, prior_mean, prior_logvar): """计算 q(z|x) 与 p(z) 两个高斯分布之间的 KL 散度。 公式:KL(q||p) = 0.5 * (logvar_p - logvar_q + (exp(logvar_q) + (mean_q - mean_p)^2) / exp(logvar_p) - 1) """ kl = 0.5 * torch.sum( prior_logvar - enc_logvar + (enc_logvar.exp() + (enc_mean - prior_mean) ** 2) / torch.exp(prior_logvar) - 1 ) return kl def encode(self, s, a, s_prime): h = self.encoder(torch.cat([s, a, s_prime], dim=-1)) return self.enc_mean(h), self.enc_logvar(h) def get_prior(self, s, a): h = self.prior(torch.cat([s, a], dim=-1)) return self.prior_mean(h), self.prior_logvar(h) def forward(self, s, a, s_prime=None, use_posterior=False): prior_mean, prior_logvar = self.get_prior(s, a) if s_prime is not None and use_posterior: enc_mean, enc_logvar = self.encode(s, a, s_prime) z = self.reparameterize(enc_mean, enc_logvar) return z, (enc_mean, enc_logvar), (prior_mean, prior_logvar) z = self.reparameterize(prior_mean, prior_logvar) return z, None, (prior_mean, prior_logvar) def predict_next_state(self, s, a, z): return self.decoder(torch.cat([s, a, z], dim=-1))5.4 编写训练循环
# 文件路径:variational_rl/train.py import torch import torch.nn.functional as F def compute_elbo_loss(model, s, a, s_prime, beta=1.0): """计算 ELBO 形式的损失:recon_loss + beta * kl_loss。 参数 beta 用于 KL 退火。训练初期 beta 从 0 开始逐渐增大, 可以避免模型一上来就牺牲重建精度来强行匹配先验。 """ z, posterior, (prior_mean, prior_logvar) = model(s, a, s_prime, use_posterior=True) enc_mean, enc_logvar = posterior s_prime_pred = model.predict_next_state(s, a, z) recon_loss = F.mse_loss(s_prime_pred, s_prime, reduction='sum') kl_loss = model.kl_gaussian(enc_mean, enc_logvar, prior_mean, prior_logvar) return recon_loss + beta * kl_loss, recon_loss, kl_loss def train_one_epoch(model, dataloader, optimizer, beta=1.0): model.train() total_loss = 0.0 total_recon = 0.0 total_kl = 0.0 for s, a, s_prime in dataloader: optimizer.zero_grad() loss, recon_loss, kl_loss = compute_elbo_loss(model, s, a, s_prime, beta=beta) loss.backward() optimizer.step() total_loss += loss.item() total_recon += recon_loss.item() total_kl += kl_loss.item() n = len(dataloader) return total_loss / n, total_recon / n, total_kl / n# 文件路径:variational_rl/main.py import torch import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset from model import VariationalDynamicsModel from train import train_one_epoch from generate_data import generate_synthetic_data def main(): state_dim, action_dim, latent_dim = 3, 2, 4 s, a, s_prime = generate_synthetic_data(num_samples=8000) dataset = TensorDataset(s, a, s_prime) dataloader = DataLoader(dataset, batch_size=256, shuffle=True) model = VariationalDynamicsModel( state_dim=state_dim, action_dim=action_dim, latent_dim=latent_dim, hidden_dim=128, ) optimizer = optim.Adam(model.parameters(), lr=1e-3) total_epochs = 30 for epoch in range(total_epochs): # KL 退火:前 10 轮 beta 从 0 线性增加到 1 beta = min(1.0, epoch / 10.0) loss, recon, kl = train_one_epoch(model, dataloader, optimizer, beta=beta) if (epoch + 1) % 5 == 0: print( f"Epoch {epoch + 1:02d}/{total_epochs} | " f"loss={loss:8.3f} | recon={recon:8.3f} | " f"kl={kl:6.3f} | beta={beta:.2f}" ) # 推理示例:不使用 s',仅用先验采样做一步预测 model.eval() with torch.no_grad(): s0 = torch.randn(1, state_dim) a0 = torch.randn(1, action_dim) z, _, _ = model(s0, a0, use_posterior=False) s1_pred = model.predict_next_state(s0, a0, z) print("预测下一状态:", s1_pred.numpy().tolist()) if __name__ == "__main__": main()5.5 运行与结果验证
cd variational_rl python main.py预期效果是:损失整体下降,重建误差(recon)逐步收敛;KL 项在 beta 较小时接近 0,随着 beta 增大而上升,最后稳定在一个较小值附近。不同机器的具体数值会有差异,重点观察两类 loss 的分量变化趋势。
需要注意,这里的数据是人为构造的,只为验证训练流程。真实项目中,把 generate_synthetic_data 替换成从 gymnasium 或仿真器采样的状态转移数据集即可。
6. 常见问题与调试思路
6.1 现象、原因与解决思路速查表
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 训练几轮后 KL 项迅速归零 | 后验坍缩(posterior collapse),解码器忽略隐变量 | KL 退火、Free Bits、增强解码器容量 |
| 重建 loss 持续偏高且不下降 | 高斯解码器表达能力不足,多峰被平均 | 使用 MDN(混合密度网络)解码器、增大隐变量维度 |
| ELBO 震荡剧烈 | 重参数化采样方差大、学习率偏高 | 降低学习率、增大 batch、梯度裁剪 |
| q(z | x) 与先验差异始终很大 | 先验网络与编码器没有同步更新 |
6.2 后验坍缩详解
后验坍缩是最常见、也最影响效果的问题。现象是 KL 项很快降到接近 0,编码器输出几乎不依赖输入数据,解码器直接“忽略”隐变量,模型退化成普通的确定性网络。
常见原因包括:解码器能力太强,不需要隐变量就能重建数据;KL 项权重过大,模型宁可牺牲重建精度也要把后验拉向先验;训练初期梯度就偏向 KL 方向。
推荐的预防手段是 KL 退火。训练开始时把 KL 权重 beta 设为 0,让模型先学会用隐变量重建数据,再逐渐提高 KL 的约束力。本文 main.py 中已经实现了这个逻辑。
6.3 预测结果“平均化”问题
如果发现不同隐变量 z 采样下预测的下一状态几乎相同,问题多半出在解码器的输出分布上。高斯解码器通常只能表达单峰分布,而真实转移往往是多峰的:同一个状态和动作下,可能走向完全不同的未来。解决办法是换成混合密度网络(MDN),让解码器输出多个高斯分布的混合,ELBO 形式也要相应改写成混合分布的 log-likelihood。这个改动虽然成本不高,但在多模态动力学建模中效果差异非常明显。
7. 工程实践与调参建议
7.1 KL 退火与 Free Bits
KL 退火只是手段,不是终点。一种更稳健的做法是 Free Bits:给 KL 每个维度设置一个下限,使隐变量即使没有足够信息量也保留一定编码能力,避免模型完全丧失对隐变量的依赖。常见实现示例如下:
free_bits = 0.5 # 每个隐维度至少保留 0.5 nats kl_per_dim = 0.5 * ( prior_logvar - enc_logvar + (enc_logvar.exp() + (enc_mean - prior_mean) ** 2) / torch.exp(prior_logvar) - 1 ) # shape: (batch, latent_dim) kl_loss = torch.maximum(kl_per_dim, torch.full_like(kl_per_dim, free_bits)).sum(dim=-1).mean()这个方案的优点是不依赖人为设定的退火曲线,模型自己决定哪些维度需要压缩、哪些维度保留结构。实际项目中建议先跑一个 KL 退火版本,再对比 Free Bits 版本,观察哪一条曲线更稳定。
7.2 训练与推理一致性
本文模型在训练时使用后验编码器 q(z | s, a, s'),推理时改用先验 p(z | s, a)。如果训练和推理的隐变量分布差距较大,会出现 rollout 时预测逐步漂移的问题。
工程上通常采用两种对策:
- 减少后验与先验的分布差距:提高 KL 权重,或者共享编码器与先验网络底层参数;
- 直接在先验分布上做 rollout 训练:让模型在隐空间里预测若干步后再重建,增强长时一致性,这也是 Dreamer 类模型的思路。
7.3 训练稳定性与模型评估
训练时建议把 recon 和 KL 分开打印,而不是只看合并后的 loss。否则很难判断是重建没学好,还是 KL 崩了。输入状态和动作也建议做标准化,尤其是机械臂、机器人这类量纲差异大的环境,否则 KL 项会被某个大数值维度带偏。
评估模型时不要只看一步预测误差,建议做多步 rollout,观察误差累积速度。误差快速发散往往意味着隐变量没有学到真正的状态结构,此时需要回到网络容量和数据质量上排查。
7.4 落地强化学习时的安全边界
如果要把这个模型用在机器人或真实环境控制中,务必遵守最小风险原则:
- 先在仿真器里完成完整训练和验证,再考虑迁移到真实环境;
- 保留一个无模型的备份策略,当模型预测方差超过阈值时自动切换;
- 对动作空间设置边界,避免规划器输出越界动作;
- 所有环境模型改动先在沙箱环境验证,再更新到线上进程。
变分推断模型的不确定性估计只在训练分布内可靠,遇到分布外状态时,模型可能“非常自信地预测错误”。真实的工程系统必须对这类场景做显式兜底。
8. 总结与下一步学习路线
这一讲的核心收获可以概括成三句话:隐变量模型无法直接做最大似然估计,所以才有变分推断;变分推断用 ELBO 把难解积分转化为可优化的重建项和 KL 项;重参数化技巧让整个流程可以用梯度下降端到端训练。本文的代码实现了一个隐空间动力学模型,它虽然简单,但已经是模型型强化学习中 latent dynamics 的核心骨架。
下一步建议按以下顺序继续深入:
- 先学归一化流(Normalizing Flows)和扩散模型,它们解决的是“近似分布表达能力不足”的问题,比高斯分布更能刻画复杂后验;
- 再读 Dreamer、PlaNet 的源码,看它们如何在隐空间做规划与训练;
- 之后可以尝试把变分推断应用到技能发现(DIAYN / VOD)或探索(VIME)中;
- 如果关注多智能体强化学习、离线强化学习和大语言模型强化学习,会看到变分推断在这些领域也经常作为“结构发现”的工具出现,原理是相通的。
最后一个建议:不要只看数学推导,必须动手改代码。把本文的 latent_dim 从 4 改成 32,观察 KL 项和重建误差的变化;把高斯解码器改成 MDN,观察多模态数据的拟合效果。只有亲手调整过这些模块,才能真正理解变分推断在强化学习中的价值。