Add final cycle report and human play CLI

This commit is contained in:
2026-07-06 00:14:24 +09:00
parent 16cc31676c
commit 19560029b5
6 changed files with 1398 additions and 1 deletions
@@ -0,0 +1,8 @@
{"cycle": 1, "event": "final_cycle_league", "league_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/league_config.yaml", "league_run_dir": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01", "target_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "target_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json", "worst_protocol_used": "warmstart_gate3"}
{"checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_205759_final-c01-long-random-exploiter/latest", "cycle": 1, "event": "final_cycle_exploiter", "exploiter": "long_random", "games": 4000, "json": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_205759_final-c01-long-random-exploiter/eval_vs_final_cycle_01_duplicate.json", "losses": 919.0, "max_steps_rate": 0.0005, "mean_game_length": 65.69525, "mean_score_diff": 55.799, "notes": "random init + shaping anneal, 1200 updates", "opened_colors_per_game": 4.92425, "play_action_rate": 0.5457588946038896, "positive_expeditions_per_game": 2.7735, "resume": null, "score_diff_ci95_high": 57.963152016622004, "score_diff_ci95_low": 53.634847983377995, "target": "final_cycle_01", "target_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "target_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json", "ties": 22.0, "train_dir": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_205759_final-c01-long-random-exploiter", "wilson_high": 0.7776365926294233, "wilson_low": 0.7513553821486798, "win_rate": 0.76475, "wins": 3059.0}
{"checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_221402_final-c01-warmstart-gate3-exploiter/latest", "cycle": 1, "event": "final_cycle_exploiter", "exploiter": "warmstart_gate3", "games": 4000, "json": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_221402_final-c01-warmstart-gate3-exploiter/eval_vs_final_cycle_01_duplicate.json", "losses": 961.0, "max_steps_rate": 0.00075, "mean_game_length": 65.13875, "mean_score_diff": 53.30925, "notes": "ladder v2 gate-3 warm start, shaping disabled, 900 updates", "opened_colors_per_game": 4.90075, "play_action_rate": 0.5480372138722999, "positive_expeditions_per_game": 2.71225, "resume": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/2026-07-05_013223_jax-ppo-ladder-v2-expert/latest", "score_diff_ci95_high": 55.52225515395847, "score_diff_ci95_low": 51.09624484604153, "target": "final_cycle_01", "target_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "target_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json", "ties": 18.0, "train_dir": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_221402_final-c01-warmstart-gate3-exploiter", "wilson_high": 0.7683246594694194, "wilson_low": 0.741685544731188, "win_rate": 0.75525, "wins": 3021.0}
{"checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_231108_final-c01-replay-exploiter-exploiter/latest", "cycle": 1, "event": "final_cycle_exploiter", "exploiter": "replay_exploiter", "games": 4000, "json": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_231108_final-c01-replay-exploiter-exploiter/eval_vs_final_cycle_01_duplicate.json", "losses": 866.0, "max_steps_rate": 0.00025, "mean_game_length": 65.467, "mean_score_diff": 59.351, "notes": "league v1 cycle-1 exploiter warm start, shaping disabled, 900 updates", "opened_colors_per_game": 4.989, "play_action_rate": 0.5882132820136853, "positive_expeditions_per_game": 2.7905, "resume": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/league/2026-07-05_052325_jax-ppo-league-v1/exploiters/2026-07-05_060727_jax-ppo-league-v1-cycle-1-exploiter/latest", "score_diff_ci95_high": 61.580508518304526, "score_diff_ci95_low": 57.12149148169547, "target": "final_cycle_01", "target_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "target_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json", "ties": 17.0, "train_dir": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/battery/2026-07-05_231108_final-c01-replay-exploiter-exploiter", "wilson_high": 0.7918317711217939, "wilson_low": 0.7661323798009079, "win_rate": 0.77925, "wins": 3117.0}
{"cycle": 1, "event": "final_cycle_battery_judgment", "passed": false, "threshold": 0.52, "worst_exploiter": "replay_exploiter", "worst_win_rate": 0.77925}
{"cycle": 1, "event": "final_cycle_guard_expert", "games": 4000, "json": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/guard_vs_heuristic_expert.json", "losses": 852.0, "max_steps_rate": 0.0, "mean_game_length": 60.38125, "mean_score_diff": 36.08475, "opened_colors_per_game": 4.15225, "passed": true, "play_action_rate": 0.5845773004045489, "positive_expeditions_per_game": 2.05975, "score_diff_ci95_high": 37.5625235448534, "score_diff_ci95_low": 34.6069764551466, "target": "final_cycle_01", "target_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "target_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json", "ties": 33.0, "wilson_high": 0.7913426526326875, "wilson_low": 0.7656224577333041, "win_rate": 0.77875, "wins": 3115.0}
{"current": "final_cycle_01", "current_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "cycle": 1, "event": "final_cycle_h2h", "games": 4000, "json": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/cycle_01/h2h_final_cycle_01_vs_repair_c01_update_500.json", "losses": 1987.0, "max_steps_rate": 0.0, "mean_game_length": 49.54275, "mean_score_diff": 0.35425, "opened_colors_per_game": 4.91275, "play_action_rate": 0.7646239835347768, "positive_expeditions_per_game": 2.422, "previous": "repair_c01_update_500", "previous_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/gates-1-2/2026-07-05_094238_gates-1-2/gate2c/league/2026-07-05_125709_jax-ppo-gates-1-2-repair-c01/snapshots/cycle_01_update_000500", "score_diff_ci95_high": 1.7891388254913125, "score_diff_ci95_low": -1.0806388254913126, "ties": 51.0, "wilson_high": 0.505993762454951, "wilson_low": 0.47502446696755335, "win_rate": 0.4905, "wins": 1962.0}
{"elapsed_seconds": 17320.886204754002, "event": "final_cycles_complete", "final_candidate": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate", "final_config": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json", "source_checkpoint": "/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/snapshots/cycle_01_update_000500", "stop_reason": "h2h_stagnation"}
@@ -0,0 +1,83 @@
# Final Cycles and Human Play - 2026-07-05
## Part A - Closing Reinforcement Cycles
Raw artifacts: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05`
Stop reason: `h2h_stagnation`
Final candidate: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate`
Final candidate config: `/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/main_ppo_config.json`
The initial 4 hour GPU budget was treated as an estimate after user confirmation,
not a hard cutoff. The run completed cycle 1 fully, then stopped by the written
H2H stagnation rule. Cycle 2 was not started.
Bottom line: the closing reinforcement cycle did not produce a materially better
policy. The cycle-1 snapshot passed the expert guard, but the strengthened
exploiter battery got substantially worse than the incoming baseline target:
worst exploiter win rate was 0.7792, far above the 0.5200 success threshold. The
adjacent H2H against `repair_c01_update_500` was statistically indistinguishable
from zero, so this recipe is stopped here. Per instruction, this report does not
recommend more training with the same observation/network/PPO recipe.
### Expert Guard
| Cycle | Passed | Win rate | Mean diff | CI low | Opened colors | Max-step |
| ---: | ---: | ---: | ---: | ---: | ---: | ---: |
| 1 | True | 0.7788 | +36.0847 | +34.6070 | 4.1523 | 0.0000 |
### Adjacent H2H
| Cycle | Current | Previous | Win rate | Mean diff | Score CI |
| ---: | --- | --- | ---: | ---: | ---: |
| 1 | `final_cycle_01` | `repair_c01_update_500` | 0.4905 | +0.3543 | [-1.0806, +1.7891] |
### Strengthened Exploiter Battery
| Cycle | Exploiter | Win rate | Mean diff | CI low | Opened colors | Max-step |
| ---: | --- | ---: | ---: | ---: | ---: | ---: |
| 1 | `long_random` | 0.7648 | +55.7990 | +53.6348 | 4.9242 | 0.0005 |
| 1 | `warmstart_gate3` | 0.7552 | +53.3092 | +51.0962 | 4.9008 | 0.0008 |
| 1 | `replay_exploiter` | 0.7792 | +59.3510 | +57.1215 | 4.9890 | 0.0003 |
### Battery Judgment
| Cycle | Worst exploiter | Worst win rate | Threshold | Passed |
| ---: | --- | ---: | ---: | ---: |
| 1 | `replay_exploiter` | 0.7792 | 0.5200 | False |
## Part B - Human Play Interface
Start a single game:
```bash
uv run --with 'jax[cuda12]' lost-cities-jax-ppo play \
--checkpoint /mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate \
--seat 0
```
Start a duplicate set with one shared shuffle and swapped seats:
```bash
uv run --with 'jax[cuda12]' lost-cities-jax-ppo play \
--checkpoint /mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate \
--seat 0 --duplicate
```
Summarize logged human games:
```bash
uv run lost-cities-jax-ppo human-play summarize \
--log-dir /mnt/2tbhdd/coolrl-lost-cities-artifacts/human-play/
```
Move syntax: `play R7 draw deck`, `discard G3 draw Y`, or `play RHS draw deck`.
The renderer shows only the human hand, both boards, all public discard piles, deck count, and current board score differential. Opponent hand and deck order are not rendered.
Every game is appended to `/mnt/2tbhdd/coolrl-lost-cities-artifacts/human-play/games.jsonl` with deck seed/index, full action list, AI top-3 policy actions/probabilities, value outputs, scoring breakdown, and optional human comment.
Validation performed before final report:
- `uv run ruff check .` passed.
- `uv run pytest -q` passed: 257 passed, 1 skipped.
- `uv run lost-cities-jax-ppo play --help` passed.
- `uv run lost-cities-jax-ppo human-play summarize --log-dir /tmp/nonexistent-human-play-log` passed.
- `scripts/librarian.sh` found no link or code-citation errors; it still exits non-zero on the known pre-existing `docs/plans/deep-cfr-selectivity.md` 500-line soft cap.
+503
View File
@@ -0,0 +1,503 @@
from __future__ import annotations
import json
import shutil
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import yaml
from lost_cities_jax.gates import ExploiterSpec
from lost_cities_jax.league import run_league
from lost_cities_jax.ppo import (
evaluate_checkpoint_match,
evaluate_checkpoint_vs_static,
load_config,
train_against_checkpoint,
)
DATE = "2026-07-05"
ROOT = Path("/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05")
REPORT_PATH = Path(f"docs/reports/final-cycles-and-human-play-{DATE}.md")
SUMMARY_PATH = Path(f"docs/reports/final-cycles-and-human-play-{DATE}-summary.jsonl")
START_CONFIG = (
"/mnt/2tbhdd/coolrl-lost-cities-artifacts/gates-1-2/2026-07-05_094238_gates-1-2/"
"gate2c/league/2026-07-05_125709_jax-ppo-gates-1-2-repair-c01/main_ppo_config.json"
)
START_CHECKPOINT = (
"/mnt/2tbhdd/coolrl-lost-cities-artifacts/gates-1-2/2026-07-05_094238_gates-1-2/"
"gate2c/league/2026-07-05_125709_jax-ppo-gates-1-2-repair-c01/snapshots/"
"cycle_01_update_000500"
)
LEAGUE_TEMPLATE = "configs/jax_ppo/league-v1.yaml"
GAMES = 2000
PASS_THRESHOLD = 0.52
GUARD_CI_LOW = 0.0
GUARD_MAX_STEPS = 0.02
@dataclass(frozen=True)
class Target:
name: str
config: str
checkpoint: str
EXPLOITERS = [
ExploiterSpec(
name="long_random",
config="configs/jax_ppo/gates-1-2-exploiter-long-random.yaml",
notes="random init + shaping anneal, 1200 updates",
),
ExploiterSpec(
name="warmstart_gate3",
config="configs/jax_ppo/gates-1-2-exploiter-warmstart.yaml",
resume=(
"/mnt/2tbhdd/coolrl-lost-cities-artifacts/ladder-v2/"
"2026-07-05_013223_jax-ppo-ladder-v2-expert/latest"
),
notes="ladder v2 gate-3 warm start, shaping disabled, 900 updates",
),
ExploiterSpec(
name="replay_exploiter",
config="configs/jax_ppo/gates-1-2-exploiter-replay.yaml",
resume=(
"/mnt/2tbhdd/coolrl-lost-cities-artifacts/league/"
"2026-07-05_052325_jax-ppo-league-v1/exploiters/"
"2026-07-05_060727_jax-ppo-league-v1-cycle-1-exploiter/latest"
),
notes="league v1 cycle-1 exploiter warm start, shaping disabled, 900 updates",
),
]
def main() -> None:
started = time.perf_counter()
ROOT.mkdir(parents=True, exist_ok=True)
REPORT_PATH.parent.mkdir(parents=True, exist_ok=True)
SUMMARY_PATH.write_text("", encoding="utf-8")
rows: list[dict[str, Any]] = []
current = Target("repair_c01_update_500", START_CONFIG, START_CHECKPOINT)
best = current
worst_protocol = "warmstart_gate3"
extra_pool: list[dict[str, Any]] = []
stop_reason = "max_cycles_exhausted"
for cycle in range(1, 3):
previous = current
league_config = write_league_config(cycle, previous, worst_protocol, extra_pool)
league_dir = run_league(league_config)
league_rows = read_jsonl(league_dir / "league_summary.jsonl")
completion = latest_row(league_rows, "league_complete")
if completion is None:
stop_reason = "league_missing_completion"
rows.append({"event": "final_cycle_error", "cycle": cycle, "reason": stop_reason})
break
current = Target(
f"final_cycle_{cycle:02d}",
str(league_dir / "main_ppo_config.json"),
completion["final_checkpoint"],
)
cycle_row = {
"event": "final_cycle_league",
"cycle": cycle,
"worst_protocol_used": worst_protocol,
"league_config": str(league_config),
"league_run_dir": str(league_dir),
"target_config": current.config,
"target_checkpoint": current.checkpoint,
}
rows.append(cycle_row)
append_jsonl(SUMMARY_PATH, cycle_row)
exploiter_member = exploiter_member_from_league(league_rows, cycle)
if exploiter_member is not None:
extra_pool.append(exploiter_member)
battery = run_battery(cycle, current)
for row in battery:
rows.append(row)
append_jsonl(SUMMARY_PATH, row)
judgment = battery_judgment(cycle, battery)
rows.append(judgment)
append_jsonl(SUMMARY_PATH, judgment)
guard = evaluate_guard(cycle, current)
rows.append(guard)
append_jsonl(SUMMARY_PATH, guard)
h2h = evaluate_h2h(cycle, current, previous)
rows.append(h2h)
append_jsonl(SUMMARY_PATH, h2h)
best = current
if not guard["passed"]:
stop_reason = "expert_guard_failed"
current = previous
best = previous
rows.append(
{
"event": "final_cycle_rollback",
"cycle": cycle,
"rolled_back_to": previous.name,
"rolled_back_checkpoint": previous.checkpoint,
"failed_checkpoint": guard["target_checkpoint"],
}
)
append_jsonl(SUMMARY_PATH, rows[-1])
break
if judgment["worst_win_rate"] <= PASS_THRESHOLD:
stop_reason = "success_exploiter_threshold"
break
if h2h["score_diff_ci95_low"] <= 0.0 <= h2h["score_diff_ci95_high"]:
stop_reason = "h2h_stagnation"
break
worst_protocol = judgment["worst_exploiter"]
final_candidate = fix_final_candidate(best)
final_row = {
"event": "final_cycles_complete",
"stop_reason": stop_reason,
"final_candidate": str(final_candidate),
"final_config": best.config,
"source_checkpoint": best.checkpoint,
"elapsed_seconds": time.perf_counter() - started,
}
rows.append(final_row)
append_jsonl(SUMMARY_PATH, final_row)
write_report(rows, final_row)
def write_league_config(
cycle: int, target: Target, worst_protocol: str, extra_pool: list[dict[str, Any]]
) -> Path:
data = yaml.safe_load(Path(LEAGUE_TEMPLATE).read_text(encoding="utf-8")) or {}
data["base_config"] = target.config
data["warm_start_checkpoint"] = target.checkpoint
run = data.setdefault("run", {})
run["experiment_name"] = f"jax-ppo-final-cycle-c{cycle:02d}"
run["artifact_root"] = str(ROOT / "league")
run["seed"] = int(run.get("seed", 20260705)) + 100 + cycle
league = data.setdefault("league", {})
league["cycles"] = 1
league["league_updates_per_cycle"] = 500
league["snapshot_interval_updates"] = 500
league["success_exploiter_win_rate"] = PASS_THRESHOLD
evaluation = data.setdefault("evaluation", {})
evaluation["games"] = GAMES
evaluation["batch_games"] = 8192
exploiter = data.setdefault("exploiter", {})
spec = exploiter_spec(worst_protocol)
exploiter["config"] = spec.config
exploiter["updates"] = load_config(spec.config).run.total_updates
if spec.resume:
exploiter["resume"] = spec.resume
else:
exploiter.pop("resume", None)
guards = data.setdefault("guards", {})
guards["expert_ci_low"] = GUARD_CI_LOW
guards["max_steps_rate"] = GUARD_MAX_STEPS
cycle_dir = ROOT / f"cycle_{cycle:02d}"
tracking = data.setdefault("tracking", {})
tracking["tracked_summary_path"] = str(cycle_dir / "league_summary.jsonl")
tracking["report_path"] = str(cycle_dir / "league_report.md")
anchors = list(data.get("anchors", []))
names = {item.get("name") for item in anchors}
for member in extra_pool:
if member["name"] not in names:
anchors.append(member)
names.add(member["name"])
data["anchors"] = anchors
path = cycle_dir / "league_config.yaml"
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(yaml.safe_dump(data, sort_keys=False), encoding="utf-8")
return path
def evaluate_guard(cycle: int, target: Target) -> dict[str, Any]:
output = ROOT / f"cycle_{cycle:02d}" / "guard_vs_heuristic_expert.json"
result = evaluate_checkpoint_vs_static(
load_config(target.config),
target.checkpoint,
"heuristic_expert",
games=GAMES,
duplicate=True,
output=output,
)
return {
"event": "final_cycle_guard_expert",
"cycle": cycle,
"target": target.name,
"target_config": target.config,
"target_checkpoint": target.checkpoint,
"json": str(output),
"passed": result["score_diff_ci95_low"] > GUARD_CI_LOW
and result["max_steps_rate"] <= GUARD_MAX_STEPS,
**summary(result),
}
def evaluate_h2h(cycle: int, target: Target, previous: Target) -> dict[str, Any]:
output = ROOT / f"cycle_{cycle:02d}" / f"h2h_{target.name}_vs_{previous.name}.json"
result = evaluate_checkpoint_match(
load_config(target.config),
target.checkpoint,
load_config(previous.config),
previous.checkpoint,
games=GAMES,
duplicate=True,
output=output,
)
return {
"event": "final_cycle_h2h",
"cycle": cycle,
"current": target.name,
"previous": previous.name,
"current_checkpoint": target.checkpoint,
"previous_checkpoint": previous.checkpoint,
"json": str(output),
**summary(result),
}
def run_battery(cycle: int, target: Target) -> list[dict[str, Any]]:
rows = []
for spec in EXPLOITERS:
target_cfg = load_config(target.config)
exploiter_cfg = load_config(spec.config)
exploiter_cfg.run.artifact_root = str(ROOT / f"cycle_{cycle:02d}" / "battery")
exploiter_cfg.run.experiment_name = f"final-c{cycle:02d}-{spec.name}-exploiter"
train_dir = train_against_checkpoint(
exploiter_cfg,
target_cfg,
target.checkpoint,
resume=spec.resume,
)
output = train_dir / f"eval_vs_{target.name}_duplicate.json"
result = evaluate_checkpoint_match(
exploiter_cfg,
train_dir / "latest",
target_cfg,
target.checkpoint,
games=GAMES,
duplicate=True,
output=output,
)
rows.append(
{
"event": "final_cycle_exploiter",
"cycle": cycle,
"exploiter": spec.name,
"target": target.name,
"target_config": target.config,
"target_checkpoint": target.checkpoint,
"train_dir": str(train_dir),
"checkpoint": str(train_dir / "latest"),
"resume": spec.resume,
"notes": spec.notes,
"json": str(output),
**summary(result),
}
)
return rows
def battery_judgment(cycle: int, rows: list[dict[str, Any]]) -> dict[str, Any]:
worst = max(rows, key=lambda row: row["win_rate"])
return {
"event": "final_cycle_battery_judgment",
"cycle": cycle,
"worst_exploiter": worst["exploiter"],
"worst_win_rate": worst["win_rate"],
"threshold": PASS_THRESHOLD,
"passed": worst["win_rate"] <= PASS_THRESHOLD,
}
def exploiter_member_from_league(rows: list[dict[str, Any]], cycle: int) -> dict[str, Any] | None:
row = latest_row(rows, "exploiter_eval")
if row is None:
return None
checkpoint = Path(row["exploiter_checkpoint"])
return {
"name": f"final_cycle_exploiter_c{cycle:02d}",
"kind": "checkpoint",
"anchor": False,
"stalling": False,
"exploiter": True,
"config": str(checkpoint.parent / "config.json"),
"checkpoint": str(checkpoint),
"recent_win_rate": 1.0 - float(row["win_rate"]),
"created_cycle": cycle,
"created_update": 0,
}
def fix_final_candidate(target: Target) -> Path:
destination = ROOT / "final_candidate"
if destination.exists():
shutil.rmtree(destination)
shutil.copytree(target.checkpoint, destination)
(ROOT / "final_candidate_config.txt").write_text(target.config + "\n", encoding="utf-8")
shutil.copy2(target.config, ROOT / "main_ppo_config.json")
return destination
def write_report(rows: list[dict[str, Any]], final: dict[str, Any]) -> None:
guards = [row for row in rows if row.get("event") == "final_cycle_guard_expert"]
h2h = [row for row in rows if row.get("event") == "final_cycle_h2h"]
exploiters = [row for row in rows if row.get("event") == "final_cycle_exploiter"]
judgments = [row for row in rows if row.get("event") == "final_cycle_battery_judgment"]
lines = [
f"# Final Cycles and Human Play - {DATE}",
"",
"## Part A - Closing Reinforcement Cycles",
"",
f"Raw artifacts: `{ROOT}`",
f"Stop reason: `{final['stop_reason']}`",
f"Final candidate: `{final['final_candidate']}`",
f"Final candidate config: `{final['final_config']}`",
"",
"### Expert Guard",
"",
"| Cycle | Passed | Win rate | Mean diff | CI low | Opened colors | Max-step |",
"| ---: | ---: | ---: | ---: | ---: | ---: | ---: |",
]
for row in guards:
lines.append(match_row(row, leading=[str(row["cycle"]), str(row["passed"])]))
lines.extend(
[
"",
"### Adjacent H2H",
"",
"| Cycle | Current | Previous | Win rate | Mean diff | Score CI |",
"| ---: | --- | --- | ---: | ---: | ---: |",
]
)
for row in h2h:
lines.append(
f"| {row['cycle']} | `{row['current']}` | `{row['previous']}` | "
f"{row['win_rate']:.4f} | {row['mean_score_diff']:+.4f} | "
f"[{row['score_diff_ci95_low']:+.4f}, {row['score_diff_ci95_high']:+.4f}] |"
)
lines.extend(
[
"",
"### Strengthened Exploiter Battery",
"",
"| Cycle | Exploiter | Win rate | Mean diff | CI low | Opened colors | Max-step |",
"| ---: | --- | ---: | ---: | ---: | ---: | ---: |",
]
)
for row in exploiters:
lines.append(match_row(row, leading=[str(row["cycle"]), f"`{row['exploiter']}`"]))
lines.extend(
[
"",
"### Battery Judgment",
"",
"| Cycle | Worst exploiter | Worst win rate | Threshold | Passed |",
"| ---: | --- | ---: | ---: | ---: |",
]
)
for row in judgments:
lines.append(
f"| {row['cycle']} | `{row['worst_exploiter']}` | {row['worst_win_rate']:.4f} | "
f"{row['threshold']:.4f} | {row['passed']} |"
)
lines.extend(human_play_usage())
REPORT_PATH.write_text("\n".join(lines) + "\n", encoding="utf-8")
def human_play_usage() -> list[str]:
return [
"",
"## Part B - Human Play Interface",
"",
"Start a single game:",
"",
"```bash",
"uv run --with 'jax[cuda12]' lost-cities-jax-ppo play \\",
" --checkpoint /mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate \\",
" --seat 0",
"```",
"",
"Start a duplicate set with one shared shuffle and swapped seats:",
"",
"```bash",
"uv run --with 'jax[cuda12]' lost-cities-jax-ppo play \\",
" --checkpoint /mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate \\",
" --seat 0 --duplicate",
"```",
"",
"Summarize logged human games:",
"",
"```bash",
"uv run lost-cities-jax-ppo human-play summarize \\",
" --log-dir /mnt/2tbhdd/coolrl-lost-cities-artifacts/human-play/",
"```",
"",
"Move syntax: `play R7 draw deck`, `discard G3 draw Y`, or `play RHS draw deck`.",
"The renderer shows only the human hand, both boards, all public discard piles, deck count, and current board score differential. Opponent hand and deck order are not rendered.",
"Every game is appended to `/mnt/2tbhdd/coolrl-lost-cities-artifacts/human-play/games.jsonl` with deck seed/index, full action list, AI top-3 policy actions/probabilities, value outputs, scoring breakdown, and optional human comment.",
]
def match_row(row: dict[str, Any], leading: list[str]) -> str:
return (
"| " + " | ".join(leading) + f" | {row['win_rate']:.4f} | {row['mean_score_diff']:+.4f} | "
f"{row['score_diff_ci95_low']:+.4f} | {row['opened_colors_per_game']:.4f} | "
f"{row['max_steps_rate']:.4f} |"
)
def summary(result: dict[str, Any]) -> dict[str, Any]:
keys = [
"games",
"wins",
"losses",
"ties",
"win_rate",
"wilson_low",
"wilson_high",
"mean_score_diff",
"score_diff_ci95_low",
"score_diff_ci95_high",
"mean_game_length",
"max_steps_rate",
"opened_colors_per_game",
"play_action_rate",
"positive_expeditions_per_game",
]
return {key: result[key] for key in keys if key in result}
def exploiter_spec(name: str) -> ExploiterSpec:
for spec in EXPLOITERS:
if spec.name == name:
return spec
raise ValueError(f"unknown exploiter protocol: {name}")
def latest_row(rows: list[dict[str, Any]], event: str) -> dict[str, Any] | None:
return next((row for row in reversed(rows) if row.get("event") == event), None)
def read_jsonl(path: Path) -> list[dict[str, Any]]:
if not path.exists():
return []
return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line]
def append_jsonl(path: Path, row: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as handle:
handle.write(json.dumps(row, sort_keys=True) + "\n")
if __name__ == "__main__":
main()
+617
View File
@@ -0,0 +1,617 @@
"""Terminal human-play client and log summarizer for JAX PPO checkpoints."""
from __future__ import annotations
import json
import math
import time
import uuid
from collections import defaultdict
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import jax
import jax.numpy as jnp
import numpy as np
from lost_cities_jax.engine import (
board_score,
current_hand_sorted,
decode_action,
legal_action_mask,
reset_from_order,
step,
)
from lost_cities_jax.obs import observation
from lost_cities_jax.ppo import (
ActorCritic,
JaxPPOConfig,
create_train_state,
load_config,
mask_logits,
restore_checkpoint,
)
from lost_cities_jax.types import (
CARDS_PER_COLOR,
DISCARD,
DRAW_DECK,
LOC_P0_BOARD,
MAX_STEPS,
N_ACTIONS,
N_CARDS,
N_COLORS,
PLAY,
State,
)
COLOR_NAMES = ("R", "G", "W", "B", "Y")
COLOR_WORDS = {
"R": 0,
"RED": 0,
"G": 1,
"GREEN": 1,
"W": 2,
"WHITE": 2,
"B": 3,
"BLUE": 3,
"Y": 4,
"YELLOW": 4,
}
DEFAULT_HUMAN_LOG_DIR = Path("/mnt/2tbhdd/coolrl-lost-cities-artifacts/human-play")
DEFAULT_HUMAN_BANK_STATE = DEFAULT_HUMAN_LOG_DIR / "shuffle_bank_state.json"
DEFAULT_HUMAN_BANK_SEED = 20260706
@dataclass(frozen=True)
class PolicyEval:
action: int
top3: list[dict[str, Any]]
value: float
def play_cli(
*,
checkpoint: str | Path,
config: str | Path | None = None,
seat: int = 0,
seed: int = DEFAULT_HUMAN_BANK_SEED,
duplicate: bool = False,
log_dir: str | Path = DEFAULT_HUMAN_LOG_DIR,
bank_state: str | Path = DEFAULT_HUMAN_BANK_STATE,
input_fn: Callable[[str], str] = input,
output_fn: Callable[[str], None] = print,
) -> list[dict[str, Any]]:
cfg = load_config(infer_config_path(checkpoint) if config is None else config)
agent = load_agent(cfg, checkpoint)
log_dir = Path(log_dir)
deck_index, deck_order = reserve_deck_order(seed, Path(bank_state))
session_id = time.strftime("%Y%m%d-%H%M%S") + "-" + uuid.uuid4().hex[:8]
seats = session_human_seats(seat, duplicate)
records = []
set_diff = 0.0
for game_index, human_seat in enumerate(seats, start=1):
output_fn("")
output_fn(f"Game {game_index}/{len(seats)} - human seat P{human_seat}")
record = play_one_game(
cfg=cfg,
agent=agent,
deck_order=deck_order,
human_seat=human_seat,
session_id=session_id,
duplicate_set_id=session_id if duplicate else None,
duplicate_game_index=game_index if duplicate else None,
deck_seed=seed,
deck_index=deck_index,
input_fn=input_fn,
output_fn=output_fn,
)
records.append(record)
set_diff += float(record["human_score_diff"])
append_human_log(log_dir, record)
if duplicate:
output_fn("")
output_fn(f"Duplicate set total human score diff: {set_diff:+.0f}")
return records
def summarize_logs(log_dir: str | Path = DEFAULT_HUMAN_LOG_DIR) -> dict[str, Any]:
records = read_human_logs(log_dir)
games = len(records)
duplicate_sets: dict[str, list[dict[str, Any]]] = defaultdict(list)
singletons = []
for record in records:
set_id = record.get("duplicate_set_id")
if set_id:
duplicate_sets[set_id].append(record)
else:
singletons.append(record)
set_diffs = [
sum(float(game["human_score_diff"]) for game in group)
for group in duplicate_sets.values()
if len(group) == 2
]
game_diffs = [float(record["human_score_diff"]) for record in records]
ai_opened = [float(record["ai_opened_colors"]) for record in records]
lengths = [float(record["game_length"]) for record in records]
swings = collect_value_swings(records)[:10]
return {
"games": games,
"duplicate_sets": len(set_diffs),
"single_games": len(singletons),
"human_game_score_diff_mean": mean(game_diffs),
"human_game_score_diff_ci95": ci95(game_diffs),
"human_duplicate_set_diff_mean": mean(set_diffs),
"human_duplicate_set_diff_ci95": ci95(set_diffs),
"ai_opened_colors_mean": mean(ai_opened),
"game_length_mean": mean(lengths),
"value_swings_top10": swings,
}
def session_human_seats(seat: int, duplicate: bool) -> list[int]:
if seat not in {0, 1}:
raise ValueError("seat must be 0 or 1")
return [seat, 1 - seat] if duplicate else [seat]
def print_summary(summary: dict[str, Any], output_fn: Callable[[str], None] = print) -> None:
output_fn(json.dumps(summary, indent=2, sort_keys=True, ensure_ascii=False))
def play_one_game(
*,
cfg: JaxPPOConfig,
agent: tuple[Any, Any],
deck_order: list[int],
human_seat: int,
session_id: str,
duplicate_set_id: str | None,
duplicate_game_index: int | None,
deck_seed: int,
deck_index: int,
input_fn: Callable[[str], str],
output_fn: Callable[[str], None],
human_action_fn: Callable[[State, int], int] | None = None,
) -> dict[str, Any]:
params, model = agent
state = reset_from_order(jnp.asarray(deck_order, dtype=jnp.int8))
actions = []
move_logs = []
while not bool(state.done) and int(state.step_count) < MAX_STEPS:
player = int(state.to_move)
policy_eval = evaluate_agent_policy(cfg, params, model, state, player)
output_fn(render_public_state(state, human_seat))
if player == human_seat:
if human_action_fn is None:
action = prompt_human_action(state, human_seat, input_fn, output_fn)
else:
action = int(human_action_fn(state, human_seat))
actor = "human"
else:
action = policy_eval.action
actor = "ai"
output_fn(f"AI: {describe_action(state, action)}")
before = state
state, _, _ = step(state, jnp.asarray(action, dtype=jnp.int32))
actions.append(int(action))
move_logs.append(
{
"ply": int(before.step_count),
"actor": actor,
"player": player,
"action": int(action),
"action_text": describe_action(before, action),
"ai_policy_player": player,
"ai_top3": policy_eval.top3,
"ai_value": policy_eval.value,
}
)
output_fn(render_public_state(state, human_seat))
output_fn(render_score_breakdown(state))
scores = np.asarray(board_score(state), dtype=np.float32)
ai_seat = 1 - human_seat
comment = input_fn("Post-game comment (optional, Enter to skip): ").strip()
return {
"schema": "lost-cities-jax-human-play-v1",
"session_id": session_id,
"duplicate_set_id": duplicate_set_id,
"duplicate_game_index": duplicate_game_index,
"deck_seed": deck_seed,
"deck_index": deck_index,
"human_seat": human_seat,
"ai_seat": ai_seat,
"deck_order": [int(card) for card in deck_order],
"actions": actions,
"moves": move_logs,
"final_scores": {"p0": float(scores[0]), "p1": float(scores[1])},
"human_score_diff": float(scores[human_seat] - scores[ai_seat]),
"ai_score_diff": float(scores[ai_seat] - scores[human_seat]),
"score_breakdown": score_breakdown(state),
"game_length": int(state.step_count),
"ai_opened_colors": int(np.asarray(state.col_len)[ai_seat].astype(bool).sum()),
"human_opened_colors": int(np.asarray(state.col_len)[human_seat].astype(bool).sum()),
"max_steps": bool(int(state.step_count) >= MAX_STEPS),
"human_comment": comment,
"created_at": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
}
def load_agent(cfg: JaxPPOConfig, checkpoint: str | Path) -> tuple[Any, ActorCritic]:
state = create_train_state(cfg, jax.random.PRNGKey(0))
state = restore_checkpoint(Path(checkpoint), state)
model = ActorCritic(cfg.network.hidden_size, cfg.network.num_layers)
return state.params, model
def infer_config_path(checkpoint: str | Path) -> Path:
path = Path(checkpoint).resolve()
candidates = []
for parent in [path, *path.parents]:
candidates.extend([parent / "main_ppo_config.json", parent / "config.json"])
for candidate in candidates:
if candidate.exists():
return candidate
msg = f"could not infer config for checkpoint; pass --config explicitly: {checkpoint}"
raise FileNotFoundError(msg)
def evaluate_agent_policy(
cfg: JaxPPOConfig,
params: Any,
model: ActorCritic,
state: State,
player: int,
) -> PolicyEval:
obs = observation(state, jnp.asarray(player, dtype=jnp.int32))[None, :]
legal = legal_action_mask(state)[None, :]
logits, value = model.apply(params, obs)
masked = mask_logits(logits, legal)
probs = np.asarray(jax.nn.softmax(masked, axis=-1)[0], dtype=np.float64)
legal_np = np.asarray(legal[0], dtype=bool)
order = np.argsort(-probs)
top = []
for action in order:
if not legal_np[action]:
continue
top.append(
{
"action": int(action),
"prob": float(probs[action]),
"text": describe_action(state, int(action)),
}
)
if len(top) == 3:
break
return PolicyEval(action=int(order[0]), top3=top, value=float(np.asarray(value)[0]))
def prompt_human_action(
state: State,
player: int,
input_fn: Callable[[str], str],
output_fn: Callable[[str], None],
) -> int:
while True:
raw = input_fn("Your move (e.g. play R7 draw deck): ").strip()
try:
action = parse_human_action(raw, state, player)
except ValueError as exc:
output_fn(f"Invalid input: {exc}")
continue
mask = np.asarray(legal_action_mask(state), dtype=bool)
if action < 0 or action >= N_ACTIONS or not mask[action]:
output_fn(f"Illegal move: {illegal_reason(state, action)}")
continue
return action
def parse_human_action(raw: str, state: State, player: int) -> int:
parts = raw.strip().split()
if len(parts) != 4 or parts[2].lower() != "draw":
msg = "expected '<play|discard> <card> draw <deck|R|G|W|B|Y>'"
raise ValueError(msg)
place_word, card_word, _, draw_word = parts
place = parse_place(place_word)
card_id = parse_card_token(card_word)
hand = [
int(card) for card in np.asarray(current_hand_sorted(state, player)) if int(card) < N_CARDS
]
matching = [idx for idx, card in enumerate(hand) if same_card_face(card, card_id)]
if not matching:
msg = f"card is not in your hand: {card_word}"
raise ValueError(msg)
hand_slot = matching[0]
draw_source = parse_draw_source(draw_word)
return encode_action(hand_slot, place, draw_source)
def parse_place(value: str) -> int:
lowered = value.lower()
if lowered == "play":
return PLAY
if lowered == "discard":
return DISCARD
msg = "placement must be 'play' or 'discard'"
raise ValueError(msg)
def parse_draw_source(value: str) -> int:
upper = value.upper()
if upper == "DECK":
return DRAW_DECK
if upper in COLOR_WORDS:
return COLOR_WORDS[upper] + 1
msg = "draw source must be deck or one of R/G/W/B/Y"
raise ValueError(msg)
def parse_card_token(value: str) -> int:
token = value.strip().upper()
if len(token) < 2:
raise ValueError("card token is too short")
color = COLOR_WORDS.get(token[0])
if color is None:
raise ValueError("card color must be R/G/W/B/Y")
suffix = token[1:]
if suffix in {"H", "HS", "W", "WAGER"}:
slot = 0
else:
try:
rank = int(suffix)
except ValueError as exc:
raise ValueError("rank must be 2..10 or HS") from exc
if rank < 2 or rank > 10:
raise ValueError("rank must be 2..10")
slot = rank + 1
return color * CARDS_PER_COLOR + slot
def same_card_face(left: int, right: int) -> bool:
return left // CARDS_PER_COLOR == right // CARDS_PER_COLOR and card_label(left) == card_label(
right
)
def encode_action(hand_slot: int, place_type: int, draw_source: int) -> int:
return hand_slot * 12 + place_type * 6 + draw_source
def illegal_reason(state: State, action: int) -> str:
if action < 0 or action >= N_ACTIONS:
return "action id outside action space"
hand_slot, place_type, draw_source = [int(x) for x in decode_action(jnp.asarray(action))]
hand = [int(card) for card in np.asarray(current_hand_sorted(state)) if int(card) < N_CARDS]
if hand_slot >= len(hand):
return "selected hand slot is empty"
card = hand[hand_slot]
color = card // CARDS_PER_COLOR
rank = card_rank(card)
player = int(state.to_move)
if place_type == PLAY:
top = int(np.asarray(state.col_top)[player, color])
if rank == 0 and top > 0:
return "handshake cannot be played after a number card"
if rank > 0 and rank <= top:
return f"number card must be above current top rank {top}"
if draw_source != DRAW_DECK:
pile_color = draw_source - 1
pile_len = int(np.asarray(state.pile_len)[pile_color])
if pile_len <= 0:
return "discard pile is empty"
if place_type == DISCARD and pile_color == color:
return "cannot draw the card you just discarded"
return "move is not legal under the current mask"
def render_public_state(state: State, player: int) -> str:
scores = np.asarray(board_score(state), dtype=np.float32)
opponent = 1 - player
lines = [
"",
f"Ply {int(state.step_count)} | to move P{int(state.to_move)} | deck left {N_CARDS - int(state.draw_ptr)}",
f"Board score: P{player} {scores[player]:+.0f} / P{opponent} {scores[opponent]:+.0f} "
f"(you {scores[player] - scores[opponent]:+.0f})",
f"Your hand: {' '.join(card_name(card) for card in np.asarray(current_hand_sorted(state, player)) if int(card) < N_CARDS)}",
"Boards:",
*render_boards(state),
"Discards:",
*render_piles(state),
]
return "\n".join(lines)
def render_boards(state: State) -> list[str]:
loc = np.asarray(state.card_loc)
lines = []
for player in range(2):
color_parts = []
for color in range(N_COLORS):
cards = [
card
for card in range(N_CARDS)
if loc[card] == LOC_P0_BOARD + player and card // CARDS_PER_COLOR == color
]
cards.sort(key=lambda card: (card_rank(card) > 0, card_rank(card), card))
color_parts.append(f"{COLOR_NAMES[color]}:[{' '.join(card_name(c) for c in cards)}]")
lines.append(f" P{player} " + " ".join(color_parts))
return lines
def render_piles(state: State) -> list[str]:
pile = np.asarray(state.pile)
pile_len = np.asarray(state.pile_len)
lines = []
for color in range(N_COLORS):
cards = [int(card) for card in pile[color, : int(pile_len[color])]]
lines.append(f" {COLOR_NAMES[color]}: {' '.join(card_name(card) for card in cards)}")
return lines
def render_score_breakdown(state: State) -> str:
breakdown = score_breakdown(state)
lines = ["Final scoring:"]
for player in range(2):
parts = []
for color in COLOR_NAMES:
item = breakdown[f"p{player}"][color]
parts.append(f"{color}:{item['score']:+.0f}")
lines.append(f" P{player} {' '.join(parts)} total {breakdown[f'p{player}']['total']:+.0f}")
return "\n".join(lines)
def score_breakdown(state: State) -> dict[str, Any]:
loc = np.asarray(state.card_loc)
col_len = np.asarray(state.col_len)
col_hs = np.asarray(state.col_hs)
data: dict[str, Any] = {}
for player in range(2):
total = 0
player_data: dict[str, Any] = {}
for color in range(N_COLORS):
ranks = [
card_rank(card)
for card in range(N_CARDS)
if loc[card] == LOC_P0_BOARD + player
and card // CARDS_PER_COLOR == color
and card_rank(card) > 0
]
length = int(col_len[player, color])
hs = int(col_hs[player, color])
rank_sum = sum(ranks)
score = 0 if length == 0 else (rank_sum - 20) * (1 + hs) + (20 if length >= 8 else 0)
total += score
player_data[COLOR_NAMES[color]] = {
"rank_sum": rank_sum,
"handshakes": hs,
"length": length,
"score": float(score),
}
player_data["total"] = float(total)
data[f"p{player}"] = player_data
return data
def describe_action(state: State, action: int) -> str:
hand_slot, place_type, draw_source = [int(x) for x in decode_action(jnp.asarray(action))]
player = int(state.to_move)
hand = [
int(card) for card in np.asarray(current_hand_sorted(state, player)) if int(card) < N_CARDS
]
card = hand[hand_slot] if hand_slot < len(hand) else -1
place = "play" if place_type == PLAY else "discard"
draw = "deck" if draw_source == DRAW_DECK else COLOR_NAMES[draw_source - 1]
return f"{place} {card_name(card)} draw {draw}"
def card_name(card: int) -> str:
if card < 0 or card >= N_CARDS:
return "?"
return COLOR_NAMES[card // CARDS_PER_COLOR] + card_label(card)
def card_label(card: int) -> str:
rank = card_rank(card)
return "HS" if rank == 0 else str(rank)
def card_rank(card: int) -> int:
slot = card % CARDS_PER_COLOR
return 0 if slot < 3 else slot - 1
def reserve_deck_order(seed: int, bank_state: Path) -> tuple[int, list[int]]:
bank_state.parent.mkdir(parents=True, exist_ok=True)
if bank_state.exists():
state = json.loads(bank_state.read_text(encoding="utf-8"))
next_index = int(state.get(str(seed), 0))
else:
state = {}
next_index = 0
state[str(seed)] = next_index + 1
bank_state.write_text(json.dumps(state, indent=2, sort_keys=True) + "\n", encoding="utf-8")
return next_index, human_deck_order(seed, next_index)
def human_deck_order(seed: int, index: int) -> list[int]:
rng = np.random.default_rng(seed + index)
return rng.permutation(N_CARDS).astype("int8").astype(int).tolist()
def append_human_log(log_dir: Path, record: dict[str, Any]) -> Path:
log_dir.mkdir(parents=True, exist_ok=True)
path = log_dir / "games.jsonl"
with path.open("a", encoding="utf-8") as handle:
handle.write(json.dumps(record, ensure_ascii=False, sort_keys=True) + "\n")
return path
def read_human_logs(log_dir: str | Path) -> list[dict[str, Any]]:
path = Path(log_dir) / "games.jsonl"
if not path.exists():
return []
return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line]
def collect_value_swings(records: list[dict[str, Any]]) -> list[dict[str, Any]]:
swings = []
for record in records:
previous = None
for move in record.get("moves", []):
value = float(move.get("ai_value", 0.0))
if previous is not None:
swings.append(
{
"session_id": record.get("session_id"),
"duplicate_game_index": record.get("duplicate_game_index"),
"ply": move.get("ply"),
"delta": value - previous,
"abs_delta": abs(value - previous),
"action_text": move.get("action_text"),
}
)
previous = value
swings.sort(key=lambda row: row["abs_delta"], reverse=True)
return swings
def mean(values: list[float]) -> float | None:
if not values:
return None
return float(np.mean(np.asarray(values, dtype=np.float64)))
def ci95(values: list[float]) -> list[float | None]:
if not values:
return [None, None]
arr = np.asarray(values, dtype=np.float64)
if arr.size <= 1:
value = float(arr[0])
return [value, value]
se = float(np.std(arr, ddof=1) / math.sqrt(arr.size))
avg = float(np.mean(arr))
return [avg - 1.959963984540054 * se, avg + 1.959963984540054 * se]
__all__ = [
"append_human_log",
"card_name",
"describe_action",
"human_deck_order",
"infer_config_path",
"parse_human_action",
"play_cli",
"play_one_game",
"read_human_logs",
"render_public_state",
"reserve_deck_order",
"session_human_seats",
"summarize_logs",
]
+40 -1
View File
@@ -1285,6 +1285,20 @@ def cli_main(argv: list[str] | None = None) -> None:
static_match_parser.add_argument("--duplicate", action="store_true")
static_match_parser.add_argument("--output")
play_parser = sub.add_parser("play")
play_parser.add_argument("--checkpoint", required=True)
play_parser.add_argument("--config")
play_parser.add_argument("--seat", type=int, choices=[0, 1], default=0)
play_parser.add_argument("--seed", type=int, default=20260706)
play_parser.add_argument("--duplicate", action="store_true")
play_parser.add_argument("--log-dir")
play_parser.add_argument("--bank-state")
human_play_parser = sub.add_parser("human-play")
human_play_sub = human_play_parser.add_subparsers(dest="human_play_command", required=True)
human_summary_parser = human_play_sub.add_parser("summarize")
human_summary_parser.add_argument("--log-dir")
league_parser = sub.add_parser("league")
league_sub = league_parser.add_subparsers(dest="league_command", required=True)
league_run_parser = league_sub.add_parser("run")
@@ -1301,7 +1315,7 @@ def cli_main(argv: list[str] | None = None) -> None:
args = parser.parse_args(argv)
cfg = (
load_config(args.config, overrides=parse_overrides(args.set))
if hasattr(args, "config") and args.command not in {"league", "gates"}
if hasattr(args, "config") and args.command not in {"league", "gates", "play", "human-play"}
else None
)
if args.command == "rollout-smoke":
@@ -1380,6 +1394,31 @@ def cli_main(argv: list[str] | None = None) -> None:
output=args.output,
)
print(json.dumps(result, indent=2, sort_keys=True))
elif args.command == "play":
from lost_cities_jax.human_play import (
DEFAULT_HUMAN_BANK_STATE,
DEFAULT_HUMAN_LOG_DIR,
play_cli,
)
play_cli(
checkpoint=args.checkpoint,
config=args.config,
seat=args.seat,
seed=args.seed,
duplicate=args.duplicate,
log_dir=args.log_dir or DEFAULT_HUMAN_LOG_DIR,
bank_state=args.bank_state or DEFAULT_HUMAN_BANK_STATE,
)
elif args.command == "human-play":
if args.human_play_command == "summarize":
from lost_cities_jax.human_play import (
DEFAULT_HUMAN_LOG_DIR,
print_summary,
summarize_logs,
)
print_summary(summarize_logs(args.log_dir or DEFAULT_HUMAN_LOG_DIR))
elif args.command == "league":
if args.league_command == "run":
from lost_cities_jax.league import run_league
+147
View File
@@ -0,0 +1,147 @@
from __future__ import annotations
import jax
import jax.numpy as jnp
import numpy as np
from lost_cities_jax.engine import board_score, reset_from_order, step
from lost_cities_jax.human_play import (
append_human_log,
evaluate_agent_policy,
human_deck_order,
parse_human_action,
play_one_game,
prompt_human_action,
render_public_state,
session_human_seats,
summarize_logs,
)
from lost_cities_jax.opponents import discard_only_action
from lost_cities_jax.ppo import (
ActorCritic,
JaxPPOConfig,
NetworkConfig,
OpponentConfig,
PPOHyperConfig,
RunConfig,
create_train_state,
)
def tiny_cfg(tmp_path) -> JaxPPOConfig:
return JaxPPOConfig(
run=RunConfig(experiment_name="human-play-test", artifact_root=str(tmp_path)),
opponent=OpponentConfig(name="discard_only"),
network=NetworkConfig(hidden_size=32, num_layers=1),
ppo=PPOHyperConfig(batch_games=4, rollout_steps=16, epochs=1, minibatches=2),
)
def test_illegal_input_reprompts_with_reason() -> None:
state = reset_from_order(jnp.asarray(list(range(60)), dtype=jnp.int8))
inputs = iter(["play R2 draw R", "play R2 draw deck"])
outputs: list[str] = []
action = prompt_human_action(
state,
0,
input_fn=lambda _prompt: next(inputs),
output_fn=outputs.append,
)
assert action == parse_human_action("play R2 draw deck", state, 0)
assert any("Illegal move" in line for line in outputs)
assert any("discard pile is empty" in line for line in outputs)
def test_public_renderer_does_not_show_opponent_hand_or_deck_order() -> None:
p0 = [0, 3, 4, 5, 6, 7, 8, 9]
p1 = [12, 15, 16, 17, 18, 19, 20, 21]
rest = [card for card in range(60) if card not in {*p0, *p1}]
state = reset_from_order(jnp.asarray(p0 + p1 + rest, dtype=jnp.int8))
rendered = render_public_state(state, 0)
assert "R2" in rendered
assert "G2" not in rendered
assert "G6" not in rendered
assert str(p0 + p1 + rest) not in rendered
def test_duplicate_session_swaps_seats() -> None:
assert session_human_seats(0, duplicate=True) == [0, 1]
assert session_human_seats(1, duplicate=True) == [1, 0]
assert session_human_seats(1, duplicate=False) == [1]
def test_human_play_summary_groups_duplicate_sets(tmp_path) -> None:
base = {
"schema": "lost-cities-jax-human-play-v1",
"session_id": "s",
"duplicate_set_id": "set-a",
"moves": [
{"ply": 0, "ai_value": 0.1, "action_text": "play R2 draw deck"},
{"ply": 1, "ai_value": 0.6, "action_text": "discard G2 draw deck"},
],
"ai_opened_colors": 3,
"game_length": 44,
}
append_human_log(tmp_path, {**base, "duplicate_game_index": 1, "human_score_diff": 10})
append_human_log(tmp_path, {**base, "duplicate_game_index": 2, "human_score_diff": -4})
summary = summarize_logs(tmp_path)
assert summary["games"] == 2
assert summary["duplicate_sets"] == 1
assert summary["human_duplicate_set_diff_mean"] == 6.0
assert summary["ai_opened_colors_mean"] == 3.0
assert summary["value_swings_top10"][0]["abs_delta"] == 0.5
def test_human_play_bot_path_matches_direct_engine_for_ten_games(tmp_path) -> None:
cfg = tiny_cfg(tmp_path)
train_state = create_train_state(cfg, jax.random.PRNGKey(7))
agent = (train_state.params, ActorCritic(cfg.network.hidden_size, cfg.network.num_layers))
for index in range(10):
deck_order = human_deck_order(991, index)
record = play_one_game(
cfg=cfg,
agent=agent,
deck_order=deck_order,
human_seat=index % 2,
session_id=f"test-{index}",
duplicate_set_id=None,
duplicate_game_index=None,
deck_seed=991,
deck_index=index,
input_fn=lambda _prompt: "",
output_fn=lambda _line: None,
human_action_fn=_discard_only_human,
)
direct_scores, direct_actions = _direct_game(cfg, agent, deck_order, index % 2)
assert record["actions"] == direct_actions
assert record["final_scores"]["p0"] == float(direct_scores[0])
assert record["final_scores"]["p1"] == float(direct_scores[1])
def _discard_only_human(state, player: int) -> int:
return int(
discard_only_action(state, jnp.asarray(player, dtype=jnp.int32), jax.random.PRNGKey(0))
)
def _direct_game(
cfg, agent, deck_order: list[int], human_seat: int
) -> tuple[np.ndarray, list[int]]:
params, model = agent
state = reset_from_order(jnp.asarray(deck_order, dtype=jnp.int8))
actions = []
while not bool(state.done):
player = int(state.to_move)
if player == human_seat:
action = _discard_only_human(state, player)
else:
action = evaluate_agent_policy(cfg, params, model, state, player).action
actions.append(action)
state, _, _ = step(state, jnp.asarray(action, dtype=jnp.int32))
return np.asarray(board_score(state), dtype=np.float32), actions