use core::mem::MaybeUninit;
use std::hint::black_box;
use std::time::Instant;
use strided_kernel::{
erased_map_into_uninit, erased_zip_into_uninit, map_into, zip_map2_into, ErasedMapOp,
ErasedRawStridedPtr, ErasedRawStridedRef, ErasedRawStridedUninitMut, ErasedZipOp, ExecContext,
KernelDType, StridedView, StridedViewMut,
};
const WARMUPS: usize = 6;
const SAMPLES: usize = 21;
fn values(len: usize, offset: f64) -> Vec<f64> {
(0..len)
.map(|index| (index % 251) as f64 * 0.00390625 - 0.5 + offset)
.collect()
}
fn repeats(len: usize) -> usize {
(1 << 20) / len.max(1) + 1
}
fn time_ms(repeats: usize, operation: &mut impl FnMut()) -> f64 {
let start = Instant::now();
for _ in 0..repeats {
operation();
}
start.elapsed().as_secs_f64() * 1e3 / repeats as f64
}
fn median(samples: &[f64]) -> f64 {
let mut ordered = samples.to_vec();
ordered.sort_by(f64::total_cmp);
ordered[ordered.len() / 2]
}
fn measure_pair(case: &str, len: usize, mut typed: impl FnMut(), mut erased: impl FnMut()) {
let repeats = repeats(len);
for _ in 0..WARMUPS {
typed();
erased();
}
let mut typed_samples = Vec::with_capacity(SAMPLES);
let mut erased_samples = Vec::with_capacity(SAMPLES);
for sample in 0..SAMPLES {
if sample % 2 == 0 {
typed_samples.push(time_ms(repeats, &mut typed));
erased_samples.push(time_ms(repeats, &mut erased));
} else {
erased_samples.push(time_ms(repeats, &mut erased));
typed_samples.push(time_ms(repeats, &mut typed));
}
}
let logs: Vec<f64> = typed_samples
.iter()
.zip(&erased_samples)
.map(|(&typed, &erased)| (erased / typed).ln())
.collect();
let mean = logs.iter().sum::<f64>() / logs.len() as f64;
let variance =
logs.iter().map(|value| (value - mean).powi(2)).sum::<f64>() / (logs.len() - 1) as f64;
let upper95 = (mean + 1.96 * (variance / logs.len() as f64).sqrt()).exp();
println!(
"SUMMARY,{case},{:.6},{:.6},{:.3},{upper95:.3}",
median(&typed_samples),
median(&erased_samples),
mean.exp()
);
}
#[derive(Clone, Copy)]
enum Layout {
Vector,
Matrix,
TransposedLhs,
}
impl Layout {
fn label(self) -> &'static str {
match self {
Self::Vector => "vector",
Self::Matrix => "matrix",
Self::TransposedLhs => "transposed_lhs",
}
}
fn describe(self, len: usize) -> (Vec<usize>, Vec<isize>, Vec<isize>) {
match self {
Self::Vector => (vec![len], vec![1], vec![1]),
Self::Matrix | Self::TransposedLhs => {
let side = (len as f64).sqrt() as usize;
let col_major = vec![1, side as isize];
let lhs = match self {
Self::TransposedLhs => vec![side as isize, 1],
_ => col_major.clone(),
};
(vec![side, side], lhs, col_major)
}
}
}
}
fn bench_zip<F>(op: ErasedZipOp, label: &str, typed_op: F, layout: Layout, len: usize)
where
F: Fn(f64, f64) -> f64 + Copy + Send + Sync,
{
let (dims, lhs_strides, strides) = layout.describe(len);
let len: usize = dims.iter().product();
let lhs = values(len, 1.0);
let rhs = values(len, 2.0);
let mut typed_out = vec![MaybeUninit::<f64>::uninit(); len];
let mut erased_out = vec![MaybeUninit::<f64>::uninit(); len];
let lhs_view = StridedView::<f64>::new(&lhs, &dims, &lhs_strides, 0).unwrap();
let rhs_view = StridedView::<f64>::new(&rhs, &dims, &strides, 0).unwrap();
let lhs_ref = ErasedRawStridedRef::from_slice(&lhs, &dims, &lhs_strides, 0).unwrap();
let rhs_ref = ErasedRawStridedRef::from_slice(&rhs, &dims, &strides, 0).unwrap();
let ctx = ExecContext::serial();
measure_pair(
&format!("zip_{label},{},{len}", layout.label()),
len,
|| {
let mut dest = StridedViewMut::new(&mut typed_out, &dims, &strides, 0).unwrap();
zip_map2_into(&mut dest, &lhs_view, &rhs_view, |a, b| {
MaybeUninit::new(typed_op(a, b))
})
.unwrap();
black_box(&typed_out);
},
|| {
let mut dest =
ErasedRawStridedUninitMut::from_uninit_slice(&mut erased_out, &dims, &strides, 0)
.unwrap();
erased_zip_into_uninit(
KernelDType::F64,
op,
&ctx,
&mut dest,
&ErasedRawStridedPtr::from_ref(&lhs_ref),
&ErasedRawStridedPtr::from_ref(&rhs_ref),
)
.unwrap();
black_box(&erased_out);
},
);
}
fn bench_map<F>(op: ErasedMapOp, label: &str, typed_op: F, len: usize)
where
F: Fn(f64) -> f64 + Copy + Send + Sync,
{
let dims = [len];
let strides = [1];
let input = values(len, 0.0);
let mut typed_out = vec![MaybeUninit::<f64>::uninit(); len];
let mut erased_out = vec![MaybeUninit::<f64>::uninit(); len];
let input_view = StridedView::<f64>::new(&input, &dims, &strides, 0).unwrap();
let input_ref = ErasedRawStridedRef::from_slice(&input, &dims, &strides, 0).unwrap();
let ctx = ExecContext::serial();
measure_pair(
&format!("map_{label},vector,{len}"),
len,
|| {
let mut dest = StridedViewMut::new(&mut typed_out, &dims, &strides, 0).unwrap();
map_into(&mut dest, &input_view, |a| MaybeUninit::new(typed_op(a))).unwrap();
black_box(&typed_out);
},
|| {
let mut dest =
ErasedRawStridedUninitMut::from_uninit_slice(&mut erased_out, &dims, &strides, 0)
.unwrap();
erased_map_into_uninit(
KernelDType::F64,
op,
&ctx,
&mut dest,
&ErasedRawStridedPtr::from_ref(&input_ref),
)
.unwrap();
black_box(&erased_out);
},
);
}
fn main() {
println!("CONFIG,warmups={WARMUPS},samples={SAMPLES},dtype=f64,context=serial");
println!("HEADER,case,layout,len,typed_ms,erased_ms,ratio,upper95");
let sizes = [1usize << 12, 1 << 16, 1 << 20];
for &len in &sizes {
for layout in [Layout::Vector, Layout::Matrix, Layout::TransposedLhs] {
bench_zip(ErasedZipOp::Add, "add", |a, b| a + b, layout, len);
bench_zip(ErasedZipOp::Multiply, "multiply", |a, b| a * b, layout, len);
}
bench_zip(
ErasedZipOp::Maximum,
"maximum",
|a, b| {
if a.is_nan() || b.is_nan() {
f64::NAN
} else if a >= b {
a
} else {
b
}
},
Layout::Vector,
len,
);
bench_map(ErasedMapOp::Negate, "negate", |a| -a, len);
bench_map(ErasedMapOp::Abs, "abs", |a: f64| a.abs(), len);
}
}