-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval_pretrained.py
More file actions
120 lines (102 loc) · 4.95 KB
/
Copy patheval_pretrained.py
File metadata and controls
120 lines (102 loc) · 4.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
#!/usr/bin/env python
"""Evaluate the SL-pretrained Padded23UNetPolicy checkpoint.
Usage:
python eval_pretrained.py # vs expander, 100 games
python eval_pretrained.py --opponent random # vs random
python eval_pretrained.py --games 500 # more games
python eval_pretrained.py --grid 15x15 # bigger grid
python eval_pretrained.py --visualize # show pygame window
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
import equinox as eqx
import jax
import jax.numpy as jnp
import jax.random as jrandom
import numpy as np
PROJECT_ROOT = Path(__file__).resolve().parent
sys.path.insert(0, str(PROJECT_ROOT / "src"))
sys.path.insert(0, str(PROJECT_ROOT / "generals-bots"))
def main():
parser = argparse.ArgumentParser(description="Evaluate SL-pretrained checkpoint")
parser.add_argument("--model-path", default="checkpoints/sl_pretrained.eqx",
help="Path to .eqx checkpoint (default: checkpoints/sl_pretrained.eqx)")
parser.add_argument("--opponent", choices=["random", "expander"], default="expander",
help="Opponent type (default: expander)")
parser.add_argument("--grid", type=str, default="10x10",
help="Grid size as HxW (default: 10x10)")
parser.add_argument("--games", type=int, default=100,
help="Number of evaluation games (default: 100)")
parser.add_argument("--visualize", action="store_true",
help="Show pygame window for game 0")
parser.add_argument("--seed", type=int, default=42,
help="Random seed (default: 42)")
args = parser.parse_args()
# Parse grid dims
parts = args.grid.split("x")
grid_dims = (int(parts[0]), int(parts[1]))
print(f"=" * 60)
print(f"Evaluating SL-pretrained model")
print(f" Model: {args.model_path}")
print(f" Opponent: {args.opponent}")
print(f" Grid: {grid_dims}")
print(f" Games: {args.games}")
print(f" Visualize: {args.visualize}")
print(f"=" * 60)
# ── Build environment (variable-size with pad_to for proper padding) ──
from generals import GeneralsEnv
env = GeneralsEnv(
grid_dims=None,
min_grid_size=grid_dims[0],
max_grid_size=grid_dims[1],
pad_to=grid_dims[1], # ensures env states are padded to max_grid_size
truncation=500,
mountain_density_range=(0.18, 0.26),
num_cities_range=(9, 11),
min_generals_distance=3,
max_generals_distance=None,
castle_val_range=(40, 51),
)
# ── Load model ────────────────────────────────────────────────────
from generals_bot.decision.model import Padded23UNetPolicy
key = jrandom.PRNGKey(args.seed)
model = Padded23UNetPolicy(key)
model = eqx.tree_deserialise_leaves(args.model_path, model)
print(f" Model loaded successfully ({Path(args.model_path).stat().st_size / 1024 / 1024:.1f} MB)")
# Params count
params, _ = eqx.partition(model, eqx.is_array)
n_params = sum(x.size for x in jax.tree.leaves(params))
print(f" Parameters: {n_params:,}")
# ── Run evaluation ────────────────────────────────────────────────
from generals_bot.decision.trainer import evaluate as evaluate_fn
from generals_bot.decision.encoding import encode_observation, encode_observation_with_belief, build_action_mask_4233, action_id_to_engine_action
# Build opponent function
if args.opponent == "random":
from generals_bot.decision.trainer import _random_action as opp_fn
else:
from generals_bot.decision.trainer import _expander_action as opp_fn
results = evaluate_fn(
network=model,
env=env,
pool=None, # will be created internally
num_games=args.games,
grid_dims=grid_dims,
seed=args.seed + 1,
visualize=args.visualize,
opponent_fn=opp_fn,
use_belief=False,
)
# ── Print results ─────────────────────────────────────────────────
print(f"\n{'=' * 60}")
print(f"RESULTS: SL-pretrained vs {args.opponent}")
print(f"{'=' * 60}")
print(f" Wins: {results['wins']:5d} ({100 * results['wins'] / max(results['total'], 1):.1f}%)")
print(f" Losses: {results['losses']:5d} ({100 * results['losses'] / max(results['total'], 1):.1f}%)")
print(f" Draws: {results['draws']:5d} ({100 * results['draws'] / max(results['total'], 1):.1f}%)")
print(f" Win Rate: {results['win_rate']:.2%}")
print(f" Total: {results['total']:5d}")
print(f"{'=' * 60}")
if __name__ == "__main__":
main()