use crate::rl::env::Env;
use crate::nn::policy::{Policy, sample, argmax};
use super::search::predict_probs_mcts;
pub fn single_solve(
env: &mut Box<dyn Env>,
policy: &Policy,
deterministic: bool,
num_mcts_searches: usize,
C: f32,
max_expand_depth: usize,
) -> ((f32, f32), Vec<usize>) {
let mut total_val = 0.0;
let mut solution = Vec::new();
let track_solution = env.track_solution();
while !env.is_final() {
let val = env.reward();
let obs = env.observe();
let masks = env.masks();
total_val += val;
let probs = if num_mcts_searches == 0 {
policy.predict(obs, masks).0
} else {
predict_probs_mcts(
env.clone(),
policy,
num_mcts_searches,
C,
max_expand_depth,
)
};
let action = if deterministic {
argmax(&probs)
} else {
sample(&probs)
};
env.step(action);
if !track_solution {
solution.push(action);
}
}
if track_solution {
solution = env.solution();
}
let val = env.reward();
total_val += val;
let success = if env.success() { 1.0 } else { 0.0 };
((success, total_val), solution)
}
pub fn solve(
env: &Box<dyn Env>,
policy: &Policy,
deterministic: bool,
num_searches: usize,
num_mcts_searches: usize,
C: f32,
max_expand_depth: usize,
) -> ((f32, f32), Vec<usize>) {
let mut best: ((f32, f32), Vec<usize>) = ((0.0, f32::NEG_INFINITY), Vec::new());
for _ in 0..num_searches {
let mut cloned_env = env.clone(); let next_val = single_solve(
&mut cloned_env,
policy,
deterministic,
num_mcts_searches,
C,
max_expand_depth,
);
if next_val.0 > best.0 {
best = next_val;
}
}
best
}