Eigen  5.0.1
 
Loading...
Searching...
No Matches
GenericPacketMathComplex.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2009-2019 Gael Guennebaud <gael.guennebaud@inria.fr>
5// Copyright (C) 2018-2025 Rasmus Munk Larsen <rmlarsen@gmail.com>
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_ARCH_GENERIC_PACKET_MATH_COMPLEX_H
13#define EIGEN_ARCH_GENERIC_PACKET_MATH_COMPLEX_H
14
15// IWYU pragma: private
16#include "../../InternalHeaderCheck.h"
17
18namespace Eigen {
19namespace internal {
20
21EIGEN_GCC_FAST_MATH_COMPLEX_VECTORIZE_WORKAROUND_PUSH
22
23//----------------------------------------------------------------------
24// Complex Arithmetic and Functions
25//----------------------------------------------------------------------
26
27template <typename Packet>
28EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet pdiv_complex(const Packet& x, const Packet& y) {
29 using RealPacket = typename unpacket_traits<Packet>::as_real;
30 using RealScalar = typename unpacket_traits<RealPacket>::type;
31 // In the following we annotate the code for the case where the inputs
32 // are a pair length-2 SIMD vectors representing a single pair of complex
33 // numbers x = a + i*b, y = c + i*d.
34 const RealPacket one = pset1<RealPacket>(RealScalar(1));
35 const RealPacket abs_y = pabs(y.v);
36 const RealPacket abs_y_flip = pcplxflip(Packet(abs_y)).v;
37
38 const RealPacket mask = pcmp_lt(abs_y, abs_y_flip); // |c| < |d|
39 RealPacket y_scaled = pselect(mask, pdiv(abs_y, abs_y_flip), one);
40 y_scaled = por(y_scaled, pandnot(y.v, abs_y)); // copy signs in case |c| == |d|
41 RealPacket denom = pmul(y.v, y_scaled);
42 denom = padd(denom, pcplxflip(Packet(denom)).v); // c * c' + d * d'
43 Packet num = pmul(x, pconj(Packet(y_scaled))); // a * c' + b * d', -a * d + b * c
44 return Packet(pdiv(num.v, denom));
45}
46
47template <typename Packet>
48EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet pmul_complex(const Packet& x, const Packet& y) {
49 // In the following we annotate the code for the case where the inputs
50 // are a pair length-2 SIMD vectors representing a single pair of complex
51 // numbers x = a + i*b, y = c + i*d.
52 Packet x_re = pdupreal(x); // a, a
53 Packet x_im = pdupimag(x); // b, b
54 Packet tmp_re = Packet(pmul(x_re.v, y.v)); // a*c, a*d
55 Packet tmp_im = Packet(pmul(x_im.v, y.v)); // b*c, b*d
56 tmp_im = pcplxflip(pconj(tmp_im)); // -b*d, d*c
57 return padd(tmp_im, tmp_re); // a*c - b*d, a*d + b*c
58}
59
60template <typename Packet>
61EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet plog_complex(const Packet& x) {
62 using RealPacket = typename unpacket_traits<Packet>::as_real;
63
64 // Real part
65 RealPacket x_flip = pcplxflip(x).v; // b, a
66 Packet x_norm = phypot_complex(x); // sqrt(a^2 + b^2), sqrt(a^2 + b^2)
67 RealPacket xlogr = plog(x_norm.v); // log(sqrt(a^2 + b^2)), log(sqrt(a^2 + b^2))
68
69 // Imag part
70 RealPacket ximg = patan2(x.v, x_flip); // atan2(a, b), atan2(b, a)
71
72 const RealPacket cst_pos_inf = pinf<RealPacket>();
73 RealPacket x_abs = pabs(x.v);
74 RealPacket is_x_pos_inf = pcmp_eq(x_abs, cst_pos_inf);
75 RealPacket is_y_pos_inf = pcplxflip(Packet(is_x_pos_inf)).v;
76 RealPacket is_any_inf = por(is_x_pos_inf, is_y_pos_inf);
77 RealPacket xreal = pselect(is_any_inf, cst_pos_inf, xlogr);
78
79 return Packet(pselect(peven_mask(xreal), xreal, ximg)); // log(sqrt(a^2 + b^2)), atan2(b, a)
80}
81
82template <typename Packet>
83EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet pexp_complex(const Packet& a) {
84 using RealPacket = typename unpacket_traits<Packet>::as_real;
85 using Scalar = typename unpacket_traits<Packet>::type;
86 using RealScalar = typename Scalar::value_type;
87 const RealPacket even_mask = peven_mask(a.v);
88 const RealPacket odd_mask = pcplxflip(Packet(even_mask)).v;
89
90 // Let a = x + iy.
91 // exp(a) = exp(x) * cis(y), plus some special edge-case handling.
92
93 // exp(x):
94 RealPacket x = pand(a.v, even_mask);
95 x = por(x, pcplxflip(Packet(x)).v);
96 RealPacket expx = pexp(x); // exp(x);
97
98 // cis(y):
99 RealPacket y = pand(odd_mask, a.v);
100 y = por(y, pcplxflip(Packet(y)).v);
101 RealPacket cisy = psincos_selector<RealPacket>(y);
102 cisy = pcplxflip(Packet(cisy)).v; // cos(y) + i * sin(y)
103
104 const RealPacket cst_pos_inf = pinf<RealPacket>();
105 const RealPacket cst_neg_inf = por(psignmask<RealPacket>(), pinf<RealPacket>());
106
107 // If x is -inf, we know that cossin(y) is bounded,
108 // so the result is (0, +/-0), where the sign of the imaginary part comes
109 // from the sign of cossin(y).
110 RealPacket cisy_sign = por(pandnot(cisy, pabs(cisy)), pset1<RealPacket>(RealScalar(1)));
111 cisy = pselect(pcmp_eq(x, cst_neg_inf), cisy_sign, cisy);
112
113 // If x is inf, and cos(y) has unknown sign (y is inf or NaN), the result
114 // is (+/-inf, NaN), where the signs are undetermined (take the sign of y).
115 RealPacket y_sign = por(pandnot(y, pabs(y)), pset1<RealPacket>(RealScalar(1)));
116 cisy = pselect(pand(pcmp_eq(x, cst_pos_inf), pisnan(cisy)), pand(y_sign, even_mask), cisy);
117
118 // If exp(x) is +inf and y is finite, replace cisy with copysign(1, cisy) to
119 // prevent inf * 0 = NaN. The vectorized sincos may compute exact zero
120 // for near-zero values like cos(pi/2), and inf * +-1 = +-inf is correct.
121 // The y=0 case is handled separately below.
122 RealPacket cisy_sign_one = por(pand(cisy, psignmask<RealPacket>()), pset1<RealPacket>(RealScalar(1)));
123 RealPacket expx_inf_y_finite = pand(pcmp_eq(expx, cst_pos_inf), pcmp_lt(pabs(y), cst_pos_inf));
124 cisy = pselect(expx_inf_y_finite, cisy_sign_one, cisy);
125
126 Packet result = Packet(pmul(expx, cisy));
127
128 // If y is +/- 0, the input is real, so take the real result for consistency.
129 result = pselect(Packet(pcmp_eq(y, pzero(y))), Packet(por(pand(expx, even_mask), pand(y, odd_mask))), result);
130
131 return result;
132}
133
134template <typename Packet>
135EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet psqrt_complex(const Packet& a) {
136 using Scalar = typename unpacket_traits<Packet>::type;
137 using RealScalar = typename Scalar::value_type;
138 using RealPacket = typename unpacket_traits<Packet>::as_real;
139
140 // Computes the principal sqrt of the complex numbers in the input.
141 //
142 // For example, for packets containing 2 complex numbers stored in interleaved format
143 // a = [a0, a1] = [x0, y0, x1, y1],
144 // where x0 = real(a0), y0 = imag(a0) etc., this function returns
145 // b = [b0, b1] = [u0, v0, u1, v1],
146 // such that b0^2 = a0, b1^2 = a1.
147 //
148 // To derive the formula for the complex square roots, let's consider the equation for
149 // a single complex square root of the number x + i*y. We want to find real numbers
150 // u and v such that
151 // (u + i*v)^2 = x + i*y <=>
152 // u^2 - v^2 + i*2*u*v = x + i*y.
153 // By equating the real and imaginary parts we get:
154 // u^2 - v^2 = x
155 // 2*u*v = y.
156 //
157 // For x >= 0, this has the numerically stable solution
158 // u = sqrt(0.5 * (x + sqrt(x^2 + y^2)))
159 // v = 0.5 * (y / u)
160 // and for x < 0,
161 // v = sign(y) * sqrt(0.5 * (-x + sqrt(x^2 + y^2)))
162 // u = 0.5 * (y / v)
163 //
164 // To avoid unnecessary over- and underflow, we compute sqrt(x^2 + y^2) as
165 // l = max(|x|, |y|) * sqrt(1 + (min(|x|, |y|) / max(|x|, |y|))^2) ,
166
167 // In the following, without loss of generality, we have annotated the code, assuming
168 // that the input is a packet of 2 complex numbers.
169 //
170 // Step 1. Compute l = [l0, l0, l1, l1], where
171 // l0 = sqrt(x0^2 + y0^2), l1 = sqrt(x1^2 + y1^2)
172 // To avoid over- and underflow, we use the stable formula for each hypotenuse
173 // l0 = (min0 == 0 ? max0 : max0 * sqrt(1 + (min0/max0)**2)),
174 // where max0 = max(|x0|, |y0|), min0 = min(|x0|, |y0|), and similarly for l1.
175
176 RealPacket a_abs = pabs(a.v); // [|x0|, |y0|, |x1|, |y1|]
177 RealPacket a_abs_flip = pcplxflip(Packet(a_abs)).v; // [|y0|, |x0|, |y1|, |x1|]
178 RealPacket a_max = pmax(a_abs, a_abs_flip);
179 RealPacket a_min = pmin(a_abs, a_abs_flip);
180 RealPacket a_min_zero_mask = pcmp_eq(a_min, pzero(a_min));
181 RealPacket a_max_zero_mask = pcmp_eq(a_max, pzero(a_max));
182 RealPacket r = pdiv(a_min, a_max);
183 const RealPacket cst_one = pset1<RealPacket>(RealScalar(1));
184 RealPacket l = pmul(a_max, psqrt(padd(cst_one, pmul(r, r)))); // [l0, l0, l1, l1]
185 // Set l to a_max if a_min is zero.
186 l = pselect(a_min_zero_mask, a_max, l);
187
188 // Step 2. Compute [rho0, *, rho1, *], where
189 // rho0 = sqrt(0.5 * (l0 + |x0|)), rho1 = sqrt(0.5 * (l1 + |x1|))
190 // We don't care about the imaginary parts computed here. They will be overwritten later.
191 const RealPacket cst_half = pset1<RealPacket>(RealScalar(0.5));
192 Packet rho;
193 rho.v = psqrt(pmul(cst_half, padd(a_abs, l)));
194
195 // Step 3. Compute [rho0, eta0, rho1, eta1], where
196 // eta0 = (y0 / rho0) / 2, and eta1 = (y1 / rho1) / 2.
197 // set eta = 0 if input is 0 + i0.
198 RealPacket eta = pandnot(pmul(cst_half, pdiv(a.v, pcplxflip(rho).v)), a_max_zero_mask);
199 RealPacket real_mask = peven_mask(a.v);
200 Packet positive_real_result;
201 // Compute result for inputs with positive real part.
202 positive_real_result.v = pselect(real_mask, rho.v, eta);
203
204 // Step 4. Compute solution for inputs with negative real part:
205 // [|eta0|, sign(y0)*rho0, |eta1|, sign(y1)*rho1]
206 // [+0.0, -0.0, ...]: the sign bit of the imaginary (odd) lanes only.
207 const RealPacket cst_imag_sign_mask = pandnot(psignmask<RealPacket>(), real_mask);
208 RealPacket imag_signs = pand(a.v, cst_imag_sign_mask);
209 Packet negative_real_result;
210 // Notice that rho is positive, so taking its absolute value is a noop.
211 negative_real_result.v = por(pabs(pcplxflip(positive_real_result).v), imag_signs);
212
213 // Step 5. Select solution branch based on the sign of the real parts.
214 Packet negative_real_mask;
215 negative_real_mask.v = pcmp_lt(pand(real_mask, a.v), pzero(a.v));
216 negative_real_mask.v = por(negative_real_mask.v, pcplxflip(negative_real_mask).v);
217 Packet result = pselect(negative_real_mask, negative_real_result, positive_real_result);
218
219 // Step 6. Handle special cases for infinities:
220 // * If z is (x,+∞), the result is (+∞,+∞) even if x is NaN
221 // * If z is (x,-∞), the result is (+∞,-∞) even if x is NaN
222 // * If z is (-∞,y), the result is (0*|y|,+∞) for finite or NaN y
223 // * If z is (+∞,y), the result is (+∞,0*|y|) for finite or NaN y
224 const RealPacket cst_pos_inf = pinf<RealPacket>();
225 Packet is_inf;
226 is_inf.v = pcmp_eq(a_abs, cst_pos_inf);
227 Packet is_real_inf;
228 is_real_inf.v = pand(is_inf.v, real_mask);
229 is_real_inf = por(is_real_inf, pcplxflip(is_real_inf));
230 // prepare packet of (+∞,0*|y|) or (0*|y|,+∞), depending on the sign of the infinite real part.
231 Packet real_inf_result;
232 real_inf_result.v = pmul(a_abs, pset1<Packet>(Scalar(RealScalar(1.0), RealScalar(0.0))).v);
233 real_inf_result.v = pselect(negative_real_mask.v, pcplxflip(real_inf_result).v, real_inf_result.v);
234 // prepare packet of (+∞,+∞) or (+∞,-∞), depending on the sign of the infinite imaginary part.
235 Packet is_imag_inf;
236 is_imag_inf.v = pandnot(is_inf.v, real_mask);
237 is_imag_inf = por(is_imag_inf, pcplxflip(is_imag_inf));
238 Packet imag_inf_result;
239 imag_inf_result.v = por(pand(cst_pos_inf, real_mask), pandnot(a.v, real_mask));
240 // unless otherwise specified, if either the real or imaginary component is nan, the entire result is nan
241 Packet result_is_nan = pisnan(result);
242 result = por(result_is_nan, result);
243
244 return pselect(is_imag_inf, imag_inf_result, pselect(is_real_inf, real_inf_result, result));
245}
246
247// \internal \returns the norm of a complex number z = x + i*y, defined as sqrt(x^2 + y^2).
248// Implemented using the hypot(a,b) algorithm from https://doi.org/10.48550/arXiv.1904.09481
249template <typename Packet>
250EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet phypot_complex(const Packet& a) {
251 using Scalar = typename unpacket_traits<Packet>::type;
252 using RealScalar = typename Scalar::value_type;
253 using RealPacket = typename unpacket_traits<Packet>::as_real;
254
255 const RealPacket cst_zero_rp = pset1<RealPacket>(static_cast<RealScalar>(0.0));
256 const RealPacket cst_minus_one_rp = pset1<RealPacket>(static_cast<RealScalar>(-1.0));
257 const RealPacket cst_two_rp = pset1<RealPacket>(static_cast<RealScalar>(2.0));
258 const RealPacket evenmask = peven_mask(a.v);
259
260 RealPacket a_abs = pabs(a.v);
261 RealPacket a_flip = pcplxflip(Packet(a_abs)).v; // |b|, |a|
262 RealPacket a_all = pselect(evenmask, a_abs, a_flip); // |a|, |a|
263 RealPacket b_all = pselect(evenmask, a_flip, a_abs); // |b|, |b|
264
265 RealPacket a2 = pmul(a.v, a.v); // |a^2, b^2|
266 RealPacket a2_flip = pcplxflip(Packet(a2)).v; // |b^2, a^2|
267 RealPacket h = psqrt(padd(a2, a2_flip)); // |sqrt(a^2 + b^2), sqrt(a^2 + b^2)|
268 RealPacket h_sq = pmul(h, h); // |a^2 + b^2, a^2 + b^2|
269 RealPacket a_sq = pselect(evenmask, a2, a2_flip); // |a^2, a^2|
270 RealPacket m_h_sq = pmul(h_sq, cst_minus_one_rp);
271 RealPacket m_a_sq = pmul(a_sq, cst_minus_one_rp);
272 RealPacket x = psub(psub(pmadd(h, h, m_h_sq), pmadd(b_all, b_all, psub(a_sq, h_sq))), pmadd(a_all, a_all, m_a_sq));
273 h = psub(h, pdiv(x, pmul(cst_two_rp, h))); // |h - x/(2*h), h - x/(2*h)|
274
275 // handle zero-case
276 RealPacket iszero = pcmp_eq(por(a_abs, a_flip), cst_zero_rp);
277
278 h = pandnot(h, iszero); // |sqrt(a^2+b^2), sqrt(a^2+b^2)|
279 return Packet(h); // |sqrt(a^2+b^2), sqrt(a^2+b^2)|
280}
281
282EIGEN_GCC_FAST_MATH_COMPLEX_VECTORIZE_WORKAROUND_POP
283
284} // end namespace internal
285} // end namespace Eigen
286
287#endif // EIGEN_ARCH_GENERIC_PACKET_MATH_COMPLEX_H