Skip to main content

radiate_core/objectives/
front.rs

1use crate::objectives::{Objective, Scored, pareto};
2#[cfg(feature = "serde")]
3use serde::{Deserialize, Serialize};
4use std::{cmp::Ordering, ops::Range};
5
6const DEFAULT_ENTROPY_BINS: usize = 20;
7
8#[derive(Debug)]
9pub struct FrontAddResult {
10    pub added_count: usize,
11    pub removed_count: usize,
12    pub comparisons: usize,
13    pub filter_count: usize,
14    pub size: usize,
15}
16
17#[derive(Clone, Default)]
18struct FrontScratch {
19    remove_buff: Vec<usize>,
20    index_buff: Vec<usize>,
21    crowding_buff: Vec<f32>,
22    filter_buff: Vec<bool>,
23}
24
25#[derive(Clone)]
26#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
27pub struct Front<T: Scored> {
28    values: Vec<T>,
29    range: Range<usize>,
30    objective: Objective,
31
32    #[cfg_attr(feature = "serde", serde(skip))]
33    scratch: FrontScratch,
34}
35
36impl<T: Scored> Front<T> {
37    pub fn new(range: Range<usize>, objective: Objective) -> Self {
38        Front {
39            values: Vec::new(),
40            range,
41            objective,
42            scratch: FrontScratch::default(),
43        }
44    }
45
46    pub fn len(&self) -> usize {
47        self.values.len()
48    }
49
50    pub fn range(&self) -> Range<usize> {
51        self.range.clone()
52    }
53
54    pub fn objective(&self) -> Objective {
55        self.objective.clone()
56    }
57
58    pub fn is_empty(&self) -> bool {
59        self.values.is_empty()
60    }
61
62    pub fn values(&self) -> &[T] {
63        &self.values
64    }
65
66    pub fn crowding_distance(&mut self) -> Option<&[f32]> {
67        let scores = self
68            .values
69            .iter()
70            .filter_map(|v| v.score())
71            .collect::<Vec<_>>();
72
73        if scores.is_empty() {
74            return None;
75        }
76
77        self.scratch.crowding_buff.clear();
78        self.scratch.crowding_buff.resize(scores.len(), 0.0);
79
80        pareto::buffered_crowding_distance(&scores, &mut self.scratch.crowding_buff);
81
82        Some(&self.scratch.crowding_buff[..])
83    }
84
85    pub fn entropy(&mut self) -> Option<f32> {
86        let scores = self
87            .values
88            .iter()
89            .filter_map(|v| v.score())
90            .collect::<Vec<_>>();
91
92        if scores.is_empty() {
93            return None;
94        }
95
96        Some(pareto::entropy(scores.as_slice(), DEFAULT_ENTROPY_BINS))
97    }
98
99    pub fn try_add_all<'a>(&mut self, items: impl Iterator<Item = &'a T>) -> FrontAddResult
100    where
101        T: Eq + Clone + 'static,
102    {
103        let mut added_count = 0;
104        let mut removed_count = 0;
105        let mut comparisons = 0;
106        let mut filter_count = 0;
107
108        for new_member in items.into_iter() {
109            self.scratch.remove_buff.clear();
110
111            let mut accept = true;
112            for (idx, existing) in self.values.iter().enumerate() {
113                if existing == new_member {
114                    accept = false;
115                    break;
116                }
117
118                match self.dom_cmp(existing, new_member) {
119                    Ordering::Greater => {
120                        // existing dominates new -> reject
121                        accept = false;
122                        comparisons += 1;
123                        break;
124                    }
125                    Ordering::Less => {
126                        // new dominates existing -> mark for removal
127                        self.scratch.remove_buff.push(idx);
128                        comparisons += 1;
129                    }
130                    Ordering::Equal => comparisons += 1,
131                }
132            }
133
134            if !accept {
135                continue;
136            }
137
138            // Remove dominated existing values efficiently (swap_remove).
139            // Need stable removal: remove in descending index order.
140            if !self.scratch.remove_buff.is_empty() {
141                self.scratch.remove_buff.sort_unstable();
142                self.scratch.remove_buff.dedup();
143
144                removed_count += self.scratch.remove_buff.len();
145
146                for &idx in self.scratch.remove_buff.iter().rev() {
147                    self.values.swap_remove(idx);
148                }
149            }
150
151            self.values.push(new_member.clone());
152            added_count += 1;
153
154            if self.values.len() > self.range.end {
155                self.fast_filter();
156                filter_count += 1;
157            }
158        }
159
160        FrontAddResult {
161            added_count,
162            removed_count,
163            comparisons,
164            filter_count,
165            size: self.values.len(),
166        }
167    }
168
169    /// Remove points with crowding distance in the top `trim` fraction.
170    /// Example: trim=0.02 removes the top 2% most isolated points.
171    #[inline]
172    pub fn remove_outliers(&mut self, trim: f32) -> Option<usize> {
173        if self.values.len() < 4 {
174            return None;
175        }
176
177        let trim = trim.clamp(0.0, 0.5);
178        if trim == 0.0 {
179            return None;
180        }
181
182        let (n, _) = self.score_dims()?;
183
184        let drop = ((n as f32) * trim).floor() as usize;
185        if drop == 0 {
186            return None;
187        }
188
189        let scores = self
190            .values
191            .iter()
192            .filter_map(|v| v.score())
193            .collect::<Vec<_>>();
194
195        self.scratch.crowding_buff.clear();
196        self.scratch.crowding_buff.resize(scores.len(), 0.0);
197
198        self.scratch.index_buff.clear();
199        self.scratch.index_buff.extend(0..scores.len());
200
201        pareto::buffered_crowding_distance(&scores, &mut self.scratch.crowding_buff);
202
203        self.scratch.index_buff.sort_unstable_by(|&i, &j| {
204            self.scratch.crowding_buff[j]
205                .partial_cmp(&self.scratch.crowding_buff[i])
206                .unwrap_or(Ordering::Equal)
207        });
208
209        self.scratch.index_buff.truncate(drop);
210        self.scratch.index_buff.sort_unstable();
211        self.scratch.index_buff.dedup();
212
213        let removed = self.scratch.index_buff.len();
214        for &idx in self.scratch.index_buff.iter().rev() {
215            self.values.swap_remove(idx);
216        }
217
218        Some(removed)
219    }
220
221    pub fn fronts(&mut self) -> Vec<Front<T>>
222    where
223        T: Clone + Eq + Send + Sync + 'static,
224    {
225        let mut fronts: Vec<Front<T>> = Vec::new();
226        for member in self.values.iter() {
227            let mut updated = false;
228
229            for front in fronts.iter_mut() {
230                let result = front.try_add_all(std::iter::once(member));
231
232                if result.added_count > 0 {
233                    updated = true;
234                    break;
235                }
236            }
237
238            if !updated {
239                let mut new_front = Front::new(self.range.clone(), self.objective.clone());
240                new_front.try_add_all(std::iter::once(member));
241                fronts.push(new_front);
242            }
243        }
244
245        fronts
246    }
247
248    fn fast_filter(&mut self) {
249        let keep = self.range.start.min(self.values.len());
250        if keep == 0 || self.values.len() <= keep {
251            return;
252        }
253
254        let scores = self
255            .values
256            .iter()
257            .filter_map(|v| v.score())
258            .collect::<Vec<_>>();
259
260        self.scratch.crowding_buff.clear();
261        self.scratch.crowding_buff.resize(scores.len(), 0.0);
262
263        self.scratch.index_buff.clear();
264        self.scratch.index_buff.extend(0..scores.len());
265
266        pareto::buffered_crowding_distance(&scores, &mut self.scratch.crowding_buff);
267
268        self.scratch
269            .index_buff
270            .select_nth_unstable_by(keep, |&a, &b| {
271                self.scratch.crowding_buff[b]
272                    .partial_cmp(&self.scratch.crowding_buff[a])
273                    .unwrap_or(Ordering::Equal)
274            });
275        self.scratch.index_buff.truncate(keep);
276
277        self.retain_indices();
278    }
279
280    #[inline]
281    fn dom_cmp(&self, one: &T, two: &T) -> Ordering {
282        let one_score = one.score();
283        let two_score = two.score();
284
285        if one_score.is_none() || two_score.is_none() {
286            return Ordering::Equal;
287        }
288
289        if let Some((a, b)) = one_score.zip(two_score) {
290            if pareto::dominance(a, b, &self.objective) {
291                return Ordering::Greater;
292            } else if pareto::dominance(b, a, &self.objective) {
293                return Ordering::Less;
294            }
295        }
296        Ordering::Equal
297    }
298
299    /// Keep only the elements at `indices`, in `self.values`' current order.
300    /// Reuses `scratch.keep_true` as scan-line scratch instead of allocating
301    /// a fresh mask on every call.
302    fn retain_indices(&mut self) {
303        self.scratch.filter_buff.clear();
304        self.scratch.filter_buff.resize(self.values.len(), false);
305
306        for &idx in self.scratch.index_buff.iter() {
307            self.scratch.filter_buff[idx] = true;
308        }
309
310        // Bind disjoint fields locally so the borrow checker sees this as
311        // two separate borrows of `self.values` and `self.scratch.keep_true`
312        // rather than one borrow of `self`.
313        let values = &mut self.values;
314        let keep_true = &self.scratch.filter_buff;
315
316        let mut idx = 0;
317        values.retain(|_| {
318            let retain = keep_true[idx];
319            idx += 1;
320            retain
321        });
322    }
323
324    #[inline]
325    fn score_dims(&self) -> Option<(usize, usize)> {
326        let n = self.values.len();
327
328        if n == 0 {
329            return None;
330        }
331
332        let first = self.values.iter().find_map(|v| v.score())?;
333        Some((n, first.len()))
334    }
335}
336
337impl<T> Default for Front<T>
338where
339    T: Scored,
340{
341    fn default() -> Self {
342        Front::new(0..0, Objective::default())
343    }
344}