86 lines
2.9 KiB
Python
86 lines
2.9 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
|
|
import jax
|
|
import jax.numpy as jnp
|
|
import numpy as np
|
|
|
|
from lost_cities_jax import (
|
|
OBS_DIM,
|
|
batched_legal_mask,
|
|
batched_obs,
|
|
batched_reset,
|
|
batched_step,
|
|
step,
|
|
)
|
|
|
|
|
|
def test_batched_step_jits_without_retrace_and_matches_vectorized_step():
|
|
batch_size = int(os.environ.get("LOST_CITIES_JAX_JIT_BATCH", "8192"))
|
|
keys = jax.random.split(jax.random.PRNGKey(123), batch_size)
|
|
states = batched_reset(keys)
|
|
masks = batched_legal_mask(states)
|
|
actions = jnp.argmax(masks, axis=1).astype(jnp.int32)
|
|
|
|
trace_count = {"value": 0}
|
|
|
|
def counted_vmap_step(batch_state, batch_action):
|
|
trace_count["value"] += 1
|
|
return jax.vmap(step, in_axes=(0, 0))(batch_state, batch_action)
|
|
|
|
counted = jax.jit(counted_vmap_step)
|
|
first = counted(states, actions)
|
|
first[1].block_until_ready()
|
|
assert trace_count["value"] == 1
|
|
|
|
second = counted(states, actions)
|
|
second[1].block_until_ready()
|
|
assert trace_count["value"] == 1
|
|
|
|
batched = batched_step(states, actions)
|
|
direct = jax.vmap(step, in_axes=(0, 0))(states, actions)
|
|
_assert_step_outputs_equal(batched, direct)
|
|
|
|
scalar_step = jax.jit(step)
|
|
sample_count = min(batch_size, 128)
|
|
for idx in range(sample_count):
|
|
scalar_state = jax.tree_util.tree_map(lambda leaf, i=idx: leaf[i], states)
|
|
scalar = scalar_step(scalar_state, actions[idx])
|
|
_assert_scalar_matches_batch(scalar, batched, idx)
|
|
|
|
|
|
def test_batched_observation_shape():
|
|
batch_size = 32
|
|
states = batched_reset(jax.random.split(jax.random.PRNGKey(456), batch_size))
|
|
players = jnp.arange(batch_size, dtype=jnp.int32) % 2
|
|
obs = batched_obs(states, players)
|
|
assert obs.shape == (batch_size, OBS_DIM)
|
|
assert obs.dtype == jnp.float32
|
|
|
|
|
|
def _assert_step_outputs_equal(left, right) -> None:
|
|
left_state, left_reward, left_done = left
|
|
right_state, right_reward, right_done = right
|
|
for left_leaf, right_leaf in zip(
|
|
jax.tree_util.tree_leaves(left_state),
|
|
jax.tree_util.tree_leaves(right_state),
|
|
strict=False,
|
|
):
|
|
np.testing.assert_array_equal(np.asarray(left_leaf), np.asarray(right_leaf))
|
|
np.testing.assert_array_equal(np.asarray(left_reward), np.asarray(right_reward))
|
|
np.testing.assert_array_equal(np.asarray(left_done), np.asarray(right_done))
|
|
|
|
|
|
def _assert_scalar_matches_batch(scalar, batch, idx: int) -> None:
|
|
scalar_state, scalar_reward, scalar_done = scalar
|
|
batch_state, batch_reward, batch_done = batch
|
|
for scalar_leaf, batch_leaf in zip(
|
|
jax.tree_util.tree_leaves(scalar_state),
|
|
jax.tree_util.tree_leaves(batch_state),
|
|
strict=False,
|
|
):
|
|
np.testing.assert_array_equal(np.asarray(scalar_leaf), np.asarray(batch_leaf[idx]))
|
|
np.testing.assert_array_equal(np.asarray(scalar_reward), np.asarray(batch_reward[idx]))
|
|
np.testing.assert_array_equal(np.asarray(scalar_done), np.asarray(batch_done[idx]))
|