von_neumann_transform 0.2.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- von_neumann_transform/__init__.py +13 -0
- von_neumann_transform/basis.py +57 -0
- von_neumann_transform/methods.py +30 -0
- von_neumann_transform/overlap.py +445 -0
- von_neumann_transform/projection.py +86 -0
- von_neumann_transform/reconstruction.py +80 -0
- von_neumann_transform/transform.py +497 -0
- von_neumann_transform-0.2.0.dist-info/METADATA +483 -0
- von_neumann_transform-0.2.0.dist-info/RECORD +11 -0
- von_neumann_transform-0.2.0.dist-info/WHEEL +4 -0
- von_neumann_transform-0.2.0.dist-info/licenses/LICENSE +201 -0
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
from .methods import BasisMethod, MatVecMethod, PrecondMethod, SolverMethod
|
|
2
|
+
from .transform import VonNeumannTransform
|
|
3
|
+
|
|
4
|
+
# optional: define an explicit public API
|
|
5
|
+
__all__ = [
|
|
6
|
+
"BasisMethod",
|
|
7
|
+
"MatVecMethod",
|
|
8
|
+
"PrecondMethod",
|
|
9
|
+
"SolverMethod",
|
|
10
|
+
"VonNeumannTransform",
|
|
11
|
+
]
|
|
12
|
+
|
|
13
|
+
version = "0.2.0"
|
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def _get_grid(
|
|
5
|
+
npoints: int,
|
|
6
|
+
w_min: float,
|
|
7
|
+
w_max: float,
|
|
8
|
+
) -> tuple[np.ndarray, float, np.ndarray, np.ndarray, int, float]:
|
|
9
|
+
|
|
10
|
+
# trimmed grid for signal in the frequency domain
|
|
11
|
+
w_grid, dw_grid = np.linspace(
|
|
12
|
+
w_min,
|
|
13
|
+
w_max,
|
|
14
|
+
npoints,
|
|
15
|
+
retstep=True,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
# grid in the von Neumann plane
|
|
19
|
+
w_span = w_grid.max() - w_grid.min()
|
|
20
|
+
t_span = 2.0 * np.pi / dw_grid
|
|
21
|
+
k = round(np.sqrt(npoints))
|
|
22
|
+
if k**2 != npoints:
|
|
23
|
+
raise ValueError(
|
|
24
|
+
"Number of points must be a perfect square "
|
|
25
|
+
"(k^2) for the von Neumann transform."
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
dw = w_span / k
|
|
29
|
+
dt = t_span / k
|
|
30
|
+
w_n_arr = w_min + (np.arange(k) + 0.5) * dw
|
|
31
|
+
t_n_arr = -t_span / 2.0 + (np.arange(k) + 0.5) * dt
|
|
32
|
+
|
|
33
|
+
# width of the basis functions
|
|
34
|
+
alpha = t_span / (2.0 * w_span)
|
|
35
|
+
|
|
36
|
+
return w_grid, t_span, w_n_arr, t_n_arr, k, alpha
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _evaluate_basis_functions(
|
|
40
|
+
w_grid: np.ndarray,
|
|
41
|
+
w_n_arr: np.ndarray,
|
|
42
|
+
t_n_arr: np.ndarray,
|
|
43
|
+
alpha: float,
|
|
44
|
+
) -> np.ndarray:
|
|
45
|
+
k = len(w_n_arr)
|
|
46
|
+
norm = (2.0 * alpha / np.pi) ** 0.25
|
|
47
|
+
|
|
48
|
+
alpha_nmo = np.zeros((k, k, k * k), dtype=np.complex128)
|
|
49
|
+
for i in range(k):
|
|
50
|
+
for j in range(k):
|
|
51
|
+
alpha_nmo[i, j] = np.exp(
|
|
52
|
+
-alpha * (w_grid - w_n_arr[i]) ** 2
|
|
53
|
+
- 1.0j * t_n_arr[j] * (w_grid - w_n_arr[i]),
|
|
54
|
+
)
|
|
55
|
+
alpha_nmo *= norm
|
|
56
|
+
|
|
57
|
+
return alpha_nmo
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
from enum import Enum, auto
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class BasisMethod(Enum):
|
|
5
|
+
DIRECT = auto()
|
|
6
|
+
FACTORISE = auto()
|
|
7
|
+
FFT = auto()
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MatVecMethod(Enum):
|
|
11
|
+
DIRECT = auto()
|
|
12
|
+
TOEPLITZ_MATMUL = auto()
|
|
13
|
+
TOEPLITZ_EINSUM = auto()
|
|
14
|
+
TOEPLITZ_BANDED = auto()
|
|
15
|
+
GAUSSIAN_STENCIL = auto()
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class PrecondMethod(Enum):
|
|
19
|
+
AUTO = auto()
|
|
20
|
+
NONE = auto()
|
|
21
|
+
CIRCULANT_DENSE = auto()
|
|
22
|
+
CIRCULANT_BANDED = auto()
|
|
23
|
+
INCOMPLETE_CHOLESKY = auto()
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class SolverMethod(Enum):
|
|
27
|
+
DIRECT = auto()
|
|
28
|
+
CG = auto()
|
|
29
|
+
BICGSTAB = auto()
|
|
30
|
+
LGMRES = auto()
|
|
@@ -0,0 +1,445 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from scipy.linalg import cho_solve_banded, cholesky_banded, solve_triangular
|
|
3
|
+
from scipy.sparse import coo_matrix, csr_matrix, tril
|
|
4
|
+
from scipy.sparse.linalg import (
|
|
5
|
+
LinearOperator,
|
|
6
|
+
aslinearoperator,
|
|
7
|
+
spsolve_triangular,
|
|
8
|
+
)
|
|
9
|
+
|
|
10
|
+
from .methods import MatVecMethod, PrecondMethod
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _chol_solve_batch(l_mat: np.ndarray, b_mat: np.ndarray) -> None:
|
|
14
|
+
solve_triangular(
|
|
15
|
+
l_mat,
|
|
16
|
+
b_mat,
|
|
17
|
+
lower=True,
|
|
18
|
+
trans=0,
|
|
19
|
+
overwrite_b=True,
|
|
20
|
+
check_finite=False,
|
|
21
|
+
) # forward
|
|
22
|
+
solve_triangular(
|
|
23
|
+
l_mat.conj().swapaxes(-1, -2),
|
|
24
|
+
b_mat,
|
|
25
|
+
lower=False,
|
|
26
|
+
trans=0,
|
|
27
|
+
overwrite_b=True,
|
|
28
|
+
check_finite=False,
|
|
29
|
+
) # backward
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _get_ovlp_direct(
|
|
33
|
+
alpha: float,
|
|
34
|
+
w_n_arr: np.ndarray,
|
|
35
|
+
t_n_arr: np.ndarray,
|
|
36
|
+
) -> np.ndarray:
|
|
37
|
+
dw = w_n_arr[1] - w_n_arr[0]
|
|
38
|
+
dt = t_n_arr[1] - t_n_arr[0]
|
|
39
|
+
k = len(w_n_arr)
|
|
40
|
+
|
|
41
|
+
# construct distance matrix
|
|
42
|
+
small_dist_mat = np.zeros((k, k), dtype=np.complex128)
|
|
43
|
+
for i in range(1, k):
|
|
44
|
+
np.fill_diagonal(small_dist_mat[:, i:], i)
|
|
45
|
+
small_dist_mat += -small_dist_mat.T
|
|
46
|
+
|
|
47
|
+
# construct sum matrix
|
|
48
|
+
small_sum_mat = np.add.outer(
|
|
49
|
+
np.arange(k, dtype=np.complex128),
|
|
50
|
+
np.arange(k, dtype=np.complex128),
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
tmp = (
|
|
54
|
+
-0.5
|
|
55
|
+
* alpha
|
|
56
|
+
* (
|
|
57
|
+
np.repeat(np.repeat(small_dist_mat**2, k, axis=0), k, axis=1)
|
|
58
|
+
* dw**2
|
|
59
|
+
)
|
|
60
|
+
)
|
|
61
|
+
tmp += -(1.0 / (8.0 * alpha)) * (
|
|
62
|
+
np.tile(small_dist_mat**2, (k, k)) * dt**2
|
|
63
|
+
)
|
|
64
|
+
tmp += (
|
|
65
|
+
0.5j
|
|
66
|
+
* (np.repeat(np.repeat(small_dist_mat, k, axis=0), k, axis=1) * dw)
|
|
67
|
+
* (np.tile(small_sum_mat, (k, k)) * dt + 2.0 * t_n_arr.min())
|
|
68
|
+
)
|
|
69
|
+
s = np.exp(tmp)
|
|
70
|
+
|
|
71
|
+
return s
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def _get_ovlp_block(
|
|
75
|
+
alpha: float,
|
|
76
|
+
w_n_arr: np.ndarray,
|
|
77
|
+
t_n_arr: np.ndarray,
|
|
78
|
+
) -> np.ndarray:
|
|
79
|
+
k = len(w_n_arr)
|
|
80
|
+
dw = w_n_arr[1] - w_n_arr[0]
|
|
81
|
+
dt = t_n_arr[1] - t_n_arr[0]
|
|
82
|
+
|
|
83
|
+
# matrices that depend only on *within-block* indices m, n
|
|
84
|
+
idx = np.arange(k, dtype=np.complex128)
|
|
85
|
+
diff_mat = idx[:, np.newaxis] - idx[np.newaxis, :] # m - n
|
|
86
|
+
diff2_mat = diff_mat**2
|
|
87
|
+
sum_mat = idx[:, np.newaxis] + idx[np.newaxis, :] # m + n
|
|
88
|
+
|
|
89
|
+
col_blocks = np.empty((k, k, k), dtype=np.complex128)
|
|
90
|
+
for j in range(k): # j = 0 … k−1 ⇒ block column index
|
|
91
|
+
col_blocks[j] = _get_ovlp_block_values(
|
|
92
|
+
alpha, dw, dt, t_n_arr.min(), j, diff2_mat, sum_mat
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
return col_blocks
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def _get_ovlp_block_values(
|
|
99
|
+
alpha: float,
|
|
100
|
+
dw: float,
|
|
101
|
+
dt: float,
|
|
102
|
+
t_min: float,
|
|
103
|
+
block_offset: int,
|
|
104
|
+
inner_diff2: np.ndarray | int,
|
|
105
|
+
inner_sum: np.ndarray,
|
|
106
|
+
) -> np.ndarray:
|
|
107
|
+
"""Evaluate entries of a block from its outer and inner offsets."""
|
|
108
|
+
exponent = (
|
|
109
|
+
-0.5 * alpha * block_offset**2 * dw**2
|
|
110
|
+
- (1.0 / (8.0 * alpha)) * inner_diff2 * dt**2
|
|
111
|
+
+ 0.5j * (-block_offset * dw) * (inner_sum * dt + 2.0 * t_min)
|
|
112
|
+
)
|
|
113
|
+
return np.exp(exponent)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def _get_ovlp_fft_band(
|
|
117
|
+
alpha: float,
|
|
118
|
+
w_n_arr: np.ndarray,
|
|
119
|
+
t_n_arr: np.ndarray,
|
|
120
|
+
bandwidth: int = 4,
|
|
121
|
+
) -> np.ndarray:
|
|
122
|
+
"""Lower bands of the real Fourier-domain circulant blocks.
|
|
123
|
+
|
|
124
|
+
The array uses SciPy's lower banded layout: entry [mode, offset, n]
|
|
125
|
+
represents the Fourier block entry [n + offset, n].
|
|
126
|
+
"""
|
|
127
|
+
k = len(w_n_arr)
|
|
128
|
+
nc = 2 * k
|
|
129
|
+
bandwidth = min(bandwidth, k - 1)
|
|
130
|
+
dw = w_n_arr[1] - w_n_arr[0]
|
|
131
|
+
dt = t_n_arr[1] - t_n_arr[0]
|
|
132
|
+
|
|
133
|
+
band_fft = np.zeros((nc, bandwidth + 1, k), dtype=np.float64)
|
|
134
|
+
for offset in range(bandwidth + 1):
|
|
135
|
+
n = np.arange(k - offset)
|
|
136
|
+
m = n + offset
|
|
137
|
+
circ = np.zeros((nc, k - offset), dtype=np.complex128)
|
|
138
|
+
for d in range(k):
|
|
139
|
+
circ[d] = _get_ovlp_block_values(
|
|
140
|
+
alpha, dw, dt, t_n_arr.min(), d, offset**2, m + n
|
|
141
|
+
)
|
|
142
|
+
circ[k + 1 :] = circ[1:k][::-1].conj()
|
|
143
|
+
# The signed block offsets pair by conjugation, so the FFT is real.
|
|
144
|
+
band_fft[:, offset, : k - offset] = np.fft.fft(circ, axis=0).real
|
|
145
|
+
|
|
146
|
+
return band_fft
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
def _get_ovlp_fft_dense(
|
|
150
|
+
alpha: float,
|
|
151
|
+
w_n_arr: np.ndarray,
|
|
152
|
+
t_n_arr: np.ndarray,
|
|
153
|
+
) -> np.ndarray:
|
|
154
|
+
k = len(w_n_arr)
|
|
155
|
+
blocks = _get_ovlp_block(alpha, w_n_arr, t_n_arr)
|
|
156
|
+
circulant = np.zeros((2 * k, k, k), dtype=np.complex128)
|
|
157
|
+
circulant[:k] = blocks
|
|
158
|
+
circulant[k + 1 :] = blocks[1:][::-1].transpose((0, 2, 1)).conj()
|
|
159
|
+
return np.fft.fft(circulant, axis=0)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def _get_ovlp_stencil(
|
|
163
|
+
alpha: float,
|
|
164
|
+
w_n_arr: np.ndarray,
|
|
165
|
+
t_n_arr: np.ndarray,
|
|
166
|
+
radius: int = 4,
|
|
167
|
+
) -> csr_matrix:
|
|
168
|
+
"""Store the exact overlap entries within a Gaussian R-neighborhood.
|
|
169
|
+
|
|
170
|
+
The outer and inner offsets are both truncated to ``[-radius, radius]``.
|
|
171
|
+
No circulant embedding or FFT is involved in the resulting matvec.
|
|
172
|
+
"""
|
|
173
|
+
k = len(w_n_arr)
|
|
174
|
+
radius = min(radius, k - 1)
|
|
175
|
+
dw = w_n_arr[1] - w_n_arr[0]
|
|
176
|
+
dt = t_n_arr[1] - t_n_arr[0]
|
|
177
|
+
t_min = t_n_arr.min()
|
|
178
|
+
|
|
179
|
+
offsets = range(-radius, radius + 1)
|
|
180
|
+
count_per_axis = sum(k - abs(offset) for offset in offsets)
|
|
181
|
+
nnz = count_per_axis**2
|
|
182
|
+
index_dtype = np.int32 if k * k <= np.iinfo(np.int32).max else np.int64
|
|
183
|
+
rows = np.empty(nnz, dtype=index_dtype)
|
|
184
|
+
cols = np.empty(nnz, dtype=index_dtype)
|
|
185
|
+
values = np.empty(nnz, dtype=np.complex128)
|
|
186
|
+
cursor = 0
|
|
187
|
+
for d in range(-radius, radius + 1):
|
|
188
|
+
p_start, p_stop = max(d, 0), min(k + d, k)
|
|
189
|
+
p_count = p_stop - p_start
|
|
190
|
+
p = np.arange(p_start, p_stop)
|
|
191
|
+
for ell in range(-radius, radius + 1):
|
|
192
|
+
m_start, m_stop = max(ell, 0), min(k + ell, k)
|
|
193
|
+
m_count = m_stop - m_start
|
|
194
|
+
m = np.arange(m_start, m_stop)
|
|
195
|
+
n = m - ell
|
|
196
|
+
weights = _get_ovlp_block_values(
|
|
197
|
+
alpha, dw, dt, t_min, d, ell**2, m + n
|
|
198
|
+
)
|
|
199
|
+
next_cursor = cursor + p_count * m_count
|
|
200
|
+
row_view = rows[cursor:next_cursor].reshape(p_count, m_count)
|
|
201
|
+
row_view[:] = p[:, None] * k + m
|
|
202
|
+
cols[cursor:next_cursor] = rows[cursor:next_cursor] - d * k - ell
|
|
203
|
+
values[cursor:next_cursor].reshape(p_count, m_count)[:] = weights
|
|
204
|
+
cursor = next_cursor
|
|
205
|
+
|
|
206
|
+
matrix = coo_matrix(
|
|
207
|
+
(values, (rows, cols)),
|
|
208
|
+
shape=(k * k, k * k),
|
|
209
|
+
).tocsr()
|
|
210
|
+
matrix.sort_indices()
|
|
211
|
+
return matrix
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
def _get_ic0_factor(matrix: csr_matrix) -> csr_matrix:
|
|
215
|
+
"""Incomplete Cholesky factor with no fill beyond the lower pattern.
|
|
216
|
+
|
|
217
|
+
The natural row-major ordering of the Gaussian stencil is retained.
|
|
218
|
+
A non-positive pivot is reported rather than silently changing the
|
|
219
|
+
preconditioner with a diagonal shift.
|
|
220
|
+
"""
|
|
221
|
+
lower = tril(matrix, format="csr")
|
|
222
|
+
lower.sort_indices()
|
|
223
|
+
indices = lower.indices
|
|
224
|
+
indptr = lower.indptr
|
|
225
|
+
values = lower.data.copy()
|
|
226
|
+
|
|
227
|
+
for i in range(lower.shape[0]):
|
|
228
|
+
start, stop = indptr[i : i + 2]
|
|
229
|
+
if start == stop or indices[stop - 1] != i:
|
|
230
|
+
raise np.linalg.LinAlgError(
|
|
231
|
+
f"IC(0) requires a diagonal entry in row {i}."
|
|
232
|
+
)
|
|
233
|
+
diagonal_position = stop - 1
|
|
234
|
+
positions = {
|
|
235
|
+
int(indices[position]): position
|
|
236
|
+
for position in range(start, diagonal_position)
|
|
237
|
+
}
|
|
238
|
+
|
|
239
|
+
for position in range(start, diagonal_position):
|
|
240
|
+
j = int(indices[position])
|
|
241
|
+
correction = 0j
|
|
242
|
+
for prior in range(indptr[j], indptr[j + 1] - 1):
|
|
243
|
+
matching = positions.get(int(indices[prior]))
|
|
244
|
+
if matching is not None:
|
|
245
|
+
correction += values[matching] * values[prior].conjugate()
|
|
246
|
+
values[position] = (values[position] - correction) / values[
|
|
247
|
+
indptr[j + 1] - 1
|
|
248
|
+
]
|
|
249
|
+
|
|
250
|
+
diagonal = values[diagonal_position]
|
|
251
|
+
pivot = (
|
|
252
|
+
diagonal.real
|
|
253
|
+
- np.vdot(
|
|
254
|
+
values[start:diagonal_position],
|
|
255
|
+
values[start:diagonal_position],
|
|
256
|
+
).real
|
|
257
|
+
)
|
|
258
|
+
if not np.isfinite(pivot) or pivot <= 0:
|
|
259
|
+
raise np.linalg.LinAlgError(
|
|
260
|
+
f"IC(0) failed at row {i}: non-positive pivot {pivot}."
|
|
261
|
+
)
|
|
262
|
+
values[diagonal_position] = np.sqrt(pivot)
|
|
263
|
+
|
|
264
|
+
return csr_matrix((values, indices, indptr), shape=lower.shape)
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def _get_ovlp_linop(
|
|
268
|
+
alpha: float,
|
|
269
|
+
w_n_arr: np.ndarray,
|
|
270
|
+
t_n_arr: np.ndarray,
|
|
271
|
+
matvec_method: MatVecMethod = MatVecMethod.TOEPLITZ_MATMUL,
|
|
272
|
+
precond_method: PrecondMethod = PrecondMethod.AUTO,
|
|
273
|
+
) -> tuple[LinearOperator, LinearOperator]:
|
|
274
|
+
k = len(w_n_arr)
|
|
275
|
+
nc = 2 * k
|
|
276
|
+
|
|
277
|
+
if matvec_method is MatVecMethod.DIRECT:
|
|
278
|
+
raise RuntimeError(
|
|
279
|
+
"DIRECT method for matrix-vector multiplication "
|
|
280
|
+
"can only be used with the get_ovlp_direct method."
|
|
281
|
+
)
|
|
282
|
+
|
|
283
|
+
if precond_method is PrecondMethod.AUTO:
|
|
284
|
+
if matvec_method in (
|
|
285
|
+
MatVecMethod.TOEPLITZ_MATMUL,
|
|
286
|
+
MatVecMethod.TOEPLITZ_EINSUM,
|
|
287
|
+
):
|
|
288
|
+
precond_method = PrecondMethod.CIRCULANT_DENSE
|
|
289
|
+
elif matvec_method is MatVecMethod.TOEPLITZ_BANDED:
|
|
290
|
+
precond_method = PrecondMethod.CIRCULANT_BANDED
|
|
291
|
+
elif matvec_method is MatVecMethod.GAUSSIAN_STENCIL:
|
|
292
|
+
precond_method = PrecondMethod.INCOMPLETE_CHOLESKY
|
|
293
|
+
else:
|
|
294
|
+
raise ValueError(f"Unknown matvec method: {matvec_method!r}")
|
|
295
|
+
elif precond_method not in (
|
|
296
|
+
PrecondMethod.NONE,
|
|
297
|
+
PrecondMethod.CIRCULANT_DENSE,
|
|
298
|
+
PrecondMethod.CIRCULANT_BANDED,
|
|
299
|
+
PrecondMethod.INCOMPLETE_CHOLESKY,
|
|
300
|
+
):
|
|
301
|
+
raise ValueError(f"Unknown preconditioner method: {precond_method!r}")
|
|
302
|
+
|
|
303
|
+
need_dense_fft = (
|
|
304
|
+
matvec_method
|
|
305
|
+
in (
|
|
306
|
+
MatVecMethod.TOEPLITZ_MATMUL,
|
|
307
|
+
MatVecMethod.TOEPLITZ_EINSUM,
|
|
308
|
+
)
|
|
309
|
+
or precond_method is PrecondMethod.CIRCULANT_DENSE
|
|
310
|
+
)
|
|
311
|
+
need_band_fft = (
|
|
312
|
+
matvec_method is MatVecMethod.TOEPLITZ_BANDED
|
|
313
|
+
or precond_method is PrecondMethod.CIRCULANT_BANDED
|
|
314
|
+
)
|
|
315
|
+
s_fft = (
|
|
316
|
+
_get_ovlp_fft_dense(alpha, w_n_arr, t_n_arr)
|
|
317
|
+
if need_dense_fft
|
|
318
|
+
else None
|
|
319
|
+
)
|
|
320
|
+
s_band = (
|
|
321
|
+
_get_ovlp_fft_band(alpha, w_n_arr, t_n_arr) if need_band_fft else None
|
|
322
|
+
)
|
|
323
|
+
|
|
324
|
+
stencil = (
|
|
325
|
+
_get_ovlp_stencil(alpha, w_n_arr, t_n_arr)
|
|
326
|
+
if (
|
|
327
|
+
matvec_method is MatVecMethod.GAUSSIAN_STENCIL
|
|
328
|
+
or precond_method is PrecondMethod.INCOMPLETE_CHOLESKY
|
|
329
|
+
)
|
|
330
|
+
else None
|
|
331
|
+
)
|
|
332
|
+
|
|
333
|
+
if matvec_method is MatVecMethod.GAUSSIAN_STENCIL:
|
|
334
|
+
assert stencil is not None
|
|
335
|
+
s_op = aslinearoperator(stencil)
|
|
336
|
+
else:
|
|
337
|
+
if matvec_method is MatVecMethod.TOEPLITZ_BANDED:
|
|
338
|
+
band_fft = s_band
|
|
339
|
+
assert band_fft is not None
|
|
340
|
+
|
|
341
|
+
def contract(x, y):
|
|
342
|
+
y.fill(0.0)
|
|
343
|
+
for offset in range(band_fft.shape[1]):
|
|
344
|
+
diagonal = band_fft[:, offset, : k - offset]
|
|
345
|
+
y[:, offset:, 0] += diagonal * x[:, : k - offset, 0]
|
|
346
|
+
if offset:
|
|
347
|
+
y[:, : k - offset, 0] += diagonal * x[:, offset:, 0]
|
|
348
|
+
|
|
349
|
+
elif matvec_method is MatVecMethod.TOEPLITZ_MATMUL:
|
|
350
|
+
dense_fft = s_fft
|
|
351
|
+
assert dense_fft is not None
|
|
352
|
+
|
|
353
|
+
def contract(x, y):
|
|
354
|
+
np.matmul(dense_fft, x, out=y)
|
|
355
|
+
|
|
356
|
+
elif matvec_method is MatVecMethod.TOEPLITZ_EINSUM:
|
|
357
|
+
dense_fft = s_fft
|
|
358
|
+
assert dense_fft is not None
|
|
359
|
+
|
|
360
|
+
def contract(x, y):
|
|
361
|
+
np.einsum("kij,kjp->kip", dense_fft, x, optimize=True, out=y)
|
|
362
|
+
|
|
363
|
+
else:
|
|
364
|
+
raise ValueError(f"Unknown matvec method: {matvec_method!r}")
|
|
365
|
+
|
|
366
|
+
x_pad = np.zeros((nc, k), dtype=np.complex128)
|
|
367
|
+
x_hat = np.empty((nc, k, 1), dtype=np.complex128)
|
|
368
|
+
y_hat = np.empty((nc, k, 1), dtype=np.complex128)
|
|
369
|
+
|
|
370
|
+
def mv(x):
|
|
371
|
+
x_pad[:k] = x.reshape(k, k)
|
|
372
|
+
np.fft.fft(x_pad, axis=0, out=x_hat[..., 0])
|
|
373
|
+
contract(x_hat, y_hat)
|
|
374
|
+
return np.fft.ifft(y_hat[..., 0], axis=0)[:k].ravel()
|
|
375
|
+
|
|
376
|
+
s_op = LinearOperator(
|
|
377
|
+
(k * k, k * k), dtype=np.complex128, matvec=mv, rmatvec=mv
|
|
378
|
+
)
|
|
379
|
+
|
|
380
|
+
if precond_method is PrecondMethod.NONE:
|
|
381
|
+
m_op = LinearOperator(
|
|
382
|
+
(k * k, k * k),
|
|
383
|
+
dtype=np.complex128,
|
|
384
|
+
matvec=lambda x: x.copy(),
|
|
385
|
+
rmatvec=lambda x: x.copy(),
|
|
386
|
+
)
|
|
387
|
+
elif precond_method is PrecondMethod.INCOMPLETE_CHOLESKY:
|
|
388
|
+
assert stencil is not None
|
|
389
|
+
factor = _get_ic0_factor(stencil)
|
|
390
|
+
upper = factor.conj().T.tocsr()
|
|
391
|
+
|
|
392
|
+
def precon(r):
|
|
393
|
+
y = spsolve_triangular(factor, r, lower=True)
|
|
394
|
+
return spsolve_triangular(upper, y, lower=False)
|
|
395
|
+
|
|
396
|
+
m_op = LinearOperator(
|
|
397
|
+
(k * k, k * k),
|
|
398
|
+
dtype=np.complex128,
|
|
399
|
+
matvec=precon,
|
|
400
|
+
rmatvec=precon,
|
|
401
|
+
)
|
|
402
|
+
else:
|
|
403
|
+
if precond_method is PrecondMethod.CIRCULANT_DENSE:
|
|
404
|
+
dense_fft = s_fft
|
|
405
|
+
assert dense_fft is not None
|
|
406
|
+
cho_dense = np.linalg.cholesky(dense_fft)
|
|
407
|
+
|
|
408
|
+
def solve_precon(r):
|
|
409
|
+
_chol_solve_batch(cho_dense, r)
|
|
410
|
+
|
|
411
|
+
else:
|
|
412
|
+
band_fft = s_band
|
|
413
|
+
assert band_fft is not None
|
|
414
|
+
cho_band = np.stack(
|
|
415
|
+
[
|
|
416
|
+
cholesky_banded(band, lower=True, check_finite=False)
|
|
417
|
+
for band in band_fft
|
|
418
|
+
]
|
|
419
|
+
)
|
|
420
|
+
|
|
421
|
+
def solve_precon(r):
|
|
422
|
+
for mode in range(nc):
|
|
423
|
+
r[mode, :, 0] = cho_solve_banded(
|
|
424
|
+
(cho_band[mode], True),
|
|
425
|
+
r[mode, :, 0],
|
|
426
|
+
check_finite=False,
|
|
427
|
+
)
|
|
428
|
+
|
|
429
|
+
r_pad = np.zeros((nc, k), dtype=np.complex128)
|
|
430
|
+
r_hat = np.empty((nc, k, 1), dtype=np.complex128)
|
|
431
|
+
|
|
432
|
+
def precon(r):
|
|
433
|
+
r_pad[:k] = r.reshape(k, k)
|
|
434
|
+
np.fft.fft(r_pad, axis=0, out=r_hat[..., 0])
|
|
435
|
+
solve_precon(r_hat)
|
|
436
|
+
return np.fft.ifft(r_hat[..., 0], axis=0)[:k].ravel()
|
|
437
|
+
|
|
438
|
+
m_op = LinearOperator(
|
|
439
|
+
(k * k, k * k),
|
|
440
|
+
dtype=np.complex128,
|
|
441
|
+
matvec=precon,
|
|
442
|
+
rmatvec=precon,
|
|
443
|
+
)
|
|
444
|
+
|
|
445
|
+
return s_op, m_op
|
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def _project_signal(
|
|
5
|
+
alpha_nmo: np.ndarray,
|
|
6
|
+
signal: np.ndarray,
|
|
7
|
+
dw: float,
|
|
8
|
+
) -> np.ndarray:
|
|
9
|
+
alpha_nm = (
|
|
10
|
+
np.einsum(
|
|
11
|
+
"nmo,o->nm",
|
|
12
|
+
alpha_nmo.conj(),
|
|
13
|
+
signal,
|
|
14
|
+
optimize=True,
|
|
15
|
+
)
|
|
16
|
+
* dw
|
|
17
|
+
)
|
|
18
|
+
return alpha_nm
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _get_signal_projection_factorise(
|
|
22
|
+
w_grid: np.ndarray,
|
|
23
|
+
w_n_arr: np.ndarray,
|
|
24
|
+
t_n_arr: np.ndarray,
|
|
25
|
+
alpha: float,
|
|
26
|
+
signal: np.ndarray,
|
|
27
|
+
) -> np.ndarray:
|
|
28
|
+
dw = w_grid[1] - w_grid[0]
|
|
29
|
+
norm = (2.0 * alpha / np.pi) ** 0.25
|
|
30
|
+
|
|
31
|
+
# (k × N): each row n is the Gaussian window centered at w_n_arr[n]
|
|
32
|
+
alpha_w = norm * np.exp(-alpha * np.subtract.outer(w_n_arr, w_grid) ** 2)
|
|
33
|
+
|
|
34
|
+
# (k × N): each row m is signal(w) multiplied by
|
|
35
|
+
# the modulation for t_n_arr[m]
|
|
36
|
+
alpha_t = np.exp(1.0j * np.outer(t_n_arr, w_grid)) * signal[np.newaxis, :]
|
|
37
|
+
|
|
38
|
+
# alpha_w @ alpha_t.T is (k × k) with
|
|
39
|
+
# alpha_nm[i, m] = \sum_o alpha_w[i, o] * alpha_t[m, o]
|
|
40
|
+
alpha_nm = alpha_w @ alpha_t.T * dw
|
|
41
|
+
|
|
42
|
+
# include extra phase shift e^{-i t_m w_n}
|
|
43
|
+
phasor = np.exp(-1.0j * np.outer(w_n_arr, t_n_arr))
|
|
44
|
+
alpha_nm *= phasor
|
|
45
|
+
|
|
46
|
+
return alpha_nm
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _get_signal_projection_fft(
|
|
50
|
+
w_grid: np.ndarray,
|
|
51
|
+
w_n_arr: np.ndarray,
|
|
52
|
+
t_n_arr: np.ndarray,
|
|
53
|
+
alpha: float,
|
|
54
|
+
signal: np.ndarray,
|
|
55
|
+
) -> np.ndarray:
|
|
56
|
+
k2 = len(w_grid)
|
|
57
|
+
dw = w_grid[1] - w_grid[0]
|
|
58
|
+
norm = (2.0 * alpha / np.pi) ** 0.25
|
|
59
|
+
|
|
60
|
+
# transform signal(w) * exp(-alpha * (w - w_i)^2)
|
|
61
|
+
# via batched IFFT
|
|
62
|
+
w_diff2 = np.subtract.outer(w_n_arr, w_grid) ** 2 # (k, k2)
|
|
63
|
+
f_tmp = norm * np.exp(-alpha * w_diff2) * signal[np.newaxis, :]
|
|
64
|
+
f_tmp = np.fft.ifft(f_tmp, axis=1) * (k2 * dw) # (k, k2)
|
|
65
|
+
|
|
66
|
+
# build the full IFFT time grid
|
|
67
|
+
t_grid = 2 * np.pi * np.fft.fftfreq(k2, d=dw) # (k2,)
|
|
68
|
+
|
|
69
|
+
# the von Neumann time grid is coarser than the FFT grid,
|
|
70
|
+
# so the closest frequency bin in the FFT grid is used
|
|
71
|
+
idx_cols = np.array(
|
|
72
|
+
[np.abs(t_grid - t).argmin() for t in t_n_arr],
|
|
73
|
+
dtype=int,
|
|
74
|
+
) # (k,)
|
|
75
|
+
|
|
76
|
+
# slice out the k columns corresponding to t_n_arr
|
|
77
|
+
alpha_nm = f_tmp[:, idx_cols]
|
|
78
|
+
|
|
79
|
+
# correct for global phase shift
|
|
80
|
+
alpha_nm *= np.exp(1.0j * t_n_arr * w_grid[0])[np.newaxis, :]
|
|
81
|
+
|
|
82
|
+
# apply the final phase correction e^{-i t_m w_n}
|
|
83
|
+
phasor = np.exp(-1.0j * np.outer(w_n_arr, t_n_arr))
|
|
84
|
+
alpha_nm *= phasor
|
|
85
|
+
|
|
86
|
+
return alpha_nm
|
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def _reconstruct_signal(
|
|
5
|
+
q_nm: np.ndarray,
|
|
6
|
+
q_nmo: np.ndarray,
|
|
7
|
+
) -> np.ndarray:
|
|
8
|
+
return np.einsum("nmo,nm->o", q_nmo, q_nm, optimize=True)
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _reconstruct_signal_factorise(
|
|
12
|
+
q_nm: np.ndarray,
|
|
13
|
+
w_grid: np.ndarray,
|
|
14
|
+
w_n_arr: np.ndarray,
|
|
15
|
+
t_n_arr: np.ndarray,
|
|
16
|
+
alpha: float,
|
|
17
|
+
) -> np.ndarray:
|
|
18
|
+
norm = (2.0 * alpha / np.pi) ** 0.25
|
|
19
|
+
|
|
20
|
+
# (k × N): each row n is the Gaussian window centered at w_n_arr[n]
|
|
21
|
+
alpha_w = norm * np.exp(-alpha * np.subtract.outer(w_n_arr, w_grid) ** 2)
|
|
22
|
+
|
|
23
|
+
# (k × N): each row m is the time modulation for t_n_arr[m]
|
|
24
|
+
alpha_t = np.exp(-1.0j * np.outer(t_n_arr, w_grid))
|
|
25
|
+
|
|
26
|
+
# extra phase shift e^{i t_m w_n}
|
|
27
|
+
|
|
28
|
+
# apply the phase shift e^{i t_m w_n}
|
|
29
|
+
phasor = np.exp(1.0j * np.outer(w_n_arr, t_n_arr))
|
|
30
|
+
tmp = q_nm * phasor
|
|
31
|
+
|
|
32
|
+
# tmp @ alpha_t is (k × N) with
|
|
33
|
+
# tmp[n, m] = \sum_m q_nm[n, m] * alpha_t[m, o]
|
|
34
|
+
tmp = tmp @ alpha_t
|
|
35
|
+
|
|
36
|
+
# sum over n: signal[o] = \sum_n tmp[n, o] * alpha_w[n, o]
|
|
37
|
+
signal = np.sum(tmp * alpha_w, axis=0)
|
|
38
|
+
|
|
39
|
+
return signal
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _reconstruct_signal_fft(
|
|
43
|
+
q_nm: np.ndarray,
|
|
44
|
+
w_grid: np.ndarray,
|
|
45
|
+
w_n_arr: np.ndarray,
|
|
46
|
+
t_n_arr: np.ndarray,
|
|
47
|
+
alpha: float,
|
|
48
|
+
) -> np.ndarray:
|
|
49
|
+
k, k2 = len(w_n_arr), len(w_grid)
|
|
50
|
+
dw = w_grid[1] - w_grid[0]
|
|
51
|
+
norm = (2.0 * alpha / np.pi) ** 0.25
|
|
52
|
+
|
|
53
|
+
# apply the phase shift e^{i t_m w_n}
|
|
54
|
+
phasor = np.exp(1.0j * np.outer(w_n_arr, t_n_arr))
|
|
55
|
+
tmp = q_nm * phasor # (k, k)
|
|
56
|
+
|
|
57
|
+
# correct for global phase shift
|
|
58
|
+
tmp *= np.exp(-1.0j * t_n_arr * w_grid[0])[np.newaxis, :]
|
|
59
|
+
|
|
60
|
+
# the von Neumann time grid is coarser than the FFT grid,
|
|
61
|
+
# so each row of f_tmp is embedded into the FFT grid
|
|
62
|
+
t_grid = 2 * np.pi * np.fft.fftfreq(k2, d=dw) # (k2,)
|
|
63
|
+
idx_cols = np.array(
|
|
64
|
+
[np.abs(t_grid - t).argmin() for t in t_n_arr],
|
|
65
|
+
dtype=int,
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
q_tmp = np.zeros((k, k2), dtype=np.complex128)
|
|
69
|
+
q_tmp[:, idx_cols] = tmp # (k, k2)
|
|
70
|
+
|
|
71
|
+
# The FFT phase e^{-i t_m (w_o - w_0)} combines with the
|
|
72
|
+
# correction above to give e^{-i t_m w_o}.
|
|
73
|
+
f_tmp = np.fft.fft(q_tmp, axis=1)
|
|
74
|
+
|
|
75
|
+
# apply the Gaussian window
|
|
76
|
+
# signal[o] = \sum_n f_tmp[n, o] * alpha_w[n, o]
|
|
77
|
+
alpha_w = norm * np.exp(-alpha * np.subtract.outer(w_n_arr, w_grid) ** 2)
|
|
78
|
+
signal = np.sum(f_tmp * alpha_w, axis=0)
|
|
79
|
+
|
|
80
|
+
return signal
|