use rayon::prelude::*;
use rayon::ThreadPoolBuilder;
use anyhow::Result;
use crate::rl::env::Env;
use crate::nn::policy::Policy;
use super::solve::solve;
pub fn evaluate(
env: &Box<dyn Env>,
policy: &Policy,
num_episodes: usize,
deterministic: bool,
num_searches: usize,
num_mcts_searches: usize,
_seed: usize, C: f32,
max_expand_depth: usize,
num_cores: usize,
) -> Result<(f32, f32)> {
let mut env = env.clone();
if num_cores <= 1 {
let mut successes = 0.0;
let mut rewards = 0.0;
for _ in 0..num_episodes {
env.reset();
let ((success, reward), _path) = solve(
&env,
&policy,
deterministic,
num_searches,
num_mcts_searches,
C,
max_expand_depth,
);
successes += success;
rewards += reward;
}
Ok((successes/(num_episodes as f32), rewards/(num_episodes as f32)))
} else {
let pool = ThreadPoolBuilder::new()
.num_threads(num_cores)
.build()
.unwrap();
let (successes, total_vals) = pool.install(|| {
(0..num_episodes)
.into_par_iter()
.map(|_| {
let mut env = env.clone();
env.reset();
let ((success, reward), _path) = solve(
&env,
&policy,
deterministic,
num_searches,
num_mcts_searches,
C,
max_expand_depth,
);
(success, reward)
})
.reduce(
|| (0.0, 0.0), |mut acc, (success, total_val)| {
acc.0 += success; acc.1 += total_val; acc }
)
});
Ok((successes / (num_episodes as f32), total_vals / (num_episodes as f32)))
}
}