space-invaders-world-model
A diffusion world model of Atari Space Invaders, trained from scratch. Given the last 4 frames and actions, a small U-Net denoises the next 64×64 frame; fed its own frames, it generates a video that your actions steer. After DIAMOND. atari.py in this repo collects the data, trains, plays and doubles as the loader.
Left: real game. Right: the model, from the same four starting frames and the same actions, never seeing a real frame again.
🚀 Usage
Play it in the browser. Each button press generates one frame:
uv run "$(hf download jgalego/space-invaders-world-model atari.py --quiet)" play --game SpaceInvaders
From Python, with atari.py on the path:
import torch
from atari import DEVICE, collect, load, rollout, to_float
model = load("jgalego/space-invaders-world-model")
frames, _, _ = collect(4, seed=0, game="SpaceInvaders") # four real frames to start from
context = to_float(frames[:4])[None].to(DEVICE)
moves = torch.tensor([[0, 0, 0] + [1] * 16], device=DEVICE) # FIRE for 16 frames
video = rollout(model, context, moves) # (1, 16, 3, 64, 64) in [-1, 1]
Actions: 0 NOOP, 1 FIRE, 2 RIGHT, 3 LEFT, 4 RIGHTFIRE, 5 LEFTFIRE.
🏋️ Training
| Model | U-Net, channels 64, 128, 256, 4.2M parameters; noise level and the last 4 actions modulate every block |
| Diffusion | EDM preconditioning, σ_data 0.5, log-normal training noise; 3 Euler steps per frame from σ=5.0 |
| Data | 200,000 frames (382 episodes) of a random policy that holds each uniformly random action for 1 to 6 steps; frameskip 4, no sticky actions |
| Steps | 50,000, batch size 64 |
| Optimizer | AdamW, lr 0.0001, gradient clipping 1.0 |
| Hardware | NVIDIA A10G, 55.9 min |
📊 Results
Rollouts on 256 windows from episodes whose seeds never appear in training. The model gets four real frames and the real actions, then only its own frames. Repeat keeps showing the last real frame. Much of an Atari frame never changes, so changed pixels scores only the pixels where the real or the predicted frame differs from the last real frame.
| Frames ahead | Model, all pixels (dB) | Repeat, all pixels (dB) | Model, changed pixels (dB) | Repeat, changed pixels (dB) |
|---|---|---|---|---|
| 1 | 35.43 | 36.56 | 18.17 | 16.0 |
| 5 | 31.07 | 30.94 | 18.35 | 15.9 |
| 15 | 27.33 | 26.77 | 16.64 | 14.54 |
⚠️ Limitations
- A random policy does not get far, so the model knows the early game best.
- 64×64 frames: small sprites are a pixel or two and can blur or vanish over long rollouts.
- Errors compound: every generated frame becomes input for the next.
- Downloads last month
- -
Collection including jgalego/space-invaders-world-model
Paper for jgalego/space-invaders-world-model
Evaluation results
- PSNR on changed pixels at 1 frames on Space Invaders, held-out random playself-reported18.170
- PSNR on changed pixels at 5 frames on Space Invaders, held-out random playself-reported18.350
- PSNR on changed pixels at 15 frames on Space Invaders, held-out random playself-reported16.640
