| 13 | |
| 14 | |
| 15 | class MCTSNode(BaseNode): |
| 16 | |
| 17 | prior: float = 1.0 |
| 18 | c_puct: float = 1.5 |
| 19 | |
| 20 | __visit_count: int = PrivateAttr(default=0) |
| 21 | __value_sum: float = PrivateAttr(default=0) |
| 22 | |
| 23 | def q_value(self) -> float: |
| 24 | if self.__visit_count == 0: |
| 25 | return 0 |
| 26 | return self.__value_sum / self.__visit_count |
| 27 | |
| 28 | def visit_count(self) -> int: |
| 29 | return self.__visit_count |
| 30 | |
| 31 | def update_visit_count(self, count: int) -> None: |
| 32 | self.__visit_count = count |
| 33 | |
| 34 | def update(self, value: float) -> None: |
| 35 | # init value |
| 36 | if self.value == -100: |
| 37 | self.value = value |
| 38 | self.__visit_count += 1 |
| 39 | self.__value_sum += value |
| 40 | |
| 41 | def update_recursive(self, value: float, start_node: Type[BaseNode]) -> None: |
| 42 | self.update(value) |
| 43 | if self.tag == start_node.tag: |
| 44 | return |
| 45 | self.parent.update_recursive(value, start_node) |
| 46 | |
| 47 | def puct(self) -> float: |
| 48 | q_value = self.q_value() if self.visit_count() > 0 else 0 |
| 49 | u_value = self.c_puct * self.prior * np.sqrt(self.parent.visit_count()) / (1 + self.visit_count()) |
| 50 | return q_value + u_value |
| 51 | |