use super::Dist;
use super::special::{gamma_p, gamma_prefactor, gamma_q, ln_gamma};
pub(super) const COUNT_TAIL: f64 = 1e-12;
pub(super) const MAX_COUNT_TERMS: usize = 100_000;
pub(super) const COUNT_HEAD_SDS: f64 = 12.0;
pub(super) const BLOCK_GROWTH: f64 = 1e-3;
pub(super) struct CountBlock {
pub(super) k: f64,
pub(super) h: f64,
pub(super) mass: f64,
pub(super) next_h: f64,
pub(super) ln_step: f64,
}
impl CountBlock {
pub(super) fn tail(&self) -> f64 {
let next_ratio = self.next_h / self.h * self.ln_step.exp();
if next_ratio < 1.0 {
self.mass * next_ratio / (1.0 - next_ratio)
} else {
f64::INFINITY
}
}
}
#[derive(Clone)]
pub(super) struct CountBlocks {
pub(super) dist: Dist,
pub(super) h_max: f64,
pub(super) k: f64,
pub(super) h: f64,
pub(super) ln_center: f64,
}
impl CountBlocks {
pub(super) fn new(dist: Dist) -> Self {
let sd = dist.std_dev();
let start = (dist.mean() - COUNT_HEAD_SDS * sd).floor().max(0.0);
let h_max = (2.0 * COUNT_HEAD_SDS * sd / (MAX_COUNT_TERMS / 2) as f64)
.ceil()
.max(1.0);
let h = Self::width(h_max, start);
CountBlocks {
dist,
h_max,
k: start,
h,
ln_center: dist.log_prob(start + 0.5 * (h - 1.0)),
}
}
pub(super) fn width(h_max: f64, k: f64) -> f64 {
(BLOCK_GROWTH * (k + 1.0)).floor().clamp(1.0, h_max)
}
pub(super) fn step(&mut self) -> CountBlock {
let (k, h) = (self.k, self.h);
let mass = h * self.ln_center.exp();
let next_h = Self::width(self.h_max, k + h);
let d = f64::midpoint(h, next_h);
let c = k + 0.5 * (h - 1.0);
let ln_step = d * self.dist.count_ratio(c + 0.5 * (d - 1.0)).ln();
self.ln_center += ln_step;
self.k += h;
self.h = next_h;
CountBlock {
k,
h,
mass,
next_h,
ln_step,
}
}
}
pub(super) fn square_sum(a: f64, d: f64, i0: f64, i1: f64) -> f64 {
if i1 <= i0 {
return 0.0;
}
let squares = |x: f64| x * (x + 1.0) * (2.0 * x + 1.0) / 6.0;
let (s0, s1) = (i0 + 1.0, i1);
let n = s1 - s0 + 1.0;
n * a * a + a * d * (s0 + s1) * n + d * d * (squares(s1) - squares(s0 - 1.0))
}
pub(super) fn block_crps(k: f64, h: f64, before: f64, after: f64, y: f64) -> f64 {
let d = (after - before) / h;
let below = |i0, i1| square_sum(before, d, i0, i1);
let above = |i0, i1| square_sum(1.0 - before, -d, i0, i1);
if y <= k {
return above(0.0, h);
}
if y >= k + h {
return below(0.0, h);
}
let n = (y - k).floor();
let f = y - k - n;
let cdf = before + d * (n + 1.0);
let upper = 1.0 - cdf;
below(0.0, n) + f * cdf * cdf + (1.0 - f) * upper * upper + above(n + 1.0, h)
}
pub(super) fn geometric_tail_crps(u: f64, rho: f64, m: f64) -> f64 {
let r2 = rho * rho;
let beyond = |t: f64| u * u * rho.powf(2.0 * t) / (1.0 - r2);
if m <= 0.0 {
return beyond(1.0);
}
let n = m.floor();
let f = m - n;
let full = n - 2.0 * u * rho * (1.0 - rho.powf(n)) / (1.0 - rho)
+ u * u * r2 * (1.0 - rho.powf(2.0 * n)) / (1.0 - r2);
let q = u * rho.powf(n + 1.0);
full + f * (1.0 - q) * (1.0 - q) + (1.0 - f) * q * q + beyond(n + 2.0)
}
pub(super) fn gamma_unit_quantile(a: f64, p: f64) -> f64 {
if p <= 0.0 {
return 0.0;
}
if p >= 1.0 {
return f64::INFINITY;
}
let a1 = a - 1.0;
let mut x = if a > 1.0 {
let pp = if p < 0.5 { p } else { 1.0 - p };
let t = (-2.0 * pp.ln()).sqrt();
let mut z = (2.30753 + t * 0.27061) / (1.0 + t * (0.99229 + t * 0.04481)) - t;
if p < 0.5 {
z = -z;
}
let wilson_hilferty = a * (1.0 - 1.0 / (9.0 * a) - z / (3.0 * a.sqrt())).powi(3);
if p < 0.5 {
let asymptote = ((p.ln() + ln_gamma(a + 1.0)) / a).exp();
if wilson_hilferty > 1e-3 {
wilson_hilferty.max(asymptote)
} else {
asymptote
}
} else {
wilson_hilferty.max(1e-3)
}
} else {
let t = 1.0 - a * (0.253 + a * 0.12);
if p < t {
(p / t).powf(1.0 / a)
} else {
1.0 - (1.0 - (p - t) / (1.0 - t)).ln()
}
};
let residual = |x: f64| {
if p < 0.5 {
gamma_p(a, x) - p
} else {
(1.0 - p) - gamma_q(a, x)
}
};
for _ in 0..100 {
if x <= 0.0 {
return 0.0;
}
let density = gamma_prefactor(a, x) / x;
if density == 0.0 || !density.is_finite() {
break;
}
let u = residual(x) / density;
let t = u / (1.0 - 0.5 * (u * (a1 / x - 1.0)).min(1.0));
x -= t;
if x <= 0.0 {
x = f64::midpoint(x, t);
}
if t.abs() < 1e-15 * x {
return x;
}
}
bisect_quantile(x, residual)
}
pub(super) fn bisect_quantile(x: f64, residual: impl Fn(f64) -> f64) -> f64 {
let (mut lo, mut hi) = (x.max(f64::MIN_POSITIVE), x.max(f64::MIN_POSITIVE));
while residual(lo) > 0.0 {
lo *= 0.5;
if lo == 0.0 {
return 0.0;
}
}
while residual(hi) < 0.0 {
hi *= 2.0;
if hi == f64::INFINITY {
return hi;
}
}
for _ in 0..200 {
let mid = f64::midpoint(lo.ln(), hi.ln()).exp();
if !(lo < mid && mid < hi) {
break;
}
if residual(mid) < 0.0 {
lo = mid;
} else {
hi = mid;
}
}
hi
}