-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbenchmark_selfplay.py
More file actions
42 lines (35 loc) · 1.69 KB
/
Copy pathbenchmark_selfplay.py
File metadata and controls
42 lines (35 loc) · 1.69 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
"""Small reproducible CUDA self-play throughput benchmark."""
from __future__ import annotations
import argparse
from pathlib import Path
from time import perf_counter
import torch
from ai.agents import MCTSAgent
from ai.network import KARDSNet
from ai.observation import ObservationEncoder
from ai.replay_buffer import ReplayBuffer
from ai.selfplay import SelfPlayRunner
from simulator.cards.loader import CardDatabase
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--episodes", type=int, default=2)
parser.add_argument("--simulations", type=int, default=64)
parser.add_argument("--max-actions", type=int, default=30)
parser.add_argument("--workers", type=int, default=1)
args = parser.parse_args()
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this benchmark")
root = Path(__file__).parent; cards = CardDatabase.from_file(root / "data/source/kards_info_cards.json")
model = KARDSNet(hidden_dim=128).to("cuda").eval()
runner = SelfPlayRunner(cards, ObservationEncoder(cards), ReplayBuffer(), args.max_actions, seed=301)
started = perf_counter()
if args.workers == 1:
runner.run(args.episodes, MCTSAgent(model, runner.encoder, args.simulations, 302),
MCTSAgent(model, runner.encoder, args.simulations, 303))
else:
runner.run_parallel(args.episodes, model, args.simulations, root / "data/source/kards_info_cards.json",
workers=args.workers, device="cuda")
elapsed = perf_counter() - started
print(f"workers={args.workers} elapsed={elapsed:.3f}s games_per_second={args.episodes / elapsed:.4f}")
if __name__ == "__main__":
main()