blog.dopana

Back

World Models(世界模型) 是强化学习(Reinforcement Learning)与人工智能领域的里程碑式工作之一。由 David Ha 和 Jürgen Schmidhuber(2018)提出,该架构允许 AI 智能体建立内部的环境空间与时间模型,从而能够在自己的“梦境(Dream)”中进行训练和决策,然后再应用到真实环境中。

本文基于开源项目 nik-55/world-models 的 PyTorch 重构版本,深入剖析 World Models 的核心架构与实现。

World Models 的 3 大核心组件#

World Models 将视觉感知、记忆与决策解耦为三个模块:

视觉组件 (V - VAE) ---> 记忆组件 (M - MDN-RNN) ---> 控制器 (C)
txt
  1. 视觉模型 (V - Variational Autoencoder): 将高维图像观测(例如 64x64x3)压缩为低维的潜在向量 (z)。
  2. 记忆模型 (M - Mixture Density Network RNN): 基于历史观测与动作,预测未来的潜在状态 (z_{t+1}) 概率分布。
  3. 控制器 (C): 一个极简的线性模型,仅根据潜在状态 (z_t) 与 RNN 隐状态 (h_t) 来输出动作 (a_t)(转向、油门、刹车)。

[!NOTE] 将视觉与记忆解耦后,Controller (C) 的参数量极小,使得我们可以轻松采用进化策略(Evolution Strategies / CMA-ES)对其进行优化,无需在整网中进行端到端反向传播。

PyTorch 代码架构 (nik-55/world-models)#

nik-55/world-models 项目使用 PyTorch 实现了 OpenAI Gym 中的 CarRacing-v0 赛车环境:

  • models/vae.py: 使用 ConvVAE 将图像压缩为潜在向量 (z \in \mathbb{R}^{32})。
  • models/mdnrnn.py: 结合 LSTM 与混合高斯模型(GMM)预测下一状态。
  • models/controller.py: 简单线性控制器。
  • trainvae.py, trainmdn.py, traincontroller.py: 各阶段独立的训练脚本。

三阶段训练流程#

  1. 随机数据采集 (Rollouts): 运行随机策略收集数万张 CarRacing 环境的观察图像。
  2. 训练 VAE 与 MDN-RNN:
    • VAE 采用无监督方式训练(Reconstruction + KL 散度损失)。
    • MDN-RNN 训练预测下一个 (z_{t+1})。
  3. 梦境训练 (Dream Training):
    • 控制器完全在 MDN-RNN 生成的虚幻梦境中训练,无需交互物理环境。

[!TIP] 完全在梦境模型中训练 Policy 可以大幅提升迭代速度,显著减少与实际环境交互的开销。

参考资料#

  • Ha, D., & Schmidhuber, J. (2018). Recurrent World Models Facilitate Policy Evolution. NeurIPS.
  • GitHub Repository: nik-55/world-models