1use crate::error::{KernelError, Result};
21use crate::types::Kernel;
22
23#[derive(Debug, Clone)]
25pub struct TaskInput {
26 pub features: Vec<f64>,
28 pub task: usize,
30}
31
32impl TaskInput {
33 pub fn new(features: Vec<f64>, task: usize) -> Self {
35 Self { features, task }
36 }
37
38 pub fn from_slice(features: &[f64], task: usize) -> Self {
40 Self {
41 features: features.to_vec(),
42 task,
43 }
44 }
45}
46
47#[derive(Debug, Clone)]
49pub struct MultiTaskConfig {
50 pub num_tasks: usize,
52 pub normalize: bool,
54}
55
56impl MultiTaskConfig {
57 pub fn new(num_tasks: usize) -> Self {
59 Self {
60 num_tasks,
61 normalize: false,
62 }
63 }
64
65 pub fn with_normalization(mut self) -> Self {
67 self.normalize = true;
68 self
69 }
70}
71
72#[derive(Debug, Clone)]
77pub struct IndexKernel {
78 task_covariance: Vec<Vec<f64>>,
80 num_tasks: usize,
82}
83
84impl IndexKernel {
85 pub fn new(task_covariance: Vec<Vec<f64>>) -> Result<Self> {
89 let num_tasks = task_covariance.len();
90 if num_tasks == 0 {
91 return Err(KernelError::InvalidParameter {
92 parameter: "task_covariance".to_string(),
93 value: "empty".to_string(),
94 reason: "must have at least one task".to_string(),
95 });
96 }
97
98 for (i, row) in task_covariance.iter().enumerate() {
100 if row.len() != num_tasks {
101 return Err(KernelError::InvalidParameter {
102 parameter: "task_covariance".to_string(),
103 value: format!("row {} has {} elements", i, row.len()),
104 reason: format!("expected {} elements (square matrix)", num_tasks),
105 });
106 }
107 }
108
109 Ok(Self {
110 task_covariance,
111 num_tasks,
112 })
113 }
114
115 pub fn identity(num_tasks: usize) -> Result<Self> {
117 let mut cov = vec![vec![0.0; num_tasks]; num_tasks];
118 for (i, row) in cov.iter_mut().enumerate() {
119 row[i] = 1.0;
120 }
121 Self::new(cov)
122 }
123
124 pub fn uniform(num_tasks: usize, correlation: f64) -> Result<Self> {
126 if !(0.0..=1.0).contains(&correlation) {
127 return Err(KernelError::InvalidParameter {
128 parameter: "correlation".to_string(),
129 value: correlation.to_string(),
130 reason: "must be in [0, 1]".to_string(),
131 });
132 }
133
134 let mut cov = vec![vec![correlation; num_tasks]; num_tasks];
135 for (i, row) in cov.iter_mut().enumerate() {
136 row[i] = 1.0;
137 }
138 Self::new(cov)
139 }
140
141 pub fn get_task_covariance(&self, task_i: usize, task_j: usize) -> Result<f64> {
143 if task_i >= self.num_tasks || task_j >= self.num_tasks {
144 return Err(KernelError::ComputationError(format!(
145 "Task index out of bounds: ({}, {}) for {} tasks",
146 task_i, task_j, self.num_tasks
147 )));
148 }
149 Ok(self.task_covariance[task_i][task_j])
150 }
151
152 pub fn num_tasks(&self) -> usize {
154 self.num_tasks
155 }
156
157 pub fn covariance_matrix(&self) -> &Vec<Vec<f64>> {
159 &self.task_covariance
160 }
161}
162
163pub struct ICMKernel {
174 base_kernel: Box<dyn Kernel>,
176 task_covariance: Vec<Vec<f64>>,
178 num_tasks: usize,
180}
181
182impl ICMKernel {
183 pub fn new(base_kernel: Box<dyn Kernel>, task_covariance: Vec<Vec<f64>>) -> Result<Self> {
189 let num_tasks = task_covariance.len();
190 if num_tasks == 0 {
191 return Err(KernelError::InvalidParameter {
192 parameter: "task_covariance".to_string(),
193 value: "empty".to_string(),
194 reason: "must have at least one task".to_string(),
195 });
196 }
197
198 for (i, row) in task_covariance.iter().enumerate() {
200 if row.len() != num_tasks {
201 return Err(KernelError::InvalidParameter {
202 parameter: "task_covariance".to_string(),
203 value: format!("row {} has {} elements", i, row.len()),
204 reason: format!("expected {} elements", num_tasks),
205 });
206 }
207 }
208
209 Ok(Self {
210 base_kernel,
211 task_covariance,
212 num_tasks,
213 })
214 }
215
216 pub fn independent(base_kernel: Box<dyn Kernel>, num_tasks: usize) -> Result<Self> {
218 let mut cov = vec![vec![0.0; num_tasks]; num_tasks];
219 for (i, row) in cov.iter_mut().enumerate() {
220 row[i] = 1.0;
221 }
222 Self::new(base_kernel, cov)
223 }
224
225 pub fn uniform(
227 base_kernel: Box<dyn Kernel>,
228 num_tasks: usize,
229 correlation: f64,
230 ) -> Result<Self> {
231 if !(0.0..=1.0).contains(&correlation) {
232 return Err(KernelError::InvalidParameter {
233 parameter: "correlation".to_string(),
234 value: correlation.to_string(),
235 reason: "must be in [0, 1]".to_string(),
236 });
237 }
238
239 let mut cov = vec![vec![correlation; num_tasks]; num_tasks];
240 for (i, row) in cov.iter_mut().enumerate() {
241 row[i] = 1.0;
242 }
243 Self::new(base_kernel, cov)
244 }
245
246 pub fn from_rank1(base_kernel: Box<dyn Kernel>, task_variances: Vec<f64>) -> Result<Self> {
250 let num_tasks = task_variances.len();
251 let mut cov = vec![vec![0.0; num_tasks]; num_tasks];
252 for i in 0..num_tasks {
253 for j in 0..num_tasks {
254 cov[i][j] = task_variances[i].sqrt() * task_variances[j].sqrt();
255 }
256 }
257 Self::new(base_kernel, cov)
258 }
259
260 pub fn compute_tasks(&self, x: &TaskInput, y: &TaskInput) -> Result<f64> {
262 if x.task >= self.num_tasks || y.task >= self.num_tasks {
263 return Err(KernelError::ComputationError(format!(
264 "Task index out of bounds: ({}, {}) for {} tasks",
265 x.task, y.task, self.num_tasks
266 )));
267 }
268
269 let k_features = self.base_kernel.compute(&x.features, &y.features)?;
270 let b_tasks = self.task_covariance[x.task][y.task];
271
272 Ok(b_tasks * k_features)
273 }
274
275 pub fn num_tasks(&self) -> usize {
277 self.num_tasks
278 }
279
280 pub fn task_covariance(&self) -> &Vec<Vec<f64>> {
282 &self.task_covariance
283 }
284
285 pub fn compute_task_matrix(&self, inputs: &[TaskInput]) -> Result<Vec<Vec<f64>>> {
287 let n = inputs.len();
288 let mut matrix = vec![vec![0.0; n]; n];
289
290 for i in 0..n {
291 for j in i..n {
292 let k = self.compute_tasks(&inputs[i], &inputs[j])?;
293 matrix[i][j] = k;
294 matrix[j][i] = k;
295 }
296 }
297
298 Ok(matrix)
299 }
300}
301
302struct LMCComponent {
304 kernel: Box<dyn Kernel>,
306 task_covariance: Vec<Vec<f64>>,
308}
309
310pub struct LMCKernel {
322 components: Vec<LMCComponent>,
324 num_tasks: usize,
326}
327
328impl LMCKernel {
329 pub fn new(num_tasks: usize) -> Self {
331 Self {
332 components: Vec::new(),
333 num_tasks,
334 }
335 }
336
337 pub fn add_component(
339 &mut self,
340 kernel: Box<dyn Kernel>,
341 task_covariance: Vec<Vec<f64>>,
342 ) -> Result<()> {
343 if task_covariance.len() != self.num_tasks {
345 return Err(KernelError::InvalidParameter {
346 parameter: "task_covariance".to_string(),
347 value: format!("{} rows", task_covariance.len()),
348 reason: format!("expected {} rows", self.num_tasks),
349 });
350 }
351
352 for (i, row) in task_covariance.iter().enumerate() {
353 if row.len() != self.num_tasks {
354 return Err(KernelError::InvalidParameter {
355 parameter: "task_covariance".to_string(),
356 value: format!("row {} has {} elements", i, row.len()),
357 reason: format!("expected {} elements", self.num_tasks),
358 });
359 }
360 }
361
362 self.components.push(LMCComponent {
363 kernel,
364 task_covariance,
365 });
366
367 Ok(())
368 }
369
370 pub fn compute_tasks(&self, x: &TaskInput, y: &TaskInput) -> Result<f64> {
372 if x.task >= self.num_tasks || y.task >= self.num_tasks {
373 return Err(KernelError::ComputationError(format!(
374 "Task index out of bounds: ({}, {}) for {} tasks",
375 x.task, y.task, self.num_tasks
376 )));
377 }
378
379 let mut result = 0.0;
380 for component in &self.components {
381 let k_features = component.kernel.compute(&x.features, &y.features)?;
382 let b_tasks = component.task_covariance[x.task][y.task];
383 result += b_tasks * k_features;
384 }
385
386 Ok(result)
387 }
388
389 pub fn num_components(&self) -> usize {
391 self.components.len()
392 }
393
394 pub fn num_tasks(&self) -> usize {
396 self.num_tasks
397 }
398
399 pub fn compute_task_matrix(&self, inputs: &[TaskInput]) -> Result<Vec<Vec<f64>>> {
401 let n = inputs.len();
402 let mut matrix = vec![vec![0.0; n]; n];
403
404 for i in 0..n {
405 for j in i..n {
406 let k = self.compute_tasks(&inputs[i], &inputs[j])?;
407 matrix[i][j] = k;
408 matrix[j][i] = k;
409 }
410 }
411
412 Ok(matrix)
413 }
414}
415
416pub struct ICMKernelWrapper {
420 inner: ICMKernel,
421}
422
423impl ICMKernelWrapper {
424 pub fn new(inner: ICMKernel) -> Self {
426 Self { inner }
427 }
428
429 pub fn inner(&self) -> &ICMKernel {
431 &self.inner
432 }
433}
434
435impl Kernel for ICMKernelWrapper {
436 fn compute(&self, x: &[f64], y: &[f64]) -> Result<f64> {
440 if x.is_empty() || y.is_empty() {
441 return Err(KernelError::ComputationError(
442 "Input must have at least task index".to_string(),
443 ));
444 }
445
446 let task_x = x[0] as usize;
447 let task_y = y[0] as usize;
448 let features_x = &x[1..];
449 let features_y = &y[1..];
450
451 let input_x = TaskInput::from_slice(features_x, task_x);
452 let input_y = TaskInput::from_slice(features_y, task_y);
453
454 self.inner.compute_tasks(&input_x, &input_y)
455 }
456
457 fn name(&self) -> &str {
458 "ICM"
459 }
460}
461
462pub struct LMCKernelWrapper {
464 inner: LMCKernel,
465}
466
467impl LMCKernelWrapper {
468 pub fn new(inner: LMCKernel) -> Self {
470 Self { inner }
471 }
472
473 pub fn inner(&self) -> &LMCKernel {
475 &self.inner
476 }
477}
478
479impl Kernel for LMCKernelWrapper {
480 fn compute(&self, x: &[f64], y: &[f64]) -> Result<f64> {
482 if x.is_empty() || y.is_empty() {
483 return Err(KernelError::ComputationError(
484 "Input must have at least task index".to_string(),
485 ));
486 }
487
488 let task_x = x[0] as usize;
489 let task_y = y[0] as usize;
490 let features_x = &x[1..];
491 let features_y = &y[1..];
492
493 let input_x = TaskInput::from_slice(features_x, task_x);
494 let input_y = TaskInput::from_slice(features_y, task_y);
495
496 self.inner.compute_tasks(&input_x, &input_y)
497 }
498
499 fn name(&self) -> &str {
500 "LMC"
501 }
502}
503
504pub struct HadamardTaskKernel {
510 kernels: Vec<ICMKernel>,
512}
513
514impl HadamardTaskKernel {
515 pub fn new() -> Self {
517 Self {
518 kernels: Vec::new(),
519 }
520 }
521
522 pub fn add_kernel(&mut self, kernel: ICMKernel) -> Result<()> {
524 if !self.kernels.is_empty() && kernel.num_tasks() != self.kernels[0].num_tasks() {
525 return Err(KernelError::InvalidParameter {
526 parameter: "num_tasks".to_string(),
527 value: kernel.num_tasks().to_string(),
528 reason: format!("expected {}", self.kernels[0].num_tasks()),
529 });
530 }
531 self.kernels.push(kernel);
532 Ok(())
533 }
534
535 pub fn compute_tasks(&self, x: &TaskInput, y: &TaskInput) -> Result<f64> {
537 if self.kernels.is_empty() {
538 return Err(KernelError::ComputationError(
539 "No component kernels added".to_string(),
540 ));
541 }
542
543 let mut result = 1.0;
544 for kernel in &self.kernels {
545 result *= kernel.compute_tasks(x, y)?;
546 }
547 Ok(result)
548 }
549
550 pub fn num_tasks(&self) -> Option<usize> {
552 self.kernels.first().map(|k| k.num_tasks())
553 }
554}
555
556impl Default for HadamardTaskKernel {
557 fn default() -> Self {
558 Self::new()
559 }
560}
561
562pub struct MultiTaskKernelBuilder {
564 num_tasks: usize,
565 base_kernels: Vec<Box<dyn Kernel>>,
566 task_covariances: Vec<Vec<Vec<f64>>>,
567}
568
569impl MultiTaskKernelBuilder {
570 pub fn new(num_tasks: usize) -> Self {
572 Self {
573 num_tasks,
574 base_kernels: Vec::new(),
575 task_covariances: Vec::new(),
576 }
577 }
578
579 pub fn add_component(
581 mut self,
582 kernel: Box<dyn Kernel>,
583 task_covariance: Vec<Vec<f64>>,
584 ) -> Self {
585 self.base_kernels.push(kernel);
586 self.task_covariances.push(task_covariance);
587 self
588 }
589
590 pub fn build_icm(self) -> Result<ICMKernel> {
592 if self.base_kernels.len() != 1 {
593 return Err(KernelError::InvalidParameter {
594 parameter: "components".to_string(),
595 value: self.base_kernels.len().to_string(),
596 reason: "ICM requires exactly one component".to_string(),
597 });
598 }
599
600 let kernel = self
601 .base_kernels
602 .into_iter()
603 .next()
604 .expect("validated non-empty");
605 let cov = self
606 .task_covariances
607 .into_iter()
608 .next()
609 .expect("validated non-empty");
610 ICMKernel::new(kernel, cov)
611 }
612
613 pub fn build_lmc(self) -> Result<LMCKernel> {
615 let mut lmc = LMCKernel::new(self.num_tasks);
616
617 for (kernel, cov) in self.base_kernels.into_iter().zip(self.task_covariances) {
618 lmc.add_component(kernel, cov)?;
619 }
620
621 Ok(lmc)
622 }
623}
624
625#[cfg(test)]
626#[allow(clippy::needless_range_loop)]
627mod tests {
628 use super::*;
629 use crate::{LinearKernel, RbfKernel, RbfKernelConfig};
630
631 #[test]
634 fn test_index_kernel_basic() {
635 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
636 let kernel = IndexKernel::new(cov).expect("unwrap");
637
638 assert_eq!(kernel.num_tasks(), 2);
639 assert!((kernel.get_task_covariance(0, 1).expect("unwrap") - 0.5).abs() < 1e-10);
640 assert!((kernel.get_task_covariance(1, 1).expect("unwrap") - 1.0).abs() < 1e-10);
641 }
642
643 #[test]
644 fn test_index_kernel_identity() {
645 let kernel = IndexKernel::identity(3).expect("unwrap");
646
647 assert!((kernel.get_task_covariance(0, 0).expect("unwrap") - 1.0).abs() < 1e-10);
648 assert!((kernel.get_task_covariance(0, 1).expect("unwrap")).abs() < 1e-10);
649 assert!((kernel.get_task_covariance(1, 2).expect("unwrap")).abs() < 1e-10);
650 }
651
652 #[test]
653 fn test_index_kernel_uniform() {
654 let kernel = IndexKernel::uniform(3, 0.5).expect("unwrap");
655
656 assert!((kernel.get_task_covariance(0, 0).expect("unwrap") - 1.0).abs() < 1e-10);
657 assert!((kernel.get_task_covariance(0, 1).expect("unwrap") - 0.5).abs() < 1e-10);
658 assert!((kernel.get_task_covariance(1, 2).expect("unwrap") - 0.5).abs() < 1e-10);
659 }
660
661 #[test]
662 fn test_index_kernel_invalid() {
663 let result = IndexKernel::new(vec![]);
665 assert!(result.is_err());
666
667 let result = IndexKernel::new(vec![vec![1.0, 0.5]]);
669 assert!(result.is_err());
670
671 let result = IndexKernel::uniform(3, 1.5);
673 assert!(result.is_err());
674 }
675
676 #[test]
677 fn test_index_kernel_out_of_bounds() {
678 let kernel = IndexKernel::identity(2).expect("unwrap");
679 assert!(kernel.get_task_covariance(2, 0).is_err());
680 }
681
682 #[test]
685 fn test_icm_kernel_basic() {
686 let base = LinearKernel::new();
687 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
688 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
689
690 assert_eq!(icm.num_tasks(), 2);
691 }
692
693 #[test]
694 fn test_icm_kernel_compute() {
695 let base = LinearKernel::new();
696 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
697 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
698
699 let x = TaskInput::new(vec![1.0, 2.0], 0);
700 let y = TaskInput::new(vec![3.0, 4.0], 1);
701
702 let k = icm.compute_tasks(&x, &y).expect("unwrap");
703 assert!((k - 5.5).abs() < 1e-10);
707 }
708
709 #[test]
710 fn test_icm_kernel_same_task() {
711 let base = LinearKernel::new();
712 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
713 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
714
715 let x = TaskInput::new(vec![1.0, 2.0], 0);
716 let y = TaskInput::new(vec![3.0, 4.0], 0);
717
718 let k = icm.compute_tasks(&x, &y).expect("unwrap");
719 assert!((k - 11.0).abs() < 1e-10);
721 }
722
723 #[test]
724 fn test_icm_kernel_independent() {
725 let base = LinearKernel::new();
726 let icm = ICMKernel::independent(Box::new(base), 3).expect("unwrap");
727
728 let x = TaskInput::new(vec![1.0], 0);
729 let y = TaskInput::new(vec![1.0], 1);
730
731 let k = icm.compute_tasks(&x, &y).expect("unwrap");
733 assert!(k.abs() < 1e-10);
734
735 let z = TaskInput::new(vec![1.0], 0);
737 let k = icm.compute_tasks(&x, &z).expect("unwrap");
738 assert!((k - 1.0).abs() < 1e-10);
739 }
740
741 #[test]
742 fn test_icm_kernel_uniform() {
743 let base = LinearKernel::new();
744 let icm = ICMKernel::uniform(Box::new(base), 2, 0.8).expect("unwrap");
745
746 let x = TaskInput::new(vec![1.0], 0);
747 let y = TaskInput::new(vec![1.0], 1);
748
749 let k = icm.compute_tasks(&x, &y).expect("unwrap");
750 assert!((k - 0.8).abs() < 1e-10);
752 }
753
754 #[test]
755 fn test_icm_kernel_rank1() {
756 let base = LinearKernel::new();
757 let variances = vec![1.0, 4.0]; let icm = ICMKernel::from_rank1(Box::new(base), variances).expect("unwrap");
759
760 let x = TaskInput::new(vec![1.0], 0);
762 let y = TaskInput::new(vec![1.0], 1);
763
764 let k = icm.compute_tasks(&x, &y).expect("unwrap");
765 assert!((k - 2.0).abs() < 1e-10);
766 }
767
768 #[test]
769 fn test_icm_kernel_matrix() {
770 let base = LinearKernel::new();
771 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
772 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
773
774 let inputs = vec![
775 TaskInput::new(vec![1.0], 0),
776 TaskInput::new(vec![1.0], 1),
777 TaskInput::new(vec![2.0], 0),
778 ];
779
780 let matrix = icm.compute_task_matrix(&inputs).expect("unwrap");
781
782 assert_eq!(matrix.len(), 3);
783 for i in 0..3 {
785 for j in 0..3 {
786 assert!(
787 (matrix[i][j] - matrix[j][i]).abs() < 1e-10,
788 "Matrix not symmetric at ({}, {})",
789 i,
790 j
791 );
792 }
793 }
794 }
795
796 #[test]
797 fn test_icm_kernel_invalid_task() {
798 let base = LinearKernel::new();
799 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
800 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
801
802 let x = TaskInput::new(vec![1.0], 0);
803 let y = TaskInput::new(vec![1.0], 5); assert!(icm.compute_tasks(&x, &y).is_err());
806 }
807
808 #[test]
811 fn test_lmc_kernel_basic() {
812 let mut lmc = LMCKernel::new(2);
813
814 let base1 = LinearKernel::new();
815 let cov1 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
816 lmc.add_component(Box::new(base1), cov1).expect("unwrap");
817
818 assert_eq!(lmc.num_tasks(), 2);
819 assert_eq!(lmc.num_components(), 1);
820 }
821
822 #[test]
823 fn test_lmc_kernel_compute() {
824 let mut lmc = LMCKernel::new(2);
825
826 let base1 = LinearKernel::new();
828 let cov1 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
829 lmc.add_component(Box::new(base1), cov1).expect("unwrap");
830
831 let base2 = RbfKernel::new(RbfKernelConfig::new(1.0)).expect("unwrap");
833 let cov2 = vec![vec![2.0, 1.0], vec![1.0, 2.0]];
834 lmc.add_component(Box::new(base2), cov2).expect("unwrap");
835
836 let x = TaskInput::new(vec![1.0, 0.0], 0);
837 let y = TaskInput::new(vec![1.0, 0.0], 1);
838
839 let k = lmc.compute_tasks(&x, &y).expect("unwrap");
840 assert!((k - 1.5).abs() < 1e-10);
844 }
845
846 #[test]
847 fn test_lmc_kernel_matrix() {
848 let mut lmc = LMCKernel::new(2);
849
850 let base = LinearKernel::new();
851 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
852 lmc.add_component(Box::new(base), cov).expect("unwrap");
853
854 let inputs = vec![TaskInput::new(vec![1.0], 0), TaskInput::new(vec![1.0], 1)];
855
856 let matrix = lmc.compute_task_matrix(&inputs).expect("unwrap");
857 assert_eq!(matrix.len(), 2);
858
859 assert!((matrix[0][1] - matrix[1][0]).abs() < 1e-10);
861 }
862
863 #[test]
864 fn test_lmc_kernel_invalid_dimensions() {
865 let mut lmc = LMCKernel::new(2);
866
867 let base = LinearKernel::new();
868 let cov = vec![
869 vec![1.0, 0.5, 0.3],
870 vec![0.5, 1.0, 0.4],
871 vec![0.3, 0.4, 1.0],
872 ];
873
874 assert!(lmc.add_component(Box::new(base), cov).is_err());
876 }
877
878 #[test]
881 fn test_icm_wrapper() {
882 let base = LinearKernel::new();
883 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
884 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
885 let wrapper = ICMKernelWrapper::new(icm);
886
887 let x = vec![0.0, 1.0, 2.0]; let y = vec![1.0, 3.0, 4.0]; let k = wrapper.compute(&x, &y).expect("unwrap");
892 assert!((k - 5.5).abs() < 1e-10);
894
895 assert_eq!(wrapper.name(), "ICM");
896 }
897
898 #[test]
899 fn test_lmc_wrapper() {
900 let mut lmc = LMCKernel::new(2);
901 let base = LinearKernel::new();
902 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
903 lmc.add_component(Box::new(base), cov).expect("unwrap");
904
905 let wrapper = LMCKernelWrapper::new(lmc);
906
907 let x = vec![0.0, 1.0]; let y = vec![1.0, 1.0]; let k = wrapper.compute(&x, &y).expect("unwrap");
911 assert!((k - 0.5).abs() < 1e-10);
912
913 assert_eq!(wrapper.name(), "LMC");
914 }
915
916 #[test]
917 fn test_wrapper_empty_input() {
918 let base = LinearKernel::new();
919 let cov = vec![vec![1.0]];
920 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
921 let wrapper = ICMKernelWrapper::new(icm);
922
923 assert!(wrapper.compute(&[], &[0.0, 1.0]).is_err());
924 }
925
926 #[test]
929 fn test_hadamard_task_kernel() {
930 let mut hadamard = HadamardTaskKernel::new();
931
932 let base1 = LinearKernel::new();
933 let cov1 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
934 let icm1 = ICMKernel::new(Box::new(base1), cov1).expect("unwrap");
935 hadamard.add_kernel(icm1).expect("unwrap");
936
937 let base2 = LinearKernel::new();
938 let cov2 = vec![vec![2.0, 1.0], vec![1.0, 2.0]];
939 let icm2 = ICMKernel::new(Box::new(base2), cov2).expect("unwrap");
940 hadamard.add_kernel(icm2).expect("unwrap");
941
942 let x = TaskInput::new(vec![1.0], 0);
943 let y = TaskInput::new(vec![1.0], 1);
944
945 let k = hadamard.compute_tasks(&x, &y).expect("unwrap");
946 assert!((k - 0.5).abs() < 1e-10);
950 }
951
952 #[test]
953 fn test_hadamard_task_kernel_empty() {
954 let hadamard = HadamardTaskKernel::new();
955 let x = TaskInput::new(vec![1.0], 0);
956 let y = TaskInput::new(vec![1.0], 0);
957
958 assert!(hadamard.compute_tasks(&x, &y).is_err());
959 }
960
961 #[test]
962 fn test_hadamard_mismatched_tasks() {
963 let mut hadamard = HadamardTaskKernel::new();
964
965 let base1 = LinearKernel::new();
966 let cov1 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
967 let icm1 = ICMKernel::new(Box::new(base1), cov1).expect("unwrap");
968 hadamard.add_kernel(icm1).expect("unwrap");
969
970 let base2 = LinearKernel::new();
972 let cov2 = vec![
973 vec![1.0, 0.5, 0.3],
974 vec![0.5, 1.0, 0.4],
975 vec![0.3, 0.4, 1.0],
976 ];
977 let icm2 = ICMKernel::new(Box::new(base2), cov2).expect("unwrap");
978
979 assert!(hadamard.add_kernel(icm2).is_err());
980 }
981
982 #[test]
985 fn test_builder_icm() {
986 let base = LinearKernel::new();
987 let cov = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
988
989 let icm = MultiTaskKernelBuilder::new(2)
990 .add_component(Box::new(base), cov)
991 .build_icm()
992 .expect("unwrap");
993
994 assert_eq!(icm.num_tasks(), 2);
995 }
996
997 #[test]
998 fn test_builder_lmc() {
999 let base1 = LinearKernel::new();
1000 let cov1 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
1001
1002 let base2 = RbfKernel::new(RbfKernelConfig::new(1.0)).expect("unwrap");
1003 let cov2 = vec![vec![2.0, 1.0], vec![1.0, 2.0]];
1004
1005 let lmc = MultiTaskKernelBuilder::new(2)
1006 .add_component(Box::new(base1), cov1)
1007 .add_component(Box::new(base2), cov2)
1008 .build_lmc()
1009 .expect("unwrap");
1010
1011 assert_eq!(lmc.num_tasks(), 2);
1012 assert_eq!(lmc.num_components(), 2);
1013 }
1014
1015 #[test]
1016 fn test_builder_icm_wrong_components() {
1017 let base1 = LinearKernel::new();
1018 let cov1 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
1019
1020 let base2 = LinearKernel::new();
1021 let cov2 = vec![vec![1.0, 0.5], vec![0.5, 1.0]];
1022
1023 let result = MultiTaskKernelBuilder::new(2)
1024 .add_component(Box::new(base1), cov1)
1025 .add_component(Box::new(base2), cov2)
1026 .build_icm();
1027
1028 assert!(result.is_err());
1029 }
1030
1031 #[test]
1034 fn test_multitask_with_rbf() {
1035 let base = RbfKernel::new(RbfKernelConfig::new(0.5)).expect("unwrap");
1036 let cov = vec![
1037 vec![1.0, 0.8, 0.6],
1038 vec![0.8, 1.0, 0.7],
1039 vec![0.6, 0.7, 1.0],
1040 ];
1041 let icm = ICMKernel::new(Box::new(base), cov).expect("unwrap");
1042
1043 let x = TaskInput::new(vec![1.0, 2.0], 0);
1045 let k = icm.compute_tasks(&x, &x).expect("unwrap");
1046 assert!((k - 1.0).abs() < 1e-10);
1047
1048 let y = TaskInput::new(vec![1.0, 2.0], 1);
1050 let k = icm.compute_tasks(&x, &y).expect("unwrap");
1051 assert!((k - 0.8).abs() < 1e-10);
1052
1053 let z = TaskInput::new(vec![1.0, 3.0], 0);
1055 let k = icm.compute_tasks(&x, &z).expect("unwrap");
1056 assert!(k > 0.5 && k < 0.7);
1058 }
1059
1060 #[test]
1061 fn test_task_input_creation() {
1062 let input = TaskInput::new(vec![1.0, 2.0, 3.0], 0);
1063 assert_eq!(input.features, vec![1.0, 2.0, 3.0]);
1064 assert_eq!(input.task, 0);
1065
1066 let input = TaskInput::from_slice(&[4.0, 5.0], 2);
1067 assert_eq!(input.features, vec![4.0, 5.0]);
1068 assert_eq!(input.task, 2);
1069 }
1070}