1use rand::prelude::*;
28#[cfg(feature = "parallel")]
29use rayon::prelude::*;
30use std::sync::atomic::{AtomicBool, Ordering};
31use std::sync::Arc;
32use std::time::Duration;
33
34use crate::error::{check_at_least, check_range, Result};
35use crate::timing::{evaluate_within, expired, Timer};
36
37#[cfg(feature = "serde")]
38use serde::{Deserialize, Serialize};
39
40#[derive(Debug, Clone)]
42#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
43pub struct BrkgaConfig {
44 pub population_size: usize,
46 pub max_generations: u32,
48 pub elite_fraction: f64,
50 pub mutant_fraction: f64,
52 pub elite_bias: f64,
54 pub time_limit: Option<Duration>,
56 pub target_fitness: Option<f64>,
58 pub stagnation_limit: Option<u32>,
60}
61
62impl Default for BrkgaConfig {
63 fn default() -> Self {
64 Self {
65 population_size: 100,
66 max_generations: 500,
67 elite_fraction: 0.2, mutant_fraction: 0.15, elite_bias: 0.7, time_limit: None,
71 target_fitness: None,
72 stagnation_limit: Some(50),
73 }
74 }
75}
76
77impl BrkgaConfig {
78 pub fn new() -> Self {
80 Self::default()
81 }
82
83 pub fn with_population_size(mut self, size: usize) -> Self {
85 self.population_size = size.max(4);
86 self
87 }
88
89 pub fn with_max_generations(mut self, gen: u32) -> Self {
91 self.max_generations = gen;
92 self
93 }
94
95 pub fn with_elite_fraction(mut self, fraction: f64) -> Self {
97 self.elite_fraction = fraction;
98 self
99 }
100
101 pub fn with_mutant_fraction(mut self, fraction: f64) -> Self {
103 self.mutant_fraction = fraction;
104 self
105 }
106
107 pub fn with_elite_bias(mut self, bias: f64) -> Self {
109 self.elite_bias = bias;
110 self
111 }
112
113 pub fn validate(&self) -> Result<()> {
121 check_range("elite_fraction", self.elite_fraction, 0.01, 0.5)?;
122 check_range("mutant_fraction", self.mutant_fraction, 0.0, 0.5)?;
123 check_range("elite_bias", self.elite_bias, 0.5, 1.0)?;
124 check_at_least("population_size", self.population_size as f64, 2.0)?;
125 Ok(())
126 }
127
128 pub fn with_time_limit(mut self, duration: Duration) -> Self {
130 self.time_limit = Some(duration);
131 self
132 }
133
134 pub fn with_target_fitness(mut self, fitness: f64) -> Self {
136 self.target_fitness = Some(fitness);
137 self
138 }
139
140 pub fn with_stagnation_limit(mut self, limit: u32) -> Self {
142 self.stagnation_limit = Some(limit);
143 self
144 }
145
146 pub fn elite_count(&self) -> usize {
148 ((self.population_size as f64) * self.elite_fraction).ceil() as usize
149 }
150
151 pub fn mutant_count(&self) -> usize {
153 ((self.population_size as f64) * self.mutant_fraction).ceil() as usize
154 }
155}
156
157#[derive(Debug, Clone)]
162pub struct RandomKeyChromosome {
163 pub keys: Vec<f64>,
165 fitness: f64,
167}
168
169impl RandomKeyChromosome {
170 pub fn new(num_keys: usize) -> Self {
172 Self {
173 keys: vec![0.0; num_keys],
174 fitness: f64::NEG_INFINITY,
175 }
176 }
177
178 pub fn random<R: Rng>(num_keys: usize, rng: &mut R) -> Self {
180 let keys: Vec<f64> = (0..num_keys).map(|_| rng.random::<f64>()).collect();
181 Self {
182 keys,
183 fitness: f64::NEG_INFINITY,
184 }
185 }
186
187 pub fn len(&self) -> usize {
189 self.keys.len()
190 }
191
192 pub fn is_empty(&self) -> bool {
194 self.keys.is_empty()
195 }
196
197 pub fn fitness(&self) -> f64 {
199 self.fitness
200 }
201
202 pub fn set_fitness(&mut self, fitness: f64) {
204 self.fitness = fitness;
205 }
206
207 pub fn biased_crossover<R: Rng>(&self, other: &Self, elite_bias: f64, rng: &mut R) -> Self {
212 let keys: Vec<f64> = self
213 .keys
214 .iter()
215 .zip(&other.keys)
216 .map(|(&elite_key, &non_elite_key)| {
217 if rng.random::<f64>() < elite_bias {
218 elite_key
219 } else {
220 non_elite_key
221 }
222 })
223 .collect();
224
225 Self {
226 keys,
227 fitness: f64::NEG_INFINITY,
228 }
229 }
230
231 pub fn decode_as_permutation(&self) -> Vec<usize> {
236 let mut indices: Vec<usize> = (0..self.keys.len()).collect();
237 indices.sort_by(|&a, &b| {
238 self.keys[a]
239 .partial_cmp(&self.keys[b])
240 .unwrap_or(std::cmp::Ordering::Equal)
241 });
242 indices
243 }
244
245 pub fn decode_as_discrete(&self, key_idx: usize, num_options: usize) -> usize {
250 if key_idx >= self.keys.len() || num_options == 0 {
251 return 0;
252 }
253 let key = self.keys[key_idx].clamp(0.0, 0.9999999);
254 (key * num_options as f64) as usize
255 }
256}
257
258pub trait BrkgaProblem: Send + Sync {
260 fn num_keys(&self) -> usize;
262
263 fn evaluate(&self, chromosome: &mut RandomKeyChromosome);
265
266 fn evaluate_parallel(&self, chromosomes: &mut [RandomKeyChromosome]) {
269 #[cfg(feature = "parallel")]
270 chromosomes.par_iter_mut().for_each(|c| {
271 self.evaluate(c);
272 });
273 #[cfg(not(feature = "parallel"))]
274 for c in chromosomes.iter_mut() {
275 self.evaluate(c);
276 }
277 }
278
279 fn initial_population<R: Rng>(&self, size: usize, rng: &mut R) -> Vec<RandomKeyChromosome> {
288 (0..size)
289 .map(|_| RandomKeyChromosome::random(self.num_keys(), rng))
290 .collect()
291 }
292
293 fn on_generation(
295 &self,
296 _generation: u32,
297 _best: &RandomKeyChromosome,
298 _population: &[RandomKeyChromosome],
299 ) {
300 }
302}
303
304#[derive(Debug, Clone)]
306pub struct BrkgaResult {
307 pub best: RandomKeyChromosome,
309 pub generations: u32,
311 pub elapsed: Duration,
313 pub target_reached: bool,
315 pub history: Vec<f64>,
317}
318
319#[derive(Debug, Clone)]
321pub struct BrkgaProgress {
322 pub generation: u32,
324 pub max_generations: u32,
326 pub best_fitness: f64,
328 pub avg_fitness: f64,
330 pub elapsed: Duration,
332 pub running: bool,
334}
335
336pub struct BrkgaRunner<P: BrkgaProblem> {
338 config: BrkgaConfig,
339 problem: P,
340 cancelled: Arc<AtomicBool>,
341}
342
343impl<P: BrkgaProblem> BrkgaRunner<P> {
344 pub fn new(config: BrkgaConfig, problem: P) -> Result<Self> {
349 Self::with_cancellation(config, problem, Arc::new(AtomicBool::new(false)))
350 }
351
352 pub fn with_cancellation(
357 config: BrkgaConfig,
358 problem: P,
359 cancelled: Arc<AtomicBool>,
360 ) -> Result<Self> {
361 config.validate()?;
362 Ok(Self {
363 config,
364 problem,
365 cancelled,
366 })
367 }
368
369 pub fn cancel_handle(&self) -> Arc<AtomicBool> {
371 self.cancelled.clone()
372 }
373
374 pub fn run(&self) -> BrkgaResult {
376 self.run_with_rng(&mut rand::rng())
377 }
378
379 pub fn run_with_progress<F>(&self, progress_callback: F) -> BrkgaResult
381 where
382 F: Fn(BrkgaProgress),
383 {
384 self.run_with_rng_and_progress(&mut rand::rng(), Some(progress_callback))
385 }
386
387 pub fn run_with_rng<R: Rng>(&self, rng: &mut R) -> BrkgaResult {
389 self.run_with_rng_and_progress::<R, fn(BrkgaProgress)>(rng, None)
390 }
391
392 pub fn run_with_rng_and_progress<R: Rng, F>(
394 &self,
395 rng: &mut R,
396 progress_callback: Option<F>,
397 ) -> BrkgaResult
398 where
399 F: Fn(BrkgaProgress),
400 {
401 let start = Timer::now();
402 let mut history = Vec::new();
403 let num_keys = self.problem.num_keys();
404
405 let mut population = self
407 .problem
408 .initial_population(self.config.population_size, rng);
409 population.truncate(self.config.population_size);
412 while population.len() < self.config.population_size {
413 population.push(RandomKeyChromosome::random(num_keys, rng));
414 }
415
416 evaluate_within(&mut population, &start, self.config.time_limit, |batch| {
418 self.problem.evaluate_parallel(batch)
419 });
420
421 population.sort_by(|a, b| {
423 b.fitness()
424 .partial_cmp(&a.fitness())
425 .unwrap_or(std::cmp::Ordering::Equal)
426 });
427
428 let mut best = population[0].clone();
429 let mut best_fitness = best.fitness();
430 let mut stagnation_count = 0u32;
431 let mut generation = 0u32;
432 let mut target_reached = false;
433
434 let elite_count = self.config.elite_count();
435 let mutant_count = self.config.mutant_count();
436
437 while generation < self.config.max_generations {
438 if self.cancelled.load(Ordering::Relaxed) {
440 break;
441 }
442
443 if expired(&start, self.config.time_limit) {
445 break;
446 }
447
448 if let Some(target) = self.config.target_fitness {
450 if best_fitness >= target {
451 target_reached = true;
452 break;
453 }
454 }
455
456 history.push(best_fitness);
458
459 let mut new_population = Vec::with_capacity(self.config.population_size);
461
462 for elite in population.iter().take(elite_count) {
464 new_population.push(elite.clone());
465 }
466
467 let mut mutants: Vec<RandomKeyChromosome> = (0..mutant_count)
469 .map(|_| RandomKeyChromosome::random(num_keys, rng))
470 .collect();
471
472 let crossover_count = self.config.population_size - elite_count - mutant_count;
474 let mut children: Vec<RandomKeyChromosome> = (0..crossover_count)
475 .map(|_| {
476 let elite_idx = rng.random_range(0..elite_count);
478 let elite_parent = &population[elite_idx];
479
480 let non_elite_idx = rng.random_range(elite_count..population.len());
482 let non_elite_parent = &population[non_elite_idx];
483
484 elite_parent.biased_crossover(non_elite_parent, self.config.elite_bias, rng)
486 })
487 .collect();
488
489 evaluate_within(&mut mutants, &start, self.config.time_limit, |batch| {
492 self.problem.evaluate_parallel(batch)
493 });
494 if expired(&start, self.config.time_limit) {
495 children.clear();
496 }
497 evaluate_within(&mut children, &start, self.config.time_limit, |batch| {
498 self.problem.evaluate_parallel(batch)
499 });
500
501 new_population.extend(mutants);
503 new_population.extend(children);
504
505 new_population.sort_by(|a, b| {
507 b.fitness()
508 .partial_cmp(&a.fitness())
509 .unwrap_or(std::cmp::Ordering::Equal)
510 });
511
512 let new_best_fitness = new_population[0].fitness();
514 if new_best_fitness > best_fitness {
515 best = new_population[0].clone();
516 best_fitness = new_best_fitness;
517 stagnation_count = 0;
518 } else {
519 stagnation_count += 1;
520 }
521
522 if let Some(limit) = self.config.stagnation_limit {
524 if stagnation_count >= limit {
525 break;
526 }
527 }
528
529 self.problem
531 .on_generation(generation, &best, &new_population);
532
533 if let Some(ref callback) = progress_callback {
535 let avg_fitness = new_population.iter().map(|c| c.fitness()).sum::<f64>()
536 / new_population.len() as f64;
537
538 callback(BrkgaProgress {
539 generation,
540 max_generations: self.config.max_generations,
541 best_fitness,
542 avg_fitness,
543 elapsed: start.elapsed(),
544 running: true,
545 });
546 }
547
548 population = new_population;
549 generation += 1;
550 }
551
552 history.push(best_fitness);
554
555 if let Some(ref callback) = progress_callback {
557 let avg_fitness = population.iter().map(|c| c.fitness()).sum::<f64>()
558 / population.len().max(1) as f64;
559
560 callback(BrkgaProgress {
561 generation,
562 max_generations: self.config.max_generations,
563 best_fitness,
564 avg_fitness,
565 elapsed: start.elapsed(),
566 running: false,
567 });
568 }
569
570 BrkgaResult {
571 best,
572 generations: generation,
573 elapsed: start.elapsed(),
574 target_reached,
575 history,
576 }
577 }
578}
579
580#[cfg(test)]
581mod tests {
582 use super::*;
583
584 struct MaxSumProblem {
586 num_keys: usize,
587 }
588
589 impl BrkgaProblem for MaxSumProblem {
590 fn num_keys(&self) -> usize {
591 self.num_keys
592 }
593
594 fn evaluate(&self, chromosome: &mut RandomKeyChromosome) {
595 let sum: f64 = chromosome.keys.iter().sum();
596 chromosome.set_fitness(sum);
597 }
598 }
599
600 #[test]
601 fn test_brkga_basic() {
602 let config = BrkgaConfig::default()
603 .with_population_size(50)
604 .with_max_generations(50);
605
606 let problem = MaxSumProblem { num_keys: 10 };
607 let runner = BrkgaRunner::new(config, problem).expect("valid config");
608 let result = runner.run();
609
610 assert!(result.best.fitness() > 5.0);
612 }
613
614 #[test]
615 fn test_random_key_chromosome() {
616 let mut rng = rand::rng();
617 let chromosome = RandomKeyChromosome::random(10, &mut rng);
618
619 assert_eq!(chromosome.len(), 10);
620 for &key in &chromosome.keys {
621 assert!((0.0..1.0).contains(&key));
622 }
623 }
624
625 #[test]
626 fn test_biased_crossover() {
627 let mut rng = rand::rng();
628 let elite = RandomKeyChromosome::random(10, &mut rng);
629 let non_elite = RandomKeyChromosome::random(10, &mut rng);
630
631 let child = elite.biased_crossover(&non_elite, 0.7, &mut rng);
632
633 assert_eq!(child.len(), 10);
634 for &key in &child.keys {
635 assert!((0.0..1.0).contains(&key));
636 }
637 }
638
639 #[test]
640 fn test_decode_as_permutation() {
641 let mut chromosome = RandomKeyChromosome::new(5);
642 chromosome.keys = vec![0.3, 0.1, 0.9, 0.5, 0.2];
643
644 let perm = chromosome.decode_as_permutation();
645
646 assert_eq!(perm, vec![1, 4, 0, 3, 2]);
648 }
649
650 #[test]
651 fn test_decode_as_discrete() {
652 let mut chromosome = RandomKeyChromosome::new(3);
653 chromosome.keys = vec![0.0, 0.5, 0.99];
654
655 assert_eq!(chromosome.decode_as_discrete(0, 6), 0);
657 assert_eq!(chromosome.decode_as_discrete(1, 6), 3);
658 assert_eq!(chromosome.decode_as_discrete(2, 6), 5);
659 }
660}