diff --git a/katrain/core/constants.py b/katrain/core/constants.py index b48a438..f391d69 100644 --- a/katrain/core/constants.py +++ b/katrain/core/constants.py @@ -2,6 +2,7 @@ PROGRAM_NAME = "KaTrain" VERSION = "1.7.0" HOMEPAGE = "https://github.com/sanderland/katrain" CONFIG_MIN_VERSION = "1.7.0" # keep config files from this version +ANALYSIS_FORMAT_VERSION = "1.0" OUTPUT_ERROR = -1 OUTPUT_KATAGO_STDERR = -0.5 diff --git a/katrain/core/game.py b/katrain/core/game.py index 92e4079..bb31576 100644 --- a/katrain/core/game.py +++ b/katrain/core/game.py @@ -18,6 +18,7 @@ from katrain.core.constants import ( PLAYER_HUMAN, VERSION, PROGRAM_NAME, + ANALYSIS_FORMAT_VERSION, ) from katrain.core.engine import KataGoEngine from katrain.core.game_node import GameNode @@ -93,7 +94,6 @@ class Game: def analyze_all_nodes(self, priority=0, analyze_fast=False, even_if_present=True): for node in self.root.nodes_in_tree: if even_if_present or not node.analysis_loaded: - print(even_if_present, node.analysis_loaded, node.move) node.clear_analysis() node.analyze(self.engines[node.next_player], priority=priority, analyze_fast=analyze_fast) @@ -321,6 +321,8 @@ class Game: x_properties[bw + "R"] = rank_label(player_info.calculated_rank) if "+" in str(self.end_result): x_properties["RE"] = self.end_result + if save_analysis: + x_properties["KTV"] = ANALYSIS_FORMAT_VERSION self.root.properties = {**root_properties, **{k: [v] for k, v in x_properties.items()}} player_names = {bw: re.sub(r"['<>:\"/\\|?*]", "", self.root.get_property("P" + bw, bw)) for bw in "BW"} base_game_name = f"katrain_{player_names['B']} vs {player_names['W']}" diff --git a/katrain/core/game_node.py b/katrain/core/game_node.py index e6005d7..341f7bb 100644 --- a/katrain/core/game_node.py +++ b/katrain/core/game_node.py @@ -7,16 +7,44 @@ from katrain.core.constants import ( SGF_INTERNAL_COMMENTS_MARKER, SGF_SEPARATOR_MARKER, PROGRAM_NAME, + ANALYSIS_FORMAT_VERSION, ) from katrain.core.lang import i18n from katrain.core.sgf_parser import Move, SGFNode -from katrain.core.utils import evaluation_class, var_to_grid +from katrain.core.utils import evaluation_class, var_to_grid, pack_floats, unpack_floats from katrain.gui.style import INFO_PV_COLOR import base64 import gzip import json +def analysis_dumps(analysis): + analysis = copy.deepcopy(analysis) + for movedict in analysis["moves"].values(): + if "ownership" in movedict: # per-move ownership rarely used + del movedict["ownership"] + ownership_data = pack_floats(analysis.pop("ownership")) + policy_data = pack_floats(analysis.pop("policy")) + main_data = json.dumps(analysis).encode("utf-8") + return [ + base64.standard_b64encode(gzip.compress(data)).decode("utf-8") + for data in [ownership_data, policy_data, main_data] + ] + + +def analysis_loads(property_array, board_squares, version): + if version > ANALYSIS_FORMAT_VERSION: + raise ValueError(f"Can not decode analysis data with version {version}, please update {PROGRAM_NAME}") + ownership_data, policy_data, main_data, *_ = [ + gzip.decompress(base64.standard_b64decode(data)) for data in property_array + ] + return { + **json.loads(main_data), + "policy": unpack_floats(policy_data, board_squares + 1), + "ownership": unpack_floats(ownership_data, board_squares), + } + + class GameNode(SGFNode): """Represents a single game node, with one or more moves and placements.""" @@ -35,7 +63,9 @@ class GameNode(SGFNode): def add_list_property(self, property: str, values: List): if property == "KT": try: - self.analysis = json.loads(gzip.decompress(base64.standard_b64decode(values[0]))) + szx, szy = self.root.board_size + version = self.root.get_property("KTV", "") + self.analysis = analysis_loads(values, szx * szy, version) self.analysis_loaded = True except Exception as e: print(f"Error in loading analysis: {e}") @@ -61,13 +91,7 @@ class GameNode(SGFNode): note = self.note.strip() if save_analysis and self.analysis_complete: try: - analysis = copy.deepcopy(self.analysis) - for movedict in analysis["moves"].values(): - if "ownership" in movedict: - del movedict["ownership"] - properties["KT"] = [ - base64.standard_b64encode(gzip.compress(json.dumps(analysis).encode("utf-8"))).decode("utf-8") - ] + properties["KT"] = analysis_dumps(self.analysis) except Exception as e: print(f"Error in saving analysis: {e}") if self.points_lost and save_comments_class is not None and eval_thresholds is not None: diff --git a/katrain/core/utils.py b/katrain/core/utils.py index a8d1e70..044aa92 100644 --- a/katrain/core/utils.py +++ b/katrain/core/utils.py @@ -57,9 +57,12 @@ def find_package_resource(path, silent_errors=False): def pack_floats(float_list): - return struct.pack('%sf' % len(float_list), *float_list) + return struct.pack(f"{len(float_list)}e", *float_list) +def unpack_floats(str, num): + return struct.unpack(f"{num}e", str) + def format_visits(n): if n < 1000: