ARTICLE DETAIL

资讯详情

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

深度强化学习中的变分推断:从ELBO到世界模型

深度强化学习中的变分推断:从ELBO到世界模型 在实际落地深度强化学习算法时很多同学会遇到一个坎模型输入是高维图像但有效信息可能只有几个关键状态奖励稀疏时智能体根本不知道应该关注什么加上随机策略天然带有高方差训练过程经常不稳定。面对这些问题变分推断Variational Inference, VI是一个绕不开的工具。伯克利 2026 春季深度强化学习课程的第 11 讲正好系统地把变分推断引入深度强化学习包括 ELBO 推导、重参数化技巧以及它在世界模型和状态表征学习中的典型用法。如果你是刚开始接触深度强化学习的读者这一讲可能比策略梯度还要抽象。但换个角度想变分推断其实就是“用一个简单的分布去逼近一个复杂的后验分布”一旦理解了这条主线后续再看 Dreamer、SVG、变分信息瓶颈等深度强化学习算法都会顺畅很多。本文围绕伯克利 2026 春季学期深度强化学习课程第 11 讲展开整理成一份适合复习和延伸阅读的教程笔记。内容包括变分推断的核心概念、ELBO 推导、与深度强化学习的结合方式、PyTorch 实现示例、常见错误排查以及实际项目中的工程建议。代码部分以常见的变分自编码器VAE为例因为它是理解变分推断与深度强化学习交叉应用的最小可运行载体。1. 背景与核心概念1.1 为什么深度强化学习需要变分推断深度强化学习解决的核心问题是智能体如何通过与环境交互学习到一个策略来最大化累计奖励。传统强化学习假设状态空间完全可观测但在真实场景中智能体往往只能拿到部分观测比如机器人只有摄像头画面、自动驾驶系统只有传感器融合结果。这时我们需要从观测中推断出真实状态而这个推断过程天然具有不确定性。另一种常见困境是图像观测维度很高但真正影响决策的语义因素很少。比如开车场景中天空、道路纹理、树木阴影都是次要信息而前方障碍物的位置、自身车速才是关键。如果直接把原始像素输入策略网络模型需要额外学习大量无关特征样本效率会明显下降。变分推断在这里的作用就是帮我们用可计算的方式得到“给定观测时状态(或隐变量)的后验分布”。它给了深度强化学习一个概率视角不仅要学习一个函数映射还要学习这个映射背后的不确定性。应用到具体方案中就表现为用变分自编码器学习低维状态表征。用变分信息瓶颈压缩策略网络的输入特征。用概率状态空间模型做规划与想象。用变分目标计算好奇心内在奖励。这也是为什么伯克利深度强化学习课程要在策略梯度、Q-Learning 之后专门安排一讲变分推断。它不是一个独立算法而是支撑“基于模型的深度强化学习”“表征学习”“探索机制”等进阶方向的基础工具。1.2 变分推断的直观理解先不谈数学公式我们可以把变分推断理解为“用简单分布拟合复杂分布”。假设我们有观测数据 (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) 往往无法解析计算。深度神经网络模型中的隐变量维度很高积分没有闭式解。于是我们不去求精确后验而是找一个容易计算的分布 (q(z))让它尽量接近真实后验。衡量“接近”程度常用 KL 散度也就是[ \min_q KL(q(z) | p(z|x)) ]这就把推断问题变成了优化问题。优化过程中我们不需要计算复杂积分只需要对 (q(z)) 采样并计算相应的对数概率配合梯度下降即可完成。这在参数化模型下非常自然因为 (q(z)) 可以是一个神经网络输出的分布参数。1.3 变分推断与 MCMC 的对比提到推断很多读者会想到 MCMC马尔可夫链蒙特卡洛。两种方法目标相同但路径不同对比项变分推断MCMC核心思路优化一个近似分布从目标分布中采样计算代价相对低适合大规模数据高需要多次迭代是否得到近似分布是直接可采样只有样本需要统计汇总适用场景深度学习、大规模模型小规模精确推断、理论研究缺点近似有偏可能欠拟合收敛慢诊断困难在深度强化学习场景中模型动辄上百万参数数据来自交互采样MCMC 的代价不可接受因此变分推断几乎成为唯一选择。1.4 课程中的关键定位伯克利 2026 春季深度强化学习课程把变分推断放在“模型学习”和“探索”之间位置非常关键。它承接了概率图模型中的极大似然估计也为后续讲解 Dreamer、Plan2Explore 等算法打下基础。在课程中变分推断并不是以独立算法形式出现而是作为“学习环境动态模型”的工具。比如一个经典流程是智能体拿到观测 (o_t)编码为隐状态 (z_t)通过动态网络预测 (z_{t1})再用解码器重构观测。整个过程的学习目标就是最大化变分下界ELBO。理解了这一点就理解了深度强化学习中变分推断的核心位置。2. 环境准备与实验基础2.1 实验环境说明本文给出的代码示例以 PyTorch 为主因为深度强化学习领域的大多数开源实现都基于 PyTorch 或 JAX而 PyTorch 在自动微分和动态采样的灵活性上更直观。需要准备的组件如下Python 3.9 或更高版本。PyTorch 2.0 或更新的稳定版本。NumPy 用于数据预处理。Matplotlib 用于可视化重构效果可选。如果需要在 Gym 环境中测试可以安装 gymnasium。版本需要根据你的项目实际情况调整本文示例以常见环境为例重点演示配置思路。如果你的 PyTorch 版本较低个别 API 可能需要替换比如torch.functional.F.kl_div在不同版本中的行为基本一致但张量设备CPU/GPU转换要留意。2.2 示例项目结构为了便于后面的实战演示我们先规划一个简单的项目结构variational_rl/ ├── main.py # 训练入口 ├── model.py # VAE / 变分推断模型 ├── trainer.py # 训练逻辑 ├── utils.py # 数据加载、可视化 └── config.py # 超参数配置这个结构很小但对理解变分推断在深度强化学习中的角色已经足够。后续如果扩展到世界模型只需要在此基础上增加一个 RNN 或 SSM 模块。2.3 核心依赖安装如果使用 pip可以执行pip install torch numpy matplotlib gymnasium如果使用 condaconda install pytorch torchvision torchaudio -c pytorch conda install numpy matplotlib gymnasium安装完成后可以快速验证 PyTorch 是否可用import torch print(torch.__version__) print(torch.cuda.is_available())这里输出的torch.cuda.is_available()如果是 False不一定是错误只能说明当前环境没有可用的 CUDA GPU。CPU 环境同样可以运行本文的示例只是训练速度会慢一些。3. 变分推断核心原理拆解3.1 从极大似然到对数边际似然在深度强化学习的模型学习中我们经常要最大化观测数据的似然。假设有一组观测数据 (x_1, x_2, ..., x_N)我们希望学习一个生成模型 (p_\theta(x|z)) 和先验 (p(z))。理论上应该最大化边际对数似然[ \log p_\theta(x) \log \int p_\theta(x|z) p(z) dz ]但直接优化这个积分是不可行的。于是我们引入一个推断网络 (q_\phi(z|x))用它来近似真实后验 (p_\theta(z|x))。3.2 ELBO 推导对任意观测样本 (x)我们可以写出[ \log p_\theta(x) \log p_\theta(x) \cdot \int q_\phi(z|x) dz ]因为 (q_\phi(z|x)) 是一个概率分布积分等于 1。继续变换[ \log p_\theta(x) \int q_\phi(z|x) \log p_\theta(x) dz ]此时我们把 (p_\theta(x)) 变成 ( \frac{p_\theta(x,z)}{p_\theta(z|x)} )得到[ \log p_\theta(x) \int q_\phi(z|x) \log \frac{p_\theta(x,z)}{p_\theta(z|x)} dz ]再拆开[ \log p_\theta(x) \int q_\phi(z|x) \log \frac{p_\theta(x,z)}{q_\phi(z|x)} dz KL(q_\phi(z|x) | p_\theta(z|x)) ]前一项就是证据下界 ELBO后一项是 KL 散度。因为 KL 散度恒大于等于零所以[ \log p_\theta(x) \ge \text{ELBO} ]而[ \text{ELBO} \mathbb{E}{q\phi(z|x)} [\log p_\theta(x|z)] - KL(q_\phi(z|x) | p(z)) ]其中第一项是重构项表示隐变量重建观测的能力第二项是正则项让推断分布不过度偏离先验。3.3 最大化 ELBO 的含义优化 ELBO 时我们同时做了两件事让解码器能从隐变量 (z) 重构出观测 (x)即重构损失越低越好。让编码器输出的分布 (q_\phi(z|x)) 与先验 (p(z)) 尽量接近防止隐变量空间退化成无序状态。在深度强化学习的世界模型中第一项帮助模型学习环境动态第二项保证隐状态空间具备连续性便于规划算法在隐空间中搜索动作。因此 ELBO 不是简单的正则化技巧而是模型学习稳定性的一部分。3.4 重参数化技巧ELBO 中包含对 (q_\phi(z|x)) 的采样操作。如果直接采样梯度无法通过随机节点反传到编码器。为了解决这个问题我们使用重参数化技巧假设 (z \sim q_\phi(z|x)) 是高斯分布即[ q_\phi(z|x) \mathcal{N}(z; \mu_\phi(x), \sigma_\phi^2(x)) ]那么我们可以这样采样[ z \mu_\phi(x) \sigma_\phi(x) \cdot \epsilon ]其中 (\epsilon \sim \mathcal{N}(0, I))。这样一来随机性被独立到 ( \epsilon ) 中(z) 与编码器参数的关系变成确定性的梯度可以通过 (z) 反向传播到 ( \mu_\phi ) 和 ( \sigma_\phi )标准反向传播算法就能正常工作。重参数化技巧是变分自编码器能够成功训练的关键也是后续几乎所有深度强化学习变分模型的标准操作。Dreamer 这样的算法在训练世界模型时同样会在 RSSM 的随机路径上使用重参数化采样。3.5 高斯分布下 KL 的计算当先验取标准正态分布 (p(z)\mathcal{N}(0, I)) 时KL 项有解析式。设编码器输出 (d) 维均值 ( \mu ) 和对数方差 ( \log \sigma^2 )KL 为[ KL -\frac{1}{2} \sum_{j1}^{d} (1 \log \sigma_j^2 - \mu_j^2 - \sigma_j^2) ]在代码中我们通常让网络输出 ( \log \sigma^2 ) 而不是直接输出 ( \sigma )这样可以保证方差恒大于零同时数值稳定性更好。3.6 变分推断与深度强化学习结合的典型接口深度强化学习算法中变分推断通常以三种形式出现第一状态表征学习。用 VAE 将高维图像压缩为低维隐向量然后策略网络和值函数网络都基于隐向量工作。这种方案常见于视觉深度强化学习它把像素级冗余信息剥离提升策略输入的信息密度。第二世界模型学习。将变分推断嵌入环境动态模型的学习通过对隐状态的预测和重构使得智能体可以基于模型想象未来轨迹。Dreamer 系列和 PlaNet 是典型代表。第三探索驱动。利用变分推断构造预测误差或信息增益。比如根据隐状态预测的不确定性计算内在奖励引导智能体访问“难以建模”的状态。这就是基于模型的探索机制典型工作包括 Plan2Explore。在伯克利课程的语境下理解 ELBO 和重参数化就能读懂这些算法的模型学习模块反之若只把这些模型当作黑盒网络很难在遇到训练不稳定时快速定位原因。4. 实战在 PyTorch 中实现变分推断模型为了让变分推断不再停留在公式层面这一节我们用 PyTorch 实现一个用于图像观测的变分自编码器并展示如何把它接入深度强化学习的状态表征模块。这个代码不是完整的深度强化学习算法而是最核心的“模型学习”部分。理解了它后面再加策略网络就很容易。4.1 项目结构按前面规划先创建文件variational_rl/ ├── config.py ├── model.py ├── trainer.py ├── utils.py └── main.py4.2 超参数配置config.py中定义基本超参数。这部分要特别注意变分推断训练对学习率比较敏感建议使用 Adam 优化器初始学习率设置在1e-3附近并配合学习率衰减。# 文件路径variational_rl/config.py class Config: # 数据参数 img_size 64 batch_size 128 # 模型参数 latent_dim 32 hidden_dim 256 # 训练参数 epochs 50 lr 1e-3 beta 1.0 # KL 项权重 # 设备 device cuda # 如果没有 GPU改成 cpu其中beta是 KL 项的权重。在变分推断里如果 KL 项太强模型容易把所有样本编码成同一区域也就是“后验坍塌”如果太弱隐变量空间没有结构无法支持后续的规划。实际深度强化学习项目中beta往往需要调参有些实现还会在训练过程中从 0 开始线性增长。4.3 变分自编码器模型model.py实现编码器、解码器和采样逻辑。编码器输入图像经过两层卷积展平后输出隐变量的均值和对数方差。注意我们并不直接输出隐变量而是先输出分布参数再通过重参数化得到样本。这个设计是变分推断的核心结构。# 文件路径variational_rl/model.py import torch import torch.nn as nn import torch.nn.functional as F class Encoder(nn.Module): 将图像观测编码为隐分布参数 def __init__(self, latent_dim32): super().__init__() self.conv1 nn.Conv2d(3, 32, kernel_size4, stride2, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size4, stride2, padding1) self.conv3 nn.Conv2d(64, 128, kernel_size4, stride2, padding1) # 经过三层卷积后特征图尺寸降低到 8*8假设输入为 64*64 self.fc_mu nn.Linear(128 * 8 * 8, latent_dim) self.fc_logvar nn.Linear(128 * 8 * 8, latent_dim) def forward(self, x): h F.relu(self.conv1(x)) h F.relu(self.conv2(h)) h F.relu(self.conv3(h)) h h.view(h.size(0), -1) mu self.fc_mu(h) logvar self.fc_logvar(h) return mu, logvar class Decoder(nn.Module): 从隐变量重构图像 def __init__(self, latent_dim32): super().__init__() self.fc nn.Linear(latent_dim, 128 * 8 * 8) self.deconv1 nn.ConvTranspose2d(128, 64, kernel_size4, stride2, padding1) self.deconv2 nn.ConvTranspose2d(64, 32, kernel_size4, stride2, padding1) self.deconv3 nn.ConvTranspose2d(32, 3, kernel_size4, stride2, padding1) def forward(self, z): h self.fc(z) h h.view(h.size(0), 128, 8, 8) h F.relu(self.deconv1(h)) h F.relu(self.deconv2(h)) x_recon torch.sigmoid(self.deconv3(h)) return x_recon class VAE(nn.Module): 变分自编码器封装编码器、解码器和采样逻辑 def __init__(self, latent_dim32): super().__init__() self.encoder Encoder(latent_dim) self.decoder Decoder(latent_dim) def reparameterize(self, mu, logvar): 重参数化采样 std torch.exp(0.5 * logvar) eps torch.randn_like(std) return mu eps * std def forward(self, x): mu, logvar self.encoder(x) z self.reparameterize(mu, logvar) x_recon self.decoder(z) return x_recon, mu, logvar def kl_divergence(mu, logvar): 计算高斯先验下的 KL 散度 return -0.5 * torch.sum(1 logvar - mu.pow(2) - logvar.exp(), dim1) def vae_loss(x, x_recon, mu, logvar, beta1.0): 变分推断总损失 重构损失 beta * KL 散度 recon_loss F.mse_loss(x_recon, x, reductionsum) / x.size(0) kl_loss kl_divergence(mu, logvar).mean() total_loss recon_loss beta * kl_loss return total_loss, recon_loss, kl_loss代码中有几个细节需要说明卷积层输出尺寸是否正确取决于输入图像大小。这里默认输入为3*64*64三层步长为 2 的卷积会把尺寸变为64 - 32 - 16 - 8。如果你的输入尺寸不同需要调整fc_mu的输入维度。重构损失使用 MSE 而不是 BCE因为我们的观测被归一化到 0-1 区间并且没有使用特殊的输出激活函数约束。MSE 在视觉重构任务中更稳定。kl_divergence返回的是每个样本 KL 值最后取均值。这里按公式逐项实现没有使用F.kl_div方便初学者对照推导过程。4.4 训练逻辑trainer.py负责封装一次训练的完整流程。为了提高可读性这里把数据加载和训练分开写。# 文件路径variational_rl/trainer.py import torch from torch.optim import Adam from model import VAE, vae_loss def train_one_epoch(model, dataloader, optimizer, beta, device): model.train() total_loss 0.0 total_recon 0.0 total_kl 0.0 for batch in dataloader: x batch[0].to(device) optimizer.zero_grad() x_recon, mu, logvar model(x) loss, recon_loss, kl_loss vae_loss(x, x_recon, mu, logvar, beta) loss.backward() optimizer.step() total_loss loss.item() * x.size(0) total_recon recon_loss.item() * x.size(0) total_kl kl_loss.item() * x.size(0) n len(dataloader.dataset) return total_loss / n, total_recon / n, total_kl / n def train_vae(model, dataloader, epochs, lr, beta, device): optimizer Adam(model.parameters(), lrlr) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs) for epoch in range(epochs): loss, recon, kl train_one_epoch(model, dataloader, optimizer, beta, device) scheduler.step() if (epoch 1) % 10 0: print(fEpoch {epoch 1}: loss{loss:.4f}, recon{recon:.4f}, kl{kl:.4f})在深度强化学习中这个训练循环往往不是独立运行而是与策略优化交替进行。但核心思想一致先让模型学会压缩观测并保留关键信息再让策略在压缩后的隐空间中决策。4.5 主程序与数据准备main.py使用 MNIST 数据集做演示。MNIST 虽然简单但足以验证变分推断流程是否跑通。如果你的目标是深度强化学习视觉环境可以把数据源替换成 Gym 的观测帧但代码结构不需要改动。# 文件路径variational_rl/main.py import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms from config import Config from model import VAE from trainer import train_vae def main(): cfg Config() device torch.device(cfg.device if torch.cuda.is_available() else cpu) transform transforms.Compose([ transforms.Resize(cfg.img_size), transforms.ToTensor(), ]) dataset datasets.MNIST( root./data, trainTrue, transformtransform, downloadTrue, ) dataloader DataLoader( dataset, batch_sizecfg.batch_size, shuffleTrue, num_workers2, ) model VAE(latent_dimcfg.latent_dim).to(device) train_vae( modelmodel, dataloaderdataloader, epochscfg.epochs, lrcfg.lr, betacfg.beta, devicedevice, ) # 保存模型权重 torch.save(model.state_dict(), vae_mnist.pth) print(训练完成模型已保存。) if __name__ __main__: main()4.6 运行与预期结果在项目根目录执行python main.py正常输出类似Epoch 10: loss128.3452, recon71.2381, kl57.1071 Epoch 20: loss104.2193, recon63.8842, kl40.3351 Epoch 30: loss96.7740, recon59.1128, kl37.6612 Epoch 40: loss92.3516, recon56.8234, kl35.5282 Epoch 50: loss89.5760, recon55.4193, kl34.1567由于随机种子、数据集下载情况、设备不同具体数值会有差异。只要损失在下降且 KL 项没有迅速坍缩到 0就说明训练基本正常。4.7 把 VAE 接入深度强化学习策略网络很多读者会问这个 VAE 和深度强化学习策略网络怎么衔接最简单的方式是两阶段训练先收集一批随机策略或专家策略的观测数据训练 VAE。冻结编码器参数将观测经过编码器得到隐变量 (z)然后策略网络以 (z) 作为输入。伪代码如下# 策略输入部分示例 def get_policy_obs(obs, encoder): with torch.no_grad(): mu, logvar encoder(obs) z mu # 推断时可直接使用均值 return z # 策略网络 # policy_net PolicyNetwork(input_dimlatent_dim, hidden_dim256, output_dimaction_dim)但两阶段训练的问题在于编码器是在静态数据集上训练的策略探索时遇到分布外观测编码器可能失效。更稳健的做法是端到端联合训练即策略梯度不仅回传到策略网络也回传到编码器。这时 ELBO 重构项可以看作辅助损失KL 正则项保证隐空间平滑。伯克利课程第 11 讲的内容实际上是让你理解这种联合训练中“模型学习”部分的来源而不是只把它当成一个预训练工具。5. 常见问题与排查思路5.1 隐变量后验坍塌问题现象重构损失正常下降但 KL 损失几乎为 0隐变量不携带信息。常见原因KL 项权重过高或解码器能力太强直接从输入绕过隐变量重构。排查步骤打印隐变量样本的标准差查看 KL 损失变化曲线检查解码器是否过于复杂。解决思路调低beta使用 KL 退火从 0 慢慢增加到目标值限制解码器容量在隐变量注入噪声。表格形式汇总如下问题现象常见原因解决思路KL 损失接近 0KL 权重过大降低 beta或使用 KL 退火KL 损失接近 0解码器过强减小解码器容量重构效果好但策略效果差表征丢失关键信息增加 latent_dim调整损失权重训练不稳定学习率过高降低学习率使用更平滑的调度器采样 z 不连续先验与后验差距大增加 KL 权重使用更复杂的先验5.2 训练不收敛问题现象损失函数振荡甚至上升。常见原因优化器设置不合理输入数据分布变化梯度爆炸。排查步骤检查 loss 曲线打印梯度范数尝试较小的学习率确认输入数据标准化到 0-1。解决思路使用 Adam 优化器对梯度裁剪在深度强化学习场景中注意数据来自策略分布变化建议使用经验回放构建相对稳定的训练集。5.3 重构图像模糊VAE 重构图像天然偏模糊这是由 ELBO 中的高斯似然假设决定的。深度强化学习中这种模糊通常不是致命问题因为策略需要的往往不是逐像素精确重构而是语义信息保留。如果确实需要清晰重构可以尝试两个方向使用更灵活的解码器分布比如 pixelCNN或者把重构损失换成感知损失比如基于预训练网络特征图的 MSE。5.4 显存不足问题现象CUDA out of memory。常见原因batch_size 太大图像分辨率太高编码器输出特征图过大。解决思路减小 batch_size缩小输入图像使用梯度累积检查是否缓存了过多无用的图计算图。5.5 设备不一致问题现象Expected all tensors to be on the same device。常见原因模型在 GPU输入数据在 CPU或复用了不同设备上的张量。解决思路在每个 batch 开始时显式调用.to(device)或者在定义优化器和模型后统一设备。x x.to(device) z z.to(device)不要只在初始化阶段做一次设备转换深度强化学习环境中数据可能来自不同线程或缓冲区需要养成“每批数据都确认设备”的习惯。6. 最佳实践与工程建议6.1 损失函数分项记录变分推断模型在训练时一定要分别记录重构损失和 KL 损失。只记录总损失会掩盖很多问题。比如 KL 项缓慢增长看起来总损失还在下降但重构损失已经失效。建议使用字典形式保存训练指标并定期打印或写入 TensorBoard。metrics { total_loss: total_loss, recon_loss: recon_loss, kl_loss: kl_loss, }6.2 隐变量维度设计隐变量维度不是越大越好。过高的 latent_dim 会让编码器容易忽略 KL 正则甚至直接复制输入过低的 latent_dim 则会丢失策略决策所需的关键信息。经验上可以从 16、32、64 这几个值开始尝试。在深度强化学习环境动态建模时还要考虑动作输入对隐状态转移的影响可能需要把动作和隐状态连接后输入转移网络。6.3 使用学习率调度变分推断优化对学习率比较敏感。训练初期重构损失占主导KL 正则容易被忽略训练后期KL 项可能增大并导致不稳定。使用余弦退火学习率能减少后期振荡。如果是在线训练则建议使用固定学习率并配合经验回放。6.4 保持数值稳定计算 logvar 时注意不要让网络输出过大或过小。可以在编码器输出后加clamp比如限制在[-10, 10]避免exp(logvar)溢出或消失。logvar torch.clamp(logvar, min-10, max10)此外在计算标准差std torch.exp(0.5 * logvar)时如果logvar非法梯度也会异常。建议在重参数化前加入小常数std torch.exp(0.5 * logvar) 1e-46.5 在深度强化学习中的部署建议如果你把变分推断模型用于真实环境的深度强化学习训练以下几点尤其重要编码器不应与策略网络完全共享参数否则策略更新会强烈影响表征分布造成训练不稳定。经验回放中的观测数据可能来自不同策略版本要尽量平衡数据分布否则变分推断模型会过度拟合近期策略的状态分布。在评估时应使用mu而不是重参数化采样得到的z减少随机性带来的策略波动。保存模型时建议同时保存编码器、解码器、策略网络和优化器状态方便断点续训。6.6 模型保存与恢复深度强化学习训练通常需要很久断点续训是必备功能。保存模型时建议使用 tar 格式包含多个组件torch.save({ vae_state_dict: vae.state_dict(), policy_state_dict: policy.state_dict(), optimizer_state_dict: optimizer.state_dict(), epoch: epoch, }, checkpoint.tar)恢复时checkpoint torch.load(checkpoint.tar) vae.load_state_dict(checkpoint[vae_state_dict]) optimizer.load_state_dict(checkpoint[optimizer_state_dict])很多深度强化学习项目中的不稳定问题最后都定位到“从错误检查点恢复导致隐变量空间不一致”。因此检查点和随机种子管理同样重要。6.7 测试阶段使用确定性映射变分推断模型训练完成后在策略网络推理阶段应关闭随机采样使用均值作为隐变量。这能降低策略输出的方差也有助于部署到真实环境时的稳定性。只有专门用于探索的策略才需要在推理阶段保持采样随机性。7. 总结与学习建议这一讲的核心内容可以浓缩为四条线第一变分推断解决的是“后验分布不可计算”的问题核心优化目标是 ELBO它由重构损失和 KL 正则组成。第二重参数化技巧让变分推断可以在深度网络中使用标准反向传播训练这是从理论走向工程的关键一步。第三深度强化学习中的变分推断不是孤立模块它常作为状态表征、世界模型和探索机制的基础组件。第四实际工程中要重点关注 KL 坍塌、损失分项记录、隐变量维度和推理时确定性映射。如果你是从零开始学习深度强化学习下一步建议按顺序做三件事手推一遍 ELBO 推导确保能从边际似然一步步写到最终损失表达式。在 PyTorch 中实现一个简单的 VAE并调低beta、调高beta观察重构损失和 KL 损失的变化。把训练好的 VAE 放进一个简单环境比如 FrozenLake 或 CartPole 的视觉版本用隐变量替换原始观测训练一个策略网络对比端到端训练与两阶段训练的效果差异。如果能完成这三步你对变分推断在深度强化学习中的作用就不再只停留在名词层面。后续再看 Dreamer、PlaNet、Plan2Explore 这类算法时会明显感觉它们的模型学习部分只是一个更复杂的 VAE 状态转移网络并没有跳出这一讲的核心框架。本文整理自伯克利 2026 春季深度强化学习课程第 11 讲的内容并结合工程实践补充了代码实现与排错经验。如果你在复现过程中遇到问题建议先从损失曲线和隐变量分布两个角度入手多数训练异常都能通过这两类信号快速定位。
返回列表