repackage
This commit is contained in:
1 parent
e8b0d2a5c4
commit
8f8d443151
17 files changed
+336
-338
No files matched your search
Whitespace-only changes.
+5
-5
@@ -3,12 +3,12 @@ import json
|
||||
import sys
|
||||
import time
|
||||
|
||||
from ai import ai_move
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_INFO
|
||||
from core.ai import ai_move
|
||||
from core.common import OUTPUT_ERROR, OUTPUT_INFO
|
||||
from bots.settings import bot_strategy_names
|
||||
from engine import EngineDiedException, KataGoEngine
|
||||
from game import Game, Move
|
||||
from sgf_parser import Move
|
||||
from core.engine import EngineDiedException, KataGoEngine
|
||||
from core.game import Game
|
||||
from core.sgf_parser import Move
|
||||
|
||||
if len(sys.argv) < 2:
|
||||
bot = "dev"
|
||||
|
||||
@@ -6,8 +6,8 @@ import sys
|
||||
import threading
|
||||
import traceback
|
||||
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_INFO
|
||||
from engine import KataGoEngine
|
||||
from core.common import OUTPUT_INFO
|
||||
from core.engine import KataGoEngine
|
||||
|
||||
PORT = int(sys.argv[1]) if len(sys.argv) > 1 else 8587
|
||||
|
||||
|
||||
+5
-5
@@ -7,11 +7,11 @@ import traceback
|
||||
from collections import defaultdict
|
||||
from concurrent.futures.thread import ThreadPoolExecutor
|
||||
|
||||
from ai import ai_move
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_INFO
|
||||
from core.ai import ai_move
|
||||
from core.common import OUTPUT_ERROR, OUTPUT_INFO
|
||||
from elote import EloCompetitor
|
||||
from engine import KataGoEngine
|
||||
from game import Game
|
||||
from core.engine import KataGoEngine
|
||||
from core.game import Game
|
||||
import json
|
||||
|
||||
DB_FILENAME = "bots/ai_performance.pickle"
|
||||
@@ -96,7 +96,7 @@ def retrieve_ais(selected_ais):
|
||||
|
||||
|
||||
test_ais = [
|
||||
# AI("Jigo", {}, {"max_visits": 100}),
|
||||
# AI("Jigo", {}, {"max_visits": 100}),
|
||||
AI("Policy", {}),
|
||||
AI("P:Local", {}),
|
||||
AI("P:Pick", {}),
|
||||
|
||||
+1
-1
@@ -239,6 +239,6 @@
|
||||
]
|
||||
},
|
||||
"debug": {
|
||||
"level": 1
|
||||
"level": 0
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
from gui.badukpan import BadukPanControls, BadukPanWidget
|
||||
from gui.controls import Controls
|
||||
from gui.kivyutils import *
|
||||
from gui.popups import LoadSGFPopup
|
||||
+7
-7
@@ -2,13 +2,13 @@ import heapq
|
||||
import math
|
||||
import random
|
||||
import time
|
||||
from typing import Any, Dict, List, Tuple
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_INFO, var_to_grid
|
||||
from engine import EngineDiedException
|
||||
from game import Game, GameNode, IllegalMoveException, Move
|
||||
from core.common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_INFO, var_to_grid
|
||||
from core.engine import EngineDiedException
|
||||
from core.game import Game, GameNode, IllegalMoveException, Move
|
||||
|
||||
|
||||
def weighted_selection_without_replacement(items: List[Tuple[float, float, int, int]], pick_n: int) -> List[Tuple[float, float, int, int]]:
|
||||
@@ -132,7 +132,7 @@ def ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode
|
||||
else:
|
||||
raise ValueError(f"Unknown AI mode {ai_mode}")
|
||||
elif "balance" in ai_mode and candidate_ai_moves[0]["move"] != "pass": # don't play suicidal to balance score - pass when it's best
|
||||
sign = cn.player_sign(cn.next_player) # TODO check
|
||||
sign = cn.player_sign(cn.next_player)
|
||||
sel_moves = [ # top move, or anything not too bad, or anything that makes you still ahead
|
||||
move
|
||||
for i, move in enumerate(candidate_ai_moves)
|
||||
@@ -140,10 +140,10 @@ def ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode
|
||||
or move["visits"] >= ai_settings["min_visits"]
|
||||
and (move["pointsLost"] < ai_settings["random_loss"] or move["pointsLost"] < ai_settings["max_loss"] and sign * move["scoreLead"] > ai_settings["target_score"])
|
||||
]
|
||||
aimove = Move.from_gtp(random.choice(sel_moves)["move"], player=cn.next_player) # TODO: could be weighted towards worse
|
||||
aimove = Move.from_gtp(random.choice(sel_moves)["move"], player=cn.next_player)
|
||||
ai_thoughts += f"Balance strategy selected moves {sel_moves} based on target score and max points lost, and randomly chose {aimove.gtp()}."
|
||||
elif "jigo" in ai_mode and candidate_ai_moves[0]["move"] != "pass":
|
||||
sign = cn.player_sign(cn.next_player) # TODO check
|
||||
sign = cn.player_sign(cn.next_player)
|
||||
jigo_move = min(candidate_ai_moves, key=lambda move: abs(sign * move["scoreLead"] - ai_settings["target_score"]))
|
||||
aimove = Move.from_gtp(jigo_move["move"], player=cn.next_player)
|
||||
ai_thoughts += f"Jigo strategy found candidate moves {candidate_ai_moves} moves and chose {aimove.gtp()} as closest to 0.5 point win"
|
||||
File renamed without changes.
@@ -6,8 +6,8 @@ import threading
|
||||
import time
|
||||
from typing import Callable, Optional
|
||||
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG
|
||||
from game_node import GameNode
|
||||
from core.common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG
|
||||
from core.game_node import GameNode
|
||||
|
||||
|
||||
class EngineDiedException(Exception):
|
||||
@@ -5,10 +5,10 @@ import threading
|
||||
from datetime import datetime
|
||||
from typing import Dict, List, Union
|
||||
|
||||
from common import var_to_grid, OUTPUT_INFO, OUTPUT_ERROR, OUTPUT_DEBUG
|
||||
from engine import KataGoEngine
|
||||
from game_node import GameNode
|
||||
from sgf_parser import SGF, Move
|
||||
from core.common import var_to_grid, OUTPUT_INFO, OUTPUT_DEBUG
|
||||
from core.engine import KataGoEngine
|
||||
from core.game_node import GameNode
|
||||
from core.sgf_parser import SGF, Move
|
||||
|
||||
|
||||
class IllegalMoveException(Exception):
|
||||
@@ -60,7 +60,7 @@ class Game:
|
||||
try:
|
||||
# for m in self.moves:
|
||||
for node in self.current_node.nodes_from_root:
|
||||
for m in node.move_with_placements: # TODO: placements are never illegal
|
||||
for m in node.move_with_placements:
|
||||
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)})")
|
||||
@@ -109,7 +109,7 @@ class Game:
|
||||
raise IllegalMoveException("Ko")
|
||||
self.prisoners += self.last_capture
|
||||
|
||||
if -1 not in neighbours(self.chains[this_chain]): # TODO: NZ?
|
||||
if -1 not in neighbours(self.chains[this_chain]): # TODO: NZ rules?
|
||||
raise IllegalMoveException("Suicide")
|
||||
|
||||
# Play a Move from the current position, raise IllegalMoveException if invalid.
|
||||
@@ -1,10 +1,9 @@
|
||||
import copy
|
||||
import math
|
||||
import random
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
from common import evaluation_class, var_to_grid
|
||||
from sgf_parser import Move, SGFNode
|
||||
from core.common import evaluation_class, var_to_grid
|
||||
from core.sgf_parser import Move, SGFNode
|
||||
|
||||
|
||||
class GameNode(SGFNode):
|
||||
+290
@@ -0,0 +1,290 @@
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import traceback
|
||||
from queue import Queue
|
||||
|
||||
from kivy.app import App
|
||||
from kivy.clock import Clock
|
||||
from kivy.core.clipboard import Clipboard
|
||||
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 kivy.uix.widget import Widget
|
||||
|
||||
from core.ai import ai_move
|
||||
from core.common import OUTPUT_INFO, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG
|
||||
from core.engine import KataGoEngine
|
||||
from core.game import Game, IllegalMoveException, KaTrainSGF
|
||||
from core.sgf_parser import Move, ParseError
|
||||
from gui.popups import NewGamePopup, ConfigPopup, 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.engine = None
|
||||
self.game = None
|
||||
self.new_game_popup = None
|
||||
self.fileselect_popup = None
|
||||
self.config_popup = None
|
||||
self.logger = lambda message, level=OUTPUT_INFO: self.log(message, level)
|
||||
|
||||
self._load_config()
|
||||
|
||||
self.debug_level = self.config("debug/level", OUTPUT_INFO)
|
||||
self.controls.ai_mode_groups["W"].values = self.controls.ai_mode_groups["B"].values = list(self.config("ai").keys())
|
||||
self.message_queue = Queue()
|
||||
|
||||
self._keyboard = Window.request_keyboard(None, self, "")
|
||||
self._keyboard.bind(on_key_down=self._on_keyboard_down)
|
||||
|
||||
def log(self, message, level=OUTPUT_INFO):
|
||||
if level == OUTPUT_ERROR:
|
||||
self.controls.set_status(f"ERROR: {message}")
|
||||
print(f"ERROR: {message}")
|
||||
elif self.debug_level >= level:
|
||||
print(message)
|
||||
|
||||
def _load_config(self):
|
||||
base_path = getattr(sys, "_MEIPASS", os.getcwd()) # 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._config_store = JsonStore(config_file, indent=4)
|
||||
self._config = dict(self._config_store)
|
||||
except Exception as e:
|
||||
self.log(f"Failed to load config {config_file}: {e}", OUTPUT_ERROR)
|
||||
sys.exit(1)
|
||||
|
||||
def save_config(self):
|
||||
for k, v in self._config.items():
|
||||
self._config_store.put(k, **v)
|
||||
|
||||
def config(self, setting, default=None):
|
||||
try:
|
||||
if "/" in setting:
|
||||
cat, key = setting.split("/")
|
||||
return self._config[cat].get(key, default)
|
||||
else:
|
||||
return self._config[setting]
|
||||
except KeyError:
|
||||
self.log(f"Missing configuration option {setting}", OUTPUT_ERROR)
|
||||
|
||||
def start(self):
|
||||
if self.engine:
|
||||
return
|
||||
self.board_gui.trainer_config = self.config("trainer")
|
||||
self.board_gui.ui_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()
|
||||
|
||||
def update_state(self, redraw_board=False): # is called after every message and on receiving analyses and config changes
|
||||
# AI and Trainer/auto-undo handlers
|
||||
cn = self.game.current_node
|
||||
auto_undo = cn.player and "undo" in self.controls.player_mode(cn.player)
|
||||
if auto_undo and cn.analysis_ready and cn.parent and cn.parent.analysis_ready:
|
||||
self.game.analyze_undo(cn, self.config("trainer")) # not via message loop
|
||||
if cn.analysis_ready and "ai" in self.controls.player_mode(cn.next_player).lower() and not cn.children and not self.game.ended and not (auto_undo and cn.auto_undo is None):
|
||||
self._do_ai_move(cn) # cn mismatch stops this if undo fired. avoid message loop here or fires repeatedly.
|
||||
|
||||
# Handle prisoners and next player display
|
||||
prisoners = self.game.prisoner_count
|
||||
top, bot = self.board_controls.black_prisoners.__self__, self.board_controls.white_prisoners.__self__ # no weakref
|
||||
if self.game.next_player == "W":
|
||||
top, bot = bot, top
|
||||
self.board_controls.mid_circles_container.clear_widgets()
|
||||
self.board_controls.mid_circles_container.add_widget(bot)
|
||||
self.board_controls.mid_circles_container.add_widget(top)
|
||||
self.board_controls.black_prisoners.text = str(prisoners["W"])
|
||||
self.board_controls.white_prisoners.text = str(prisoners["B"])
|
||||
|
||||
# update engine status dot
|
||||
if not self.engine or not self.engine.katago_process or self.engine.katago_process.poll() is not None:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_down_col")
|
||||
elif len(self.engine.queries) >= 4:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_busy_col")
|
||||
elif len(self.engine.queries) >= 2:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_little_busy_col")
|
||||
elif len(self.engine.queries) == 0:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_ready_col")
|
||||
else:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_almost_done_col")
|
||||
# redraw
|
||||
if redraw_board:
|
||||
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)
|
||||
self.controls.update_evaluation()
|
||||
|
||||
def _message_loop_thread(self):
|
||||
while True:
|
||||
game, msg, *args = self.message_queue.get()
|
||||
try:
|
||||
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)
|
||||
self.update_state()
|
||||
except Exception as e:
|
||||
self.log(f"Exception in processing message {msg} {args}: {e}", OUTPUT_ERROR)
|
||||
traceback.print_exc()
|
||||
|
||||
def __call__(self, message, *args):
|
||||
if self.game:
|
||||
self.message_queue.put([self.game.game_id, message, *args])
|
||||
|
||||
def _do_new_game(self, move_tree=None, analyze_fast=False):
|
||||
self.engine.on_new_game() # clear queries
|
||||
self.game = Game(self, self.engine, self.config("game"), move_tree=move_tree, analyze_fast=analyze_fast)
|
||||
self.controls.select_mode("analyze" if move_tree and len(move_tree.nodes_in_tree) > 1 else "play")
|
||||
self.controls.graph.initialize_from_game(self.game.root)
|
||||
self.update_state(redraw_board=True)
|
||||
|
||||
def _do_ai_move(self, node=None):
|
||||
if node is None or self.game.current_node == node:
|
||||
mode = self.controls.ai_mode(self.game.current_node.next_player)
|
||||
settings = self.config(f"ai/{mode}")
|
||||
if settings:
|
||||
ai_move(self.game, mode, settings)
|
||||
|
||||
def _do_undo(self, n_times=1):
|
||||
self.game.undo(n_times)
|
||||
|
||||
def _do_redo(self, n_times=1):
|
||||
self.game.redo(n_times)
|
||||
|
||||
def _do_switch_branch(self, direction):
|
||||
self.game.switch_branch(direction)
|
||||
|
||||
def _do_play(self, coords):
|
||||
try:
|
||||
self.game.play(Move(coords, player=self.game.next_player))
|
||||
except IllegalMoveException as e:
|
||||
self.controls.set_status(f"Illegal Move: {str(e)}")
|
||||
|
||||
def _do_analyze_extra(self, mode):
|
||||
self.game.analyze_extra(mode)
|
||||
|
||||
def _do_analyze_sgf_popup(self):
|
||||
if not self.fileselect_popup:
|
||||
self.fileselect_popup = Popup(title="Double Click SGF file to analyze", size_hint=(0.8, 0.8)).__self__
|
||||
popup_contents = LoadSGFPopup()
|
||||
self.fileselect_popup.add_widget(popup_contents)
|
||||
popup_contents.filesel.path = os.path.expanduser(self.config("sgf/sgf_load"))
|
||||
|
||||
def readfile(files, _mouse):
|
||||
self.fileselect_popup.dismiss()
|
||||
try:
|
||||
move_tree = KaTrainSGF.parse_file(files[0])
|
||||
except ParseError as e:
|
||||
self.log(f"Failed to load SGF. Parse Error: {e}", OUTPUT_ERROR)
|
||||
return
|
||||
self._do_new_game(move_tree=move_tree, analyze_fast=popup_contents.fast.active)
|
||||
|
||||
popup_contents.filesel.on_submit = readfile
|
||||
self.fileselect_popup.open()
|
||||
|
||||
def _do_new_game_popup(self):
|
||||
if not self.new_game_popup:
|
||||
self.new_game_popup = Popup(title="New Game", size_hint=(0.5, 0.6)).__self__
|
||||
popup_contents = NewGamePopup(self, self.new_game_popup, {k: v[0] for k, v in self.game.root.properties.items() if len(v) == 1})
|
||||
self.new_game_popup.add_widget(popup_contents)
|
||||
self.new_game_popup.open()
|
||||
|
||||
def _do_config_popup(self):
|
||||
if not self.config_popup:
|
||||
self.config_popup = Popup(title="Edit Settings", size_hint=(0.9, 0.9)).__self__
|
||||
popup_contents = ConfigPopup(self, self.config_popup, dict(self._config), ignore_cats=("trainer", "ai"))
|
||||
self.config_popup.add_widget(popup_contents)
|
||||
self.config_popup.open()
|
||||
|
||||
def _do_output_sgf(self):
|
||||
for pl in Move.PLAYERS:
|
||||
if not self.game.root.get_property(f"P{pl}"):
|
||||
_, model_file = os.path.split(self.engine.config["model"])
|
||||
self.game.root.set_property(
|
||||
f"P{pl}", f"AI {self.controls.ai_mode(pl)} (KataGo { os.path.splitext(model_file)[0]})" if "ai" in self.controls.player_mode(pl) else "Player"
|
||||
)
|
||||
msg = self.game.write_sgf(
|
||||
self.config("sgf/sgf_save"),
|
||||
trainer_config=self.config("trainer"),
|
||||
save_feedback=self.config("sgf/save_feedback"),
|
||||
eval_thresholds=self.config("trainer/eval_thresholds"),
|
||||
)
|
||||
self.log(msg, OUTPUT_INFO)
|
||||
self.controls.set_status(msg)
|
||||
|
||||
def load_sgf_from_clipboard(self):
|
||||
clipboard = Clipboard.paste()
|
||||
if not clipboard:
|
||||
self.controls.set_status(f"Ctrl-V pressed but clipboard is empty.")
|
||||
return
|
||||
try:
|
||||
move_tree = KaTrainSGF.parse(clipboard)
|
||||
except Exception as e:
|
||||
self.controls.set_status(f"Failed to imported game from clipboard: {e}\nClipboard contents: {clipboard[:50]}...")
|
||||
return
|
||||
move_tree.nodes_in_tree[-1].analyze(self.engine, analyze_fast=False) # speed up result for looking at end of game
|
||||
self._do_new_game(move_tree=move_tree, analyze_fast=True)
|
||||
self("redo", 999)
|
||||
self.log("Imported game from clipboard.", OUTPUT_INFO)
|
||||
|
||||
def on_touch_up(self, touch):
|
||||
if self.board_gui.collide_point(*touch.pos) or self.board_controls.collide_point(*touch.pos):
|
||||
if touch.button == "scrollup":
|
||||
self("redo")
|
||||
elif touch.button == "scrolldown":
|
||||
self("undo")
|
||||
return super().on_touch_up(touch)
|
||||
|
||||
def _on_keyboard_down(self, keyboard, keycode, text, modifiers):
|
||||
if isinstance(App.get_running_app().root_window.children[0], Popup):
|
||||
return # if in new game or load, don't allow keyboard shortcuts
|
||||
|
||||
shortcuts = {
|
||||
"q": self.controls.show_children,
|
||||
"w": self.controls.eval,
|
||||
"e": self.controls.hints,
|
||||
"r": self.controls.ownership,
|
||||
"t": self.controls.policy,
|
||||
"enter": ("ai-move",),
|
||||
"a": self.controls.analyze_extra,
|
||||
"s": self.controls.analyze_equalize,
|
||||
"d": self.controls.analyze_sweep,
|
||||
"right": ("switch-branch", 1),
|
||||
"left": ("switch-branch", -1),
|
||||
}
|
||||
if keycode[1] in shortcuts.keys():
|
||||
shortcut = shortcuts[keycode[1]]
|
||||
if isinstance(shortcut, Widget):
|
||||
shortcut.trigger_action(duration=0)
|
||||
else:
|
||||
self(*shortcut)
|
||||
elif keycode[1] == "tab":
|
||||
self.controls.switch_mode()
|
||||
elif keycode[1] == "spacebar":
|
||||
self("play", None) # pass
|
||||
elif keycode[1] in ["`", "~", "p"]:
|
||||
self.controls_box.hidden = not self.controls_box.hidden
|
||||
elif keycode[1] in ["up", "z"]:
|
||||
self("undo", 1 + ("shift" in modifiers) * 9 + ("ctrl" in modifiers) * 999)
|
||||
elif keycode[1] in ["down", "x"]:
|
||||
self("redo", 1 + ("shift" in modifiers) * 9 + ("ctrl" in modifiers) * 999)
|
||||
elif keycode[1] == "n" and "ctrl" in modifiers:
|
||||
self("new-game-popup")
|
||||
elif keycode[1] == "l" and "ctrl" in modifiers:
|
||||
self("analyze-sgf-popup")
|
||||
elif keycode[1] == "s" and "ctrl" in modifiers:
|
||||
self("output-sgf")
|
||||
elif keycode[1] == "c" and "ctrl" in modifiers:
|
||||
Clipboard.copy(self.game.root.sgf())
|
||||
self.controls.set_status("Copied SGF to clipboard.")
|
||||
elif keycode[1] == "v" and "ctrl" in modifiers:
|
||||
self.load_sgf_from_clipboard()
|
||||
return True
|
||||
File renamed without changes.
+3
-4
@@ -3,14 +3,13 @@ import math
|
||||
|
||||
from kivy.graphics.context_instructions import Color
|
||||
from kivy.graphics.vertex_instructions import Ellipse, Line, Rectangle
|
||||
from kivy.properties import ListProperty
|
||||
from kivy.uix.boxlayout import BoxLayout
|
||||
from kivy.uix.widget import Widget
|
||||
|
||||
from common import OUTPUT_DEBUG, evaluation_class
|
||||
from game import Move
|
||||
from core.common import OUTPUT_DEBUG, evaluation_class
|
||||
from core.game import Move
|
||||
from gui.kivyutils import draw_circle, draw_text
|
||||
from common import var_to_grid
|
||||
from core.common import var_to_grid
|
||||
from kivy.core.window import Window
|
||||
|
||||
|
||||
|
||||
+4
-4
@@ -1,5 +1,5 @@
|
||||
from collections import defaultdict
|
||||
from typing import Dict, List, DefaultDict, Tuple, Set
|
||||
from typing import Dict, List, DefaultDict, Tuple
|
||||
|
||||
from kivy.clock import Clock
|
||||
from kivy.uix.boxlayout import BoxLayout
|
||||
@@ -7,9 +7,9 @@ from kivy.uix.gridlayout import GridLayout
|
||||
from kivy.uix.label import Label
|
||||
from kivy.uix.popup import Popup
|
||||
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR
|
||||
from engine import KataGoEngine
|
||||
from game import Game, GameNode
|
||||
from core.common import OUTPUT_DEBUG, OUTPUT_ERROR
|
||||
from core.engine import KataGoEngine
|
||||
from core.game import Game, GameNode
|
||||
from gui.kivyutils import (
|
||||
BackgroundLabel,
|
||||
LabelledCheckBox,
|
||||
|
||||
+4
-297
@@ -2,299 +2,11 @@ from kivy.config import Config # isort:skip
|
||||
|
||||
Config.set("input", "mouse", "mouse,multitouch_on_demand") # isort:skip # no red dots on right click
|
||||
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import threading
|
||||
import traceback
|
||||
from queue import Queue
|
||||
from typing import Optional
|
||||
|
||||
from core.main import KaTrainGui
|
||||
import signal, sys, traceback
|
||||
from kivy.app import App
|
||||
from kivy.core.clipboard import Clipboard
|
||||
from kivy.core.window import Window
|
||||
from kivy.storage.jsonstore import JsonStore
|
||||
from kivy.uix.popup import Popup
|
||||
|
||||
from ai import ai_move
|
||||
from common import OUTPUT_DEBUG, OUTPUT_ERROR, OUTPUT_EXTRA_DEBUG, OUTPUT_INFO
|
||||
from engine import KataGoEngine
|
||||
from game import Game, IllegalMoveException, KaTrainSGF, Move
|
||||
from core.common import OUTPUT_DEBUG
|
||||
from gui import *
|
||||
from gui.popups import ConfigPopup, NewGamePopup
|
||||
from sgf_parser import ParseError
|
||||
|
||||
|
||||
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.engine = None # type: Optional[KataGoEngine]
|
||||
self.game = None
|
||||
self.new_game_popup = None
|
||||
self.fileselect_popup = None
|
||||
self.config_popup = None
|
||||
self.logger = lambda message, level=OUTPUT_INFO: self.log(message, level)
|
||||
|
||||
self._load_config()
|
||||
|
||||
self.debug_level = self.config("debug/level", OUTPUT_INFO)
|
||||
self.controls.ai_mode_groups["W"].values = self.controls.ai_mode_groups["B"].values = list(self.config("ai").keys())
|
||||
self.message_queue = Queue()
|
||||
|
||||
self._keyboard = Window.request_keyboard(None, self, "")
|
||||
self._keyboard.bind(on_key_down=self._on_keyboard_down)
|
||||
|
||||
def log(self, message, level=OUTPUT_INFO):
|
||||
if level == OUTPUT_ERROR:
|
||||
self.controls.set_status(f"ERROR: {message}")
|
||||
print(f"ERROR: {message}")
|
||||
elif self.debug_level >= level:
|
||||
print(message)
|
||||
|
||||
def _load_config(self):
|
||||
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._config_store = JsonStore(config_file, indent=4)
|
||||
self._config = dict(self._config_store)
|
||||
except Exception as e:
|
||||
self.log(f"Failed to load config {config_file}: {e}", OUTPUT_ERROR)
|
||||
sys.exit(1)
|
||||
|
||||
def save_config(self):
|
||||
for k, v in self._config.items():
|
||||
self._config_store.put(k, **v)
|
||||
|
||||
def config(self, setting, default=None):
|
||||
try:
|
||||
if "/" in setting:
|
||||
cat, key = setting.split("/")
|
||||
return self._config[cat].get(key, default)
|
||||
else:
|
||||
return self._config[setting]
|
||||
except KeyError:
|
||||
self.log(f"Missing configuration option {setting}", OUTPUT_ERROR)
|
||||
|
||||
def start(self):
|
||||
if self.engine:
|
||||
return
|
||||
self.board_gui.trainer_config = self.config("trainer") # TODO: could be cleaner
|
||||
self.board_gui.ui_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()
|
||||
|
||||
def update_state(self, redraw_board=False): # is called after every message and on receiving analyses and config changes
|
||||
# AI and Trainer/auto-undo handlers
|
||||
cn = self.game.current_node
|
||||
auto_undo = cn.player and "undo" in self.controls.player_mode(cn.player)
|
||||
if auto_undo and cn.analysis_ready and cn.parent and cn.parent.analysis_ready:
|
||||
self.game.analyze_undo(cn, self.config("trainer")) # not via message loop
|
||||
if cn.analysis_ready and "ai" in self.controls.player_mode(cn.next_player).lower() and not cn.children and not self.game.ended and not (auto_undo and cn.auto_undo is None):
|
||||
self._do_ai_move(cn) # cn mismatch stops this if undo fired. avoid message loop here or fires repeatedly.
|
||||
|
||||
# Handle prisoners and next player display
|
||||
prisoners = self.game.prisoner_count
|
||||
top, bot = self.board_controls.black_prisoners.__self__, self.board_controls.white_prisoners.__self__ # no weakref
|
||||
if self.game.next_player == "W":
|
||||
top, bot = bot, top
|
||||
self.board_controls.mid_circles_container.clear_widgets()
|
||||
self.board_controls.mid_circles_container.add_widget(bot)
|
||||
self.board_controls.mid_circles_container.add_widget(top)
|
||||
self.board_controls.black_prisoners.text = str(prisoners["W"])
|
||||
self.board_controls.white_prisoners.text = str(prisoners["B"])
|
||||
|
||||
# update engine status dot
|
||||
if not self.engine or not self.engine.katago_process or self.engine.katago_process.poll() is not None:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_down_col")
|
||||
elif len(self.engine.queries) >= 4:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_busy_col")
|
||||
elif len(self.engine.queries) >= 2:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_little_busy_col")
|
||||
elif len(self.engine.queries) == 0:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_ready_col")
|
||||
else:
|
||||
self.board_controls.engine_status_col = self.config("board_ui/engine_almost_done_col")
|
||||
# redraw
|
||||
if redraw_board:
|
||||
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)
|
||||
self.controls.update_evaluation()
|
||||
|
||||
def _message_loop_thread(self):
|
||||
while True:
|
||||
game, msg, *args = self.message_queue.get()
|
||||
try:
|
||||
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)
|
||||
self.update_state()
|
||||
except Exception as e:
|
||||
self.log(f"Exception in processing message {msg} {args}: {e}", OUTPUT_ERROR)
|
||||
traceback.print_exc()
|
||||
|
||||
def __call__(self, message, *args):
|
||||
# curframe = inspect.currentframe() # TODO remove
|
||||
# calframe = inspect.getouterframes(curframe, 2)
|
||||
# print('caller name:', calframe[1])
|
||||
if self.game:
|
||||
self.message_queue.put([self.game.game_id, message, *args])
|
||||
|
||||
def _do_new_game(self, move_tree=None, analyze_fast=False):
|
||||
self.engine.on_new_game() # clear queries
|
||||
self.game = Game(self, self.engine, self.config("game"), move_tree=move_tree, analyze_fast=analyze_fast)
|
||||
self.controls.select_mode("analyze" if move_tree and len(move_tree.nodes_in_tree) > 1 else "play")
|
||||
self.controls.graph.initialize_from_game(self.game.root)
|
||||
self.update_state(redraw_board=True)
|
||||
|
||||
def _do_ai_move(self, node=None):
|
||||
if node is None or self.game.current_node == node:
|
||||
mode = self.controls.ai_mode(self.game.current_node.next_player)
|
||||
settings = self.config(f"ai/{mode}")
|
||||
if settings:
|
||||
ai_move(self.game, mode, settings)
|
||||
|
||||
def _do_undo(self, n_times=1):
|
||||
self.game.undo(n_times)
|
||||
|
||||
def _do_redo(self, n_times=1):
|
||||
self.game.redo(n_times)
|
||||
|
||||
def _do_switch_branch(self, direction):
|
||||
self.game.switch_branch(direction)
|
||||
|
||||
def _do_play(self, coords):
|
||||
try:
|
||||
self.game.play(Move(coords, player=self.game.next_player))
|
||||
except IllegalMoveException as e:
|
||||
self.controls.set_status(f"Illegal Move: {str(e)}")
|
||||
|
||||
def _do_analyze_extra(self, mode):
|
||||
self.game.analyze_extra(mode)
|
||||
|
||||
def _do_analyze_sgf_popup(self):
|
||||
if not self.fileselect_popup:
|
||||
self.fileselect_popup = Popup(title="Double Click SGF file to analyze", size_hint=(0.8, 0.8)).__self__
|
||||
popup_contents = LoadSGFPopup()
|
||||
self.fileselect_popup.add_widget(popup_contents)
|
||||
popup_contents.filesel.path = os.path.expanduser(self.config("sgf/sgf_load"))
|
||||
|
||||
def readfile(files, _mouse):
|
||||
self.fileselect_popup.dismiss()
|
||||
try:
|
||||
move_tree = KaTrainSGF.parse_file(files[0])
|
||||
except ParseError as e:
|
||||
self.log(f"Failed to load SGF. Parse Error: {e}", OUTPUT_ERROR)
|
||||
return
|
||||
self._do_new_game(move_tree=move_tree, analyze_fast=popup_contents.fast.active)
|
||||
|
||||
popup_contents.filesel.on_submit = readfile
|
||||
self.fileselect_popup.open()
|
||||
|
||||
def _do_new_game_popup(self):
|
||||
if not self.new_game_popup:
|
||||
self.new_game_popup = Popup(title="New Game", size_hint=(0.5, 0.6)).__self__
|
||||
popup_contents = NewGamePopup(self, self.new_game_popup, {k: v[0] for k, v in self.game.root.properties.items() if len(v) == 1})
|
||||
self.new_game_popup.add_widget(popup_contents)
|
||||
self.new_game_popup.open()
|
||||
|
||||
def _do_config_popup(self):
|
||||
if not self.config_popup:
|
||||
self.config_popup = Popup(title="Edit Settings", size_hint=(0.9, 0.9)).__self__
|
||||
popup_contents = ConfigPopup(self, self.config_popup, dict(self._config), ignore_cats=("trainer", "ai"))
|
||||
self.config_popup.add_widget(popup_contents)
|
||||
self.config_popup.open()
|
||||
|
||||
def _do_output_sgf(self):
|
||||
for pl in Move.PLAYERS:
|
||||
if not self.game.root.get_property(f"P{pl}"):
|
||||
_, model_file = os.path.split(self.engine.config["model"])
|
||||
self.game.root.set_property(
|
||||
f"P{pl}", f"AI {self.controls.ai_mode(pl)} (KataGo { os.path.splitext(model_file)[0]})" if "ai" in self.controls.player_mode(pl) else "Player"
|
||||
)
|
||||
msg = self.game.write_sgf(
|
||||
self.config("sgf/sgf_save"),
|
||||
trainer_config=self.config("trainer"),
|
||||
save_feedback=self.config("sgf/save_feedback"),
|
||||
eval_thresholds=self.config("trainer/eval_thresholds"),
|
||||
)
|
||||
self.log(msg, OUTPUT_INFO)
|
||||
self.controls.set_status(msg)
|
||||
|
||||
def load_sgf_from_clipboard(self):
|
||||
clipboard = Clipboard.paste()
|
||||
if not clipboard:
|
||||
self.controls.set_status(f"Ctrl-V pressed but clipboard is empty.")
|
||||
return
|
||||
try:
|
||||
move_tree = KaTrainSGF.parse(clipboard)
|
||||
except Exception as e:
|
||||
self.controls.set_status(f"Failed to imported game from clipboard: {e}\nClipboard contents: {clipboard[:50]}...")
|
||||
return
|
||||
move_tree.nodes_in_tree[-1].analyze(self.engine, analyze_fast=False) # speed up result for looking at end of game
|
||||
self._do_new_game(move_tree=move_tree, analyze_fast=True)
|
||||
self("redo", 999)
|
||||
self.log("Imported game from clipboard.", OUTPUT_INFO)
|
||||
|
||||
def on_touch_up(self, touch):
|
||||
if self.board_gui.collide_point(*touch.pos) or self.board_controls.collide_point(*touch.pos):
|
||||
if touch.button == "scrollup":
|
||||
self("redo")
|
||||
elif touch.button == "scrolldown":
|
||||
self("undo")
|
||||
return super().on_touch_up(touch)
|
||||
|
||||
def _on_keyboard_down(self, keyboard, keycode, text, modifiers):
|
||||
if isinstance(App.get_running_app().root_window.children[0], Popup):
|
||||
return # if in new game or load, don't allow keyboard shortcuts
|
||||
|
||||
shortcuts = {
|
||||
"q": self.controls.show_children,
|
||||
"w": self.controls.eval,
|
||||
"e": self.controls.hints,
|
||||
"r": self.controls.ownership,
|
||||
"t": self.controls.policy,
|
||||
"enter": ("ai-move",),
|
||||
"a": self.controls.analyze_extra,
|
||||
"s": self.controls.analyze_equalize,
|
||||
"d": self.controls.analyze_sweep,
|
||||
"right": ("switch-branch", 1),
|
||||
"left": ("switch-branch", -1),
|
||||
}
|
||||
if keycode[1] in shortcuts.keys():
|
||||
shortcut = shortcuts[keycode[1]]
|
||||
if isinstance(shortcut, Widget):
|
||||
shortcut.trigger_action(duration=0)
|
||||
else:
|
||||
self(*shortcut)
|
||||
elif keycode[1] == "tab":
|
||||
self.controls.switch_mode()
|
||||
elif keycode[1] == "spacebar":
|
||||
self("play", None) # pass
|
||||
elif keycode[1] in ["`", "~", "p"]:
|
||||
self.controls_box.hidden = not self.controls_box.hidden
|
||||
elif keycode[1] in ["up", "z"]:
|
||||
self("undo", 1 + ("shift" in modifiers) * 9 + ("ctrl" in modifiers) * 999)
|
||||
elif keycode[1] in ["down", "x"]:
|
||||
self("redo", 1 + ("shift" in modifiers) * 9 + ("ctrl" in modifiers) * 999)
|
||||
elif keycode[1] == "n" and "ctrl" in modifiers:
|
||||
self("new-game-popup")
|
||||
elif keycode[1] == "l" and "ctrl" in modifiers:
|
||||
self("analyze-sgf-popup")
|
||||
elif keycode[1] == "s" and "ctrl" in modifiers:
|
||||
self("output-sgf")
|
||||
elif keycode[1] == "c" and "ctrl" in modifiers:
|
||||
Clipboard.copy(self.game.root.sgf())
|
||||
self.controls.set_status("Copied SGF to clipboard.")
|
||||
elif keycode[1] == "v" and "ctrl" in modifiers:
|
||||
self.load_sgf_from_clipboard()
|
||||
return True
|
||||
|
||||
|
||||
class KaTrainApp(App):
|
||||
@@ -313,10 +25,7 @@ class KaTrainApp(App):
|
||||
if getattr(self, "gui", None) and self.gui.engine:
|
||||
self.gui.engine.shutdown()
|
||||
|
||||
def signal_handler(self, signal, frame):
|
||||
import sys
|
||||
import traceback
|
||||
|
||||
def signal_handler(self, *args):
|
||||
if self.gui.debug_level >= OUTPUT_DEBUG:
|
||||
print("TRACEBACKS")
|
||||
for threadId, stack in sys._current_frames().items():
|
||||
@@ -330,8 +39,6 @@ class KaTrainApp(App):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# with open("katrain.kv", encoding="utf-8") as f: # avoid windows using another encoding
|
||||
# Builder.load_string(f.read())
|
||||
app = KaTrainApp()
|
||||
signal.signal(signal.SIGINT, app.signal_handler)
|
||||
try:
|
||||
|
||||
+1
-2
@@ -1,7 +1,6 @@
|
||||
import pytest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from game import Game, IllegalMoveException, Move
|
||||
from core.game import Game, IllegalMoveException, Move
|
||||
|
||||
|
||||
class MockKaTrain:
|
||||
|
||||
Reference in new issue
Block a user