跳转至

隐空间动力学世界模型

隐空间动力学是世界模型中最经典、最适合新手入门的一条线:先把高维图像压缩成 latent state,再在 latent state 里预测未来、规划或训练策略。


1. 为什么要用隐空间?

图像太大,未来细节太多。对决策来说,很多像素都不重要:

  • 背景纹理变化通常不影响机械臂抓杯子。
  • 赛车游戏里天空云朵不影响转向。
  • 机器人只需要知道物体大概位置、速度、接触关系。

所以世界模型常常不直接预测像素,而是学习一个紧凑状态:

\[ z_t = e_\theta(o_t) \]

再预测:

\[ z_{t+1} = f_\theta(z_t, a_t) \]

1.1 好的 Latent State 应该满足什么?

一个好的 latent state 不只是压缩图像。它要服务预测和控制:

要求 含义
可预测 给定 \(z_t,a_t\) 能预测 \(z_{t+1}\)
任务相关 保留奖励、目标、接触、速度等信息
紧凑 比原始图像小得多
稳定 相邻帧 latent 不应无意义跳变
可规划 在 latent 中 rollout 的轨迹和真实世界对应

如果 latent 只保留颜色纹理,不保留物体位置,那么重建图像可能不错,但控制会很差。


2. World Models:VAE + RNN + Controller

David Ha 和 Jürgen Schmidhuber 的 World Models 是现代神经世界模型的经典起点。

结构可以概括为:

graph LR
    A[图像观测] --> B[V: VAE 编码器]
    B --> C[低维 latent z_t]
    C --> D[M: RNN/MDN-RNN]
    E[动作 a_t] --> D
    D --> F[预测未来 latent]
    F --> G[C: Controller]
    G --> H[动作]

三个模块:

模块 作用
V 把图像压缩成 latent vector
M 用 RNN 预测 latent dynamics
C 根据 latent 和 hidden state 输出动作

它的重要思想是:策略可以在世界模型“梦出来”的环境中训练,再迁移到真实环境。

2.1 VAE:把图像压缩成概率 Latent

VAE 编码器不只输出一个向量,而是输出分布参数:

\[ q_\phi(z_t\mid o_t)=\mathcal{N}(\mu_\phi(o_t),\sigma_\phi(o_t)^2) \]

然后用重参数化采样:

\[ z_t=\mu_\phi+\sigma_\phi\odot\epsilon,\quad \epsilon\sim\mathcal{N}(0,I) \]

解码器重建图像:

\[ \hat{o}_t=d_\psi(z_t) \]

VAE 损失:

\[ \mathcal{L}_{\text{VAE}} =\underbrace{\|o_t-\hat{o}_t\|^2}_{\text{重建}} +\beta\underbrace{D_{\text{KL}}(q_\phi(z_t\mid o_t)\|p(z))}_{\text{约束 latent 分布}} \]

直觉:

  • 重建项让 latent 保留图像信息。
  • KL 项让 latent 分布规整,方便采样和预测。

2.2 MDN-RNN:预测多种未来

World Models 的 M 模块使用 MDN-RNN。RNN 记住历史:

\[ h_t=\text{RNN}(h_{t-1},z_t,a_t) \]

MDN(Mixture Density Network)输出一个混合高斯分布:

\[ p(z_{t+1}\mid h_t,a_t)=\sum_{i=1}^{K}\pi_i\mathcal{N}(\mu_i,\sigma_i^2) \]

为什么要混合高斯?因为未来可能有多种。比如赛车遇到弯道,可以向左修正,也可以先减速再转向。混合分布比单个均值更能表达多未来。


3. PlaNet:从像素学习 latent dynamics,再在线规划

PlaNet 的核心是 RSSM(Recurrent State-Space Model)。它结合:

  • 确定性隐藏状态:记忆长期信息。
  • 随机 latent state:表示不确定性和多未来。

形式上:

