use rayon::prelude::*;
fn sorted_profile_row_parallel<K>(n: usize, anchor: usize, pair_key: &K) -> Vec<f64>
where
K: Fn(usize, usize) -> f64 + Sync,
{
let mut profile: Vec<f64> = (0..n)
.into_par_iter()
.map(|row| pair_key(anchor, row))
.collect();
profile.par_sort_by(f64::total_cmp);
profile
}
fn sorted_profile_serial<K>(n: usize, anchor: usize, pair_key: &K) -> Vec<f64>
where
K: Fn(usize, usize) -> f64 + Sync,
{
let mut profile: Vec<f64> = (0..n).map(|row| pair_key(anchor, row)).collect();
profile.sort_by(f64::total_cmp);
profile
}
pub(crate) fn sorted_profile_cmp(a: &[f64], b: &[f64]) -> std::cmp::Ordering {
for (x, y) in a.iter().zip(b.iter()) {
let ordering = x.total_cmp(y);
if !ordering.is_eq() {
return ordering;
}
}
std::cmp::Ordering::Equal
}
pub(crate) fn resolve_sorted_profile_tie<K, F>(
n: usize,
tied: &[usize],
pair_key: K,
on_profile_builds: &mut F,
) -> Vec<usize>
where
K: Fn(usize, usize) -> f64 + Sync,
F: FnMut(usize),
{
if tied.len() <= 1 {
return tied.to_vec();
}
on_profile_builds(tied.len());
let profiles: Vec<Vec<f64>> = if tied.len() >= rayon::current_num_threads() {
tied.par_iter()
.map(|&i| sorted_profile_serial(n, i, &pair_key))
.collect()
} else {
tied.iter()
.map(|&i| sorted_profile_row_parallel(n, i, &pair_key))
.collect()
};
let mut least = 0usize;
for candidate in 1..profiles.len() {
if sorted_profile_cmp(&profiles[candidate], &profiles[least]).is_lt() {
least = candidate;
}
}
(0..tied.len())
.filter(|&k| sorted_profile_cmp(&profiles[k], &profiles[least]).is_eq())
.map(|k| tied[k])
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extremum_then_refine_matches_the_incumbent_comparator_scan() {
let points: Vec<[f64; 2]> = vec![
[-1.0, -1.0],
[-1.0, 1.0],
[1.0, -1.0],
[1.0, 1.0],
[0.0, 0.0],
[0.25, -0.5],
[-0.75, 0.125],
];
let n = points.len();
let pair_key = |i: usize, j: usize| {
let dx = points[i][0] - points[j][0];
let dy = points[i][1] - points[j][1];
dx * dx + dy * dy
};
let profile = |i: usize| sorted_profile_serial(n, i, &pair_key);
for tied in [
vec![0_usize, 1, 2, 3],
vec![0, 1, 2, 3, 4],
vec![4, 5, 6],
vec![5],
vec![1, 3],
(0..n).collect::<Vec<usize>>(),
] {
let pivot = tied
.iter()
.copied()
.reduce(|a, b| {
if sorted_profile_cmp(&profile(a), &profile(b)).is_lt() {
a
} else {
b
}
})
.expect("non-empty tie class");
let expected: Vec<usize> = tied
.iter()
.copied()
.filter(|&i| sorted_profile_cmp(&profile(i), &profile(pivot)).is_eq())
.collect();
let mut builds = 0usize;
let got = resolve_sorted_profile_tie(n, &tied, &pair_key, &mut |b| builds += b);
assert_eq!(got, expected, "tie class {tied:?} resolved differently");
assert_eq!(
builds,
if tied.len() > 1 { tied.len() } else { 0 },
"tie class {tied:?} built {builds} profiles"
);
}
}
#[test]
fn a_singleton_tie_class_builds_no_profile() {
let mut builds = 0usize;
let got =
resolve_sorted_profile_tie(1_000, &[7], |i, j| (i as f64) - (j as f64), &mut |b| {
builds += b
});
assert_eq!(got, vec![7]);
assert_eq!(
builds, 0,
"a lone candidate needs no profile to win a refinement"
);
}
#[test]
fn a_real_tie_class_builds_exactly_one_profile_per_candidate() {
let mut builds = 0usize;
let tied = [0_usize, 1, 2, 3, 4];
let resolved =
resolve_sorted_profile_tie(64, &tied, |i, j| ((i * j) % 7) as f64, &mut |b| {
builds += b
});
assert_eq!(builds, tied.len());
assert!(!resolved.is_empty());
assert!(resolved.iter().all(|index| tied.contains(index)));
}
}