use crate::{Distribution, Value};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum TrialState {
Running,
Complete,
Pruned,
Failed,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ParamRecord {
pub name: String,
pub distribution: Distribution,
pub value: Value,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Trial {
pub number: usize,
pub params: Vec<ParamRecord>,
pub intermediate_values: Vec<(usize, f64)>,
pub value: Option<f64>,
pub state: TrialState,
}
impl Trial {
pub fn new(number: usize) -> Self {
Trial {
number,
params: Vec::new(),
intermediate_values: Vec::new(),
value: None,
state: TrialState::Running,
}
}
pub fn record(&mut self, name: &str, distribution: Distribution, value: Value) {
if let Some(existing) = self.params.iter_mut().find(|p| p.name == name) {
existing.distribution = distribution;
existing.value = value;
} else {
self.params.push(ParamRecord {
name: name.to_string(),
distribution,
value,
});
}
}
pub fn param_value(&self, name: &str) -> Option<&Value> {
self.params.iter().find(|p| p.name == name).map(|p| &p.value)
}
pub fn value_at_step(&self, step: usize) -> Option<f64> {
self.intermediate_values
.iter()
.find(|(s, _)| *s == step)
.map(|(_, v)| *v)
}
pub fn value_at_or_after(&self, step: usize) -> Option<f64> {
self.intermediate_values
.iter()
.filter(|(s, _)| *s >= step)
.min_by_key(|(s, _)| *s)
.map(|(_, v)| *v)
}
pub fn last_intermediate(&self) -> Option<(usize, f64)> {
self.intermediate_values.last().copied()
}
}