use super::options::ForwardPassPerfOptions;
use super::samples::{AxisRange, BucketedSamples, StoreStats, WithOptions, median_ratio};
#[derive(Clone, Debug)]
pub(crate) struct CorrectionBuckets {
samples: BucketedSamples<CorrectionObservation>,
min_observations: usize,
min_faster_correction_factor: Option<f64>,
max_slower_correction_factor: Option<f64>,
}
#[derive(Clone, Copy, Debug)]
struct CorrectionObservation {
correction_factor: f64,
}
impl WithOptions for CorrectionBuckets {
fn with_options(options: &ForwardPassPerfOptions, axis_ranges: &[AxisRange]) -> Self {
Self {
samples: BucketedSamples::new_fixed(options, axis_ranges),
min_observations: options.min_observations,
min_faster_correction_factor: options.min_faster_correction_factor,
max_slower_correction_factor: options.max_slower_correction_factor,
}
}
}
impl StoreStats for CorrectionBuckets {
fn observation_count(&self) -> usize {
self.samples.total_observations
}
fn is_ready(&self) -> bool {
self.samples.total_observations >= self.min_observations
}
}
impl CorrectionBuckets {
pub(crate) fn add_observation(&mut self, x: Vec<f64>, observed_ms: f64, native_ms: f64) {
if native_ms.is_finite() && native_ms > 0.0 && observed_ms.is_finite() && observed_ms > 0.0
{
let correction_factor = observed_ms / native_ms;
let lower_bounded_correction_factor = self
.min_faster_correction_factor
.map_or(correction_factor, |min_factor| {
correction_factor.max(min_factor)
});
let bounded_correction_factor = self
.max_slower_correction_factor
.map_or(lower_bounded_correction_factor, |max_factor| {
lower_bounded_correction_factor.min(max_factor)
});
self.samples.add(
x,
CorrectionObservation {
correction_factor: bounded_correction_factor,
},
);
}
}
pub(crate) fn correction_factor_for(&self, x: &[f64]) -> f64 {
if !self.is_ready() {
return 1.0;
}
let Some(key) = self.samples.bucket_key_if_in_bounds(x) else {
return 1.0;
};
let Some(bucket) = self.samples.buckets.get(&key) else {
return 1.0;
};
median_ratio(
bucket
.iter()
.map(|(_, observation)| observation.correction_factor),
)
.unwrap_or(1.0)
}
pub(crate) fn ready_bucket_count(&self) -> usize {
if self.is_ready() {
self.samples.buckets.len()
} else {
0
}
}
pub(crate) fn correction_factors(&self) -> Vec<f64> {
if !self.is_ready() {
return Vec::new();
}
self.samples
.buckets
.values()
.filter_map(|bucket| {
median_ratio(
bucket
.iter()
.map(|(_, observation)| observation.correction_factor),
)
})
.collect()
}
}