use hyperopt_core::{Distribution, Sampler, StudyState, Trial, Value};
use rand::rngs::StdRng;
use rand::SeedableRng;
use crate::random::sample_value;
pub struct GridSampler {
grids: Vec<(String, Vec<Value>)>,
fallback_rng: StdRng,
}
impl GridSampler {
pub fn new() -> Self {
GridSampler {
grids: Vec::new(),
fallback_rng: StdRng::seed_from_u64(0),
}
}
pub fn add_grid(mut self, name: &str, values: Vec<Value>) -> Self {
self.grids.push((name.to_string(), values));
self
}
pub fn add_float_grid(self, name: &str, values: &[f64]) -> Self {
self.add_grid(name, values.iter().map(|v| Value::Float(*v)).collect())
}
pub fn add_int_grid(self, name: &str, values: &[i64]) -> Self {
self.add_grid(name, values.iter().map(|v| Value::Int(*v)).collect())
}
pub fn add_categorical_grid(self, name: &str, values: &[&str]) -> Self {
self.add_grid(
name,
values.iter().map(|v| Value::Categorical(v.to_string())).collect(),
)
}
pub fn grid_size(&self) -> usize {
self.grids
.iter()
.map(|(_, v)| v.len().max(1))
.product::<usize>()
.max(1)
}
fn index_for(&self, name: &str, trial_number: usize) -> Option<usize> {
let mut suffix_product = 1usize;
let mut target_len = None;
let mut divisor = 1usize;
for (grid_name, values) in self.grids.iter().rev() {
let radix = values.len().max(1);
if grid_name == name {
target_len = Some(values.len());
divisor = suffix_product;
}
suffix_product = suffix_product.saturating_mul(radix);
}
let len = target_len?;
if len == 0 {
return None;
}
let total = self.grid_size();
let combo = trial_number % total;
Some((combo / divisor) % len)
}
}
impl Default for GridSampler {
fn default() -> Self {
Self::new()
}
}
impl Sampler for GridSampler {
fn suggest(
&mut self,
_study_state: &StudyState,
trial: &Trial,
param_name: &str,
distribution: &Distribution,
) -> Value {
if let Some(idx) = self.index_for(param_name, trial.number) {
if let Some((_, values)) = self.grids.iter().find(|(n, _)| n == param_name) {
if let Some(v) = values.get(idx) {
return v.clone();
}
}
}
sample_value(&mut self.fallback_rng, distribution)
}
}