1use 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#[derive(Debug)]
15pub struct KernelPCA {
16 n_components: usize,
18 kernel: KernelType,
20 tol: Float,
22 max_iter: usize,
24 fit_inverse_transform: bool,
26 eigenvalues_: Option<Array1<Float>>,
28 eigenvectors_: Option<Array2<Float>>,
30 alphas_: Option<Array2<Float>>,
32 x_fit_: Option<Array2<Float>>,
34 k_fit_cols_: Option<Array1<Float>>,
36 k_fit_all_: Option<Float>,
38}
39
40impl Default for KernelPCA {
41 fn default() -> Self {
42 Self::new()
43 }
44}
45
46impl KernelPCA {
47 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 pub fn with_n_components(mut self, n_components: usize) -> Self {
66 self.n_components = n_components;
67 self
68 }
69
70 pub fn with_kernel(mut self, kernel: KernelType) -> Self {
72 self.kernel = kernel;
73 self
74 }
75
76 pub fn with_tol(mut self, tol: Float) -> Self {
78 self.tol = tol;
79 self
80 }
81
82 pub fn with_max_iter(mut self, max_iter: usize) -> Self {
84 self.max_iter = max_iter;
85 self
86 }
87
88 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 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 let mut k_centered = k.clone();
106
107 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 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 k_centered + *k_fit_all
126 } else if let Some(k_train) = k_train {
127 let n = k_train.nrows();
129
130 let row_means = k_train
132 .mean_axis(Axis(1))
133 .expect("mean should not fail on non-empty array");
134
135 let col_means = k_train
137 .mean_axis(Axis(0))
138 .expect("mean should not fail on non-empty array");
139
140 let grand_mean = k_train
142 .mean()
143 .expect("mean should not fail on non-empty array");
144
145 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 k.clone()
157 }
158 }
159
160 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 let kernel = create_kernel(self.kernel.clone())?;
173
174 let k = kernel.compute_matrix(x, x);
176
177 self.k_fit_cols_ = Some(
179 k.mean_axis(Axis(0))
180 .expect("mean should not fail on non-empty array"),
181 ); self.k_fit_all_ = Some(k.mean().expect("mean should not fail on non-empty array"));
183
184 let k_centered = self.center_kernel_matrix(&k, Some(&k));
186
187 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 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 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 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 let n_keep = std::cmp::min(self.n_components, n_valid);
223
224 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 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 if self.fit_inverse_transform {
240 self.x_fit_ = Some(x.clone());
241 }
242
243 Ok(self)
244 }
245
246 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 let kernel = create_kernel(self.kernel.clone())?;
268
269 let k = kernel.compute_matrix(x, x_fit);
271
272 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 let mut k_centered = Array2::zeros((n_test, n_train));
288 for i in 0..n_test {
289 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 k_centered[[i, j]] = k[[i, j]] - k_fit_cols[j] - test_row_mean + *k_fit_all;
298 }
299 }
300
301 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 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 let kernel = create_kernel(self.kernel.clone())?;
325
326 let k = kernel.compute_matrix(x, x);
328
329 let k_centered = self.center_kernel_matrix(&k, Some(&k));
331
332 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 pub fn eigenvalues(&self) -> Option<&Array1<Float>> {
345 self.eigenvalues_.as_ref()
346 }
347
348 pub fn eigenvectors(&self) -> Option<&Array2<Float>> {
350 self.eigenvectors_.as_ref()
351 }
352
353 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
362pub 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 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 assert!(kpca.eigenvalues().is_some());
442 let eigenvals = kpca.eigenvalues().expect("operation should succeed");
443 assert_eq!(eigenvals.len(), 2);
444
445 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 kpca.fit(&x_train).expect("model fitting should succeed");
489
490 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 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); let result = kpca.fit(&x);
548 assert!(result.is_err());
549 }
550}