somatize_runtime/executors/
pbt.rs1use 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#[derive(Debug, Clone)]
22pub struct PbtConfig {
23 pub population_size: usize,
25 pub generations: usize,
27 pub exploit: ExploitStrategy,
29 pub explore: ExploreStrategy,
31 pub search_space: SearchSpace,
33 pub train_steps_per_generation: usize,
37}
38
39#[derive(Debug, Clone)]
41pub struct PopulationMember {
42 pub id: String,
44 pub params: HashMap<String, serde_json::Value>,
46 pub state: Value,
49 pub fitness: Option<f64>,
52}
53
54pub trait PbtExecutor: Send + Sync {
56 fn train(&self, member: &PopulationMember) -> Result<Value>;
58 fn evaluate(&self, member: &PopulationMember) -> Result<f64>;
60}
61
62pub struct FnPbtExecutor<T, E> {
64 pub train_fn: T,
66 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
83pub struct PbtRunner {
85 event_bus: Arc<EventBus>,
86}
87
88impl PbtRunner {
89 pub fn new(event_bus: Arc<EventBus>) -> Self {
91 Self { event_bus }
92 }
93
94 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 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 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 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 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 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 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 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 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
302fn 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
334fn 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 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 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 assert!(result.iter().all(|m| m.fitness.is_some()));
464 }
465}