Files
coorl-lost-cities/benchmarks/throughput.py
T

57 lines
1.9 KiB
Python

from __future__ import annotations
import argparse
import time
from functools import partial
import jax
import jax.numpy as jnp
from lost_cities_jax.engine import legal_action_mask, reset, step
@partial(jax.jit, static_argnames=("steps",))
def rollout(states, rng, *, steps: int):
def body(carry, _):
state, key = carry
key, action_key = jax.random.split(key)
mask = jax.vmap(legal_action_mask)(state)
logits = jnp.where(mask, 0.0, -1.0e9)
actions = jax.random.categorical(action_key, logits, axis=1).astype(jnp.int32)
state, _, _ = jax.vmap(step, in_axes=(0, 0))(state, actions)
return (state, key), None
(states, rng), _ = jax.lax.scan(body, (states, rng), xs=None, length=steps)
return states, rng
def main() -> None:
parser = argparse.ArgumentParser(description="Lost Cities JAX random-policy throughput")
parser.add_argument("--batch-size", type=int, default=8192)
parser.add_argument("--steps", type=int, default=256)
parser.add_argument("--warmup-steps", type=int, default=32)
args = parser.parse_args()
key = jax.random.PRNGKey(0)
reset_keys = jax.random.split(key, args.batch_size)
states = jax.jit(jax.vmap(reset))(reset_keys)
warm_states, key = rollout(states, jax.random.PRNGKey(1), steps=args.warmup_steps)
jax.tree_util.tree_leaves(warm_states)[0].block_until_ready()
start = time.perf_counter()
states, _ = rollout(states, key, steps=args.steps)
jax.tree_util.tree_leaves(states)[0].block_until_ready()
elapsed = time.perf_counter() - start
transitions = args.batch_size * args.steps
print(f"backend={jax.default_backend()}")
print(f"batch_size={args.batch_size}")
print(f"steps={args.steps}")
print(f"elapsed_sec={elapsed:.6f}")
print(f"steps_per_sec={transitions / elapsed:.2f}")
if __name__ == "__main__":
main()