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
55pub(super) const AUTOSAVE_INTERVAL: u64 = 50;
59
60#[derive(Debug, Default)]
61pub(super) struct CalibrationStats {
62 entries: HashMap<CalibrationKey, RunningCalibration>,
63 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}