Skip to main content

radiate_selectors/
nsga3.rs

1use radiate_core::{Chromosome, Objective, Optimize, Phenotype, Select, pareto};
2use radiate_utils::Matrix;
3use std::cmp::Ordering;
4use std::sync::{Arc, OnceLock};
5
6const EPS: f32 = 1e-12;
7
8/// Reference directions plus their precomputed self-dot norms, so
9/// `nearest_reference_direction` never recomputes `dot(dir, dir)` per call.
10#[derive(Debug)]
11pub struct ReferenceDirs {
12    dirs: Vec<Vec<f32>>,
13    norms: Vec<f32>,
14}
15
16impl ReferenceDirs {
17    pub fn new(dirs: Vec<Vec<f32>>) -> Self {
18        let norms = dirs.iter().map(|d| dot(d, d)).collect();
19        Self { dirs, norms }
20    }
21
22    fn build(dims: usize, partitions: usize) -> Self {
23        Self::new(pareto::das_dennis(dims, partitions))
24    }
25}
26
27#[derive(Debug, Clone)]
28pub struct NSGA3Selector {
29    ref_dirs: Arc<OnceLock<Arc<ReferenceDirs>>>,
30    partitions: usize,
31}
32
33impl NSGA3Selector {
34    pub fn new(partitions: usize) -> Self {
35        Self {
36            ref_dirs: Arc::new(OnceLock::new()),
37            partitions,
38        }
39    }
40
41    pub fn partitions(&self) -> usize {
42        self.partitions
43    }
44
45    fn reference_dirs(&self, dims: usize) -> Arc<ReferenceDirs> {
46        self.ref_dirs
47            .get_or_init(|| Arc::new(ReferenceDirs::build(dims, self.partitions)))
48            .clone()
49    }
50}
51
52impl<C: Chromosome> Select<C> for NSGA3Selector {
53    fn select(
54        &self,
55        population: &[Phenotype<C>],
56        objective: &Objective,
57        count: usize,
58    ) -> Vec<usize> {
59        if population.is_empty() || count == 0 {
60            return Vec::new();
61        }
62
63        let scores = population
64            .iter()
65            .filter_map(|p| p.score())
66            .map(|score| to_minimization_space(score.as_ref(), objective))
67            .collect::<Matrix<f32>>();
68
69        let score_rows = scores.iter().collect::<Vec<&[f32]>>();
70
71        let min_objective = minimization_objective(objective.dims());
72        let ranks = pareto::rank(&score_rows, &min_objective);
73        let fronts = fronts_from_ranks(&ranks);
74
75        let ref_dirs = self.reference_dirs(objective.dims());
76
77        let mut selected = Vec::with_capacity(count);
78        let mut front_idx = 0;
79
80        while front_idx < fronts.len() && selected.len() + fronts[front_idx].len() <= count {
81            selected.extend_from_slice(&fronts[front_idx]);
82            front_idx += 1;
83        }
84
85        if selected.len() < count && front_idx < fronts.len() {
86            let remaining = count - selected.len();
87
88            selected.extend(niching_fill(
89                &scores,
90                &ref_dirs,
91                &selected,
92                &fronts[front_idx],
93                remaining,
94            ));
95        }
96
97        selected.into_iter().take(count).collect()
98    }
99}
100
101#[inline]
102fn minimization_objective(dims: usize) -> Objective {
103    Objective::Multi(vec![Optimize::Minimize; dims])
104}
105
106#[inline]
107pub fn to_minimization_space(score: &[f32], objective: &Objective) -> Vec<f32> {
108    match objective {
109        Objective::Single(opt) => {
110            if *opt == Optimize::Minimize {
111                score.to_vec()
112            } else {
113                score.iter().map(|&x| -x).collect()
114            }
115        }
116        Objective::Multi(opts) => score
117            .iter()
118            .zip(opts.iter())
119            .map(|(&x, opt)| if *opt == Optimize::Minimize { x } else { -x })
120            .collect(),
121    }
122}
123
124#[inline]
125pub fn fronts_from_ranks(ranks: &[usize]) -> Vec<Vec<usize>> {
126    if ranks.is_empty() {
127        return Vec::new();
128    }
129
130    let max_rank = *ranks.iter().max().unwrap_or(&0);
131    let mut fronts = vec![Vec::<usize>::new(); max_rank + 1];
132
133    for (idx, &rank) in ranks.iter().enumerate() {
134        fronts[rank].push(idx);
135    }
136
137    while fronts.last().is_some_and(|front| front.is_empty()) {
138        fronts.pop();
139    }
140
141    fronts
142}
143
144#[derive(Debug, Clone)]
145pub struct ObjectiveBounds {
146    ideal: Vec<f32>,
147    nadir: Vec<f32>,
148}
149
150impl ObjectiveBounds {
151    pub fn from_scores(scores: &Matrix<f32>) -> Self {
152        if scores.is_empty() {
153            return Self {
154                ideal: Vec::new(),
155                nadir: Vec::new(),
156            };
157        }
158
159        let dims = scores.cols();
160        let mut ideal = vec![f32::INFINITY; dims];
161        let mut nadir = vec![f32::NEG_INFINITY; dims];
162
163        for row in scores.iter() {
164            for dim in 0..dims {
165                ideal[dim] = ideal[dim].min(row[dim]);
166                nadir[dim] = nadir[dim].max(row[dim]);
167            }
168        }
169
170        Self { ideal, nadir }
171    }
172
173    pub fn normalize(&self, score: &[f32]) -> Vec<f32> {
174        score
175            .iter()
176            .enumerate()
177            .map(|(dim, &value)| {
178                let den = self.nadir[dim] - self.ideal[dim];
179
180                if !den.is_finite() || den.abs() <= EPS {
181                    0.0
182                } else {
183                    (value - self.ideal[dim]) / den
184                }
185            })
186            .collect()
187    }
188}
189
190#[derive(Clone, Copy, Debug)]
191struct Association {
192    idx: usize,
193    niche: usize,
194    distance: f32,
195}
196
197/// Given:
198/// - `already_selected`: indices chosen from earlier fronts
199/// - `last_front`: indices in the partial front
200///
201/// Returns additional indices from `last_front` using NSGA-III niching.
202pub fn niching_fill(
203    scores: &Matrix<f32>,
204    ref_dirs: &ReferenceDirs,
205    already_selected: &[usize],
206    last_front: &[usize],
207    remaining: usize,
208) -> Vec<usize> {
209    if remaining == 0 || last_front.is_empty() || ref_dirs.dirs.is_empty() {
210        return Vec::new();
211    }
212
213    let bounds = ObjectiveBounds::from_scores(scores);
214    let mut niche_count = vec![0usize; ref_dirs.dirs.len()];
215
216    for &idx in already_selected {
217        let normalized = bounds.normalize(&scores[idx]);
218        let (niche, _) = nearest_reference_direction(&normalized, ref_dirs);
219        niche_count[niche] += 1;
220    }
221
222    let mut candidates = last_front
223        .iter()
224        .map(|&idx| {
225            let normalized = bounds.normalize(&scores[idx]);
226            let (niche, distance) = nearest_reference_direction(&normalized, ref_dirs);
227
228            Association {
229                idx,
230                niche,
231                distance,
232            }
233        })
234        .collect::<Vec<_>>();
235
236    let mut picked = Vec::with_capacity(remaining);
237
238    while picked.len() < remaining && !candidates.is_empty() {
239        let niche = least_crowded_candidate_niche(&candidates, &niche_count);
240        let candidate_idx = closest_candidate_in_niche(&candidates, niche);
241
242        let selected = candidates.swap_remove(candidate_idx);
243
244        picked.push(selected.idx);
245        niche_count[selected.niche] += 1;
246    }
247
248    picked
249}
250
251#[inline]
252fn least_crowded_candidate_niche(candidates: &[Association], niche_count: &[usize]) -> usize {
253    candidates
254        .iter()
255        .map(|candidate| candidate.niche)
256        .min_by_key(|&niche| niche_count[niche])
257        .unwrap()
258}
259
260#[inline]
261fn closest_candidate_in_niche(candidates: &[Association], niche: usize) -> usize {
262    candidates
263        .iter()
264        .enumerate()
265        .filter(|(_, candidate)| candidate.niche == niche)
266        .min_by(|(_, a), (_, b)| {
267            a.distance
268                .partial_cmp(&b.distance)
269                .unwrap_or(Ordering::Equal)
270        })
271        .map(|(idx, _)| idx)
272        .unwrap()
273}
274
275#[inline]
276pub fn nearest_reference_direction(point: &[f32], refs: &ReferenceDirs) -> (usize, f32) {
277    let mut best = (0usize, f32::INFINITY);
278
279    for (idx, direction) in refs.dirs.iter().enumerate() {
280        let direction_norm = refs.norms[idx];
281
282        if direction_norm <= EPS || !direction_norm.is_finite() {
283            continue;
284        }
285
286        let projection = dot(point, direction) / direction_norm;
287        let distance = perpendicular_distance(point, direction, projection);
288
289        if distance < best.1 {
290            best = (idx, distance);
291        }
292    }
293
294    best
295}
296
297#[inline]
298fn perpendicular_distance(point: &[f32], direction: &[f32], projection: f32) -> f32 {
299    point
300        .iter()
301        .zip(direction)
302        .map(|(&p, &d)| {
303            let diff = p - projection * d;
304            diff * diff
305        })
306        .sum::<f32>()
307        .sqrt()
308}
309
310#[inline]
311fn dot(a: &[f32], b: &[f32]) -> f32 {
312    a.iter().zip(b).map(|(&x, &y)| x * y).sum()
313}