diff --git a/mcts/node.py b/mcts/node.py index 3efae2e..fce7a21 100644 --- a/mcts/node.py +++ b/mcts/node.py @@ -472,8 +472,11 @@ def get_analysis_from_status_list(self, mode: str, """ out = "" if mode == "cgos": + # 探索結果が未反映 (node_visits == 0) の場合はゼロ除算を避け、 + # ニューラルネットワークの評価値で代用する cgos_dict = { - "winrate" : float(self.node_value_sum) / self.node_visits, + "winrate" : float(self.node_value_sum) / self.node_visits \ + if self.node_visits > 0 else float(self.raw_value), "visits" : self.node_visits, "moves" : [] } diff --git a/mcts/tree.py b/mcts/tree.py index f10ce6c..862b85e 100644 --- a/mcts/tree.py +++ b/mcts/tree.py @@ -162,6 +162,10 @@ def search(self, board: GoBoard, color: Stone, time_manager: TimeManager, \ break if len(analysis_query) > 0 and interval == 0: + # 未反映のミニバッチ (最大 batch_size - 1 プレイアウト) を + # 反映してから解析結果を出力する + if len(self.batch_queue.node_index) > 0: + self.process_mini_batch(board) root = self.node[self.current_root] mode = analysis_query.get("mode", "lz") sys.stdout.write(root.get_analysis(board, mode, self.get_pv_lists))