From f872204b13206be8a4a23694368a4951dd468a51 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:09 +0900 Subject: [PATCH] Document JAX engine and benchmark --- README.md | 137 +++++++++++++++++++++++++++++++++++++++ benchmarks/throughput.py | 56 ++++++++++++++++ 2 files changed, 193 insertions(+) create mode 100644 benchmarks/throughput.py diff --git a/README.md b/README.md index 8834e1c..b514728 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,143 @@ uv run lost-cities-classic-gui --mode pvc --bot safe-heuristic The GUI uses the in-process Cython game engine. +## JAX Rules Engine + +`lost_cities_jax` is a standalone pure rules simulator for one two-player +Lost Cities round. It does not contain neural networks, PPO, CFR, MCTS, match +wrappers, bots, or rule variants. + +Public API: + +```python +from lost_cities_jax import ( + OBS_DIM, + N_ACTIONS, + State, + batched_legal_mask, + batched_obs, + batched_reset, + batched_step, + board_score, + legal_action_mask, + observation, + reset, + reset_from_order, + score, + step, +) +``` + +Core functions are pure JAX functions: + +- `reset(rng) -> State` +- `reset_from_order(deck_order) -> State` +- `legal_action_mask(state) -> bool[96]` +- `step(state, action) -> (State, float32[2], bool)` +- `score(state) -> float32[2]` +- `board_score(state) -> float32[2]` +- `observation(state, player) -> float32[454]` + +The batched exports are `jax.jit(jax.vmap(...))` wrappers. Illegal actions and +done-state actions are defined as no-op transitions with zero reward; training +code should still sample only from `legal_action_mask`. + +### Rule Summary + +One round uses 60 cards: five colors, each with three handshakes and ranks +2 through 10. Each player starts with eight cards, then each ply must place one +hand card to the matching expedition or discard pile and draw one card from the +deck or a discard pile. A player may not draw the card they just discarded. + +Expedition numbers must be strictly increasing. Handshakes may be played only +before any number in that color. The round ends immediately when the final deck +card is drawn, or at `MAX_STEPS == 400`; forced termination is scored exactly +like natural termination. + +Scoring per player/color: + +```text +empty column: 0 +non-empty: (sum(number ranks) - 20) * (1 + handshake_count) +bonus: +20 if total column length >= 8, not multiplied +``` + +### Encodings + +Cards: + +| Field | Encoding | +| --- | --- | +| `card_id` | `color * 12 + slot` | +| `color` | `0..4` | +| `slot 0..2` | handshake | +| `slot 3..11` | ranks `2..10`, with `rank = slot - 1` | + +Actions (`N_ACTIONS == 96`): + +```text +action_id = hand_slot * 12 + place_type * 6 + draw_source +hand_slot = 0..7, current player's hand sorted by card_id +place_type = 0 play, 1 discard +draw_source = 0 deck, 1..5 discard pile color 0..4 +``` + +Observation (`OBS_DIM == 454`): + +- 60 cards x 7 one-hot channels: + my hand, my board, opponent board, discard top, discard non-top, + opponent public hand, unknown. +- 34 scalar features: + remaining deck `/44`, opponent unknown hand count `/8`, step count `/400`, + current-player then opponent `col_top /10`, `col_hs /3`, `col_len /12`, + and current board score difference `(player - opponent) /780`. + +### Verification + +```bash +uv run pytest -q tests/lost_cities_jax +uv run pytest -q +uv run ruff check . +``` + +Large differential profiles: + +```bash +# CI profile: 100,000 random legal-policy games +CI=1 uv run pytest -q tests/lost_cities_jax/test_differential.py + +# Full profile: 1,000,000 random legal-policy games +uv run pytest -q tests/lost_cities_jax/test_differential.py --full +``` + +Throughput benchmark: + +```bash +flock -n .compute.lock uv run python benchmarks/throughput.py +``` + +Measured on 2026-07-04 with CPU JAX backend: + +```text +backend=cpu +batch_size=8192 +steps=256 +elapsed_sec=4.944725 +steps_per_sec=424119.01 +``` + +### DECISIONS.md + +- Explicit `deck_order` dealing uses the first eight cards for player 0 and + the next eight for player 1. The remaining cards are drawn from index 16. + This is equivalent under a uniform shuffle and is fixed by tests. +- After a legal terminal transition, `to_move` is advanced to the next player, + but `done=True` makes all later steps complete no-ops. +- Terminal reward is emitted only on the transition that reaches `done=True`. + Done-state no-op steps return zero reward. +- Observation scalar normalization is implementation-defined as documented + above and locked by the exported `OBS_DIM`. + ## Basic Usage ```python diff --git a/benchmarks/throughput.py b/benchmarks/throughput.py new file mode 100644 index 0000000..6fb6672 --- /dev/null +++ b/benchmarks/throughput.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +import argparse +import time +from functools import partial + +import jax +import jax.numpy as jnp + +from lost_cities_jax.engine import legal_action_mask, reset, step + + +@partial(jax.jit, static_argnames=("steps",)) +def rollout(states, rng, *, steps: int): + def body(carry, _): + state, key = carry + key, action_key = jax.random.split(key) + mask = jax.vmap(legal_action_mask)(state) + logits = jnp.where(mask, 0.0, -1.0e9) + actions = jax.random.categorical(action_key, logits, axis=1).astype(jnp.int32) + state, _, _ = jax.vmap(step, in_axes=(0, 0))(state, actions) + return (state, key), None + + (states, rng), _ = jax.lax.scan(body, (states, rng), xs=None, length=steps) + return states, rng + + +def main() -> None: + parser = argparse.ArgumentParser(description="Lost Cities JAX random-policy throughput") + parser.add_argument("--batch-size", type=int, default=8192) + parser.add_argument("--steps", type=int, default=256) + parser.add_argument("--warmup-steps", type=int, default=32) + args = parser.parse_args() + + key = jax.random.PRNGKey(0) + reset_keys = jax.random.split(key, args.batch_size) + states = jax.jit(jax.vmap(reset))(reset_keys) + + warm_states, key = rollout(states, jax.random.PRNGKey(1), steps=args.warmup_steps) + jax.tree_util.tree_leaves(warm_states)[0].block_until_ready() + + start = time.perf_counter() + states, _ = rollout(states, key, steps=args.steps) + jax.tree_util.tree_leaves(states)[0].block_until_ready() + elapsed = time.perf_counter() - start + + transitions = args.batch_size * args.steps + print(f"backend={jax.default_backend()}") + print(f"batch_size={args.batch_size}") + print(f"steps={args.steps}") + print(f"elapsed_sec={elapsed:.6f}") + print(f"steps_per_sec={transitions / elapsed:.2f}") + + +if __name__ == "__main__": + main()