use std::collections::HashMap;
use crate::tensor::{Device, Result, Tensor, TensorOptions};
use super::prefetch::GovernorCtl;
const SLAB_BYTES: usize = 64 << 20;
const MARGIN_BYTES: u64 = 512 << 20;
pub(crate) const FLOW_RESERVE_BATCHES: u64 = 16;
pub(crate) const VRAM_POOL_DEFAULT: bool = true;
pub(crate) fn vram_pool_env_off() -> bool {
std::env::var("FLODL_VRAM_POOL")
.map(|v| v.eq_ignore_ascii_case("off") || v == "0")
.unwrap_or(false)
}
pub(crate) fn flow_reserve_bytes(in_flight_depth: u64, batch_bytes: u64) -> u64 {
in_flight_depth.min(FLOW_RESERVE_BATCHES) * batch_bytes
}
#[cfg_attr(not(feature = "cuda"), allow(dead_code))]
struct Slab {
tensors: Vec<Tensor>,
used: usize,
}
pub(crate) struct VramSamplePool {
enabled: bool,
device: Device,
decided: bool,
budget: usize,
bytes: usize,
rows_per_slab: usize,
row_bytes: usize,
slabs: Vec<Slab>,
slots: HashMap<usize, (u32, u32)>,
full_logged: bool,
hit_rows: usize,
miss_rows: usize,
captured_rows: usize,
}
impl VramSamplePool {
pub(crate) fn new(device: Device, enabled: bool) -> Self {
VramSamplePool {
enabled: enabled && device.is_cuda(),
device,
decided: false,
budget: 0,
bytes: 0,
rows_per_slab: 0,
row_bytes: 0,
slabs: Vec::new(),
slots: HashMap::new(),
full_logged: false,
hit_rows: 0,
miss_rows: 0,
captured_rows: 0,
}
}
pub(crate) fn active(&self) -> bool {
self.budget > 0
}
pub(crate) fn maybe_install(&mut self, governor: &GovernorCtl, batch: &[Tensor]) {
if !self.enabled || self.decided {
return;
}
if !governor
.honest_resize_done
.load(std::sync::atomic::Ordering::Relaxed)
{
return;
}
let batch_bytes: u64 = batch
.iter()
.map(|t| (t.numel() as u64) * t.dtype().element_size() as u64)
.sum();
let target = governor
.target
.load(std::sync::atomic::Ordering::Relaxed)
.max(1) as u64;
self.install_with_reserve(flow_reserve_bytes(target, batch_bytes));
}
pub(crate) fn install_with_reserve(&mut self, reserve_bytes: u64) {
if !self.enabled || self.decided {
return;
}
self.decided = true;
let Ok((used, total)) =
crate::tensor::cuda_memory_info_idx(self.device.index() as i32)
else {
return; };
let free = total.saturating_sub(used);
let reserve = reserve_bytes + MARGIN_BYTES;
let budget = free.saturating_sub(reserve);
if (budget as usize) < SLAB_BYTES {
crate::verbose!(
"vram-pool: dormant on {:?} | free {}MB - reserve {}MB leaves no slab",
self.device,
free >> 20,
reserve >> 20,
);
return;
}
self.budget = budget as usize;
crate::verbose!(
"vram-pool: {:?} budget {}MB (free {}MB - in-flight reserve {}MB)",
self.device,
budget >> 20,
free >> 20,
reserve >> 20,
);
}
#[cfg_attr(not(feature = "cuda"), allow(dead_code))]
pub(crate) fn partition(&mut self, indices: &[usize]) -> (Vec<usize>, Vec<usize>) {
if !self.active() || self.slots.is_empty() {
return (Vec::new(), (0..indices.len()).collect());
}
let mut hits = Vec::new();
let mut misses = Vec::new();
for (pos, idx) in indices.iter().enumerate() {
if self.slots.contains_key(idx) {
hits.push(pos);
} else {
misses.push(pos);
}
}
self.hit_rows += hits.len();
self.miss_rows += misses.len();
(hits, misses)
}
#[cfg_attr(not(feature = "cuda"), allow(dead_code))]
pub(crate) fn gather(&self, indices: &[usize], positions: &[usize]) -> Result<Vec<Tensor>> {
let mut per_slab: HashMap<u32, (Vec<i64>, Vec<i64>)> = HashMap::new();
for (rank, &pos) in positions.iter().enumerate() {
let &(slab, row) = self
.slots
.get(&indices[pos])
.expect("gather called on a non-hit index");
let entry = per_slab.entry(slab).or_default();
entry.0.push(row as i64);
entry.1.push(rank as i64);
}
let n_positions = self.slabs[0].tensors.len();
let mut out = Vec::with_capacity(n_positions);
for p in 0..n_positions {
let mut pieces = Vec::with_capacity(per_slab.len());
let mut ranks = Vec::with_capacity(positions.len());
for (&slab, (rows, rs)) in &per_slab {
let rows_t =
Tensor::from_i64(rows, &[rows.len() as i64], self.device)?;
pieces.push(self.slabs[slab as usize].tensors[p].index_select(0, &rows_t)?);
ranks.extend_from_slice(rs);
}
let cat = if pieces.len() == 1 {
pieces.pop().expect("one piece")
} else {
Tensor::cat_many(&pieces.iter().collect::<Vec<_>>(), 0)?
};
let mut inv = vec![0i64; ranks.len()];
for (j, &r) in ranks.iter().enumerate() {
inv[r as usize] = j as i64;
}
let inv_t = Tensor::from_i64(&inv, &[inv.len() as i64], self.device)?;
out.push(cat.index_select(0, &inv_t)?);
}
Ok(out)
}
#[cfg_attr(not(feature = "cuda"), allow(dead_code))]
pub(crate) fn capture(
&mut self,
sample_indices: &[usize],
tensors: &[Tensor],
) -> Result<()> {
if !self.active() {
return Ok(());
}
let mut fresh: Vec<usize> = Vec::new();
let mut seen = std::collections::HashSet::new();
for (row, &idx) in sample_indices.iter().enumerate() {
if !self.slots.contains_key(&idx) && seen.insert(idx) {
fresh.push(row);
}
}
if fresh.is_empty() {
return Ok(());
}
if self.rows_per_slab == 0 {
let row_bytes: usize = tensors
.iter()
.map(|t| {
let numel: i64 = t.shape().iter().skip(1).product::<i64>().max(1);
numel as usize * t.dtype().element_size()
})
.sum();
self.row_bytes = row_bytes.max(1);
self.rows_per_slab = (SLAB_BYTES / self.row_bytes).max(1);
}
let mut cursor = 0;
while cursor < fresh.len() {
let space = self
.slabs
.last()
.map(|s| self.rows_per_slab - s.used)
.unwrap_or(0);
if space == 0 {
let slab_bytes = self.rows_per_slab * self.row_bytes;
if self.bytes + slab_bytes > self.budget || self.slabs.len() >= u32::MAX as usize {
if !self.full_logged {
self.full_logged = true;
crate::verbose!(
"vram-pool: {:?} full | {} rows in {} slab(s), {}MB",
self.device,
self.slots.len(),
self.slabs.len(),
self.bytes >> 20,
);
}
return Ok(());
}
let mut slab_tensors = Vec::with_capacity(tensors.len());
for t in tensors {
let mut shape: Vec<i64> = vec![self.rows_per_slab as i64];
shape.extend(t.shape().iter().skip(1));
slab_tensors.push(Tensor::empty(
&shape,
TensorOptions { dtype: t.dtype(), device: self.device },
)?);
}
self.slabs.push(Slab { tensors: slab_tensors, used: 0 });
self.bytes += slab_bytes;
continue;
}
let take = space.min(fresh.len() - cursor);
let chunk = &fresh[cursor..cursor + take];
let rows: Vec<i64> = chunk.iter().map(|&r| r as i64).collect();
let rows_t = Tensor::from_i64(&rows, &[rows.len() as i64], self.device)?;
let slab_id = (self.slabs.len() - 1) as u32;
let slab = self.slabs.last_mut().expect("tail slab exists");
for (p, t) in tensors.iter().enumerate() {
let src = t.index_select(0, &rows_t)?;
slab.tensors[p]
.narrow(0, slab.used as i64, take as i64)?
.copy_(&src, true)?;
}
for (off, &row) in chunk.iter().enumerate() {
self.slots
.insert(sample_indices[row], (slab_id, (slab.used + off) as u32));
}
slab.used += take;
self.captured_rows += take;
cursor += take;
}
Ok(())
}
pub(crate) fn evict_one_slab(&mut self) -> bool {
let Some(slab) = self.slabs.pop() else {
return false;
};
let slab_id = self.slabs.len() as u32;
self.slots.retain(|_, &mut (s, _)| s != slab_id);
self.bytes -= self.rows_per_slab * self.row_bytes;
self.budget = self.bytes;
self.full_logged = false;
crate::verbose!(
"vram-pool: {:?} evicted a slab under memory pressure | {} rows retained, budget now {}MB",
self.device,
self.slots.len(),
self.budget >> 20,
);
drop(slab);
crate::tensor::cuda_empty_cache();
true
}
pub(crate) fn epoch_report(&mut self) {
if !self.active() && self.hit_rows + self.miss_rows == 0 {
return;
}
let seen = self.hit_rows + self.miss_rows;
if seen == 0 {
return;
}
crate::verbose!(
"vram-pool: {:?} served {}/{} rows on-device ({}MB H2D saved), {} captured, {} pooled",
self.device,
self.hit_rows,
seen,
(self.hit_rows * self.row_bytes) >> 20,
self.captured_rows,
self.slots.len(),
);
self.hit_rows = 0;
self.miss_rows = 0;
self.captured_rows = 0;
}
#[cfg(test)]
pub(crate) fn pooled_rows(&self) -> usize {
self.slots.len()
}
#[cfg(test)]
pub(crate) fn set_budget_for_test(&mut self, bytes: usize) {
self.enabled = true;
self.decided = true;
self.budget = bytes;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tensor::DType;
fn test_pool(budget: usize) -> VramSamplePool {
let mut pool = VramSamplePool::new(Device::CPU, false);
pool.set_budget_for_test(budget);
pool
}
fn make_batch(indices: &[usize]) -> Vec<Tensor> {
let n = indices.len();
let data: Vec<f32> = indices
.iter()
.flat_map(|&i| std::iter::repeat_n(i as f32, 4))
.collect();
let labels: Vec<i64> = indices.iter().map(|&i| i as i64).collect();
vec![
Tensor::from_f32(&data, &[n as i64, 4], Device::CPU).unwrap(),
Tensor::from_i64(&labels, &[n as i64], Device::CPU).unwrap(),
]
}
#[test]
fn capture_then_gather_roundtrip() {
let mut pool = test_pool(1 << 30);
let indices = [7usize, 3, 11, 5];
let batch = make_batch(&indices);
pool.capture(&indices, &batch).unwrap();
assert_eq!(pool.pooled_rows(), 4);
let want = [11usize, 7, 5];
let (hits, misses) = pool.partition(&want);
assert_eq!(hits, vec![0, 1, 2]);
assert!(misses.is_empty());
let out = pool.gather(&want, &hits).unwrap();
assert_eq!(out[0].shape(), &[3, 4]);
let data = out[0].to_f32_vec().unwrap();
assert_eq!(&data[0..4], &[11.0; 4]);
assert_eq!(&data[4..8], &[7.0; 4]);
assert_eq!(&data[8..12], &[5.0; 4]);
let labels = out[1].to_i64_vec().unwrap();
assert_eq!(labels, vec![11, 7, 5]);
}
#[test]
fn partition_splits_hits_and_misses_in_order() {
let mut pool = test_pool(1 << 30);
let batch = make_batch(&[1, 2]);
pool.capture(&[1, 2], &batch).unwrap();
let (hits, misses) = pool.partition(&[9, 1, 8, 2]);
assert_eq!(hits, vec![1, 3]);
assert_eq!(misses, vec![0, 2]);
}
#[test]
fn budget_declines_admissions_and_dedups() {
let mut pool = test_pool(1);
let batch = make_batch(&[1, 2]);
pool.capture(&[1, 2], &batch).unwrap();
assert_eq!(pool.pooled_rows(), 0);
let mut pool = test_pool(1 << 30);
let dup = [4usize, 4, 4];
let batch = make_batch(&dup);
pool.capture(&dup, &batch).unwrap();
assert_eq!(pool.pooled_rows(), 1);
}
#[test]
fn eviction_forgets_newest_slab_and_latches_budget() {
let mut pool = test_pool(1 << 30);
let indices: Vec<usize> = (0..10).collect();
let batch = make_batch(&indices);
pool.capture(&indices, &batch).unwrap();
assert_eq!(pool.pooled_rows(), 10);
assert!(pool.evict_one_slab());
assert_eq!(pool.pooled_rows(), 0);
pool.capture(&indices, &batch).unwrap();
assert_eq!(pool.pooled_rows(), 0);
assert!(!pool.evict_one_slab());
}
#[test]
fn dormant_pool_is_pass_through() {
let mut pool = VramSamplePool::new(Device::CPU, false);
let (hits, misses) = pool.partition(&[1, 2, 3]);
assert!(hits.is_empty());
assert_eq!(misses, vec![0, 1, 2]);
let batch = make_batch(&[1, 2, 3]);
pool.capture(&[1, 2, 3], &batch).unwrap();
assert_eq!(pool.pooled_rows(), 0);
assert!(!pool.active());
}
#[test]
fn label_dtype_survives_the_pool() {
let mut pool = test_pool(1 << 30);
let batch = make_batch(&[42]);
pool.capture(&[42], &batch).unwrap();
let out = pool.gather(&[42], &[0]).unwrap();
assert_eq!(out[1].dtype(), DType::Int64);
assert_eq!(out[0].dtype(), DType::Float32);
}
}