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