use std::cell::RefCell;
use std::path::{Path, PathBuf};
use std::sync::LazyLock;
use anyhow::{anyhow, Context, Result};
use mlx_native::MlxBuffer;
fn dump_dir_env() -> Option<PathBuf> {
match std::env::var("HF2Q_VIT_DUMP") {
Ok(s) if !s.is_empty() => Some(PathBuf::from(s)),
_ => None,
}
}
fn dtype_audit_env() -> bool {
matches!(
std::env::var("HF2Q_VIT_DUMP_DTYPE_AUDIT").as_deref(),
Ok("1")
)
}
thread_local! {
static COLLECTOR: RefCell<Option<Vec<(String, MlxBuffer)>>> = const {
RefCell::new(None)
};
static AUDIT_COLLECTOR: RefCell<Vec<AuditEntry>> = const {
RefCell::new(Vec::new())
};
}
#[derive(Debug, Clone)]
pub struct AuditEntry {
pub name: String,
pub dtype: String,
pub shape: Vec<usize>,
}
pub fn with_dump_collector<F, R>(f: F) -> Result<(R, Vec<(String, MlxBuffer)>)>
where
F: FnOnce() -> Result<R>,
{
COLLECTOR.with(|c| {
if c.borrow().is_some() {
return Err(anyhow!(
"with_dump_collector: re-entrant arming not supported"
));
}
*c.borrow_mut() = Some(Vec::new());
Ok(())
})?;
let result = f();
let collected = COLLECTOR.with(|c| c.borrow_mut().take().unwrap_or_default());
let r = result?;
Ok((r, collected))
}
pub fn record(name: &str, buffer: &MlxBuffer) {
COLLECTOR.with(|c| {
if let Some(v) = c.borrow_mut().as_mut() {
v.push((name.to_string(), buffer.clone()));
}
});
}
pub fn is_armed() -> bool {
COLLECTOR.with(|c| c.borrow().is_some())
}
pub fn is_dtype_audit_armed() -> bool {
static AUDIT_ARMED: LazyLock<bool> = LazyLock::new(dtype_audit_env);
*AUDIT_ARMED
}
pub fn record_audit(name: &str, buffer: &MlxBuffer) {
if !is_dtype_audit_armed() {
return;
}
if !is_armed() {
return;
}
let dtype = format!("{:?}", buffer.dtype());
let shape = buffer.shape().to_vec();
AUDIT_COLLECTOR.with(|c| {
c.borrow_mut().push(AuditEntry {
name: name.to_string(),
dtype,
shape,
});
});
}
pub fn drain_audit_entries() -> Vec<AuditEntry> {
AUDIT_COLLECTOR.with(|c| std::mem::take(&mut *c.borrow_mut()))
}
pub fn write_dtype_audit(dir: &Path, entries: &[AuditEntry]) -> Result<()> {
use std::io::Write;
let path = dir.join("_dtype_audit.json");
let mut f =
std::fs::File::create(&path).with_context(|| format!("create {}", path.display()))?;
f.write_all(b"[\n")
.with_context(|| format!("write {}", path.display()))?;
for (i, e) in entries.iter().enumerate() {
let shape_str = e
.shape
.iter()
.map(|n| n.to_string())
.collect::<Vec<_>>()
.join(",");
let line = format!(
" {{\"name\":\"{}\",\"dtype\":\"{}\",\"shape\":[{}]}}",
e.name, e.dtype, shape_str
);
f.write_all(line.as_bytes())
.with_context(|| format!("write entry to {}", path.display()))?;
if i + 1 < entries.len() {
f.write_all(b",\n")
.with_context(|| format!("write {}", path.display()))?;
} else {
f.write_all(b"\n")
.with_context(|| format!("write {}", path.display()))?;
}
}
f.write_all(b"]\n")
.with_context(|| format!("write {}", path.display()))?;
Ok(())
}
pub fn record_f32(name: &str, data: &[f32], shape: Vec<usize>) {
COLLECTOR.with(|c| {
if let Some(_v) = c.borrow_mut().as_mut() {
CPU_MIRRORS.with(|m| {
m.borrow_mut().push(CpuMirror {
name: name.to_string(),
data: data.to_vec(),
shape,
});
});
}
});
}
pub struct CpuMirror {
pub name: String,
pub data: Vec<f32>,
pub shape: Vec<usize>,
}
thread_local! {
static CPU_MIRRORS: RefCell<Vec<CpuMirror>> = const { RefCell::new(Vec::new()) };
}
pub fn drain_cpu_mirrors() -> Vec<CpuMirror> {
CPU_MIRRORS.with(|m| std::mem::take(&mut *m.borrow_mut()))
}
pub fn resolve_dump_dir() -> Result<Option<PathBuf>> {
let Some(dir) = dump_dir_env() else {
return Ok(None);
};
if !dir.exists() {
std::fs::create_dir_all(&dir)
.with_context(|| format!("create dump dir {}", dir.display()))?;
} else if !dir.is_dir() {
return Err(anyhow!(
"HF2Q_VIT_DUMP={} exists but is not a directory",
dir.display()
));
}
Ok(Some(dir))
}
pub fn write_dump_gpu(dir: &Path, name: &str, buffer: &MlxBuffer) -> Result<()> {
use mlx_native::DType;
if buffer.dtype() != DType::F32 {
return Err(anyhow!(
"write_dump_gpu({name}): expected F32, got {:?}",
buffer.dtype()
));
}
let slice: &[f32] = buffer
.as_slice::<f32>()
.map_err(|e| anyhow!("write_dump_gpu({name}): as_slice: {e}"))?;
write_dump_inner(dir, name, slice, buffer.shape())
}
pub fn write_dump_cpu(dir: &Path, mirror: &CpuMirror) -> Result<()> {
write_dump_inner(dir, &mirror.name, &mirror.data, &mirror.shape)
}
fn write_dump_inner(dir: &Path, name: &str, data: &[f32], shape: &[usize]) -> Result<()> {
let bin_path = dir.join(format!("{name}.bin"));
let json_path = dir.join(format!("{name}.json"));
use std::io::Write;
let mut file = std::fs::File::create(&bin_path)
.with_context(|| format!("create {}", bin_path.display()))?;
let mut buf = Vec::with_capacity(data.len() * 4);
for v in data {
buf.extend_from_slice(&v.to_le_bytes());
}
file.write_all(&buf)
.with_context(|| format!("write {}", bin_path.display()))?;
let shape_str = shape
.iter()
.map(|n| n.to_string())
.collect::<Vec<_>>()
.join(",");
let json = format!(
"{{\"name\":\"{}\",\"dtype\":\"f32\",\"shape\":[{}],\"n_elements\":{}}}\n",
name,
shape_str,
data.len()
);
let mut jf = std::fs::File::create(&json_path)
.with_context(|| format!("create {}", json_path.display()))?;
jf.write_all(json.as_bytes())
.with_context(|| format!("write {}", json_path.display()))?;
Ok(())
}
#[allow(dead_code)]
pub fn env_var_name() -> &'static str {
"HF2Q_VIT_DUMP"
}
#[allow(dead_code)]
pub static DUMP_DIR_ONESHOT: LazyLock<Option<PathBuf>> = LazyLock::new(dump_dir_env);
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn collector_unarmed_no_op() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
record_f32("00_test_unarmed", &[1.0, 2.0, 3.0], vec![3]);
let drained = drain_cpu_mirrors();
assert!(
drained.is_empty(),
"CPU mirrors should be empty when collector is unarmed"
);
}
#[test]
fn collector_armed_collects_cpu_mirror() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let result: Result<()> = with_dump_collector(|| {
record_f32("00_test_armed", &[1.0, 2.0, 3.0, 4.0], vec![2, 2]);
Ok(())
})
.map(|((), _)| ());
result.expect("armed collect");
let mirrors = drain_cpu_mirrors();
assert_eq!(mirrors.len(), 1);
assert_eq!(mirrors[0].name, "00_test_armed");
assert_eq!(mirrors[0].shape, vec![2, 2]);
assert_eq!(mirrors[0].data, vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn write_dump_inner_round_trip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let tmp = std::env::temp_dir().join(format!("vit_dump_test_{}", std::process::id()));
std::fs::create_dir_all(&tmp).expect("mkdir tmp");
let data = vec![1.5_f32, -2.25, 3.75, 0.0];
let shape = vec![2, 2];
write_dump_inner(&tmp, "01_round_trip", &data, &shape).expect("write");
let bin = std::fs::read(tmp.join("01_round_trip.bin")).expect("read bin");
assert_eq!(bin.len(), data.len() * 4);
let f0 = f32::from_le_bytes([bin[0], bin[1], bin[2], bin[3]]);
assert_eq!(f0, 1.5);
let json = std::fs::read_to_string(tmp.join("01_round_trip.json")).expect("read json");
assert!(json.contains("\"shape\":[2,2]"));
assert!(json.contains("\"dtype\":\"f32\""));
assert!(json.contains("\"n_elements\":4"));
std::fs::remove_dir_all(&tmp).ok();
}
}