Skip to main content

tenshift_core/sample/
mod.rs

1//! Sample  -  the unit of data that flows through the pipeline.
2
3mod ops;
4mod sample_type;
5mod tensor;
6
7pub use sample_type::{Sample, SampleMetadata};
8pub use tensor::{DType, Tensor};
9
10#[cfg(test)]
11mod tests {
12    use super::*;
13    use crate::error::Error;
14    use std::sync::{Arc, OnceLock};
15
16    #[test]
17    fn sample_builder() {
18        let sample = Sample::new()
19            .with("x", Tensor::f32(&[1.0, 2.0, 3.0], vec![3]))
20            .with("y", Tensor::i64(&[42], vec![1]));
21
22        assert_eq!(sample.len(), 2);
23        assert_eq!(sample.get("x").unwrap().shape(), &[3]);
24        assert_eq!(sample.get("y").unwrap().try_as_i64().unwrap(), &[42]);
25    }
26
27    #[test]
28    fn sample_with_replaces_duplicate_field() {
29        let sample = Sample::new()
30            .with("x", Tensor::f32(&[1.0], vec![1]))
31            .with("x", Tensor::f32(&[2.0], vec![1]));
32
33        assert_eq!(sample.len(), 1);
34        assert_eq!(sample.get("x").unwrap().try_as_f32().unwrap(), &[2.0]);
35    }
36
37    #[test]
38    fn sample_insert_replaces_duplicate_field() {
39        let mut sample = Sample::new().with("a", Tensor::i64(&[1], vec![1]));
40        sample.insert("a", Tensor::i64(&[99], vec![1]));
41
42        assert_eq!(sample.len(), 1);
43        assert_eq!(sample.get("a").unwrap().try_as_i64().unwrap(), &[99]);
44    }
45
46    #[test]
47    fn sample_preserves_insertion_order() {
48        let sample = Sample::new()
49            .with("c", Tensor::i64(&[3], vec![1]))
50            .with("a", Tensor::i64(&[1], vec![1]))
51            .with("b", Tensor::i64(&[2], vec![1]));
52
53        let names: Vec<&str> = sample.field_names().collect();
54        assert_eq!(names, &["c", "a", "b"]);
55    }
56
57    #[test]
58    fn sample_remove_works() {
59        let mut sample = Sample::new()
60            .with("x", Tensor::f32(&[1.0], vec![1]))
61            .with("y", Tensor::i64(&[0], vec![1]));
62
63        let removed = sample.remove("x");
64        assert!(removed.is_some());
65        assert_eq!(sample.len(), 1);
66        assert!(sample.get("x").is_none());
67        assert!(sample.get("y").is_some());
68    }
69
70    #[test]
71    fn sample_remove_nonexistent_returns_none() {
72        let mut sample = Sample::new().with("x", Tensor::f32(&[1.0], vec![1]));
73        assert!(sample.remove("y").is_none());
74        assert_eq!(sample.len(), 1);
75    }
76
77    #[test]
78    fn sample_contains() {
79        let sample = Sample::new().with("x", Tensor::f32(&[1.0], vec![1]));
80        assert!(sample.contains("x"));
81        assert!(!sample.contains("y"));
82    }
83
84    #[test]
85    fn tensor_zero_copy_clone() {
86        let t1 = Tensor::f32(&[1.0, 2.0], vec![2]);
87        let t2 = t1.clone();
88        assert!(Arc::ptr_eq(&t1.data, &t2.data));
89    }
90
91    #[cfg(feature = "uring")]
92    #[test]
93    fn bytes_tensor_accepts_completed_wireshift_buffer_without_copying_to_vec() {
94        let pool = wireshift::BufferPool::new(8, 1).unwrap();
95        let mut owned = pool.acquire().unwrap();
96        owned.as_mut_slice()[..4].copy_from_slice(b"ring");
97        let completed = owned
98            .set_filled_len(4)
99            .unwrap()
100            .into_submitted()
101            .into_completed(4)
102            .unwrap();
103
104        let tensor = Tensor::bytes_from_completed(completed);
105        assert_eq!(tensor.as_bytes(), b"ring");
106        assert_eq!(tensor.shape(), &[4]);
107    }
108
109    #[test]
110    fn i32_tensor_round_trip() {
111        let tensor = Tensor::i32(&[1, -2, 3], vec![3]);
112        assert_eq!(tensor.dtype(), DType::I32);
113        assert_eq!(tensor.try_as_i32().unwrap(), &[1, -2, 3]);
114    }
115
116    #[test]
117    fn f64_tensor_round_trip() {
118        let tensor = Tensor::f64(&[1.5, -2.5, 3.5], vec![3]);
119        assert_eq!(tensor.dtype(), DType::F64);
120        assert_eq!(tensor.try_as_f64().unwrap(), &[1.5, -2.5, 3.5]);
121    }
122
123    #[test]
124    fn metadata() {
125        let sample = Sample::new()
126            .with("data", Tensor::u8(vec![0], vec![1]))
127            .with_metadata("train/001.jpg", 0);
128
129        assert_eq!(sample.metadata().unwrap().source, "train/001.jpg");
130    }
131
132    #[test]
133    fn tensor_from_bytes_validates_shape() {
134        let error = Tensor::from_bytes(vec![1, 2, 3], DType::I64, vec![1]).unwrap_err();
135        assert!(matches!(error, Error::InvalidConfig { .. }));
136    }
137
138    #[test]
139    fn misaligned_f32_view_copies_to_aligned_cache() {
140        let mut raw = vec![0_u8];
141        raw.extend_from_slice(&1.5_f32.to_le_bytes());
142        raw.extend_from_slice(&2.5_f32.to_le_bytes());
143        let cache = OnceLock::new();
144
145        let values =
146            ops::cast_numeric_slice::<f32, 4>(&raw[1..], &cache, f32::from_le_bytes).unwrap();
147        assert_eq!(values, &[1.5, 2.5]);
148    }
149
150    #[test]
151    fn misaligned_i64_view_copies_to_aligned_cache() {
152        let mut raw = vec![0_u8];
153        raw.extend_from_slice(&7_i64.to_le_bytes());
154        raw.extend_from_slice(&9_i64.to_le_bytes());
155        let cache = OnceLock::new();
156
157        let values =
158            ops::cast_numeric_slice::<i64, 8>(&raw[1..], &cache, i64::from_le_bytes).unwrap();
159        assert_eq!(values, &[7, 9]);
160    }
161
162    #[test]
163    fn bytes_tensor_shape_matches_input_length() {
164        let tensor = Tensor::bytes(vec![1, 2, 3, 4]);
165        assert_eq!(tensor.shape(), &[4]);
166        assert_eq!(tensor.byte_len(), 4);
167    }
168
169    #[test]
170    fn empty_sample_reports_empty() {
171        let sample = Sample::new();
172        assert!(sample.is_empty());
173        assert_eq!(sample.len(), 0);
174    }
175
176    #[test]
177    fn try_as_f32_wrong_dtype_returns_error() {
178        let tensor = Tensor::i64(&[1], vec![1]);
179        assert!(tensor.try_as_f32().is_err());
180    }
181
182    #[test]
183    fn try_as_f64_wrong_dtype_returns_error() {
184        let tensor = Tensor::f32(&[1.0], vec![1]);
185        assert!(tensor.try_as_f64().is_err());
186    }
187
188    #[test]
189    fn try_as_i32_wrong_dtype_returns_error() {
190        let tensor = Tensor::f32(&[1.0], vec![1]);
191        assert!(tensor.try_as_i32().is_err());
192    }
193
194    #[test]
195    fn try_as_i64_wrong_dtype_returns_error() {
196        let tensor = Tensor::f32(&[1.0], vec![1]);
197        assert!(tensor.try_as_i64().is_err());
198    }
199
200    #[test]
201    fn try_as_f32_correct_dtype_succeeds() {
202        let tensor = Tensor::f32(&[1.0, 2.0], vec![2]);
203        assert_eq!(tensor.try_as_f32().unwrap(), &[1.0, 2.0]);
204    }
205
206    #[test]
207    fn try_as_f64_correct_dtype_succeeds() {
208        let tensor = Tensor::f64(&[3.5, 2.75], vec![2]);
209        assert_eq!(tensor.try_as_f64().unwrap(), &[3.5, 2.75]);
210    }
211
212    #[test]
213    fn cast_numeric_slice_misaligned_length_returns_error() {
214        let data = vec![1_u8, 2, 3];
215        let cache = OnceLock::new();
216        assert!(ops::cast_numeric_slice::<f32, 4>(&data, &cache, f32::from_le_bytes).is_err());
217    }
218
219    #[test]
220    fn dtype_display_is_human_readable() {
221        assert_eq!(DType::F32.to_string(), "f32");
222        assert_eq!(DType::F64.to_string(), "f64");
223        assert_eq!(DType::I64.to_string(), "i64");
224        assert_eq!(DType::Bytes.to_string(), "bytes");
225    }
226
227    #[test]
228    fn num_elements() {
229        let tensor = Tensor::f32(&vec![0.0; 3 * 28 * 28], vec![3, 28, 28]);
230        assert_eq!(tensor.num_elements(), Some(2352));
231    }
232}