-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathplay.py
More file actions
191 lines (162 loc) · 7.35 KB
/
Copy pathplay.py
File metadata and controls
191 lines (162 loc) · 7.35 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
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
"""
play.py — Play against a trained model
Example:
python play.py --checkpoint checkpoint_0190.pt --human-first
"""
import argparse
import torch
import numpy as np
import openvino as ov
import onnxruntime as ort
from model import AlphaNet
from mcts import Connect4, run_mcts_simulations, print_board
class OpenVINOModel:
"""A wrapper to make an OpenVINO model behave like a PyTorch model for inference."""
def __init__(self, model_path: str, device: str = "AUTO"):
core = ov.Core()
# Read and compile the model for the target device
ov_model = core.read_model(model=model_path)
# Auto-discover the best hardware accelerator to use
available = core.available_devices
if "NPU" in available:
self.ov_device = "NPU"
elif "GPU" in available:
self.ov_device = "GPU"
else:
self.ov_device = "CPU"
self.compiled_model = core.compile_model(model=ov_model, device_name=self.ov_device)
# Get output tensors by their friendly names from the ONNX export
self.policy_output = self.compiled_model.output("policy")
self.value_output = self.compiled_model.output("value")
def __call__(self, x: torch.Tensor):
# Convert torch.Tensor to a numpy array for OpenVINO
x_np = x.cpu().numpy()
result = self.compiled_model([x_np])
# Convert results back to torch.Tensor to match the original PyTorch model output
return torch.from_numpy(result[self.policy_output]), \
torch.from_numpy(result[self.value_output])
def eval(self):
# This method is required to mimic the PyTorch model interface
pass
class ONNXRuntimeModel:
"""A wrapper to make an ONNX Runtime model behave like a PyTorch model for inference."""
def __init__(self, model_path: str, provider: str = "CPUExecutionProvider"):
self.session = ort.InferenceSession(model_path, providers=[provider])
self.input_name = self.session.get_inputs()[0].name
self.ov_device = "RTX 4070 (ORT)" if provider == "CUDAExecutionProvider" else "CPU (ORT)"
def __call__(self, x: torch.Tensor):
x_np = x.cpu().numpy()
result = self.session.run(None, {self.input_name: x_np})
return torch.from_numpy(result[0]), torch.from_numpy(result[1])
def eval(self):
pass
def get_human_move(game: Connect4) -> int:
"""Get a valid column choice from the human player."""
valid_moves = game.get_valid_moves()
while True:
try:
move_str = input(f"Enter column ({', '.join(map(str, valid_moves))}): ")
move = int(move_str)
if move in valid_moves:
return move
else:
print("Invalid column. Please choose a valid, non-full column.")
except ValueError:
print("Invalid input. Please enter a number.")
def main():
parser = argparse.ArgumentParser(description="Play Connect 4 against a trained AlphaZero model.")
parser.add_argument("--model", type=str, required=True, help="Path to the model file (.pt for PyTorch, .onnx for ONNX/OpenVINO).")
parser.add_argument("--simulations", type=int, default=800, help="Number of MCTS simulations per AI move.")
parser.add_argument("--human-first", action="store_true", help="Set this flag for the human to play first as 'X'.")
parser.add_argument("--backend", type=str, default="auto", choices=["auto", "pytorch", "openvino", "onnx-gpu", "onnx-cpu"],
help="Inference backend to use (default: auto).")
args = parser.parse_args()
# Determine backend and device
backend = args.backend
if backend == "auto":
if args.model.endswith(".onnx"):
# RATIONALE: For single-move inference (batch size 1), discrete GPUs (RTX 4070)
# are slower than CPUs/NPUs due to PCIe latency. We prefer the NPU or ONNX-CPU
# for the smoothest UI experience.
core = ov.Core()
if "NPU" in core.available_devices:
backend = "openvino"
else:
backend = "onnx-cpu"
else:
backend = "pytorch"
print(f"Initializing backend: {backend}")
if backend == "openvino":
model = OpenVINOModel(args.model)
print(f"Using OpenVINO for inference on {model.ov_device}.")
device = torch.device("cpu")
elif backend == "onnx-gpu":
model = ONNXRuntimeModel(args.model, provider="CUDAExecutionProvider")
print(f"Using ONNX Runtime for inference on {model.ov_device}.")
device = torch.device("cpu")
elif backend == "onnx-cpu":
model = ONNXRuntimeModel(args.model, provider="CPUExecutionProvider")
print(f"Using ONNX Runtime for inference on {model.ov_device}.")
device = torch.device("cpu")
else: # pytorch
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using PyTorch on device: {device}")
model = AlphaNet().to(device)
try:
checkpoint = torch.load(args.model, map_location=device, weights_only=True)
model.load_state_dict(checkpoint['model_state_dict'])
except (FileNotFoundError, KeyError):
# Fallback for raw state_dict
model.load_state_dict(torch.load(args.model, map_location=device, weights_only=True))
model.eval()
game = Connect4()
human_player = 1 if args.human_first else -1
while True:
print_board(game)
if game.current_player == human_player:
# ── Human Move Evaluation ──
print("Analyzing best options...")
mcts_probs = run_mcts_simulations(
game, model, device,
num_sims=args.simulations,
temperature=0,
add_dirichlet_noise=False,
)
max_p = np.max(mcts_probs)
move = get_human_move(game)
# Assessment Logic
p_move = mcts_probs[move]
if p_move >= 0.95 * max_p:
score, comment = 5, "Brilliant! (Best Move)"
elif p_move >= 0.70 * max_p:
score, comment = 4, "Strong Move"
elif p_move >= 0.30 * max_p:
score, comment = 3, "Decent"
elif p_move >= 0.05 * max_p:
score, comment = 2, "Inaccurate"
else:
score, comment = 1, "Blunder!"
print(f"Assessment: {'★' * score}{'☆' * (5-score)} — {comment}")
else:
print("AI is thinking...")
# For AI moves, use MCTS with temperature=0 to be greedy
mcts_probs = run_mcts_simulations(
game, model, device,
num_sims=args.simulations,
temperature=0,
add_dirichlet_noise=False, # No exploration needed for play
)
move = int(np.argmax(mcts_probs))
print(f"AI chooses column {move}")
r, c = game.play(move)
if game.check_win(r, c):
print_board(game)
winner_char = 'You' if game.current_player != human_player else 'The AI'
print(f"Game over. {winner_char} won!")
break
if not game.get_valid_moves():
print_board(game)
print("Game over. It's a draw!")
break
if __name__ == "__main__":
main()