Skip to main content

ruda_model/data/dataloader/
multithread.rs

1use ruda_dataset::Dataset;
2use ruda_dataset::transform::PartialDataset;
3use ruda_tensor::api::backend::Backend;
4use rand::distr::{Distribution, StandardUniform};
5use rand::rngs::StdRng;
6use rand::{Rng, SeedableRng};
7
8use super::batcher::Batcher;
9use super::{BatchDataLoader, BatchStrategy, DataLoader, DataLoaderIterator, Progress};
10use std::sync::{
11    Arc, OnceLock,
12    atomic::{AtomicBool, Ordering},
13    mpsc,
14};
15use std::thread;
16
17const MAX_QUEUED_ITEMS: usize = 100;
18
19type RngSeed = <StdRng as SeedableRng>::Seed;
20
21/// A multi-threaded data loader that can be used to iterate over a dataset.
22///
23/// Dropping an iterator stops prefetching and joins its workers, waiting for any
24/// dataset read or batch construction already in progress to return.
25pub struct MultiThreadDataLoader<B: Backend, I, O> {
26    // Configuration parameters needed for initialization
27    strategy: Box<dyn BatchStrategy<I>>,
28    dataset: Arc<dyn Dataset<I>>,
29    batcher: Arc<dyn Batcher<B, I, O>>,
30    device: B::Device,
31    seed: Option<RngSeed>,
32    num_threads: usize,
33
34    // The lazily initialized data loaders
35    dataloaders: OnceLock<Vec<BatchDataLoader<B, I, O>>>,
36}
37
38/// A message that can be sent between threads.
39#[derive(Debug)]
40pub enum Message<O> {
41    /// A batch of items.
42    Batch(usize, O, Progress),
43
44    /// The thread is done.
45    Done,
46}
47
48struct MultiThreadsDataloaderIterator<O> {
49    num_done: usize,
50    workers: Vec<thread::JoinHandle<()>>,
51    receiver: Option<mpsc::Receiver<Message<O>>>,
52    progresses: Vec<Progress>,
53    cancelled: Arc<AtomicBool>,
54}
55
56impl<B: Backend, I, O> MultiThreadDataLoader<B, I, O>
57where
58    I: Send + Sync + Clone + 'static,
59    O: Send + 'static,
60{
61    /// Creates a new multi-threaded batch data loader.
62    ///
63    /// # Arguments
64    ///
65    /// * `strategy` - The batch strategy.
66    /// * `dataset` - The dataset.
67    /// * `batcher` - The batcher.
68    /// * `num_threads` - The number of threads.
69    /// * `device`  - The device to use when loading a batch.
70    /// * `rng`     - The rng determining if the dataset is shuffled each time a dataloader
71    ///   iterator is created.
72    ///
73    /// # Returns
74    ///
75    /// The multi-threaded batch data loader.
76    pub fn new(
77        strategy: Box<dyn BatchStrategy<I>>,
78        dataset: Arc<dyn Dataset<I>>,
79        batcher: Arc<dyn Batcher<B, I, O>>,
80        num_threads: usize,
81        device: B::Device,
82        rng: Option<rand::rngs::StdRng>,
83    ) -> Self {
84        let mut seed = None;
85        if let Some(mut rng) = rng {
86            // RNG stream splitting (not state cloning): derive a new seed from the RNG's output.
87            // This is exactly what `rng.fork()` does.
88            let mut s = RngSeed::default();
89            rng.fill_bytes(&mut s);
90
91            seed = Some(s);
92        }
93        Self::from_seed(strategy, dataset, batcher, num_threads, device, seed)
94    }
95
96    fn from_seed(
97        strategy: Box<dyn BatchStrategy<I>>,
98        dataset: Arc<dyn Dataset<I>>,
99        batcher: Arc<dyn Batcher<B, I, O>>,
100        num_threads: usize,
101        device: B::Device,
102        seed: Option<RngSeed>,
103    ) -> Self {
104        Self {
105            strategy,
106            dataset,
107            batcher,
108            num_threads,
109            device,
110            seed,
111            dataloaders: OnceLock::new(),
112        }
113    }
114
115    /// Force initialization if needed.
116    fn initialize(&self) -> &[BatchDataLoader<B, I, O>] {
117        self.dataloaders
118            .get_or_init(|| {
119                let mut dataset = self.dataset.clone();
120                if let Some(seed) = self.seed.as_ref() {
121                    // Pre-shuffle the dataset before split if shuffle is enabled.
122                    // This ensures that each thread gets a uniform random sample of the dataset.
123                    let mut rng = StdRng::from_seed(*seed);
124                    dataset = Arc::new(ruda_dataset::transform::ShuffledDataset::new(
125                        dataset, &mut rng,
126                    ));
127                }
128
129                let datasets = match self.strategy.batch_size() {
130                    Some(batch_size) => {
131                        PartialDataset::split_chunks(dataset, self.num_threads, batch_size)
132                    }
133                    None => PartialDataset::split(dataset, self.num_threads),
134                };
135
136                // Create more rngs from the first one, one for each new dataloader.
137                let mut rng = self.seed.map(StdRng::from_seed);
138                let rngs = (0..self.num_threads).map(|_| {
139                    rng.as_mut().map(|rng| {
140                        StdRng::seed_from_u64(Distribution::sample(&StandardUniform, rng))
141                    })
142                });
143
144                datasets
145                    .into_iter()
146                    .zip(rngs)
147                    .map(|(dataset, rng)| {
148                        let strategy = self.strategy.clone_dyn();
149                        BatchDataLoader::new(
150                            strategy,
151                            Arc::new(dataset),
152                            self.batcher.clone(),
153                            self.device.clone(),
154                            rng,
155                        )
156                    })
157                    .collect()
158            })
159            .as_ref()
160    }
161}
162
163impl<B: Backend, I, O> DataLoader<B, O> for MultiThreadDataLoader<B, I, O>
164where
165    I: Send + Sync + Clone + 'static,
166    O: Send + 'static + std::fmt::Debug,
167{
168    fn iter<'a>(&'a self) -> Box<dyn DataLoaderIterator<O> + 'a> {
169        // This will initialize the loader if it hasn't been initialized yet
170        let dataloaders = self.initialize();
171
172        let (sender, receiver) = mpsc::sync_channel::<Message<O>>(MAX_QUEUED_ITEMS);
173
174        let progresses = dataloaders
175            .iter()
176            .map(|dataloader| Progress::new(0, dataloader.num_items()))
177            .collect();
178        let mut iterator = MultiThreadsDataloaderIterator::new(
179            receiver,
180            Vec::with_capacity(dataloaders.len()),
181            progresses,
182        );
183
184        for (index, dataloader) in dataloaders.iter().enumerate() {
185            let dataloader_cloned = dataloader.clone();
186            let sender_cloned = sender.clone();
187            let cancelled = iterator.cancelled.clone();
188            let worker = std::thread::Builder::new()
189                .name(std::format!("dataloader-{index}"))
190                .spawn(move || {
191                    if cancelled.load(Ordering::Relaxed) {
192                        return;
193                    }
194                    let mut iterator = dataloader_cloned.iter();
195                    while !cancelled.load(Ordering::Relaxed) {
196                        let Some(item) = iterator.next() else {
197                            break;
198                        };
199                        let progress = iterator.progress();
200                        if sender_cloned
201                            .send(Message::Batch(index, item, progress))
202                            .is_err()
203                        {
204                            return;
205                        }
206                    }
207                    sender_cloned.send(Message::Done).ok();
208                })
209                .expect("failed to spawn data loader worker");
210            iterator.workers.push(worker);
211        }
212        drop(sender);
213
214        Box::new(iterator)
215    }
216
217    fn num_items(&self) -> usize {
218        // For num_items, we can directly use the dataset size without
219        // necessarily initializing the full loader
220        self.dataset.len()
221    }
222
223    fn to_device(&self, device: &B::Device) -> Arc<dyn DataLoader<B, O>> {
224        Arc::new(Self::from_seed(
225            self.strategy.clone_dyn(),
226            self.dataset.clone(),
227            self.batcher.clone(),
228            self.num_threads,
229            device.clone(),
230            self.seed,
231        ))
232    }
233
234    fn slice(&self, start: usize, end: usize) -> Arc<dyn DataLoader<B, O>> {
235        let dataloader = Self::from_seed(
236            self.strategy.clone_dyn(),
237            Arc::new(PartialDataset::new(self.dataset.clone(), start, end)),
238            self.batcher.clone(),
239            self.num_threads,
240            self.device.clone(),
241            self.seed,
242        );
243        Arc::new(dataloader)
244    }
245}
246
247impl<O> MultiThreadsDataloaderIterator<O> {
248    pub fn new(
249        receiver: mpsc::Receiver<Message<O>>,
250        workers: Vec<thread::JoinHandle<()>>,
251        progresses: Vec<Progress>,
252    ) -> Self {
253        MultiThreadsDataloaderIterator {
254            num_done: 0,
255            workers,
256            receiver: Some(receiver),
257            progresses,
258            cancelled: Arc::new(AtomicBool::new(false)),
259        }
260    }
261
262    fn shutdown(&mut self) -> thread::Result<()> {
263        self.cancelled.store(true, Ordering::Relaxed);
264        // Disconnect before joining so workers blocked on a full queue can exit.
265        drop(self.receiver.take());
266        let mut result = Ok(());
267        for worker in self.workers.drain(..) {
268            if let Err(payload) = worker.join() {
269                if result.is_ok() {
270                    result = Err(payload);
271                }
272            }
273        }
274        result
275    }
276}
277
278impl<O> Drop for MultiThreadsDataloaderIterator<O> {
279    fn drop(&mut self) {
280        // Worker panics are propagated by next(), not during unwinding in Drop.
281        let _ = self.shutdown();
282    }
283}
284impl<O: std::fmt::Debug> DataLoaderIterator<O> for MultiThreadsDataloaderIterator<O> {
285    fn progress(&self) -> Progress {
286        let mut items_total = 0;
287        let mut items_processed = 0;
288
289        for progress in self.progresses.iter() {
290            items_total += progress.items_total;
291            items_processed += progress.items_processed;
292        }
293
294        Progress::new(items_processed, items_total)
295    }
296}
297
298impl<O: std::fmt::Debug> Iterator for MultiThreadsDataloaderIterator<O> {
299    type Item = O;
300
301    fn next(&mut self) -> Option<O> {
302        if self.workers.is_empty() {
303            return None;
304        }
305
306        loop {
307            let item = match self.receiver.as_ref()?.recv() {
308                Ok(item) => item,
309                Err(_) => {
310                    if let Err(payload) = self.shutdown() {
311                        std::panic::resume_unwind(payload);
312                    }
313                    panic!("data loader workers disconnected before reporting completion");
314                }
315            };
316
317            match item {
318                Message::Batch(index, item, progress) => {
319                    if let Some(current) = self.progresses.get_mut(index) {
320                        *current = progress;
321                    }
322                    return Some(item);
323                }
324                Message::Done => {
325                    self.num_done += 1;
326                }
327            };
328
329            if self.num_done == self.workers.len() {
330                if let Err(payload) = self.shutdown() {
331                    std::panic::resume_unwind(payload);
332                }
333                return None;
334            }
335        }
336    }
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342    use crate::data::dataloader::FixBatchStrategy;
343    use crate::data::dataloader::batcher::TestBatcher;
344    use crate::data::dataset::FakeDataset;
345    use ruda_dataset::InMemDataset;
346    use std::collections::HashSet;
347
348    #[test]
349    fn test_multi_thread_batch_dataloader() {
350        let batcher = Arc::new(TestBatcher::new());
351        let dataset = Arc::new(FakeDataset::<String>::new(27));
352        let dataloader_single_thread = BatchDataLoader::new(
353            Box::new(FixBatchStrategy::new(5)),
354            dataset.clone(),
355            batcher.clone(),
356            Default::default(),
357            None,
358        );
359        let dataloader_multi_thread = MultiThreadDataLoader::new(
360            Box::new(FixBatchStrategy::new(5)),
361            dataset,
362            batcher,
363            4,
364            Default::default(),
365            None,
366        );
367
368        let mut items_single_thread = HashSet::new();
369        let mut items_multi_thread = HashSet::new();
370
371        for items in dataloader_single_thread.iter() {
372            for item in items {
373                items_single_thread.insert(item);
374            }
375        }
376
377        for items in dataloader_multi_thread.iter() {
378            for item in items {
379                items_multi_thread.insert(item);
380            }
381        }
382
383        assert_eq!(items_single_thread, items_multi_thread);
384    }
385
386    #[test]
387    fn test_multi_thread_batch_dataloader_shuffle() {
388        let num_classes = 2;
389        let class_size = 100;
390        let batch_size = 10;
391
392        // Items is a deliberately ordered dataset.
393        let mut items = Vec::new();
394        for class in 0..num_classes {
395            items.extend(vec![class; class_size]);
396        }
397
398        {
399            // Unshuffled multithreaded loader
400            let dataset = Arc::new(InMemDataset::new(items.clone()));
401            let batcher = Arc::new(TestBatcher::new());
402
403            let loader = MultiThreadDataLoader::new(
404                Box::new(FixBatchStrategy::new(batch_size)),
405                dataset,
406                batcher,
407                num_classes,
408                Default::default(),
409                // No rng means no shuffling.
410                None,
411            );
412
413            for batch in loader.iter() {
414                let mut batch_items = HashSet::new();
415                for item in batch {
416                    batch_items.insert(item);
417                }
418
419                // Since the dataset is not shuffled, we expect each batch to contain the same item.
420                assert_eq!(batch_items.len(), 1);
421            }
422        }
423
424        {
425            // Shuffled multithreaded loader
426            let dataset = Arc::new(InMemDataset::new(items.clone()));
427            let batcher = Arc::new(TestBatcher::new());
428
429            let loader = MultiThreadDataLoader::new(
430                Box::new(FixBatchStrategy::new(batch_size)),
431                dataset.clone(),
432                batcher.clone(),
433                num_classes,
434                Default::default(),
435                // The rng enables shuffling.
436                Some(StdRng::seed_from_u64(42)),
437            );
438
439            for batch in loader.iter() {
440                let mut batch_items = HashSet::new();
441                for item in batch {
442                    batch_items.insert(item);
443                }
444
445                // Since the dataset is shuffled, we expect to see all items.
446                assert_eq!(batch_items.len(), num_classes);
447            }
448        }
449    }
450
451    #[test]
452    fn test_multi_thread_batch_dataloader_incomplete_batches() {
453        let batcher = Arc::new(TestBatcher::new());
454        let dataset = Arc::new(FakeDataset::<String>::new(27));
455        let dataloader_single_thread = BatchDataLoader::new(
456            Box::new(FixBatchStrategy::new(5)),
457            dataset.clone(),
458            batcher.clone(),
459            Default::default(),
460            None,
461        );
462        let dataloader_multi_thread = MultiThreadDataLoader::new(
463            Box::new(FixBatchStrategy::new(5)),
464            dataset,
465            batcher,
466            4,
467            Default::default(),
468            None,
469        );
470
471        let mut items_single_thread = HashSet::new();
472        let mut items_multi_thread = HashSet::new();
473
474        let mut single_thread_cnt = 0;
475        let mut multi_thread_cnt = 0;
476        for items in dataloader_single_thread.iter() {
477            items_single_thread.insert(items);
478            single_thread_cnt += 1;
479        }
480
481        for items in dataloader_multi_thread.iter() {
482            items_multi_thread.insert(items);
483            multi_thread_cnt += 1;
484        }
485
486        assert_eq!(single_thread_cnt, multi_thread_cnt);
487        assert_eq!(items_single_thread, items_multi_thread);
488    }
489}