use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct Costs {
pub matmul_flops: u64,
pub matmul_calls: u64,
pub elem_ops: u64,
pub elem_calls: u64,
pub transcendental: u64,
pub transcendental_vec: u64,
pub bytes_moved: u64,
pub copies: u64,
pub copy_bytes: u64,
}
macro_rules! counters {
($($name:ident),* $(,)?) => {
$(static $name: AtomicU64 = AtomicU64::new(0);)*
};
}
counters!(
MATMUL_FLOPS,
MATMUL_CALLS,
ELEM_OPS,
ELEM_CALLS,
TRANSCENDENTAL,
TRANSCENDENTAL_VEC,
BYTES_MOVED,
COPIES,
COPY_BYTES
);
static ON: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub fn start() -> bool {
for c in [
&MATMUL_FLOPS,
&MATMUL_CALLS,
&ELEM_OPS,
&ELEM_CALLS,
&TRANSCENDENTAL,
&TRANSCENDENTAL_VEC,
&BYTES_MOVED,
&COPIES,
©_BYTES,
] {
c.store(0, Ordering::Relaxed);
}
ON.swap(true, Ordering::Relaxed)
}
pub fn stop() -> Costs {
ON.store(false, Ordering::Relaxed);
Costs {
matmul_flops: MATMUL_FLOPS.load(Ordering::Relaxed),
matmul_calls: MATMUL_CALLS.load(Ordering::Relaxed),
elem_ops: ELEM_OPS.load(Ordering::Relaxed),
elem_calls: ELEM_CALLS.load(Ordering::Relaxed),
transcendental: TRANSCENDENTAL.load(Ordering::Relaxed),
transcendental_vec: TRANSCENDENTAL_VEC.load(Ordering::Relaxed),
bytes_moved: BYTES_MOVED.load(Ordering::Relaxed),
copies: COPIES.load(Ordering::Relaxed),
copy_bytes: COPY_BYTES.load(Ordering::Relaxed),
}
}
#[inline]
fn on() -> bool {
ON.load(Ordering::Relaxed)
}
pub fn matmul(batch: u64, m: u64, k: u64, n: u64) {
if !on() {
return;
}
MATMUL_CALLS.fetch_add(batch, Ordering::Relaxed);
MATMUL_FLOPS.fetch_add(2 * batch * m * k * n, Ordering::Relaxed);
BYTES_MOVED.fetch_add(4 * batch * (m * k + k * n + m * n), Ordering::Relaxed);
}
pub fn elementwise(n: u64, reads: u64, writes: u64) {
if !on() {
return;
}
ELEM_CALLS.fetch_add(1, Ordering::Relaxed);
ELEM_OPS.fetch_add(n, Ordering::Relaxed);
BYTES_MOVED.fetch_add(4 * n * (reads + writes), Ordering::Relaxed);
}
pub fn transcendental_vector(n: u64) {
if !on() {
return;
}
TRANSCENDENTAL_VEC.fetch_add(n, Ordering::Relaxed);
}
pub fn transcendental_scalar(n: u64) {
if !on() {
return;
}
TRANSCENDENTAL.fetch_add(n, Ordering::Relaxed);
}
pub fn copy(n: u64) {
if !on() {
return;
}
COPIES.fetch_add(1, Ordering::Relaxed);
COPY_BYTES.fetch_add(4 * n * 2, Ordering::Relaxed);
BYTES_MOVED.fetch_add(4 * n * 2, Ordering::Relaxed);
}
impl Costs {
#[must_use]
pub const fn weighted(&self) -> u64 {
const ELEM_WEIGHT: u64 = 264; const TRANS_WEIGHT: u64 = 8800; const TRANS_VEC_WEIGHT: u64 = 244; self.matmul_flops
+ self.elem_ops * ELEM_WEIGHT
+ self.transcendental * TRANS_WEIGHT
+ self.transcendental_vec * TRANS_VEC_WEIGHT
}
#[must_use]
pub fn report(&self, label: &str) -> String {
format!(
"{label:<22} matmul {:>7.1} GF in {:>4} calls | elem {:>7.1} M in {:>3} calls | \
transc {:>6.1} M scalar / {:>6.1} M vec | moved {:>7.1} MB | copies {:>3} ({:>6.1} MB) | weighted {:>8.1} G",
self.matmul_flops as f64 / 1e9,
self.matmul_calls,
self.elem_ops as f64 / 1e6,
self.elem_calls,
self.transcendental as f64 / 1e6,
self.transcendental_vec as f64 / 1e6,
self.bytes_moved as f64 / 1e6,
self.copies,
self.copy_bytes as f64 / 1e6,
self.weighted() as f64 / 1e9,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
static SERIAL: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn guard() -> std::sync::MutexGuard<'static, ()> {
SERIAL
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[test]
fn counting_is_off_until_asked() {
let _g = guard();
start();
let _ = stop();
matmul(1, 10, 10, 10);
transcendental_scalar(99);
assert_eq!(stop(), Costs::default(), "counted while disabled");
}
#[test]
fn a_matmul_costs_two_flops_per_multiply_add() {
let _g = guard();
start();
matmul(1, 2, 3, 4);
let c = stop();
assert_eq!(c.matmul_flops, 2 * 2 * 3 * 4);
assert_eq!(c.matmul_calls, 1);
assert_eq!(c.bytes_moved, 4 * (2 * 3 + 3 * 4 + 2 * 4));
}
#[test]
fn the_weighting_stops_a_bad_trade_reading_as_a_win() {
let _g = guard();
start();
elementwise(12_000_000, 1, 1);
let before = stop();
start();
transcendental_scalar(3_000_000);
let after = stop();
assert!(
after.weighted() > before.weighted(),
"3 M transcendentals ({}) should outweigh 12 M elementwise visits ({})",
after.weighted(),
before.weighted()
);
}
#[test]
fn counters_are_reproducible_across_runs() {
let _g = guard();
let run = || {
start();
matmul(12, 1024, 64, 1024);
elementwise(3_145_728, 1, 1);
transcendental_scalar(3_145_728);
copy(786_432);
stop()
};
assert_eq!(run(), run(), "counters are not deterministic");
}
}