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]
|