jgalego's picture
Add model and results
12fb5b1 verified
Raw History Blame Contribute Delete
22.9 kB
# /// 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/<slug>/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()