checkpoint
This commit is contained in:
commit
d0deefd8cf
35 changed files with 2543 additions and 0 deletions
28
tests/test_self_play.py
Normal file
28
tests/test_self_play.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
import jax
|
||||
import jax.numpy as jnp
|
||||
import jax.random as jrandom
|
||||
from generals.core.env import GeneralsEnv
|
||||
|
||||
from general_bots_training.network import PolicyValueNetwork
|
||||
from general_bots_training.opponents import SelfPlayOpponent
|
||||
from general_bots_training.rollout import make_collect_rollout
|
||||
|
||||
|
||||
def test_self_play_collects_both_player_perspectives():
|
||||
key = jrandom.PRNGKey(0)
|
||||
key, network_key, pool_key, state_key = jrandom.split(key, 4)
|
||||
env = GeneralsEnv(grid_dims=(4, 4), truncation=20, pool_size=8)
|
||||
pool, _ = env.reset(pool_key)
|
||||
states = jax.vmap(env.init_state)(jrandom.split(state_key, 2))
|
||||
states = states._replace(pool_idx=jnp.arange(2, dtype=states.pool_idx.dtype))
|
||||
network = PolicyValueNetwork(network_key)
|
||||
|
||||
_, transitions, (_, last_next_obs) = make_collect_rollout(env, 1, SelfPlayOpponent())(
|
||||
states, pool, network, key
|
||||
)
|
||||
|
||||
assert transitions["obs"].shape == (1, 4, 14, 4, 4)
|
||||
assert transitions["action"].shape == (1, 4, 5)
|
||||
assert transitions["player"][0].tolist() == [0, 0, 1, 1]
|
||||
assert transitions["done"].shape == (1, 4)
|
||||
assert last_next_obs.shape == (4, 14, 4, 4)
|
||||
Loading…
Reference in a new issue