Skip to main content

sklears_svm/
kernel_pca.rs

1//! Kernel Principal Component Analysis (Kernel PCA) for non-linear dimensionality reduction
2//!
3//! Kernel PCA extends classical PCA to non-linear data by first mapping the data
4//! to a higher-dimensional feature space using kernel functions, then performing
5//! PCA in that space. This is particularly useful as a preprocessing step for SVMs.
6
7use scirs2_core::ndarray::{s, Array1, Array2, Axis};
8use scirs2_linalg::compat::eigh;
9use sklears_core::{error::SklearsError, types::Float};
10
11use crate::kernels::{create_kernel, Kernel, KernelType};
12
13/// Kernel Principal Component Analysis for non-linear dimensionality reduction
14#[derive(Debug)]
15pub struct KernelPCA {
16    /// Number of components to keep
17    n_components: usize,
18    /// Kernel function to use
19    kernel: KernelType,
20    /// Tolerance for eigenvalue computation
21    tol: Float,
22    /// Maximum number of iterations for iterative solvers
23    max_iter: usize,
24    /// Whether to center the kernel matrix
25    fit_inverse_transform: bool,
26    /// Eigenvalues (computed during fit)
27    eigenvalues_: Option<Array1<Float>>,
28    /// Eigenvectors (computed during fit)
29    eigenvectors_: Option<Array2<Float>>,
30    /// Centered kernel matrix eigenvectors for inverse transform
31    alphas_: Option<Array2<Float>>,
32    /// Training data for inverse transform
33    x_fit_: Option<Array2<Float>>,
34    /// Mean of kernel matrix columns (for centering)
35    k_fit_cols_: Option<Array1<Float>>,
36    /// Mean of all kernel matrix elements (for centering)
37    k_fit_all_: Option<Float>,
38}
39
40impl Default for KernelPCA {
41    fn default() -> Self {
42        Self::new()
43    }
44}
45
46impl KernelPCA {
47    /// Create a new KernelPCA instance with default parameters
48    pub fn new() -> Self {
49        Self {
50            n_components: 2,
51            kernel: KernelType::Rbf { gamma: 1.0 },
52            tol: 1e-10,
53            max_iter: 1000,
54            fit_inverse_transform: false,
55            eigenvalues_: None,
56            eigenvectors_: None,
57            alphas_: None,
58            x_fit_: None,
59            k_fit_cols_: None,
60            k_fit_all_: None,
61        }
62    }
63
64    /// Set the number of components
65    pub fn with_n_components(mut self, n_components: usize) -> Self {
66        self.n_components = n_components;
67        self
68    }
69
70    /// Set the kernel function
71    pub fn with_kernel(mut self, kernel: KernelType) -> Self {
72        self.kernel = kernel;
73        self
74    }
75
76    /// Set the tolerance for eigenvalue computation
77    pub fn with_tol(mut self, tol: Float) -> Self {
78        self.tol = tol;
79        self
80    }
81
82    /// Set the maximum number of iterations
83    pub fn with_max_iter(mut self, max_iter: usize) -> Self {
84        self.max_iter = max_iter;
85        self
86    }
87
88    /// Enable fitting of inverse transform
89    pub fn with_fit_inverse_transform(mut self, fit_inverse_transform: bool) -> Self {
90        self.fit_inverse_transform = fit_inverse_transform;
91        self
92    }
93
94    /// Center the kernel matrix
95    fn center_kernel_matrix(
96        &self,
97        k: &Array2<Float>,
98        k_train: Option<&Array2<Float>>,
99    ) -> Array2<Float> {
100        let n_test = k.nrows();
101        let n_train = k.ncols();
102
103        if let (Some(k_fit_cols), Some(k_fit_all)) = (&self.k_fit_cols_, &self.k_fit_all_) {
104            // For transform: center using training statistics
105            let mut k_centered = k.clone();
106
107            // Subtract column means from training data
108            for i in 0..n_test {
109                for j in 0..n_train {
110                    k_centered[[i, j]] -= k_fit_cols[j];
111                }
112            }
113
114            // Subtract column means (compute for test data)
115            let test_row_means = k
116                .mean_axis(Axis(1))
117                .expect("mean should not fail on non-empty array");
118            for i in 0..n_test {
119                for j in 0..n_train {
120                    k_centered[[i, j]] -= test_row_means[i];
121                }
122            }
123
124            // Add back grand mean
125            k_centered + *k_fit_all
126        } else if let Some(k_train) = k_train {
127            // For fit: compute centering statistics
128            let n = k_train.nrows();
129
130            // Compute row means
131            let row_means = k_train
132                .mean_axis(Axis(1))
133                .expect("mean should not fail on non-empty array");
134
135            // Compute column means
136            let col_means = k_train
137                .mean_axis(Axis(0))
138                .expect("mean should not fail on non-empty array");
139
140            // Compute grand mean
141            let grand_mean = k_train
142                .mean()
143                .expect("mean should not fail on non-empty array");
144
145            // Center the matrix
146            let mut k_centered = Array2::<Float>::zeros((n, n));
147            for i in 0..n {
148                for j in 0..n {
149                    k_centered[[i, j]] = k_train[[i, j]] - row_means[i] - col_means[j] + grand_mean;
150                }
151            }
152
153            k_centered
154        } else {
155            // Fallback: no centering
156            k.clone()
157        }
158    }
159
160    /// Fit the kernel PCA model to the data
161    pub fn fit(&mut self, x: &Array2<Float>) -> Result<&mut Self, SklearsError> {
162        let n_samples = x.nrows();
163
164        if self.n_components > n_samples {
165            return Err(SklearsError::InvalidInput(format!(
166                "n_components ({}) cannot be larger than n_samples ({})",
167                self.n_components, n_samples
168            )));
169        }
170
171        // Create kernel instance
172        let kernel = create_kernel(self.kernel.clone())?;
173
174        // Compute kernel matrix
175        let k = kernel.compute_matrix(x, x);
176
177        // Store training statistics for centering
178        self.k_fit_cols_ = Some(
179            k.mean_axis(Axis(0))
180                .expect("mean should not fail on non-empty array"),
181        ); // Column means
182        self.k_fit_all_ = Some(k.mean().expect("mean should not fail on non-empty array"));
183
184        // Center the kernel matrix
185        let k_centered = self.center_kernel_matrix(&k, Some(&k));
186
187        // Compute eigenvalue decomposition using scirs2-linalg
188        // eigh returns (eigenvalues, eigenvectors) for symmetric matrices
189        let (eigenvalues, eigenvectors) =
190            eigh(&k_centered.view(), scirs2_linalg::compat::UPLO::Lower).map_err(|e| {
191                SklearsError::NumericalError(format!("Failed to compute eigendecomposition: {}", e))
192            })?;
193
194        // Sort eigenvalues and eigenvectors in descending order
195        let mut indices: Vec<usize> = (0..n_samples).collect();
196        indices.sort_by(|&i, &j| {
197            eigenvalues[j]
198                .partial_cmp(&eigenvalues[i])
199                .unwrap_or(std::cmp::Ordering::Equal)
200        });
201
202        // Reorder based on sorted indices
203        let sorted_eigenvalues = Array1::from_shape_fn(n_samples, |i| eigenvalues[indices[i]]);
204        let mut sorted_eigenvectors = Array2::zeros((n_samples, n_samples));
205        for i in 0..n_samples {
206            for j in 0..n_samples {
207                sorted_eigenvectors[[j, i]] = eigenvectors[[j, indices[i]]];
208            }
209        }
210
211        // Keep only positive eigenvalues above tolerance
212        let mut n_valid = 0;
213        for i in 0..n_samples {
214            if sorted_eigenvalues[i] > self.tol {
215                n_valid += 1;
216            } else {
217                break;
218            }
219        }
220
221        // Limit to requested number of components
222        let n_keep = std::cmp::min(self.n_components, n_valid);
223
224        // Store results
225        self.eigenvalues_ = Some(sorted_eigenvalues.slice(s![..n_keep]).to_owned());
226        self.eigenvectors_ = Some(sorted_eigenvectors.slice(s![.., ..n_keep]).to_owned());
227
228        // Normalize eigenvectors by square root of eigenvalues
229        let mut alphas = sorted_eigenvectors.slice(s![.., ..n_keep]).to_owned();
230        for i in 0..n_keep {
231            let norm = sorted_eigenvalues[i].sqrt();
232            if norm > self.tol {
233                alphas.column_mut(i).mapv_inplace(|x| x / norm);
234            }
235        }
236        self.alphas_ = Some(alphas);
237
238        // Store training data if inverse transform is needed
239        if self.fit_inverse_transform {
240            self.x_fit_ = Some(x.clone());
241        }
242
243        Ok(self)
244    }
245
246    /// Transform data to the kernel PCA space
247    pub fn transform(&self, x: &Array2<Float>) -> Result<Array2<Float>, SklearsError> {
248        let alphas = self
249            .alphas_
250            .as_ref()
251            .ok_or_else(|| SklearsError::NotFitted {
252                operation: "transform".to_string(),
253            })?;
254
255        let x_fit = self
256            .x_fit_
257            .as_ref()
258            .ok_or_else(|| SklearsError::NotFitted {
259                operation: "transform with stored training data".to_string(),
260            })?;
261
262        let n_test = x.nrows();
263        let n_train = x_fit.nrows();
264        let n_components = alphas.ncols();
265
266        // Create kernel instance
267        let kernel = create_kernel(self.kernel.clone())?;
268
269        // Compute kernel matrix between test and training data (shape: n_test x n_train)
270        let k = kernel.compute_matrix(x, x_fit);
271
272        // Center the kernel matrix using training statistics
273        let k_fit_cols = self
274            .k_fit_cols_
275            .as_ref()
276            .ok_or_else(|| SklearsError::NotFitted {
277                operation: "transform with centering statistics".to_string(),
278            })?;
279        let k_fit_all = self
280            .k_fit_all_
281            .as_ref()
282            .ok_or_else(|| SklearsError::NotFitted {
283                operation: "transform with centering statistics".to_string(),
284            })?;
285
286        // Center the kernel matrix
287        let mut k_centered = Array2::zeros((n_test, n_train));
288        for i in 0..n_test {
289            // Compute row mean for test sample i
290            let test_row_mean = k
291                .row(i)
292                .mean()
293                .expect("mean should not fail on non-empty array");
294
295            for j in 0..n_train {
296                // Apply centering: K_ij - mean_j(train) - mean_i(test) + grand_mean(train)
297                k_centered[[i, j]] = k[[i, j]] - k_fit_cols[j] - test_row_mean + *k_fit_all;
298            }
299        }
300
301        // Project onto principal components
302        let mut x_transformed = Array2::zeros((n_test, n_components));
303        for i in 0..n_test {
304            for j in 0..n_components {
305                x_transformed[[i, j]] = k_centered.row(i).dot(&alphas.column(j));
306            }
307        }
308
309        Ok(x_transformed)
310    }
311
312    /// Fit the model and transform the data
313    pub fn fit_transform(&mut self, x: &Array2<Float>) -> Result<Array2<Float>, SklearsError> {
314        self.fit(x)?;
315
316        let alphas = self
317            .alphas_
318            .as_ref()
319            .expect("alphas_ not available - model not fitted");
320        let n_samples = x.nrows();
321        let n_components = alphas.ncols();
322
323        // Create kernel instance
324        let kernel = create_kernel(self.kernel.clone())?;
325
326        // Compute kernel matrix
327        let k = kernel.compute_matrix(x, x);
328
329        // Center the kernel matrix
330        let k_centered = self.center_kernel_matrix(&k, Some(&k));
331
332        // Project onto principal components
333        let mut x_transformed = Array2::zeros((n_samples, n_components));
334        for i in 0..n_samples {
335            for j in 0..n_components {
336                x_transformed[[i, j]] = k_centered.row(i).dot(&alphas.column(j));
337            }
338        }
339
340        Ok(x_transformed)
341    }
342
343    /// Get the eigenvalues
344    pub fn eigenvalues(&self) -> Option<&Array1<Float>> {
345        self.eigenvalues_.as_ref()
346    }
347
348    /// Get the eigenvectors
349    pub fn eigenvectors(&self) -> Option<&Array2<Float>> {
350        self.eigenvectors_.as_ref()
351    }
352
353    /// Get explained variance ratio
354    pub fn explained_variance_ratio(&self) -> Option<Array1<Float>> {
355        self.eigenvalues_.as_ref().map(|eigenvals| {
356            let total_var = eigenvals.sum();
357            eigenvals.mapv(|x| x / total_var)
358        })
359    }
360}
361
362/// Builder pattern for KernelPCA
363pub struct KernelPCABuilder {
364    kernel_pca: KernelPCA,
365}
366
367impl KernelPCABuilder {
368    pub fn new() -> Self {
369        Self {
370            kernel_pca: KernelPCA::new(),
371        }
372    }
373
374    pub fn n_components(mut self, n_components: usize) -> Self {
375        self.kernel_pca = self.kernel_pca.with_n_components(n_components);
376        self
377    }
378
379    pub fn kernel(mut self, kernel: KernelType) -> Self {
380        self.kernel_pca = self.kernel_pca.with_kernel(kernel);
381        self
382    }
383
384    pub fn tol(mut self, tol: Float) -> Self {
385        self.kernel_pca = self.kernel_pca.with_tol(tol);
386        self
387    }
388
389    pub fn max_iter(mut self, max_iter: usize) -> Self {
390        self.kernel_pca = self.kernel_pca.with_max_iter(max_iter);
391        self
392    }
393
394    pub fn fit_inverse_transform(mut self, fit_inverse_transform: bool) -> Self {
395        self.kernel_pca = self
396            .kernel_pca
397            .with_fit_inverse_transform(fit_inverse_transform);
398        self
399    }
400
401    pub fn build(self) -> KernelPCA {
402        self.kernel_pca
403    }
404}
405
406impl Default for KernelPCABuilder {
407    fn default() -> Self {
408        Self::new()
409    }
410}
411
412#[allow(non_snake_case)]
413#[cfg(test)]
414mod tests {
415    use super::*;
416    use approx::assert_abs_diff_eq;
417    use scirs2_core::ndarray::Array;
418
419    #[test]
420    fn test_kernel_pca_basic() {
421        // Create simple 2D data that should be separable with kernel PCA
422        let x = Array::from_shape_vec(
423            (6, 2),
424            vec![
425                0.0, 0.0, 1.0, 1.0, 2.0, 2.0, -1.0, -1.0, -2.0, -2.0, 0.5, -0.5,
426            ],
427        )
428        .expect("operation should succeed");
429
430        let mut kpca = KernelPCA::new()
431            .with_n_components(2)
432            .with_kernel(KernelType::Rbf { gamma: 1.0 });
433
434        let result = kpca.fit_transform(&x);
435        assert!(result.is_ok());
436
437        let transformed = result.expect("operation should succeed");
438        assert_eq!(transformed.shape(), &[6, 2]);
439
440        // Check that eigenvalues are available
441        assert!(kpca.eigenvalues().is_some());
442        let eigenvals = kpca.eigenvalues().expect("operation should succeed");
443        assert_eq!(eigenvals.len(), 2);
444
445        // Eigenvalues should be positive and in descending order
446        for i in 0..eigenvals.len() {
447            assert!(eigenvals[i] > 0.0);
448            if i > 0 {
449                assert!(eigenvals[i - 1] >= eigenvals[i]);
450            }
451        }
452    }
453
454    #[test]
455    fn test_kernel_pca_polynomial_kernel() {
456        let x = Array::from_shape_vec((4, 2), vec![1.0, 1.0, 2.0, 2.0, -1.0, -1.0, -2.0, -2.0])
457            .expect("operation should succeed");
458
459        let mut kpca = KernelPCA::new()
460            .with_n_components(2)
461            .with_kernel(KernelType::Polynomial {
462                gamma: 1.0,
463                degree: 2.0,
464                coef0: 1.0,
465            });
466
467        let result = kpca.fit_transform(&x);
468        assert!(result.is_ok());
469
470        let transformed = result.expect("operation should succeed");
471        assert_eq!(transformed.shape(), &[4, 2]);
472    }
473
474    #[test]
475    fn test_kernel_pca_transform_new_data() {
476        let x_train = Array::from_shape_vec((4, 2), vec![0.0, 0.0, 1.0, 1.0, -1.0, -1.0, 0.0, 1.0])
477            .expect("array shape mismatch");
478
479        let x_test = Array::from_shape_vec((2, 2), vec![0.5, 0.5, -0.5, -0.5])
480            .expect("array shape mismatch");
481
482        let mut kpca = KernelPCA::new()
483            .with_n_components(2)
484            .with_kernel(KernelType::Rbf { gamma: 1.0 })
485            .with_fit_inverse_transform(true);
486
487        // Fit on training data
488        kpca.fit(&x_train).expect("model fitting should succeed");
489
490        // Transform test data
491        let result = kpca.transform(&x_test);
492        assert!(result.is_ok());
493
494        let transformed = result.expect("operation should succeed");
495        assert_eq!(transformed.shape(), &[2, 2]);
496    }
497
498    #[test]
499    fn test_kernel_pca_builder() {
500        let kpca = KernelPCABuilder::new()
501            .n_components(3)
502            .kernel(KernelType::Linear)
503            .tol(1e-12)
504            .max_iter(500)
505            .fit_inverse_transform(true)
506            .build();
507
508        assert_eq!(kpca.n_components, 3);
509        assert_eq!(kpca.tol, 1e-12);
510        assert_eq!(kpca.max_iter, 500);
511        assert!(kpca.fit_inverse_transform);
512    }
513
514    #[test]
515    fn test_kernel_pca_explained_variance() {
516        let x = Array::from_shape_vec(
517            (5, 2),
518            vec![0.0, 0.0, 1.0, 1.0, 2.0, 2.0, -1.0, -1.0, -2.0, -2.0],
519        )
520        .expect("operation should succeed");
521
522        let mut kpca = KernelPCA::new()
523            .with_n_components(2)
524            .with_kernel(KernelType::Rbf { gamma: 1.0 });
525
526        kpca.fit_transform(&x)
527            .expect("fit_transform should succeed");
528
529        let explained_var_ratio = kpca.explained_variance_ratio();
530        assert!(explained_var_ratio.is_some());
531
532        let ratios = explained_var_ratio.expect("operation should succeed");
533        let sum: Float = ratios.sum();
534        assert_abs_diff_eq!(sum, 1.0, epsilon = 1e-10);
535
536        // First component should explain more variance than the second
537        assert!(ratios[0] >= ratios[1]);
538    }
539
540    #[test]
541    fn test_kernel_pca_invalid_n_components() {
542        let x = Array::from_shape_vec((3, 2), vec![0.0, 0.0, 1.0, 1.0, -1.0, -1.0])
543            .expect("array shape mismatch");
544
545        let mut kpca = KernelPCA::new().with_n_components(5); // More components than samples
546
547        let result = kpca.fit(&x);
548        assert!(result.is_err());
549    }
550}