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
21pub struct MultiThreadDataLoader<B: Backend, I, O> {
26 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 dataloaders: OnceLock<Vec<BatchDataLoader<B, I, O>>>,
36}
37
38#[derive(Debug)]
40pub enum Message<O> {
41 Batch(usize, O, Progress),
43
44 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 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 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 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 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 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 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 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 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 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 let mut items = Vec::new();
394 for class in 0..num_classes {
395 items.extend(vec![class; class_size]);
396 }
397
398 {
399 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 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 assert_eq!(batch_items.len(), 1);
421 }
422 }
423
424 {
425 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 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 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}