Skip to content

Commit d043895

Browse files
committed
implement time-dependent Hamiltonian specification for numba backend
1 parent ce7900e commit d043895

3 files changed

Lines changed: 199 additions & 18 deletions

File tree

‎jftools/short_iterative_lanczos.py‎

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -280,6 +280,15 @@ def _is_csr_matrix(H):
280280
return csr_array is not None and isinstance(H, csr_array)
281281

282282

283+
def _is_numba_sum_operator(H):
284+
if not isinstance(H, (tuple, list)) or len(H) < 2:
285+
return False
286+
for term in H[1:]:
287+
if not isinstance(term, (tuple, list)) or len(term) != 2 or not callable(term[1]):
288+
return False
289+
return True
290+
291+
283292
def _select_backend(H, backend):
284293
backend = backend.strip().lower()
285294

@@ -299,7 +308,7 @@ def _select_backend(H, backend):
299308
if backend != "auto":
300309
raise ValueError("Unknown backend value '%s'. Valid values are 'python', 'numba', 'cython', 'auto'." % backend)
301310

302-
if have_numba_backend and (_is_dense_matrix(H) or _is_csr_matrix(H)):
311+
if have_numba_backend and (_is_dense_matrix(H) or _is_csr_matrix(H) or _is_numba_sum_operator(H)):
303312
return "numba"
304313
if have_cython_backend:
305314
return "cython"

‎jftools/short_iterative_lanczos_numba.py‎

Lines changed: 100 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -1,18 +1,47 @@
11
import numpy as np
2-
from numba import njit, objmode
2+
from numba import literal_unroll, njit, objmode
33
from numba.core import types
4+
from numba.core.dispatcher import Dispatcher
45
from numba.extending import overload
56
from numba_lapack import dstevd, zgemm
67
from scipy import sparse as sp
78

89
_CALLABLE_OPERATOR_REGISTRY = {}
910

11+
1012
def _register_callable_operator(H):
1113
handle = id(H)
1214
_CALLABLE_OPERATOR_REGISTRY[handle] = H
1315
return handle
1416

1517

18+
def _evaluate_operator_coefficient(coeff_handle, t):
19+
return complex(_CALLABLE_OPERATOR_REGISTRY[coeff_handle](t))
20+
21+
22+
def _evaluate_operator_coefficient_numba(coeff_spec, t):
23+
raise NotImplementedError
24+
25+
26+
@overload(_evaluate_operator_coefficient_numba)
27+
def _overload_evaluate_operator_coefficient_numba(coeff_spec, t):
28+
if isinstance(coeff_spec, types.Integer):
29+
def impl(coeff_spec, t):
30+
with objmode(coeff='complex128'):
31+
coeff = _evaluate_operator_coefficient(coeff_spec, t)
32+
return coeff
33+
34+
return impl
35+
36+
if isinstance(coeff_spec, types.Dispatcher):
37+
def impl(coeff_spec, t):
38+
return complex(coeff_spec(t))
39+
40+
return impl
41+
42+
raise TypeError("Coefficient spec must be a callable handle or a Numba-jitted function")
43+
44+
1645
@njit(cache=True)
1746
def _vdot_numba(a, b):
1847
out = 0.0 + 0.0j
@@ -44,34 +73,51 @@ def _axpy_numba(scale, x, y):
4473
y[idx] += scale * x[idx]
4574

4675

47-
def _apply_H_operator_numba(H, t, x, y):
76+
@njit(cache=True)
77+
def _scal_numba(scale, x):
78+
for idx in range(x.shape[0]):
79+
x[idx] *= scale
80+
81+
82+
def _apply_H_operator_numba(H, t, x, y, alpha, beta):
4883
raise NotImplementedError
4984

5085

