Skip to main content

radiate_core/fitness/
novelty.rs

1use crate::{
2    BatchFitnessFunction, BatchedFn, CosineDistance, EuclideanDistance, FitnessFunction,
3    HammingDistance, diversity::Distance, math::knn::KNN,
4};
5use radiate_utils::WindowBuffer;
6use std::sync::{Arc, RwLock};
7
8const DEFAULT_ARCHIVE_SIZE: usize = 1000;
9const DEFAULT_K: usize = 15;
10const DEFAULT_THRESHOLD: f32 = 0.5;
11
12pub trait Novelty<T>: Send + Sync {
13    fn description(&self, member: &T) -> Vec<f32>;
14
15    /// Compute descriptors for a whole batch. Default fans out to `description`.
16    /// Override this on your own concrete `Novelty` impl if you can vectorise the
17    /// batch path (shared setup, SIMD, GPU, etc.) — the closure blanket impl
18    /// always takes the default.
19    fn batch_description(&self, members: &[T]) -> Vec<Vec<f32>> {
20        members.iter().map(|m| self.description(m)).collect()
21    }
22}
23
24impl<T, F> Novelty<T> for F
25where
26    F: Fn(&T) -> Vec<f32> + Send + Sync,
27{
28    fn description(&self, member: &T) -> Vec<f32> {
29        self(member)
30    }
31}
32
33impl<T, F> Novelty<T> for BatchedFn<F>
34where
35    F: Fn(&[T]) -> Vec<Vec<f32>> + Send + Sync,
36{
37    fn description(&self, member: &T) -> Vec<f32> {
38        (self.0)(std::slice::from_ref(member))
39            .into_iter()
40            .next()
41            .unwrap_or_default()
42    }
43
44    fn batch_description(&self, members: &[T]) -> Vec<Vec<f32>> {
45        (self.0)(members)
46    }
47}
48
49#[derive(Clone)]
50pub struct NoveltySearch<T> {
51    pub behavior: Arc<dyn Novelty<T>>,
52    pub archive: Arc<RwLock<WindowBuffer<Vec<f32>>>>,
53    pub k: usize,
54    pub threshold: f32,
55    pub distance_fn: Arc<dyn Distance<Vec<f32>>>,
56}
57
58impl<T> NoveltySearch<T> {
59    pub fn new<N>(behavior: N) -> Self
60    where
61        N: Novelty<T> + Send + Sync + 'static,
62    {
63        NoveltySearch {
64            behavior: Arc::new(behavior),
65            archive: Arc::new(RwLock::new(WindowBuffer::with_capacity(
66                DEFAULT_ARCHIVE_SIZE,
67            ))),
68            k: DEFAULT_K,
69            threshold: DEFAULT_THRESHOLD,
70            distance_fn: Arc::new(EuclideanDistance),
71        }
72    }
73
74    pub fn from_batch_fn<F>(f: F) -> Self
75    where
76        F: Fn(&[T]) -> Vec<Vec<f32>> + Send + Sync + 'static,
77        T: 'static,
78    {
79        Self::new(BatchedFn(f))
80    }
81
82    pub fn k(mut self, k: usize) -> Self {
83        self.k = k;
84        self
85    }
86
87    pub fn threshold(mut self, threshold: f32) -> Self {
88        self.threshold = threshold;
89        self
90    }
91
92    pub fn archive_size(mut self, size: usize) -> Self {
93        self.archive = Arc::new(RwLock::new(WindowBuffer::with_capacity(size)));
94        self
95    }
96
97    pub fn cosine_distance(mut self) -> Self {
98        self.distance_fn = Arc::new(CosineDistance);
99        self
100    }
101
102    pub fn euclidean_distance(mut self) -> Self {
103        self.distance_fn = Arc::new(EuclideanDistance);
104        self
105    }
106
107    pub fn hamming_distance(mut self) -> Self {
108        self.distance_fn = Arc::new(HammingDistance);
109        self
110    }
111
112    fn novelty_score(&self, descriptor: &Vec<f32>, archive: &WindowBuffer<Vec<f32>>) -> f32 {
113        let slice = archive.values();
114        let mut knn = KNN::new(slice, Arc::clone(&self.distance_fn));
115        let query = knn.query_point(descriptor, self.k);
116
117        let min_dist = query.min_distance;
118        let max_dist = query.max_distance;
119        let range = max_dist - min_dist;
120
121        if range < f32::EPSILON {
122            return if min_dist < f32::EPSILON { 0.0 } else { 0.5 };
123        }
124
125        let avg_dist = query.average_distance();
126        (avg_dist - min_dist) / range
127    }
128
129    fn evaluate_internal(&self, individual: &T) -> f32 {
130        let description = self.behavior.description(individual);
131        let mut archive = self.archive.write().unwrap();
132
133        if archive.is_empty() {
134            archive.push(description);
135            return 0.5;
136        }
137
138        let novelty = self.novelty_score(&description, &archive);
139        if novelty > self.threshold || archive.len() < self.k {
140            archive.push(description);
141        }
142
143        novelty
144    }
145
146    fn evaluate_batch_internal(&self, individuals: &[T]) -> Vec<f32> {
147        let descriptions = self.behavior.batch_description(individuals);
148        let mut archive = self.archive.write().unwrap();
149
150        if archive.is_empty() {
151            let result = vec![0.5; descriptions.len()];
152            for desc in descriptions {
153                archive.push(desc);
154            }
155
156            return result;
157        }
158
159        // Score every descriptor against the same pre-batch archive snapshot —
160        // no archive mutations happen during scoring, so individual-N's score
161        // does not depend on its position in the batch.
162        let mut scores = Vec::with_capacity(descriptions.len());
163        for desc in descriptions.into_iter() {
164            let score = self.novelty_score(&desc, &archive);
165
166            if score > self.threshold || archive.len() < self.k {
167                archive.push(desc);
168            }
169
170            scores.push(score);
171        }
172
173        scores
174    }
175}
176
177impl<T> FitnessFunction<T, f32> for NoveltySearch<T>
178where
179    T: Send + Sync,
180{
181    fn evaluate(&self, individual: T) -> f32 {
182        self.evaluate_internal(&individual)
183    }
184}
185
186impl<T> FitnessFunction<&T, f32> for NoveltySearch<T>
187where
188    T: Send + Sync,
189{
190    fn evaluate(&self, individual: &T) -> f32 {
191        self.evaluate_internal(individual)
192    }
193}
194
195impl<T> BatchFitnessFunction<T, f32> for NoveltySearch<T>
196where
197    T: Send + Sync,
198{
199    fn evaluate(&self, individuals: Vec<T>) -> Vec<f32> {
200        self.evaluate_batch_internal(&individuals)
201    }
202}
203
204#[cfg(test)]
205mod tests {
206    use super::*;
207    use crate::{BatchFitnessFunction, FitnessFunction};
208
209    fn make_ns(k: usize, threshold: f32) -> NoveltySearch<Vec<f32>> {
210        NoveltySearch::new(|v: &Vec<f32>| v.clone())
211            .k(k)
212            .threshold(threshold)
213            .archive_size(100)
214    }
215
216    fn seed(ns: &NoveltySearch<Vec<f32>>, points: impl IntoIterator<Item = Vec<f32>>) {
217        let mut archive = ns.archive.write().unwrap();
218        for p in points {
219            archive.push(p);
220        }
221    }
222
223    fn archive_view_len(ns: &NoveltySearch<Vec<f32>>) -> usize {
224        ns.archive.read().unwrap().values().len()
225    }
226
227    fn eval(ns: &NoveltySearch<Vec<f32>>, v: Vec<f32>) -> f32 {
228        <NoveltySearch<Vec<f32>> as FitnessFunction<Vec<f32>, f32>>::evaluate(ns, v)
229    }
230
231    fn eval_batch(ns: &NoveltySearch<Vec<f32>>, vs: Vec<Vec<f32>>) -> Vec<f32> {
232        <NoveltySearch<Vec<f32>> as BatchFitnessFunction<Vec<f32>, f32>>::evaluate(ns, vs)
233    }
234
235    #[test]
236    fn empty_archive_single_eval_scores_half_and_seeds_one_entry() {
237        let ns = make_ns(3, 0.5);
238        let score = eval(&ns, vec![1.0, 0.0]);
239        assert_eq!(score, 0.5);
240        assert_eq!(archive_view_len(&ns), 1);
241    }
242
243    #[test]
244    fn empty_archive_batch_eval_scores_half_and_seeds_every_entry() {
245        let ns = make_ns(3, 0.5);
246        let scores = eval_batch(
247            &ns,
248            vec![vec![0.0], vec![1.0], vec![2.0], vec![3.0], vec![4.0]],
249        );
250
251        assert_eq!(scores, vec![0.5; 5]);
252        assert_eq!(archive_view_len(&ns), 5);
253    }
254
255    #[test]
256    fn identical_to_sole_archive_point_scores_zero() {
257        let ns = make_ns(3, 0.99);
258        seed(&ns, [vec![1.0, 0.0]]);
259
260        // distance is clamped to 1e-12 inside KNN, which is < f32::EPSILON,
261        // so the degenerate branch returns 0.0.
262        let score = eval(&ns, vec![1.0, 0.0]);
263        assert_eq!(score, 0.0);
264    }
265
266    #[test]
267    fn different_from_sole_archive_point_scores_half_neutral() {
268        let ns = make_ns(3, 0.99);
269        seed(&ns, [vec![0.0]]);
270
271        // single archive point => min == max, range == 0, min > epsilon => 0.5
272        let score = eval(&ns, vec![5.0]);
273        assert_eq!(score, 0.5);
274    }
275
276    #[test]
277    fn bootstrap_admits_first_k_individuals_under_strict_threshold() {
278        // threshold 0.99 means scores effectively can't pass on merit.
279        // archive.len() < k path must still admit the first k.
280        let ns = make_ns(3, 0.99);
281
282        for _ in 0..3 {
283            eval(&ns, vec![1.0]);
284        }
285        assert_eq!(archive_view_len(&ns), 3);
286
287        // 4th identical individual: archive.len() == k, score won't beat 0.99 → not added.
288        eval(&ns, vec![1.0]);
289        assert_eq!(archive_view_len(&ns), 3);
290    }
291
292    #[test]
293    fn post_bootstrap_score_below_threshold_does_not_admit() {
294        let ns = make_ns(2, 0.2);
295        // 3 seeded points → past bootstrap (archive.len() >= k=2).
296        seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
297
298        // query=[2.0]: 2-NN are [0.0]@2 and [5.0]@3, global_max=8 (to [10.0]).
299        // avg=2.5, range=6 → score = 0.5/6 ≈ 0.0833 < 0.2 → not added.
300        let score = eval(&ns, vec![2.0]);
301        assert!(
302            (score - 0.083_333).abs() < 1e-4,
303            "expected ≈0.0833, got {score}"
304        );
305        assert_eq!(archive_view_len(&ns), 3);
306    }
307
308    #[test]
309    fn novelty_score_matches_normalization_formula() {
310        // threshold = -1 so admission never filters; we want to inspect the score itself.
311        let ns = make_ns(2, -1.0);
312        seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
313
314        // query=[12.0]: 2-NN are [10.0]@2 and [5.0]@7, global_max=12 (to [0.0]).
315        // avg=4.5, range=10 → score = 2.5/10 = 0.25 exact.
316        let score = eval(&ns, vec![12.0]);
317        assert!((score - 0.25).abs() < 1e-5, "expected 0.25, got {score}");
318    }
319
320    #[test]
321    fn novelty_score_with_k_equal_to_archive_size_uses_all_points() {
322        let ns = make_ns(3, -1.0);
323        seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
324
325        // k >= n branch: cluster contains all archive points sorted ascending.
326        // query=[4.0] distances: 4, 1, 6 → min=1, max=6, avg=11/3.
327        // score = (11/3 - 1) / 5 = (8/3) / 5 = 8/15 ≈ 0.5333.
328        let score = eval(&ns, vec![4.0]);
329        assert!(
330            (score - 8.0 / 15.0).abs() < 1e-5,
331            "expected 8/15 ≈ 0.5333, got {score}"
332        );
333    }
334
335    #[test]
336    fn novelty_score_always_in_unit_interval() {
337        let ns = make_ns(2, -1.0);
338        seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
339
340        for x in [-100.0, -1.0, 0.0, 2.5, 5.0, 7.5, 10.0, 12.0, 100.0] {
341            // fresh archive each iteration so writes from earlier queries don't drift.
342            let ns = make_ns(2, -1.0);
343            seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
344
345            let score = eval(&ns, vec![x]);
346            assert!(
347                (0.0..=1.0).contains(&score),
348                "score {score} out of [0,1] for x={x}"
349            );
350        }
351    }
352
353    #[test]
354    fn score_above_threshold_admits_to_archive() {
355        let ns = make_ns(2, 0.2);
356        seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
357        assert_eq!(archive_view_len(&ns), 3);
358
359        // query=[12.0] → score 0.25 > 0.2 → admitted.
360        let score = eval(&ns, vec![12.0]);
361        assert!(score > 0.2, "expected > threshold, got {score}");
362        assert_eq!(archive_view_len(&ns), 4);
363    }
364
365    #[test]
366    fn archive_window_caps_at_configured_size() {
367        // archive_size=5, threshold=-1 (always admit).
368        let ns = NoveltySearch::new(|v: &Vec<f32>| v.clone())
369            .k(1)
370            .threshold(-1.0)
371            .archive_size(5);
372
373        for i in 0..40 {
374            eval(&ns, vec![i as f32 * 100.0]);
375        }
376
377        // The k-NN sees archive.values(), which is the live window.
378        let archive = ns.archive.read().unwrap();
379        assert!(
380            archive.values().len() <= 5,
381            "archive view exceeds window cap: {}",
382            archive.values().len()
383        );
384    }
385
386    #[test]
387    fn fitness_function_ref_variant_evaluates_and_admits() {
388        let ns = make_ns(1, 0.5);
389        let ind = vec![1.0, 0.0];
390        let score = eval(&ns, ind);
391        assert_eq!(score, 0.5);
392        assert_eq!(archive_view_len(&ns), 1);
393    }
394
395    #[test]
396    fn batch_eval_returns_one_score_per_individual() {
397        let ns = make_ns(3, 0.5);
398        let scores = eval_batch(&ns, vec![vec![0.0], vec![5.0], vec![10.0]]);
399        assert_eq!(scores.len(), 3);
400        for (i, &s) in scores.iter().enumerate() {
401            assert!((0.0..=1.0).contains(&s), "scores[{i}] = {s} out of [0,1]");
402        }
403    }
404
405    #[test]
406    fn batch_eval_admits_via_running_archive_size_for_bootstrap() {
407        // Scoring is against a frozen pre-batch snapshot, but the admission pass
408        // walks the batch with the live archive size so bootstrap (`archive.len()
409        // < k`) still works mid-batch when the pre-batch archive is undersized.
410        let ns = make_ns(2, -1.0); // threshold=-1 → score-based admission also passes.
411        seed(&ns, [vec![0.0], vec![10.0]]);
412        let initial = archive_view_len(&ns);
413
414        let scores = eval_batch(&ns, vec![vec![5.0], vec![20.0], vec![-5.0]]);
415        assert_eq!(scores.len(), 3);
416        assert_eq!(archive_view_len(&ns), initial + 3);
417    }
418
419    #[test]
420    fn batch_eval_does_not_score_against_intra_batch_additions() {
421        // If the batch were "online" (admitting earlier members before scoring
422        // later ones), then a duplicate entry later in the batch would see its
423        // earlier copy as a near-zero-distance neighbour and score ~0.
424        // True batch should score the duplicate against only the pre-batch
425        // archive — yielding the same score as the original.
426        let ns = make_ns(2, -1.0);
427        seed(&ns, [vec![0.0], vec![10.0]]);
428
429        let scores = eval_batch(&ns, vec![vec![5.0], vec![5.0]]);
430        assert_eq!(scores.len(), 2);
431        assert!(
432            (scores[0] - scores[1]).abs() < 1e-6,
433            "duplicate batch members should score identically: {scores:?}"
434        );
435    }
436
437    #[test]
438    fn from_batch_fn_routes_batch_through_user_closure_and_falls_back_per_item() {
439        use std::sync::atomic::{AtomicUsize, Ordering};
440
441        let batch_calls = Arc::new(AtomicUsize::new(0));
442        let total_seen = Arc::new(AtomicUsize::new(0));
443
444        let ns: NoveltySearch<Vec<f32>> = {
445            let batch_calls = Arc::clone(&batch_calls);
446            let total_seen = Arc::clone(&total_seen);
447            NoveltySearch::from_batch_fn(move |members: &[Vec<f32>]| {
448                batch_calls.fetch_add(1, Ordering::Relaxed);
449                total_seen.fetch_add(members.len(), Ordering::Relaxed);
450                members.iter().map(|v| v.clone()).collect()
451            })
452            .k(3)
453            .threshold(0.5)
454            .archive_size(100)
455        };
456
457        // Single-eval routes through the per-item fallback, which calls the
458        // batch closure with a 1-element slice.
459        let _ = eval(&ns, vec![1.0, 0.0]);
460        assert_eq!(batch_calls.load(Ordering::Relaxed), 1);
461        assert_eq!(total_seen.load(Ordering::Relaxed), 1);
462
463        // Batch eval calls the closure once with the full slice — the fast path.
464        let _ = eval_batch(&ns, vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0]]);
465        assert_eq!(batch_calls.load(Ordering::Relaxed), 2);
466        assert_eq!(total_seen.load(Ordering::Relaxed), 5);
467    }
468
469    #[test]
470    fn clone_shares_archive_with_original() {
471        let ns = make_ns(3, 0.5);
472        let twin = ns.clone();
473
474        eval(&ns, vec![1.0]);
475        eval(&ns, vec![2.0]);
476
477        // Same Arc<RwLock<...>> backing both handles.
478        assert_eq!(archive_view_len(&twin), 2);
479    }
480
481    #[test]
482    fn cosine_distance_identical_direction_scores_zero_in_degenerate_case() {
483        let ns = NoveltySearch::new(|v: &Vec<f32>| v.clone())
484            .k(3)
485            .threshold(0.99)
486            .archive_size(100)
487            .cosine_distance();
488        seed(&ns, [vec![1.0, 0.0]]);
489
490        // Same direction, different magnitude → cosine distance 0 (clamped to 1e-12).
491        let score = eval(&ns, vec![100.0, 0.0]);
492        assert_eq!(score, 0.0);
493    }
494
495    #[test]
496    fn concurrent_evaluation_does_not_panic_or_deadlock() {
497        use std::thread;
498
499        let ns = Arc::new(
500            NoveltySearch::new(|v: &Vec<f32>| v.clone())
501                .k(5)
502                .threshold(0.3)
503                .archive_size(200),
504        );
505
506        let handles: Vec<_> = (0..8)
507            .map(|i| {
508                let ns = Arc::clone(&ns);
509                thread::spawn(move || {
510                    for j in 0..50 {
511                        let v = (i * 50 + j) as f32;
512                        eval(&ns, vec![v, v * 0.5]);
513                    }
514                })
515            })
516            .collect();
517
518        for h in handles {
519            h.join().expect("worker thread panicked");
520        }
521
522        // 8 threads × 50 evals = 400 attempts; capped by archive_size=200.
523        let archive = ns.archive.read().unwrap();
524        assert!(archive.values().len() <= 200);
525    }
526}