-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathevaluate.py
More file actions
27 lines (26 loc) · 1.78 KB
/
Copy pathevaluate.py
File metadata and controls
27 lines (26 loc) · 1.78 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
"""Evaluate a checkpoint against RandomAgent or RuleBasedAgent."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from ai.agents import NeuralAgent, RandomAgent, RuleBasedAgent
from ai.evaluator import Evaluator
from ai.network import KARDSNet
from ai.observation import ObservationEncoder
from simulator.cards.loader import CardDatabase
ROOT = Path(__file__).parent
def main() -> None:
parser = argparse.ArgumentParser(); parser.add_argument("--model", required=True); parser.add_argument("--games", type=int, default=20); parser.add_argument("--opponent", choices=("random", "rule", "model"), default="random"); parser.add_argument("--opponent-model"); parser.add_argument("--deck-pool", help="JSON list of {main, ally} evaluation decks"); parser.add_argument("--seed", type=int, default=20260730); parser.add_argument("--report")
args = parser.parse_args(); cards = CardDatabase.from_file(ROOT / "data/source/kards_info_cards.json"); encoder = ObservationEncoder(cards)
candidate = NeuralAgent(KARDSNet.load_checkpoint(args.model), encoder)
if args.opponent == "random": opponent = RandomAgent()
elif args.opponent == "rule": opponent = RuleBasedAgent()
else:
if not args.opponent_model: parser.error("--opponent-model is required when --opponent model")
opponent = NeuralAgent(KARDSNet.load_checkpoint(args.opponent_model), encoder)
raw_pool = json.loads(args.deck_pool) if args.deck_pool else []
deck_pool = [(item["main"], item.get("ally")) for item in raw_pool] or None
evaluator = Evaluator(cards, encoder, seed=args.seed); report = evaluator.evaluate(candidate, opponent, args.games, deck_pool=deck_pool)
if args.report: evaluator.save_report(report, args.report)
print(report)
if __name__ == "__main__": main()