selfplay
This commit is contained in:
1 parent
1a3e21610b
commit
4a0ae3963c
11 files changed
+280
-67
No files matched your search
@@ -8,6 +8,8 @@ gtp.log
|
||||
log.txt
|
||||
*.sgf
|
||||
sgfout
|
||||
sgf_selfplay
|
||||
log*
|
||||
my
|
||||
|
||||
# debug
|
||||
|
||||
@@ -2,16 +2,17 @@ import heapq
|
||||
import math
|
||||
import random
|
||||
import time
|
||||
from typing import Dict
|
||||
|
||||
import numpy as np
|
||||
|
||||
from common import OUTPUT_INFO, var_to_grid
|
||||
from game import Move
|
||||
from common import OUTPUT_INFO, var_to_grid, OUTPUT_DEBUG, OUTPUT_ERROR
|
||||
from game import Move, Game, IllegalMoveException
|
||||
|
||||
|
||||
def weighted_selection_without_replacement(items, m):
|
||||
"""For a list of arrays where the first element is a weight, returns random items with those weights, without replacement."""
|
||||
elt = [(math.log(random.random()) / item[0], item) for item in items]
|
||||
elt = [(math.log(random.random()) / item[0], item) for item in items] # magic
|
||||
return [e[1] for e in heapq.nlargest(m, elt)] # NB fine if too small
|
||||
|
||||
|
||||
@@ -19,15 +20,12 @@ def dirichlet_noise(num, dir_alpha=0.3):
|
||||
return np.random.dirichlet([dir_alpha] * num)
|
||||
|
||||
|
||||
def ai_move(game, ai_settings):
|
||||
def ai_move(game: Game, ai_mode: str, ai_settings: Dict):
|
||||
cn = game.current_node
|
||||
while not cn.analysis_ready:
|
||||
game.katrain.controls.set_status("Thinking...") # TODO: non blocking somehow?
|
||||
time.sleep(0.01)
|
||||
# select move
|
||||
time.sleep(0.001)
|
||||
ai_mode = ai_mode.lower()
|
||||
candidate_ai_moves = cn.candidate_moves
|
||||
ai_mode = game.katrain.controls.ai_mode(cn.next_player)
|
||||
|
||||
if ("policy" in ai_mode or "p+" in ai_mode) and cn.policy:
|
||||
policy_moves = cn.policy_ranking
|
||||
size = game.board_size
|
||||
@@ -40,34 +38,43 @@ def ai_move(game, ai_settings):
|
||||
noise_str = ai_settings["noise_strength"]
|
||||
d_noise = dirichlet_noise(len(legal_policy_moves))
|
||||
noisy_policy_moves = [(mv, (1 - noise_str) * pol + noise_str * noise) for ((mv, pol), noise) in zip(legal_policy_moves, d_noise)]
|
||||
aimove = max(noisy_policy_moves, key=lambda mp: mp[1])[0]
|
||||
best = max(noisy_policy_moves, key=lambda mp: mp[1])
|
||||
aimove = best[0]
|
||||
game.katrain.log(f"Noisy policy strategy (strength={noise_str:.2f}) generated move {aimove.gtp()} with value {best[1]}", OUTPUT_DEBUG)
|
||||
if "local" in ai_mode or "tenuki" in ai_mode or "pick" in ai_mode and cn.single_move and cn.single_move.coords:
|
||||
var = ai_settings["local_stddev"] ** 2
|
||||
n_moves = int(ai_settings["pick_frac"] * len(legal_policy_moves) + ai_settings["pick_n"])
|
||||
mx, my = cn.single_move.coords
|
||||
top_5_pass = any([polmove[0].is_pass for polmove in policy_moves[:5]]) # dont make it jump around for the last few sensible non pass moves
|
||||
if not top_5_pass:
|
||||
if "local" in ai_mode:
|
||||
if not top_5_pass and cn.single_move: # otherwise falls through to random by policy move
|
||||
if "local" in ai_mode or "tenuki" in ai_mode:
|
||||
var = ai_settings["local_stddev"] ** 2
|
||||
if cn.single_move.coords is not None:
|
||||
mx, my = cn.single_move.coords
|
||||
weighted_coords = [
|
||||
(math.exp(-0.5 * ((x - mx) ** 2 + (y - my) ** 2) / var), x, y, policy_grid[y][x])
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > 0
|
||||
]
|
||||
game.katrain.log(f"Generated weights based on gaussian with var {var} around {mx},{my}", OUTPUT_DEBUG)
|
||||
if "tenuki" in ai_mode:
|
||||
weighted_coords = [ (1-w,x,y,p) for w,x,y,p in weighted_coords]
|
||||
else:
|
||||
weighted_coords = [(1, *aimove.coords, 1)]
|
||||
game.katrain.log(f"Local strategy: opponent passed but it's not our top move, playing top policy move {aimove}", OUTPUT_DEBUG)
|
||||
else: # if "pick" in ai_mode -> even
|
||||
weighted_coords = [(1, x, y, policy_grid[y][x]) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0]
|
||||
pick_moves = weighted_selection_without_replacement(weighted_coords, n_moves)
|
||||
if pick_moves:
|
||||
best = max(pick_moves, key=lambda m: m[3])
|
||||
aimove = Move((best[1], best[2]), player=cn.next_player)
|
||||
game.katrain.log(
|
||||
f"{aimove} was top from pick moves starting with {[Move((best[1], best[2]), player=cn.next_player).gtp() for best in pick_moves[:10]]} out of {len(pick_moves)} "
|
||||
)
|
||||
game.katrain.log(f"Pick policy strategy (n={n_moves}) generated move {aimove.gtp()} with weight {best[0]} and value {best[3]}", OUTPUT_DEBUG)
|
||||
else:
|
||||
game.katrain.log(f"Pick policy strategy failed to find legal moves, so is passing", OUTPUT_DEBUG)
|
||||
aimove = Move(None, player=cn.next_player) # pass
|
||||
else:
|
||||
weighted_coords = [(policy_grid[y][x], x, y, policy_grid[y][x]) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0]
|
||||
aimove = Move(weighted_selection_without_replacement(weighted_coords, 1)[0][1:3], player=cn.next_player) # just take a random move by policy w/o noise
|
||||
game.katrain.log(f"Pick policy strategy found pass in top 5 moves so chose {aimove} as weighted-by-policy move", OUTPUT_DEBUG)
|
||||
elif "balance" in ai_mode and candidate_ai_moves[0]["move"] != "pass": # don't play suicidal to balance score - pass when it's best
|
||||
sign = cn.player_sign(cn.next_player) # TODO check
|
||||
sel_moves = [ # top move, or anything not too bad, or anything that makes you still ahead
|
||||
@@ -82,13 +89,21 @@ def ai_move(game, ai_settings):
|
||||
)
|
||||
]
|
||||
aimove = Move.from_gtp(random.choice(sel_moves)["move"], player=cn.next_player) # TODO: could be weighted towards worse
|
||||
game.katrain.log(f"Balance strategy considered {len(sel_moves)} moves and chose {aimove} randomly", OUTPUT_DEBUG)
|
||||
elif "jigo" in ai_mode and candidate_ai_moves[0]["move"] != "pass":
|
||||
sign = cn.player_sign(cn.next_player) # TODO check
|
||||
jigo_move = min(candidate_ai_moves, key=lambda move: abs(sign * move["scoreLead"] - 0.5))
|
||||
aimove = Move.from_gtp(jigo_move["move"], player=cn.next_player)
|
||||
game.katrain.log(f"Jigo strategy found {len(candidate_ai_moves)} moves and chose {aimove} as closest to 0.5 point win", OUTPUT_DEBUG)
|
||||
else:
|
||||
if "default" not in ai_mode:
|
||||
if "default" not in ai_mode and 'katago' not in ai_mode:
|
||||
game.katrain.log(f"Unknown AI mode {ai_mode} or policy missing, using default.", OUTPUT_INFO)
|
||||
aimove = Move.from_gtp(candidate_ai_moves[0]["move"], player=cn.next_player)
|
||||
print("COORDS", aimove.coords)
|
||||
game.katrain.log(f"Default strategy found {len(candidate_ai_moves)} moves and chose {aimove} as top move", OUTPUT_DEBUG)
|
||||
|
||||
try:
|
||||
game.play(aimove)
|
||||
except IllegalMoveException as e:
|
||||
game.katrain.log(f"AI Strategy {ai_mode} generated illegal move {aimove}: {e}", OUTPUT_ERROR)
|
||||
|
||||
return aimove
|
||||
Binary file not shown.
+2
-2
@@ -1,10 +1,10 @@
|
||||
{
|
||||
"engine": {
|
||||
"katago": "../KataGo/cpp/katago",
|
||||
"katago": "KataGo/katago-bs",
|
||||
"model": " models/b15-1.3.2.txt.gz",
|
||||
"config": "KataGo/analysis_config.cfg",
|
||||
"threads": 8,
|
||||
"max_visits": 5,
|
||||
"max_visits": 50,
|
||||
"max_time": 3.0,
|
||||
"enable_ownership": true
|
||||
},
|
||||
|
||||
@@ -6,7 +6,7 @@ import threading
|
||||
import time
|
||||
from typing import Callable, Optional
|
||||
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG
|
||||
from game_node import GameNode
|
||||
|
||||
|
||||
@@ -31,6 +31,7 @@ class KataGoEngine:
|
||||
self.query_counter = 0
|
||||
self.katago_process = None
|
||||
self.base_priority = 0
|
||||
self._lock = threading.Lock()
|
||||
|
||||
try:
|
||||
self.katrain.log(f"Starting KataGo with {self.command}", OUTPUT_DEBUG)
|
||||
@@ -47,7 +48,7 @@ class KataGoEngine:
|
||||
self.queries = {}
|
||||
|
||||
def shutdown(self, finish=False):
|
||||
process = getattr(self, "katago_process")
|
||||
process = getattr(self, "katago_process", None)
|
||||
if finish and process:
|
||||
while self.queries and process.poll() is None:
|
||||
time.sleep(0.1)
|
||||
@@ -78,25 +79,25 @@ class KataGoEngine:
|
||||
else:
|
||||
callback, start_time, next_move = self.queries[analysis["id"]]
|
||||
time_taken = time.time() - start_time
|
||||
self.katrain.log(
|
||||
f"[{time_taken:.1f}][{analysis['id']}] KataGo Analysis Received: {analysis.keys()} {line[:80]}...", OUTPUT_DEBUG,
|
||||
)
|
||||
self.katrain.log(f"[{time_taken:.1f}][{analysis['id']}] KataGo Analysis Received: {analysis.keys()} {line[:80]}...", OUTPUT_EXTRA_DEBUG)
|
||||
callback(analysis)
|
||||
del self.queries[analysis["id"]]
|
||||
if getattr(self.katrain, "update_state", None): # easier mocking etc
|
||||
self.katrain.update_state()
|
||||
|
||||
def send_query(self, query, callback, next_move):
|
||||
with self._lock:
|
||||
self.query_counter += 1
|
||||
if "id" not in query:
|
||||
query["id"] = f"QUERY:{str(self.query_counter)}"
|
||||
self.queries[query["id"]] = (callback, time.time(), next_move)
|
||||
if self.katago_process:
|
||||
self.katrain.log(f"Sending query {query['id']}: {str(query)}", OUTPUT_DEBUG)
|
||||
self.katrain.log(f"Sending query {query['id']}: {str(query)}", OUTPUT_EXTRA_DEBUG)
|
||||
self.katago_process.stdin.write((json.dumps(query) + "\n").encode())
|
||||
self.katago_process.stdin.flush()
|
||||
|
||||
def request_analysis(
|
||||
self, analysis_node: GameNode, callback: Callable, visits: int = None, time_limit=True, priority: int = 0, ownership: Optional[bool] = None, next_move=None,
|
||||
self, analysis_node: GameNode, callback: Callable, visits: int = None, time_limit=True, priority: int = 0, ownership: Optional[bool] = None, next_move=None
|
||||
):
|
||||
moves = [m for node in analysis_node.nodes_from_root for m in node.move_with_placements]
|
||||
if next_move:
|
||||
|
||||
@@ -1,14 +1,11 @@
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import List
|
||||
from typing import List, Union, Dict
|
||||
import threading
|
||||
|
||||
from kivy.clock import Clock
|
||||
|
||||
from common import OUTPUT_DEBUG, OUTPUT_INFO
|
||||
from engine import KataGoEngine
|
||||
from game_node import GameNode
|
||||
from sgf_parser import SGF, Move
|
||||
|
||||
@@ -26,9 +23,11 @@ class Game:
|
||||
|
||||
DEFAULT_PROPERTIES = {"GM": 1, "FF": 4, "RU": "JP", "AP": "KaTrain:https://github.com/sanderland/katrain"}
|
||||
|
||||
def __init__(self, katrain, engine, config, move_tree=None):
|
||||
def __init__(self, katrain, engine: Union[Dict, KataGoEngine], config: Dict, move_tree: GameNode = None):
|
||||
self.katrain = katrain
|
||||
self.engine = engine
|
||||
if isinstance(engine, KataGoEngine):
|
||||
engine = {"B": engine, "W": engine}
|
||||
self.engines = engine
|
||||
self.config = config
|
||||
self.game_id = datetime.strftime(datetime.now(), "%Y-%m-%d %H %M %S")
|
||||
|
||||
@@ -39,19 +38,17 @@ class Game:
|
||||
if handicap and not self.root.placements:
|
||||
self.place_handicap_stones(handicap)
|
||||
else:
|
||||
board_size = config["init_size"]
|
||||
self.komi = self.config["init_komi"].get(str(board_size), 6.5)
|
||||
board_size = config.get("init_size", 19)
|
||||
self.komi = self.config.get("init_komi", {}).get(str(board_size), 6.5)
|
||||
self.root = GameNode(properties={**Game.DEFAULT_PROPERTIES, **{"SZ": board_size, "KM": self.komi, "DT": self.game_id}})
|
||||
|
||||
self.current_node = self.root
|
||||
self._init_chains()
|
||||
|
||||
Clock.schedule_once(lambda _dt: self.analyze_all_nodes(-1_000_000), -1) # return faster
|
||||
threading.Thread(target=lambda: self.analyze_all_nodes(-1_000_000), daemon=True).start() # return faster, but bypass Kivy Clock
|
||||
|
||||
def analyze_all_nodes(self, priority=0):
|
||||
self.engine.on_new_game()
|
||||
for node in self.root.nodes_in_tree:
|
||||
node.analyze(self.engine, priority=priority)
|
||||
node.analyze(self.engines[node.next_player], priority=priority)
|
||||
|
||||
# -- move tree functions --
|
||||
def _init_chains(self):
|
||||
@@ -127,7 +124,7 @@ class Game:
|
||||
raise
|
||||
played_node = self.current_node.play(move)
|
||||
self.current_node = played_node
|
||||
played_node.analyze(self.engine)
|
||||
played_node.analyze(self.engines[played_node.next_player])
|
||||
return played_node
|
||||
|
||||
def undo(self, n_times=1):
|
||||
@@ -190,7 +187,7 @@ class Game:
|
||||
return sum(self.chains, [])
|
||||
|
||||
@property
|
||||
def game_ended(self):
|
||||
def ended(self):
|
||||
return self.current_node.parent and self.current_node.is_pass and self.current_node.parent.is_pass
|
||||
|
||||
@property
|
||||
@@ -201,9 +198,9 @@ class Game:
|
||||
return "\n".join("".join(Move.PLAYERS[self.chains[c][0].player] if c >= 0 else "-" for c in line) for line in self.board) + f"\ncaptures: {self.prisoner_count}"
|
||||
|
||||
def write_sgf(self, path=None):
|
||||
black = re.sub(r"['<>:\"/\\|?*]", "", self.root.get_first("PB"))
|
||||
white = re.sub(r"['<>:\"/\\|?*]", "", self.root.get_first("PW"))
|
||||
white = self.root.get_first("PW")
|
||||
black, white = self.root.get_first("PB"), self.root.get_first("PW")
|
||||
black = re.sub(r"['<>:\"/\\|?*]", "", black or "Black")
|
||||
white = re.sub(r"['<>:\"/\\|?*]", "", white or "White")
|
||||
game_name = f"katrain_{black} vs {white} {self.game_id}"
|
||||
file_name = os.path.join(path, f"{game_name}.sgf")
|
||||
os.makedirs(os.path.dirname(file_name), exist_ok=True)
|
||||
@@ -218,15 +215,16 @@ class Game:
|
||||
self.katrain.controls.set_status("Wait for initial analysis to complete before doing a board-sweep or refinement", self.current_node)
|
||||
return
|
||||
|
||||
engine = self.engines[cn.next_player]
|
||||
if mode == "extra":
|
||||
visits = cn.analysis["root"]["visits"] + self.engine.config["max_visits"]
|
||||
visits = cn.analysis["root"]["visits"] + engine.config["max_visits"]
|
||||
self.katrain.controls.set_status(f"Performing additional analysis to {visits} visits")
|
||||
cn.analyze(self.engine, visits=visits, priority=-1_000, time_limit=False)
|
||||
cn.analyze(engine, visits=visits, priority=-1_000, time_limit=False)
|
||||
return
|
||||
elif mode == "sweep":
|
||||
board_size_x, board_size_y = self.board_size
|
||||
analyze_moves = [Move(coords=(x, y), player=cn.next_player) for x in range(board_size_x) for y in range(board_size_y) if (x, y) not in stones]
|
||||
visits = int(self.engine.config["max_visits"] * self.config["sweep_visits_frac"] + 0.5)
|
||||
visits = int(engine.config["max_visits"] * self.config["sweep_visits_frac"] + 0.5)
|
||||
self.katrain.controls.set_status(f"Refining analysis of entire board to {visits} visits")
|
||||
priority = -1_000_000_000
|
||||
else: # mode=='equalize':
|
||||
@@ -235,7 +233,7 @@ class Game:
|
||||
self.katrain.controls.set_status(f"Equalizing analysis of candidate moves to {visits} visits")
|
||||
priority = -1_000
|
||||
for move in analyze_moves:
|
||||
cn.analyze(self.engine, priority, visits=visits, refine_move=move, time_limit=False) # explicitly requested so take as long as you need
|
||||
cn.analyze(engine, priority, visits=visits, refine_move=move, time_limit=False) # explicitly requested so take as long as you need
|
||||
|
||||
def analyze_undo(self, node, train_config):
|
||||
move = node.single_move
|
||||
|
||||
+1
-1
@@ -159,7 +159,7 @@ class BadukPanWidget(Widget):
|
||||
# stones
|
||||
current_node = katrain.game.current_node
|
||||
next_player = katrain.game.next_player
|
||||
game_ended = katrain.game.game_ended
|
||||
game_ended = katrain.game.ended
|
||||
full_eval_on = katrain.controls.eval.active
|
||||
has_stone = {}
|
||||
drawn_stone = {}
|
||||
|
||||
+4
-12
@@ -8,16 +8,7 @@ from kivy.uix.label import Label
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR
|
||||
from engine import KataGoEngine
|
||||
from game import Game, GameNode
|
||||
from gui.kivyutils import (
|
||||
LabelledCheckBox,
|
||||
LabelledFloatInput,
|
||||
LabelledIntInput,
|
||||
LabelledObjectInputArea,
|
||||
LabelledSpinner,
|
||||
LabelledTextInput,
|
||||
ScaledLightLabel,
|
||||
StyledButton,
|
||||
)
|
||||
from gui.kivyutils import LabelledCheckBox, LabelledFloatInput, LabelledIntInput, LabelledObjectInputArea, LabelledSpinner, LabelledTextInput, ScaledLightLabel, StyledButton
|
||||
|
||||
|
||||
class InputParseError(Exception):
|
||||
@@ -156,8 +147,9 @@ class ConfigPopup(QuickConfigGui):
|
||||
self.katrain.log(f"Restarting Engine after {engine_updates} settings change")
|
||||
self.katrain.controls.set_status(f"Restarting Engine after {engine_updates} settings change")
|
||||
old_engine = self.katrain.engine
|
||||
self.katrain.engine = KataGoEngine(self.katrain, self.config["engine"])
|
||||
self.katrain.game.engine = self.katrain.engine
|
||||
new_engine = KataGoEngine(self.katrain, self.config["engine"])
|
||||
self.katrain.engine = {"B": new_engine, "W": new_engine}
|
||||
self.katrain.game.engine = new_engine
|
||||
if getattr(old_engine, "katago_process"):
|
||||
old_engine.shutdown(finish=True)
|
||||
else:
|
||||
|
||||
+1
-1
@@ -2,7 +2,7 @@
|
||||
#:import ew kivy.uix.effectwidget
|
||||
|
||||
|
||||
#:set AI_MODES ['Default','Balance','Jigo','Policy','P+Pick','P+Local','P+Noise']
|
||||
#:set AI_MODES ['Default','Balance','Jigo','Policy','P+Local','P+Tenuki','P+Pick','P+Noise']
|
||||
#:set PLAYER_MODES ['Human', 'Teach','AI:']
|
||||
#:set PLAYER_MODE_VALUES ['human','human+undo','ai']
|
||||
#:set BUTTON_COLOR [0.23, 0.30, 0.35, 1]
|
||||
|
||||
+3
-2
@@ -90,7 +90,7 @@ class KaTrainGui(BoxLayout):
|
||||
if auto_undo and cn.analysis_ready and cn.parent and cn.parent.analysis_ready:
|
||||
self.game.analyze_undo(cn, self.config("trainer")) # not via message loop
|
||||
|
||||
if cn.analysis_ready and "ai" in self.controls.player_mode(cn.next_player) and not cn.children and not self.game.game_ended and not (auto_undo and cn.auto_undo is None):
|
||||
if cn.analysis_ready and "ai" in self.controls.player_mode(cn.next_player) and not cn.children and not self.game.ended and not (auto_undo and cn.auto_undo is None):
|
||||
self("ai-move", cn) # cn mismatch stops this if undo fired
|
||||
|
||||
# Handle prisoners and next player display
|
||||
@@ -127,6 +127,7 @@ class KaTrainGui(BoxLayout):
|
||||
self.message_queue.put([self.game.game_id, message, *args])
|
||||
|
||||
def _do_new_game(self, move_tree=None):
|
||||
self.engine.on_new_game() # clear queries
|
||||
self.game = Game(self, self.engine, self.config("game"), move_tree=move_tree)
|
||||
self.controls.select_mode("analyze" if move_tree and len(move_tree.nodes_in_tree) > 1 else "play")
|
||||
self.controls.graph.initialize_from_game(self.game.root)
|
||||
@@ -134,7 +135,7 @@ class KaTrainGui(BoxLayout):
|
||||
|
||||
def _do_ai_move(self, node=None):
|
||||
if node is None or self.game.current_node == node:
|
||||
ai_move(self.game, self.config("ai"))
|
||||
ai_move(self.game, self.controls.ai_mode(self.game.current_node.next_player), self.config("ai"))
|
||||
|
||||
def _do_undo(self, n_times=1):
|
||||
self.game.undo(n_times)
|
||||
|
||||
+204
@@ -0,0 +1,204 @@
|
||||
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")
|
||||
Reference in new issue
Block a user