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            for member in &mut population {
127                match executor.evaluate(member) {
128                    Ok(fitness) => member.fitness = Some(fitness),
129                    Err(e) => {
130                        tracing::warn!("PBT evaluate failed for {}: {e}", member.id);
131                        member.fitness = Some(f64::NEG_INFINITY);
132                    }
133                }
134            }
135
136            // Sort by fitness (descending)
137            population.sort_by(|a, b| {
138                b.fitness
139                    .unwrap_or(f64::NEG_INFINITY)
140                    .partial_cmp(&a.fitness.unwrap_or(f64::NEG_INFINITY))
141                    .unwrap_or(std::cmp::Ordering::Equal)
142            });
143
144            let best_fitness = population[0].fitness.unwrap_or(0.0);
145            let mean_fitness =
146                population.iter().filter_map(|m| m.fitness).sum::<f64>() / population.len() as f64;
147
148            // Stage 3: Exploit/Explore
149            self.evolve(
150                &mut population,
151                config,
152                generation,
153                &study_id,
154                &mut rng_state,
155            );
156
157            self.event_bus.emit(Event::GenerationCompleted {
158                study_id: study_id.clone(),
159                generation,
160                best_fitness,
161                mean_fitness,
162            });
163        }
164
165        // Final sort
166        population.sort_by(|a, b| {
167            b.fitness
168                .unwrap_or(f64::NEG_INFINITY)
169                .partial_cmp(&a.fitness.unwrap_or(f64::NEG_INFINITY))
170                .unwrap_or(std::cmp::Ordering::Equal)
171        });
172
173        Ok(population)
174    }
175
176    fn initialize_population(
177        &self,
178        config: &PbtConfig,
179        rng_state: &mut u64,
180    ) -> Vec<PopulationMember> {
181        let mut population = Vec::with_capacity(config.population_size);
182
183        for i in 0..config.population_size {
184            let params = sample_params(&config.search_space, rng_state);
185            population.push(PopulationMember {
186                id: format!("member_{i}"),
187                params,
188                state: Value::Empty,
189                fitness: None,
190            });
191        }
192
193        population
194    }
195
196    fn evolve(
197        &self,
198        population: &mut [PopulationMember],
199        config: &PbtConfig,
200        generation: usize,
201        study_id: &str,
202        rng_state: &mut u64,
203    ) {
204        let n = population.len();
205        if n < 2 {
206            return;
207        }
208
209        let cutoff = match &config.exploit {
210            ExploitStrategy::Truncation { fraction } => {
211                let c = ((n as f64) * fraction).ceil() as usize;
212                c.max(1).min(n / 2)
213            }
214            ExploitStrategy::Binary { .. } => n / 2,
215            _ => n / 2,
216        };
217
218        // Exploit: bottom performers copy from top
219        match &config.exploit {
220            ExploitStrategy::Truncation { .. } => {
221                for i in 0..cutoff {
222                    let bottom_idx = n - 1 - i;
223                    let top_idx = i;
224                    if bottom_idx <= top_idx {
225                        break;
226                    }
227
228                    let donor_id = population[top_idx].id.clone();
229                    let replaced_id = population[bottom_idx].id.clone();
230
231                    population[bottom_idx].params = population[top_idx].params.clone();
232                    population[bottom_idx].state = population[top_idx].state.clone();
233
234                    self.event_bus.emit(Event::MemberExploited {
235                        study_id: study_id.to_string(),
236                        generation,
237                        replaced_id,
238                        donor_id,
239                    });
240                }
241            }
242            ExploitStrategy::Binary { .. } => {
243                for i in cutoff..n {
244                    *rng_state = hash_u64(*rng_state, i as u64, generation as u64);
245                    let opponent = (*rng_state as usize) % cutoff;
246                    let my_fitness = population[i].fitness.unwrap_or(f64::NEG_INFINITY);
247                    let opp_fitness = population[opponent].fitness.unwrap_or(f64::NEG_INFINITY);
248                    if my_fitness < opp_fitness {
249                        let donor_id = population[opponent].id.clone();
250                        let replaced_id = population[i].id.clone();
251                        population[i].params = population[opponent].params.clone();
252                        population[i].state = population[opponent].state.clone();
253
254                        self.event_bus.emit(Event::MemberExploited {
255                            study_id: study_id.to_string(),
256                            generation,
257                            replaced_id,
258                            donor_id,
259                        });
260                    }
261                }
262            }
263            _ => {}
264        }
265
266        // Explore: mutate exploited members' hyperparameters
267        match &config.explore {
268            ExploreStrategy::Perturbation { factor } => {
269                for member in population[(n - cutoff)..].iter_mut() {
270                    perturb_params(&mut member.params, *factor, rng_state);
271                }
272            }
273            ExploreStrategy::Resample => {
274                for member in population[(n - cutoff)..].iter_mut() {
275                    member.params = sample_params(&config.search_space, rng_state);
276                }
277            }
278            _ => {}
279        }
280    }
281}
282
283/// Sample random parameters from a search space.
284fn sample_params(space: &SearchSpace, rng_state: &mut u64) -> HashMap<String, serde_json::Value> {
285    let mut params = HashMap::new();
286
287    for (dim_idx, dim) in space.dimensions.iter().enumerate() {
288        *rng_state = hash_u64(*rng_state, dim_idx as u64, 0);
289        let value = match dim {
290            SearchDimension::Float { low, high, .. } => {
291                let t = pseudo_random(*rng_state);
292                let v = low + t * (high - low);
293                serde_json::Value::from(v)
294            }
295            SearchDimension::Int { low, high, .. } => {
296                let t = pseudo_random(*rng_state);
297                let range = (*high - *low + 1) as f64;
298                let v = *low + (t * range) as i64;
299                serde_json::Value::from(v.min(*high))
300            }
301            SearchDimension::Categorical { choices, .. } => {
302                let t = pseudo_random(*rng_state);
303                let idx = (t * choices.len() as f64) as usize;
304                let idx = idx.min(choices.len() - 1);
305                choices[idx].clone()
306            }
307            _ => continue,
308        };
309        params.insert(dim.name().to_string(), value);
310    }
311
312    params
313}
314
315/// Perturb numeric parameters by a random factor in [1-factor, 1+factor].
316fn perturb_params(
317    params: &mut HashMap<String, serde_json::Value>,
318    factor: f64,
319    rng_state: &mut u64,
320) {
321    for (i, value) in params.values_mut().enumerate() {
322        if let Some(v) = value.as_f64() {
323            *rng_state = hash_u64(*rng_state, i as u64, 999);
324            let t = pseudo_random(*rng_state);
325            let perturbation = 1.0 + (t * 2.0 - 1.0) * factor;
326            *value = serde_json::Value::from(v * perturbation);
327        }
328    }
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334    use somatize_core::search::Scale;
335
336    fn test_config() -> PbtConfig {
337        let mut space = SearchSpace::new();
338        space.add(SearchDimension::Float {
339            name: "lr".into(),
340            low: 0.001,
341            high: 1.0,
342            scale: Scale::Log,
343            default: None,
344        });
345
346        PbtConfig {
347            population_size: 6,
348            generations: 3,
349            exploit: ExploitStrategy::Truncation { fraction: 0.33 },
350            explore: ExploreStrategy::Perturbation { factor: 0.2 },
351            search_space: space,
352            train_steps_per_generation: 10,
353        }
354    }
355
356    #[test]
357    fn pbt_basic_run() {
358        let bus = Arc::new(EventBus::new(256));
359        let runner = PbtRunner::new(bus);
360
361        let executor = FnPbtExecutor {
362            train_fn: |member: &PopulationMember| {
363                let lr = member
364                    .params
365                    .get("lr")
366                    .and_then(|v| v.as_f64())
367                    .unwrap_or(0.01);
368                Ok(Value::json(serde_json::json!({"lr": lr})))
369            },
370            eval_fn: |member: &PopulationMember| {
371                let lr = member
372                    .params
373                    .get("lr")
374                    .and_then(|v| v.as_f64())
375                    .unwrap_or(0.01);
376                Ok(-(lr - 0.1).abs())
377            },
378        };
379
380        let config = test_config();
381        let result = runner.run(&config, &executor).unwrap();
382
383        assert_eq!(result.len(), 6);
384        assert!(result.iter().all(|m| m.fitness.is_some()));
385        // Sorted by fitness descending
386        assert!(result[0].fitness.unwrap() >= result.last().unwrap().fitness.unwrap());
387    }
388
389    #[test]
390    fn pbt_emits_events() {
391        let bus = Arc::new(EventBus::new(256));
392        let mut rx = bus.subscribe();
393        let runner = PbtRunner::new(bus);
394
395        let executor = FnPbtExecutor {
396            train_fn: |_: &PopulationMember| Ok(Value::Empty),
397            eval_fn: |_: &PopulationMember| Ok(1.0),
398        };
399
400        let config = test_config();
401        runner.run(&config, &executor).unwrap();
402
403        let mut events = Vec::new();
404        while let Ok(e) = rx.try_recv() {
405            events.push(e);
406        }
407
408        let gen_started = events
409            .iter()
410            .filter(|e| matches!(e, Event::GenerationStarted { .. }))
411            .count();
412        let gen_completed = events
413            .iter()
414            .filter(|e| matches!(e, Event::GenerationCompleted { .. }))
415            .count();
416        assert_eq!(gen_started, 3);
417        assert_eq!(gen_completed, 3);
418    }
419
420    #[test]
421    fn pbt_population_evolves() {
422        let bus = Arc::new(EventBus::new(64));
423        let runner = PbtRunner::new(bus);
424
425        let executor = FnPbtExecutor {
426            train_fn: |_: &PopulationMember| Ok(Value::Empty),
427            eval_fn: |member: &PopulationMember| {
428                let lr = member
429                    .params
430                    .get("lr")
431                    .and_then(|v| v.as_f64())
432                    .unwrap_or(0.5);
433                // Fitness = -|lr - 0.1| (best at lr=0.1)
434                Ok(-(lr - 0.1).abs())
435            },
436        };
437
438        let mut config = test_config();
439        config.generations = 10;
440        let result = runner.run(&config, &executor).unwrap();
441
442        assert_eq!(result.len(), 6);
443        // All should have fitness
444        assert!(result.iter().all(|m| m.fitness.is_some()));
445    }
446}