Skip to main content

sim_lib_numbers_stats/
decision.rs

1//! Bounded, deterministic inference primitives for sequential study decisions.
2
3use super::{BootstrapControl, BootstrapEffectInterval, StatsError, StatsResult, exact_quantile};
4use crate::SeededSampler;
5
6/// A two-sided Clopper--Pearson interval for a finite binary count.
7#[derive(Clone, Copy, Debug, PartialEq)]
8pub struct BinaryInterval {
9    /// Observed successes.
10    pub successes: u64,
11    /// Observed trials.
12    pub trials: u64,
13    /// Declared central confidence mass.
14    pub confidence_level: f64,
15    /// Exact lower endpoint.
16    pub lower: f64,
17    /// Exact upper endpoint.
18    pub upper: f64,
19}
20
21/// Computes an exact equal-tailed finite-count binary interval.
22///
23/// The endpoints invert binomial tail probabilities and use no normal or
24/// other large-sample approximation.
25pub 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/// One independent cluster, retaining its stable identity and paired rows.
63#[derive(Clone, Debug, PartialEq)]
64pub struct ClusterSample {
65    /// Stable cluster identity; identities must be unique.
66    pub id: u64,
67    /// `(baseline, candidate)` rows belonging to this cluster.
68    pub pairs: Vec<(f64, f64)>,
69}
70
71/// Deterministically bootstraps paired candidate-minus-baseline effects.
72pub 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
86/// Deterministically resamples whole independent clusters with replacement.
87///
88/// `minimum_clusters` is the caller-declared independence floor. Cluster rows
89/// remain together; clusters are sorted by stable id before seeded sampling,
90/// making within-cluster row order irrelevant.
91pub 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/// A pre-registered sequential look and its allocated false-elimination mass.
134#[derive(Clone, Copy, Debug, PartialEq)]
135pub struct RegisteredLook {
136    /// Cumulative sample count at which this look is legal.
137    pub samples: usize,
138    /// Positive alpha allocated to this look.
139    pub alpha: f64,
140}
141
142/// An alpha-spent contract for bounded observations in `[0, 1]`.
143#[derive(Clone, Debug, PartialEq)]
144pub struct RegisteredLookSequence {
145    looks: Vec<RegisteredLook>,
146    total_budget: f64,
147}
148
149/// A Hoeffding interval valid at its pre-registered look under the sealed budget.
150#[derive(Clone, Copy, Debug, PartialEq)]
151pub struct SequentialInterval {
152    /// Registered cumulative sample count.
153    pub samples: usize,
154    /// Arithmetic mean of admitted observations.
155    pub mean: f64,
156    /// Lower endpoint clipped to zero.
157    pub lower: f64,
158    /// Upper endpoint clipped to one.
159    pub upper: f64,
160    /// Alpha spent at this look.
161    pub alpha_spent: f64,
162    /// Caller-declared total false-elimination budget.
163    pub total_budget: f64,
164}
165
166impl RegisteredLookSequence {
167    /// Seals an increasing, unique set of looks whose alpha does not exceed the budget.
168    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    /// Evaluates exactly one registered look; optional peeking is therefore excluded by construction.
211    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/// One weighted raw point supplied to isotonic regression.
243#[derive(Clone, Copy, Debug, PartialEq)]
244pub struct IsotonicPoint {
245    /// Tested level, strictly increasing after canonical sorting.
246    pub level: f64,
247    /// Raw response.
248    pub value: f64,
249    /// Positive observation weight.
250    pub weight: f64,
251}
252
253/// Threshold crossing evidence, including censoring beyond the tested range.
254#[derive(Clone, Copy, Debug, PartialEq)]
255pub enum ThresholdReadout {
256    /// The first tested level whose fit reaches the threshold.
257    Observed {
258        /// First tested level reaching the threshold.
259        level: f64,
260    },
261    /// Every fitted point is already at or above the threshold.
262    BelowTestedRange,
263    /// No fitted point reaches the threshold.
264    AboveTestedRange,
265}
266
267/// Inspectable weighted pool-adjacent-violators fit.
268#[derive(Clone, Debug, PartialEq)]
269pub struct IsotonicFit {
270    /// Canonically sorted raw points.
271    pub raw: Vec<IsotonicPoint>,
272    /// Nondecreasing fitted values aligned with `raw`.
273    pub fitted: Vec<f64>,
274    /// Trapezoidal area divided by the tested level span, available for two or more levels.
275    pub normalized_area: Option<f64>,
276}
277
278impl IsotonicFit {
279    /// Reads a threshold crossing and preserves left/right censoring.
280    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
303/// Fits a weighted nondecreasing curve with pool-adjacent-violators.
304pub 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}