1use super::{BootstrapControl, BootstrapEffectInterval, StatsError, StatsResult, exact_quantile};
4use crate::SeededSampler;
5
6#[derive(Clone, Copy, Debug, PartialEq)]
8pub struct BinaryInterval {
9 pub successes: u64,
11 pub trials: u64,
13 pub confidence_level: f64,
15 pub lower: f64,
17 pub upper: f64,
19}
20
21pub fn exact_binary_interval(
26 successes: u64,
27 trials: u64,
28 confidence_level: f64,
29) -> StatsResult<BinaryInterval> {
30 confidence(confidence_level)?;
31 if trials == 0 {
32 return Err(StatsError::ZeroTotal {
33 label: "binary trials",
34 });
35 }
36 if successes > trials {
37 return Err(StatsError::InvalidControl {
38 field: "successes",
39 reason: "must not exceed trials",
40 });
41 }
42 let alpha = (1.0 - confidence_level) / 2.0;
43 let lower = if successes == 0 {
44 0.0
45 } else {
46 bisect_probability(|p| binomial_upper_tail(successes, trials, p), alpha)
47 };
48 let upper = if successes == trials {
49 1.0
50 } else {
51 bisect_probability(|p| binomial_cdf(successes, trials, p), alpha)
52 };
53 Ok(BinaryInterval {
54 successes,
55 trials,
56 confidence_level,
57 lower,
58 upper,
59 })
60}
61
62#[derive(Clone, Debug, PartialEq)]
64pub struct ClusterSample {
65 pub id: u64,
67 pub pairs: Vec<(f64, f64)>,
69}
70
71pub fn paired_bootstrap_interval(
73 pairs: &[(f64, f64)],
74 control: BootstrapControl,
75) -> StatsResult<BootstrapEffectInterval> {
76 control_parts(control)?;
77 if pairs.is_empty() {
78 return Err(StatsError::EmptyInput {
79 metric: "paired bootstrap",
80 });
81 }
82 let effects = pair_effects(pairs, "paired bootstrap")?;
83 bootstrap_effects(&effects, control, pairs.len(), pairs.len(), 0)
84}
85
86pub fn clustered_bootstrap_interval(
92 clusters: &[ClusterSample],
93 minimum_clusters: usize,
94 control: BootstrapControl,
95) -> StatsResult<BootstrapEffectInterval> {
96 control_parts(control)?;
97 if minimum_clusters < 2 {
98 return Err(StatsError::InvalidControl {
99 field: "minimum_clusters",
100 reason: "must be at least two",
101 });
102 }
103 if clusters.len() < minimum_clusters {
104 return Err(StatsError::InsufficientInput {
105 metric: "clustered bootstrap independent clusters",
106 minimum: minimum_clusters,
107 actual: clusters.len(),
108 });
109 }
110 let mut ordered = clusters.iter().collect::<Vec<_>>();
111 ordered.sort_by_key(|cluster| cluster.id);
112 if ordered.windows(2).any(|pair| pair[0].id == pair[1].id) {
113 return Err(StatsError::InvalidControl {
114 field: "cluster ids",
115 reason: "must be unique",
116 });
117 }
118 let mut cluster_effects = Vec::with_capacity(ordered.len());
119 let mut rows = 0usize;
120 for cluster in ordered {
121 let effects = pair_effects(&cluster.pairs, "clustered bootstrap")?;
122 rows = rows
123 .checked_add(effects.len())
124 .ok_or(StatsError::WorkLimitExceeded {
125 required: u64::MAX,
126 limit: control.max_work,
127 })?;
128 cluster_effects.push(effects.iter().sum::<f64>() / effects.len() as f64);
129 }
130 bootstrap_effects(&cluster_effects, control, rows, rows, clusters.len())
131}
132
133#[derive(Clone, Copy, Debug, PartialEq)]
135pub struct RegisteredLook {
136 pub samples: usize,
138 pub alpha: f64,
140}
141
142#[derive(Clone, Debug, PartialEq)]
144pub struct RegisteredLookSequence {
145 looks: Vec<RegisteredLook>,
146 total_budget: f64,
147}
148
149#[derive(Clone, Copy, Debug, PartialEq)]
151pub struct SequentialInterval {
152 pub samples: usize,
154 pub mean: f64,
156 pub lower: f64,
158 pub upper: f64,
160 pub alpha_spent: f64,
162 pub total_budget: f64,
164}
165
166impl RegisteredLookSequence {
167 pub fn new(looks: Vec<RegisteredLook>, total_budget: f64) -> StatsResult<Self> {
169 if !total_budget.is_finite() || !(0.0..1.0).contains(&total_budget) {
170 return Err(StatsError::InvalidControl {
171 field: "total_budget",
172 reason: "must be finite and strictly between zero and one",
173 });
174 }
175 if looks.is_empty() {
176 return Err(StatsError::EmptyInput {
177 metric: "registered looks",
178 });
179 }
180 let mut previous = 0;
181 let mut spent = 0.0;
182 for look in &looks {
183 if look.samples == 0 || look.samples <= previous {
184 return Err(StatsError::InvalidControl {
185 field: "registered looks",
186 reason: "sample counts must be positive and strictly increasing",
187 });
188 }
189 if !look.alpha.is_finite() || look.alpha <= 0.0 || look.alpha >= 1.0 {
190 return Err(StatsError::InvalidControl {
191 field: "look alpha",
192 reason: "must be finite and strictly between zero and one",
193 });
194 }
195 previous = look.samples;
196 spent += look.alpha;
197 }
198 if spent > total_budget + f64::EPSILON * looks.len() as f64 {
199 return Err(StatsError::InvalidControl {
200 field: "look alpha",
201 reason: "sum must not exceed total_budget",
202 });
203 }
204 Ok(Self {
205 looks,
206 total_budget,
207 })
208 }
209
210 pub fn interval(&self, observations: &[f64]) -> StatsResult<SequentialInterval> {
212 let look = self
213 .looks
214 .iter()
215 .find(|look| look.samples == observations.len())
216 .ok_or(StatsError::InvalidControl {
217 field: "observations",
218 reason: "sample count is not a registered look",
219 })?;
220 for (index, value) in observations.iter().enumerate() {
221 if !value.is_finite() || !(0.0..=1.0).contains(value) {
222 return Err(StatsError::NonFinite {
223 metric: "sequential bounded observation",
224 index: Some(index),
225 value: *value,
226 });
227 }
228 }
229 let mean = observations.iter().sum::<f64>() / observations.len() as f64;
230 let radius = ((2.0 / look.alpha).ln() / (2.0 * observations.len() as f64)).sqrt();
231 Ok(SequentialInterval {
232 samples: observations.len(),
233 mean,
234 lower: (mean - radius).max(0.0),
235 upper: (mean + radius).min(1.0),
236 alpha_spent: look.alpha,
237 total_budget: self.total_budget,
238 })
239 }
240}
241
242#[derive(Clone, Copy, Debug, PartialEq)]
244pub struct IsotonicPoint {
245 pub level: f64,
247 pub value: f64,
249 pub weight: f64,
251}
252
253#[derive(Clone, Copy, Debug, PartialEq)]
255pub enum ThresholdReadout {
256 Observed {
258 level: f64,
260 },
261 BelowTestedRange,
263 AboveTestedRange,
265}
266
267#[derive(Clone, Debug, PartialEq)]
269pub struct IsotonicFit {
270 pub raw: Vec<IsotonicPoint>,
272 pub fitted: Vec<f64>,
274 pub normalized_area: Option<f64>,
276}
277
278impl IsotonicFit {
279 pub fn threshold(&self, threshold: f64) -> StatsResult<ThresholdReadout> {
281 if !threshold.is_finite() {
282 return Err(StatsError::NonFinite {
283 metric: "isotonic threshold",
284 index: None,
285 value: threshold,
286 });
287 }
288 if self.fitted[0] >= threshold {
289 return Ok(ThresholdReadout::BelowTestedRange);
290 }
291 Ok(self
292 .fitted
293 .iter()
294 .position(|value| *value >= threshold)
295 .map_or(ThresholdReadout::AboveTestedRange, |index| {
296 ThresholdReadout::Observed {
297 level: self.raw[index].level,
298 }
299 }))
300 }
301}
302
303pub fn fit_isotonic(points: &[IsotonicPoint]) -> StatsResult<IsotonicFit> {
305 if points.is_empty() {
306 return Err(StatsError::EmptyInput {
307 metric: "isotonic points",
308 });
309 }
310 let mut raw = points.to_vec();
311 for (index, point) in raw.iter().enumerate() {
312 for (metric, value) in [
313 ("isotonic level", point.level),
314 ("isotonic value", point.value),
315 ("isotonic weight", point.weight),
316 ] {
317 if !value.is_finite() {
318 return Err(StatsError::NonFinite {
319 metric,
320 index: Some(index),
321 value,
322 });
323 }
324 }
325 if point.weight <= 0.0 {
326 return Err(StatsError::InvalidControl {
327 field: "isotonic weight",
328 reason: "must be positive",
329 });
330 }
331 }
332 raw.sort_by(|a, b| a.level.total_cmp(&b.level));
333 if raw.windows(2).any(|pair| pair[0].level == pair[1].level) {
334 return Err(StatsError::InvalidControl {
335 field: "isotonic levels",
336 reason: "must be unique",
337 });
338 }
339 let mut blocks: Vec<(usize, usize, f64, f64)> = Vec::new();
340 for (index, point) in raw.iter().enumerate() {
341 blocks.push((index, index + 1, point.weight, point.weight * point.value));
342 while blocks.len() >= 2 {
343 let n = blocks.len();
344 if blocks[n - 2].3 / blocks[n - 2].2 <= blocks[n - 1].3 / blocks[n - 1].2 {
345 break;
346 }
347 let right = blocks.pop().expect("right block");
348 let left = blocks.pop().expect("left block");
349 blocks.push((left.0, right.1, left.2 + right.2, left.3 + right.3));
350 }
351 }
352 let mut fitted = vec![0.0; raw.len()];
353 for (start, end, weight, sum) in blocks {
354 fitted[start..end].fill(sum / weight);
355 }
356 let normalized_area = (raw.len() >= 2).then(|| {
357 let span = raw.last().expect("nonempty").level - raw[0].level;
358 raw.windows(2)
359 .enumerate()
360 .map(|(index, pair)| {
361 (pair[1].level - pair[0].level) * (fitted[index] + fitted[index + 1]) / 2.0
362 })
363 .sum::<f64>()
364 / span
365 });
366 Ok(IsotonicFit {
367 raw,
368 fitted,
369 normalized_area,
370 })
371}
372
373fn pair_effects(pairs: &[(f64, f64)], metric: &'static str) -> StatsResult<Vec<f64>> {
374 if pairs.is_empty() {
375 return Err(StatsError::EmptyInput { metric });
376 }
377 pairs
378 .iter()
379 .enumerate()
380 .map(|(index, (baseline, candidate))| {
381 if !baseline.is_finite() {
382 return Err(StatsError::NonFinite {
383 metric,
384 index: Some(index * 2),
385 value: *baseline,
386 });
387 }
388 if !candidate.is_finite() {
389 return Err(StatsError::NonFinite {
390 metric,
391 index: Some(index * 2 + 1),
392 value: *candidate,
393 });
394 }
395 Ok(candidate - baseline)
396 })
397 .collect()
398}
399
400fn bootstrap_effects(
401 effects: &[f64],
402 control: BootstrapControl,
403 baseline_samples: usize,
404 candidate_samples: usize,
405 cluster_count: usize,
406) -> StatsResult<BootstrapEffectInterval> {
407 let required = u64::try_from(effects.len())
408 .ok()
409 .and_then(|n| n.checked_mul(control.resamples as u64))
410 .ok_or(StatsError::WorkLimitExceeded {
411 required: u64::MAX,
412 limit: control.max_work,
413 })?;
414 if required > control.max_work {
415 return Err(StatsError::WorkLimitExceeded {
416 required,
417 limit: control.max_work,
418 });
419 }
420 let mut rng = SeededSampler::new(control.seed);
421 let mut estimates = Vec::with_capacity(control.resamples);
422 for _ in 0..control.resamples {
423 estimates.push(
424 (0..effects.len())
425 .map(|_| effects[rng.index_multiply_high(effects.len())])
426 .sum::<f64>()
427 / effects.len() as f64,
428 );
429 }
430 let tail = (1.0 - control.confidence_level) / 2.0;
431 Ok(BootstrapEffectInterval {
432 point_effect: effects.iter().sum::<f64>() / effects.len() as f64,
433 lower: exact_quantile(&estimates, tail).map_err(|_| StatsError::InvalidControl {
434 field: "bootstrap quantile",
435 reason: "internal quantile must remain valid",
436 })?,
437 upper: exact_quantile(&estimates, 1.0 - tail).map_err(|_| StatsError::InvalidControl {
438 field: "bootstrap quantile",
439 reason: "internal quantile must remain valid",
440 })?,
441 confidence_level: control.confidence_level,
442 seed: control.seed,
443 resamples: control.resamples,
444 baseline_samples,
445 candidate_samples,
446 exclusions: 0,
447 cluster_count,
448 admitted_work: required,
449 })
450}
451
452fn control_parts(control: BootstrapControl) -> StatsResult<()> {
453 if control.resamples < 2 {
454 return Err(StatsError::InvalidControl {
455 field: "resamples",
456 reason: "must be at least two",
457 });
458 }
459 confidence(control.confidence_level)
460}
461
462fn confidence(value: f64) -> StatsResult<()> {
463 if !value.is_finite() || !(0.0..1.0).contains(&value) {
464 return Err(StatsError::InvalidControl {
465 field: "confidence_level",
466 reason: "must be finite and strictly between zero and one",
467 });
468 }
469 Ok(())
470}
471
472fn binomial_cdf(k: u64, n: u64, p: f64) -> f64 {
473 (0..=k).map(|i| binomial_probability(i, n, p)).sum()
474}
475fn binomial_upper_tail(k: u64, n: u64, p: f64) -> f64 {
476 (k..=n).map(|i| binomial_probability(i, n, p)).sum()
477}
478fn binomial_probability(k: u64, n: u64, p: f64) -> f64 {
479 if p == 0.0 {
480 return f64::from(k == 0);
481 }
482 if p == 1.0 {
483 return f64::from(k == n);
484 }
485 let log_choose = (1..=k.min(n - k))
486 .map(|i| ((n + 1 - i) as f64 / i as f64).ln())
487 .sum::<f64>();
488 (log_choose + k as f64 * p.ln() + (n - k) as f64 * (-p).ln_1p()).exp()
489}
490fn bisect_probability(mut tail: impl FnMut(f64) -> f64, target: f64) -> f64 {
491 let increasing = tail(0.0) < tail(1.0);
492 let (mut low, mut high) = (0.0, 1.0);
493 for _ in 0..80 {
494 let mid = (low + high) / 2.0;
495 if (tail(mid) < target) == increasing {
496 low = mid;
497 } else {
498 high = mid;
499 }
500 }
501 (low + high) / 2.0
502}