prospicio_prob/
mixture.rs1use prospicio_core::{Error, Result};
4use prospicio_math::roots::bisect;
5
6use crate::dist::SeverityDist;
7use crate::distribution::{Distribution, check_probability};
8use crate::severity::Severity;
9
10pub struct Mixture {
31 weights: Vec<f64>,
32 components: Vec<Box<dyn Severity + Send + Sync>>,
33 dists: Option<Vec<SeverityDist>>,
36}
37
38impl std::fmt::Debug for Mixture {
39 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
40 f.debug_struct("Mixture")
41 .field("weights", &self.weights)
42 .field("components", &self.components.len())
43 .finish()
44 }
45}
46
47impl Mixture {
48 pub fn new(parts: Vec<(f64, Box<dyn Severity + Send + Sync>)>) -> Result<Self> {
51 if parts.is_empty() {
52 return Err(Error::InvalidParameter {
53 name: "components",
54 value: 0.0,
55 reason: "must not be empty",
56 });
57 }
58 if let Some((w, _)) = parts.iter().find(|(w, _)| !(w.is_finite() && *w > 0.0)) {
59 return Err(Error::InvalidParameter {
60 name: "weights",
61 value: *w,
62 reason: "must be finite and positive",
63 });
64 }
65 let total: f64 = parts.iter().map(|(w, _)| w).sum();
66 if (total - 1.0).abs() > 1e-12 {
67 return Err(Error::InvalidParameter {
68 name: "weights",
69 value: total,
70 reason: "must sum to 1",
71 });
72 }
73 let (weights, components) = parts.into_iter().unzip();
74 Ok(Self {
75 weights,
76 components,
77 dists: None,
78 })
79 }
80
81 pub fn from_dists(parts: Vec<(f64, SeverityDist)>) -> Result<Self> {
85 let dists: Vec<SeverityDist> = parts.iter().map(|(_, d)| d.clone()).collect();
86 let mut m = Self::new(
87 parts
88 .into_iter()
89 .map(|(w, d)| (w, Box::new(d) as Box<dyn Severity + Send + Sync>))
90 .collect(),
91 )?;
92 m.dists = Some(dists);
93 Ok(m)
94 }
95
96 pub fn dists(&self) -> Option<&[SeverityDist]> {
99 self.dists.as_deref()
100 }
101
102 pub fn weights(&self) -> &[f64] {
104 &self.weights
105 }
106
107 fn sum(&self, f: impl Fn(&(dyn Severity + Send + Sync)) -> f64) -> f64 {
108 self.weights
109 .iter()
110 .zip(&self.components)
111 .map(|(w, c)| w * f(c.as_ref()))
112 .sum()
113 }
114}
115
116impl Distribution for Mixture {
117 fn mean(&self) -> f64 {
118 self.sum(|c| c.mean())
119 }
120
121 fn variance(&self) -> f64 {
123 let m = self.mean();
124 self.sum(|c| c.variance() + c.mean() * c.mean()) - m * m
125 }
126
127 fn cdf(&self, x: f64) -> f64 {
128 self.sum(|c| c.cdf(x))
129 }
130
131 fn survival(&self, x: f64) -> f64 {
132 self.sum(|c| c.survival(x))
133 }
134
135 fn is_parallel_safe(&self) -> bool {
136 self.components.iter().all(|c| c.is_parallel_safe())
137 }
138
139 fn quantile(&self, p: f64) -> Result<f64> {
142 check_probability(p)?;
143 if p == 1.0 {
144 return Ok(f64::INFINITY);
145 }
146 let qs = self
147 .components
148 .iter()
149 .map(|c| c.quantile(p))
150 .collect::<Result<Vec<_>>>()?;
151 let lo = qs.iter().copied().fold(f64::INFINITY, f64::min);
152 let hi = qs.iter().copied().fold(f64::NEG_INFINITY, f64::max);
153 if lo == hi {
154 return Ok(lo);
155 }
156 let below = |x: f64| {
157 if p <= 0.5 {
158 self.cdf(x) < p
159 } else {
160 self.survival(x) > 1.0 - p
161 }
162 };
163 Ok(bisect(lo, hi, below))
164 }
165}
166
167impl Severity for Mixture {
168 fn lev(&self, limit: f64) -> f64 {
169 self.sum(|c| c.lev(limit))
170 }
171
172 fn stop_loss(&self, retention: f64) -> f64 {
173 self.sum(|c| c.stop_loss(retention))
174 }
175
176 fn layer(&self, limit: f64, attachment: f64) -> f64 {
177 self.sum(|c| c.layer(limit, attachment))
178 }
179
180 fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
181 self.sum(|c| c.layer_second_moment(limit, attachment))
182 }
183}
184
185#[cfg(test)]
186mod tests {
187 use super::*;
188 use crate::{Gamma, Lognormal};
189
190 #[test]
191 fn mixture_of_one_is_the_component_and_quantiles_invert() {
192 let g = Gamma::new(2.0, 50.0).unwrap();
193 let m = Mixture::new(vec![(1.0, Box::new(g) as Box<dyn Severity + Send + Sync>)]).unwrap();
194 assert_eq!(m.variance(), g.variance());
195 let two = Mixture::new(vec![
196 (0.3, Box::new(g) as Box<dyn Severity + Send + Sync>),
197 (0.7, Box::new(Lognormal::new(6.0, 1.5).unwrap())),
198 ])
199 .unwrap();
200 for p in [1e-6, 0.3, 0.5, 0.95, 1.0 - 1e-9] {
201 let x = two.quantile(p).unwrap();
202 let err = if p <= 0.5 {
203 two.cdf(x) / p - 1.0
204 } else {
205 two.survival(x) / (1.0 - p) - 1.0
206 };
207 assert!(err.abs() < 1e-9, "{p} {err}");
208 }
209 assert!(Mixture::new(vec![(0.5, Box::new(g) as Box<dyn Severity + Send + Sync>)]).is_err());
210 }
211}