\[ h_t = f(h_{t-1}, z_{t-1}, a_{t-1}) \]
\[ z_t \sim p(z_t \mid h_t) \]

观测、奖励从 \(h_t,z_t\) 解码出来。

3.1 RSSM 的核心结构

RSSM 把状态拆成两部分:

确定性状态 h_t:像 RNN hidden state,记住历史
随机状态 z_t:表示当前不确定状态

每一步有两个分布:

分布 来源 作用
Prior \(p(z_t\mid h_t)\) 只根据过去和动作预测当前 latent
Posterior \(q(z_t\mid h_t,o_t)\) 看到真实观测后修正当前 latent

训练时有真实观测,所以用 posterior;想象未来时没有真实观测,所以用 prior。

graph LR
    A[h_t] --> B[Prior p z_t]
    C[o_t] --> D[Encoder]
    A --> E[Posterior q z_t]
    D --> E
    E --> F[Decoder/Reward]
    E --> G[Transition to h_next]

3.2 RSSM 的训练损失

典型 RSSM 世界模型损失包含:

\[ \mathcal{L} = \mathcal{L}_{\text{obs}} + \mathcal{L}_{\text{reward}} + \mathcal{L}_{\text{continue}} + \beta \mathcal{L}_{\text{KL}} \]

其中:

损失 作用
\(\mathcal{L}_{\text{obs}}\) 从 latent 重建图像或观测
\(\mathcal{L}_{\text{reward}}\) 预测奖励
\(\mathcal{L}_{\text{continue}}\) 预测 episode 是否继续
\(\mathcal{L}_{\text{KL}}\) 让 prior 接近 posterior

KL 项通常是:

\[ D_{\text{KL}}(q(z_t\mid h_t,o_t)\|p(z_t\mid h_t)) \]

它的意思是:模型只靠历史和动作预测出来的 prior,应该接近看到真实图像后得到的 posterior。这样将来没有真实观测时,模型也能在 latent 中继续想象。

PlaNet 怎么选动作?

PlaNet 不直接训练一个策略,而是用 CEM 在模型里搜索动作序列:

  1. 从当前 latent state 出发。
  2. 随机采样很多候选动作序列。
  3. 在世界模型里 rollout。
  4. 计算每条轨迹的预测奖励。
  5. 保留最好的动作序列并迭代优化。
  6. 执行第一个动作,然后重新观测和规划。

这就是 MPC 思路。


4. Dreamer:在想象中训练策略

Dreamer 延续 RSSM,但把重点从“每步在线搜索”转向“在模型想象轨迹中训练 actor-critic”。

graph TD
    A[真实环境收集数据] --> B[训练世界模型]
    B --> C[从真实 latent 出发想象未来轨迹]
    C --> D[训练 actor]
    C --> E[训练 critic]
    D --> F[真实环境执行策略]
    F --> A

4.1 想象 rollout

从真实数据中的某个 latent state 出发:

\[ z_t, a_t, z_{t+1}, a_{t+1}, ... \]

后面的状态不来自真实环境,而来自世界模型预测。这样可以用很少真实交互生成大量训练信号。

更具体地说:

z_t 来自真实观测编码
actor 选择 a_t
world model 预测 z_{t+1}, r_t, continue_t
actor 再根据 z_{t+1} 选择 a_{t+1}
重复 H 步

这条 imagined trajectory 不需要真实环境参与,所以数据效率高。

4.2 Actor-Critic

  • actor 学会选择动作。
  • critic 学会估计未来回报。
  • 世界模型提供“可微分的未来”。

这让 Dreamer 的数据效率很高,尤其适合图像输入控制任务。

4.3 Lambda Return

Dreamer 用 critic 估计 imagined trajectory 的回报。有限 horizon 的目标可以写成:

\[ G_t = r_t + \gamma V(z_{t+1}) \]

但只看一步太短视。多步 return 又方差大。\(\lambda\)-return 在两者之间折中:

