19#ifndef EIGEN_STRUCTURED_MATRIX_UTILS_H
20#define EIGEN_STRUCTURED_MATRIX_UTILS_H
23#include "./InternalHeaderCheck.h"
31struct StructuredShape {};
35constexpr Index structured_direct_threshold() {
return 32; }
40constexpr Index structured_scalar_threshold() {
return 16; }
45using structured_exponent_type = numext::int64_t;
51template <
typename Scalar>
52typename NumTraits<Scalar>::Real structured_component_magnitude_impl(
const Scalar& z, std::true_type) {
53 using RealScalar =
typename NumTraits<Scalar>::Real;
54 using Binary = binary_floating_point_traits<RealScalar>;
55 return numext::bit_cast<RealScalar>(
56 larger_magnitude_bits<RealScalar>(Binary::magnitude(numext::real(z)), Binary::magnitude(numext::imag(z))));
58template <
typename Scalar>
59typename NumTraits<Scalar>::Real structured_component_magnitude_impl(
const Scalar& z, std::false_type) {
60 return numext::maxi(numext::abs(numext::real(z)), numext::abs(numext::imag(z)));
62template <
typename Scalar>
63typename NumTraits<Scalar>::Real structured_component_magnitude(
const Scalar& z) {
64 using RealScalar =
typename NumTraits<Scalar>::Real;
65 return structured_component_magnitude_impl(z,
66 bool_constant < complex_array_access<Scalar>::value &&
67 use_subnormal_preserving_scaling<RealScalar, RealScalar>::value > ());
84template <typename Scalar, bool IsComplex = NumTraits<Scalar>::IsComplex>
85struct structured_balance_impl {
86 using RealScalar =
typename NumTraits<Scalar>::Real;
87 template <
typename Exponent>
88 static Scalar run(
const Scalar& z, Exponent& exponent) {
89 const RealScalar mag = structured_component_magnitude(z);
90 if (numext::is_exactly_zero_no_flush(mag) || !(numext::isfinite)(mag))
return z;
91 const int e = frexp_exponent_preserving_subnormals(mag);
93 return apply_exponent(z, -e);
95 static Scalar apply_exponent(
const Scalar& z,
int e) {
96 return Scalar(ldexp_preserving_subnormals(numext::real(z), e), ldexp_preserving_subnormals(numext::imag(z), e));
100template <
typename Scalar>
101struct structured_balance_impl<Scalar, false> {
102 template <
typename Exponent>
103 static Scalar run(
const Scalar& x, Exponent& exponent) {
104 if (numext::is_exactly_zero_no_flush(x) || !(numext::isfinite)(x))
return x;
105 const int e = frexp_exponent_preserving_subnormals(x);
107 return apply_exponent(x, -e);
109 static Scalar apply_exponent(
const Scalar& x,
int e) {
return ldexp_preserving_subnormals(x, e); }
112template <
typename Scalar,
typename Exponent>
113Scalar structured_balance(
const Scalar& z, Exponent& exponent) {
114 return structured_balance_impl<Scalar>::run(z, exponent);
121template <
typename Scalar,
typename Exponent>
122Scalar structured_ldexp_clamped(
const Scalar& z, Exponent exponent) {
123 constexpr Exponent kMaxExponent = Exponent(1) << 24;
124 const int e =
static_cast<int>(numext::mini(numext::maxi(exponent, -kMaxExponent), kMaxExponent));
125 return structured_balance_impl<Scalar>::apply_exponent(z, e);
134template <
typename RealVectorType>
135std::vector<Index> structured_svd_permutation(
const RealVectorType& mods) {
136 using RealScalar =
typename RealVectorType::Scalar;
137 std::vector<Index> perm;
138 perm.reserve(
static_cast<std::size_t
>(mods.size()));
139 for (Index k = 0; k < mods.size(); ++k) perm.push_back(k);
140 std::stable_sort(perm.begin(), perm.end(), [&mods](Index a, Index b) {
141 const RealScalar ka = mods[a], kb = mods[b];
143 return std::isgreater(ka, kb) || (!(numext::isnan)(ka) && (numext::isnan)(kb));
162#ifndef EIGEN_AVOID_THREAD_LOCAL
163template <
typename RealScalar>
164FFT<RealScalar>& structured_fft_engine() {
165 static thread_local FFT<RealScalar> fft;
169template <
typename RealScalar>
170FFT<RealScalar> structured_fft_engine() {
171 return FFT<RealScalar>();
183inline Index fft_next_good_size(Index n) {
185 for (Index m = n;; ++m) {
187 while (r % 2 == 0) r /= 2;
188 while (r % 3 == 0) r /= 3;
189 while (r % 5 == 0) r /= 5;
190 if (r == 1)
return m;
199template <typename Scalar, bool IsComplex = NumTraits<Scalar>::IsComplex>
200struct structured_scalar_part_impl {
201 template <
typename Xpr>
202 static const Xpr& run(
const Xpr& xpr) {
205 static const Scalar& run_scalar(
const Scalar& x) {
return x; }
208template <
typename Scalar>
209struct structured_scalar_part_impl<Scalar, false> {
210 template <
typename Xpr>
211 static typename Xpr::RealReturnType run(
const Xpr& xpr) {
214 static Scalar run_scalar(
const std::complex<Scalar>& x) {
return numext::real(x); }
234template <
typename Xpr>
235bool structured_exponent_bound_finite(
const Xpr& x,
int& e) {
236 using ScalarTraits = NumTraits<typename Xpr::Scalar>;
237 using RealScalar =
typename ScalarTraits::Real;
242 if (x.size() == 0)
return true;
244 if (ScalarTraits::IsComplex)
247 m = x.realView().cwiseAbs().maxCoeff();
249 m = x.cwiseAbs().maxCoeff();
253 m = safe_scaling<RealScalar>::recover_flushed_max_coeff(x, m);
255 if (!(numext::isfinite)(m))
return false;
256 if (!numext::is_exactly_zero_no_flush(m)) {
257 e = frexp_exponent_preserving_subnormals(m);
258 if (ScalarTraits::IsComplex) ++e;
265template <
typename Xpr>
266int structured_exponent_bound(
const Xpr& x) {
268 structured_exponent_bound_finite(x, e);
275template <typename Xpr, std::enable_if_t<!NumTraits<typename Xpr::Scalar>::IsComplex,
bool> =
true>
276void structured_ldexp_entries_packet(Xpr& M,
int e) {
277 M.array() = M.array().ldexp(e);
280template <typename Xpr, std::enable_if_t<complex_array_access<typename Xpr::Scalar>::value,
bool> =
true>
281void structured_ldexp_entries_packet(Xpr& M,
int e) {
282 M.realView().array() = M.realView().array().ldexp(e);
285template <typename Xpr, std::enable_if_t<NumTraits<typename Xpr::Scalar>::IsComplex &&
286 !complex_array_access<typename Xpr::Scalar>::value,
288void structured_ldexp_entries_packet(Xpr& M,
int e) {
289 using Scalar =
typename Xpr::Scalar;
290 M = M.unaryExpr([e](
const Scalar& z) {
return structured_ldexp_clamped(z, Index(e)); });
293template <
typename Xpr>
294void structured_ldexp_entries_impl(Xpr& M,
int e,
int, std::false_type) {
295 structured_ldexp_entries_packet(M, e);
298template <
typename Xpr>
299void structured_ldexp_entries_impl(Xpr& M,
int e,
int bound, std::true_type) {
300 using Scalar =
typename Xpr::Scalar;
301 using RealScalar =
typename NumTraits<Scalar>::Real;
306 const int componentBound = bound - (NumTraits<Scalar>::IsComplex ? 1 : 0);
307 const int recovery = safe_scaling<RealScalar>::subnormal_recovery_exponent();
308 if (componentBound - 1 >= recovery && componentBound - 1 + e >= recovery)
309 structured_ldexp_entries_packet(M, e);
311 M = M.unaryExpr(scale_by_exponent_op<RealScalar>(e));
324template <
typename Xpr>
325void structured_ldexp_entries(Xpr& M,
int e,
int bound) {
327 using Scalar =
typename Xpr::Scalar;
328 using RealScalar =
typename NumTraits<Scalar>::Real;
329 structured_ldexp_entries_impl(M, e, bound,
330 bool_constant<use_subnormal_preserving_scaling<RealScalar, Scalar>::value>());
336template <
typename Xpr>
337void structured_ldexp_entries(Xpr& M,
int e) {
339 structured_ldexp_entries(M, e, structured_exponent_bound(M));
342template <
typename Xpr>
343void structured_ldexp_entries_exact_impl(Xpr& M,
int e, std::false_type) {
344 structured_ldexp_entries_packet(M, e);
346template <
typename Xpr>
347void structured_ldexp_entries_exact_impl(Xpr& M,
int e, std::true_type) {
348 using RealScalar =
typename NumTraits<typename Xpr::Scalar>::Real;
349 M = M.unaryExpr(scale_by_exponent_op<RealScalar>(e));
357template <
typename Xpr>
358void structured_ldexp_entries_exact(Xpr& M,
int e) {
360 using Scalar =
typename Xpr::Scalar;
361 using RealScalar =
typename NumTraits<Scalar>::Real;
362 structured_ldexp_entries_exact_impl(M, e,
363 bool_constant<use_subnormal_preserving_scaling<RealScalar, Scalar>::value>());
371template <
typename ComplexVectorType>
372ComplexVectorType structured_reverse_symbol(
const ComplexVectorType& symbol) {
373 const Index p = symbol.size();
374 ComplexVectorType reversed(p);
376 reversed[0] = symbol[0];
377 reversed.tail(p - 1) = symbol.tail(p - 1).reverse();
392template <
typename Scalar,
typename Dest,
typename Rhs>
393void structured_fft_apply(Dest& dst,
const Matrix<std::complex<
typename NumTraits<Scalar>::Real>, Dynamic, 1>& symbol,
394 Index outSize,
const Rhs& rhs,
const Scalar& alpha) {
395 using RealScalar =
typename NumTraits<Scalar>::Real;
396 using Complex = std::complex<RealScalar>;
397 using ComplexVector = Matrix<Complex, Dynamic, 1>;
399 const Index p = symbol.size();
400 eigen_assert(rhs.rows() <= p && outSize <= p);
403 dst.row(0) += alpha * structured_scalar_part_impl<Scalar>::run(Complex(symbol.coeff(0)) *
404 rhs.row(0).template cast<Complex>());
408 auto&& fft = structured_fft_engine<RealScalar>();
409 ComplexVector xt = ComplexVector::Zero(p);
410 ComplexVector xf(p), yt(p);
411 for (Index k = 0; k < rhs.cols(); ++k) {
412 xt.head(rhs.rows()) = rhs.col(k).template cast<Complex>();
414 xf.array() *= symbol.array();
416 dst.col(k) += alpha * structured_scalar_part_impl<Scalar>::run(yt.head(outSize));
426template <
typename SymbolType,
typename ModsType,
typename RealScalar>
427SymbolType structured_pinv_symbol(
const SymbolType& symbol,
const ModsType& mods,
const RealScalar& tol) {
428 using Complex =
typename SymbolType::Scalar;
429 const auto w = (mods.array() < tol).select(RealScalar(0), mods.array().inverse()).template cast<Complex>().eval();
430 return (symbol.array().conjugate() * w * w).matrix();
438template <
typename SymbolType,
typename ModsType>
439typename ModsType::Scalar structured_rank_threshold(
const SymbolType& symbol,
const ModsType& mods) {
440 using RealScalar =
typename ModsType::Scalar;
441 const RealScalar factor = RealScalar(mods.size()) * NumTraits<RealScalar>::epsilon();
442 RealScalar tol = factor * mods.maxCoeff();
443 if (!(numext::isfinite)(tol)) tol = (RealScalar(2) * factor) * (symbol * RealScalar(0.5)).cwiseAbs().maxCoeff();
444 return numext::maxi(tol, (std::numeric_limits<RealScalar>::min)());
458template <
typename Op,
typename Rhs>
459struct structured_product_impl : generic_product_impl_base<Op, Rhs, structured_product_impl<Op, Rhs>> {
460 using Scalar =
typename Product<Op, Rhs>::Scalar;
462 template <
typename Dest>
463 static void evalTo(Dest& dst,
const Op& lhs,
const Rhs& rhs) {
465 scaleAndAddTo(dst, lhs, rhs, Scalar(1));
468 template <
typename Dest>
469 static void scaleAndAddTo(Dest& dst,
const Op& lhs,
const Rhs& rhs,
const Scalar& alpha) {
470 using RhsNested =
typename nested_eval<Rhs, Op::RowsAtCompileTime>::type;
471 RhsNested actualRhs(rhs);
472 lhs.addProduct(dst, actualRhs, alpha);
Namespace containing all symbols from the Eigen library.