sgf fixes

This commit is contained in:
Sander Land committed 2020-06-05 17:47:57 +02:00
1 parent 20f4598410
commit cd9bf18411
3 files changed
+50 -35

No files matched your search

+1
View File
@@ -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
+29 -22
View File
@@ -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):
+20 -13
View File
@@ -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)