Skip to main content

probl_engine/
complex.rs

1//! Finite complex scalars. These are data, never probability weights.
2//!
3//! Keeping the numerical kernel independent of Value and the interpreter lets
4//! a future quantum register use dense arrays of the same scalars.
5
6use crate::error::{OpError, OpResult};
7use std::f64::consts::{FRAC_PI_2, LN_2, LN_10, SQRT_2};
8
9#[derive(Clone, Copy, Debug, PartialEq)]
10pub struct Complex {
11    re: f64,
12    im: f64,
13}
14
15impl Complex {
16    pub fn new(re: f64, im: f64) -> OpResult<Self> {
17        if !re.is_finite() || !im.is_finite() {
18            return Err(OpError::new("complex components and results must be finite"));
19        }
20        // Equality and world merging identify signed zeros. Branch cuts use
21        // the side selected by positive zero in each input component.
22        Ok(Self {
23            re: if re == 0.0 { 0.0 } else { re },
24            im: if im == 0.0 { 0.0 } else { im },
25        })
26    }
27
28    pub fn re(self) -> f64 {
29        self.re
30    }
31
32    pub fn im(self) -> f64 {
33        self.im
34    }
35
36    pub fn conjugate(self) -> Self {
37        Self {
38            im: if self.im == 0.0 { 0.0 } else { -self.im },
39            ..self
40        }
41    }
42
43    pub fn negated(self) -> Self {
44        Self {
45            re: if self.re == 0.0 { 0.0 } else { -self.re },
46            im: if self.im == 0.0 { 0.0 } else { -self.im },
47        }
48    }
49
50    pub fn plus(self, other: Self) -> OpResult<Self> {
51        Self::new(self.re + other.re, self.im + other.im)
52    }
53
54    pub fn minus(self, other: Self) -> OpResult<Self> {
55        Self::new(self.re - other.re, self.im - other.im)
56    }
57
58    pub fn times(self, other: Self) -> OpResult<Self> {
59        let (re, er) = products(self.re, other.re, -self.im, other.im);
60        let (im, ei) = products(self.re, other.im, self.im, other.re);
61        Self::new(libm::scalbn(re, er), libm::scalbn(im, ei))
62    }
63
64    pub fn divided_by(self, other: Self) -> OpResult<Self> {
65        if other.re == 0.0 && other.im == 0.0 {
66            return Err(OpError::new("division by zero"));
67        }
68        let (den, ed) = products(other.re, other.re, other.im, other.im);
69        let (re, er) = products(self.re, other.re, self.im, other.im);
70        let (im, ei) = products(self.im, other.re, -self.re, other.im);
71        Self::new(libm::scalbn(re / den, er - ed), libm::scalbn(im / den, ei - ed))
72    }
73
74    pub fn powi(self, exponent: i64) -> OpResult<Self> {
75        let one = Self { re: 1.0, im: 0.0 };
76        let mut base = if exponent < 0 { one.divided_by(self)? } else { self };
77        let mut n = exponent.unsigned_abs();
78        let mut result = one;
79        // At most 64 iterations, even for i64::MIN.
80        while n != 0 {
81            if n & 1 != 0 {
82                result = result.times(base)?;
83            }
84            n >>= 1;
85            if n != 0 {
86                base = base.times(base)?;
87            }
88        }
89        Ok(result)
90    }
91
92    pub fn pow_integer(self, exponent: &probl_number::Integer) -> OpResult<Self> {
93        if let Some(n) = exponent.to_i64() {
94            return self.powi(n);
95        }
96        let one = Self { re: 1.0, im: 0.0 };
97        let mut base = if exponent.is_negative() {
98            one.divided_by(self)?
99        } else {
100            self
101        };
102        let mut result = one;
103        for i in 0..exponent.bits() {
104            if exponent.magnitude_bit(i) {
105                result = result.times(base)?;
106            }
107            if i + 1 < exponent.bits() {
108                base = base.times(base)?;
109            }
110        }
111        Ok(result)
112    }
113
114    pub fn abs(self) -> f64 {
115        libm::hypot(self.re, self.im)
116    }
117
118    pub fn abs2(self) -> f64 {
119        let (m, e) = products(self.re, self.re, self.im, self.im);
120        libm::scalbn(m, e)
121    }
122
123    pub fn arg(self) -> f64 {
124        libm::atan2(self.im, self.re)
125    }
126
127    /// Principal square root: nonnegative real part, and +i on the negative
128    /// real axis. Scale before taking the norm so even MAX + MAX*i works.
129    pub fn sqrt(self) -> OpResult<Self> {
130        let m = self.re.abs().max(self.im.abs());
131        if m == 0.0 {
132            return Ok(self);
133        }
134        let x = self.re / m;
135        let y = self.im / m;
136        let t = libm::sqrt(m) * libm::sqrt((libm::hypot(x, y) + x.abs()) / 2.0);
137        if self.re >= 0.0 {
138            Self::new(t, self.im / (2.0 * t))
139        } else {
140            Self::new(self.im.abs() / (2.0 * t), t.copysign(self.im))
141        }
142    }
143
144    /// Principal cube root; unlike the real cbrt, the negative axis has phase pi/3.
145    pub fn cbrt(self) -> OpResult<Self> {
146        if self.im == 0.0 && self.re >= 0.0 {
147            return Self::new(libm::cbrt(self.re), 0.0);
148        }
149        let r = crate::math::exp(log_hypot(self.re, self.im) / 3.0);
150        let theta = self.arg() / 3.0;
151        Self::new(r * libm::cos(theta), r * libm::sin(theta))
152    }
153
154    pub fn ln(self) -> OpResult<Self> {
155        if self.re == 0.0 && self.im == 0.0 {
156            return Err(OpError::new("`ln` isn't defined for complex zero"));
157        }
158        Self::new(log_hypot(self.re, self.im), self.arg())
159    }
160
161    pub fn log2(self) -> OpResult<Self> {
162        let z = self.ln()?;
163        Self::new(z.re / LN_2, z.im / LN_2)
164    }
165
166    pub fn log10(self) -> OpResult<Self> {
167        let z = self.ln()?;
168        Self::new(z.re / LN_10, z.im / LN_10)
169    }
170
171    pub fn log1p(self) -> OpResult<Self> {
172        let x = self.re;
173        let y = self.im;
174        if x.abs() < 0.5 && y.abs() < 0.5 {
175            // Keep the low bits that forming 1 + z would discard.
176            Self::new(0.5 * libm::log1p(x * (2.0 + x) + y * y), libm::atan2(y, 1.0 + x))
177        } else {
178            Self::new(1.0 + x, y)?.ln()
179        }
180    }
181
182    pub fn exp(self) -> OpResult<Self> {
183        Self::new(
184            exp_times(self.re, libm::cos(self.im)),
185            exp_times(self.re, libm::sin(self.im)),
186        )
187    }
188
189    pub fn exp2(self) -> OpResult<Self> {
190        let theta = self.im * LN_2;
191        let component = |factor| {
192            if factor == 0.0 {
193                0.0
194            } else if self.re > 1000.0 {
195                (libm::exp2(1000.0) * factor) * libm::exp2(self.re - 1000.0)
196            } else {
197                libm::exp2(self.re) * factor
198            }
199        };
200        Self::new(component(libm::cos(theta)), component(libm::sin(theta)))
201    }
202
203    pub fn expm1(self) -> OpResult<Self> {
204        if self.re.abs() < 0.5 && self.im.abs() < 0.5 {
205            let s = libm::sin(self.im / 2.0);
206            Self::new(
207                libm::expm1(self.re) * libm::cos(self.im) - 2.0 * s * s,
208                crate::math::exp(self.re) * libm::sin(self.im),
209            )
210        } else {
211            let z = self.exp()?;
212            Self::new(z.re - 1.0, z.im)
213        }
214    }
215
216    pub fn sin(self) -> OpResult<Self> {
217        Self::new(
218            cosh_times(self.im, libm::sin(self.re)),
219            sinh_times(self.im, libm::cos(self.re)),
220        )
221    }
222
223    pub fn cos(self) -> OpResult<Self> {
224        Self::new(
225            cosh_times(self.im, libm::cos(self.re)),
226            -sinh_times(self.im, libm::sin(self.re)),
227        )
228    }
229
230    pub fn tan(self) -> OpResult<Self> {
231        let s = libm::sin(self.re);
232        let c = libm::cos(self.re);
233        if self.im.abs() > 20.0 {
234            // Divide the double-angle formula by exp(2*|y|). Taking sin/cos
235            // of x first also avoids overflow when forming the angle 2*x.
236            let t = crate::math::exp(-2.0 * self.im.abs());
237            let den = 1.0 + 2.0 * (c * c - s * s) * t + t * t;
238            Self::new(4.0 * s * c * t / den, ((1.0 - t * t) / den).copysign(self.im))
239        } else {
240            let sh = libm::sinh(self.im);
241            let den = c * c + sh * sh;
242            Self::new(s * c / den, sh * libm::cosh(self.im) / den)
243        }
244    }
245
246    pub fn sinh(self) -> OpResult<Self> {
247        Self::new(
248            sinh_times(self.re, libm::cos(self.im)),
249            cosh_times(self.re, libm::sin(self.im)),
250        )
251    }
252
253    pub fn cosh(self) -> OpResult<Self> {
254        Self::new(
255            cosh_times(self.re, libm::cos(self.im)),
256            sinh_times(self.re, libm::sin(self.im)),
257        )
258    }
259
260    pub fn tanh(self) -> OpResult<Self> {
261        let z = Self {
262            re: self.im,
263            im: self.re,
264        }
265        .tan()?;
266        Self::new(z.im, z.re)
267    }
268
269    pub fn asin(self) -> OpResult<Self> {
270        self.inverse_sin_cos(false)
271    }
272
273    pub fn acos(self) -> OpResult<Self> {
274        self.inverse_sin_cos(true)
275    }
276
277    fn inverse_sin_cos(self, cosine: bool) -> OpResult<Self> {
278        let x = self.re.abs();
279        let y = self.im.abs();
280        let (d, im) = if x.max(y) > 1e150 {
281            // The omitted terms are O(1/|z|^2); avoid squaring huge values.
282            (y, log_hypot(x, y) + LN_2)
283        } else {
284            let r = libm::hypot(x + 1.0, y);
285            let s = libm::hypot(x - 1.0, y);
286            let a = r / 2.0 + s / 2.0;
287            // a = max(x, 1) + correction. Rationalize the hypot differences
288            // and compute sqrt(correction) directly, preserving tiny y even
289            // when y*y underflows or a rounds to 1 (including near z = +/-1).
290            let correction_root = if y == 0.0 {
291                0.0
292            } else {
293                libm::hypot(y / libm::sqrt(r + x + 1.0), y / libm::sqrt(s + (x - 1.0).abs())) / SQRT_2
294            };
295            let amx_root = libm::hypot(libm::sqrt((1.0 - x).max(0.0)), correction_root);
296            let am1_root = libm::hypot(libm::sqrt((x - 1.0).max(0.0)), correction_root);
297            (libm::sqrt(a + x) * amx_root, 2.0 * libm::asinh(am1_root / SQRT_2))
298        };
299        if cosine {
300            Self::new(libm::atan2(d, self.re), -im.copysign(self.im))
301        } else {
302            Self::new(libm::atan2(self.re, d), im.copysign(self.im))
303        }
304    }
305
306    pub fn asinh(self) -> OpResult<Self> {
307        // Keep signs during internal rotations; only public results canonicalize
308        // zero. In particular, do not switch the side of an inverse's cut.
309        let z = Self {
310            re: -self.im,
311            im: self.re,
312        }
313        .asin()?;
314        Self::new(z.im, -z.re)
315    }
316
317    pub fn acosh(self) -> OpResult<Self> {
318        let z = self.acos()?;
319        Self::new(z.im.abs(), z.re.copysign(self.im))
320    }
321
322    pub fn atanh(self) -> OpResult<Self> {
323        let x = self.re.abs();
324        let y = self.im;
325        if x == 1.0 && y == 0.0 {
326            return Err(OpError::new("`atanh` isn't defined at complex +1 or -1"));
327        }
328        let m = x.max(y.abs());
329        if m > 1e150 {
330            let rx = x / m;
331            let ry = y / m;
332            return Self::new(
333                ((rx / (rx * rx + ry * ry)) / m).copysign(self.re),
334                FRAC_PI_2.copysign(y),
335            );
336        }
337        let h = libm::hypot(1.0 - x, y);
338        let q = (4.0 * x / h) / h;
339        let re = if q.is_finite() {
340            0.25 * libm::log1p(q)
341        } else {
342            0.5 * (log_hypot(1.0 + x, y) - log_hypot(1.0 - x, y))
343        };
344        Self::new(
345            re.copysign(self.re),
346            0.5 * libm::atan2(2.0 * y, (1.0 - x) * (1.0 + x) - y * y),
347        )
348    }
349
350    pub fn atan(self) -> OpResult<Self> {
351        let z = Self {
352            re: -self.im,
353            im: self.re,
354        }
355        .atanh()?;
356        Self::new(z.im, -z.re)
357    }
358}
359
360/// ln(hypot(x,y)) without overflowing the norm or losing tiny offsets from 1.
361fn log_hypot(x: f64, y: f64) -> f64 {
362    let m = x.abs().max(y.abs());
363    let n = x.abs().min(y.abs());
364    if (0.5..1.5).contains(&m) {
365        0.5 * libm::log1p((m - 1.0) * (m + 1.0) + n * n)
366    } else if m == 0.0 {
367        f64::NEG_INFINITY
368    } else {
369        let r = n / m;
370        libm::log(m) + 0.5 * libm::log1p(r * r)
371    }
372}
373
374/// exp(x)*factor, with |factor| <= 1. The modulus can overflow even when
375/// both final components fit, e.g. exp(710 + pi/4*i).
376fn exp_times(x: f64, factor: f64) -> f64 {
377    if factor == 0.0 {
378        0.0
379    } else if x > 700.0 {
380        (crate::math::exp(700.0) * factor) * crate::math::exp(x - 700.0)
381    } else {
382        crate::math::exp(x) * factor
383    }
384}
385
386fn cosh_times(x: f64, factor: f64) -> f64 {
387    if x.abs() > 20.0 {
388        exp_times(x.abs() - LN_2, factor)
389    } else {
390        libm::cosh(x) * factor
391    }
392}
393
394fn sinh_times(x: f64, factor: f64) -> f64 {
395    if x.abs() > 20.0 {
396        exp_times(x.abs() - LN_2, factor) * x.signum()
397    } else {
398        libm::sinh(x) * factor
399    }
400}
401
402/// a*b + c*d as mantissa * 2^exponent. Separate exponents avoid overflowing
403/// or underflowing intermediate products, notably the denominator of division.
404/// Scale each product separately so disparate components (1e308 + 1e-308 i)
405/// survive multiplication by one.
406fn products(a: f64, b: f64, c: f64, d: f64) -> (f64, i32) {
407    let product = |x, y| {
408        let (x, ex) = libm::frexp(x);
409        let (y, ey) = libm::frexp(y);
410        (x * y, ex + ey)
411    };
412    let (x, ex) = product(a, b);
413    let (y, ey) = product(c, d);
414    if x == 0.0 {
415        return (y, ey);
416    }
417    if y == 0.0 {
418        return (x, ex);
419    }
420    let e = ex.max(ey);
421    (libm::scalbn(x, ex - e) + libm::scalbn(y, ey - e), e)
422}