policy moves display
This commit is contained in:
1 parent
875b752190
commit
0276e4b2a4
3 files changed
+59
-23
No files matched your search
+42
-14
@@ -420,24 +420,52 @@ class BadukPanWidget(Widget):
|
||||
if katrain.analysis_controls.policy.active and policy:
|
||||
policy_grid = var_to_grid(policy, (board_size_x, board_size_y))
|
||||
best_move_policy = max(*policy)
|
||||
colors = EVAL_COLORS[self.trainer_config["theme"]]
|
||||
text_lb = 0.01 * 0.01
|
||||
for y in range(board_size_y - 1, -1, -1):
|
||||
for x in range(board_size_x):
|
||||
if policy_grid[y][x] > 0:
|
||||
polsize = 1.1 * math.sqrt(policy_grid[y][x])
|
||||
policy_circle_color = (
|
||||
*POLICY_COLOR,
|
||||
POLICY_ALPHA + TOP_POLICY_ALPHA * (policy_grid[y][x] == best_move_policy),
|
||||
)
|
||||
move_policy = policy_grid[y][x]
|
||||
pol_order = 5 - int(-math.log10(max(1e-9, move_policy - 1e-9)))
|
||||
if pol_order >= 0:
|
||||
if move_policy > text_lb:
|
||||
Color(0.95, 0.75, 0.47, 1)
|
||||
draw_circle(
|
||||
(self.gridpos_x[x], self.gridpos_y[y]),
|
||||
self.stone_size * HINT_SCALE * 0.98,
|
||||
[0.95, 0.75, 0.47, 1],
|
||||
)
|
||||
scale = 0.95
|
||||
else:
|
||||
scale = 0.5
|
||||
draw_circle(
|
||||
(self.gridpos_x[x], self.gridpos_y[y]), polsize * self.stone_size, policy_circle_color
|
||||
(self.gridpos_x[x], self.gridpos_y[y]),
|
||||
HINT_SCALE * self.stone_size * scale,
|
||||
(*colors[pol_order][:3], POLICY_ALPHA),
|
||||
)
|
||||
polsize = math.sqrt(policy[-1])
|
||||
if move_policy > text_lb:
|
||||
Color(*BLACK)
|
||||
draw_text(
|
||||
pos=(self.gridpos_x[x], self.gridpos_y[y]),
|
||||
text=f"{100 * move_policy :.2f}"[:4] + "%",
|
||||
font_name="Roboto",
|
||||
halign="center",
|
||||
)
|
||||
if move_policy == best_move_policy:
|
||||
Color(*TOP_MOVE_BORDER_COLOR[:3], POLICY_ALPHA)
|
||||
Line(
|
||||
circle=(self.gridpos_x[x], self.gridpos_y[y], self.stone_size - dp(1.2),),
|
||||
width=dp(2),
|
||||
)
|
||||
|
||||
with pass_btn.canvas.after:
|
||||
draw_circle(
|
||||
(pass_btn.pos[0] + pass_btn.width / 2, pass_btn.pos[1] + pass_btn.height / 2),
|
||||
polsize * pass_btn.height / 2,
|
||||
POLICY_COLOR,
|
||||
)
|
||||
move_policy = policy[-1]
|
||||
pol_order = 5 - int(-math.log10(max(1e-9, move_policy - 1e-9)))
|
||||
if pol_order >= 0:
|
||||
draw_circle(
|
||||
(pass_btn.pos[0] + pass_btn.width / 2, pass_btn.pos[1] + pass_btn.height / 2),
|
||||
pass_btn.height / 2,
|
||||
(*colors[pol_order][:3], GHOST_ALPHA),
|
||||
)
|
||||
|
||||
# pass circle
|
||||
passed = len(nodes) > 1 and current_node.is_pass
|
||||
@@ -472,7 +500,7 @@ class BadukPanWidget(Widget):
|
||||
)
|
||||
|
||||
def draw_hover_contents(self, *_args):
|
||||
ghost_alpha = POLICY_ALPHA
|
||||
ghost_alpha = GHOST_ALPHA
|
||||
katrain = self.katrain
|
||||
game_ended = katrain.game.end_result
|
||||
current_node = katrain.game.current_node
|
||||
|
||||
@@ -614,14 +614,22 @@ class ScrollableLabel(ScrollView, BackgroundMixin):
|
||||
pass
|
||||
|
||||
|
||||
def draw_text(pos, text, font_name=None, markup=False, **kw):
|
||||
def cached_text_texture(text, font_name, markup, _cache={}, **kwargs):
|
||||
args = (text, font_name, markup, *[(k, v) for k, v in kwargs.items()])
|
||||
texture = _cache.get(args)
|
||||
if texture:
|
||||
return texture
|
||||
label_cls = CoreMarkupLabel if markup else CoreLabel
|
||||
label = label_cls(text=text, bold=True, font_name=font_name or i18n.font_name, **kw)
|
||||
label = label_cls(text=text, bold=True, font_name=font_name or i18n.font_name, **kwargs)
|
||||
label.refresh()
|
||||
texture = _cache[args] = label.texture
|
||||
return texture
|
||||
|
||||
|
||||
def draw_text(pos, text, font_name=None, markup=False, **kwargs):
|
||||
texture = cached_text_texture(text, font_name, markup, **kwargs)
|
||||
Rectangle(
|
||||
texture=label.texture,
|
||||
pos=(pos[0] - label.texture.size[0] / 2, pos[1] - label.texture.size[1] / 2),
|
||||
size=label.texture.size,
|
||||
texture=texture, pos=(pos[0] - texture.size[0] / 2, pos[1] - texture.size[1] / 2), size=texture.size,
|
||||
)
|
||||
|
||||
|
||||
@@ -631,7 +639,7 @@ def draw_circle(pos, r, col):
|
||||
|
||||
|
||||
# direct cache to texture, bypassing resource_find
|
||||
def cached_texture(path,_cache={}):
|
||||
def cached_texture(path, _cache={}):
|
||||
tex = _cache.get(path)
|
||||
if not tex:
|
||||
tex = _cache[path] = Image(resource_find(path)).texture
|
||||
|
||||
@@ -58,15 +58,15 @@ STONE_TEXT_COLORS = {"W": BLACK, "B": WHITE}
|
||||
|
||||
# board
|
||||
LINE_COLOR = [0, 0, 0]
|
||||
POLICY_COLOR = [0.9, 0.2, 0.8]
|
||||
STARPOINT_SIZE = 0.1
|
||||
BOARD_COLOR = [0.85, 0.68, 0.40, 1]
|
||||
STONE_SIZE = 0.505 # texture edge is transparent
|
||||
|
||||
GHOST_ALPHA = 0.6
|
||||
POLICY_ALPHA = 0.5
|
||||
|
||||
HINTS_LO_ALPHA = 0.6
|
||||
HINTS_ALPHA = 0.8
|
||||
POLICY_ALPHA = 0.6
|
||||
TOP_POLICY_ALPHA = 0.3
|
||||
TOP_MOVE_BORDER_COLOR = [10 / 255, 200 / 255, 250 / 255, 1.0]
|
||||
CHILD_SCALE = 0.95
|
||||
HINT_SCALE = 0.98
|
||||
|
||||
Reference in new issue
Block a user