use std::sync::Arc;
use super::prefetch::PrefetchWorker;
use super::sample_cache::SampleCache;
use super::sampler::{RandomSampler, Sampler, SequentialSampler};
use super::{Batch, BatchDataSet, DataSet, DataSetAdapter};
use crate::tensor::{Device, Result, Tensor, TensorError};
const VRAM_MAX_USAGE: f64 = 0.90;
fn can_fit_resident(n: usize, per_sample_bytes: usize, device: Device) -> bool {
if !device.is_cuda() {
return true;
}
let total_bytes = per_sample_bytes as u64 * n as u64;
let idx = device.index() as i32;
match crate::tensor::cuda_memory_info_idx(idx) {
Ok((used, total)) => {
let cap = (total as f64 * VRAM_MAX_USAGE) as u64;
let budget = cap.saturating_sub(used);
total_bytes < budget
}
Err(_) => false, }
}
const BOOTSTRAP_PREFETCH: usize = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ReserveSource {
Bare,
Auto,
User,
}
pub(crate) fn initial_fill_target(full_depth: usize, source: ReserveSource) -> usize {
let divisor = match source {
ReserveSource::Bare => 3,
ReserveSource::Auto => 2,
ReserveSource::User => 1,
};
(full_depth / divisor).max(1)
}
pub(crate) use super::budget::{
prefetch_depth_from_vram, ring_slots_from_ram, sample_cache_budget, RING_SLOTS_WITH_CACHE,
};
#[cfg(test)]
pub(crate) use super::budget::RING_SLOTS_FALLBACK;
#[derive(Clone)]
pub(crate) struct PickCtx {
pub(crate) augment: usize,
pub(crate) seed: u64,
pub(crate) transform: Option<crate::data::TransformFn>,
}
pub struct DataLoaderBuilder {
dataset: Box<dyn BatchDataSet>,
batch_size: usize,
device: Device,
sampler: Option<Box<dyn Sampler>>,
prefetch_depth: Option<usize>,
seed: u64,
drop_last: bool,
force_streaming: bool,
names: Option<Vec<String>>,
vram_max_usage: f64,
ram_max_usage: f64,
activation_reserve: Option<usize>,
pub(crate) sample_cache: Option<Arc<SampleCache>>,
sample_cache_enabled: bool,
disk_stage_bytes: u64,
disk_stage_dir: Option<std::path::PathBuf>,
vram_pool_enabled: bool,
no_shuffle: bool,
augment: usize,
transform: Option<crate::data::TransformFn>,
}
impl DataLoaderBuilder {
fn new(dataset: Box<dyn BatchDataSet>) -> Self {
DataLoaderBuilder {
dataset,
batch_size: 0,
device: Device::CPU,
sampler: None,
prefetch_depth: None,
seed: 42,
drop_last: true,
force_streaming: false,
names: None,
vram_max_usage: 0.90,
ram_max_usage: 0.50,
activation_reserve: None,
sample_cache: None,
sample_cache_enabled: true,
disk_stage_bytes: 0,
disk_stage_dir: None,
vram_pool_enabled: super::vram_pool::VRAM_POOL_DEFAULT,
no_shuffle: false,
augment: 1,
transform: None,
}
}
pub fn batch_size(mut self, batch_size: usize) -> Self {
self.batch_size = batch_size;
self
}
pub fn device(mut self, device: Device) -> Self {
self.device = device;
self
}
pub fn seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
pub fn shuffle(mut self, shuffle: bool) -> Self {
self.no_shuffle = !shuffle;
self
}
pub fn sampler(mut self, sampler: Box<dyn Sampler>) -> Self {
self.sampler = Some(sampler);
self
}
pub fn augment(mut self, k: usize) -> Self {
self.augment = k.max(1);
self
}
pub fn transform(
mut self,
f: impl Fn(Vec<Tensor>, &[crate::data::PickKey]) -> Result<Vec<Tensor>>
+ Send
+ Sync
+ 'static,
) -> Self {
self.transform = Some(crate::data::TransformFn::new(f));
self
}
pub fn prefetch(mut self, depth: usize) -> Self {
self.prefetch_depth = Some(depth);
self
}
pub fn vram_max_usage(mut self, max_usage: f64) -> Self {
self.vram_max_usage = max_usage.clamp(0.50, 0.99);
self
}
pub fn ram_max_usage(mut self, max_usage: f64) -> Self {
self.ram_max_usage = max_usage.clamp(0.0, 0.90);
self
}
pub fn sample_cache(mut self, enabled: bool) -> Self {
self.sample_cache_enabled = enabled;
self
}
pub fn vram_pool(mut self, enabled: bool) -> Self {
self.vram_pool_enabled = enabled;
self
}
pub fn disk_stage(mut self, gb: u64) -> Self {
self.disk_stage_bytes = gb.saturating_mul(1 << 30);
self
}
pub fn disk_stage_dir(mut self, dir: impl Into<std::path::PathBuf>) -> Self {
self.disk_stage_dir = Some(dir.into());
self
}
pub fn activation_reserve(mut self, bytes: usize) -> Self {
self.activation_reserve = Some(bytes);
self
}
pub fn streaming(mut self) -> Self {
self.force_streaming = true;
self
}
pub fn names(mut self, names: &[&str]) -> Self {
self.names = Some(names.iter().map(|s| s.to_string()).collect());
self
}
pub fn drop_last(mut self, drop_last: bool) -> Self {
self.drop_last = drop_last;
self
}
pub fn build(self) -> Result<DataLoader> {
if self.dataset.is_empty() {
return Err(TensorError::new("DataLoader: empty dataset"));
}
if self.batch_size == 0 {
return Err(TensorError::new("DataLoader: batch_size must be > 0"));
}
let DataLoaderBuilder {
dataset,
batch_size,
device,
sampler,
prefetch_depth,
seed,
drop_last,
force_streaming,
names,
vram_max_usage,
ram_max_usage,
activation_reserve,
sample_cache,
sample_cache_enabled,
disk_stage_bytes,
disk_stage_dir,
vram_pool_enabled,
no_shuffle,
augment,
transform,
} = self;
let vram_pool_enabled =
vram_pool_enabled && !super::vram_pool::vram_pool_env_off();
if augment > 1 && sampler.is_some() {
return Err(TensorError::new(
"DataLoader: augment(k) composes with the built-in samplers only \
(the schedule becomes a shuffle of len()*k picks). A custom \
sampler owns its index stream; emit repeated indices from it \
directly if you need multiplicity, or drop the custom sampler.",
));
}
let sample_cache = if sample_cache_enabled {
sample_cache
} else {
None
};
if disk_stage_bytes > 0 && sample_cache.is_none() {
return Err(TensorError::new(
"DataLoader: disk_stage requires the sample layer — a DataSet-backed \
loader with the sample cache enabled. Opaque BatchDataSet loaders \
have no per-sample access to stage; sample_cache(false) disables \
the tier the stage overflows from.",
));
}
let n = dataset.len();
let sample = dataset.get_batch(&[0])?;
if sample.is_empty() {
return Err(TensorError::new(
"DataLoader: dataset returned empty tensor list",
));
}
let num_tensors = sample.len();
let per_sample_bytes: usize = sample.iter().map(|t| t.nbytes()).sum();
drop(sample);
let names = match names {
Some(ref n) if n.len() != num_tensors => {
return Err(TensorError::new(&format!(
"DataLoader: names count ({}) does not match dataset tensor count ({})",
n.len(),
num_tensors,
)));
}
Some(n) => n,
None => (0..num_tensors).map(|i| i.to_string()).collect(),
};
let use_resident = !force_streaming && can_fit_resident(n, per_sample_bytes, device);
let dataset: Arc<dyn BatchDataSet> = Arc::from(dataset);
let shuffle = sampler.is_none() && !no_shuffle;
let picks = n * augment;
let pick_ctx = PickCtx {
augment,
seed,
transform,
};
let sampler = sampler.unwrap_or_else(|| -> Box<dyn Sampler> {
if no_shuffle {
Box::new(SequentialSampler::new(picks))
} else {
Box::new(RandomSampler::new(picks, seed))
}
});
let user_set_depth = prefetch_depth.is_some();
let streaming_depth = prefetch_depth.unwrap_or(BOOTSTRAP_PREFETCH);
if use_resident {
match build_resident(Arc::clone(&dataset), batch_size, device, sampler, drop_last, names.clone(), pick_ctx.clone()) {
Ok(loader) => Ok(loader),
Err(e) if device.is_cuda() && e.is_cuda_oom() => {
let sampler: Box<dyn Sampler> = if shuffle {
Box::new(RandomSampler::new(picks, seed))
} else {
Box::new(SequentialSampler::new(picks))
};
crate::tensor::cuda_empty_cache();
build_streaming(dataset, batch_size, device, sampler, drop_last, streaming_depth, per_sample_bytes, vram_max_usage, ram_max_usage, user_set_depth, activation_reserve, sample_cache, disk_stage_bytes, &disk_stage_dir, vram_pool_enabled, names, pick_ctx)
}
Err(e) => Err(e),
}
} else {
build_streaming(dataset, batch_size, device, sampler, drop_last, streaming_depth, per_sample_bytes, vram_max_usage, ram_max_usage, user_set_depth, activation_reserve, sample_cache, disk_stage_bytes, &disk_stage_dir, vram_pool_enabled, names, pick_ctx)
}
}
}
fn build_resident(
dataset: Arc<dyn BatchDataSet>,
batch_size: usize,
device: Device,
sampler: Box<dyn Sampler>,
drop_last: bool,
names: Vec<String>,
pick_ctx: PickCtx,
) -> Result<DataLoader> {
let n = dataset.len();
let all_indices: Vec<usize> = (0..n).collect();
let tensors = dataset.get_batch(&all_indices)?;
if tensors.is_empty() {
return Err(TensorError::new(
"DataLoader: dataset returned empty tensor list",
));
}
let gpu_data = if device.is_cuda() {
let mut on_device = Vec::with_capacity(tensors.len());
for t in &tensors {
let pinned = t.pin_memory()?;
on_device.push(pinned.to_device(device)?);
}
on_device
} else {
tensors
};
Ok(DataLoader {
inner: LoaderInner::Resident(ResidentLoader {
gpu_data,
_dataset: dataset,
device,
batch_size,
sampler,
drop_last,
names,
pick_ctx,
}),
})
}
#[allow(clippy::too_many_arguments)]
fn build_streaming(
dataset: Arc<dyn BatchDataSet>,
batch_size: usize,
device: Device,
sampler: Box<dyn Sampler>,
drop_last: bool,
prefetch_depth: usize,
per_sample_bytes: usize,
vram_max_usage: f64,
ram_max_usage: f64,
user_set_depth: bool,
activation_reserve: Option<usize>,
sample_cache: Option<Arc<SampleCache>>,
disk_stage_bytes: u64,
disk_stage_dir: &Option<std::path::PathBuf>,
vram_pool_enabled: bool,
names: Vec<String>,
pick_ctx: PickCtx,
) -> Result<DataLoader> {
if disk_stage_bytes > 0 {
if let Some(cache) = &sample_cache {
let dir = disk_stage_dir
.clone()
.unwrap_or_else(std::env::temp_dir);
cache.attach_disk(super::sample_cache::DiskStage::create(
&dir,
disk_stage_bytes,
dataset.len(),
)?);
}
}
let worker = PrefetchWorker::new(
Arc::clone(&dataset),
device,
prefetch_depth,
vram_pool_enabled,
pick_ctx.augment,
);
let (reserve, reserve_source) = match activation_reserve {
Some(bytes) => (bytes, ReserveSource::User),
None => (0, ReserveSource::Bare),
};
Ok(DataLoader {
inner: LoaderInner::Streaming(StreamingLoader {
_dataset: dataset,
batch_size,
device,
sampler,
drop_last,
worker,
names,
per_sample_bytes,
vram_max_usage,
ram_max_usage,
sample_cache,
user_set_depth,
activation_reserve: reserve,
reserve_source,
governor: Arc::new(super::prefetch::GovernorCtl::new(prefetch_depth)),
pick_ctx,
}),
})
}
pub struct DataLoader {
pub(crate) inner: LoaderInner,
}
pub(crate) enum LoaderInner {
Resident(ResidentLoader),
Streaming(StreamingLoader),
}
impl DataLoader {
#[allow(dead_code)]
pub(crate) fn inner(&self) -> &LoaderInner {
&self.inner
}
}
impl DataLoader {
pub fn from_dataset<D: DataSet + 'static>(dataset: D) -> DataLoaderBuilder {
let cache = Arc::new(SampleCache::new(dataset.len()));
let mut builder = DataLoaderBuilder::new(Box::new(DataSetAdapter::with_cache(
dataset,
Arc::clone(&cache),
)));
builder.sample_cache = Some(cache);
builder
}
pub fn from_batch_dataset<D: BatchDataSet + 'static>(dataset: D) -> DataLoaderBuilder {
DataLoaderBuilder::new(Box::new(dataset))
}
pub fn epoch(&mut self, epoch: usize) -> EpochIterator<'_> {
match &mut self.inner {
LoaderInner::Resident(loader) => loader.epoch(epoch),
LoaderInner::Streaming(loader) => loader.epoch(epoch),
}
}
pub fn len(&self) -> usize {
match &self.inner {
LoaderInner::Resident(l) => l.sampler.len(),
LoaderInner::Streaming(l) => l.sampler.len(),
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn num_batches(&self) -> usize {
let (n, bs, dl) = match &self.inner {
LoaderInner::Resident(l) => (l.sampler.len(), l.batch_size, l.drop_last),
LoaderInner::Streaming(l) => (l.sampler.len(), l.batch_size, l.drop_last),
};
if dl { n / bs } else { n.div_ceil(bs) }
}
pub fn batch_size(&self) -> usize {
match &self.inner {
LoaderInner::Resident(l) => l.batch_size,
LoaderInner::Streaming(l) => l.batch_size,
}
}
pub fn device(&self) -> Device {
match &self.inner {
LoaderInner::Resident(l) => l.device,
LoaderInner::Streaming(l) => l.device,
}
}
pub fn is_resident(&self) -> bool {
matches!(&self.inner, LoaderInner::Resident(_))
}
pub fn names(&self) -> &[String] {
match &self.inner {
LoaderInner::Resident(l) => &l.names,
LoaderInner::Streaming(l) => &l.names,
}
}
pub fn prefetch_depth(&self) -> usize {
match &self.inner {
LoaderInner::Resident(_) => 0,
LoaderInner::Streaming(l) => l.worker.prefetch_depth(),
}
}
pub fn set_prefetch_depth(&mut self, depth: usize) {
match &mut self.inner {
LoaderInner::Resident(_) => {}
LoaderInner::Streaming(l) => {
l.worker.set_prefetch_depth(depth.max(1));
l.governor
.target
.store(depth.max(1), std::sync::atomic::Ordering::Relaxed);
l.user_set_depth = true;
}
}
}
pub fn set_activation_reserve(&mut self, bytes: usize) {
if let LoaderInner::Streaming(l) = &mut self.inner {
l.activation_reserve = bytes;
l.reserve_source = ReserveSource::User;
}
}
pub(crate) fn set_activation_reserve_auto(&mut self, bytes: usize) {
if let LoaderInner::Streaming(l) = &mut self.inner {
if l.reserve_source == ReserveSource::Bare {
l.activation_reserve = bytes;
l.reserve_source = ReserveSource::Auto;
}
}
}
pub fn auto_resize(&mut self) -> usize {
match &mut self.inner {
LoaderInner::Resident(_) => 0,
LoaderInner::Streaming(l) => {
use std::sync::atomic::Ordering;
let reserve = if l.governor.honest_resize_done.load(Ordering::Relaxed) {
0
} else {
l.activation_reserve
};
let depth = prefetch_depth_from_vram(
l.per_sample_bytes, l.batch_size, l.device, l.vram_max_usage, reserve,
);
let depth = depth.max(1);
l.worker.set_prefetch_depth(depth);
l.governor.target.store(depth, Ordering::Relaxed);
l.user_set_depth = true;
depth
}
}
}
}
pub(crate) struct ResidentLoader {
gpu_data: Vec<Tensor>,
_dataset: Arc<dyn BatchDataSet>,
device: Device,
batch_size: usize,
sampler: Box<dyn Sampler>,
drop_last: bool,
names: Vec<String>,
pick_ctx: PickCtx,
}
impl ResidentLoader {
fn epoch(&mut self, epoch: usize) -> EpochIterator<'_> {
let picks = self.sampler.indices(epoch);
let n = picks.len();
let bs = self.batch_size;
let mut batch_ranges = Vec::new();
let mut start = 0;
while start < n {
let end = (start + bs).min(n);
if self.drop_last && (end - start) < bs {
break;
}
batch_ranges.push((start, end - start));
start = end;
}
let k = self.pick_ctx.augment.max(1) as i64;
let i64_indices: Vec<i64> = picks.iter().map(|&i| i as i64 / k).collect();
let perm = match Tensor::from_i64(
&i64_indices,
&[i64_indices.len() as i64],
self.device,
) {
Ok(t) => t,
Err(e) => {
return EpochIterator {
inner: EpochIteratorInner::Failed(Some(TensorError::new(&format!(
"resident loader: failed to upload the epoch permutation: {e}"
)))),
}
}
};
EpochIterator {
inner: EpochIteratorInner::Resident(ResidentEpochIter {
data: &self.gpu_data,
perm,
batch_ranges,
pos: 0,
names: &self.names,
picks,
pick_ctx: &self.pick_ctx,
epoch,
}),
}
}
}
pub(crate) struct StreamingLoader {
_dataset: Arc<dyn BatchDataSet>,
batch_size: usize,
device: Device,
sampler: Box<dyn Sampler>,
drop_last: bool,
worker: PrefetchWorker,
names: Vec<String>,
per_sample_bytes: usize,
vram_max_usage: f64,
ram_max_usage: f64,
sample_cache: Option<Arc<SampleCache>>,
user_set_depth: bool,
activation_reserve: usize,
reserve_source: ReserveSource,
governor: Arc<super::prefetch::GovernorCtl>,
pick_ctx: PickCtx,
}
impl StreamingLoader {
fn epoch(&mut self, epoch: usize) -> EpochIterator<'_> {
use std::sync::atomic::Ordering;
if !self.user_set_depth {
let full = prefetch_depth_from_vram(
self.per_sample_bytes, self.batch_size, self.device, self.vram_max_usage, 0,
);
let target = if self.governor.honest_resize_done.load(Ordering::Relaxed) {
full.max(1)
} else {
let reserved = prefetch_depth_from_vram(
self.per_sample_bytes,
self.batch_size,
self.device,
self.vram_max_usage,
self.activation_reserve,
);
initial_fill_target(reserved, self.reserve_source)
};
self.worker.set_prefetch_depth(full.max(2));
self.governor.begin_epoch(target);
} else {
self.governor.begin_epoch(self.worker.prefetch_depth());
}
let indices = self.sampler.indices(epoch);
let n = indices.len();
let bs = self.batch_size;
let num_batches = if self.drop_last {
n / bs
} else {
n.div_ceil(bs)
};
let mem = crate::sys::mem_info().map(|m| m.available_bytes);
let ring_slots = if self.device.is_cuda() {
let sized = ring_slots_from_ram(
self.per_sample_bytes,
bs,
self.ram_max_usage,
mem,
num_batches,
);
if self.sample_cache.is_some() {
sized.min(RING_SLOTS_WITH_CACHE)
} else {
sized
}
} else {
0
};
if let Some(cache) = &self.sample_cache {
if let Some(available) = mem {
let ring_bytes = (ring_slots as u64)
.saturating_mul(self.per_sample_bytes.saturating_mul(bs) as u64);
let budget = sample_cache_budget(
available,
cache.bytes() as u64,
ring_bytes,
self.ram_max_usage,
);
cache.set_budget(usize::try_from(budget).unwrap_or(usize::MAX));
}
}
let batch_rx = self.worker.start_epoch(
indices,
bs,
self.drop_last,
Arc::clone(&self.governor),
ring_slots,
);
EpochIterator {
inner: EpochIteratorInner::Streaming(StreamingEpochIter {
batch_rx,
remaining: num_batches,
names: &self.names,
governor: Arc::clone(&self.governor),
adaptive: !self.user_set_depth,
per_sample_bytes: self.per_sample_bytes,
batch_size: self.batch_size,
device: self.device,
vram_max_usage: self.vram_max_usage,
pick_ctx: &self.pick_ctx,
epoch,
}),
}
}
}
pub struct EpochIterator<'a> {
inner: EpochIteratorInner<'a>,
}
enum EpochIteratorInner<'a> {
Resident(ResidentEpochIter<'a>),
Streaming(StreamingEpochIter<'a>),
Failed(Option<TensorError>),
}
struct ResidentEpochIter<'a> {
data: &'a [Tensor],
perm: Tensor,
batch_ranges: Vec<(usize, usize)>,
pos: usize,
names: &'a [String],
picks: Vec<usize>,
pick_ctx: &'a PickCtx,
epoch: usize,
}
struct StreamingEpochIter<'a> {
batch_rx: std::sync::mpsc::Receiver<Result<super::prefetch::PrefetchedBatch>>,
remaining: usize,
names: &'a [String],
governor: Arc<super::prefetch::GovernorCtl>,
adaptive: bool,
per_sample_bytes: usize,
batch_size: usize,
device: Device,
vram_max_usage: f64,
pick_ctx: &'a PickCtx,
epoch: usize,
}
impl Drop for StreamingEpochIter<'_> {
fn drop(&mut self) {
self.governor
.abandoned
.store(true, std::sync::atomic::Ordering::Relaxed);
}
}
impl<'a> Iterator for EpochIterator<'a> {
type Item = Result<Batch>;
fn next(&mut self) -> Option<Self::Item> {
match &mut self.inner {
EpochIteratorInner::Resident(iter) => iter.next(),
EpochIteratorInner::Streaming(iter) => iter.next(),
EpochIteratorInner::Failed(err) => err.take().map(Err),
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
match &self.inner {
EpochIteratorInner::Resident(iter) => {
let remaining = iter.batch_ranges.len() - iter.pos;
(remaining, Some(remaining))
}
EpochIteratorInner::Streaming(iter) => {
(iter.remaining, Some(iter.remaining))
}
EpochIteratorInner::Failed(err) => {
let n = usize::from(err.is_some());
(n, Some(n))
}
}
}
}
impl ExactSizeIterator for EpochIterator<'_> {}
impl<'a> ResidentEpochIter<'a> {
fn next(&mut self) -> Option<Result<Batch>> {
if self.pos >= self.batch_ranges.len() {
return None;
}
let (start, len) = self.batch_ranges[self.pos];
self.pos += 1;
let batch_perm = match self.perm.narrow(0, start as i64, len as i64) {
Ok(p) => p,
Err(e) => return Some(Err(e)),
};
let mut tensors = Vec::with_capacity(self.data.len());
for t in self.data {
match t.index_select(0, &batch_perm) {
Ok(selected) => tensors.push(selected),
Err(e) => return Some(Err(e)),
}
}
if let Some(ref f) = self.pick_ctx.transform {
let batch_picks = &self.picks[start..start + len];
tensors = match crate::data::apply_transform(
f,
tensors,
batch_picks,
self.pick_ctx.augment,
self.epoch,
self.pick_ctx.seed,
) {
Ok(t) => t,
Err(e) => return Some(Err(e)),
};
}
Some(Ok(Batch::new(tensors, self.names.to_vec())))
}
}
impl StreamingEpochIter<'_> {
fn next(&mut self) -> Option<Result<Batch>> {
use std::sync::atomic::Ordering;
if self.remaining == 0 {
return None;
}
self.remaining -= 1;
match self.batch_rx.recv() {
Ok(Ok(batch)) => {
#[cfg(feature = "cuda")]
if let Some(ref event) = batch.ready_event {
if let Err(e) = event.synchronize() {
return Some(Err(e));
}
match crate::tensor::cuda_stream::CudaStream::current(self.device) {
Ok(cur) => {
for t in &batch.tensors {
if let Err(e) = t.record_stream(&cur) {
return Some(Err(e));
}
}
}
Err(e) => return Some(Err(e)),
}
}
self.governor.consumed.fetch_add(1, Ordering::Relaxed);
let run_consumed =
self.governor.run_consumed.fetch_add(1, Ordering::Relaxed) + 1;
if run_consumed >= 2
&& !self.governor.honest_resize_done.load(Ordering::Relaxed)
{
self.governor.honest_resize_done.store(true, Ordering::Relaxed);
if self.adaptive {
let depth = prefetch_depth_from_vram(
self.per_sample_bytes,
self.batch_size,
self.device,
self.vram_max_usage,
0,
);
self.governor.target.store(depth.max(1), Ordering::Relaxed);
}
}
let tensors = if let Some(ref f) = self.pick_ctx.transform {
match crate::data::apply_transform(
f,
batch.tensors,
&batch.picks,
self.pick_ctx.augment,
self.epoch,
self.pick_ctx.seed,
) {
Ok(t) => t,
Err(e) => return Some(Err(e)),
}
} else {
batch.tensors
};
Some(Ok(Batch::new(tensors, self.names.to_vec())))
}
Ok(Err(e)) => Some(Err(e)),
Err(_) => {
self.remaining = 0;
Some(Err(TensorError::new(
"DataLoader: prefetch worker stopped unexpectedly \
(dataset errors are reported per-batch, so this is \
likely a flodl bug — please report it)",
)))
}
}
}
}
#[cfg(test)]
#[path = "loader_tests.rs"]
mod tests;