Skip to main content

qualia_core_db/solvers/learning/classification/
knn.rs

1//! k-Nearest-Neighbours classifier (ISL ch 2/4) — a lazy learner: store the
2//! training set, classify a query by majority vote of its `k` nearest neighbours
3//! (squared Euclidean). Kernel-class `AllPairs` (query↔train distances).
4
5use crate::solvers::learning::LearningError;
6
7/// A fitted (stored) k-NN classifier.
8#[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    /// Store the training data. Fails closed on shape mismatch or `k` out of range.
19    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    /// Predict the class of one query row by majority vote of the `k` nearest
42    /// training points (ties broken toward the lower class label).
43    pub fn predict_row(&self, q: &[f64]) -> usize {
44        // Indices of all training rows, sorted by distance to q, take k.
45        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        // Tally votes among the k nearest.
52        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        // argmax votes, lowest label on a tie.
58        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    /// Predict classes for a row-major `m × p` query matrix.
70    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        // Two clusters: class 0 near origin, class 1 near (10,10).
90        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}