☰
强化学习中的变分推断:原理、应用与PyTorch实战
2026/9/28 5:09:27 网站建设 项目流程

最近在整理深度强化学习课程笔记时,最容易被卡住的一个点就是“为什么要学变分推断”。前面刚把策略梯度、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.md

3. 变分推断核心原理拆解

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_rl

5.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_prime

5.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(zx) 与先验差异始终很大先验网络与编码器没有同步更新

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,观察多模态数据的拟合效果。只有亲手调整过这些模块,才能真正理解变分推断在强化学习中的价值。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询