use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum CardinalityClass {
Bounded(u64),
BoundedAtMost(u64),
Unbounded,
Continuous {
intervals: Vec<Interval>,
measure: ProductMeasure,
},
ContinuousAtMost {
intervals: Vec<Interval>,
measure_at_most: ProductMeasure,
},
Hybrid(Hybrid),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Hybrid {
pub discrete_axes: Vec<u64>,
pub continuous_axes: Vec<Interval>,
pub measure: ProductMeasure,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Interval {
pub lo: f64,
pub hi: f64,
pub lo_open: bool,
pub hi_open: bool,
}
impl Interval {
pub fn closed(lo: f64, hi: f64) -> Self {
Self {
lo,
hi,
lo_open: false,
hi_open: false,
}
}
pub fn half_open(lo: f64, hi: f64) -> Self {
Self {
lo,
hi,
lo_open: false,
hi_open: true,
}
}
pub fn open(lo: f64, hi: f64) -> Self {
Self {
lo,
hi,
lo_open: true,
hi_open: true,
}
}
pub fn is_bounded(&self) -> bool {
self.lo.is_finite() && self.hi.is_finite()
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum ProductMeasure {
Uniform,
Named(MeasureName),
Product(Vec<ProductMeasure>),
}
impl ProductMeasure {
pub fn is_integrable(&self, intervals: &[Interval]) -> bool {
match self {
ProductMeasure::Uniform => intervals.iter().all(Interval::is_bounded),
ProductMeasure::Named(name) => name.is_proper_probability_measure(),
ProductMeasure::Product(children) => {
if children.len() != intervals.len() {
return false;
}
children
.iter()
.zip(intervals.iter())
.all(|(m, i)| m.is_integrable(std::slice::from_ref(i)))
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum MeasureName {
Normal,
Exponential,
Pareto,
Beta,
LogNormal,
Gamma,
Uniform01,
}
impl MeasureName {
pub fn is_proper_probability_measure(self) -> bool {
true
}
pub fn parameter_names(self) -> &'static [&'static str] {
match self {
MeasureName::Normal | MeasureName::LogNormal => &["mean", "stddev"],
MeasureName::Exponential => &["rate"],
MeasureName::Pareto => &["scale", "shape"],
MeasureName::Beta => &["alpha", "beta"],
MeasureName::Gamma => &["shape", "scale"],
MeasureName::Uniform01 => &[],
}
}
pub fn default_params(self) -> &'static [f64] {
match self {
MeasureName::Normal | MeasureName::LogNormal => &[0.0, 1.0],
MeasureName::Exponential => &[1.0],
MeasureName::Pareto | MeasureName::Beta | MeasureName::Gamma => &[1.0, 1.0],
MeasureName::Uniform01 => &[],
}
}
pub fn resolve_params(self, params: &[f64]) -> Result<Vec<f64>, String> {
let names = self.parameter_names();
if params.is_empty() {
return Ok(self.default_params().to_vec());
}
if params.len() != names.len() {
return Err(format!(
"{self:?} takes {} parameter(s) ({}); found {}",
names.len(),
names.join(", "),
params.len()
));
}
if let Some(bad) = params.iter().find(|p| !p.is_finite()) {
return Err(format!("{self:?}: parameter {bad} is not finite"));
}
Ok(params.to_vec())
}
pub fn text(self) -> &'static str {
match self {
MeasureName::Normal => "normal",
MeasureName::Exponential => "exponential",
MeasureName::Pareto => "pareto",
MeasureName::Beta => "beta",
MeasureName::LogNormal => "log_normal",
MeasureName::Gamma => "gamma",
MeasureName::Uniform01 => "uniform01",
}
}
pub fn from_text(name: &str) -> Option<Self> {
[
MeasureName::Normal,
MeasureName::Exponential,
MeasureName::Pareto,
MeasureName::Beta,
MeasureName::LogNormal,
MeasureName::Gamma,
MeasureName::Uniform01,
]
.into_iter()
.find(|m| m.text() == name)
}
pub fn support(self, params: &[f64]) -> Interval {
let p = self
.resolve_params(params)
.unwrap_or_else(|_| self.default_params().to_vec());
let inf = f64::INFINITY;
match self {
MeasureName::Normal => Interval::open(-inf, inf),
MeasureName::Exponential | MeasureName::Gamma => Interval {
lo: 0.0,
hi: inf,
lo_open: false,
hi_open: true,
},
MeasureName::Pareto => Interval {
lo: p[0],
hi: inf,
lo_open: false,
hi_open: true,
},
MeasureName::LogNormal => Interval {
lo: 0.0,
hi: inf,
lo_open: true,
hi_open: true,
},
MeasureName::Beta | MeasureName::Uniform01 => Interval::closed(0.0, 1.0),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bounded_interval_is_bounded() {
assert!(Interval::closed(0.0, 1.0).is_bounded());
assert!(Interval::open(-1.0, 1.0).is_bounded());
}
#[test]
fn unbounded_interval_is_not_bounded() {
let i = Interval {
lo: 0.0,
hi: f64::INFINITY,
lo_open: false,
hi_open: true,
};
assert!(!i.is_bounded());
}
#[test]
fn uniform_integrable_on_bounded_interval() {
let m = ProductMeasure::Uniform;
assert!(m.is_integrable(&[Interval::closed(0.0, 1.0)]));
}
#[test]
fn uniform_not_integrable_on_unbounded_interval() {
let m = ProductMeasure::Uniform;
let unbounded = Interval {
lo: 0.0,
hi: f64::INFINITY,
lo_open: false,
hi_open: true,
};
assert!(!m.is_integrable(&[unbounded]));
}
#[test]
fn named_measure_always_integrable() {
let m = ProductMeasure::Named(MeasureName::Normal);
let unbounded = Interval {
lo: f64::NEG_INFINITY,
hi: f64::INFINITY,
lo_open: true,
hi_open: true,
};
assert!(m.is_integrable(&[unbounded]));
}
#[test]
fn product_measure_requires_matching_arity() {
let m = ProductMeasure::Product(vec![ProductMeasure::Uniform, ProductMeasure::Uniform]);
assert!(m.is_integrable(&[Interval::closed(0.0, 1.0), Interval::closed(0.0, 1.0)]));
assert!(!m.is_integrable(&[Interval::closed(0.0, 1.0)]));
}
#[test]
fn cardinality_class_round_trip_serde() {
let c = CardinalityClass::Continuous {
intervals: vec![Interval::closed(0.0, 1.0)],
measure: ProductMeasure::Uniform,
};
let json = serde_json::to_string(&c).unwrap();
let back: CardinalityClass = serde_json::from_str(&json).unwrap();
assert_eq!(c, back);
}
}