From ac54f981891eb0d1a037411f5144d1452a3c51b4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Sat, 4 Jul 2026 19:38:04 +0900 Subject: [PATCH] Add JAX engine verification tests --- reference/__init__.py | 1 + reference/lost_cities_ref.py | 247 ++++++++++++++++++ tests/lost_cities_jax/__init__.py | 1 + tests/lost_cities_jax/conftest.py | 26 ++ tests/lost_cities_jax/helpers.py | 129 +++++++++ tests/lost_cities_jax/test_batched.py | 85 ++++++ tests/lost_cities_jax/test_differential.py | 94 +++++++ tests/lost_cities_jax/test_engine_core.py | 62 +++++ tests/lost_cities_jax/test_mask_bruteforce.py | 104 ++++++++ tests/lost_cities_jax/test_properties.py | 66 +++++ tests/lost_cities_jax/test_scoring.py | 31 +++ 11 files changed, 846 insertions(+) create mode 100644 reference/__init__.py create mode 100644 reference/lost_cities_ref.py create mode 100644 tests/lost_cities_jax/__init__.py create mode 100644 tests/lost_cities_jax/conftest.py create mode 100644 tests/lost_cities_jax/helpers.py create mode 100644 tests/lost_cities_jax/test_batched.py create mode 100644 tests/lost_cities_jax/test_differential.py create mode 100644 tests/lost_cities_jax/test_engine_core.py create mode 100644 tests/lost_cities_jax/test_mask_bruteforce.py create mode 100644 tests/lost_cities_jax/test_properties.py create mode 100644 tests/lost_cities_jax/test_scoring.py diff --git a/reference/__init__.py b/reference/__init__.py new file mode 100644 index 0000000..6a634a1 --- /dev/null +++ b/reference/__init__.py @@ -0,0 +1 @@ +"""Reference implementations used by tests.""" diff --git a/reference/lost_cities_ref.py b/reference/lost_cities_ref.py new file mode 100644 index 0000000..45e7887 --- /dev/null +++ b/reference/lost_cities_ref.py @@ -0,0 +1,247 @@ +"""Independent pure-Python Lost Cities rules reference. + +This module intentionally uses ordinary Python containers rather than the JAX +state representation. It is deterministic under an explicit ``deck_order`` and +the shared flat action encoding. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from random import Random + +N_PLAYERS = 2 +N_COLORS = 5 +CARDS_PER_COLOR = 12 +N_CARDS = N_COLORS * CARDS_PER_COLOR +HAND_SIZE = 8 +INITIAL_DEAL = N_PLAYERS * HAND_SIZE +MAX_STEPS = 400 +N_ACTIONS = HAND_SIZE * 2 * 6 + +LOC_DECK = 0 +LOC_P0_HAND = 1 +LOC_P1_HAND = 2 +LOC_P0_BOARD = 3 +LOC_P1_BOARD = 4 +LOC_DISCARD = 5 + +NO_CARD = -1 +PLAY = 0 +DISCARD = 1 +DRAW_DECK = 0 + + +@dataclass +class RefState: + deck_order: list[int] + draw_ptr: int + card_loc: list[int] + hand_public: list[bool] + board: list[list[list[int]]] + piles: list[list[int]] + to_move: int + just_discarded: int + step_count: int + done: bool + + +def reset(seed: int | None = None) -> RefState: + rng = Random(seed) + order = list(range(N_CARDS)) + rng.shuffle(order) + return reset_from_order(order) + + +def reset_from_order(deck_order: list[int] | tuple[int, ...]) -> RefState: + order = [int(card) for card in deck_order] + card_loc = [LOC_DECK] * N_CARDS + for card in order[:HAND_SIZE]: + card_loc[card] = LOC_P0_HAND + for card in order[HAND_SIZE:INITIAL_DEAL]: + card_loc[card] = LOC_P1_HAND + return RefState( + deck_order=order, + draw_ptr=INITIAL_DEAL, + card_loc=card_loc, + hand_public=[False] * N_CARDS, + board=[[[] for _ in range(N_COLORS)] for _ in range(N_PLAYERS)], + piles=[[] for _ in range(N_COLORS)], + to_move=0, + just_discarded=NO_CARD, + step_count=0, + done=False, + ) + + +def clone_state(state: RefState) -> RefState: + return RefState( + deck_order=list(state.deck_order), + draw_ptr=state.draw_ptr, + card_loc=list(state.card_loc), + hand_public=list(state.hand_public), + board=[[list(col) for col in player] for player in state.board], + piles=[list(pile) for pile in state.piles], + to_move=state.to_move, + just_discarded=state.just_discarded, + step_count=state.step_count, + done=state.done, + ) + + +def decode_action(action: int) -> tuple[int, int, int]: + action = int(action) + hand_slot = action // 12 + rem = action % 12 + place_type = rem // 6 + draw_source = rem % 6 + return hand_slot, place_type, draw_source + + +def hand_cards(state: RefState, player: int | None = None) -> list[int]: + if player is None: + player = state.to_move + hand_loc = LOC_P0_HAND + player + return sorted(card for card, loc in enumerate(state.card_loc) if loc == hand_loc) + + +def legal_action_mask(state: RefState) -> list[bool]: + mask = [False] * N_ACTIONS + if state.done: + return mask + + player = state.to_move + hand = hand_cards(state, player) + for hand_slot in range(HAND_SIZE): + if hand_slot >= len(hand): + continue + card = hand[hand_slot] + for place_type in (PLAY, DISCARD): + if place_type == PLAY: + place_ok = can_play_card(state, player, card) + else: + place_ok = True + if not place_ok: + continue + for draw_source in range(6): + if _can_draw_after_place(state, card, place_type, draw_source): + mask[hand_slot * 12 + place_type * 6 + draw_source] = True + return mask + + +def step(state: RefState, action: int) -> tuple[RefState, list[float], bool]: + if state.done or action < 0 or action >= N_ACTIONS: + return clone_state(state), [0.0, 0.0], state.done + if not legal_action_mask(state)[action]: + return clone_state(state), [0.0, 0.0], state.done + + next_state = clone_state(state) + player = next_state.to_move + hand_slot, place_type, draw_source = decode_action(action) + card = hand_cards(next_state, player)[hand_slot] + color = card_color(card) + + next_state.hand_public[card] = False + if place_type == PLAY: + next_state.board[player][color].append(card) + next_state.card_loc[card] = LOC_P0_BOARD + player + next_state.just_discarded = NO_CARD + else: + next_state.piles[color].append(card) + next_state.card_loc[card] = LOC_DISCARD + next_state.just_discarded = card + + if draw_source == DRAW_DECK: + drawn = next_state.deck_order[next_state.draw_ptr] + next_state.draw_ptr += 1 + public = False + else: + src = draw_source - 1 + drawn = next_state.piles[src].pop() + public = True + + next_state.card_loc[drawn] = LOC_P0_HAND + player + next_state.hand_public[drawn] = public + next_state.just_discarded = NO_CARD + next_state.step_count += 1 + next_state.done = (draw_source == DRAW_DECK and next_state.draw_ptr >= N_CARDS) or ( + next_state.step_count >= MAX_STEPS + ) + next_state.to_move = 1 - player + + reward = board_score(next_state) if next_state.done else [0.0, 0.0] + return next_state, reward, next_state.done + + +def can_play_card(state: RefState, player: int, card: int) -> bool: + column = state.board[player][card_color(card)] + if is_handshake(card): + return not any(not is_handshake(played) for played in column) + top_rank = 0 + for played in column: + if not is_handshake(played): + top_rank = rank(played) + return rank(card) > top_rank + + +def board_score(state: RefState) -> list[float]: + return [float(sum(score_column(column) for column in player)) for player in state.board] + + +def score(state: RefState) -> list[float]: + return board_score(state) + + +def score_column(column: list[int]) -> int: + if not column: + return 0 + handshakes = sum(1 for card in column if is_handshake(card)) + rank_sum = sum(rank(card) for card in column if not is_handshake(card)) + value = (rank_sum - 20) * (1 + handshakes) + if len(column) >= 8: + value += 20 + return value + + +def card_color(card: int) -> int: + return int(card) // CARDS_PER_COLOR + + +def card_slot(card: int) -> int: + return int(card) % CARDS_PER_COLOR + + +def is_handshake(card: int) -> bool: + return card_slot(card) < 3 + + +def rank(card: int) -> int: + return card_slot(card) - 1 + + +def _can_draw_after_place(state: RefState, card: int, place_type: int, draw_source: int) -> bool: + if draw_source == DRAW_DECK: + return True + src = draw_source - 1 + same_discard_pile = place_type == DISCARD and card_color(card) == src + after_len = len(state.piles[src]) + int(same_discard_pile) + if after_len == 0: + return False + after_top = card if same_discard_pile else state.piles[src][-1] + just_discarded = card if place_type == DISCARD else state.just_discarded + return after_top != just_discarded + + +__all__ = [ + "N_ACTIONS", + "RefState", + "board_score", + "clone_state", + "decode_action", + "hand_cards", + "legal_action_mask", + "reset", + "reset_from_order", + "score", + "step", +] diff --git a/tests/lost_cities_jax/__init__.py b/tests/lost_cities_jax/__init__.py new file mode 100644 index 0000000..f19f885 --- /dev/null +++ b/tests/lost_cities_jax/__init__.py @@ -0,0 +1 @@ +"""Tests for the standalone JAX Lost Cities engine.""" diff --git a/tests/lost_cities_jax/conftest.py b/tests/lost_cities_jax/conftest.py new file mode 100644 index 0000000..2bf451d --- /dev/null +++ b/tests/lost_cities_jax/conftest.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +import os + + +def pytest_addoption(parser): + parser.addoption( + "--full", + action="store_true", + default=False, + help="run full Lost Cities JAX differential profile", + ) + + +def pytest_configure(config): + config.addinivalue_line("markers", "slow: longer randomized Lost Cities JAX checks") + + +def game_count(config, env_name: str, local: int, ci: int, full: int) -> int: + if env_name in os.environ: + return int(os.environ[env_name]) + if config.getoption("--full"): + return full + if os.environ.get("CI"): + return ci + return local diff --git a/tests/lost_cities_jax/helpers.py b/tests/lost_cities_jax/helpers.py new file mode 100644 index 0000000..27fe5a5 --- /dev/null +++ b/tests/lost_cities_jax/helpers.py @@ -0,0 +1,129 @@ +from __future__ import annotations + +import random + +import jax +import jax.numpy as jnp +import numpy as np + +from lost_cities_jax import reset_from_order +from lost_cities_jax.types import ( + CARDS_PER_COLOR, + LOC_DISCARD, + LOC_P0_BOARD, + LOC_P0_HAND, + MAX_PILE_SIZE, + N_CARDS, + N_COLORS, + NO_CARD, + State, +) + + +def shuffled_order(seed: int) -> list[int]: + order = list(range(N_CARDS)) + random.Random(seed).shuffle(order) + return order + + +def hs(color: int, index: int = 0) -> int: + return color * CARDS_PER_COLOR + index + + +def num(color: int, rank: int) -> int: + return color * CARDS_PER_COLOR + rank + 1 + + +def state_with_columns( + p0_columns: list[list[int]], + p1_columns: list[list[int]] | None = None, +) -> State: + if p1_columns is None: + p1_columns = [[] for _ in range(N_COLORS)] + + state = reset_from_order(jnp.arange(N_CARDS, dtype=jnp.int8)) + card_loc = jnp.zeros((N_CARDS,), dtype=jnp.int8) + col_top = jnp.zeros((2, N_COLORS), dtype=jnp.int8) + col_hs = jnp.zeros((2, N_COLORS), dtype=jnp.int8) + col_len = jnp.zeros((2, N_COLORS), dtype=jnp.int8) + + for player, columns in enumerate([p0_columns, p1_columns]): + for color, cards in enumerate(columns): + if not cards: + continue + idx = jnp.asarray(cards, dtype=jnp.int32) + card_loc = card_loc.at[idx].set(jnp.int8(LOC_P0_BOARD + player)) + slots = [card % CARDS_PER_COLOR for card in cards] + ranks = [slot - 1 for slot in slots if slot >= 3] + col_top = col_top.at[player, color].set(max(ranks, default=0)) + col_hs = col_hs.at[player, color].set(sum(slot < 3 for slot in slots)) + col_len = col_len.at[player, color].set(len(cards)) + + return state._replace( + card_loc=card_loc, + col_top=col_top, + col_hs=col_hs, + col_len=col_len, + ) + + +def first_deck_draw_action(mask) -> int: + values = np.asarray(mask, dtype=bool) + for action, legal in enumerate(values): + if legal and action % 6 == 0: + return action + raise AssertionError("no legal deck-draw action") + + +def assert_state_consistent(state: State) -> None: + loc = np.asarray(state.card_loc) + public = np.asarray(state.hand_public) + pile = np.asarray(state.pile) + pile_len = np.asarray(state.pile_len) + col_top = np.asarray(state.col_top) + col_hs = np.asarray(state.col_hs) + col_len = np.asarray(state.col_len) + + assert loc.shape == (N_CARDS,) + assert np.all((0 <= loc) & (loc <= LOC_DISCARD)) + assert sum(np.bincount(loc, minlength=LOC_DISCARD + 1)) == N_CARDS + assert int(np.asarray(state.just_discarded)) == NO_CARD + + for player in range(2): + assert int(np.sum(loc == LOC_P0_HAND + player)) == 8 + for color in range(N_COLORS): + cards = [ + card + for card in range(N_CARDS) + if loc[card] == LOC_P0_BOARD + player and card // CARDS_PER_COLOR == color + ] + slots = [card % CARDS_PER_COLOR for card in cards] + ranks = [slot - 1 for slot in slots if slot >= 3] + assert int(col_len[player, color]) == len(cards) + assert int(col_hs[player, color]) == sum(slot < 3 for slot in slots) + assert int(col_top[player, color]) == max(ranks, default=0) + assert int(col_hs[player, color]) + len(ranks) == int(col_len[player, color]) + + for color in range(N_COLORS): + length = int(pile_len[color]) + assert 0 <= length < MAX_PILE_SIZE + used = list(map(int, pile[color, :length])) + assert all(card != NO_CARD for card in used) + assert len(set(used)) == len(used) + assert np.all(pile[color, length:] == NO_CARD) + discard_cards = [ + card + for card in range(N_CARDS) + if loc[card] == LOC_DISCARD and card // CARDS_PER_COLOR == color + ] + assert sorted(used) == sorted(discard_cards) + + in_hand = (loc == LOC_P0_HAND) | (loc == LOC_P0_HAND + 1) + assert not np.any(public & ~in_hand) + + +def assert_states_equal(left: State, right: State) -> None: + for left_leaf, right_leaf in zip( + jax.tree_util.tree_leaves(left), jax.tree_util.tree_leaves(right), strict=False + ): + np.testing.assert_array_equal(np.asarray(left_leaf), np.asarray(right_leaf)) diff --git a/tests/lost_cities_jax/test_batched.py b/tests/lost_cities_jax/test_batched.py new file mode 100644 index 0000000..11583af --- /dev/null +++ b/tests/lost_cities_jax/test_batched.py @@ -0,0 +1,85 @@ +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])) diff --git a/tests/lost_cities_jax/test_differential.py b/tests/lost_cities_jax/test_differential.py new file mode 100644 index 0000000..97e4b06 --- /dev/null +++ b/tests/lost_cities_jax/test_differential.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +import json +import random + +import jax +import jax.numpy as jnp +import pytest + +from lost_cities_jax import legal_action_mask, reset_from_order, score, step +from reference import lost_cities_ref as ref +from tests.lost_cities_jax.conftest import game_count +from tests.lost_cities_jax.helpers import shuffled_order + +JIT_MASK = jax.jit(legal_action_mask) +JIT_STEP = jax.jit(step) +JIT_SCORE = jax.jit(score) + + +@pytest.mark.slow +def test_jax_matches_reference_random_legal_policy(pytestconfig, tmp_path): + games = game_count( + pytestconfig, + env_name="LOST_CITIES_JAX_DIFF_GAMES", + local=100, + ci=100_000, + full=1_000_000, + ) + policy_rng = random.Random(20260704) + + for game_idx in range(games): + deck_order = shuffled_order(10_000_000 + game_idx) + jax_state = reset_from_order(jnp.asarray(deck_order, dtype=jnp.int8)) + ref_state = ref.reset_from_order(deck_order) + actions: list[int] = [] + + for step_idx in range(500): + jax_mask = list(map(bool, JIT_MASK(jax_state).tolist())) + ref_mask = ref.legal_action_mask(ref_state) + if jax_mask != ref_mask: + _dump_failure(tmp_path, game_idx, step_idx, deck_order, actions) + diff = [ + idx + for idx, pair in enumerate(zip(jax_mask, ref_mask, strict=False)) + if pair[0] != pair[1] + ] + pytest.fail(f"legal mask mismatch game={game_idx} step={step_idx} diff={diff[:20]}") + + if bool(jax_state.done) != ref_state.done: + _dump_failure(tmp_path, game_idx, step_idx, deck_order, actions) + pytest.fail(f"done mismatch game={game_idx} step={step_idx}") + + if ref_state.done: + jax_score = [float(value) for value in JIT_SCORE(jax_state).tolist()] + ref_score = ref.score(ref_state) + if jax_score != ref_score: + _dump_failure(tmp_path, game_idx, step_idx, deck_order, actions) + pytest.fail( + f"score mismatch game={game_idx} step={step_idx}: {jax_score} != {ref_score}" + ) + break + + legal_actions = [idx for idx, legal in enumerate(ref_mask) if legal] + action = legal_actions[policy_rng.randrange(len(legal_actions))] + actions.append(action) + jax_state, jax_reward, jax_done = JIT_STEP(jax_state, jnp.int32(action)) + ref_state, ref_reward, ref_done = ref.step(ref_state, action) + if bool(jax_done) != ref_done or [float(v) for v in jax_reward.tolist()] != ref_reward: + _dump_failure(tmp_path, game_idx, step_idx, deck_order, actions) + pytest.fail(f"step result mismatch game={game_idx} step={step_idx}") + else: + _dump_failure(tmp_path, game_idx, 500, deck_order, actions) + pytest.fail(f"game did not finish within guard loop: game={game_idx}") + + +def _dump_failure( + tmp_path, + game_idx: int, + step_idx: int, + deck_order: list[int], + actions: list[int], +) -> None: + path = tmp_path / f"lost_cities_jax_diff_failure_{game_idx}_{step_idx}.json" + path.write_text( + json.dumps( + { + "game_idx": game_idx, + "step_idx": step_idx, + "deck_order": deck_order, + "actions": actions, + }, + indent=2, + ) + ) diff --git a/tests/lost_cities_jax/test_engine_core.py b/tests/lost_cities_jax/test_engine_core.py new file mode 100644 index 0000000..c3da7f6 --- /dev/null +++ b/tests/lost_cities_jax/test_engine_core.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import jax.numpy as jnp +import numpy as np + +from lost_cities_jax import legal_action_mask, reset_from_order, score, step +from lost_cities_jax.engine import current_hand_sorted +from lost_cities_jax.types import CARDS_PER_COLOR, N_CARDS +from tests.lost_cities_jax.helpers import assert_states_equal, first_deck_draw_action + + +def test_discarded_card_cannot_be_drawn_from_same_pile(): + state = reset_from_order(jnp.arange(N_CARDS, dtype=jnp.int8)) + hand = np.asarray(current_hand_sorted(state)) + first_card_color = int(hand[0]) // CARDS_PER_COLOR + discard_then_same_pile = 0 * 12 + 1 * 6 + (first_card_color + 1) + discard_then_deck = 0 * 12 + 1 * 6 + + mask = np.asarray(legal_action_mask(state), dtype=bool) + + assert not mask[discard_then_same_pile] + assert mask[discard_then_deck] + + +def test_illegal_action_is_noop_with_zero_reward(): + state = reset_from_order(jnp.arange(N_CARDS, dtype=jnp.int8)) + hand = np.asarray(current_hand_sorted(state)) + color = int(hand[0]) // CARDS_PER_COLOR + illegal_action = 0 * 12 + 1 * 6 + (color + 1) + + next_state, reward, done = step(state, jnp.int32(illegal_action)) + + assert_states_equal(next_state, state) + np.testing.assert_array_equal(np.asarray(reward), np.zeros(2, dtype=np.float32)) + assert bool(done) is False + + +def test_done_state_step_is_noop_and_has_no_repeated_terminal_reward(): + state = reset_from_order(jnp.arange(N_CARDS, dtype=jnp.int8)) + while not bool(state.done): + action = first_deck_draw_action(legal_action_mask(state)) + state, _, _ = step(state, jnp.int32(action)) + + after_done, reward, done = step(state, jnp.int32(0)) + + assert_states_equal(after_done, state) + np.testing.assert_array_equal(np.asarray(reward), np.zeros(2, dtype=np.float32)) + assert bool(done) is True + + +def test_deck_only_draws_finish_in_exactly_44_plies(): + state = reset_from_order(jnp.arange(N_CARDS, dtype=jnp.int8)) + plies = 0 + while not bool(state.done): + action = first_deck_draw_action(legal_action_mask(state)) + state, reward, done = step(state, jnp.int32(action)) + plies += 1 + + assert plies == 44 + assert int(state.draw_ptr) == N_CARDS + assert bool(done) is True + np.testing.assert_array_equal(np.asarray(reward), np.asarray(score(state))) diff --git a/tests/lost_cities_jax/test_mask_bruteforce.py b/tests/lost_cities_jax/test_mask_bruteforce.py new file mode 100644 index 0000000..5adb720 --- /dev/null +++ b/tests/lost_cities_jax/test_mask_bruteforce.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import random + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from lost_cities_jax import legal_action_mask, reset_from_order, step +from lost_cities_jax.types import ( + CARDS_PER_COLOR, + DISCARD, + DRAW_DECK, + HAND_SIZE, + LOC_P0_HAND, + N_ACTIONS, + N_CARDS, + PLAY, +) +from tests.lost_cities_jax.conftest import game_count +from tests.lost_cities_jax.helpers import shuffled_order + +JIT_MASK = jax.jit(legal_action_mask) +JIT_STEP = jax.jit(step) + + +@pytest.mark.slow +def test_legal_mask_matches_slow_bruteforce_checker(pytestconfig): + state_count = game_count( + pytestconfig, + env_name="LOST_CITIES_JAX_MASK_STATES", + local=1_000, + ci=10_000, + full=10_000, + ) + rng = random.Random(31337) + state = reset_from_order(jnp.asarray(shuffled_order(30_000_000), dtype=jnp.int8)) + next_seed = 30_000_001 + + for _ in range(state_count): + jax_mask = np.asarray(JIT_MASK(state), dtype=bool) + slow_mask = _slow_mask(state) + np.testing.assert_array_equal(jax_mask, slow_mask) + + if bool(state.done): + state = reset_from_order(jnp.asarray(shuffled_order(next_seed), dtype=jnp.int8)) + next_seed += 1 + continue + + legal_actions = np.flatnonzero(jax_mask) + action = int(legal_actions[rng.randrange(len(legal_actions))]) + state, _, _ = JIT_STEP(state, jnp.int32(action)) + + +def _slow_mask(state) -> np.ndarray: + mask = np.zeros(N_ACTIONS, dtype=bool) + if bool(state.done): + return mask + + loc = np.asarray(state.card_loc) + pile = np.asarray(state.pile) + pile_len = np.asarray(state.pile_len) + col_top = np.asarray(state.col_top) + player = int(state.to_move) + just_discarded = int(state.just_discarded) + hand_loc = LOC_P0_HAND + player + hand = sorted(card for card in range(N_CARDS) if loc[card] == hand_loc) + + for hand_slot in range(HAND_SIZE): + if hand_slot >= len(hand): + continue + card = hand[hand_slot] + color = card // CARDS_PER_COLOR + slot = card % CARDS_PER_COLOR + is_handshake = slot < 3 + rank = slot - 1 + + for place_type in (PLAY, DISCARD): + if place_type == PLAY: + place_ok = ( + col_top[player, color] == 0 if is_handshake else rank > col_top[player, color] + ) + else: + place_ok = True + if not place_ok: + continue + + for draw_source in range(6): + if draw_source == DRAW_DECK: + draw_ok = True + else: + src = draw_source - 1 + same_discard_pile = place_type == DISCARD and src == color + after_len = int(pile_len[src]) + int(same_discard_pile) + if after_len == 0: + draw_ok = False + else: + after_top = card if same_discard_pile else int(pile[src, after_len - 1]) + just = card if place_type == DISCARD else just_discarded + draw_ok = after_top != just + mask[hand_slot * 12 + place_type * 6 + draw_source] = draw_ok + + return mask diff --git a/tests/lost_cities_jax/test_properties.py b/tests/lost_cities_jax/test_properties.py new file mode 100644 index 0000000..0e5c5ea --- /dev/null +++ b/tests/lost_cities_jax/test_properties.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +import random + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from lost_cities_jax import legal_action_mask, reset_from_order, score, step +from tests.lost_cities_jax.conftest import game_count +from tests.lost_cities_jax.helpers import ( + assert_state_consistent, + assert_states_equal, + shuffled_order, +) + +JIT_MASK = jax.jit(legal_action_mask) +JIT_STEP = jax.jit(step) + + +@pytest.mark.slow +def test_random_trajectories_preserve_state_invariants(pytestconfig): + games = game_count( + pytestconfig, + env_name="LOST_CITIES_JAX_PROPERTY_GAMES", + local=50, + ci=1_000, + full=20_000, + ) + rng = random.Random(777) + + for game_idx in range(games): + state = reset_from_order(jnp.asarray(shuffled_order(20_000_000 + game_idx), dtype=jnp.int8)) + assert_state_consistent(state) + + for _ in range(500): + if bool(state.done): + after_done, reward, done = JIT_STEP(state, jnp.int32(0)) + assert_states_equal(after_done, state) + np.testing.assert_array_equal(np.asarray(reward), np.zeros(2, dtype=np.float32)) + assert bool(done) is True + break + + mask = np.asarray(JIT_MASK(state), dtype=bool) + legal_actions = np.flatnonzero(mask) + assert len(legal_actions) > 0 + action = int(legal_actions[rng.randrange(len(legal_actions))]) + state, reward, done = JIT_STEP(state, jnp.int32(action)) + assert_state_consistent(state) + if bool(done): + np.testing.assert_array_equal(np.asarray(reward), np.asarray(score(state))) + else: + pytest.fail(f"game did not terminate within guard loop: game={game_idx}") + + +def test_forced_termination_scores_current_board(): + state = reset_from_order(jnp.arange(60, dtype=jnp.int8))._replace( + step_count=jnp.asarray(399, dtype=jnp.int32) + ) + action = int(np.flatnonzero(np.asarray(legal_action_mask(state), dtype=bool))[0]) + next_state, reward, done = step(state, jnp.int32(action)) + + assert bool(done) is True + assert int(next_state.step_count) == 400 + np.testing.assert_array_equal(np.asarray(reward), np.asarray(score(next_state))) diff --git a/tests/lost_cities_jax/test_scoring.py b/tests/lost_cities_jax/test_scoring.py new file mode 100644 index 0000000..d1d18c6 --- /dev/null +++ b/tests/lost_cities_jax/test_scoring.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import pytest + +from lost_cities_jax import board_score +from reference import lost_cities_ref as ref +from tests.lost_cities_jax.helpers import hs, num, state_with_columns + + +@pytest.mark.parametrize( + ("cards", "expected"), + [ + ([num(0, 2), num(0, 3), num(0, 7), num(0, 8), num(0, 10)], 10), + ( + [hs(0, 0), hs(0, 1), num(0, 4), num(0, 5), num(0, 6), num(0, 7), num(0, 8), num(0, 10)], + 80, + ), + ([hs(0), num(0, 4), num(0, 6), num(0, 7)], -6), + ([hs(0)], -40), + ([], 0), + ([hs(0, 0), hs(0, 1), hs(0, 2), *[num(0, rank) for rank in range(2, 11)]], 156), + ([hs(0, 0), hs(0, 1), hs(0, 2), num(0, 2)], -72), + ], +) +def test_board_score_examples(cards, expected): + state = state_with_columns([cards, [], [], [], []]) + assert float(board_score(state)[0]) == expected + + ref_state = ref.reset_from_order(list(range(60))) + ref_state.board[0][0] = list(cards) + assert ref.board_score(ref_state)[0] == expected