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::timing::{evaluate_within, expired, Timer};
35
36#[cfg(feature = "serde")]
37use serde::{Deserialize, Serialize};
38
39#[derive(Debug, Clone)]
41#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
42pub struct BrkgaConfig {
43 pub population_size: usize,
45 pub max_generations: u32,
47 pub elite_fraction: f64,
49 pub mutant_fraction: f64,
51 pub elite_bias: f64,
53 pub time_limit: Option<Duration>,
55 pub target_fitness: Option<f64>,
57 pub stagnation_limit: Option<u32>,
59}
60
61impl Default for BrkgaConfig {
62 fn default() -> Self {
63 Self {
64 population_size: 100,
65 max_generations: 500,
66 elite_fraction: 0.2, mutant_fraction: 0.15, elite_bias: 0.7, time_limit: None,
70 target_fitness: None,
71 stagnation_limit: Some(50),
72 }
73 }
74}
75
76impl BrkgaConfig {
77 pub fn new() -> Self {
79 Self::default()
80 }
81
82 pub fn with_population_size(mut self, size: usize) -> Self {
84 self.population_size = size.max(4);
85 self
86 }
87
88 pub fn with_max_generations(mut self, gen: u32) -> Self {
90 self.max_generations = gen;
91 self
92 }
93
94 pub fn with_elite_fraction(mut self, fraction: f64) -> Self {
96 self.elite_fraction = fraction.clamp(0.01, 0.5);
97 self
98 }
99
100 pub fn with_mutant_fraction(mut self, fraction: f64) -> Self {
102 self.mutant_fraction = fraction.clamp(0.0, 0.5);
103 self
104 }
105
106 pub fn with_elite_bias(mut self, bias: f64) -> Self {
108 self.elite_bias = bias.clamp(0.5, 1.0);
109 self
110 }
111
112 pub fn with_time_limit(mut self, duration: Duration) -> Self {
114 self.time_limit = Some(duration);
115 self
116 }
117
118 pub fn with_target_fitness(mut self, fitness: f64) -> Self {
120 self.target_fitness = Some(fitness);
121 self
122 }
123
124 pub fn with_stagnation_limit(mut self, limit: u32) -> Self {
126 self.stagnation_limit = Some(limit);
127 self
128 }
129
130 pub fn elite_count(&self) -> usize {
132 ((self.population_size as f64) * self.elite_fraction).ceil() as usize
133 }
134
135 pub fn mutant_count(&self) -> usize {
137 ((self.population_size as f64) * self.mutant_fraction).ceil() as usize
138 }
139}
140
141#[derive(Debug, Clone)]
146pub struct RandomKeyChromosome {
147 pub keys: Vec<f64>,
149 fitness: f64,
151}
152
153impl RandomKeyChromosome {
154 pub fn new(num_keys: usize) -> Self {
156 Self {
157 keys: vec![0.0; num_keys],
158 fitness: f64::NEG_INFINITY,
159 }
160 }
161
162 pub fn random<R: Rng>(num_keys: usize, rng: &mut R) -> Self {
164 let keys: Vec<f64> = (0..num_keys).map(|_| rng.random::<f64>()).collect();
165 Self {
166 keys,
167 fitness: f64::NEG_INFINITY,
168 }
169 }
170
171 pub fn len(&self) -> usize {
173 self.keys.len()
174 }
175
176 pub fn is_empty(&self) -> bool {
178 self.keys.is_empty()
179 }
180
181 pub fn fitness(&self) -> f64 {
183 self.fitness
184 }
185
186 pub fn set_fitness(&mut self, fitness: f64) {
188 self.fitness = fitness;
189 }
190
191 pub fn biased_crossover<R: Rng>(&self, other: &Self, elite_bias: f64, rng: &mut R) -> Self {
196 let keys: Vec<f64> = self
197 .keys
198 .iter()
199 .zip(&other.keys)
200 .map(|(&elite_key, &non_elite_key)| {
201 if rng.random::<f64>() < elite_bias {
202 elite_key
203 } else {
204 non_elite_key
205 }
206 })
207 .collect();
208
209 Self {
210 keys,
211 fitness: f64::NEG_INFINITY,
212 }
213 }
214
215 pub fn decode_as_permutation(&self) -> Vec<usize> {
220 let mut indices: Vec<usize> = (0..self.keys.len()).collect();
221 indices.sort_by(|&a, &b| {
222 self.keys[a]
223 .partial_cmp(&self.keys[b])
224 .unwrap_or(std::cmp::Ordering::Equal)
225 });
226 indices
227 }
228
229 pub fn decode_as_discrete(&self, key_idx: usize, num_options: usize) -> usize {
234 if key_idx >= self.keys.len() || num_options == 0 {
235 return 0;
236 }
237 let key = self.keys[key_idx].clamp(0.0, 0.9999999);
238 (key * num_options as f64) as usize
239 }
240}
241
242pub trait BrkgaProblem: Send + Sync {
244 fn num_keys(&self) -> usize;
246
247 fn evaluate(&self, chromosome: &mut RandomKeyChromosome);
249
250 fn evaluate_parallel(&self, chromosomes: &mut [RandomKeyChromosome]) {
253 #[cfg(feature = "parallel")]
254 chromosomes.par_iter_mut().for_each(|c| {
255 self.evaluate(c);
256 });
257 #[cfg(not(feature = "parallel"))]
258 for c in chromosomes.iter_mut() {
259 self.evaluate(c);
260 }
261 }
262
263 fn initial_population<R: Rng>(&self, size: usize, rng: &mut R) -> Vec<RandomKeyChromosome> {
272 (0..size)
273 .map(|_| RandomKeyChromosome::random(self.num_keys(), rng))
274 .collect()
275 }
276
277 fn on_generation(
279 &self,
280 _generation: u32,
281 _best: &RandomKeyChromosome,
282 _population: &[RandomKeyChromosome],
283 ) {
284 }
286}
287
288#[derive(Debug, Clone)]
290pub struct BrkgaResult {
291 pub best: RandomKeyChromosome,
293 pub generations: u32,
295 pub elapsed: Duration,
297 pub target_reached: bool,
299 pub history: Vec<f64>,
301}
302
303#[derive(Debug, Clone)]
305pub struct BrkgaProgress {
306 pub generation: u32,
308 pub max_generations: u32,
310 pub best_fitness: f64,
312 pub avg_fitness: f64,
314 pub elapsed: Duration,
316 pub running: bool,
318}
319
320pub struct BrkgaRunner<P: BrkgaProblem> {
322 config: BrkgaConfig,
323 problem: P,
324 cancelled: Arc<AtomicBool>,
325}
326
327impl<P: BrkgaProblem> BrkgaRunner<P> {
328 pub fn new(config: BrkgaConfig, problem: P) -> Self {
330 Self {
331 config,
332 problem,
333 cancelled: Arc::new(AtomicBool::new(false)),
334 }
335 }
336
337 pub fn with_cancellation(config: BrkgaConfig, problem: P, cancelled: Arc<AtomicBool>) -> Self {
339 Self {
340 config,
341 problem,
342 cancelled,
343 }
344 }
345
346 pub fn cancel_handle(&self) -> Arc<AtomicBool> {
348 self.cancelled.clone()
349 }
350
351 pub fn run(&self) -> BrkgaResult {
353 self.run_with_rng(&mut rand::rng())
354 }
355
356 pub fn run_with_progress<F>(&self, progress_callback: F) -> BrkgaResult
358 where
359 F: Fn(BrkgaProgress),
360 {
361 self.run_with_rng_and_progress(&mut rand::rng(), Some(progress_callback))
362 }
363
364 pub fn run_with_rng<R: Rng>(&self, rng: &mut R) -> BrkgaResult {
366 self.run_with_rng_and_progress::<R, fn(BrkgaProgress)>(rng, None)
367 }
368
369 pub fn run_with_rng_and_progress<R: Rng, F>(
371 &self,
372 rng: &mut R,
373 progress_callback: Option<F>,
374 ) -> BrkgaResult
375 where
376 F: Fn(BrkgaProgress),
377 {
378 let start = Timer::now();
379 let mut history = Vec::new();
380 let num_keys = self.problem.num_keys();
381
382 let mut population = self
384 .problem
385 .initial_population(self.config.population_size, rng);
386 population.truncate(self.config.population_size);
389 while population.len() < self.config.population_size {
390 population.push(RandomKeyChromosome::random(num_keys, rng));
391 }
392
393 evaluate_within(&mut population, &start, self.config.time_limit, |batch| {
395 self.problem.evaluate_parallel(batch)
396 });
397
398 population.sort_by(|a, b| {
400 b.fitness()
401 .partial_cmp(&a.fitness())
402 .unwrap_or(std::cmp::Ordering::Equal)
403 });
404
405 let mut best = population[0].clone();
406 let mut best_fitness = best.fitness();
407 let mut stagnation_count = 0u32;
408 let mut generation = 0u32;
409 let mut target_reached = false;
410
411 let elite_count = self.config.elite_count();
412 let mutant_count = self.config.mutant_count();
413
414 while generation < self.config.max_generations {
415 if self.cancelled.load(Ordering::Relaxed) {
417 break;
418 }
419
420 if expired(&start, self.config.time_limit) {
422 break;
423 }
424
425 if let Some(target) = self.config.target_fitness {
427 if best_fitness >= target {
428 target_reached = true;
429 break;
430 }
431 }
432
433 history.push(best_fitness);
435
436 let mut new_population = Vec::with_capacity(self.config.population_size);
438
439 for elite in population.iter().take(elite_count) {
441 new_population.push(elite.clone());
442 }
443
444 let mut mutants: Vec<RandomKeyChromosome> = (0..mutant_count)
446 .map(|_| RandomKeyChromosome::random(num_keys, rng))
447 .collect();
448
449 let crossover_count = self.config.population_size - elite_count - mutant_count;
451 let mut children: Vec<RandomKeyChromosome> = (0..crossover_count)
452 .map(|_| {
453 let elite_idx = rng.random_range(0..elite_count);
455 let elite_parent = &population[elite_idx];
456
457 let non_elite_idx = rng.random_range(elite_count..population.len());
459 let non_elite_parent = &population[non_elite_idx];
460
461 elite_parent.biased_crossover(non_elite_parent, self.config.elite_bias, rng)
463 })
464 .collect();
465
466 evaluate_within(&mut mutants, &start, self.config.time_limit, |batch| {
469 self.problem.evaluate_parallel(batch)
470 });
471 if expired(&start, self.config.time_limit) {
472 children.clear();
473 }
474 evaluate_within(&mut children, &start, self.config.time_limit, |batch| {
475 self.problem.evaluate_parallel(batch)
476 });
477
478 new_population.extend(mutants);
480 new_population.extend(children);
481
482 new_population.sort_by(|a, b| {
484 b.fitness()
485 .partial_cmp(&a.fitness())
486 .unwrap_or(std::cmp::Ordering::Equal)
487 });
488
489 let new_best_fitness = new_population[0].fitness();
491 if new_best_fitness > best_fitness {
492 best = new_population[0].clone();
493 best_fitness = new_best_fitness;
494 stagnation_count = 0;
495 } else {
496 stagnation_count += 1;
497 }
498
499 if let Some(limit) = self.config.stagnation_limit {
501 if stagnation_count >= limit {
502 break;
503 }
504 }
505
506 self.problem
508 .on_generation(generation, &best, &new_population);
509
510 if let Some(ref callback) = progress_callback {
512 let avg_fitness = new_population.iter().map(|c| c.fitness()).sum::<f64>()
513 / new_population.len() as f64;
514
515 callback(BrkgaProgress {
516 generation,
517 max_generations: self.config.max_generations,
518 best_fitness,
519 avg_fitness,
520 elapsed: start.elapsed(),
521 running: true,
522 });
523 }
524
525 population = new_population;
526 generation += 1;
527 }
528
529 history.push(best_fitness);
531
532 if let Some(ref callback) = progress_callback {
534 let avg_fitness = population.iter().map(|c| c.fitness()).sum::<f64>()
535 / population.len().max(1) as f64;
536
537 callback(BrkgaProgress {
538 generation,
539 max_generations: self.config.max_generations,
540 best_fitness,
541 avg_fitness,
542 elapsed: start.elapsed(),
543 running: false,
544 });
545 }
546
547 BrkgaResult {
548 best,
549 generations: generation,
550 elapsed: start.elapsed(),
551 target_reached,
552 history,
553 }
554 }
555}
556
557#[cfg(test)]
558mod tests {
559 use super::*;
560
561 struct MaxSumProblem {
563 num_keys: usize,
564 }
565
566 impl BrkgaProblem for MaxSumProblem {
567 fn num_keys(&self) -> usize {
568 self.num_keys
569 }
570
571 fn evaluate(&self, chromosome: &mut RandomKeyChromosome) {
572 let sum: f64 = chromosome.keys.iter().sum();
573 chromosome.set_fitness(sum);
574 }
575 }
576
577 #[test]
578 fn test_brkga_basic() {
579 let config = BrkgaConfig::default()
580 .with_population_size(50)
581 .with_max_generations(50);
582
583 let problem = MaxSumProblem { num_keys: 10 };
584 let runner = BrkgaRunner::new(config, problem);
585 let result = runner.run();
586
587 assert!(result.best.fitness() > 5.0);
589 }
590
591 #[test]
592 fn test_random_key_chromosome() {
593 let mut rng = rand::rng();
594 let chromosome = RandomKeyChromosome::random(10, &mut rng);
595
596 assert_eq!(chromosome.len(), 10);
597 for &key in &chromosome.keys {
598 assert!((0.0..1.0).contains(&key));
599 }
600 }
601
602 #[test]
603 fn test_biased_crossover() {
604 let mut rng = rand::rng();
605 let elite = RandomKeyChromosome::random(10, &mut rng);
606 let non_elite = RandomKeyChromosome::random(10, &mut rng);
607
608 let child = elite.biased_crossover(&non_elite, 0.7, &mut rng);
609
610 assert_eq!(child.len(), 10);
611 for &key in &child.keys {
612 assert!((0.0..1.0).contains(&key));
613 }
614 }
615
616 #[test]
617 fn test_decode_as_permutation() {
618 let mut chromosome = RandomKeyChromosome::new(5);
619 chromosome.keys = vec![0.3, 0.1, 0.9, 0.5, 0.2];
620
621 let perm = chromosome.decode_as_permutation();
622
623 assert_eq!(perm, vec![1, 4, 0, 3, 2]);
625 }
626
627 #[test]
628 fn test_decode_as_discrete() {
629 let mut chromosome = RandomKeyChromosome::new(3);
630 chromosome.keys = vec![0.0, 0.5, 0.99];
631
632 assert_eq!(chromosome.decode_as_discrete(0, 6), 0);
634 assert_eq!(chromosome.decode_as_discrete(1, 6), 3);
635 assert_eq!(chromosome.decode_as_discrete(2, 6), 5);
636 }
637}