mod probability;
pub(crate) mod propagation;
mod rational;
mod sampling;
mod weight;
use std::{
collections::{BTreeMap, btree_map},
error::Error,
fmt::{self, Display, Formatter},
iter::FusedIterator
};
pub use probability::*;
pub use propagation::{
BuildError, DistributionPlan,
cost::{Class, Cost},
meter::{Budget, Dimension, Progress, Unobserved, Usage}
};
pub use rational::*;
pub use sampling::Sampled;
#[cfg(feature = "serde")]
use serde::{Deserialize, Deserializer, Serialize, de};
pub use weight::*;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(Serialize))]
pub struct Distribution
{
weights: BTreeMap<i32, Weight>,
total: Weight
}
static ZERO: Weight = Weight::ZERO;
impl Distribution
{
pub fn point(outcome: i32) -> Self
{
Self {
weights: BTreeMap::from([(outcome, Weight::ONE)]),
total: Weight::ONE
}
}
pub fn from_weights(
weights: impl IntoIterator<Item = (i32, Weight)>
) -> Result<Self, EmptyDistributionError>
{
let mut merged = BTreeMap::<i32, Weight>::new();
for (outcome, weight) in weights
{
if !weight.is_zero()
{
*merged.entry(outcome).or_default() += weight;
}
}
if merged.is_empty()
{
return Err(EmptyDistributionError)
}
let total = merged.values().sum();
Ok(Self {
weights: merged,
total
})
}
}
impl Distribution
{
pub fn get(&self, outcome: i32) -> &Weight
{
self.weights.get(&outcome).unwrap_or(&ZERO)
}
#[inline]
pub fn total(&self) -> &Weight { &self.total }
#[allow(clippy::len_without_is_empty)]
#[inline]
pub fn len(&self) -> usize { self.weights.len() }
pub fn min(&self) -> i32
{
*self
.weights
.keys()
.next()
.expect("distribution is never empty")
}
pub fn max(&self) -> i32
{
*self
.weights
.keys()
.next_back()
.expect("distribution is never empty")
}
pub fn probability(&self, outcome: i32) -> Probability
{
Probability::new_unchecked(
self.get(outcome).clone(),
self.total.clone()
)
}
pub fn cdf(&self, outcome: i32) -> Probability
{
let cumulative = self.weights.range(..=outcome).map(|(_, w)| w).sum();
Probability::new_unchecked(cumulative, self.total.clone())
}
pub fn quantile(&self, p: &Probability) -> i32
{
let target = p.numerator() * &self.total;
let mut cumulative = Weight::ZERO;
for (&outcome, weight) in &self.weights
{
cumulative += weight;
if &cumulative * p.denominator() >= target
{
return outcome
}
}
unreachable!("the total weight reaches every probability")
}
pub fn mean(&self) -> Rational
{
let mut positive = Weight::ZERO;
let mut negative = Weight::ZERO;
for (&outcome, weight) in &self.weights
{
let magnitude = Weight::from(outcome.unsigned_abs());
if outcome > 0
{
positive += magnitude * weight;
}
else
{
negative += magnitude * weight;
}
}
match positive.checked_sub(&negative)
{
Some(difference) =>
{
Rational::new_unchecked(false, difference, self.total.clone())
},
None =>
{
let difference = negative
.checked_sub(&positive)
.expect("negative exceeds positive");
Rational::new_unchecked(true, difference, self.total.clone())
}
}
}
}
impl Distribution
{
pub fn to_f64(&self) -> BTreeMap<i32, f64>
{
self.weights
.iter()
.map(|(&outcome, weight)| {
(outcome, weight.ratio_to_f64(&self.total))
})
.collect()
}
}
impl Distribution
{
pub fn iter(&self) -> DistributionIter<'_>
{
DistributionIter(self.weights.iter())
}
}
#[derive(Debug, Clone)]
pub struct DistributionIter<'a>(btree_map::Iter<'a, i32, Weight>);
impl<'a> Iterator for DistributionIter<'a>
{
type Item = (i32, &'a Weight);
fn next(&mut self) -> Option<Self::Item>
{
self.0.next().map(|(&outcome, weight)| (outcome, weight))
}
fn size_hint(&self) -> (usize, Option<usize>) { self.0.size_hint() }
}
impl DoubleEndedIterator for DistributionIter<'_>
{
fn next_back(&mut self) -> Option<Self::Item>
{
self.0
.next_back()
.map(|(&outcome, weight)| (outcome, weight))
}
}
impl ExactSizeIterator for DistributionIter<'_> {}
impl FusedIterator for DistributionIter<'_> {}
impl<'a> IntoIterator for &'a Distribution
{
type Item = (i32, &'a Weight);
type IntoIter = DistributionIter<'a>;
fn into_iter(self) -> Self::IntoIter { self.iter() }
}
impl IntoIterator for Distribution
{
type Item = (i32, Weight);
type IntoIter = btree_map::IntoIter<i32, Weight>;
fn into_iter(self) -> Self::IntoIter { self.weights.into_iter() }
}
impl Display for Distribution
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result
{
for (outcome, weight) in self
{
writeln!(f, "{}: {}", outcome, weight)?;
}
Ok(())
}
}
#[cfg(feature = "serde")]
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct UncheckedDistribution
{
weights: BTreeMap<i32, Weight>,
total: Weight
}
#[cfg(feature = "serde")]
impl<'de> Deserialize<'de> for Distribution
{
fn deserialize<D: Deserializer<'de>>(
deserializer: D
) -> Result<Self, D::Error>
{
use de::Error as _;
let UncheckedDistribution { weights, total } =
UncheckedDistribution::deserialize(deserializer)?;
if weights.is_empty()
{
return Err(D::Error::custom("distribution has no outcomes"))
}
if let Some((outcome, _)) = weights.iter().find(|(_, w)| w.is_zero())
{
return Err(D::Error::custom(format_args!(
"distribution weighs outcome {outcome} as zero"
)))
}
let sum: Weight = weights.values().sum();
if sum != total
{
return Err(D::Error::custom(format_args!(
"distribution total {total} is not the sum of its weights, \
{sum}"
)))
}
Ok(Self { weights, total })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct EmptyDistributionError;
impl Display for EmptyDistributionError
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result
{
write!(f, "distribution has no outcome with nonzero weight")
}
}
impl Error for EmptyDistributionError {}