外观
4.5 DreamerV2 与 DreamerV3:离散状态与跨任务稳健性
DreamerV1 使用连续高斯随机状态。DreamerV2 把这部分改成多组分类变量,并针对 Atari 调整了训练方法;DreamerV3 进一步处理不同任务间奖励与价值尺度差异,目标是减少逐任务调参。

图 4.5-1:DreamerV3 的实验环境拼图同时呈现控制、Atari、DMLab、ProcGen 与 Minecraft,直观看到“一套配置跨领域”的实际任务跨度。 出处:Danijar Hafner;Jurgis Pasukonis;Jimmy Ba;Timothy Lillicrap,Mastering Diverse Domains through World Models(2023),Figure 2。
本节讨论 DreamerV2 [Hafner et al., 2021] 和 DreamerV3 [Hafner et al., 2023]。DreamerV2 使用多组分类分布组成离散随机状态,并在 Atari 上取得较强结果。DreamerV3 通过 symlog、two-hot 分布回归、回报尺度归一化等设计,提高同一套超参数跨领域使用的稳定性;原文报告覆盖 8 个领域、150 多项任务。
我们将从最基础的概率论概念起步,逐步推导出离散空间的梯度反向传播技巧,并深入理解旨在解决值域缩放问题的对数变换技巧。
连续空间的局限性与离散表征的崛起

图 4.5-2:VQ-VAE 的原图把 ImageNet 原图与离散码本重建并排,给出分类潜变量能够保留视觉结构的早期证据。 出处:Aaron van den Oord;Oriol Vinyals;Koray Kavukcuoglu,Neural Discrete Representation Learning(2017),Figure 2。
DreamerV1 用高斯分布表示 RSSM 的随机状态,并通过重参数化获得路径梯度。DreamerV2 的实验发现,分组分类状态在 Atari 训练中更有效;这是经验设计选择,不意味着连续表示原则上不能描述离散概念。
分组分类变量让每组状态从有限类别中选择,并通过多组组合形成较大的表示空间。相比对角高斯,它改变了分布形状、采样方式与 KL 计算;实际收益来自这些建模与优化差异,而不是预先规定每个类别对应“猫”“狗”等人类语义。
具体来说,DreamerV2 用离散的分类分布表示随机状态。网络输出一个形状为

图 4.5-3:DreamerV2 世界模型图把分类随机状态的后验、先验、奖励预测和图像重建沿时间展开。 出处:Danijar Hafner;Timothy Lillicrap;Mohammad Norouzi;Jimmy Ba,Mastering Atari with Discrete World Models(2021),Figure 2。
分类分布的数学表达
假设我们有一个离散随机变量
当我们对
直接从分类分布采样得到离散索引,普通计算图没有从该样本回到 logits 的路径导数。因此不能像高斯重参数化那样直接把下游梯度传给采样器参数。
解决梯度回传:直通估计器 (Straight-Through Estimator)
在微积分中,如果一个函数的导数处处为零,我们就无法通过它来回传任何有用的误差信号。这正是深度学习在面对离散变量时长期以来的痛点。
直通估计器(Straight-Through Estimator, STE)采用不同的前向与反向规则:前向使用离散样本,反向则用连续概率的梯度近似离散采样的梯度。
说明
STE 不是离散采样真实导数的无偏估计,而是一种有偏近似;它通常比基于采样回报的似然比估计方差低,但效果取决于任务与参数化。
让我们用数学语言更精确地描述它。假设神经网络输出了一组未归一化的对数概率(Logits),记为
根据
如果我们在反向传播时切断(Detach)
- 在前向计算时,由于
和 相互抵消,计算图的实际值为 ,保证了输出是严格离散的独热向量。 - 在反向计算时,梯度流经
时,被截断的部分不会产生梯度,所有的梯度都会流向剩下的 。而 是通过可导的 Softmax 函数得到的。
我们将通过代码展示这一过程。
python
import torch
import torch.nn.functional as F
import torch.distributions as D
def straight_through_sample(logits, num_classes):
"""
(实现直通估计器的独热采样)
"""
# [1] 计算连续的概率分布
probs = F.softmax(logits, dim=-1)
# [2] 根据概率进行真实的分类分布采样
dist = D.Categorical(probs=probs)
indices = dist.sample()
# [3] 将采样结果转换为独热向量
z_one_hot = F.one_hot(indices, num_classes=num_classes).float()
# [4] 应用直通估计器技巧 (z_one_hot - probs).detach() + probs
# 前向传播时等于 z_one_hot,反向传播时等价于 probs
z_sample = z_one_hot + probs - probs.detach()
return z_sample若使用
KL 散度平衡 (KL Balancing)

