qualia_core_db/solvers/learning/classification/
knn.rs1use crate::solvers::learning::LearningError;
6
7#[derive(Debug, Clone)]
9pub struct KnnClassifier {
10 x: Vec<f64>,
11 y: Vec<usize>,
12 n: usize,
13 p: usize,
14 k: usize,
15}
16
17impl KnnClassifier {
18 pub fn fit(
20 x: &[f64],
21 y: &[usize],
22 n: usize,
23 p: usize,
24 k: usize,
25 ) -> Result<Self, LearningError> {
26 if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
27 return Err(LearningError::InvalidDimension);
28 }
29 if k == 0 || k > n {
30 return Err(LearningError::InsufficientData);
31 }
32 Ok(Self {
33 x: x.to_vec(),
34 y: y.to_vec(),
35 n,
36 p,
37 k,
38 })
39 }
40
41 pub fn predict_row(&self, q: &[f64]) -> usize {
44 let mut idx: Vec<usize> = (0..self.n).collect();
46 idx.sort_by(|&a, &b| {
47 let da = self.sq_dist(a, q);
48 let db = self.sq_dist(b, q);
49 da.partial_cmp(&db).unwrap_or(core::cmp::Ordering::Equal)
50 });
51 let max_label = *self.y.iter().max().unwrap_or(&0);
53 let mut votes = vec![0usize; max_label + 1];
54 for &i in idx.iter().take(self.k) {
55 votes[self.y[i]] += 1;
56 }
57 let mut best = 0;
59 let mut best_v = 0;
60 for (label, &v) in votes.iter().enumerate() {
61 if v > best_v {
62 best_v = v;
63 best = label;
64 }
65 }
66 best
67 }
68
69 pub fn predict(&self, x: &[f64], m: usize) -> Vec<usize> {
71 (0..m)
72 .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
73 .collect()
74 }
75
76 #[inline]
77 fn sq_dist(&self, train_row: usize, q: &[f64]) -> f64 {
78 let row = &self.x[train_row * self.p..(train_row + 1) * self.p];
79 row.iter().zip(q).map(|(a, b)| (a - b) * (a - b)).sum()
80 }
81}
82
83#[cfg(test)]
84mod tests {
85 use super::*;
86
87 #[test]
88 fn classifies_by_nearest_neighbours() {
89 let x = [
91 0.0, 0.0, 0.5, 0.3, -0.2, 0.1, 10.0, 10.0, 9.7, 10.2, 10.3, 9.8,
92 ];
93 let y = [0, 0, 0, 1, 1, 1];
94 let knn = KnnClassifier::fit(&x, &y, 6, 2, 3).unwrap();
95 assert_eq!(knn.predict_row(&[0.1, 0.1]), 0);
96 assert_eq!(knn.predict_row(&[10.1, 9.9]), 1);
97 }
98
99 #[test]
100 fn k_one_is_the_single_nearest() {
101 let x = [0.0, 0.0, 5.0, 5.0];
102 let y = [7, 3];
103 let knn = KnnClassifier::fit(&x, &y, 2, 2, 1).unwrap();
104 assert_eq!(knn.predict_row(&[0.4, 0.4]), 7);
105 assert_eq!(knn.predict_row(&[4.6, 4.6]), 3);
106 }
107
108 #[test]
109 fn guards() {
110 assert_eq!(
111 KnnClassifier::fit(&[1.0, 2.0], &[0], 1, 2, 5).unwrap_err(),
112 LearningError::InsufficientData
113 );
114 assert_eq!(
115 KnnClassifier::fit(&[1.0, 2.0], &[0, 1], 2, 2, 1).unwrap_err(),
116 LearningError::InvalidDimension
117 );
118 }
119}