Skip to main content

prospicio_prob/
dist.rs

1//! [`Dist`], the closed enum of every native distribution, for the places
2//! where the family is chosen at run time: the Python and R bindings,
3//! serialization and model outputs (`docs/design/distributions.md`,
4//! "Static vs dynamic dispatch").
5//!
6//! Hot loops stay generic over `D: Distribution` and monomorphize; `Dist`
7//! dispatches once per call with a `match`, not through a vtable.
8
9use std::sync::Arc;
10
11use prospicio_core::{Result, StreamRng};
12
13use crate::distribution::Distribution;
14use crate::evt::Gpd;
15use crate::severity::Severity;
16use crate::{
17    Custom, Gamma, Grid, LogAffinePareto, Loglogistic, Lognormal, Mixture, Pareto, PiecewisePareto,
18    Sampled, Tweedie, Weibull,
19};
20
21/// Any native univariate distribution.
22///
23/// Every variant is a [`Distribution`]; all but [`Dist::Sampled`] are also
24/// a [`Severity`], which [`Dist::as_severity`] exposes. A `Mixture` is held
25/// in an [`Arc`] (its components are trait objects), so cloning a `Dist`
26/// never copies one.
27///
28/// ```
29/// use prospicio_prob::{Dist, Distribution, Lognormal, Sampled};
30///
31/// let dists = vec![
32///     Dist::from(Lognormal::from_mean_cv(1000.0, 0.5).unwrap()),
33///     Dist::from(Sampled::new(vec![800.0, 1000.0, 1200.0]).unwrap()),
34/// ];
35/// for d in &dists {
36///     assert!((d.mean() - 1000.0).abs() < 1e-9);
37/// }
38/// // Limited expected values exist for the parametric family only.
39/// assert!(dists[0].as_severity().is_some());
40/// assert!(dists[1].as_severity().is_none());
41/// ```
42#[derive(Debug, Clone)]
43pub enum Dist {
44    Lognormal(Lognormal),
45    Pareto(Pareto),
46    PiecewisePareto(PiecewisePareto),
47    LogAffinePareto(LogAffinePareto),
48    GeneralizedPareto(Gpd),
49    Gamma(Gamma),
50    Tweedie(Tweedie),
51    Weibull(Weibull),
52    Loglogistic(Loglogistic),
53    Mixture(Arc<Mixture>),
54    Grid(Grid),
55    Sampled(Sampled),
56    /// A user-defined severity (a Python or R callback): the slow path.
57    Custom(Custom),
58}
59
60/// Calls `$call` on the inner distribution, whichever it is.
61macro_rules! each {
62    ($self:ident, $d:ident => $call:expr) => {
63        match $self {
64            Dist::Lognormal($d) => $call,
65            Dist::Pareto($d) => $call,
66            Dist::PiecewisePareto($d) => $call,
67            Dist::LogAffinePareto($d) => $call,
68            Dist::GeneralizedPareto($d) => $call,
69            Dist::Gamma($d) => $call,
70            Dist::Tweedie($d) => $call,
71            Dist::Weibull($d) => $call,
72            Dist::Loglogistic($d) => $call,
73            Dist::Mixture($d) => $call,
74            Dist::Grid($d) => $call,
75            Dist::Sampled($d) => $call,
76            Dist::Custom($d) => $call,
77        }
78    };
79}
80
81impl Dist {
82    /// Short name of the family: `"lognormal"`, `"pareto"`,
83    /// `"piecewise_pareto"`, `"log_affine_pareto"`, `"generalized_pareto"`,
84    /// `"gamma"`, `"tweedie"`, `"weibull"`, `"loglogistic"`, `"mixture"`,
85    /// `"grid"`, `"sampled"` or `"custom"`.
86    pub fn family(&self) -> &'static str {
87        match self {
88            Self::Lognormal(_) => "lognormal",
89            Self::Pareto(_) => "pareto",
90            Self::PiecewisePareto(_) => "piecewise_pareto",
91            Self::LogAffinePareto(_) => "log_affine_pareto",
92            Self::GeneralizedPareto(_) => "generalized_pareto",
93            Self::Gamma(_) => "gamma",
94            Self::Tweedie(_) => "tweedie",
95            Self::Weibull(_) => "weibull",
96            Self::Loglogistic(_) => "loglogistic",
97            Self::Mixture(_) => "mixture",
98            Self::Grid(_) => "grid",
99            Self::Sampled(_) => "sampled",
100            Self::Custom(_) => "custom",
101        }
102    }
103
104    /// The distribution as a [`Severity`] (limited expected values, layers),
105    /// or `None` for [`Dist::Sampled`]: draws have no exact layer moments
106    /// (`docs/design/distributions.md`).
107    pub fn as_severity(&self) -> Option<&(dyn Severity + Send + Sync)> {
108        Some(match self {
109            Self::Lognormal(d) => d,
110            Self::Pareto(d) => d,
111            Self::PiecewisePareto(d) => d,
112            Self::LogAffinePareto(d) => d,
113            Self::GeneralizedPareto(d) => d,
114            Self::Gamma(d) => d,
115            Self::Tweedie(d) => d,
116            Self::Weibull(d) => d,
117            Self::Loglogistic(d) => d,
118            Self::Mixture(d) => d.as_ref(),
119            Self::Grid(d) => d,
120            Self::Custom(d) => d,
121            Self::Sampled(_) => return None,
122        })
123    }
124}
125
126impl Distribution for Dist {
127    fn mean(&self) -> f64 {
128        each!(self, d => d.mean())
129    }
130
131    fn variance(&self) -> f64 {
132        each!(self, d => d.variance())
133    }
134
135    fn cdf(&self, x: f64) -> f64 {
136        each!(self, d => d.cdf(x))
137    }
138
139    fn survival(&self, x: f64) -> f64 {
140        each!(self, d => d.survival(x))
141    }
142
143    fn quantile(&self, p: f64) -> Result<f64> {
144        each!(self, d => d.quantile(p))
145    }
146
147    fn sample(&self, rng: &mut StreamRng, n: usize) -> Vec<f64> {
148        each!(self, d => d.sample(rng, n))
149    }
150
151    fn is_parallel_safe(&self) -> bool {
152        each!(self, d => d.is_parallel_safe())
153    }
154}
155
156macro_rules! from_family {
157    ($($variant:ident($ty:ty)),* $(,)?) => {
158        $(
159            impl From<$ty> for Dist {
160                fn from(d: $ty) -> Self {
161                    Self::$variant(d)
162                }
163            }
164        )*
165    };
166}
167
168from_family!(
169    Lognormal(Lognormal),
170    Pareto(Pareto),
171    PiecewisePareto(PiecewisePareto),
172    LogAffinePareto(LogAffinePareto),
173    GeneralizedPareto(Gpd),
174    Gamma(Gamma),
175    Tweedie(Tweedie),
176    Weibull(Weibull),
177    Loglogistic(Loglogistic),
178    Grid(Grid),
179    Sampled(Sampled),
180    Custom(Custom),
181);
182
183/// A [`Dist`] known to be a [`Severity`]: every variant but
184/// [`Dist::Sampled`].
185///
186/// This is what the bindings accept wherever a loss severity is required
187/// (collective models, layers, copula marginals, mixture components). It
188/// dispatches by `match`, like [`Dist`], with no vtable.
189///
190/// ```
191/// use prospicio_prob::{Dist, Lognormal, Sampled, Severity, SeverityDist};
192///
193/// let ln = Lognormal::from_mean_cv(1000.0, 0.5).unwrap();
194/// let s = SeverityDist::try_from(Dist::from(ln)).unwrap();
195/// assert_eq!(s.lev(800.0), ln.lev(800.0));
196///
197/// let draws = Dist::from(Sampled::new(vec![1.0, 2.0]).unwrap());
198/// assert!(SeverityDist::try_from(draws).is_err());
199/// ```
200#[derive(Debug, Clone)]
201pub struct SeverityDist(Dist);
202
203impl SeverityDist {
204    /// The distribution.
205    pub fn dist(&self) -> &Dist {
206        &self.0
207    }
208
209    /// The distribution, by value.
210    pub fn into_dist(self) -> Dist {
211        self.0
212    }
213}
214
215impl TryFrom<Dist> for SeverityDist {
216    /// The distribution back, when it is [`Dist::Sampled`].
217    type Error = Dist;
218
219    fn try_from(d: Dist) -> std::result::Result<Self, Dist> {
220        match d {
221            Dist::Sampled(_) => Err(d),
222            d => Ok(Self(d)),
223        }
224    }
225}
226
227impl From<SeverityDist> for Dist {
228    fn from(s: SeverityDist) -> Self {
229        s.0
230    }
231}
232
233/// Calls `$call` on the inner severity, whichever it is.
234macro_rules! each_severity {
235    ($self:ident, $d:ident => $call:expr) => {
236        match &$self.0 {
237            Dist::Lognormal($d) => $call,
238            Dist::Pareto($d) => $call,
239            Dist::PiecewisePareto($d) => $call,
240            Dist::LogAffinePareto($d) => $call,
241            Dist::GeneralizedPareto($d) => $call,
242            Dist::Gamma($d) => $call,
243            Dist::Tweedie($d) => $call,
244            Dist::Weibull($d) => $call,
245            Dist::Loglogistic($d) => $call,
246            Dist::Mixture($d) => $call,
247            Dist::Grid($d) => $call,
248            Dist::Custom($d) => $call,
249            Dist::Sampled(_) => unreachable!("SeverityDist never holds Sampled"),
250        }
251    };
252}
253
254impl Distribution for SeverityDist {
255    fn mean(&self) -> f64 {
256        self.0.mean()
257    }
258
259    fn variance(&self) -> f64 {
260        self.0.variance()
261    }
262
263    fn cdf(&self, x: f64) -> f64 {
264        self.0.cdf(x)
265    }
266
267    fn survival(&self, x: f64) -> f64 {
268        self.0.survival(x)
269    }
270
271    fn quantile(&self, p: f64) -> Result<f64> {
272        self.0.quantile(p)
273    }
274
275    fn sample(&self, rng: &mut StreamRng, n: usize) -> Vec<f64> {
276        self.0.sample(rng, n)
277    }
278
279    fn is_parallel_safe(&self) -> bool {
280        self.0.is_parallel_safe()
281    }
282}
283
284impl Severity for SeverityDist {
285    fn lev(&self, limit: f64) -> f64 {
286        each_severity!(self, d => d.lev(limit))
287    }
288
289    fn stop_loss(&self, retention: f64) -> f64 {
290        each_severity!(self, d => d.stop_loss(retention))
291    }
292
293    fn layer(&self, limit: f64, attachment: f64) -> f64 {
294        each_severity!(self, d => d.layer(limit, attachment))
295    }
296
297    fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
298        each_severity!(self, d => d.layer_second_moment(limit, attachment))
299    }
300
301    fn layer_variance(&self, limit: f64, attachment: f64) -> f64 {
302        each_severity!(self, d => d.layer_variance(limit, attachment))
303    }
304}
305
306impl From<Mixture> for Dist {
307    fn from(d: Mixture) -> Self {
308        Self::Mixture(Arc::new(d))
309    }
310}
311
312impl From<Arc<Mixture>> for Dist {
313    fn from(d: Arc<Mixture>) -> Self {
314        Self::Mixture(d)
315    }
316}
317
318#[cfg(test)]
319mod tests {
320    use super::*;
321
322    fn all() -> Vec<Dist> {
323        let ln = Lognormal::from_mean_cv(1000.0, 0.8).unwrap();
324        vec![
325            ln.into(),
326            Pareto::new(500.0, 2.5).unwrap().into(),
327            Gamma::new(2.0, 500.0).unwrap().into(),
328            Weibull::new(1.5, 1000.0).unwrap().into(),
329            Loglogistic::new(4.0, 900.0).unwrap().into(),
330            Tweedie::new(1000.0, 2.0, 1.5).unwrap().into(),
331            Gpd::new(0.2, 300.0).unwrap().into(),
332            Mixture::new(vec![
333                (0.5, Box::new(ln) as Box<dyn Severity + Send + Sync>),
334                (0.5, Box::new(Gamma::new(2.0, 500.0).unwrap())),
335            ])
336            .unwrap()
337            .into(),
338            Sampled::new(vec![1.0, 2.0, 3.0, 10.0]).unwrap().into(),
339            Custom::new(
340                "exponential",
341                Arc::new(|x: f64| Ok((1.0 - (-x / 500.0).exp()).max(0.0))),
342                None,
343                false,
344            )
345            .unwrap()
346            .into(),
347        ]
348    }
349
350    #[test]
351    fn dispatch_matches_the_family() {
352        let ln = Lognormal::from_mean_cv(1000.0, 0.8).unwrap();
353        let d = Dist::from(ln);
354        assert_eq!(d.family(), "lognormal");
355        assert_eq!(d.mean(), ln.mean());
356        assert_eq!(d.variance(), ln.variance());
357        assert_eq!(d.cdf(700.0), ln.cdf(700.0));
358        assert_eq!(d.quantile(0.9).unwrap(), ln.quantile(0.9).unwrap());
359        let s = d.as_severity().unwrap();
360        assert_eq!(s.lev(1500.0), ln.lev(1500.0));
361        assert_eq!(s.layer(500.0, 1000.0), ln.layer(500.0, 1000.0));
362        let (mut a, mut b) = (StreamRng::new(1, 2), StreamRng::new(1, 2));
363        assert_eq!(d.sample(&mut a, 5), ln.sample(&mut b, 5));
364    }
365
366    #[test]
367    fn every_variant_is_a_distribution_and_severities_have_layers() {
368        for d in all() {
369            assert!(d.mean().is_finite(), "{}", d.family());
370            assert_eq!(d.is_parallel_safe(), d.family() != "custom");
371            let q = d.quantile(0.5).unwrap();
372            assert!(d.cdf(q) >= 0.5 - 1e-9, "{}", d.family());
373            match d.as_severity() {
374                Some(s) => {
375                    let m = s.lev(f64::INFINITY);
376                    assert!((m / d.mean() - 1.0).abs() < 1e-6, "{}", d.family());
377                }
378                None => assert_eq!(d.family(), "sampled"),
379            }
380            // Clones share a mixture rather than copying it.
381            let c = d.clone();
382            assert_eq!(c.mean(), d.mean());
383        }
384    }
385
386    #[test]
387    fn severity_dist_matches_as_severity_and_rejects_sampled() {
388        for d in all() {
389            let family = d.family();
390            match SeverityDist::try_from(d.clone()) {
391                Ok(s) => {
392                    let r = d.as_severity().unwrap();
393                    assert_eq!(s.lev(700.0), r.lev(700.0), "{family}");
394                    assert_eq!(s.stop_loss(700.0), r.stop_loss(700.0), "{family}");
395                    assert_eq!(s.layer(500.0, 200.0), r.layer(500.0, 200.0), "{family}");
396                    assert_eq!(
397                        s.layer_second_moment(500.0, 200.0),
398                        r.layer_second_moment(500.0, 200.0),
399                        "{family}"
400                    );
401                    assert_eq!(s.mean(), d.mean(), "{family}");
402                    assert_eq!(s.dist().family(), family);
403                }
404                Err(back) => {
405                    assert_eq!(family, "sampled");
406                    assert_eq!(back.family(), "sampled");
407                }
408            }
409        }
410    }
411}