prospicio_prob/
distortion.rs1use prospicio_core::{Error, Result};
11use prospicio_math::special::{norm_cdf, norm_quantile};
12
13use crate::distribution::check_probability;
14
15#[derive(Debug, Clone, Copy, PartialEq)]
47pub enum Distortion {
48 Tvar(f64),
49 Wang(f64),
50 ProportionalHazard(f64),
51 DualPower(f64),
52 Exponential(f64),
53}
54
55impl Distortion {
56 pub fn tvar(p: f64) -> Result<Self> {
58 check_probability(p)?;
59 Ok(Self::Tvar(p))
60 }
61
62 pub fn wang(lambda: f64) -> Result<Self> {
64 if !lambda.is_finite() || lambda < 0.0 {
65 return Err(invalid("lambda", lambda, "must be finite and non-negative"));
66 }
67 Ok(Self::Wang(lambda))
68 }
69
70 pub fn proportional_hazard(rho: f64) -> Result<Self> {
72 if !(rho > 0.0 && rho <= 1.0) {
73 return Err(invalid("rho", rho, "must be in (0, 1]"));
74 }
75 Ok(Self::ProportionalHazard(rho))
76 }
77
78 pub fn dual_power(beta: f64) -> Result<Self> {
80 if !beta.is_finite() || beta < 1.0 {
81 return Err(invalid("beta", beta, "must be finite and at least 1"));
82 }
83 Ok(Self::DualPower(beta))
84 }
85
86 pub fn exponential(k: f64) -> Result<Self> {
89 if !(k.is_finite() && k > 0.0) {
90 return Err(invalid("k", k, "must be finite and positive"));
91 }
92 Ok(Self::Exponential(k))
93 }
94
95 pub fn g(&self, s: f64) -> f64 {
97 if s <= 0.0 {
98 return 0.0;
99 }
100 if s >= 1.0 {
101 return 1.0;
102 }
103 match *self {
104 Self::Tvar(p) if p >= 1.0 => 1.0,
105 Self::Tvar(p) => (s / (1.0 - p)).min(1.0),
106 Self::Wang(lambda) => norm_cdf(norm_quantile(s) + lambda),
107 Self::ProportionalHazard(rho) => s.powf(rho),
108 Self::DualPower(beta) => -((-s).ln_1p() * beta).exp_m1(),
109 Self::Exponential(k) => (-k * s).exp_m1() / (-k).exp_m1(),
110 }
111 }
112
113 pub fn weights(&self, n: usize) -> Vec<f64> {
124 let nf = n as f64;
125 (0..n)
126 .map(|i| self.g((n - i) as f64 / nf) - self.g((n - i - 1) as f64 / nf))
127 .collect()
128 }
129
130 pub fn apply_sorted(&self, sorted: &[f64]) -> f64 {
137 debug_assert!(!sorted.is_empty(), "draws must not be empty");
138 debug_assert!(
139 sorted.windows(2).all(|w| w[0] <= w[1]),
140 "draws must be sorted ascending"
141 );
142 self.weights(sorted.len())
143 .iter()
144 .zip(sorted)
145 .map(|(w, x)| w * x)
146 .sum()
147 }
148
149 pub fn apply_discrete(&self, values: &[f64], probs: &[f64]) -> f64 {
163 debug_assert_eq!(values.len(), probs.len());
164 debug_assert!(
165 values.windows(2).all(|w| w[0] <= w[1]),
166 "values must be sorted ascending"
167 );
168 let mut above = 0.0; let mut g_above = 0.0;
170 let mut total = 0.0;
171 for (x, p) in values.iter().zip(probs).rev() {
172 let at_or_above = above + p;
173 let g_at_or_above = self.g(at_or_above);
174 total += (g_at_or_above - g_above) * x;
175 above = at_or_above;
176 g_above = g_at_or_above;
177 }
178 total
179 }
180}
181
182fn invalid(name: &'static str, value: f64, reason: &'static str) -> Error {
183 Error::InvalidParameter {
184 name,
185 value,
186 reason,
187 }
188}
189
190#[cfg(test)]
191mod tests {
192 use super::*;
193
194 #[test]
196 fn exponential_spectral_on_a_uniform() {
197 let n = 100_000;
198 let u: Vec<f64> = (0..n).map(|i| (i as f64 + 0.5) / n as f64).collect();
199 for k in [0.5f64, 3.0, 20.0] {
200 let want = (k.exp() * (k - 1.0) + 1.0) / (k * k.exp_m1());
201 let got = Distortion::exponential(k).unwrap().apply_sorted(&u);
202 assert!((got - want).abs() < 1e-6, "{k}: {got} vs {want}");
203 }
204 assert!(Distortion::exponential(0.0).is_err());
205 }
206 use crate::risk::tvar_sorted;
207
208 const X: [f64; 5] = [10.0, 20.0, 30.0, 40.0, 50.0];
209
210 fn all() -> Vec<Distortion> {
211 vec![
212 Distortion::tvar(0.7).unwrap(),
213 Distortion::wang(0.5).unwrap(),
214 Distortion::proportional_hazard(0.6).unwrap(),
215 Distortion::dual_power(2.5).unwrap(),
216 ]
217 }
218
219 #[test]
220 fn tvar_matches_the_risk_module() {
221 for i in 0..=100 {
222 let p = i as f64 / 100.0;
223 let d = Distortion::tvar(p).unwrap().apply_sorted(&X);
224 let t = tvar_sorted(&X, p).unwrap();
225 assert!((d - t).abs() <= 1e-12 * t, "p = {p}: {d} vs {t}");
226 }
227 }
228
229 #[test]
230 fn identity_parameters_give_the_mean() {
231 for d in [
232 Distortion::tvar(0.0).unwrap(),
233 Distortion::wang(0.0).unwrap(),
234 Distortion::proportional_hazard(1.0).unwrap(),
235 Distortion::dual_power(1.0).unwrap(),
236 ] {
237 assert!((d.apply_sorted(&X) - 30.0).abs() < 1e-12, "{d:?}");
238 }
239 }
240
241 #[test]
242 fn weights_are_a_probability_vector_rising_to_the_tail() {
243 for d in all() {
244 let w = d.weights(1000);
245 assert!((w.iter().sum::<f64>() - 1.0).abs() < 1e-12, "{d:?}");
246 assert!(w.iter().all(|&w| w >= 0.0));
247 assert!(w.windows(2).all(|p| p[1] >= p[0] - 1e-15), "{d:?}");
249 }
250 }
251
252 #[test]
253 fn coherence_properties() {
254 let mean = 30.0;
255 for d in all() {
256 let r = d.apply_sorted(&X);
257 assert!(r >= mean && r <= 50.0, "{d:?}: {r}");
258 let shifted: Vec<f64> = X.iter().map(|x| 2.0 * x + 7.0).collect();
260 assert!((d.apply_sorted(&shifted) - (2.0 * r + 7.0)).abs() < 1e-12);
261 }
262 let w = |l| Distortion::wang(l).unwrap().apply_sorted(&X);
264 assert!(w(0.2) < w(0.5) && w(0.5) < w(1.0));
265 let ph = |r| Distortion::proportional_hazard(r).unwrap().apply_sorted(&X);
266 assert!(ph(0.9) < ph(0.5));
267 let dp = |b| Distortion::dual_power(b).unwrap().apply_sorted(&X);
268 assert!(dp(1.5) < dp(3.0));
269 }
270
271 #[test]
272 fn discrete_agrees_with_equal_weights_and_merges_ties() {
273 let p = [0.2; 5];
274 for d in all() {
275 let a = d.apply_discrete(&X, &p);
276 let b = d.apply_sorted(&X);
277 assert!((a - b).abs() < 1e-12, "{d:?}");
278 }
279 let d = Distortion::dual_power(2.0).unwrap();
281 let ties = d.apply_sorted(&[1.0, 5.0, 5.0, 5.0]);
282 let atom = d.apply_discrete(&[1.0, 5.0], &[0.25, 0.75]);
283 assert!((ties - atom).abs() < 1e-15);
284 assert!((atom - (0.9375 * 5.0 + 0.0625)).abs() < 1e-15);
286 }
287
288 #[test]
289 fn g_closed_forms() {
290 let s = 0.3;
291 assert_eq!(Distortion::tvar(0.8).unwrap().g(s), 1.0);
292 assert!((Distortion::tvar(0.4).unwrap().g(s) - 0.5).abs() < 1e-15);
293 assert!((Distortion::proportional_hazard(0.5).unwrap().g(s) - s.sqrt()).abs() < 1e-15);
294 assert!((Distortion::dual_power(2.0).unwrap().g(s) - 0.51).abs() < 1e-15);
295 assert!((Distortion::wang(1.0).unwrap().g(0.5) - norm_cdf(1.0)).abs() < 1e-15);
297 assert_eq!(Distortion::tvar(1.0).unwrap().g(1e-300), 1.0);
298 for d in all() {
299 assert_eq!(d.g(0.0), 0.0);
300 assert_eq!(d.g(1.0), 1.0);
301 }
302 }
303
304 #[test]
305 fn rejects_bad_parameters() {
306 assert!(Distortion::tvar(1.5).is_err());
307 assert!(Distortion::wang(-0.1).is_err());
308 assert!(Distortion::wang(f64::INFINITY).is_err());
309 assert!(Distortion::proportional_hazard(0.0).is_err());
310 assert!(Distortion::proportional_hazard(1.5).is_err());
311 assert!(Distortion::dual_power(0.5).is_err());
312 }
313}