# 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: 1. A batched rollout driver using `batched_reset`, `batched_step`, `batched_legal_mask`, and `batched_obs`. 2. Three static pure-JAX opponent policies: `discard_only`, `heuristic_balanced`, and `heuristic_cautious`. 3. A 3 x 512 MLP actor-critic with 96-action policy logits and scalar value. 4. Standard PPO training against one static opponent at a time. 5. Duplicate evaluation on a fixed 10,000-deck shuffle bank. 6. 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: ```bash 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: ```bash uv run pytest -q tests/lost_cities_jax ``` Current benchmark evidence on RTX 3090: ```text backend=gpu batch_size=8192 steps=256 steps_per_sec=3956598.38 ``` ## Implementation Shape Prefer a small number of files while keeping testable boundaries: ```text 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: - `flax` for model modules and train state. - `optax` for Adam and PPO losses. - `orbax-checkpoint` for checkpointing. Add them through `uv add`, not pip/conda/poetry. ## Static Opponents All opponent policies are pure JAX functions: ```python 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: ```python 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_steps` termination 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_DIM` observation 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: ```yaml 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: ```text tanh((learner_score - opponent_score) / terminal_scale) ``` Potential shaping: ```text 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_order` permutations 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: ```bash uv run --with 'jax[cuda12]' lost-cities-jax-ppo train --config uv run --with 'jax[cuda12]' lost-cities-jax-ppo eval --config \ --checkpoint /latest --games 10000 --duplicate --output ``` 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: ```text /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 1. Add static opponent policies. 2. Add random-policy batched rollout. 3. Log canary metrics: returns, game length, max-step rate, `play_action_rate`, opened colors, positive expeditions. 4. Run CPU tests and one GPU random rollout. 5. Record the random baseline in README. Exit criteria: - deterministic smoke rollout passes; - canary metrics look finite and stable; - `play_action_rate` and game length are logged before any PPO code is judged. ### Phase 2: PPO Against `discard_only` 1. Add MLP actor-critic and masked action sampling. 2. Add GAE and PPO update. 3. Add Orbax checkpoints and resume. 4. Add CLI and the discard-only config under configs/jax_ppo. 5. Train on GPU. First-run checks: - with shaping enabled, `play_action_rate` rises above random baseline within the first few hundred thousand environment steps; - average game length converges near natural 44-ply games; - `max_steps` termination rate goes to zero. Then anneal shaping and run duplicate evaluation for gate 1. ### Phase 3: Static Ladder 1. Reuse the same config and checkpoint flow. 2. Change only the opponent config for `heuristic_balanced`. 3. Train/evaluate until gate 2 passes or diagnostics point to variance. 4. Repeat for `heuristic_cautious`. No self-play work starts before gate 3 passes. ## Commands Smoke rollout: ```bash uv run lost-cities-jax-ppo rollout-smoke --config configs/jax_ppo/discard-only.yaml ``` GPU training in tmux: ```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" ``` Duplicate evaluation: ```bash 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//latest \ --opponent discard_only \ --games 10000 \ --duplicate \ --output /mnt/2tbhdd/coolrl-lost-cities-artifacts/jax-ppo-static-opponents//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/2tbhdd` and keep only `latest` plus 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.