polydat_core/iteration/comprehension/
measure.rs1use 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#[derive(Debug, Clone, PartialEq)]
23pub enum AxisMeasure {
24 Uniform,
26 Named {
28 name: MeasureName,
30 params: Vec<f64>,
32 },
33}
34
35impl AxisMeasure {
36 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 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 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 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 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 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
152fn 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 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 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}