# coolrl-lost-cities JAX PPO training and on-device browser play for the two-player card game Lost Cities. The project provides a pure JAX rules engine, PPO training and evaluation tools, and a static web client that runs the shipped final policy with WebGPU when available and WebAssembly otherwise. The active path is **JAX + PPO**. Deep CFR and ISMCTS are retained only as historical research implementations; see [legacy notes](docs/legacy.md). ## Quick start Install the project, then run a complete CPU-sized sanity check: ```bash uv sync uv run lost-cities-jax-ppo rollout-smoke --config configs/jax_ppo/smoke.yaml uv run lost-cities-jax-ppo train --config configs/jax_ppo/smoke.yaml ``` The final command prints a run directory containing `latest`, `config.json`, and `metrics.jsonl` under `runs/tmp/jax-ppo-artifacts/`. ## Train and evaluate a PPO policy The committed configurations describe the opponent and training budget. For a GPU run, use CUDA JAX and keep generated artifacts outside git: ```bash flock -n .compute.lock uv run --with 'jax[cuda12]' lost-cities-jax-ppo train \ --config configs/jax_ppo/balanced.yaml \ --set run.artifact_root=runs/jax-ppo ``` Evaluate a saved checkpoint against a fixed opponent. Duplicate evaluation swaps seats over the same shuffled games: ```bash uv run lost-cities-jax-ppo eval \ --config configs/jax_ppo/balanced.yaml \ --checkpoint runs/jax-ppo//latest \ --opponent heuristic_balanced \ --games 10000 \ --duplicate ``` Run `uv run lost-cities-jax-ppo --help` for training against saved opponents, league runs, gates, human-play logs, and evaluation variants. ## Browser client The final verified JAX PPO policy is shipped as a 3.1 MB static ONNX asset. No server or local checkpoint is required to play it: ```bash cd web npm ci npm run dev ``` `npm run build` produces a fully static site. The model path is deployment-base aware, so the build can be served from GitHub Pages, GitLab Pages, or a normal web root. Pushes to `main` deploy that build to both configured Pages hosts. See [web/README.md](web/README.md) for model replacement, tests, and the cross-runtime parity fixture. ## 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 ``` Observed differential results on 2026-07-04 with CPU JAX backend: ```text CI=1 ... test_differential.py 1 passed in 247.61s (0:04:07) ... test_differential.py --full 1 passed in 2514.14s (0:41:54) elapsed=41:54.48 ``` Throughput benchmark: ```bash flock -n .compute.lock uv run python benchmarks/throughput.py # Optional CUDA check without making CUDA a project dependency: flock -n .compute.lock uv run --with 'jax[cuda12]' python benchmarks/throughput.py ``` Measured on 2026-07-04 with CPU JAX backend: ```text backend=cpu batch_size=8192 steps=256 elapsed_sec=4.655526 steps_per_sec=450465.14 ``` Measured on 2026-07-04 with CUDA JAX backend on RTX 3090, using the optional `uv run --with 'jax[cuda12]' ...` command: ```text backend=gpu batch_size=8192 steps=256 elapsed_sec=0.530039 steps_per_sec=3956598.38 ``` ### 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`. ## JAX PPO Static-Opponent Ladder The first training stack above `lost_cities_jax` is exposed as: ```bash uv run lost-cities-jax-ppo --help ``` CPU smoke: ```bash uv run lost-cities-jax-ppo rollout-smoke --config configs/jax_ppo/smoke.yaml uv run lost-cities-jax-ppo train --config configs/jax_ppo/smoke.yaml ``` GPU training uses optional CUDA JAX, keeping CUDA wheels out of the default project dependency set: ```bash tmux new-session -s coolrl-jax-ppo-discard \ -c /home/coolguy/dev/coolrl-lost-cities \ "flock -n .compute.lock uv run --with 'jax[cuda12]' lost-cities-jax-ppo train \ --config configs/jax_ppo/discard-only.yaml" ``` Random-policy baseline vs `discard_only`, measured on 2026-07-04 with 8192 games x 400 plies on GPU: ```text return_mean=-55.582763671875 game_length_mean=69.6085205078125 max_steps_rate=0.0 play_action_rate=0.28854578733444214 opened_colors_mean=4.942626953125 positive_expeditions_mean=0.41796875 ``` Static-opponent gate results, measured on 2026-07-04 with 10,000 fixed shuffles and duplicate seat-swapped evaluation: | Gate | Opponent | Result | Win rate (Wilson 95%) | Mean score diff | Mean length | Positive expeditions/game | | --- | --- | --- | --- | ---: | ---: | ---: | | 1 | `discard_only` | PASS | 1.00000 [0.99981, 1.00000] | 204.56335 | 82.45495 | 3.2783 | | 2 | `heuristic_balanced` | PASS | 0.98655 [0.98486, 0.98806] | 116.83400 | 167.81450 | 3.9573 | | 3 | `heuristic_cautious` | PASS | 0.95955 [0.95673, 0.96219] | 142.89930 | 185.38690 | 4.0946 | Large PPO artifacts are written under `/mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/`. Generated checkpoints and evaluation JSON are not committed to git. The full run summary is in [docs/reports/jax-ppo-static-opponent-ladder-2026-07-04.md](docs/reports/jax-ppo-static-opponent-ladder-2026-07-04.md). ## Basic Usage ```python from coolrl_lost_cities.games.classic import GameState, build_bot, classic_config state = GameState.new_game(classic_config(seed=1)) bot = build_bot("random", seed=1) while not state.terminal: state.apply_action(bot.act(state)) print(state.total_score(0), state.total_score(1)) ``` See [classic port notes](docs/archive/classic-port-notes.md) for the current direction.