\[ G_t^\lambda = r_t + \gamma\left((1-\lambda)V(z_{t+1})+\lambda G_{t+1}^\lambda\right) \]

直觉:

  • \(\lambda=0\):主要信 critic,一步 bootstrap。
  • \(\lambda=1\):主要信 rollout 的多步奖励。
  • 中间值:稳定性和长期性折中。

4.4 Actor 和 Critic 的训练

critic 学:

\[ \mathcal{L}_{V} = \|V_\psi(z_t)-G_t^\lambda\|^2 \]

actor 学:

\[ \max_\theta \mathbb{E}[G_t^\lambda] \]

也就是让 actor 在世界模型想象出的未来里选择高回报动作。

4.5 DreamerV3

DreamerV3 的重点是通用性和稳定性。它使用一套配置解决很多不同任务,包括连续控制、Atari、Minecraft 等。它的价值不只是算法细节,而是说明世界模型可以成为跨领域 RL 的通用框架。

4.6 Dreamer 训练循环

repeat:
  1. 用当前 actor 和真实环境交互,收集数据
  2. 从 replay buffer 采样序列
  3. 训练 RSSM:重建观测、预测奖励、预测 continue、KL 正则
  4. 从真实 latent 出发,用 RSSM 想象未来
  5. 在想象轨迹上训练 critic
  6. 在想象轨迹上训练 actor

如果你能把这 6 步讲清楚,就基本理解 Dreamer 了。


5. TD-MPC 与 TD-MPC2

TD-MPC 系列走的是 latent planning + value learning 路线。

核心思想:

  • 不一定需要重建图像。
  • 学一个适合控制的 latent dynamics。
  • 用价值函数指导 MPC,避免纯 rollout 太短视。

TD-MPC2 进一步强调规模化:一个较大的 agent 可以跨多个任务、embodiment 和动作空间训练。

5.1 Decoder-Free 是什么意思?

Dreamer 通常训练观测解码器重建图像;TD-MPC 系列更强调不一定要重建图像。它直接学习对控制有用的 latent:

观测 -> latent
latent + action -> next latent
latent -> reward/value

优点:

  • 不把算力浪费在无关像素。
  • latent 更贴近控制目标。
  • 适合连续控制和规划。

风险:

  • latent 不可视化,调试更难。
  • 如果训练目标设计不好,可能丢掉关键状态。

6. 这些方法怎么比较?

方法 世界模型形式 决策方式 学习价值
World Models VAE + RNN 简单控制器/梦中训练 入门概念最清楚
PlaNet RSSM CEM/MPC 在线规划 理解 latent planning
Dreamer RSSM 想象中训练 actor-critic 现代世界模型 RL 主线
TD-MPC2 decoder-free latent model value-guided MPC 控制和规模化更强

7. 关键概念速查

概念 含义
Latent State 从观测压缩出来的内部状态
RSSM 结合确定性记忆和随机状态的动力学模型
Imagination 在世界模型中 rollout 未来轨迹
CEM 采样候选动作序列并迭代保留好样本
MPC 每次只执行规划结果的第一步,然后重新规划
Decoder-Free 不重建像素,只学习对控制有用的 latent

8. 学习建议

  1. 先读 World Models,看懂 V/M/C 三模块。
  2. 再读 PlaNet,看懂 RSSM 和 CEM 规划。
  3. 再读 Dreamer,看懂 imagined rollout 和 actor-critic。
  4. 最后看 TD-MPC2,理解“世界模型不一定要生成图像”。

9. 常见错误

现象 可能原因
重建图像很清楚,但控制差 latent 记住纹理,没记住速度/接触/奖励相关状态
imagined rollout 很快漂移 prior 没学好,KL 太弱或 horizon 太长
actor 在模型里很强,真实环境失败 actor 利用了世界模型漏洞
KL loss 崩掉 posterior/prior 不平衡,latent 被忽略或过度随机
reward 预测准但策略差 短期奖励准,长期动态不准