1use std::cmp::Ordering;
6use std::fmt;
7use std::ops::{Add, AddAssign, Mul};
8
9#[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 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 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 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 pub fn to_f64(self) -> f64 {
52 ldexp(self.mant, self.exp)
53 }
54
55 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 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 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 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 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
189fn 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 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
202fn 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 let half = e / 2;
215 m * pow2(half as i32) * pow2((e - half) as i32)
216}
217
218fn 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 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 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 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 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}