LCOV - code coverage report
Current view: top level - src/pairinteraction - diagonalization.py (source / functions) Hit Total Coverage
Test: coverage.info Lines: 41 48 85.4 %
Date: 2026-07-28 15:38:42 Functions: 3 3 100.0 %

          Line data    Source code
       1             : # SPDX-FileCopyrightText: 2025 PairInteraction Developers
       2             : # SPDX-License-Identifier: LGPL-3.0-or-later
       3           1 : from __future__ import annotations
       4             : 
       5           1 : from typing import TYPE_CHECKING, Any, Literal, TypeAlias, TypeVar
       6             : 
       7           1 : from pairinteraction import _backend
       8           1 : from pairinteraction.enums import get_cpp_float_type
       9           1 : from pairinteraction.units import QuantityScalar
      10             : 
      11             : if TYPE_CHECKING:
      12             :     from collections.abc import Callable, Sequence
      13             : 
      14             :     from pairinteraction.enums import FloatType
      15             :     from pairinteraction.system import SystemBase
      16             :     from pairinteraction.units import PintFloat
      17             : 
      18             :     Quantity = TypeVar("Quantity", bound="float | PintFloat")
      19             : 
      20             : 
      21           1 : Diagonalizer = Literal["eigen", "lapacke_evd", "lapacke_evr", "feast"]
      22           1 : UnionCPPDiagonalizer: TypeAlias = "_backend.DiagonalizerInterfaceReal | _backend.DiagonalizerInterfaceComplex"
      23           1 : UnionCPPDiagonalizerType: TypeAlias = "type[_backend.DiagonalizerInterfaceReal | _backend.DiagonalizerInterfaceComplex]"
      24             : 
      25           1 : _DiagonalizerDict: dict[str, dict[Diagonalizer, UnionCPPDiagonalizerType]] = {
      26             :     "real": {
      27             :         "eigen": _backend.DiagonalizerEigenReal,
      28             :         "lapacke_evd": _backend.DiagonalizerLapackeEvdReal,
      29             :         "lapacke_evr": _backend.DiagonalizerLapackeEvrReal,
      30             :         "feast": _backend.DiagonalizerFeastReal,
      31             :     },
      32             :     "complex": {
      33             :         "eigen": _backend.DiagonalizerEigenComplex,
      34             :         "lapacke_evd": _backend.DiagonalizerLapackeEvdComplex,
      35             :         "lapacke_evr": _backend.DiagonalizerLapackeEvrComplex,
      36             :         "feast": _backend.DiagonalizerFeastComplex,
      37             :     },
      38             : }
      39             : 
      40             : 
      41           1 : def diagonalize(
      42             :     systems: Sequence[SystemBase[Any]],
      43             :     diagonalizer: Diagonalizer = "eigen",
      44             :     float_type: FloatType = "float64",
      45             :     rtol: float = 1e-6,
      46             :     sort_by_energy: bool = True,
      47             :     energy_range: tuple[Quantity | None, Quantity | None] = (None, None),
      48             :     energy_range_unit: str | None = None,
      49             :     m0: int | None = None,
      50             : ) -> None:
      51             :     """Diagonalize a list of systems in parallel using the C++ backend.
      52             : 
      53             :     A convenience function for diagonalizing a list of systems in parallel using the C++ backend.
      54             :     This is much faster than diagonalizing each system individually on the Python side.
      55             : 
      56             :     Examples:
      57             :         >>> import pairinteraction as pi
      58             :         >>> ket = pi.KetAtom("Rb", n=60, l=0, m=0.5)
      59             :         >>> basis = pi.BasisAtom("Rb", n=(58, 63), l=(0, 3))
      60             :         >>> systems = [pi.SystemAtom(basis).set_magnetic_field([0, 0, b], unit="gauss") for b in range(1, 4)]
      61             :         >>> print(systems[0])
      62             :         SystemAtom(BasisAtom('Rb', n=(58, 63), l=(0, 3)), is_diagonal=False)
      63             :         >>> pi.diagonalize(systems)
      64             :         >>> print(systems[0])
      65             :         SystemAtom(BasisAtom(|Rb:58,S_1/2,-1/2⟩ ... |Rb:63,F_5/2,5/2⟩), is_diagonal=True)
      66             : 
      67             :     Args:
      68             :         systems: A list of `SystemAtom` or `SystemPair` objects, which will get diagonalized inplace.
      69             :         diagonalizer: The diagonalizer method to use. Defaults to "eigen".
      70             :         float_type: The floating point precision to use for the diagonalization. Defaults to "float64".
      71             :         rtol: The relative tolerance allowed for eigenenergies. The error in eigenenergies is bounded
      72             :             by rtol * ||H||, where ||H|| is the norm of the Hamiltonian matrix. Defaults to 1e-6.
      73             :         sort_by_energy: Whether to sort the resulting basis by energy. Defaults to True.
      74             :         energy_range: A tuple specifying an energy range, in which eigenvlaues should be calculated.
      75             :             Specifying a range can speed up the diagonalization process (depending on the diagonalizer method).
      76             :             The accuracy of the eigenenergies is not affected by this, but not all eigenenergies will be calculated.
      77             :             Defaults to (None, None), i.e. calculate all eigenenergies.
      78             :         energy_range_unit: The unit in which the energy_range is given. Defaults to None assumes pint objects.
      79             :         m0: The search subspace size for the FEAST diagonalizer. Defaults to None.
      80             : 
      81             :     """
      82           1 :     cpp_systems = [s._cpp for s in systems]
      83           1 :     cpp_diagonalize_fct = get_cpp_diagonalize_function(systems[0])
      84           1 :     cpp_diagonalizer = get_cpp_diagonalizer(diagonalizer, systems[0], float_type, m0=m0)
      85             : 
      86           1 :     energy_range_au: list[float | None] = [None, None]
      87           1 :     for i, energy in enumerate(energy_range):
      88           1 :         if energy is not None:
      89           1 :             energy_range_au[i] = QuantityScalar.convert_user_to_au(energy, energy_range_unit, "energy")
      90             : 
      91           1 :     cpp_diagonalize_fct(cpp_systems, cpp_diagonalizer, energy_range_au[0], energy_range_au[1], rtol, sort_by_energy)
      92             : 
      93           1 :     for system, cpp_system in zip(systems, cpp_systems, strict=True):
      94           1 :         system._cpp = cpp_system
      95             : 
      96             : 
      97           1 : def get_cpp_diagonalize_function(system: SystemBase[Any]) -> Callable[..., None]:
      98           1 :     if isinstance(system._cpp, _backend.SystemAtomReal):
      99           1 :         return _backend.diagonalizeSystemAtomReal
     100           1 :     if isinstance(system._cpp, _backend.SystemAtomComplex):
     101           1 :         return _backend.diagonalizeSystemAtomComplex
     102           1 :     if isinstance(system._cpp, _backend.SystemPairReal):
     103           1 :         return _backend.diagonalizeSystemPairReal
     104           1 :     if isinstance(system._cpp, _backend.SystemPairComplex):
     105           1 :         return _backend.diagonalizeSystemPairComplex
     106           0 :     raise TypeError(
     107             :         f"system must be of type SystemAtomReal, SystemPairReal, SystemAtomComplex, or SystemPairComplex, "
     108             :         f"not {type(system)}"
     109             :     )
     110             : 
     111             : 
     112           1 : def get_cpp_diagonalizer(
     113             :     diagonalizer: Diagonalizer,
     114             :     system: SystemBase[Any],
     115             :     float_type: FloatType,
     116             :     m0: int | None = None,
     117             : ) -> UnionCPPDiagonalizer:
     118           1 :     if diagonalizer == "feast" and m0 is None:
     119           0 :         raise ValueError("m0 must be specified for the 'feast' diagonalizer")
     120           1 :     if diagonalizer != "feast" and m0 is not None:
     121           0 :         raise ValueError("m0 must not be specified if the diagonalizer is not 'feast'")
     122             : 
     123           1 :     if isinstance(system._cpp, (_backend.SystemAtomReal, _backend.SystemPairReal)):
     124           1 :         type_ = "real"
     125           1 :     elif isinstance(system._cpp, (_backend.SystemAtomComplex, _backend.SystemPairComplex)):
     126           1 :         type_ = "complex"
     127             :     else:
     128           0 :         raise TypeError(
     129             :             f"system must be of type SystemAtomReal, SystemPairReal, SystemAtomComplex, or SystemPairComplex, "
     130             :             f"not {type(system)}"
     131             :         )
     132             : 
     133           1 :     try:
     134           1 :         diagonalizer_class = _DiagonalizerDict[type_][diagonalizer]
     135           0 :     except KeyError:
     136           0 :         raise ValueError(
     137             :             f"Unknown diagonalizer '{diagonalizer}', should be one of {list(_DiagonalizerDict[type_].keys())}"
     138             :         ) from None
     139             : 
     140           1 :     cpp_float_type = get_cpp_float_type(float_type)
     141           1 :     if diagonalizer == "feast":
     142           0 :         return diagonalizer_class(m0=m0, float_type=cpp_float_type)  # type: ignore [call-arg]
     143           1 :     return diagonalizer_class(float_type=cpp_float_type)  # type: ignore [call-arg]

Generated by: LCOV version 1.16