Download atari.py from jgalego/space-invaders-world-model: direct link, hf CLI and curl.
- Browser
- Download file 22.9 kB
-
https://huggingface.co/jgalego/space-invaders-world-model/resolve/main/atari.py
- Command line
-
hf download hf://jgalego/space-invaders-world-model/atari.py
-
curl -L -o atari.py https://huggingface.co/jgalego/space-invaders-world-model/resolve/main/atari.py
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) | |
| 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() | |
| 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) | |
| 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() | |