use crate::matrix::clustering::{Kmeans, KmeansArgs};
use crate::matrix::traits::MatOps;
use nalgebra::{DMatrix, DVector};
use rand::rngs::StdRng;
use rand::SeedableRng;
use rayon::prelude::*;
use std::borrow::Cow;
#[derive(Clone, Debug)]
pub struct AaArgs {
pub k: usize,
pub max_iter: usize,
pub fw_iters: usize,
pub tol: f32,
pub seed: u64,
pub subsample: Option<usize>,
}
impl Default for AaArgs {
fn default() -> Self {
Self {
k: 10,
max_iter: 50,
fw_iters: 30,
tol: 1e-4,
seed: 42,
subsample: None,
}
}
}
pub struct AaResult {
pub alpha: DMatrix<f32>,
pub theta: DMatrix<f32>,
pub rss: f32,
}
pub fn archetypal_analysis(z: &DMatrix<f32>, args: &AaArgs) -> AaResult {
let fit = subsample_rows(z, args.subsample, args.seed);
let alpha = fit_archetypes(&fit, args).0;
let theta = assign_theta(z, &alpha, args.fw_iters);
let rss = reconstruction_rss(z, &theta, &alpha);
AaResult { alpha, theta, rss }
}
pub fn select_archetype_k(z: &DMatrix<f32>, k_range: &[usize], args: &AaArgs) -> (usize, AaResult) {
assert!(!k_range.is_empty(), "select_archetype_k: empty k_range");
let fit = subsample_rows(z, args.subsample, args.seed);
let mut fits: Vec<(f32, DMatrix<f32>)> = k_range
.iter()
.map(|&k| {
let (alpha, rss) = fit_archetypes(&fit, &AaArgs { k, ..args.clone() });
log::info!("archetypal K-sweep: k={k} fit-RSS={rss:.4}");
(rss, alpha)
})
.collect();
let rss_fit: Vec<f32> = fits.iter().map(|(r, _)| *r).collect();
let bi = elbow_index(k_range, &rss_fit);
log::info!("archetypal K-sweep selected k={}", k_range[bi]);
let (_, alpha) = fits.swap_remove(bi);
let theta = assign_theta(z, &alpha, args.fw_iters);
let rss = reconstruction_rss(z, &theta, &alpha);
(k_range[bi], AaResult { alpha, theta, rss })
}
pub struct AnchorResult {
pub alpha: DMatrix<f32>,
pub theta: DMatrix<f32>,
pub anchors: Vec<usize>,
pub rss: f32,
}
#[derive(Clone, Copy, Debug)]
pub struct AnchorOpts {
pub fw_iters: usize,
pub min_anchor_cells: usize,
}
fn spa_anchors(rho: &DMatrix<f32>, k: usize) -> (Vec<usize>, Vec<f32>) {
let d = rho.nrows();
let kk = k.min(d);
let mut r = rho.clone(); let mut anchors = Vec::with_capacity(kk);
let mut residuals = Vec::with_capacity(kk);
for _ in 0..kk {
let (best, best_sq) = (0..d).map(|i| (i, r.row(i).norm_squared())).fold(
(0usize, -1.0f32),
|(bi, bn), (i, n)| {
if n > bn {
(i, n)
} else {
(bi, bn)
}
},
);
if best_sq <= f32::EPSILON {
break; }
anchors.push(best);
residuals.push(best_sq.sqrt());
let u = r.row(best).transpose() / best_sq.sqrt(); let ru = &r * &u; r.ger(-1.0, &ru, &u, 1.0); }
(anchors, residuals)
}
pub fn anchor_topics(
z: &DMatrix<f32>,
rho: &DMatrix<f32>,
k: usize,
opts: AnchorOpts,
) -> AnchorResult {
let (anchors, _) = spa_anchors(rho, k);
finalize_anchors(z, rho, anchors, opts)
}
pub fn select_anchor_topics(
z: &DMatrix<f32>,
rho: &DMatrix<f32>,
k_range: &[usize],
opts: AnchorOpts,
) -> (usize, AnchorResult) {
assert!(!k_range.is_empty(), "select_anchor_topics: empty k_range");
let kmax = *k_range.iter().max().unwrap();
let (anchors_full, residuals) = spa_anchors(rho, kmax);
let rss_curve: Vec<f32> = k_range
.iter()
.map(|&k| residuals.get(k).copied().unwrap_or(0.0))
.collect();
let bi = elbow_index(k_range, &rss_curve);
let k = k_range[bi].min(anchors_full.len()).max(1);
log::info!("anchor K-sweep selected k={k}");
let anchors = anchors_full[..k].to_vec();
let res = finalize_anchors(z, rho, anchors, opts);
(res.anchors.len(), res)
}
fn finalize_anchors(
z: &DMatrix<f32>,
rho: &DMatrix<f32>,
mut anchors: Vec<usize>,
opts: AnchorOpts,
) -> AnchorResult {
let mut alpha = rho.select_rows(anchors.iter());
let mut theta = assign_theta(z, &alpha, opts.fw_iters);
while opts.min_anchor_cells > 0 && anchors.len() > 2 {
let support = anchor_support(&theta);
let (weak, &n) = support
.iter()
.enumerate()
.min_by_key(|(_, &n)| n)
.expect("non-empty anchors");
if n >= opts.min_anchor_cells {
break;
}
log::info!(
"anchor guard: dropping topic {weak} (anchor row {}): {n} cells < min {}",
anchors[weak],
opts.min_anchor_cells
);
anchors.remove(weak);
alpha = rho.select_rows(anchors.iter());
theta = assign_theta(z, &alpha, opts.fw_iters);
}
let rss = reconstruction_rss(z, &theta, &alpha);
AnchorResult {
alpha,
theta,
anchors,
rss,
}
}
fn anchor_support(theta: &DMatrix<f32>) -> Vec<usize> {
let k = theta.ncols();
let mut counts = vec![0usize; k];
for i in 0..theta.nrows() {
let row = theta.row(i);
let mut best = 0usize;
let mut best_v = f32::NEG_INFINITY;
for j in 0..k {
if row[j] > best_v {
best_v = row[j];
best = j;
}
}
counts[best] += 1;
}
counts
}
pub fn topic_dictionary(rho: &DMatrix<f32>, alpha: &DMatrix<f32>) -> DMatrix<f32> {
(rho * alpha.centre_columns().transpose()).log_softmax_columns()
}
fn fit_archetypes(fit: &DMatrix<f32>, args: &AaArgs) -> (DMatrix<f32>, f32) {
let m = fit.nrows();
let k = args.k.min(m.max(1));
let mut alpha = init_archetypes(fit, k, args.seed);
let mut prev_rss = f32::INFINITY;
let mut rss = f32::INFINITY;
for it in 0..args.max_iter {
let theta = assign_theta(fit, &alpha, args.fw_iters);
let gtx = theta.tr_mul(fit); let gtg = theta.tr_mul(&theta); let g = >x - >g * α
let vertices = best_vertices(fit, &g, k);
for kk in 0..k {
let z_star = fit.row(vertices[kk]).transpose(); let d = &z_star - alpha.row(kk).transpose(); let denom = d.norm_squared() * gtg[(kk, kk)];
if denom <= f32::EPSILON {
continue;
}
let gamma = (d.dot(&g.row(kk).transpose()) / denom).clamp(0.0, 1.0);
let new_row = alpha.row(kk).transpose() + gamma * d;
alpha.row_mut(kk).copy_from(&new_row.transpose());
}
rss = reconstruction_rss(fit, &theta, &alpha);
let rel = (prev_rss - rss).abs() / prev_rss.max(f32::EPSILON);
log::debug!("archetypal fit k={k} iter={it} rss={rss:.4} rel={rel:.2e}");
if rel < args.tol {
break;
}
prev_rss = rss;
}
(alpha, rss)
}
fn init_archetypes(z: &DMatrix<f32>, k: usize, seed: u64) -> DMatrix<f32> {
let (m, h) = (z.nrows(), z.ncols());
let membership = z.kmeans_rows(KmeansArgs {
num_clusters: k,
max_iter: 100,
});
let mut alpha = DMatrix::<f32>::zeros(k, h);
let mut counts = vec![0usize; k];
for (i, &c) in membership.iter().enumerate() {
let c = c.min(k - 1);
let sum = alpha.row(c).transpose() + z.row(i).transpose();
alpha.row_mut(c).copy_from(&sum.transpose());
counts[c] += 1;
}
let mut rng = StdRng::seed_from_u64(seed);
use rand::RngExt;
for (c, &count) in counts.iter().enumerate() {
if count > 0 {
let avg = alpha.row(c) / count as f32;
alpha.row_mut(c).copy_from(&avg);
} else {
let r = rng.random_range(0..m);
alpha.row_mut(c).copy_from(&z.row(r));
}
}
alpha
}
fn assign_theta(z: &DMatrix<f32>, alpha: &DMatrix<f32>, fw_iters: usize) -> DMatrix<f32> {
let n = z.nrows();
let k = alpha.nrows();
let rows: Vec<DVector<f32>> = (0..n)
.into_par_iter()
.map(|i| simplex_lsq(alpha, &z.row(i).transpose(), fw_iters))
.collect();
let mut theta = DMatrix::<f32>::zeros(n, k);
for (i, row) in rows.into_iter().enumerate() {
theta.row_mut(i).copy_from(&row.transpose());
}
theta
}
pub fn simplex_lsq(alpha: &DMatrix<f32>, x: &DVector<f32>, fw_iters: usize) -> DVector<f32> {
let k = alpha.nrows();
let mut a = DVector::<f32>::from_element(k, 1.0 / k as f32);
for _ in 0..fw_iters {
let pred = alpha.tr_mul(&a); let r = x - &pred; let grad = alpha * &r * (-2.0); let j = argmin(&grad);
let mut d = -&a;
d[j] += 1.0; let ad = alpha.tr_mul(&d); let denom = ad.norm_squared();
if denom <= f32::EPSILON {
break;
}
let gamma = (r.dot(&ad) / denom).clamp(0.0, 1.0);
a += gamma * d;
}
a
}
fn best_vertices(z: &DMatrix<f32>, g: &DMatrix<f32>, k: usize) -> Vec<usize> {
let init = || (vec![f32::NEG_INFINITY; k], vec![0usize; k]);
let (_, idx) = (0..z.nrows())
.into_par_iter()
.fold(init, |(mut mv, mut mi), n| {
let gv = g * z.row(n).transpose(); for kk in 0..k {
if gv[kk] > mv[kk] {
mv[kk] = gv[kk];
mi[kk] = n;
}
}
(mv, mi)
})
.reduce(init, |(mut amv, mut ami), (bmv, bmi)| {
for kk in 0..k {
if bmv[kk] > amv[kk] {
amv[kk] = bmv[kk];
ami[kk] = bmi[kk];
}
}
(amv, ami)
});
idx
}
fn reconstruction_rss(z: &DMatrix<f32>, theta: &DMatrix<f32>, alpha: &DMatrix<f32>) -> f32 {
let pred = theta * alpha;
(z - pred).norm_squared()
}
fn argmin(v: &DVector<f32>) -> usize {
let mut best = 0;
let mut best_v = f32::INFINITY;
for (i, &x) in v.iter().enumerate() {
if x < best_v {
best_v = x;
best = i;
}
}
best
}
fn subsample_rows(z: &DMatrix<f32>, cap: Option<usize>, seed: u64) -> Cow<'_, DMatrix<f32>> {
let n = z.nrows();
match cap {
Some(m) if m < n => {
let mut rng = StdRng::seed_from_u64(seed);
let idx = rand::seq::index::sample(&mut rng, n, m).into_vec();
let mut out = DMatrix::<f32>::zeros(m, z.ncols());
for (r, &i) in idx.iter().enumerate() {
out.row_mut(r).copy_from(&z.row(i));
}
Cow::Owned(out)
}
_ => Cow::Borrowed(z),
}
}
fn elbow_index(ks: &[usize], rss: &[f32]) -> usize {
let n = ks.len();
if n < 3 {
return 0;
}
let kf: Vec<f32> = ks.iter().map(|&k| k as f32).collect();
let (kmin, kmax) = (kf[0], kf[n - 1]);
let (rmin, rmax) = rss
.iter()
.fold((f32::INFINITY, f32::NEG_INFINITY), |(lo, hi), &r| {
(lo.min(r), hi.max(r))
});
let kspan = (kmax - kmin).max(f32::EPSILON);
let rspan = (rmax - rmin).max(f32::EPSILON);
let xs: Vec<f32> = kf.iter().map(|&k| (k - kmin) / kspan).collect();
let ys: Vec<f32> = rss.iter().map(|&r| (r - rmin) / rspan).collect();
let (x0, y0) = (xs[0], ys[0]);
let (x1, y1) = (xs[n - 1], ys[n - 1]);
let den = ((y1 - y0).powi(2) + (x1 - x0).powi(2))
.sqrt()
.max(f32::EPSILON);
let mut best = 0;
let mut best_d = f32::NEG_INFINITY;
for i in 0..n {
let num = ((y1 - y0) * xs[i] - (x1 - x0) * ys[i] + x1 * y0 - y1 * x0).abs();
let d = num / den;
if d > best_d {
best_d = d;
best = i;
}
}
best
}
#[cfg(test)]
mod tests {
use super::*;
use rand::RngExt;
fn planted(n: usize, k: usize, h: usize, seed: u64) -> (DMatrix<f32>, DMatrix<f32>) {
let mut rng = StdRng::seed_from_u64(seed);
let mut a = DMatrix::<f32>::zeros(k, h);
for kk in 0..k {
for hh in 0..h {
a[(kk, hh)] = if hh % k == kk {
5.0
} else {
rng.random_range(-0.2..0.2)
};
}
}
let mut z = DMatrix::<f32>::zeros(n, h);
for i in 0..n {
let theta = if i < k {
let mut e = DVector::<f32>::zeros(k);
e[i] = 1.0;
e
} else {
let mut w: Vec<f32> = (0..k)
.map(|_| -rng.random_range(0.0f32..1.0).ln())
.collect();
let s: f32 = w.iter().sum();
w.iter_mut().for_each(|x| *x /= s);
DVector::from_vec(w)
};
let row = a.tr_mul(&theta); z.row_mut(i).copy_from(&row.transpose());
}
(z, a)
}
#[test]
fn recovers_planted_archetypes() {
let (k, h, n) = (4, 8, 600);
let (z, a_true) = planted(n, k, h, 1);
let res = archetypal_analysis(
&z,
&AaArgs {
k,
max_iter: 100,
fw_iters: 50,
tol: 1e-6,
seed: 7,
subsample: None,
},
);
for kt in 0..k {
let mut best = f32::INFINITY;
for kr in 0..k {
let d = (a_true.row(kt) - res.alpha.row(kr)).norm();
best = best.min(d);
}
assert!(
best < 0.75,
"true archetype {kt} unmatched (min dist {best})"
);
}
}
#[test]
fn theta_rows_are_simplex() {
let (k, h, n) = (3, 6, 300);
let (z, _) = planted(n, k, h, 2);
let res = archetypal_analysis(
&z,
&AaArgs {
k,
..Default::default()
},
);
for i in 0..n {
let s: f32 = res.theta.row(i).sum();
assert!((s - 1.0).abs() < 1e-3, "row {i} sums to {s}");
assert!(
res.theta.row(i).iter().all(|&x| x >= -1e-5),
"row {i} negative"
);
}
}
#[test]
fn sweep_picks_planted_k() {
let (k, h, n) = (4, 8, 500);
let (z, _) = planted(n, k, h, 3);
let krange: Vec<usize> = (2..=8).collect();
let (best, _) = select_archetype_k(
&z,
&krange,
&AaArgs {
fw_iters: 40,
..Default::default()
},
);
assert!(
(3..=5).contains(&best),
"elbow picked k={best}, expected ~4"
);
}
#[test]
fn spa_recovers_planted_anchors() {
let (k, h, n) = (4, 8, 500);
let (z, a_true) = planted(n, k, h, 5);
let (anchors, resid) = spa_anchors(&z, k);
let got: std::collections::BTreeSet<usize> = anchors.iter().copied().collect();
let want: std::collections::BTreeSet<usize> = (0..k).collect();
assert_eq!(
got, want,
"SPA anchors {anchors:?} != planted vertices 0..{k}"
);
assert!(
resid.windows(2).all(|w| w[0] >= w[1] - 1e-4),
"residuals not monotone: {resid:?}"
);
let res = anchor_topics(
&z,
&z,
k,
AnchorOpts {
fw_iters: 50,
min_anchor_cells: 0,
},
);
for kt in 0..k {
let best = (0..k)
.map(|kr| (a_true.row(kt) - res.alpha.row(kr)).norm())
.fold(f32::INFINITY, f32::min);
assert!(
best < 0.1,
"planted archetype {kt} unmatched (min dist {best})"
);
}
for i in 0..n {
let s: f32 = res.theta.row(i).sum();
assert!((s - 1.0).abs() < 1e-3, "θ row {i} sums to {s}");
assert!(
res.theta.row(i).iter().all(|&x| x >= -1e-5),
"θ row {i} negative"
);
}
}
#[test]
fn anchor_sweep_picks_planted_k() {
let (k, h, n) = (4, 8, 500);
let (z, _) = planted(n, k, h, 6);
let krange: Vec<usize> = (2..=8).collect();
let (best, _) = select_anchor_topics(
&z,
&z,
&krange,
AnchorOpts {
fw_iters: 40,
min_anchor_cells: 0,
},
);
assert!(
(3..=5).contains(&best),
"anchor elbow picked k={best}, expected ~4"
);
}
#[test]
fn guard_drops_singleton_outlier_anchor() {
let (k, h, n) = (3, 8, 300);
let (z3, _) = planted(n, k, h, 9);
let mut z = z3.insert_row(n, 0.0); z[(n, 0)] = 50.0;
let res = anchor_topics(
&z,
&z,
k + 1, AnchorOpts {
fw_iters: 50,
min_anchor_cells: 10,
},
);
assert_eq!(res.anchors.len(), k, "guard should drop the outlier anchor");
assert!(
!res.anchors.contains(&n),
"outlier row {n} survived as an anchor: {:?}",
res.anchors
);
}
}