use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
pub fn row_trust_scores(
assignments: ArrayView2<f64>,
atom_trust: ArrayView1<f64>,
) -> Result<(Array1<f64>, Array2<f64>), String> {
let (n_rows, n_atoms) = assignments.dim();
if n_atoms != atom_trust.len() {
return Err(format!(
"trust score assignments/atom_trust mismatch: {} assignment columns vs {} trust entries",
n_atoms,
atom_trust.len()
));
}
let mut per_atom = Array2::<f64>::zeros((n_rows, n_atoms));
let mut row = Array1::<f64>::zeros(n_rows);
for i in 0..n_rows {
let mut denom = 0.0_f64;
for k in 0..n_atoms {
let w = assignments[[i, k]].max(0.0);
denom += w;
}
if denom > 0.0 {
let mut r = 0.0_f64;
for k in 0..n_atoms {
let w = assignments[[i, k]].max(0.0);
let credit = (w / denom) * atom_trust[k];
per_atom[[i, k]] = credit;
r += credit;
}
row[i] = r;
}
}
Ok((row, per_atom))
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn routing_weighted_average_of_atom_trust() {
let assignments = array![[3.0, 1.0], [0.0, 2.0]];
let atom_trust = array![0.2, 0.8];
let (row, per_atom) = row_trust_scores(assignments.view(), atom_trust.view()).unwrap();
assert!((per_atom[[0, 0]] - 0.15).abs() < 1e-12);
assert!((per_atom[[0, 1]] - 0.20).abs() < 1e-12);
assert!((row[0] - 0.35).abs() < 1e-12);
assert!((per_atom[[1, 0]]).abs() < 1e-12);
assert!((per_atom[[1, 1]] - 0.8).abs() < 1e-12);
assert!((row[1] - 0.8).abs() < 1e-12);
}
#[test]
fn negative_assignments_are_clipped_out() {
let assignments = array![[-5.0, 1.0]];
let atom_trust = array![0.3, 0.9];
let (row, per_atom) = row_trust_scores(assignments.view(), atom_trust.view()).unwrap();
assert!((per_atom[[0, 0]]).abs() < 1e-12);
assert!((per_atom[[0, 1]] - 0.9).abs() < 1e-12);
assert!((row[0] - 0.9).abs() < 1e-12);
}
#[test]
fn dead_row_yields_zero_not_nan() {
let assignments = array![[0.0, -1.0]];
let atom_trust = array![0.3, 0.9];
let (row, per_atom) = row_trust_scores(assignments.view(), atom_trust.view()).unwrap();
assert_eq!(row[0], 0.0);
assert_eq!(per_atom[[0, 0]], 0.0);
assert_eq!(per_atom[[0, 1]], 0.0);
}
#[test]
fn column_mismatch_errors() {
let assignments = array![[1.0, 2.0, 3.0]];
let atom_trust = array![0.3, 0.9];
assert!(row_trust_scores(assignments.view(), atom_trust.view()).is_err());
}
}