1#[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
24pub 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], pub support_vectors: &'a [f32], 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}