use core::f64;
use std::sync::{atomic::{AtomicU8, Ordering}, Arc, Mutex};
use crate::{Game, GameEvaluator, Node, NodeRef};
const INFINITY : f64 = 1e300;
struct MctsNodeData<const N: usize>{
policy: [f64; N],
mask: [bool; N],
score: f64,
n: usize,
finish: bool
}
impl<const N: usize> MctsNodeData<N>{
pub fn new() -> Self{
MctsNodeData {
policy: [1./N as f64; N],
mask: [false; N],
score: 0.0,
n: 0,
finish: false
}
}
#[inline]
pub fn get_value(&self) -> f64{
if self.n != 0 { self.score / self.n as f64 } else { 0.0 }
}
#[inline]
pub fn get_n(&self) -> usize{
self.n
}
#[inline]
pub fn get_policy(&self, index: usize) -> f64{
self.policy[index]
}
#[inline]
pub fn get_mask(&self, index: usize) -> bool{
self.mask[index]
}
#[inline]
pub fn is_finish(&self) -> bool{
self.finish
}
#[inline]
pub fn add_score(&mut self, score: f64){
self.score += score;
self.n += 1;
}
}
type MctsNode<const N: usize> = Node<MctsNodeData<N>, N>;
type MctsNodeRef<const N: usize> = NodeRef<MctsNodeData<N>, N>;
#[derive(Debug)]
pub struct MctsState(pub AtomicU8);
impl MctsState{
pub const USABLE: u8 = 0;
pub const AWAITING_SIMULATION: u8 = 1;
pub const LOCKED: u8 = 2;
}
#[derive(Debug)]
pub enum MctsError{
InvalidState(u8),
InvalidEvaluationCount(usize, usize),
SearchAlreadyOver,
ActionOutOfRange(usize, usize),
InvalidAction(usize),
UnexploredAction
}
pub type SelectionFunction<const N: usize> = fn(value: f64, policy: f64, n_visits: f64, parent_n_visits: f64, exploration_coef: f64) -> f64;
pub struct MctsConfig<const N: usize>{
pub exploration_coef: f64,
pub selection_function: SelectionFunction<N>
}
impl<const N: usize> MctsConfig<N>{
pub const DEFAULT: MctsConfig<N> = MctsConfig{
exploration_coef: std::f64::consts::SQRT_2,
selection_function: default_selection_score::<N>
};
}
pub struct Mcts<T: Game<N>, const N: usize>{
game: T,
root: Option<MctsNodeRef<N>>,
coef: f64,
state: MctsState,
latent: Option<(T, MctsNodeRef<N>)>,
selection_function: SelectionFunction<N>
}
pub fn default_selection_score<const N: usize>(value: f64, policy: f64, n_visits: f64, parent_n_visits: f64, exploration_coef: f64) -> f64{
value + exploration_coef * policy * parent_n_visits.sqrt() / (1.+n_visits)
}
pub fn ucb1<const N: usize>(value: f64, policy: f64, n_visits: f64, parent_n_visits: f64, exploration_coef: f64) -> f64{
value + exploration_coef * policy * (parent_n_visits.ln() / n_visits).sqrt()
}
impl<T: Game<N>, const N: usize> Mcts<T, N>{
pub const VICTORY_SCORE: f64 = 1.0;
pub const DEFEAT_SCORE: f64 = -1.0;
pub const EQUALITY_SCORE: f64 = 0.0;
#[inline]
pub fn new() -> Self{
Self::from_config(&MctsConfig::DEFAULT)
}
#[inline]
pub fn from_config(config: &MctsConfig<N>) -> Self{
Self::from_game_with_config(T::new(), config)
}
#[inline]
pub fn from_game(game: T) -> Self{
Mcts::from_game_with_config(game, &MctsConfig::DEFAULT)
}
#[inline]
pub fn from_game_with_config(game: T, config: &MctsConfig<N>) -> Self{
Mcts {
game: game,
root: None,
coef: config.exploration_coef,
state: MctsState(AtomicU8::new(MctsState::USABLE)),
latent: None,
selection_function: config.selection_function
}
}
pub fn get_game(&self) -> &T{
&self.game
}
#[inline]
pub fn get_state(&self) -> u8{
self.state.0.load(Ordering::SeqCst)
}
#[inline]
fn get_selection_score(&self, node: &MctsNode<N>, index: usize) -> f64{
if let Some(child) = node.get_child(index){
let node_child = &*child.lock().unwrap();
if node_child.get().is_finish() {
-INFINITY
}
else{
(self.selection_function) (
node_child.get().get_value(),
node.get().get_policy(index),
node_child.get().get_n() as f64,
node.get().get_n() as f64,
self.coef
)
}
}
else{
if node.get().get_mask(index) && !node.get().is_finish() { INFINITY * (1. + node.get().get_policy(index))} else { -INFINITY }
}
}
#[inline]
fn selection(&self) -> (Option<MctsNodeRef<N>>, usize, T){
let mut game: T = self.game.clone();
let mut node = match &self.root {
Some(root) => Arc::clone(root),
None => return (None, 0, game)
};
loop {
let scores : [f64; N] = std::array::from_fn(|index| self.get_selection_score(&node.lock().unwrap(), index));
let index = scores.iter().enumerate().max_by(|a, b| (a.1).total_cmp(b.1)).unwrap().0;
game.play(index);
if node.lock().unwrap().get_child(index).is_none() {
return (Some(node), index, game);
}
let next = node.lock().unwrap().get_child(index).unwrap();
node = next;
}
}
#[inline]
fn expansion(&mut self, node: &Option<MctsNodeRef<N>>, index: usize) -> MctsNodeRef<N>{
if let Some(node) = node{
Node::add_child(&node, index, MctsNodeData::new())
}
else{
let node_ref = Arc::new(Mutex::new(
MctsNode::new(None, MctsNodeData::new())
));
self.root = Some(Arc::clone(&node_ref));
node_ref
}
}
#[inline]
fn simulation(&mut self, node: &mut MctsNode<N>, game: &T, evaluator: &dyn GameEvaluator<T, N>){
let data = node.get_mut();
if let Some(score) = game.get_result(){
data.add_score(-score);
data.finish = true;
}
else{
let (score, policy) = evaluator.evaluate(game.get_state());
data.add_score(-score);
data.policy=policy;
data.mask=game.get_actions();
}
}
#[inline]
fn simulation_from_data(&mut self, node: &mut MctsNode<N>, game: &T, evaluation: (f64, [f64; N])){
let data = node.get_mut();
if let Some(score) = game.get_result(){
data.add_score(-score);
data.finish = true;
}
else{
let (score, policy) = evaluation;
data.add_score(-score);
data.policy=policy;
data.mask=game.get_actions();
}
}
#[inline]
fn backpropagation(&mut self, node_ref: &MctsNodeRef<N>){
let mut current_ref_opt: Option<Arc<Mutex<Node<MctsNodeData<N>, N>>>>;
let mut score: f64;
let mut finish: bool;
{
let node = &*node_ref.lock().unwrap();
score = node.get().get_value();
finish = node.get().is_finish();
current_ref_opt = node.get_parent();
}
while let Some(current_ref) = current_ref_opt {
score = -score;
let current = &mut *current_ref.lock().unwrap();
if !finish{
current.get_mut().add_score(score);
}
else {
if score == -Self::VICTORY_SCORE {
current.get_mut().score = current.get().n as f64 * score;
current.get_mut().finish = true;
}
else{
let mut max_score = f64::MIN;
let current_is_finish = (0..N).all(|index|{
!current.get().get_mask(index) || {
if let Some(child_ref) = current.get_child(index){
let child = &*child_ref.lock().unwrap();
let value = child.get().get_value();
if value > max_score {
max_score = value;
}
child.get().is_finish()
}
else{ false }
}
});
if current_is_finish {
current.get_mut().score = current.get().n as f64 * -max_score;
current.get_mut().finish = true;
}
else{
current.get_mut().add_score(score);
finish = false;
}
}
}
current_ref_opt = current.get_parent();
}
}
#[inline]
pub fn iterate(&mut self, evaluator: &dyn GameEvaluator<T, N>) -> Result<(), MctsError> {
match self.state.0.compare_exchange(MctsState::USABLE, MctsState::LOCKED, Ordering::SeqCst, Ordering::SeqCst){
Ok(_) => {}
Err(current_state) => { return Err(MctsError::InvalidState(current_state)); }
}
if let Some(root) = self.root.as_ref(){
if root.lock().unwrap().get().is_finish(){
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
return Ok(());
}
}
let (node, index, game) = self.selection();
let child_ref = self.expansion(&node, index);
self.simulation(&mut *child_ref.lock().unwrap(), &game, evaluator);
self.backpropagation(&child_ref);
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
#[inline]
pub fn start_iteration(&mut self) -> Result<T::State, MctsError>{
match self.state.0.compare_exchange(MctsState::USABLE, MctsState::LOCKED, Ordering::SeqCst, Ordering::SeqCst){
Ok(_) => {}
Err(current_state) => { return Err(MctsError::InvalidState(current_state)); }
}
if let Some(root) = self.root.as_ref(){
if root.lock().unwrap().get().is_finish(){
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
return Err(MctsError::SearchAlreadyOver)
}
}
let (node, index, game) = self.selection();
let child_ref = self.expansion(&node, index);
let game_state = game.get_state();
self.latent = Some((game, child_ref));
self.state.0.store(MctsState::AWAITING_SIMULATION, Ordering::Relaxed);
Ok(game_state)
}
#[inline]
pub fn apply_simulation(&mut self, evaluation : (f64, [f64; N])) -> Result<(), MctsError>{
match self.state.0.compare_exchange(MctsState::AWAITING_SIMULATION, MctsState::LOCKED, Ordering::SeqCst, Ordering::SeqCst){
Ok(_) => {}
Err(current_state) => { return Err(MctsError::InvalidState(current_state)); }
}
let (game, child_ref) = self.latent.take().unwrap();
self.simulation_from_data(&mut *child_ref.lock().unwrap(), &game, evaluation);
self.backpropagation(&child_ref);
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
#[inline]
pub fn is_finish(&self) -> bool{
if let Some(root) = &self.root{
root.lock().unwrap().get().is_finish()
}
else{ false }
}
#[inline]
pub fn get_score(&self) -> f64{
if let Some(root) = &self.root {
-root.lock().unwrap().get().get_value()
}
else{ Self::EQUALITY_SCORE }
}
#[inline]
fn statistics_from_root(root: &MctsNode<N>) -> [f64; N]{
let scores: [f64; N] = std::array::from_fn(|index|{
if let Some(child_ref) = root.get_child(index){
let score = child_ref.lock().unwrap().get().get_value();
if score != 1.{ (score + 1.) / (1. - score) + f64::MIN_POSITIVE } else{ f64::MAX / N as f64 }
}
else if root.get().get_mask(index){ f64::MIN_POSITIVE }
else { 0.0 }
});
let total: f64 = scores.iter().sum();
let total: f64 = if total == 0.0 { f64::MIN_POSITIVE } else{ total };
let scores = scores.map(|x| x/total);
scores
}
#[inline]
pub fn get_statistics(&self) -> [f64; N]{
if let Some(root_ref) = &self.root {
Self::statistics_from_root(&*root_ref.lock().unwrap())
}
else{
[1./N as f64; N]
}
}
#[inline]
pub fn get_result(&self) -> (f64, [f64; N]){
if let Some(root_ref) = &self.root {
let root = &*root_ref.lock().unwrap();
(-root.get().get_value(), Self::statistics_from_root(root))
}
else{
(0.0, [1./N as f64; N])
}
}
#[inline]
pub fn count_visit(&self) -> usize{
if let Some(root_ref) = &self.root{
let root = &*root_ref.lock().unwrap();
root.get().get_n()
}
else { 0 }
}
#[inline]
pub fn play(&mut self, action: usize) -> Result<(), MctsError>{
match self.state.0.compare_exchange(MctsState::USABLE, MctsState::LOCKED, Ordering::SeqCst, Ordering::SeqCst){
Ok(_) => {}
Err(current_state) => { return Err(MctsError::InvalidState(current_state)); }
}
if action >= N {
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
return Err(MctsError::ActionOutOfRange(action, N));
}
let new_root;
if let Some(root_ref) = &self.root{
let root = &*root_ref.lock().unwrap();
if let Some(child_ref) = root.get_child(action){
{
let child = &mut *child_ref.lock().unwrap();
child.detach();
}
new_root = Some(child_ref);
}
else if root.get().get_mask(action){
new_root = None;
}
else{
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
return Err(MctsError::InvalidAction(action));
}
}
else{
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
return Err(MctsError::UnexploredAction);
}
self.game.play(action);
self.root = new_root;
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
}
#[cfg(test)]
mod tests {
use crate::{test_utils::{compare_array, GameEvaluatorTest, GameEvaluatorTest2, GameTest}, Game, GameEvaluator, Mcts, MctsError};
#[test]
fn test_selection_empty(){
let mcts = Mcts::<GameTest, 4>::new();
assert!(mcts.selection().0.is_none())
}
#[test]
fn test_expansion_empty(){
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
mcts.expansion(&None, 0);
assert!(mcts.root.is_some());
}
#[test]
fn test_selection_root(){
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
mcts.expansion(&None, 0);
let (node, _index, _game) = mcts.selection();
let node = &*node.as_ref().unwrap().lock().unwrap();
assert!(node.is_root());
assert!(node.get_child(0).is_none());
assert!(node.get_child(1).is_none());
assert!(node.get_child(2).is_none());
assert!(node.get_child(3).is_none());
}
#[test]
fn test_simulation_root(){
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest = GameEvaluatorTest::new();
let (node, index, game) = mcts.selection();
let child = mcts.expansion(&node, index);
let child = &mut *child.lock().unwrap();
mcts.simulation(child, &game, &evaluator);
assert!(child.is_root());
assert!(child.get_child(0).is_none());
assert!(child.get_child(1).is_none());
assert!(child.get_child(2).is_none());
assert!(child.get_child(3).is_none());
assert!(!child.get().is_finish());
assert_eq!(child.get().get_policy(0), 0.2);
assert_eq!(child.get().get_policy(1), 0.7);
assert_eq!(child.get().get_policy(2), 0.06);
assert_eq!(child.get().get_policy(3), 0.04);
}
#[test]
fn test_iteration_empty() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest = GameEvaluatorTest::new();
mcts.iterate(&evaluator)?;
{
let node_ref = mcts.root.as_ref().unwrap();
let node = &*node_ref.lock().unwrap();
assert!(node.is_root());
assert!(!node.get().is_finish());
assert_eq!(node.get().mask, [true, true, true, true]);
assert_eq!(node.get().policy, [0.2, 0.7, 0.06, 0.04]);
assert_eq!(node.get().score, 0.0);
assert_eq!(node.get().n, 1);
}
assert!(!mcts.is_finish());
Ok(())
}
#[test]
fn test_iteration_root_with_no_child() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest = GameEvaluatorTest::new();
mcts.iterate(&evaluator)?;
mcts.iterate(&evaluator)?;
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert_eq!(root.get().score, 0.2);
assert_eq!(root.get().n, 2);
assert!(root.get_child(0).is_none());
assert!(root.get_child(1).is_some());
assert!(root.get_child(2).is_none());
assert!(root.get_child(3).is_none());
}
mcts.iterate(&evaluator)?;
mcts.iterate(&evaluator)?;
mcts.iterate(&evaluator)?;
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert_eq!(root.get().score, 0.0);
assert_eq!(root.get().n, 5);
assert!(root.get_child(0).is_some());
assert!(root.get_child(1).is_some());
assert!(root.get_child(2).is_some());
assert!(root.get_child(3).is_some());
}
assert!(!mcts.is_finish());
Ok(())
}
#[test]
fn test_iteration_root_with_child() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest = GameEvaluatorTest::new();
for _ in 0..7{
mcts.iterate(&evaluator)?;
}
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert_eq!(root.get().score, 0.0);
assert_eq!(root.get().n, 7);
}
assert!(!mcts.is_finish());
Ok(())
}
#[test]
fn test_iteration_victory_1() -> Result<(), MctsError>{
let mut game: GameTest = GameTest::new();
game.play(3);
game.play(1);
game.play(2);
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::from_game(game);
let evaluator: GameEvaluatorTest = GameEvaluatorTest::new();
mcts.iterate(&evaluator)?;
mcts.iterate(&evaluator)?;
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert!(root.get().is_finish());
assert_eq!(root.get().get_value(), 1.0);
}
assert!(mcts.is_finish());
assert_eq!(mcts.get_result(), (-1.0, [1.0, 0., 0., 0.]));
Ok(())
}
#[test]
fn test_iteration_victory_2() -> Result<(), MctsError>{
let mut game: GameTest = GameTest::new();
game.play(0);
game.play(3);
game.play(1);
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::from_game(game);
let evaluator: GameEvaluatorTest = GameEvaluatorTest::new();
mcts.iterate(&evaluator)?;
mcts.iterate(&evaluator)?;
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert!(root.get().is_finish());
assert_eq!(root.get().get_value(), -1.0);
}
assert!(mcts.is_finish());
assert_eq!(mcts.get_result(), (1.0, [0., 0., 1., 0.]));
Ok(())
}
#[test]
fn test_iteration_end_1() -> Result<(), MctsError>{
let mut game: GameTest = GameTest::new();
game.play(3);
game.play(1);
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::from_game(game);
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
for _ in 0..4{
mcts.iterate(&evaluator)?;
}
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert!(root.get().is_finish());
assert_eq!(root.get().get_value(), -1.0);
}
assert!(mcts.is_finish());
let result = mcts.get_result();
assert_eq!(result.0, 1.0);
assert!(compare_array(&result.1, &[0., 0., 1., 0.]));
Ok(())
}
#[test]
fn test_iteration_total() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
for _ in 0..18{
mcts.iterate(&evaluator)?;
}
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert!(root.get().is_finish());
assert_eq!(root.get().get_value(), -1.0);
}
assert!(mcts.is_finish());
let result = mcts.get_result();
assert_eq!(result.0, 1.0);
assert!(compare_array(&result.1, &[0., 0., 0., 1.]));
Ok(())
}
#[test]
fn test_iteration_total_2() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
for _ in 0..18{
let game = mcts.start_iteration()?;
mcts.apply_simulation(evaluator.evaluate(game))?;
}
{
let root_ref = mcts.root.as_ref().unwrap();
let root = &*root_ref.lock().unwrap();
assert!(root.get().is_finish());
assert_eq!(root.get().get_value(), -1.0);
}
assert!(mcts.is_finish());
let result = mcts.get_result();
assert_eq!(result.0, 1.0);
assert!(compare_array(&result.1, &[0., 0., 0., 1.]));
Ok(())
}
#[test]
fn test_play_and_count_visit() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
assert_eq!(mcts.count_visit(), 0);
for _ in 0..6{
mcts.iterate(&evaluator)?;
}
assert_eq!(mcts.count_visit(), 6);
mcts.play(3)?;
assert_eq!(mcts.count_visit(), 2);
assert!(mcts.root.unwrap().lock().unwrap().is_root());
Ok(())
}
#[test]
fn test_play_empty() -> Result<(), MctsError>{
let mut mcts: Mcts<GameTest, 4> = Mcts::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
mcts.iterate(&evaluator)?;
mcts.play(3)?;
Ok(())
}
}