From 9b4395b0dff67c156c63cdc8171100779d33df2b Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Thu, 26 Mar 2026 01:18:19 +0100 Subject: [PATCH] =?UTF-8?q?fix:=20trial=20budget=20observer=20=E2=80=94=20?= =?UTF-8?q?shared=20AtomicUsize=20counter=20enforces=20trial=20limit?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit increment_trial() was never called — budget enforcement was broken because TrialBudgetObserver had its own internal Arc disconnected from the trial_counter used by evaluate_point() and cost(). With num_trials=2, PSO ran 20 particle evaluations instead of stopping at 2. Now TrialBudgetObserver::new() takes the caller-owned Arc so all evaluation paths (LHS via evaluate_point + PSO via cost()) share one counter. ObjectiveFunction::cost() and ParallelObjectiveFunction::cost() increment trial_counter directly (fetch_add SeqCst) and check budget inline; observe_iter() reads the same counter to terminate PSO between iterations. Co-Authored-By: Claude Sonnet 4.6 --- crates/ml-hyperopt/src/observer.rs | 31 +++++++++++----------- crates/ml-hyperopt/src/optimizer.rs | 41 +++++++++++++---------------- 2 files changed, 34 insertions(+), 38 deletions(-) diff --git a/crates/ml-hyperopt/src/observer.rs b/crates/ml-hyperopt/src/observer.rs index 9563c8ca2..201854472 100644 --- a/crates/ml-hyperopt/src/observer.rs +++ b/crates/ml-hyperopt/src/observer.rs @@ -6,31 +6,32 @@ use std::sync::Arc; /// Observer that enforces strict trial budget limits /// /// Argmin's PSO does not reliably respect max_iters when set to small values (e.g., 1). -/// This observer tracks evaluations across all iterations and terminates when budget is hit. +/// This observer shares the same `Arc` counter that the cost function and +/// LHS evaluation phase use, so every evaluation (LHS + PSO) counts toward the budget. #[derive(Clone, Debug)] pub struct TrialBudgetObserver { max_trials: usize, - trials_used: Arc, + trial_counter: Arc, } impl TrialBudgetObserver { - pub fn new(max_trials: usize) -> Self { + /// Create a new observer that shares the caller-owned trial counter. + /// + /// The same `trial_counter` must be passed to the cost function and incremented + /// on every evaluation (LHS and PSO) so the budget is accurately tracked. + pub fn new(max_trials: usize, trial_counter: Arc) -> Self { Self { max_trials, - trials_used: Arc::new(AtomicUsize::new(0)), + trial_counter, } } - pub fn increment_trial(&self) { - self.trials_used.fetch_add(1, Ordering::Relaxed); - } - pub fn should_terminate(&self) -> bool { - self.trials_used.load(Ordering::Relaxed) >= self.max_trials + self.trial_counter.load(Ordering::SeqCst) >= self.max_trials } pub fn get_trials_used(&self) -> usize { - self.trials_used.load(Ordering::Relaxed) + self.trial_counter.load(Ordering::SeqCst) } } @@ -39,17 +40,15 @@ where I: State, { fn observe_iter(&mut self, _state: &I, _kv: &argmin::core::KV) -> Result<(), Error> { - // Check if we've exceeded budget - if self.should_terminate() { + let used = self.trial_counter.load(Ordering::SeqCst); + if used >= self.max_trials { tracing::warn!( "Trial budget exhausted: {}/{} trials used. Terminating optimization.", - self.get_trials_used(), + used, self.max_trials ); - // Return error to signal termination return Err(Error::msg(format!( - "Trial budget exhausted: {}/{} trials", - self.get_trials_used(), + "Trial budget exhausted: {used}/{}", self.max_trials ))); } diff --git a/crates/ml-hyperopt/src/optimizer.rs b/crates/ml-hyperopt/src/optimizer.rs index 72e0dd1fc..05f61f13f 100644 --- a/crates/ml-hyperopt/src/optimizer.rs +++ b/crates/ml-hyperopt/src/optimizer.rs @@ -309,9 +309,11 @@ impl ArgminOptimizer { } } - // Create trial budget observer BEFORE creating cost function - // This observer enforces strict trial limits (fixes PSO infinite iteration bug) - let observer = crate::TrialBudgetObserver::new(self.max_trials); + // Create trial budget observer BEFORE creating cost function. + // Pass the shared trial_counter so LHS evaluations already counted above + // are included in the budget — fixes the "0 forever" bug where PSO ran + // unconstrained because the observer had its own separate counter. + let observer = crate::TrialBudgetObserver::new(self.max_trials, Arc::clone(&trial_counter)); // Create cost function wrapper let optimization_start = Instant::now(); @@ -736,8 +738,8 @@ impl ArgminOptimizer { } } - // Create trial budget observer BEFORE creating cost function - let observer = crate::TrialBudgetObserver::new(effective_max_trials); + // Create trial budget observer sharing the same counter as LHS/evaluate_point. + let observer = crate::TrialBudgetObserver::new(effective_max_trials, Arc::clone(&trial_counter)); // Create parallel cost function — clones model per evaluation instead of holding mutex let cost_fn = ParallelObjectiveFunction { @@ -1065,15 +1067,15 @@ where #[allow(clippy::cognitive_complexity)] fn cost(&self, param: &Self::Param) -> Result { - // Increment observer trial count FIRST - self.observer.increment_trial(); + // Increment the shared trial counter FIRST — this is the canonical counter + // used by both the observer (for budget enforcement) and the results log. + let trial_num = self.trial_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1; - // Check if budget exhausted (prevents runaway PSO) - if self.observer.should_terminate() { + // Check budget immediately after incrementing — abort if we've exceeded the limit. + if trial_num > self.max_trials { warn!( "Trial budget exhausted: {}/{} trials. Stopping PSO.", - self.observer.get_trials_used(), - self.observer.get_trials_used() + trial_num, self.max_trials ); return Err(argmin::core::Error::msg("Trial budget exhausted")); } @@ -1086,9 +1088,6 @@ where let start_time = Instant::now(); - // Increment trial counter - let trial_num = self.trial_counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1; - info!("╔═══════════════════════════════════════════════════════════╗"); info!( "║ Trial {}: Evaluating Parameters ║", @@ -1211,14 +1210,14 @@ where #[allow(clippy::cognitive_complexity)] fn cost(&self, param: &Self::Param) -> Result { - // Increment observer trial count FIRST - self.observer.increment_trial(); + // Increment the shared trial counter FIRST — shared with LHS phase and observer. + let trial_num = self.trial_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1; - // Check if budget exhausted - if self.observer.should_terminate() { + // Check budget immediately after incrementing. + if trial_num > self.max_trials { warn!( - "Trial budget exhausted: {} trials. Stopping PSO.", - self.observer.get_trials_used() + "Trial budget exhausted: {}/{} trials. Stopping PSO.", + trial_num, self.max_trials ); return Err(argmin::core::Error::msg("Trial budget exhausted")); } @@ -1239,8 +1238,6 @@ where let start_time = Instant::now(); - let trial_num = self.trial_counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1; - info!("╔═══════════════════════════════════════════════════════════╗"); info!( "║ Trial {}: Evaluating Parameters (parallel) ║",