5186
@overload(_apply_H_operator_numba)
52-
def _overload_apply_H_operator_numba(H, t, x, y):
87+
def _overload_apply_H_operator_numba(H, t, x, y, alpha, beta):
5388
if isinstance(H, types.BaseTuple) and len(H) == 3:
54-
def impl(H, t, x, y):
89+
def impl(H, t, x, y, alpha, beta):
5590
data, indices, indptr = H
5691
for row in range(indptr.shape[0] - 1):
5792
acc = 0.0 + 0.0j
5893
for pos in range(indptr[row], indptr[row + 1]):
5994
acc += data[pos] * x[indices[pos]]
60-
y[row] = acc
95+
y[row] = alpha * acc + beta * y[row]
6196

6297
elif isinstance(H, types.Array):
63-
def impl(H, t, x, y):
98+
def impl(H, t, x, y, alpha, beta):
6499
charN = np.uint8(ord("N"))
65-
zgemm(charN, charN, H.shape[0], 1, H.shape[1], 1.0 + 0.0j, H,
66-
H.shape[0], x, x.shape[0], 0.0 + 0.0j, y, y.shape[0])
100+
zgemm(charN, charN, H.shape[0], 1, H.shape[1], alpha, H,
101+
H.shape[0], x, x.shape[0], beta, y, y.shape[0])
102+
103+
elif isinstance(H, types.BaseTuple) and len(H) == 2:
104+
def impl(H, t, x, y, alpha, beta):
105+
H0, H_terms = H
106+
_apply_H_operator_numba(H0, t, x, y, alpha, beta)
107+
for term in literal_unroll(H_terms):
108+
Hk, coeff_spec = term
109+
coeff = _evaluate_operator_coefficient_numba(coeff_spec, t)
110+
_apply_H_operator_numba(Hk, t, x, y, alpha * coeff, 1.0 + 0.0j)
67111

68112
elif isinstance(H, types.Integer):
69-
def impl(H, t, x, y):
113+
def impl(H, t, x, y, alpha, beta):
114+
if alpha != 1.0 or beta != 0.0:
115+
raise ValueError("Scaling not supported for callable operators in Numba Lanczos backend")
70116
with objmode():
71117
_CALLABLE_OPERATOR_REGISTRY[H](t, x, y)
72118

73119
else:
74-
raise TypeError("Numba Lanczos operator must be dense array, CSR tuple, or callable handle")
120+
raise TypeError("Numba Lanczos operator must be dense array, CSR tuple, specialized sum tuple, or callable handle")
75121

76122
return impl
77123

@@ -115,7 +161,7 @@ def _step_numba(operator, t, HT, config, scratch):
115161

116162
for step in range(1, max_lanczos_steps + 1):
117163
step_count = step
118-
_apply_H_operator_numba(operator, t, phia[step - 1], phia[step])
164+
_apply_H_operator_numba(operator, t, phia[step - 1], phia[step], 1.0 + 0.0j, 0.0 + 0.0j)
119165
prefacs[step] = prefacs[step - 1]
120166
phinorm = prefacs[step] * _vnorm_numba(phia[step])
121167

@@ -138,12 +184,12 @@ def _step_numba(operator, t, HT, config, scratch):
138184
beta[step - 1] = prefacs[step] * phinorm
139185
if phinorm <= breakdown_tol:
140186
prefacs[step] = 1.0
141-
phia[step] = 0.
187+
phia[step][:] = 0.0 + 0.0j
142188
exact_complete = True
143189
else:
144190
prefacs[step] = 1.0 / phinorm
145191
if abs(np.log10(prefacs[step])) > 4.0:
146-
phia[step] *= prefacs[step]
192+
_scal_numba(prefacs[step], phia[step])
147193
prefacs[step] = 1.0
148194

149195
prev_coeff[:] = curr_coeff
@@ -205,6 +251,27 @@ def _normalize_static_operator(H):
205251
return (H_dense.shape[0], H_dense)
206252

207253

