diff --git a/katrain/core/ai.py b/katrain/core/ai.py index 422480b..601482e 100644 --- a/katrain/core/ai.py +++ b/katrain/core/ai.py @@ -151,6 +151,7 @@ def generate_ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, x_ai_thoughts = ( f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. " ) + n_moves = int(max(n_moves,0.5 * len(legal_policy_moves))) elif ai_mode in [AI_INFLUENCE, AI_TERRITORY]: weighted_coords, x_ai_thoughts = generate_influence_territory_weights( ai_mode, ai_settings, policy_grid, size diff --git a/katrain/core/engine.py b/katrain/core/engine.py index f2b1132..919c888 100644 --- a/katrain/core/engine.py +++ b/katrain/core/engine.py @@ -1,15 +1,17 @@ import copy import json import subprocess -import traceback import threading import time +import traceback from typing import Callable, Optional + from kivy.utils import platform -from katrain.core.utils import find_package_resource -from katrain.core.lang import i18n -from katrain.core.constants import OUTPUT_ERROR, OUTPUT_KATAGO_STDERR, OUTPUT_DEBUG, OUTPUT_EXTRA_DEBUG + +from katrain.core.constants import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG, OUTPUT_KATAGO_STDERR from katrain.core.game_node import GameNode +from katrain.core.lang import i18n +from katrain.core.utils import find_package_resource class EngineDiedException(Exception): @@ -27,22 +29,25 @@ class KataGoEngine: def get_rules(node): return KataGoEngine.RULESETS.get(str(node.ruleset).lower(), "japanese") - def __init__(self, katrain, config): + def __init__(self, katrain, config, override_command=None): self.katrain = katrain - executable = config["katago"].strip() - if not executable: - if platform == "win": - executable = "katrain/KataGo/katago.exe" - elif platform == "linux": - executable = "katrain/KataGo/katago" - else: # e.g. MacOS after brewing - executable = "katago" + if override_command: + self.command = override_command + else: + executable = config["katago"].strip() + if not executable: + if platform == "win": + executable = "katrain/KataGo/katago.exe" + elif platform == "linux": + executable = "katrain/KataGo/katago" + else: # e.g. MacOS after brewing + executable = "katago" - model = find_package_resource(config["model"]) - cfg = find_package_resource(config["config"]) - exe = find_package_resource(executable) + model = find_package_resource(config["model"]) + cfg = find_package_resource(config["config"]) + exe = find_package_resource(executable) + self.command = f'"{exe}" analysis -model "{model}" -config "{cfg}" -analysis-threads {config["threads"]}' - self.command = f'"{exe}" analysis -model "{model}" -config "{cfg}" -analysis-threads {config["threads"]}' self.queries = {} # outstanding query id -> start time and callback self.config = config self.query_counter = 0 @@ -64,11 +69,11 @@ class KataGoEngine: except (FileNotFoundError, PermissionError, OSError) as e: if not self.config["katago"].strip(): self.katrain.log( - i18n._("Starting default Kata failed").format(command=self.comment, error=e), OUTPUT_ERROR, + i18n._("Starting default Kata failed").format(command=self.command, error=e), OUTPUT_ERROR, ) else: self.katrain.log( - i18n._("Starting Kata failed").format(command=self.comment, error=e), OUTPUT_ERROR, + i18n._("Starting Kata failed").format(command=self.command, error=e), OUTPUT_ERROR, ) self.analysis_thread = threading.Thread(target=self._analysis_read_thread, daemon=True).start() self.stderr_thread = threading.Thread(target=self._read_stderr_thread, daemon=True).start() @@ -82,10 +87,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 + 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}") + 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): diff --git a/katrain/core/game.py b/katrain/core/game.py index d847496..296a3a4 100644 --- a/katrain/core/game.py +++ b/katrain/core/game.py @@ -26,7 +26,14 @@ class Game: DEFAULT_PROPERTIES = {"GM": 1, "FF": 4, "AP": f"KaTrain:{HOMEPAGE}", "CA": "UTF-8"} - def __init__(self, katrain, engine: Union[Dict, KataGoEngine], move_tree: GameNode = None, analyze_fast=False): + def __init__( + self, + katrain, + engine: Union[Dict, KataGoEngine], + move_tree: GameNode = None, + analyze_fast=False, + game_properties: Optional[Dict] = None, + ): self.katrain = katrain if not isinstance(engine, Dict): engine = {"B": engine, "W": engine} @@ -43,7 +50,11 @@ class Game: board_size = katrain.config("game/size") self.komi = katrain.config("game/komi") self.root = GameNode( - properties={**Game.DEFAULT_PROPERTIES, **{"SZ": board_size, "KM": self.komi, "DT": self.game_id}} + properties={ + **Game.DEFAULT_PROPERTIES, + **{"SZ": board_size, "KM": self.komi, "DT": self.game_id}, + **(game_properties or {}), + } ) handicap = katrain.config("game/handicap") if handicap: @@ -272,18 +283,14 @@ class Game: ) def write_sgf( - self, - path: str, - trainer_config: Optional[Dict] = None, - save_feedback: Optional[List] = None, - eval_thresholds: Optional[List] = None, + self, path: str, trainer_config: Optional[Dict] = None, ): if trainer_config is None: trainer_config = self.katrain.config("trainer") - if save_feedback is None: - save_feedback = self.katrain.config("trainer/save_feedback") - if eval_thresholds is None: - eval_thresholds = self.katrain.config("trainer/eval_thresholds") + save_feedback = trainer_config["save_feedback"] + eval_thresholds = trainer_config["eval_thresholds"] + + print(trainer_config, save_feedback, eval_thresholds) def player_name(player_info): return f"{i18n._(player_info.player_type)} ({i18n._(player_info.player_subtype)})" @@ -299,12 +306,12 @@ class Game: os.makedirs(os.path.dirname(file_name), exist_ok=True) show_dots_for = { - bw: trainer_config.get("eval_show_ai", True) or pl.human for bw, pl in self.katrain.players_info.items() + bw: trainer_config.get("eval_show_ai", True) or self.katrain.players_info[bw].human for bw in "BW" } sgf = self.root.sgf( save_comments_player=show_dots_for, save_comments_class=save_feedback, eval_thresholds=eval_thresholds ) - with open(file_name, "w", encoding='utf-8') as f: + with open(file_name, "w", encoding="utf-8") as f: f.write(sgf) return i18n._("sgf written").format(file_name=file_name)