Files
coorl-lost-cities/tests/lost_cities_jax/test_human_play.py
T

148 lines
5.0 KiB
Python

from __future__ import annotations
import jax
import jax.numpy as jnp
import numpy as np
from lost_cities_jax.engine import board_score, reset_from_order, step
from lost_cities_jax.human_play import (
append_human_log,
evaluate_agent_policy,
human_deck_order,
parse_human_action,
play_one_game,
prompt_human_action,
render_public_state,
session_human_seats,
summarize_logs,
)
from lost_cities_jax.opponents import discard_only_action
from lost_cities_jax.ppo import (
ActorCritic,
JaxPPOConfig,
NetworkConfig,
OpponentConfig,
PPOHyperConfig,
RunConfig,
create_train_state,
)
def tiny_cfg(tmp_path) -> JaxPPOConfig:
return JaxPPOConfig(
run=RunConfig(experiment_name="human-play-test", artifact_root=str(tmp_path)),
opponent=OpponentConfig(name="discard_only"),
network=NetworkConfig(hidden_size=32, num_layers=1),
ppo=PPOHyperConfig(batch_games=4, rollout_steps=16, epochs=1, minibatches=2),
)
def test_illegal_input_reprompts_with_reason() -> None:
state = reset_from_order(jnp.asarray(list(range(60)), dtype=jnp.int8))
inputs = iter(["play R2 draw R", "play R2 draw deck"])
outputs: list[str] = []
action = prompt_human_action(
state,
0,
input_fn=lambda _prompt: next(inputs),
output_fn=outputs.append,
)
assert action == parse_human_action("play R2 draw deck", state, 0)
assert any("Illegal move" in line for line in outputs)
assert any("discard pile is empty" in line for line in outputs)
def test_public_renderer_does_not_show_opponent_hand_or_deck_order() -> None:
p0 = [0, 3, 4, 5, 6, 7, 8, 9]
p1 = [12, 15, 16, 17, 18, 19, 20, 21]
rest = [card for card in range(60) if card not in {*p0, *p1}]
state = reset_from_order(jnp.asarray(p0 + p1 + rest, dtype=jnp.int8))
rendered = render_public_state(state, 0)
assert "R2" in rendered
assert "G2" not in rendered
assert "G6" not in rendered
assert str(p0 + p1 + rest) not in rendered
def test_duplicate_session_swaps_seats() -> None:
assert session_human_seats(0, duplicate=True) == [0, 1]
assert session_human_seats(1, duplicate=True) == [1, 0]
assert session_human_seats(1, duplicate=False) == [1]
def test_human_play_summary_groups_duplicate_sets(tmp_path) -> None:
base = {
"schema": "lost-cities-jax-human-play-v1",
"session_id": "s",
"duplicate_set_id": "set-a",
"moves": [
{"ply": 0, "ai_value": 0.1, "action_text": "play R2 draw deck"},
{"ply": 1, "ai_value": 0.6, "action_text": "discard G2 draw deck"},
],
"ai_opened_colors": 3,
"game_length": 44,
}
append_human_log(tmp_path, {**base, "duplicate_game_index": 1, "human_score_diff": 10})
append_human_log(tmp_path, {**base, "duplicate_game_index": 2, "human_score_diff": -4})
summary = summarize_logs(tmp_path)
assert summary["games"] == 2
assert summary["duplicate_sets"] == 1
assert summary["human_duplicate_set_diff_mean"] == 6.0
assert summary["ai_opened_colors_mean"] == 3.0
assert summary["value_swings_top10"][0]["abs_delta"] == 0.5
def test_human_play_bot_path_matches_direct_engine_for_ten_games(tmp_path) -> None:
cfg = tiny_cfg(tmp_path)
train_state = create_train_state(cfg, jax.random.PRNGKey(7))
agent = (train_state.params, ActorCritic(cfg.network.hidden_size, cfg.network.num_layers))
for index in range(10):
deck_order = human_deck_order(991, index)
record = play_one_game(
cfg=cfg,
agent=agent,
deck_order=deck_order,
human_seat=index % 2,
session_id=f"test-{index}",
duplicate_set_id=None,
duplicate_game_index=None,
deck_seed=991,
deck_index=index,
input_fn=lambda _prompt: "",
output_fn=lambda _line: None,
human_action_fn=_discard_only_human,
)
direct_scores, direct_actions = _direct_game(cfg, agent, deck_order, index % 2)
assert record["actions"] == direct_actions
assert record["final_scores"]["p0"] == float(direct_scores[0])
assert record["final_scores"]["p1"] == float(direct_scores[1])
def _discard_only_human(state, player: int) -> int:
return int(
discard_only_action(state, jnp.asarray(player, dtype=jnp.int32), jax.random.PRNGKey(0))
)
def _direct_game(
cfg, agent, deck_order: list[int], human_seat: int
) -> tuple[np.ndarray, list[int]]:
params, model = agent
state = reset_from_order(jnp.asarray(deck_order, dtype=jnp.int8))
actions = []
while not bool(state.done):
player = int(state.to_move)
if player == human_seat:
action = _discard_only_human(state, player)
else:
action = evaluate_agent_policy(cfg, params, model, state, player).action
actions.append(action)
state, _, _ = step(state, jnp.asarray(action, dtype=jnp.int32))
return np.asarray(board_score(state), dtype=np.float32), actions