use std::fs::File;
use std::io::{Read, Write};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Mutex, OnceLock};
use crate::tensor::{Result, Tensor, TensorError};
pub(crate) struct SampleCache {
slots: Vec<std::sync::RwLock<Option<Vec<Tensor>>>>,
bytes: AtomicUsize,
budget: AtomicUsize,
disk: OnceLock<DiskStage>,
#[cfg(all(debug_assertions, not(test)))]
purity_probed: AtomicBool,
}
impl SampleCache {
pub(crate) fn new(n: usize) -> Self {
let mut slots = Vec::with_capacity(n);
slots.resize_with(n, || std::sync::RwLock::new(None));
SampleCache {
slots,
bytes: AtomicUsize::new(0),
budget: AtomicUsize::new(0),
disk: OnceLock::new(),
#[cfg(all(debug_assertions, not(test)))]
purity_probed: AtomicBool::new(false),
}
}
pub(crate) fn attach_disk(&self, stage: DiskStage) {
let _ = self.disk.set(stage);
}
#[cfg(test)]
pub(crate) fn disk(&self) -> Option<&DiskStage> {
self.disk.get()
}
pub(crate) fn bytes(&self) -> usize {
self.bytes.load(Ordering::Relaxed)
}
pub(crate) fn contains_ram(&self, index: usize) -> bool {
self.slots
.get(index)
.is_some_and(|s| Self::read_slot(s).is_some())
}
#[cfg(test)]
pub(crate) fn cached_count(&self) -> usize {
self.slots
.iter()
.filter(|s| Self::read_slot(s).is_some())
.count()
}
pub(crate) fn resident_indices(&self) -> Vec<usize> {
self.slots
.iter()
.enumerate()
.filter(|(_, s)| Self::read_slot(s).is_some())
.map(|(i, _)| i)
.collect()
}
pub(crate) fn evict(&self, index: usize) -> usize {
let Some(slot) = self.slots.get(index) else {
return 0;
};
let mut guard = slot.write().unwrap_or_else(|p| p.into_inner());
let Some(sample) = guard.take() else {
return 0;
};
let freed = crate::data::budget::retained_cost_estimate(&sample);
self.bytes.fetch_sub(freed, Ordering::Relaxed);
freed
}
fn read_slot(
slot: &std::sync::RwLock<Option<Vec<Tensor>>>,
) -> Option<std::sync::RwLockReadGuard<'_, Option<Vec<Tensor>>>> {
let guard = slot.read().unwrap_or_else(|p| p.into_inner());
if guard.is_some() {
Some(guard)
} else {
None
}
}
pub(crate) fn set_budget(&self, bytes: usize) {
self.budget.store(bytes, Ordering::Relaxed);
}
pub(crate) fn get_or_fetch(
&self,
index: usize,
mut fetch: impl FnMut() -> Result<Vec<Tensor>>,
) -> Result<Vec<Tensor>> {
if let Some(staged) = self.lookup(index) {
return staged;
}
let sample = fetch()?;
#[cfg(all(debug_assertions, not(test)))]
if !self.purity_probed.swap(true, Ordering::Relaxed) {
if let Ok(second) = fetch() {
crate::data::assert_fetch_pure("DataSet::get", &sample, &second);
}
}
self.admit(index, &sample);
Ok(sample)
}
pub(crate) fn lookup(&self, index: usize) -> Option<Result<Vec<Tensor>>> {
if let Some(slot) = self.slots.get(index) {
if let Some(guard) = Self::read_slot(slot) {
let hit = guard.as_ref().expect("read_slot returns occupied");
return Some(Ok(hit.clone()));
}
}
self.disk.get().and_then(|stage| stage.read(index))
}
pub(crate) fn admit(&self, index: usize, sample: &[Tensor]) {
let mut ram_admitted = false;
if let Some(slot) = self.slots.get(index) {
let estimate = crate::data::budget::retained_cost_estimate(sample);
if self.bytes.load(Ordering::Relaxed) + estimate
<= self.budget.load(Ordering::Relaxed)
{
if let Ok((rows, cost)) = crate::data::budget::retain_rows(sample) {
let mut guard = slot.write().unwrap_or_else(|p| p.into_inner());
if guard.is_none() {
*guard = Some(rows);
self.bytes.fetch_add(cost, Ordering::Relaxed);
ram_admitted = true;
}
}
}
}
if !ram_admitted {
if let Some(stage) = self.disk.get() {
stage.admit(index, sample);
}
}
}
}
static STAGE_SEQ: AtomicUsize = AtomicUsize::new(0);
pub(crate) struct DiskStage {
#[cfg(unix)]
reader: File,
writer: Mutex<DiskWriter>,
offsets: Vec<OnceLock<(u64, u64)>>,
budget: u64,
failed: AtomicBool,
path: PathBuf,
}
struct DiskWriter {
file: File,
offset: u64,
}
impl DiskStage {
pub(crate) fn create(dir: &Path, budget_bytes: u64, n: usize) -> Result<Self> {
std::fs::create_dir_all(dir).map_err(|e| {
TensorError::new(&format!(
"DataLoader: disk_stage directory {} cannot be created: {e}",
dir.display()
))
})?;
warn_if_ram_backed(dir);
let seq = STAGE_SEQ.fetch_add(1, Ordering::Relaxed);
let path = dir.join(format!("flodl-stage-{}-{seq}.pack", std::process::id()));
let file = File::options()
.write(true)
.create_new(true)
.open(&path)
.map_err(|e| {
TensorError::new(&format!(
"DataLoader: disk_stage pack file {} cannot be created: {e}",
path.display()
))
})?;
#[cfg(unix)]
let reader = File::open(&path).map_err(|e| {
TensorError::new(&format!(
"DataLoader: disk_stage pack file {} cannot be reopened: {e}",
path.display()
))
})?;
let mut offsets = Vec::with_capacity(n);
offsets.resize_with(n, OnceLock::new);
Ok(DiskStage {
#[cfg(unix)]
reader,
writer: Mutex::new(DiskWriter { file, offset: 0 }),
offsets,
budget: budget_bytes,
failed: AtomicBool::new(false),
path,
})
}
#[cfg(test)]
pub(crate) fn bytes(&self) -> u64 {
self.writer.lock().map(|w| w.offset).unwrap_or(0)
}
#[cfg(test)]
pub(crate) fn staged_count(&self) -> usize {
self.offsets.iter().filter(|s| s.get().is_some()).count()
}
#[cfg(test)]
pub(crate) fn path(&self) -> &Path {
&self.path
}
pub(crate) fn admit(&self, index: usize, sample: &[Tensor]) {
if self.failed.load(Ordering::Relaxed) {
return;
}
let Some(slot) = self.offsets.get(index) else {
return;
};
if slot.get().is_some() {
return;
}
let encoded = match encode_sample(sample) {
Ok(b) => b,
Err(e) => {
self.fail(&format!("sample encode failed: {e}"));
return;
}
};
let len = encoded.len() as u64;
let offset = {
let mut w = match self.writer.lock() {
Ok(w) => w,
Err(_) => return, };
if w.offset.saturating_add(len) > self.budget {
return; }
let offset = w.offset;
if let Err(e) = w.file.write_all(&encoded) {
drop(w);
self.fail(&format!("pack-file write failed: {e}"));
return;
}
w.offset += len;
offset
};
let _ = slot.set((offset, len));
}
pub(crate) fn read(&self, index: usize) -> Option<Result<Vec<Tensor>>> {
let &(offset, len) = self.offsets.get(index)?.get()?;
let mut buf = vec![0u8; len as usize];
#[cfg(unix)]
{
use std::os::unix::fs::FileExt;
if let Err(e) = self.reader.read_exact_at(&mut buf, offset) {
return Some(Err(TensorError::new(&format!(
"DataLoader: disk_stage read failed at {}: {e}",
self.path.display()
))));
}
}
#[cfg(not(unix))]
{
use std::io::{Seek, SeekFrom};
let mut w = match self.writer.lock() {
Ok(w) => w,
Err(_) => {
return Some(Err(TensorError::new(
"DataLoader: disk_stage writer poisoned",
)))
}
};
let end = w.offset;
let read = w
.file
.seek(SeekFrom::Start(offset))
.and_then(|_| w.file.read_exact(&mut buf))
.and_then(|_| w.file.seek(SeekFrom::Start(end)).map(|_| ()));
if let Err(e) = read {
return Some(Err(TensorError::new(&format!(
"DataLoader: disk_stage read failed at {}: {e}",
self.path.display()
))));
}
}
Some(decode_sample(&buf))
}
fn fail(&self, why: &str) {
if !self.failed.swap(true, Ordering::Relaxed) {
eprintln!(
"flodl data: disk_stage disabled ({why}); training continues source-backed, \
already-staged samples keep serving"
);
}
}
}
impl Drop for DiskStage {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.path);
}
}
fn encode_sample(tensors: &[Tensor]) -> Result<Vec<u8>> {
let mut buf = Vec::new();
buf.extend_from_slice(&(tensors.len() as u32).to_le_bytes());
for t in tensors {
crate::nn::checkpoint::write_tensor_data(&mut buf, t)?;
}
Ok(buf)
}
fn decode_sample(bytes: &[u8]) -> Result<Vec<Tensor>> {
let mut r = std::io::Cursor::new(bytes);
let mut n4 = [0u8; 4];
r.read_exact(&mut n4)
.map_err(|e| TensorError::new(&format!("disk_stage: corrupt sample header: {e}")))?;
let n = u32::from_le_bytes(n4) as usize;
let mut out = Vec::with_capacity(n);
for _ in 0..n {
out.push(crate::nn::checkpoint::read_tensor_data(&mut r)?);
}
Ok(out)
}
fn warn_if_ram_backed(dir: &Path) {
let Ok(mounts) = std::fs::read_to_string("/proc/mounts") else {
return; };
let target = dir.canonicalize().unwrap_or_else(|_| dir.to_path_buf());
if let Some(fstype) = fs_type_for(&target, &mounts) {
if fstype == "tmpfs" || fstype == "ramfs" {
eprintln!(
"flodl data: disk_stage directory {} is on {fstype} (RAM-backed): the stage \
will spend RAM, not disk. Point .disk_stage_dir() at a real drive.",
target.display()
);
}
}
}
fn fs_type_for(path: &Path, mounts: &str) -> Option<String> {
let mut best: Option<(&str, &str)> = None;
for line in mounts.lines() {
let mut fields = line.split_whitespace();
let (Some(_dev), Some(mount_point), Some(fstype)) =
(fields.next(), fields.next(), fields.next())
else {
continue;
};
if path.starts_with(mount_point)
&& best.is_none_or(|(b, _)| mount_point.len() > b.len())
{
best = Some((mount_point, fstype));
}
}
best.map(|(_, t)| t.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tensor::Device;
fn sample(value: f32) -> Vec<Tensor> {
vec![Tensor::from_f32(&[value], &[1], Device::CPU).unwrap()]
}
#[test]
fn dormant_cache_is_pure_pass_through() {
let cache = SampleCache::new(4);
let mut fetches = 0;
for _ in 0..3 {
let s = cache
.get_or_fetch(1, || {
fetches += 1;
Ok(sample(1.0))
})
.unwrap();
assert_eq!(s[0].to_f64_vec().unwrap(), vec![1.0]);
}
assert_eq!(fetches, 3, "budget 0: nothing retained, every call fetches");
assert_eq!(cache.bytes(), 0);
assert_eq!(cache.cached_count(), 0);
}
#[test]
fn admit_materializes_oversized_views_and_prices_honestly() {
let cache = SampleCache::new(4);
cache.set_budget(1 << 20);
let base = Tensor::from_f32(
&(0..64).map(|i| i as f32).collect::<Vec<_>>(),
&[8, 8],
Device::CPU,
)
.unwrap();
let row = base.select(0, 1).unwrap();
assert!(row.storage_nbytes() >= 256);
cache.admit(0, std::slice::from_ref(&row));
assert_eq!(cache.bytes(), 32, "charged logical bytes, not the view");
let got = cache.lookup(0).unwrap().unwrap();
assert_eq!(
got[0].storage_nbytes(),
32,
"stored copy owns its storage (base buffer not pinned)"
);
assert_eq!(
got[0].to_f64_vec().unwrap(),
row.to_f64_vec().unwrap(),
"materialized bytes match the view"
);
assert_eq!(cache.evict(0), 32, "eviction frees what admission charged");
assert_eq!(cache.bytes(), 0);
}
#[test]
fn admits_until_budget_then_serves_hits() {
let cache = SampleCache::new(4);
cache.set_budget(8);
let mut fetches = 0;
for idx in 0..4 {
let _ = cache
.get_or_fetch(idx, || {
fetches += 1;
Ok(sample(idx as f32))
})
.unwrap();
}
assert_eq!(fetches, 4, "first pass: all misses");
assert_eq!(cache.cached_count(), 2, "admission stopped at the budget");
assert_eq!(cache.bytes(), 8);
for idx in 0..4 {
let s = cache
.get_or_fetch(idx, || {
fetches += 1;
Ok(sample(idx as f32))
})
.unwrap();
assert_eq!(s[0].to_f64_vec().unwrap(), vec![idx as f64]);
}
assert_eq!(fetches, 6, "two hits, two re-fetches");
}
#[test]
fn evict_frees_room_for_readmission() {
let cache = SampleCache::new(4);
cache.set_budget(8);
cache.admit(0, &sample(0.0));
cache.admit(1, &sample(1.0));
assert_eq!(cache.bytes(), 8);
cache.admit(2, &sample(2.0)); assert_eq!(cache.cached_count(), 2);
assert_eq!(cache.evict(0), 4);
assert_eq!(cache.evict(0), 0);
assert_eq!(cache.evict(99), 0, "out-of-range is a no-op");
assert_eq!(cache.bytes(), 4);
assert!(!cache.contains_ram(0));
assert_eq!(cache.resident_indices(), vec![1]);
cache.admit(2, &sample(2.0));
assert!(cache.contains_ram(2));
assert_eq!(cache.evict(1), 4);
cache.admit(0, &sample(0.5));
assert!(cache.contains_ram(0));
let hit = cache.lookup(0).unwrap().unwrap();
assert_eq!(hit[0].to_f64_vec().unwrap(), vec![0.5]);
assert_eq!(cache.bytes(), 8);
}
#[test]
fn budget_shrink_keeps_retained_content() {
let cache = SampleCache::new(2);
cache.set_budget(64);
let _ = cache.get_or_fetch(0, || Ok(sample(7.0))).unwrap();
assert_eq!(cache.cached_count(), 1);
cache.set_budget(0);
let mut fetched = false;
let s = cache
.get_or_fetch(0, || {
fetched = true;
Ok(sample(0.0))
})
.unwrap();
assert!(!fetched, "still a hit after budget shrink");
assert_eq!(s[0].to_f64_vec().unwrap(), vec![7.0]);
let _ = cache.get_or_fetch(1, || Ok(sample(1.0))).unwrap();
assert_eq!(cache.cached_count(), 1, "no new admission at budget 0");
}
#[test]
fn fetch_errors_pass_through_and_cache_nothing() {
let cache = SampleCache::new(2);
cache.set_budget(64);
let out = cache.get_or_fetch(0, || {
Err(crate::tensor::TensorError::new("io failure"))
});
assert!(out.is_err());
assert_eq!(cache.cached_count(), 0);
let _ = cache.get_or_fetch(0, || Ok(sample(1.0))).unwrap();
assert_eq!(cache.cached_count(), 1);
}
#[test]
fn out_of_range_index_is_fetch_only() {
let cache = SampleCache::new(1);
cache.set_budget(64);
let s = cache.get_or_fetch(9, || Ok(sample(3.0))).unwrap();
assert_eq!(s[0].to_f64_vec().unwrap(), vec![3.0]);
assert_eq!(cache.bytes(), 0);
}
fn stage_dir() -> std::path::PathBuf {
std::env::temp_dir().join("flodl-stage-tests")
}
fn mixed_sample(v: f32, i: i64) -> Vec<Tensor> {
vec![
Tensor::from_f32(&[v, v + 1.0], &[2], Device::CPU).unwrap(),
Tensor::from_i64(&[i], &[1], Device::CPU).unwrap(),
]
}
#[test]
fn disk_stage_round_trips_samples_and_cleans_up() {
let stage = DiskStage::create(&stage_dir(), 1 << 20, 4).unwrap();
let path = stage.path().to_path_buf();
assert!(path.exists());
stage.admit(0, &mixed_sample(1.0, 10));
stage.admit(2, &mixed_sample(3.0, 30));
assert_eq!(stage.staged_count(), 2);
assert!(stage.bytes() > 0);
let s0 = stage.read(0).unwrap().unwrap();
assert_eq!(s0[0].to_f64_vec().unwrap(), vec![1.0, 2.0]);
assert_eq!(s0[1].to_i64_vec().unwrap(), vec![10]);
let s2 = stage.read(2).unwrap().unwrap();
assert_eq!(s2[0].to_f64_vec().unwrap(), vec![3.0, 4.0]);
assert_eq!(s2[1].to_i64_vec().unwrap(), vec![30]);
assert!(stage.read(1).is_none());
assert!(stage.read(9).is_none());
let before = stage.bytes();
stage.admit(0, &mixed_sample(9.0, 99));
assert_eq!(stage.bytes(), before);
drop(stage);
assert!(!path.exists(), "pack file removed on drop");
}
#[test]
fn disk_stage_budget_declines_without_failing() {
let stage = DiskStage::create(&stage_dir(), 8, 4).unwrap();
stage.admit(0, &mixed_sample(1.0, 1));
assert_eq!(stage.staged_count(), 0);
assert!(stage.read(0).is_none());
assert!(!stage.failed.load(Ordering::Relaxed));
}
#[test]
fn cache_cascades_ram_then_disk_then_source() {
let cache = SampleCache::new(3);
cache.set_budget(4);
cache.attach_disk(DiskStage::create(&stage_dir(), 1 << 20, 3).unwrap());
let mut fetches = 0;
for idx in 0..3 {
let s = cache
.get_or_fetch(idx, || {
fetches += 1;
Ok(sample(idx as f32))
})
.unwrap();
assert_eq!(s[0].to_f64_vec().unwrap(), vec![idx as f64]);
}
assert_eq!(fetches, 3, "first pass: all misses");
assert_eq!(cache.cached_count(), 1, "RAM took one");
assert_eq!(cache.disk().unwrap().staged_count(), 2, "disk took the rest");
for idx in 0..3 {
let s = cache
.get_or_fetch(idx, || {
fetches += 1;
Ok(sample(-1.0))
})
.unwrap();
assert_eq!(s[0].to_f64_vec().unwrap(), vec![idx as f64]);
}
assert_eq!(fetches, 3, "second pass: zero source fetches");
}
#[test]
fn fs_type_longest_mount_prefix_wins() {
let mounts = "\
/dev/nvme0n1p2 / ext4 rw,relatime 0 0
tmpfs /tmp tmpfs rw,nosuid,nodev 0 0
/dev/sda1 /data ssdfs rw 0 0
tmpfs /data/scratch tmpfs rw 0 0
";
let t = |p: &str| fs_type_for(Path::new(p), mounts);
assert_eq!(t("/tmp/flodl").as_deref(), Some("tmpfs"));
assert_eq!(t("/data/set").as_deref(), Some("ssdfs"));
assert_eq!(t("/data/scratch/x").as_deref(), Some("tmpfs"));
assert_eq!(t("/home/u").as_deref(), Some("ext4"));
}
}