Document JAX engine and benchmark
This commit is contained in:
@@ -0,0 +1,56 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user