use memra_engine::Engine;
use memra_engine::hybrid::HybridModel;
use memra_engine::forward::argmax;
use memra_gguf::GgufFile;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let path = std::env::args().nth(1).expect("usage: depth-profile <model> [depth] [n]");
let depth: usize = std::env::args().nth(2).and_then(|s| s.parse().ok()).unwrap_or(512);
let n: usize = std::env::args().nth(3).and_then(|s| s.parse().ok()).unwrap_or(32);
let e = Engine::new(0)?;
let g = GgufFile::open(&path)?;
let model = HybridModel::load(&e, &g)?;
let prompt: Vec<u32> = (0..depth).map(|i| (100 + (i * 7) % 900) as u32).collect();
let mut cache = memra_engine::cache::Cache::new(&e, &model.cfg, depth + n + 72)?;
let t_prime = std::time::Instant::now();
let mut ll: Vec<f32> = if depth >= memra_engine::hybrid_forward::PRIME_MIN_T {
let (l, _h, _hiddens) = model.prime_cache(&e, &prompt, &mut cache, 0)?;
l
} else {
let mut l = Vec::new();
for &t in &prompt { l = model.decode_step(&e, t, &mut cache)?; }
l
};
for _ in 0..64 { let nx = argmax(&ll) as u32; ll = model.decode_step(&e, nx, &mut cache)?; }
e.stream().synchronize()?;
println!("primed depth={} (+64 warmup) in {:.2}s", cache.pos, t_prime.elapsed().as_secs_f64());
{
let mut scratch = e.uninit(1_944_448)?; for _ in 0..64 {
let mut v = scratch.slice_mut(0..1_944_448);
e.memset_zeros_view(&mut v)?;
}
e.stream().synchronize()?;
}
let t0 = std::time::Instant::now();
for _ in 0..n { let nx = argmax(&ll) as u32; ll = model.decode_step(&e, nx, &mut cache)?; }
e.stream().synchronize()?;
let dt = t0.elapsed().as_secs_f64();
println!("decode {n} toks @d{}..{}: {:.2} tok/s ({:.1} us/tok)",
depth + 4, depth + 4 + n, n as f64 / dt, dt * 1e6 / n as f64);
Ok(())
}