Skip to main content

hyperopt_pruners/
median.rs

1use crate::{is_worse, median};
2use hyperopt_core::{Pruner, StudyState, Trial};
3
4/// Prunes a trial when its latest intermediate value is worse than the median
5/// of other trials' values at the same step.
6///
7/// Mirrors Optuna's `MedianPruner`:
8/// - No pruning until at least `n_startup_trials` trials have completed (so the
9///   median is meaningful).
10/// - No pruning before `n_warmup_steps` steps have elapsed within a trial (give
11///   every trial a chance to get past a noisy start).
12/// - At the current step, compare against the median of the values reported at
13///   that same step by trials that reached it; prune if strictly worse.
14#[derive(Debug, Clone)]
15pub struct MedianPruner {
16    n_startup_trials: usize,
17    n_warmup_steps: usize,
18    min_trials_at_step: usize,
19}
20
21impl Default for MedianPruner {
22    fn default() -> Self {
23        MedianPruner {
24            n_startup_trials: 5,
25            n_warmup_steps: 0,
26            min_trials_at_step: 1,
27        }
28    }
29}
30
31impl MedianPruner {
32    /// A median pruner with default gates (`n_startup_trials = 5`,
33    /// `n_warmup_steps = 0`, `min_trials_at_step = 1`).
34    pub fn new() -> Self {
35        Self::default()
36    }
37
38    /// Number of completed trials required before any pruning happens.
39    pub fn n_startup_trials(mut self, n: usize) -> Self {
40        self.n_startup_trials = n;
41        self
42    }
43
44    /// Number of steps a trial must run before it becomes eligible for pruning.
45    pub fn n_warmup_steps(mut self, n: usize) -> Self {
46        self.n_warmup_steps = n;
47        self
48    }
49
50    /// Minimum number of comparison values required at a step before the median
51    /// is trusted enough to prune against.
52    pub fn min_trials_at_step(mut self, n: usize) -> Self {
53        self.min_trials_at_step = n;
54        self
55    }
56}
57
58impl Pruner for MedianPruner {
59    fn should_prune(&self, study_state: &StudyState, trial: &Trial) -> bool {
60        if study_state.n_completed() < self.n_startup_trials {
61            return false;
62        }
63        let Some((step, value)) = trial.last_intermediate() else {
64            return false;
65        };
66        if step < self.n_warmup_steps {
67            return false;
68        }
69        let others = study_state.intermediate_values_at(step);
70        if others.len() < self.min_trials_at_step {
71            return false;
72        }
73        match median(&others) {
74            Some(m) => is_worse(study_state.direction(), value, m),
75            None => false,
76        }
77    }
78}