use std::collections::HashMap;
use ndarray::Array2;
use num_traits::Float;
use crate::algorithm::APAlgorithm;
use crate::preference::place_preference;
use crate::similarity::{calculate_similarity, Similarity};
use crate::Preference;
pub type Idx = usize;
pub type Cluster = Vec<Idx>;
pub type ClusterMap = HashMap<Idx, Cluster>;
pub type ClusterResults = (bool, ClusterMap);
#[derive(Debug, Clone)]
pub struct AffinityPropagation<F> {
damping: F,
threads: usize,
convergence_iter: usize,
max_iterations: usize,
}
impl<F> Default for AffinityPropagation<F>
where
F: Float + Send + Sync,
{
fn default() -> Self {
Self {
damping: F::from(0.5).unwrap(),
threads: 4,
convergence_iter: 10,
max_iterations: 100,
}
}
}
impl<F> AffinityPropagation<F>
where
F: Float + Send + Sync,
{
pub fn new(damping: F, threads: usize, convergence_iter: usize, max_iterations: usize) -> Self {
assert!(
damping > F::from(0.).unwrap() && damping < F::from(1.).unwrap(),
"invalid damping value provided"
);
Self {
damping,
threads,
max_iterations,
convergence_iter,
}
}
pub fn predict<S>(&self, x: &Array2<F>, s: S, p: Preference<F>) -> ClusterResults
where
S: Similarity<F>,
{
let mut s = calculate_similarity(x, s);
assert!(s.is_square(), "similarity dim must be NxN");
place_preference(&mut s, p);
self.predict_parallel(s)
}
pub fn predict_precalculated(&self, s: Array2<F>, p: Preference<F>) -> ClusterResults {
let mut s = s;
assert!(s.is_square(), "similarity dim must be NxN");
place_preference(&mut s, p);
self.predict_parallel(s)
}
fn predict_parallel(&self, s: Array2<F>) -> ClusterResults {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(self.threads)
.build()
.unwrap();
pool.scope(move |_| {
let mut calculation = APAlgorithm::new(self.damping, s);
calculation.predict(self.convergence_iter, self.max_iterations)
})
}
}
#[cfg(test)]
mod test {
use ndarray::{arr2, Array2};
use crate::similarity::calculate_similarity;
use crate::Preference::Value;
use crate::{AffinityPropagation, NegCosine, NegEuclidean, Preference};
fn test_data() -> Array2<f32> {
arr2(&[
[3., 4., 3., 2., 1.],
[4., 3., 5., 1., 1.],
[3., 5., 3., 3., 3.],
[2., 1., 3., 3., 2.],
[1., 1., 3., 2., 3.],
])
}
#[test]
fn simple() {
let x: Array2<f32> = arr2(&[[1., 1., 1.], [2., 2., 2.], [3., 3., 3.]]);
let ap = AffinityPropagation::default();
let (converged, results) = ap.predict(&x, NegEuclidean::default(), Preference::Value(-10.));
assert!(converged);
assert_eq!(1, results.len());
assert!(results.contains_key(&1));
}
#[test]
fn zero() {
let x: Array2<f32> = arr2(&[[]]);
let ap = AffinityPropagation::default();
let (converged, _) = ap.predict(&x, NegCosine::default(), Value(-10.));
assert!(!converged);
}
#[test]
fn one() {
let x: Array2<f32> = arr2(&[[0., 1., 0.]]);
let ap = AffinityPropagation::default();
let (converged, _) = ap.predict(&x, NegCosine::default(), Value(-10.));
assert!(!converged);
}
#[test]
fn with_cosine() {
let x: Array2<f32> = arr2(&[[0., 1., 0.], [2., 3., 2.], [3., 2., 3.]]);
let ap = AffinityPropagation::default();
let (converged, results) = ap.predict(&x, NegCosine::default(), Value(-10.));
assert!(converged);
assert_eq!(1, results.len());
assert!(results.contains_key(&0));
}
#[test]
fn with_cosine_unconverged() {
let x: Array2<f32> = arr2(&[[1., 1., 1.], [2., 2., 2.], [3., 3., 3.]]);
let ap = AffinityPropagation::default();
let (converged, _) = ap.predict(&x, NegCosine::default(), Value(-10.));
assert!(!converged);
}
#[test]
fn with_parameters() {
let x: Array2<f32> = arr2(&[[1., 2., 1.], [2., 3., 2.], [3., 2., 3.]]);
let ap = AffinityPropagation::new(0.5, 2, 10, 100);
let (converged, results) = ap.predict(&x, NegCosine::default(), Preference::Median);
assert!(converged);
assert_eq!(2, results.len());
assert!(results.contains_key(&0) && results.contains_key(&2));
}
#[test]
fn larger() {
let x: Array2<f32> = test_data();
let ap = AffinityPropagation::default();
let (converged, results) = ap.predict(&x, NegEuclidean::default(), Value(-10.));
assert!(converged);
assert_eq!(1, results.len());
assert!(results.contains_key(&0));
}
#[test]
fn test_pre_calculated() {
let x: Array2<f32> = test_data();
let s = calculate_similarity(&x, NegEuclidean::default());
let ap = AffinityPropagation::default();
let (converged, results) = ap.predict_precalculated(s, Value(-10.));
assert!(converged);
assert_eq!(1, results.len());
assert!(results.contains_key(&0));
}
}