useful and descriptive commit message

This commit is contained in:
Sander Land committed 2020-04-15 17:13:33 +02:00
1 parent d0f49b5f72
commit 15276c9062
14 files changed
+338 -443

No files matched your search

+1 -1
View File
@@ -1,4 +1,4 @@
This repository includes binaries for 'katago', which is Copyright 2019 David J Wu et al.
This repository includes binaries for 'KataGo', which is Copyright 2019 David J Wu et al.
For on related licenses for these binaries and libraries see https://github.com/lightvector/KataGo
Aside from the above, the license for all other content in this repository is as follows:
+1 -1
View File
@@ -1,6 +1,6 @@
# KaTrain v1.0
This repository contains tool for playing go with AI feedback.
This repository contains tool for analyzing and playing go with AI feedback.
The idea is to give immediate feedback on the many large mistakes we make in terms of inefficient moves.
It is based on the KataGo AI and relies heavily on score estimation rather than win rate.
+16 -31
View File
@@ -1,56 +1,41 @@
{
"analysis": {
"pass_visits": 100,
"pass_visits_fast": 25,
"engine": {
"command": "KataGo/katago analysis -model models/b15-1.3.2.txt.gz -config KataGo/analysis_config.cfg -analysis-threads 8",
"visits": 2000,
"visits_fast": 500,
"analyze_all_visits": 500,
"analyze_all_visits_fast": 100
"visits_fast": 500
},
"board": {
"game": {
"size": 19,
"komi_19": 6.5,
"komi_13": 6.5,
"komi_9": 6.5
"komi_9": 6.5,
"balance_play_target_score": 2,
"balance_play_randomize_eval": 1,
"balance_play_min_eval": 2,
"balance_play_min_visits": 20,
"undo_point_threshold": 1.5,
"num_undo_prompts": 1,
"sgf_load": "~/Downloads",
"sgf_save": "./sgfout"
},
"ui": {
"board_ui": {
"size_min": 1,
"size_max": 15,
"stones": {"B": [0.05, 0.05, 0.05], "W": [0.95, 0.95, 0.95] },
"outline": {"B": [0.3,0.3,0.3,0.5], "W": [0.7, 0.7, 0.7,0.5] },
"ghost_alpha": 0.5,
"eval_colors": [[0.537, 0.129, 0.42], [1, 0, 0], [1, 0.95, 0], [0.117, 0.588, 0]],
"undo_alpha": 0.5,
"undo_scale": 0.95,
"eval_knots": [0, 0.5, 0.875, 1],
"eval_bounds": [1,12],
"eval_colors": [[0.537, 0.129, 0.42], [1, 0, 0], [1, 0.95, 0], [0.5, 0.6, 0], [0.117, 0.588, 0]],
"eval_thresholds": [10,5,1.5,0.5],
"board_margin": 1.5,
"starpoint_size": 0.1,
"stone_size": 0.475,
"board_color": [0.85, 0.68, 0.40],
"line_color": [0,0,0],
"min_eval_temperature": 0.5
},
"engine": {
"command": "KataGo/katago analysis -model models/b15-1.3.2.txt.gz -config KataGo/analysis_config.cfg -analysis-threads 8"
},
"trainer": {
"balance_play_target_score": 2,
"balance_play_randomize_eval": 0.95,
"balance_play_min_eval": 0.875,
"balance_play_min_visits": 20,
"undo_eval_threshold": 0.875,
"undo_point_threshold": 1,
"num_undo_prompts": 1,
"sgf_show_best_move_threshold": 0.95,
"dont_lock_undos": false,
"eval_off_show_last": 3
},
"debug": {
"level": 1
},
"sgf": {
"load": "~/Downloads",
"save": "./sgfout"
}
}
+4
View File
@@ -0,0 +1,4 @@
OUTPUT_ERROR = -1
OUTPUT_INFO = 0
OUTPUT_DEBUG = 1
OUTPUT_EXTRA_DEBUG = 2
+57 -56
View File
@@ -1,73 +1,74 @@
import os,sys,shlex
import subprocess, threading
import shlex, json, time, copy
from .katrain import OUTPUT_ERROR, OUTPUT_DEBUG
import copy
import json
import os
import shlex
import subprocess
import sys
import threading
import time
from typing import Callable
from constants import OUTPUT_DEBUG, OUTPUT_ERROR
from game_node import GameNode
class KataGoEngine:
def __init__(self,config,logger):
"""Starts and communicates with the KataGO analysis engine"""
def __init__(self, katrain, config):
self.command = os.path.join(config["command"])
self.logger = logger
self.katrain = katrain
if "win" not in sys.platform:
self.command = shlex.split(self.command)
self.kata = None
self.query_time = {}
self.outstanding_analysis_queries = [] # allows faster interaction while kata is starting
# engine main loop
def _engine_thread(self):
self.queries = {}
self.config = config
self.visits = [config["visits"], config["visits_fast"]]
self.fast = True
self.query_counter = 0
self.katago_process = None
try:
self.kata = subprocess.Popen(self.command, stdin=subprocess.PIPE, stdout=subprocess.PIPE)
self.katago_process = subprocess.Popen(self.command, stdin=subprocess.PIPE, stdout=subprocess.PIPE)
self.analysis_thread = threading.Thread(target=self._analysis_read_thread, daemon=True).start()
except FileNotFoundError:
self.logger(f"Starting kata with command '{self.command}' failed. If you are on Mac or Linux, please edit configuration file (config.json) to point to the correct KataGo executable.",OUTPUT_ERROR) # fmt off
self.analysis_thread = threading.Thread(target=self._analysis_read_thread, daemon=True).start()
self.katrain.log(
f"Starting kata with command '{self.command}' failed. If you are on Mac or Linux, please edit configuration file (config.json) to point to the correct KataGo executable.",
OUTPUT_ERROR,
)
# analysis thread
def _analysis_read_thread(self):
while True:
while self.outstanding_analysis_queries:
self._send_analysis_query(self.outstanding_analysis_queries.pop(0))
line = self.kata.stdout.readline()
if not line: # occasionally happens?
return
try:
analysis = json.loads(line)
except json.JSONDecodeError as e:
print(f"JSON decode error: '{e}' encountered after receiving input '{line}'")
return
if self.debug:
print(f"[{time.time()-self.query_time.get(analysis['id'],0):.1f}] kata analysis received:", line[:80], "...")
line = self.katago_process.stdout.readline()
analysis = json.loads(line)
if "error" in analysis:
if "AA" not in analysis["id"]: # silently drop illegal moves from analysis all
print(analysis)
self.logger(f"ERROR IN KATA ANALYSIS: {analysis['error']}")
self.katrain.log(f"ERROR IN KATA ANALYSIS: {analysis['error']}")
else:
self.board.store_analysis(analysis)
self.update_evaluation()
callback, start_time = self.queries[analysis["id"]]
time_taken = time.time() - start_time
self.katrain.log(f"[{time_taken:.1f}][{analysis['id']}] KataGo Analysis Received:", line[:80], "...")
callback(analysis)
self.katrain.update_evaluation() # TODO: ??
def _send_analysis_query(self, query):
self.query_time[query["id"]] = time.time()
query = {"rules": "japanese", "komi": self.komi, "boardXSize": self.board_size, "boardYSize": self.board_size, "analyzeTurns": [len(query["moves"])], **query}
if self.kata:
self.kata.stdin.write((json.dumps(query) + "\n").encode())
self.kata.stdin.flush()
else: # early on / root / etc
self.outstanding_analysis_queries.append(copy.copy(query))
def _request_analysis(self, analysis_node: KaTrainSGFNode, faster=False, min_visits=0, priority=0):
faster_fac = 5 if faster else 1
node_id = analysis_node.id
fast = self.ai_fast.active
def request_analysis(self, analysis_node: GameNode, callback: Callable, faster=False, min_visits=0, priority=0):
query_id = f"QUERY:{str(self.query_counter)}"
self.query_counter += 1
visits = 100 # TODO / fast = self.ai_fast.active
if faster:
visits /= 5
moves = [m for node in analysis_node.nodes_from_root for m in node.move_with_placements]
query = {
"id": str(node_id),
"moves": [[m.player, m.gtp()] for node in analysis_node.nodes_from_root for m in node.move_with_placements],
"id": query_id,
"moves": [[m.player, m.gtp()] for m in moves],
"includeOwnership": True,
"maxVisits": max(min_visits, self.visits[fast][1] // faster_fac),
"maxVisits": max(min_visits, visits),
"priority": priority,
"rules": "japanese",
"komi": analysis_node.komi,
"boardXSize": analysis_node.board_size,
"boardYSize": analysis_node.board_size,
"analyzeTurns": [len(moves)],
}
if self.debug:
print(f"sending query for move {node_id}: {str(query)[:80]}")
self._send_analysis_query(query)
query.update({"id": f"PASS_{node_id}", "maxVisits": self.visits[fast][0] // faster_fac, "includeOwnership": False})
query["moves"] += [[analysis_node.next_player, "pass"]]
self._send_analysis_query(query)
self.queries[query_id] = (callback, time.time())
if self.katago_process:
self.katrain.log(f"Sending query {query_id}: {str(query)[:80]}", OUTPUT_DEBUG)
self.katago_process.stdin.write((json.dumps(query) + "\n").encode())
self.katago_process.stdin.flush()
+38 -82
View File
@@ -1,44 +1,41 @@
import os, random
import os
import random
from datetime import datetime
from typing import List
from game_node import GameNode
from sgf_parser import Move, SGF
from typing import List
from sgf_parser import SGF, Move
class IllegalMoveException(Exception):
pass
class KaTrainSGF(SGF):
_NODE_CLASS = GameNode
class Game:
"""Represents a game of go, including an implementation of capture rules."""
GAME_COUNTER = 0
def __init__(self, katrain, engine, analysis_options, board_options, board_size=None, move_tree=None):
def __init__(self, katrain, engine, config, board_size=None, move_tree=None):
Game.GAME_COUNTER += 1
self.katrain = katrain
self.engine = engine
self.analysis_options = analysis_options
self.board_options = board_options
self.board_size = board_size or self.board_options.get('size',19)
self.komi = self.board_options.get(f"komi_{self.board_size}",6.5)
self.config = config
self.board_size = board_size or self.config.get("size", 19)
self.komi = self.config.get(f"komi_{self.board_size}", 6.5)
self.game_id = datetime.strftime(datetime.now(), "%Y-%m-%d %H %M %S")
self.visits = [
[analysis_options["pass_visits"], analysis_options["visits"], analysis_options["analyze_all_visits"]],
[analysis_options["pass_visits_fast"], analysis_options["visits_fast"], analysis_options["analyze_all_visits_fast"]],
]
#self.train_settings = Config.get("trainer")
if move_tree:
self.root = move_tree
else:
self.root = GameNode(properties={"RU": "JP", "SZ": self.board_size, "KM": self.komi,
"PC": "KaTrain: https://github.com/sanderland/katrain",
"DT": self.game_id})
self.root = GameNode(properties={"RU": "JP", "SZ": self.board_size, "KM": self.komi, "PC": "KaTrain: https://github.com/sanderland/katrain", "DT": self.game_id})
self.current_node = self.root
self._node_by_id = {m.id: m for m in self.root.nodes_in_tree}
for node in self.root.nodes_in_tree:
node.analyze(self.engine)
self._init_chains()
# -- move tree functions --
@@ -106,13 +103,13 @@ class Game:
raise IllegalMoveException(f"Move {move} outside of board coordinates")
played_node = self.current_node.play(move)
try:
self._validate_move_and_update_chains(played_node.move, ignore_ko)
self._validate_move_and_update_chains(played_node.single_move, ignore_ko)
except IllegalMoveException:
self.current_node.children = [m for m in self.current_node.children if m != played_node]
self._init_chains() # restore
raise
self._node_by_id[played_node.id] = played_node
self.current_node = played_node
played_node.analyze(self.engine)
return played_node
def undo(self):
@@ -139,44 +136,12 @@ class Game:
if n_handicaps % 2 == 1:
stones.append((middle, middle))
stones += [(near, middle), (far, middle), (middle, near), (middle, far)]
self.root["AB"] = [Move(stone).sgf(board_size=self.board_size) for stone in stones[:n_handicaps]]
# @property
# def moves(self) -> list: # flat list of moves to current, including placements
# return sum([node.move_with_placements for node in self.current_node.nodes_from_root],[])
self.root.set_property("AB", [Move(stone).sgf(board_size=self.board_size) for stone in stones[:n_handicaps]])
@property
def next_player(self):
return self.current_node.next_player
def store_analysis(self, json):
if json["id"].startswith("AA:"): # board sweep analyze all
_, move_id, gtpcoords = json["id"].split(":")
move = self._node_by_id.get(int(move_id))
if not move.analysis:
return # should have been prevented, but better not to crash
cur_analysis = [d for d in move.analysis if d["move"] == gtpcoords]
move_analysis = {k: v for k, v in json["moveInfos"][0].items() if k not in {"move", "pv"}}
move_analysis["visits"] = sum(d["visits"] for d in json["moveInfos"]) # TODO: ??
if cur_analysis:
if cur_analysis[0]["visits"] < move_analysis["visits"]:
cur_analysis[0].update(move_analysis)
else:
move.analysis.append({"move": gtpcoords, **move_analysis})
return
if json["id"].startswith("PASS_"):
move_id = int(json["id"].lstrip("PASS_"))
is_pass = True
else:
move_id = int(json["id"])
is_pass = False
move = self._node_by_id.get(move_id)
if move: # else this should be old
move.set_analysis(json, is_pass)
else:
print("WARNING: ORPHANED ANALYSIS FOUND - RECENT NEW GAME?")
@property
def stones(self):
return sum(self.chains, [])
@@ -200,41 +165,34 @@ class Game:
return f"SGF with analysis written to {file_name}"
def ai_move(self):
ts = self.train_settings
while not self.parent.game.current_node.analysis_ready:
self.info.text = "Thinking..."
self.ai_thinking = True
time.sleep(0.05)
self.ai_thinking = False
if not self.current_node.analysis_ready:
return # TODO: hook/wait?
# select move
current_move = self.parent.game.current_node
ai_moves = self.current_node.ai_moves
pos_moves = [
(d["move"], float(d["scoreLead"]), d["evaluation"]) for i, d in enumerate(current_move.ai_moves) if i == 0 or int(d["visits"]) >= ts["balance_play_min_visits"]
]
[d["move"], d["scoreLead"], d["pointsLost"]] for i, d in enumerate(ai_moves) if i == 0 or int(d["visits"]) >= self.config["balance_play_min_visits"]
] # TODO: lcb based ?
sel_moves = pos_moves[:1]
# don't play suicidal to balance score - pass when it's best
if self.ai_balance.active and pos_moves[0][0] != "pass":
if self.katrain.controls.ai_balance.active and pos_moves[0][0] != "pass": # TODO: settings where they belong?
sel_moves = [
(move, score, move_eval)
for move, score, move_eval in pos_moves
if move_eval > ts["balance_play_randomize_eval"]
and -current_move.player_sign * score > 0
or move_eval > ts["balance_play_min_eval"]
and -current_move.player_sign * score > ts["balance_play_target_score"]
(move, score, points_lost)
for move, score, points_lost in pos_moves
if points_lost < self.config["balance_play_randomize_eval"]
or points_lost < self.config["balance_play_min_eval"]
and -self.current_node.move.player_sign * score > self.config["balance_play_target_score"]
] or sel_moves
aimove = Move.from_gtp(random.choice(sel_moves)[0], player=self.parent.game.next_player)
if len(sel_moves) > 1:
aimove.x_comment["ai"] = "AI Balance on, moves considered: " + ", ".join(f"{move} ({aimove.format_score(score)})" for move, score, _ in sel_moves) + "\n"
aimove = Move.from_gtp(random.choice(sel_moves)[0], player=self.next_player)
self.play(aimove)
def num_undos(self, move):
if self.train_settings["num_undo_prompts"] < 1:
return int(move.undo_threshold < self.train_settings["num_undo_prompts"])
if self.config["num_undo_prompts"] < 1:
return int(move.undo_threshold < self.config["num_undo_prompts"])
else:
return self.train_settings["num_undo_prompts"]
return self.config["num_undo_prompts"]
def analyze_extra(self,mode):
def analyze_extra(self, mode):
stones = {s.coords for s in self.parent.game.stones}
current_move = self.current_node
if not current_move.analysis:
@@ -248,8 +206,7 @@ class Game:
self._request_analysis(current_move, min_visits=visits, priority=self.game_counter - 1_000)
return
elif mode == "sweep":
analyze_moves = [Move(coords=(x, y)).gtp() for x in range(self.parent.game_size) for y in
range(self.parent.game_size) if (x, y) not in stones]
analyze_moves = [Move(coords=(x, y)).gtp() for x in range(self.parent.game_size) for y in range(self.parent.game_size) if (x, y) not in stones]
visits = self.visits[self.ai_fast.active][2]
self.info.text = f"Refining analysis of entire board to {visits} visits"
priority = self.game_counter - 1_000_000_000
@@ -263,10 +220,9 @@ class Game:
self._send_analysis_query(
{
"id": f"AA:{current_move.id}:{gtpcoords}",
"moves": [[m.bw_player(), m.gtp()] for m in played_moves] + [
[current_move.bw_player(True), gtpcoords]],
"moves": [[m.bw_player(), m.gtp()] for m in played_moves] + [[current_move.bw_player(True), gtpcoords]],
"includeOwnership": False,
"maxVisits": visits,
"priority": priority,
}
)
)
+27 -97
View File
@@ -1,21 +1,17 @@
import copy
import random
from typing import Dict, List, Optional
from sgf_parser import SGFNode
class GameNode(SGFNode):
_node_id_counter = -1
"""Represents a single game node, with one or more moves and placements."""
def __init__(self, parent=None, properties=None, move=None):
super().__init__(parent=parent, properties=properties, move=move)
GameNode._node_id_counter += 1
self.id = GameNode._node_id_counter
self.analysis = None
self.pass_analysis = None
self.ownership = None
self.x_comment = {}
self.auto_undid = False
self.move_number = 0
self.undo_threshold = random.random() # for fractional undos, store the random threshold in the move itself for consistency
@@ -28,134 +24,68 @@ class GameNode(SGFNode):
properties["SQ"] = best_sq
comment = self.comment(sgf=True)
if comment:
properties["C"] = properties.get("C","") + comment
properties["C"] = properties.get("C", "") + comment
return properties
def update_top_move_evaluation(self): # a move's outdated analysis
if self.analysis and self.parent and self.parent.analysis:
for move_dict in self.parent.analysis:
if move_dict["move"] == self.gtp():
move_dict["outdatedScoreLead"] = move_dict["scoreLead"]
move_dict["scoreLead"] = self.analysis[0]["scoreLead"]
self.parent.update_top_move_evaluation()
return
# various analysis functions
def set_analysis(self, analysis_blob, is_pass):
if is_pass:
self.pass_analysis = analysis_blob["moveInfos"]
else:
self.analysis = analysis_blob["moveInfos"]
self.ownership = analysis_blob["ownership"]
# if self.children: # TODO: fix when rootInfos comes in
# self.children[0].update_top_move_evaluation()
# self.update_top_move_evaluation()
def analyze(self, engine):
engine.request_analysis(self, lambda result: self.set_analysis(result))
def set_analysis(self, analysis_blob):
self.analysis = analysis_blob["moveInfos"] # TODO: fix when rootInfos comes in
self.ownership = analysis_blob["ownership"]
@property
def analysis_ready(self):
return self.analysis is not None and self.pass_analysis is not None
return self.analysis is not None
def format_score(self, score=None):
score = score or self.score
return f"{'B' if score >= 0 else 'W'}+{abs(score):.1f}"
def comment(self, sgf=False, eval=False, hints=False):
move = self.move
if not self.parent or not move: # root
single_move = self.single_move
if not self.parent or not single_move: # root
return ""
if eval and not sgf and self.children: # show undos and on previous move as well while playing
text = "".join(f"Auto undid move {m.gtp()} ({-self.temperature_stats[2] * (1-m.evaluation):.1f} pt)\n" for m in self.children if m.auto_undid)
if text:
text += "\n"
else:
text = ""
text += f"Move {self.depth}: {move.player} {move.gtp()}\n"
text += "\n".join(self.x_comment.values())
text = f"Move {self.depth}: {single_move.player} {single_move.gtp()}\n"
if self.analysis_ready:
score, _, temperature = self.temperature_stats
score = self.score
if sgf:
text += f"Score: {self.format_score(score)}\n"
if self.parent and self.parent.analysis_ready:
prev_best_score, prev_worst_score, prev_temperature = self.parent.temperature_stats
if sgf or hints:
text += f"Top move was {self.parent.analysis[0]['move']} ({self.format_score(prev_best_score)})\n"
text += f"Pass score was {self.format_score(prev_worst_score)}\n"
text += f"Previous temperature: {prev_temperature:.1f}\n"
if prev_temperature < 0.5:
text += f"Previous temperature ({prev_temperature:.1f}) too low for evaluation\n"
elif not move.is_pass and self.parent.analysis[0]["move"] != move.gtp():
if sgf: # shown in stats anyway
text += f"Evaluation: {self.evaluation:.1%} efficient\n"
outdated_evaluation, outdated_details = self.outdated_evaluation
if outdated_evaluation and outdated_evaluation > self.evaluation and outdated_evaluation > self.evaluation + 0.05:
text += f"(Was considered last move as {outdated_evaluation:.0%})\n"
points_lost = self.player_sign(self.parent.next_player) * (prev_best_score - score)
text += f"Top move was {self.parent.analysis[0]['move']} ({self.format_score(self.parent.analysis[0]['scoreLead'])})\n"
elif self.parent.analysis[0]["move"] != single_move.gtp():
points_lost = self.points_lost
if points_lost > 0.5:
text += f"Estimated point loss: {points_lost:.1f}\n"
if eval or sgf: # show undos on move itself in both sgf and while playing
undids = [m.gtp() + (f"({m.evaluation_info[0]:.1%} efficient)" if m.evaluation_info[0] else "") for m in self.parent.children if m != self]
if undids:
text += "Other attempted move(s): " + ", ".join(undids) + "\n"
else:
text = "No analysis available" if sgf else "Analyzing move..."
return text
# returns evaluation, temperature scale or None, None when not ready
@property
def evaluation_info(self):
if self.parent and self.parent.analysis_ready and self.analysis_ready:
return self.evaluation, self.parent.temperature_stats[2]
else:
return None, None
# needing own analysis ready
@property
def temperature_stats(self):
best = self.analysis[0]["scoreLead"]
worst = self.pass_analysis[0]["scoreLead"]
return best, worst, max(self.player_sign(self.next_player) * (best - worst), 0)
def points_lost(self) -> Optional[float]:
single_move = self.single_move
if single_move and self.parent and self.analysis_ready and self.parent.analysis_ready:
parent_score = self.parent.score
score = self.score
return self.player_sign(single_move.player) * (parent_score - score)
@property
def score(self):
return self.temperature_stats[0]
return self.analysis[0]["scoreLead"] # TODO: update for rootInfo
@staticmethod
def player_sign(player):
return {"B": 1, "W": -1, None: 0}[player]
# need parent analysis ready
@property
def evaluation(self):
best, worst, temp = self.parent.temperature_stats
return self.player_sign(self.parent.next_player) * (self.score - worst) / temp if temp > 0 else None
@property
def outdated_evaluation(self):
def outdated_score(move_dict):
return move_dict.get("outdatedScoreLead") or move_dict["scoreLead"]
prev_analysis_current_move = [d for d in self.parent.analysis if d["move"] == self.move.gtp()]
if prev_analysis_current_move:
best_score = outdated_score(self.parent.analysis[0])
worst_score = self.parent.pass_analysis[0]["scoreLead"]
prev_temp = max(self.player_sign(self.parent.next_player) * (best_score - worst_score), 0)
score = outdated_score(prev_analysis_current_move[0])
return (self.player_sign(self.parent.next_player) * (score - worst_score) / prev_temp if prev_temp > 0 else None), prev_analysis_current_move
else:
return None, None
@property
def ai_moves(self):
def ai_moves(self) -> List[Dict]:
if not self.analysis_ready:
return []
_, worst_score, temperature = self.temperature_stats
analysis = copy.copy(self.analysis) # not deep, so eval is saved, but avoids race conditions
for d in analysis:
if temperature > 0.5:
d["evaluation"] = self.player_sign(self.next_player) * (d["scoreLead"] - worst_score) / temperature
else:
d["evaluation"] = int(self.player_sign(self.next_player) * d["scoreLead"] >= self.player_sign(self.next_player) * self.analysis[0]["scoreLead"])
return analysis
d["pointsLost"] = self.player_sign(self.next_player) * (analysis[0]["scoreLead"] - d["scoreLead"]) # TODO: update for rootInfo
return analysis
+4
View File
@@ -0,0 +1,4 @@
from gui.badukpan import BadukPanWidget
from gui.controls import Controls
from gui.kivyutils import BWCheckBoxHint, CensorableLabel, CensorableScoreLabel, CheckBoxHint
from gui.popups import LoadSGFPopup
+72 -63
View File
@@ -1,39 +1,35 @@
import math
from kivy.graphics.context_instructions import Color
from kivy.graphics.vertex_instructions import Line, Rectangle, Ellipse
from kivy.graphics.vertex_instructions import Ellipse, Line, Rectangle
from kivy.uix.widget import Widget
from gui.controls import Config
from constants import OUTPUT_DEBUG
from gui.kivyutils import draw_circle, draw_text
from sgf_parser import Move
STONE_COLORS = Config.get("ui")["stones"]
OUTLINE_COLORS = Config.get("ui").get("outline", [None, None])
GHOST_ALPHA = Config.get("ui")["ghost_alpha"]
class BadukPanWidget(Widget):
def __init__(self, **kwargs):
super(BadukPanWidget, self).__init__(**kwargs)
self.config = {}
self.ghost_stone = []
self.gridpos = []
self.grid_size = 0
self.stone_size = 0
self.last_eval = 0
self.EVAL_COLORS = Config.get("ui")["eval_colors"]
self.EVAL_KNOTS = Config.get("ui")["eval_knots"]
self.EVAL_BOUNDS = Config.get("ui")["eval_bounds"]
# stone placement functions
def _find_closest(self, pos):
return sorted([(abs(p - pos), i) for i, p in enumerate(self.gridpos)])[0]
def on_touch_down(self, touch):
if not self.gridpos:
return
xd, xp = self._find_closest(touch.x)
yd, yp = self._find_closest(touch.y)
prev_ghost = self.ghost_stone
if self.engine.ready and max(yd, xd) < self.grid_size / 2 and (xp, yp) not in [m.coords for m in self.engine.board.stones]:
if max(yd, xd) < self.grid_size / 2 and (xp, yp) not in [m.coords for m in self.parent.game.stones]:
self.ghost_stone = (xp, yp)
else:
self.ghost_stone = None
@@ -44,22 +40,22 @@ class BadukPanWidget(Widget):
return self.on_touch_down(touch)
def on_touch_up(self, touch):
if not self.gridpos:
return
katrain = self.parent
if self.ghost_stone:
self.engine.action("play", self.ghost_stone)
katrain("play", self.ghost_stone)
else:
xd, xp = self._find_closest(touch.x)
yd, yp = self._find_closest(touch.y)
nodes_here = [node for node in self.engine.board.current_node.nodes_from_root if node.move and node.move.coords==(xp,yp)]
nodes_here = [node for node in katrain.game.current_node.nodes_from_root if node.single_move and node.single_move.coords == (xp, yp)]
if nodes_here and max(yd, xd) < self.grid_size / 2: # load old comment
if self.engine.debug:
print("\nAnalysis:\n", nodes_here[-1].analysis)
print("\nParent Analysis:\n", nodes_here[-1].parent.analysis)
if nodes_here[-1].parent.pass_analysis:
print("\nParent Pass Analysis:\n", nodes_here[-1].parent.pass_analysis[0])
if not self.engine.ai_lock.active:
self.engine.info.text = nodes_here[-1].comment(sgf=True)
self.engine.show_evaluation_stats(nodes_here[-1])
katrain.log(f"\nAnalysis:\n{nodes_here[-1].analysis}", OUTPUT_DEBUG)
katrain.log(f"\nParent Analysis:\n{nodes_here[-1].parent.analysis}", OUTPUT_DEBUG)
if not katrain.ai_lock.active:
katrain.info.text = nodes_here[-1].comment(sgf=True)
katrain.show_evaluation_stats(nodes_here[-1])
self.ghost_stone = None
self.draw_board_contents() # remove ghost
@@ -83,96 +79,109 @@ class BadukPanWidget(Widget):
Color(*innercol)
Line(circle=(self.gridpos[x], self.gridpos[y], stone_size * 0.45 / 0.85), width=0.125 * stone_size) # 1.75
def _eval_spectrum(self, score):
score = max(0, score)
for i in range(len(self.EVAL_KNOTS) - 1):
if self.EVAL_KNOTS[i] <= score < self.EVAL_KNOTS[i + 1]:
t = (score - self.EVAL_KNOTS[i]) / (self.EVAL_KNOTS[i + 1] - self.EVAL_KNOTS[i])
return [a + t * (b - a) for a, b in zip(self.EVAL_COLORS[i], self.EVAL_COLORS[i + 1])]
return self.EVAL_COLORS[-1]
def _eval_spectrum(self, points_lost):
EVAL_COLORS = self.config["eval_colors"]
EVAL_THRESHOLDS = self.config["eval_thresholds"]
i = 0
while points_lost < EVAL_THRESHOLDS[i] and i < len(EVAL_COLORS):
i += 1
return EVAL_COLORS[i]
def draw_board(self, *args):
if not self.config:
return
katrain = self.parent
board_size = katrain.game.board_size
self.canvas.before.clear()
with self.canvas.before:
# board
sz = min(self.width, self.height)
Color(*Config.get("ui")["board_color"])
board = Rectangle(pos=(0, 0), size=(sz, sz))
Color(*self.config["board_color"])
board_rectangle = Rectangle(pos=(0, 0), size=(sz, sz))
# grid lines
margin = Config.get("ui")["board_margin"]
self.grid_size = board.size[0] / (self.engine.board_size - 1 + 1.5 * margin)
self.stone_size = self.grid_size * Config.get("ui")["stone_size"]
self.gridpos = [math.floor((margin + i) * self.grid_size + 0.5) for i in range(self.engine.board_size)]
margin = self.config["board_margin"]
self.grid_size = board_rectangle.size[0] / (board_size - 1 + 1.5 * margin)
self.stone_size = self.grid_size * self.config["stone_size"]
self.gridpos = [math.floor((margin + i) * self.grid_size + 0.5) for i in range(board_size)]
line_color = Config.get("ui")["line_color"]
line_color = self.config["line_color"]
Color(*line_color)
lo, hi = self.gridpos[0], self.gridpos[-1]
for i in range(self.engine.board_size):
for i in range(board_size):
Line(points=[(self.gridpos[i], lo), (self.gridpos[i], hi)])
Line(points=[(lo, self.gridpos[i]), (hi, self.gridpos[i])])
# star points
star_point_pos = 3 if self.engine.board_size <= 11 else 4
starpt_size = self.grid_size * Config.get("ui")["starpoint_size"]
for x in [star_point_pos - 1, self.engine.board_size - star_point_pos, int(self.engine.board_size / 2)]:
for y in [star_point_pos - 1, self.engine.board_size - star_point_pos, int(self.engine.board_size / 2)]:
star_point_pos = 3 if board_size <= 11 else 4
starpt_size = self.grid_size * self.config["starpoint_size"]
for x in [star_point_pos - 1, board_size - star_point_pos, int(board_size / 2)]:
for y in [star_point_pos - 1, board_size - star_point_pos, int(board_size / 2)]:
draw_circle((self.gridpos[x], self.gridpos[y]), starpt_size, line_color)
# coordinates
Color(0.25, 0.25, 0.25)
for i in range(self.engine.board_size):
for i in range(board_size):
draw_text(pos=(self.gridpos[i], lo / 2), text=Move.GTP_COORD[i], font_size=self.grid_size / 1.5)
draw_text(pos=(lo / 2, self.gridpos[i]), text=str(i + 1), font_size=self.grid_size / 1.5)
def draw_board_contents(self, *args):
if not self.config:
return
stone_color = self.config["stones"]
outline_color = self.config["outline"]
ghost_alpha = self.config["ghost_alpha"]
katrain = self.parent
board_size = katrain.game.board_size
self.canvas.clear()
with self.canvas:
# stones
current_node = self.engine.board.current_node
next_player = self.engine.board.next_player
full_eval_on = {p: self.engine.eval.active(p) for p in Move.PLAYERS} # TODO: map?
current_node = katrain.game.current_node
next_player = katrain.game.next_player
full_eval_on = {p: katrain.controls.eval.active(p) for p in Move.PLAYERS} # TODO: map? TODO: settings here
has_stone = {}
for m in self.engine.board.stones:
for m in katrain.game.stones:
has_stone[m.coords] = m.player
show_n_eval = Config.get("trainer")["eval_off_show_last"]
nodes = self.engine.board.current_node.nodes_from_root
show_n_eval = self.config["eval_off_show_last"]
nodes = katrain.game.current_node.nodes_from_root
for i, node in enumerate(nodes):
eval, evalsize = node.evaluation_info
eval = node.points_lost
evalsize = 1
for m in node.move_with_placements:
if has_stone[m.coords]: # skip captures, draw over repeat plays
move_eval_on = full_eval_on[m.player] or i >= len(nodes) - show_n_eval
evalcol = self._eval_spectrum(eval) if move_eval_on and eval and evalsize > Config.get("ui").get("min_eval_temperature", 0.5) else None
inner = STONE_COLORS[m.opponent] if (m is current_node) else None
self.draw_stone(m.coords[0], m.coords[1], STONE_COLORS[m.player], OUTLINE_COLORS[m.player], inner, evalcol, evalsize)
evalcol = self._eval_spectrum(eval) if move_eval_on and eval and evalsize > self.config.get("min_eval_temperature", 0.5) else None
inner = stone_color[m.opponent] if (m is current_node) else None
self.draw_stone(m.coords[0], m.coords[1], stone_color[m.player], outline_color[m.player], inner, evalcol, evalsize)
# ownership - allow one move out of date for smooth animation
ownership = current_node.ownership or (current_node.parent and current_node.parent.ownership)
if self.engine.ownership.active and ownership:
if katrain.controls.ownership.active and ownership:
rsz = self.grid_size * 0.2
ix = 0
for y in range(self.engine.board_size - 1, -1, -1):
for x in range(self.engine.board_size):
for y in range(board_size - 1, -1, -1):
for x in range(board_size):
ix_owner = "B" if ownership[ix] > 0 else "W"
if ix_owner != (has_stone.get((x, y), -1)):
Color(*STONE_COLORS[ix_owner], abs(ownership[ix]))
Color(*stone_color[ix_owner], abs(ownership[ix]))
Rectangle(pos=(self.gridpos[x] - rsz / 2, self.gridpos[y] - rsz / 2), size=(rsz, rsz))
ix = ix + 1
# children of current moves in undo / review
undo_coords = set()
alpha = Config.get("ui")["undo_alpha"]
alpha = self.config["undo_alpha"]
for child_node in current_node.children:
eval_info = child_node.evaluation_info
m = child_node.move
m = child_node.single_move
if m and m.coords[0] is not None:
undo_coords.add(m.coords)
evalcol = (*self._eval_spectrum(eval_info[0]), alpha) if eval_info[0] else None
scale = Config.get("ui").get("undo_scale", 0.95)
self.draw_stone(m.coords[0], m.coords[1], (*STONE_COLORS[m.player][:3], alpha), None, None, evalcol, self.EVAL_BOUNDS[1], scale=scale)
scale = self.config.get("undo_scale", 0.95)
self.draw_stone(m.coords[0], m.coords[1], (*stone_color[m.player][:3], alpha), None, None, evalcol, self.EVAL_BOUNDS[1], scale=scale)
# hints
if self.engine.hints.active(next_player):
if katrain.controls.hints.active(next_player):
hint_moves = current_node.ai_moves
for i, d in enumerate(hint_moves):
move = Move.from_gtp(d["move"])
@@ -188,17 +197,17 @@ class BadukPanWidget(Widget):
# hover next move ghost stone
if self.ghost_stone:
self.draw_stone(*self.ghost_stone, (*STONE_COLORS[next_player], GHOST_ALPHA))
self.draw_stone(*self.ghost_stone, (*stone_color[next_player], ghost_alpha))
# pass circle
passed = len(nodes) > 1 and current_node.is_pass
if passed:
if self.engine.board.game_ended:
if katrain.game.game_ended:
text = "game\nend"
else:
text = "pass"
Color(0.45, 0.05, 0.45, 0.5)
center = self.gridpos[int(self.engine.board_size / 2)]
center = self.gridpos[int(board_size / 2)]
Ellipse(pos=(center - self.grid_size * 1.5, center - self.grid_size * 1.5), size=(self.grid_size * 3, self.grid_size * 3))
Color(0.15, 0.15, 0.15)
draw_text(pos=(center, center), text=text, font_size=self.grid_size * 0.66, halign="center", outline_color=[0.95, 0.95, 0.95])
draw_text(pos=(center, center), text=text, font_size=self.grid_size * 0.66, halign="center", outline_color=[0.95, 0.95, 0.95])
+1 -3
View File
@@ -14,8 +14,6 @@ from kivy.uix.gridlayout import GridLayout
from kivy.uix.label import Label
class Controls(GridLayout):
def __init__(self, **kwargs):
super(Controls, self).__init__(**kwargs)
@@ -36,7 +34,7 @@ class Controls(GridLayout):
# handles showing completed analysis and triggered actions like auto undo and ai move
def update_evaluation(self):
current_node = self.parent.game.current_node
move = current_node.move
move = current_node.single_move
self.score.set_prisoners(self.parent.game.prisoner_count)
current_player_is_human_or_both_robots = True # move not self.ai_auto.active(current_node.player) or self.ai_auto.active(1 - current_node.player) # TODO FIX
+1
View File
@@ -1,4 +1,5 @@
from kivy.uix.boxlayout import BoxLayout
class LoadSGFPopup(BoxLayout):
pass
+14 -14
View File
@@ -160,23 +160,23 @@
path: os.path.expanduser("~")
BoxLayout:
orientation: 'horizontal'
Label:
text: "Analyze Extra Fast"
Checkbox:
id: fast
color: (0.95, 0.95, 0.95)
Label:
text: "Rewind to Start"
Checkbox:
id: rewind
active: True
color: (0.95, 0.95, 0.95)
Label:
text: "Analyze Extra Fast"
Checkbox:
id: fast
color: (0.95, 0.95, 0.95)
Label:
text: "Rewind to Start"
Checkbox:
id: rewind
active: True
color: (0.95, 0.95, 0.95)
<BadukPanWidget>:
size: self.parent.height, self.parent.height
engine: self.parent.controls
<EngineControls>
<Controls>
cols: 1
rows: 9
info: info
@@ -329,7 +329,7 @@
on_press: root.restart(19)
<KaTrainGui>:
board: board
board_gui: board_gui
controls: controls
canvas.before:
Color:
@@ -338,7 +338,7 @@
pos: self.pos
size: root.size
BadukPanWidget:
id: board
id: board_gui
size_hint: 1 - controls.size_hint[0], 1
Controls:
id: controls
+53 -55
View File
@@ -1,29 +1,30 @@
import os
import signal
from kivy.app import App
from kivy.core.window import Window
from kivy.uix.boxlayout import BoxLayout
from .engine import KataGoEngine
import os, sys, threading
from kivy.storage.jsonstore import JsonStore
from kivy.clock import Clock
import sys
import threading
from queue import Queue
from game import Game, IllegalMoveException, KaTrainSGF, Move
from game_node import GameNode
from kivy.uix.popup import Popup
from gui.popups import LoadSGFPopup
OUTPUT_ERROR = -1
OUTPUT_INFO = 0
OUTPUT_DEBUG = 1
OUTPUT_EXTRA_DEBUG = 2
from kivy.app import App
from kivy.clock import Clock
from kivy.core.window import Window
from kivy.storage.jsonstore import JsonStore
from kivy.uix.boxlayout import BoxLayout
from kivy.uix.popup import Popup
from constants import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG, OUTPUT_INFO
from engine import KataGoEngine
from game import Game, GameNode, IllegalMoveException, KaTrainSGF, Move
from gui import BadukPanWidget, BWCheckBoxHint, CensorableLabel, CensorableScoreLabel, CheckBoxHint, Controls, LoadSGFPopup
class KaTrainGui(BoxLayout):
"""Top level class responsible for tying everything together"""
def __init__(self, **kwargs):
super(KaTrainGui, self).__init__(**kwargs)
self.debug_level = 0
self._load_config()
self.debug_level = self.config("debug/level",OUTPUT_INFO)
self.debug_level = self.config("debug/level", OUTPUT_INFO)
self.logger = lambda message, level=OUTPUT_INFO: self.log(message, level)
self.engine = None
@@ -34,7 +35,7 @@ class KaTrainGui(BoxLayout):
self._keyboard.bind(on_key_down=self._on_keyboard_down)
def log(self, message, level=OUTPUT_INFO):
if level==OUTPUT_ERROR:
if level == OUTPUT_ERROR:
self.controls.set_status(f"ERROR: {message}")
print(f"ERROR: {message}")
elif self.debug_level >= level:
@@ -44,25 +45,28 @@ class KaTrainGui(BoxLayout):
base_path = getattr(sys, "_MEIPASS", os.path.dirname(os.path.abspath(__file__))) # for pyinstaller
config_file = sys.argv[1] if len(sys.argv) > 1 else os.path.join(base_path, "config.json")
try:
self.log(f"Using config file {config_file}",OUTPUT_INFO)
self.log(f"Using config file {config_file}", OUTPUT_INFO)
self._config_store = JsonStore(config_file)
except FileNotFoundError:
self.log(f"Config file {config_file} not found",OUTPUT_ERROR)
self.log(f"Config file {config_file} not found", OUTPUT_ERROR)
except Exception as e:
self.log(f"Failed to load config {config_file}: {e}", OUTPUT_ERROR)
def config(self,setting,default=None):
def config(self, setting, default=None):
try:
if '/' in setting:
cat, key = setting.split('/')
return self._config_store.get(cat).get(key,default)
if "/" in setting:
cat, key = setting.split("/")
return self._config_store.get(cat).get(key, default)
else:
return self._config_store.get(setting)
except Exception:
self.log(f"Missing configuration option {setting}",OUTPUT_ERROR)
self.log(f"Missing configuration option {setting}", OUTPUT_ERROR)
def start(self):
if self.engine:
return
self.engine = KataGoEngine(self.config("engine"),self.logger)
self.board_gui.config = self.config("board_ui")
self.engine = KataGoEngine(self, self.config("engine"))
threading.Thread(target=self._message_loop_thread, daemon=True).start()
self._do_new_game()
@@ -70,21 +74,29 @@ class KaTrainGui(BoxLayout):
while True:
game, msg, *args = self.message_queue.get()
try:
self.log(f"Message Loop Received {msg}: {args} for Game {game}",OUTPUT_EXTRA_DEBUG)
self.log(f"Message Loop Received {msg}: {args} for Game {game}", OUTPUT_EXTRA_DEBUG)
if game != self.game.game_id:
self.log(f"Message skipped as it is outdated (current game is {self.game.game_id}", OUTPUT_EXTRA_DEBUG)
continue
getattr(self, f"_do_{msg.replace('-','_')}")(*args)
except Exception as e:
self.log(f"Exception in Engine thread: {e}",OUTPUT_ERROR)
self.log(f"Exception in Engine thread: {e}", OUTPUT_ERROR)
raise
def __call__(self, message, *args):
if self.game:
self.message_queue.put([self.game.game_id, message, *args])
def _do_new_game(self,board_size=None):
self.game = Game(self,self.engine,self.config("analysis"),self.config("board"),board_size=board_size)
def _do_new_game(self, board_size=None):
self.game = Game(self, self.engine, self.config("game"), board_size=board_size)
# TODO controls reset -> Controls
# if self.ai_lock.active:
# self.ai_lock.checkbox._do_press()
# for el in [self.ai_lock.checkbox, self.hints.black, self.hints.white, self.ai_auto.black, self.ai_auto.white, self.auto_undo.black, self.auto_undo.white, self.ai_move]:
# el.disabled = False
self.redraw(include_board=True)
def _do_aimove(self):
self.game.ai_move()
@@ -109,22 +121,9 @@ class KaTrainGui(BoxLayout):
self.game.switch_branch(direction)
self.update_evaluation()
def _do_init(self, board_size=None, komi=None, move_tree=None):
self.game_counter += 1 # prioritize newer games
self.game_size = board_size or 19
self.komi = float(komi or self.config.get("board").get(f"komi_{board_size}", 6.5))
self.game = Game(board_size, move_tree)
self._request_analysis(self.game.root, priority=self.game_counter)
self.redraw(include_board=True)
self.ready = True
if self.ai_lock.active:
self.ai_lock.checkbox._do_press()
for el in [self.ai_lock.checkbox, self.hints.black, self.hints.white, self.ai_auto.black, self.ai_auto.white, self.auto_undo.black, self.auto_undo.white, self.ai_move]:
el.disabled = False
def play(self, move: Move, faster=False, analysis_priority=None):
try:
next_node = self.board.play(move)
next_node = self.board_gui.play(move)
except IllegalMoveException as e:
self.info.text = f"Illegal move: {str(e)}"
return
@@ -134,12 +133,12 @@ class KaTrainGui(BoxLayout):
return next_node
def _do_play(self, *args):
self.play(Move(args[0], player=self.board.next_player))
self.game.play(Move(args[0], player=self.game.next_player))
self.redraw()
def _do_analyze_extra(self, mode):
self.game.analyze_extra(mode)
def analyze_movetree(self, root, faster=False):
self._do_init(root["SZ"], root["KM"])
self.parent.game.root = root
@@ -164,16 +163,14 @@ class KaTrainGui(BoxLayout):
def output_sgf(self):
for pl in Move.PLAYERS:
if self.parent.game.root[f"P{pl}"] not in ["KaTrain","Player",None,""]:
if self.parent.game.root[f"P{pl}"] not in ["KaTrain", "Player", None, ""]:
self.parent.game.root[f"P{pl}"] = "KaTrain" if self.ai_auto.active(pl) else "Player"
return self.parent.game.write_sgf(self.komi)
def redraw(self, include_board=False):
if include_board:
Clock.schedule_once(self.board.draw_board, -1) # main thread needs to do this
Clock.schedule_once(self.board.draw_board_contents, -1)
Clock.schedule_once(self.board_gui.draw_board, -1) # main thread needs to do this
Clock.schedule_once(self.board_gui.draw_board_contents, -1)
def _on_keyboard_down(self, keyboard, keycode, text, modifiers):
if keycode[1] == "up":
@@ -207,7 +204,7 @@ class KaTrainGui(BoxLayout):
self.controls.ai_balance.label.trigger_action(duration=0)
elif keycode[1] == "o":
self.controls.ownership.label.trigger_action(duration=0)
elif keycode[1] == "l": # ctrl-l?
elif keycode[1] == "l": # ctrl-l?
self.controls.load.trigger_action(duration=0)
elif keycode[1] == "k":
self.controls.save.trigger_action(duration=0)
@@ -221,22 +218,23 @@ class KaTrainApp(App):
return self.gui
def on_start(self):
self.gui.restart()
self.gui.start()
signal.signal(signal.SIGINT, self.signal_handler)
def signal_handler(self, signal, frame):
import sys
import traceback
if self.gui.controls.debug:
if self.gui.debug_level >= OUTPUT_DEBUG:
print("TRACEBACKS")
for threadId, stack in sys._current_frames().items():
print(f"\n# ThreadID: {threadId}")
for filename, lineno, name, line in traceback.extract_stack(stack):
print("File: filename}, line {lineno}, in {name}")
print(f"\tFile: {filename}, line {lineno}, in {name}")
if line:
print(f" {line.strip()}")
print(f"\t\t{line.strip()}")
sys.exit(0)
if __name__ == "__main__":
KaTrainApp().run()
+49 -40
View File
@@ -1,5 +1,6 @@
import re
import copy
import re
from collections import defaultdict
from typing import Any, Dict, List, Optional, Tuple
@@ -22,7 +23,7 @@ class Move:
def from_sgf(sgf_coords, board_size, player="B"):
if sgf_coords == "" or Move.SGF_COORD.index(sgf_coords[0]) == board_size: # some servers use [tt] for pass
return Move(coords=None, player=player)
return Move(coords=(Move.SGF_COORD.index(sgf_coords[0]), board_size - Move.SGF_COORD.index(sgf_coords[1]) - 1), player=player)
return Move(coords=(Move.SGF_COORD.index(sgf_coords[0]), board_size - Move.SGF_COORD.index(sgf_coords[1]) - 1), player=player,)
def __init__(self, coords: Optional[Tuple[int, int]] = None, player: str = "B"):
self.player = player
@@ -31,10 +32,8 @@ class Move:
def __repr__(self):
return f"Move({self.player}{self.gtp()})"
# def __hash__(self):
# return self.__repr__().__hash__()
# def __eq__(self, other):
# return self.coords == other.coords and self.player == other.player
def __eq__(self, other):
return self.coords == other.coords and self.player == other.player
def gtp(self):
if self.is_pass:
@@ -56,26 +55,25 @@ class Move:
class SGFNode:
CAST_FIELDS = {"KM": float, "SZ": int, "HA": int} # cast property to this type
LIST_FIELDS = ["AB", "AW", "TW", "TB", "MA", "SQ", "CR", "TR", "LN", "AR", "LB"] # cast these properties to lists
# TODO: all are potential lists? what a headache!
def __init__(self, parent=None, properties=None, move=None):
self.children = []
self.properties = copy.copy(properties) if properties is not None else {}
self.properties = defaultdict(list)
if properties:
for k, v in properties.items():
self.add_property(k, v)
self.parent = parent
if self.parent:
self.parent.children.append(self)
if parent and move:
self.properties[move.player] = move.sgf(self.board_size)
self.add_property(move.player, move.sgf(self.board_size))
@property
def sgf_properties(self) -> Dict:
"""For hooking into in a subclass and overriding/formatting any additional properties to be output"""
return self.properties
return copy.deepcopy(self.properties)
def sgf(self) -> str:
sgf_str = "".join([f"{k}[{']['.join(v) if isinstance(v,list) else v}]" for k, v in self.sgf_properties.items()])
sgf_str = "".join([prop + "".join(f"[{v}]" for v in values) for prop, values in self.sgf_properties.items() if values])
if self.children:
children = [c.sgf() for c in self.children]
if len(children) == 1:
@@ -84,20 +82,20 @@ class SGFNode:
sgf_str += "(;" + ")(;".join(children) + ")"
return f"(;{sgf_str})" if self.is_root else sgf_str
def __setitem__(self, prop: str, value: Any):
if prop in self.LIST_FIELDS and isinstance(value, str): # lists (placements, IGS marked dead stones)
self.properties[prop] = self.properties.get(prop, []) + re.split(r"\]\s*\[", value)
elif prop in self.CAST_FIELDS:
self.properties[prop] = self.CAST_FIELDS[prop](value)
else:
self.properties[prop] = value
def add_property(self, property: str, values: Any):
"""Add some values to the property. If not a list, it will be made into a single-value list."""
if not isinstance(values, list):
values = [values]
self.properties[property] += values
def __getitem__(self, property) -> Any:
return self.properties.get(property)
def get(self, property, default) -> Any:
def get(self, property, default=None) -> Any:
"""Get the list of values for a property."""
return self.properties.get(property, default)
def get_first(self, property, default) -> Any:
"""Get the first value of the property, typically when exactly one is expected."""
return self.properties.get(property, [default])[0]
@property
def parent(self) -> Optional["SGFNode"]:
return self._parent
@@ -125,22 +123,33 @@ class SGFNode:
@property
def board_size(self) -> int:
return self.root.get("SZ", 19)
return int(self.root.get_first("SZ", 19))
@property
def move(self) -> Optional[Move]:
for pl in Move.PLAYERS:
if self[pl]:
return Move.from_sgf(self[pl], player=pl, board_size=self.board_size)
def komi(self) -> float:
return float(self.root.get_first("KM", 6.5))
@property
def moves(self) -> List[Move]:
"""Returns all moves in the node."""
return [Move.from_sgf(move, player=pl, board_size=self.board_size) for pl in Move.PLAYERS for move in self.get(pl, [])]
@property
def placements(self) -> List[Move]:
return [Move.from_sgf(self[pl], player=pl, board_size=self.board_size) for pl in Move.PLAYERS for sgf in self.get("A" + pl, [])]
"""Returns all placements (AB/AW) in the node."""
return [Move.from_sgf(sgf_coords, player=pl, board_size=self.board_size) for pl in Move.PLAYERS for sgf_coords in self.get("A" + pl, [])]
@property
def move_with_placements(self) -> List[Move]:
move = self.move
return self.placements + ([move] if move else [])
"""Returns all moves (B/W) and placements (AB/AW) in the node."""
return self.placements + self.moves
@property
def single_move(self) -> Optional[Move]:
"""Returns the single move for the node if one exists, or None if no moves (or multiple ones) exist."""
moves = self.moves
if len(moves) == 1:
return moves[0]
@property
def is_root(self) -> bool:
@@ -148,7 +157,7 @@ class SGFNode:
@property
def is_pass(self) -> bool:
return not self.placements and self.move and self.move.is_pass
return not self.placements and self.single_move and self.single_move.is_pass
@property
def empty(self) -> bool:
@@ -165,14 +174,13 @@ class SGFNode:
def play(self, move) -> "SGFNode":
"""Either find an existing child or create a new one with the given move."""
for c in self.children:
if c.move == move:
return c.move
if c.single_move == move:
return c
return self.__class__(parent=self, move=move)
@property
def next_player(self):
m = self.move
if m and m.player == "B" or "AB" in self.properties:
if self.get("B") or self.get("AB"):
return "W"
return "B"
@@ -220,8 +228,9 @@ class SGF:
if not current_move.empty: # ignore ; that generate empty nodes
current_move = self._NODE_CLASS(parent=current_move)
else:
prop, value = match[1], match[2].strip()[1:-1]
current_move[prop] = value
property, value = match[1], match[2].strip()[1:-1]
values = re.split(r"\]\s*\[", value)
current_move.add_property(property, values)
if self.ix < len(self.contents):
raise ParseError(f"Parse Error: unexpected character at {self.contents[self.ix - 25:self.ix]}>{self.contents[self.ix]}<{self.contents[self.ix + 1:self.ix + 25]}")
raise ParseError("Parse Error: expected ')' at end of input.")