use hyperopt_core::{Pruner, StudyState, Trial};
#[derive(Debug, Clone)]
pub struct SuccessiveHalvingPruner {
min_resource: usize,
reduction_factor: usize,
min_early_stopping_rate: u32,
}
impl Default for SuccessiveHalvingPruner {
fn default() -> Self {
SuccessiveHalvingPruner {
min_resource: 1,
reduction_factor: 4,
min_early_stopping_rate: 0,
}
}
}
impl SuccessiveHalvingPruner {
pub fn new() -> Self {
Self::default()
}
pub fn min_resource(mut self, r: usize) -> Self {
self.min_resource = r.max(1);
self
}
pub fn reduction_factor(mut self, eta: usize) -> Self {
self.reduction_factor = eta.max(2);
self
}
pub fn min_early_stopping_rate(mut self, s: u32) -> Self {
self.min_early_stopping_rate = s;
self
}
fn rung_resource_for(&self, step: usize) -> Option<usize> {
let eta = self.reduction_factor as u64;
let first = (self.min_resource as u64) * eta.pow(self.min_early_stopping_rate);
if (step as u64) < first {
return None;
}
let mut rung = first;
loop {
let next = rung.saturating_mul(eta);
if next <= step as u64 {
rung = next;
} else {
break;
}
}
Some(rung as usize)
}
}
impl Pruner for SuccessiveHalvingPruner {
fn should_prune(&self, study_state: &StudyState, trial: &Trial) -> bool {
let Some((step, value)) = trial.last_intermediate() else {
return false;
};
let Some(rung) = self.rung_resource_for(step) else {
return false;
};
let peers = study_state.values_at_or_after(rung);
let eta = self.reduction_factor;
if peers.len() + 1 < eta {
return false;
}
let direction = study_state.direction();
let better_than_current = peers
.iter()
.filter(|&&p| direction.is_better(p, value))
.count();
let total = peers.len() + 1;
let top_k = (total / eta).max(1);
let rank = better_than_current; rank >= top_k
}
}