use std::cmp::Ordering;
use std::collections::hash_map;
use std::collections::HashMap;
use crate::set::HpoSet;
use crate::utils::Combinations;
pub mod cluster;
use cluster::Cluster;
use cluster::ClusterVec;
#[derive(Debug, Default)]
struct DistanceMatrix(HashMap<(usize, usize), f32>);
impl DistanceMatrix {
fn iter(&'_ self) -> hash_map::Iter<'_, (usize, usize), f32> {
self.0.iter()
}
fn insert(&mut self, k: (usize, usize), v: f32) -> Option<f32> {
self.0.insert(k, v)
}
fn is_empty(&self) -> bool {
self.0.is_empty()
}
fn retain<F>(&mut self, f: F)
where
F: FnMut(&(usize, usize), &mut f32) -> bool,
{
self.0.retain(f);
}
fn get(&self, k: &(usize, usize)) -> Option<&f32> {
self.0.get(k)
}
}
pub struct Linkage<'a> {
sets: Vec<Option<HpoSet<'a>>>,
distance_matrix: DistanceMatrix,
initial_len: usize,
clusters: ClusterVec,
}
impl<'a> Linkage<'a> {
pub fn union<T, F>(sets: T, distance: F) -> Self
where
T: IntoIterator<Item = HpoSet<'a>>,
F: Fn(Combinations<HpoSet<'_>>) -> Vec<f32>,
{
let mut s = Self::new(sets, &distance);
s.cluster_set_unions(&distance);
s
}
pub fn single<T, F>(sets: T, distance: F) -> Self
where
T: IntoIterator<Item = HpoSet<'a>>,
F: Fn(Combinations<HpoSet<'_>>) -> Vec<f32>,
{
fn f32_min(v1: Option<&f32>, v2: Option<&f32>) -> f32 {
if v1.expect("v1 must be `Some`") < v2.expect("v2 must be `Some`") {
*v1.expect("v1 must be `Some`")
} else {
*v2.expect("v2 must be `Some`")
}
}
let mut linkage = Self::new(sets, &distance);
linkage.arithmetic_cluster(f32_min);
linkage
}
pub fn complete<T, F>(sets: T, distance: F) -> Self
where
T: IntoIterator<Item = HpoSet<'a>>,
F: Fn(Combinations<HpoSet<'_>>) -> Vec<f32>,
{
fn f32_max(v1: Option<&f32>, v2: Option<&f32>) -> f32 {
if v1.expect("v1 must be `Some`") > v2.expect("v2 must be `Some`") {
*v1.expect("v1 must be `Some`")
} else {
*v2.expect("v2 must be `Some`")
}
}
let mut linkage = Self::new(sets, &distance);
linkage.arithmetic_cluster(f32_max);
linkage
}
pub fn average<T, F>(sets: T, distance: F) -> Self
where
T: IntoIterator<Item = HpoSet<'a>>,
F: Fn(Combinations<HpoSet<'_>>) -> Vec<f32>,
{
fn mean(v1: Option<&f32>, v2: Option<&f32>) -> f32 {
(v1.expect("v1 must be `Some`") + v2.expect("v2 must be `Some`")) / 2.0
}
let mut linkage = Self::new(sets, &distance);
linkage.arithmetic_cluster(mean);
linkage
}
pub fn cluster(&'_ self) -> cluster::Iter<'_> {
self.clusters.iter()
}
pub fn into_cluster(self) -> cluster::IntoIter {
self.clusters.into_iter()
}
pub fn indicies(&self) -> Vec<usize> {
let mut res = Vec::with_capacity(self.initial_len);
for cluster in &self.clusters {
if cluster.lhs() < self.initial_len {
res.push(cluster.lhs());
}
if cluster.rhs() < self.initial_len {
res.push(cluster.rhs());
}
}
res
}
fn new<T, F>(sets: T, distance: F) -> Self
where
T: IntoIterator<Item = HpoSet<'a>>,
F: Fn(Combinations<HpoSet<'_>>) -> Vec<f32>,
{
let sets: Vec<Option<HpoSet<'a>>> = sets.into_iter().map(Some).collect();
let len = sets.len();
let mut s = Self {
sets,
distance_matrix: DistanceMatrix::default(),
initial_len: len,
clusters: ClusterVec::with_capacity(len),
};
s.calculate_initial_distances(distance);
s
}
fn calculate_initial_distances<F: Fn(Combinations<HpoSet<'a>>) -> Vec<f32>>(
&mut self,
func: F,
) {
let similarities = func(Combinations::new(&self.sets));
let index: Vec<Option<usize>> = (0..self.sets.len()).map(Some).collect();
for ((idx1, idx2), sim) in Combinations::new(&index).zip(similarities.into_iter()) {
self.distance_matrix.insert((*idx1, *idx2), sim);
}
}
fn closest_clusters(&self) -> ((usize, usize), f32) {
self.distance_matrix
.iter()
.reduce(|max, elmt| if elmt.1 < max.1 { elmt } else { max })
.map(|elmt| (*elmt.0, *elmt.1))
.expect("distance matrix is not empty")
}
fn new_cluster(&mut self, key: (usize, usize), dist: f32) {
self.clusters.push(Cluster::new(
key.0,
key.1,
dist,
self.size_of_cluster(key.0, key.1),
));
}
fn cluster_set_unions<F>(&mut self, func: F)
where
F: Fn(Combinations<HpoSet<'a>>) -> Vec<f32>,
{
loop {
if self.distance_matrix.is_empty() {
return;
}
let (key, dist) = self.closest_clusters();
self.new_cluster(key, dist);
let mut newset = self.sets[key.0]
.take()
.expect("set is part of distance matrix and must exist");
let set2 = self.sets[key.1]
.take()
.expect("set is part of distance matrix and must exist");
newset.extend(&set2);
self.sets.push(Some(newset));
self.distance_matrix.retain(|(idx1, idx2), _| {
idx1 != &key.0 && idx1 != &key.1 && idx2 != &key.0 && idx2 != &key.1
});
let mut new_combinations = Combinations::new(&self.sets);
new_combinations.set_to_last();
let mut distances = func(new_combinations).into_iter();
let last_index = self.sets.len() - 1;
for (idx, set) in self.sets[..last_index].iter().enumerate() {
if set.is_some() {
self.distance_matrix.insert(
(idx, last_index),
distances.next().expect("distance score must be present"),
);
}
}
}
}
fn arithmetic_cluster<F>(&mut self, func: F)
where
F: Fn(Option<&f32>, Option<&f32>) -> f32,
{
loop {
if self.distance_matrix.is_empty() {
return;
}
let (key, dist) = self.closest_clusters();
self.new_cluster(key, dist);
let x = self.sets[key.0].take();
self.sets[key.1].take();
let new_idx = self.sets.len();
for (idx, set) in self.sets.iter().enumerate() {
if idx == key.0 || idx == key.1 {
continue;
}
if set.is_some() {
let distance = match (idx.cmp(&key.0), idx.cmp(&key.1)) {
(Ordering::Less, Ordering::Less) => func(
self.distance_matrix.get(&(idx, key.0)),
self.distance_matrix.get(&(idx, key.1)),
),
(Ordering::Less, Ordering::Greater) => func(
self.distance_matrix.get(&(idx, key.0)),
self.distance_matrix.get(&(key.1, idx)),
),
(Ordering::Greater, Ordering::Less) => func(
self.distance_matrix.get(&(key.0, idx)),
self.distance_matrix.get(&(idx, key.1)),
),
(Ordering::Greater, Ordering::Greater) => func(
self.distance_matrix.get(&(key.0, idx)),
self.distance_matrix.get(&(key.1, idx)),
),
(Ordering::Equal, _) | (_, Ordering::Equal) => {
unreachable!("Cannot be reached")
}
};
self.distance_matrix.insert((idx, new_idx), distance);
}
}
self.distance_matrix.retain(|(idx1, idx2), _| {
idx1 != &key.0 && idx1 != &key.1 && idx2 != &key.0 && idx2 != &key.1
});
self.sets.push(x);
}
}
fn size_of_cluster(&self, idx1: usize, idx2: usize) -> usize {
(if idx1 < self.initial_len {
1
} else {
self.clusters
.get(idx1 - self.initial_len)
.expect("idx is guaranteed to be in cluster")
.len()
}) + (if idx2 < self.initial_len {
1
} else {
self.clusters
.get(idx2 - self.initial_len)
.expect("idx is guaranteed to be in cluster")
.len()
})
}
pub fn iter(&'_ self) -> cluster::Iter<'_> {
self.cluster()
}
}
impl<'a> IntoIterator for &'a Linkage<'a> {
type Item = &'a Cluster;
type IntoIter = cluster::Iter<'a>;
fn into_iter(self) -> Self::IntoIter {
self.iter()
}
}
impl IntoIterator for Linkage<'_> {
type Item = Cluster;
type IntoIter = cluster::IntoIter;
fn into_iter(self) -> Self::IntoIter {
self.into_cluster()
}
}