use {
crate::{
core::{
data_source::{
DataSource
},
shape::{
Shape
},
split_data_source::{
SplitDataSource
}
}
}
};
pub struct DataSet< I, O >
where I: DataSource,
O: DataSource
{
input_data: I,
expected_output_data: O
}
impl< I, O > DataSet< I, O >
where I: DataSource,
O: DataSource
{
pub fn new( input_data: I, expected_output_data: O ) -> Self {
assert_eq!(
input_data.len(),
expected_output_data.len(),
"The training input data has {} samples which is not equal to the amount of samples in the expected output data where we have {} samples",
input_data.len(),
expected_output_data.len()
);
DataSet {
input_data,
expected_output_data
}
}
pub fn len( &self ) -> usize {
self.input_data.len()
}
pub fn is_empty( &self ) -> bool {
self.len() == 0
}
pub fn input_shape( &self ) -> Shape {
self.input_data.shape()
}
pub fn output_shape( &self ) -> Shape {
self.expected_output_data.shape()
}
pub fn input_data( &self ) -> &I {
&self.input_data
}
pub fn expected_output_data( &self ) -> &O {
&self.expected_output_data
}
pub fn clone_and_split( self, split_at: f32 )
-> (
DataSet< SplitDataSource< I >, SplitDataSource< O > >,
DataSet< SplitDataSource< I >, SplitDataSource< O > >
)
where I: Clone, O: Clone
{
assert!( split_at >= 0.0 );
assert!( split_at <= 1.0 );
let left_count = (self.len() as f32 * split_at) as usize;
let left_range = 0..left_count;
let right_range = left_count..self.len();
let right = DataSet {
input_data: SplitDataSource::new( self.input_data.clone(), right_range.clone() ),
expected_output_data: SplitDataSource::new( self.expected_output_data.clone(), right_range )
};
let left = DataSet {
input_data: SplitDataSource::new( self.input_data, left_range.clone() ),
expected_output_data: SplitDataSource::new( self.expected_output_data, left_range )
};
(left, right)
}
}
#[cfg(test)]
mod tests {
use {
crate::{
core::{
shape::{
Shape
},
slice_source::{
SliceSource
}
}
},
super::{
DataSet
}
};
#[test]
fn test_data_set_basics() {
let input_data: Vec< u32 > = vec![ 1, 2, 3, 4 ];
let output_data: Vec< u32 > = vec![ 10, 20 ];
let inputs = SliceSource::from( Shape::new_2d( 1, 2 ), input_data );
let outputs = SliceSource::from( Shape::new_1d( 1 ), output_data );
let data_set = DataSet::new( inputs, outputs );
assert_eq!( data_set.len(), 2 );
assert_eq!( data_set.is_empty(), false );
assert_eq!( data_set.input_shape(), Shape::new_2d( 1, 2 ) );
}
}