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 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 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 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 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 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 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
283fn 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
315fn 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 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 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 assert!(result.iter().all(|m| m.fitness.is_some()));
445 }
446}