qualia_core_db/solvers/learning/active/
committee.rs1use super::{argsort_desc, ActiveError};
9use crate::solvers::statistics::information::{entropy, kl_divergence};
10
11pub fn vote_entropy(votes: &[usize], n_classes: usize) -> Result<f64, ActiveError> {
15 if votes.is_empty() || n_classes == 0 {
16 return Err(ActiveError::InsufficientData);
17 }
18 let mut counts = vec![0.0f64; n_classes];
19 for &v in votes {
20 if v >= n_classes {
21 return Err(ActiveError::InvalidDimension);
22 }
23 counts[v] += 1.0;
24 }
25 let total = votes.len() as f64;
26 for c in counts.iter_mut() {
27 *c /= total;
28 }
29 entropy(&counts).ok_or(ActiveError::InvalidDimension)
30}
31
32pub fn consensus(members: &[Vec<f64>]) -> Result<Vec<f64>, ActiveError> {
35 if members.is_empty() {
36 return Err(ActiveError::InsufficientData);
37 }
38 let n_classes = members[0].len();
39 if n_classes == 0 || members.iter().any(|m| m.len() != n_classes) {
40 return Err(ActiveError::InvalidDimension);
41 }
42 let mut mean = vec![0.0; n_classes];
43 for m in members {
44 for (acc, &p) in mean.iter_mut().zip(m) {
45 *acc += p;
46 }
47 }
48 let k = members.len() as f64;
49 for v in mean.iter_mut() {
50 *v /= k;
51 }
52 Ok(mean)
53}
54
55pub fn consensus_entropy(members: &[Vec<f64>]) -> Result<f64, ActiveError> {
58 entropy(&consensus(members)?).ok_or(ActiveError::InvalidDimension)
59}
60
61pub fn average_kl_disagreement(members: &[Vec<f64>]) -> Result<f64, ActiveError> {
65 let cons = consensus(members)?;
66 let mut sum = 0.0;
67 for m in members {
68 sum += kl_divergence(m, &cons).ok_or(ActiveError::InvalidDimension)?;
69 }
70 Ok(sum / members.len() as f64)
71}
72
73pub fn rank_by_disagreement(pool: &[Vec<Vec<f64>>]) -> Result<Vec<usize>, ActiveError> {
76 if pool.is_empty() {
77 return Err(ActiveError::InsufficientData);
78 }
79 let scores: Result<Vec<f64>, ActiveError> =
80 pool.iter().map(|c| average_kl_disagreement(c)).collect();
81 Ok(argsort_desc(&scores?))
82}
83
84#[cfg(test)]
85mod tests {
86 use super::*;
87 const EPS: f64 = 1e-9;
88
89 #[test]
90 fn vote_entropy_max_on_even_split() {
91 let v = vote_entropy(&[0, 1], 2).unwrap();
93 assert!((v - 1.0).abs() < 1e-9);
94 assert!(vote_entropy(&[1, 1, 1], 2).unwrap().abs() < EPS);
96 }
97
98 #[test]
99 fn disagreeing_committee_scores_above_agreeing() {
100 let agree = vec![vec![0.9, 0.1], vec![0.85, 0.15]];
102 let disagree = vec![vec![0.95, 0.05], vec![0.05, 0.95]];
104 let a = average_kl_disagreement(&agree).unwrap();
105 let d = average_kl_disagreement(&disagree).unwrap();
106 assert!(d > a, "disagreement {d} should exceed agreement {a}");
107 }
108
109 #[test]
110 fn consensus_is_the_mean_distribution() {
111 let m = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
112 let c = consensus(&m).unwrap();
113 assert!((c[0] - 0.5).abs() < EPS && (c[1] - 0.5).abs() < EPS);
114 assert!((consensus_entropy(&m).unwrap() - 1.0).abs() < 1e-9);
116 }
117
118 #[test]
119 fn pool_ranking_puts_the_split_sample_first() {
120 let agree = vec![vec![0.9, 0.1], vec![0.88, 0.12]];
121 let split = vec![vec![0.95, 0.05], vec![0.05, 0.95]];
122 let pool = vec![agree, split];
123 let r = rank_by_disagreement(&pool).unwrap();
124 assert_eq!(r[0], 1);
125 }
126
127 #[test]
128 fn fails_closed() {
129 assert_eq!(consensus(&[]).unwrap_err(), ActiveError::InsufficientData);
130 assert_eq!(
131 vote_entropy(&[5], 2).unwrap_err(),
132 ActiveError::InvalidDimension
133 );
134 }
135}