Skip to main content

polydat_core/iteration/comprehension/
measure.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! A continuous axis's measure as the sampler maps onto it
5//! (comprehension_forms.md §10.2 R2): a sampling strategy draws a
6//! point in `[0, 1)` per continuous axis, and the axis's measure
7//! carries it onto the axis's interval, affinely for `Uniform` and
8//! by the inverse CDF for a named measure. A named measure on an
9//! interval narrower than its support is the measure restricted to
10//! that interval: the unit point is placed between the CDF values
11//! of the two ends and then inverted, so the draws keep the
12//! measure's shape inside the interval.
13
14use crate::iteration::comprehension::cardinality::{Interval, MeasureName, ProductMeasure};
15use crate::numeric::special::{
16    inv_regularized_beta, inv_regularized_gamma_p, normal_cdf, probit, regularized_beta,
17    regularized_gamma_p,
18};
19
20/// The measure of one continuous axis, with its parameters resolved
21/// (see [`MeasureName::parameter_names`]).
22#[derive(Debug, Clone, PartialEq)]
23pub enum AxisMeasure {
24    /// Lebesgue measure scaled to the interval.
25    Uniform,
26    /// A named distribution with its parameters.
27    Named {
28        /// The distribution.
29        name: MeasureName,
30        /// Its parameters, in the order of [`MeasureName::parameter_names`].
31        params: Vec<f64>,
32    },
33}
34
35impl AxisMeasure {
36    /// The measure of the `axis`-th continuous axis of `measure`: a
37    /// `Product` is indexed, `Uniform` and `Named` apply to every
38    /// axis. Named measures here carry no parameters, so they take
39    /// the standard ones; a `Source::Distribution` supplies its
40    /// own through [`AxisMeasure::named`].
41    pub fn from_product(measure: &ProductMeasure, axis: usize) -> Result<Self, String> {
42        match measure {
43            ProductMeasure::Uniform => Ok(AxisMeasure::Uniform),
44            ProductMeasure::Named(name) => Self::named(*name, &[]),
45            ProductMeasure::Product(children) => match children.get(axis) {
46                Some(child) => Self::from_product(child, 0),
47                None => Err(format!(
48                    "product measure has {} axes; axis {axis} requested",
49                    children.len()
50                )),
51            },
52        }
53    }
54
55    /// A named measure with `params` resolved against the measure's
56    /// parameter table.
57    pub fn named(name: MeasureName, params: &[f64]) -> Result<Self, String> {
58        Ok(AxisMeasure::Named {
59            name,
60            params: name.resolve_params(params)?,
61        })
62    }
63
64    /// The CDF of a named measure at `x`; `None` for `Uniform`,
65    /// whose CDF depends on the interval.
66    pub fn cdf(&self, x: f64) -> Option<f64> {
67        let AxisMeasure::Named { name, params } = self else {
68            return None;
69        };
70        if x.is_nan() {
71            return Some(f64::NAN);
72        }
73        Some(match name {
74            MeasureName::Normal => normal_cdf((x - params[0]) / params[1]),
75            MeasureName::Exponential => {
76                if x <= 0.0 {
77                    0.0
78                } else {
79                    1.0 - (-params[0] * x).exp()
80                }
81            }
82            MeasureName::Pareto => {
83                let (scale, shape) = (params[0], params[1]);
84                if x <= scale {
85                    0.0
86                } else {
87                    1.0 - (scale / x).powf(shape)
88                }
89            }
90            MeasureName::Beta => regularized_beta(x, params[0], params[1]),
91            MeasureName::LogNormal => {
92                if x <= 0.0 {
93                    0.0
94                } else {
95                    normal_cdf((x.ln() - params[0]) / params[1])
96                }
97            }
98            MeasureName::Gamma => {
99                if x <= 0.0 {
100                    0.0
101                } else {
102                    regularized_gamma_p(params[0], x / params[1])
103                }
104            }
105            MeasureName::Uniform01 => x.clamp(0.0, 1.0),
106        })
107    }
108
109    /// The quantile of a named measure at `p ∈ [0, 1]`; `None` for
110    /// `Uniform`.
111    pub fn quantile(&self, p: f64) -> Option<f64> {
112        let AxisMeasure::Named { name, params } = self else {
113            return None;
114        };
115        Some(match name {
116            MeasureName::Normal => params[0] + params[1] * probit(p),
117            MeasureName::Exponential => -(1.0 - p).ln() / params[0],
118            MeasureName::Pareto => params[0] / (1.0 - p).powf(1.0 / params[1]),
119            MeasureName::Beta => inv_regularized_beta(p, params[0], params[1]),
120            MeasureName::LogNormal => (params[0] + params[1] * probit(p)).exp(),
121            MeasureName::Gamma => params[1] * inv_regularized_gamma_p(p, params[0]),
122            MeasureName::Uniform01 => p.clamp(0.0, 1.0),
123        })
124    }
125
126    /// Carry a unit point `u ∈ [0, 1)` onto `interval` under this
127    /// measure. `Uniform` is affine; a named measure inverts its CDF
128    /// between the CDF values of the interval's ends. The result
129    /// lies in the interval, and an open end is never returned: a
130    /// point that lands on it moves to the next float inside.
131    pub fn map_unit(&self, u: f64, interval: &Interval) -> f64 {
132        let x = match self.cdf(interval.lo) {
133            None => interval.lo + u * (interval.hi - interval.lo),
134            Some(f_lo) => {
135                let f_hi = self.cdf(interval.hi).unwrap_or(1.0);
136                let p = f_lo + u * (f_hi - f_lo);
137                let q = self.quantile(p).unwrap_or(interval.lo);
138                if q.is_nan() { interval.lo } else { q }
139            }
140        };
141        inside(x.clamp(interval.lo, interval.hi), interval)
142    }
143
144    /// The point at an end of `interval` under this measure: the end
145    /// itself, or the next float inside when the end is open.
146    pub fn endpoint(&self, interval: &Interval, upper: bool) -> f64 {
147        let x = if upper { interval.hi } else { interval.lo };
148        inside(x, interval)
149    }
150}
151
152/// `x` moved off an open end of `interval` onto the next float
153/// inside; an infinite end has no next float and stays.
154fn inside(x: f64, interval: &Interval) -> f64 {
155    if interval.lo_open && x == interval.lo && x.is_finite() {
156        x.next_up().min(interval.hi)
157    } else if interval.hi_open && x == interval.hi && x.is_finite() {
158        x.next_down().max(interval.lo)
159    } else {
160        x
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    fn named(name: MeasureName, params: &[f64]) -> AxisMeasure {
169        AxisMeasure::named(name, params).unwrap()
170    }
171
172    #[test]
173    fn uniform_is_affine() {
174        let iv = Interval::closed(2.0, 4.0);
175        assert_eq!(AxisMeasure::Uniform.map_unit(0.0, &iv), 2.0);
176        assert_eq!(AxisMeasure::Uniform.map_unit(0.5, &iv), 3.0);
177        assert_eq!(AxisMeasure::Uniform.map_unit(0.25, &iv), 2.5);
178    }
179
180    #[test]
181    fn every_named_measure_inverts_its_own_cdf() {
182        let all = [
183            named(MeasureName::Normal, &[10.0, 2.0]),
184            named(MeasureName::Exponential, &[0.5]),
185            named(MeasureName::Pareto, &[2.0, 3.0]),
186            named(MeasureName::Beta, &[2.0, 5.0]),
187            named(MeasureName::LogNormal, &[0.0, 0.5]),
188            named(MeasureName::Gamma, &[2.5, 2.0]),
189            named(MeasureName::Uniform01, &[]),
190        ];
191        for m in &all {
192            for p in [0.05, 0.3, 0.5, 0.8, 0.95] {
193                let x = m.quantile(p).unwrap();
194                let back = m.cdf(x).unwrap();
195                // The probit is the least exact piece (4.5e-4).
196                assert!((back - p).abs() < 1e-3, "{m:?} p={p} x={x} back={back}");
197            }
198        }
199    }
200
201    #[test]
202    fn a_named_measure_on_its_full_support_is_its_quantile() {
203        let m = named(MeasureName::Normal, &[0.0, 1.0]);
204        let iv = Interval::open(f64::NEG_INFINITY, f64::INFINITY);
205        assert_eq!(m.map_unit(0.5, &iv), m.quantile(0.5).unwrap());
206        assert_eq!(m.map_unit(0.975, &iv), m.quantile(0.975).unwrap());
207    }
208
209    #[test]
210    fn a_named_measure_on_a_narrower_interval_is_restricted_to_it() {
211        // Exponential(1) on [0, 1]: u maps to F⁻¹(u · F(1)).
212        let m = named(MeasureName::Exponential, &[1.0]);
213        let iv = Interval::closed(0.0, 1.0);
214        let f1 = 1.0 - (-1.0f64).exp();
215        for u in [0.0, 0.25, 0.5, 0.9] {
216            let expected = -(1.0 - u * f1).ln();
217            let got = m.map_unit(u, &iv);
218            assert!((got - expected).abs() < 1e-12, "u={u} got={got}");
219            assert!((0.0..=1.0).contains(&got));
220        }
221    }
222
223    #[test]
224    fn open_ends_are_never_returned() {
225        let iv = Interval::open(0.0, 1.0);
226        let u = AxisMeasure::Uniform;
227        assert!(u.map_unit(0.0, &iv) > 0.0);
228        assert!(u.endpoint(&iv, false) > 0.0);
229        assert!(u.endpoint(&iv, true) < 1.0);
230        assert_eq!(u.endpoint(&Interval::closed(0.0, 1.0), true), 1.0);
231    }
232
233    #[test]
234    fn a_product_measure_is_indexed_per_axis() {
235        let pm = ProductMeasure::Product(vec![
236            ProductMeasure::Uniform,
237            ProductMeasure::Named(MeasureName::Exponential),
238        ]);
239        assert_eq!(
240            AxisMeasure::from_product(&pm, 0).unwrap(),
241            AxisMeasure::Uniform
242        );
243        assert_eq!(
244            AxisMeasure::from_product(&pm, 1).unwrap(),
245            named(MeasureName::Exponential, &[1.0])
246        );
247        assert!(AxisMeasure::from_product(&pm, 2).is_err());
248    }
249
250    #[test]
251    fn parameters_are_checked_against_the_table() {
252        assert!(AxisMeasure::named(MeasureName::Normal, &[1.0]).is_err());
253        assert!(AxisMeasure::named(MeasureName::Normal, &[1.0, f64::NAN]).is_err());
254        assert_eq!(
255            AxisMeasure::named(MeasureName::Gamma, &[]).unwrap(),
256            named(MeasureName::Gamma, &[1.0, 1.0])
257        );
258    }
259}