use std::{sync::atomic::{AtomicU8, Ordering}, time::{SystemTime, UNIX_EPOCH}};
use rand::{rngs::StdRng, SeedableRng};
use crate::{utils, Game, GameEvaluator, Mcts, MctsConfig, MctsError, MctsState, SelectionFunction};
pub struct MctsBatchConfig<const N: usize>{
pub exploration_coef: f64,
pub selection_function: SelectionFunction<N>,
pub seed: Option<u64>,
}
impl<const N: usize> MctsBatchConfig<N>{
pub const DEFAULT: MctsBatchConfig<N> = MctsBatchConfig::<N>{
exploration_coef: MctsConfig::<N>::DEFAULT.exploration_coef,
selection_function: MctsConfig::<N>::DEFAULT.selection_function,
seed: None
};
}
#[allow(type_alias_bounds)]
type History<T: Game<N>, const N: usize> = Vec<(T, f64, [f64; N])>;
pub struct MctsBatch<T: Game<N>, const N: usize>{
instances: Vec<Option<(Mcts<T, N>, History<T, N>)>>,
count: usize,
rand: StdRng,
state: MctsState,
config: MctsConfig<N>
}
impl<T: Game<N>, const N: usize> MctsBatch<T, N>{
#[inline]
pub fn new() -> Self{
MctsBatch::from_config(&MctsBatchConfig::<N>::DEFAULT)
}
#[inline]
pub fn from_config(config: &MctsBatchConfig<N>) -> Self{
MctsBatch {
instances: Vec::new(),
count: 0,
rand: SeedableRng::seed_from_u64(
if let Some(seed) = &config.seed {
*seed
} else {
(SystemTime::now().duration_since(UNIX_EPOCH).expect("").as_nanos()%u64::MAX as u128) as u64
}
),
state: MctsState(AtomicU8::new(MctsState::USABLE)),
config: MctsConfig { exploration_coef: config.exploration_coef, selection_function: config.selection_function }
}
}
#[inline]
pub fn get_count(&self) -> usize{
self.count
}
#[inline]
pub fn get_state(&self) -> u8{
self.state.0.load(Ordering::SeqCst)
}
pub fn populate(&mut self, mut n: 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)); }
}
let mut empty_place: usize = self.instances.len() - self.get_count();
self.count += n;
if empty_place > 0 && n > 0{
for opt in &mut self.instances{
if opt.is_none(){
*opt = Some((Mcts::<T, N>::new(), Vec::new()));
empty_place -= 1;
n -= 1;
}
if n == 0 { return Ok(()); }
if empty_place == 0{ break; }
}
}
for _ in 0..n{
self.instances.push(Some((Mcts::<T, N>::from_config(&self.config), Vec::new())));
}
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
pub fn populate_from_game(&mut self, mut games: Vec<T>) -> 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)); }
}
let mut n: usize = games.len();
let mut empty_place: usize = self.instances.len() - self.get_count();
self.count += n;
if empty_place > 0 && n > 0{
for opt in &mut self.instances{
if opt.is_none(){
*opt = Some((Mcts::<T, N>::from_game_with_config(games.pop().unwrap(), &self.config), Vec::new()));
empty_place -= 1;
n -= 1;
}
if n == 0 { return Ok(()); }
if empty_place == 0{ break; }
}
}
for _ in 0..n{
self.instances.push(Some((Mcts::<T, N>::from_game(games.pop().unwrap()), Vec::new())));
}
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
pub fn clear(&mut self) -> 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)); }
}
self.instances.clear();
self.count = 0;
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
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)); }
}
for opt in &mut self.instances{
if let Some((mcts, _history)) = opt{
if !mcts.get_game().is_finish() && !mcts.is_finish() {
mcts.iterate(evaluator)?;
}
}
}
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
pub fn start_iteration(&mut self) -> Result<Vec<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)); }
}
let mut game_states: Vec<T::State> = Vec::with_capacity(self.get_count());
for opt in &mut self.instances{
if let Some((mcts, _history)) = opt{
if !mcts.get_game().is_finish() && !mcts.is_finish(){
let state = mcts.start_iteration()?;
if !mcts.get_game().is_finish() && !mcts.is_finish() {
game_states.push(state);
}
}
}
}
self.state.0.store(MctsState::AWAITING_SIMULATION, Ordering::Relaxed);
Ok(game_states)
}
pub fn apply_simulation(&mut self, mut evaluations : Vec<(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)); }
}
if evaluations.len() != self.get_count() {
self.state.0.store(MctsState::AWAITING_SIMULATION, Ordering::Relaxed);
return Err(MctsError::InvalidEvaluationCount(self.get_count(), evaluations.len()));
}
for opt in &mut self.instances.iter_mut().rev(){
if let Some((mcts, _history)) = opt{
if mcts.get_game().is_finish() || mcts.is_finish() { continue; }
mcts.apply_simulation(evaluations.pop().unwrap())?;
}
}
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(())
}
pub fn next(&mut self) -> Result<Vec<History<T, N>>, 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)); }
}
let mut result = Vec::new();
for opt in &mut self.instances{
if let Some((mcts, _history)) = &opt{
if mcts.get_game().is_finish() || mcts.is_finish(){
{
let (mcts, history) = opt.as_mut().unwrap();
let mut score: f64 = -mcts.get_game().get_result().unwrap();
for (_game, value, _policy) in history.iter_mut().rev(){
*value=score;
score = -score;
}
}
result.push(opt.take().unwrap().1);
self.count -= 1;
}
else if mcts.count_visit() > 0{
let (mcts, history) = opt.as_mut().unwrap();
let (value, policy) = mcts.get_result();
history.push((mcts.get_game().clone(), value, policy));
let action: usize = utils::sample(&policy, &mut self.rand);
mcts.play(action)?;
}
}
}
self.state.0.store(MctsState::USABLE, Ordering::Relaxed);
Ok(result)
}
}
#[cfg(test)]
mod tests {
use crate::{test_utils::{GameEvaluatorTest2, GameTest}, Game, MctsBatch, MctsError};
#[test]
fn test_batch_iterate_1() -> Result<(), MctsError>{
let mut manager = MctsBatch::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
let a: GameTest = GameTest::new();
let b: GameTest = GameTest::new();
manager.populate_from_game(vec![a, b])?;
let mut i: usize = 0;
let mut result = Vec::new();
while result.is_empty() {
manager.iterate(&evaluator)?;
result = manager.next()?;
i += 1;
}
assert_eq!(result.len(), 2);
assert_eq!(i, 4+1);
Ok(())
}
#[test]
fn test_batch_iterate_2() -> Result<(), MctsError>{
let mut manager = MctsBatch::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
let a: GameTest = GameTest::new();
let mut b: GameTest = GameTest::new();
b.play(4);
manager.populate_from_game(vec![a, b])?;
assert_eq!(manager.get_count(), 2);
let mut i: usize = 0;
let mut result = Vec::new();
while result.is_empty() {
manager.iterate(&evaluator)?;
result = manager.next()?;
i += 1;
}
assert_eq!(result.len(), 1);
assert_eq!(i, 3+1);
assert_eq!(manager.get_count(), 1);
manager.iterate(&evaluator)?;
result = manager.next()?;
assert_eq!(result.len(), 1);
assert_eq!(manager.get_count(), 0);
Ok(())
}
#[test]
fn test_batch_iterate_3() -> Result<(), MctsError>{
let mut manager = MctsBatch::<GameTest, 4>::new();
let evaluator: GameEvaluatorTest2 = GameEvaluatorTest2::new();
manager.populate_from_game(vec![GameTest::new()])?;
let mut i: usize = 0;
let mut result = Vec::new();
while result.is_empty() {
for _ in 0..18{
manager.iterate(&evaluator)?;
}
result = manager.next()?;
i += 1;
}
assert_eq!(result.len(), 1);
assert_eq!(i, 4+1);
assert_eq!((result[0][0].0).get_actions(), [true, true, true, true]);
assert_eq!((result[0][0].1), 1.0);
Ok(())
}
}