refactor ai
This commit is contained in:
1 parent
5a67d3cc6f
commit
a322980f79
9 files changed
+159
-184
No files matched your search
@@ -22,7 +22,7 @@ jobs:
|
||||
pip3 install pytest wheel twine polib
|
||||
|
||||
- name: Run tests
|
||||
run: pytest -v tests
|
||||
run: pytest -v -s tests
|
||||
|
||||
- name: Run I18N conversion
|
||||
run: python i18n.py
|
||||
|
||||
@@ -27,7 +27,7 @@ jobs:
|
||||
pip3 install pytest wheel polib
|
||||
|
||||
- name: Run tests
|
||||
run: pytest tests
|
||||
run: pytest -v -s tests
|
||||
|
||||
- name: Run I18N conversion
|
||||
run: python i18n.py
|
||||
|
||||
+2
-2
@@ -39,7 +39,7 @@ from kivy.lang import Builder
|
||||
from kivy.resources import resource_add_path
|
||||
from kivy.uix.popup import Popup
|
||||
from kivy.uix.screenmanager import Screen
|
||||
from katrain.core.ai import ai_move
|
||||
from katrain.core.ai import generate_ai_move
|
||||
from kivy.core.window import Window
|
||||
|
||||
from katrain.core.lang import DEFAULT_LANGUAGE, i18n
|
||||
@@ -231,7 +231,7 @@ class KaTrainGui(Screen, KaTrainBase):
|
||||
mode = self.next_player_info.strategy
|
||||
settings = self.config(f"ai/{mode}")
|
||||
if settings is not None:
|
||||
ai_move(self.game, mode, settings)
|
||||
generate_ai_move(self.game, mode, settings)
|
||||
else:
|
||||
self.log(f"AI Mode {mode} not found!", OUTPUT_ERROR)
|
||||
|
||||
|
||||
+4
-5
@@ -1,7 +1,6 @@
|
||||
{
|
||||
"engine": {
|
||||
"katago": "",
|
||||
"_hint_katago": "Path to your katago executable",
|
||||
"model": "katrain/models/g170e-b15c192-s1672170752-d466197061.bin.gz",
|
||||
"config": "katrain/KataGo/analysis_config.cfg",
|
||||
"threads": 12,
|
||||
@@ -9,7 +8,6 @@
|
||||
"fast_visits": 50,
|
||||
"max_time": 3.0,
|
||||
"wide_root_noise": 0.0,
|
||||
"_hint_wide_root_noise": "A higher value here (typically 0.05-0.1)\nmakes the analysis explore more moves\nat the cost of some strength.",
|
||||
"_enable_ownership": true
|
||||
},
|
||||
"general": {
|
||||
@@ -18,7 +16,7 @@
|
||||
"anim_pv_time": 0.5,
|
||||
"debug_level": 0,
|
||||
"lang": "en",
|
||||
"version": "1.1.2"
|
||||
"version": "1.2.0"
|
||||
},
|
||||
"timer": {
|
||||
"byo_length": 30,
|
||||
@@ -94,7 +92,8 @@
|
||||
"pick_override": 0.95,
|
||||
"stddev": 1.5,
|
||||
"pick_n": 15,
|
||||
"pick_frac": 0.0
|
||||
"pick_frac": 0.0,
|
||||
"endgame": 0.5
|
||||
},
|
||||
"ai:p:tenuki": {
|
||||
"pick_override": 0.85,
|
||||
@@ -120,7 +119,7 @@
|
||||
"endgame": 0.4
|
||||
},
|
||||
"ai:p:rank": {
|
||||
"kyu": 4.0
|
||||
"kyu_rank": 4.0
|
||||
}
|
||||
}
|
||||
}
|
||||
+89
-141
@@ -20,9 +20,8 @@ from katrain.core.constants import (
|
||||
AI_TENUKI,
|
||||
AI_TERRITORY,
|
||||
AI_PICK,
|
||||
AI_RANK,
|
||||
AI_RANK,
|
||||
)
|
||||
from katrain.core.engine import EngineDiedException
|
||||
from katrain.core.game import Game, GameNode, Move
|
||||
|
||||
|
||||
@@ -42,157 +41,115 @@ def fmt_moves(moves: List[Tuple[float, Move]]):
|
||||
return ", ".join(f"{mv.gtp()} ({p:.2%})" for p, mv in moves)
|
||||
|
||||
|
||||
def ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode]:
|
||||
def policy_weighted_move(policy_moves, lower_bound, weaken_fac):
|
||||
lower_bound, weaken_fac = max(0, lower_bound), max(0.01, weaken_fac)
|
||||
weighted_coords = [(pv, pv ** (1 / weaken_fac), move) for pv, move in policy_moves if pv > lower_bound and not move.is_pass]
|
||||
if weighted_coords:
|
||||
top = weighted_selection_without_replacement(weighted_coords, 1)[0]
|
||||
ai_thoughts = f"Playing policy-weighted random move {top[2].gtp()} ({top[0]:.1%}) from {len(weighted_coords)} moves above lower_bound of {lower_bound:.1%}."
|
||||
else:
|
||||
top = policy_moves[0]
|
||||
ai_thoughts = f"Playing top policy move because no non-pass move > above lower_bound of {lower_bound:.1%}."
|
||||
return top[2], ai_thoughts
|
||||
|
||||
|
||||
def generate_influence_territory_weights(ai_mode, ai_settings, policy_grid, size):
|
||||
thr_line = ai_settings["threshold"] - 1 # zero-based
|
||||
if ai_mode == AI_INFLUENCE:
|
||||
weight = lambda x, y: (1 / ai_settings["line_weight"]) ** (max(0, thr_line - min(size[0] - 1 - x, x)) + max(0, thr_line - min(size[1] - 1 - y, y)))
|
||||
else:
|
||||
weight = lambda x, y: (1 / ai_settings["line_weight"]) ** (max(0, min(size[0] - 1 - x, x, size[1] - 1 - y, y) - thr_line))
|
||||
weighted_coords = [(policy_grid[y][x] * weight(x, y), weight(x, y), x, y) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0]
|
||||
ai_thoughts = f"Generated weights for {ai_mode} according to weight factor {ai_settings['line_weight']} and distance from {thr_line + 1}th line. "
|
||||
return weighted_coords, ai_thoughts
|
||||
|
||||
|
||||
def generate_local_tenuki_weights(ai_mode, ai_settings, policy_grid, cn, size):
|
||||
var = ai_settings["stddev"] ** 2
|
||||
mx, my = cn.move.coords
|
||||
weighted_coords = [(policy_grid[y][x], math.exp(-0.5 * ((x - mx) ** 2 + (y - my) ** 2) / var), x, y) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0]
|
||||
ai_thoughts = f"Generated weights based on one minus gaussian with variance {var} around coordinates {mx},{my}. "
|
||||
if ai_mode == AI_TENUKI:
|
||||
weighted_coords = [(p, 1 - w, x, y) for p, w, x, y in weighted_coords]
|
||||
ai_thoughts = f"Generated weights based on one minus gaussian with variance {var} around coordinates {mx},{my}. "
|
||||
return weighted_coords, ai_thoughts
|
||||
|
||||
|
||||
def generate_ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode]:
|
||||
cn = game.current_node
|
||||
while not cn.analysis_ready:
|
||||
time.sleep(0.01)
|
||||
engine = game.engines[cn.next_player]
|
||||
if engine.katago_process.poll() is not None: # TODO: clean up
|
||||
raise EngineDiedException(f"Engine for {cn.next_player} ({engine.config}) died")
|
||||
game.engines[cn.next_player].check_alive(exception_if_dead=True)
|
||||
|
||||
ai_thoughts = ""
|
||||
if (ai_mode in AI_STRATEGIES_POLICY) and cn.policy: # pure policy based move
|
||||
policy_moves = cn.policy_ranking
|
||||
pass_policy = cn.policy[-1]
|
||||
top_5_pass = any(
|
||||
[polmove[1].is_pass for polmove in policy_moves[:5]]
|
||||
) # dont make it jump around for the last few sensible non pass moves
|
||||
# dont make it jump around for the last few sensible non pass moves
|
||||
top_5_pass = any([polmove[1].is_pass for polmove in policy_moves[:5]])
|
||||
|
||||
size = game.board_size
|
||||
policy_grid = var_to_grid(cn.policy, size) # type: List[List[float]]
|
||||
top_policy_move = policy_moves[0][1]
|
||||
ai_thoughts += f"Using policy based strategy, base top 5 moves are {fmt_moves(policy_moves[:5])}. "
|
||||
len_legal_policy_moves = len([(pol, mv) for pol, mv in policy_moves if not mv.is_pass if pol > 0])
|
||||
if ai_mode == AI_POLICY and cn.depth <= ai_settings["opening_moves"]:
|
||||
if (ai_mode == AI_POLICY and cn.depth <= ai_settings["opening_moves"]) or (ai_mode in [AI_LOCAL, AI_TENUKI] and not cn.move or cn.move.coords is None):
|
||||
ai_mode = AI_WEIGHTED
|
||||
ai_thoughts += f"Switching to weighted strategy in the opening {int(ai_settings['opening_moves'])} moves. "
|
||||
ai_thoughts += f"Strategy override, using policy-weighted strategy instead. "
|
||||
ai_settings = {"pick_override": 0.9, "weaken_fac": 1, "lower_bound": 0.02}
|
||||
if ai_mode == AI_RANK:
|
||||
ai_settings = {"pick_override": (0.8*(1-((size[0]*size[1])-len_legal_policy_moves)/(size[0]*size[1])*.5)), "kyu": ai_settings["kyu"] }
|
||||
|
||||
if top_5_pass:
|
||||
aimove = top_policy_move
|
||||
ai_thoughts += "Playing top one because one of them is pass."
|
||||
elif ai_mode == AI_POLICY:
|
||||
aimove = top_policy_move
|
||||
ai_thoughts += f"Playing top policy move {aimove.gtp()}."
|
||||
elif policy_moves[0][0] > ai_settings["pick_override"]:
|
||||
aimove = top_policy_move
|
||||
ai_thoughts += (
|
||||
f"Top policy move has weight > {ai_settings['pick_override']:.1%}, so overriding other strategies."
|
||||
)
|
||||
elif ai_mode == AI_WEIGHTED:
|
||||
lower_bound = max(0, ai_settings["lower_bound"]) * 2 # compensate for first halving in loop
|
||||
weaken_fac = max(0.01, ai_settings["weaken_fac"])
|
||||
weighted_coords = []
|
||||
while not weighted_coords and lower_bound > 1e-6: # fix edge case where no moves are > lb
|
||||
lower_bound /= 2
|
||||
weighted_coords = [
|
||||
(policy_grid[y][x], policy_grid[y][x] ** (1 / weaken_fac), x, y)
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > lower_bound
|
||||
]
|
||||
top = weighted_selection_without_replacement(weighted_coords, 1)
|
||||
if top:
|
||||
best = top[0]
|
||||
policy_value = best[0]
|
||||
coords = best[2:]
|
||||
else: # weighted or pick-based
|
||||
legal_policy_moves = [(pol, mv) for pol, mv in policy_moves if not mv.is_pass and pol > 0]
|
||||
board_squares = size[0] * size[1]
|
||||
if ai_mode == AI_RANK: # calibrated, override from 0.8 at start to ~0.4 at full board
|
||||
override = 0.8 * (1 - 0.5 * (board_squares - len(legal_policy_moves)) / board_squares)
|
||||
else:
|
||||
policy_value = pass_policy
|
||||
coords = None
|
||||
aimove = Move(coords, player=cn.next_player) # just take a random move by policy w/o noise
|
||||
ai_thoughts += f"Playing policy-weighted random move {aimove.gtp()} ({policy_value:.1%})" + (
|
||||
" because no other moves were found."
|
||||
if not top
|
||||
else f" because strategy is weighted (lower bound={lower_bound:.2%}, num moves > lb={len(weighted_coords)})."
|
||||
)
|
||||
elif ai_mode in AI_STRATEGIES_PICK:
|
||||
legal_policy_moves = [(pol, mv) for pol, mv in policy_moves if not mv.is_pass if pol > 0]
|
||||
if ai_mode!=AI_RANK:
|
||||
n_moves = int(ai_settings["pick_frac"] * len(legal_policy_moves) + ai_settings["pick_n"])
|
||||
if ai_mode in [AI_INFLUENCE, AI_TERRITORY]:
|
||||
override = ai_settings["pick_override"]
|
||||
|
||||
thr_line = ai_settings["threshold"] - 1 # zero-based
|
||||
if cn.depth >= ai_settings["endgame"] * size[0] * size[1]:
|
||||
weighted_coords = [
|
||||
(policy_grid[y][x], 1, x, y)
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > 0
|
||||
]
|
||||
ai_thoughts += (
|
||||
f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. "
|
||||
)
|
||||
else:
|
||||
if ai_mode == AI_INFLUENCE:
|
||||
weight = lambda x, y: (1 / ai_settings["line_weight"]) ** (
|
||||
max(0, thr_line - min(size[0] - 1 - x, x)) + max(0, thr_line - min(size[1] - 1 - y, y))
|
||||
)
|
||||
else:
|
||||
weight = lambda x, y: (1 / ai_settings["line_weight"]) ** (
|
||||
max(0, min(size[0] - 1 - x, x, size[1] - 1 - y, y) - thr_line)
|
||||
)
|
||||
weighted_coords = [
|
||||
(policy_grid[y][x] * weight(x, y), weight(x, y), x, y)
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > 0
|
||||
]
|
||||
ai_thoughts += f"Generated weights for {ai_mode} according to weight factor {ai_settings['line_weight']} and distance from {thr_line+1}th line. "
|
||||
elif ai_mode in [AI_LOCAL, AI_TENUKI]:
|
||||
var = ai_settings["stddev"] ** 2
|
||||
if not cn.move or cn.move.coords is None:
|
||||
weighted_coords = [(1, 1, *top_policy_move.coords)] # if "pick" in ai_mode -> even
|
||||
ai_thoughts += f"No previous non-pass move, faking weights to play top policy move. "
|
||||
else:
|
||||
mx, my = cn.move.coords
|
||||
weighted_coords = [
|
||||
(policy_grid[y][x], math.exp(-0.5 * ((x - mx) ** 2 + (y - my) ** 2) / var), x, y)
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > 0
|
||||
]
|
||||
if ai_mode == AI_TENUKI:
|
||||
if cn.depth < ai_settings["endgame"] * size[0] * size[1]:
|
||||
weighted_coords = [(p, 1 - w, x, y) for p, w, x, y in weighted_coords]
|
||||
ai_thoughts += f"Generated weights based on one minus gaussian with variance {var} around coordinates {mx},{my}. "
|
||||
else:
|
||||
weighted_coords = [(p, 1, x, y) for p, w, x, y in weighted_coords]
|
||||
ai_thoughts += f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. "
|
||||
else:
|
||||
ai_thoughts += (
|
||||
f"Generated weights based on gaussian with variance {var} around coordinates {mx},{my}. "
|
||||
)
|
||||
elif ai_mode == AI_PICK:
|
||||
weighted_coords = [
|
||||
(policy_grid[y][x], 1, x, y)
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > 0
|
||||
]
|
||||
elif ai_mode == AI_RANK:
|
||||
n_moves = int(round((size[0]*size[1])/361*10**(-0.05737*ai_settings["kyu"] + 1.9482)))
|
||||
weighted_coords = [
|
||||
(policy_grid[y][x], 1, x, y)
|
||||
for x in range(size[0])
|
||||
for y in range(size[1])
|
||||
if policy_grid[y][x] > 0
|
||||
]
|
||||
else:
|
||||
raise ValueError(f"Unknown AI mode {ai_mode}")
|
||||
pick_moves = weighted_selection_without_replacement(weighted_coords, n_moves)
|
||||
ai_thoughts += f"Picked {min(n_moves,len(weighted_coords))} random moves according to weights. "
|
||||
if pick_moves:
|
||||
new_top = [(p, Move((x, y), player=cn.next_player)) for p, wt, x, y in heapq.nlargest(5, pick_moves)]
|
||||
aimove = new_top[0][1]
|
||||
ai_thoughts += f"Top 5 among these were {fmt_moves(new_top)} and picked top {aimove.gtp()}. "
|
||||
if new_top[0][0] < pass_policy:
|
||||
ai_thoughts += f"But found pass ({pass_policy:.2%} to be higher rated than {aimove.gtp()} ({new_top[0][0]:.2%}) so will play top policy move instead."
|
||||
aimove = top_policy_move
|
||||
else:
|
||||
if policy_moves[0][0] > override:
|
||||
aimove = top_policy_move
|
||||
ai_thoughts += f"Pick policy strategy {ai_mode} failed to find legal moves, so is playing top policy move {aimove.gtp()}."
|
||||
else:
|
||||
raise ValueError(f"Unknown AI mode {ai_mode}")
|
||||
ai_thoughts += f"Top policy move has weight > {override:.1%}, so overriding other strategies."
|
||||
elif ai_mode == AI_WEIGHTED:
|
||||
aimove, ai_thoughts = policy_weighted_move(policy_moves, ai_settings["lower_bound"], ai_settings["weaken_fac"])
|
||||
elif ai_mode in AI_STRATEGIES_PICK:
|
||||
|
||||
if ai_mode != AI_RANK:
|
||||
n_moves = int(ai_settings["pick_frac"] * len(legal_policy_moves) + ai_settings["pick_n"])
|
||||
else:
|
||||
n_moves = int(round(board_squares / 361 * 10 ** (-0.05737 * ai_settings["kyu_rank"] + 1.9482)))
|
||||
|
||||
if ai_mode in [AI_INFLUENCE, AI_TERRITORY, AI_LOCAL, AI_TENUKI]:
|
||||
if cn.depth > ai_settings["endgame"] * board_squares:
|
||||
weighted_coords = [(pol, 1, *mv.coords) for pol, mv in legal_policy_moves]
|
||||
x_ai_thoughts = f"Generated equal weights as move number >= {ai_settings['endgame'] * size[0] * size[1]}. "
|
||||
elif ai_mode in [AI_INFLUENCE, AI_TERRITORY]:
|
||||
weighted_coords, x_ai_thoughts = generate_influence_territory_weights(ai_mode, ai_settings, policy_grid, size)
|
||||
else: # ai_mode in [AI_LOCAL, AI_TENUKI]
|
||||
weighted_coords, x_ai_thoughts = generate_local_tenuki_weights(ai_mode, ai_settings, policy_grid, cn, size)
|
||||
ai_thoughts += x_ai_thoughts
|
||||
else: # ai_mode in [AI_PICK, AI_RANK]:
|
||||
weighted_coords = [(policy_grid[y][x], 1, x, y) for x in range(size[0]) for y in range(size[1]) if policy_grid[y][x] > 0]
|
||||
|
||||
pick_moves = weighted_selection_without_replacement(weighted_coords, n_moves)
|
||||
ai_thoughts += f"Picked {min(n_moves,len(weighted_coords))} random moves according to weights. "
|
||||
|
||||
if pick_moves:
|
||||
new_top = [(p, Move((x, y), player=cn.next_player)) for p, wt, x, y in heapq.nlargest(5, pick_moves)]
|
||||
aimove = new_top[0][1]
|
||||
ai_thoughts += f"Top 5 among these were {fmt_moves(new_top)} and picked top {aimove.gtp()}. "
|
||||
if new_top[0][0] < pass_policy:
|
||||
ai_thoughts += f"But found pass ({pass_policy:.2%} to be higher rated than {aimove.gtp()} ({new_top[0][0]:.2%}) so will play top policy move instead."
|
||||
aimove = top_policy_move
|
||||
else:
|
||||
aimove = top_policy_move
|
||||
ai_thoughts += f"Pick policy strategy {ai_mode} failed to find legal moves, so is playing top policy move {aimove.gtp()}."
|
||||
else:
|
||||
raise ValueError(f"Unknown Policy-based AI mode {ai_mode}")
|
||||
else: # Engine based move
|
||||
candidate_ai_moves = cn.candidate_moves
|
||||
top_cand = Move.from_gtp(candidate_ai_moves[0]["move"], player=cn.next_player)
|
||||
@@ -202,21 +159,12 @@ def ai_move(game: Game, ai_mode: str, ai_settings: Dict) -> Tuple[Move, GameNode
|
||||
else:
|
||||
if ai_mode == AI_JIGO:
|
||||
sign = cn.player_sign(cn.next_player)
|
||||
jigo_move = min(
|
||||
candidate_ai_moves, key=lambda move: abs(sign * move["scoreLead"] - ai_settings["target_score"])
|
||||
)
|
||||
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 {len(candidate_ai_moves)} candidate moves (best {top_cand.gtp()}) and chose {aimove.gtp()} as closest to 0.5 point win"
|
||||
elif ai_mode == AI_SCORELOSS:
|
||||
c = ai_settings["strength"]
|
||||
moves = [
|
||||
(
|
||||
d["pointsLost"],
|
||||
math.exp(min(200, -c * max(0, d["pointsLost"]))),
|
||||
Move.from_gtp(d["move"], player=cn.next_player),
|
||||
)
|
||||
for d in candidate_ai_moves
|
||||
]
|
||||
moves = [(d["pointsLost"], math.exp(min(200, -c * max(0, d["pointsLost"]))), Move.from_gtp(d["move"], player=cn.next_player),) for d in candidate_ai_moves]
|
||||
topmove = weighted_selection_without_replacement(moves, 1)[0]
|
||||
aimove = topmove[2]
|
||||
ai_thoughts += f"ScoreLoss strategy found {len(candidate_ai_moves)} candidate moves (best {top_cand.gtp()}) and chose {aimove.gtp()} (weight {topmove[1]:.3f}, point loss {topmove[0]:.1f}) based on score weights."
|
||||
|
||||
@@ -45,13 +45,13 @@ class KaTrainBase:
|
||||
|
||||
"""Settings, logging, and players functionality, so other classes like bots who need a katrain instance can be used without a GUI"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.debug_level = 0
|
||||
def __init__(self, force_package_config=False,debug_level=0, **kwargs):
|
||||
self.debug_level = debug_level
|
||||
self.game = None
|
||||
|
||||
self.logger = lambda message, level=OUTPUT_INFO: self.log(message, level)
|
||||
self.config_file = self._load_config()
|
||||
self.debug_level = self.config("general/debug_level", OUTPUT_INFO)
|
||||
self.config_file = self._load_config(force_package_config=force_package_config)
|
||||
self.debug_level = debug_level or self.config("general/debug_level", OUTPUT_INFO)
|
||||
|
||||
Config.set("kivy", "log_level", "error")
|
||||
if self.debug_level >= OUTPUT_DEBUG:
|
||||
@@ -68,37 +68,40 @@ class KaTrainBase:
|
||||
elif self.debug_level >= level:
|
||||
print(message)
|
||||
|
||||
def _load_config(self):
|
||||
def _load_config(self,force_package_config):
|
||||
if len(sys.argv) > 1 and sys.argv[1].endswith(".json"):
|
||||
config_file = os.path.abspath(sys.argv[1])
|
||||
self.log(f"Using command line config file {config_file}", OUTPUT_INFO)
|
||||
else:
|
||||
user_config_file = find_package_resource(self.USER_CONFIG_FILE)
|
||||
package_config_file = find_package_resource(self.PACKAGE_CONFIG_FILE)
|
||||
try:
|
||||
if not os.path.exists(user_config_file):
|
||||
os.makedirs(os.path.split(user_config_file)[0], exist_ok=True)
|
||||
shutil.copyfile(package_config_file, user_config_file)
|
||||
config_file = user_config_file
|
||||
self.log(f"Copied package config to local file {config_file}", OUTPUT_INFO)
|
||||
else: # user file exists
|
||||
version = JsonStore(user_config_file, indent=4).get("general")["version"]
|
||||
if version < CONFIG_MIN_VERSION:
|
||||
backup = user_config_file + f".{version}.backup"
|
||||
shutil.copyfile(user_config_file, backup)
|
||||
shutil.copyfile(package_config_file, user_config_file)
|
||||
self.log(
|
||||
f"Copied package config file to {user_config_file} as user file is outdated (<{CONFIG_MIN_VERSION}). Old version stored as {backup}",
|
||||
OUTPUT_INFO,
|
||||
)
|
||||
config_file = user_config_file
|
||||
self.log(f"Using user config file {config_file}", OUTPUT_INFO)
|
||||
except Exception as e:
|
||||
if force_package_config:
|
||||
config_file = package_config_file
|
||||
self.log(
|
||||
f"Using package config file {config_file} (exception {e} occurred when finding or creating user config)",
|
||||
OUTPUT_INFO,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
if not os.path.exists(user_config_file):
|
||||
os.makedirs(os.path.split(user_config_file)[0], exist_ok=True)
|
||||
shutil.copyfile(package_config_file, user_config_file)
|
||||
config_file = user_config_file
|
||||
self.log(f"Copied package config to local file {config_file}", OUTPUT_INFO)
|
||||
else: # user file exists
|
||||
version = JsonStore(user_config_file, indent=4).get("general")["version"]
|
||||
if version < CONFIG_MIN_VERSION:
|
||||
backup = user_config_file + f".{version}.backup"
|
||||
shutil.copyfile(user_config_file, backup)
|
||||
shutil.copyfile(package_config_file, user_config_file)
|
||||
self.log(
|
||||
f"Copied package config file to {user_config_file} as user file is outdated (<{CONFIG_MIN_VERSION}). Old version stored as {backup}",
|
||||
OUTPUT_INFO,
|
||||
)
|
||||
config_file = user_config_file
|
||||
self.log(f"Using user config file {config_file}", OUTPUT_INFO)
|
||||
except Exception as e:
|
||||
config_file = package_config_file
|
||||
self.log(
|
||||
f"Using package config file {config_file} (exception {e} occurred when finding or creating user config)",
|
||||
OUTPUT_INFO,
|
||||
)
|
||||
try:
|
||||
self._config_store = JsonStore(config_file, indent=4)
|
||||
except Exception as e:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
VERSION = "1.1.2"
|
||||
VERSION = "1.2.0"
|
||||
HOMEPAGE = "https://github.com/sanderland/katrain"
|
||||
CONFIG_MIN_VERSION = "1.1.2"
|
||||
CONFIG_MIN_VERSION = "1.2.0"
|
||||
|
||||
PLAYER_HUMAN, PLAYER_AI = "player:human", "player:ai"
|
||||
PLAYER_TYPES = [PLAYER_HUMAN, PLAYER_AI]
|
||||
@@ -33,13 +33,13 @@ AI_STRATEGIES_RECOMMENDED_ORDER = [
|
||||
AI_SCORELOSS,
|
||||
AI_POLICY,
|
||||
AI_WEIGHTED,
|
||||
AI_RANK,
|
||||
AI_PICK,
|
||||
AI_LOCAL,
|
||||
AI_TENUKI,
|
||||
AI_TERRITORY,
|
||||
AI_INFLUENCE,
|
||||
AI_JIGO,
|
||||
AI_RANK,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -82,6 +82,12 @@ class KataGoEngine:
|
||||
self.shutdown(finish=False)
|
||||
self.start()
|
||||
|
||||
def check_alive(self,exception_if_dead=False):
|
||||
ok = self.katago_process and self.katago_process.poll() is None
|
||||
if not ok and exception_if_dead:
|
||||
raise EngineDiedException(f"Engine died (process {self.katago_process}, poll {self.katago_process and self.katago_process.poll()}) config {self.config}")
|
||||
return ok
|
||||
|
||||
def shutdown(self, finish=False):
|
||||
process = self.katago_process
|
||||
if finish and process:
|
||||
|
||||
+22
-3
@@ -1,8 +1,27 @@
|
||||
import pytest
|
||||
|
||||
from katrain.core.constants import AI_STRATEGIES_RECOMMENDED_ORDER, AI_STRATEGIES
|
||||
from katrain.core.ai import generate_ai_move
|
||||
from katrain.core.constants import AI_STRATEGIES_RECOMMENDED_ORDER, AI_STRATEGIES, OUTPUT_INFO
|
||||
from katrain.core.base_katrain import KaTrainBase
|
||||
from katrain.core.engine import KataGoEngine
|
||||
from katrain.core.game import Game
|
||||
from katrain.core.constants import AI_STRATEGIES
|
||||
|
||||
|
||||
class TestAI:
|
||||
def test_order(self):
|
||||
assert set(AI_STRATEGIES_RECOMMENDED_ORDER) == set(AI_STRATEGIES)
|
||||
|
||||
def test_ai_strategies(self):
|
||||
katrain = KaTrainBase(force_package_config=True, debug_level=0)
|
||||
engine = KataGoEngine(katrain, katrain.config("engine"))
|
||||
game = Game(katrain, engine)
|
||||
|
||||
n_rounds = 3
|
||||
for _ in range(n_rounds):
|
||||
for strategy in AI_STRATEGIES:
|
||||
settings = katrain.config(f"ai/{strategy}")
|
||||
move, played_node = generate_ai_move(game, strategy, settings)
|
||||
katrain.log(f"Testing strategy {strategy} -> {move}", OUTPUT_INFO)
|
||||
assert move.coords is not None
|
||||
assert played_node == game.current_node
|
||||
|
||||
assert game.current_node.depth == len(AI_STRATEGIES) * n_rounds
|
||||
Reference in new issue
Block a user