use crate::error::ProbError;
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct BetaHyperparameters {
pub alpha: f64,
pub beta: f64,
}
impl BetaHyperparameters {
pub fn from_moments(mean: f64, variance: f64) -> Result<Self, ProbError> {
if !(mean.is_finite() && mean > 0.0 && mean < 1.0) {
return Err(ProbError::InvalidPrior {
message: "beta_from_moments: mean must be finite and strictly inside (0, 1)",
});
}
if !variance.is_finite() || !(variance > 0.0) {
return Err(ProbError::InvalidPrior {
message: "beta_from_moments: variance must be finite and > 0",
});
}
let kappa = mean * (1.0 - mean) / variance - 1.0;
if !(kappa > 0.0) {
return Err(ProbError::InvalidPrior {
message: "beta_from_moments: variance must be < mean * (1 - mean) for a Beta \
distribution to have these moments",
});
}
let alpha = mean * kappa;
let beta = (1.0 - mean) * kappa;
if !alpha.is_finite() || !beta.is_finite() || !(alpha > 0.0) || !(beta > 0.0) {
return Err(ProbError::Numerical {
message: "beta_from_moments: alpha/beta non-finite or non-positive".into(),
});
}
Ok(Self { alpha, beta })
}
pub fn from_mean_and_ess(mean: f64, ess: f64) -> Result<Self, ProbError> {
if !(mean.is_finite() && mean > 0.0 && mean < 1.0) {
return Err(ProbError::InvalidPrior {
message: "beta_from_mean_and_ess: mean must be finite and strictly inside (0, 1)",
});
}
if !ess.is_finite() || ess < 0.0 {
return Err(ProbError::InvalidPrior {
message: "beta_from_mean_and_ess: ess must be finite and >= 0",
});
}
let total = ess + 2.0;
let alpha = mean * total;
let beta = (1.0 - mean) * total;
if !alpha.is_finite() || !beta.is_finite() || !(alpha > 0.0) || !(beta > 0.0) {
return Err(ProbError::Numerical {
message: "beta_from_mean_and_ess: alpha/beta non-finite or non-positive".into(),
});
}
Ok(Self { alpha, beta })
}
#[must_use]
pub fn mean(&self) -> f64 {
self.alpha / (self.alpha + self.beta)
}
#[must_use]
pub fn variance(&self) -> f64 {
let total = self.alpha + self.beta;
(self.alpha * self.beta) / (total * total * (total + 1.0))
}
#[must_use]
pub fn ess(&self) -> f64 {
self.alpha + self.beta - 2.0
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct GammaHyperparameters {
pub shape: f64,
pub rate: f64,
}
impl GammaHyperparameters {
pub fn from_moments(mean: f64, variance: f64) -> Result<Self, ProbError> {
if !mean.is_finite() || !(mean > 0.0) {
return Err(ProbError::InvalidPrior {
message: "gamma_from_moments: mean must be finite and > 0",
});
}
if !variance.is_finite() || !(variance > 0.0) {
return Err(ProbError::InvalidPrior {
message: "gamma_from_moments: variance must be finite and > 0",
});
}
let shape = mean * mean / variance;
let rate = mean / variance;
if !shape.is_finite() || !rate.is_finite() || !(shape > 0.0) || !(rate > 0.0) {
return Err(ProbError::Numerical {
message: "gamma_from_moments: shape/rate non-finite or non-positive".into(),
});
}
Ok(Self { shape, rate })
}
pub fn from_mean_and_ess(mean: f64, ess: f64) -> Result<Self, ProbError> {
if !mean.is_finite() || !(mean > 0.0) {
return Err(ProbError::InvalidPrior {
message: "gamma_from_mean_and_ess: mean must be finite and > 0",
});
}
if !ess.is_finite() || ess < 0.0 {
return Err(ProbError::InvalidPrior {
message: "gamma_from_mean_and_ess: ess must be finite and >= 0",
});
}
let shape = ess + 1.0;
let rate = shape / mean;
if !shape.is_finite() || !rate.is_finite() || !(rate > 0.0) {
return Err(ProbError::Numerical {
message: "gamma_from_mean_and_ess: shape/rate non-finite or non-positive".into(),
});
}
Ok(Self { shape, rate })
}
#[must_use]
pub fn mean(&self) -> f64 {
self.shape / self.rate
}
#[must_use]
pub fn variance(&self) -> f64 {
self.shape / (self.rate * self.rate)
}
#[must_use]
pub fn ess(&self) -> f64 {
self.shape - 1.0
}
}
#[cfg(test)]
mod tests {
use super::*;
const TOL: f64 = 1e-9;
#[test]
fn beta_from_moments_round_trips_input_moments() {
let h = BetaHyperparameters::from_moments(0.3, 0.02).unwrap();
assert!((h.alpha - 2.85).abs() < TOL);
assert!((h.beta - 6.65).abs() < TOL);
assert!((h.mean() - 0.3).abs() < TOL);
assert!((h.variance() - 0.02).abs() < TOL);
assert!((h.ess() - 7.5).abs() < TOL);
}
#[test]
fn beta_from_moments_can_report_negative_ess() {
let h = BetaHyperparameters::from_moments(0.5, 0.24).unwrap();
assert!(h.alpha > 0.0 && h.beta > 0.0);
assert!(h.ess() < 0.0);
assert!((h.mean() - 0.5).abs() < TOL);
assert!((h.variance() - 0.24).abs() < 1e-6);
}
#[test]
fn beta_from_moments_rejects_variance_at_or_above_support_bound() {
assert!(BetaHyperparameters::from_moments(0.5, 0.25).is_err());
assert!(BetaHyperparameters::from_moments(0.5, 0.3).is_err());
}
#[test]
fn beta_from_moments_rejects_mean_outside_open_interval() {
assert!(BetaHyperparameters::from_moments(0.0, 0.01).is_err());
assert!(BetaHyperparameters::from_moments(1.0, 0.01).is_err());
assert!(BetaHyperparameters::from_moments(-0.1, 0.01).is_err());
assert!(BetaHyperparameters::from_moments(1.1, 0.01).is_err());
}
#[test]
fn beta_from_moments_rejects_nonfinite_inputs() {
assert!(BetaHyperparameters::from_moments(f64::NAN, 0.02).is_err());
assert!(BetaHyperparameters::from_moments(0.3, f64::NAN).is_err());
assert!(BetaHyperparameters::from_moments(0.3, f64::INFINITY).is_err());
assert!(BetaHyperparameters::from_moments(f64::INFINITY, 0.02).is_err());
}
#[test]
fn beta_from_moments_rejects_nonpositive_variance() {
assert!(BetaHyperparameters::from_moments(0.3, 0.0).is_err());
assert!(BetaHyperparameters::from_moments(0.3, -0.01).is_err());
}
#[test]
fn beta_from_mean_and_ess_zero_is_beta_1_1_strength_at_requested_mean() {
let h = BetaHyperparameters::from_mean_and_ess(0.3, 0.0).unwrap();
assert!((h.alpha - 0.6).abs() < TOL);
assert!((h.beta - 1.4).abs() < TOL);
assert!((h.mean() - 0.3).abs() < TOL);
assert!(h.ess().abs() < TOL);
assert!(h.alpha > 0.0 && h.beta > 0.0);
}
#[test]
fn beta_from_mean_and_ess_zero_at_mean_half_is_exactly_beta_1_1() {
let h = BetaHyperparameters::from_mean_and_ess(0.5, 0.0).unwrap();
assert!((h.alpha - 1.0).abs() < TOL);
assert!((h.beta - 1.0).abs() < TOL);
}
#[test]
fn beta_from_mean_and_ess_matches_any_nonnegative_request() {
let h = BetaHyperparameters::from_mean_and_ess(0.5, 10.0).unwrap();
assert!((h.alpha - 6.0).abs() < TOL);
assert!((h.beta - 6.0).abs() < TOL);
assert!((h.mean() - 0.5).abs() < TOL);
assert!((h.ess() - 10.0).abs() < TOL);
}
#[test]
fn beta_from_mean_and_ess_rejects_mean_outside_open_interval() {
assert!(BetaHyperparameters::from_mean_and_ess(0.0, 1.0).is_err());
assert!(BetaHyperparameters::from_mean_and_ess(1.0, 1.0).is_err());
}
#[test]
fn beta_from_mean_and_ess_rejects_negative_ess() {
assert!(BetaHyperparameters::from_mean_and_ess(0.3, -1.0).is_err());
}
#[test]
fn beta_from_mean_and_ess_rejects_nonfinite_inputs() {
assert!(BetaHyperparameters::from_mean_and_ess(f64::NAN, 1.0).is_err());
assert!(BetaHyperparameters::from_mean_and_ess(0.3, f64::NAN).is_err());
assert!(BetaHyperparameters::from_mean_and_ess(0.3, f64::INFINITY).is_err());
}
#[test]
fn gamma_from_moments_round_trips_input_moments() {
let h = GammaHyperparameters::from_moments(4.0, 2.0).unwrap();
assert!((h.shape - 8.0).abs() < TOL);
assert!((h.rate - 2.0).abs() < TOL);
assert!((h.mean() - 4.0).abs() < TOL);
assert!((h.variance() - 2.0).abs() < TOL);
assert!((h.ess() - 7.0).abs() < TOL);
}
#[test]
fn gamma_from_moments_can_report_negative_ess() {
let h = GammaHyperparameters::from_moments(4.0, 32.0).unwrap();
assert!(h.shape > 0.0 && h.rate > 0.0);
assert!(h.ess() < 0.0);
assert!((h.mean() - 4.0).abs() < TOL);
assert!((h.variance() - 32.0).abs() < 1e-6);
}
#[test]
fn gamma_from_moments_rejects_nonpositive_mean() {
assert!(GammaHyperparameters::from_moments(0.0, 1.0).is_err());
assert!(GammaHyperparameters::from_moments(-1.0, 1.0).is_err());
}
#[test]
fn gamma_from_moments_rejects_nonpositive_variance() {
assert!(GammaHyperparameters::from_moments(4.0, 0.0).is_err());
assert!(GammaHyperparameters::from_moments(4.0, -1.0).is_err());
}
#[test]
fn gamma_from_moments_rejects_nonfinite_inputs() {
assert!(GammaHyperparameters::from_moments(f64::NAN, 2.0).is_err());
assert!(GammaHyperparameters::from_moments(4.0, f64::NAN).is_err());
assert!(GammaHyperparameters::from_moments(4.0, f64::INFINITY).is_err());
}
#[test]
fn gamma_from_mean_and_ess_zero_is_reference_exponential_at_requested_mean() {
let h = GammaHyperparameters::from_mean_and_ess(4.0, 0.0).unwrap();
assert!((h.shape - 1.0).abs() < TOL);
assert!((h.rate - 0.25).abs() < TOL);
assert!((h.mean() - 4.0).abs() < TOL);
assert!(h.ess().abs() < TOL);
assert!(h.shape > 0.0 && h.rate > 0.0);
}
#[test]
fn gamma_from_mean_and_ess_matches_any_nonnegative_request() {
let h = GammaHyperparameters::from_mean_and_ess(4.0, 7.0).unwrap();
assert!((h.shape - 8.0).abs() < TOL);
assert!((h.rate - 2.0).abs() < TOL);
assert!((h.mean() - 4.0).abs() < TOL);
assert!((h.ess() - 7.0).abs() < TOL);
}
#[test]
fn gamma_from_mean_and_ess_rejects_nonpositive_mean() {
assert!(GammaHyperparameters::from_mean_and_ess(0.0, 1.0).is_err());
assert!(GammaHyperparameters::from_mean_and_ess(-1.0, 1.0).is_err());
}
#[test]
fn gamma_from_mean_and_ess_rejects_negative_ess() {
assert!(GammaHyperparameters::from_mean_and_ess(4.0, -1.0).is_err());
}
#[test]
fn gamma_from_mean_and_ess_rejects_nonfinite_inputs() {
assert!(GammaHyperparameters::from_mean_and_ess(f64::NAN, 2.0).is_err());
assert!(GammaHyperparameters::from_mean_and_ess(4.0, f64::NAN).is_err());
assert!(GammaHyperparameters::from_mean_and_ess(4.0, f64::INFINITY).is_err());
}
}