use std::collections::{HashMap, HashSet};
use ndarray::{Array1, Array2, ArrayView, Axis, Dim, Zip};
use num_traits::Float;
use crate::{Cluster, ClusterMap, ClusterResults, Idx};
pub(crate) struct APAlgorithm<F> {
similarity: Array2<F>,
responsibility: Array2<F>,
availability: Array2<F>,
tmp: Array2<F>,
damping: F,
inv_damping: F,
neg_inf: F,
zero: F,
idx: Array1<F>,
}
impl<F> APAlgorithm<F>
where
F: Float + Send + Sync,
{
pub(crate) fn new(damping: F, s: Array2<F>) -> Self {
let s_dim = s.dim();
let zero = F::from(0.).unwrap();
Self {
similarity: s,
responsibility: Array2::zeros(s_dim),
availability: Array2::zeros(s_dim),
tmp: Array2::zeros(s_dim),
damping,
inv_damping: F::from(1.).unwrap() - damping,
neg_inf: F::from(-1.).unwrap() * F::infinity(),
zero,
idx: Self::generate_idx(zero, s_dim.0),
}
}
pub(crate) fn predict(
&mut self,
convergence_iter: usize,
max_iterations: usize,
) -> ClusterResults {
let mut has_converged = false;
for _ in 0..convergence_iter {
self.update();
}
let mut final_exemplars = self.generate_exemplars();
for _ in convergence_iter..max_iterations {
self.update();
let sol_map = self.generate_exemplars();
if !sol_map.is_empty()
&& final_exemplars.len() == sol_map.len()
&& final_exemplars.iter().all(|k| sol_map.contains(k))
{
has_converged = true;
break;
}
final_exemplars = sol_map;
}
(has_converged, self.generate_exemplar_map(final_exemplars))
}
fn update(&mut self) {
self.update_r();
self.update_a();
}
fn generate_exemplars(&self) -> HashSet<Idx> {
let values: Vec<Option<Idx>> = Vec::from_iter(
Zip::from(&self.responsibility.diag())
.and(&self.availability.diag())
.and(&self.idx)
.par_map_collect(|&r, &a, &i: &F| {
if r + a > self.zero {
return Some(i.to_usize().unwrap());
}
None
}),
);
HashSet::from_iter(values.into_iter().flatten())
}
fn generate_exemplar_map(&self, sol_map: HashSet<Idx>) -> ClusterMap {
let mut exemplar_map = HashMap::from_iter(sol_map.into_iter().map(|x| (x, Cluster::new())));
if exemplar_map.is_empty() {
return exemplar_map;
}
let max_results = Zip::from(&self.idx)
.and(self.similarity.axis_iter(Axis(1)))
.par_map_collect(|&i, col| {
let i = i.to_usize().unwrap();
if exemplar_map.contains_key(&i) {
return (i, i);
}
let mut col_data: Vec<(usize, &F)> = col.into_iter().enumerate().collect();
col_data.sort_by(|&v1, &v2| v2.1.partial_cmp(v1.1).unwrap());
for v in col_data.iter() {
if exemplar_map.contains_key(&v.0) {
return (v.0, i);
}
}
unreachable!()
});
max_results
.into_iter()
.for_each(|max_val| exemplar_map.get_mut(&max_val.0).unwrap().push(max_val.1));
exemplar_map
}
fn generate_idx(zero: F, x_dim: usize) -> Array1<F> {
Array1::range(zero, F::from(x_dim).unwrap(), F::from(1.).unwrap())
}
fn update_r(&mut self) {
Zip::from(&mut self.tmp)
.and(&self.similarity)
.and(&self.availability)
.par_for_each(|t, &s, &a| *t = s + a);
let combined =
Zip::from(self.tmp.axis_iter(Axis(1))).par_map_collect(|col| Self::max_argmax(col));
let max_idx: Array1<usize> = combined.iter().map(|c| c.0).collect();
let max1: Array1<F> = combined.iter().map(|c| c.1).collect();
Zip::from(self.tmp.axis_iter_mut(Axis(1)))
.and(&max_idx)
.par_for_each(|mut t, &m| {
t[m] = self.neg_inf;
});
let max2 =
Zip::from(self.tmp.axis_iter(Axis(1))).par_map_collect(|col| Self::max_argmax(col).1);
Zip::from(self.tmp.axis_iter_mut(Axis(0)))
.and(self.similarity.axis_iter(Axis(0)))
.and(&max1)
.par_for_each(|mut t, s, &m| t.iter_mut().zip(s.iter()).for_each(|(t, s)| *t = *s - m));
Zip::from(self.tmp.axis_iter_mut(Axis(0)))
.and(self.similarity.axis_iter(Axis(0)))
.and(&max_idx)
.and(&max2)
.par_for_each(|mut t, s, &m_idx, &m2| t[m_idx] = s[m_idx] - m2);
self.tmp.par_map_inplace(|v| *v = *v * self.inv_damping);
self.responsibility
.par_map_inplace(|v| *v = *v * self.damping);
Zip::from(&mut self.responsibility)
.and(&self.tmp)
.par_for_each(|r, &t| *r = *r + t);
}
fn update_a(&mut self) {
Zip::from(&mut self.tmp)
.and(&self.responsibility)
.par_for_each(|t, &r| *t = r);
self.tmp.par_map_inplace(|v| {
if *v < self.zero {
*v = self.zero;
}
});
Zip::from(&mut self.tmp.diag_mut())
.and(self.responsibility.diag())
.par_for_each(|t, &r| *t = r);
let sum = self.tmp.sum_axis(Axis(0));
Zip::from(self.tmp.axis_iter_mut(Axis(0)))
.and(&sum)
.par_for_each(|mut t, &s| t.par_map_inplace(|t| *t = *t - s));
let tmp_diag = self.tmp.diag().to_owned();
self.tmp.par_map_inplace(|v| {
if *v < self.zero {
*v = self.zero;
}
});
Zip::from(self.tmp.diag_mut())
.and(&tmp_diag)
.par_for_each(|t, d| *t = *d);
self.tmp.par_map_inplace(|v| *v = *v * self.inv_damping);
self.availability
.par_map_inplace(|v| *v = *v * self.damping);
Zip::from(&mut self.availability)
.and(&self.tmp)
.par_for_each(|a, &t| *a = *a - t);
}
fn max_argmax(data: ArrayView<F, Dim<[usize; 1]>>) -> (usize, F) {
let mut max_pos = 0;
let mut max: F = data[0];
data.iter().enumerate().for_each(|(idx, &val)| {
if val > max {
max = val;
max_pos = idx;
}
});
(max_pos, max)
}
}
#[cfg(test)]
mod test {
use std::collections::{HashMap, HashSet};
use ndarray::{arr2, Array2};
use rayon::ThreadPool;
use crate::algorithm::APAlgorithm;
use crate::preference::place_preference;
use crate::Preference::Value;
fn pool(t: usize) -> ThreadPool {
rayon::ThreadPoolBuilder::new()
.num_threads(t)
.build()
.unwrap()
}
fn test_data() -> Array2<f32> {
arr2(&[
[0., -7., -6., -12., -17.],
[-7., 0., -17., -17., -22.],
[-6., -17., 0., -18., -21.],
[-12., -17., -18., 0., -3.],
[-17., -22., -21., -3., 0.],
])
}
#[test]
fn valid_select_exemplars() {
pool(2).scope(move |_| {
let mut sim = test_data();
place_preference(&mut sim, Value(-22.));
let mut calc: APAlgorithm<f32> = APAlgorithm::new(0., sim);
calc.update();
let exemplars = calc.generate_exemplars();
let actual: HashSet<usize> = HashSet::from([0]);
assert!(
actual.len() == exemplars.len() && actual.iter().all(|v| exemplars.contains(v))
);
});
}
#[test]
fn valid_gather_members() {
pool(2).scope(move |_| {
let mut sim = test_data();
place_preference(&mut sim, Value(-22.));
let mut calc: APAlgorithm<f32> = APAlgorithm::new(0., sim);
calc.update();
let exemplars = calc.generate_exemplar_map(calc.generate_exemplars());
let actual: HashMap<usize, Vec<usize>> = HashMap::from([(0, vec![0, 1, 2, 3, 4])]);
assert!(
actual.len() == exemplars.len()
&& actual.iter().all(|(idx, values)| {
let v: HashSet<usize> =
HashSet::from_iter(values.iter().map(|v| v.clone()));
let a: HashSet<usize> = HashSet::from_iter(
exemplars.get(idx).unwrap().iter().map(|v| v.clone()),
);
return v.len() == a.len() && v.iter().all(|p| v.contains(p));
})
);
});
}
}