World Models探訪:PyTorchで実装する世界モデル学習
Ha & SchmidhuberによるWorld Modelsのアーキテクチャをnik-55のPyTorch実装(VAE、MDN-RNN、Controller)をもとに解説します。
World Models(世界モデル)は、強化学習(Reinforcement Learning)および人工知能の分野における最も重要なマイルストーンの1つです。David Ha と Jürgen Schmidhuber (2018) によって提案されたこの手法は、AIエージェントが環境の空間的・時系列的な内部表現を自律的に学習し、現実世界で行動する前に「夢の中(シミュレーション)」で訓練や意思決定を行えるようにします。
本記事では、nik-55/world-models ↗ によるシンプルかつ堅牢な PyTorch 実装を参考に、World Models の構造と仕組みを解説します。
World Modelsを構成する3つのコンポーネント#
World Models は、視覚・記憶・制御を3つのモジュールに分離しています。
視覚コンポーネント (V - VAE) ---> 記憶コンポーネント (M - MDN-RNN) ---> コントローラー (C)txt- Vision Model (V - Variational Autoencoder): 高次元の画像観測(例: 64x64x3)を低次元の潜在ベクトル (z) に圧縮します。
- Memory Model (M - Mixture Density Network RNN): 過去の観測と行動の履歴に基づき、未来の潜在状態 (z_{t+1}) の確率分布を予測します。
- Controller (C): 潜在状態 (z_t) と RNN の隠れ状態 (h_t) から行動 (a_t)(ステアリング、アクセル、ブレーキ)を決定する軽量な線形モデルです。
[!NOTE] 視覚と記憶を分離することで Controller (C) のパラメータ数を極限まで減らすことができ、勾配降下法ではなく進化戦略(CMA-ES 等)を用いて高速にポリシーを最適化できます。
PyTorchによる実装構造 (nik-55/world-models)#
nik-55/world-models ↗ リポジトリでは、OpenAI Gym の CarRacing-v0 タスクを PyTorch で実装しています。
models/vae.py: 画像観測を潜在ベクトル (z \in \mathbb{R}^{32}) に圧縮する ConvVAE。models/mdnrnn.py: LSTM と 混合ガウスモデル(GMM)を組み合わせた MDN-RNN。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 * stdpython3段階のトレーニングパイプライン#
- ランダム収集 (Rollouts): ランダムエージェントを環境で走らせ、多数の観測画像データを収集します。
- VAE & MDN-RNN の学習:
- VAE は再構成誤差と KL ダイバージェンス損失で非指導学習します。
- MDN-RNN は次の状態 (z_{t+1}) の分布を予測するように学習します。
- 夢の中での学習 (Dream Training):
- 実際の環境を呼び出さず、MDN-RNN が生成する「夢の世界」の中で Controller を訓練します。
[!TIP] 物理環境やシミュレータにアクセスせず「夢の中」でポリシーを訓練することで、学習速度を大幅に向上させることができます。
参考文献#
- Ha, D., & Schmidhuber, J. (2018). Recurrent World Models Facilitate Policy Evolution. NeurIPS.
- GitHub Repository: nik-55/world-models ↗