11import numpy as np
2- from numba import njit , objmode
2+ from numba import literal_unroll , njit , objmode
33from numba .core import types
4+ from numba .core .dispatcher import Dispatcher
45from numba .extending import overload
56from numba_lapack import dstevd , zgemm
67from scipy import sparse as sp
78
89_CALLABLE_OPERATOR_REGISTRY = {}
910
11+
1012def _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 )
1746def _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+
208275def _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
0 commit comments