use crate::error::Error;
use cuda_core::Stream;
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct BenchOptions {
pub warmup: Duration,
pub rep: Duration,
pub min_reps: usize,
pub max_reps: usize,
pub clear_l2: bool,
}
impl Default for BenchOptions {
fn default() -> Self {
Self {
warmup: Duration::from_millis(25),
rep: Duration::from_millis(100),
min_reps: 5,
max_reps: 1000,
clear_l2: true,
}
}
}
#[derive(Debug, Clone)]
pub struct Measurement {
times_ms: Vec<f32>,
}
impl Measurement {
#[cfg(all(test, feature = "experimental-tune"))]
pub(crate) fn from_times_ms(times_ms: Vec<f32>) -> Self {
Self { times_ms }
}
pub fn reps(&self) -> usize {
self.times_ms.len()
}
pub fn times_ms(&self) -> &[f32] {
&self.times_ms
}
pub fn min_ms(&self) -> f32 {
self.times_ms.iter().copied().fold(f32::INFINITY, f32::min)
}
pub fn mean_ms(&self) -> f32 {
self.times_ms.iter().sum::<f32>() / self.times_ms.len() as f32
}
pub fn median_ms(&self) -> f32 {
self.quantile_ms(0.5)
}
pub fn quantile_ms(&self, q: f32) -> f32 {
let mut sorted = self.times_ms.clone();
sorted.sort_by(|a, b| a.total_cmp(b));
let q = q.clamp(0.0, 1.0);
let pos = q * (sorted.len() - 1) as f32;
let lo = pos.floor() as usize;
let hi = pos.ceil() as usize;
if lo == hi {
sorted[lo]
} else {
let frac = pos - lo as f32;
sorted[lo] * (1.0 - frac) + sorted[hi] * frac
}
}
}
struct L2Clear {
dptr: cuda_core::sys::CUdeviceptr,
num_bytes: usize,
stream: Arc<Stream>,
}
impl L2Clear {
fn new(stream: &Arc<Stream>) -> Self {
let l2 = stream.device().l2_cache_size_bytes().unwrap_or(0);
let num_bytes = (l2 * 2).clamp(64 << 20, 512 << 20);
let dptr = unsafe { cuda_core::malloc_async(num_bytes, stream) };
Self {
dptr,
num_bytes,
stream: stream.clone(),
}
}
fn clear(&self) -> Result<(), Error> {
unsafe { cuda_core::memset_d8_async(self.dptr, 0, self.num_bytes, &self.stream) }?;
Ok(())
}
}
impl Drop for L2Clear {
fn drop(&mut self) {
unsafe { cuda_core::free_async(self.dptr, &self.stream) };
}
}
fn time_one<F>(stream: &Arc<Stream>, f: &mut F) -> Result<f32, Error>
where
F: FnMut(&Arc<Stream>) -> Result<(), Error>,
{
let device = stream.device();
let start = device.new_event()?;
let end = device.new_event()?;
start.record(stream)?;
f(stream)?;
end.record(stream)?;
end.synchronize()?;
Ok(start.elapsed_time(&end)?)
}
pub fn do_bench<F>(
stream: &Arc<Stream>,
opts: &BenchOptions,
mut f: F,
) -> Result<Measurement, Error>
where
F: FnMut(&Arc<Stream>) -> Result<(), Error>,
{
let warmup_start = Instant::now();
f(stream)?;
while warmup_start.elapsed() < opts.warmup {
f(stream)?;
}
let est_ms = time_one(stream, &mut f)?.max(1e-4);
let target_ms = opts.rep.as_secs_f64() * 1e3;
let reps = ((target_ms / est_ms as f64).round() as usize).clamp(opts.min_reps, opts.max_reps);
let l2 = opts.clear_l2.then(|| L2Clear::new(stream));
let mut times_ms = Vec::with_capacity(reps);
for _ in 0..reps {
if let Some(l2) = &l2 {
l2.clear()?;
}
times_ms.push(time_one(stream, &mut f)?);
}
Ok(Measurement { times_ms })
}
pub fn do_bench_paired<A, B>(
stream: &Arc<Stream>,
opts: &BenchOptions,
mut a: A,
mut b: B,
) -> Result<(Measurement, Measurement), Error>
where
A: FnMut(&Arc<Stream>) -> Result<(), Error>,
B: FnMut(&Arc<Stream>) -> Result<(), Error>,
{
let warmup_start = Instant::now();
a(stream)?;
b(stream)?;
while warmup_start.elapsed() < opts.warmup {
a(stream)?;
b(stream)?;
}
let est_ms = time_one(stream, &mut a)?.max(1e-4);
let target_ms = opts.rep.as_secs_f64() * 1e3;
let reps =
(((target_ms / est_ms as f64) / 2.0).round() as usize).clamp(opts.min_reps, opts.max_reps);
let l2 = opts.clear_l2.then(|| L2Clear::new(stream));
let (mut times_a, mut times_b) = (Vec::with_capacity(reps), Vec::with_capacity(reps));
for _ in 0..reps {
if let Some(l2) = &l2 {
l2.clear()?;
}
times_a.push(time_one(stream, &mut a)?);
if let Some(l2) = &l2 {
l2.clear()?;
}
times_b.push(time_one(stream, &mut b)?);
}
Ok((
Measurement { times_ms: times_a },
Measurement { times_ms: times_b },
))
}
#[cfg(test)]
mod tests {
use super::*;
fn result(times: &[f32]) -> Measurement {
Measurement {
times_ms: times.to_vec(),
}
}
#[test]
fn quantiles_interpolate_and_clamp() {
let r = result(&[4.0, 1.0, 3.0, 2.0]);
assert_eq!(r.min_ms(), 1.0);
assert_eq!(r.median_ms(), 2.5);
assert_eq!(r.quantile_ms(0.0), 1.0);
assert_eq!(r.quantile_ms(1.0), 4.0);
assert_eq!(r.quantile_ms(-1.0), 1.0);
assert_eq!(r.quantile_ms(2.0), 4.0);
assert!((r.mean_ms() - 2.5).abs() < 1e-6);
}
#[test]
fn single_sample_quantiles_are_that_sample() {
let r = result(&[7.0]);
assert_eq!(r.median_ms(), 7.0);
assert_eq!(r.quantile_ms(0.25), 7.0);
}
#[test]
fn median_is_robust_to_one_outlier() {
let r = result(&[1.0, 1.0, 1.0, 1.0, 100.0]);
assert_eq!(r.median_ms(), 1.0);
assert!(r.mean_ms() > 20.0);
}
}