Skip to main content

embedded_dsp/
svm.rs

1//! Support Vector Machine (SVM) Classifier (Linear, Polynomial, RBF, Sigmoid kernels).
2
3#[allow(unused_imports)]
4use crate::math::FloatMath;
5use crate::types::*;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
8pub enum SvmKernelType {
9    Linear,
10    Polynomial,
11    Rbf,
12    Sigmoid,
13}
14
15#[inline(always)]
16fn pow_core(x: f32, p: i32) -> f32 {
17    let mut res = 1.0f32;
18    for _ in 0..p {
19        res *= x;
20    }
21    res
22}
23
24/// SVM Classifier Instance structure for f32.
25pub struct SvmInstanceF32<'a> {
26    pub num_vector_dim: usize,
27    pub num_support_vectors: usize,
28    pub intercept: f32,
29    pub dual_coefs: &'a [f32],      // Size num_support_vectors
30    pub support_vectors: &'a [f32], // Size num_support_vectors * num_vector_dim
31    pub kernel_type: SvmKernelType,
32    pub gamma: f32,
33    pub coef0: f32,
34    pub degree: i32,
35}
36
37impl<'a> SvmInstanceF32<'a> {
38    pub fn predict(&self, input: &[f32], result: &mut i32) -> Status {
39        if input.len() < self.num_vector_dim {
40            return Status::LengthError;
41        }
42
43        let mut sum = self.intercept;
44        for i in 0..self.num_support_vectors {
45            let sv = &self.support_vectors[i * self.num_vector_dim..(i + 1) * self.num_vector_dim];
46            let alpha = self.dual_coefs[i];
47
48            let kernel_val = match self.kernel_type {
49                SvmKernelType::Linear => {
50                    let mut dot = 0.0f32;
51                    for d in 0..self.num_vector_dim {
52                        dot += sv[d] * input[d];
53                    }
54                    dot
55                }
56                SvmKernelType::Polynomial => {
57                    let mut dot = 0.0f32;
58                    for d in 0..self.num_vector_dim {
59                        dot += sv[d] * input[d];
60                    }
61                    pow_core(self.gamma * dot + self.coef0, self.degree)
62                }
63                SvmKernelType::Rbf => {
64                    let mut dist_sq = 0.0f32;
65                    for d in 0..self.num_vector_dim {
66                        let diff = input[d] - sv[d];
67                        dist_sq += diff * diff;
68                    }
69                    (-self.gamma * dist_sq).exp()
70                }
71                SvmKernelType::Sigmoid => {
72                    let mut dot = 0.0f32;
73                    for d in 0..self.num_vector_dim {
74                        dot += sv[d] * input[d];
75                    }
76                    (self.gamma * dot + self.coef0).tanh()
77                }
78            };
79            sum += alpha * kernel_val;
80        }
81
82        *result = if sum >= 0.0 { 1 } else { -1 };
83        Status::Success
84    }
85}
86
87pub fn svm_predict_f32(instance: &SvmInstanceF32, input: &[f32], result: &mut i32) -> Status {
88    instance.predict(input, result)
89}