use num_bigint::BigInt;
use num_traits::{One, Zero};
use crate::api::context::Context;
use crate::base::errors::SymplexError;
use crate::base::interval::Interval;
use crate::domains::stats::Distribution;
use crate::domains::stats::data::Q;
fn invalid(op: &'static str, reason: impl Into<String>) -> SymplexError {
SymplexError::invalid_argument(op, reason)
}
fn qu(n: usize) -> Q {
Q::from_integer(BigInt::from(n))
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LabelTable {
rows: Vec<Vec<Option<usize>>>,
n_raters: usize,
n_categories: usize,
}
impl LabelTable {
pub fn new(rows: Vec<Vec<Option<usize>>>, n_categories: usize) -> Result<Self, SymplexError> {
let op = "LabelTable::new";
if n_categories == 0 {
return Err(invalid(op, "needs at least one category"));
}
let n_raters = match rows.first() {
Some(r) => r.len(),
None => return Err(invalid(op, "a label table needs at least one item")),
};
if n_raters == 0 {
return Err(invalid(op, "a label table needs at least one rater"));
}
for (i, r) in rows.iter().enumerate() {
if r.len() != n_raters {
return Err(invalid(
op,
format!(
"item {i} has {} cells but the table has {n_raters} raters",
r.len()
),
));
}
if let Some(l) = r.iter().flatten().find(|&&l| l >= n_categories) {
return Err(invalid(
op,
format!("item {i} has label {l} but there are {n_categories} categories"),
));
}
}
Ok(Self {
rows,
n_raters,
n_categories,
})
}
pub fn from_rows(rows: &[&[Option<usize>]], n_categories: usize) -> Result<Self, SymplexError> {
Self::new(rows.iter().map(|r| r.to_vec()).collect(), n_categories)
}
pub fn complete(rows: &[&[usize]], n_categories: usize) -> Result<Self, SymplexError> {
Self::new(
rows.iter()
.map(|r| r.iter().map(|&l| Some(l)).collect())
.collect(),
n_categories,
)
}
pub fn n_items(&self) -> usize {
self.rows.len()
}
pub fn n_raters(&self) -> usize {
self.n_raters
}
pub fn n_categories(&self) -> usize {
self.n_categories
}
pub fn rows(&self) -> &[Vec<Option<usize>>] {
&self.rows
}
pub fn item(&self, i: usize) -> Option<&[Option<usize>]> {
self.rows.get(i).map(|r| r.as_slice())
}
pub fn rater(&self, j: usize) -> Option<Vec<Option<usize>>> {
(j < self.n_raters).then(|| self.rows.iter().map(|r| r[j]).collect())
}
pub fn to_counts(&self) -> Vec<Vec<Vec<usize>>> {
self.rows
.iter()
.map(|r| {
r.iter()
.map(|l| {
let mut c = vec![0usize; self.n_categories];
if let Some(l) = l {
c[*l] = 1;
}
c
})
.collect()
})
.collect()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Vote {
pub winner: Option<usize>,
pub tied: Vec<usize>,
pub counts: Vec<usize>,
}
fn vote_from_counts(counts: Vec<usize>) -> Vote {
let top = counts.iter().copied().max().unwrap_or(0);
let tied: Vec<usize> = if top == 0 {
Vec::new()
} else {
counts
.iter()
.enumerate()
.filter_map(|(c, &n)| (n == top).then_some(c))
.collect()
};
let winner = (tied.len() == 1).then(|| tied[0]);
Vote {
winner,
tied,
counts,
}
}
pub fn majority_vote(labels: &[Option<usize>]) -> Vote {
let n = labels.iter().flatten().max().map_or(0, |m| m + 1);
let mut counts = vec![0usize; n];
for &l in labels.iter().flatten() {
counts[l] += 1;
}
vote_from_counts(counts)
}
pub fn majority_votes(table: &LabelTable) -> Vec<Vote> {
table
.rows
.iter()
.map(|r| {
let mut counts = vec![0usize; table.n_categories];
for &l in r.iter().flatten() {
counts[l] += 1;
}
vote_from_counts(counts)
})
.collect()
}
pub fn plurality(labels: &[Option<usize>], threshold: &Q) -> Result<Vote, SymplexError> {
if threshold < &Q::zero() || threshold > &Q::one() {
return Err(invalid("plurality", "the threshold must lie in [0, 1]"));
}
let mut vote = majority_vote(labels);
if let Some(w) = vote.winner {
let cast: usize = vote.counts.iter().sum();
if qu(vote.counts[w]) < threshold * qu(cast) {
vote.winner = None;
}
}
Ok(vote)
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WeightedVote {
pub winner: Option<usize>,
pub tied: Vec<usize>,
pub scores: Vec<Q>,
}
pub fn weighted_vote(
labels: &[Option<usize>],
weights: &[Q],
) -> Result<WeightedVote, SymplexError> {
let op = "weighted_vote";
if labels.len() != weights.len() {
return Err(invalid(op, "one weight per rater is needed"));
}
if weights.iter().any(|w| w < &Q::zero()) {
return Err(invalid(op, "weights must be non-negative"));
}
let n = labels.iter().flatten().max().map_or(0, |m| m + 1);
let mut scores = vec![Q::zero(); n];
for (l, w) in labels.iter().zip(weights) {
if let Some(l) = l {
scores[*l] += w;
}
}
let top = scores.iter().max().cloned();
let tied: Vec<usize> = match top {
Some(t) if !t.is_zero() => scores
.iter()
.enumerate()
.filter_map(|(c, s)| (*s == t).then_some(c))
.collect(),
_ => Vec::new(),
};
let winner = (tied.len() == 1).then(|| tied[0]);
Ok(WeightedVote {
winner,
tied,
scores,
})
}
#[derive(Clone, Debug, PartialEq)]
pub enum DawidSkeneInit {
MajorityVote,
Posteriors(Vec<Vec<f64>>),
}
#[derive(Clone, Debug, PartialEq)]
pub struct DawidSkeneOpts {
pub max_iter: usize,
pub tol: f64,
pub smoothing: f64,
pub init: DawidSkeneInit,
}
impl Default for DawidSkeneOpts {
fn default() -> Self {
Self {
max_iter: 1000,
tol: 1e-10,
smoothing: 0.0,
init: DawidSkeneInit::MajorityVote,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct DawidSkene {
pub posteriors: Vec<Vec<f64>>,
pub confusion: Vec<Vec<Vec<f64>>>,
pub priors: Vec<f64>,
pub iterations: usize,
pub converged: bool,
pub log_likelihood: f64,
}
impl DawidSkene {
pub fn labels(&self) -> Vec<usize> {
self.posteriors
.iter()
.map(|row| {
row.iter()
.enumerate()
.fold((0usize, f64::NEG_INFINITY), |(bi, bv), (i, &v)| {
if v > bv { (i, v) } else { (bi, bv) }
})
.0
})
.collect()
}
}
fn ln_or_neg_inf(p: f64) -> f64 {
if p > 0.0 { p.ln() } else { f64::NEG_INFINITY }
}
fn ds_m_step(
counts: &[Vec<Vec<usize>>],
t: &[Vec<f64>],
j: usize,
smoothing: f64,
) -> (Vec<f64>, Vec<Vec<Vec<f64>>>) {
let n_items = counts.len();
let n_raters = counts[0].len();
let priors: Vec<f64> = (0..j)
.map(|c| t.iter().map(|row| row[c]).sum::<f64>() / n_items as f64)
.collect();
let confusion = (0..n_raters)
.map(|k| {
(0..j)
.map(|c| {
let mut num: Vec<f64> = (0..j)
.map(|l| {
counts
.iter()
.zip(t)
.map(|(item, row)| row[c] * item[k][l] as f64)
.sum::<f64>()
+ smoothing
})
.collect();
let den: f64 = num.iter().sum();
if den > 0.0 {
for v in &mut num {
*v /= den;
}
num
} else {
vec![1.0 / j as f64; j]
}
})
.collect()
})
.collect();
(priors, confusion)
}
fn ds_log_posteriors(item: &[Vec<usize>], priors: &[f64], confusion: &[Vec<Vec<f64>>]) -> Vec<f64> {
let mut lp: Vec<f64> = priors.iter().map(|&p| ln_or_neg_inf(p)).collect();
for (k, labels) in item.iter().enumerate() {
for (l, &c) in labels.iter().enumerate() {
if c > 0 {
for (j, v) in lp.iter_mut().enumerate() {
*v += c as f64 * ln_or_neg_inf(confusion[k][j][l]);
}
}
}
}
lp
}
fn ds_e_step(
counts: &[Vec<Vec<usize>>],
priors: &[f64],
confusion: &[Vec<Vec<f64>>],
prev: &[Vec<f64>],
) -> Vec<Vec<f64>> {
counts
.iter()
.zip(prev)
.map(|(item, old)| {
let lp = ds_log_posteriors(item, priors, confusion);
let m = lp.iter().copied().fold(f64::NEG_INFINITY, f64::max);
if !m.is_finite() {
return old.to_vec();
}
let w: Vec<f64> = lp.iter().map(|&v| (v - m).exp()).collect();
let s: f64 = w.iter().sum();
w.into_iter().map(|v| v / s).collect()
})
.collect()
}
fn ds_log_likelihood(
counts: &[Vec<Vec<usize>>],
priors: &[f64],
confusion: &[Vec<Vec<f64>>],
) -> f64 {
counts
.iter()
.map(|item| {
let lp = ds_log_posteriors(item, priors, confusion);
let m = lp.iter().copied().fold(f64::NEG_INFINITY, f64::max);
if !m.is_finite() {
return f64::NEG_INFINITY;
}
m + lp.iter().map(|&v| (v - m).exp()).sum::<f64>().ln()
})
.sum()
}
fn ds_initial(
counts: &[Vec<Vec<usize>>],
j: usize,
init: &DawidSkeneInit,
) -> Result<Vec<Vec<f64>>, SymplexError> {
let op = "dawid_skene";
match init {
DawidSkeneInit::MajorityVote => Ok(counts
.iter()
.map(|item| {
let totals: Vec<usize> = (0..j).map(|l| item.iter().map(|r| r[l]).sum()).collect();
let v = vote_from_counts(totals);
if v.tied.is_empty() {
vec![1.0 / j as f64; j]
} else {
let share = 1.0 / v.tied.len() as f64;
(0..j)
.map(|c| if v.tied.contains(&c) { share } else { 0.0 })
.collect()
}
})
.collect()),
DawidSkeneInit::Posteriors(t) => {
if t.len() != counts.len() || t.iter().any(|r| r.len() != j) {
return Err(invalid(
op,
"the initial posteriors must be items × categories",
));
}
t.iter()
.map(|row| {
if row.iter().any(|v| !v.is_finite() || *v < 0.0) {
return Err(invalid(
op,
"initial posteriors must be finite and non-negative",
));
}
let s: f64 = row.iter().sum();
if s <= 0.0 {
return Err(invalid(op, "an initial posterior row sums to zero"));
}
Ok(row.iter().map(|v| v / s).collect())
})
.collect()
}
}
}
pub fn dawid_skene_counts(
counts: &[Vec<Vec<usize>>],
n_categories: usize,
opts: &DawidSkeneOpts,
) -> Result<DawidSkene, SymplexError> {
let op = "dawid_skene";
let j = n_categories;
if j < 2 {
return Err(invalid(op, "needs at least two categories"));
}
let n_raters = counts.first().map_or(0, |i| i.len());
if counts.is_empty() || n_raters == 0 {
return Err(invalid(op, "needs at least one item and one rater"));
}
if counts
.iter()
.any(|i| i.len() != n_raters || i.iter().any(|r| r.len() != j))
{
return Err(invalid(
op,
"the counts must be items × raters × categories",
));
}
if opts.max_iter == 0 {
return Err(invalid(op, "max_iter must be positive"));
}
if !(opts.tol > 0.0 && opts.tol.is_finite()) {
return Err(invalid(op, "tol must be a positive finite number"));
}
if !(opts.smoothing >= 0.0 && opts.smoothing.is_finite()) {
return Err(invalid(
op,
"smoothing must be a non-negative finite number",
));
}
let mut t = ds_initial(counts, j, &opts.init)?;
let mut iterations = 0;
let mut converged = false;
for it in 1..=opts.max_iter {
let (priors, confusion) = ds_m_step(counts, &t, j, opts.smoothing);
let next = ds_e_step(counts, &priors, &confusion, &t);
let delta = next
.iter()
.zip(&t)
.flat_map(|(a, b)| a.iter().zip(b).map(|(x, y)| (x - y).abs()))
.fold(0.0, f64::max);
t = next;
iterations = it;
if delta < opts.tol {
converged = true;
break;
}
}
let (priors, confusion) = ds_m_step(counts, &t, j, opts.smoothing);
let log_likelihood = ds_log_likelihood(counts, &priors, &confusion);
Ok(DawidSkene {
posteriors: t,
confusion,
priors,
iterations,
converged,
log_likelihood,
})
}
pub fn dawid_skene(table: &LabelTable, opts: &DawidSkeneOpts) -> Result<DawidSkene, SymplexError> {
dawid_skene_counts(&table.to_counts(), table.n_categories, opts)
}
#[derive(Clone, Debug, PartialEq)]
pub struct BradleyTerryOpts {
pub max_iter: usize,
pub tol: f64,
}
impl Default for BradleyTerryOpts {
fn default() -> Self {
Self {
max_iter: 10_000,
tol: 1e-12,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct BradleyTerry {
pub strengths: Vec<f64>,
pub iterations: usize,
pub converged: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct PairwiseOutcome {
pub winner: usize,
pub loser: usize,
}
pub fn wins_matrix(
outcomes: &[PairwiseOutcome],
n: usize,
) -> Result<Vec<Vec<usize>>, SymplexError> {
let op = "wins_matrix";
let mut w = vec![vec![0usize; n]; n];
for outcome in outcomes {
let (a, b) = (outcome.winner, outcome.loser);
if a >= n || b >= n {
return Err(invalid(
op,
format!("player index out of range in ({a}, {b}), n = {n}"),
));
}
if a == b {
return Err(invalid(op, format!("player {a} cannot play itself")));
}
w[a][b] += 1;
}
Ok(w)
}
fn strongly_connected(wins: &[Vec<usize>]) -> bool {
let n = wins.len();
let reach = |forward: bool| -> bool {
let mut seen = vec![false; n];
let mut stack = vec![0usize];
seen[0] = true;
while let Some(i) = stack.pop() {
for j in 0..n {
let edge = if forward { wins[i][j] } else { wins[j][i] };
if edge > 0 && !seen[j] {
seen[j] = true;
stack.push(j);
}
}
}
seen.iter().all(|&s| s)
};
reach(true) && reach(false)
}
pub fn bradley_terry(
wins: &[Vec<usize>],
opts: &BradleyTerryOpts,
) -> Result<BradleyTerry, SymplexError> {
let op = "bradley_terry";
let n = wins.len();
if n < 2 || wins.iter().any(|r| r.len() != n) {
return Err(invalid(
op,
"the wins matrix must be square with at least two players",
));
}
if (0..n).any(|i| wins[i][i] != 0) {
return Err(invalid(op, "the diagonal of the wins matrix must be zero"));
}
if opts.max_iter == 0 {
return Err(invalid(op, "max_iter must be positive"));
}
if !(opts.tol > 0.0 && opts.tol.is_finite()) {
return Err(invalid(op, "tol must be a positive finite number"));
}
if !strongly_connected(wins) {
return Err(invalid(
op,
"the beat graph is not strongly connected (Ford 1957): the maximum-likelihood strengths do not exist",
));
}
let total_wins: Vec<f64> = wins
.iter()
.map(|r| r.iter().sum::<usize>() as f64)
.collect();
let mut p = vec![1.0 / n as f64; n];
let mut iterations = 0;
let mut converged = false;
for it in 1..=opts.max_iter {
let mut next: Vec<f64> = (0..n)
.map(|i| {
let denom: f64 = (0..n)
.filter(|&j| j != i)
.map(|j| (wins[i][j] + wins[j][i]) as f64 / (p[i] + p[j]))
.sum();
total_wins[i] / denom
})
.collect();
let s: f64 = next.iter().sum();
for v in &mut next {
*v /= s;
}
let delta = next
.iter()
.zip(&p)
.map(|(a, b)| (a - b).abs())
.fold(0.0, f64::max);
p = next;
iterations = it;
if delta < opts.tol {
converged = true;
break;
}
}
Ok(BradleyTerry {
strengths: p,
iterations,
converged,
})
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Accuracy {
pub correct: usize,
pub answered: usize,
pub accuracy: Option<Q>,
}
pub fn worker_accuracy(labels: &[Option<usize>], gold: &[usize]) -> Result<Accuracy, SymplexError> {
if labels.len() != gold.len() {
return Err(invalid(
"worker_accuracy",
"one gold label per item is needed",
));
}
let answered = labels.iter().flatten().count();
let correct = labels
.iter()
.zip(gold)
.filter(|(l, g)| **l == Some(**g))
.count();
let accuracy = (answered > 0).then(|| qu(correct) / qu(answered));
Ok(Accuracy {
correct,
answered,
accuracy,
})
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CategoryMetrics {
pub category: usize,
pub true_positives: usize,
pub false_positives: usize,
pub false_negatives: usize,
pub precision: Option<Q>,
pub recall: Option<Q>,
pub f1: Option<Q>,
}
pub fn category_metrics(
labels: &[Option<usize>],
gold: &[usize],
n_categories: usize,
) -> Result<Vec<CategoryMetrics>, SymplexError> {
let op = "category_metrics";
if labels.len() != gold.len() {
return Err(invalid(op, "one gold label per item is needed"));
}
if labels
.iter()
.flatten()
.chain(gold)
.any(|&l| l >= n_categories)
{
return Err(invalid(op, format!("a label is not below {n_categories}")));
}
let ratio = |num: usize, den: usize| (den > 0).then(|| qu(num) / qu(den));
Ok((0..n_categories)
.map(|c| {
let (mut tp, mut fp, mut fneg) = (0usize, 0usize, 0usize);
for (l, &g) in labels.iter().zip(gold) {
if let Some(l) = l {
match (*l == c, g == c) {
(true, true) => tp += 1,
(true, false) => fp += 1,
(false, true) => fneg += 1,
(false, false) => {}
}
}
}
CategoryMetrics {
category: c,
true_positives: tp,
false_positives: fp,
false_negatives: fneg,
precision: ratio(tp, tp + fp),
recall: ratio(tp, tp + fneg),
f1: ratio(2 * tp, 2 * tp + fp + fneg),
}
})
.collect())
}
pub fn gold_screening(
table: &LabelTable,
gold: &[Option<usize>],
threshold: &Q,
) -> Result<Vec<Option<bool>>, SymplexError> {
let op = "gold_screening";
if gold.len() != table.n_items() {
return Err(invalid(op, "one gold entry per item is needed"));
}
if gold.iter().flatten().any(|&g| g >= table.n_categories) {
return Err(invalid(op, "a gold label is not a category"));
}
if threshold < &Q::zero() || threshold > &Q::one() {
return Err(invalid(op, "the threshold must lie in [0, 1]"));
}
Ok((0..table.n_raters)
.map(|j| {
let (labels, golds): (Vec<Option<usize>>, Vec<usize>) = table
.rows
.iter()
.zip(gold)
.filter_map(|(r, g)| g.map(|g| (r[j], g)))
.unzip();
let (mut answered, mut correct) = (0usize, 0usize);
for (l, g) in labels.iter().zip(&golds) {
if let Some(l) = l {
answered += 1;
if l == g {
correct += 1;
}
}
}
(answered > 0).then(|| qu(correct) >= threshold * qu(answered))
})
.collect())
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum IntervalMethod {
Wilson,
ClopperPearson,
AgrestiCoull,
Wald,
}
fn binomial_tail(n: usize, k: usize, p: f64, upper: bool) -> f64 {
let (lp, lq) = (p.ln(), (1.0 - p).ln());
let mut log_c = 0.0; let mut tail = 0.0;
for i in 0..=n {
if i > 0 {
log_c += ((n - i + 1) as f64).ln() - (i as f64).ln();
}
if if upper { i >= k } else { i <= k } {
tail += (log_c + i as f64 * lp + (n - i) as f64 * lq).exp();
}
}
tail.min(1.0)
}
fn bisect_unit(f: impl Fn(f64) -> f64, increasing: bool) -> f64 {
let (mut lo, mut hi) = (0.0f64, 1.0f64);
for _ in 0..200 {
let mid = 0.5 * (lo + hi);
if mid <= lo || mid >= hi {
break;
}
let v = f(mid);
if (v < 0.0) == increasing {
lo = mid;
} else {
hi = mid;
}
}
0.5 * (lo + hi)
}
pub fn proportion_interval(
successes: usize,
trials: usize,
confidence: f64,
method: IntervalMethod,
) -> Result<Interval<f64>, SymplexError> {
let op = "proportion_interval";
if trials == 0 {
return Err(invalid(op, "needs at least one trial"));
}
if successes > trials {
return Err(invalid(op, "more successes than trials"));
}
if !(confidence > 0.0 && confidence < 1.0) {
return Err(invalid(
op,
"the confidence must lie strictly between 0 and 1",
));
}
let alpha = 1.0 - confidence;
let (k, n) = (successes as f64, trials as f64);
let unit = Interval::closed(0.0, 1.0);
let clip = |ci: Interval<f64>| ci.map(|v| unit.clamp_to_closure(v));
if method == IntervalMethod::ClopperPearson {
let half = alpha / 2.0;
let lo = if successes == 0 {
0.0
} else {
bisect_unit(|p| binomial_tail(trials, successes, p, true) - half, true)
};
let hi = if successes == trials {
1.0
} else {
bisect_unit(|p| binomial_tail(trials, successes, p, false) - half, false)
};
return Ok(Interval::closed(lo, hi));
}
let ctx = Context::new();
let z = Distribution::normal(ctx.int(0), ctx.int(1)).quantile_f64(1.0 - alpha / 2.0)?;
let p = k / n;
Ok(clip(match method {
IntervalMethod::Wald => {
let half = z * (p * (1.0 - p) / n).sqrt();
Interval::closed(p - half, p + half)
}
IntervalMethod::Wilson => {
let z2 = z * z;
let denom = 1.0 + z2 / n;
let centre = (p + z2 / (2.0 * n)) / denom;
let half = z * (p * (1.0 - p) / n + z2 / (4.0 * n * n)).sqrt() / denom;
Interval::closed(centre - half, centre + half)
}
IntervalMethod::AgrestiCoull => {
let z2 = z * z;
let n_t = n + z2;
let p_t = (k + z2 / 2.0) / n_t;
let half = z * (p_t * (1.0 - p_t) / n_t).sqrt();
Interval::closed(p_t - half, p_t + half)
}
IntervalMethod::ClopperPearson => Interval::closed(0.0, 1.0),
}))
}