12 KiB
Plan: JAX PPO Static-Opponent Ladder
Status: Static-opponent ladder passed on 2026-07-04.
Owner: Codex implements; operator reviews training gates.
Scope: A compact PPO training stack on top of lost_cities_jax, using GPU
via optional CUDA JAX execution.
Goal
Build the first learning layer above the JAX Lost Cities rules engine:
- A batched rollout driver using
batched_reset,batched_step,batched_legal_mask, andbatched_obs. - Three static pure-JAX opponent policies:
discard_only,heuristic_balanced, andheuristic_cautious. - A 3 x 512 MLP actor-critic with 96-action policy logits and scalar value.
- Standard PPO training against one static opponent at a time.
- Duplicate evaluation on a fixed 10,000-deck shuffle bank.
- Orbax checkpoints, metrics logging, and a CLI:
lost-cities-jax-ppo train --config ....
The near-term objective is not league self-play. It is to pass the static-opponent diagnostic ladder and preserve enough artifacts that a later self-play failure can be localized cleanly.
Non-Goals
- No league self-play, snapshot pool, Elo, or opponent matchmaking in this phase.
- No MCTS, CFR, search, or neural opponent ensemble.
- No multi-round match wrapper.
- No rule variants, expanded colors, or extra players.
- No large binary checkpoints committed to git. Generated artifacts are stored outside the code repo, with small manifests and summaries committed.
GPU Execution
Keep the project dependency portable. Do not make CUDA wheels mandatory in
pyproject.toml.
Use optional CUDA JAX for training and GPU benchmarks:
flock -n .compute.lock uv run --with 'jax[cuda12]' lost-cities-jax-ppo train \
--config configs/jax_ppo/discard-only.yaml
The CPU path must still work for tests:
uv run pytest -q tests/lost_cities_jax
Current benchmark evidence on RTX 3090:
backend=gpu
batch_size=8192
steps=256
steps_per_sec=3956598.38
Implementation Shape
Prefer a small number of files while keeping testable boundaries:
src/lost_cities_jax/
ppo.py # config, network, rollout, PPO update, train loop
opponents.py # pure-JAX static opponent policy functions
ppo_cli.py # argparse CLI
configs/jax_ppo/
discard-only.yaml
balanced.yaml
cautious.yaml
If the PPO implementation remains readable in one main file, keep it there. Split only when a file becomes hard to test or review.
Dependencies likely needed:
flaxfor model modules and train state.optaxfor Adam and PPO losses.orbax-checkpointfor checkpointing.
Add them through uv add, not pip/conda/poetry.
Static Opponents
All opponent policies are pure JAX functions:
policy_fn(state: State, player: jax.Array) -> jax.Array # int32 action
They must sample no Python-side randomness. If tie-breaking needs randomness, pass a JAX key explicitly:
policy_fn(state, player, rng) -> action
Policies:
discard_only: always discard a legal hand slot and draw from deck when legal. This opponent should make free expedition building easy.heuristic_balanced: prefer legal plays that improve expedition prospects, avoid obviously toxic openings, draw useful discard tops when available.heuristic_cautious: stricter opening threshold, fewer negative expedition commitments, more conservative discard/draw behavior.
Opponent policies are not learning targets. They are diagnostic fixtures.
Rollout Driver
Run 8192 games in parallel by default.
The learner controls one fixed seat per rollout batch. The opponent occupies the other seat. Because games alternate turns, each environment step chooses:
- learner action from the actor policy when
state.to_move == learner_seat; - static opponent action otherwise.
Done states remain no-op through the engine, so the rollout can keep a fixed
time axis. Use MAX_STEPS == 400 as the scan length unless a shorter config
value is explicitly introduced.
The first milestone is random-policy rollout with the full dashboard. This is not throwaway; it defines the baseline canaries.
Log canaries:
- episode return;
- score difference;
- game length distribution;
max_stepstermination rate;play_action_rate;- opened color count;
- positive expedition count per game;
- average entropy under legal-action masking.
Record the random-policy baseline in README.md before PPO training starts.
PPO Details
Network:
- Input:
OBS_DIMobservation vector. - Body: MLP 512 -> 512 -> 512, ReLU.
- Policy head: 96 logits.
- Value head: scalar.
- Illegal actions are masked to a large negative value before sampling and before log-prob/loss computation.
Training defaults:
ppo:
batch_games: 8192
rollout_steps: 400
gamma: 1.0
gae_lambda: 0.95
clip_epsilon: 0.2
entropy_coef: 0.01
value_coef: 0.5
max_grad_norm: 0.5
learning_rate: 0.0003
epochs: 4
minibatches: 128
reward:
terminal_scale: 50.0
potential_shaping_initial: 1.0
potential_shaping_final: 0.0
potential_shaping_anneal_steps: 5_000_000
Terminal reward:
tanh((learner_score - opponent_score) / terminal_scale)
Potential shaping:
coef(t) * (board_score_diff_after - board_score_diff_before)
The shaping coefficient anneals linearly from initial to final. Expose all
four shaping fields in config. This shaping exists to prevent the previous
play-action-rate collapse by giving immediate credit for board progress; its
schedule is a controlled experiment variable, not a hidden constant.
Evaluation Gates
Evaluation uses a fixed shuffle bank:
- Generate 10,000 explicit
deck_orderpermutations from a fixed seed. - For each deck, play twice:
- learner as player 0, opponent as player 1;
- opponent as player 0, learner as player 1.
- Aggregate duplicate-pair results.
Report:
- win rate;
- Wilson confidence interval;
- mean score difference;
- mean game length;
- positive expedition count per game;
- opened colors;
play_action_rate.
Gates:
| Gate | Opponent | Pass condition | Diagnosis if failed |
|---|---|---|---|
| 1 | discard_only |
win rate >= 90% and >= 2 positive expeditions/game | Audit reward pipeline, observations, and masks. Self-play is irrelevant. |
| 2 | heuristic_balanced |
mean score difference > 0 | Learner works; tune shaping schedule and batch/variance. |
| 3 | heuristic_cautious |
mean score difference > 0 | Same as gate 2; static ladder still not cleared. |
Only after all three gates pass should league self-play begin.
2026-07-04 Gate Results
All three static-opponent gates passed on an RTX 3090 using optional CUDA JAX:
uv run --with 'jax[cuda12]' lost-cities-jax-ppo train --config <config>
uv run --with 'jax[cuda12]' lost-cities-jax-ppo eval --config <config> \
--checkpoint <run>/latest --games 10000 --duplicate --output <eval.json>
Training used batch_games=8192, rollout_steps=400, total_updates=250,
seed=20260704, and artifact root
/mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/.
The implementation commit used for the runs was 4c0c2e9.
| 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 |
Artifact directories:
- Gate 1:
/mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/2026-07-04_223401_jax-ppo-discard-only/ - Gate 2:
/mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/2026-07-04_224749_jax-ppo-balanced/ - Gate 3:
/mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/2026-07-04_230150_jax-ppo-cautious/
Final training rollout canaries:
| Opponent | Final train return_mean | Final play_action_rate | Final max_steps_rate |
|---|---|---|---|
discard_only |
0.99281 | 0.59007 | 0.00000 |
heuristic_balanced |
0.93545 | 0.37737 | 0.00415 |
heuristic_cautious |
0.97101 | 0.37473 | 0.00391 |
The balanced and cautious policies win decisively but produce longer games than the discard-only diagnostic. That is not a gate failure, but the next training phase should keep game length and forced-end rate as canaries.
Artifact Policy
Generated training output stays out of git. Use:
/mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/
Per gate, preserve:
- checkpoint directory;
- resolved config;
- duplicate evaluation JSON;
- metrics JSONL;
- short summary markdown with command, git commit, seed, and pass/fail result.
Commit only small, durable metadata to the code repo:
- config files;
- evaluator/training code;
- README baseline and gate summaries;
- optional dated summary under docs/reports.
If binary artifact tracking is later required, add DVC/Git LFS explicitly instead of committing large checkpoint files directly.
Work Order
Phase 1: Rollout Dashboard
- Add static opponent policies.
- Add random-policy batched rollout.
- Log canary metrics:
returns, game length, max-step rate,
play_action_rate, opened colors, positive expeditions. - Run CPU tests and one GPU random rollout.
- Record the random baseline in README.
Exit criteria:
- deterministic smoke rollout passes;
- canary metrics look finite and stable;
play_action_rateand game length are logged before any PPO code is judged.
Phase 2: PPO Against discard_only
- Add MLP actor-critic and masked action sampling.
- Add GAE and PPO update.
- Add Orbax checkpoints and resume.
- Add CLI and the discard-only config under configs/jax_ppo.
- Train on GPU.
First-run checks:
- with shaping enabled,
play_action_raterises above random baseline within the first few hundred thousand environment steps; - average game length converges near natural 44-ply games;
max_stepstermination rate goes to zero.
Then anneal shaping and run duplicate evaluation for gate 1.
Phase 3: Static Ladder
- Reuse the same config and checkpoint flow.
- Change only the opponent config for
heuristic_balanced. - Train/evaluate until gate 2 passes or diagnostics point to variance.
- Repeat for
heuristic_cautious.
No self-play work starts before gate 3 passes.
Commands
Smoke rollout:
uv run lost-cities-jax-ppo rollout-smoke --config configs/jax_ppo/discard-only.yaml
GPU training in tmux:
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"
Duplicate evaluation:
uv run --with 'jax[cuda12]' lost-cities-jax-ppo eval \
--config configs/jax_ppo/discard-only.yaml \
--checkpoint /mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/<run>/latest \
--opponent discard_only \
--games 10000 \
--duplicate \
--output /mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents/<run>/eval_duplicate.json
Risks
- The engine is fast on GPU, but Python logging/evaluation can dominate if metrics are copied every step. Aggregate in JAX and transfer per rollout.
- Static heuristics can accidentally become too weak or too strong. Keep them deterministic and versioned by config.
- Potential shaping can teach score-chasing artifacts if it never anneals. Treat the coefficient schedule as part of the experiment identity.
- Checkpoint volume can grow quickly. Store large artifacts under
/mnt/2tbhddand keep onlylatestplus gate checkpoints unless a run is explicitly archived.
Current Next Action
The static-opponent ladder is complete. The next eligible phase is snapshot-pool league self-play design, using the three passed checkpoints as initial anchors.