breakout-world-model
A diffusion world model of Atari Breakout, 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. breakout.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 https://huggingface.co/jgalego/breakout-world-model/resolve/main/breakout.py play
From Python, with breakout.py on the path:
import torch
from breakout import DEVICE, collect, load, rollout, to_float
model = load("jgalego/breakout-world-model")
frames, _, _ = collect(4, seed=0) # four real frames to start from
context = to_float(frames[:4])[None].to(DEVICE)
moves = torch.tensor([[0, 0, 0] + [1] + [2] * 15], device=DEVICE) # FIRE, then RIGHT
video = rollout(model, context, moves) # (1, 16, 3, 64, 64) in [-1, 1]
Actions: 0 NOOP, 1 FIRE, 2 RIGHT, 3 LEFT.
🏋️ 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 (600 episodes) of a random policy that holds each 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. Most of a Breakout frame never changes, so changed pixels scores only the pixels where the real or the predicted frame differs from the last real frame. That is where the ball, paddle and bricks are.
| Frames ahead | Model, all pixels (dB) | Repeat, all pixels (dB) | Model, changed pixels (dB) | Repeat, changed pixels (dB) |
|---|---|---|---|---|
| 1 | 41.74 | 37.19 | 24.49 | 11.38 |
| 5 | 38.61 | 34.87 | 20.7 | 11.13 |
| 15 | 36.32 | 33.44 | 18.68 | 10.99 |
⚠️ Limitations
- A random policy rarely clears bricks or survives long, so the model knows the early game best.
- 64×64 frames: the ball is 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
- 3
Collection including jgalego/breakout-world-model
Paper for jgalego/breakout-world-model
Evaluation results
- PSNR on changed pixels at 1 frames on Breakout, held-out random playself-reported24.490
- PSNR on changed pixels at 5 frames on Breakout, held-out random playself-reported20.700
- PSNR on changed pixels at 15 frames on Breakout, held-out random playself-reported18.680
