探索 World Models:基于 PyTorch 实现的世界模型
深入解析 Ha & Schmidhuber 提出的 World Models 架构,基于 nik-55 的 PyTorch 实现(VAE、MDN-RNN 与 Controller)。
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- 视觉模型 (V - Variational Autoencoder): 将高维图像观测(例如 64x64x3)压缩为低维的潜在向量 (z)。
- 记忆模型 (M - Mixture Density Network RNN): 基于历史观测与动作,预测未来的潜在状态 (z_{t+1}) 概率分布。
- 控制器 (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: 各阶段独立的训练脚本。
models/vae.py
import torch
import torch.nn as nn
import torch.nn.functional as F
class VAE(nn.Module):
def __init__(self, img_channels=3, latent_size=32):
super(VAE, self).__init__()
self.latent_size = latent_size
# Encoder
self.encoder = nn.Sequential(
nn.Conv2d(img_channels, 32, 4, stride=2),
nn.ReLU(),
nn.Conv2d(32, 64, 4, stride=2),
nn.ReLU(),
nn.Conv2d(64, 128, 4, stride=2),
nn.ReLU(),
nn.Conv2d(128, 256, 4, stride=2),
nn.ReLU()
)
self.fc_mu = nn.Linear(256 * 2 * 2, latent_size)
self.fc_logvar = nn.Linear(256 * 2 * 2, latent_size)
# Decoder
self.fc_decoder = nn.Linear(latent_size, 1024)
self.decoder = nn.Sequential(
nn.ConvTranspose2d(1024, 128, 5, stride=2),
nn.ReLU(),
nn.ConvTranspose2d(128, 64, 5, stride=2),
nn.ReLU(),
nn.ConvTranspose2d(64, 32, 6, stride=2),
nn.ReLU(),
nn.ConvTranspose2d(32, img_channels, 6, stride=2),
nn.Sigmoid()
)
def reparameterize(self, mu, logvar):
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std)
return mu + eps * stdpython三阶段训练流程#
- 随机数据采集 (Rollouts): 运行随机策略收集数万张 CarRacing 环境的观察图像。
- 训练 VAE 与 MDN-RNN:
- VAE 采用无监督方式训练(Reconstruction + KL 散度损失)。
- MDN-RNN 训练预测下一个 (z_{t+1})。
- 梦境训练 (Dream Training):
- 控制器完全在 MDN-RNN 生成的虚幻梦境中训练,无需交互物理环境。
[!TIP] 完全在梦境模型中训练 Policy 可以大幅提升迭代速度,显著减少与实际环境交互的开销。
参考资料#
- Ha, D., & Schmidhuber, J. (2018). Recurrent World Models Facilitate Policy Evolution. NeurIPS.
- GitHub Repository: nik-55/world-models ↗