use burn_dataset::Dataset;
use burn_dataset::transform::PartialDataset;
use burn_tensor::Device;
use rand::distr::{Distribution, StandardUniform};
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
use super::batcher::Batcher;
use super::{BatchDataLoader, BatchStrategy, DataLoader, DataLoaderIterator, Progress};
use std::sync::{Arc, OnceLock, mpsc, mpsc::SyncSender};
use std::thread;
const MAX_QUEUED_ITEMS: usize = 100;
type RngSeed = <StdRng as SeedableRng>::Seed;
pub struct MultiThreadDataLoader<I, O> {
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
device: Device,
seed: Option<RngSeed>,
num_threads: usize,
dataloaders: OnceLock<Vec<BatchDataLoader<I, O>>>,
workers: OnceLock<WorkerPool<O>>,
}
#[derive(Debug)]
pub enum Message<O> {
Batch(usize, O, Progress),
Done,
Error(usize, String),
}
struct MultiThreadsDataloaderIterator<O> {
num_done: usize,
num_workers: usize,
receiver: mpsc::Receiver<Message<O>>,
progresses: Vec<Progress>,
}
type WorkerCommand<O> = SyncSender<Message<O>>;
struct WorkerPool<O> {
senders: Vec<mpsc::Sender<WorkerCommand<O>>>,
handles: Vec<thread::JoinHandle<()>>,
item_counts: Vec<usize>,
}
impl<O> Drop for WorkerPool<O> {
fn drop(&mut self) {
self.senders.clear();
for handle in self.handles.drain(..) {
let _ = handle.join();
}
}
}
impl<I, O> MultiThreadDataLoader<I, O>
where
I: Send + Sync + Clone + 'static,
O: Send + 'static,
{
pub fn new(
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
num_threads: usize,
device: Device,
rng: Option<rand::rngs::StdRng>,
) -> Self {
let mut seed = None;
if let Some(mut rng) = rng {
let mut s = RngSeed::default();
rng.fill_bytes(&mut s);
seed = Some(s);
}
Self::from_seed(strategy, dataset, batcher, num_threads, device, seed)
}
fn from_seed(
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
num_threads: usize,
device: Device,
seed: Option<RngSeed>,
) -> Self {
Self {
strategy,
dataset,
batcher,
num_threads,
device,
seed,
dataloaders: OnceLock::new(),
workers: OnceLock::new(),
}
}
fn initialize(&self) -> &[BatchDataLoader<I, O>] {
self.dataloaders
.get_or_init(|| {
let mut dataset = self.dataset.clone();
if let Some(seed) = self.seed.as_ref() {
let mut rng = StdRng::from_seed(*seed);
dataset = Arc::new(burn_dataset::transform::ShuffledDataset::new(
dataset, &mut rng,
));
}
let datasets = match self.strategy.batch_size() {
Some(batch_size) => {
PartialDataset::split_chunks(dataset, self.num_threads, batch_size)
}
None => PartialDataset::split(dataset, self.num_threads),
};
let mut rng = self.seed.map(StdRng::from_seed);
let rngs = (0..self.num_threads).map(|_| {
rng.as_mut().map(|rng| {
StdRng::seed_from_u64(Distribution::sample(&StandardUniform, rng))
})
});
datasets
.into_iter()
.zip(rngs)
.map(|(dataset, rng)| {
let strategy = self.strategy.clone_dyn();
BatchDataLoader::new(
strategy,
Arc::new(dataset),
self.batcher.clone(),
self.device.clone(),
rng,
)
})
.collect()
})
.as_ref()
}
fn workers(&self) -> &WorkerPool<O> {
self.workers.get_or_init(|| {
let dataloaders = self.initialize();
let item_counts: Vec<usize> = dataloaders.iter().map(|d| d.num_items()).collect();
let mut senders = Vec::with_capacity(dataloaders.len());
let mut handles = Vec::with_capacity(dataloaders.len());
for (index, dataloader) in dataloaders.iter().enumerate() {
let dataloader = dataloader.clone();
let (command_sender, command_receiver) = mpsc::channel::<WorkerCommand<O>>();
let handle = thread::Builder::new()
.name(std::format!("dataloader-{index}"))
.spawn(move || {
while let Ok(sender) = command_receiver.recv() {
let mut iterator = dataloader.iter();
loop {
match iterator.next() {
Some(Ok(item)) => {
let progress = iterator.progress();
if sender
.send(Message::Batch(index, item, progress))
.is_err()
{
break;
}
}
None => break,
Some(Err(dataset_err)) => {
sender
.send(Message::Error(index, dataset_err.to_string()))
.ok();
break;
}
}
}
sender.send(Message::Done).ok();
}
})
.unwrap();
senders.push(command_sender);
handles.push(handle);
}
WorkerPool {
senders,
handles,
item_counts,
}
})
}
}
impl<I, O> DataLoader<O> for MultiThreadDataLoader<I, O>
where
I: Send + Sync + Clone + 'static,
O: Send + 'static + std::fmt::Debug,
{
fn iter<'a>(&'a self) -> Box<dyn DataLoaderIterator<O> + 'a> {
let workers = self.workers();
let (sender, receiver) = mpsc::sync_channel::<Message<O>>(MAX_QUEUED_ITEMS);
let unit: Option<String> = Some("items".to_string());
let mut progresses = Vec::with_capacity(workers.senders.len());
for (command_sender, &num_items) in workers.senders.iter().zip(workers.item_counts.iter()) {
progresses.push(Progress::new(0, num_items, unit.clone()));
command_sender
.send(sender.clone())
.expect("Dataloader worker thread should be alive");
}
let num_workers = workers.senders.len();
drop(sender);
Box::new(MultiThreadsDataloaderIterator::new(
receiver,
num_workers,
progresses,
))
}
fn num_items(&self) -> usize {
self.dataset.len()
}
fn to_device(&self, device: &Device) -> Arc<dyn DataLoader<O>> {
Arc::new(Self::from_seed(
self.strategy.clone_dyn(),
self.dataset.clone(),
self.batcher.clone(),
self.num_threads,
device.clone(),
self.seed,
))
}
fn slice(&self, start: usize, end: usize) -> Arc<dyn DataLoader<O>> {
let dataloader = Self::from_seed(
self.strategy.clone_dyn(),
Arc::new(PartialDataset::new(self.dataset.clone(), start, end)),
self.batcher.clone(),
self.num_threads,
self.device.clone(),
self.seed,
);
Arc::new(dataloader)
}
}
impl<O> MultiThreadsDataloaderIterator<O> {
pub fn new(
receiver: mpsc::Receiver<Message<O>>,
num_workers: usize,
progresses: Vec<Progress>,
) -> Self {
MultiThreadsDataloaderIterator {
num_done: 0,
num_workers,
receiver,
progresses,
}
}
}
impl<O: std::fmt::Debug> DataLoaderIterator<O> for MultiThreadsDataloaderIterator<O> {
fn progress(&self) -> Progress {
let mut items_total = 0;
let mut items_processed = 0;
let unit: Option<String> = Some("items".to_string());
for progress in self.progresses.iter() {
items_total += progress.items_total;
items_processed += progress.items_processed;
}
Progress::new(items_processed, items_total, unit)
}
}
impl<O: std::fmt::Debug> Iterator for MultiThreadsDataloaderIterator<O> {
type Item = Result<O, burn_dataset::DatasetError>;
fn next(&mut self) -> Option<Self::Item> {
if self.num_workers == 0 {
return None;
}
loop {
match self.receiver.recv() {
Ok(Message::Batch(index, item, progress)) => {
if let Some(current) = self.progresses.get_mut(index) {
*current = progress;
}
return Some(Ok(item));
}
Ok(Message::Done) => {
self.num_done += 1;
if self.num_done == self.num_workers {
return None;
}
}
Ok(Message::Error(index, msg)) => {
return Some(Err(burn_dataset::DatasetError::new(std::io::Error::other(
format!("dataloader worker {index} failed: {msg}"),
))));
}
Err(_) => return None,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::data::dataloader::FixBatchStrategy;
use crate::data::dataloader::batcher::TestBatcher;
use crate::data::dataset::FakeDataset;
use burn_dataset::DatasetError;
use burn_dataset::InMemDataset;
use std::collections::HashSet;
use std::sync::mpsc;
use std::time::Duration;
struct FlakyDataset {
len: usize,
fail_at: usize,
}
impl Dataset<usize> for FlakyDataset {
fn get(&self, index: usize) -> Result<usize, DatasetError> {
assert!(index < self.len, "index out of bounds");
if index == self.fail_at {
return Err(DatasetError::new(std::io::Error::other(
"simulated dataset failure",
)));
}
Ok(index)
}
fn len(&self) -> usize {
self.len
}
}
#[test]
fn test_multi_thread_batch_dataloader() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let dataloader_single_thread = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset.clone(),
batcher.clone(),
Default::default(),
None,
);
let dataloader_multi_thread = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
4,
Default::default(),
None,
);
let mut items_single_thread = HashSet::new();
let mut items_multi_thread = HashSet::new();
for items in dataloader_single_thread.iter().map(Result::unwrap) {
for item in items {
items_single_thread.insert(item);
}
}
for items in dataloader_multi_thread.iter().map(Result::unwrap) {
for item in items {
items_multi_thread.insert(item);
}
}
assert_eq!(items_single_thread, items_multi_thread);
}
#[test]
fn test_multi_thread_batch_dataloader_shuffle() {
let num_classes = 2;
let class_size = 100;
let batch_size = 10;
let mut items = Vec::new();
for class in 0..num_classes {
items.extend(vec![class; class_size]);
}
{
let dataset = Arc::new(InMemDataset::new(items.clone()));
let batcher = Arc::new(TestBatcher::new());
let loader = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(batch_size)),
dataset,
batcher,
num_classes,
Default::default(),
None,
);
for batch in loader.iter().map(Result::unwrap) {
let mut batch_items = HashSet::new();
for item in batch {
batch_items.insert(item);
}
assert_eq!(batch_items.len(), 1);
}
}
{
let dataset = Arc::new(InMemDataset::new(items.clone()));
let batcher = Arc::new(TestBatcher::new());
let loader = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(batch_size)),
dataset.clone(),
batcher.clone(),
num_classes,
Default::default(),
Some(StdRng::seed_from_u64(42)),
);
for batch in loader.iter().map(Result::unwrap) {
let mut batch_items = HashSet::new();
for item in batch {
batch_items.insert(item);
}
assert_eq!(batch_items.len(), num_classes);
}
}
}
#[test]
fn test_multi_thread_batch_dataloader_incomplete_batches() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let dataloader_single_thread = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset.clone(),
batcher.clone(),
Default::default(),
None,
);
let dataloader_multi_thread = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
4,
Default::default(),
None,
);
let mut items_single_thread = HashSet::new();
let mut items_multi_thread = HashSet::new();
let mut single_thread_cnt = 0;
let mut multi_thread_cnt = 0;
for items in dataloader_single_thread.iter().map(Result::unwrap) {
items_single_thread.insert(items);
single_thread_cnt += 1;
}
for items in dataloader_multi_thread.iter().map(Result::unwrap) {
items_multi_thread.insert(items);
multi_thread_cnt += 1;
}
assert_eq!(single_thread_cnt, multi_thread_cnt);
assert_eq!(items_single_thread, items_multi_thread);
}
#[test]
fn test_multi_thread_batch_dataloader_multiple_epochs() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let expected: HashSet<_> = dataset.iter().map(Result::unwrap).collect();
let dataloader = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
4,
Default::default(),
None,
);
for _epoch in 0..3 {
let mut items = HashSet::new();
for batch in dataloader.iter().map(Result::unwrap) {
for item in batch {
items.insert(item);
}
}
assert_eq!(items, expected);
}
}
#[test]
fn test_multi_thread_batch_dataloader_resumes_after_early_drop() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let expected: HashSet<_> = dataset.iter().map(Result::unwrap).collect();
let dataloader = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
4,
Default::default(),
None,
);
{
let mut iterator = dataloader.iter();
let _ = iterator.next();
}
let mut items = HashSet::new();
for batch in dataloader.iter().map(Result::unwrap) {
for item in batch {
items.insert(item);
}
}
assert_eq!(items, expected);
}
#[test]
fn test_multi_thread_batch_dataloader_propagates_worker_error_instead_of_hanging() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FlakyDataset {
len: 40,
fail_at: 20,
});
let dataloader = MultiThreadDataLoader::new(
Box::new(FixBatchStrategy::new(1)),
dataset,
batcher,
4,
Default::default(),
None,
);
let (done_tx, done_rx) = mpsc::channel();
let handle = thread::spawn(move || {
let saw_error = dataloader.iter().any(|batch| batch.is_err());
done_tx.send(saw_error).ok();
});
let saw_error = done_rx
.recv_timeout(Duration::from_secs(10))
.expect("dataloader hung instead of reporting the worker error");
assert!(
saw_error,
"expected the dataloader to yield an Err on a worker error"
);
let _ = handle.join();
}
}