Skip to main content

affinityprop/
preference.rs

1use ndarray::{Array1, Array2, Zip};
2use num_traits::Float;
3
4/// Preference is the value representing the degree to which a data point will act as its own exemplar,
5/// with lower (more negative) values yielding fewer clusters.
6///
7/// - Median: Use median similarity value as preference
8/// - List: Use provided preference list
9/// - Value: Assign all members the same preference value
10#[derive(Debug, Clone)]
11pub enum Preference<'a, F>
12where
13    F: Float + Send + Sync,
14{
15    /// Use the median pairwise similarity as the preference
16    Median,
17    /// Use a list of preferences, one per input
18    List(&'a Array1<F>),
19    /// Use a single value as the preference for all inputs
20    Value(F),
21}
22
23pub(crate) fn place_preference<F>(s: &mut Array2<F>, p: Preference<F>)
24where
25    F: Float + Send + Sync,
26{
27    let s_dim = s.dim();
28    let preference = match p {
29        Preference::Value(pref) => pref,
30        Preference::Median => median(s),
31        Preference::List(l) => {
32            assert_eq!(
33                s_dim.0,
34                l.len(),
35                "Preference list length does not match input length!"
36            );
37            Zip::from(l)
38                .and(s.diag_mut())
39                .par_for_each(|pref, s_pos| *s_pos = *pref);
40            return;
41        }
42    };
43    s.diag_mut().par_map_inplace(|v| *v = preference);
44}
45
46/// Computed simply - collect values into vector, sort, and return value at len() / 2
47fn median<F>(x: &Array2<F>) -> F
48where
49    F: Float + Send + Sync,
50{
51    let mut sorted_values = Vec::new();
52    let x_dim_0 = x.dim().0 as usize;
53    for i in 0..x_dim_0 {
54        for j in (i + 1)..x_dim_0 {
55            sorted_values.push(x[[i, j]]);
56        }
57    }
58    sorted_values.sort_by(|a, b| a.partial_cmp(b).unwrap());
59    sorted_values[sorted_values.len() / 2]
60}
61
62#[cfg(test)]
63mod test {
64    use ndarray::{arr1, arr2, Array2};
65    use rayon::ThreadPool;
66
67    use crate::preference::{median, place_preference};
68    use crate::Preference::List;
69
70    fn pool(t: usize) -> ThreadPool {
71        rayon::ThreadPoolBuilder::new()
72            .num_threads(t)
73            .build()
74            .unwrap()
75    }
76
77    fn test_data() -> Array2<f32> {
78        arr2(&[
79            [0., -5., -6., -12., -17.],
80            [-5., 0., -17., -17., -22.],
81            [-6., -17., 0., -18., -21.],
82            [-12., -17., -18., 0., -3.],
83            [-17., -22., -21., -3., 0.],
84        ])
85    }
86
87    #[test]
88    fn valid_median() {
89        assert_eq!(-17., median(&test_data()));
90    }
91
92    #[test]
93    fn provided_preference_list() {
94        pool(2).scope(move |_| {
95            let mut sim = test_data();
96            let pref_list = arr1(&[-1., -2., -3., -4., -5.]);
97            place_preference(&mut sim, List(&pref_list));
98        });
99    }
100
101    #[test]
102    #[should_panic]
103    fn invalid_preference_list() {
104        pool(2).scope(move |_| {
105            let mut sim = test_data();
106            let pref_list = arr1(&[-1., -2., -3.]);
107            place_preference(&mut sim, List(&pref_list));
108        });
109    }
110}