use ndarray::Array2;
use crate::duplicates::PopulationCleaner;
use crate::helpers::linalg::cross_euclidean_distances;
#[derive(Debug, Clone)]
pub struct CloseDuplicatesCleaner {
pub epsilon: f64,
}
impl CloseDuplicatesCleaner {
pub fn new(epsilon: f64) -> Self {
Self { epsilon }
}
}
impl PopulationCleaner for CloseDuplicatesCleaner {
fn remove(&self, population: Array2<f64>, reference: Option<&Array2<f64>>) -> Array2<f64> {
let ref_array = reference.unwrap_or(&population);
let n = population.nrows();
let num_cols = population.ncols();
let dists_sq = cross_euclidean_distances(&population, ref_array);
let eps_sq = self.epsilon;
let mut keep = vec![true; n];
if let Some(ref_pop) = reference {
for i in 0..n {
for j in 0..ref_pop.nrows() {
if dists_sq[(i, j)] <= eps_sq {
keep[i] = false;
break;
}
}
}
} else {
for i in 0..n {
if !keep[i] {
continue;
}
for j in (i + 1)..n {
if dists_sq[(i, j)] < eps_sq {
keep[j] = false;
}
}
}
}
let kept_rows: Vec<_> = population
.outer_iter()
.enumerate()
.filter_map(|(i, row)| if keep[i] { Some(row.to_owned()) } else { None })
.collect();
let data_flat: Vec<f64> = kept_rows.into_iter().flatten().collect();
Array2::<f64>::from_shape_vec((data_flat.len() / num_cols, num_cols), data_flat)
.expect("Failed to create deduplicated Array2")
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn test_close_duplicates_cleaner_without_reference() {
let population = array![
[1.0, 2.0, 3.0],
[1.05, 2.05, 3.05], [4.0, 5.0, 6.0]
];
let epsilon = 0.1;
let cleaner = CloseDuplicatesCleaner::new(epsilon);
let cleaned = cleaner.remove(population, None);
assert_eq!(cleaned.nrows(), 2);
}
#[test]
fn test_close_duplicates_cleaner_with_reference() {
let population = array![[1.0, 2.0, 3.0], [10.0, 10.0, 10.0]];
let reference = array![
[1.01, 2.01, 3.01] ];
let epsilon = 0.05;
let cleaner = CloseDuplicatesCleaner::new(epsilon);
let cleaned = cleaner.remove(population, Some(&reference));
assert_eq!(cleaned.nrows(), 1);
assert_eq!(cleaned.row(0).to_vec(), vec![10.0, 10.0, 10.0]);
}
}