Khám Phá World Models: Học Mô Hình Thế Giới Với PyTorch
Tìm hiểu kiến trúc World Models của Ha & Schmidhuber thông qua bản cài đặt PyTorch tinh gọn của nik-55 (VAE, MDN-RNN, Controller).
World Models là một trong những cột mốc quan trọng bậc nhất trong lĩnh vực Học máy Tăng cường (Reinforcement Learning - RL) và Trí tuệ Nhân tạo. Được giới thiệu bởi David Ha và Jürgen Schmidhuber (2018), kiến trúc này cho phép tác nhân AI tự học một mô hình ảo đại diện cho môi trường xung quanh, từ đó tập luyện và đưa ra quyết định bên trong “mơ ước” (dream environment) trước khi hành động ngoài thực tế.
Trong bài viết này, chúng ta sẽ tìm hiểu về World Models thông qua bản triển khai PyTorch hiện đại và tinh gọn từ repository nik-55/world-models ↗.
3 Thành Phần Cốt Lõi Của World Models#
Hệ thống World Model bao gồm ba module chính được kết nối nhịp nhàng:
Visual Component (V - VAE) ---> Memory Component (M - MDN-RNN) ---> Controller (C)txt- Vision Model (V - Variational Autoencoder): Nén khung hình quan sát góc nhìn rộng (VD: 64x64x3) thành một vector không gian tiềm ẩn (latent vector (z)) có kích thước nhỏ gọn.
- Memory Model (M - Mixture Density Network RNN): Dự đoán trạng thái tiềm ẩn tương lai (z_{t+1}) dựa trên lịch sử quan sát và hành động quá quứ.
- Controller (C): Một mô hình tuyến tính đơn giản chịu trách nhiệm chọn lựa hành động (a_t) dựa trên (z_t) và trạng thái ẩn (h_t) của RNN.
[!NOTE] Việc phân tách riêng phần quan sát thị giác (V) và trí nhớ (M) giúp Controller (C) vô cùng nhỏ gọn, có thể dễ dàng tối ưu hóa bằng các thuật toán tiến hóa (Evolution Strategies - ES) hoặc CMA-ES mà không cần lan truyền ngược (backpropagation) toàn bộ mạng.
Cài Đặt và Cấu Trúc Mã Nguồn PyTorch (nik-55/world-models)#
Dự án nik-55/world-models ↗ tái hiện lại bài toán đua xe CarRacing-v0 trong OpenAI Gym sử dụng PyTorch với cấu trúc vô cùng sáng tỏ:
models/vae.py: Kiến trúc ConvVAE nén quan sát hình ảnh thành vector tiềm ẩn (z \in \mathbb{R}^{32}).models/mdnrnn.py: Kết hợp LSTM với Mixture Density Network (MDN) để dự đoán phân phối xác suất của trạng thái tiềm ẩn tiếp theo.models/controller.py: Mạng tuyến tính quyết định tay lái (steering), ga (gas), và phanh (brake).trainvae.py,trainmdn.py,traincontroller.py: Tách biệt hoàn toàn các giai đoạn huấn luyện.
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 * stdpythonQuy Trình Huấn Luyện 3 Giai Đoạn#
- Thu thập dữ liệu ngẫu nhiên (Rollouts): Cho một tác nhân hành động ngẫu nhiên để thu thập hàng nghìn khung hình quan sát trong môi trường CarRacing.
- Huấn luyện VAE & MDN-RNN:
- VAE được huấn luyện unsupervised bằng hàm mất mát ELU (Reconstruction + KL Divergence).
- MDN-RNN được huấn luyện để dự đoán (z_{t+1}) với hỗn hợp Gauss (Gaussian Mixture Model).
- Tối ưu Controller trong Môi Trường Mơ (Dream Training):
- Controller tương tác trực tiếp với MDN-RNN mà không cần mở môi trường Gym thực tế.
[!TIP] Việc tập luyện bên trong môi trường mơ mộng (Mental Simulation) giúp tăng tốc độ huấn luyện gấp hàng trăm lần so với tương tác trực tiếp với môi trường vật lý.
Tài liệu tham khảo#
- Ha, D., & Schmidhuber, J. (2018). Recurrent World Models Facilitate Policy Evolution. NeurIPS.
- GitHub Repository: nik-55/world-models ↗