在基于模型的强化学习和规划任务中,世界模型承担着预测环境状态转移的职责。但一个常见的问题是:单步预测看起来还挺准,多步滚动预测却迅速崩坏。机器人想象中走出三步就撞墙,智能体在想象轨迹里规划的动作到了真实环境完全失效。这种误差累积并非偶然,它根植于模型的结构与训练方式。Dreamer系列算法中采用的RSSM(Recurrent State-Supported Model)配合KL平衡策略,是目前公认最有效的解决方案之一。本文将从问题根源讲起,逐步展开这两个关键技术的原理与实现。

一、误差累积到底是怎么发生的
想象一下自回归文本生成中的错误传播:模型输出的每一个词都会成为下一步的输入,一旦某一步出错,错误就会被放大。世界模型的情况更严重,因为它面对的是连续、高维、部分可观测的环境。
具体来说,误差累积有三条主要路径。第一是分布偏移:训练时模型看到的输入来自真实数据,推理时输入却来自自己上一步的预测输出。预测值哪怕只偏差一点点,落在训练数据从未覆盖的区域,模型在这个区域的预测质量就无从保证。第二是随机性建模失败:真实环境往往带有内在随机性,如果模型强行用确定性映射去拟合,会把噪声当成信号记住,多步滚动后噪声被反复放大。第三是复合误差的数学本质:假设每步预测误差是独立的,n步之后的累积误差量级会以接近线性甚至更快的速度增长,而如果误差之间存在相关性(实际上几乎总是存在),增长会更快。
早期的一些确定性世界模型,比如简单的RNN预测器,把这三条路径全部踩了一遍。它们既无法表达环境的不确定性,又不得不面对自反馈带来的分布偏移,长时序预测质量自然惨不忍睹。要打破这个局面,需要从状态表示的架构层面重新设计,这就是RSSM登场的原因。
二、RSSM:确定性与随机性的双轨设计
RSSM的全称是Recurrent State-Space Model,它是PlaNet提出、Dreamer系列沿用至今的核心组件。它的核心思想是:把隐状态拆成确定性部分和随机性部分,两者并行演化、互相支持。
确定性路径由GRU或LSTM驱动,负责承载长期记忆。这条路径的特点是误差不会随着随机采样而进一步放大,为整个状态序列提供了稳定的骨架。随机路径则用一个条件高斯分布(或 categorical离散分布)刻画当前状态下的不确定性,让模型能够显式地说出“我不确定下一刻会发生什么”。两部分拼接后共同决定观测重构和奖励预测。
import torch
import torch.nn as nn
class RSSM(nn.Module):
def __init__(self, deter_size=200, stoch_size=30, hidden=200, action_dim=6):
super().__init__()
self.deter_size = deter_size
self.stoch_size = stoch_size
# 确定性路径:GRU单元
self.cell = nn.GRUCell(hidden, deter_size)
# 先验:由确定性状态预测随机状态分布
self.prior_net = nn.Sequential(
nn.Linear(deter_size, hidden), nn.ELU(),
nn.Linear(hidden, 2 * stoch_size)) # 输出均值与对数方差
# 后验:由确定性状态和当前观测编码共同推断
self.post_net = nn.Sequential(
nn.Linear(deter_size + 1024, hidden), nn.ELU(),
nn.Linear(hidden, 2 * stoch_size))
def forward(self, prev_state, prev_action, embed=None):
h = self.cell(torch.cat([prev_state['stoch'], prev_action], -1),
prev_state['deter'])
prior = self.prior_net(h)
mean, log_var = prior.chunk(2, -1)
std = torch.softplus(log_var) + 0.1
if embed is not None: # 训练阶段用后验修正
post = self.post_net(torch.cat([h, embed], -1))
mean, log_var = post.chunk(2, -1)
std = torch.softplus(log_var) + 0.1
stoch = mean + std * torch.randn_like(std) # 重参数化采样
return {'deter': h, 'stoch': stoch,
'mean': mean, 'std': std}
这个双轨设计为什么能缓解误差累积?关键在于推理阶段的行为。做长时序想象时,随机状态的采样来源从后验切换为先验,而先验的输入只有确定性隐状态h。GRU的隐藏状态作为一条连续、无采样噪声的通路,保证了即使随机部分的采样出现波动,整个轨迹的宏观走势依然稳定。实验表明,纯随机状态空间模型在长时序预测中会因反复采样而退化,纯确定性模型则无法表达多模态的未来,RSSM恰好取了两者的长处。
另外值得强调的是后验分布的训练作用。训练时编码器能“偷看”当前观测,给出更准的随机状态分布,模型通过拉近先验与后验的距离,学会在没有观测的情况下做出接近真相的预测。这就自然引出了训练目标中的KL项,也引出了下一个话题:这个KL项如果直接优化,会出大问题。
三、KL平衡:被忽视的训练失衡问题
RSSM的训练损失中有一项KL散度,用于约束先验分布贴近后验分布。直觉上,直接最小化两者之间的KL散度就完事了。但实际训练中会出现一个隐蔽的失衡:KL散度的梯度会同时流向先验和后验两侧,而“移动后验去迁就先验”往往比“移动先验去学习后验”更容易被优化。
后果是什么?后验分布会主动向一个偷懒的、接近先验的分布坍缩,观测信息根本编码不进去,模型的隐状态变成几乎不依赖真实观测的空壳。表征质量崩塌,下游的重构、奖励预测、策略学习全部遭殃。这就是所谓后验主导梯度的问题。
KL平衡的解法朴素而有效:给KL项的两个方向分配不同的权重。用大的权重(典型值如0.8)驱动先验去学习后验,用小的权重(如0.1)允许后验轻微地靠近先验。用公式表达就是:
def kl_balancing_loss(post, prior, alpha_p=0.8, alpha_q=0.1):
# KL(N(post) || N(prior)) 拆成两项分别加权
kl_stop_post = kl_divergence(
dist_stop_grad(post), prior) # 梯度只流向先验
kl_stop_prior = kl_divergence(
post, dist_stop_grad(prior)) # 梯度只流向后验
return alpha_p * kl_stop_post + alpha_q * kl_stop_prior
def dist_stop_grad(d):
return d.detach() # 停止梯度传播的辅助函数
这里的技巧是分别对两侧做stop gradient,构造两个单向的KL项再加权求和。DreamerV2的消融实验显示,仅引入这一改动,就能显著提升后验分布的信息量,策略的最终得分也有可观提升。权重比例并非唯一,0.8对0.1是常见配置,具体任务上可以在0.5到1.0的区间内搜索先验侧权重。
顺带一提,KL平衡的价值不限于世界模型。任何涉及变分推断、需要约束两个分布靠近的场景,比如表征学习中的蒸馏目标,都可能遇到某一侧“躺平”的失衡问题,加权解耦梯度方向是一个通用思路。
四、实践建议与常见坑
把RSSM和KL平衡用起来时,有几个工程细节值得注意。首先是自由比特(free bits)技巧:对KL项设置一个下限阈值,低于阈值的部分不计入损失,防止模型为了压低KL而彻底放弃随机性建模。其次,DreamerV3推荐用两档离散分布(twohot编码)配合离散隐状态,连续高斯在数值稳定性上更容易翻车。再者,训练早期先验很弱,后验修正量大是正常现象,可以通过监控先验与后验之间的KL值随训练的收敛曲线来判断模型是否学到了有效的转移规律。
最后提醒一点,RSSM缓解但并未彻底消除误差累积。在做超长时序规划时,仍可结合短时程重规划、模型预测控制式的滚动优化等策略,让想象与现实定期对齐。结构设计与训练技巧的组合拳,才是应对误差累积的完整答案。
RSSMKL Balancing世界模型修改时间:2026-09-08 06:22:53