Skip to main content

probl_engine/
weight.rs

1//! Weights: non-negative numbers with a binary mantissa and a 64-bit
2//! exponent, so products of many small probabilities (long chains of
3//! observations) can't underflow to zero.
4
5use std::cmp::Ordering;
6use std::fmt;
7use std::ops::{Add, AddAssign, Mul};
8
9/// `mant × 2^exp`, with `mant` in [0.5, 1), or zero.
10#[derive(Clone, Copy, PartialEq)]
11pub struct Weight {
12    mant: f64,
13    exp: i64,
14}
15
16impl Weight {
17    pub const ZERO: Weight = Weight { mant: 0.0, exp: 0 };
18    pub const ONE: Weight = Weight { mant: 0.5, exp: 1 };
19
20    /// A weight from a finite, non-negative number.
21    pub fn new(x: f64) -> Weight {
22        debug_assert!(x >= 0.0 && x.is_finite(), "invalid weight {x}");
23        if x <= 0.0 || !x.is_finite() {
24            return Weight::ZERO;
25        }
26        let (mant, exp) = frexp(x);
27        Weight { mant, exp }
28    }
29
30    /// e^l: a weight from its natural logarithm, even one far outside the
31    /// range of an `f64`. A logarithm of −∞ (or NaN) gives zero.
32    pub fn from_ln(l: f64) -> Weight {
33        if l.is_nan() || l == f64::NEG_INFINITY {
34            return Weight::ZERO;
35        }
36        debug_assert!(l.is_finite(), "invalid log weight {l}");
37        // l = k ln 2 + r with |r| ≤ ln 2 / 2, so e^l = e^r × 2^k. ln 2 is
38        // split in two, the first part short enough for k ln 2 to be exact.
39        const LN2_HI: f64 = 6.931_471_803_691_238e-1;
40        const LN2_LO: f64 = 1.908_214_929_270_587_7e-10;
41        let k = libm::round(l / std::f64::consts::LN_2);
42        let r = (l - k * LN2_HI) - k * LN2_LO;
43        normalize(crate::math::exp(r), k as i64)
44    }
45
46    pub fn is_zero(self) -> bool {
47        self.mant == 0.0
48    }
49
50    /// The weight as an `f64`; tiny weights round to 0.
51    pub fn to_f64(self) -> f64 {
52        ldexp(self.mant, self.exp)
53    }
54
55    /// Multiply by a non-negative finite factor, such as a probability.
56    pub fn scale(self, factor: f64) -> Weight {
57        if self.is_zero() || factor <= 0.0 || !factor.is_finite() {
58            return Weight::ZERO;
59        }
60        let (m, e) = frexp(factor);
61        normalize(self.mant * m, self.exp + e)
62    }
63
64    /// `self − other`, or zero if `other` is larger.
65    pub fn saturating_sub(self, other: Weight) -> Weight {
66        if other >= self {
67            return Weight::ZERO;
68        }
69        if other.is_zero() {
70            return self;
71        }
72        let diff = self.exp - other.exp;
73        if diff > 1100 {
74            return self;
75        }
76        normalize(self.mant - ldexp(other.mant, -diff), self.exp)
77    }
78
79    /// `self / other` as an `f64` (for probabilities and shares).
80    pub fn ratio(self, other: Weight) -> f64 {
81        if self.is_zero() {
82            return 0.0;
83        }
84        if other.is_zero() {
85            return f64::INFINITY;
86        }
87        ldexp(self.mant / other.mant, self.exp - other.exp)
88    }
89
90    /// Base-10 logarithm, for printing weights too small for an `f64`.
91    pub fn log10(self) -> f64 {
92        if self.is_zero() {
93            return f64::NEG_INFINITY;
94        }
95        libm::log10(self.mant) + self.exp as f64 * std::f64::consts::LOG10_2
96    }
97
98    /// The natural logarithm, even of a weight far outside the range of an
99    /// `f64`: −∞ for zero.
100    pub fn ln(self) -> f64 {
101        if self.is_zero() {
102            return f64::NEG_INFINITY;
103        }
104        libm::log(self.mant) + self.exp as f64 * std::f64::consts::LN_2
105    }
106
107    pub fn sum(weights: impl IntoIterator<Item = Weight>) -> Weight {
108        weights.into_iter().fold(Weight::ZERO, |a, b| a + b)
109    }
110}
111
112impl Default for Weight {
113    fn default() -> Weight {
114        Weight::ZERO
115    }
116}
117
118impl Add for Weight {
119    type Output = Weight;
120
121    fn add(self, other: Weight) -> Weight {
122        if self.is_zero() {
123            return other;
124        }
125        if other.is_zero() {
126            return self;
127        }
128        let (big, small) = if self.exp >= other.exp {
129            (self, other)
130        } else {
131            (other, self)
132        };
133        let diff = big.exp - small.exp;
134        if diff > 1100 {
135            return big;
136        }
137        normalize(big.mant + ldexp(small.mant, -diff), big.exp)
138    }
139}
140
141impl AddAssign for Weight {
142    fn add_assign(&mut self, other: Weight) {
143        *self = *self + other;
144    }
145}
146
147impl Mul for Weight {
148    type Output = Weight;
149
150    fn mul(self, other: Weight) -> Weight {
151        if self.is_zero() || other.is_zero() {
152            return Weight::ZERO;
153        }
154        normalize(self.mant * other.mant, self.exp + other.exp)
155    }
156}
157
158impl PartialOrd for Weight {
159    fn partial_cmp(&self, other: &Weight) -> Option<Ordering> {
160        Some(match (self.is_zero(), other.is_zero()) {
161            (true, true) => Ordering::Equal,
162            (true, false) => Ordering::Less,
163            (false, true) => Ordering::Greater,
164            (false, false) => self.exp.cmp(&other.exp).then(self.mant.total_cmp(&other.mant)),
165        })
166    }
167}
168
169impl fmt::Debug for Weight {
170    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
171        let x = self.to_f64();
172        if x == 0.0 && !self.is_zero() {
173            let l = self.log10();
174            write!(f, "{:.3}e{}", libm::pow(10.0, l - l.floor()), l.floor())
175        } else {
176            write!(f, "{x}")
177        }
178    }
179}
180
181fn normalize(m: f64, e: i64) -> Weight {
182    if m == 0.0 {
183        return Weight::ZERO;
184    }
185    let (mm, ee) = frexp(m);
186    Weight { mant: mm, exp: e + ee }
187}
188
189/// Split a positive finite `x` into a mantissa in [0.5, 1) and an exponent.
190fn frexp(x: f64) -> (f64, i64) {
191    let bits = x.to_bits();
192    let raw_exp = ((bits >> 52) & 0x7ff) as i64;
193    if raw_exp == 0 {
194        // Subnormal: scale into the normal range first.
195        let (m, e) = frexp(x * 2f64.powi(64));
196        return (m, e - 64);
197    }
198    let mant = f64::from_bits((bits & !(0x7ff << 52)) | (1022 << 52));
199    (mant, raw_exp - 1022)
200}
201
202/// `m × 2^e`, saturating to 0 or infinity.
203fn ldexp(m: f64, e: i64) -> f64 {
204    if m == 0.0 {
205        return 0.0;
206    }
207    if e > 2100 {
208        return f64::INFINITY;
209    }
210    if e < -2200 {
211        return 0.0;
212    }
213    // Two steps keep each factor representable.
214    let half = e / 2;
215    m * pow2(half as i32) * pow2((e - half) as i32)
216}
217
218/// 2^k, as `2f64.powi(k)` gives it, without its loop: exact from 2^-1023
219/// to 2^1023, infinite above, and zero below (its intermediate 2^1024
220/// overflows). A `powi` kept for the rare cases would still be called every
221/// time: the compiler computes both sides of a choice between them.
222fn pow2(k: i32) -> f64 {
223    match k {
224        1024.. => f64::INFINITY,
225        -1022..=1023 => f64::from_bits(((k + 1023) as u64) << 52),
226        -1023 => f64::from_bits(1 << 51),
227        _ => 0.0,
228    }
229}
230
231#[cfg(test)]
232mod tests {
233    use super::*;
234
235    #[test]
236    fn powers_of_two_are_exact() {
237        for k in -1100..=1100 {
238            assert_eq!(pow2(k).to_bits(), 2f64.powi(k).to_bits(), "2^{k}");
239        }
240    }
241
242    #[test]
243    fn arithmetic() {
244        let w = Weight::new(0.75);
245        assert_eq!(w.to_f64(), 0.75);
246        assert_eq!(w.scale(0.5).to_f64(), 0.375);
247        assert_eq!((w + Weight::new(0.25)).to_f64(), 1.0);
248        assert_eq!(Weight::ONE.to_f64(), 1.0);
249        assert_eq!(Weight::new(0.3).ratio(Weight::new(0.6)), 0.5);
250        assert_eq!(Weight::new(0.5).saturating_sub(Weight::new(0.25)).to_f64(), 0.25);
251        assert!(Weight::new(0.5).saturating_sub(Weight::new(0.75)).is_zero());
252        assert!(Weight::new(1e-300) < Weight::new(1e-200));
253        assert!(Weight::ZERO < Weight::new(1e-300));
254    }
255
256    #[test]
257    fn no_underflow() {
258        // A thousand observations with probability 10^-5 each.
259        let mut w = Weight::ONE;
260        for _ in 0..1000 {
261            w = w.scale(1e-5);
262        }
263        assert!(!w.is_zero());
264        assert_eq!(w.to_f64(), 0.0);
265        assert!((w.log10() + 5000.0).abs() < 1e-6);
266        // Ratios of tiny weights are still exact enough.
267        let a = w.scale(0.25);
268        assert!((a.ratio(w) - 0.25).abs() < 1e-15);
269        let b = w + w;
270        assert!((b.ratio(w) - 2.0).abs() < 1e-15);
271    }
272
273    #[test]
274    fn weights_from_logarithms() {
275        assert_eq!(Weight::from_ln(0.0), Weight::ONE);
276        assert!((Weight::from_ln(0.75f64.ln()).to_f64() - 0.75).abs() < 1e-15);
277        assert!((Weight::from_ln(3.5).to_f64() - 3.5f64.exp()).abs() < 1e-13);
278        assert!(Weight::from_ln(f64::NEG_INFINITY).is_zero());
279        assert!(Weight::from_ln(f64::NAN).is_zero());
280        // Far below an f64: e^-4234.1 = 10^-1838.8…
281        let w = Weight::from_ln(-4234.102082009147);
282        assert!(!w.is_zero());
283        assert!((w.log10() - -4234.102082009147 / std::f64::consts::LN_10).abs() < 1e-10);
284        // Multiplying adds logarithms.
285        let (a, b) = (Weight::from_ln(-1000.25), Weight::from_ln(-2000.5));
286        assert!(((a * b).ratio(Weight::from_ln(-3000.75)) - 1.0).abs() < 1e-12);
287    }
288
289    #[test]
290    fn subnormal_inputs() {
291        let tiny = f64::MIN_POSITIVE / 1024.0;
292        let w = Weight::new(tiny);
293        assert!((w.to_f64() - tiny).abs() < 1e-320);
294        assert!((Weight::ONE.scale(tiny).ratio(Weight::new(f64::MIN_POSITIVE)) - 1.0 / 1024.0).abs() < 1e-15);
295    }
296}