use scirs2_core::ndarray::{Array2, ArrayView2};
use sklears_core::types::Float;
use std::collections::HashMap;
fn pairwise_distances(x: &ArrayView2<Float>) -> Array2<Float> {
let n = x.nrows();
let mut distances = Array2::zeros((n, n));
for i in 0..n {
for j in i + 1..n {
let dist = x
.row(i)
.iter()
.zip(x.row(j).iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<Float>()
.sqrt();
distances[(i, j)] = dist;
distances[(j, i)] = dist;
}
}
distances
}
fn find_k_nearest_neighbors(distances: &Array2<Float>, k: usize) -> Vec<Vec<usize>> {
let n = distances.nrows();
let mut neighbors = Vec::with_capacity(n);
for i in 0..n {
let mut indexed_distances: Vec<(usize, Float)> = (0..n)
.filter(|&j| i != j)
.map(|j| (j, distances[(i, j)]))
.collect();
indexed_distances.sort_by(|a, b| a.1.partial_cmp(&b.1).expect("operation should succeed"));
let k_neighbors: Vec<usize> = indexed_distances
.into_iter()
.take(k.min(n - 1))
.map(|(idx, _)| idx)
.collect();
neighbors.push(k_neighbors);
}
neighbors
}
pub fn trustworthiness(
x_original: &ArrayView2<Float>,
x_embedded: &ArrayView2<Float>,
k: usize,
) -> Float {
if x_original.nrows() != x_embedded.nrows() {
panic!("Original and embedded data must have same number of samples");
}
let n = x_original.nrows();
if k >= n {
return 1.0; }
let orig_distances = pairwise_distances(x_original);
let emb_distances = pairwise_distances(x_embedded);
let orig_neighbors = find_k_nearest_neighbors(&orig_distances, k);
let emb_neighbors = find_k_nearest_neighbors(&emb_distances, k);
let mut orig_ranks: Vec<HashMap<usize, usize>> = Vec::with_capacity(n);
for i in 0..n {
let mut rank_map = HashMap::new();
let mut indexed_distances: Vec<(usize, Float)> = (0..n)
.filter(|&j| i != j)
.map(|j| (j, orig_distances[(i, j)]))
.collect();
indexed_distances.sort_by(|a, b| a.1.partial_cmp(&b.1).expect("operation should succeed"));
for (rank, (idx, _)) in indexed_distances.iter().enumerate() {
rank_map.insert(*idx, rank + 1);
}
orig_ranks.push(rank_map);
}
let mut trustworthiness_sum = 0.0;
for i in 0..n {
let orig_k_neighbors: std::collections::HashSet<usize> =
orig_neighbors[i].iter().cloned().collect();
for &j in &emb_neighbors[i] {
if !orig_k_neighbors.contains(&j) {
let rank_j = orig_ranks[i].get(&j).unwrap_or(&n);
trustworthiness_sum += (*rank_j as Float - k as Float).max(0.0);
}
}
}
let max_sum = (n as Float) * k as Float * (2.0 * n as Float - 3.0 * k as Float - 1.0) / 2.0;
1.0 - (2.0 / max_sum) * trustworthiness_sum
}
pub fn continuity(
x_original: &ArrayView2<Float>,
x_embedded: &ArrayView2<Float>,
k: usize,
) -> Float {
if x_original.nrows() != x_embedded.nrows() {
panic!("Original and embedded data must have same number of samples");
}
let n = x_original.nrows();
if k >= n {
return 1.0; }
let orig_distances = pairwise_distances(x_original);
let emb_distances = pairwise_distances(x_embedded);
let orig_neighbors = find_k_nearest_neighbors(&orig_distances, k);
let emb_neighbors = find_k_nearest_neighbors(&emb_distances, k);
let mut emb_ranks: Vec<HashMap<usize, usize>> = Vec::with_capacity(n);
for i in 0..n {
let mut rank_map = HashMap::new();
let mut indexed_distances: Vec<(usize, Float)> = (0..n)
.filter(|&j| i != j)
.map(|j| (j, emb_distances[(i, j)]))
.collect();
indexed_distances.sort_by(|a, b| a.1.partial_cmp(&b.1).expect("operation should succeed"));
for (rank, (idx, _)) in indexed_distances.iter().enumerate() {
rank_map.insert(*idx, rank + 1);
}
emb_ranks.push(rank_map);
}
let mut continuity_sum = 0.0;
for i in 0..n {
let emb_k_neighbors: std::collections::HashSet<usize> =
emb_neighbors[i].iter().cloned().collect();
for &j in &orig_neighbors[i] {
if !emb_k_neighbors.contains(&j) {
let rank_j = emb_ranks[i].get(&j).unwrap_or(&n);
continuity_sum += (*rank_j as Float - k as Float).max(0.0);
}
}
}
let max_sum = (n as Float) * k as Float * (2.0 * n as Float - 3.0 * k as Float - 1.0) / 2.0;
1.0 - (2.0 / max_sum) * continuity_sum
}
pub fn neighborhood_hit_rate(
x_original: &ArrayView2<Float>,
x_embedded: &ArrayView2<Float>,
k: usize,
) -> Float {
if x_original.nrows() != x_embedded.nrows() {
panic!("Original and embedded data must have same number of samples");
}
let n = x_original.nrows();
if k >= n {
return 1.0; }
let orig_distances = pairwise_distances(x_original);
let emb_distances = pairwise_distances(x_embedded);
let orig_neighbors = find_k_nearest_neighbors(&orig_distances, k);
let emb_neighbors = find_k_nearest_neighbors(&emb_distances, k);
let mut total_hits = 0;
let total_possible = n * k;
for i in 0..n {
let orig_set: std::collections::HashSet<usize> =
orig_neighbors[i].iter().cloned().collect();
let emb_set: std::collections::HashSet<usize> = emb_neighbors[i].iter().cloned().collect();
let intersection_size = orig_set.intersection(&emb_set).count();
total_hits += intersection_size;
}
total_hits as Float / total_possible as Float
}
pub fn local_continuity_meta_criterion(
x_original: &ArrayView2<Float>,
x_embedded: &ArrayView2<Float>,
k: usize,
) -> Float {
let trust = trustworthiness(x_original, x_embedded, k);
let cont = continuity(x_original, x_embedded, k);
if trust + cont == 0.0 {
0.0
} else {
2.0 * trust * cont / (trust + cont)
}
}
pub fn normalized_stress(x_original: &ArrayView2<Float>, x_embedded: &ArrayView2<Float>) -> Float {
if x_original.nrows() != x_embedded.nrows() {
panic!("Original and embedded data must have same number of samples");
}
let orig_distances = pairwise_distances(x_original);
let emb_distances = pairwise_distances(x_embedded);
let n = x_original.nrows();
let mut numerator = 0.0;
let mut denominator = 0.0;
for i in 0..n {
for j in i + 1..n {
let d_orig = orig_distances[(i, j)];
let d_emb = emb_distances[(i, j)];
numerator += (d_orig - d_emb).powi(2);
denominator += d_orig.powi(2);
}
}
if denominator == 0.0 {
0.0
} else {
numerator / denominator
}
}
pub fn mean_relative_rank_error(
x_original: &ArrayView2<Float>,
x_embedded: &ArrayView2<Float>,
) -> Float {
if x_original.nrows() != x_embedded.nrows() {
panic!("Original and embedded data must have same number of samples");
}
let orig_distances = pairwise_distances(x_original);
let emb_distances = pairwise_distances(x_embedded);
let n = x_original.nrows();
let mut total_error = 0.0;
let mut total_pairs = 0;
for i in 0..n {
let mut orig_pairs: Vec<(usize, Float)> = (0..n)
.filter(|&j| i != j)
.map(|j| (j, orig_distances[(i, j)]))
.collect();
orig_pairs.sort_by(|a, b| a.1.partial_cmp(&b.1).expect("operation should succeed"));
let mut orig_rank_map = HashMap::new();
for (rank, (idx, _)) in orig_pairs.iter().enumerate() {
orig_rank_map.insert(*idx, rank);
}
let mut emb_pairs: Vec<(usize, Float)> = (0..n)
.filter(|&j| i != j)
.map(|j| (j, emb_distances[(i, j)]))
.collect();
emb_pairs.sort_by(|a, b| a.1.partial_cmp(&b.1).expect("operation should succeed"));
let mut emb_rank_map = HashMap::new();
for (rank, (idx, _)) in emb_pairs.iter().enumerate() {
emb_rank_map.insert(*idx, rank);
}
for j in 0..n {
if i != j {
let orig_rank = *orig_rank_map.get(&j).expect("operation should succeed") as Float;
let emb_rank = *emb_rank_map.get(&j).expect("operation should succeed") as Float;
let n_minus_1 = (n - 1) as Float;
total_error += (orig_rank - emb_rank).abs() / n_minus_1;
total_pairs += 1;
}
}
}
total_error / total_pairs as Float
}
#[derive(Debug, Clone)]
pub struct QualityReport {
pub trustworthiness: Float,
pub continuity: Float,
pub neighborhood_hit_rate: Float,
pub local_continuity_meta_criterion: Float,
pub normalized_stress: Float,
pub mean_relative_rank_error: Float,
pub k_neighbors: usize,
}
impl std::fmt::Display for QualityReport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
writeln!(f, "Manifold Embedding Quality Report")?;
writeln!(f, "=================================")?;
writeln!(f, "Local Metrics (k={})", self.k_neighbors)?;
writeln!(
f,
" Trustworthiness: {:.4}",
self.trustworthiness
)?;
writeln!(f, " Continuity: {:.4}", self.continuity)?;
writeln!(
f,
" Neighborhood Hit Rate: {:.4}",
self.neighborhood_hit_rate
)?;
writeln!(
f,
" LCMC (Harmonic Mean): {:.4}",
self.local_continuity_meta_criterion
)?;
writeln!(f, "Global Metrics")?;
writeln!(
f,
" Normalized Stress: {:.4}",
self.normalized_stress
)?;
writeln!(
f,
" Mean Relative Rank Error: {:.4}",
self.mean_relative_rank_error
)?;
Ok(())
}
}
pub fn quality_report(
x_original: &ArrayView2<Float>,
x_embedded: &ArrayView2<Float>,
k: Option<usize>,
) -> QualityReport {
let n = x_original.nrows();
let k_neighbors = k.unwrap_or_else(|| std::cmp::min(10, n - 1).max(1));
QualityReport {
trustworthiness: trustworthiness(x_original, x_embedded, k_neighbors),
continuity: continuity(x_original, x_embedded, k_neighbors),
neighborhood_hit_rate: neighborhood_hit_rate(x_original, x_embedded, k_neighbors),
local_continuity_meta_criterion: local_continuity_meta_criterion(
x_original,
x_embedded,
k_neighbors,
),
normalized_stress: normalized_stress(x_original, x_embedded),
mean_relative_rank_error: mean_relative_rank_error(x_original, x_embedded),
k_neighbors,
}
}