Line data Source code
1 : // SPDX-FileCopyrightText: 2024 PairInteraction Developers 2 : // SPDX-License-Identifier: LGPL-3.0-or-later 3 : 4 : #pragma once 5 : 6 : #include "pairinteraction/enums/FloatType.hpp" 7 : #include "pairinteraction/utils/eigen_assertion.hpp" 8 : #include "pairinteraction/utils/eigen_compat.hpp" 9 : #include "pairinteraction/utils/traits.hpp" 10 : 11 : #include <Eigen/Dense> 12 : #include <Eigen/SparseCore> 13 : #include <complex> 14 : #include <optional> 15 : 16 : namespace pairinteraction { 17 : template <typename Scalar> 18 : struct EigenSystemH { 19 : static_assert(traits::NumTraits<Scalar>::from_floating_point_v); 20 : using real_t = typename traits::NumTraits<Scalar>::real_t; 21 : Eigen::SparseMatrix<Scalar, Eigen::RowMajor> eigenvectors; 22 : Eigen::VectorX<real_t> eigenvalues; 23 : }; 24 : 25 : template <typename Scalar> 26 : class DiagonalizerInterface { 27 : public: 28 : static_assert(traits::NumTraits<Scalar>::from_floating_point_v); 29 : 30 : using real_t = typename traits::NumTraits<Scalar>::real_t; 31 : 32 : DiagonalizerInterface(FloatType float_type); 33 113 : virtual ~DiagonalizerInterface() = default; 34 : virtual EigenSystemH<Scalar> eigh(const Eigen::SparseMatrix<Scalar, Eigen::RowMajor> &matrix, 35 : double rtol) const = 0; 36 : virtual EigenSystemH<Scalar> eigh(const Eigen::SparseMatrix<Scalar, Eigen::RowMajor> &matrix, 37 : std::optional<real_t> min_eigenvalue, 38 : std::optional<real_t> max_eigenvalue, double rtol) const; 39 : 40 : protected: 41 : FloatType float_type; 42 : template <typename ScalarLim> 43 : Eigen::MatrixX<ScalarLim> subtract_mean(const Eigen::MatrixX<Scalar> &matrix, real_t &shift, 44 : double rtol) const; 45 : template <typename RealLim> 46 : Eigen::VectorX<real_t> add_mean(const Eigen::VectorX<RealLim> &eigenvalues, real_t shift) const; 47 : }; 48 : 49 : extern template class DiagonalizerInterface<double>; 50 : extern template class DiagonalizerInterface<std::complex<double>>; 51 : 52 : extern template Eigen::MatrixX<float> 53 : DiagonalizerInterface<double>::subtract_mean(const Eigen::MatrixX<double> &matrix, double &shift, 54 : double rtol) const; 55 : extern template Eigen::MatrixX<std::complex<float>> 56 : DiagonalizerInterface<std::complex<double>>::subtract_mean( 57 : const Eigen::MatrixX<std::complex<double>> &matrix, double &shift, double rtol) const; 58 : 59 : extern template Eigen::VectorX<double> 60 : DiagonalizerInterface<double>::add_mean(const Eigen::VectorX<float> &shifted_eigenvalues, 61 : double shift) const; 62 : extern template Eigen::VectorX<double> DiagonalizerInterface<std::complex<double>>::add_mean( 63 : const Eigen::VectorX<float> &shifted_eigenvalues, double shift) const; 64 : 65 : extern template Eigen::MatrixX<double> 66 : DiagonalizerInterface<double>::subtract_mean(const Eigen::MatrixX<double> &matrix, double &shift, 67 : double rtol) const; 68 : extern template Eigen::MatrixX<std::complex<double>> 69 : DiagonalizerInterface<std::complex<double>>::subtract_mean( 70 : const Eigen::MatrixX<std::complex<double>> &matrix, double &shift, double rtol) const; 71 : 72 : extern template Eigen::VectorX<double> 73 : DiagonalizerInterface<double>::add_mean(const Eigen::VectorX<double> &shifted_eigenvalues, 74 : double shift) const; 75 : extern template Eigen::VectorX<double> DiagonalizerInterface<std::complex<double>>::add_mean( 76 : const Eigen::VectorX<double> &shifted_eigenvalues, double shift) const; 77 : } // namespace pairinteraction