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
}
}
#[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);
}
}