use std::hint::black_box;
use std::sync::Arc;
use std::time::{Duration, Instant};
use arrow::array::{Array as _, ArrayRef, BooleanArray, Int64Array, Scalar as ArrowScalar};
use arrow::compute::kernels::cmp;
use arrowmetal::CompareOp;
const N: usize = 10_000_000;
const WARMUP: usize = 3;
const TIMED: usize = 5;
fn measure(mut f: impl FnMut()) -> (Duration, Duration) {
for _ in 0..WARMUP {
f();
}
let mut times = Vec::with_capacity(TIMED);
for _ in 0..TIMED {
let t = Instant::now();
f();
times.push(t.elapsed());
}
times.sort();
(times[0], times[TIMED / 2])
}
fn ms(d: Duration) -> f64 {
d.as_secs_f64() * 1000.0
}
fn data(n: usize) -> Int64Array {
let mut state = 0x2545_F491_4F6C_DD1Du64;
(0..n)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
Some((state % 2_000_000) as i64 - 1_000_000)
})
.collect()
}
fn main() {
println!("ArrowMetal {} on {}", arrowmetal::version(), arrowmetal::device_name());
println!("rustc {}, arrow-rs 59, {N} Int64 rows, no nulls", rustc_version());
println!("best of {TIMED} after {WARMUP} warm-up iterations, wall time, single-threaded\n");
let values = data(N);
let values_ref: ArrayRef = Arc::new(values.clone());
let arrow_mask: BooleanArray =
cmp::gt(&values, &ArrowScalar::new(&Int64Array::from(vec![0i64]))).unwrap();
let selected = arrow_mask.true_count();
println!("filter predicate: x > 0, {selected} of {N} rows kept ({:.1}%)\n",
100.0 * selected as f64 / N as f64);
let gpu = arrowmetal::Array::from_arrow(&values).unwrap();
let gpu_mask = gpu.compare_scalar(CompareOp::Gt, 0i64).unwrap();
let arrow_sum = arrow::compute::sum(&values).unwrap();
let gpu_sum = gpu.sum().unwrap().unwrap().as_i64().unwrap();
assert_eq!(arrow_sum, gpu_sum, "sum disagrees; the timing below would be meaningless");
let arrow_filtered = arrow::compute::filter(&values, &arrow_mask).unwrap();
let gpu_filtered = gpu.filter(&gpu_mask).unwrap().to_arrow().unwrap();
assert_eq!(&arrow_filtered, &gpu_filtered, "filter disagrees");
println!("both libraries agree on both answers\n");
let (a_sum, a_sum_med) = measure(|| {
black_box(arrow::compute::sum(black_box(&values)));
});
let (g_sum, g_sum_med) = measure(|| {
black_box(black_box(&gpu).sum().unwrap());
});
let (e_sum, e_sum_med) = measure(|| {
let a = arrowmetal::Array::from_arrow(black_box(&values)).unwrap();
black_box(a.sum().unwrap());
});
let (a_filter, a_filter_med) = measure(|| {
black_box(arrow::compute::filter(black_box(&values), black_box(&arrow_mask)).unwrap());
});
let (g_filter, g_filter_med) = measure(|| {
black_box(black_box(&gpu).filter(black_box(&gpu_mask)).unwrap());
});
let (e_filter, e_filter_med) = measure(|| {
let a = arrowmetal::Array::from_arrow(black_box(&values)).unwrap();
let m = a.compare_scalar(CompareOp::Gt, 0i64).unwrap();
black_box(a.filter(&m).unwrap().to_arrow().unwrap());
});
let (a_cmp_filter, a_cmp_filter_med) = measure(|| {
let m = cmp::gt(black_box(&values), &ArrowScalar::new(&Int64Array::from(vec![0i64]))).unwrap();
black_box(arrow::compute::filter(black_box(&values), &m).unwrap());
});
let (imp, imp_med) = measure(|| {
black_box(arrowmetal::Array::from_arrow(black_box(&values)).unwrap());
});
let (exp, exp_med) = measure(|| {
black_box(black_box(&gpu).to_arrow().unwrap());
});
println!("| Operation, 10M Int64 | arrow-rs | ArrowMetal, kernel | ArrowMetal, end to end |");
println!("|---|---|---|---|");
println!(
"| `sum` | {:.2} ms | {:.2} ms | {:.2} ms |",
ms(a_sum),
ms(g_sum),
ms(e_sum)
);
println!(
"| `filter` (mask ready) | {:.2} ms | {:.2} ms | -- |",
ms(a_filter),
ms(g_filter)
);
println!(
"| compare + `filter` | {:.2} ms | -- | {:.2} ms |",
ms(a_cmp_filter),
ms(e_filter)
);
let values_alignment = {
let data = values.to_data();
let p = data.buffers()[0].as_ptr() as usize;
1usize << p.trailing_zeros()
};
println!("\nSupporting numbers (best of {TIMED}):");
println!(
" import (arrow-rs -> ArrowMetal): {:.3} ms [values buffer aligned to {values_alignment} B; \
copy-free needs 16384]",
ms(imp)
);
println!(" export (ArrowMetal -> arrow-rs, no copy): {:.3} ms", ms(exp));
println!("\nMedians, for comparison with the bests above:");
println!(" arrow-rs sum {:.2} ms | ArrowMetal kernel {:.2} ms | end to end {:.2} ms",
ms(a_sum_med), ms(g_sum_med), ms(e_sum_med));
println!(" arrow-rs filter {:.2} ms | ArrowMetal kernel {:.2} ms", ms(a_filter_med), ms(g_filter_med));
println!(" arrow-rs cmp+filter {:.2} ms | ArrowMetal end to end {:.2} ms",
ms(a_cmp_filter_med), ms(e_filter_med));
println!(" import {:.3} ms | export {:.3} ms", ms(imp_med), ms(exp_med));
println!("\nvalues live: {} rows, {} filtered", values_ref.len(), gpu_filtered.len());
}
fn rustc_version() -> String {
std::process::Command::new("rustc")
.arg("--version")
.output()
.ok()
.and_then(|o| String::from_utf8(o.stdout).ok())
.map(|s| s.trim().to_string())
.unwrap_or_else(|| "unknown".into())
}