57 lines
1.9 KiB
Python
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()
|