Skip to main content

prospicio_prob/
custom.rs

1//! [`Custom`], a user-defined loss severity given by its cdf: the "slow
2//! path" of `docs/design/distributions.md`, through which a Python or R
3//! function enters the native calculations.
4//!
5//! The cdf is the only function a user must supply; a quantile function
6//! is optional and makes sampling much faster. Everything else (the mean,
7//! the variance, limited expected values, layer moments) is computed by
8//! Gauss–Legendre quadrature of the survival function on panels between
9//! the distribution's own quantiles, so it is as accurate as the cdf is
10//! smooth. The integrals stop at the `1 - 1e-12` quantile: the probability
11//! above it is ignored, which is negligible unless the tail is so heavy
12//! that the mean barely exists.
13//!
14//! A callback may come from a language that must not be entered from
15//! several threads at once (R must only ever be called on its main thread),
16//! so a `Custom` built with `parallel_safe = false` reports so through
17//! [`Distribution::is_parallel_safe`], and the parallel simulations run
18//! single-threaded on the calling thread when they meet one.
19
20use std::fmt;
21use std::sync::{Arc, Mutex};
22
23use prospicio_core::{Error, Result};
24use prospicio_math::integrate::gauss_legendre;
25use prospicio_math::roots::bisect_log;
26
27use crate::distribution::{Distribution, check_probability};
28use crate::severity::Severity;
29
30/// A user function of one variable. An `Err` carries the message of the
31/// failure (a Python exception, an R error).
32pub type Callback = Arc<dyn Fn(f64) -> std::result::Result<f64, String> + Send + Sync>;
33
34/// Probabilities whose quantiles bound the integration panels.
35const PANEL_PROBS: [f64; 16] = [
36    0.01, 0.1, 0.25, 0.5, 0.75, 0.9, 0.95, 0.99, 0.999, 1e-4, 1e-5, 1e-6, 1e-7, 1e-8, 1e-10, 1e-12,
37];
38
39/// Gauss–Legendre pieces per panel.
40const PIECES: usize = 8;
41
42/// A non-negative loss severity defined by a user's cdf (and optionally
43/// quantile) function.
44///
45/// ```
46/// use std::sync::Arc;
47/// use prospicio_prob::{Custom, Distribution, Severity};
48///
49/// // An exponential with mean 100, given only by its cdf.
50/// let cdf = Arc::new(|x: f64| Ok((1.0 - (-x / 100.0).exp()).max(0.0)));
51/// let d = Custom::new("exponential", cdf, None, true).unwrap();
52/// assert!((d.mean() - 100.0).abs() < 1e-6);
53/// // E[min(X, 50)] = 100 (1 - e^-0.5).
54/// assert!((d.lev(50.0) - 100.0 * (1.0 - (-0.5f64).exp())).abs() < 1e-9);
55/// assert!((d.quantile(0.5).unwrap() - 100.0 * 2f64.ln()).abs() < 1e-9);
56/// ```
57#[derive(Clone)]
58pub struct Custom {
59    name: String,
60    cdf: Callback,
61    quantile: Option<Callback>,
62    parallel_safe: bool,
63    /// Panel boundaries: 0, then quantiles up to the `1 - 1e-12` one.
64    breaks: Vec<f64>,
65    mean: f64,
66    second_moment: f64,
67    /// The first callback failure after construction.
68    error: Arc<Mutex<Option<String>>>,
69}
70
71impl fmt::Debug for Custom {
72    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73        f.debug_struct("Custom")
74            .field("name", &self.name)
75            .field("has_quantile", &self.quantile.is_some())
76            .field("parallel_safe", &self.parallel_safe)
77            .field("mean", &self.mean)
78            .field("upper", &self.upper())
79            .finish()
80    }
81}
82
83fn callback_error(name: &str, what: &str, x: f64, msg: &str) -> Error {
84    Error::Data(format!(
85        "custom distribution {name:?}: {what}({x}) failed: {msg}"
86    ))
87}
88
89impl Custom {
90    /// A severity from its `cdf` and, optionally, its `quantile` function.
91    ///
92    /// `parallel_safe` says whether the callbacks may run on several
93    /// threads at once (true for a pure Rust closure; false for an R
94    /// function). Construction evaluates the cdf at the panel quantiles and
95    /// computes the mean and second moment, so a cdf that fails, leaves
96    /// `[0, 1]`, decreases or never reaches `1 - 1e-12` is reported here.
97    pub fn new(
98        name: impl Into<String>,
99        cdf: Callback,
100        quantile: Option<Callback>,
101        parallel_safe: bool,
102    ) -> Result<Self> {
103        let mut d = Self {
104            name: name.into(),
105            cdf,
106            quantile,
107            parallel_safe,
108            breaks: vec![0.0],
109            mean: f64::NAN,
110            second_moment: f64::NAN,
111            error: Arc::new(Mutex::new(None)),
112        };
113        d.try_cdf(0.0)?;
114        let mut breaks = vec![0.0];
115        for p in PANEL_PROBS {
116            let p = if p < 0.01 { 1.0 - p } else { p };
117            let q = d.try_quantile(p)?;
118            if !q.is_finite() {
119                return Err(Error::Data(format!(
120                    "custom distribution {:?}: its cdf does not reach {p} at any finite loss",
121                    d.name
122                )));
123            }
124            if q > *breaks.last().unwrap() {
125                breaks.push(q);
126            } else if q < *breaks.last().unwrap() {
127                return Err(Error::Data(format!(
128                    "custom distribution {:?}: quantiles must increase with p, but q({p}) = {q} \
129                     is below an earlier quantile",
130                    d.name
131                )));
132            }
133        }
134        d.breaks = breaks;
135        // Both moments in one pass: E[X] = ∫ S, E[X²] = ∫ 2x S.
136        let upper = d.upper();
137        d.mean = d.integrate(0.0, upper, |_, s| s)?;
138        d.second_moment = d.integrate(0.0, upper, |x, s| 2.0 * x * s)?;
139        Ok(d)
140    }
141
142    /// The name given at construction.
143    pub fn name(&self) -> &str {
144        &self.name
145    }
146
147    /// Whether a quantile function was supplied (otherwise quantiles invert
148    /// the cdf by bisection, about a hundred cdf calls each).
149    pub fn has_quantile(&self) -> bool {
150        self.quantile.is_some()
151    }
152
153    /// The `1 - 1e-12` quantile, where the integrals stop.
154    pub fn upper(&self) -> f64 {
155        *self.breaks.last().unwrap()
156    }
157
158    /// The first callback failure since construction, if any. A failure
159    /// inside a calculation makes that value NaN; this says why.
160    pub fn error(&self) -> Option<String> {
161        self.error.lock().map(|e| e.clone()).unwrap_or(None)
162    }
163
164    fn record(&self, e: &Error) {
165        if let Ok(mut slot) = self.error.lock()
166            && slot.is_none()
167        {
168            *slot = Some(e.to_string());
169        }
170    }
171
172    fn try_cdf(&self, x: f64) -> Result<f64> {
173        let v = (self.cdf)(x).map_err(|m| callback_error(&self.name, "cdf", x, &m))?;
174        if !(0.0..=1.0).contains(&v) {
175            return Err(callback_error(
176                &self.name,
177                "cdf",
178                x,
179                &format!("returned {v}, outside [0, 1]"),
180            ));
181        }
182        Ok(v)
183    }
184
185    fn try_quantile(&self, p: f64) -> Result<f64> {
186        check_probability(p)?;
187        if let Some(q) = &self.quantile {
188            let v = q(p).map_err(|m| callback_error(&self.name, "quantile", p, &m))?;
189            if v.is_nan() || v < 0.0 {
190                return Err(callback_error(
191                    &self.name,
192                    "quantile",
193                    p,
194                    &format!("returned {v}; losses are non-negative"),
195                ));
196            }
197            return Ok(v);
198        }
199        if self.try_cdf(0.0)? >= p {
200            return Ok(0.0);
201        }
202        // Bracket by doubling, then bisect on a log scale.
203        let mut hi = 1.0;
204        while self.try_cdf(hi)? < p {
205            hi *= 2.0;
206            if !hi.is_finite() {
207                return Ok(f64::INFINITY);
208            }
209        }
210        let mut lo = hi;
211        loop {
212            lo *= 0.5;
213            if lo == 0.0 {
214                return Ok(0.0);
215            }
216            if self.try_cdf(lo)? < p {
217                break;
218            }
219        }
220        let mut failure = None;
221        let x = bisect_log(lo, hi, |x| match self.try_cdf(x) {
222            Ok(c) => c < p,
223            Err(e) => {
224                failure.get_or_insert(e);
225                false
226            }
227        });
228        match failure {
229            Some(e) => Err(e),
230            None => Ok(x),
231        }
232    }
233
234    /// `∫_lo^hi g(x, S(x)) dx` over the panels, clipped to `[0, upper]`.
235    fn integrate(&self, lo: f64, hi: f64, g: impl Fn(f64, f64) -> f64) -> Result<f64> {
236        let (lo, hi) = (lo.max(0.0), hi.min(self.upper()));
237        let mut total = 0.0;
238        for w in self.breaks.windows(2) {
239            let (a, b) = (w[0].max(lo), w[1].min(hi));
240            if a >= b {
241                continue;
242            }
243            let step = (b - a) / PIECES as f64;
244            for k in 0..PIECES {
245                let x0 = a + step * k as f64;
246                let x1 = if k + 1 == PIECES { b } else { x0 + step };
247                total += gauss_legendre(|x| self.try_cdf(x).map(|c| g(x, 1.0 - c)), x0, x1)?;
248            }
249        }
250        Ok(total)
251    }
252
253    fn or_nan(&self, r: Result<f64>) -> f64 {
254        r.unwrap_or_else(|e| {
255            self.record(&e);
256            f64::NAN
257        })
258    }
259}
260
261impl Distribution for Custom {
262    fn mean(&self) -> f64 {
263        self.mean
264    }
265
266    fn variance(&self) -> f64 {
267        self.second_moment - self.mean * self.mean
268    }
269
270    fn cdf(&self, x: f64) -> f64 {
271        if x < 0.0 {
272            return 0.0;
273        }
274        self.or_nan(self.try_cdf(x))
275    }
276
277    fn quantile(&self, p: f64) -> Result<f64> {
278        check_probability(p)?;
279        self.try_quantile(p).inspect_err(|e| self.record(e))
280    }
281
282    fn is_parallel_safe(&self) -> bool {
283        self.parallel_safe
284    }
285}
286
287impl Severity for Custom {
288    fn lev(&self, limit: f64) -> f64 {
289        if limit <= 0.0 {
290            return limit;
291        }
292        if limit == f64::INFINITY {
293            return self.mean;
294        }
295        self.or_nan(self.integrate(0.0, limit, |_, s| s))
296    }
297
298    /// `∫_r^∞ S`, integrated directly so far retentions keep their
299    /// precision.
300    fn stop_loss(&self, retention: f64) -> f64 {
301        if retention <= 0.0 {
302            return self.mean - retention;
303        }
304        self.or_nan(self.integrate(retention, f64::INFINITY, |_, s| s))
305    }
306
307    /// `∫_a^{a+l} 2 (x - a) S(x) dx`.
308    fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
309        let a = attachment.max(0.0);
310        self.or_nan(self.integrate(a, a + limit, |x, s| 2.0 * (x - a) * s))
311    }
312}
313
314#[cfg(test)]
315mod tests {
316    use super::*;
317    use crate::Lognormal;
318
319    fn from<D: Distribution + Send + Sync + 'static>(d: D, with_quantile: bool) -> Custom {
320        let d = Arc::new(d);
321        let c = d.clone();
322        let cdf: Callback = Arc::new(move |x| Ok(c.cdf(x)));
323        let quantile: Option<Callback> = with_quantile.then(|| {
324            let q = d.clone();
325            Arc::new(move |p| q.quantile(p).map_err(|e| e.to_string())) as Callback
326        });
327        Custom::new("test", cdf, quantile, true).unwrap()
328    }
329
330    #[test]
331    fn moments_and_layers_match_a_closed_form() {
332        let ln = Lognormal::from_mean_cv(1000.0, 1.0).unwrap();
333        for with_quantile in [true, false] {
334            let c = from(ln, with_quantile);
335            assert!((c.mean() / ln.mean() - 1.0).abs() < 1e-7, "{with_quantile}");
336            assert!((c.variance() / ln.variance() - 1.0).abs() < 1e-4);
337            for l in [100.0, 1000.0, 5000.0] {
338                assert!((c.lev(l) / ln.lev(l) - 1.0).abs() < 1e-9, "lev {l}");
339                assert!(
340                    (c.stop_loss(l) / ln.stop_loss(l) - 1.0).abs() < 1e-6,
341                    "sl {l}"
342                );
343            }
344            let (m, v) = (c.layer(2000.0, 1000.0), c.layer_variance(2000.0, 1000.0));
345            assert!((m / ln.layer(2000.0, 1000.0) - 1.0).abs() < 1e-8);
346            assert!((v / ln.layer_variance(2000.0, 1000.0) - 1.0).abs() < 1e-8);
347            for p in [0.001, 0.5, 0.99] {
348                let q = c.quantile(p).unwrap();
349                assert!((q / ln.quantile(p).unwrap() - 1.0).abs() < 1e-12, "q {p}");
350            }
351        }
352    }
353
354    #[test]
355    fn a_point_mass_at_zero_is_allowed() {
356        // 30% no-claim, else exponential(mean 10).
357        let cdf: Callback = Arc::new(|x| {
358            Ok(if x < 0.0 {
359                0.0
360            } else {
361                0.3 + 0.7 * (1.0 - (-x / 10.0).exp())
362            })
363        });
364        let c = Custom::new("zero-inflated", cdf, None, true).unwrap();
365        assert!((c.mean() - 7.0).abs() < 1e-7);
366        assert_eq!(c.quantile(0.2).unwrap(), 0.0);
367        assert_eq!(c.cdf(-1.0), 0.0);
368    }
369
370    #[test]
371    fn bad_callbacks_are_reported() {
372        let failing: Callback = Arc::new(|x| {
373            if x > 50.0 {
374                Err("boom".into())
375            } else {
376                Ok(x / 100.0)
377            }
378        });
379        let e = Custom::new("bad", failing, None, true)
380            .unwrap_err()
381            .to_string();
382        assert!(e.contains("boom"), "{e}");
383
384        let outside: Callback = Arc::new(|_| Ok(1.5));
385        assert!(Custom::new("bad", outside, None, true).is_err());
386
387        let never: Callback = Arc::new(|x| Ok(0.5 * (1.0 - (-x).exp())));
388        let e = Custom::new("bad", never, None, true)
389            .unwrap_err()
390            .to_string();
391        assert!(e.contains("does not reach"), "{e}");
392
393        // A failure after construction gives NaN and is kept.
394        let late: Callback = Arc::new(|x| {
395            if x == 12345.0 {
396                Err("late".into())
397            } else {
398                Ok(1.0 - (-x).exp())
399            }
400        });
401        let c = Custom::new("late", late, None, false).unwrap();
402        assert!(c.error().is_none());
403        assert!(c.cdf(12345.0).is_nan());
404        assert!(c.error().unwrap().contains("late"));
405        assert!(!c.is_parallel_safe());
406    }
407}