use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::mpsc;
use std::sync::Arc;
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::tensor::{Device, Result, Tensor, TensorError};
use super::vram_pool::VramSamplePool;
use super::BatchDataSet;
pub(crate) struct GovernorCtl {
pub(crate) target: AtomicUsize,
pub(crate) sent: AtomicUsize,
pub(crate) consumed: AtomicUsize,
pub(crate) run_consumed: AtomicUsize,
pub(crate) honest_resize_done: AtomicBool,
pub(crate) abandoned: AtomicBool,
}
impl GovernorCtl {
pub(crate) fn new(initial_target: usize) -> Self {
GovernorCtl {
target: AtomicUsize::new(initial_target.max(1)),
sent: AtomicUsize::new(0),
consumed: AtomicUsize::new(0),
run_consumed: AtomicUsize::new(0),
honest_resize_done: AtomicBool::new(false),
abandoned: AtomicBool::new(false),
}
}
pub(crate) fn begin_epoch(&self, target: usize) {
self.abandoned.store(false, Ordering::Relaxed);
self.target.store(target.max(1), Ordering::Relaxed);
}
pub(crate) fn reset_flight_counters(&self) {
self.sent.store(0, Ordering::Relaxed);
self.consumed.store(0, Ordering::Relaxed);
}
}
fn governor_gate(gov: &GovernorCtl) -> bool {
loop {
if gov.abandoned.load(Ordering::Relaxed) {
return false;
}
let sent = gov.sent.load(Ordering::Relaxed);
let consumed = gov.consumed.load(Ordering::Relaxed);
let target = gov.target.load(Ordering::Relaxed).max(1);
if sent.saturating_sub(consumed) < target {
return true;
}
thread::sleep(Duration::from_millis(1));
}
}
pub(crate) const OOM_RETRY_ATTEMPTS: usize = 10;
pub(crate) const OOM_RETRY_SLEEP: Duration = Duration::from_millis(100);
pub(crate) fn retry_on_oom<T>(
pool: &mut VramSamplePool,
mut attempt: impl FnMut(&mut VramSamplePool) -> Result<T>,
mut on_oom: impl FnMut(&mut VramSamplePool, usize),
) -> Result<T> {
let mut result = attempt(pool);
for i in 0..OOM_RETRY_ATTEMPTS {
match &result {
Err(e) if e.is_cuda_oom() => {
on_oom(pool, i);
result = attempt(pool);
}
_ => break,
}
}
result
}
pub(crate) struct PrefetchedBatch {
pub tensors: Vec<Tensor>,
pub picks: Vec<usize>,
#[cfg(feature = "cuda")]
pub ready_event: Option<crate::tensor::cuda_event::CudaEvent>,
}
pub(crate) enum WorkerCmd {
StartEpoch {
indices: Vec<usize>,
batch_size: usize,
drop_last: bool,
batch_tx: mpsc::SyncSender<Result<PrefetchedBatch>>,
governor: Arc<GovernorCtl>,
ring_slots: usize,
},
StartDistributedEpoch {
batch_tx: mpsc::SyncSender<Result<PrefetchedBatch>>,
},
LoadBatch {
indices: Vec<usize>,
},
InstallVramPool {
reserve_bytes: u64,
},
Stop,
}
pub(crate) struct PrefetchWorker {
cmd_tx: mpsc::Sender<WorkerCmd>,
handle: Option<JoinHandle<()>>,
prefetch_depth: usize,
stop: Arc<AtomicBool>,
}
impl PrefetchWorker {
pub fn new(
dataset: Arc<dyn BatchDataSet>,
device: Device,
prefetch_depth: usize,
vram_pool: bool,
augment: usize,
) -> Self {
let (cmd_tx, cmd_rx) = mpsc::channel::<WorkerCmd>();
let stop = Arc::new(AtomicBool::new(false));
let worker_stop = stop.clone();
let handle = thread::spawn(move || {
worker_loop(dataset, device, cmd_rx, vram_pool, augment, &worker_stop);
});
PrefetchWorker {
cmd_tx,
handle: Some(handle),
prefetch_depth,
stop,
}
}
pub fn start_epoch(
&self,
indices: Vec<usize>,
batch_size: usize,
drop_last: bool,
governor: Arc<GovernorCtl>,
ring_slots: usize,
) -> mpsc::Receiver<Result<PrefetchedBatch>> {
let (batch_tx, batch_rx) =
mpsc::sync_channel::<Result<PrefetchedBatch>>(self.prefetch_depth);
let _ = self.cmd_tx.send(WorkerCmd::StartEpoch {
indices,
batch_size,
drop_last,
batch_tx,
governor,
ring_slots,
});
batch_rx
}
pub fn start_distributed_epoch(&self) -> mpsc::Receiver<Result<PrefetchedBatch>> {
let (batch_tx, batch_rx) =
mpsc::sync_channel::<Result<PrefetchedBatch>>(self.prefetch_depth);
let _ = self.cmd_tx.send(WorkerCmd::StartDistributedEpoch { batch_tx });
batch_rx
}
pub fn load_batch(&self, indices: Vec<usize>) {
let _ = self.cmd_tx.send(WorkerCmd::LoadBatch { indices });
}
pub fn install_vram_pool_budget(&self, reserve_bytes: u64) {
let _ = self.cmd_tx.send(WorkerCmd::InstallVramPool { reserve_bytes });
}
pub fn prefetch_depth(&self) -> usize {
self.prefetch_depth
}
pub fn set_prefetch_depth(&mut self, depth: usize) {
self.prefetch_depth = depth;
}
}
impl Drop for PrefetchWorker {
fn drop(&mut self) {
self.stop.store(true, Ordering::Relaxed);
let _ = self.cmd_tx.send(WorkerCmd::Stop);
if let Some(h) = self.handle.take() {
let _ = h.join();
}
}
}
fn worker_loop(
dataset: Arc<dyn BatchDataSet>,
device: Device,
cmd_rx: mpsc::Receiver<WorkerCmd>,
vram_pool: bool,
augment: usize,
stop: &AtomicBool,
) {
#[cfg(feature = "cuda")]
let copy_stream = if device.is_cuda() {
crate::tensor::cuda_stream::CudaStream::new(device, false).ok()
} else {
None
};
let mut pool = VramSamplePool::new(device, vram_pool);
let mut dist_tx: Option<mpsc::SyncSender<Result<PrefetchedBatch>>> = None;
#[cfg(all(debug_assertions, not(test)))]
let mut purity_probed = false;
#[cfg(all(debug_assertions, not(test)))]
let probe_purity = |probe: &[usize]| {
let samples = crate::data::picks_to_samples(probe, augment);
if let (Ok(a), Ok(b)) = (dataset.get_batch(&samples), dataset.get_batch(&samples)) {
crate::data::assert_fetch_pure("BatchDataSet::get_batch", &a, &b);
}
};
for cmd in &cmd_rx {
match cmd {
WorkerCmd::StartEpoch {
indices,
batch_size,
drop_last,
batch_tx,
governor,
ring_slots,
} => {
dist_tx = None;
governor.reset_flight_counters();
#[cfg(all(debug_assertions, not(test)))]
if !purity_probed && !indices.is_empty() {
purity_probed = true;
probe_purity(&indices[..batch_size.min(indices.len())]);
}
if ring_slots > 0 {
run_two_stage_epoch(
&dataset,
device,
indices,
batch_size,
drop_last,
&batch_tx,
&governor,
ring_slots,
&mut pool,
augment,
stop,
#[cfg(feature = "cuda")]
copy_stream.as_ref(),
);
} else {
run_single_stage_epoch(
&dataset,
device,
&indices,
batch_size,
drop_last,
&batch_tx,
&governor,
&mut pool,
augment,
stop,
#[cfg(feature = "cuda")]
copy_stream.as_ref(),
);
}
pool.epoch_report();
}
WorkerCmd::StartDistributedEpoch { batch_tx } => {
pool.epoch_report();
dist_tx = Some(batch_tx);
}
WorkerCmd::InstallVramPool { reserve_bytes } => {
pool.install_with_reserve(reserve_bytes);
}
WorkerCmd::LoadBatch { indices } => {
#[cfg(all(debug_assertions, not(test)))]
if !purity_probed && !indices.is_empty() {
purity_probed = true;
probe_purity(&indices);
}
if let Some(ref tx) = dist_tx {
let result = if device.is_cuda() {
retry_on_oom(
&mut pool,
|pool| {
fetch_and_transfer(
&*dataset,
&indices,
augment,
device,
pool,
#[cfg(feature = "cuda")]
copy_stream.as_ref(),
)
},
|pool, attempt| {
crate::tensor::cuda_empty_cache();
thread::sleep(OOM_RETRY_SLEEP);
if attempt >= OOM_RETRY_ATTEMPTS / 2 {
pool.evict_one_slab();
}
},
)
} else {
fetch_and_transfer(
&*dataset,
&indices,
augment,
device,
&mut pool,
#[cfg(feature = "cuda")]
copy_stream.as_ref(),
)
};
match send_or_stop(tx, result, stop) {
SendOutcome::Sent => {}
SendOutcome::Disconnected => dist_tx = None, SendOutcome::Stopping => return,
}
}
}
WorkerCmd::Stop => {
pool.epoch_report();
break;
}
}
}
}
enum SendOutcome {
Sent,
Disconnected,
Stopping,
}
fn send_or_stop(
tx: &mpsc::SyncSender<Result<PrefetchedBatch>>,
result: Result<PrefetchedBatch>,
stop: &AtomicBool,
) -> SendOutcome {
let mut pending = result;
loop {
match tx.try_send(pending) {
Ok(()) => return SendOutcome::Sent,
Err(mpsc::TrySendError::Disconnected(_)) => return SendOutcome::Disconnected,
Err(mpsc::TrySendError::Full(v)) => {
if stop.load(Ordering::Relaxed) {
return SendOutcome::Stopping;
}
pending = v;
thread::sleep(Duration::from_millis(1));
}
}
}
}
#[allow(clippy::too_many_arguments)]
fn run_single_stage_epoch(
dataset: &Arc<dyn BatchDataSet>,
device: Device,
indices: &[usize],
batch_size: usize,
drop_last: bool,
batch_tx: &mpsc::SyncSender<Result<PrefetchedBatch>>,
governor: &GovernorCtl,
pool: &mut VramSamplePool,
augment: usize,
stop: &AtomicBool,
#[cfg(feature = "cuda")] copy_stream: Option<&crate::tensor::cuda_stream::CudaStream>,
) {
let n = indices.len();
let mut start = 0;
while start < n {
let end = (start + batch_size).min(n);
if drop_last && (end - start) < batch_size {
break;
}
if !governor_gate(governor) {
break; }
let batch_picks = &indices[start..end];
let samples = crate::data::picks_to_samples(batch_picks, augment);
start = end;
let result = guarded_get_batch(dataset.as_ref(), &samples).and_then(|tensors| {
pooled_transfer_with_retry(
batch_picks,
&samples,
&tensors,
device,
governor,
pool,
#[cfg(feature = "cuda")]
copy_stream,
)
});
match send_or_stop(batch_tx, result, stop) {
SendOutcome::Sent => {}
SendOutcome::Disconnected | SendOutcome::Stopping => break,
}
governor.sent.fetch_add(1, Ordering::Relaxed);
}
}
#[allow(clippy::too_many_arguments)]
fn run_two_stage_epoch(
dataset: &Arc<dyn BatchDataSet>,
device: Device,
indices: Vec<usize>,
batch_size: usize,
drop_last: bool,
batch_tx: &mpsc::SyncSender<Result<PrefetchedBatch>>,
governor: &GovernorCtl,
ring_slots: usize,
pool: &mut VramSamplePool,
augment: usize,
stop: &AtomicBool,
#[cfg(feature = "cuda")] copy_stream: Option<&crate::tensor::cuda_stream::CudaStream>,
) {
let (ring_tx, ring_rx) =
mpsc::sync_channel::<Result<(Vec<usize>, Vec<Tensor>)>>(ring_slots);
let reader_dataset = Arc::clone(dataset);
let reader = thread::spawn(move || {
reader_loop(reader_dataset, indices, batch_size, drop_last, ring_tx, augment);
});
loop {
if !governor_gate(governor) {
break; }
let cpu_batch = match ring_rx.recv() {
Ok(b) => b,
Err(_) => break, };
let result = match cpu_batch {
Ok((batch_picks, tensors)) => {
let samples = crate::data::picks_to_samples(&batch_picks, augment);
pooled_transfer_with_retry(
&batch_picks,
&samples,
&tensors,
device,
governor,
pool,
#[cfg(feature = "cuda")]
copy_stream,
)
}
Err(e) => Err(e),
};
match send_or_stop(batch_tx, result, stop) {
SendOutcome::Sent => {}
SendOutcome::Disconnected | SendOutcome::Stopping => break,
}
governor.sent.fetch_add(1, Ordering::Relaxed);
}
drop(ring_rx);
let _ = reader.join();
}
fn reader_loop(
dataset: Arc<dyn BatchDataSet>,
indices: Vec<usize>,
batch_size: usize,
drop_last: bool,
ring_tx: mpsc::SyncSender<Result<(Vec<usize>, Vec<Tensor>)>>,
augment: usize,
) {
let n = indices.len();
let mut start = 0;
while start < n {
let end = (start + batch_size).min(n);
if drop_last && (end - start) < batch_size {
break;
}
let batch_picks = &indices[start..end];
let samples = crate::data::picks_to_samples(batch_picks, augment);
let result = guarded_get_batch(dataset.as_ref(), &samples)
.map(|tensors| (batch_picks.to_vec(), tensors));
start = end;
if ring_tx.send(result).is_err() {
break;
}
}
}
fn oom_backoff(governor: &GovernorCtl) {
let t = governor.target.load(Ordering::Relaxed);
governor.target.store((t / 2).max(1), Ordering::Relaxed);
crate::tensor::cuda_empty_cache();
thread::sleep(OOM_RETRY_SLEEP);
}
fn pooled_transfer_with_retry(
picks: &[usize],
samples: &[usize],
tensors: &[Tensor],
device: Device,
governor: &GovernorCtl,
pool: &mut VramSamplePool,
#[cfg(feature = "cuda")] copy_stream: Option<&crate::tensor::cuda_stream::CudaStream>,
) -> Result<PrefetchedBatch> {
if !device.is_cuda() {
return transfer_batch(
picks,
samples,
tensors,
device,
pool,
#[cfg(feature = "cuda")]
copy_stream,
);
}
pool.maybe_install(governor, tensors);
retry_on_oom(
pool,
|pool| {
transfer_batch(
picks,
samples,
tensors,
device,
pool,
#[cfg(feature = "cuda")]
copy_stream,
)
},
|pool, _attempt| {
let at_floor = governor.target.load(Ordering::Relaxed) <= 1;
oom_backoff(governor);
if at_floor {
pool.evict_one_slab();
}
},
)
}
fn guarded_get_batch(dataset: &dyn BatchDataSet, samples: &[usize]) -> Result<Vec<Tensor>> {
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| dataset.get_batch(samples))) {
Ok(r) => r,
Err(payload) => {
let msg = payload
.downcast_ref::<&str>()
.map(|s| s.to_string())
.or_else(|| payload.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "non-string panic payload".into());
Err(TensorError::new(&format!(
"dataset panicked in get_batch: {msg}"
)))
}
}
}
fn fetch_and_transfer(
dataset: &dyn BatchDataSet,
picks: &[usize],
augment: usize,
device: Device,
pool: &mut VramSamplePool,
#[cfg(feature = "cuda")] copy_stream: Option<&crate::tensor::cuda_stream::CudaStream>,
) -> Result<PrefetchedBatch> {
let samples = crate::data::picks_to_samples(picks, augment);
let tensors = guarded_get_batch(dataset, &samples)?;
transfer_batch(
picks,
&samples,
&tensors,
device,
pool,
#[cfg(feature = "cuda")]
copy_stream,
)
}
fn transfer_batch(
picks: &[usize],
samples: &[usize],
tensors: &[Tensor],
device: Device,
pool: &mut VramSamplePool,
#[cfg(feature = "cuda")] copy_stream: Option<&crate::tensor::cuda_stream::CudaStream>,
) -> Result<PrefetchedBatch> {
if !device.is_cuda() {
return Ok(PrefetchedBatch {
tensors: tensors.to_vec(),
picks: picks.to_vec(),
#[cfg(feature = "cuda")]
ready_event: None,
});
}
#[cfg(feature = "cuda")]
{
use crate::tensor::cuda_event::{CudaEvent, CudaEventFlags};
use crate::tensor::cuda_stream::StreamGuard;
if let Some(stream) = copy_stream {
let _guard = StreamGuard::new(stream);
let on_device =
assemble_on_device(samples, tensors, device, pool, true)?;
let event = CudaEvent::new(CudaEventFlags::DisableTiming)?;
event.record_on(stream)?;
return Ok(PrefetchedBatch {
tensors: on_device,
picks: picks.to_vec(),
ready_event: Some(event),
});
}
let on_device =
assemble_on_device(samples, tensors, device, pool, false)?;
Ok(PrefetchedBatch {
tensors: on_device,
picks: picks.to_vec(),
ready_event: None,
})
}
#[cfg(not(feature = "cuda"))]
{
let _ = samples;
let _ = pool;
Ok(PrefetchedBatch {
tensors: tensors.to_vec(),
picks: picks.to_vec(),
})
}
}
#[cfg(feature = "cuda")]
fn assemble_on_device(
indices: &[usize],
tensors: &[Tensor],
device: Device,
pool: &mut VramSamplePool,
async_copy: bool,
) -> Result<Vec<Tensor>> {
let upload = |t: &Tensor| -> Result<Tensor> {
let pinned = t.pin_memory()?;
if async_copy {
pinned.to_device_async(device)
} else {
pinned.to_device(device)
}
};
let (hits, misses) = pool.partition(indices);
let uploaded: Vec<Tensor> = if misses.len() == indices.len() {
tensors.iter().map(&upload).collect::<Result<_>>()?
} else if !misses.is_empty() {
let rows: Vec<i64> = misses.iter().map(|&p| p as i64).collect();
let rows_t = Tensor::from_i64(&rows, &[rows.len() as i64], Device::CPU)?;
tensors
.iter()
.map(|t| upload(&t.index_select(0, &rows_t)?))
.collect::<Result<_>>()?
} else {
Vec::new()
};
if hits.is_empty() {
pool.capture(indices, &uploaded)?;
return Ok(uploaded);
}
let gathered = pool.gather(indices, &hits)?;
if misses.is_empty() {
return Ok(gathered);
}
let n = indices.len();
let mut map = vec![0i64; n];
for (k, &pos) in hits.iter().enumerate() {
map[pos] = k as i64;
}
for (m, &pos) in misses.iter().enumerate() {
map[pos] = (hits.len() + m) as i64;
}
let map_t = Tensor::from_i64(&map, &[n as i64], device)?;
let mut out = Vec::with_capacity(tensors.len());
for (g, u) in gathered.iter().zip(uploaded.iter()) {
out.push(Tensor::cat_many(&[g, u], 0)?.index_select(0, &map_t)?);
}
let miss_samples: Vec<usize> = misses.iter().map(|&p| indices[p]).collect();
pool.capture(&miss_samples, &uploaded)?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
struct TinyBatch;
impl BatchDataSet for TinyBatch {
fn len(&self) -> usize {
64
}
fn get_batch(&self, indices: &[usize]) -> Result<Vec<Tensor>> {
let vals: Vec<f32> = indices.iter().map(|&i| i as f32).collect();
Ok(vec![Tensor::from_f32(
&vals,
&[vals.len() as i64],
Device::CPU,
)?])
}
}
#[test]
fn drop_unwedges_worker_blocked_on_full_channel() {
let w = PrefetchWorker::new(Arc::new(TinyBatch), Device::CPU, 1, false, 1);
let rx = w.start_distributed_epoch();
w.load_batch(vec![0]); w.load_batch(vec![1]); w.load_batch(vec![2]);
thread::sleep(Duration::from_millis(200));
let (done_tx, done_rx) = mpsc::channel();
thread::spawn(move || {
drop(w);
let _ = done_tx.send(());
});
done_rx
.recv_timeout(Duration::from_secs(30))
.expect("PrefetchWorker::drop hung — teardown latch failed");
drop(rx);
}
}