Files
katrain-qt/selfplay.py
T
2020-04-25 01:20:53 +02:00

205 lines
6.8 KiB
Python

import threading
import time, sys
import random
import traceback
from collections import defaultdict
import pickle
from concurrent.futures.thread import ThreadPoolExecutor
from game import Game
from ai import ai_move
from engine import KataGoEngine
from common import OUTPUT_ERROR, OUTPUT_INFO, OUTPUT_DEBUG
from elote import EloCompetitor
DB_FILENAME = "ai_performance.pickle"
class Logger:
def log(self, msg, level):
if level <= OUTPUT_DEBUG:
print(msg)
if level <= OUTPUT_ERROR:
print(msg, file=sys.stderr)
logger = Logger()
class AI:
DEFAULT_ENGINE_SETTINGS = {
"katago": "KataGo/katago-bs",
"model": " models/b15-1.3.2.txt.gz",
"config": "KataGo/analysis_config.cfg",
"threads": 8,
"max_visits": 1,
"max_time": 300.0,
"enable_ownership": False,
}
DEFAULT_SETTINGS = {
"balance_target_score": 2,
"balance_random_loss": 1,
"balance_max_loss": 5,
"balance_min_visits": 20,
"noise_strength": 0.8,
"pick_n": 10,
"pick_frac": 0.2,
"local_stddev": 10,
}
ENGINES = []
LOCK = threading.Lock()
def __init__(self, strategy, ai_settings, engine_settings={}):
self.elo_comp = EloCompetitor(initial_rating=1000)
self.strategy = strategy
self.ai_settings = {**AI.DEFAULT_SETTINGS, **ai_settings}
self.engine_settings = {**AI.DEFAULT_ENGINE_SETTINGS, **engine_settings}
fmt_settings = [f"{k}={v}" for k, v in {**ai_settings, **engine_settings}.items()]
self.name = f"{strategy}({ ','.join(fmt_settings) })"
def get_engine(self): # factory
with AI.LOCK:
for existing_engine_settings, engine in AI.ENGINES:
if existing_engine_settings == self.engine_settings:
return engine
engine = KataGoEngine(logger, self.engine_settings)
AI.ENGINES.append((self.engine_settings, engine))
print("Creating new engine for", self.engine_settings, "now have", len(AI.ENGINES), "engines up")
return engine
def __eq__(self, other):
return self.strategy == other.strategy and self.ai_settings == other.ai_settings and self.engine_settings == other.engine_settings
try:
with open(DB_FILENAME, "rb") as f:
ai_database, all_results = pickle.load(f)
except FileNotFoundError:
ai_database = []
all_results = []
def add_ai(ai):
if ai not in ai_database:
ai_database.append(ai)
print(f"Adding {ai.name}")
else:
print(f"AI {ai.name} already in DB")
def retrieve_ais(selected_ais):
return [ai for ai in ai_database if ai in selected_ais]
add_ai(AI("KataGo", {}, {"max_visits": 50}))
add_ai(AI("Jigo", {}, {"max_visits": 50}))
add_ai(AI("P+Noise", {"noise_strength": 0.9}))
add_ai(AI("P+Noise", {"noise_strength": 0.8}))
add_ai(AI("P+Noise", {"noise_strength": 0.7}))
add_ai(AI("Policy",{}))
add_ai(AI("P+Local", {'local_stddev':1}))
add_ai(AI("P+Local", {'local_stddev':5}))
add_ai(AI("P+Local", {'local_stddev':10}))
add_ai(AI("P+Pick", {'pick_frac':0.2,'pick_n':10}))
add_ai(AI("P+Pick", {'pick_frac':0.3,'pick_n':10}))
new_ais1 = [AI("P+Pick", {'pick_frac':0.3,'pick_n':20}),
AI("P+Pick", {'pick_frac':0.4,'pick_n':20}),
AI("P+Local", {'local_stddev':1}),
AI("KataGo", {}, {"max_visits": 50}),
AI("Jigo", {}, {"max_visits": 50})]
new_ais = [ AI("P+Pick", {'pick_frac':0.4,'pick_n':20}),
AI("P+Local", {'local_stddev':10}),
AI("P+Local", {'local_stddev':5}),
AI("Policy", {})]
new_ais = [ AI("P+Local", {'local_stddev':1,'pick_frac':0.1}),
AI("P+Local", {'local_stddev':1,'pick_frac':0.05}),
AI("P+Pick", {'pick_frac': 0.4, 'pick_n': 20}),
AI("P+Noise", {"noise_strength": 0.8}),
AI("P+Tenuki", {'local_stddev':1}),
AI("P+Tenuki", {'local_stddev':5}),
AI("P+Tenuki", {'local_stddev':10})
]
new_ais1 = [AI("Policy", {}),
AI("Policy", {},{'model':'b10-1.3.txt.gz'}),
AI("Policy", {},{'model':'g170-b30c320x2-s2846858752-d829865719.bin.gz'}),
AI("Policy", {},{'model':'g170-b40c256x2-s2990766336-d830712531.bin.gz'}),
AI("Policy", {}, {'model': 'g170e-b20c256x2-s3761649408-d809581368.bin.gz'}),
]
# AI("KataGo", {}, {"max_visits": 50})]
for ai in new_ais:
add_ai(ai)
N_GAMES = 2
ais_to_test = retrieve_ais(new_ais)
#ais_to_test = ai_database
#ais_to_test = [ai for ai in ai_database if 'visits' not in ai.name]
results = defaultdict(list)
def play_games(black: AI, white: AI, n: int=N_GAMES):
players = {"B": black, "W": white}
engines = {"B": black.get_engine(), "W": white.get_engine()}
tag = f"{black.name} vs {white.name}"
try:
for i in range(n):
game = Game(logger, engines, {})
game.root.add_property("PW", [white.name])
game.root.add_property("PB", [black.name])
game.game_id += f"_{int(random.random()*1e6)}"
start_time = time.time()
while not game.ended:
p = game.current_node.next_player
move = ai_move(game, players[p].strategy, players[p].ai_settings)
while not game.current_node.analysis_ready:
time.sleep(0.001)
print(f"{tag}\tGame {i+1} finished in {time.time()-start_time:.1f}s {game.current_node.format_score()} -> {game.write_sgf('sgf_selfplay/')}", file=sys.stderr)
score = game.current_node.score
if score > 0.3:
black.elo_comp.beat(white.elo_comp)
elif score > -0.3:
black.elo_comp.tied(white.elo_comp)
results[tag].append(score)
all_results.append((black.name, white.name, score))
except Exception as e:
print(e,file=sys.stderr)
traceback.print_tb(file=sys.stderr)
def fmt_score(score):
return f"{'B' if score >= 0 else 'W'}+{abs(score):.1f}"
print(len(ais_to_test),"ais to test")
with ThreadPoolExecutor(max_workers=16) as threadpool:
for b in ais_to_test:
for w in ais_to_test:
if b is not w:
threadpool.submit(play_games, b, w)
print("POOL EXIT")
print("---- RESULTS ----")
for k, v in results.items():
b_win = sum([s > 0.3 for s in v])
w_win = sum([s < -0.3 for s in v])
print(f"{b_win} {k} {w_win} : {list(map(fmt_score,v))}")
print("---- ELO ----")
for ai in sorted(ai_database, key=lambda a: -a.elo_comp.rating):
print(f"{'*' if ai in ais_to_test else ' '} {ai.name}: ELO {ai.elo_comp.rating:.1f}")
print(f"{'*' if ai in ais_to_test else ' '} {ai.name}: ELO {ai.elo_comp.rating:.1f}", file=sys.stderr)
with open(DB_FILENAME, "wb") as f:
pickle.dump((ai_database, all_results), f)
print(f"Done! saving {len(all_results)} to pickle")