node / move split
This commit is contained in:
1 parent
dca6ad97e8
commit
be8b8f8fa2
4 files changed
+226
-205
No files matched your search
@@ -2,26 +2,21 @@ import os
|
||||
import random
|
||||
from datetime import datetime
|
||||
import copy
|
||||
import sgfparser
|
||||
from sgfparser import Move, SGFNode, SGF
|
||||
from typing import List
|
||||
|
||||
class IllegalMoveException(Exception):
|
||||
pass
|
||||
|
||||
# TODO: split sgf node vs move?
|
||||
|
||||
class Move(sgfparser.Move):
|
||||
GTP_COORD = "ABCDEFGHJKLMNOPQRSTUVWYXYZ"
|
||||
PLAYERS = "BW"
|
||||
SGF_COORD = [chr(i) for i in range(97, 123)]
|
||||
_move_id_counter = -1
|
||||
class KaTrainSGFNode(SGFNode):
|
||||
_node_id_counter = -1
|
||||
|
||||
def __init__(self,parent=None, properties=None, move=None):
|
||||
super().__init__(parent=parent,properties=properties,move=move)
|
||||
KaTrainSGFNode._node_id_counter += 1
|
||||
self.id = KaTrainSGFNode._node_id_counter
|
||||
|
||||
def __init__(self, player=0, coords=None, gtpcoords=None, sgfcoords=None, robot=False):
|
||||
super().__init__()
|
||||
Move._move_id_counter += 1
|
||||
self.id = Move._move_id_counter
|
||||
self.player = player
|
||||
self.coords = coords or (gtpcoords and self.gtp2ix(gtpcoords)) or self.sgf2ix(sgfcoords)
|
||||
self.robot = robot
|
||||
self.analysis = None
|
||||
self.pass_analysis = None
|
||||
self.ownership = None
|
||||
@@ -30,31 +25,14 @@ class Move(sgfparser.Move):
|
||||
self.move_number = 0
|
||||
self.undo_threshold = random.random() # for fractional undos, store the random threshold in the move itself for consistency
|
||||
|
||||
def __repr__(self):
|
||||
return f"{Move.PLAYERS[self.player]}{self.gtp()}"
|
||||
|
||||
def __eq__(self, other):
|
||||
return self.coords == other.coords and self.player == other.player
|
||||
|
||||
def __hash__(self):
|
||||
return self.gtp().__hash__()
|
||||
|
||||
def play(self, move):
|
||||
try:
|
||||
return self.children[self.children.index(move)]
|
||||
except ValueError:
|
||||
move.parent = self
|
||||
move.move_number = self.move_number + 1
|
||||
self.children.append(move)
|
||||
return move
|
||||
|
||||
@property
|
||||
def moves_in_tree(self):
|
||||
return [self] + sum([c.moves_in_tree for c in self.children],[])
|
||||
|
||||
@property
|
||||
def is_pass(self):
|
||||
return self.coords[0] is None
|
||||
def sgf_properties(self):
|
||||
best_sq = []
|
||||
properties = copy.copy(super().sgf_properties)
|
||||
if 'SQ' not in properties:
|
||||
properties['SQ'] = best_sq
|
||||
properties['C'] += self.comment(sgf=True)
|
||||
return properties
|
||||
|
||||
def update_top_move_evaluation(self): # a move's outdated analysis
|
||||
if self.analysis and self.parent and self.parent.analysis:
|
||||
@@ -85,7 +63,8 @@ class Move(sgfparser.Move):
|
||||
return f"{'B' if score >= 0 else 'W'}+{abs(score):.1f}"
|
||||
|
||||
def comment(self, sgf=False, eval=False, hints=False):
|
||||
if not self.parent: # root
|
||||
move = self.move
|
||||
if not self.parent or not move: # root
|
||||
return ""
|
||||
|
||||
if eval and not sgf and self.children: # show undos and on previous move as well while playing
|
||||
@@ -95,7 +74,7 @@ class Move(sgfparser.Move):
|
||||
else:
|
||||
text = ""
|
||||
|
||||
text += f"Move {self.move_number}: {self.bw_player()} {self.gtp()} {'(AI Move)' if self.robot else ''}\n"
|
||||
text += f"Move {move.player} {move.gtp()}\n"
|
||||
text += "\n".join(self.x_comment.values())
|
||||
|
||||
if self.analysis_ready:
|
||||
@@ -110,7 +89,7 @@ class Move(sgfparser.Move):
|
||||
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 self.is_pass and self.parent.analysis[0]["move"] != self.gtp():
|
||||
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
|
||||
@@ -184,37 +163,6 @@ class Move(sgfparser.Move):
|
||||
d["evaluation"] = int(-self.player_sign * d["scoreLead"] >= -self.player_sign * self.analysis[0]["scoreLead"])
|
||||
return analysis
|
||||
|
||||
# various output and conversion functions
|
||||
@staticmethod
|
||||
def gtp2ix(gtpmove):
|
||||
if "pass" in gtpmove:
|
||||
return None, None
|
||||
return Move.GTP_COORD.index(gtpmove[0]), int(gtpmove[1:]) - 1
|
||||
|
||||
@staticmethod
|
||||
def sgf2ix(sgfmove_with_board_size):
|
||||
sgfmove, board_size = sgfmove_with_board_size
|
||||
if sgfmove == "" or Move.SGF_COORD.index(sgfmove[0]) == board_size: # some servers use [tt] for pass
|
||||
return None, None
|
||||
return Move.SGF_COORD.index(sgfmove[0]), board_size - Move.SGF_COORD.index(sgfmove[1]) - 1
|
||||
|
||||
def gtp(self):
|
||||
if self.is_pass:
|
||||
return "pass"
|
||||
return Move.GTP_COORD[self.coords[0]] + str(self.coords[1] + 1)
|
||||
|
||||
def sgfcoords(self, board_size):
|
||||
return f"{Move.SGF_COORD[self.coords[0]]}{Move.SGF_COORD[board_size - self.coords[1] - 1]}"
|
||||
|
||||
def bw_player(self, next_move=False):
|
||||
return Move.PLAYERS[1 - self.player if next_move else self.player]
|
||||
|
||||
def sgf(self, board_size):
|
||||
if self.is_pass:
|
||||
return f"{self.bw_player()}[]"
|
||||
else:
|
||||
return f"{self.bw_player()}[{self.sgfcoords(board_size)}]"
|
||||
|
||||
|
||||
class Board:
|
||||
def __init__(self, board_size=19, move_tree=None):
|
||||
@@ -223,24 +171,24 @@ class Board:
|
||||
if move_tree:
|
||||
self.root = move_tree
|
||||
else:
|
||||
self.root = Move(1, (None, None)) # root is 1=white so black is first
|
||||
self.current_move = self.root
|
||||
self.all_moves = {m.id:m for m in self.root.moves_in_tree}
|
||||
self.root = KaTrainSGFNode(properties={'RU':'JP','SZ':board_size}) # TODO: Komi, etc?
|
||||
self.current_node = self.root
|
||||
self._node_by_id = {m.id:m for m in self.root.nodes_in_tree}
|
||||
self._init_chains()
|
||||
|
||||
# -- move tree functions --
|
||||
def _init_chains(self):
|
||||
self.board = [[-1 for _x in range(self.board_size)] for _y in range(self.board_size)] # board pos -> chain id
|
||||
self.chains = [] # chain id -> chain
|
||||
self.prisoners = []
|
||||
self.last_capture = []
|
||||
self.board = [[-1 for _x in range(self.board_size)] for _y in range(self.board_size)] # type: List[list[int]] # board pos -> chain id
|
||||
self.chains = [] # type: List[List[Move]] # chain id -> chain
|
||||
self.prisoners = [] # type: List[Move]
|
||||
self.last_capture = [] # type: List[Move]
|
||||
try:
|
||||
for m in self.moves:
|
||||
self._validate_move_and_update_chains(m, True) # ignore ko since we didn't know if it was forced
|
||||
except IllegalMoveException as e:
|
||||
raise Exception(f"Unexpected illegal move ({str(e)})")
|
||||
|
||||
def _validate_move_and_update_chains(self, move, ignore_ko):
|
||||
def _validate_move_and_update_chains(self, move: Move, ignore_ko: bool):
|
||||
def neighbours(moves):
|
||||
return {
|
||||
self.board[m.coords[1] + dy][m.coords[0] + dx]
|
||||
@@ -286,34 +234,34 @@ class Board:
|
||||
raise IllegalMoveException("Suicide")
|
||||
|
||||
# Play a Move from the current position, raise IllegalMoveException if invalid.
|
||||
def play(self, move, ignore_ko=False):
|
||||
def play(self, move: Move, ignore_ko: bool=False):
|
||||
if not move.is_pass and not (0 <= move.coords[0] < self.board_size and 0 <= move.coords[1] < self.board_size):
|
||||
raise IllegalMoveException(f"Move {move} outside of board coordinates")
|
||||
played_move = self.current_move.play(move)
|
||||
played_node = self.current_node.play(move)
|
||||
try:
|
||||
self._validate_move_and_update_chains(played_move, ignore_ko)
|
||||
self._validate_move_and_update_chains(played_node.move, ignore_ko)
|
||||
except IllegalMoveException:
|
||||
self.current_move.children = [m for m in self.current_move.children if m != played_move]
|
||||
self.current_node.children = [m for m in self.current_node.children if m != played_node]
|
||||
self._init_chains() # restore
|
||||
raise
|
||||
self.all_moves[played_move.id] = played_move
|
||||
self.current_move = played_move
|
||||
return played_move
|
||||
self._node_by_id[played_node.id] = played_node
|
||||
self.current_node = played_node
|
||||
return played_node
|
||||
|
||||
def undo(self):
|
||||
if self.current_move is not self.root:
|
||||
self.current_move = self.current_move.parent
|
||||
if self.current_node is not self.root:
|
||||
self.current_node = self.current_node.parent
|
||||
self._init_chains()
|
||||
|
||||
def redo(self):
|
||||
if self.current_move.children:
|
||||
self.play(self.current_move.children[-1])
|
||||
if self.current_node.children:
|
||||
self.play(self.current_node.children[-1])
|
||||
|
||||
def switch_branch(self, direction):
|
||||
cm = self.current_move # avoid race conditions
|
||||
cm = self.current_node # avoid race conditions
|
||||
if cm.parent and len(cm.parent.children) > 1:
|
||||
ix = cm.parent.children.index(cm)
|
||||
self.current_move = cm.parent.children[(ix + direction) % len(cm.parent.children)]
|
||||
self.current_node = cm.parent.children[(ix + direction) % len(cm.parent.children)]
|
||||
self._init_chains()
|
||||
|
||||
def place_handicap_stones(self, n_handicaps):
|
||||
@@ -324,25 +272,20 @@ class Board:
|
||||
if n_handicaps % 2 == 1:
|
||||
stones.append((middle, middle))
|
||||
stones += [(near, middle), (far, middle), (middle, near), (middle, far)]
|
||||
self.root['AB'] =[ Move(player=0, coords=stone).sgf() for stone in stones[:n_handicaps] ]
|
||||
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
|
||||
moves = []
|
||||
p = self.current_move
|
||||
while p is not self.root: # NB == is wrong here
|
||||
moves.append(p)
|
||||
p = p.parent
|
||||
return moves[::-1]
|
||||
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],[])
|
||||
|
||||
@property
|
||||
def current_player(self):
|
||||
return 1 - self.current_move.player
|
||||
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.all_moves.get(int(move_id))
|
||||
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]
|
||||
@@ -361,7 +304,7 @@ class Board:
|
||||
else:
|
||||
move_id = int(json["id"])
|
||||
is_pass = False
|
||||
move = self.all_moves.get(move_id)
|
||||
move = self._node_by_id.get(move_id)
|
||||
if move: # else this should be old
|
||||
move.set_analysis(json, is_pass)
|
||||
else:
|
||||
@@ -373,55 +316,24 @@ class Board:
|
||||
|
||||
@property
|
||||
def game_ended(self):
|
||||
return self.current_move.parent and self.current_move.is_pass and self.current_move.parent.is_pass
|
||||
return self.current_node.parent and self.current_node.is_pass and self.current_node.parent.is_pass
|
||||
|
||||
@property
|
||||
def prisoner_count(self):
|
||||
return [sum([m.player == player for m in self.prisoners]) for player in [0, 1]]
|
||||
return [sum([m.player == player for m in self.prisoners]) for player in Move.PLAYERS]
|
||||
|
||||
def __repr__(self):
|
||||
return "\n".join("".join(Move.PLAYERS[self.chains[c][0].player] if c >= 0 else "-" for c in line) for line in self.board) + f"\ncaptures: {self.prisoner_count}"
|
||||
|
||||
def write_sgf(self, komi, train_settings, file_name=None):
|
||||
def sgfify(mvs, comment=""):
|
||||
return f"(;GM[1]FF[4]SZ[{self.board_size}]KM[{komi}]RU[JP];" + ";".join(mvs) + (f"C[{comment}])" if comment else "")
|
||||
|
||||
def format_move(move, prev_move):
|
||||
undos = [m for m in prev_move.children if m != move]
|
||||
undo_cr = "".join(f"MA[{u.sgfcoords(self.board_size)}]" for u in undos if not u.is_pass)
|
||||
if (
|
||||
prev_move.analysis
|
||||
and prev_move.analysis[0]["move"] != "pass"
|
||||
and (move.evaluation_info[0] or 0.0) < train_settings["sgf_show_best_move_threshold"]
|
||||
and prev_move.analysis[0]["move"] != move.gtp()
|
||||
):
|
||||
best_sq = "".join(
|
||||
f"SQ[{Move(gtpcoords=mv['move'], player=0).sgfcoords(self.board_size)}]"
|
||||
for mv in prev_move.analysis
|
||||
if move.player_sign * mv["scoreLead"] >= move.player_sign * prev_move.analysis[0]["scoreLead"] - 0.5
|
||||
and mv["visits"] >= train_settings["balance_play_min_visits"]
|
||||
and mv["move"] != "pass"
|
||||
)
|
||||
else:
|
||||
best_sq = ""
|
||||
return move.sgf(self.board_size) + f"C[{move.comment(sgf=True)}]{undo_cr}{best_sq}"
|
||||
|
||||
moves = self.moves
|
||||
sgfmoves_small = [mv.sgf(self.board_size) for mv in moves]
|
||||
sgfmoves = [format_move(mv, pmv) for mv, pmv in zip(moves, [self.root] + moves[:-1])]
|
||||
|
||||
def write_sgf(self, file_name=None):
|
||||
file_name = file_name or f"sgfout/katrain_{self.game_id}.sgf"
|
||||
try:
|
||||
os.makedirs(os.path.dirname(file_name))
|
||||
except FileExistsError:
|
||||
pass
|
||||
os.makedirs(os.path.dirname(file_name),exist_ok=True)
|
||||
with open(file_name, "w") as f:
|
||||
f.write(sgfify(sgfmoves))
|
||||
return sgfify(sgfmoves_small, f"SGF with analysis written to {file_name}")
|
||||
f.write(self.root.sgf())
|
||||
return f"SGF with analysis written to {file_name}"
|
||||
|
||||
|
||||
|
||||
class SGF(sgfparser.SGF):
|
||||
_MOVE_CLASS = Move
|
||||
class KaTrainSGF(SGF):
|
||||
_MOVE_CLASS = KaTrainSGFNode
|
||||
|
||||
|
||||
+13
-13
@@ -19,7 +19,7 @@ from kivy.uix.gridlayout import GridLayout
|
||||
from kivy.uix.label import Label
|
||||
from kivy.uix.popup import Popup
|
||||
|
||||
from board import Board, IllegalMoveException, Move, SGF
|
||||
from board import Board, IllegalMoveException, SGFNode, KaTrainSGF
|
||||
|
||||
BASE_PATH = getattr(sys, "_MEIPASS", os.path.dirname(os.path.abspath(__file__))) # for pyinstaller
|
||||
|
||||
@@ -120,7 +120,7 @@ class EngineControls(GridLayout):
|
||||
|
||||
# handles showing completed analysis and triggered actions like auto undo and ai move
|
||||
def update_evaluation(self):
|
||||
current_move = self.board.current_move
|
||||
current_move = self.board.current_node
|
||||
self.score.set_prisoners(self.board.prisoner_count)
|
||||
current_player_is_human_or_both_robots = not self.ai_auto.active(current_move.player) or self.ai_auto.active(1 - current_move.player)
|
||||
if current_player_is_human_or_both_robots and current_move is not self.board.root:
|
||||
@@ -151,7 +151,7 @@ class EngineControls(GridLayout):
|
||||
self.update_evaluation()
|
||||
return
|
||||
# ai player doesn't technically need parent ready, but don't want to override waiting for undo
|
||||
current_move = self.board.current_move # this effectively checks undo didn't just happen
|
||||
current_move = self.board.current_node # this effectively checks undo didn't just happen
|
||||
if self.ai_auto.active(1 - current_move.player) and not self.board.game_ended:
|
||||
if current_move.children:
|
||||
self.info.text = "AI paused since moves were undone. Press 'AI Move' or choose a move for the AI to continue playing."
|
||||
@@ -161,17 +161,17 @@ class EngineControls(GridLayout):
|
||||
|
||||
# engine action functions
|
||||
def _do_play(self, *args):
|
||||
self.play(Move(player=self.board.current_player, coords=args[0]))
|
||||
self.play(SGFNode(player=self.board.next_player, coords=args[0]))
|
||||
|
||||
def _do_aimove(self):
|
||||
ts = self.train_settings
|
||||
while not self.board.current_move.analysis_ready:
|
||||
while not self.board.current_node.analysis_ready:
|
||||
self.info.text = "Thinking..."
|
||||
self.ai_thinking = True
|
||||
time.sleep(0.05)
|
||||
self.ai_thinking = False
|
||||
# select move
|
||||
current_move = self.board.current_move
|
||||
current_move = self.board.current_node
|
||||
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"]
|
||||
]
|
||||
@@ -186,7 +186,7 @@ class EngineControls(GridLayout):
|
||||
or move_eval > ts["balance_play_min_eval"]
|
||||
and -current_move.player_sign * score > ts["balance_play_target_score"]
|
||||
] or sel_moves
|
||||
aimove = Move(player=self.board.current_player, gtpcoords=random.choice(sel_moves)[0], robot=True)
|
||||
aimove = SGFNode(player=self.board.next_player, gtpcoords=random.choice(sel_moves)[0], robot=True)
|
||||
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"
|
||||
self.play(aimove)
|
||||
@@ -200,11 +200,11 @@ class EngineControls(GridLayout):
|
||||
def _do_undo(self):
|
||||
if (
|
||||
self.ai_lock.active
|
||||
and self.auto_undo.active(self.board.current_move.player)
|
||||
and len(self.board.current_move.parent.children) > self.num_undos(self.board.current_move)
|
||||
and self.auto_undo.active(self.board.current_node.player)
|
||||
and len(self.board.current_node.parent.children) > self.num_undos(self.board.current_node)
|
||||
and not self.train_settings.get("dont_lock_undos")
|
||||
):
|
||||
self.info.text = f"Can't undo this move more than {self.num_undos(self.board.current_move)} time(s) when locked"
|
||||
self.info.text = f"Can't undo this move more than {self.num_undos(self.board.current_node)} time(s) when locked"
|
||||
return
|
||||
self.board.undo()
|
||||
self.update_evaluation()
|
||||
@@ -232,7 +232,7 @@ class EngineControls(GridLayout):
|
||||
|
||||
def _do_analyze_extra(self, mode):
|
||||
stones = {s.coords for s in self.board.stones}
|
||||
current_move = self.board.current_move
|
||||
current_move = self.board.current_node
|
||||
if not current_move.analysis:
|
||||
self.info.text = "Wait for initial analysis to complete before doing a board-sweep or refinement"
|
||||
return
|
||||
@@ -244,7 +244,7 @@ class EngineControls(GridLayout):
|
||||
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.board_size) for y in range(self.board_size) if (x, y) not in stones]
|
||||
analyze_moves = [SGFNode(coords=(x, y)).gtp() for x in range(self.board_size) for y in range(self.board_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
|
||||
@@ -279,7 +279,7 @@ class EngineControls(GridLayout):
|
||||
try:
|
||||
root = SGF.parse(sgf)
|
||||
except:
|
||||
root = Move()
|
||||
root = SGFNode()
|
||||
if root.empty():
|
||||
fileselect_popup = Popup(title="Double Click SGF file to analyze", size_hint=(0.8, 0.8))
|
||||
fc = FileChooserListView(multiselect=False, path=os.path.expanduser(Config.get("sgf")["load"]), filters=["*.sgf"])
|
||||
|
||||
+2
-2
@@ -131,7 +131,7 @@ class BadukPanWidget(Widget):
|
||||
# stones
|
||||
moves = self.engine.board.moves
|
||||
last_move = moves[-1] if moves else self.engine.board.root
|
||||
current_player = self.engine.board.current_player
|
||||
current_player = self.engine.board.next_player
|
||||
full_eval_on = [self.engine.eval.active(0), self.engine.eval.active(1)]
|
||||
has_stone = {}
|
||||
last_few_moves = self.engine.board.moves[-Config.get("trainer").get("eval_off_show_last", 3) :]
|
||||
@@ -171,7 +171,7 @@ class BadukPanWidget(Widget):
|
||||
if self.engine.hints.active(current_player):
|
||||
hint_moves = last_move.ai_moves
|
||||
for i, d in enumerate(hint_moves):
|
||||
move = Move(gtpcoords=d["move"])
|
||||
move = SGFNode(gtpcoords=d["move"])
|
||||
c = [*self._eval_spectrum(d["evaluation"]), 0.5]
|
||||
if move.coords[0] is not None and move.coords not in undo_coords:
|
||||
if i == 0:
|
||||
|
||||
+155
-46
@@ -1,5 +1,6 @@
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional
|
||||
import copy
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
class ParseError(Exception):
|
||||
@@ -7,64 +8,168 @@ class ParseError(Exception):
|
||||
|
||||
|
||||
class Move:
|
||||
GTP_COORD = "ABCDEFGHJKLMNOPQRSTUVWYXYZ"
|
||||
PLAYERS = 'BW'
|
||||
SGF_COORD = [chr(i) for i in range(97, 123)]
|
||||
|
||||
@staticmethod
|
||||
def from_gtp(gtp_coords, player='B'):
|
||||
if "pass" in gtp_coords:
|
||||
Move(coords=None, player=player)
|
||||
return Move(coords=(Move.GTP_COORD.index(gtp_coords[0]), int(gtp_coords[1:]) - 1), player=player)
|
||||
|
||||
@staticmethod
|
||||
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)
|
||||
|
||||
def __init__(self, coords: Optional[Tuple[int,int]]=None, player: str='B'):
|
||||
self.player = player
|
||||
self.coords = coords
|
||||
|
||||
# def __repr__(self):
|
||||
# return f"{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 gtp(self):
|
||||
if self.is_pass:
|
||||
return "pass"
|
||||
return Move.GTP_COORD[self.coords[0]] + str(self.coords[1] + 1)
|
||||
|
||||
def sgf(self, board_size):
|
||||
if self.is_pass:
|
||||
return ""
|
||||
return f"{Move.SGF_COORD[self.coords[0]]}{Move.SGF_COORD[board_size - self.coords[1] - 1]}"
|
||||
|
||||
@property
|
||||
def is_pass(self):
|
||||
return self.coords is None
|
||||
|
||||
@property
|
||||
def opponent(self):
|
||||
return 'W' if self.player=='B' else 'B'
|
||||
|
||||
|
||||
class SGFNode:
|
||||
CAST_FIELDS = {"KM": float, "SZ": int, "HA": int} # cast property to this type
|
||||
LIST_FIELDS = ["AB", "AW", "TW", "TB"] # cast these properties to lists
|
||||
LIST_FIELDS = ["AB", "AW", "TW", "TB", "MA", "SQ", "CR", "TR", "LN","AR","LB"] # cast these properties to lists
|
||||
# TODO: all are potential lists?
|
||||
|
||||
def __init__(self):
|
||||
self.parent = None
|
||||
def __init__(self, parent=None, properties=None, move=None):
|
||||
self.children = []
|
||||
self.properties = {}
|
||||
self.properties = copy.copy(properties) if properties is not None else {}
|
||||
self.parent = parent
|
||||
if self.parent:
|
||||
self.parent.children.append(self)
|
||||
if parent and move:
|
||||
properties[move.player] = move.sgf() # NB needs root['SZ']
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"(;{self._node_sgf()})"
|
||||
@property
|
||||
def sgf_properties(self) -> Dict:
|
||||
"""For hooking into in a subclass and overriding/formatting any additional properties to be output"""
|
||||
return self.properties
|
||||
|
||||
def empty(self) -> bool:
|
||||
return not self.children and not self.properties
|
||||
|
||||
def add_child(self,child_branch):
|
||||
self.children.append(child_branch)
|
||||
child_branch.parent = self
|
||||
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()])
|
||||
if self.children:
|
||||
children = [c.sgf() for c in self.children]
|
||||
if len(children) == 1:
|
||||
sgf_str += ";" + children[0]
|
||||
else:
|
||||
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] = re.split(r"\]\s*\[", value)
|
||||
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 __getitem__(self, ix) -> Any:
|
||||
return self.properties.get(ix)
|
||||
def __getitem__(self, property) -> Any:
|
||||
return self.properties.get(property)
|
||||
|
||||
def _node_sgf(self) -> str:
|
||||
move_props = "".join([f"{k}[{']['.join(v) if isinstance(v,list) else v}]" for k, v in self.properties.items()])
|
||||
if not self.children:
|
||||
return move_props
|
||||
else:
|
||||
children = [c._node_sgf() for c in self.children]
|
||||
if len(children) == 1:
|
||||
return move_props + ";" + children[0]
|
||||
else:
|
||||
return move_props + "(;" + ")(;".join(children) + ")"
|
||||
def get(self, property, default) -> Any:
|
||||
return self.properties.get(property, default)
|
||||
|
||||
@property
|
||||
def parent(self) -> Optional["SGFNode"]:
|
||||
return self._parent
|
||||
|
||||
@parent.setter
|
||||
def parent(self, parent_node):
|
||||
self._parent = parent_node
|
||||
self._root = None
|
||||
|
||||
@property
|
||||
def root(self) -> "SGFNode": # cached root property
|
||||
if self._root is None:
|
||||
self._root = self.parent.root if self.parent else self
|
||||
return self._root
|
||||
|
||||
@property
|
||||
def board_size(self) -> int:
|
||||
return self.root['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)
|
||||
|
||||
@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,[]) ]
|
||||
|
||||
@property
|
||||
def move_with_placements(self) -> List[Move]:
|
||||
return self.placements + (self.move or [])
|
||||
|
||||
@property
|
||||
def is_root(self):
|
||||
return self.parent is None
|
||||
|
||||
@property
|
||||
def empty(self) -> bool:
|
||||
return not self.children and not self.properties
|
||||
|
||||
@property
|
||||
def nodes_in_tree(self):
|
||||
return [self] + sum([c.nodes_in_tree for c in self.children],[])
|
||||
|
||||
@property
|
||||
def nodes_from_root(self):
|
||||
return [self] if self.is_root else self.parent.nodes_from_root + [self]
|
||||
|
||||
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
|
||||
return SGFNode(parent=self, move=move)
|
||||
|
||||
@property
|
||||
def next_player(self):
|
||||
m = self.move
|
||||
if m and m.player=='B' or 'AB' in self.properties:
|
||||
return 'W'
|
||||
return 'B'
|
||||
|
||||
|
||||
class SGF:
|
||||
_MOVE_CLASS = Move
|
||||
|
||||
def __init__(self, contents):
|
||||
self.contents = contents
|
||||
try:
|
||||
self.ix = self.contents.index("(") + 1
|
||||
except ValueError:
|
||||
raise ParseError("Parse error: Expected '('")
|
||||
self.root = self._parse_branch()
|
||||
_MOVE_CLASS = SGFNode
|
||||
|
||||
@staticmethod
|
||||
def parse(input_str) -> Move:
|
||||
def parse(input_str) -> SGFNode:
|
||||
return SGF(input_str).root
|
||||
|
||||
@staticmethod
|
||||
def parse_file(filename, encoding=None) -> Move:
|
||||
def parse_file(filename, encoding=None) -> SGFNode:
|
||||
with open(filename, "rb") as f:
|
||||
bin_contents = f.read()
|
||||
if not encoding:
|
||||
@@ -76,24 +181,28 @@ class SGF:
|
||||
decoded = bin_contents.decode(encoding=encoding)
|
||||
return SGF.parse(decoded)
|
||||
|
||||
def _parse_branch(self) -> Move:
|
||||
move_tree = self._MOVE_CLASS()
|
||||
current_move = move_tree
|
||||
def __init__(self, contents):
|
||||
self.contents = contents
|
||||
try:
|
||||
self.ix = self.contents.index("(") + 1
|
||||
except ValueError:
|
||||
raise ParseError("Parse error: Expected '('")
|
||||
self.root = SGFNode()
|
||||
self._parse_branch(self.root)
|
||||
|
||||
def _parse_branch(self,current_move: SGFNode):
|
||||
while self.ix < len(self.contents): # https://xkcd.com/1171/
|
||||
match = re.match(r"\s*(?:\(|\)|;|(?:(\w+)((?:\[.*?(?<!\\)\]\s*)+)))", self.contents[self.ix :], re.DOTALL)
|
||||
if not match:
|
||||
break
|
||||
self.ix += len(match[0])
|
||||
if match[0] == ")":
|
||||
return move_tree
|
||||
return
|
||||
if match[0] == "(":
|
||||
current_move.add_child(self._parse_branch())
|
||||
self._parse_branch(SGFNode(parent=current_move))
|
||||
elif match[0] == ";":
|
||||
if not current_move.empty(): # ignore ; that generate empty nodes
|
||||
next_move = self._MOVE_CLASS()
|
||||
current_move.add_child(next_move)
|
||||
current_move = next_move
|
||||
if not current_move.empty: # ignore ; that generate empty nodes
|
||||
current_move = self._MOVE_CLASS(parent=current_move)
|
||||
else:
|
||||
prop, value = match[1], match[2].strip()[1:-1]
|
||||
current_move[prop] = value
|
||||
|
||||
Reference in new issue
Block a user