1use prospicio_core::{Error, Result};
4
5use crate::distribution::{Distribution, check_probability};
6use crate::severity::Severity;
7
8#[derive(Debug, Clone, PartialEq)]
32pub struct Grid {
33 step: f64,
34 probs: Vec<f64>,
35}
36
37#[derive(Debug, Clone, PartialEq)]
40pub struct DiscretizationReport {
41 pub method: Discretization,
42 pub step: f64,
43 pub points: usize,
44 pub tail_mass: f64,
47 pub source_mean: f64,
49 pub grid_mean: f64,
51}
52
53impl DiscretizationReport {
54 pub fn mean_error(&self) -> f64 {
56 self.grid_mean - self.source_mean
57 }
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq)]
62pub enum Discretization {
63 LocalMoment,
67 Rounding,
69 Lower,
72}
73
74impl Grid {
75 pub fn new(step: f64, probs: Vec<f64>) -> Result<Self> {
81 if !step.is_finite() || step <= 0.0 {
82 return Err(Error::InvalidParameter {
83 name: "step",
84 value: step,
85 reason: "must be finite and positive",
86 });
87 }
88 if probs.is_empty() {
89 return Err(Error::InvalidParameter {
90 name: "probs",
91 value: 0.0,
92 reason: "must not be empty",
93 });
94 }
95 if let Some(&bad) = probs.iter().find(|p| !p.is_finite() || **p < 0.0) {
96 return Err(Error::InvalidParameter {
97 name: "probs",
98 value: bad,
99 reason: "must all be finite and non-negative",
100 });
101 }
102 let total: f64 = probs.iter().sum();
103 if (total - 1.0).abs() > 1e-9 {
104 return Err(Error::InvalidParameter {
105 name: "probs",
106 value: total,
107 reason: "must sum to 1",
108 });
109 }
110 Ok(Self { step, probs })
111 }
112
113 pub fn local_moment<D: Severity>(
124 source: &D,
125 step: f64,
126 points: usize,
127 ) -> Result<(Self, DiscretizationReport)> {
128 check_shape(step, points)?;
129 let lev: Vec<f64> = (0..points).map(|j| source.lev(j as f64 * step)).collect();
130 let mean = source.mean();
135 let tail_from = (0..points)
136 .find(|&j| j as f64 * step >= mean)
137 .unwrap_or(points);
138 let sl: Vec<f64> = (tail_from.saturating_sub(1)..points)
139 .map(|j| source.stop_loss(j as f64 * step))
140 .collect();
141 let sl = |j: usize| sl[j + 1 - tail_from.max(1)];
142 let probs = if points == 1 {
143 vec![1.0]
144 } else {
145 let mut probs = Vec::with_capacity(points);
146 probs.push(1.0 - lev[1] / step);
147 for j in 1..points - 1 {
148 let f = if j > tail_from {
151 sl(j - 1) - 2.0 * sl(j) + sl(j + 1)
152 } else {
153 2.0 * lev[j] - lev[j - 1] - lev[j + 1]
154 };
155 probs.push((f / step).max(0.0));
156 }
157 probs.push(if points - 1 > tail_from {
158 (sl(points - 2) - sl(points - 1)) / step
159 } else {
160 (lev[points - 1] - lev[points - 2]) / step
161 });
162 probs
163 };
164 Self::finish(source, step, probs, Discretization::LocalMoment)
165 }
166
167 pub fn rounding<D: Distribution>(
171 source: &D,
172 step: f64,
173 points: usize,
174 ) -> Result<(Self, DiscretizationReport)> {
175 check_shape(step, points)?;
176 let edges: Vec<f64> = (0..points.saturating_sub(1))
177 .map(|j| source.cdf((j as f64 + 0.5) * step))
178 .collect();
179 Self::finish(source, step, cells(&edges), Discretization::Rounding)
180 }
181
182 pub fn lower<D: Distribution>(
186 source: &D,
187 step: f64,
188 points: usize,
189 ) -> Result<(Self, DiscretizationReport)> {
190 check_shape(step, points)?;
191 let edges: Vec<f64> = (1..points).map(|j| source.cdf(j as f64 * step)).collect();
192 Self::finish(source, step, cells(&edges), Discretization::Lower)
193 }
194
195 fn finish<D: Distribution>(
196 source: &D,
197 step: f64,
198 probs: Vec<f64>,
199 method: Discretization,
200 ) -> Result<(Self, DiscretizationReport)> {
201 let points = probs.len();
202 let grid = Self::new(step, probs)?;
203 let report = DiscretizationReport {
204 method,
205 step,
206 points,
207 tail_mass: 1.0 - source.cdf((points - 1) as f64 * step),
208 source_mean: source.mean(),
209 grid_mean: grid.mean(),
210 };
211 Ok((grid, report))
212 }
213
214 pub fn step(&self) -> f64 {
216 self.step
217 }
218
219 pub fn probs(&self) -> &[f64] {
221 &self.probs
222 }
223
224 pub fn len(&self) -> usize {
226 self.probs.len()
227 }
228
229 pub fn is_empty(&self) -> bool {
231 false
232 }
233
234 pub fn x(&self, j: usize) -> f64 {
236 j as f64 * self.step
237 }
238
239 pub fn distortion(&self, d: &crate::Distortion) -> f64 {
250 let values: Vec<f64> = (0..self.len()).map(|j| self.x(j)).collect();
251 d.apply_discrete(&values, &self.probs)
252 }
253
254 pub fn map(&self, mut f: impl FnMut(f64) -> f64) -> Result<(Self, bool)> {
283 let mut probs = vec![0.0; 1];
284 let mut exact = true;
285 let add = |probs: &mut Vec<f64>, k: usize, p: f64| {
286 if probs.len() <= k {
287 probs.resize(k + 1, 0.0);
288 }
289 probs[k] += p;
290 };
291 for (x, p) in self.points() {
292 if p == 0.0 {
293 continue;
294 }
295 let y = f(x);
296 if !y.is_finite() || y < 0.0 {
297 return Err(Error::InvalidParameter {
298 name: "f",
299 value: y,
300 reason: "must map every point with mass to a finite, non-negative value",
301 });
302 }
303 let at = y / self.step;
304 let nearest = at.round();
305 if (at - nearest).abs() <= 1e-9 * nearest.max(1.0) {
306 add(&mut probs, nearest as usize, p);
307 } else {
308 exact = false;
309 let k = at.floor();
310 let upper = at - k;
311 add(&mut probs, k as usize, p * (1.0 - upper));
312 add(&mut probs, k as usize + 1, p * upper);
313 }
314 }
315 Ok((Self::new(self.step, probs)?, exact))
317 }
318
319 fn points(&self) -> impl Iterator<Item = (f64, f64)> + '_ {
320 self.probs.iter().enumerate().map(|(j, &p)| (self.x(j), p))
321 }
322}
323
324fn cells(edges: &[f64]) -> Vec<f64> {
327 let mut probs = Vec::with_capacity(edges.len() + 1);
328 let mut below = 0.0;
329 for &f in edges {
330 probs.push((f - below).max(0.0));
331 below = f;
332 }
333 probs.push(1.0 - below);
334 probs
335}
336
337fn check_shape(step: f64, points: usize) -> Result<()> {
338 if !step.is_finite() || step <= 0.0 {
339 return Err(Error::InvalidParameter {
340 name: "step",
341 value: step,
342 reason: "must be finite and positive",
343 });
344 }
345 if points == 0 {
346 return Err(Error::InvalidParameter {
347 name: "points",
348 value: 0.0,
349 reason: "must be positive",
350 });
351 }
352 Ok(())
353}
354
355impl Distribution for Grid {
356 fn mean(&self) -> f64 {
357 self.points().map(|(x, p)| x * p).sum()
358 }
359
360 fn variance(&self) -> f64 {
361 let mean = self.mean();
362 self.points()
363 .map(|(x, p)| (x - mean) * (x - mean) * p)
364 .sum()
365 }
366
367 fn cdf(&self, x: f64) -> f64 {
368 if x < 0.0 {
369 return 0.0;
370 }
371 let last = ((x / self.step).floor() as usize).min(self.len() - 1);
372 self.probs[..=last].iter().sum::<f64>().min(1.0)
373 }
374
375 fn survival(&self, x: f64) -> f64 {
378 if x < 0.0 {
379 return 1.0;
380 }
381 let last = ((x / self.step).floor() as usize).min(self.len() - 1);
382 self.probs[last + 1..].iter().sum::<f64>().min(1.0)
383 }
384
385 fn quantile(&self, p: f64) -> Result<f64> {
386 check_probability(p)?;
387 if p == 0.0 {
388 return Ok(0.0);
389 }
390 let mut total = 0.0;
391 for (j, &q) in self.probs.iter().enumerate() {
392 total += q;
393 if q > 0.0 && total >= p {
394 return Ok(self.x(j));
395 }
396 }
397 let last = self.probs.iter().rposition(|&q| q > 0.0).unwrap_or(0);
399 Ok(self.x(last))
400 }
401}
402
403impl Severity for Grid {
404 fn lev(&self, limit: f64) -> f64 {
405 if limit <= 0.0 {
406 return limit;
407 }
408 self.points().map(|(x, p)| x.min(limit) * p).sum()
409 }
410
411 fn stop_loss(&self, retention: f64) -> f64 {
412 if retention <= 0.0 {
413 return self.mean() - retention;
414 }
415 self.points()
416 .map(|(x, p)| (x - retention).max(0.0) * p)
417 .sum()
418 }
419
420 fn layer_second_moment(&self, limit: f64, attachment: f64) -> f64 {
421 self.points()
422 .map(|(x, p)| {
423 let y = (x - attachment).max(0.0).min(limit);
424 y * y * p
425 })
426 .sum()
427 }
428}
429
430#[cfg(test)]
431mod tests {
432 use super::*;
433 use crate::Lognormal;
434
435 #[test]
436 fn map_keeps_mass_and_mean_off_the_points() {
437 let x = Grid::new(2.0, vec![0.1, 0.2, 0.3, 0.25, 0.15]).unwrap();
438 let f = |v: f64| 0.7 * v + 0.3;
439 let (y, exact) = x.map(f).unwrap();
440 assert!(!exact);
441 assert!((y.probs().iter().sum::<f64>() - 1.0).abs() < 1e-15);
442 let want: f64 = (0..x.len()).map(|j| f(x.x(j)) * x.probs()[j]).sum();
443 assert!((y.mean() - want).abs() < 1e-14);
444 assert_eq!(y.step(), 2.0);
445 }
446
447 #[test]
448 fn map_can_grow_the_grid_and_rejects_negative_values() {
449 let x = Grid::new(1.0, vec![0.5, 0.5]).unwrap();
450 let (y, exact) = x.map(|v| 3.0 * v).unwrap();
451 assert!(exact);
452 assert_eq!(y.probs(), [0.5, 0.0, 0.0, 0.5]);
453 assert!(x.map(|v| v - 0.5).is_err());
454 }
455
456 fn sev() -> Lognormal {
457 Lognormal::new(7.0, 0.5).unwrap()
458 }
459
460 #[test]
461 fn local_moment_preserves_the_limited_mean() {
462 let s = sev();
463 for (step, n) in [(100.0, 200), (250.0, 20), (50.0, 30)] {
464 let (g, r) = Grid::local_moment(&s, step, n).unwrap();
465 let lev = s.lev((n - 1) as f64 * step);
466 assert!((g.mean() - lev).abs() < 1e-9 * lev, "h {step}, n {n}");
467 assert!(g.probs().iter().all(|&p| p >= 0.0));
468 assert_eq!(r.grid_mean, g.mean());
469 assert!((r.mean_error() + s.stop_loss((n - 1) as f64 * step)).abs() < 1e-9);
470 }
471 }
472
473 #[test]
474 fn probabilities_sum_to_one_and_tail_is_reported() {
475 let s = sev();
476 for method in [Grid::rounding::<Lognormal>, Grid::lower::<Lognormal>] {
477 let (g, r) = method(&s, 100.0, 20).unwrap();
478 assert!((g.probs().iter().sum::<f64>() - 1.0).abs() < 1e-12);
479 assert_eq!(r.points, 20);
480 assert!((r.tail_mass - (1.0 - s.cdf(1900.0))).abs() < 1e-15);
481 }
482 let (_, r) = Grid::local_moment(&s, 100.0, 20).unwrap();
483 assert!(r.tail_mass > 0.01, "a short grid truncates");
484 }
485
486 #[test]
487 fn rounding_and_lower_cells() {
488 let s = sev();
489 let (g, _) = Grid::rounding(&s, 100.0, 50).unwrap();
490 assert!((g.probs()[0] - s.cdf(50.0)).abs() < 1e-15);
491 assert!((g.probs()[10] - (s.cdf(1050.0) - s.cdf(950.0))).abs() < 1e-15);
492 let (g, _) = Grid::lower(&s, 100.0, 50).unwrap();
493 assert!((g.probs()[0] - s.cdf(100.0)).abs() < 1e-15);
494 assert!((g.probs()[10] - (s.cdf(1100.0) - s.cdf(1000.0))).abs() < 1e-15);
495 }
496
497 #[test]
498 fn lower_is_a_stochastic_lower_bound() {
499 let s = sev();
500 let (g, _) = Grid::lower(&s, 100.0, 60).unwrap();
501 for x in (0..60).map(|j| j as f64 * 100.0 + 50.0) {
502 assert!(g.cdf(x) >= s.cdf(x) - 1e-15, "at {x}");
503 }
504 assert!(g.mean() <= s.mean());
505 }
506
507 #[test]
508 fn methods_converge_as_the_step_shrinks() {
509 let s = sev();
510 let mut prev = f64::INFINITY;
511 for step in [400.0, 100.0, 25.0] {
512 let n = (40_000.0 / step) as usize;
513 let (g, _) = Grid::rounding(&s, step, n).unwrap();
514 let err = (g.mean() - s.mean()).abs();
515 assert!(err < prev);
516 prev = err;
517 }
518 let (g, _) = Grid::local_moment(&s, 25.0, 1600).unwrap();
519 assert!((g.quantile(0.99).unwrap() - s.quantile(0.99).unwrap()).abs() <= 25.0);
520 }
521
522 #[test]
523 fn grid_distribution_methods() {
524 let g = Grid::new(10.0, vec![0.5, 0.25, 0.25]).unwrap();
525 assert_eq!(g.mean(), 7.5);
526 assert_eq!(g.variance(), 0.5 * 56.25 + 0.25 * 6.25 + 0.25 * 156.25);
527 assert_eq!(g.cdf(-1.0), 0.0);
528 assert_eq!(g.cdf(0.0), 0.5);
529 assert_eq!(g.cdf(15.0), 0.75);
530 assert_eq!(g.cdf(1e9), 1.0);
531 assert_eq!(g.quantile(0.0), Ok(0.0));
532 assert_eq!(g.quantile(0.5), Ok(0.0));
533 assert_eq!(g.quantile(0.6), Ok(10.0));
534 assert_eq!(g.quantile(1.0), Ok(20.0));
535 assert_eq!(g.lev(15.0), 0.25 * 10.0 + 0.25 * 15.0);
536 assert_eq!(g.stop_loss(15.0), 0.25 * 5.0);
537 assert_eq!(g.layer(5.0, 10.0), 0.25 * 5.0);
538 }
539
540 #[test]
541 fn quantile_skips_points_without_mass() {
542 let g = Grid::new(1.0, vec![0.5, 0.0, 0.5]).unwrap();
543 assert_eq!(g.quantile(0.5), Ok(0.0));
544 assert_eq!(g.quantile(0.50001), Ok(2.0));
545 }
546
547 #[test]
548 fn rejects_bad_input() {
549 assert!(Grid::new(0.0, vec![1.0]).is_err());
550 assert!(Grid::new(1.0, vec![]).is_err());
551 assert!(Grid::new(1.0, vec![0.5, 0.4]).is_err());
552 assert!(Grid::new(1.0, vec![1.5, -0.5]).is_err());
553 assert!(Grid::local_moment(&sev(), 100.0, 0).is_err());
554 assert!(Grid::rounding(&sev(), f64::NAN, 10).is_err());
555 let (g, _) = Grid::local_moment(&sev(), 100.0, 1).unwrap();
556 assert_eq!(g.probs(), [1.0]);
557 }
558}