图 4.5-4:DreamerV2 消融曲线把离散状态、KL balancing 和 actor 梯度选择对 Atari 表现的影响分开。 出处:Danijar Hafner;Timothy Lillicrap;Mohammad Norouzi;Jimmy Ba,Mastering Atari with Discrete World Models(2021),Figure 5。
在变分自编码器架构中,我们需要计算后验分布(由编码器结合当前观测得出)与先验分布(由动力学模型基于过去状态预测得出)之间的 KL 散度(Kullback-Leibler divergence)。在优化过程中,我们希望这两个分布相互靠近。
对于一般的损失函数:
在这个极小化过程中,存在两个变量:先验
DreamerV2 使用 KL 平衡(KL Balancing)分别控制先验与后验从 KL 项接收的梯度强度:让先验更积极地拟合停止梯度后的后验,同时减弱后验仅为迁就先验而改变的趋势。
这一策略通过对 KL 散度进行切断分离(Stop-gradient,在数学中通常表示为
当
DreamerV3:走向通用与健壮性
DreamerV3 保留离散世界模型,并重点处理跨任务数值尺度差异与训练稳定性。
在传统的强化学习研究中,对于不同的任务,研究人员往往需要针对性地调节学习率、奖励缩放系数以及神经网络的初始化权重。例如,在某些雅达利游戏中,得分可能是以万为单位的(如 10000 分),而在其他连续控制任务中,奖励可能被严格限制在
奖励尺度增大时,价值目标也会增大,平方误差会让大残差主导梯度。DreamerV3 组合多项尺度处理与正则设计,使同一组超参数能覆盖论文评测中的多种领域。
对称对数变换 (Symlog Transformation)
为了压缩具有巨大范围的实数值,一个直观的想法是使用对数函数
- 它未定义于负数区域,而环境的奖励完全可能是负数。
- 当
接近 0 时, 会趋向于负无穷,这在数值计算中是不可接受的。
为了应对这两个问题,DreamerV3 引入了对称对数(Symlog)函数。对于任意实数
这个变换包含两步:
使对数输入至少为 1,因此 时输出为 0。 是符号函数。它保留了原始数值的符号。
相应的逆变换(Symexp)为:
symlog 保留符号,并把大幅值压缩到对数尺度。例如
双热编码 (Two-Hot Encoding)
DreamerV3 还把标量预测写成有限支撑上的分布回归,并用 two-hot 目标训练。
给定一个目标值
对于任意目标值
令相邻点之间距离为
赋予左侧桶

图 4.5-5:目标越靠近某个桶,该桶获得的概率越大;两个权重之和为 1,并且桶中心的加权和严格等于原目标 y。
其余桶的概率为 0。若
代码实现:健壮的世界模型组件
下面实现 symlog/symexp 与 two-hot 编码。它们是 DreamerV3 数值处理的一部分,不构成完整算法。
python
class SymlogTransform:
"""
(对称对数变换及其逆变换)
"""
@staticmethod
def forward(x):
return torch.sign(x) * torch.log1p(torch.abs(x))
@staticmethod
def inverse(y):
return torch.sign(y) * torch.expm1(torch.abs(y))
def two_hot_encode(target, min_val=-20.0, max_val=20.0, num_bins=255):
"""
(将连续目标值转化为双热编码分布)
"""
# [1] 构建等距的离散网格 (Bins)
bins = torch.linspace(min_val, max_val, num_bins, device=target.device)
# [2] 限制目标值的范围以防越界
target = torch.clamp(target, min_val, max_val)
# [3] 计算目标值在网格上的相对位置
# 这里通过减去最小值并除以网格间距得到浮点索引
step = (max_val - min_val) / (num_bins - 1)
index_float = (target - min_val) / step
# [4] 找到相邻的左右两个桶的整数索引
left_idx = torch.floor(index_float).long()
right_idx = torch.ceil(index_float).long()
# [5] 计算分配给右侧桶的权重(距离左侧桶越远,右侧权重越大)
weight_right = index_float - left_idx.float()
weight_left = 1.0 - weight_right
# 处理恰好落在网格点上的情况:全部权重放在该网格点
mask = (left_idx == right_idx)
weight_left[mask] = 1.0
weight_right[mask] = 0.0
# [6] 将权重散布到一个全零张量中形成 Two-hot 分布
# 获取 target 的 batch 大小
batch_shape = target.shape
two_hot = torch.zeros(*batch_shape, num_bins, device=target.device)
# 使用 scatter_ 填充概率值
# 注意:需要增加一个维度以便于 scatter 操作
two_hot.scatter_(-1, left_idx.unsqueeze(-1), weight_left.unsqueeze(-1))
two_hot.scatter_add_(-1, right_idx.unsqueeze(-1), weight_right.unsqueeze(-1))
return two_hot小结
- DreamerV2 用多组分类分布替换连续高斯随机状态,并用 STE 提供有偏的近似梯度。
- KL 平衡分别调节先验与后验从 KL 项接收的梯度,避免两者以同样强度相互迁就。
- DreamerV3 用 symlog、two-hot 分布回归和归一化等组件减小跨任务尺度差异。这些设计提高了论文评测中的稳健性,但不消除模型误差或所有优化风险。

