use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashSet};
use distances::Number;
use super::{Edge, EdgeSet, Ratios, Vertex, VertexSet};
use crate::{Cluster, Dataset, Instance};
pub type MetaMLScorer = Box<fn(Ratios) -> f64>;
struct VertexWrapper<'a, U: Number> {
pub cluster: &'a Vertex<U>,
pub score: f64,
}
impl<'a, U: Number> PartialEq for VertexWrapper<'a, U> {
fn eq(&self, other: &Self) -> bool {
self.score == other.score
}
}
impl<'a, U: Number> Eq for VertexWrapper<'a, U> {}
impl<'a, U: Number> Ord for VertexWrapper<'a, U> {
fn cmp(&self, other: &Self) -> Ordering {
match self.score.partial_cmp(&other.score).unwrap_or(Ordering::Equal) {
Ordering::Equal => {
match self.cluster.offset().cmp(&other.cluster.offset()) {
Ordering::Equal => {
self.cluster.cardinality().cmp(&other.cluster.cardinality())
}
ord => ord.reverse(),
}
}
ord => ord,
}
}
}
impl<'a, U: Number> PartialOrd for VertexWrapper<'a, U> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
fn score_clusters<'a, U: Number>(
root: &'a Vertex<U>,
scoring_function: &super::MetaMLScorer,
) -> BinaryHeap<VertexWrapper<'a, U>> {
let mut scored_clusters: BinaryHeap<VertexWrapper<'a, U>> = BinaryHeap::new();
for cluster in root.subtree() {
let score = scoring_function(cluster.ratios());
scored_clusters.push(VertexWrapper { cluster, score });
}
scored_clusters
}
pub fn select_clusters<'a, U: Number>(
root: &'a Vertex<U>,
scoring_function: &MetaMLScorer,
min_depth: usize,
) -> Result<VertexSet<'a, U>, String> {
let mut cluster_set: HashSet<&'a Vertex<U>> = HashSet::new();
let mut scored_clusters = score_clusters(root, scoring_function);
scored_clusters.retain(|item| item.cluster.depth() >= min_depth || item.cluster.is_leaf());
while !scored_clusters.is_empty() {
let Some(wrapper) = scored_clusters.pop() else {
return Err("Invalid ClusterWrapper passed to `get_clusterset`".to_string());
};
let best = wrapper.cluster;
scored_clusters.retain(|item| !item.cluster.is_ancestor_of(best) && !item.cluster.is_descendant_of(best));
cluster_set.insert(best);
}
Ok(cluster_set)
}
#[allow(clippy::implicit_hasher)]
pub fn detect_edges<'a, I: Instance, U: Number, D: Dataset<I, U>>(
clusters: &VertexSet<'a, U>,
data: &D,
) -> EdgeSet<'a, U> {
let mut edges = HashSet::new();
for (i, c1) in clusters.iter().enumerate() {
for (j, c2) in clusters.iter().enumerate().skip(i + 1) {
if i != j {
let distance = c1.distance_to_other(data, c2);
if distance <= c1.radius() + c2.radius() {
edges.insert(Edge::new(c1, c2, distance));
}
}
}
}
edges
}
#[cfg(test)]
mod tests {
use crate::{Cluster, PartitionCriteria, Tree, VecDataset};
use distances::number::Float;
use distances::Number;
use rand::SeedableRng;
use super::*;
use crate::chaoda::pretrained_models;
pub fn gen_dataset(
cardinality: usize,
dimensionality: usize,
seed: u64,
metric: fn(&Vec<f32>, &Vec<f32>) -> f32,
) -> VecDataset<Vec<f32>, f32, usize> {
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
let data = symagen::random_data::random_tabular(cardinality, dimensionality, -1., 1., &mut rng);
let name = "test".to_string();
VecDataset::new(name, data, metric, false)
}
pub fn euclidean<T: Number, F: Float>(x: &Vec<T>, y: &Vec<T>) -> F {
distances::vectors::euclidean(x, y)
}
#[test]
fn scoring() {
let data = gen_dataset(1000, 10, 42, euclidean);
let partition_criteria: PartitionCriteria<f32> = PartitionCriteria::default();
let raw_tree = Tree::new(data, Some(42))
.partition(&partition_criteria, Some(42))
.normalize_ratios();
let root = raw_tree.root();
let mut priority_queue = score_clusters(&root, &pretrained_models::get_meta_ml_scorers()[0].1);
assert_eq!(priority_queue.len(), root.subtree().len());
let mut prev_value: f64;
let mut curr_value: f64;
prev_value = priority_queue.pop().unwrap().score;
while !priority_queue.is_empty() {
curr_value = priority_queue.pop().unwrap().score;
assert!(prev_value >= curr_value);
prev_value = curr_value;
}
let cluster_set = select_clusters(&root, &pretrained_models::get_meta_ml_scorers()[0].1, 4).unwrap();
for i in &cluster_set {
for j in &cluster_set {
if i != j {
assert!(!i.is_descendant_of(j) && !i.is_ancestor_of(j));
}
}
}
for i in &root.subtree() {
let mut ancestor_of = false;
let mut descendant_of = false;
for j in &cluster_set {
if i.is_ancestor_of(j) {
ancestor_of = true;
}
if i.is_descendant_of(j) {
descendant_of = true
}
}
assert!(ancestor_of || descendant_of || cluster_set.contains(i))
}
}
}