Add JAX PPO ladder v2 expert pass

This commit is contained in:
2026-07-05 02:06:22 +09:00
parent 7fe0bfdcfe
commit 6037b650f3
11 changed files with 854 additions and 3 deletions
+40
View File
@@ -0,0 +1,40 @@
run:
experiment_name: jax-ppo-ladder-v2-balanced
seed: 20260704
learner_seat: 0
total_updates: 250
log_every: 1
checkpoint_every: 10
artifact_root: /mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2
opponent:
name: heuristic_balanced
network:
hidden_size: 512
num_layers: 3
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: 5000000
evaluation:
games: 10000
duplicate: true
shuffle_bank_seed: 20260704
batch_games: 8192
@@ -0,0 +1,40 @@
run:
experiment_name: jax-ppo-ladder-v2-discard-only
seed: 20260704
learner_seat: 0
total_updates: 250
log_every: 1
checkpoint_every: 10
artifact_root: /mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2
opponent:
name: discard_only
network:
hidden_size: 512
num_layers: 3
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: 5000000
evaluation:
games: 10000
duplicate: true
shuffle_bank_seed: 20260704
batch_games: 8192
+40
View File
@@ -0,0 +1,40 @@
run:
experiment_name: jax-ppo-ladder-v2-expert
seed: 20260704
learner_seat: 0
total_updates: 250
log_every: 1
checkpoint_every: 10
artifact_root: /mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2
opponent:
name: heuristic_expert
network:
hidden_size: 512
num_layers: 3
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: 5000000
evaluation:
games: 10000
duplicate: true
shuffle_bank_seed: 20260704
batch_games: 8192
+40
View File
@@ -0,0 +1,40 @@
run:
experiment_name: jax-ppo-ladder-v2-exploiter
seed: 20260705
learner_seat: 0
total_updates: 250
log_every: 1
checkpoint_every: 10
artifact_root: /mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2
opponent:
name: ladder_v2_gate3_checkpoint_frozen
network:
hidden_size: 512
num_layers: 3
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: 5000000
evaluation:
games: 10000
duplicate: true
shuffle_bank_seed: 20260704
batch_games: 8192
@@ -0,0 +1,139 @@
# Lost Cities JAX PPO Snapshot-Pool League Spec
**Status:** ready for implementation after ladder v2.
**Last updated:** 2026-07-05.
**Warm start:** `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_013223_jax-ppo-ladder-v2-expert/latest`.
## Preconditions
The static-opponent ladder v2 has passed:
- `discard_only`: win rate 1.0000, mean score diff +182.0670.
- `heuristic_balanced`: win rate 0.9409, mean score diff +119.9250.
- `heuristic_expert`: win rate 0.8382, mean score diff +43.3099, score-diff
CI95 lower bound +42.7065.
Canary warnings remain:
- Gate-3 opened colors are 4.5509/game, not the desired 2-3 range.
- The v2 gate-3 checkpoint is exploitable: a 250-update PPO exploiter reached
win rate 0.83135 and mean score diff +54.7721 against it.
## Anchor Semantics
Anchors are not certificates of strength. They serve two narrower purposes:
- Fix the Elo scale so curves remain comparable over time.
- Provide style diversity so the league does not train only against recent
policy snapshots.
Strength certification is tracked on separate axes:
- Duplicate mean score difference against `heuristic_expert`.
- Exploiter win-rate trend under the fixed exploiter protocol.
- Later human play, once the automated diagnostics are stable.
## Initial Pool
Permanent anchors are never removed:
- Static anchors: `discard_only`, `heuristic_balanced`, `heuristic_cautious`,
`heuristic_expert`.
- Learned ladder anchors:
- Gate 1: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_010429_jax-ppo-ladder-v2-discard-only/latest`
- Gate 2: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_011827_jax-ppo-ladder-v2-balanced/latest`
- Gate 3: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_013223_jax-ppo-ladder-v2-expert/latest`
The learning policy starts from the v2 gate-3 checkpoint. Shaping coefficient is
fixed at 0 by default. If warm-start collapse is observed, the implementation
may expose a config switch for random initialization plus shaping anneal, but
that is not the default path.
## Opponent Sampling
For each episode:
- Randomize the learner seat.
- With probability 50%, play mirror self-play against the current policy copy.
- With probability 50%, sample from the pool.
Pool sampling is PFSP-like:
- Track recent win rate versus each pool member.
- Weight non-anchor snapshots by `(1 - win_rate)^2`.
- Mix in a small uniform component to avoid starvation.
- Give every permanent anchor only a small floor probability, such as 2-3% per
anchor, so stalling or weak anchors do not waste most throughput.
- Allocate the remaining probability mass to snapshots according to the PFSP
weights.
## Snapshot Lifecycle
- Add the current policy to the pool at a configurable interval.
- If the pool size cap is reached, remove the oldest non-anchor snapshot.
- Never remove permanent anchors.
- Store large checkpoints and raw evaluation JSON under
`/mnt/2tbhdd/coolrl-lost-cities-artifacts/league/`.
- Store one-line tracked JSON summaries under `docs/reports/` so results do
not disappear with artifact cleanup.
## Evaluation
At each snapshot interval, evaluate the current policy with fixed shuffle-bank
duplicate play against:
- The six ladder anchors.
- `heuristic_expert`.
- A configurable set of recent snapshots.
Update logistic Elo from the pairwise results, using the heuristic anchors as
fixed reference points. Elo is a scale diagnostic, not a strength certificate.
Run the exploiter protocol periodically against the current best checkpoint:
- Randomly initialized PPO.
- Same 250-update budget used in ladder v2 unless explicitly overridden.
- Duplicate evaluation versus the frozen target checkpoint.
- Track exploiter win rate and mean score difference over time.
## Canary Metrics
Log these metrics at every evaluation:
- Opened colors per game.
- Play action rate.
- Mean and quantile game length.
- Max-steps rate.
- Positive expeditions per game.
- Duplicate score-diff distribution against `heuristic_expert`.
Canaries are diagnostics only. Do not put them directly into the reward.
## Regression Guards
- `heuristic_expert` replaces `heuristic_cautious` as the main regression
opponent.
- If win rate versus `heuristic_expert` falls below 90% of the v2 gate-3 value,
flag the snapshot.
- If max-steps rate rises above the gate-3 baseline of 0.0 by a material
amount, flag the snapshot.
- If exploiter win rate rises above the v2 baseline of 0.83135, flag the
snapshot.
## Stop Conditions
Stop the league run when either:
- Elo is statistically flat over a configured recent-snapshot window and
exploiter win rate is at or below the configured threshold, for example 0.55.
- The wall-clock budget is exhausted.
## Decisions
- The earlier `heuristic_cautious` remains a permanent anchor for scale and
style diversity only. It is a stalling anchor and no longer certifies
strength.
- `discard_only` is also kept only as a low-probability anchor and regression
sanity check.
- The league spec file was not present in the repository when this update was
requested, so this tracked plan file is the authoritative spec location.
@@ -0,0 +1 @@
{"date":"2026-07-05","status":"pass_with_canary_warnings","expert_bot":{"mirror_max_steps_rate":0.0,"mirror_opened_colors_per_game":2.256,"mirror_play_action_rate":0.3117858091947914,"vs_discard_mean_score_diff":8.9716,"vs_balanced_mean_score_diff":54.36725,"vs_cautious_mean_score_diff":51.9397},"gates":{"discard_only":{"win_rate":1.0,"wilson_low":0.9998079639438954,"mean_score_diff":182.067,"opened_colors_per_game":4.2163,"max_steps_rate":0.0},"heuristic_balanced":{"win_rate":0.9409,"wilson_low":0.9375464241082592,"mean_score_diff":119.925,"opened_colors_per_game":4.9862,"max_steps_rate":0.1859},"heuristic_expert":{"win_rate":0.83815,"wilson_low":0.8329806740267004,"mean_score_diff":43.3099,"score_diff_ci95_low":42.706527573078894,"opened_colors_per_game":4.5509,"max_steps_rate":0.0}},"exploiter":{"target":"ladder_v2_gate3","win_rate":0.83135,"wilson_low":0.8260970642435626,"mean_score_diff":54.7721,"score_diff_ci95_low":53.98710920050297,"opened_colors_per_game":4.99135,"max_steps_rate":0.0001},"artifact_root":"/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/"}
+114
View File
@@ -0,0 +1,114 @@
# JAX PPO Ladder v2 Report - 2026-07-05
**Status:** PASS for the ordered ladder gates; canary warnings remain.
**Hardware:** RTX 3090 via `uv run --with 'jax[cuda12]'`.
**Artifact roots:**
- Expert bot: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/expert-bot/2026-07-05/`
- Ladder v2: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/`
## Protocol
Part A added a pure-JAX `heuristic_expert` opponent and evaluated it with a
fixed shuffle bank plus duplicate play. Part B trained PPO from random
initialization for each gate with the same 250-update configuration used by the
original ladder, changing only the opponent and artifact root:
- Gate 1: `discard_only`
- Gate 2: `heuristic_balanced`
- Gate 3: `heuristic_expert`
All evaluations below used 10,000 deck orders with duplicate seat-swapped play,
for 20,000 evaluated games.
## Part A - Expert Bot Gate
| Check | Result | Win rate | Mean diff | CI95 diff | Length | Max-step | Opened colors | Play rate |
| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
| Expert mirror, 1,000 duplicate pairs | PASS | 0.4815 | 0.0000 | [-0.6956, +0.6956] | 45.1560 | 0.0000 | 2.2560 | 0.3118 |
| Expert vs `discard_only` | PASS | 0.6643 | +8.9716 | [+8.7886, +9.1546] | 45.2125 | 0.0000 | 1.8971 | 0.2828 |
| Expert vs `heuristic_balanced` | PASS | 0.9961 | +54.3673 | [+54.0791, +54.6554] | 64.9181 | 0.0000 | 2.8003 | 0.2728 |
| Expert vs `heuristic_cautious` | PASS | 0.9972 | +51.9397 | [+51.6960, +52.1834] | 75.2329 | 0.0000 | 2.6090 | 0.2240 |
The new script bot is non-stalling in mirror play, symmetric under duplicate
seat swap, and beats all three previous static bots by a clear positive score
margin. Its own opened-color profile is in the intended 2-3 color band.
## Part B - Ladder v2
| Gate | Opponent | Result | Win rate (Wilson 95%) | Mean diff | CI95 diff | Length | Max-step | Opened colors | Play rate | Positive exp. |
| --- | --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
| 1 | `discard_only` | PASS | 1.0000 [0.9998, 1.0000] | +182.0670 | [+181.6390, +182.4950] | 82.1558 | 0.0000 | 4.2163 | 0.6398 | 3.1845 |
| 2 | `heuristic_balanced` | PASS | 0.9409 [0.9375, 0.9441] | +119.9250 | [+118.9950, +120.8550] | 151.2468 | 0.1859 | 4.9862 | 0.2937 | 3.5316 |
| 3 | `heuristic_expert` | PASS | 0.8382 [0.8330, 0.8432] | +43.3099 | [+42.7065, +43.9133] | 61.4025 | 0.0000 | 4.5509 | 0.6661 | 2.6121 |
Gate 1 passed the original win-rate and positive-expedition conditions. Gate 2
passed the mean-score-difference condition. Gate 3 passed the v2 condition:
duplicate mean score difference is positive and its 95% confidence lower bound
is above zero.
## Canary Findings
The ladder did not produce the desired 2-3 color selectivity in the learned PPO
policies. The gate-3 policy still opens 4.5509 colors/game against the expert.
This is lower than the original gate-3 checkpoint but still far outside the
target band. Treat the selectivity hypothesis as not confirmed.
Gate 2 also has a high forced-end canary: `max_steps_rate=0.1859` in duplicate
evaluation versus `heuristic_balanced`. It passed its score gate, but the
balanced opponent still permits long games and 5-color farming.
## Exploiter Baseline
After gate 3 passed, a new PPO exploiter was trained from random initialization
against the frozen v2 gate-3 checkpoint with the same 250-update budget and
duplicate-evaluated against that checkpoint.
| Target checkpoint | Exploiter win rate (Wilson 95%) | Mean diff | CI95 diff | Length | Max-step | Opened colors | Play rate | Positive exp. |
| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
| Old gate-3 checkpoint | 0.6712 [0.6646, 0.6776] | +24.1371 | n/a | 48.0017 | 0.0000 | 4.8971 | 0.7461 | 2.1287 |
| New v2 gate-3 checkpoint | 0.8314 [0.8261, 0.8365] | +54.7721 | [+53.9871, +55.5571] | 60.8744 | 0.0001 | 4.9914 | 0.7028 | 2.8921 |
The new gate-3 checkpoint beats `heuristic_expert`, but is more exploitable
under this protocol than the old checkpoint. That does not invalidate the gate,
but it makes exploiter reduction a primary objective for the league phase.
## Old vs New Gate-3 Canary Comparison
| Metric | Old gate-3 vs cautious | New gate-3 vs expert | Interpretation |
| --- | ---: | ---: | --- |
| Mean score diff | +142.8993 | +43.3099 | New gate is harder and less suspiciously clean. |
| Win rate | 0.9596 | 0.8382 | New result is less inflated. |
| Opened colors | 4.9896 | 4.5509 | Improved, but still not selective enough. |
| Max-step rate | not recorded in old eval | 0.0000 | New expert evaluation does not stall. |
| Play action rate | 0.2763 | 0.6661 | New policy plays much more actively. |
| Exploiter win rate | 0.6712 | 0.8314 | New checkpoint is currently more exploitable. |
## Artifacts
| Item | Path |
| --- | --- |
| Expert mirror JSON | `/mnt/2tbhdd/coolrl-lost-cities-artifacts/expert-bot/2026-07-05/expert-mirror-1000-duplicate.json` |
| Expert vs old bots JSONs | `/mnt/2tbhdd/coolrl-lost-cities-artifacts/expert-bot/2026-07-05/` |
| Gate 1 run | `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_010429_jax-ppo-ladder-v2-discard-only/` |
| Gate 2 run | `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_011827_jax-ppo-ladder-v2-balanced/` |
| Gate 3 run | `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_013223_jax-ppo-ladder-v2-expert/` |
| Exploiter run | `/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_014713_jax-ppo-ladder-v2-exploiter/` |
## Decisions
- The requested `lost-cities-league-selfplay-spec.md` file was not present in
the repository. I created `docs/plans/lost-cities-league-selfplay-spec.md` as
the tracked league spec location, following the repository docs routing for
active plans.
- Gate 2's high max-step rate is treated as a canary warning rather than a
gate failure because the v2 instruction keeps gate 2's original score-based
pass condition.
## Next
Proceed to league implementation only with the updated interpretation: anchors
fix the Elo scale and provide style diversity, but they do not certify strength.
The league's main success criterion should be lowering exploiter win rate from
the v2 baseline of `0.83135` while preserving positive duplicate performance
against `heuristic_expert`.
+6
View File
@@ -19,9 +19,12 @@ from lost_cities_jax.engine import (
)
from lost_cities_jax.obs import observation
from lost_cities_jax.opponents import (
HeuristicExpertConfig,
discard_only_action,
heuristic_balanced_action,
heuristic_cautious_action,
heuristic_expert_action,
make_heuristic_expert_policy,
policy_by_name,
random_legal_action,
)
@@ -36,6 +39,7 @@ __all__ = [
"MAX_STEPS",
"N_ACTIONS",
"OBS_DIM",
"HeuristicExpertConfig",
"State",
"batched_legal_mask",
"batched_obs",
@@ -48,7 +52,9 @@ __all__ = [
"discard_only_action",
"heuristic_balanced_action",
"heuristic_cautious_action",
"heuristic_expert_action",
"legal_action_mask",
"make_heuristic_expert_policy",
"observation",
"policy_by_name",
"random_legal_action",
+291 -1
View File
@@ -3,15 +3,20 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
import jax
import jax.numpy as jnp
from lost_cities_jax.engine import current_hand_sorted, legal_action_mask
from lost_cities_jax.engine import board_score, current_hand_sorted, legal_action_mask
from lost_cities_jax.types import (
CARDS_PER_COLOR,
DECK_DRAWS,
DISCARD,
DRAW_DECK,
LOC_DISCARD,
LOC_P0_BOARD,
LOC_P0_HAND,
N_ACTIONS,
N_CARDS,
N_COLORS,
@@ -27,6 +32,37 @@ ACTION_PLACE_TYPE = (ACTION_IDS % 12) // 6
ACTION_DRAW_SOURCE = ACTION_IDS % 6
@dataclass(frozen=True)
class HeuristicExpertConfig:
target_open_colors: int = 3
fourth_color_margin: float = 38.0
extra_color_margin: float = 58.0
base_open_threshold: float = 12.0
late_open_threshold: float = 26.0
future_inflow_scale: float = 0.16
handshake_risk_weight: float = 0.85
open_value_scale: float = 1.15
followup_base: float = 44.0
followup_rank_weight: float = 2.4
followup_gap_penalty: float = 8.0
hs_followup_base: float = 22.0
new_open_base: float = 18.0
play_bias: float = 8.0
defensive_sacrifice_bonus: float = 0.0
discard_base: float = 20.0
discard_rank_penalty: float = 1.05
discard_own_open_penalty: float = 16.0
discard_own_potential_penalty: float = 0.42
denial_immediate_penalty: float = 24.0
denial_near_penalty: float = 8.0
denial_hs_multiplier: float = 0.8
deck_draw_base: float = 4.5
clock_draw_weight: float = 3.0
useful_pile_draw_base: float = 13.0
useful_pile_draw_rank_weight: float = 1.0
useless_pile_draw_penalty: float = 36.0
def policy_by_name(name: str) -> ActionPolicy:
normalized = name.replace("-", "_")
if normalized == "discard_only":
@@ -35,6 +71,8 @@ def policy_by_name(name: str) -> ActionPolicy:
return heuristic_balanced_action
if normalized == "heuristic_cautious":
return heuristic_cautious_action
if normalized == "heuristic_expert":
return heuristic_expert_action
msg = f"unknown opponent policy: {name}"
raise ValueError(msg)
@@ -60,6 +98,23 @@ def heuristic_cautious_action(state: State, player: jax.Array, rng: jax.Array) -
return _heuristic_action(state, player, cautious=True)
def heuristic_expert_action(state: State, player: jax.Array, rng: jax.Array) -> jax.Array:
del rng
return _heuristic_expert_action(state, player, HeuristicExpertConfig())
def make_heuristic_expert_policy(
config: HeuristicExpertConfig | None = None,
) -> ActionPolicy:
cfg = HeuristicExpertConfig() if config is None else config
def policy(state: State, player: jax.Array, rng: jax.Array) -> jax.Array:
del rng
return _heuristic_expert_action(state, player, cfg)
return policy
def random_legal_action(state: State, player: jax.Array, rng: jax.Array) -> jax.Array:
del player
mask = legal_action_mask(state)
@@ -111,6 +166,238 @@ def _heuristic_action(state: State, player: jax.Array, *, cautious: bool) -> jax
return _argmax_legal(mask, scores)
def _heuristic_expert_action(
state: State, player: jax.Array, cfg: HeuristicExpertConfig
) -> jax.Array:
player = player.astype(jnp.int32)
opponent = 1 - player
mask = legal_action_mask(state)
hand = current_hand_sorted(state, player).astype(jnp.int32)
action_cards = hand[ACTION_HAND_SLOT]
valid_card = action_cards < N_CARDS
safe_cards = jnp.clip(action_cards, 0, N_CARDS - 1)
color = safe_cards // CARDS_PER_COLOR
slot = safe_cards % CARDS_PER_COLOR
is_handshake = slot < 3
rank_i = jnp.where(is_handshake, 0, slot - 1)
rank = rank_i.astype(jnp.float32)
hand_numeric_sum, hand_hs, hand_high_count = _hand_color_features(hand)
unseen_numeric = _unseen_numeric_sum(state, player, hand)
deck_frac = (N_CARDS - state.draw_ptr).astype(jnp.float32) / DECK_DRAWS
future_inflow = unseen_numeric * deck_frac * cfg.future_inflow_scale
multiplier = 1.0 + hand_hs
hand_total = hand_numeric_sum + future_inflow
downside = jnp.maximum(20.0 - hand_numeric_sum, 0.0)
open_value = (
(hand_total - 20.0) * multiplier
- downside * (1.0 + hand_hs * cfg.handshake_risk_weight)
+ hand_high_count * 2.5
)
progress = (state.draw_ptr - 16).astype(jnp.float32) / DECK_DRAWS
active_colors = jnp.sum(state.col_len[player] > 0).astype(jnp.float32)
open_threshold = (
cfg.base_open_threshold
+ progress * cfg.late_open_threshold
+ jnp.maximum(active_colors - 2.0, 0.0) * 9.0
)
color_open_value = open_value[color]
concentration_margin = jnp.where(
active_colors < cfg.target_open_colors,
0.0,
jnp.where(active_colors < 4.0, cfg.fourth_color_margin, cfg.extra_color_margin),
)
ev_open_allowed = color_open_value > (open_threshold + concentration_margin)
tempo_open_allowed = (
(active_colors < cfg.target_open_colors)
& ~is_handshake
& (rank >= 3.0)
& (rank <= 6.0)
& (color_open_value > -4.0)
)
fourth_tempo_allowed = (
(active_colors < 4.0) & ~is_handshake & (rank >= 8.0) & (color_open_value > 8.0)
)
open_allowed = ev_open_allowed | tempo_open_allowed | fourth_tempo_allowed
hs_open_allowed = (
open_allowed
& (hand_numeric_sum[color] >= 18.0)
& (hand_numeric_sum[color] + future_inflow[color] > 28.0)
& (progress < 0.55)
)
hs_followup_allowed = (
(hand_numeric_sum[color] >= 18.0)
& (hand_numeric_sum[color] + future_inflow[color] > 28.0)
& (progress < 0.55)
)
top = state.col_top[player, color].astype(jnp.float32)
hs_count = state.col_hs[player, color].astype(jnp.float32)
gap = jnp.maximum(rank - top - 1.0, 0.0)
followup_number_score = (
cfg.followup_base
+ rank * cfg.followup_rank_weight
- gap * cfg.followup_gap_penalty
+ jnp.maximum(color_open_value, -12.0) * 0.12
)
followup_hs_score = (
cfg.hs_followup_base
+ jnp.maximum(color_open_value, -10.0) * 0.10
- hs_count * 4.0
- progress * 10.0
)
new_open_number_score = (
cfg.new_open_base + color_open_value * cfg.open_value_scale - rank * 0.35
)
new_open_hs_score = (
cfg.new_open_base - 14.0 + color_open_value * 0.9 + hand_hs[color] * 1.0 - progress * 8.0
)
own_open = state.col_len[player, color] > 0
discard_danger = _discard_danger(state, opponent, safe_cards, cfg)
defensive_sacrifice = (
(discard_danger > 22.0)
& ~own_open
& (active_colors < cfg.target_open_colors)
& (rank_i <= 5)
& (color_open_value > -18.0)
)
new_color_allowed = jnp.where(is_handshake, hs_open_allowed, open_allowed)
play_score = jnp.where(
own_open,
jnp.where(
is_handshake,
jnp.where(hs_followup_allowed, followup_hs_score, -60.0),
followup_number_score,
),
jnp.where(
new_color_allowed | defensive_sacrifice,
jnp.where(is_handshake, new_open_hs_score, new_open_number_score),
-60.0,
),
)
play_score = jnp.where(
defensive_sacrifice, play_score + cfg.defensive_sacrifice_bonus, play_score
)
discard_score = (
cfg.discard_base
- rank * cfg.discard_rank_penalty
- own_open.astype(jnp.float32) * cfg.discard_own_open_penalty
- jnp.maximum(color_open_value, 0.0) * cfg.discard_own_potential_penalty
- discard_danger
)
discard_score = jnp.where(is_handshake, discard_score - 5.0, discard_score)
place_score = jnp.where(ACTION_PLACE_TYPE == PLAY, play_score + cfg.play_bias, discard_score)
draw_score = _expert_draw_score(state, player, ACTION_DRAW_SOURCE, open_value, cfg)
scores = jnp.where(valid_card, place_score + draw_score, -1.0e9)
scores = scores - ACTION_IDS.astype(jnp.float32) * 1.0e-4
return _argmax_legal(mask, scores)
def _hand_color_features(hand: jax.Array) -> tuple[jax.Array, jax.Array, jax.Array]:
valid = hand < N_CARDS
safe = jnp.clip(hand, 0, N_CARDS - 1)
colors = safe // CARDS_PER_COLOR
slots = safe % CARDS_PER_COLOR
is_hs = (slots < 3) & valid
ranks = jnp.where((slots >= 3) & valid, slots - 1, 0).astype(jnp.float32)
numeric_sum = []
hs_count = []
high_count = []
for color in range(N_COLORS):
in_color = colors == color
numeric_sum.append(jnp.sum(jnp.where(in_color, ranks, 0.0)))
hs_count.append(jnp.sum(jnp.where(in_color & is_hs, 1.0, 0.0)))
high_count.append(jnp.sum(jnp.where(in_color & (ranks >= 7.0), 1.0, 0.0)))
return jnp.stack(numeric_sum), jnp.stack(hs_count), jnp.stack(high_count)
def _unseen_numeric_sum(state: State, player: jax.Array, hand: jax.Array) -> jax.Array:
ids = jnp.arange(N_CARDS, dtype=jnp.int32)
colors = ids // CARDS_PER_COLOR
slots = ids % CARDS_PER_COLOR
ranks = jnp.where(slots >= 3, slots - 1, 0).astype(jnp.float32)
loc = state.card_loc.astype(jnp.int32)
visible = (
(loc == LOC_DISCARD)
| (loc == LOC_P0_BOARD)
| (loc == (LOC_P0_BOARD + 1))
| (loc == (LOC_P0_HAND + player.astype(jnp.int32)))
)
visible = visible | state.hand_public
unseen_sum = []
for color in range(N_COLORS):
unseen_sum.append(jnp.sum(jnp.where((colors == color) & ~visible, ranks, 0.0)))
return jnp.stack(unseen_sum)
def _discard_danger(
state: State, opponent: jax.Array, card: jax.Array, cfg: HeuristicExpertConfig
) -> jax.Array:
color = card // CARDS_PER_COLOR
slot = card % CARDS_PER_COLOR
is_hs = slot < 3
rank = jnp.where(is_hs, 0, slot - 1)
opp_top = state.col_top[opponent, color].astype(jnp.int32)
opp_len = state.col_len[opponent, color].astype(jnp.int32)
opp_hs = state.col_hs[opponent, color].astype(jnp.float32)
immediate = jnp.where(is_hs, opp_top == 0, (opp_len > 0) & (rank > opp_top))
near = jnp.where(is_hs, opp_top == 0, rank > opp_top)
multiplier = 1.0 + opp_hs * cfg.denial_hs_multiplier
return (
immediate.astype(jnp.float32) * cfg.denial_immediate_penalty * multiplier
+ (~immediate & near).astype(jnp.float32) * cfg.denial_near_penalty * multiplier
)
def _expert_draw_score(
state: State,
player: jax.Array,
draw_source: jax.Array,
open_value: jax.Array,
cfg: HeuristicExpertConfig,
) -> jax.Array:
scores = board_score(state)
opponent = 1 - player.astype(jnp.int32)
score_diff = scores[player.astype(jnp.int32)] - scores[opponent]
ahead = jnp.tanh(score_diff / 40.0)
deck_score = jnp.where(
draw_source == DRAW_DECK,
cfg.deck_draw_base + ahead * cfg.clock_draw_weight,
0.0,
)
pile_color = jnp.clip(draw_source - 1, 0, N_COLORS - 1)
pile_len = state.pile_len[pile_color].astype(jnp.int32)
top_idx = jnp.maximum(pile_len - 1, 0)
top_card = jnp.clip(state.pile[pile_color, top_idx].astype(jnp.int32), 0, N_CARDS - 1)
slot = top_card % CARDS_PER_COLOR
is_hs = slot < 3
rank = jnp.where(is_hs, 0, slot - 1)
color = top_card // CARDS_PER_COLOR
own_top = state.col_top[player.astype(jnp.int32), color].astype(jnp.int32)
own_open = state.col_len[player.astype(jnp.int32), color] > 0
can_play_open = jnp.where(is_hs, own_top == 0, rank > own_top)
worth_opening = open_value[color] > 18.0
high_rank_plan = (~is_hs) & (rank >= 7) & (open_value[color] > -4.0)
useful = (pile_len > 0) & can_play_open & (own_open | worth_opening | high_rank_plan)
pile_score = jnp.where(
draw_source > 0,
jnp.where(
useful,
cfg.useful_pile_draw_base
+ rank.astype(jnp.float32) * cfg.useful_pile_draw_rank_weight
- ahead * cfg.clock_draw_weight,
-cfg.useless_pile_draw_penalty,
),
0.0,
)
return deck_score + pile_score
def _draw_source_score(state: State, player: jax.Array, draw_source: jax.Array) -> jax.Array:
deck_score = jnp.where(draw_source == DRAW_DECK, 1.0, 0.0)
pile_color = jnp.clip(draw_source - 1, 0, N_COLORS - 1)
@@ -138,9 +425,12 @@ def _argmax_legal(mask: jax.Array, scores: jax.Array) -> jax.Array:
__all__ = [
"ActionPolicy",
"HeuristicExpertConfig",
"discard_only_action",
"heuristic_balanced_action",
"heuristic_cautious_action",
"heuristic_expert_action",
"make_heuristic_expert_policy",
"policy_by_name",
"random_legal_action",
]
+55 -1
View File
@@ -675,9 +675,11 @@ def evaluate_checkpoint_vs_static(
duplicate=duplicate,
shuffle_bank_seed=learner_cfg.evaluation.shuffle_bank_seed,
batch_games=learner_cfg.evaluation.batch_games,
output=output,
output=None,
)
result["opponent"] = opponent_name
if output:
_write_json(Path(output), result)
return result
@@ -702,6 +704,32 @@ def evaluate_static_mirror(
)
def evaluate_static_match(
learner_name: str,
opponent_name: str,
*,
games: int,
duplicate: bool = True,
shuffle_bank_seed: int = 20260704,
batch_games: int = 8192,
output: str | Path | None = None,
) -> dict:
result = evaluate_policy_match(
policy_by_name(learner_name),
policy_by_name(opponent_name),
games=games,
duplicate=duplicate,
shuffle_bank_seed=shuffle_bank_seed,
batch_games=batch_games,
output=None,
)
result["learner"] = learner_name
result["opponent"] = opponent_name
if output:
_write_json(Path(output), result)
return result
def evaluate_policy_match(
learner_policy,
opponent_policy,
@@ -745,6 +773,8 @@ def evaluate_policy_match(
ties = int(np.sum(diff_arr == 0))
wilson_low, wilson_high = wilson_interval(wins, total_games)
quantiles = np.quantile(diff_arr, [0.0, 0.05, 0.25, 0.5, 0.75, 0.95, 1.0])
diff_std = float(np.std(diff_arr, ddof=1)) if total_games > 1 else 0.0
diff_se = diff_std / math.sqrt(total_games) if total_games > 0 else 0.0
hist_counts, hist_edges = np.histogram(
diff_arr,
bins=[-400, -300, -200, -100, -50, 0, 50, 100, 200, 300, 400, 600],
@@ -759,6 +789,9 @@ def evaluate_policy_match(
"wilson_low": wilson_low,
"wilson_high": wilson_high,
"mean_score_diff": float(np.mean(diff_arr)),
"score_diff_std": diff_std,
"score_diff_ci95_low": float(np.mean(diff_arr) - 1.959963984540054 * diff_se),
"score_diff_ci95_high": float(np.mean(diff_arr) + 1.959963984540054 * diff_se),
"score_diff_sum": float(np.sum(diff_arr)),
"score_diff_quantiles": {
"min": float(quantiles[0]),
@@ -1015,6 +1048,15 @@ def cli_main(argv: list[str] | None = None) -> None:
static_mirror_parser.add_argument("--duplicate", action="store_true")
static_mirror_parser.add_argument("--output")
static_match_parser = sub.add_parser("eval-static-match")
static_match_parser.add_argument("--learner", required=True)
static_match_parser.add_argument("--opponent", required=True)
static_match_parser.add_argument("--games", type=int, default=10_000)
static_match_parser.add_argument("--batch-games", type=int, default=8192)
static_match_parser.add_argument("--shuffle-bank-seed", type=int, default=20260704)
static_match_parser.add_argument("--duplicate", action="store_true")
static_match_parser.add_argument("--output")
args = parser.parse_args(argv)
cfg = (
load_config(args.config, overrides=parse_overrides(args.set))
@@ -1086,6 +1128,17 @@ def cli_main(argv: list[str] | None = None) -> None:
output=args.output,
)
print(json.dumps(result, indent=2, sort_keys=True))
elif args.command == "eval-static-match":
result = evaluate_static_match(
args.learner,
args.opponent,
games=args.games,
duplicate=args.duplicate,
shuffle_bank_seed=args.shuffle_bank_seed,
batch_games=args.batch_games,
output=args.output,
)
print(json.dumps(result, indent=2, sort_keys=True))
def _create_run_dir(cfg: JaxPPOConfig) -> Path:
@@ -1166,6 +1219,7 @@ __all__ = [
"evaluate_checkpoint_match",
"evaluate_checkpoint_vs_static",
"evaluate_policy_match",
"evaluate_static_match",
"evaluate_static_mirror",
"load_config",
"random_rollout",
+88 -1
View File
@@ -13,6 +13,7 @@ from lost_cities_jax.opponents import (
discard_only_action,
heuristic_balanced_action,
heuristic_cautious_action,
heuristic_expert_action,
)
from lost_cities_jax.ppo import (
JaxPPOConfig,
@@ -65,7 +66,12 @@ def tiny_config(tmp_path) -> JaxPPOConfig:
def test_static_opponents_return_legal_actions():
state = reset(jax.random.PRNGKey(0))
mask = np.asarray(legal_action_mask(state), dtype=bool)
for fn in [discard_only_action, heuristic_balanced_action, heuristic_cautious_action]:
for fn in [
discard_only_action,
heuristic_balanced_action,
heuristic_cautious_action,
heuristic_expert_action,
]:
action = int(fn(state, jnp.int32(0), jax.random.PRNGKey(1)))
assert mask[action]
@@ -112,6 +118,66 @@ def test_cautious_uses_own_board_when_playing_as_p1():
assert hand[hand_slot] == playable
def test_expert_opens_strong_ev_hand():
strong_cards = [
_hs_card(0, 0),
_rank_card(0, 7),
_rank_card(0, 8),
_rank_card(0, 9),
_rank_card(0, 10),
_rank_card(1, 2),
_rank_card(2, 2),
_rank_card(3, 2),
]
state = _manual_state(p0_hand=strong_cards, to_move=0)
action = int(heuristic_expert_action(state, jnp.int32(0), jax.random.PRNGKey(1)))
hand_slot, place_type, _ = [int(x) for x in decode_action(jnp.asarray(action))]
hand = [int(x) for x in current_hand_sorted(state, 0)]
assert place_type == PLAY
assert hand[hand_slot] // CARDS_PER_COLOR == 0
def test_expert_rejects_weak_new_open():
weak_cards = [
_rank_card(0, 2),
_rank_card(0, 3),
_rank_card(0, 4),
_rank_card(1, 2),
_rank_card(2, 2),
_rank_card(3, 2),
_rank_card(4, 2),
_rank_card(4, 3),
]
state = _manual_state(p0_hand=weak_cards, to_move=0)
action = int(heuristic_expert_action(state, jnp.int32(0), jax.random.PRNGKey(1)))
_, place_type, _ = [int(x) for x in decode_action(jnp.asarray(action))]
assert place_type != PLAY
def test_expert_avoids_discarding_immediately_useful_opponent_card():
dangerous = _rank_card(0, 8)
safe = _rank_card(4, 2)
state = _manual_state(
p0_hand=[
dangerous,
safe,
_rank_card(1, 2),
_rank_card(1, 3),
_rank_card(2, 2),
_rank_card(2, 3),
_rank_card(3, 2),
_rank_card(3, 3),
],
p1_board=[_rank_card(0, 7)],
to_move=0,
)
action = int(heuristic_expert_action(state, jnp.int32(0), jax.random.PRNGKey(1)))
hand_slot, place_type, _ = [int(x) for x in decode_action(jnp.asarray(action))]
hand = [int(x) for x in current_hand_sorted(state, 0)]
if place_type != PLAY:
assert hand[hand_slot] != dangerous
@pytest.mark.parametrize(
"policy_name", ["discard_only", "heuristic_balanced", "heuristic_cautious"]
)
@@ -128,6 +194,23 @@ def test_static_policy_duplicate_mirror_score_diff_is_zero(policy_name):
assert result["wins"] == result["losses"]
def test_expert_duplicate_mirror_regression_stats():
result = evaluate_static_mirror(
"heuristic_expert",
games=1000,
duplicate=True,
shuffle_bank_seed=20260704,
batch_games=1000,
)
assert result["games"] == 2000
assert abs(result["mean_score_diff"]) <= 1.0e-6
assert result["wins"] == result["losses"]
assert result["max_steps_rate"] < 0.05
assert 2.0 <= result["opened_colors_per_game"] <= 3.5
assert 0.30 <= result["play_action_rate"] <= 0.60
assert result["mean_game_length"] <= 70.0
def test_gate3_checkpoint_duplicate_self_mirror_score_diff_is_zero():
checkpoint = Path(GATE3_CHECKPOINT)
if not checkpoint.exists():
@@ -189,6 +272,10 @@ def _rank_card(color: int, rank: int) -> int:
return color * CARDS_PER_COLOR + rank + 1
def _hs_card(color: int, slot: int = 0) -> int:
return color * CARDS_PER_COLOR + slot
def _manual_state(
*,
p0_hand: list[int] | None = None,