torsh_data/dataloader/simple.rs
1//! Simple DataLoader API
2//!
3//! This module provides convenient functions for creating DataLoaders with common
4//! configurations without needing to use the builder pattern.
5
6use super::core::DataLoader;
7use crate::{
8 collate::DefaultCollate,
9 dataset::Dataset,
10 sampler::{BatchingSampler, RandomSampler, SequentialSampler},
11};
12use torsh_core::error::Result;
13
14/// Simplified DataLoader type for common use cases with sequential sampling
15pub type SimpleDataLoader<D> = DataLoader<D, BatchingSampler<SequentialSampler>, DefaultCollate>;
16
17/// Simplified DataLoader type for common use cases with random sampling
18pub type SimpleRandomDataLoader<D> = DataLoader<D, BatchingSampler<RandomSampler>, DefaultCollate>;
19
20/// Create a simple DataLoader with basic settings (sequential sampling)
21///
22/// This is a convenience function for quickly creating a DataLoader with sequential
23/// sampling, which is useful for evaluation or when deterministic order is desired.
24///
25/// # Arguments
26///
27/// * `dataset` - The dataset to iterate over
28/// * `batch_size` - Number of samples per batch
29/// * `_shuffle` - Currently ignored (for API compatibility), always uses sequential sampling
30///
31/// # Returns
32///
33/// A Result containing the configured DataLoader or an error
34///
35/// # Examples
36///
37/// ```rust,ignore
38/// use torsh_data::dataloader::simple::simple_dataloader;
39/// use torsh_data::dataset::TensorDataset;
40///
41/// let dataset = TensorDataset::new(vec![1, 2, 3, 4, 5]);
42/// let dataloader = simple_dataloader(dataset, 2, false)?;
43///
44/// for batch in dataloader.iter() {
45/// // Process batch sequentially
46/// }
47/// ```
48///
49/// # Note
50///
51/// This function always creates a DataLoader with sequential sampling for type consistency.
52/// If you need random sampling, use `simple_random_dataloader` instead.
53pub fn simple_dataloader<D: Dataset>(
54 dataset: D,
55 batch_size: usize,
56 _shuffle: bool, // Note: This function always uses sequential sampling for type consistency
57) -> Result<SimpleDataLoader<D>> {
58 DataLoader::builder(dataset)
59 .batch_size(batch_size)
60 .shuffle(false) // Sequential sampling
61 .build()
62}
63
64/// Create a simple DataLoader with random sampling (shuffled)
65///
66/// This is a convenience function for quickly creating a DataLoader with random
67/// sampling, which is useful for training scenarios where data order randomization
68/// is important.
69///
70/// # Arguments
71///
72/// * `dataset` - The dataset to iterate over
73/// * `batch_size` - Number of samples per batch
74/// * `generator` - Optional random seed for reproducible shuffling
75///
76/// # Returns
77///
78/// A Result containing the configured DataLoader with random sampling or an error
79///
80/// # Examples
81///
82/// ```rust,ignore
83/// use torsh_data::dataloader::simple::simple_random_dataloader;
84/// use torsh_data::dataset::TensorDataset;
85///
86/// let dataset = TensorDataset::new(vec![1, 2, 3, 4, 5]);
87/// let dataloader = simple_random_dataloader(dataset, 2, Some(42))?;
88///
89/// for batch in dataloader.iter() {
90/// // Process batch in random order
91/// }
92/// ```
93pub fn simple_random_dataloader<D: Dataset>(
94 dataset: D,
95 batch_size: usize,
96 generator: Option<u64>,
97) -> Result<SimpleRandomDataLoader<D>> {
98 let mut builder = DataLoader::builder(dataset)
99 .batch_size(batch_size)
100 .shuffle(true); // Random sampling
101
102 if let Some(seed) = generator {
103 builder = builder.generator(seed);
104 }
105
106 builder.build_with_random_sampling()
107}
108
109/// Create a simple DataLoader with automatic sampling strategy
110///
111/// This function automatically chooses between sequential and random sampling
112/// based on the shuffle parameter.
113///
114/// # Arguments
115///
116/// * `dataset` - The dataset to iterate over
117/// * `batch_size` - Number of samples per batch
118/// * `shuffle` - Whether to use random sampling (true) or sequential sampling (false)
119/// * `generator` - Optional random seed for reproducible shuffling (only used when shuffle=true)
120///
121/// # Returns
122///
123/// Either a SimpleDataLoader or SimpleRandomDataLoader depending on shuffle setting
124///
125/// # Examples
126///
127/// ```rust,ignore
128/// use torsh_data::dataloader::simple::{simple_dataloader, simple_random_dataloader};
129/// use torsh_data::dataset::TensorDataset;
130///
131/// let dataset = TensorDataset::new(vec![1, 2, 3, 4, 5]);
132///
133/// // Sequential sampling
134/// let sequential_loader = simple_dataloader(dataset, 2, false)?;
135/// for batch in sequential_loader.iter() {
136/// // Process batch
137/// }
138/// ```
139// Note: This function was removed due to lifetime issues with returning boxed iterators
140// Use simple_dataloader() or simple_random_dataloader() directly instead
141
142/// Configuration for simple DataLoader creation
143///
144/// This struct provides a more structured way to configure simple DataLoaders
145/// while maintaining the convenience of the simple API.
146#[derive(Debug, Clone)]
147pub struct SimpleConfig {
148 /// Number of samples per batch
149 pub batch_size: usize,
150 /// Whether to shuffle the data
151 pub shuffle: bool,
152 /// Number of worker threads (0 for single-threaded)
153 pub num_workers: usize,
154 /// Whether to drop the last incomplete batch
155 pub drop_last: bool,
156 /// Optional random seed for reproducible results
157 pub generator: Option<u64>,
158}
159
160impl Default for SimpleConfig {
161 fn default() -> Self {
162 Self {
163 batch_size: 1,
164 shuffle: false,
165 num_workers: 0,
166 drop_last: false,
167 generator: None,
168 }
169 }
170}
171
172impl SimpleConfig {
173 /// Create a new simple configuration
174 pub fn new() -> Self {
175 Self::default()
176 }
177
178 /// Set batch size
179 pub fn batch_size(mut self, batch_size: usize) -> Self {
180 self.batch_size = batch_size;
181 self
182 }
183
184 /// Set shuffle
185 pub fn shuffle(mut self, shuffle: bool) -> Self {
186 self.shuffle = shuffle;
187 self
188 }
189
190 /// Set number of workers
191 pub fn num_workers(mut self, num_workers: usize) -> Self {
192 self.num_workers = num_workers;
193 self
194 }
195
196 /// Set drop last
197 pub fn drop_last(mut self, drop_last: bool) -> Self {
198 self.drop_last = drop_last;
199 self
200 }
201
202 /// Set random generator seed
203 pub fn generator(mut self, seed: u64) -> Self {
204 self.generator = Some(seed);
205 self
206 }
207}
208
209/// Create a DataLoader with simple configuration
210///
211/// This function provides a middle ground between the simple functions and the full builder,
212/// allowing for more configuration while maintaining simplicity.
213///
214/// # Arguments
215///
216/// * `dataset` - The dataset to iterate over
217/// * `config` - Configuration for the DataLoader
218///
219/// # Returns
220///
221/// A boxed DataLoader trait object configured according to the provided config
222///
223/// # Examples
224///
225/// ```rust,ignore
226/// use torsh_data::dataloader::simple::{simple_configured_dataloader, SimpleConfig};
227/// use torsh_data::dataset::TensorDataset;
228///
229/// let dataset = TensorDataset::new(vec![1, 2, 3, 4, 5]);
230/// let config = SimpleConfig::new()
231/// .batch_size(2)
232/// .shuffle(true)
233/// .num_workers(2)
234/// .generator(42);
235///
236/// let dataloader = simple_configured_dataloader(dataset, config)?;
237/// ```
238// Note: simple_configured_dataloader was removed due to lifetime issues
239// Use DataLoader::builder() for advanced configuration or simple_dataloader()/simple_random_dataloader() for basic usage
240
241#[cfg(test)]
242mod tests {
243 use super::*;
244 use crate::dataset::TensorDataset;
245
246 #[test]
247 fn test_simple_dataloader() {
248 // Create a tensor with 5 samples (first dimension is number of samples)
249 let tensor = torsh_tensor::creation::ones::<f32>(&[5]).expect("operation should succeed");
250 let dataset = TensorDataset::from_tensor(tensor);
251 let dataloader =
252 simple_dataloader(dataset, 2, false).expect("simple dataloader should succeed");
253
254 assert_eq!(dataloader.len(), 3); // 5 items, batch size 2 = 3 batches
255 assert!(!dataloader.is_empty());
256 }
257
258 #[test]
259 fn test_simple_random_dataloader() {
260 // Create a tensor with 5 samples (first dimension is number of samples)
261 let tensor = torsh_tensor::creation::ones::<f32>(&[5]).expect("operation should succeed");
262 let dataset = TensorDataset::from_tensor(tensor);
263 let dataloader =
264 simple_random_dataloader(dataset, 2, Some(42)).expect("operation should succeed");
265
266 assert_eq!(dataloader.len(), 3);
267 assert!(!dataloader.is_empty());
268 }
269
270 #[test]
271 fn test_simple_random_dataloader_no_seed() {
272 // Create a tensor with 5 samples (first dimension is number of samples)
273 let tensor = torsh_tensor::creation::ones::<f32>(&[5]).expect("operation should succeed");
274 let dataset = TensorDataset::from_tensor(tensor);
275 let dataloader = simple_random_dataloader(dataset, 2, None)
276 .expect("simple random dataloader should succeed");
277
278 assert_eq!(dataloader.len(), 3);
279 assert!(!dataloader.is_empty());
280 }
281
282 // Tests for removed functions - use simple_dataloader/simple_random_dataloader directly
283
284 #[test]
285 fn test_simple_config() {
286 let config = SimpleConfig::new()
287 .batch_size(4)
288 .shuffle(true)
289 .num_workers(2)
290 .drop_last(true)
291 .generator(42);
292
293 assert_eq!(config.batch_size, 4);
294 assert!(config.shuffle);
295 assert_eq!(config.num_workers, 2);
296 assert!(config.drop_last);
297 assert_eq!(config.generator, Some(42));
298 }
299
300 #[test]
301 fn test_simple_config_defaults() {
302 let config = SimpleConfig::new();
303
304 assert_eq!(config.batch_size, 1);
305 assert!(!config.shuffle);
306 assert_eq!(config.num_workers, 0);
307 assert!(!config.drop_last);
308 assert_eq!(config.generator, None);
309 }
310
311 #[test]
312 fn test_simple_configured_dataloader_sequential() {
313 use torsh_core::device::DeviceType;
314 use torsh_tensor::Tensor;
315
316 let tensor = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0, 5.0], vec![5], DeviceType::Cpu)
317 .expect("Tensor should succeed");
318 let dataset = TensorDataset::from_tensor(tensor);
319 let _config = SimpleConfig::new()
320 .batch_size(2)
321 .shuffle(false)
322 .drop_last(false);
323
324 let dataloader =
325 simple_dataloader(dataset, 2, false).expect("simple dataloader should succeed");
326 assert_eq!(dataloader.len(), 3);
327 }
328
329 #[test]
330 fn test_simple_configured_dataloader_random() {
331 use torsh_core::device::DeviceType;
332 use torsh_tensor::Tensor;
333
334 let tensor = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0, 5.0], vec![5], DeviceType::Cpu)
335 .expect("Tensor should succeed");
336 let _dataset = TensorDataset::from_tensor(tensor);
337 let config = SimpleConfig::new()
338 .batch_size(2)
339 .shuffle(true)
340 .generator(42);
341
342 // Test that config is built correctly since simple_configured_dataloader was removed
343 // Use DataLoader::builder() for configuration instead
344 assert_eq!(config.batch_size, 2);
345 }
346
347 #[test]
348 fn test_empty_dataset_simple() {
349 let dataset: TensorDataset<f32> = TensorDataset::new(vec![]);
350 let dataloader =
351 simple_dataloader(dataset, 2, false).expect("simple dataloader should succeed");
352
353 assert_eq!(dataloader.len(), 0);
354 assert!(dataloader.is_empty());
355 }
356}