From a322980f79fdce778abb1925e9f52202a6cead16 Mon Sep 17 00:00:00 2001 From: Sander Land Date: Thu, 4 Jun 2020 22:03:41 +0200 Subject: [PATCH] refactor ai --- .github/workflows/release.yaml | 2 +- .github/workflows/test.yaml | 2 +- katrain/__main__.py | 4 +- katrain/config.json | 9 +- katrain/core/ai.py | 230 +++++++++++++-------------------- katrain/core/base_katrain.py | 59 +++++---- katrain/core/constants.py | 6 +- katrain/core/engine.py | 6 + tests/test_ai.py | 25 +++- 9 files changed, 159 insertions(+), 184 deletions(-) diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index 35a11bd..21240c7 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -22,7 +22,7 @@ jobs: pip3 install pytest wheel twine polib - name: Run tests - run: pytest -v tests + run: pytest -v -s tests - name: Run I18N conversion run: python i18n.py diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index 2fb049f..a07564e 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -27,7 +27,7 @@ jobs: pip3 install pytest wheel polib - name: Run tests - run: pytest tests + run: pytest -v -s tests - name: Run I18N conversion run: python i18n.py diff --git a/katrain/__main__.py b/katrain/__main__.py index 2675c44..6c752d0 100644 --- a/katrain/__main__.py +++ b/katrain/__main__.py @@ -39,7 +39,7 @@ from kivy.lang import Builder from kivy.resources import resource_add_path from kivy.uix.popup import Popup from kivy.uix.screenmanager import Screen -from katrain.core.ai import ai_move +from katrain.core.ai import generate_ai_move from kivy.core.window import Window from katrain.core.lang import DEFAULT_LANGUAGE, i18n @@ -231,7 +231,7 @@ class KaTrainGui(Screen, KaTrainBase): mode = self.next_player_info.strategy settings = self.config(f"ai/{mode}") if settings is not None: - ai_move(self.game, mode, settings) + generate_ai_move(self.game, mode, settings) else: self.log(f"AI Mode {mode} not found!", OUTPUT_ERROR) diff --git a/katrain/config.json b/katrain/config.json index c1a9509..6ff6777 100644 --- a/katrain/config.json +++ b/katrain/config.json @@ -1,7 +1,6 @@ { "engine": { "katago": "", - "_hint_katago": "Path to your katago executable", "model": "katrain/models/g170e-b15c192-s1672170752-d466197061.bin.gz", "config": "katrain/KataGo/analysis_config.cfg", "threads": 12, @@ -9,7 +8,6 @@ "fast_visits": 50, "max_time": 3.0, "wide_root_noise": 0.0, - "_hint_wide_root_noise": "A higher value here (typically 0.05-0.1)\nmakes the analysis explore more moves\nat the cost of some strength.", "_enable_ownership": true }, "general": { @@ -18,7 +16,7 @@ "anim_pv_time": 0.5, "debug_level": 0, "lang": "en", - "version": "1.1.2" + "version": "1.2.0" }, "timer": { "byo_length": 30, @@ -94,7 +92,8 @@ "pick_override": 0.95, "stddev": 1.5, "pick_n": 15, - "pick_frac": 0.0 + "pick_frac": 0.0, + "endgame": 0.5 }, "ai:p:tenuki": { "pick_override": 0.85, @@ -120,7 +119,7 @@ "endgame": 0.4 }, "ai:p:rank": { - "kyu": 4.0 + "kyu_rank": 4.0 } } } diff --git a/katrain/core/ai.py b/katrain/core/ai.py index 833fced..464cfe9 100644 --- a/katrain/core/ai.py +++ b/katrain/core/ai.py @@ -20,9 +20,8 @@ from katrain.core.constants import ( AI_TENUKI, AI_TERRITORY, AI_PICK, - AI_RANK, + AI_RANK, ) -from katrain.core.engine import EngineDiedException from katrain.core.game import Game, GameNode, Move @@ -42,157 +41,115 @@ def fmt_moves(moves: List[Tuple[float, Move]]): return ", ".join(f"{mv.gtp()} ({p:.2%})" for p, mv in moves) -def ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode]: +def policy_weighted_move(policy_moves, lower_bound, weaken_fac): + lower_bound, weaken_fac = max(0, lower_bound), max(0.01, weaken_fac) + weighted_coords = [(pv, pv ** (1 / weaken_fac), move) for pv, move in policy_moves if pv > lower_bound and not move.is_pass] + if weighted_coords: + top = weighted_selection_without_replacement(weighted_coords, 1)[0] + ai_thoughts = f"Playing policy-weighted random move {top[2].gtp()} ({top[0]:.1%}) from {len(weighted_coords)} moves above lower_bound of {lower_bound:.1%}." + else: + top = policy_moves[0] + ai_thoughts = f"Playing top policy move because no non-pass move > above lower_bound of {lower_bound:.1%}." + return top[2], ai_thoughts + + +def generate_influence_territory_weights(ai_mode, ai_settings, policy_grid, size): + thr_line = ai_settings["threshold"] - 1 # zero-based + if ai_mode == AI_INFLUENCE: + weight = lambda x, y: (1 / ai_settings["line_weight"]) ** (max(0, thr_line - min(size[0] - 1 - x, x)) + max(0, thr_line - min(size[1] - 1 - y, y))) + else: + weight = lambda x, y: (1 / ai_settings["line_weight"]) ** (max(0, min(size[0] - 1 - x, x, size[1] - 1 - y, y) - thr_line)) + weighted_coords = [(policy_grid[y][x] * weight(x, y), weight(x, y), x, y) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0] + ai_thoughts = f"Generated weights for {ai_mode} according to weight factor {ai_settings['line_weight']} and distance from {thr_line + 1}th line. " + return weighted_coords, ai_thoughts + + +def generate_local_tenuki_weights(ai_mode, ai_settings, policy_grid, cn, size): + var = ai_settings["stddev"] ** 2 + mx, my = cn.move.coords + weighted_coords = [(policy_grid[y][x], math.exp(-0.5 * ((x - mx) ** 2 + (y - my) ** 2) / var), x, y) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0] + ai_thoughts = f"Generated weights based on one minus gaussian with variance {var} around coordinates {mx},{my}. " + if ai_mode == AI_TENUKI: + weighted_coords = [(p, 1 - w, x, y) for p, w, x, y in weighted_coords] + ai_thoughts = f"Generated weights based on one minus gaussian with variance {var} around coordinates {mx},{my}. " + return weighted_coords, ai_thoughts + + +def generate_ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode]: cn = game.current_node while not cn.analysis_ready: time.sleep(0.01) - engine = game.engines[cn.next_player] - if engine.katago_process.poll() is not None: # TODO: clean up - raise EngineDiedException(f"Engine for {cn.next_player} ({engine.config}) died") + game.engines[cn.next_player].check_alive(exception_if_dead=True) + ai_thoughts = "" if (ai_mode in AI_STRATEGIES_POLICY) and cn.policy: # pure policy based move policy_moves = cn.policy_ranking pass_policy = cn.policy[-1] - top_5_pass = any( - [polmove[1].is_pass for polmove in policy_moves[:5]] - ) # dont make it jump around for the last few sensible non pass moves + # dont make it jump around for the last few sensible non pass moves + top_5_pass = any([polmove[1].is_pass for polmove in policy_moves[:5]]) size = game.board_size policy_grid = var_to_grid(cn.policy, size) # type: List[List[float]] top_policy_move = policy_moves[0][1] ai_thoughts += f"Using policy based strategy, base top 5 moves are {fmt_moves(policy_moves[:5])}. " - len_legal_policy_moves = len([(pol, mv) for pol, mv in policy_moves if not mv.is_pass if pol > 0]) - if ai_mode == AI_POLICY and cn.depth <= ai_settings["opening_moves"]: + if (ai_mode == AI_POLICY and cn.depth <= ai_settings["opening_moves"]) or (ai_mode in [AI_LOCAL, AI_TENUKI] and not cn.move or cn.move.coords is None): ai_mode = AI_WEIGHTED - ai_thoughts += f"Switching to weighted strategy in the opening {int(ai_settings['opening_moves'])} moves. " + ai_thoughts += f"Strategy override, using policy-weighted strategy instead. " ai_settings = {"pick_override": 0.9, "weaken_fac": 1, "lower_bound": 0.02} - if ai_mode == AI_RANK: - ai_settings = {"pick_override": (0.8*(1-((size[0]*size[1])-len_legal_policy_moves)/(size[0]*size[1])*.5)), "kyu": ai_settings["kyu"] } + if top_5_pass: aimove = top_policy_move ai_thoughts += "Playing top one because one of them is pass." elif ai_mode == AI_POLICY: aimove = top_policy_move ai_thoughts += f"Playing top policy move {aimove.gtp()}." - elif policy_moves[0][0] > ai_settings["pick_override"]: - aimove = top_policy_move - ai_thoughts += ( - f"Top policy move has weight > {ai_settings['pick_override']:.1%}, so overriding other strategies." - ) - elif ai_mode == AI_WEIGHTED: - lower_bound = max(0, ai_settings["lower_bound"]) * 2 # compensate for first halving in loop - weaken_fac = max(0.01, ai_settings["weaken_fac"]) - weighted_coords = [] - while not weighted_coords and lower_bound > 1e-6: # fix edge case where no moves are > lb - lower_bound /= 2 - weighted_coords = [ - (policy_grid[y][x], policy_grid[y][x] ** (1 / weaken_fac), x, y) - for x in range(size[0]) - for y in range(size[1]) - if policy_grid[y][x] > lower_bound - ] - top = weighted_selection_without_replacement(weighted_coords, 1) - if top: - best = top[0] - policy_value = best[0] - coords = best[2:] + else: # weighted or pick-based + legal_policy_moves = [(pol, mv) for pol, mv in policy_moves if not mv.is_pass and pol > 0] + board_squares = size[0] * size[1] + if ai_mode == AI_RANK: # calibrated, override from 0.8 at start to ~0.4 at full board + override = 0.8 * (1 - 0.5 * (board_squares - len(legal_policy_moves)) / board_squares) else: - policy_value = pass_policy - coords = None - aimove = Move(coords, player=cn.next_player) # just take a random move by policy w/o noise - ai_thoughts += f"Playing policy-weighted random move {aimove.gtp()} ({policy_value:.1%})" + ( - " because no other moves were found." - if not top - else f" because strategy is weighted (lower bound={lower_bound:.2%}, num moves > lb={len(weighted_coords)})." - ) - elif ai_mode in AI_STRATEGIES_PICK: - legal_policy_moves = [(pol, mv) for pol, mv in policy_moves if not mv.is_pass if pol > 0] - if ai_mode!=AI_RANK: - n_moves = int(ai_settings["pick_frac"] * len(legal_policy_moves) + ai_settings["pick_n"]) - if ai_mode in [AI_INFLUENCE, AI_TERRITORY]: + override = ai_settings["pick_override"] - thr_line = ai_settings["threshold"] - 1 # zero-based - if cn.depth >= ai_settings["endgame"] * size[0] * size[1]: - weighted_coords = [ - (policy_grid[y][x], 1, x, y) - for x in range(size[0]) - for y in range(size[1]) - if policy_grid[y][x] > 0 - ] - ai_thoughts += ( - f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. " - ) - else: - if ai_mode == AI_INFLUENCE: - weight = lambda x, y: (1 / ai_settings["line_weight"]) ** ( - max(0, thr_line - min(size[0] - 1 - x, x)) + max(0, thr_line - min(size[1] - 1 - y, y)) - ) - else: - weight = lambda x, y: (1 / ai_settings["line_weight"]) ** ( - max(0, min(size[0] - 1 - x, x, size[1] - 1 - y, y) - thr_line) - ) - weighted_coords = [ - (policy_grid[y][x] * weight(x, y), weight(x, y), x, y) - for x in range(size[0]) - for y in range(size[1]) - if policy_grid[y][x] > 0 - ] - ai_thoughts += f"Generated weights for {ai_mode} according to weight factor {ai_settings['line_weight']} and distance from {thr_line+1}th line. " - elif ai_mode in [AI_LOCAL, AI_TENUKI]: - var = ai_settings["stddev"] ** 2 - if not cn.move or cn.move.coords is None: - weighted_coords = [(1, 1, *top_policy_move.coords)] # if "pick" in ai_mode -> even - ai_thoughts += f"No previous non-pass move, faking weights to play top policy move. " - else: - mx, my = cn.move.coords - weighted_coords = [ - (policy_grid[y][x], math.exp(-0.5 * ((x - mx) ** 2 + (y - my) ** 2) / var), x, y) - for x in range(size[0]) - for y in range(size[1]) - if policy_grid[y][x] > 0 - ] - if ai_mode == AI_TENUKI: - if cn.depth < ai_settings["endgame"] * size[0] * size[1]: - weighted_coords = [(p, 1 - w, x, y) for p, w, x, y in weighted_coords] - ai_thoughts += f"Generated weights based on one minus gaussian with variance {var} around coordinates {mx},{my}. " - else: - weighted_coords = [(p, 1, x, y) for p, w, x, y in weighted_coords] - ai_thoughts += f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. " - else: - ai_thoughts += ( - f"Generated weights based on gaussian with variance {var} around coordinates {mx},{my}. " - ) - elif ai_mode == AI_PICK: - weighted_coords = [ - (policy_grid[y][x], 1, x, y) - for x in range(size[0]) - for y in range(size[1]) - if policy_grid[y][x] > 0 - ] - elif ai_mode == AI_RANK: - n_moves = int(round((size[0]*size[1])/361*10**(-0.05737*ai_settings["kyu"] + 1.9482))) - weighted_coords = [ - (policy_grid[y][x], 1, x, y) - for x in range(size[0]) - for y in range(size[1]) - if policy_grid[y][x] > 0 - ] - else: - raise ValueError(f"Unknown AI mode {ai_mode}") - pick_moves = weighted_selection_without_replacement(weighted_coords, n_moves) - ai_thoughts += f"Picked {min(n_moves,len(weighted_coords))} random moves according to weights. " - if pick_moves: - new_top = [(p, Move((x, y), player=cn.next_player)) for p, wt, x, y in heapq.nlargest(5, pick_moves)] - aimove = new_top[0][1] - ai_thoughts += f"Top 5 among these were {fmt_moves(new_top)} and picked top {aimove.gtp()}. " - if new_top[0][0] < pass_policy: - ai_thoughts += f"But found pass ({pass_policy:.2%} to be higher rated than {aimove.gtp()} ({new_top[0][0]:.2%}) so will play top policy move instead." - aimove = top_policy_move - else: + if policy_moves[0][0] > override: aimove = top_policy_move - ai_thoughts += f"Pick policy strategy {ai_mode} failed to find legal moves, so is playing top policy move {aimove.gtp()}." - else: - raise ValueError(f"Unknown AI mode {ai_mode}") + ai_thoughts += f"Top policy move has weight > {override:.1%}, so overriding other strategies." + elif ai_mode == AI_WEIGHTED: + aimove, ai_thoughts = policy_weighted_move(policy_moves, ai_settings["lower_bound"], ai_settings["weaken_fac"]) + elif ai_mode in AI_STRATEGIES_PICK: + + if ai_mode != AI_RANK: + n_moves = int(ai_settings["pick_frac"] * len(legal_policy_moves) + ai_settings["pick_n"]) + else: + n_moves = int(round(board_squares / 361 * 10 ** (-0.05737 * ai_settings["kyu_rank"] + 1.9482))) + + if ai_mode in [AI_INFLUENCE, AI_TERRITORY, AI_LOCAL, AI_TENUKI]: + if cn.depth > ai_settings["endgame"] * board_squares: + weighted_coords = [(pol, 1, *mv.coords) for pol, mv in legal_policy_moves] + x_ai_thoughts = f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. " + elif ai_mode in [AI_INFLUENCE, AI_TERRITORY]: + weighted_coords, x_ai_thoughts = generate_influence_territory_weights(ai_mode, ai_settings, policy_grid, size) + else: # ai_mode in [AI_LOCAL, AI_TENUKI] + weighted_coords, x_ai_thoughts = generate_local_tenuki_weights(ai_mode, ai_settings, policy_grid, cn, size) + ai_thoughts += x_ai_thoughts + else: # ai_mode in [AI_PICK, AI_RANK]: + weighted_coords = [(policy_grid[y][x], 1, x, y) 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) + ai_thoughts += f"Picked {min(n_moves,len(weighted_coords))} random moves according to weights. " + + if pick_moves: + new_top = [(p, Move((x, y), player=cn.next_player)) for p, wt, x, y in heapq.nlargest(5, pick_moves)] + aimove = new_top[0][1] + ai_thoughts += f"Top 5 among these were {fmt_moves(new_top)} and picked top {aimove.gtp()}. " + if new_top[0][0] < pass_policy: + ai_thoughts += f"But found pass ({pass_policy:.2%} to be higher rated than {aimove.gtp()} ({new_top[0][0]:.2%}) so will play top policy move instead." + aimove = top_policy_move + else: + aimove = top_policy_move + ai_thoughts += f"Pick policy strategy {ai_mode} failed to find legal moves, so is playing top policy move {aimove.gtp()}." + else: + raise ValueError(f"Unknown Policy-based AI mode {ai_mode}") else: # Engine based move candidate_ai_moves = cn.candidate_moves top_cand = Move.from_gtp(candidate_ai_moves[0]["move"], player=cn.next_player) @@ -202,21 +159,12 @@ def ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode else: if ai_mode == AI_JIGO: sign = cn.player_sign(cn.next_player) - jigo_move = min( - candidate_ai_moves, key=lambda move: abs(sign * move["scoreLead"] - ai_settings["target_score"]) - ) + jigo_move = min(candidate_ai_moves, key=lambda move: abs(sign * move["scoreLead"] - ai_settings["target_score"])) aimove = Move.from_gtp(jigo_move["move"], player=cn.next_player) ai_thoughts += f"Jigo strategy found {len(candidate_ai_moves)} candidate moves (best {top_cand.gtp()}) and chose {aimove.gtp()} as closest to 0.5 point win" elif ai_mode == AI_SCORELOSS: c = ai_settings["strength"] - moves = [ - ( - d["pointsLost"], - math.exp(min(200, -c * max(0, d["pointsLost"]))), - Move.from_gtp(d["move"], player=cn.next_player), - ) - for d in candidate_ai_moves - ] + moves = [(d["pointsLost"], math.exp(min(200, -c * max(0, d["pointsLost"]))), Move.from_gtp(d["move"], player=cn.next_player),) for d in candidate_ai_moves] topmove = weighted_selection_without_replacement(moves, 1)[0] aimove = topmove[2] ai_thoughts += f"ScoreLoss strategy found {len(candidate_ai_moves)} candidate moves (best {top_cand.gtp()}) and chose {aimove.gtp()} (weight {topmove[1]:.3f}, point loss {topmove[0]:.1f}) based on score weights." diff --git a/katrain/core/base_katrain.py b/katrain/core/base_katrain.py index a9a2bdd..2d04655 100644 --- a/katrain/core/base_katrain.py +++ b/katrain/core/base_katrain.py @@ -45,13 +45,13 @@ class KaTrainBase: """Settings, logging, and players functionality, so other classes like bots who need a katrain instance can be used without a GUI""" - def __init__(self, **kwargs): - self.debug_level = 0 + def __init__(self, force_package_config=False,debug_level=0, **kwargs): + self.debug_level = debug_level self.game = None self.logger = lambda message, level=OUTPUT_INFO: self.log(message, level) - self.config_file = self._load_config() - self.debug_level = self.config("general/debug_level", OUTPUT_INFO) + self.config_file = self._load_config(force_package_config=force_package_config) + self.debug_level = debug_level or self.config("general/debug_level", OUTPUT_INFO) Config.set("kivy", "log_level", "error") if self.debug_level >= OUTPUT_DEBUG: @@ -68,37 +68,40 @@ class KaTrainBase: elif self.debug_level >= level: print(message) - def _load_config(self): + def _load_config(self,force_package_config): if len(sys.argv) > 1 and sys.argv[1].endswith(".json"): config_file = os.path.abspath(sys.argv[1]) self.log(f"Using command line config file {config_file}", OUTPUT_INFO) else: user_config_file = find_package_resource(self.USER_CONFIG_FILE) package_config_file = find_package_resource(self.PACKAGE_CONFIG_FILE) - try: - if not os.path.exists(user_config_file): - os.makedirs(os.path.split(user_config_file)[0], exist_ok=True) - shutil.copyfile(package_config_file, user_config_file) - config_file = user_config_file - self.log(f"Copied package config to local file {config_file}", OUTPUT_INFO) - else: # user file exists - version = JsonStore(user_config_file, indent=4).get("general")["version"] - if version < CONFIG_MIN_VERSION: - backup = user_config_file + f".{version}.backup" - shutil.copyfile(user_config_file, backup) - shutil.copyfile(package_config_file, user_config_file) - self.log( - f"Copied package config file to {user_config_file} as user file is outdated (<{CONFIG_MIN_VERSION}). Old version stored as {backup}", - OUTPUT_INFO, - ) - config_file = user_config_file - self.log(f"Using user config file {config_file}", OUTPUT_INFO) - except Exception as e: + if force_package_config: config_file = package_config_file - self.log( - f"Using package config file {config_file} (exception {e} occurred when finding or creating user config)", - OUTPUT_INFO, - ) + else: + try: + if not os.path.exists(user_config_file): + os.makedirs(os.path.split(user_config_file)[0], exist_ok=True) + shutil.copyfile(package_config_file, user_config_file) + config_file = user_config_file + self.log(f"Copied package config to local file {config_file}", OUTPUT_INFO) + else: # user file exists + version = JsonStore(user_config_file, indent=4).get("general")["version"] + if version < CONFIG_MIN_VERSION: + backup = user_config_file + f".{version}.backup" + shutil.copyfile(user_config_file, backup) + shutil.copyfile(package_config_file, user_config_file) + self.log( + f"Copied package config file to {user_config_file} as user file is outdated (<{CONFIG_MIN_VERSION}). Old version stored as {backup}", + OUTPUT_INFO, + ) + config_file = user_config_file + self.log(f"Using user config file {config_file}", OUTPUT_INFO) + except Exception as e: + config_file = package_config_file + self.log( + f"Using package config file {config_file} (exception {e} occurred when finding or creating user config)", + OUTPUT_INFO, + ) try: self._config_store = JsonStore(config_file, indent=4) except Exception as e: diff --git a/katrain/core/constants.py b/katrain/core/constants.py index c4244dd..1ae368b 100644 --- a/katrain/core/constants.py +++ b/katrain/core/constants.py @@ -1,6 +1,6 @@ -VERSION = "1.1.2" +VERSION = "1.2.0" HOMEPAGE = "https://github.com/sanderland/katrain" -CONFIG_MIN_VERSION = "1.1.2" +CONFIG_MIN_VERSION = "1.2.0" PLAYER_HUMAN, PLAYER_AI = "player:human", "player:ai" PLAYER_TYPES = [PLAYER_HUMAN, PLAYER_AI] @@ -33,13 +33,13 @@ AI_STRATEGIES_RECOMMENDED_ORDER = [ AI_SCORELOSS, AI_POLICY, AI_WEIGHTED, + AI_RANK, AI_PICK, AI_LOCAL, AI_TENUKI, AI_TERRITORY, AI_INFLUENCE, AI_JIGO, - AI_RANK, ] diff --git a/katrain/core/engine.py b/katrain/core/engine.py index ecf6748..5897002 100644 --- a/katrain/core/engine.py +++ b/katrain/core/engine.py @@ -82,6 +82,12 @@ class KataGoEngine: self.shutdown(finish=False) self.start() + def check_alive(self,exception_if_dead=False): + ok = self.katago_process and self.katago_process.poll() is None + if not ok and exception_if_dead: + raise EngineDiedException(f"Engine died (process {self.katago_process}, poll {self.katago_process and self.katago_process.poll()}) config {self.config}") + return ok + def shutdown(self, finish=False): process = self.katago_process if finish and process: diff --git a/tests/test_ai.py b/tests/test_ai.py index f70efa1..3d67138 100644 --- a/tests/test_ai.py +++ b/tests/test_ai.py @@ -1,8 +1,27 @@ -import pytest - -from katrain.core.constants import AI_STRATEGIES_RECOMMENDED_ORDER, AI_STRATEGIES +from katrain.core.ai import generate_ai_move +from katrain.core.constants import AI_STRATEGIES_RECOMMENDED_ORDER, AI_STRATEGIES, OUTPUT_INFO +from katrain.core.base_katrain import KaTrainBase +from katrain.core.engine import KataGoEngine +from katrain.core.game import Game +from katrain.core.constants import AI_STRATEGIES class TestAI: def test_order(self): assert set(AI_STRATEGIES_RECOMMENDED_ORDER) == set(AI_STRATEGIES) + + def test_ai_strategies(self): + katrain = KaTrainBase(force_package_config=True, debug_level=0) + engine = KataGoEngine(katrain, katrain.config("engine")) + game = Game(katrain, engine) + + n_rounds = 3 + for _ in range(n_rounds): + for strategy in AI_STRATEGIES: + settings = katrain.config(f"ai/{strategy}") + move, played_node = generate_ai_move(game, strategy, settings) + katrain.log(f"Testing strategy {strategy} -> {move}", OUTPUT_INFO) + assert move.coords is not None + assert played_node == game.current_node + + assert game.current_node.depth == len(AI_STRATEGIES) * n_rounds