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 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 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 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 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 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 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 seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
297
298 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 let ns = make_ns(2, -1.0);
312 seed(&ns, [vec![0.0], vec![5.0], vec![10.0]]);
313
314 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 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 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 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 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 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 let ns = make_ns(2, -1.0); 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 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 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 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 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 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 let archive = ns.archive.read().unwrap();
524 assert!(archive.values().len() <= 200);
525 }
526}