Skip to main content

somatize_runtime/executors/
pbt.rs

1//! Population-Based Training runner.
2//!
3//! PBT is a cyclic evolutionary process where each generation:
4//! 1. **Train**: each population member trains for N steps
5//! 2. **Evaluate**: each member is evaluated to produce a fitness score
6//! 3. **Exploit/Explore**: underperformers copy top performers, then mutate hyperparameters
7//!
8//! Each generation's training phase uses the existing sampler infrastructure.
9
10use crate::event_bus::EventBus;
11use crate::sampler::{hash_u64, pseudo_random};
12use somatize_core::error::Result;
13use somatize_core::event::Event;
14use somatize_core::search::{SearchDimension, SearchSpace};
15use somatize_core::strategy::{ExploitStrategy, ExploreStrategy};
16use somatize_core::value::Value;
17use std::collections::HashMap;
18use std::sync::Arc;
19
20/// Configuration for a PBT run.
21#[derive(Debug, Clone)]
22pub struct PbtConfig {
23    /// Number of members evolved together.
24    pub population_size: usize,
25    /// Train → evaluate → exploit/explore cycles to run.
26    pub generations: usize,
27    /// How underperformers copy top performers (truncation, binary tournament).
28    pub exploit: ExploitStrategy,
29    /// How copied hyperparameters are mutated afterwards.
30    pub explore: ExploreStrategy,
31    /// Dimensions the initial population and `Resample` mutations draw from.
32    pub search_space: SearchSpace,
33    /// Advisory length of one generation's training phase. The runner calls
34    /// [`PbtExecutor::train`] once per member per generation; what a "step"
35    /// means is the executor's to interpret.
36    pub train_steps_per_generation: usize,
37}
38
39/// A single member of the population.
40#[derive(Debug, Clone)]
41pub struct PopulationMember {
42    /// Stable member identifier (`member_<i>`).
43    pub id: String,
44    /// Current hyperparameters — copied on exploit, mutated on explore.
45    pub params: HashMap<String, serde_json::Value>,
46    /// Trained state, updated after each generation's training phase and
47    /// copied along with `params` when a top performer is exploited.
48    pub state: Value,
49    /// Latest evaluation score, higher is better; `None` before the first
50    /// evaluation (a failed evaluation records `f64::NEG_INFINITY`).
51    pub fitness: Option<f64>,
52}
53
54/// Trait for the training + evaluation callback.
55pub trait PbtExecutor: Send + Sync {
56    /// Train a member for one generation. Returns updated state.
57    fn train(&self, member: &PopulationMember) -> Result<Value>;
58    /// Evaluate a member. Returns fitness score (higher = better).
59    fn evaluate(&self, member: &PopulationMember) -> Result<f64>;
60}
61
62/// Function-based PBT executor for convenience.
63pub struct FnPbtExecutor<T, E> {
64    /// Closure backing [`PbtExecutor::train`].
65    pub train_fn: T,
66    /// Closure backing [`PbtExecutor::evaluate`].
67    pub eval_fn: E,
68}
69
70impl<T, E> PbtExecutor for FnPbtExecutor<T, E>
71where
72    T: Fn(&PopulationMember) -> Result<Value> + Send + Sync,
73    E: Fn(&PopulationMember) -> Result<f64> + Send + Sync,
74{
75    fn train(&self, member: &PopulationMember) -> Result<Value> {
76        (self.train_fn)(member)
77    }
78    fn evaluate(&self, member: &PopulationMember) -> Result<f64> {
79        (self.eval_fn)(member)
80    }
81}
82
83/// Orchestrates the PBT evolutionary cycle.
84pub struct PbtRunner {
85    event_bus: Arc<EventBus>,
86}
87
88impl PbtRunner {
89    /// A runner emitting generation and exploit events on `event_bus`.
90    pub fn new(event_bus: Arc<EventBus>) -> Self {
91        Self { event_bus }
92    }
93
94    /// Run the full PBT evolutionary process.
95    ///
96    /// Returns the final population sorted by fitness (best first).
97    pub fn run(
98        &self,
99        config: &PbtConfig,
100        executor: &dyn PbtExecutor,
101    ) -> Result<Vec<PopulationMember>> {
102        let study_id = somatize_core::util::timestamp_id("pbt");
103        let mut rng_state: u64 = 42;
104
105        // Initialize population with random params
106        let mut population = self.initialize_population(config, &mut rng_state);
107
108        for generation in 0..config.generations {
109            self.event_bus.emit(Event::GenerationStarted {
110                study_id: study_id.clone(),
111                generation,
112                population_size: population.len(),
113            });
114
115            // Stage 1: Train
116            for member in &mut population {
117                match executor.train(member) {
118                    Ok(new_state) => member.state = new_state,
119                    Err(e) => {
120                        tracing::warn!("PBT train failed for {}: {e}", member.id);
121                    }
122                }
123            }
124
125            // Stage 2: Evaluate
126            //
127            // One member that cannot be evaluated is scored at negative
128            // infinity and exploited away — a flaky run should not cost
129            // the population. Every member failing is a different thing:
130            // the callback is wrong, and a generation where nothing could
131            // be scored has no signal to evolve on. Reporting success
132            // there hands back a population of -inf that looks ranked.
133            let mut first_failure: Option<String> = None;
134            let mut failures = 0usize;
135            for member in &mut population {
136                match executor.evaluate(member) {
137                    Ok(fitness) => member.fitness = Some(fitness),
138                    Err(e) => {
139                        tracing::warn!("PBT evaluate failed for {}: {e}", member.id);
140                        first_failure.get_or_insert_with(|| e.to_string());
141                        failures += 1;
142                        member.fitness = Some(f64::NEG_INFINITY);
143                    }
144                }
145            }
146            if failures == population.len() {
147                return Err(somatize_core::error::SomaError::Other(format!(
148                    "PBT generation {generation}: no member could be evaluated, \
149                     so there is no fitness to evolve on. The first failure was: \
150                     {}",
151                    first_failure.unwrap_or_else(|| "unreported".into())
152                )));
153            }
154
155            // Sort by fitness (descending)
156            population.sort_by(|a, b| {
157                b.fitness
158                    .unwrap_or(f64::NEG_INFINITY)
159                    .partial_cmp(&a.fitness.unwrap_or(f64::NEG_INFINITY))
160                    .unwrap_or(std::cmp::Ordering::Equal)
161            });
162
163            let best_fitness = population[0].fitness.unwrap_or(0.0);
164            let mean_fitness =
165                population.iter().filter_map(|m| m.fitness).sum::<f64>() / population.len() as f64;
166
167            // Stage 3: Exploit/Explore
168            self.evolve(
169                &mut population,
170                config,
171                generation,
172                &study_id,
173                &mut rng_state,
174            );
175
176            self.event_bus.emit(Event::GenerationCompleted {
177                study_id: study_id.clone(),
178                generation,
179                best_fitness,
180                mean_fitness,
181            });
182        }
183
184        // Final sort
185        population.sort_by(|a, b| {
186            b.fitness
187                .unwrap_or(f64::NEG_INFINITY)
188                .partial_cmp(&a.fitness.unwrap_or(f64::NEG_INFINITY))
189                .unwrap_or(std::cmp::Ordering::Equal)
190        });
191
192        Ok(population)
193    }
194
195    fn initialize_population(
196        &self,
197        config: &PbtConfig,
198        rng_state: &mut u64,
199    ) -> Vec<PopulationMember> {
200        let mut population = Vec::with_capacity(config.population_size);
201
202        for i in 0..config.population_size {
203            let params = sample_params(&config.search_space, rng_state);
204            population.push(PopulationMember {
205                id: format!("member_{i}"),
206                params,
207                state: Value::Empty,
208                fitness: None,
209            });
210        }
211
212        population
213    }
214
215    fn evolve(
216        &self,
217        population: &mut [PopulationMember],
218        config: &PbtConfig,
219        generation: usize,
220        study_id: &str,
221        rng_state: &mut u64,
222    ) {
223        let n = population.len();
224        if n < 2 {
225            return;
226        }
227
228        let cutoff = match &config.exploit {
229            ExploitStrategy::Truncation { fraction } => {
230                let c = ((n as f64) * fraction).ceil() as usize;
231                c.max(1).min(n / 2)
232            }
233            ExploitStrategy::Binary { .. } => n / 2,
234            _ => n / 2,
235        };
236
237        // Exploit: bottom performers copy from top
238        match &config.exploit {
239            ExploitStrategy::Truncation { .. } => {
240                for i in 0..cutoff {
241                    let bottom_idx = n - 1 - i;
242                    let top_idx = i;
243                    if bottom_idx <= top_idx {
244                        break;
245                    }
246
247                    let donor_id = population[top_idx].id.clone();
248                    let replaced_id = population[bottom_idx].id.clone();
249
250                    population[bottom_idx].params = population[top_idx].params.clone();
251                    population[bottom_idx].state = population[top_idx].state.clone();
252
253                    self.event_bus.emit(Event::MemberExploited {
254                        study_id: study_id.to_string(),
255                        generation,
256                        replaced_id,
257                        donor_id,
258                    });
259                }
260            }
261            ExploitStrategy::Binary { .. } => {
262                for i in cutoff..n {
263                    *rng_state = hash_u64(*rng_state, i as u64, generation as u64);
264                    let opponent = (*rng_state as usize) % cutoff;
265                    let my_fitness = population[i].fitness.unwrap_or(f64::NEG_INFINITY);
266                    let opp_fitness = population[opponent].fitness.unwrap_or(f64::NEG_INFINITY);
267                    if my_fitness < opp_fitness {
268                        let donor_id = population[opponent].id.clone();
269                        let replaced_id = population[i].id.clone();
270                        population[i].params = population[opponent].params.clone();
271                        population[i].state = population[opponent].state.clone();
272
273                        self.event_bus.emit(Event::MemberExploited {
274                            study_id: study_id.to_string(),
275                            generation,
276                            replaced_id,
277                            donor_id,
278                        });
279                    }
280                }
281            }
282            _ => {}
283        }
284
285        // Explore: mutate exploited members' hyperparameters
286        match &config.explore {
287            ExploreStrategy::Perturbation { factor } => {
288                for member in population[(n - cutoff)..].iter_mut() {
289                    perturb_params(&mut member.params, *factor, rng_state);
290                }
291            }
292            ExploreStrategy::Resample => {
293                for member in population[(n - cutoff)..].iter_mut() {
294                    member.params = sample_params(&config.search_space, rng_state);
295                }
296            }
297            _ => {}
298        }
299    }
300}
301
302/// Sample random parameters from a search space.
303fn sample_params(space: &SearchSpace, rng_state: &mut u64) -> HashMap<String, serde_json::Value> {
304    let mut params = HashMap::new();
305
306    for (dim_idx, dim) in space.dimensions.iter().enumerate() {
307        *rng_state = hash_u64(*rng_state, dim_idx as u64, 0);
308        let value = match dim {
309            SearchDimension::Float { low, high, .. } => {
310                let t = pseudo_random(*rng_state);
311                let v = low + t * (high - low);
312                serde_json::Value::from(v)
313            }
314            SearchDimension::Int { low, high, .. } => {
315                let t = pseudo_random(*rng_state);
316                let range = (*high - *low + 1) as f64;
317                let v = *low + (t * range) as i64;
318                serde_json::Value::from(v.min(*high))
319            }
320            SearchDimension::Categorical { choices, .. } => {
321                let t = pseudo_random(*rng_state);
322                let idx = (t * choices.len() as f64) as usize;
323                let idx = idx.min(choices.len() - 1);
324                choices[idx].clone()
325            }
326            _ => continue,
327        };
328        params.insert(dim.name().to_string(), value);
329    }
330
331    params
332}
333
334/// Perturb numeric parameters by a random factor in [1-factor, 1+factor].
335fn perturb_params(
336    params: &mut HashMap<String, serde_json::Value>,
337    factor: f64,
338    rng_state: &mut u64,
339) {
340    for (i, value) in params.values_mut().enumerate() {
341        if let Some(v) = value.as_f64() {
342            *rng_state = hash_u64(*rng_state, i as u64, 999);
343            let t = pseudo_random(*rng_state);
344            let perturbation = 1.0 + (t * 2.0 - 1.0) * factor;
345            *value = serde_json::Value::from(v * perturbation);
346        }
347    }
348}
349
350#[cfg(test)]
351mod tests {
352    use super::*;
353    use somatize_core::search::Scale;
354
355    fn test_config() -> PbtConfig {
356        let mut space = SearchSpace::new();
357        space.add(SearchDimension::Float {
358            name: "lr".into(),
359            low: 0.001,
360            high: 1.0,
361            scale: Scale::Log,
362            default: None,
363        });
364
365        PbtConfig {
366            population_size: 6,
367            generations: 3,
368            exploit: ExploitStrategy::Truncation { fraction: 0.33 },
369            explore: ExploreStrategy::Perturbation { factor: 0.2 },
370            search_space: space,
371            train_steps_per_generation: 10,
372        }
373    }
374
375    #[test]
376    fn pbt_basic_run() {
377        let bus = Arc::new(EventBus::new(256));
378        let runner = PbtRunner::new(bus);
379
380        let executor = FnPbtExecutor {
381            train_fn: |member: &PopulationMember| {
382                let lr = member
383                    .params
384                    .get("lr")
385                    .and_then(|v| v.as_f64())
386                    .unwrap_or(0.01);
387                Ok(Value::json(serde_json::json!({"lr": lr})))
388            },
389            eval_fn: |member: &PopulationMember| {
390                let lr = member
391                    .params
392                    .get("lr")
393                    .and_then(|v| v.as_f64())
394                    .unwrap_or(0.01);
395                Ok(-(lr - 0.1).abs())
396            },
397        };
398
399        let config = test_config();
400        let result = runner.run(&config, &executor).unwrap();
401
402        assert_eq!(result.len(), 6);
403        assert!(result.iter().all(|m| m.fitness.is_some()));
404        // Sorted by fitness descending
405        assert!(result[0].fitness.unwrap() >= result.last().unwrap().fitness.unwrap());
406    }
407
408    #[test]
409    fn pbt_emits_events() {
410        let bus = Arc::new(EventBus::new(256));
411        let mut rx = bus.subscribe();
412        let runner = PbtRunner::new(bus);
413
414        let executor = FnPbtExecutor {
415            train_fn: |_: &PopulationMember| Ok(Value::Empty),
416            eval_fn: |_: &PopulationMember| Ok(1.0),
417        };
418
419        let config = test_config();
420        runner.run(&config, &executor).unwrap();
421
422        let mut events = Vec::new();
423        while let Ok(e) = rx.try_recv() {
424            events.push(e);
425        }
426
427        let gen_started = events
428            .iter()
429            .filter(|e| matches!(e, Event::GenerationStarted { .. }))
430            .count();
431        let gen_completed = events
432            .iter()
433            .filter(|e| matches!(e, Event::GenerationCompleted { .. }))
434            .count();
435        assert_eq!(gen_started, 3);
436        assert_eq!(gen_completed, 3);
437    }
438
439    #[test]
440    fn pbt_population_evolves() {
441        let bus = Arc::new(EventBus::new(64));
442        let runner = PbtRunner::new(bus);
443
444        let executor = FnPbtExecutor {
445            train_fn: |_: &PopulationMember| Ok(Value::Empty),
446            eval_fn: |member: &PopulationMember| {
447                let lr = member
448                    .params
449                    .get("lr")
450                    .and_then(|v| v.as_f64())
451                    .unwrap_or(0.5);
452                // Fitness = -|lr - 0.1| (best at lr=0.1)
453                Ok(-(lr - 0.1).abs())
454            },
455        };
456
457        let mut config = test_config();
458        config.generations = 10;
459        let result = runner.run(&config, &executor).unwrap();
460
461        assert_eq!(result.len(), 6);
462        // All should have fitness
463        assert!(result.iter().all(|m| m.fitness.is_some()));
464    }
465}