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()