Skip to content

Commit 56a13d3

Browse files
improve
1 parent c3ef7fd commit 56a13d3

2 files changed

Lines changed: 19 additions & 5 deletions

File tree

src/rydstate/radial/radial_ket.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -525,8 +525,11 @@ def calc_matrix_element(
525525
The radial matrix element in the desired unit.
526526
527527
"""
528+
ket1, ket2 = self, other
529+
if hash(self) > hash(other):
530+
ket1, ket2 = other, self
528531
radial_matrix_element_au = calc_matrix_element_au_cached(
529-
self, other, k_radial, integration_method=integration_method
532+
ket1, ket2, k_radial, integration_method=integration_method
530533
)
531534

532535
if unit == "a.u.":
@@ -537,7 +540,7 @@ def calc_matrix_element(
537540
return radial_matrix_element.to(unit).magnitude
538541

539542

540-
@lru_cache(maxsize=100_000)
543+
@lru_cache(maxsize=1000_000)
541544
def calc_matrix_element_au_cached(
542545
ket1: RadialKet,
543546
ket2: RadialKet,

src/rydstate/radial/radial_matrix_element.py

Lines changed: 14 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
import numpy as np
77
import scipy.integrate
8+
from numba import njit
89

910
if 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

Comments
 (0)