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.
@@ -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