mlprep_train_test_split_seeded/
train_test_split_seeded.rs1use matten::Tensor;
15use matten_mlprep::train_test_split_seeded;
16
17fn main() {
18 let x = Tensor::new(vec![10.0, 20.0, 30.0, 40.0, 50.0], &[5, 1]);
19 let (train, test) = train_test_split_seeded(&x, 0.6, 7).expect("valid split"); println!("train {:?}: {:?}", train.shape(), train.as_slice());
21 println!("test {:?}: {:?}", test.shape(), test.as_slice());
22
23 assert_eq!(train.shape(), &[3, 1]);
24 assert_eq!(test.shape(), &[2, 1]);
25
26 let (train2, test2) = train_test_split_seeded(&x, 0.6, 7).expect("valid split");
28 assert_eq!(train.as_slice(), train2.as_slice());
29 assert_eq!(test.as_slice(), test2.as_slice());
30 println!("same seed -> reproduced split: OK");
31
32 let (train3, _) = train_test_split_seeded(&x, 0.6, 8).expect("valid split");
34 println!(
35 "different seed -> train {:?}: {:?}",
36 train3.shape(),
37 train3.as_slice()
38 );
39
40 println!("train_test_split_seeded: OK");
41}