use super::{BatchStrategy, DataLoader, DataLoaderIterator, Progress, batcher::Batcher};
use burn_dataset::{
Dataset,
transform::{PartialDataset, ShuffledDataset},
};
use burn_tensor::Device;
use rand::SeedableRng;
use std::ops::DerefMut;
use std::sync::Arc;
pub struct BatchDataLoader<I, O> {
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
device: Device,
rng: Option<Arc<spin::Mutex<rand::rngs::StdRng>>>,
}
impl<I, O> Clone for BatchDataLoader<I, O> {
fn clone(&self) -> Self {
Self {
strategy: self.strategy.clone_dyn(),
dataset: self.dataset.clone(),
batcher: self.batcher.clone(),
device: self.device.clone(),
rng: self.rng.clone(),
}
}
}
impl<I, O> BatchDataLoader<I, O> {
pub fn new(
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
device: Device,
rng: Option<rand::rngs::StdRng>,
) -> Self {
Self {
strategy,
dataset,
batcher,
device,
rng: rng.map(|rng| Arc::new(spin::Mutex::new(rng))),
}
}
}
struct BatchDataloaderIterator<I, O> {
current_index: usize,
len: usize,
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
device: Device,
}
impl<I, O> DataLoader<O> for BatchDataLoader<I, O>
where
I: Send + Sync + Clone + 'static,
O: Send + 'static,
{
fn iter<'a>(&'a self) -> Box<dyn DataLoaderIterator<O> + 'a> {
let dataset = match &self.rng {
Some(rng) => Arc::new(ShuffledDataset::new(
self.dataset.clone(),
rng.lock().deref_mut(),
)),
None => self.dataset.clone(),
};
Box::new(BatchDataloaderIterator::new(
self.strategy.clone_dyn(),
dataset,
self.batcher.clone(),
self.device.clone(),
))
}
fn num_items(&self) -> usize {
self.dataset.len()
}
fn to_device(&self, device: &Device) -> Arc<dyn DataLoader<O>> {
let rng = self.rng.as_ref().map(|rng| {
let mut rng = rng.lock();
rng.fork()
});
Arc::new(Self::new(
self.strategy.clone_dyn(),
self.dataset.clone(),
self.batcher.clone(),
device.clone(),
rng,
))
}
fn slice(&self, start: usize, end: usize) -> Arc<dyn DataLoader<O>> {
let rng = self.rng.as_ref().map(|rng| {
let mut rng = rng.lock();
rng.fork()
});
let dataloader = Self::new(
self.strategy.clone_dyn(),
Arc::new(PartialDataset::new(self.dataset.clone(), start, end)),
self.batcher.clone(),
self.device.clone(),
rng,
);
Arc::new(dataloader)
}
}
impl<I, O> BatchDataloaderIterator<I, O> {
pub fn new(
strategy: Box<dyn BatchStrategy<I>>,
dataset: Arc<dyn Dataset<I>>,
batcher: Arc<dyn Batcher<I, O>>,
device: Device,
) -> Self {
let len = dataset.len();
BatchDataloaderIterator {
current_index: 0,
len,
strategy,
dataset,
batcher,
device,
}
}
}
impl<I, O> Iterator for BatchDataloaderIterator<I, O> {
type Item = Result<O, burn_dataset::DatasetError>;
fn next(&mut self) -> Option<Self::Item> {
while self.current_index < self.len {
let chunk_size = self
.strategy
.batch_size()
.unwrap_or(1)
.min(self.len - self.current_index);
let indexes = (self.current_index..self.current_index + chunk_size).collect();
let items = match self.dataset.get_many(indexes) {
Ok(items) => items,
Err(err) => return Some(Err(err)),
};
self.current_index += chunk_size;
for item in items {
self.strategy.add(item);
}
if let Some(items) = self.strategy.batch(false) {
return Some(Ok(self.batcher.batch(items, &self.device)));
}
}
if let Some(items) = self.strategy.batch(true) {
return Some(Ok(self.batcher.batch(items, &self.device)));
}
None
}
}
impl<I, O> DataLoaderIterator<O> for BatchDataloaderIterator<I, O> {
fn progress(&self) -> Progress {
let unit: Option<String> = Some("items".to_string());
Progress::new(self.current_index, self.len, unit)
}
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use super::*;
use crate::data::dataloader::FixBatchStrategy;
use crate::data::dataloader::batcher::TestBatcher;
use crate::data::dataset::FakeDataset;
#[test]
fn test_batch_dataloader() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let dataloader = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset.clone(),
batcher,
Default::default(),
None,
);
let mut items_dataset = HashSet::new();
let mut items_dataloader = HashSet::new();
for item in dataset.iter().map(Result::unwrap) {
items_dataset.insert(item);
}
for items in dataloader.iter().map(Result::unwrap) {
for item in items {
items_dataloader.insert(item);
}
}
assert_eq!(items_dataset, items_dataloader);
}
#[test]
fn test_batch_dataloader_slice() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let dataloader = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset.clone(),
batcher,
Default::default(),
None,
);
let dataloader_slice = dataloader.slice(5, 15);
let mut items_dataloader = HashSet::new();
let mut items_dataloader_slice = HashSet::new();
let mut idx = 0;
for items in dataloader.iter().map(Result::unwrap) {
for item in items {
if (5..15).contains(&idx) {
items_dataloader.insert(item);
}
idx += 1;
}
}
for items in dataloader_slice.iter().map(Result::unwrap) {
for item in items {
items_dataloader_slice.insert(item);
}
}
assert_eq!(items_dataloader, items_dataloader_slice);
}
#[test]
fn test_batch_dataloader_incomplete_last_batch() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(27));
let dataloader = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
Default::default(),
None,
);
let batch_sizes: Vec<usize> = dataloader
.iter()
.map(Result::unwrap)
.map(|items| items.len())
.collect();
assert_eq!(batch_sizes, vec![5, 5, 5, 5, 5, 2]);
}
#[test]
fn test_batch_dataloader_exact_multiple_batches() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(25));
let dataloader = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
Default::default(),
None,
);
let batch_sizes: Vec<usize> = dataloader
.iter()
.map(Result::unwrap)
.map(|items| items.len())
.collect();
assert_eq!(batch_sizes, vec![5, 5, 5, 5, 5]);
}
#[test]
fn test_batch_dataloader_empty_dataset() {
let batcher = Arc::new(TestBatcher::new());
let dataset = Arc::new(FakeDataset::<String>::new(0));
let dataloader = BatchDataLoader::new(
Box::new(FixBatchStrategy::new(5)),
dataset,
batcher,
Default::default(),
None,
);
assert!(dataloader.iter().next().is_none());
}
}