254+
def _normalize_sum_operator(H):
255+
dim, H0 = _normalize_static_operator(H[0])
256+
H_terms = []
257+
handles = []
258+
for term in H[1:]:
259+
if not isinstance(term, (tuple, list)) or len(term) != 2:
260+
raise TypeError("Time-dependent Numba operator terms must be (H_k, f_k) pairs")
261+
Hk, fk = term
262+
Hk_dim, Hk_norm = _normalize_static_operator(Hk)
263+
if Hk_dim != dim:
264+
raise ValueError("All H_k operators must match the dimension of H_0")
265+
if isinstance(fk, Dispatcher):
266+
coeff_spec = fk
267+
else:
268+
coeff_spec = _register_callable_operator(fk)
269+
handles.append(coeff_spec)
270+
H_terms.append((Hk_norm, coeff_spec))
271+
272+
return dim, (H0, tuple(H_terms)), tuple(handles)
273+
274+
208275
def _allocate_scratch(maxsteps, dim):
209276
return (
210277
np.empty(maxsteps + 1, dtype=np.float64), # alpha diagonal terms
@@ -231,18 +298,25 @@ def __init__(self, H, maxsteps, target_convg, debug=0, do_full_order=False):
231298
self.breakdown_tol = 1e-14
232299
self.config = (maxsteps, target_convg, do_full_order, self.breakdown_tol)
233300
self.H = H
301+
self._registry_handles = []
234302

235-
if callable(H):
303+
if _is_numba_sum_operator_input(H):
304+
self.dim, self.operator, handles = _normalize_sum_operator(H)
305+
self._registry_handles.extend(handles)
306+
self.scratch = _allocate_scratch(maxsteps, self.dim)
307+
elif callable(H):
236308
self.operator = _register_callable_operator(H)
309+
self._registry_handles.append(self.operator)
237310
self.dim = None
238311
self.scratch = None
239312
else:
240313
self.dim, self.operator = _normalize_static_operator(H)
241314
self.scratch = _allocate_scratch(maxsteps, self.dim)
242315

243316
def __del__(self):
244-
if callable(self.H) and self.operator in _CALLABLE_OPERATOR_REGISTRY:
245-
del _CALLABLE_OPERATOR_REGISTRY[self.operator]
317+
for handle in self._registry_handles:
318+
if handle in _CALLABLE_OPERATOR_REGISTRY:
319+
del _CALLABLE_OPERATOR_REGISTRY[handle]
246320

247321
def propagate(self, phi0, ts, maxHT=None):
248322
phi0 = np.asarray(phi0, dtype=np.complex128)
@@ -268,4 +342,13 @@ def propagate(self, phi0, ts, maxHT=None):
268342

269343
_propagate_numba(self.operator, ts, use_maxht, maxht_value, self.config, self.scratch, out)
270344

271-
return out
345+
return out
346+
347+
348+
def _is_numba_sum_operator_input(H):
349+
if not isinstance(H, (tuple, list)) or len(H) < 2:
350+
return False
351+
for term in H[1:]:
352+
if not isinstance(term, (tuple, list)) or len(term) != 2 or not callable(term[1]):
353+
return False
354+
return True

‎tests/test_all.py‎

Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
import numpy as np
44
import pytest
5+
from numba import njit
56
from scipy.sparse import diags
67
from scipy.sparse.linalg import LinearOperator, expm_multiply
78

@@ -265,6 +266,94 @@ def Hfun(t, phi, Hphi):
265266
assert np.allclose(phi, phi_ref, rtol=2e-4, atol=1e-6)
266267

267268

269+
def test_short_iterative_lanczos_numba_h0_sum_fk_hk_dense_exact():
270+
if not jftools.short_iterative_lanczos.have_numba_backend:
271+
pytest.skip("numba backend not available")
272+
273+
n = 3
274+
h0_diag = np.array([0.8, -0.6, 0.2], dtype=float)
275+
h1_diag = np.array([0.4, 0.1, -0.3], dtype=float)
276+
h2_diag = np.array([-0.15, 0.25, 0.05], dtype=float)
277+
omega1 = 2.3
278+
omega2 = 1.7
279+
phi0 = _normalized_random_state(n, seed=53)
280+
ts = np.linspace(0.0, 0.6, 5)
281+
282+
H0 = np.diag(h0_diag).astype(complex)
283+
H1 = np.diag(h1_diag).astype(complex)
284+
H2 = np.diag(h2_diag).astype(complex)
285+
286+
@njit(cache=True)
287+
def f1(t):
288+
return np.cos(omega1 * t)
289+
290+
@njit(cache=True)
291+
def f2(t):
292+
return np.sin(omega2 * t)
293+
294+
H_numba = (H0, (H1, f1), (H2, f2))
295+
prop = jftools.short_iterative_lanczos.lanczos_timeprop(H_numba, maxsteps=8, target_convg=1e-13, backend="auto")
296+
assert prop.backend == "numba"
297+
298+
out = prop.propagate(phi0, ts, maxHT=2e-4)
299+
300+
for t, phi in zip(ts, out):
301+
theta = h0_diag * t
302+
theta += (h1_diag / omega1) * np.sin(omega1 * t)
303+
theta += (h2_diag / omega2) * (1.0 - np.cos(omega2 * t))
304+
phi_ref = np.exp(-1j * theta) * phi0
305+
assert np.allclose(phi, phi_ref, rtol=2e-4, atol=1e-6)
306+
307+
308+
def test_short_iterative_lanczos_numba_h0_sum_fk_hk_csr_matches_callable():
309+
if not jftools.short_iterative_lanczos.have_numba_backend:
310+
pytest.skip("numba backend not available")
311+
312+
n = 20
313+
H0 = _make_chain_hamiltonian(n).astype(complex)
314+
H1 = diags([np.linspace(-1.0, 1.0, n)], [0], shape=(n, n), format="csr").astype(complex)
315+
H2 = diags([0.2 * np.ones(n - 1), 0.2 * np.ones(n - 1)], [-1, 1], shape=(n, n), format="csr").astype(complex)
316+
phi0 = _normalized_random_state(n, seed=54)
317+
ts = np.linspace(0.0, 0.4, 5)
318+
319+
@njit(cache=True)
320+
def f1(t):
321+
return np.cos(1.9 * t)
322+
323+
@njit(cache=True)
324+
def f2(t):
325+
return 0.4 * np.sin(0.8 * t)
326+
327+
def Hfun(t, phi, Hphi):
328+
Hphi[:] = H0.dot(phi)
329+
Hphi[:] += np.cos(1.9 * t) * H1.dot(phi)
330+
Hphi[:] += 0.4 * np.sin(0.8 * t) * H2.dot(phi)
331+
return Hphi
332+
333+
H_numba = (H0, (H1, f1), (H2, f2))
334+
out_numba = jftools.short_iterative_lanczos.sesolve_lanczos(
335+
H_numba,
336+
phi0,
337+
ts,
338+
maxsteps=14,
339+
target_convg=1e-12,
340+
maxHT=0.05,
341+
backend="numba",
342+
)
343+
out_python = jftools.short_iterative_lanczos.sesolve_lanczos(
344+
Hfun,
345+
phi0,
346+
ts,
347+
maxsteps=14,
348+
target_convg=1e-12,
349+
maxHT=0.05,
350+
backend="python",
351+
)
352+
353+
for phi_numba, phi_python in zip(out_numba, out_python):
354+
assert np.allclose(phi_numba, phi_python, rtol=5e-9, atol=5e-10)
355+
356+
268357
@pytest.mark.parametrize("backend", ["python", "numba", "cython"])
269358
def test_short_iterative_lanczos_full_basis_completion_stops_cleanly(backend):
270359
if backend == "numba" and not jftools.short_iterative_lanczos.have_numba_backend:

0 commit comments

Comments
 (0)