diff --git a/src/parameter.rs b/src/parameter.rs index fb57794..3fa820c 100644 --- a/src/parameter.rs +++ b/src/parameter.rs @@ -37,7 +37,7 @@ static NEXT_PARAM_ID: AtomicU64 = AtomicU64::new(0); /// /// Each parameter is assigned a unique `ParamId` at creation time. Cloning a parameter /// copies its `ParamId`, so clones refer to the same logical parameter. -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct ParamId(u64); impl ParamId { diff --git a/src/study.rs b/src/study.rs index 0042a23..894ac5e 100644 --- a/src/study.rs +++ b/src/study.rs @@ -1,6 +1,7 @@ //! Study implementation for managing optimization trials. use core::any::Any; +use core::fmt; #[cfg(feature = "async")] use core::future::Future; use core::ops::ControlFlow; @@ -1465,6 +1466,86 @@ where } } +impl Study +where + V: PartialOrd + Clone + fmt::Display, +{ + /// Returns a human-readable summary of the study. + /// + /// The summary includes: + /// - Optimization direction and total trial count + /// - Breakdown by state (complete, pruned) when applicable + /// - Best trial value and parameters (if any completed trials exist) + /// + /// # Examples + /// + /// ``` + /// use optimizer::parameter::{FloatParam, Parameter}; + /// use optimizer::{Direction, Study}; + /// + /// let study: Study = Study::new(Direction::Minimize); + /// let x = FloatParam::new(0.0, 10.0).name("x"); + /// + /// let mut trial = study.create_trial(); + /// let _ = x.suggest(&mut trial).unwrap(); + /// study.complete_trial(trial, 0.42); + /// + /// let summary = study.summary(); + /// assert!(summary.contains("Minimize")); + /// assert!(summary.contains("0.42")); + /// ``` + #[must_use] + pub fn summary(&self) -> String { + use fmt::Write; + + let trials = self.completed_trials.read(); + let n_complete = trials + .iter() + .filter(|t| t.state == TrialState::Complete) + .count(); + let n_pruned = trials + .iter() + .filter(|t| t.state == TrialState::Pruned) + .count(); + + let direction_str = match self.direction { + Direction::Minimize => "Minimize", + Direction::Maximize => "Maximize", + }; + + let mut s = format!("Study: {direction_str} | {n} trials", n = trials.len()); + if n_pruned > 0 { + let _ = write!(s, " ({n_complete} complete, {n_pruned} pruned)"); + } + + drop(trials); + + if let Ok(best) = self.best_trial() { + let _ = write!(s, "\nBest value: {} (trial #{})", best.value, best.id); + if !best.params.is_empty() { + s.push_str("\nBest parameters:"); + let mut params: Vec<_> = best.params.iter().collect(); + params.sort_by_key(|(id, _)| *id); + for (id, value) in params { + let label = best.param_labels.get(id).map_or("?", String::as_str); + let _ = write!(s, "\n {label} = {value}"); + } + } + } + + s + } +} + +impl fmt::Display for Study +where + V: PartialOrd + Clone + fmt::Display, +{ + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.summary()) + } +} + // Specialized implementation for Study that provides deprecated `_with_sampler` aliases. // // For Study, the generic methods from `impl Study` (like `optimize()`, diff --git a/tests/integration.rs b/tests/integration.rs index fa63f24..f10a999 100644 --- a/tests/integration.rs +++ b/tests/integration.rs @@ -1892,3 +1892,75 @@ fn test_enqueue_counted_in_n_trials() { // All 5 trials count, including the 2 enqueued ones assert_eq!(study.n_trials(), 5); } + +// ============================================================================= +// Test: Study summary and Display +// ============================================================================= + +#[test] +fn test_summary_with_completed_trials() { + let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(1)); + let x = FloatParam::new(0.0, 10.0).name("x"); + + study + .optimize(5, |trial| { + let val = x.suggest(trial)?; + Ok::<_, Error>(val * val) + }) + .unwrap(); + + let summary = study.summary(); + assert!(summary.contains("Minimize")); + assert!(summary.contains("5 trials")); + assert!(summary.contains("Best value:")); + assert!(summary.contains("x = ")); +} + +#[test] +fn test_summary_no_completed_trials() { + let study: Study = Study::new(Direction::Maximize); + let summary = study.summary(); + assert!(summary.contains("Maximize")); + assert!(summary.contains("0 trials")); + assert!(!summary.contains("Best value:")); +} + +#[test] +fn test_summary_with_pruned_trials() { + let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(1)); + let x = FloatParam::new(0.0, 10.0).name("x"); + + // Manually create some complete and pruned trials + for _ in 0..3 { + let mut trial = study.create_trial(); + let val = x.suggest(&mut trial).unwrap(); + study.complete_trial(trial, val); + } + for _ in 0..2 { + let mut trial = study.create_trial(); + let _ = x.suggest(&mut trial).unwrap(); + study.prune_trial(trial); + } + + let summary = study.summary(); + // Should show breakdown when there are pruned trials + if study.n_pruned_trials() > 0 { + assert!(summary.contains("complete")); + assert!(summary.contains("pruned")); + } +} + +#[test] +fn test_display_matches_summary() { + let study: Study = Study::with_sampler(Direction::Minimize, RandomSampler::with_seed(1)); + let x = FloatParam::new(0.0, 10.0).name("x"); + + study + .optimize(3, |trial| { + let val = x.suggest(trial)?; + Ok::<_, Error>(val) + }) + .unwrap(); + + assert_eq!(format!("{study}"), study.summary()); +}