1use crate::event_bus::EventBus;
11use crate::pruner::{MedianPruner, PercentilePruner, Pruner, TrialMetricHistory};
12use crate::sampler::Sampler;
13use chrono::Utc;
14use somatize_core::error::Result;
15use somatize_core::event::{Event, MetricRecord};
16use somatize_core::study::{Direction, PruningStrategy, Study, Trial, TrialState};
17use somatize_core::tracking::Tracker;
18use std::collections::HashMap;
19use std::sync::{Arc, Mutex};
20use std::time::Instant;
21
22#[derive(Debug, Clone)]
24pub enum TrialOutcome {
25 Completed(Vec<MetricRecord>),
27 Pruned {
29 step: usize,
31 reason: String,
33 },
34}
35
36#[derive(Clone)]
48pub struct TrialContext {
49 study_id: String,
50 trial_id: String,
51 objective: Option<(String, Direction)>,
53 pruner: Option<Arc<dyn Pruner>>,
54 history: Arc<Vec<TrialMetricHistory>>,
56 event_bus: Arc<EventBus>,
57 shared: Arc<Mutex<TrialShared>>,
58}
59
60#[derive(Default)]
61struct TrialShared {
62 metrics: Vec<MetricRecord>,
63 pruned: Option<(usize, String)>,
64}
65
66impl TrialContext {
67 pub fn report(&self, name: &str, value: f64, step: usize) -> bool {
71 let record = MetricRecord {
72 name: name.to_string(),
73 value,
74 step,
75 timestamp: Utc::now(),
76 };
77 {
78 let mut shared = self.lock_shared();
79 shared.metrics.push(record.clone());
80 }
81 self.event_bus.emit(Event::TrialMetric {
82 study_id: self.study_id.clone(),
83 trial_id: self.trial_id.clone(),
84 metric: record,
85 });
86
87 if self.should_prune() {
88 return true;
89 }
90 if let (Some(pruner), Some((obj_name, direction))) = (&self.pruner, &self.objective)
91 && name == obj_name
92 && let Some(reason) =
93 pruner.should_prune(obj_name, direction.normalize(value), step, &self.history)
94 {
95 self.lock_shared().pruned = Some((step, reason));
96 return true;
97 }
98 false
99 }
100
101 pub fn should_prune(&self) -> bool {
103 self.lock_shared().pruned.is_some()
104 }
105
106 pub fn metrics(&self) -> Vec<MetricRecord> {
108 self.lock_shared().metrics.clone()
109 }
110
111 pub fn trial_id(&self) -> &str {
113 &self.trial_id
114 }
115
116 fn lock_shared(&self) -> std::sync::MutexGuard<'_, TrialShared> {
117 match self.shared.lock() {
118 Ok(guard) => guard,
119 Err(poisoned) => poisoned.into_inner(),
120 }
121 }
122
123 fn take_results(&self) -> (Vec<MetricRecord>, Option<(usize, String)>) {
124 let mut shared = self.lock_shared();
125 (std::mem::take(&mut shared.metrics), shared.pruned.take())
126 }
127}
128
129pub trait TrialExecutor: Send + Sync {
134 fn execute_trial(
137 &self,
138 params: &HashMap<String, serde_json::Value>,
139 ctx: &TrialContext,
140 ) -> Result<TrialOutcome>;
141}
142
143pub struct FnTrialExecutor<F>(pub F);
145
146impl<F> TrialExecutor for FnTrialExecutor<F>
147where
148 F: Fn(&HashMap<String, serde_json::Value>, &TrialContext) -> Result<TrialOutcome> + Send + Sync,
149{
150 fn execute_trial(
151 &self,
152 params: &HashMap<String, serde_json::Value>,
153 ctx: &TrialContext,
154 ) -> Result<TrialOutcome> {
155 (self.0)(params, ctx)
156 }
157}
158
159pub struct StudyRunner {
161 event_bus: Arc<EventBus>,
162 tracker: Option<Arc<dyn Tracker>>,
163}
164
165impl StudyRunner {
166 pub fn new(event_bus: Arc<EventBus>) -> Self {
169 Self {
170 event_bus,
171 tracker: None,
172 }
173 }
174
175 pub fn with_tracker(mut self, tracker: Arc<dyn Tracker>) -> Self {
177 self.tracker = Some(tracker);
178 self
179 }
180
181 pub fn run(
188 &self,
189 study: &mut Study,
190 sampler: &mut dyn Sampler,
191 executor: &dyn TrialExecutor,
192 ) -> Result<()> {
193 sampler.prepare(&study.search_space);
194 let total = sampler.n_trials().unwrap_or(0) * study.seeds.len().max(1);
196 if total > 0 {
197 study.planned_trials = Some(total);
198 }
199 let direction = study.primary_direction().unwrap_or(Direction::Maximize);
200 if study.created_at.is_none() {
201 study.created_at = Some(Utc::now());
202 }
203
204 for trial in &study.trials {
206 if let Some(value) = study.objective_value(trial) {
207 sampler.record_result(&trial.params, direction.normalize(value));
208 }
209 }
210
211 self.event_bus.emit(Event::StudyStarted {
212 study_id: study.id.clone(),
213 name: study.name.clone(),
214 total_trials: total,
215 });
216
217 let pruner = build_pruner(&study.pruning);
218 let mut trial_index = study.trials.len();
219
220 let seeds = study.seeds.clone();
225 let n_seeds = seeds.len().max(1);
226 let mut current_config: Option<HashMap<String, serde_json::Value>> = None;
227
228 loop {
229 let config_index = trial_index / n_seeds;
230 let seed_slot = trial_index % n_seeds;
231
232 let base = if seed_slot == 0 || current_config.is_none() {
233 if seed_slot > 0 {
234 let mut prev = study.trials[trial_index - 1].params.clone();
237 prev.remove("seed");
238 Some(prev)
239 } else {
240 sampler.sample(&study.search_space, config_index)?
241 }
242 } else {
243 current_config.clone()
244 };
245 let Some(base) = base else { break };
246 current_config = Some(base.clone());
247
248 let mut params = base;
249 if !seeds.is_empty() {
250 params.insert("seed".to_string(), serde_json::json!(seeds[seed_slot]));
251 }
252 for (name, value) in &study.frozen {
255 params.insert(name.clone(), value.clone());
256 }
257
258 let trial_id = format!("trial_{trial_index:04}");
259 let mut trial = Trial::new(trial_id.clone(), params.clone());
260 trial.state = TrialState::Running;
261 trial.started_at = Some(Utc::now());
262
263 self.event_bus.emit(Event::TrialStarted {
264 study_id: study.id.clone(),
265 trial_id: trial_id.clone(),
266 params: serde_json::json!(params),
267 });
268
269 let ctx = TrialContext {
271 study_id: study.id.clone(),
272 trial_id: trial_id.clone(),
273 objective: objective_metric(study).map(|name| (name, direction)),
274 pruner: pruner.clone(),
275 history: Arc::new(normalized_histories(study, direction)),
276 event_bus: self.event_bus.clone(),
277 shared: Arc::new(Mutex::new(TrialShared::default())),
278 };
279
280 let start = Instant::now();
281 let outcome = executor.execute_trial(¶ms, &ctx);
282 let (reported, pruned) = ctx.take_results();
283 trial.duration_ms = Some(start.elapsed().as_millis() as u64);
284 trial.finished_at = Some(Utc::now());
285 trial.metrics = reported;
286
287 match (outcome, pruned) {
288 (Ok(_), Some((step, reason)))
291 | (Ok(TrialOutcome::Pruned { step, reason }), None) => {
292 trial.state = TrialState::Pruned {
293 step,
294 reason: reason.clone(),
295 };
296 self.event_bus.emit(Event::TrialPruned {
297 study_id: study.id.clone(),
298 trial_id: trial_id.clone(),
299 step,
300 reason,
301 });
302 }
303 (Ok(TrialOutcome::Completed(final_metrics)), None) => {
304 for metric in &final_metrics {
305 self.event_bus.emit(Event::TrialMetric {
306 study_id: study.id.clone(),
307 trial_id: trial_id.clone(),
308 metric: metric.clone(),
309 });
310 }
311 trial.metrics.extend(final_metrics.clone());
312 trial.state = TrialState::Completed;
313
314 self.event_bus.emit(Event::TrialCompleted {
315 study_id: study.id.clone(),
316 trial_id: trial_id.clone(),
317 final_metrics,
318 });
319 }
320 (Err(e), _) => {
321 trial.state = TrialState::Failed {
322 error: e.to_string(),
323 };
324 self.event_bus.emit(Event::TrialFailed {
325 study_id: study.id.clone(),
326 trial_id: trial_id.clone(),
327 error: e.to_string(),
328 });
329 }
330 }
331
332 study.trials.push(trial);
333
334 if let Some(value) = study.objective_value(study.trials.last().unwrap()) {
336 sampler.record_result(¶ms, direction.normalize(value));
337 }
338
339 if let Some(best) = study.best_trial()
341 && best.id == trial_id
342 && let Some(val) = study.best_value()
343 {
344 self.event_bus.emit(Event::BestUpdated {
345 study_id: study.id.clone(),
346 trial_id: trial_id.clone(),
347 value: val,
348 params: serde_json::json!(params),
349 });
350 }
351
352 let completed = study.trials.iter().filter(|t| t.is_terminal()).count();
353 self.event_bus.emit(Event::StudyProgress {
354 study_id: study.id.clone(),
355 completed,
356 total,
357 best_value: study.best_value().unwrap_or(f64::NAN),
358 });
359
360 study.updated_at = Some(Utc::now());
361 self.save_study(study);
362 trial_index += 1;
363 }
364
365 let best_trial_id = study.best_trial().map(|t| t.id.clone()).unwrap_or_default();
366 let best_value = study.best_value().unwrap_or(f64::NAN);
367
368 self.event_bus.emit(Event::StudyCompleted {
369 study_id: study.id.clone(),
370 best_trial_id,
371 best_value,
372 });
373 study.updated_at = Some(Utc::now());
374 self.save_study(study);
375 self.event_bus.flush_sinks();
378
379 Ok(())
380 }
381
382 fn save_study(&self, study: &Study) {
383 if let Some(tracker) = &self.tracker
384 && let Err(e) = tracker.save_study(study)
385 {
386 tracing::warn!("tracking: failed to persist study: {e}");
387 }
388 }
389}
390
391fn objective_metric(study: &Study) -> Option<String> {
395 study
396 .objectives
397 .first()
398 .map(|o| o.metric.clone())
399 .or_else(|| {
400 study
401 .composite
402 .as_ref()
403 .and_then(|c| c.terms.first().map(|(name, _)| name.clone()))
404 })
405}
406
407fn build_pruner(strategy: &PruningStrategy) -> Option<Arc<dyn Pruner>> {
408 match strategy {
409 PruningStrategy::None | PruningStrategy::Hyperband => None,
410 PruningStrategy::Median { n_warmup_steps } => {
411 Some(Arc::new(MedianPruner::new(*n_warmup_steps)))
412 }
413 PruningStrategy::Percentile {
414 percentile,
415 n_warmup_steps,
416 } => Some(Arc::new(PercentilePruner::new(
417 *percentile,
418 *n_warmup_steps,
419 ))),
420 }
421}
422
423fn normalized_histories(study: &Study, direction: Direction) -> Vec<TrialMetricHistory> {
426 study
427 .trials
428 .iter()
429 .filter(|t| t.is_complete())
430 .map(|t| TrialMetricHistory {
431 trial_id: t.id.clone(),
432 metrics: t
433 .metrics
434 .iter()
435 .map(|m| MetricRecord {
436 name: m.name.clone(),
437 value: direction.normalize(m.value),
438 step: m.step,
439 timestamp: m.timestamp,
440 })
441 .collect(),
442 })
443 .collect()
444}
445
446#[cfg(test)]
447mod tests {
448 use super::*;
449 use crate::sampler::{BayesianSampler, GridSampler, RandomSampler};
450 use crate::study_io::StudyIo;
451 use chrono::Utc;
452 use somatize_core::error::SomaError;
453 use somatize_core::search::{Scale, SearchDimension, SearchSpace};
454 use somatize_core::study::{Direction, Objective, SearchStrategy};
455
456 fn sample_space() -> SearchSpace {
457 let mut space = SearchSpace::new();
458 space.add(SearchDimension::Float {
459 name: "lr".into(),
460 low: 0.001,
461 high: 0.1,
462 scale: Scale::Log,
463 default: None,
464 });
465 space.add(SearchDimension::Categorical {
466 name: "activation".into(),
467 choices: vec![serde_json::json!("relu"), serde_json::json!("tanh")],
468 });
469 space
470 }
471
472 fn make_executor() -> FnTrialExecutor<
474 impl Fn(&HashMap<String, serde_json::Value>, &TrialContext) -> Result<TrialOutcome>,
475 > {
476 FnTrialExecutor(
477 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
478 let lr = params["lr"].as_f64().unwrap();
479 let f1 = (1.0 - (lr - 0.01).abs() * 10.0).max(0.0);
480 Ok(TrialOutcome::Completed(vec![MetricRecord {
481 name: "f1".into(),
482 value: f1,
483 step: 0,
484 timestamp: Utc::now(),
485 }]))
486 },
487 )
488 }
489
490 #[test]
491 fn study_runner_grid_search() {
492 let bus = Arc::new(EventBus::new(256));
493 let mut rx = bus.subscribe();
494 let runner = StudyRunner::new(bus);
495
496 let space = sample_space();
497 let mut study = Study::new(
498 "grid_test",
499 space,
500 SearchStrategy::Grid { points_per_dim: 3 },
501 vec![Objective {
502 metric: "f1".into(),
503 direction: Direction::Maximize,
504 }],
505 );
506
507 let mut sampler = GridSampler::new(3);
508 let executor = make_executor();
509
510 runner.run(&mut study, &mut sampler, &executor).unwrap();
511
512 assert_eq!(study.trials.len(), 6);
514 assert!(study.trials.iter().all(|t| t.is_complete()));
515
516 let best = study.best_trial().unwrap();
518 let best_lr = best.params["lr"].as_f64().unwrap();
519 assert!(
520 (best_lr - 0.01).abs() < 0.05,
521 "best lr should be near 0.01, got {best_lr}"
522 );
523
524 let mut events = Vec::new();
526 while let Ok(e) = rx.try_recv() {
527 events.push(e);
528 }
529 assert!(
530 events
531 .iter()
532 .any(|e| matches!(e, Event::StudyStarted { .. }))
533 );
534 assert!(
535 events
536 .iter()
537 .any(|e| matches!(e, Event::TrialStarted { .. }))
538 );
539 assert!(
540 events
541 .iter()
542 .any(|e| matches!(e, Event::TrialCompleted { .. }))
543 );
544 assert!(
545 events
546 .iter()
547 .any(|e| matches!(e, Event::BestUpdated { .. }))
548 );
549 assert!(
550 events
551 .iter()
552 .any(|e| matches!(e, Event::StudyCompleted { .. }))
553 );
554 }
555
556 #[test]
557 fn study_runner_random_search() {
558 let bus = Arc::new(EventBus::new(256));
559 let runner = StudyRunner::new(bus);
560
561 let space = sample_space();
562 let mut study = Study::new(
563 "random_test",
564 space,
565 SearchStrategy::Random {
566 n_trials: 20,
567 seed: Some(42),
568 },
569 vec![Objective {
570 metric: "f1".into(),
571 direction: Direction::Maximize,
572 }],
573 );
574
575 let mut sampler = RandomSampler::new(20, Some(42));
576 let executor = make_executor();
577
578 runner.run(&mut study, &mut sampler, &executor).unwrap();
579
580 assert_eq!(study.trials.len(), 20);
581 assert!(study.best_trial().is_some());
582 }
583
584 #[test]
585 fn study_runner_handles_failed_trials() {
586 let bus = Arc::new(EventBus::new(256));
587 let runner = StudyRunner::new(bus);
588
589 let mut space = SearchSpace::new();
590 space.add(SearchDimension::Float {
591 name: "x".into(),
592 low: 0.0,
593 high: 1.0,
594 scale: Scale::Linear,
595 default: None,
596 });
597
598 let mut study = Study::new(
599 "fail_test",
600 space,
601 SearchStrategy::Random {
602 n_trials: 5,
603 seed: None,
604 },
605 vec![Objective {
606 metric: "f1".into(),
607 direction: Direction::Maximize,
608 }],
609 );
610
611 let executor = FnTrialExecutor(
613 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
614 let x = params["x"].as_f64().unwrap();
615 if x > 0.5 {
616 Err(SomaError::Other("too high".into()))
617 } else {
618 Ok(TrialOutcome::Completed(vec![MetricRecord {
619 name: "f1".into(),
620 value: x,
621 step: 0,
622 timestamp: Utc::now(),
623 }]))
624 }
625 },
626 );
627
628 let mut sampler = RandomSampler::new(5, Some(42));
629 runner.run(&mut study, &mut sampler, &executor).unwrap();
630
631 assert_eq!(study.trials.len(), 5);
632 let failed = study
634 .trials
635 .iter()
636 .filter(|t| matches!(t.state, TrialState::Failed { .. }))
637 .count();
638 assert!(failed > 0, "should have some failed trials");
639 }
640
641 #[test]
642 fn study_runner_handles_pruned_trials() {
643 let bus = Arc::new(EventBus::new(256));
644 let runner = StudyRunner::new(bus);
645
646 let mut space = SearchSpace::new();
647 space.add(SearchDimension::Float {
648 name: "x".into(),
649 low: 0.0,
650 high: 1.0,
651 scale: Scale::Linear,
652 default: None,
653 });
654
655 let mut study = Study::new(
656 "prune_test",
657 space,
658 SearchStrategy::Random {
659 n_trials: 3,
660 seed: None,
661 },
662 vec![Objective {
663 metric: "f1".into(),
664 direction: Direction::Maximize,
665 }],
666 );
667
668 let executor = FnTrialExecutor(
670 |_params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
671 Ok(TrialOutcome::Pruned {
672 step: 5,
673 reason: "below median".into(),
674 })
675 },
676 );
677
678 let mut sampler = RandomSampler::new(3, Some(42));
679 runner.run(&mut study, &mut sampler, &executor).unwrap();
680
681 assert!(
682 study
683 .trials
684 .iter()
685 .all(|t| matches!(t.state, TrialState::Pruned { .. }))
686 );
687 }
688
689 fn one_dim_space() -> SearchSpace {
690 let mut space = SearchSpace::new();
691 space.add(SearchDimension::Float {
692 name: "x".into(),
693 low: 0.0,
694 high: 1.0,
695 scale: Scale::Linear,
696 default: None,
697 });
698 space
699 }
700
701 #[test]
702 fn grid_study_started_reports_real_total() {
703 let bus = Arc::new(EventBus::new(256));
704 let mut rx = bus.subscribe();
705 let runner = StudyRunner::new(bus);
706
707 let mut study = Study::new(
708 "grid_total",
709 sample_space(),
710 SearchStrategy::Grid { points_per_dim: 3 },
711 vec![Objective {
712 metric: "f1".into(),
713 direction: Direction::Maximize,
714 }],
715 );
716 let mut sampler = GridSampler::new(3);
717 runner
718 .run(&mut study, &mut sampler, &make_executor())
719 .unwrap();
720
721 let mut started_total = None;
722 while let Ok(e) = rx.try_recv() {
723 if let Event::StudyStarted { total_trials, .. } = e {
724 started_total = Some(total_trials);
725 }
726 }
727 assert_eq!(started_total, Some(6));
729 }
730
731 struct SpySampler {
733 inner: RandomSampler,
734 recorded: Vec<f64>,
735 }
736
737 impl Sampler for SpySampler {
738 fn sample(
739 &mut self,
740 space: &SearchSpace,
741 trial_index: usize,
742 ) -> Result<Option<HashMap<String, serde_json::Value>>> {
743 self.inner.sample(space, trial_index)
744 }
745 fn n_trials(&self) -> Option<usize> {
746 self.inner.n_trials()
747 }
748 fn record_result(&mut self, _params: &HashMap<String, serde_json::Value>, value: f64) {
749 self.recorded.push(value);
750 }
751 }
752
753 #[test]
754 fn sampler_receives_feedback_per_completed_trial() {
755 let bus = Arc::new(EventBus::new(256));
756 let runner = StudyRunner::new(bus);
757
758 let mut study = Study::new(
759 "feedback",
760 one_dim_space(),
761 SearchStrategy::Random {
762 n_trials: 4,
763 seed: Some(1),
764 },
765 vec![Objective {
766 metric: "loss".into(),
767 direction: Direction::Minimize,
768 }],
769 );
770 let executor = FnTrialExecutor(
771 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
772 Ok(TrialOutcome::Completed(vec![MetricRecord {
773 name: "loss".into(),
774 value: params["x"].as_f64().unwrap(),
775 step: 0,
776 timestamp: Utc::now(),
777 }]))
778 },
779 );
780
781 let mut sampler = SpySampler {
782 inner: RandomSampler::new(4, Some(1)),
783 recorded: Vec::new(),
784 };
785 runner.run(&mut study, &mut sampler, &executor).unwrap();
786
787 assert_eq!(
788 sampler.recorded.len(),
789 4,
790 "one feedback per completed trial"
791 );
792 assert!(sampler.recorded.iter().all(|v| *v <= 0.0));
794 }
795
796 fn pruning_study(direction: Direction) -> Study {
797 let metric = match direction {
798 Direction::Maximize => "f1",
799 Direction::Minimize => "loss",
800 };
801 Study::new(
802 "pruning",
803 one_dim_space(),
804 SearchStrategy::Random {
805 n_trials: 4,
806 seed: Some(7),
807 },
808 vec![Objective {
809 metric: metric.into(),
810 direction,
811 }],
812 )
813 .with_pruning(PruningStrategy::Median { n_warmup_steps: 2 })
814 }
815
816 fn run_pruning_case(direction: Direction) -> Study {
819 use std::sync::atomic::{AtomicUsize, Ordering};
820
821 let bus = Arc::new(EventBus::new(256));
822 let runner = StudyRunner::new(bus);
823 let mut study = pruning_study(direction);
824 let counter = AtomicUsize::new(0);
825
826 let executor = FnTrialExecutor(
827 move |_params: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
828 let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
829 let metric = match direction {
830 Direction::Maximize => "f1",
831 Direction::Minimize => "loss",
832 };
833 for step in 0..10 {
834 let value = match (direction, good) {
835 (Direction::Maximize, true) => 0.5 + step as f64 * 0.05,
836 (Direction::Maximize, false) => 0.01,
837 (Direction::Minimize, true) => 1.0 - step as f64 * 0.05,
838 (Direction::Minimize, false) => 10.0,
839 };
840 if ctx.report(metric, value, step) {
841 return Ok(TrialOutcome::Pruned {
842 step,
843 reason: "stopped by pruner".into(),
844 });
845 }
846 }
847 Ok(TrialOutcome::Completed(vec![]))
848 },
849 );
850
851 let mut sampler = RandomSampler::new(4, Some(7));
852 runner.run(&mut study, &mut sampler, &executor).unwrap();
853 study
854 }
855
856 fn assert_bad_trials_pruned(study: &Study) {
857 let pruned: Vec<_> = study
858 .trials
859 .iter()
860 .filter(|t| matches!(t.state, TrialState::Pruned { .. }))
861 .collect();
862 assert_eq!(pruned.len(), 3, "all trials after the first get pruned");
863 for t in pruned {
864 if let TrialState::Pruned { step, .. } = &t.state {
865 assert_eq!(*step, 2, "pruned right after warmup, not at the end");
866 }
867 }
868 assert!(study.trials[0].is_complete());
869 }
870
871 #[test]
872 fn median_pruner_stops_bad_trials_maximize() {
873 assert_bad_trials_pruned(&run_pruning_case(Direction::Maximize));
874 }
875
876 #[test]
877 fn median_pruner_stops_bad_trials_minimize() {
878 assert_bad_trials_pruned(&run_pruning_case(Direction::Minimize));
880 }
881
882 #[test]
883 fn resume_continues_without_repeating_grid_params() {
884 let executor = make_executor();
885
886 let bus = Arc::new(EventBus::new(256));
888 let runner = StudyRunner::new(bus);
889 let mut full = Study::new(
890 "full",
891 sample_space(),
892 SearchStrategy::Grid { points_per_dim: 3 },
893 vec![Objective {
894 metric: "f1".into(),
895 direction: Direction::Maximize,
896 }],
897 );
898 runner
899 .run(&mut full, &mut GridSampler::new(3), &executor)
900 .unwrap();
901 assert_eq!(full.trials.len(), 6);
902
903 let mut resumed = full.clone();
905 resumed.trials.truncate(3);
906
907 let bus = Arc::new(EventBus::new(256));
908 let mut rx = bus.subscribe();
909 let runner = StudyRunner::new(bus);
910 runner
911 .run(&mut resumed, &mut GridSampler::new(3), &executor)
912 .unwrap();
913
914 assert_eq!(resumed.trials.len(), 6);
915 for i in 3..6 {
918 assert_eq!(resumed.trials[i].params, full.trials[i].params);
919 assert_eq!(resumed.trials[i].id, format!("trial_{i:04}"));
920 }
921 let mut started = 0;
923 while let Ok(e) = rx.try_recv() {
924 if matches!(e, Event::TrialStarted { .. }) {
925 started += 1;
926 }
927 }
928 assert_eq!(started, 3);
929 }
930
931 #[test]
932 fn frozen_params_reach_every_trial() {
933 let bus = Arc::new(EventBus::new(256));
934 let runner = StudyRunner::new(bus);
935
936 let mut study = Study::new(
937 "frozen",
938 one_dim_space(),
939 SearchStrategy::Random {
940 n_trials: 3,
941 seed: Some(2),
942 },
943 vec![Objective {
944 metric: "f1".into(),
945 direction: Direction::Maximize,
946 }],
947 );
948 study
949 .frozen
950 .insert("batch_size".into(), serde_json::json!(64));
951
952 let executor = FnTrialExecutor(
953 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
954 assert_eq!(params["batch_size"], serde_json::json!(64));
955 Ok(TrialOutcome::Completed(vec![MetricRecord {
956 name: "f1".into(),
957 value: 0.5,
958 step: 0,
959 timestamp: Utc::now(),
960 }]))
961 },
962 );
963 let mut sampler = RandomSampler::new(3, Some(2));
964 runner.run(&mut study, &mut sampler, &executor).unwrap();
965
966 for t in &study.trials {
967 assert_eq!(t.params["batch_size"], serde_json::json!(64));
968 }
969 }
970
971 #[test]
972 fn composite_objective_selects_best_trial() {
973 use somatize_core::study::{CompositeObjective, Scalarizer};
974
975 let bus = Arc::new(EventBus::new(256));
976 let runner = StudyRunner::new(bus);
977
978 let mut study = Study::new(
979 "composite",
980 one_dim_space(),
981 SearchStrategy::Grid { points_per_dim: 5 },
982 vec![],
983 )
984 .with_composite(CompositeObjective {
985 terms: vec![("x".into(), 1.0), ("x_sq".into(), -1.0)],
987 direction: Direction::Maximize,
988 scalarizer: Scalarizer::WeightedSum,
989 });
990
991 let executor = FnTrialExecutor(
992 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
993 let x = params["x"].as_f64().unwrap();
994 let now = Utc::now();
995 Ok(TrialOutcome::Completed(vec![
996 MetricRecord {
997 name: "x".into(),
998 value: x,
999 step: 0,
1000 timestamp: now,
1001 },
1002 MetricRecord {
1003 name: "x_sq".into(),
1004 value: x * x,
1005 step: 0,
1006 timestamp: now,
1007 },
1008 ]))
1009 },
1010 );
1011 let mut sampler = GridSampler::new(5);
1012 runner.run(&mut study, &mut sampler, &executor).unwrap();
1013
1014 let best_x = study.best_trial().unwrap().params["x"].as_f64().unwrap();
1015 assert!(
1016 (best_x - 0.5).abs() < 1e-9,
1017 "optimum of x - x² is 0.5, got {best_x}"
1018 );
1019 }
1020
1021 #[test]
1022 fn tracker_persists_study_after_every_trial() {
1023 use somatize_core::tracking::{RunKind, Tracker as _};
1024
1025 let root = tempfile::tempdir().unwrap();
1026 let tracker = Arc::new(
1027 crate::tracking::LocalTracker::create(root.path(), RunKind::Study, "t").unwrap(),
1028 );
1029 let run_dir = tracker.run_dir().to_path_buf();
1030
1031 let bus = Arc::new(EventBus::new(256));
1032 bus.add_sink(tracker.sink());
1033 let runner = StudyRunner::new(bus).with_tracker(tracker);
1034
1035 let mut study = Study::new(
1036 "persisted",
1037 one_dim_space(),
1038 SearchStrategy::Random {
1039 n_trials: 3,
1040 seed: Some(3),
1041 },
1042 vec![Objective {
1043 metric: "f1".into(),
1044 direction: Direction::Maximize,
1045 }],
1046 );
1047 let mut sampler = RandomSampler::new(3, Some(3));
1048 runner
1049 .run(&mut study, &mut sampler, &make_simple_executor())
1050 .unwrap();
1051
1052 let loaded = Study::load(run_dir.join("study.json")).unwrap();
1054 assert_eq!(loaded.trials.len(), 3);
1055 assert!(loaded.updated_at.is_some());
1056 let events = std::fs::read_to_string(run_dir.join("events.jsonl")).unwrap();
1058 assert!(events.contains("TrialCompleted"));
1059 }
1060
1061 fn make_simple_executor() -> FnTrialExecutor<
1062 impl Fn(&HashMap<String, serde_json::Value>, &TrialContext) -> Result<TrialOutcome>,
1063 > {
1064 FnTrialExecutor(
1065 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1066 Ok(TrialOutcome::Completed(vec![MetricRecord {
1067 name: "f1".into(),
1068 value: params["x"].as_f64().unwrap_or(0.5),
1069 step: 0,
1070 timestamp: Utc::now(),
1071 }]))
1072 },
1073 )
1074 }
1075
1076 use crate::pruner::TrialMetricHistory as History;
1079 use std::sync::atomic::{AtomicUsize, Ordering};
1080
1081 struct SpyPruner {
1083 calls: AtomicUsize,
1084 prune_from_step: usize,
1085 }
1086
1087 impl Pruner for SpyPruner {
1088 fn should_prune(
1089 &self,
1090 _metric: &str,
1091 _value: f64,
1092 step: usize,
1093 _history: &[History],
1094 ) -> Option<String> {
1095 self.calls.fetch_add(1, Ordering::SeqCst);
1096 (step >= self.prune_from_step).then(|| "spy".to_string())
1097 }
1098 }
1099
1100 fn make_ctx(pruner: Option<Arc<SpyPruner>>) -> (TrialContext, Arc<EventBus>) {
1101 let bus = Arc::new(EventBus::new(64));
1102 let ctx = TrialContext {
1103 study_id: "s".into(),
1104 trial_id: "trial_0000".into(),
1105 objective: Some(("f1".to_string(), Direction::Maximize)),
1106 pruner: pruner.map(|p| p as Arc<dyn Pruner>),
1107 history: Arc::new(Vec::new()),
1108 event_bus: bus.clone(),
1109 shared: Arc::new(Mutex::new(TrialShared::default())),
1110 };
1111 (ctx, bus)
1112 }
1113
1114 #[test]
1115 fn report_after_prune_is_sticky_and_skips_the_pruner() {
1116 let pruner = Arc::new(SpyPruner {
1117 calls: AtomicUsize::new(0),
1118 prune_from_step: 0,
1119 });
1120 let (ctx, _bus) = make_ctx(Some(pruner.clone()));
1121
1122 assert!(!ctx.should_prune());
1123 assert!(ctx.report("f1", 0.1, 0), "pruned immediately");
1124 assert!(ctx.should_prune());
1125 let calls_after_prune = pruner.calls.load(Ordering::SeqCst);
1126
1127 assert!(ctx.report("f1", 0.9, 1));
1129 assert!(ctx.report("f1", 0.9, 2));
1130 assert_eq!(pruner.calls.load(Ordering::SeqCst), calls_after_prune);
1131 assert_eq!(ctx.metrics().len(), 3);
1134 }
1135
1136 #[test]
1137 fn non_objective_metric_never_consults_the_pruner() {
1138 let pruner = Arc::new(SpyPruner {
1139 calls: AtomicUsize::new(0),
1140 prune_from_step: 0,
1141 });
1142 let (ctx, _bus) = make_ctx(Some(pruner.clone()));
1143
1144 assert!(!ctx.report("train_loss", 99.0, 0));
1145 assert!(!ctx.report("lr", 0.001, 0));
1146 assert_eq!(pruner.calls.load(Ordering::SeqCst), 0);
1147 assert!(ctx.report("f1", 0.0, 0), "objective metric does consult");
1148 assert_eq!(pruner.calls.load(Ordering::SeqCst), 1);
1149 }
1150
1151 #[test]
1152 fn report_emits_trial_metric_events() {
1153 let (ctx, bus) = make_ctx(None);
1154 let mut rx = bus.subscribe();
1155 ctx.report("f1", 0.5, 3);
1156 ctx.report("aux", 1.5, 3);
1157
1158 let mut got = Vec::new();
1159 while let Ok(e) = rx.try_recv() {
1160 if let Event::TrialMetric {
1161 study_id,
1162 trial_id,
1163 metric,
1164 } = e
1165 {
1166 assert_eq!(study_id, "s");
1167 assert_eq!(trial_id, "trial_0000");
1168 got.push((metric.name, metric.value, metric.step));
1169 }
1170 }
1171 assert_eq!(
1172 got,
1173 vec![("f1".to_string(), 0.5, 3), ("aux".to_string(), 1.5, 3)]
1174 );
1175 }
1176
1177 #[test]
1178 fn trial_context_accessors_and_cross_thread_clone() {
1179 let (ctx, _bus) = make_ctx(None);
1180 assert_eq!(ctx.trial_id(), "trial_0000");
1181
1182 let clone = ctx.clone();
1185 let handle = std::thread::spawn(move || {
1186 clone.report("from_thread", 1.0, 0);
1187 });
1188 ctx.report("from_main", 2.0, 0);
1189 handle.join().unwrap();
1190
1191 let names: Vec<String> = ctx.metrics().into_iter().map(|m| m.name).collect();
1192 assert_eq!(names.len(), 2);
1193 assert!(names.contains(&"from_thread".to_string()));
1194 assert!(names.contains(&"from_main".to_string()));
1195 }
1196
1197 #[test]
1200 fn percentile_pruning_works_through_the_runner() {
1201 use std::sync::atomic::AtomicUsize;
1204 let bus = Arc::new(EventBus::new(256));
1205 let runner = StudyRunner::new(bus);
1206 let mut study = Study::new(
1207 "percentile",
1208 one_dim_space(),
1209 SearchStrategy::Random {
1210 n_trials: 4,
1211 seed: Some(7),
1212 },
1213 vec![Objective {
1214 metric: "f1".into(),
1215 direction: Direction::Maximize,
1216 }],
1217 )
1218 .with_pruning(PruningStrategy::Percentile {
1219 percentile: 50.0,
1220 n_warmup_steps: 2,
1221 });
1222 let counter = AtomicUsize::new(0);
1223 let executor = FnTrialExecutor(
1224 move |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1225 let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
1226 for step in 0..10 {
1227 let value = if good { 0.5 + step as f64 * 0.05 } else { 0.01 };
1228 if ctx.report("f1", value, step) {
1229 return Ok(TrialOutcome::Pruned {
1230 step,
1231 reason: "stopped".into(),
1232 });
1233 }
1234 }
1235 Ok(TrialOutcome::Completed(vec![]))
1236 },
1237 );
1238 let mut sampler = RandomSampler::new(4, Some(7));
1239 runner.run(&mut study, &mut sampler, &executor).unwrap();
1240 assert_bad_trials_pruned(&study);
1241 }
1242
1243 #[test]
1244 fn pruner_verdict_wins_over_completed_outcome() {
1245 use std::sync::atomic::AtomicUsize;
1249 let bus = Arc::new(EventBus::new(256));
1250 let mut rx = bus.subscribe();
1251 let runner = StudyRunner::new(bus);
1252 let mut study = pruning_study(Direction::Maximize);
1253
1254 let counter = AtomicUsize::new(0);
1255 let executor = FnTrialExecutor(
1256 move |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1257 let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
1258 for step in 0..10 {
1259 let value = if good { 0.5 + step as f64 * 0.05 } else { 0.01 };
1260 ctx.report("f1", value, step); }
1262 Ok(TrialOutcome::Completed(vec![MetricRecord {
1263 name: "sneaky".into(),
1264 value: 1.0,
1265 step: 0,
1266 timestamp: Utc::now(),
1267 }]))
1268 },
1269 );
1270 let mut sampler = RandomSampler::new(4, Some(7));
1271 runner.run(&mut study, &mut sampler, &executor).unwrap();
1272
1273 let pruned: Vec<_> = study
1274 .trials
1275 .iter()
1276 .filter(|t| matches!(t.state, TrialState::Pruned { .. }))
1277 .collect();
1278 assert_eq!(pruned.len(), 3);
1279 for t in &pruned {
1280 if let TrialState::Pruned { step, reason } = &t.state {
1281 assert_eq!(*step, 2, "pruner's step, not the executor's");
1282 assert!(reason.contains("median"), "pruner's reason: {reason}");
1283 }
1284 assert!(
1285 !t.metrics.iter().any(|m| m.name == "sneaky"),
1286 "final metrics of an overridden Completed are dropped"
1287 );
1288 }
1289 let mut completed = 0;
1291 let mut pruned_events = 0;
1292 while let Ok(e) = rx.try_recv() {
1293 match e {
1294 Event::TrialCompleted { .. } => completed += 1,
1295 Event::TrialPruned { .. } => pruned_events += 1,
1296 _ => {}
1297 }
1298 }
1299 assert_eq!(completed, 1);
1300 assert_eq!(pruned_events, 3);
1301 }
1302
1303 #[derive(Default)]
1305 struct SpyTracker {
1306 saves: Mutex<Vec<usize>>, fail: bool,
1308 }
1309
1310 impl somatize_core::tracking::Tracker for SpyTracker {
1311 fn run_id(&self) -> &str {
1312 "spy"
1313 }
1314 fn run_dir(&self) -> &std::path::Path {
1315 std::path::Path::new("/nonexistent")
1316 }
1317 fn sink(&self) -> Arc<dyn somatize_core::tracking::EventSink> {
1318 struct Null;
1319 impl somatize_core::tracking::EventSink for Null {
1320 fn record(&self, _event: &Event) {}
1321 }
1322 Arc::new(Null)
1323 }
1324 fn save_manifest(&self, _m: &somatize_core::tracking::RunManifest) -> Result<()> {
1325 Ok(())
1326 }
1327 fn save_artifact(&self, _p: &str, _b: &[u8]) -> Result<()> {
1328 Ok(())
1329 }
1330 fn save_study(&self, study: &Study) -> Result<()> {
1331 if self.fail {
1332 return Err(somatize_core::SomaError::Other("disk gone".into()));
1333 }
1334 self.saves.lock().unwrap().push(study.trials.len());
1335 Ok(())
1336 }
1337 fn heartbeat(&self) -> Result<()> {
1338 Ok(())
1339 }
1340 fn finalize(&self, _s: somatize_core::tracking::RunState) -> Result<()> {
1341 Ok(())
1342 }
1343 }
1344
1345 #[test]
1346 fn study_is_saved_after_every_trial_with_monotonic_growth() {
1347 let bus = Arc::new(EventBus::new(256));
1348 let tracker = Arc::new(SpyTracker::default());
1349 let runner = StudyRunner::new(bus).with_tracker(tracker.clone());
1350
1351 let mut study = Study::new(
1352 "persist",
1353 one_dim_space(),
1354 SearchStrategy::Random {
1355 n_trials: 3,
1356 seed: Some(3),
1357 },
1358 vec![Objective {
1359 metric: "f1".into(),
1360 direction: Direction::Maximize,
1361 }],
1362 );
1363 let mut sampler = RandomSampler::new(3, Some(3));
1364 runner
1365 .run(&mut study, &mut sampler, &make_simple_executor())
1366 .unwrap();
1367
1368 let saves = tracker.saves.lock().unwrap().clone();
1371 assert_eq!(saves, vec![1, 2, 3, 3]);
1372 }
1373
1374 #[test]
1375 fn failing_tracker_never_fails_the_study() {
1376 let bus = Arc::new(EventBus::new(256));
1377 let tracker = Arc::new(SpyTracker {
1378 fail: true,
1379 ..Default::default()
1380 });
1381 let runner = StudyRunner::new(bus).with_tracker(tracker);
1382
1383 let mut study = Study::new(
1384 "resilient",
1385 one_dim_space(),
1386 SearchStrategy::Random {
1387 n_trials: 3,
1388 seed: Some(3),
1389 },
1390 vec![Objective {
1391 metric: "f1".into(),
1392 direction: Direction::Maximize,
1393 }],
1394 );
1395 let mut sampler = RandomSampler::new(3, Some(3));
1396 runner
1397 .run(&mut study, &mut sampler, &make_simple_executor())
1398 .unwrap();
1399 assert_eq!(study.trials.len(), 3);
1400 }
1401
1402 struct OrderSpySampler {
1404 inner: GridSampler,
1405 log: Vec<String>,
1406 }
1407
1408 impl Sampler for OrderSpySampler {
1409 fn prepare(&mut self, space: &SearchSpace) {
1410 self.inner.prepare(space);
1411 }
1412 fn sample(
1413 &mut self,
1414 space: &SearchSpace,
1415 trial_index: usize,
1416 ) -> Result<Option<HashMap<String, serde_json::Value>>> {
1417 self.log.push(format!("sample:{trial_index}"));
1418 self.inner.sample(space, trial_index)
1419 }
1420 fn n_trials(&self) -> Option<usize> {
1421 self.inner.n_trials()
1422 }
1423 fn record_result(&mut self, _params: &HashMap<String, serde_json::Value>, value: f64) {
1424 self.log.push(format!("record:{value:.2}"));
1425 }
1426 }
1427
1428 #[test]
1429 fn resume_replays_history_into_the_sampler_before_sampling() {
1430 let executor = make_executor();
1431
1432 let bus = Arc::new(EventBus::new(256));
1434 let runner = StudyRunner::new(bus);
1435 let mut study = Study::new(
1436 "replay",
1437 sample_space(),
1438 SearchStrategy::Grid { points_per_dim: 3 },
1439 vec![Objective {
1440 metric: "f1".into(),
1441 direction: Direction::Maximize,
1442 }],
1443 );
1444 runner
1445 .run(&mut study, &mut GridSampler::new(3), &executor)
1446 .unwrap();
1447 study.trials.truncate(3);
1448
1449 let bus = Arc::new(EventBus::new(256));
1450 let runner = StudyRunner::new(bus);
1451 let mut spy = OrderSpySampler {
1452 inner: GridSampler::new(3),
1453 log: Vec::new(),
1454 };
1455 runner.run(&mut study, &mut spy, &executor).unwrap();
1456
1457 assert_eq!(
1460 spy.log.iter().filter(|e| e.starts_with("record:")).count(),
1461 3 + 3
1462 );
1463 assert!(
1464 spy.log[..3].iter().all(|e| e.starts_with("record:")),
1465 "history replay must precede sampling: {:?}",
1466 &spy.log[..4]
1467 );
1468 assert_eq!(spy.log[3], "sample:3", "sampling resumes at the next index");
1469 }
1470
1471 #[test]
1472 fn frozen_param_overrides_a_sampled_dimension() {
1473 let bus = Arc::new(EventBus::new(256));
1474 let runner = StudyRunner::new(bus);
1475
1476 let mut study = Study::new(
1478 "frozen-collision",
1479 one_dim_space(),
1480 SearchStrategy::Random {
1481 n_trials: 3,
1482 seed: Some(2),
1483 },
1484 vec![Objective {
1485 metric: "f1".into(),
1486 direction: Direction::Maximize,
1487 }],
1488 );
1489 study.frozen.insert("x".into(), serde_json::json!(0.75));
1490
1491 let executor = FnTrialExecutor(
1492 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1493 assert_eq!(params["x"], serde_json::json!(0.75));
1494 Ok(TrialOutcome::Completed(vec![MetricRecord {
1495 name: "f1".into(),
1496 value: 0.5,
1497 step: 0,
1498 timestamp: Utc::now(),
1499 }]))
1500 },
1501 );
1502 let mut sampler = RandomSampler::new(3, Some(2));
1503 runner.run(&mut study, &mut sampler, &executor).unwrap();
1504
1505 for t in &study.trials {
1506 assert_eq!(t.params["x"], serde_json::json!(0.75));
1507 }
1508 }
1509
1510 #[test]
1511 fn pruning_watches_first_composite_term_when_no_objectives() {
1512 use somatize_core::study::{CompositeObjective, Scalarizer};
1513 use std::sync::atomic::AtomicUsize;
1514
1515 let bus = Arc::new(EventBus::new(256));
1516 let runner = StudyRunner::new(bus);
1517 let mut study = Study::new(
1518 "composite-pruning",
1519 one_dim_space(),
1520 SearchStrategy::Random {
1521 n_trials: 4,
1522 seed: Some(7),
1523 },
1524 vec![], )
1526 .with_composite(CompositeObjective {
1527 terms: vec![("f1".into(), 1.0), ("aux".into(), 0.1)],
1528 direction: Direction::Maximize,
1529 scalarizer: Scalarizer::WeightedSum,
1530 })
1531 .with_pruning(PruningStrategy::Median { n_warmup_steps: 2 });
1532
1533 let counter = AtomicUsize::new(0);
1534 let executor = FnTrialExecutor(
1535 move |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1536 let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
1537 for step in 0..10 {
1538 let value = if good { 0.5 + step as f64 * 0.05 } else { 0.01 };
1539 if ctx.report("f1", value, step) {
1540 return Ok(TrialOutcome::Pruned {
1541 step,
1542 reason: "stopped".into(),
1543 });
1544 }
1545 }
1546 Ok(TrialOutcome::Completed(vec![MetricRecord {
1547 name: "aux".into(),
1548 value: 0.0,
1549 step: 0,
1550 timestamp: Utc::now(),
1551 }]))
1552 },
1553 );
1554 let mut sampler = RandomSampler::new(4, Some(7));
1555 runner.run(&mut study, &mut sampler, &executor).unwrap();
1556 assert_bad_trials_pruned(&study);
1557 }
1558
1559 #[test]
1560 fn planned_trials_is_stamped_and_progress_completes() {
1561 let bus = Arc::new(EventBus::new(256));
1562 let runner = StudyRunner::new(bus);
1563 let mut study = Study::new(
1564 "stamped",
1565 sample_space(),
1566 SearchStrategy::Grid { points_per_dim: 3 },
1567 vec![Objective {
1568 metric: "f1".into(),
1569 direction: Direction::Maximize,
1570 }],
1571 );
1572 assert_eq!(study.total_trials(), None);
1573 runner
1574 .run(&mut study, &mut GridSampler::new(3), &make_executor())
1575 .unwrap();
1576 assert_eq!(study.planned_trials, Some(6));
1577 assert_eq!(study.total_trials(), Some(6));
1578 assert!((study.progress() - 1.0).abs() < f64::EPSILON);
1579 }
1580
1581 #[test]
1582 fn all_failed_study_completes_with_nan_best() {
1583 let bus = Arc::new(EventBus::new(256));
1584 let mut rx = bus.subscribe();
1585 let runner = StudyRunner::new(bus);
1586 let mut study = Study::new(
1587 "doomed",
1588 one_dim_space(),
1589 SearchStrategy::Random {
1590 n_trials: 3,
1591 seed: Some(1),
1592 },
1593 vec![Objective {
1594 metric: "f1".into(),
1595 direction: Direction::Maximize,
1596 }],
1597 );
1598 let executor = FnTrialExecutor(
1599 |_p: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1600 Err(SomaError::Other("cuda out of memory".into()))
1601 },
1602 );
1603 let mut sampler = RandomSampler::new(3, Some(1));
1604 runner.run(&mut study, &mut sampler, &executor).unwrap();
1605
1606 for (i, t) in study.trials.iter().enumerate() {
1608 assert_eq!(t.id, format!("trial_{i:04}"));
1609 match &t.state {
1610 TrialState::Failed { error } => assert!(error.contains("cuda out of memory")),
1611 other => panic!("expected Failed, got {other:?}"),
1612 }
1613 }
1614 assert!(study.best_trial().is_none());
1615
1616 let mut failed_events = 0;
1617 let mut completed_nan = false;
1618 while let Ok(e) = rx.try_recv() {
1619 match e {
1620 Event::TrialFailed { error, .. } => {
1621 assert!(error.contains("cuda out of memory"));
1622 failed_events += 1;
1623 }
1624 Event::StudyCompleted {
1625 best_trial_id,
1626 best_value,
1627 ..
1628 } => {
1629 assert!(best_trial_id.is_empty());
1630 assert!(best_value.is_nan());
1631 completed_nan = true;
1632 }
1633 _ => {}
1634 }
1635 }
1636 assert_eq!(failed_events, 3);
1637 assert!(completed_nan);
1638 }
1639
1640 #[test]
1641 fn best_updated_fires_once_when_trials_worsen() {
1642 use std::sync::atomic::AtomicUsize;
1643 let bus = Arc::new(EventBus::new(256));
1644 let mut rx = bus.subscribe();
1645 let runner = StudyRunner::new(bus);
1646 let mut study = Study::new(
1647 "worsening",
1648 one_dim_space(),
1649 SearchStrategy::Random {
1650 n_trials: 3,
1651 seed: Some(1),
1652 },
1653 vec![Objective {
1654 metric: "f1".into(),
1655 direction: Direction::Maximize,
1656 }],
1657 );
1658 let counter = AtomicUsize::new(0);
1659 let executor = FnTrialExecutor(
1660 move |_p: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1661 let i = counter.fetch_add(1, Ordering::SeqCst);
1662 Ok(TrialOutcome::Completed(vec![MetricRecord {
1663 name: "f1".into(),
1664 value: 0.9 - i as f64 * 0.2, step: 0,
1666 timestamp: Utc::now(),
1667 }]))
1668 },
1669 );
1670 let mut sampler = RandomSampler::new(3, Some(1));
1671 runner.run(&mut study, &mut sampler, &executor).unwrap();
1672
1673 let mut best_events = Vec::new();
1674 while let Ok(e) = rx.try_recv() {
1675 if let Event::BestUpdated { value, .. } = e {
1676 best_events.push(value);
1677 }
1678 }
1679 assert_eq!(best_events, vec![0.9], "only the first trial is ever best");
1680 }
1681
1682 #[test]
1687 fn reported_and_final_metrics_are_concatenated_not_deduped() {
1688 let bus = Arc::new(EventBus::new(256));
1689 let mut rx = bus.subscribe();
1690 let runner = StudyRunner::new(bus);
1691 let mut study = Study::new(
1692 "dup",
1693 one_dim_space(),
1694 SearchStrategy::Random {
1695 n_trials: 1,
1696 seed: Some(1),
1697 },
1698 vec![Objective {
1699 metric: "f1".into(),
1700 direction: Direction::Maximize,
1701 }],
1702 );
1703 let executor = FnTrialExecutor(
1704 |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1705 ctx.report("f1", 0.4, 0);
1706 Ok(TrialOutcome::Completed(vec![MetricRecord {
1707 name: "f1".into(),
1708 value: 0.6,
1709 step: 1,
1710 timestamp: Utc::now(),
1711 }]))
1712 },
1713 );
1714 let mut sampler = RandomSampler::new(1, Some(1));
1715 runner.run(&mut study, &mut sampler, &executor).unwrap();
1716
1717 let f1_records: Vec<f64> = study.trials[0]
1718 .metrics
1719 .iter()
1720 .filter(|m| m.name == "f1")
1721 .map(|m| m.value)
1722 .collect();
1723 assert_eq!(f1_records, vec![0.4, 0.6]);
1724
1725 let mut metric_events = 0;
1726 while let Ok(e) = rx.try_recv() {
1727 if matches!(e, Event::TrialMetric { .. }) {
1728 metric_events += 1;
1729 }
1730 }
1731 assert_eq!(metric_events, 2);
1732 }
1733
1734 #[test]
1735 fn sampler_feedback_values_are_exact_and_skip_non_completed() {
1736 let bus = Arc::new(EventBus::new(256));
1737 let runner = StudyRunner::new(bus);
1738 let mut study = Study::new(
1739 "feedback-exact",
1740 one_dim_space(),
1741 SearchStrategy::Random {
1742 n_trials: 3,
1743 seed: Some(1),
1744 },
1745 vec![Objective {
1746 metric: "loss".into(),
1747 direction: Direction::Minimize,
1748 }],
1749 );
1750 use std::sync::atomic::AtomicUsize;
1751 let counter = AtomicUsize::new(0);
1752 let executor = FnTrialExecutor(
1753 move |_p: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| match counter
1754 .fetch_add(1, Ordering::SeqCst)
1755 {
1756 0 => Ok(TrialOutcome::Completed(vec![MetricRecord {
1757 name: "loss".into(),
1758 value: 0.25,
1759 step: 0,
1760 timestamp: Utc::now(),
1761 }])),
1762 1 => Ok(TrialOutcome::Pruned {
1763 step: 1,
1764 reason: "bad".into(),
1765 }),
1766 _ => Err(SomaError::Other("boom".into())),
1767 },
1768 );
1769 let mut spy = SpySampler {
1770 inner: RandomSampler::new(3, Some(1)),
1771 recorded: Vec::new(),
1772 };
1773 runner.run(&mut study, &mut spy, &executor).unwrap();
1774
1775 assert_eq!(spy.recorded, vec![-0.25]);
1777 }
1778
1779 #[test]
1780 fn timestamps_backfilled_and_monotonic() {
1781 let bus = Arc::new(EventBus::new(256));
1782 let tracker = Arc::new(SpyTracker::default());
1783 let runner = StudyRunner::new(bus).with_tracker(tracker);
1784 let mut study = Study::new(
1785 "ts",
1786 one_dim_space(),
1787 SearchStrategy::Random {
1788 n_trials: 2,
1789 seed: Some(1),
1790 },
1791 vec![Objective {
1792 metric: "f1".into(),
1793 direction: Direction::Maximize,
1794 }],
1795 );
1796 study.created_at = None; let mut sampler = RandomSampler::new(2, Some(1));
1798 runner
1799 .run(&mut study, &mut sampler, &make_simple_executor())
1800 .unwrap();
1801
1802 assert!(study.created_at.is_some(), "backfilled");
1803 assert!(study.updated_at.is_some());
1804 for t in &study.trials {
1805 assert!(t.started_at.is_some());
1806 assert!(t.finished_at.unwrap() >= t.started_at.unwrap());
1807 }
1808 }
1809
1810 #[test]
1811 fn bayesian_through_runner_improves_over_time() {
1812 let bus = Arc::new(EventBus::new(1024));
1816 let runner = StudyRunner::new(bus);
1817 let mut study = Study::new(
1818 "tpe",
1819 one_dim_space(),
1820 SearchStrategy::Bayesian {
1821 n_trials: 30,
1822 n_startup: 8,
1823 seed: Some(42),
1824 },
1825 vec![Objective {
1826 metric: "score".into(),
1827 direction: Direction::Maximize,
1828 }],
1829 );
1830 let executor = FnTrialExecutor(
1831 |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1832 let x = params["x"].as_f64().unwrap();
1833 Ok(TrialOutcome::Completed(vec![MetricRecord {
1834 name: "score".into(),
1835 value: 1.0 - (x - 0.7).abs(), step: 0,
1837 timestamp: Utc::now(),
1838 }]))
1839 },
1840 );
1841 let mut sampler = BayesianSampler::new(30, 8, Some(42));
1842 runner.run(&mut study, &mut sampler, &executor).unwrap();
1843
1844 let scores: Vec<f64> = study
1845 .trials
1846 .iter()
1847 .map(|t| t.last_metric("score").unwrap())
1848 .collect();
1849 let head: f64 = scores[..8].iter().sum::<f64>() / 8.0; let tail: f64 = scores[20..].iter().sum::<f64>() / 10.0;
1851 assert!(
1852 tail > head,
1853 "TPE with feedback must beat its random startup: head={head:.3} tail={tail:.3}"
1854 );
1855 }
1856
1857 #[test]
1858 fn study_progress_tracking() {
1859 let bus = Arc::new(EventBus::new(256));
1860 let mut rx = bus.subscribe();
1861 let runner = StudyRunner::new(bus);
1862
1863 let mut space = SearchSpace::new();
1864 space.add(SearchDimension::Float {
1865 name: "x".into(),
1866 low: 0.0,
1867 high: 1.0,
1868 scale: Scale::Linear,
1869 default: None,
1870 });
1871
1872 let mut study = Study::new(
1873 "progress_test",
1874 space,
1875 SearchStrategy::Random {
1876 n_trials: 3,
1877 seed: None,
1878 },
1879 vec![Objective {
1880 metric: "f1".into(),
1881 direction: Direction::Maximize,
1882 }],
1883 );
1884
1885 let executor = FnTrialExecutor(
1886 |_params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1887 Ok(TrialOutcome::Completed(vec![MetricRecord {
1888 name: "f1".into(),
1889 value: 0.5,
1890 step: 0,
1891 timestamp: Utc::now(),
1892 }]))
1893 },
1894 );
1895
1896 let mut sampler = RandomSampler::new(3, Some(42));
1897 runner.run(&mut study, &mut sampler, &executor).unwrap();
1898
1899 let mut progress_events = Vec::new();
1901 while let Ok(e) = rx.try_recv() {
1902 if let Event::StudyProgress {
1903 completed, total, ..
1904 } = e
1905 {
1906 progress_events.push((completed, total));
1907 }
1908 }
1909
1910 assert_eq!(progress_events.len(), 3);
1911 assert_eq!(progress_events[0], (1, 3));
1912 assert_eq!(progress_events[1], (2, 3));
1913 assert_eq!(progress_events[2], (3, 3));
1914 }
1915}