use std::time::Duration;
#[derive(Clone)]
pub struct Metrics {
durations: Vec<Duration>,
warmup_duration: Option<Duration>,
}
impl Metrics {
pub fn new() -> Metrics {
Metrics {
durations: Vec::new(),
warmup_duration: None,
}
}
pub fn add_step_duration(&mut self, duration: Duration) {
if self.warmup_duration.is_some() {
self.durations.push(duration);
} else {
self.warmup_duration = Some(duration);
}
}
pub fn warmup_duration(&self) -> Option<Duration> {
self.warmup_duration
}
pub fn step_durations(&self) -> &[Duration] {
&self.durations
}
pub fn token_count(&self) -> usize {
if self.warmup_duration.is_none() {
0
} else {
1 + self.durations.len()
}
}
pub fn total_duration(&self) -> Duration {
self.durations.iter().sum::<Duration>() + self.warmup_duration.unwrap_or(Duration::ZERO)
}
pub fn total_main_duration(&self) -> Duration {
self.durations.iter().sum()
}
pub fn mean_duration(&self) -> f32 {
let total_ms = self.total_main_duration().as_secs_f64() * 1000.0;
(total_ms / self.durations.len() as f64) as f32
}
pub fn tokens_per_second(&self) -> f32 {
self.durations.len() as f32 / self.total_main_duration().as_secs_f32()
}
}
impl Default for Metrics {
fn default() -> Self {
Metrics::new()
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::Metrics;
macro_rules! assert_approx_eq {
($a:expr, $b:expr, $threshold: expr) => {
let a = $a;
let b = $b;
let threshold = $threshold;
assert!(
(a - b).abs() < threshold,
"values {} and {} not approximately equal",
a,
b
)
};
}
#[test]
fn test_metrics() {
let ms = Duration::from_millis;
let mut metrics = Metrics::new();
metrics.add_step_duration(ms(200));
metrics.add_step_duration(ms(110));
metrics.add_step_duration(ms(90));
assert_eq!(metrics.warmup_duration(), Some(ms(200)));
assert_eq!(metrics.step_durations(), &[ms(110), ms(90)]);
assert_eq!(metrics.total_duration(), ms(400));
assert_eq!(metrics.total_main_duration(), ms(200));
assert_approx_eq!(metrics.mean_duration(), 100.0, 1e-5);
assert_eq!(metrics.token_count(), 3);
assert_approx_eq!(metrics.tokens_per_second(), 10.0, 1e-5);
}
}