hyperopt_pruners/
median.rs1use crate::{is_worse, median};
2use hyperopt_core::{Pruner, StudyState, Trial};
3
4#[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 pub fn new() -> Self {
35 Self::default()
36 }
37
38 pub fn n_startup_trials(mut self, n: usize) -> Self {
40 self.n_startup_trials = n;
41 self
42 }
43
44 pub fn n_warmup_steps(mut self, n: usize) -> Self {
46 self.n_warmup_steps = n;
47 self
48 }
49
50 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}