Skip to main content

multivector/
calibration.rs

1use super::*;
2
3const MIN_OBSERVATIONS: u64 = 5;
4const ALPHA: f64 = 0.2;
5
6#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
7#[serde(tag = "kind", content = "operator", rename_all = "snake_case")]
8pub enum CalibrationTarget {
9    Channel(PhysicalOperator),
10    Fusion(FusionOperator),
11    RerankMaxsim,
12    Context(ContextOperator),
13}
14
15impl CalibrationTarget {
16    fn sort_key(self) -> String {
17        match self {
18            Self::Channel(operator) => format!("channel/{}", operator.as_str()),
19            Self::Fusion(operator) => format!("fusion/{operator:?}"),
20            Self::RerankMaxsim => "rerank/maxsim".into(),
21            Self::Context(operator) => format!("context/{operator:?}"),
22        }
23    }
24}
25
26#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
27pub struct CalibrationKey {
28    pub target: CalibrationTarget,
29    pub dimension_bucket: usize,
30    pub corpus_bucket: usize,
31    pub selectivity_bucket: u8,
32}
33
34#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
35pub struct CalibrationEntry {
36    pub key: CalibrationKey,
37    pub observations: u64,
38    pub mean_ms_per_cost_unit: f64,
39    pub variance_ms_per_cost_unit: f64,
40    pub p90_ms_per_cost_unit: f64,
41}
42
43#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
44pub struct CalibrationSnapshot {
45    pub entries: Vec<CalibrationEntry>,
46}
47
48#[derive(Clone, Debug, Default)]
49struct RunningCalibration {
50    observations: u64,
51    mean: f64,
52    variance: f64,
53}
54
55/// Number of observations between automatic calibration flushes to disk.
56/// Low enough to avoid losing a full warm-up on crashes; high enough not
57/// to dominate retrieve() latency with file I/O.
58pub(super) const AUTOSAVE_INTERVAL: u64 = 50;
59
60#[derive(Debug, Default)]
61pub(super) struct CalibrationStats {
62    entries: HashMap<CalibrationKey, RunningCalibration>,
63    /// Total observations since the last save, used to trigger autosave.
64    pub(super) since_save: u64,
65}
66
67impl CalibrationStats {
68    pub(super) fn observe(&mut self, key: CalibrationKey, cost_units: f64, elapsed_ms: f64) {
69        if !cost_units.is_finite()
70            || cost_units <= 0.0
71            || !elapsed_ms.is_finite()
72            || elapsed_ms < 0.0
73        {
74            return;
75        }
76        let sample = elapsed_ms / cost_units;
77        let entry = self.entries.entry(key).or_default();
78        entry.observations += 1;
79        self.since_save += 1;
80        if entry.observations == 1 {
81            entry.mean = sample;
82            return;
83        }
84        let delta = sample - entry.mean;
85        entry.mean += ALPHA * delta;
86        entry.variance = (1.0 - ALPHA) * (entry.variance + ALPHA * delta * delta);
87    }
88
89    pub(super) fn snapshot(&self) -> CalibrationSnapshot {
90        let mut entries = self
91            .entries
92            .iter()
93            .map(|(&key, value)| CalibrationEntry {
94                key,
95                observations: value.observations,
96                mean_ms_per_cost_unit: value.mean,
97                variance_ms_per_cost_unit: value.variance,
98                p90_ms_per_cost_unit: value.mean + 1.282 * value.variance.sqrt(),
99            })
100            .collect::<Vec<_>>();
101        entries.sort_by_key(|entry| {
102            (
103                entry.key.target.sort_key(),
104                entry.key.dimension_bucket,
105                entry.key.corpus_bucket,
106                entry.key.selectivity_bucket,
107            )
108        });
109        CalibrationSnapshot { entries }
110    }
111
112    pub(super) fn restore(&mut self, snapshot: &CalibrationSnapshot) {
113        self.entries.clear();
114        for entry in &snapshot.entries {
115            self.entries.insert(
116                entry.key,
117                RunningCalibration {
118                    observations: entry.observations,
119                    mean: entry.mean_ms_per_cost_unit,
120                    variance: entry.variance_ms_per_cost_unit,
121                },
122            );
123        }
124        self.since_save = 0;
125    }
126}
127
128impl CalibrationSnapshot {
129    pub(super) fn estimate_ms(&self, key: CalibrationKey, cost_units: f64) -> Option<f64> {
130        let entry = self
131            .entries
132            .iter()
133            .find(|entry| entry.key == key && entry.observations >= MIN_OBSERVATIONS)?;
134        Some(entry.p90_ms_per_cost_unit * cost_units)
135    }
136}
137
138pub(super) fn key(
139    target: CalibrationTarget,
140    dimension: usize,
141    documents: usize,
142    selectivity: Option<f32>,
143) -> CalibrationKey {
144    CalibrationKey {
145        target,
146        dimension_bucket: bucket(dimension, &[128, 256, 512, 1_024, 2_048]),
147        corpus_bucket: bucket(documents, &[1_000, 10_000, 100_000, 1_000_000]),
148        selectivity_bucket: match selectivity.unwrap_or(1.0) {
149            value if value <= 0.01 => 1,
150            value if value <= 0.05 => 5,
151            value if value <= 0.2 => 20,
152            value if value <= 0.5 => 50,
153            _ => 100,
154        },
155    }
156}
157
158fn bucket(value: usize, boundaries: &[usize]) -> usize {
159    boundaries
160        .iter()
161        .copied()
162        .find(|boundary| value <= *boundary)
163        .unwrap_or(usize::MAX)
164}
165
166#[cfg(test)]
167mod tests {
168    use super::*;
169
170    #[test]
171    fn calibration_is_gated_and_stratified() {
172        let dense = key(
173            CalibrationTarget::Channel(PhysicalOperator::ExactDense),
174            128,
175            10_000,
176            Some(0.1),
177        );
178        let filtered = key(
179            CalibrationTarget::Channel(PhysicalOperator::ExactDense),
180            128,
181            10_000,
182            Some(0.01),
183        );
184        let mut stats = CalibrationStats::default();
185        for _ in 0..4 {
186            stats.observe(dense, 100.0, 2.0);
187        }
188        assert_eq!(stats.snapshot().estimate_ms(dense, 200.0), None);
189        stats.observe(dense, 100.0, 2.0);
190        let snapshot = stats.snapshot();
191        assert_eq!(snapshot.estimate_ms(dense, 200.0), Some(4.0));
192        assert_eq!(snapshot.estimate_ms(filtered, 200.0), None);
193    }
194
195    #[test]
196    fn p90_penalizes_variable_operators() {
197        let key = key(
198            CalibrationTarget::Channel(PhysicalOperator::HnswDense),
199            768,
200            1_000_000,
201            None,
202        );
203        let mut stats = CalibrationStats::default();
204        for elapsed in [1.0, 1.0, 1.0, 1.0, 5.0] {
205            stats.observe(key, 1.0, elapsed);
206        }
207        let entry = &stats.snapshot().entries[0];
208        assert!(entry.p90_ms_per_cost_unit > entry.mean_ms_per_cost_unit);
209    }
210}