feat: add constraint handling with feasibility-aware trial ranking

Add constraint support so that optimization problems with constraints
(e.g., "model size < 100MB") prefer feasible solutions. Convention:
constraint value <= 0.0 means feasible.

- Add constraint_values field to Trial with set_constraints/getter
- Add constraints field to CompletedTrial with is_feasible() method
- Propagate constraints through complete_trial/prune_trial
- Make best_trial() and top_trials() constraint-aware: feasible trials
  rank above infeasible, infeasible ranked by total violation
This commit is contained in:
Manuel Raimann
2026-02-11 18:49:08 +01:00
parent aa9734587a
commit 09480a973b
5 changed files with 181 additions and 34 deletions
+2
View File
@@ -178,6 +178,7 @@ impl BohbSampler {
intermediate_values: trial.intermediate_values.clone(), intermediate_values: trial.intermediate_values.clone(),
state: trial.state, state: trial.state,
user_attrs: trial.user_attrs.clone(), user_attrs: trial.user_attrs.clone(),
constraints: trial.constraints.clone(),
}) })
}) })
.collect() .collect()
@@ -396,6 +397,7 @@ mod tests {
intermediate_values, intermediate_values,
state: TrialState::Complete, state: TrialState::Complete,
user_attrs: HashMap::new(), user_attrs: HashMap::new(),
constraints: Vec::new(),
} }
} }
+13
View File
@@ -41,6 +41,9 @@ pub struct CompletedTrial<V = f64> {
pub state: TrialState, pub state: TrialState,
/// User-defined attributes stored during the trial. /// User-defined attributes stored during the trial.
pub user_attrs: HashMap<String, AttrValue>, pub user_attrs: HashMap<String, AttrValue>,
/// Constraint values for this trial (<=0.0 means feasible).
#[cfg_attr(feature = "serde", serde(default))]
pub constraints: Vec<f64>,
} }
impl<V> CompletedTrial<V> { impl<V> CompletedTrial<V> {
@@ -61,6 +64,7 @@ impl<V> CompletedTrial<V> {
intermediate_values: Vec::new(), intermediate_values: Vec::new(),
state: TrialState::Complete, state: TrialState::Complete,
user_attrs: HashMap::new(), user_attrs: HashMap::new(),
constraints: Vec::new(),
} }
} }
@@ -83,6 +87,7 @@ impl<V> CompletedTrial<V> {
intermediate_values, intermediate_values,
state: TrialState::Complete, state: TrialState::Complete,
user_attrs, user_attrs,
constraints: Vec::new(),
} }
} }
@@ -127,6 +132,14 @@ impl<V> CompletedTrial<V> {
}) })
} }
/// Returns `true` if all constraints are satisfied (values <= 0.0).
///
/// A trial with no constraints is considered feasible.
#[must_use]
pub fn is_feasible(&self) -> bool {
self.constraints.iter().all(|&c| c <= 0.0)
}
/// Gets a user attribute by key. /// Gets a user attribute by key.
#[must_use] #[must_use]
pub fn user_attr(&self, key: &str) -> Option<&AttrValue> { pub fn user_attr(&self, key: &str) -> Option<&AttrValue> {
+43 -34
View File
@@ -343,18 +343,11 @@ where
/// Returns the trial ID of the current best trial from the given slice. /// Returns the trial ID of the current best trial from the given slice.
#[cfg(feature = "tracing")] #[cfg(feature = "tracing")]
fn best_id(&self, trials: &[CompletedTrial<V>]) -> Option<u64> { fn best_id(&self, trials: &[CompletedTrial<V>]) -> Option<u64> {
let direction = self.direction;
trials trials
.iter() .iter()
.filter(|t| t.state == TrialState::Complete) .filter(|t| t.state == TrialState::Complete)
.max_by(|a, b| { .max_by(|a, b| Self::compare_trials(a, b, direction))
let ordering = a.value.partial_cmp(&b.value);
match self.direction {
Direction::Minimize => {
ordering.map_or(core::cmp::Ordering::Equal, core::cmp::Ordering::reverse)
}
Direction::Maximize => ordering.unwrap_or(core::cmp::Ordering::Equal),
}
})
.map(|t| t.id) .map(|t| t.id)
} }
@@ -461,6 +454,7 @@ where
trial.user_attrs().clone(), trial.user_attrs().clone(),
); );
completed.state = TrialState::Complete; completed.state = TrialState::Complete;
completed.constraints = trial.constraint_values().to_vec();
self.completed_trials.write().push(completed); self.completed_trials.write().push(completed);
} }
@@ -570,6 +564,7 @@ where
trial.user_attrs().clone(), trial.user_attrs().clone(),
); );
completed.state = TrialState::Pruned; completed.state = TrialState::Pruned;
completed.constraints = trial.constraint_values().to_vec();
self.completed_trials.write().push(completed); self.completed_trials.write().push(completed);
} }
@@ -636,12 +631,46 @@ where
.count() .count()
} }
/// Compares two completed trials using constraint-aware ranking.
///
/// 1. Feasible trials always rank above infeasible trials.
/// 2. Among feasible trials, rank by objective value (respecting direction).
/// 3. Among infeasible trials, rank by total constraint violation (lower is better).
fn compare_trials(
a: &CompletedTrial<V>,
b: &CompletedTrial<V>,
direction: Direction,
) -> core::cmp::Ordering {
match (a.is_feasible(), b.is_feasible()) {
(true, false) => core::cmp::Ordering::Greater,
(false, true) => core::cmp::Ordering::Less,
(false, false) => {
let va: f64 = a.constraints.iter().map(|c| c.max(0.0)).sum();
let vb: f64 = b.constraints.iter().map(|c| c.max(0.0)).sum();
vb.partial_cmp(&va).unwrap_or(core::cmp::Ordering::Equal)
}
(true, true) => {
let ordering = a.value.partial_cmp(&b.value);
match direction {
Direction::Minimize => {
ordering.map_or(core::cmp::Ordering::Equal, core::cmp::Ordering::reverse)
}
Direction::Maximize => ordering.unwrap_or(core::cmp::Ordering::Equal),
}
}
}
}
/// Returns the trial with the best objective value. /// Returns the trial with the best objective value.
/// ///
/// The "best" trial depends on the optimization direction: /// The "best" trial depends on the optimization direction:
/// - `Direction::Minimize`: Returns the trial with the lowest objective value. /// - `Direction::Minimize`: Returns the trial with the lowest objective value.
/// - `Direction::Maximize`: Returns the trial with the highest objective value. /// - `Direction::Maximize`: Returns the trial with the highest objective value.
/// ///
/// When constraints are present, feasible trials always rank above infeasible
/// trials. Among infeasible trials, those with lower total constraint violation
/// are preferred.
///
/// # Errors /// # Errors
/// ///
/// Returns `Error::NoCompletedTrials` if no trials have been completed. /// Returns `Error::NoCompletedTrials` if no trials have been completed.
@@ -675,25 +704,12 @@ where
V: Clone, V: Clone,
{ {
let trials = self.completed_trials.read(); let trials = self.completed_trials.read();
let direction = self.direction;
let best = trials let best = trials
.iter() .iter()
.filter(|t| t.state == TrialState::Complete) .filter(|t| t.state == TrialState::Complete)
.max_by(|a, b| { .max_by(|a, b| Self::compare_trials(a, b, direction))
// For Minimize, we want the smallest value to be "max" in ordering
// For Maximize, we want the largest value to be "max" in ordering
let ordering = a.value.partial_cmp(&b.value);
match self.direction {
Direction::Minimize => {
// Reverse ordering: smaller values are "greater" for max_by
ordering.map_or(core::cmp::Ordering::Equal, core::cmp::Ordering::reverse)
}
Direction::Maximize => {
// Normal ordering: larger values are "greater" for max_by
ordering.unwrap_or(core::cmp::Ordering::Equal)
}
}
})
.ok_or(crate::Error::NoCompletedTrials)?; .ok_or(crate::Error::NoCompletedTrials)?;
Ok(best.clone()) Ok(best.clone())
@@ -752,21 +768,14 @@ where
V: Clone, V: Clone,
{ {
let trials = self.completed_trials.read(); let trials = self.completed_trials.read();
let direction = self.direction;
let mut completed: Vec<_> = trials let mut completed: Vec<_> = trials
.iter() .iter()
.filter(|t| t.state == TrialState::Complete) .filter(|t| t.state == TrialState::Complete)
.cloned() .cloned()
.collect(); .collect();
completed.sort_by(|a, b| match self.direction { // Sort best-first: reverse the compare_trials ordering (which is designed for max_by)
Direction::Minimize => a completed.sort_by(|a, b| Self::compare_trials(b, a, direction));
.value
.partial_cmp(&b.value)
.unwrap_or(core::cmp::Ordering::Equal),
Direction::Maximize => b
.value
.partial_cmp(&a.value)
.unwrap_or(core::cmp::Ordering::Equal),
});
completed.truncate(n); completed.truncate(n);
completed completed
} }
+19
View File
@@ -89,6 +89,8 @@ pub struct Trial {
user_attrs: HashMap<String, AttrValue>, user_attrs: HashMap<String, AttrValue>,
/// Pre-filled parameter values from enqueue (used instead of sampling). /// Pre-filled parameter values from enqueue (used instead of sampling).
fixed_params: HashMap<ParamId, ParamValue>, fixed_params: HashMap<ParamId, ParamValue>,
/// Constraint values for this trial (<=0.0 means feasible).
constraint_values: Vec<f64>,
} }
impl core::fmt::Debug for Trial { impl core::fmt::Debug for Trial {
@@ -105,6 +107,7 @@ impl core::fmt::Debug for Trial {
.field("has_pruner", &self.pruner.is_some()) .field("has_pruner", &self.pruner.is_some())
.field("user_attrs", &self.user_attrs) .field("user_attrs", &self.user_attrs)
.field("fixed_params", &self.fixed_params) .field("fixed_params", &self.fixed_params)
.field("constraint_values", &self.constraint_values)
.finish() .finish()
} }
} }
@@ -144,6 +147,7 @@ impl Trial {
pruner: None, pruner: None,
user_attrs: HashMap::new(), user_attrs: HashMap::new(),
fixed_params: HashMap::new(), fixed_params: HashMap::new(),
constraint_values: Vec::new(),
} }
} }
@@ -175,6 +179,7 @@ impl Trial {
pruner: Some(pruner), pruner: Some(pruner),
user_attrs: HashMap::new(), user_attrs: HashMap::new(),
fixed_params: HashMap::new(), fixed_params: HashMap::new(),
constraint_values: Vec::new(),
} }
} }
@@ -292,6 +297,20 @@ impl Trial {
&self.user_attrs &self.user_attrs
} }
/// Sets constraint values for this trial.
///
/// Each value represents a constraint; a value <= 0.0 means the constraint
/// is satisfied (feasible). A value > 0.0 means the constraint is violated.
pub fn set_constraints(&mut self, values: Vec<f64>) {
self.constraint_values = values;
}
/// Returns the constraint values for this trial.
#[must_use]
pub fn constraint_values(&self) -> &[f64] {
&self.constraint_values
}
/// Sets the trial state to Complete. /// Sets the trial state to Complete.
pub(crate) fn set_complete(&mut self) { pub(crate) fn set_complete(&mut self) {
self.state = TrialState::Complete; self.state = TrialState::Complete;
+104
View File
@@ -2141,3 +2141,107 @@ fn test_into_iterator_preserves_insertion_order() {
let ids: Vec<u64> = (&study).into_iter().map(|t| t.id).collect(); let ids: Vec<u64> = (&study).into_iter().map(|t| t.id).collect();
assert_eq!(ids, vec![0, 1, 2]); assert_eq!(ids, vec![0, 1, 2]);
} }
// =============================================================================
// Tests: Constraint handling
// =============================================================================
#[test]
fn test_is_feasible_all_satisfied() {
let study: Study<f64> = Study::new(Direction::Minimize);
let mut trial = study.create_trial();
trial.set_constraints(vec![-1.0, 0.0, -0.5]);
study.complete_trial(trial, 1.0);
let completed = study.best_trial().unwrap();
assert!(completed.is_feasible());
}
#[test]
fn test_is_feasible_one_violated() {
let study: Study<f64> = Study::new(Direction::Minimize);
let mut trial = study.create_trial();
trial.set_constraints(vec![-1.0, 0.5, -0.5]);
study.complete_trial(trial, 1.0);
let completed = study.best_trial().unwrap();
assert!(!completed.is_feasible());
}
#[test]
fn test_is_feasible_empty_constraints() {
let study: Study<f64> = Study::new(Direction::Minimize);
let trial = study.create_trial();
study.complete_trial(trial, 1.0);
let completed = study.best_trial().unwrap();
assert!(completed.is_feasible());
}
#[test]
fn test_best_trial_prefers_feasible() {
let study: Study<f64> = Study::new(Direction::Minimize);
// Infeasible trial with better objective
let mut trial1 = study.create_trial();
trial1.set_constraints(vec![1.0]);
study.complete_trial(trial1, 0.1);
// Feasible trial with worse objective
let mut trial2 = study.create_trial();
trial2.set_constraints(vec![-1.0]);
study.complete_trial(trial2, 100.0);
let best = study.best_trial().unwrap();
assert_eq!(best.id, 1); // feasible trial wins
assert_eq!(best.value, 100.0);
}
#[test]
fn test_best_trial_feasible_by_objective() {
let study: Study<f64> = Study::new(Direction::Minimize);
// Feasible, worse objective
let mut trial1 = study.create_trial();
trial1.set_constraints(vec![-1.0]);
study.complete_trial(trial1, 10.0);
// Feasible, better objective
let mut trial2 = study.create_trial();
trial2.set_constraints(vec![-0.5]);
study.complete_trial(trial2, 2.0);
let best = study.best_trial().unwrap();
assert_eq!(best.id, 1); // lower objective wins among feasible
assert_eq!(best.value, 2.0);
}
#[test]
fn test_top_trials_ranks_feasible_above_infeasible() {
let study: Study<f64> = Study::new(Direction::Minimize);
// Infeasible, low violation
let mut t0 = study.create_trial();
t0.set_constraints(vec![0.5]);
study.complete_trial(t0, 1.0);
// Feasible, worst objective among feasible
let mut t1 = study.create_trial();
t1.set_constraints(vec![-1.0]);
study.complete_trial(t1, 50.0);
// Feasible, best objective among feasible
let mut t2 = study.create_trial();
t2.set_constraints(vec![-0.1]);
study.complete_trial(t2, 5.0);
// Infeasible, high violation
let mut t3 = study.create_trial();
t3.set_constraints(vec![3.0]);
study.complete_trial(t3, 0.5);
let top = study.top_trials(4);
let ids: Vec<u64> = top.iter().map(|t| t.id).collect();
// Feasible sorted by objective first (5.0, 50.0), then infeasible by violation (0.5, 3.0)
assert_eq!(ids, vec![2, 1, 0, 3]);
}