MCPcopy Create free account
hub / github.com/MARIO-Math-Reasoning/Super_MARIO / MCTSNode

Class MCTSNode

mcts_math/nodes/mcts_node.py:15–50  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

13
14
15class 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

Callers 1

create_nodeMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected