1use prospicio_core::Result;
6
7use crate::distribution::{Distribution, check_probability};
8use crate::pareto::{invalid, power_integral, raw_integral};
9use crate::severity::Severity;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum Truncation {
15 LastPiece,
18 WholeDistribution,
20}
21
22#[derive(Debug, Clone, PartialEq)]
49pub struct PiecewisePareto {
50 t: Vec<f64>,
51 alpha: Vec<f64>,
52 log_s: Vec<f64>,
54 truncation: Option<(f64, Truncation)>,
55}
56
57impl PiecewisePareto {
58 pub fn new(t: Vec<f64>, alpha: Vec<f64>) -> Result<Self> {
62 if t.is_empty() || t.len() != alpha.len() {
63 return Err(invalid(
64 "alpha",
65 alpha.len() as f64,
66 "must be non-empty and have one alpha per threshold",
67 ));
68 }
69 for (i, &x) in t.iter().enumerate() {
70 if !x.is_finite() || x <= 0.0 {
71 return Err(invalid("t", x, "must be finite and positive"));
72 }
73 if i > 0 && x <= t[i - 1] {
74 return Err(invalid("t", x, "must be strictly increasing"));
75 }
76 }
77 for &a in &alpha {
78 if !a.is_finite() || a < 0.0 {
79 return Err(invalid("alpha", a, "must be finite and non-negative"));
80 }
81 }
82 let last = alpha[alpha.len() - 1];
83 if last <= 0.0 {
84 return Err(invalid("alpha", last, "the last alpha must be positive"));
85 }
86 let mut log_s = Vec::with_capacity(t.len());
87 let mut acc = 0.0;
88 for k in 0..t.len() {
89 if k > 0 {
90 acc += alpha[k - 1] * (t[k - 1] / t[k]).ln();
91 }
92 log_s.push(acc);
93 }
94 Ok(Self {
95 t,
96 alpha,
97 log_s,
98 truncation: None,
99 })
100 }
101
102 pub fn truncated(self, truncation: f64, kind: Truncation) -> Result<Self> {
105 let last = self.t[self.t.len() - 1];
106 if !truncation.is_finite() || truncation <= last {
107 return Err(invalid(
108 "truncation",
109 truncation,
110 "must be finite and above the last threshold",
111 ));
112 }
113 Ok(Self {
114 truncation: Some((truncation, kind)),
115 ..self
116 })
117 }
118
119 pub fn thresholds(&self) -> &[f64] {
121 &self.t
122 }
123
124 pub fn alphas(&self) -> &[f64] {
126 &self.alpha
127 }
128
129 pub fn truncation(&self) -> Option<(f64, Truncation)> {
131 self.truncation
132 }
133
134 fn survival_at(&self, x: f64) -> f64 {
136 match self.truncation {
137 Some((tr, Truncation::WholeDistribution)) => {
138 if x >= tr {
139 return 0.0;
140 }
141 if x < self.t[0] {
142 return 1.0;
143 }
144 let (k, k_tr) = (self.piece(x), self.piece(tr));
146 let log_ratio = if k == k_tr {
147 self.alpha[k] * (x / tr).ln()
148 } else {
149 self.log_survival(tr) - self.log_survival(x)
150 };
151 let (_, one_minus) = self.whole_mass(tr);
152 self.log_survival(x).exp() * -log_ratio.exp_m1() / one_minus
153 }
154 _ => self.base_survival(x),
155 }
156 }
157
158 fn piece(&self, x: f64) -> usize {
160 self.t.partition_point(|&t| t <= x) - 1
161 }
162
163 fn last_piece_truncation(&self) -> Option<f64> {
165 match self.truncation {
166 Some((tr, Truncation::LastPiece)) => Some(tr),
167 _ => None,
168 }
169 }
170
171 fn last_piece_mass(&self, tr: f64) -> (f64, f64) {
173 let n = self.t.len() - 1;
174 let log_q = self.alpha[n] * (self.t[n] / tr).ln();
175 (log_q.exp(), -log_q.exp_m1())
176 }
177
178 fn log_survival(&self, x: f64) -> f64 {
180 let k = self.piece(x);
181 self.log_s[k] + self.alpha[k] * (self.t[k] / x).ln()
182 }
183
184 fn whole_mass(&self, tr: f64) -> (f64, f64) {
186 let log = self.log_survival(tr);
187 (log.exp(), -log.exp_m1())
188 }
189
190 fn base_survival(&self, x: f64) -> f64 {
193 if x < self.t[0] {
194 return 1.0;
195 }
196 let k = self.piece(x);
197 match self.last_piece_truncation() {
198 Some(tr) if k == self.t.len() - 1 => {
199 if x >= tr {
200 return 0.0;
201 }
202 let (_, one_minus_q) = self.last_piece_mass(tr);
204 let rest = -(self.alpha[k] * (x / tr).ln()).exp_m1();
205 self.log_s[k].exp() * (self.t[k] / x).powf(self.alpha[k]) * rest / one_minus_q
206 }
207 _ => self.log_survival(x).exp(),
208 }
209 }
210
211 fn integral(&self, k: i32, a: f64, b: f64) -> f64 {
213 match self.truncation {
214 Some((tr, Truncation::WholeDistribution)) => {
215 let (a, b) = (a.min(tr), b.min(tr));
216 if a >= b {
217 return 0.0;
218 }
219 let (s_tr, one_minus) = self.whole_mass(tr);
220 (self.base_integral(k, a, b) - s_tr * power_integral(k, a, b)) / one_minus
221 }
222 _ => self.base_integral(k, a, b),
223 }
224 }
225
226 fn base_integral(&self, k: i32, a: f64, b: f64) -> f64 {
229 if a >= b {
230 return 0.0;
231 }
232 let n = self.t.len();
233 let mut sum = if a < self.t[0] {
234 power_integral(k, a, b.min(self.t[0]))
235 } else {
236 0.0
237 };
238 for j in 0..n {
239 let lo = a.max(self.t[j]);
240 let mut hi = if j + 1 < n { b.min(self.t[j + 1]) } else { b };
241 let truncated = if j + 1 == n {
242 self.last_piece_truncation()
243 } else {
244 None
245 };
246 if let Some(tr) = truncated {
247 hi = hi.min(tr);
248 }
249 if lo >= hi {
250 continue;
251 }
252 let s_j = self.log_s[j].exp();
253 if s_j == 0.0 {
254 break;
255 }
256 let piece = raw_integral(k, self.t[j], self.alpha[j], lo, hi);
257 sum += match truncated {
258 Some(tr) => {
259 let (q, one_minus_q) = self.last_piece_mass(tr);
260 s_j * (piece - q * power_integral(k, lo, hi)) / one_minus_q
261 }
262 None => s_j * piece,
263 };
264 }
265 sum
266 }
267
268 fn base_inverse(&self, s: f64) -> f64 {
270 let log = s.ln();
271 let k = self.log_s.partition_point(|&l| l >= log) - 1;
273 match self.last_piece_truncation() {
274 Some(tr) if k == self.t.len() - 1 => {
275 let (q, one_minus_q) = self.last_piece_mass(tr);
276 let rel = (log - self.log_s[k]).exp();
277 (self.t[k] * (q + rel * one_minus_q).powf(-1.0 / self.alpha[k])).min(tr)
278 }
279 _ => self.t[k] * ((self.log_s[k] - log) / self.alpha[k]).exp(),
280 }
281 }
282}
283
284impl Distribution for PiecewisePareto {
285 fn mean(&self) -> f64 {
286 self.integral(0, 0.0, f64::INFINITY)
287 }
288
289 fn variance(&self) -> f64 {
290 let m = self.mean();
291 if m == f64::INFINITY {
292 return f64::INFINITY;
293 }
294 2.0 * self.integral(1, 0.0, f64::INFINITY) - m * m
295 }
296
297 fn cdf(&self, x: f64) -> f64 {
298 1.0 - self.survival_at(x)
299 }
300
301 fn survival(&self, x: f64) -> f64 {
302 self.survival_at(x)
303 }
304
305 fn quantile(&self, p: f64) -> Result<f64> {
307 check_probability(p)?;
308 let s = match self.truncation {
309 Some((tr, Truncation::WholeDistribution)) => {
310 let (s_tr, one_minus) = self.whole_mass(tr);
311 return Ok(self.base_inverse(s_tr + (1.0 - p) * one_minus).min(tr));
312 }
313 Some((tr, Truncation::LastPiece)) if p == 1.0 => return Ok(tr),
314 _ => 1.0 - p,
315 };
316 if s <= 0.0 {
317 return Ok(f64::INFINITY);
318 }
319 Ok(self.base_inverse(s))
320 }
321}
322
323impl Severity for PiecewisePareto {
324 fn lev(&self, limit: f64) -> f64 {
325 if limit <= 0.0 {
326 return limit;
327 }
328 self.integral(0, 0.0, limit)
329 }
330
331 fn stop_loss(&self, retention: f64) -> f64 {
332 if retention <= 0.0 {
333 return self.mean() - retention;
334 }
335 self.integral(0, retention, f64::INFINITY)
336 }
337
338 fn layer(&self, limit: f64, attachment: f64) -> f64 {
339 let a = attachment.max(0.0);
340 self.integral(0, a, a + limit)
341 }
342
343 fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
345 let a = attachment.max(0.0);
346 let b = a + limit;
347 2.0 * (self.integral(1, a, b) - a * self.integral(0, a, b))
348 }
349}
350
351#[cfg(test)]
352mod tests {
353 use super::*;
354 use crate::Pareto;
355
356 fn close(a: f64, b: f64, rel: f64) -> bool {
357 (a - b).abs() <= rel * b.abs().max(1e-300)
358 }
359
360 fn example() -> PiecewisePareto {
361 PiecewisePareto::new(vec![1000.0, 2000.0, 3000.0], vec![1.0, 1.5, 2.0]).unwrap()
362 }
363
364 #[test]
365 fn one_piece_is_a_pareto() {
366 let pp = PiecewisePareto::new(vec![500.0], vec![1.7]).unwrap();
367 let p = Pareto::new(500.0, 1.7).unwrap();
368 let ppt = pp.clone().truncated(9000.0, Truncation::LastPiece).unwrap();
369 let ppw = pp
370 .clone()
371 .truncated(9000.0, Truncation::WholeDistribution)
372 .unwrap();
373 let pt = p.truncated(9000.0).unwrap();
374 for (a, b) in [(1000.0, 0.0), (4000.0, 1000.0), (f64::INFINITY, 2000.0)] {
375 assert!(close(pp.layer(a, b), p.layer(a, b), 1e-14));
376 assert!(close(
377 pp.layer_second_moment(a.min(1e5), b),
378 p.layer_second_moment(a.min(1e5), b),
379 1e-13
380 ));
381 for q in [&ppt, &ppw] {
383 assert!(close(q.layer(a, b), pt.layer(a, b), 1e-13));
384 assert!(close(
385 q.layer_second_moment(a, b),
386 pt.layer_second_moment(a, b),
387 1e-12
388 ));
389 }
390 }
391 for x in [100.0, 500.0, 700.0, 8999.0] {
392 assert!(close(ppt.survival(x), pt.survival(x), 1e-13));
393 assert!(close(ppw.survival(x), pt.survival(x), 1e-13));
394 }
395 }
396
397 #[test]
398 fn survival_is_continuous_with_the_stated_alphas() {
399 let pp = example();
400 assert_eq!(pp.survival(999.0), 1.0);
401 assert!(close(pp.survival(2000.0), 0.5, 1e-15));
402 assert!(close(
403 pp.survival(3000.0),
404 0.5 * (2.0f64 / 3.0).powf(1.5),
405 1e-15
406 ));
407 for &t in pp.thresholds() {
408 let below = pp.survival(t * (1.0 - 1e-12));
409 assert!(close(below, pp.survival(t), 1e-10), "{t}");
410 }
411 for (x, alpha) in [(1500.0, 1.0), (2500.0, 1.5), (9000.0, 2.0)] {
413 let h = x * 1e-6;
414 let d = (pp.survival(x + h).ln() - pp.survival(x - h).ln()) / (2.0 * h);
415 assert!(close(-x * d, alpha, 1e-6));
416 }
417 }
418
419 #[test]
420 fn severity_identities() {
421 for pp in [
422 example(),
423 example().truncated(5000.0, Truncation::LastPiece).unwrap(),
424 example()
425 .truncated(5000.0, Truncation::WholeDistribution)
426 .unwrap(),
427 PiecewisePareto::new(vec![100.0, 200.0, 400.0], vec![0.5, 0.0, 3.0]).unwrap(),
428 ] {
429 for d in [500.0, 1000.0, 2500.0, 4000.0] {
430 assert!(
431 close(pp.lev(d) + pp.stop_loss(d), pp.mean(), 1e-12),
432 "{pp:?} at {d}"
433 );
434 assert!(close(
435 pp.layer(1700.0, d),
436 pp.stop_loss(d) - pp.stop_loss(d + 1700.0),
437 1e-10
438 ));
439 }
440 let m2 = pp.layer_second_moment(f64::INFINITY, 0.0);
441 if pp.variance().is_finite() {
442 assert!(close(m2 - pp.mean() * pp.mean(), pp.variance(), 1e-10));
443 } else {
444 assert_eq!(m2, f64::INFINITY);
445 }
446 }
447 let heavy = PiecewisePareto::new(vec![1.0, 2.0], vec![3.0, 0.9]).unwrap();
448 assert_eq!(heavy.mean(), f64::INFINITY);
449 assert_eq!(heavy.variance(), f64::INFINITY);
450 }
451
452 #[test]
453 fn cdf_quantile_round_trip() {
454 for pp in [
455 example(),
456 example().truncated(5000.0, Truncation::LastPiece).unwrap(),
457 example()
458 .truncated(5000.0, Truncation::WholeDistribution)
459 .unwrap(),
460 PiecewisePareto::new(vec![100.0, 200.0, 400.0], vec![0.5, 0.0, 3.0]).unwrap(),
461 ] {
462 for q in [0.0, 0.1, 0.25, 0.5, 0.75, 0.9, 0.999] {
463 let x = pp.quantile(q).unwrap();
464 assert!(close(pp.cdf(x), q, 1e-12) || q == 0.0, "{pp:?} {q}");
465 }
466 assert_eq!(pp.quantile(0.0).unwrap(), pp.thresholds()[0]);
467 }
468 assert_eq!(example().quantile(1.0).unwrap(), f64::INFINITY);
469 let t = example().truncated(5000.0, Truncation::LastPiece).unwrap();
470 assert_eq!(t.quantile(1.0).unwrap(), 5000.0);
471 assert_eq!(t.survival(5000.0), 0.0);
472 }
473
474 #[test]
475 fn steep_pieces_do_not_underflow() {
476 let pp = PiecewisePareto::new(vec![1.0, 10.0, 100.0], vec![200.0, 150.0, 2.0]).unwrap();
477 assert_eq!(pp.survival(200.0), 0.0);
479 assert!(pp.mean().is_finite());
480 let x = pp.quantile(0.5).unwrap();
481 assert!(close(pp.cdf(x), 0.5, 1e-12));
482 }
483
484 #[test]
485 fn rejects_bad_parameters() {
486 assert!(PiecewisePareto::new(vec![], vec![]).is_err());
487 assert!(PiecewisePareto::new(vec![1.0, 2.0], vec![1.0]).is_err());
488 assert!(PiecewisePareto::new(vec![2.0, 1.0], vec![1.0, 1.0]).is_err());
489 assert!(PiecewisePareto::new(vec![1.0, 2.0], vec![1.0, 0.0]).is_err());
490 assert!(PiecewisePareto::new(vec![1.0, 2.0], vec![-1.0, 1.0]).is_err());
491 assert!(PiecewisePareto::new(vec![0.0, 2.0], vec![1.0, 1.0]).is_err());
492 assert!(example().truncated(2500.0, Truncation::LastPiece).is_err());
493 }
494}