use crate::tree::*;
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
pub struct Action(usize);
impl Action{
pub fn new(action: usize) -> Self{
Action(action)
}
pub fn action(&self) -> usize{
self.0
}
}
#[derive(Copy, Clone, Debug, PartialEq)]
enum MctsNodeState{
Active,
Terminal(f32)
}
#[derive(Debug, Copy, Clone, PartialEq)]
struct MctsNode<const N: usize>{
children_scores: [f32; N],
children_visits: [i32; N],
children_policies: [f32; N],
visits: i32,
state: MctsNodeState,
action: Action,
}
pub trait SelectionFunction: Fn(f32, i32, i32, f32, f32) -> f32 {}
impl<T: Fn(f32, i32, i32, f32, f32) -> f32> SelectionFunction for T {}
pub trait ScoreTransformer: FnMut(f32) -> f32 {}
impl<T: FnMut(f32) -> f32> ScoreTransformer for T {}
pub trait ResettableBuffer {
fn push(&mut self, value: Action);
fn clear(&mut self);
}
#[macro_export]
macro_rules! impl_resettable_buffer {
($type:ty) => {
impl ResettableBuffer for $type {
fn push(&mut self, value: Action) {
self.push(value);
}
fn clear(&mut self) {
self.clear();
}
}
};
}
impl_resettable_buffer!(Vec<Action>);
impl<const N: usize> MctsNode<N> {
fn new(policies: [f32; N], action: Action, state: MctsNodeState) -> Self {
MctsNode{
children_scores: [0.; N],
children_visits: [0; N],
children_policies: policies,
visits: 0,
action,
state
}
}
fn best_child(&self, score_f: &impl SelectionFunction, c: f32) -> usize {
self.children_scores.iter()
.zip(self.children_visits.iter())
.zip(self.children_policies.iter())
.enumerate()
.filter_map(
|(i, ((c_scores, c_visits), policy))| {
if *policy <= 0. { None } else { Some((i, score_f(*c_scores, *c_visits, self.visits, *policy, c))) }
}
).max_by(
|(_a, a), (_b, b)| { a.total_cmp(b) }
).expect("Error during the selection of the best action.").0
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
pub struct MctsNodeId(NodeId);
pub struct Engine<const N: usize> {
tree: Tree<MctsNode<N>, N>
}
#[derive(Copy, Clone, Debug, PartialEq)]
pub enum SelectionResult {
Empty,
Active(MctsNodeId, Action),
Terminal(MctsNodeId, f32)
}
#[derive(Copy, Clone, Debug, PartialEq)]
pub enum StateEvaluation<const N: usize> {
Active(f32, [f32; N]),
Terminal(f32)
}
impl<const N: usize> StateEvaluation<N>{
pub fn score(&self) -> f32{
match self {
StateEvaluation::Active(score, _) => *score,
StateEvaluation::Terminal(score) => *score
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum MctsEngineError{
UnknownError,
InvalidMctsNode,
ChildAlreadyExists,
InvalidAction,
SelectionIsTerminal
}
impl<const N: usize> Engine<N> {
pub fn new() -> Self {
Engine::<N> {
tree: Tree::<MctsNode<N>, N>::new()
}
}
pub fn with_capacity(capacity: usize) -> Self {
Engine::<N> {
tree: Tree::<MctsNode<N>, N>::with_capacity(capacity)
}
}
pub fn select(&self, path_out: &mut impl ResettableBuffer, score_f: &impl SelectionFunction, c: f32) -> SelectionResult{
path_out.clear();
if let Some(mut current_id) = self.tree.root(){
loop {
let current = self.tree.get(current_id).unwrap();
match current.data().state {
MctsNodeState::Terminal(score) => {
return SelectionResult::Terminal(MctsNodeId(current_id), score);
}
MctsNodeState::Active => {
let action = current.data().best_child(score_f, c);
path_out.push(Action(action));
if let Some(child_id) = current.child(action){
current_id = child_id;
}
else{
return SelectionResult::Active(MctsNodeId(current_id), Action(action));
}
}
}
}
}
else{ SelectionResult::Empty }
}
pub fn expand(&mut self, evaluation: StateEvaluation<N>, selection: SelectionResult) -> Result<MctsNodeId, MctsEngineError>{
let (state, policies) = match evaluation {
StateEvaluation::Active(_, policies) => (MctsNodeState::Active, policies),
StateEvaluation::Terminal(score) => (MctsNodeState::Terminal(score), [0.; N])
};
match selection {
SelectionResult::Empty => {
match self.tree.set_root(MctsNode::new(policies, Action(0), state)) {
Err(TreeError::RootAlreadyExists) => Err(MctsEngineError::ChildAlreadyExists),
Err(_) => Err(MctsEngineError::UnknownError),
Ok(node) => Ok(MctsNodeId(node)),
}
},
SelectionResult::Active(child_id, action) => {
match self.tree.add(child_id.0, action.0, MctsNode::new(policies, action, state)) {
Err(TreeError::ParentDoesntExist) => Err(MctsEngineError::InvalidMctsNode),
Err(TreeError::ChildAlreadyExists) => Err(MctsEngineError::ChildAlreadyExists),
Err(_) => Err(MctsEngineError::UnknownError),
Ok(node) => Ok(MctsNodeId(node))
}
},
SelectionResult::Terminal(_, _) => Err(MctsEngineError::SelectionIsTerminal)
}
}
pub fn backpropagate(&mut self, node: MctsNodeId, score: f32, mut score_updater: impl ScoreTransformer) -> Result<(), MctsEngineError> {
let mut current = Some(node.0);
let mut score = score;
while let Some(node) = current {
let node = self.tree.get_mut(node).map_err(|_| MctsEngineError::InvalidMctsNode)?;
node.data_mut().visits += 1;
let action = node.data().action.0;
current = node.parent();
if let Some(parent) = current {
let parent = self.tree.get_mut(parent).map_err(|_| MctsEngineError::InvalidMctsNode)?;
parent.data_mut().children_visits[action] += 1;
parent.data_mut().children_scores[action] += score;
}
score = score_updater(score);
}
Ok(())
}
pub fn update(&mut self, evaluation: StateEvaluation<N>, selection: SelectionResult, score_updater: impl ScoreTransformer) -> Result<(), MctsEngineError> {
let score = evaluation.score();
let node = match selection {
SelectionResult::Empty => self.expand(evaluation, selection)?,
SelectionResult::Active(_child_id, _action) => self.expand(evaluation, selection)?,
SelectionResult::Terminal(child_id, _score) => child_id
};
self.backpropagate(node, score, score_updater)?;
Ok(())
}
pub fn scores(&self) -> [i32; N]{
if let Some(root) = self.tree.root() {
self.tree.get(root).unwrap().data().children_visits
}
else {
[0; N]
}
}
pub fn commit_action(&mut self, action: Action) {
if let Some(root) = self.tree.root() {
let child = self.tree.child(root, action.0).unwrap();
if let Some(new_root_id) = child {
self.tree.move_root_to(new_root_id).unwrap();
let new_root_data = self.tree.data(new_root_id).unwrap();
if self.tree.allocated_nodes() > 2*new_root_data.visits as usize {
self.tree.compact();
}
}
else{
self.tree.clear();
}
}
}
}
#[cfg(test)]
mod tests {
use crate::{puct, negate_score};
use super::*;
#[test]
fn test_action_creation() {
let action = Action::new(42);
assert_eq!(action.action(), 42);
}
#[test]
fn test_action_equality() {
let a1 = Action::new(10);
let a2 = Action::new(10);
let a3 = Action::new(11);
assert_eq!(a1, a2, "Two actions with the same index should be equal");
assert_ne!(a1, a3, "Actions with different indices should not be equal");
}
#[test]
fn test_action_hashability() {
use std::collections::HashSet;
let mut set = HashSet::new();
let a1 = Action::new(5);
let a2 = Action::new(5);
set.insert(a1);
assert!(set.contains(&a2), "Action should be hashable and work in a HashSet");
assert_eq!(set.len(), 1, "HashSet should handle identical Actions correctly");
}
fn dummy_select(_score: f32, _node_visits: i32, _parent_visits: i32, policy: f32, _c: f32) -> f32 {
policy
}
fn identity_score(s: f32) -> f32 {
s
}
const N: usize = 2;
#[test]
fn test_engine_initialization() {
let engine = Engine::<N>::new();
assert!(engine.tree.root().is_none());
}
#[test]
fn test_select_on_empty_tree_returns_empty() {
let engine = Engine::<N>::new();
let mut path = Vec::new();
let selection = engine.select(&mut path, &dummy_select, 1.0);
assert_eq!(selection, SelectionResult::Empty);
assert!(path.is_empty());
}
#[test]
fn test_update_empty_creates_root() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let selection = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(1.0, [0.7, 0.3]);
let result = engine.update(evaluation, selection, negate_score);
assert!(result.is_ok());
let root_id = engine.tree.root().unwrap();
let root = engine.tree.get(root_id).unwrap();
assert_eq!(root.data().visits, 1);
assert_eq!(root.data().state, MctsNodeState::Active);
}
#[test]
fn test_select_active_node_and_expand_child() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let selection = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.7, 0.3]);
engine.update(evaluation, selection, negate_score).unwrap();
let selection2 = engine.select(&mut path, &dummy_select, 1.0);
match selection2 {
SelectionResult::Active(id, action) => {
assert_eq!(id.0, engine.tree.root().unwrap());
assert_eq!(action, Action(0));
},
_ => panic!("Expected Active selection"),
}
assert_eq!(path, vec![Action(0)]);
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
let result = engine.update(evaluation, selection2, negate_score);
assert!(result.is_ok());
}
#[test]
fn test_backpropagate_values_correctly() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel_empty = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.7, 0.3]);
engine.update(evaluation, sel_empty, negate_score).unwrap();
let sel_active = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
engine.update(evaluation, sel_active, negate_score).unwrap();
let root_id = engine.tree.root().unwrap();
let root = engine.tree.get(root_id).unwrap();
assert_eq!(root.data().visits, 2);
assert_eq!(root.data().children_visits[0], 1);
assert_eq!(root.data().children_scores[0], 1.0);
}
#[test]
fn test_terminal_node_selection_and_state() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel_empty = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Terminal(42.0);
engine.update(evaluation, sel_empty, identity_score).unwrap();
let root_id = engine.tree.root().unwrap();
let root = engine.tree.get(root_id).unwrap();
assert_eq!(root.data().state, MctsNodeState::Terminal(42.0));
let sel_terminal = engine.select(&mut path, &dummy_select, 1.0);
match sel_terminal {
SelectionResult::Terminal(id, score) => {
assert_eq!(id.0, root_id);
assert_eq!(score, 42.0);
},
_ => panic!("Expected Terminal selection"),
}
}
#[test]
#[should_panic(expected = "Error during the selection of the best action.")]
fn test_best_child_panics_if_no_legal_moves() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel_empty = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.0, 0.0]);
engine.update(evaluation, sel_empty, negate_score).unwrap();
engine.select(&mut path, &dummy_select, 1.0);
}
#[test]
fn test_expand_on_terminal_returns_error() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel_empty = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Terminal(1.0);
engine.update(evaluation, sel_empty, identity_score).unwrap();
let sel_terminal = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
let result = engine.expand(evaluation, sel_terminal);
assert!(matches!(result, Err(MctsEngineError::SelectionIsTerminal)));
}
#[test]
fn test_adding_existing_child_returns_error() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel_empty = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.7, 0.3]);
engine.update(evaluation, sel_empty, negate_score).unwrap();
let sel_active = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
engine.update(evaluation, sel_active, negate_score).unwrap();
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
let result = engine.update(evaluation, sel_active, negate_score);
assert!(matches!(result, Err(MctsEngineError::ChildAlreadyExists)));
}
#[test]
fn test_scores_empty_engine() {
let engine = Engine::<N>::new();
assert_eq!(engine.scores(), [0; N]);
}
#[test]
fn test_scores_after_root_expansion() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.5, 0.5]);
engine.update(evaluation, sel, identity_score).unwrap();
assert_eq!(engine.scores(), [0; N]);
}
#[test]
fn test_scores_after_multiple_updates() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel_empty = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.7, 0.3]);
engine.update(evaluation, sel_empty, identity_score).unwrap();
for _ in 0..3 {
let sel = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
engine.update(evaluation, sel, identity_score).unwrap();
}
let sel = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(1.0, [0.5, 0.5]);
engine.update(evaluation, sel, identity_score).unwrap();
let expected = [4, 0];
assert_eq!(engine.scores(), expected);
}
#[test]
#[allow(unused_assignments)]
fn test_scores_are_independent_copies() {
let mut engine = Engine::<N>::new();
let mut path = Vec::new();
let sel = engine.select(&mut path, &dummy_select, 1.0);
let evaluation = StateEvaluation::Active(0.0, [0.5, 0.5]);
engine.update(evaluation, sel, identity_score).unwrap();
let mut scores = engine.scores();
scores[0] = 999;
assert_ne!(engine.scores()[0], 999);
assert_eq!(engine.scores()[0], 0);
}
#[test]
fn test_scnerario_1(){
let c = std::f32::consts::SQRT_2;
let mut path = Vec::new();
let mut engine = Engine::<3>::new();
let selection = engine.select(&mut path, &puct, c);
let evaluation = StateEvaluation::Active(0.1, [0.5, 0.3, 0.2]);
let node = engine.expand(evaluation, selection).unwrap();
engine.backpropagate(node, 0.1, negate_score).unwrap();
assert_eq!(engine.scores(), [0, 0, 0]);
assert_eq!(path.len(), 0);
let selection = engine.select(&mut path, &puct, c);
let evaluation = StateEvaluation::Active(0.5, [1., 0., 0.]);
let node = engine.expand(evaluation, selection).unwrap();
engine.backpropagate(node, 0.5, negate_score).unwrap();
assert_eq!(engine.scores(), [1, 0, 0]);
assert_eq!(path.as_slice(), &[Action(0)]);
let selection = engine.select(&mut path, &puct, c);
let evaluation = StateEvaluation::Active(1., [1., 0., 0.]);
let node = engine.expand(evaluation, selection).unwrap();
engine.backpropagate(node, 1.0, negate_score).unwrap();
assert_eq!(engine.scores(), [2, 0, 0]);
assert_eq!(path.as_slice(), &[Action(0), Action(0)]);
let selection = engine.select(&mut path, &puct, c);
let evaluation = StateEvaluation::Terminal(-0.9);
let node = engine.expand(evaluation, selection).unwrap();
engine.backpropagate(node, -0.9, negate_score).unwrap();
assert_eq!(engine.scores(), [2, 1, 0]);
assert_eq!(path.as_slice(), &[Action(1)]);
let selection = engine.select(&mut path, &puct, c);
let evaluation = StateEvaluation::Terminal(-0.9);
let node = engine.expand(evaluation, selection).unwrap();
engine.backpropagate(node, -0.9, negate_score).unwrap();
assert_eq!(engine.scores(), [2, 1, 1]);
assert_eq!(path.as_slice(), &[Action(2)]);
}
fn setup_engine() -> Engine<3> {
let mut engine = Engine::<3>::new();
let mut path = Vec::new();
let sel = engine.select(&mut path, &|_,_,_,_,_| 0.0, 1.0);
let evaluation = StateEvaluation::Active(0., [1., 0., 0.]);
let node = engine.expand(evaluation, sel).unwrap();
engine.backpropagate(node, 0.0, negate_score).unwrap();
let sel = engine.select(&mut path, &|_,_,_,_,_| 0.0, 1.0);
let evaluation = StateEvaluation::Active(1., [1., 0., 0.]);
let node = engine.expand(evaluation, sel).unwrap();
engine.backpropagate(node, 1.0, negate_score).unwrap();
engine
}
#[test]
fn test_commit_action_success() {
let mut engine = setup_engine();
engine.tree.root().unwrap();
engine.commit_action(Action(0));
let new_root = engine.tree.root().unwrap();
assert_eq!(engine.tree.data(new_root).unwrap().visits, 1);
assert_eq!(engine.tree.get(new_root).unwrap().data().action, Action(0));
}
#[test]
fn test_commit_action_unexplored_leads_to_clear() {
let mut engine = setup_engine();
engine.commit_action(Action(2));
assert!(engine.tree.root().is_none());
}
#[test]
fn test_commit_action_no_root_does_nothing() {
let mut engine = Engine::<3>::new();
engine.commit_action(Action(0));
assert!(engine.tree.root().is_none());
}
#[test]
#[should_panic(expected = "index out of bounds")]
fn test_commit_action_out_of_bounds_panics() {
let mut engine = setup_engine();
engine.commit_action(Action(5));
}
#[test]
fn test_compact_trigger() {
let mut engine = setup_engine();
let root = engine.tree.root().unwrap();
engine.tree.get_mut(root).unwrap().data_mut().visits = 100;
let initial_nodes = engine.tree.allocated_nodes();
engine.commit_action(Action(0));
assert!(engine.tree.allocated_nodes() <= initial_nodes);
}
}