Skip to main content

radiate_gp/regression/
data.rs

1use radiate_core::random_provider;
2use radiate_utils::{Float, Matrix};
3
4#[derive(Default, Clone)]
5pub struct DataSet<T> {
6    features: Matrix<T>,
7    labels: Matrix<T>,
8}
9
10impl<T> DataSet<T> {
11    pub fn new(inputs: Vec<Vec<T>>, outputs: Vec<Vec<T>>) -> Self {
12        let features = Matrix::from(inputs);
13        let labels = Matrix::from(outputs);
14
15        DataSet { features, labels }
16    }
17
18    pub fn iter(&self) -> impl Iterator<Item = (&[T], &[T])> {
19        self.features.iter().zip(self.labels.iter())
20    }
21
22    pub fn len(&self) -> usize {
23        self.features.rows()
24    }
25
26    pub fn is_empty(&self) -> bool {
27        self.features.is_empty()
28    }
29
30    pub fn shape(&self) -> (usize, usize, usize) {
31        let num_samples = self.features.rows();
32        let input_dim = if num_samples > 0 {
33            self.features.cols()
34        } else {
35            0
36        };
37        let output_dim = if num_samples > 0 {
38            self.labels.cols()
39        } else {
40            0
41        };
42
43        (num_samples, input_dim, output_dim)
44    }
45
46    pub fn append(mut self, features: Vec<T>, labels: Vec<T>) -> Self {
47        self.features.append_row(features);
48        self.labels.append_row(labels);
49        self
50    }
51}
52
53impl<T: Clone> DataSet<T> {
54    #[inline]
55    pub fn features(&self) -> Vec<Vec<T>> {
56        (0..self.features.rows())
57            .map(|i| self.features.row(i).to_vec())
58            .collect()
59    }
60
61    #[inline]
62    pub fn labels(&self) -> Vec<Vec<T>> {
63        (0..self.labels.rows())
64            .map(|i| self.labels.row(i).to_vec())
65            .collect()
66    }
67
68    pub fn shuffle(self) -> Self {
69        let mut indices: Vec<usize> = (0..self.len()).collect();
70        random_provider::shuffle(&mut indices);
71
72        let features = self.features.sort_by_indices(&indices);
73        let labels = self.labels.sort_by_indices(&indices);
74
75        DataSet { features, labels }
76    }
77
78    #[inline]
79    pub fn split(self, ratio: f32) -> (Self, Self) {
80        let ratio = ratio.clamp(0.0, 1.0);
81        let split = (self.len() as f32 * ratio).round() as usize;
82        let (features_left, features_right) = self.features.split_at_row(split);
83        let (labels_left, labels_right) = self.labels.split_at_row(split);
84
85        (
86            DataSet {
87                features: features_left,
88                labels: labels_left,
89            },
90            DataSet {
91                features: features_right,
92                labels: labels_right,
93            },
94        )
95    }
96}
97
98impl<F: Float> DataSet<F> {
99    pub fn standardize(mut self) -> Self {
100        self.features.standardize();
101        self
102    }
103
104    pub fn normalize(mut self) -> Self {
105        self.features.normalize();
106        self
107    }
108}
109
110impl<T> From<(Vec<Vec<T>>, Vec<Vec<T>>)> for DataSet<T> {
111    fn from(data: (Vec<Vec<T>>, Vec<Vec<T>>)) -> Self {
112        DataSet::new(data.0, data.1)
113    }
114}
115
116impl<T> From<(Matrix<T>, Matrix<T>)> for DataSet<T> {
117    fn from(data: (Matrix<T>, Matrix<T>)) -> Self {
118        DataSet {
119            features: data.0,
120            labels: data.1,
121        }
122    }
123}