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.

Real frames (left) and the model's rollout from the same start and actions (right)

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
Safetensors
Model size
4.18M params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

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 play
    self-reported
    24.490
  • PSNR on changed pixels at 5 frames on Breakout, held-out random play
    self-reported
    20.700
  • PSNR on changed pixels at 15 frames on Breakout, held-out random play
    self-reported
    18.680