diff --git a/mctx/_src/policies.py b/mctx/_src/policies.py index f170db9..f47fec7 100644 --- a/mctx/_src/policies.py +++ b/mctx/_src/policies.py @@ -115,6 +115,7 @@ def muzero_policy( action_weights = summary.visit_probs action_logits = _apply_temperature( _get_logits_from_probs(action_weights), temperature) + action_logits = _mask_invalid_actions(action_logits, invalid_actions) action = jax.random.categorical(rng_key, action_logits) return base.PolicyOutput( action=action, @@ -368,6 +369,7 @@ def stochastic_muzero_policy( action_weights = summary.visit_probs action_logits = _apply_temperature( _get_logits_from_probs(action_weights), temperature) + action_logits = _mask_invalid_actions(action_logits, invalid_actions) action = jax.random.categorical(rng_key, action_logits) return base.PolicyOutput( action=action, action_weights=action_weights, search_tree=search_tree) diff --git a/mctx/_src/tests/policies_test.py b/mctx/_src/tests/policies_test.py index 549dca9..f6b1d97 100644 --- a/mctx/_src/tests/policies_test.py +++ b/mctx/_src/tests/policies_test.py @@ -16,6 +16,7 @@ import functools from absl.testing import absltest +from absl.testing import parameterized import jax import jax.numpy as jnp import mctx @@ -359,5 +360,56 @@ def test_stochastic_muzero_policy(self): policy_output.action_weights) +class PolicyActionMaskTest(parameterized.TestCase): + + @parameterized.product( + stochastic=[False, True], + use_jit=[False, True], + temperature=[0.0, 1.0, 100.0], + ) + def test_sampled_actions_respect_mask(self, stochastic, use_jit, temperature): + batch_size, num_actions = 32, 4 + valid_actions = jnp.arange(batch_size) % num_actions + invalid_actions = jnp.arange(num_actions)[None, :] != valid_actions[:, None] + root = mctx.RootFnOutput( + prior_logits=jnp.zeros((batch_size, num_actions)), + value=jnp.zeros(batch_size), + embedding=jnp.zeros((batch_size, 1)), + ) + rewards = jnp.zeros_like(root.prior_logits) + if stochastic: + decision_fn, chance_fn = _make_bandit_decision_and_chance_fns(rewards, 2) + policy = functools.partial( + mctx.stochastic_muzero_policy, + decision_recurrent_fn=decision_fn, + chance_recurrent_fn=chance_fn, + ) + else: + policy = functools.partial( + mctx.muzero_policy, + recurrent_fn=_make_bandit_recurrent_fn( + rewards, dummy_embedding=jnp.zeros_like(root.embedding) + ), + ) + policy = functools.partial( + policy, + params=(), + root=root, + num_simulations=2, + invalid_actions=invalid_actions, + dirichlet_fraction=0.0, + temperature=temperature, + ) + if use_jit: + policy = jax.jit(policy) + + output = policy(rng_key=jax.random.PRNGKey(0)) + + np.testing.assert_array_equal(output.action, valid_actions) + np.testing.assert_allclose( + output.action_weights, jax.nn.one_hot(valid_actions, num_actions) + ) + + if __name__ == "__main__": absltest.main()