coolrl-lost-cities
Focused Lost Cities extraction from the legacy coolrl repository.
The current implementation starts with the classic two-player card game:
- classic 5-expedition rules by default
- Python/Cython game engine
- env wrapper
- random, discard-only, and safe-heuristic bots
- core rule, scoring, mask, env, canonical-state, bot, and GUI smoke tests
Training code, Deep CFR, learned-policy evaluation, desktop GUI, and an on-device web client now live alongside the original rules port.
Development
uv run pytest tests/games/classic
uv run lost-cities-classic
For future GUI work, install the optional GUI dependencies:
uv sync --extra gui
Run the classic pygame GUI:
uv run lost-cities-classic-gui --mode pvc --bot safe-heuristic
The GUI uses the in-process Cython game engine.
On-device Web Client
The web/ app runs its TypeScript rules engine and exported JAX PPO policy
entirely in the browser. It prefers WebGPU and falls back to WebAssembly.
Export a local Orbax checkpoint and start Vite:
uv run --with onnx scripts/export_jax_ppo_onnx.py \
--checkpoint /path/to/checkpoint \
--output web/public/models/jax-ppo.onnx
cd web
npm install
npm run dev
See web/README.md for tests and model parity tooling.
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:
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) -> Statereset_from_order(deck_order) -> Statelegal_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:
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):
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 opponentcol_top /10,col_hs /3,col_len /12, and current board score difference(player - opponent) /780.
Verification
uv run pytest -q tests/lost_cities_jax
uv run pytest -q
uv run ruff check .
Large differential profiles:
# 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:
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:
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:
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:
backend=gpu
batch_size=8192
steps=256
elapsed_sec=0.530039
steps_per_sec=3956598.38
DECISIONS.md
- Explicit
deck_orderdealing 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_moveis advanced to the next player, butdone=Truemakes 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:
uv run lost-cities-jax-ppo --help
CPU smoke:
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:
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:
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.
Basic Usage
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 for the current direction.