# /// script # requires-python = ">=3.12" # dependencies = [ # "ale-py", # "gradio", # "gymnasium", # "huggingface_hub", # "jinja2", # "numpy", # "pillow", # "safetensors", # "torch", # ] # /// """Train a diffusion world model of an Atari game (Space Invaders, Ms. Pac-Man) and play it. After DIAMOND (https://arxiv.org/abs/2405.12399): a small U-Net denoises the next 64x64 frame from the last four frames and actions, with EDM preconditioning and a few Euler steps at inference. Fed its own frames, it generates a video that the actions steer. """ # Torch modules only expose forward and keep one attribute per layer. # pylint: disable=too-few-public-methods,too-many-instance-attributes import argparse import json import math import os import shutil import time from dataclasses import asdict, dataclass from functools import partial from pathlib import Path import ale_py import gradio as gr import gymnasium as gym import numpy as np import torch import torch.nn.functional as F from huggingface_hub import EvalResult, HfApi, ModelCard, ModelCardData, snapshot_download from huggingface_hub.errors import RepositoryNotFoundError from PIL import Image from safetensors.torch import load_file, save_file from torch import nn # ALE name: (display name, slug for the model repo) GAMES = { "SpaceInvaders": ("Space Invaders", "space-invaders"), "MsPacman": ("Ms. Pac-Man", "ms-pacman"), } DEVICE = "cuda" if torch.cuda.is_available() else "cpu" HERE = Path(__file__).parent SIZE = 64 EVAL_SEED = 1_000_000 # evaluation episodes never share a seed with training episodes SIGMA_DATA = 0.5 HORIZONS = (1, 5, 15) # frames ahead; at frameskip 4, 15 frames is one second of play GIF_FRAMES = 60 gym.register_envs(ale_py) @dataclass class Config: """Architecture and sampler.""" game: str = "SpaceInvaders" actions: tuple = () # action names, in ALE order context: int = 4 channels: tuple = (64, 128, 256) steps: int = 3 # Euler steps per frame sigma_min: float = 2e-3 sigma_max: float = 5.0 rho: float = 7.0 # Data def repo_for(game): """The model repo of a game.""" return f"jgalego/{GAMES[game][1]}-world-model" def make_env(game): """The ALE environment of a game: minimal action set, frameskip 4, no sticky actions.""" return gym.make(f"ALE/{game}-v5", frameskip=4, repeat_action_probability=0.0) def action_names(game): """The game's action names, in ALE order.""" env = make_env(game) names = tuple(env.unwrapped.get_action_meanings()) env.close() return names def collect(frames, seed, game): """Frames (n, 3, 64, 64) uint8, actions (n,) and episode ids (n,) from a random policy. The policy holds each uniformly random action for 1 to 6 steps. actions[i] turns frame i into frame i + 1; the last frame of an episode gets a NOOP that no window uses. """ env = make_env(game) rng = np.random.default_rng(seed) obs, episode, action, hold = env.reset(seed=seed)[0], 0, 0, 0 out = {"frames": [], "actions": [], "episodes": []} while len(out["frames"]) < frames: if hold == 0: action, hold = int(rng.integers(env.action_space.n)), int(rng.integers(1, 7)) hold -= 1 for key, value in zip(out, (resize(obs), action, episode)): out[key].append(value) obs, _, terminated, truncated, _ = env.step(action) if terminated or truncated: for key, value in zip(out, (resize(obs), 0, episode)): out[key].append(value) obs, episode = env.reset()[0], episode + 1 env.close() return ( torch.from_numpy(np.stack(out["frames"])), torch.tensor(out["actions"]), torch.tensor(out["episodes"]), ) def resize(obs): """A 210x160 RGB frame as a 3x64x64 uint8 array.""" return np.asarray(Image.fromarray(obs).resize((SIZE, SIZE), Image.BOX)).transpose(2, 0, 1) def starts(episodes, length): """Indices i where frames i to i + length - 1 belong to one episode.""" return torch.nonzero(episodes[: -length + 1] == episodes[length - 1 :])[:, 0] def to_float(frames): """uint8 frames to [-1, 1].""" return frames.float() / 127.5 - 1 def window(frames, actions, index, context): """Context frames (b, c, 3, 64, 64), their actions (b, c) and the next frame (b, 3, 64, 64).""" rows = index[:, None] + torch.arange(context + 1, device=index.device) clip = to_float(frames[rows]) return clip[:, :-1], actions[rows[:, :-1]], clip[:, -1] # Model class Block(nn.Module): """Residual conv block with scale and shift from the conditioning vector.""" def __init__(self, inputs, outputs, cond): super().__init__() self.norm1, self.norm2 = nn.GroupNorm(32, inputs), nn.GroupNorm(32, outputs) self.conv1 = nn.Conv2d(inputs, outputs, 3, padding=1) self.conv2 = nn.Conv2d(outputs, outputs, 3, padding=1) self.film = nn.Linear(cond, 2 * outputs) self.skip = nn.Conv2d(inputs, outputs, 1) if inputs != outputs else nn.Identity() def forward(self, x, cond): """x (b, inputs, h, w) to (b, outputs, h, w).""" hidden = self.conv1(F.silu(self.norm1(x))) scale, shift = self.film(cond)[:, :, None, None].chunk(2, dim=1) hidden = self.conv2(F.silu(self.norm2(hidden) * (1 + scale) + shift)) return self.skip(x) + hidden class Denoiser(nn.Module): """U-Net over the noisy next frame stacked with the context frames.""" def __init__(self, config, cond=256): super().__init__() self.config = config width = list(config.channels) self.action = nn.Embedding(len(config.actions), 64) self.embed = nn.Sequential( nn.Linear(64 * (config.context + 1), cond), nn.SiLU(), nn.Linear(cond, cond) ) self.stem = nn.Conv2d(3 * (config.context + 1), width[0], 3, padding=1) self.down = nn.ModuleList(Block(a, b, cond) for a, b in zip(width[:1] + width, width)) self.middle = Block(width[-1], width[-1], cond) self.up = nn.ModuleList( Block(2 * width[i], width[max(i - 1, 0)], cond) for i in range(len(width)) ) self.out = nn.Sequential( nn.GroupNorm(32, width[0]), nn.SiLU(), nn.Conv2d(width[0], 3, 3, padding=1) ) def forward(self, noisy, noise, context, actions): """Raw network output for the preconditioned input; `noise` is log(sigma) / 4.""" freqs = torch.exp(torch.linspace(0, math.log(1000), 32, device=noise.device)) angles = noise[:, None] * freqs cond = torch.cat([angles.sin(), angles.cos(), self.action(actions).flatten(1)], dim=1) cond = self.embed(cond) x = self.stem(torch.cat([noisy, context.flatten(1, 2)], dim=1)) skips = [] for level, block in enumerate(self.down): x = block(x, cond) skips.append(x) if level < len(self.down) - 1: x = F.avg_pool2d(x, 2) x = self.middle(x, cond) for block, skip in zip(reversed(self.up), reversed(skips)): if x.shape[-1] != skip.shape[-1]: x = F.interpolate(x, scale_factor=2, mode="nearest") x = block(torch.cat([x, skip], dim=1), cond) return self.out(x) def denoise(self, noisy, sigma, context, actions): """EDM denoiser: the clean frame estimate at noise level sigma (b,).""" c_skip, c_out, c_in = precondition(sigma) return c_skip * noisy + c_out * self(c_in * noisy, sigma.log() / 4, context, actions) def precondition(sigma): """EDM skip, output and input scales for noise levels sigma (b,).""" sigma = sigma[:, None, None, None] total = sigma**2 + SIGMA_DATA**2 return SIGMA_DATA**2 / total, sigma * SIGMA_DATA / total.sqrt(), 1 / total.sqrt() @torch.no_grad() def sample(model, context, actions): """The next frame in [-1, 1] by Euler steps on the Karras noise schedule.""" config = model.config ramp = torch.linspace(0, 1, config.steps, device=context.device) low, high = config.sigma_min ** (1 / config.rho), config.sigma_max ** (1 / config.rho) sigmas = (high + ramp * (low - high)) ** config.rho x = torch.randn_like(context[:, -1]) * sigmas[0] for i, sigma in enumerate(sigmas): denoised = model.denoise(x, sigma.expand(len(x)), context, actions) following = sigmas[i + 1] if i + 1 < len(sigmas) else 0.0 x = denoised + (x - denoised) * following / sigma return x.clamp(-1, 1) @torch.no_grad() def rollout(model, context, actions): """Generate frames from context (b, c, 3, 64, 64) with actions (b, c - 1 + steps). Returns (b, steps, 3, 64, 64) in [-1, 1]. """ size = model.config.context frames = [] for step in range(actions.shape[1] - size + 1): with torch.autocast(DEVICE, torch.bfloat16, enabled=DEVICE == "cuda"): frame = sample(model, context, actions[:, step : step + size]).float() frames.append(frame) context = torch.cat([context[:, 1:], frame[:, None]], dim=1) return torch.stack(frames, dim=1) # Evaluation def psnr(predicted, truth, mask): """PSNR in dB per step for frames (b, t, 3, h, w) in [-1, 1], pooled over masked pixels.""" error = ((predicted - truth) / 2).pow(2) * mask mse = error.sum((0, 2, 3, 4)) / mask.sum((0, 2, 3, 4)).clamp_min(1) return 10 * torch.log10(1 / mse.clamp_min(1e-10)) def evaluate_model(model, frames, actions, episodes, count): """PSNR of rollouts against real frames, with a repeat-the-last-frame baseline. Much of an Atari frame is static, so PSNR is also given over changed pixels: those where the real or the predicted frame differs from the last context frame. """ size, horizon = model.config.context, max(HORIZONS) valid = starts(episodes, size + horizon) index = valid[torch.linspace(0, len(valid) - 1, min(count, len(valid))).long()] rows = index[:, None] + torch.arange(size + horizon) clip = to_float(frames[rows].to(DEVICE)) last = clip[:, size - 1 : size].expand_as(clip[:, size:]) guesses = { "model": rollout(model, clip[:, :size], actions[rows[:, : size - 1 + horizon]].to(DEVICE)), "repeat": last, } return { "n": len(index), **{name: score(guess, clip[:, size:], last) for name, guess in guesses.items()}, } def score(guess, truth, last): """PSNR over all pixels and over changed pixels at each horizon.""" changed = ((truth - last).abs() > 0.1) | ((guess - last).abs() > 0.1) parts = {"all": torch.ones_like(truth), "changed": changed} return { part: {f"{h}": round(psnr(guess, truth, mask)[h - 1].item(), 2) for h in HORIZONS} for part, mask in parts.items() } def gif(model, frames, actions, episodes, path): """Real frames next to a rollout from the same start and actions, 4x upscaled.""" size = model.config.context valid = starts(episodes, size + GIF_FRAMES) index = valid[len(valid) // 2] # mid-play, past intros where nothing moves clip = to_float(frames[index : index + size + GIF_FRAMES].to(DEVICE)) moves = actions[index : index + size - 1 + GIF_FRAMES].to(DEVICE) predicted = rollout(model, clip[None, :size], moves[None])[0] pairs = torch.cat([clip[size:], predicted], dim=-1) images = [picture(pair, 4) for pair in pairs] images[0].save(path, save_all=True, append_images=images[1:], duration=66, loop=0) def picture(frame, scale): """A [-1, 1] frame (3, h, w) as an upscaled PIL image.""" pixels = ((frame + 1) * 127.5).round().clamp(0, 255).byte().permute(1, 2, 0).cpu().numpy() image = Image.fromarray(pixels) return image.resize((image.width * scale, image.height * scale), Image.NEAREST) # Training def edm_loss(model, context, actions, clean, generator): """EDM loss on the next frame, with noise levels from a log-normal.""" sigma = (torch.randn(len(clean), device=DEVICE, generator=generator) * 1.2 - 0.4).exp() c_skip, c_out, c_in = precondition(sigma) noisy = clean + sigma[:, None, None, None] * torch.randn_like(clean) with torch.autocast(DEVICE, torch.bfloat16, enabled=DEVICE == "cuda"): output = model(c_in * noisy, sigma.log() / 4, context, actions) return (output.float() - (clean - c_skip * noisy) / c_out).pow(2).mean() def save_state(folder, state): """Write resumable training state to folder/state.pt, replacing it in one step.""" folder = Path(folder) folder.mkdir(parents=True, exist_ok=True) torch.save(state, folder / "state.pt.tmp") os.replace(folder / "state.pt.tmp", folder / "state.pt") def training_state(step, model, optimizer, generator): """Everything needed to continue training exactly where it stopped.""" return { "step": step, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "generator": generator.get_state(), "rng": torch.get_rng_state(), "cuda_rng": torch.cuda.get_rng_state() if DEVICE == "cuda" else None, } def resume(folder, model, optimizer, generator): """Restore training state from folder/state.pt if there is one; return its step, else 0.""" if not folder or not (Path(folder) / "state.pt").exists(): return 0 state = torch.load(Path(folder) / "state.pt", map_location="cpu") model.load_state_dict(state["model"]) optimizer.load_state_dict(state["optimizer"]) generator.set_state(state["generator"]) torch.set_rng_state(state["rng"]) if DEVICE == "cuda": torch.cuda.set_rng_state(state["cuda_rng"]) print(f"resumed at step {state['step']}", flush=True) return state["step"] def fit(model, frames, actions, episodes, args): """AdamW on random windows of the collected play, resuming from --checkpoint if set.""" optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-2) valid = starts(episodes, model.config.context + 1).to(DEVICE) generator = torch.Generator(DEVICE).manual_seed(args.seed) first = resume(args.checkpoint, model, optimizer, generator) start = time.time() for step in range(first + 1, args.steps + 1): pick = torch.randint(0, len(valid), (args.batch_size,), device=DEVICE, generator=generator) batch = window(frames, actions, valid[pick], model.config.context) loss = edm_loss(model, *batch, generator) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() if step % args.log_every == 0: log = { "step": step, "loss": round(loss.item(), 4), "steps_s": round((step - first) / (time.time() - start), 1), } print(json.dumps(log), flush=True) if args.checkpoint and (step % args.save_every == 0 or step == args.steps): save_state(args.checkpoint, training_state(step, model, optimizer, generator)) def train(args): """Collect random play, fit, score on held-out episodes, save and optionally push.""" torch.manual_seed(args.seed) frames, actions, episodes = collect(args.frames, args.seed, args.game) config = Config( game=args.game, actions=action_names(args.game), context=args.context, channels=tuple(args.channels), steps=args.sample_steps, ) model = Denoiser(config).to(DEVICE) training = { **{key: value for key, value in vars(args).items() if key not in ("run", "push")}, "episodes": int(episodes[-1]) + 1, "params": sum(param.numel() for param in model.parameters()), "device": torch.cuda.get_device_name() if DEVICE == "cuda" else "cpu", } print(json.dumps(training, indent=1), flush=True) start = time.time() fit(model, *(tensor.to(DEVICE) for tensor in (frames, actions, episodes)), args) training["minutes"] = round((time.time() - start) / 60, 1) model.eval() test = collect(args.eval_frames, EVAL_SEED, args.game) result = evaluate_model(model, *test, args.eval_windows) print(json.dumps(result, indent=1)) folder = Path(args.output) save(model, folder, training, result) gif(model, *test, folder / "results" / "rollout.gif") if args.push: push(folder, args.repo) def save(model, folder, training, result): """Write weights, config, code and results to a folder.""" (folder / "results").mkdir(parents=True, exist_ok=True) state = {key: value.cpu().contiguous() for key, value in model.state_dict().items()} save_file(state, folder / "model.safetensors") files = { "config.json": asdict(model.config), "results/train.json": training, "results/eval.json": result, } for name, value in files.items(): (folder / name).write_text(json.dumps(value, indent=1) + "\n", encoding="utf-8") if Path(__file__).resolve() != (folder / "atari.py").resolve(): shutil.copy(__file__, folder / "atari.py") def push(folder, repo): """Upload a folder to a private model repo, creating it if needed.""" api = HfApi() api.create_repo(repo, private=True, exist_ok=True) api.upload_folder(folder_path=folder, repo_id=repo, commit_message="Add model and results") def load(name, device=DEVICE): """Load a trained model from a Hub repo id or a local folder.""" folder = Path(name) if Path(name).is_dir() else Path(snapshot_download(name)) config = Config(**json.loads((folder / "config.json").read_text(encoding="utf-8"))) model = Denoiser(config) model.load_state_dict(load_file(folder / "model.safetensors")) return model.to(device).eval() # Play def play(args): """A browser page with one button per action; every press generates the next frame.""" model = load(args.model) size = model.config.context frames, actions, _ = collect(size, args.seed, model.config.game) first = (to_float(frames[:size]).to(DEVICE), actions[: size - 1].to(DEVICE)) def act(action, state): context, past = state moves = torch.cat([past, torch.tensor([action], device=DEVICE)]) frame = rollout(model, context[None], moves[None])[0, 0] state = (torch.cat([context[1:], frame[None]]), moves[1:]) return picture(frame, 6), state with gr.Blocks(title=f"{GAMES[model.config.game][0]} world model") as demo: state = gr.State(first) screen = gr.Image(picture(first[0][-1], 6), label="Generated frame") with gr.Row(): for action, name in enumerate(model.config.actions): button = gr.Button(name) button.click(partial(act, action), state, [screen, state], api_name=name.lower()) start = picture(first[0][-1], 6) gr.Button("Reset").click(lambda: (start, first), None, [screen, state], api_name="reset") demo.launch() def card(args): """Render card.jinja into card//README.md with the config and results in the repo.""" display, slug = GAMES[args.game] files = {} try: folder = Path(snapshot_download(args.repo, allow_patterns=["config.json", "results/*"])) files = { path.relative_to(folder).as_posix(): json.loads(path.read_text(encoding="utf-8")) for path in folder.glob("**/*.json") } except RepositoryNotFoundError: pass evaluation = files.get("results/eval.json") data = ModelCardData( model_name=args.repo.split("/")[1], license="apache-2.0", library_name="pytorch", tags=["world-model", "diffusion", "atari", slug, "video-generation"], eval_results=[ EvalResult( task_type="other", task_name="Next-frame rollout", dataset_type=f"ale/{slug}", dataset_name=f"{display}, held-out random play", metric_type=f"psnr_changed_{h}", metric_name=f"PSNR on changed pixels at {h} frames", metric_value=evaluation["model"]["changed"][f"{h}"], ) for h in HORIZONS ] if evaluation else None, ) rendered = ModelCard.from_template( data, template_path=HERE / "card.jinja", repo=args.repo, game=args.game, display=display, actions=action_names(args.game), config=files.get("config.json"), train=files.get("results/train.json"), evaluation=evaluation, horizons=[f"{h}" for h in HORIZONS], ) (HERE / "card" / slug).mkdir(parents=True, exist_ok=True) rendered.save(HERE / "card" / slug / "README.md") def main(): """Parse arguments and run a command.""" parser = argparse.ArgumentParser(description=__doc__) commands = parser.add_subparsers(dest="command", required=True) train_parser = commands.add_parser("train") train_parser.add_argument("--frames", type=int, default=200_000) train_parser.add_argument("--eval-frames", type=int, default=5_000) train_parser.add_argument("--eval-windows", type=int, default=256) train_parser.add_argument("--context", type=int, default=4) train_parser.add_argument("--channels", type=int, nargs="+", default=[64, 128, 256]) train_parser.add_argument("--sample-steps", type=int, default=3) train_parser.add_argument("--steps", type=int, default=50_000) train_parser.add_argument("--batch-size", type=int, default=64) train_parser.add_argument("--lr", type=float, default=1e-4) train_parser.add_argument("--seed", type=int, default=0) train_parser.add_argument("--log-every", type=int, default=500) train_parser.add_argument("--checkpoint", help="folder for resumable state, e.g. a bucket") train_parser.add_argument("--save-every", type=int, default=2500) train_parser.add_argument("--output", default="out") train_parser.add_argument("--repo", help="defaults to the game's repo") train_parser.add_argument("--push", action="store_true") train_parser.set_defaults(run=train) play_parser = commands.add_parser("play") play_parser.add_argument("--model", help="repo id or folder; defaults to the game's repo") play_parser.add_argument("--seed", type=int, default=EVAL_SEED) play_parser.set_defaults(run=play) card_parser = commands.add_parser("card") card_parser.add_argument("--repo", help="defaults to the game's repo") card_parser.set_defaults(run=card) for command in (train_parser, play_parser, card_parser): command.add_argument("--game", choices=GAMES, required=True) args = parser.parse_args() for name in ("repo", "model"): if getattr(args, name, "") is None: setattr(args, name, repo_for(args.game)) args.run(args) if __name__ == "__main__": main()