use num_bigint::BigInt;
use num_traits::{One, Zero};
use crate::base::errors::SymplexError;
use crate::domains::stats::data::{self, 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))
}
fn qi(n: i64) -> Q {
Q::from_integer(BigInt::from(n))
}
fn sum(values: impl IntoIterator<Item = Q>) -> Q {
values.into_iter().fold(Q::zero(), |acc, x| acc + x)
}
fn square(x: &Q) -> Q {
x * x
}
fn check_pair(op: &'static str, a: &[Q], b: &[Q]) -> Result<(), SymplexError> {
if a.len() != b.len() {
return Err(invalid(
op,
format!("the two raters rated {} and {} items", a.len(), b.len()),
));
}
if a.is_empty() {
return Err(invalid(op, "needs at least one item"));
}
Ok(())
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RatingTable {
rows: Vec<Vec<Option<Q>>>,
n_raters: usize,
}
impl RatingTable {
pub fn new(rows: Vec<Vec<Option<Q>>>) -> Result<Self, SymplexError> {
let op = "RatingTable::new";
let n_raters = match rows.first() {
Some(r) => r.len(),
None => return Err(invalid(op, "a rating table needs at least one item")),
};
if n_raters == 0 {
return Err(invalid(op, "a rating table needs at least one rater"));
}
if let Some((i, r)) = rows.iter().enumerate().find(|(_, r)| r.len() != n_raters) {
return Err(invalid(
op,
format!(
"item {i} has {} cells but the table has {n_raters} raters",
r.len()
),
));
}
Ok(Self { rows, n_raters })
}
pub fn from_rows(rows: &[&[Option<Q>]]) -> Result<Self, SymplexError> {
Self::new(rows.iter().map(|r| r.to_vec()).collect())
}
pub fn from_i64(rows: &[&[i64]]) -> Result<Self, SymplexError> {
Self::new(
rows.iter()
.map(|r| r.iter().map(|&x| Some(qi(x))).collect())
.collect(),
)
}
pub fn from_i64_missing(rows: &[&[Option<i64>]]) -> Result<Self, SymplexError> {
Self::new(
rows.iter()
.map(|r| r.iter().map(|x| x.map(qi)).collect())
.collect(),
)
}
pub fn from_raters_i64(raters: &[&[Option<i64>]]) -> Result<Self, SymplexError> {
let op = "RatingTable::from_raters_i64";
let n_items = match raters.first() {
Some(r) => r.len(),
None => return Err(invalid(op, "a rating table needs at least one rater")),
};
if raters.iter().any(|r| r.len() != n_items) {
return Err(invalid(op, "every rater must have one cell per item"));
}
if n_items == 0 {
return Err(invalid(op, "a rating table needs at least one item"));
}
Self::new(
(0..n_items)
.map(|i| raters.iter().map(|r| r[i].map(qi)).collect())
.collect(),
)
}
pub fn n_items(&self) -> usize {
self.rows.len()
}
pub fn n_raters(&self) -> usize {
self.n_raters
}
pub fn rows(&self) -> &[Vec<Option<Q>>] {
&self.rows
}
pub fn item(&self, i: usize) -> Option<&[Option<Q>]> {
self.rows.get(i).map(|r| r.as_slice())
}
pub fn rater(&self, j: usize) -> Option<Vec<Option<Q>>> {
(j < self.n_raters).then(|| self.rows.iter().map(|r| r[j].clone()).collect())
}
pub fn get(&self, i: usize, j: usize) -> Option<&Q> {
self.rows
.get(i)
.and_then(|r| r.get(j))
.and_then(|c| c.as_ref())
}
pub fn is_complete(&self) -> bool {
self.rows.iter().all(|r| r.iter().all(|c| c.is_some()))
}
pub fn categories(&self) -> Vec<Q> {
let mut v: Vec<Q> = self.rows.iter().flatten().flatten().cloned().collect();
v.sort();
v.dedup();
v
}
pub fn count_table(&self, categories: &[Q]) -> Result<Vec<Vec<usize>>, SymplexError> {
let op = "RatingTable::count_table";
check_categories(op, categories)?;
self.rows
.iter()
.map(|row| {
let mut counts = vec![0usize; categories.len()];
for x in row.iter().flatten() {
let c = categories
.iter()
.position(|k| k == x)
.ok_or_else(|| invalid(op, format!("rating {x} is not a category")))?;
counts[c] += 1;
}
Ok(counts)
})
.collect()
}
pub fn paired_ratings(&self, j1: usize, j2: usize) -> Result<(Vec<Q>, Vec<Q>), SymplexError> {
if j1 >= self.n_raters || j2 >= self.n_raters {
return Err(invalid(
"RatingTable::paired_ratings",
format!(
"rater index out of range (the table has {} raters)",
self.n_raters
),
));
}
Ok(self
.rows
.iter()
.filter_map(|r| match (&r[j1], &r[j2]) {
(Some(x), Some(y)) => Some((x.clone(), y.clone())),
_ => None,
})
.unzip())
}
pub fn complete_rows(&self) -> Result<Vec<Vec<Q>>, SymplexError> {
self.rows
.iter()
.enumerate()
.map(|(i, r)| {
r.iter()
.cloned()
.map(|c| {
c.ok_or_else(|| {
invalid(
"RatingTable::complete_rows",
format!("item {i} has a missing rating"),
)
})
})
.collect()
})
.collect()
}
}
fn check_categories(op: &'static str, categories: &[Q]) -> Result<(), SymplexError> {
if categories.is_empty() {
return Err(invalid(op, "needs at least one category"));
}
for (i, c) in categories.iter().enumerate() {
if categories[..i].contains(c) {
return Err(invalid(op, format!("category {c} is listed twice")));
}
}
Ok(())
}
pub fn category_frequencies(table: &RatingTable) -> Vec<(Q, usize)> {
let all: Vec<Q> = table.rows.iter().flatten().flatten().cloned().collect();
data::frequencies(&all)
}
pub fn confusion_matrix(
a: &[Q],
b: &[Q],
categories: &[Q],
) -> Result<Vec<Vec<usize>>, SymplexError> {
let op = "confusion_matrix";
check_pair(op, a, b)?;
check_categories(op, categories)?;
let index = |x: &Q| {
categories
.iter()
.position(|k| k == x)
.ok_or_else(|| invalid(op, format!("rating {x} is not a category")))
};
let k = categories.len();
let mut m = vec![vec![0usize; k]; k];
for (x, y) in a.iter().zip(b) {
m[index(x)?][index(y)?] += 1;
}
Ok(m)
}
fn observed_categories(a: &[Q], b: &[Q]) -> Vec<Q> {
let mut v: Vec<Q> = a.iter().chain(b).cloned().collect();
v.sort();
v.dedup();
v
}
pub fn percent_agreement(a: &[Q], b: &[Q]) -> Result<Q, SymplexError> {
check_pair("percent_agreement", a, b)?;
let agree = a.iter().zip(b).filter(|(x, y)| x == y).count();
Ok(qu(agree) / qu(a.len()))
}
pub fn pairwise_percent_agreement(table: &RatingTable) -> Result<Q, SymplexError> {
let op = "pairwise_percent_agreement";
let m = table.n_raters();
if m < 2 {
return Err(invalid(op, "needs at least two raters"));
}
let mut total = Q::zero();
let mut pairs = 0usize;
for j1 in 0..m {
for j2 in j1 + 1..m {
let (a, b) = table.paired_ratings(j1, j2)?;
if a.is_empty() {
continue;
}
total += percent_agreement(&a, &b)?;
pairs += 1;
}
}
if pairs == 0 {
return Err(invalid(op, "no two raters rated a common item"));
}
Ok(total / qu(pairs))
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct KappaResult {
pub kappa: Q,
pub observed: Q,
pub expected: Q,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Weights {
Unweighted,
Linear,
Quadratic,
Custom(Vec<Vec<Q>>),
}
impl Weights {
fn matrix(&self, op: &'static str, k: usize) -> Result<Vec<Vec<Q>>, SymplexError> {
let dist = |i: usize, j: usize| qu(i.abs_diff(j));
match self {
Weights::Unweighted => Ok((0..k)
.map(|i| {
(0..k)
.map(|j| if i == j { Q::zero() } else { Q::one() })
.collect()
})
.collect()),
Weights::Linear | Weights::Quadratic => {
if k < 2 {
return Err(invalid(op, "weighted κ needs at least two categories"));
}
let span = qu(k - 1);
Ok((0..k)
.map(|i| {
(0..k)
.map(|j| {
let d = dist(i, j) / &span;
if *self == Weights::Linear {
d
} else {
square(&d)
}
})
.collect()
})
.collect())
}
Weights::Custom(w) => {
if w.len() != k || w.iter().any(|r| r.len() != k) {
return Err(invalid(
op,
format!("the weight matrix must be {k} × {k} like the confusion matrix"),
));
}
Ok(w.clone())
}
}
}
}
pub fn kappa_from_confusion(
table: &[Vec<usize>],
weights: &Weights,
) -> Result<KappaResult, SymplexError> {
let op = "cohen_kappa";
let k = table.len();
if k == 0 || table.iter().any(|r| r.len() != k) {
return Err(invalid(
op,
"the confusion matrix must be square and non-empty",
));
}
let n: usize = table.iter().flatten().sum();
if n == 0 {
return Err(invalid(op, "the confusion matrix is empty"));
}
let w = weights.matrix(op, k)?;
let row: Vec<usize> = table.iter().map(|r| r.iter().sum()).collect();
let col: Vec<usize> = (0..k).map(|j| table.iter().map(|r| r[j]).sum()).collect();
let nq = qu(n);
let mut d_obs = Q::zero();
let mut d_exp = Q::zero();
for (i, r) in table.iter().enumerate() {
for (j, &cell) in r.iter().enumerate() {
d_obs += &w[i][j] * qu(cell);
d_exp += &w[i][j] * qu(row[i] * col[j]);
}
}
d_obs /= &nq;
d_exp /= square(&nq);
if d_exp.is_zero() {
return Err(invalid(
op,
"the expected disagreement is zero (a single category), κ is undefined",
));
}
Ok(KappaResult {
kappa: Q::one() - &d_obs / &d_exp,
observed: Q::one() - d_obs,
expected: Q::one() - d_exp,
})
}
pub fn cohen_kappa(a: &[Q], b: &[Q]) -> Result<KappaResult, SymplexError> {
weighted_kappa(a, b, &Weights::Unweighted)
}
pub fn weighted_kappa(a: &[Q], b: &[Q], weights: &Weights) -> Result<KappaResult, SymplexError> {
check_pair("cohen_kappa", a, b)?;
let table = confusion_matrix(a, b, &observed_categories(a, b))?;
kappa_from_confusion(&table, weights)
}
pub fn scott_pi(a: &[Q], b: &[Q]) -> Result<Q, SymplexError> {
let op = "scott_pi";
check_pair(op, a, b)?;
let n = a.len();
let p_o = percent_agreement(a, b)?;
let two_n = qu(2 * n);
let p_e = sum(observed_categories(a, b).iter().map(|c| {
let pooled = a.iter().filter(|x| *x == c).count() + b.iter().filter(|x| *x == c).count();
square(&(qu(pooled) / &two_n))
}));
let denom = Q::one() - &p_e;
if denom.is_zero() {
return Err(invalid(op, "a single category, π is undefined"));
}
Ok((p_o - p_e) / denom)
}
pub fn fleiss_kappa(counts: &[Vec<usize>]) -> Result<Q, SymplexError> {
let op = "fleiss_kappa";
let n_items = counts.len();
let k = counts.first().map_or(0, |r| r.len());
if n_items == 0 || k == 0 || counts.iter().any(|r| r.len() != k) {
return Err(invalid(
op,
"the count table must be rectangular and non-empty",
));
}
let n: usize = counts[0].iter().sum();
if n < 2 {
return Err(invalid(op, "needs at least two raters per item"));
}
if counts.iter().any(|r| r.iter().sum::<usize>() != n) {
return Err(invalid(
op,
"every item must be rated by the same number of raters",
));
}
let per_item = qu(n * (n - 1));
let p_bar = sum(counts.iter().map(|r| {
let sq: usize = r.iter().map(|&c| c * c).sum();
qu(sq - n) / &per_item
})) / qu(n_items);
let total = qu(n_items * n);
let p_e = sum((0..k).map(|j| {
let col: usize = counts.iter().map(|r| r[j]).sum();
square(&(qu(col) / &total))
}));
let denom = Q::one() - &p_e;
if denom.is_zero() {
return Err(invalid(op, "a single category, κ is undefined"));
}
Ok((p_bar - p_e) / denom)
}
pub fn fleiss_kappa_ratings(table: &RatingTable) -> Result<Q, SymplexError> {
if !table.is_complete() {
return Err(invalid("fleiss_kappa", "Fleiss' κ needs a complete table"));
}
fleiss_kappa(&table.count_table(&table.categories())?)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Level {
Nominal,
Ordinal,
Interval,
Ratio,
}
pub fn krippendorff_alpha(table: &RatingTable, level: Level) -> Result<Q, SymplexError> {
let op = "krippendorff_alpha";
let cats = table.categories();
let v = cats.len();
if v < 2 {
return Err(invalid(op, "needs at least two distinct values"));
}
let counts = table.count_table(&cats)?;
let mut o = vec![vec![Q::zero(); v]; v];
let mut pairable_items = 0usize;
for row in &counts {
let m: usize = row.iter().sum();
if m < 2 {
continue;
}
pairable_items += 1;
let denom = qu(m - 1);
for (c, &n_uc) in row.iter().enumerate() {
if n_uc == 0 {
continue;
}
for (k, &n_uk) in row.iter().enumerate() {
let pairs = n_uc * (n_uk - usize::from(c == k));
if pairs > 0 {
o[c][k] += qu(pairs) / &denom;
}
}
}
}
if pairable_items == 0 {
return Err(invalid(
op,
"needs at least one item rated by two or more raters",
));
}
let n_c: Vec<Q> = o.iter().map(|r| sum(r.iter().cloned())).collect();
let n = sum(n_c.iter().cloned());
let mut prefix = vec![Q::zero()];
for x in &n_c {
let last = prefix[prefix.len() - 1].clone();
prefix.push(last + x);
}
let delta = |c: usize, k: usize| -> Q {
match level {
Level::Nominal => {
if c == k {
Q::zero()
} else {
Q::one()
}
}
Level::Interval => square(&(&cats[c] - &cats[k])),
Level::Ratio => {
let s = &cats[c] + &cats[k];
if s.is_zero() {
Q::zero()
} else {
square(&((&cats[c] - &cats[k]) / s))
}
}
Level::Ordinal => {
let (lo, hi) = (c.min(k), c.max(k));
let between = &prefix[hi + 1] - &prefix[lo];
square(&(between - (&n_c[lo] + &n_c[hi]) / qi(2)))
}
}
};
let mut d_o = Q::zero();
let mut d_e = Q::zero();
for c in 0..v {
for k in 0..v {
let d = delta(c, k);
if d.is_zero() {
continue;
}
d_o += &o[c][k] * &d;
let expected = &n_c[c] * &n_c[k] - if c == k { n_c[c].clone() } else { Q::zero() };
d_e += expected * d;
}
}
d_e /= n - Q::one();
if d_e.is_zero() {
return Err(invalid(
op,
"every pairable value is identical, α is undefined",
));
}
Ok(Q::one() - d_o / d_e)
}
pub fn gwet_ac1(table: &RatingTable, categories: Option<&[Q]>) -> Result<Q, SymplexError> {
let op = "gwet_ac1";
let cats: Vec<Q> = match categories {
Some(c) => c.to_vec(),
None => table.categories(),
};
let kcat = cats.len();
if kcat < 2 {
return Err(invalid(op, "needs at least two categories"));
}
let counts = table.count_table(&cats)?;
let mut pa_sum = Q::zero();
let mut pairable = 0usize;
let mut rated = 0usize;
let mut pi = vec![Q::zero(); kcat];
for row in &counts {
let r: usize = row.iter().sum();
if r == 0 {
continue;
}
rated += 1;
for (k, &c) in row.iter().enumerate() {
pi[k] += qu(c) / qu(r);
}
if r >= 2 {
pairable += 1;
let agree: usize = row.iter().map(|&c| c * (c.saturating_sub(1))).sum();
pa_sum += qu(agree) / qu(r * (r - 1));
}
}
if pairable == 0 {
return Err(invalid(
op,
"needs at least one item rated by two or more raters",
));
}
let p_a = pa_sum / qu(pairable);
let p_e = sum(pi.iter().map(|p| {
let p = p / qu(rated);
&p * (Q::one() - &p)
})) / qu(kcat - 1);
let denom = Q::one() - &p_e;
if denom.is_zero() {
return Err(invalid(op, "the chance agreement is 1, AC₁ is undefined"));
}
Ok((p_a - p_e) / denom)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum IccForm {
Icc1,
Icc1Average,
Icc2Single,
Icc2Average,
Icc3Single,
Icc3Average,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct IccAnova {
pub msr: Q,
pub msc: Q,
pub mse: Q,
pub msw: Q,
}
pub fn icc_anova(table: &RatingTable) -> Result<IccAnova, SymplexError> {
let op = "icc";
let x = table
.complete_rows()
.map_err(|_| invalid(op, "the intraclass correlation needs a complete table"))?;
let (n, k) = (x.len(), table.n_raters());
if n < 2 {
return Err(invalid(op, "needs at least two items"));
}
if k < 2 {
return Err(invalid(op, "needs at least two raters"));
}
let grand = sum(x.iter().flatten().cloned()) / qu(n * k);
let row_means: Vec<Q> = x.iter().map(|r| sum(r.iter().cloned()) / qu(k)).collect();
let col_means: Vec<Q> = (0..k)
.map(|j| sum(x.iter().map(|r| r[j].clone())) / qu(n))
.collect();
let ssr = qu(k) * sum(row_means.iter().map(|m| square(&(m - &grand))));
let ssc = qu(n) * sum(col_means.iter().map(|m| square(&(m - &grand))));
let sst = sum(x.iter().flatten().map(|v| square(&(v - &grand))));
let sse = &sst - &ssr - &ssc;
let ssw = &ssc + &sse;
Ok(IccAnova {
msr: ssr / qu(n - 1),
msc: ssc / qu(k - 1),
mse: sse / qu((n - 1) * (k - 1)),
msw: ssw / qu(n * (k - 1)),
})
}
pub fn icc(table: &RatingTable, form: IccForm) -> Result<Q, SymplexError> {
let a = icc_anova(table)?;
let (n, k) = (qu(table.n_items()), qu(table.n_raters()));
let k1 = &k - Q::one();
let (num, den) = match form {
IccForm::Icc1 => (&a.msr - &a.msw, &a.msr + &k1 * &a.msw),
IccForm::Icc1Average => (&a.msr - &a.msw, a.msr.clone()),
IccForm::Icc2Single => (
&a.msr - &a.mse,
&a.msr + &k1 * &a.mse + &k * (&a.msc - &a.mse) / &n,
),
IccForm::Icc2Average => (&a.msr - &a.mse, &a.msr + (&a.msc - &a.mse) / &n),
IccForm::Icc3Single => (&a.msr - &a.mse, &a.msr + &k1 * &a.mse),
IccForm::Icc3Average => (&a.msr - &a.mse, a.msr.clone()),
};
if den.is_zero() {
return Err(invalid(
"icc",
"the denominator is zero (no variance between items), the ICC is undefined",
));
}
Ok(num / den)
}
pub fn kendall_w(table: &RatingTable) -> Result<Q, SymplexError> {
let op = "kendall_w";
let x = table
.complete_rows()
.map_err(|_| invalid(op, "Kendall's W needs a complete table"))?;
let (n, m) = (x.len(), table.n_raters());
if n < 2 {
return Err(invalid(op, "needs at least two items"));
}
if m < 2 {
return Err(invalid(op, "needs at least two raters"));
}
let mut rank_sums = vec![Q::zero(); n];
let mut ties = BigInt::zero();
for j in 0..m {
let col: Vec<Q> = x.iter().map(|r| r[j].clone()).collect();
for (i, r) in data::ranks(&col).into_iter().enumerate() {
rank_sums[i] += r;
}
for t in data::tie_sizes(&col) {
let t = BigInt::from(t);
ties += &t * &t * &t - &t;
}
}
let mean = sum(rank_sums.iter().cloned()) / qu(n);
let s = sum(rank_sums.iter().map(|r| square(&(r - &mean))));
let (nb, mb) = (BigInt::from(n), BigInt::from(m));
let denom = &mb * &mb * (&nb * &nb * &nb - &nb) - &mb * ties;
if denom.is_zero() {
return Err(invalid(op, "every rater ties all items, W is undefined"));
}
Ok(qi(12) * s / Q::from_integer(denom))
}