use ndarray::{Array2, ArrayView2, Axis};
const AUX_DISCRETE_MAX_LEVELS: usize = 64;
const AUX_LEVEL_DEDUP_TOL: f64 = 1.0e-12;
#[derive(Debug, Clone)]
pub struct AuxRichnessMetrics {
pub aux_observed: bool,
pub n_nonfinite_aux: usize,
pub aux_dim: usize,
pub latent_dim: usize,
pub n_rows: usize,
pub constant_columns: Vec<usize>,
pub aux_is_discrete: bool,
pub n_distinct_levels: usize,
pub jacobian_rank: usize,
pub jacobian_rank_estimated: bool,
}
pub fn aux_richness_metrics(aux: ArrayView2<f64>, latents: ArrayView2<f64>) -> AuxRichnessMetrics {
let (n, aux_dim) = aux.dim();
let (n_z, latent_dim) = latents.dim();
assert_eq!(n, n_z, "aux and latents must share row count");
let mut n_nonfinite_aux: usize = 0;
for &v in aux.iter() {
if !v.is_finite() {
n_nonfinite_aux += 1;
}
}
let aux_observed = n_nonfinite_aux == 0;
let mut constant_columns: Vec<usize> = Vec::new();
if aux_observed && n >= 1 {
for j in 0..aux_dim {
let col = aux.column(j);
let mean: f64 = col.sum() / n as f64;
let mut var = 0.0_f64;
for &v in col.iter() {
let d = v - mean;
var += d * d;
}
var /= n as f64;
if var <= 1.0e-24 {
constant_columns.push(j);
}
}
}
let (aux_is_discrete, n_distinct_levels) = if aux_observed && n >= 1 {
let mut discrete = true;
for &v in aux.iter() {
if (v - v.round()).abs() > 0.0 {
discrete = false;
break;
}
}
if discrete {
for j in 0..aux_dim {
let col = aux.column(j);
let mut sorted: Vec<f64> = col.iter().copied().collect();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
sorted.dedup_by(|a, b| (*a - *b).abs() < AUX_LEVEL_DEDUP_TOL);
if sorted.len() > AUX_DISCRETE_MAX_LEVELS {
discrete = false;
break;
}
}
}
if discrete {
let mut keys: Vec<Vec<i64>> = Vec::with_capacity(n);
for i in 0..n {
let mut row = Vec::with_capacity(aux_dim);
for j in 0..aux_dim {
row.push(aux[[i, j]].round() as i64);
}
keys.push(row);
}
keys.sort();
keys.dedup();
(true, keys.len())
} else {
(false, 0)
}
} else {
(false, 0)
};
let need_rows = aux_dim.max(latent_dim) + 1;
let mut jacobian_rank_estimated = false;
let mut jacobian_rank: usize = usize::MAX;
let z_finite = latents.iter().all(|v| v.is_finite());
if aux_observed && z_finite && n >= need_rows && aux_dim >= 1 && latent_dim >= 1 {
let mut a_c = aux.to_owned();
let mut z_c = latents.to_owned();
let a_mean = a_c
.mean_axis(Axis(0))
.expect("the n >= need_rows >= 1 guard above rules out an empty axis");
let z_mean = z_c
.mean_axis(Axis(0))
.expect("the n >= need_rows >= 1 guard above rules out an empty axis");
for mut row in a_c.rows_mut() {
row -= &a_mean;
}
for mut row in z_c.rows_mut() {
row -= &z_mean;
}
let ata = a_c.t().dot(&a_c);
let atz = a_c.t().dot(&z_c);
if let Ok(b_hat) = pinv_solve(ata.view(), atz.view())
&& let Ok(rank) = matrix_rank(b_hat.view(), 1.0e-8)
{
jacobian_rank = rank;
jacobian_rank_estimated = true;
}
}
AuxRichnessMetrics {
aux_observed,
n_nonfinite_aux,
aux_dim,
latent_dim,
n_rows: n,
constant_columns,
aux_is_discrete,
n_distinct_levels,
jacobian_rank,
jacobian_rank_estimated,
}
}
fn pinv_solve(a: ArrayView2<f64>, b: ArrayView2<f64>) -> Result<Array2<f64>, String> {
let (m, n) = a.dim();
assert_eq!(m, n, "pinv_solve expects a square normal-equation matrix");
let (eigvals, eigvecs) = symmetric_eigen_lower(a)?;
let max_abs = eigvals.iter().fold(0.0_f64, |acc, &v| acc.max(v.abs()));
let tol = 1.0e-12 * max_abs.max(1.0);
let k = eigvals.len();
let mut inv_diag = vec![0.0_f64; k];
for i in 0..k {
if eigvals[i].abs() > tol {
inv_diag[i] = 1.0 / eigvals[i];
}
}
let vtb = eigvecs.t().dot(&b);
let mut dvtb = vtb.clone();
for i in 0..k {
let scale = inv_diag[i];
for j in 0..dvtb.ncols() {
dvtb[[i, j]] *= scale;
}
}
Ok(eigvecs.dot(&dvtb))
}
fn symmetric_eigen_lower(a: ArrayView2<f64>) -> Result<(Vec<f64>, Array2<f64>), String> {
let (values, vectors) = gam_linalg::faer_ndarray::FaerEigh::eigh(&a, faer::Side::Lower)
.map_err(|error| format!("identifiability eigendecomposition: {error}"))?;
Ok((values.to_vec(), vectors))
}
fn matrix_rank(m: ArrayView2<f64>, tol: f64) -> Result<usize, String> {
let gram = m.t().dot(&m);
let (eigvals, _) = symmetric_eigen_lower(gram.view())?;
let mut rank = 0usize;
for &lam in eigvals.iter() {
if lam.max(0.0).sqrt() > tol {
rank += 1;
}
}
Ok(rank)
}
#[derive(Debug, Clone)]
pub struct JacobianSparsityMetrics {
pub n_samples: usize,
pub p_features: usize,
pub latent_dim: usize,
pub mean_sparsity: f64,
pub max_abs: f64,
pub ranks: Vec<usize>,
}
pub fn jacobian_sparsity_metrics(
jacobians_flat: ArrayView2<f64>,
n_samples: usize,
zero_threshold: f64,
) -> Result<JacobianSparsityMetrics, String> {
let (np_rows, latent_dim) = jacobians_flat.dim();
assert!(np_rows % n_samples == 0, "rows not divisible by n_samples");
let p_features = np_rows / n_samples;
let mut max_abs = 0.0_f64;
for &v in jacobians_flat.iter() {
let a = v.abs();
if a > max_abs {
max_abs = a;
}
}
let cutoff = zero_threshold * max_abs;
let mut total_near_zero: usize = 0;
let total_entries = np_rows * latent_dim;
if max_abs > 0.0 {
for &v in jacobians_flat.iter() {
if v.abs() < cutoff {
total_near_zero += 1;
}
}
} else {
total_near_zero = total_entries;
}
let mean_sparsity = if total_entries > 0 {
total_near_zero as f64 / total_entries as f64
} else {
0.0
};
let mut ranks = Vec::with_capacity(n_samples);
for s in 0..n_samples {
let start = s * p_features;
let end = start + p_features;
let view = jacobians_flat.slice(ndarray::s![start..end, ..]);
ranks.push(matrix_rank(view, cutoff)?);
}
Ok(JacobianSparsityMetrics {
n_samples,
p_features,
latent_dim,
mean_sparsity,
max_abs,
ranks,
})
}
#[derive(Debug, Clone)]
pub struct AnchorConsistencyMetrics {
pub n_rows: usize,
pub n_atoms: usize,
pub n_anchors: usize,
pub anchors_per_atom: Vec<usize>,
}
fn anchor_consistency_metrics(
assignments: ArrayView2<f64>,
anchor_dominance: f64,
) -> AnchorConsistencyMetrics {
let (n, k) = assignments.dim();
let mut anchors_per_atom = vec![0_usize; k];
let mut n_anchors = 0_usize;
for i in 0..n {
let row = assignments.row(i);
let mut mass = 0.0_f64;
let mut max_val = 0.0_f64;
let mut max_j = 0_usize;
for j in 0..k {
let a = row[j].abs();
mass += a;
if a > max_val {
max_val = a;
max_j = j;
}
}
if mass > 0.0 && max_val / mass >= anchor_dominance {
n_anchors += 1;
anchors_per_atom[max_j] += 1;
}
}
AnchorConsistencyMetrics {
n_rows: n,
n_atoms: k,
n_anchors,
anchors_per_atom,
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AnchorConsistencyPreconditions {
pub enough_anchors_total: bool,
pub anchors_cover_all_atoms: bool,
}
#[derive(Debug, Clone)]
pub struct AnchorConsistencyReport {
pub metrics: AnchorConsistencyMetrics,
pub anchor_dominance: f64,
pub anchor_fraction: f64,
pub preconditions: AnchorConsistencyPreconditions,
pub violations: Vec<String>,
pub recommendations: Vec<String>,
pub uncovered_atoms: Vec<usize>,
}
impl AnchorConsistencyReport {
pub fn passes(&self) -> bool {
self.preconditions.enough_anchors_total && self.preconditions.anchors_cover_all_atoms
}
}
pub const ANCHOR_DOMINANCE_DEFAULT: f64 = f64::from_bits(0.5_f64.to_bits() + 1);
pub fn anchor_consistency_report(
assignments: ArrayView2<f64>,
anchor_dominance: Option<f64>,
) -> Result<AnchorConsistencyReport, String> {
if let Some(((row, atom), value)) = assignments
.indexed_iter()
.find(|(_, value)| !value.is_finite())
{
return Err(format!(
"assignments must be finite; entry ({row}, {atom}) is {value}"
));
}
let anchor_dominance = anchor_dominance.unwrap_or(ANCHOR_DOMINANCE_DEFAULT);
if !(anchor_dominance > 0.5 && anchor_dominance <= 1.0) {
return Err(format!(
"anchor_dominance must be in (0.5, 1]; got {anchor_dominance}"
));
}
let (_, k) = assignments.dim();
if k < 1 {
return Err("assignments must have at least one atom column".to_string());
}
let metrics = anchor_consistency_metrics(assignments, anchor_dominance);
let anchor_fraction = metrics.n_anchors as f64 / metrics.n_rows.max(1) as f64;
let mut violations = Vec::new();
let mut recommendations = Vec::new();
let mut uncovered_atoms = Vec::new();
let preconditions = if k == 1 {
AnchorConsistencyPreconditions {
enough_anchors_total: true,
anchors_cover_all_atoms: true,
}
} else {
let enough_anchors = metrics.n_anchors >= k;
if !enough_anchors {
violations.push(format!(
"Only {} anchor row(s) (dominance >= {:.2}) found in a K={}-atom \
model; need at least {}. The recovered atoms are identified only \
up to a linear transformation in atom space.",
metrics.n_anchors, anchor_dominance, k, k
));
recommendations.push(format!(
"Reduce K to <= {}, sharpen the assignment prior (e.g. lower \
temperature / stronger IBP concentration), or collect more \
anchor-like rows where a single atom dominates.",
metrics.n_anchors.max(1)
));
}
uncovered_atoms = metrics
.anchors_per_atom
.iter()
.enumerate()
.filter_map(|(j, &count)| (count == 0).then_some(j))
.collect();
let cover_ok = uncovered_atoms.is_empty();
if !cover_ok {
violations.push(format!(
"Atom(s) {:?} have zero anchor rows; they are not individually \
identifiable and may be redundant or merged with neighbours.",
uncovered_atoms
));
recommendations.push(format!(
"Prune the {} uncovered atom(s) (refit with K={}) or strengthen \
the per-atom sparsity prior so that each atom acquires a \
dominant region.",
uncovered_atoms.len(),
(k - uncovered_atoms.len()).max(1)
));
}
AnchorConsistencyPreconditions {
enough_anchors_total: enough_anchors,
anchors_cover_all_atoms: cover_ok,
}
};
Ok(AnchorConsistencyReport {
metrics,
anchor_dominance,
anchor_fraction,
preconditions,
violations,
recommendations,
uncovered_atoms,
})
}
pub fn concat_decoder_blocks(blocks: &[ArrayView2<f64>]) -> Result<Array2<f64>, String> {
if blocks.is_empty() {
return Err("concat_decoder_blocks: empty block list".into());
}
let p = blocks[0].ncols();
for (i, b) in blocks.iter().enumerate() {
if b.ncols() != p {
return Err(format!(
"concat_decoder_blocks: block {} has {} cols, expected {}",
i,
b.ncols(),
p
));
}
}
let total_k: usize = blocks.iter().map(|b| b.nrows()).sum();
let mut out = Array2::<f64>::zeros((p, total_k));
let mut col = 0_usize;
for b in blocks {
for k in 0..b.nrows() {
for row in 0..p {
out[[row, col]] = b[[k, row]];
}
col += 1;
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn aux_richness_passes_on_rich_2d_aux() {
let aux = array![
[0.0, 0.0],
[0.0, 1.0],
[1.0, 0.0],
[1.0, 1.0],
[2.0, 0.0],
[2.0, 1.0],
[0.0, 2.0],
[1.0, 2.0],
[2.0, 2.0],
];
let lat = array![
[0.10, 0.05],
[0.02, 1.01],
[1.05, 0.04],
[1.01, 1.02],
[2.03, 0.07],
[2.04, 1.01],
[0.05, 2.02],
[1.02, 2.01],
[2.01, 2.05],
];
let m = aux_richness_metrics(aux.view(), lat.view());
assert!(m.aux_observed);
assert_eq!(m.aux_dim, 2);
assert_eq!(m.latent_dim, 2);
assert!(m.constant_columns.is_empty());
assert!(m.aux_is_discrete);
assert!(m.n_distinct_levels >= 3);
assert!(m.jacobian_rank_estimated);
assert_eq!(m.jacobian_rank, 2);
}
#[test]
fn aux_richness_flags_constant_aux() {
let aux = Array2::<f64>::zeros((20, 1));
let mut lat = Array2::<f64>::zeros((20, 2));
for i in 0..20 {
lat[[i, 0]] = i as f64;
lat[[i, 1]] = (i as f64).cos();
}
let m = aux_richness_metrics(aux.view(), lat.view());
assert_eq!(m.aux_dim, 1);
assert_eq!(m.latent_dim, 2);
assert_eq!(m.constant_columns, vec![0_usize]);
}
#[test]
fn aux_richness_flags_nonfinite_aux() {
let mut aux = Array2::<f64>::zeros((10, 1));
aux[[3, 0]] = f64::NAN;
let lat = Array2::<f64>::zeros((10, 1));
let m = aux_richness_metrics(aux.view(), lat.view());
assert!(!m.aux_observed);
assert_eq!(m.n_nonfinite_aux, 1);
}
#[test]
fn jacobian_sparsity_passes_on_diagonal() {
let j = array![
[1.0_f64, 0.0, 0.0],
[0.0, 1.0, 0.0],
[0.0, 0.0, 1.0],
[0.0, 0.0, 0.0]
];
let m = jacobian_sparsity_metrics(j.view(), 1, 1.0e-3).expect("sparsity metrics");
assert_eq!(m.p_features, 4);
assert_eq!(m.latent_dim, 3);
assert!(m.mean_sparsity > 0.5);
assert_eq!(m.ranks, vec![3_usize]);
}
#[test]
fn jacobian_sparsity_dense_has_low_sparsity() {
let mut j = Array2::<f64>::zeros((4, 3));
for i in 0..4 {
for k in 0..3 {
j[[i, k]] = 1.0 + 0.1 * (i + k) as f64;
}
}
let m = jacobian_sparsity_metrics(j.view(), 1, 1.0e-3).expect("sparsity metrics");
assert!(m.mean_sparsity < 0.1);
}
#[test]
fn anchor_consistency_three_clusters() {
let mut a = Array2::<f64>::from_elem((9, 3), 0.01);
for i in 0..3 {
a[[i, 0]] = 1.0;
}
for i in 3..6 {
a[[i, 1]] = 1.0;
}
for i in 6..9 {
a[[i, 2]] = 1.0;
}
let m = anchor_consistency_metrics(a.view(), 0.95);
assert_eq!(m.n_atoms, 3);
assert_eq!(m.n_anchors, 9);
assert_eq!(m.anchors_per_atom, vec![3, 3, 3]);
}
#[test]
fn anchor_consistency_uniform_has_zero_anchors() {
let a = Array2::<f64>::from_elem((10, 4), 0.25);
let m = anchor_consistency_metrics(a.view(), 0.95);
assert_eq!(m.n_anchors, 0);
assert_eq!(m.anchors_per_atom, vec![0, 0, 0, 0]);
}
#[test]
fn anchor_consistency_report_owns_the_pass_fail_verdict() {
let a = array![[1.0_f64, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]];
let report = anchor_consistency_report(a.view(), None).unwrap();
assert_eq!(report.anchor_dominance, ANCHOR_DOMINANCE_DEFAULT);
assert!(report.passes());
assert_eq!(
report.preconditions,
AnchorConsistencyPreconditions {
enough_anchors_total: true,
anchors_cover_all_atoms: true,
}
);
assert!(report.uncovered_atoms.is_empty());
}
#[test]
fn anchor_consistency_report_derives_thresholds_from_atom_count() {
let a = Array2::<f64>::from_elem((7, 4), 0.25);
let report = anchor_consistency_report(a.view(), Some(0.95)).unwrap();
assert!(!report.passes());
assert_eq!(
report.preconditions,
AnchorConsistencyPreconditions {
enough_anchors_total: false,
anchors_cover_all_atoms: false,
}
);
assert_eq!(report.uncovered_atoms, vec![0, 1, 2, 3]);
assert_eq!(report.violations.len(), 2);
assert_eq!(report.recommendations.len(), report.violations.len());
assert!(report.violations[0].contains("need at least 4"));
}
#[test]
fn anchor_consistency_report_rejects_invalid_dominance() {
let a = Array2::<f64>::ones((2, 2));
let error = anchor_consistency_report(a.view(), Some(0.0)).unwrap_err();
assert!(error.contains("anchor_dominance must be in (0.5, 1]"));
}
#[test]
fn default_anchor_rule_is_theorem_derived_strict_majority() {
let tied = array![[0.5_f64, 0.5], [0.5, 0.5]];
let tied_report = anchor_consistency_report(tied.view(), None).unwrap();
assert_eq!(tied_report.metrics.n_anchors, 0);
let majority = array![[0.5_f64.next_up(), 0.5], [0.5, 0.5_f64.next_up()]];
let majority_report = anchor_consistency_report(majority.view(), None).unwrap();
assert_eq!(majority_report.metrics.n_anchors, 2);
assert!(majority_report.passes());
}
#[test]
fn anchor_consistency_report_rejects_non_finite_assignments() {
for value in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
let assignments = array![[1.0_f64, 0.0], [0.0, value]];
let error = anchor_consistency_report(assignments.view(), None).unwrap_err();
assert!(error.contains("assignments must be finite"));
assert!(error.contains("(1, 1)"));
}
}
#[test]
fn anchor_consistency_report_distinguishes_count_from_atom_coverage() {
let a = array![
[1.0_f64, 0.0, 0.0],
[1.0, 0.0, 0.0],
[1.0, 0.0, 0.0],
[0.0, 1.0, 0.0],
];
let report = anchor_consistency_report(a.view(), None).unwrap();
assert!(report.preconditions.enough_anchors_total);
assert!(!report.preconditions.anchors_cover_all_atoms);
assert_eq!(report.uncovered_atoms, vec![2]);
assert!(!report.passes());
assert_eq!(report.violations.len(), 1);
}
}