Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion rustuna_importance/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,4 +12,4 @@ pub(crate) mod test_utils;
pub use common::{
get_param_importances, get_param_importances_with, ImportanceEvaluator, ImportanceOptions,
};
pub use ped_anova::PedAnovaImportanceEvaluator;
pub use ped_anova::{PedAnovaBuilder, PedAnovaImportanceEvaluator};
192 changes: 178 additions & 14 deletions rustuna_importance/src/ped_anova.rs
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,148 @@ impl Default for PedAnovaImportanceEvaluator {
}
}

/// Builder for [`PedAnovaImportanceEvaluator`], following the API style of
/// [`std::thread::Builder`].
///
/// # Examples
///
/// ```
/// use rustuna_importance::PedAnovaImportanceEvaluator;
///
/// let evaluator = PedAnovaImportanceEvaluator::builder()
/// .target_quantile(0.2)
/// .region_quantile(0.9)
/// .evaluate_on_local(false)
/// .n_steps(100)
/// .prior_weight(2.0)
/// .min_n_trials_in_regime(3)
/// .build()
/// .unwrap();
/// ```
pub struct PedAnovaBuilder {
target_quantile: f64,
region_quantile: f64,
evaluate_on_local: bool,
n_steps: usize,
prior_weight: f64,
min_n_trials_in_regime: usize,
}
impl Default for PedAnovaBuilder {
fn default() -> Self {
Self::new()
}
}
impl PedAnovaBuilder {
/// Creates a builder with the default configuration.
pub fn new() -> Self {
Self {
target_quantile: 0.1,
region_quantile: 1.0,
evaluate_on_local: true,
n_steps: 50,
prior_weight: 1.0,
min_n_trials_in_regime: 2,
}
}

/// Sets the top fraction of completed trials used as the target region.
///
/// For example, `0.1` evaluates which parameters were important for achieving the top
/// 10% of observed objective values.
pub fn target_quantile(self, target_quantile: f64) -> Self {
Self {
target_quantile,
..self
}
}

/// Sets the reference region against which the target region is compared.
pub fn region_quantile(self, region_quantile: f64) -> Self {
Self {
region_quantile,
..self
}
}

/// Sets whether the reference density is estimated from the explored region (`true`)
/// or from the full search space (`false`).
pub fn evaluate_on_local(self, evaluate_on_local: bool) -> Self {
Self {
evaluate_on_local,
..self
}
}

/// Sets the number of grid steps used by the Parzen estimators.
pub fn n_steps(self, n_steps: usize) -> Self {
Self { n_steps, ..self }
}

/// Sets the prior weight used by the Parzen estimators.
pub fn prior_weight(self, prior_weight: f64) -> Self {
Self {
prior_weight,
..self
}
}

/// Sets the minimum number of completed trials for a search regime to contribute to
/// the importance.
pub fn min_n_trials_in_regime(self, min_n_trials_in_regime: usize) -> Self {
Self {
min_n_trials_in_regime,
..self
}
}

/// Builds the evaluator.
///
/// # Errors
///
/// Returns an error when the configuration is invalid.
pub fn build(self) -> Result<PedAnovaImportanceEvaluator> {
if !(0.0 < self.target_quantile
&& self.target_quantile < self.region_quantile
&& self.region_quantile <= 1.0)
{
return Err(Error::with_reason(
ErrorKind::ImportanceEvaluatorError,
"condition 0.0 < `target_quantile` < `region_quantile` <= 1.0 must be satisfied",
));
}
Ok(PedAnovaImportanceEvaluator {
target_quantile: self.target_quantile,
region_quantile: self.region_quantile,
evaluate_on_local: self.evaluate_on_local,
n_steps: self.n_steps,
prior_weight: self.prior_weight,
min_n_trials_in_regime: self.min_n_trials_in_regime,
})
}
}

