use std::collections::BTreeMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Mutex, OnceLock};
use std::time::Duration;
#[derive(Default, Clone)]
struct Row {
calls: u64,
gpu_ns: u128,
cpu_ns: u128,
refused: u64,
worst: f32,
}
fn table() -> &'static Mutex<BTreeMap<(usize, usize, usize), Row>> {
static T: OnceLock<Mutex<BTreeMap<(usize, usize, usize), Row>>> = OnceLock::new();
T.get_or_init(|| Mutex::new(BTreeMap::new()))
}
pub fn on() -> bool {
static ON: OnceLock<bool> = OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_MM_AB").as_deref() == Ok("1"))
}
static SEEN: AtomicU64 = AtomicU64::new(0);
#[allow(clippy::too_many_arguments)]
pub fn record(
b: usize,
rows: usize,
cols: usize,
gpu_took_it: bool,
gpu: Duration,
cpu: Duration,
g: &[f32],
c: &[f32],
) {
SEEN.fetch_add(1, Ordering::Relaxed);
let mut t = table().lock().unwrap();
let e = t.entry((b, rows, cols)).or_default();
e.calls += 1;
e.cpu_ns += cpu.as_nanos();
if gpu_took_it {
e.gpu_ns += gpu.as_nanos();
let scale = c.iter().fold(0f32, |m, &v| m.max(v.abs())).max(1e-6);
let d = g
.iter()
.zip(c)
.fold(0f32, |m, (&a, &b)| m.max((a - b).abs()))
/ scale;
e.worst = e.worst.max(d);
} else {
e.refused += 1;
}
}
pub fn report() -> String {
if SEEN.load(Ordering::Relaxed) == 0 {
return "no q4tp matmat calls were eligible for the device arm".into();
}
let t = table().lock().unwrap();
let mut rows: Vec<_> = t.iter().collect();
rows.sort_by_key(|(_, r)| std::cmp::Reverse(r.cpu_ns));
let mut s = String::from(
"\n q4tp matmat, both arms per call (CMF_MM_AB=1)\n\
\x20 b rows cols calls gpu ms cpu ms ratio worst\n",
);
let (mut tg, mut tc) = (0u128, 0u128);
for ((b, r, c), e) in rows {
let took = e.calls - e.refused;
let g = e.gpu_ns as f64 / 1e6;
let cp = e.cpu_ns as f64 / 1e6;
tg += e.gpu_ns;
tc += e.cpu_ns;
let ratio = if took > 0 && g > 0.0 {
format!("{:.2}x", cp / g)
} else {
"refused".into()
};
s.push_str(&format!(
" {b:>5} {r:>8} {c:>7} {:>6} {g:>9.1} {cp:>9.1} {ratio:>7} {:>6.4}\n",
e.calls, e.worst
));
}
s.push_str(&format!(
" total gpu {:.1} ms, cpu {:.1} ms — the device arm is {:.2}x the host's\n",
tg as f64 / 1e6,
tc as f64 / 1e6,
tc as f64 / (tg as f64).max(1.0),
));
s
}
pub fn split_on() -> bool {
static ON: OnceLock<bool> = OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_MM_SPLIT").as_deref() == Ok("1"))
}
#[derive(Default, Clone)]
struct Split {
calls: u64,
kernel_ns: u128,
rb_ns: u128,
}
fn splits() -> &'static Mutex<BTreeMap<(usize, usize, usize), Split>> {
static T: OnceLock<Mutex<BTreeMap<(usize, usize, usize), Split>>> = OnceLock::new();
T.get_or_init(|| Mutex::new(BTreeMap::new()))
}
pub fn split_note(b: usize, rows: usize, cols: usize, kernel: Duration, rb: Duration) {
let mut t = splits().lock().unwrap();
let e = t.entry((b, rows, cols)).or_default();
e.calls += 1;
e.kernel_ns += kernel.as_nanos();
e.rb_ns += rb.as_nanos();
}
pub fn split_report() -> String {
let t = splits().lock().unwrap();
if t.is_empty() {
return "no split-timed calls".into();
}
let mut s = String::from(
"\n device arm, kernel vs readback (CMF_MM_SPLIT=1)\n\
\x20 b rows cols calls kernel ms readbk ms rb MB/call\n",
);
for ((b, r, c), e) in t.iter() {
s.push_str(&format!(
" {b:>5} {r:>8} {c:>7} {:>6} {:>10.1} {:>10.1} {:>10.1}\n",
e.calls,
e.kernel_ns as f64 / 1e6,
e.rb_ns as f64 / 1e6,
(b * r * 4) as f64 / 1e6,
));
}
s
}