take program flags from command line
This commit is contained in:
+21
-17
@@ -4,6 +4,7 @@ from machine_player import MachinePlayer
|
||||
from human_player import HumanPlayer
|
||||
from player_provider import PlayerProvider
|
||||
import json
|
||||
import argparse
|
||||
|
||||
float_formatter = "{:.3f}".format
|
||||
np.set_printoptions(formatter={'float_kind': float_formatter})
|
||||
@@ -56,11 +57,15 @@ def play(player_provider: PlayerProvider, k_max=10000, with_debug=False):
|
||||
move += 1
|
||||
|
||||
|
||||
do_training = 0
|
||||
do_human_player = 0
|
||||
do_load_values = 1
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--train", default=False, help="Train before play")
|
||||
parser.add_argument("--num", default=5, help="Number of games to play")
|
||||
parser.add_argument("--human", default=False, help="X-player is human")
|
||||
parser.add_argument("--load", default=True, help="Load experience from previous trainings")
|
||||
args = parser.parse_args()
|
||||
|
||||
if do_load_values:
|
||||
if args.load:
|
||||
try:
|
||||
with open("values_x.json", "r") as fp:
|
||||
x_values = json.load(fp)
|
||||
@@ -70,19 +75,17 @@ if do_load_values:
|
||||
except Exception:
|
||||
x_values = {}
|
||||
o_values = {}
|
||||
else:
|
||||
else:
|
||||
x_values = {}
|
||||
o_values = {}
|
||||
|
||||
p_hx = HumanPlayer(mark='X', name="Jens")
|
||||
p_mx = MachinePlayer(mark='X', params=MachinePlayer.Params(p_explore=0.1, alpha=0.1), values=x_values)
|
||||
p_mo = MachinePlayer(mark='O', params=MachinePlayer.Params(p_explore=0.1, alpha=0.1), values=o_values)
|
||||
p_mx = MachinePlayer(mark='X', params=MachinePlayer.Params(p_explore=0.1, alpha=0.1), values=x_values)
|
||||
p_mo = MachinePlayer(mark='O', params=MachinePlayer.Params(p_explore=0.1, alpha=0.1), values=o_values)
|
||||
|
||||
K = 5
|
||||
N = 1000
|
||||
if do_training:
|
||||
for k in range(0, K):
|
||||
play(PlayerProvider(p_mx, p_mo), N, False)
|
||||
N = args.num
|
||||
if args.train:
|
||||
for n in range(0, N):
|
||||
play(PlayerProvider(p_mx, p_mo), 1000, False)
|
||||
|
||||
# Values after training
|
||||
# Convert and write JSON object to file
|
||||
@@ -92,12 +95,13 @@ if do_training:
|
||||
with open("values_o.json", "w") as fp:
|
||||
json.dump(p_mo.values, fp, indent=0)
|
||||
|
||||
print(f"Iteration: {k*N:8d}")
|
||||
print(f"Iteration: {n*N:8d}")
|
||||
|
||||
if do_human_player:
|
||||
if args.human:
|
||||
p_hx = HumanPlayer(mark='X', name="Jens")
|
||||
pp = PlayerProvider(p_hx, p_mo)
|
||||
else:
|
||||
else:
|
||||
pp = PlayerProvider(p_mx, p_mo)
|
||||
|
||||
play(pp, 10, True)
|
||||
play(pp, 10, True)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user