55
66import numpy as np
77import scipy .integrate
8+ from numba import njit
89
910if TYPE_CHECKING :
1011 from rydstate .units import NDArray
@@ -49,7 +50,9 @@ def calc_radial_matrix_element_from_w_z(
4950
5051 """
5152 # Find overlapping grid range
52- z_min = max (z1 [0 ], z2 [0 ])
53+ min_ind1 = first_nonzero (w1 )
54+ min_ind2 = first_nonzero (w2 )
55+ z_min = max (z1 [min_ind1 ], z2 [min_ind2 ])
5356 z_max = min (z1 [- 1 ], z2 [- 1 ])
5457 if z_max <= z_min :
5558 logger .debug ("No overlapping grid points between states, returning 0" )
@@ -61,7 +64,7 @@ def calc_radial_matrix_element_from_w_z(
6164 ind = int ((z_min - z1 [0 ]) / dz + 0.5 )
6265 z1 = z1 [ind :]
6366 w1 = w1 [ind :]
64- elif z2 [0 ] < z_min - dz / 2 :
67+ if z2 [0 ] < z_min - dz / 2 :
6568 ind = int ((z_min - z2 [0 ]) / dz + 0.5 )
6669 z2 = z2 [ind :]
6770 w2 = w2 [ind :]
@@ -70,7 +73,7 @@ def calc_radial_matrix_element_from_w_z(
7073 ind = int ((z1 [- 1 ] - z_max ) / dz + 0.5 )
7174 z1 = z1 [:- ind ]
7275 w1 = w1 [:- ind ]
73- elif z2 [- 1 ] > z_max + dz / 2 :
76+ if z2 [- 1 ] > z_max + dz / 2 :
7477 ind = int ((z2 [- 1 ] - z_max ) / dz + 0.5 )
7578 z2 = z2 [:- ind ]
7679 w2 = w2 [:- ind ]
@@ -120,3 +123,11 @@ def _integrate(integrand: NDArray, dz: float, method: INTEGRATION_METHODS) -> fl
120123 raise ValueError (f"Invalid integration method: { method } " )
121124
122125 return float (value )
126+
127+
128+ @njit (cache = True )
129+ def first_nonzero (a : np .ndarray ) -> int :
130+ for i in range (a .size ):
131+ if a [i ] != 0 :
132+ return i
133+ return - 1
0 commit comments