impl PedAnovaImportanceEvaluator {
/// Returns a builder for creating an evaluator with an explicit configuration.
///
/// This is the counterpart of [`std::thread::Builder`]: settings are configured by
/// chaining methods and the evaluator is created with [`PedAnovaBuilder::build`].
///
/// # Examples
///
/// ```
/// use rustuna_importance::PedAnovaImportanceEvaluator;
///
/// let evaluator = PedAnovaImportanceEvaluator::builder()
/// .target_quantile(0.3)
/// .n_steps(100)
/// .build()
/// .unwrap();
/// # let _ = evaluator;
/// ```
pub fn builder() -> PedAnovaBuilder {
PedAnovaBuilder::new()
}

/// Creates a PED-ANOVA evaluator.
///
/// `target_quantile` selects the top fraction of completed trials used as the target region.
Expand All @@ -105,20 +246,11 @@ impl PedAnovaImportanceEvaluator {
region_quantile: f64,
evaluate_on_local: bool,
) -> Result<Self> {
if !(0.0 < target_quantile && target_quantile < region_quantile && region_quantile <= 1.0) {
return Err(Error::with_reason(
ErrorKind::ImportanceEvaluatorError,
"condition 0.0 < `target_quantile` < `region_quantile` <= 1.0 must be satisfied",
));
}
Ok(Self {
target_quantile,
region_quantile,
evaluate_on_local,
n_steps: 50,
prior_weight: 1.0,
min_n_trials_in_regime: 2,
})
Self::builder()
.target_quantile(target_quantile)
.region_quantile(region_quantile)
.evaluate_on_local(evaluate_on_local)
.build()
}

fn get_top_quantile_trials<'a>(
Expand Down Expand Up @@ -599,6 +731,38 @@ mod tests {
Ok(())
}

#[test]
fn test_builder() -> Result<()> {
// A full configuration through the builder.
let evaluator = PedAnovaImportanceEvaluator::builder()
.target_quantile(0.3)
.region_quantile(0.9)
.evaluate_on_local(false)
.n_steps(100)
.prior_weight(2.0)
.min_n_trials_in_regime(3)
.build()?;
let study = test_utils::get_study(42, 20, ObjectiveType::Single, Direction::Minimize)?;
assert!(!evaluator.evaluate(&study)?.is_empty());

// The same configuration as `new`, with the internal defaults for the rest.
let from_new = PedAnovaImportanceEvaluator::new(0.3, 1.0, true)?;
let from_builder = PedAnovaImportanceEvaluator::builder()
.target_quantile(0.3)
.region_quantile(1.0)
.evaluate_on_local(true)
.build()?;
assert_eq!(from_new.evaluate(&study)?, from_builder.evaluate(&study)?);

// Invalid configurations are rejected by `build`.
assert!(PedAnovaImportanceEvaluator::builder()
.target_quantile(0.5)
.region_quantile(0.3)
.build()
.is_err());
Ok(())
}

#[test]
fn test_conditional() -> Result<()> {
let study = test_utils::get_study(42, 20, ObjectiveType::Conditional, Direction::Minimize)?;
Expand Down
9 changes: 6 additions & 3 deletions rustuna_pyo3/src/importance.rs
Original file line number Diff line number Diff line change
Expand Up @@ -56,9 +56,12 @@ impl PyPedAnovaImportanceEvaluator {
region_quantile: f64,
evaluate_on_local: bool,
) -> PyResult<Self> {
let evaluator =
PedAnovaImportanceEvaluator::new(target_quantile, region_quantile, evaluate_on_local)
.map_err(err_to_exceptions)?;
let evaluator = PedAnovaImportanceEvaluator::builder()
.target_quantile(target_quantile)
.region_quantile(region_quantile)
.evaluate_on_local(evaluate_on_local)
.build()
.map_err(err_to_exceptions)?;
if region_quantile != 1.0 && !evaluate_on_local {
PyErr::warn(
py,
Expand Down
Loading