use super::traits::*;
use core::fmt;
use core::ops::{Add, Mul};
use num_traits::{One, Zero};
use ordered_float::OrderedFloat;
#[derive(Clone, Copy, Debug, PartialEq, PartialOrd, Eq, Ord, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct LogWeight(OrderedFloat<f64>);
impl LogWeight {
pub const INFINITY: Self = Self(OrderedFloat(f64::INFINITY));
pub fn new(value: f64) -> Self {
Self(OrderedFloat(value))
}
pub fn from_probability(p: f64) -> Self {
if p == 0.0 {
Self::INFINITY
} else {
Self::new(-p.ln())
}
}
pub fn to_probability(&self) -> f64 {
if self.0.is_infinite() {
0.0
} else {
(-*self.0).exp()
}
}
}
impl fmt::Display for LogWeight {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.0.is_infinite() {
write!(f, "∞")
} else {
let value = self.0;
write!(f, "{value}")
}
}
}
impl Zero for LogWeight {
fn zero() -> Self {
Self::INFINITY
}
fn is_zero(&self) -> bool {
self.0.is_infinite()
}
}
impl One for LogWeight {
fn one() -> Self {
Self::new(0.0)
}
}
impl Add for LogWeight {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
if <Self as num_traits::Zero>::is_zero(&self) {
rhs
} else if <Self as num_traits::Zero>::is_zero(&rhs) {
self
} else {
let a = -*self.0;
let b = -*rhs.0;
Self::new(-(a.max(b) + (1.0 + (-(a - b).abs()).exp()).ln()))
}
}
}
impl Mul for LogWeight {
type Output = Self;
fn mul(self, rhs: Self) -> Self::Output {
if <Self as num_traits::Zero>::is_zero(&self) || <Self as num_traits::Zero>::is_zero(&rhs) {
Self::zero()
} else {
Self(self.0 + rhs.0)
}
}
}
impl Semiring for LogWeight {
type Value = f64;
fn new(value: Self::Value) -> Self {
Self::new(value)
}
fn value(&self) -> &Self::Value {
&self.0
}
fn properties() -> SemiringProperties {
SemiringProperties {
left_semiring: true,
right_semiring: true,
commutative: true,
idempotent: false,
path: false,
}
}
fn approx_eq(&self, other: &Self, epsilon: f64) -> bool {
if <Self as num_traits::Zero>::is_zero(self) && <Self as num_traits::Zero>::is_zero(other) {
true
} else {
(self.0 - other.0).abs() < epsilon
}
}
}
impl DivisibleSemiring for LogWeight {
fn divide(&self, other: &Self) -> Option<Self> {
if <Self as num_traits::Zero>::is_zero(other) {
None
} else if <Self as num_traits::Zero>::is_zero(self) {
Some(Self::zero())
} else {
Some(Self(self.0 - other.0))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use num_traits::{One, Zero};
#[test]
fn test_log_weight_creation() {
let w = LogWeight::new(2.0);
assert_eq!(*w.value(), 2.0);
}
#[test]
fn test_log_zero_one() {
let zero = LogWeight::zero();
let one = LogWeight::one();
assert!(Semiring::is_zero(&zero));
assert!(Semiring::is_one(&one));
assert!(zero.value().is_infinite());
assert_eq!(*one.value(), 0.0);
}
#[test]
fn test_log_addition() {
let w1 = LogWeight::new(1.0);
let w2 = LogWeight::new(2.0);
let result = w1.plus(&w2);
assert!(result.approx_eq(&LogWeight::new(0.6867), 0.001));
}
#[test]
fn test_log_multiplication() {
let w1 = LogWeight::new(1.0);
let w2 = LogWeight::new(2.0);
let result = w1.times(&w2);
assert_eq!(*result.value(), 3.0); }
#[test]
fn test_log_zero_operations() {
let w = LogWeight::new(2.0);
let zero = LogWeight::zero();
assert_eq!(w.plus(&zero), w);
assert_eq!(zero.plus(&w), w);
assert!(Semiring::is_zero(&w.times(&zero)));
assert!(Semiring::is_zero(&zero.times(&w)));
}
#[test]
fn test_log_one_operations() {
let w = LogWeight::new(2.0);
let one = LogWeight::one();
let mul_result = w.times(&one);
assert_eq!(mul_result, w);
}
#[test]
fn test_log_display() {
let w = LogWeight::new(2.5);
let zero = LogWeight::zero();
assert_eq!(format!("{w}"), "2.5");
assert_eq!(format!("{zero}"), "∞");
}
#[test]
fn test_log_division() {
let w1 = LogWeight::new(5.0);
let w2 = LogWeight::new(3.0);
let result = w1.divide(&w2).unwrap();
assert_eq!(*result.value(), 2.0);
let zero = LogWeight::zero();
assert!(w1.divide(&zero).is_none());
}
#[test]
fn test_log_from_to_probability() {
let prob = 0.5;
let log_weight = LogWeight::from_probability(prob);
assert!((log_weight.value() - (-prob.ln())).abs() < 1e-10);
assert!((log_weight.to_probability() - prob).abs() < 1e-10);
let zero_log = LogWeight::from_probability(0.0);
assert!(Semiring::is_zero(&zero_log));
assert_eq!(zero_log.to_probability(), 0.0);
let small_prob = 1e-100;
let small_log = LogWeight::from_probability(small_prob);
assert!((small_log.to_probability() - small_prob).abs() < small_prob * 1e-10);
}
#[test]
fn test_log_properties() {
let props = LogWeight::properties();
assert!(props.left_semiring);
assert!(props.right_semiring);
assert!(props.commutative);
assert!(!props.idempotent);
assert!(!props.path);
}
#[test]
fn test_log_approx_eq() {
let w1 = LogWeight::new(2.000_001);
let w2 = LogWeight::new(2.0);
assert!(w1.approx_eq(&w2, 0.001));
assert!(!w1.approx_eq(&w2, 0.000_000_1));
}
#[test]
fn test_log_operator_overloads() {
let w1 = LogWeight::new(1.0);
let w2 = LogWeight::new(2.0);
let sum = w1 + w2;
assert!(sum.approx_eq(&LogWeight::new(0.6867), 0.001));
assert_eq!(w1 * w2, LogWeight::new(3.0));
}
#[test]
fn test_log_identity_laws() {
let w = LogWeight::new(2.0);
let zero = LogWeight::zero();
let one = LogWeight::one();
assert_eq!(w + zero, w);
assert_eq!(zero + w, w);
assert_eq!(w * one, w);
assert_eq!(one * w, w);
assert!(Semiring::is_zero(&(w * zero)));
assert!(Semiring::is_zero(&(zero * w)));
}
#[test]
fn test_log_semiring_axioms() {
let a = LogWeight::new(1.0);
let b = LogWeight::new(2.0);
let c = LogWeight::new(3.0);
let tolerance = 1e-10;
assert!(((a + b) + c).approx_eq(&(a + (b + c)), tolerance));
assert_eq!((a * b) * c, a * (b * c));
assert_eq!(a + b, b + a);
assert_eq!(a * b, b * a);
assert!(((a + b) * c).approx_eq(&((a * c) + (b * c)), tolerance));
}
}