Files
rustytorch/crates/tooling/rtx-kernel-bench/src/metrics.rs
T
2026-03-04 00:08:42 +00:00

286 lines
8.1 KiB
Rust

//! Benchmark Metrics and Statistics
//!
//! Statistical analysis of benchmark measurements.
use serde::{Deserialize, Serialize};
use std::time::Duration;
use crate::benchmark::Measurement;
/// Timing statistics from benchmark measurements
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TimingStats {
/// Mean duration
pub mean: Duration,
/// Median duration
pub median: Duration,
/// Standard deviation
pub std_dev: Duration,
/// Minimum duration
pub min: Duration,
/// Maximum duration
pub max: Duration,
/// 5th percentile
pub p5: Duration,
/// 95th percentile
pub p95: Duration,
/// 99th percentile
pub p99: Duration,
/// Coefficient of variation (std_dev / mean)
pub cv: f64,
}
impl TimingStats {
/// Compute timing statistics from durations
pub fn from_durations(durations: &[Duration]) -> Self {
if durations.is_empty() {
return Self::zero();
}
let mut sorted: Vec<_> = durations.to_vec();
sorted.sort();
let n = sorted.len();
let sum: Duration = sorted.iter().sum();
let mean = sum / n as u32;
let median = if n.is_multiple_of(2) {
(sorted[n / 2 - 1] + sorted[n / 2]) / 2
} else {
sorted[n / 2]
};
let min = sorted[0];
let max = sorted[n - 1];
// Percentiles
let p5 = sorted[(n as f64 * 0.05) as usize];
let p95 = sorted[(n as f64 * 0.95).min((n - 1) as f64) as usize];
let p99 = sorted[(n as f64 * 0.99).min((n - 1) as f64) as usize];
// Standard deviation
let mean_nanos = mean.as_nanos() as f64;
let variance: f64 = sorted
.iter()
.map(|d| {
let diff = d.as_nanos() as f64 - mean_nanos;
diff * diff
})
.sum::<f64>()
/ n as f64;
let std_dev_nanos = variance.sqrt();
let std_dev = Duration::from_nanos(std_dev_nanos as u64);
// Coefficient of variation
let cv = if mean_nanos > 0.0 {
std_dev_nanos / mean_nanos
} else {
0.0
};
Self {
mean,
median,
std_dev,
min,
max,
p5,
p95,
p99,
cv,
}
}
/// Create zero stats
fn zero() -> Self {
Self {
mean: Duration::ZERO,
median: Duration::ZERO,
std_dev: Duration::ZERO,
min: Duration::ZERO,
max: Duration::ZERO,
p5: Duration::ZERO,
p95: Duration::ZERO,
p99: Duration::ZERO,
cv: 0.0,
}
}
}
/// Memory statistics
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct MemoryStats {
/// Peak memory usage (bytes)
pub peak_bytes: u64,
/// Average memory usage (bytes)
pub avg_bytes: u64,
/// Minimum memory usage (bytes)
pub min_bytes: u64,
/// Maximum memory usage (bytes)
pub max_bytes: u64,
}
impl MemoryStats {
/// Compute memory statistics from measurements
pub fn from_measurements(measurements: &[Measurement]) -> Self {
let memory_values: Vec<_> = measurements.iter().filter_map(|m| m.memory_bytes).collect();
if memory_values.is_empty() {
return Self::default();
}
let sum: u64 = memory_values.iter().sum();
let avg = sum / memory_values.len() as u64;
let min = *memory_values.iter().min().unwrap_or(&0);
let max = *memory_values.iter().max().unwrap_or(&0);
Self {
peak_bytes: max,
avg_bytes: avg,
min_bytes: min,
max_bytes: max,
}
}
}
/// Combined metrics from benchmark measurements
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Metrics {
/// Timing statistics
pub timing: TimingStats,
/// Memory statistics
pub memory: MemoryStats,
/// Number of measurements
pub sample_count: usize,
/// Whether measurements are statistically stable (low CV)
pub is_stable: bool,
}
impl Metrics {
/// Compute metrics from measurements
pub fn from_measurements(measurements: &[Measurement]) -> Self {
let durations: Vec<_> = measurements.iter().map(|m| m.duration).collect();
let timing = TimingStats::from_durations(&durations);
let memory = MemoryStats::from_measurements(measurements);
// Consider stable if CV < 10%
let is_stable = timing.cv < 0.10;
Self {
timing,
memory,
sample_count: measurements.len(),
is_stable,
}
}
/// Get a summary string
pub fn summary(&self) -> String {
format!(
"{:.2}ms ± {:.2}ms (n={}{})",
self.timing.mean.as_secs_f64() * 1000.0,
self.timing.std_dev.as_secs_f64() * 1000.0,
self.sample_count,
if self.is_stable { "" } else { ", unstable" }
)
}
/// Get detailed timing string
pub fn timing_detail(&self) -> String {
format!(
"mean={:.3}ms, median={:.3}ms, min={:.3}ms, max={:.3}ms, p95={:.3}ms",
self.timing.mean.as_secs_f64() * 1000.0,
self.timing.median.as_secs_f64() * 1000.0,
self.timing.min.as_secs_f64() * 1000.0,
self.timing.max.as_secs_f64() * 1000.0,
self.timing.p95.as_secs_f64() * 1000.0,
)
}
}
/// Compute speedup between two timing results
pub fn compute_speedup(baseline: &TimingStats, optimized: &TimingStats) -> f64 {
let baseline_ns = baseline.mean.as_nanos() as f64;
let optimized_ns = optimized.mean.as_nanos() as f64;
if optimized_ns > 0.0 {
baseline_ns / optimized_ns
} else {
0.0
}
}
/// Determine if a speedup is statistically significant
pub fn is_significant(speedup: f64, baseline_cv: f64, optimized_cv: f64) -> bool {
// Simple heuristic: speedup must be greater than combined uncertainty
let combined_cv = (baseline_cv.powi(2) + optimized_cv.powi(2)).sqrt();
let uncertainty = 1.0 + 2.0 * combined_cv; // 2 sigma
speedup > uncertainty || speedup < 1.0 / uncertainty
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_timing_stats() {
let durations = vec![
Duration::from_millis(10),
Duration::from_millis(12),
Duration::from_millis(11),
Duration::from_millis(9),
Duration::from_millis(10),
];
let stats = TimingStats::from_durations(&durations);
assert!(stats.mean >= Duration::from_millis(10));
assert!(stats.mean <= Duration::from_millis(11));
assert_eq!(stats.min, Duration::from_millis(9));
assert_eq!(stats.max, Duration::from_millis(12));
}
#[test]
fn test_empty_durations() {
let stats = TimingStats::from_durations(&[]);
assert_eq!(stats.mean, Duration::ZERO);
}
#[test]
fn test_metrics_stability() {
// Very stable measurements
let stable_measurements: Vec<_> = (0..100)
.map(|_| Measurement {
duration: Duration::from_millis(10),
memory_bytes: None,
throughput: None,
})
.collect();
let metrics = Metrics::from_measurements(&stable_measurements);
assert!(metrics.is_stable);
// Unstable measurements
let unstable_measurements: Vec<_> = (0..10)
.map(|i| Measurement {
duration: Duration::from_millis(10 + i * 10),
memory_bytes: None,
throughput: None,
})
.collect();
let metrics = Metrics::from_measurements(&unstable_measurements);
assert!(!metrics.is_stable);
}
#[test]
fn test_compute_speedup() {
let baseline = TimingStats::from_durations(&[Duration::from_millis(100)]);
let optimized = TimingStats::from_durations(&[Duration::from_millis(50)]);
let speedup = compute_speedup(&baseline, &optimized);
assert!((speedup - 2.0).abs() < 0.01);
}
}