1mod 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}