From 20864053522074bbcd074afdc880f8b4a8668e1b Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Thu, 12 Mar 2026 21:55:56 +0100 Subject: [PATCH] fix(hyperopt): auto-bump trials to n_initial+1 instead of crashing When trials=0 (or any value <= n_initial), both RL and supervised hyperopt binaries now auto-bump to n_initial+1 instead of bailing. Previously the RL binary bumped to 5 which equalled n_initial=5, triggering "trials must be greater than n_initial" error. The supervised binary lacked the bump entirely and just crashed. Co-Authored-By: Claude Opus 4.6 --- crates/ml/examples/hyperopt_baseline_rl.rs | 20 ++++++++----------- .../examples/hyperopt_baseline_supervised.rs | 14 +++++++------ 2 files changed, 16 insertions(+), 18 deletions(-) diff --git a/crates/ml/examples/hyperopt_baseline_rl.rs b/crates/ml/examples/hyperopt_baseline_rl.rs index a73ba6091..a347c4f0d 100644 --- a/crates/ml/examples/hyperopt_baseline_rl.rs +++ b/crates/ml/examples/hyperopt_baseline_rl.rs @@ -509,19 +509,15 @@ fn main() -> Result<()> { ); } - // Enforce minimum 5 trials for meaningful hyperopt - if args.trials < 5 { - info!("Trials {} below minimum — bumping to 5", args.trials); - args.trials = 5; - } - - // Verify trials > n_initial (ArgminOptimizer requirement) - if args.trials <= args.n_initial { - anyhow::bail!( - "trials ({}) must be greater than n_initial ({})", - args.trials, - args.n_initial + // Enforce trials > n_initial (ArgminOptimizer needs at least n_initial + // LHS exploration rounds + 1 TPE-guided trial to be meaningful) + let min_trials = args.n_initial + 1; + if args.trials < min_trials { + info!( + "Trials {} below minimum (n_initial {} + 1) — bumping to {}", + args.trials, args.n_initial, min_trials ); + args.trials = min_trials; } // Create output directory diff --git a/crates/ml/examples/hyperopt_baseline_supervised.rs b/crates/ml/examples/hyperopt_baseline_supervised.rs index f40ea97fa..f003eb8d6 100644 --- a/crates/ml/examples/hyperopt_baseline_supervised.rs +++ b/crates/ml/examples/hyperopt_baseline_supervised.rs @@ -613,7 +613,7 @@ fn main() -> Result<()> { metrics_server::start_metrics_server(9094); training_metrics::set_active_workers(1.0); - let args = Args::parse(); + let mut args = Args::parse(); // Signal hyperopt mode active — will be cleared at exit let hyperopt_model_label = args.model.clone(); @@ -654,12 +654,14 @@ fn main() -> Result<()> { info!("Models to optimize: {:?}", models); - if args.trials <= args.n_initial { - anyhow::bail!( - "trials ({}) must be greater than n_initial ({})", - args.trials, - args.n_initial + // Enforce trials > n_initial (ArgminOptimizer needs LHS exploration + ≥1 TPE trial) + let min_trials = args.n_initial + 1; + if args.trials < min_trials { + info!( + "Trials {} below minimum (n_initial {} + 1) — bumping to {}", + args.trials, args.n_initial, min_trials ); + args.trials = min_trials; } if let Some(parent) = args.output.parent() {