makewfs 1.0.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.
makewfs/__about__.py ADDED
@@ -0,0 +1,3 @@
1
+ """Package version metadata."""
2
+
3
+ __version__ = "1.0.0"
makewfs/__init__.py ADDED
@@ -0,0 +1,15 @@
1
+ """Configuration-driven adaptive-optics wavefront-sensor image simulation."""
2
+
3
+ from .__about__ import __version__
4
+ from .api import WavefrontSensor, simulate
5
+ from .config import Config, ConfigError, WFSConfig, load_config
6
+
7
+ __all__ = [
8
+ "Config",
9
+ "ConfigError",
10
+ "WFSConfig",
11
+ "WavefrontSensor",
12
+ "__version__",
13
+ "load_config",
14
+ "simulate",
15
+ ]
makewfs/api.py ADDED
@@ -0,0 +1,210 @@
1
+ """Small public facade for configured wavefront-sensor simulations."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Iterable, Iterator
6
+ from pathlib import Path
7
+ from time import perf_counter
8
+ from typing import Any, cast
9
+
10
+ import numpy as np
11
+ from numpy.typing import ArrayLike, NDArray
12
+
13
+ from .backend import ArrayBackend, cpu_backend, cupy_backend
14
+ from .config import WFSConfig, load_config
15
+ from .detector import DetectorAdapter
16
+ from .provenance import metadata as build_metadata
17
+ from .sensors.base import OpticalResult, SensorEngine
18
+ from .sensors.pyramid import PyramidEngine
19
+ from .sensors.shack_hartmann import ShackHartmannEngine
20
+ from .wavefront import iter_phase_samples
21
+
22
+
23
+ class WavefrontSensor:
24
+ """Configured wavefront-sensor facade.
25
+
26
+ Construct once and call :meth:`expose` for each closed-loop residual OPD.
27
+ ``photon_rate`` exposes the deterministic optical result for workflows that
28
+ use another detector or need an ideal reference image.
29
+ """
30
+
31
+ def __init__(self, config: WFSConfig, *, _backend: ArrayBackend | None = None) -> None:
32
+ self.config = config
33
+ self.backend = _backend or (
34
+ cupy_backend() if config.numerics.device == "gpu" else cpu_backend()
35
+ )
36
+ self.engine: SensorEngine
37
+ if config.sensor.kind == "shack_hartmann":
38
+ self.engine = ShackHartmannEngine(config, backend=self.backend)
39
+ elif config.sensor.kind == "pyramid":
40
+ self.engine = PyramidEngine(config, backend=self.backend)
41
+ else:
42
+ raise NotImplementedError(f"unsupported sensor kind {config.sensor.kind!r}")
43
+ self.detector = DetectorAdapter(
44
+ config.detector,
45
+ self.engine.output_shape,
46
+ device="cpu" if self.backend.is_cpu else "gpu",
47
+ )
48
+ self._metadata_base = build_metadata(
49
+ self.config,
50
+ sensor_kind=self.config.sensor.kind,
51
+ launched_rate=0.0,
52
+ captured_rate=0.0,
53
+ opd_rms_m=0.0,
54
+ seed=None,
55
+ source_states=self.engine.source_states,
56
+ file_digests=self.engine.file_digests,
57
+ )
58
+
59
+ @classmethod
60
+ def from_toml(cls, path: str | Path) -> WavefrontSensor:
61
+ """Build a sensor from a validated TOML file."""
62
+ return cls(load_config(path))
63
+
64
+ def _render(self, wavefront: ArrayLike) -> OpticalResult:
65
+ return self.engine.render(cast(NDArray[np.float64], wavefront))
66
+
67
+ def _opd_rms(self, opd: Any) -> Any:
68
+ """Reduce OPD RMS on-device before the batched metadata crossing."""
69
+ return self.backend.sqrt(self.backend.mean(opd**2))
70
+
71
+ def _frame_metadata(
72
+ self,
73
+ *,
74
+ launched_rate: float,
75
+ captured_rate: float,
76
+ opd_rms_m: float,
77
+ seed: int | None,
78
+ ) -> dict[str, Any]:
79
+ """Copy cached static provenance and fill the per-frame values."""
80
+ result = dict(self._metadata_base)
81
+ result.update(
82
+ {
83
+ "wfs_launched_photons_s": float(launched_rate),
84
+ "wfs_captured_photons_s": float(captured_rate),
85
+ "wfs_input_opd_rms_m": float(opd_rms_m),
86
+ "wfs_seed": seed if seed is not None else "internal",
87
+ }
88
+ )
89
+ return result
90
+
91
+ def photon_rate(self, wavefront: ArrayLike) -> Any:
92
+ """Return the ideal native-pixel photon rate before detector noise."""
93
+ return self._render(wavefront).photon_rate
94
+
95
+ def reference(self) -> Any:
96
+ """Return the ideal image for a zero dynamic OPD."""
97
+ dtype = np.dtype(self.config.numerics.dtype)
98
+ return self.photon_rate(self.backend.zeros(self.config.input.shape, dtype=dtype))
99
+
100
+ def expose(self, wavefront: ArrayLike, *, seed: int | None = None) -> Any:
101
+ """Render one wavefront and expose it through the configured detector."""
102
+ total_start = perf_counter()
103
+ optical_start = total_start
104
+ result = self._render(wavefront)
105
+ captured_rate, opd_rms = self.backend.scalars(
106
+ result.captured_rate_per_s, self._opd_rms(result.opd_m)
107
+ )
108
+ optical_elapsed = perf_counter() - optical_start
109
+ frame_metadata = self._frame_metadata(
110
+ launched_rate=result.launched_rate_per_s,
111
+ captured_rate=captured_rate,
112
+ opd_rms_m=opd_rms,
113
+ seed=seed,
114
+ )
115
+ detector_start = perf_counter()
116
+ frame = self.detector.expose(
117
+ result.photon_rate,
118
+ metadata=frame_metadata,
119
+ seed=seed,
120
+ spectral_photon_rate=(
121
+ None if result.spectral_photon_rate is None else result.spectral_photon_rate
122
+ ),
123
+ spectral_wavelengths_m=result.spectral_wavelengths_m,
124
+ )
125
+ detector_elapsed = perf_counter() - detector_start
126
+ frame.metadata["wfs_optical_render_s"] = optical_elapsed
127
+ frame.metadata["wfs_detector_expose_s"] = detector_elapsed
128
+ frame.metadata["wfs_total_expose_s"] = perf_counter() - total_start
129
+ return frame
130
+
131
+ def expose_many(
132
+ self, phases: Iterable[ArrayLike], seeds: Iterable[int | None] | None = None
133
+ ) -> Iterator[Any]:
134
+ """Yield one detector frame per phase sample without stacking the stream."""
135
+ seed_iter = iter(seeds) if seeds is not None else None
136
+ for phase in phases:
137
+ seed = next(seed_iter) if seed_iter is not None else None
138
+ yield self.expose(phase, seed=seed)
139
+
140
+ def expose_integrated(
141
+ self, phase_samples: ArrayLike | Iterable[ArrayLike], *, seed: int | None = None
142
+ ) -> Any:
143
+ """Expose one detector frame after uniformly averaging temporal OPD samples."""
144
+ total_start = perf_counter()
145
+ optical_start = total_start
146
+ rates: list[Any] = []
147
+ spectral_rates: list[Any] = []
148
+ spectral_wavelengths_m: tuple[float, ...] | None = None
149
+ opds: list[Any] = []
150
+ launched = 0.0
151
+ captured: Any = 0.0
152
+ for sample in iter_phase_samples(
153
+ phase_samples,
154
+ self.config.input.shape,
155
+ backend=self.backend,
156
+ ):
157
+ result = self._render(sample)
158
+ rates.append(result.photon_rate)
159
+ if result.spectral_photon_rate is not None:
160
+ spectral_rates.append(result.spectral_photon_rate)
161
+ if spectral_wavelengths_m is None:
162
+ spectral_wavelengths_m = result.spectral_wavelengths_m
163
+ elif spectral_wavelengths_m != result.spectral_wavelengths_m:
164
+ raise RuntimeError("spectral wavelength nodes changed within one exposure")
165
+ opds.append(result.opd_m)
166
+ launched = result.launched_rate_per_s
167
+ captured += result.captured_rate_per_s
168
+ if not rates:
169
+ raise ValueError("phase_samples must contain at least one sample")
170
+ average_rate = self.backend.mean(self.backend.stack(rates), axis=0)
171
+ average_spectral_rate = (
172
+ None
173
+ if not spectral_rates
174
+ else self.backend.mean(self.backend.stack(spectral_rates), axis=0)
175
+ )
176
+ average_opd = self.backend.mean(self.backend.stack(opds), axis=0)
177
+ captured_rate, opd_rms = self.backend.scalars(
178
+ captured / len(rates), self._opd_rms(average_opd)
179
+ )
180
+ frame_metadata = self._frame_metadata(
181
+ launched_rate=launched,
182
+ captured_rate=captured_rate,
183
+ opd_rms_m=opd_rms,
184
+ seed=seed,
185
+ )
186
+ frame_metadata["wfs_temporal_samples"] = len(rates)
187
+ optical_elapsed = perf_counter() - optical_start
188
+ detector_start = perf_counter()
189
+ frame = self.detector.expose(
190
+ average_rate,
191
+ metadata=frame_metadata,
192
+ seed=seed,
193
+ spectral_photon_rate=average_spectral_rate,
194
+ spectral_wavelengths_m=spectral_wavelengths_m,
195
+ )
196
+ frame.metadata["wfs_optical_render_s"] = optical_elapsed
197
+ frame.metadata["wfs_detector_expose_s"] = perf_counter() - detector_start
198
+ frame.metadata["wfs_total_expose_s"] = perf_counter() - total_start
199
+ return frame
200
+
201
+
202
+ def simulate(
203
+ wavefront: ArrayLike, config: WFSConfig | str | Path, *, seed: int | None = None
204
+ ) -> Any:
205
+ """One-shot convenience wrapper around :class:`WavefrontSensor`."""
206
+ resolved = load_config(config) if isinstance(config, (str, Path)) else config
207
+ return WavefrontSensor(resolved).expose(wavefront, seed=seed)
208
+
209
+
210
+ __all__ = ["WavefrontSensor", "simulate"]
makewfs/backend.py ADDED
@@ -0,0 +1,396 @@
1
+ """Array and FFT primitives used by the portable optical kernels.
2
+
3
+ The CPU release uses :mod:`numpy` arrays and :mod:`scipy.fft`, but sensor
4
+ mathematics calls this small backend object rather than allocating through
5
+ NumPy directly. That boundary is deliberately private today; it gives a
6
+ future CuPy implementation one place to provide array creation, reductions,
7
+ FFT, and interpolation semantics without changing the optical equations.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from dataclasses import dataclass
13
+ from typing import Any, cast
14
+
15
+ import numpy as np
16
+ from numpy.typing import NDArray
17
+
18
+
19
+ @dataclass(frozen=True)
20
+ class ArrayBackend:
21
+ """Numerical namespace for one optical array backend.
22
+
23
+ ``xp`` is an Array API-compatible namespace. The CPU instance is the only
24
+ supported instance in the public release; a future device backend may
25
+ provide the corresponding namespace and override the SciPy-only helpers.
26
+ Methods named ``scalar`` and ``to_host`` are explicit host-boundary points
27
+ for metadata and configuration diagnostics, rather than accidental scalar
28
+ extraction in a sensor kernel.
29
+ """
30
+
31
+ xp: Any
32
+ name: str = "cpu"
33
+
34
+ @property
35
+ def is_cpu(self) -> bool:
36
+ """Whether this backend uses host NumPy/SciPy arrays."""
37
+ return self.name == "cpu"
38
+
39
+ def asarray(self, value: Any, *, dtype: Any | None = None) -> Any:
40
+ """Convert a value using this backend's array namespace."""
41
+ return self.xp.asarray(value, dtype=dtype)
42
+
43
+ def zeros(self, shape: Any, *, dtype: Any) -> Any:
44
+ """Allocate a zero-filled array on this backend."""
45
+ return self.xp.zeros(shape, dtype=dtype)
46
+
47
+ def zeros_like(self, value: Any) -> Any:
48
+ """Allocate an array matching ``value`` on this backend."""
49
+ return self.xp.zeros_like(value)
50
+
51
+ def astype(self, value: Any, dtype: Any) -> Any:
52
+ """Cast an array without changing its backend."""
53
+ return value.astype(dtype)
54
+
55
+ def full(self, shape: Any, value: Any, *, dtype: Any) -> Any:
56
+ """Allocate a constant-filled array on this backend."""
57
+ return self.xp.full(shape, value, dtype=dtype)
58
+
59
+ def empty(self, shape: Any, *, dtype: Any) -> Any:
60
+ """Allocate an uninitialized array on this backend."""
61
+ return self.xp.empty(shape, dtype=dtype)
62
+
63
+ def arange(self, *args: Any, **kwargs: Any) -> Any:
64
+ """Create a backend array of evenly spaced values."""
65
+ return self.xp.arange(*args, **kwargs)
66
+
67
+ def meshgrid(self, *args: Any, **kwargs: Any) -> Any:
68
+ """Create backend coordinate grids."""
69
+ return self.xp.meshgrid(*args, **kwargs)
70
+
71
+ def repeat(self, value: Any, repeats: Any, *, axis: int) -> Any:
72
+ """Repeat values along one backend axis."""
73
+ return self.xp.repeat(value, repeats, axis=axis)
74
+
75
+ def tile(self, value: Any, reps: Any) -> Any:
76
+ """Tile a backend array."""
77
+ return self.xp.tile(value, reps)
78
+
79
+ def stack(self, values: Any, *, axis: int = 0) -> Any:
80
+ """Stack backend arrays."""
81
+ return self.xp.stack(values, axis=axis)
82
+
83
+ def sum(self, value: Any, *, axis: Any = None) -> Any:
84
+ """Reduce a backend array by summation."""
85
+ return self.xp.sum(value, axis=axis)
86
+
87
+ def mean(self, value: Any, *, axis: Any = None) -> Any:
88
+ """Reduce a backend array by mean."""
89
+ return self.xp.mean(value, axis=axis)
90
+
91
+ def average(self, value: Any, *, weights: Any) -> Any:
92
+ """Compute a weighted backend average."""
93
+ return self.xp.average(value, weights=weights)
94
+
95
+ def any(self, value: Any) -> Any:
96
+ """Backend reduction testing whether any element is true."""
97
+ return self.xp.any(value)
98
+
99
+ def all(self, value: Any) -> Any:
100
+ """Backend reduction testing whether all elements are true."""
101
+ return self.xp.all(value)
102
+
103
+ def isfinite(self, value: Any) -> Any:
104
+ """Elementwise finite-value test."""
105
+ return self.xp.isfinite(value)
106
+
107
+ def abs(self, value: Any) -> Any:
108
+ """Elementwise absolute value."""
109
+ return self.xp.abs(value)
110
+
111
+ def exp(self, value: Any) -> Any:
112
+ """Elementwise exponential."""
113
+ return self.xp.exp(value)
114
+
115
+ def sqrt(self, value: Any) -> Any:
116
+ """Elementwise square root."""
117
+ return self.xp.sqrt(value)
118
+
119
+ def hypot(self, left: Any, right: Any) -> Any:
120
+ """Elementwise Euclidean norm."""
121
+ return self.xp.hypot(left, right)
122
+
123
+ def cos(self, value: Any) -> Any:
124
+ """Elementwise cosine."""
125
+ return self.xp.cos(value)
126
+
127
+ def sin(self, value: Any) -> Any:
128
+ """Elementwise sine."""
129
+ return self.xp.sin(value)
130
+
131
+ def arctan2(self, left: Any, right: Any) -> Any:
132
+ """Elementwise two-argument arctangent."""
133
+ return self.xp.arctan2(left, right)
134
+
135
+ def mod(self, value: Any, divisor: Any) -> Any:
136
+ """Elementwise remainder."""
137
+ return self.xp.mod(value, divisor)
138
+
139
+ def where(self, condition: Any, left: Any, right: Any) -> Any:
140
+ """Select values elementwise on the backend."""
141
+ return self.xp.where(condition, left, right)
142
+
143
+ def ptp(self, value: Any) -> Any:
144
+ """Backend peak-to-peak reduction."""
145
+ return self.xp.ptp(value)
146
+
147
+ def argmax(self, value: Any) -> Any:
148
+ """Return the flat index of a backend array maximum."""
149
+ return self.xp.argmax(value)
150
+
151
+ def fftfreq(self, value: int) -> Any:
152
+ """Return backend FFT frequency bins."""
153
+ return self.xp.fft.fftfreq(value)
154
+
155
+ def fftshift(self, value: Any, *, axes: Any = None) -> Any:
156
+ """Shift zero frequency to the center."""
157
+ return self.xp.fft.fftshift(value, axes=axes)
158
+
159
+ def ifftshift(self, value: Any, *, axes: Any = None) -> Any:
160
+ """Undo a centered FFT shift."""
161
+ return self.xp.fft.ifftshift(value, axes=axes)
162
+
163
+ def centered_fft2(self, array: Any, *, workers: int = 1) -> Any:
164
+ """Perform a centered, unitary two-dimensional FFT."""
165
+ axes = (-2, -1)
166
+ if self.is_cpu:
167
+ from scipy import fft
168
+
169
+ transformed = fft.fftshift(
170
+ fft.fft2(
171
+ fft.ifftshift(array, axes=axes),
172
+ axes=axes,
173
+ workers=workers,
174
+ norm="ortho",
175
+ overwrite_x=True,
176
+ ),
177
+ axes=axes,
178
+ )
179
+ else: # pragma: no cover - reserved for a future device backend
180
+ transformed = self.fftshift(
181
+ self.xp.fft.fft2(self.ifftshift(array, axes=axes), axes=axes, norm="ortho"),
182
+ axes=axes,
183
+ )
184
+ return transformed
185
+
186
+ def centered_fft_intensity(
187
+ self, array: Any, *, workers: int = 1, overwrite_input: bool = False
188
+ ) -> Any:
189
+ """Return centered unitary FFT intensity without an irrelevant input roll.
190
+
191
+ Translating an entrance field changes only Fourier phase, so an input
192
+ ``ifftshift`` cannot affect intensity. Shack-Hartmann propagation uses
193
+ this identity to avoid one detector-batch-sized array permutation.
194
+ """
195
+ axes = (-2, -1)
196
+ height, width = array.shape[-2:]
197
+ if height % 2 == 0 and width % 2 == 0:
198
+ working = array if overwrite_input else self.xp.array(array, copy=True)
199
+ working[..., ::2, 1::2] *= -1
200
+ working[..., 1::2, ::2] *= -1
201
+ if self.is_cpu:
202
+ from scipy import fft
203
+
204
+ transformed = fft.fft2(
205
+ working,
206
+ axes=axes,
207
+ workers=workers,
208
+ norm="ortho",
209
+ overwrite_x=overwrite_input,
210
+ )
211
+ else: # pragma: no cover - GPU optional
212
+ transformed = self.xp.fft.fft2(working, axes=axes, norm="ortho")
213
+ return self.abs(transformed) ** 2
214
+ if self.is_cpu:
215
+ from scipy import fft
216
+
217
+ transformed = fft.fftshift(
218
+ fft.fft2(array, axes=axes, workers=workers, norm="ortho"),
219
+ axes=axes,
220
+ )
221
+ else: # pragma: no cover - GPU optional
222
+ transformed = self.fftshift(self.xp.fft.fft2(array, axes=axes, norm="ortho"), axes=axes)
223
+ return self.abs(transformed) ** 2
224
+
225
+ def centered_ifft2(self, array: Any, *, workers: int = 1) -> Any:
226
+ """Perform a centered, unitary two-dimensional inverse FFT."""
227
+ axes = (-2, -1)
228
+ if self.is_cpu:
229
+ from scipy import fft
230
+
231
+ transformed = fft.fftshift(
232
+ fft.ifft2(
233
+ fft.ifftshift(array, axes=axes),
234
+ axes=axes,
235
+ workers=workers,
236
+ norm="ortho",
237
+ overwrite_x=True,
238
+ ),
239
+ axes=axes,
240
+ )
241
+ else: # pragma: no cover - reserved for a future device backend
242
+ transformed = self.fftshift(
243
+ self.xp.fft.ifft2(self.ifftshift(array, axes=axes), axes=axes, norm="ortho"),
244
+ axes=axes,
245
+ )
246
+ return transformed
247
+
248
+ def map_coordinates(self, array: Any, coordinates: Any, *, order: int, mode: str) -> Any:
249
+ """Interpolate coordinates, using SciPy only for the CPU backend."""
250
+ if self.is_cpu:
251
+ from scipy.ndimage import map_coordinates
252
+
253
+ return map_coordinates(array, coordinates, order=order, mode=mode)
254
+ from cupyx.scipy.ndimage import map_coordinates # pragma: no cover - GPU optional
255
+
256
+ device_coordinates = (
257
+ self.xp.stack(coordinates, axis=0)
258
+ if isinstance(coordinates, (list, tuple))
259
+ else coordinates
260
+ )
261
+ return map_coordinates(array, device_coordinates, order=order, mode=mode)
262
+
263
+ def convolve(self, array: Any, kernel: Any) -> Any:
264
+ """Convolve a batch of arrays with a backend-compatible kernel."""
265
+ if self.is_cpu:
266
+ from scipy.ndimage import convolve
267
+
268
+ return convolve(array, kernel, mode="constant", cval=0.0)
269
+ from cupyx.scipy.ndimage import convolve # pragma: no cover - GPU optional
270
+
271
+ return convolve(array, kernel, mode="constant", cval=0.0)
272
+
273
+ def gaussian_filter(self, array: Any, sigma: Any) -> Any:
274
+ """Apply a Gaussian filter through the CPU numerical backend."""
275
+ if self.is_cpu:
276
+ from scipy.ndimage import gaussian_filter
277
+
278
+ return gaussian_filter(array, sigma=sigma, mode="constant")
279
+ from cupyx.scipy.ndimage import gaussian_filter # pragma: no cover - GPU optional
280
+
281
+ return gaussian_filter(array, sigma=sigma, mode="constant")
282
+
283
+ def next_fast_length(self, value: int) -> int:
284
+ """Return an FFT-friendly length using this backend's implementation."""
285
+ if self.is_cpu:
286
+ from scipy.fft import next_fast_len
287
+ else: # pragma: no cover - GPU optional
288
+ from cupyx.scipy.fft import next_fast_len
289
+
290
+ return int(next_fast_len(value))
291
+
292
+ def scalar(self, value: Any) -> float:
293
+ """Extract one host scalar at an explicit metadata/geometry boundary."""
294
+ item = value.item() if hasattr(value, "item") else value
295
+ return float(item)
296
+
297
+ def scalars(self, *values: Any) -> tuple[float, ...]:
298
+ """Extract several scalars with one device synchronization."""
299
+ if self.is_cpu:
300
+ return tuple(self.scalar(value) for value in values)
301
+ packed = self.xp.stack([self.xp.asarray(value) for value in values])
302
+ return tuple(float(value) for value in self.xp.asnumpy(packed))
303
+
304
+ def to_host(self, value: Any) -> NDArray[Any]:
305
+ """Copy an array to host NumPy storage at an explicit boundary."""
306
+ if self.is_cpu:
307
+ return cast(NDArray[Any], value)
308
+ return np.asarray(self.xp.asnumpy(value))
309
+
310
+
311
+ _CPU_BACKEND = ArrayBackend(np, name="cpu")
312
+
313
+
314
+ def cpu_backend() -> ArrayBackend:
315
+ """Return the shared CPU backend instance."""
316
+ return _CPU_BACKEND
317
+
318
+
319
+ def cupy_backend() -> ArrayBackend:
320
+ """Return the private optional CuPy backend.
321
+
322
+ CuPy is intentionally imported lazily so the core package remains usable
323
+ without CUDA. Callers should treat this as experimental and keep the
324
+ explicit host transfer before the ``getframes`` detector adapter.
325
+ """
326
+ try:
327
+ import cupy
328
+ except ImportError as exc: # pragma: no cover - depends on optional install
329
+ raise ImportError("the private CuPy backend requires makewfs[gpu]") from exc
330
+ return ArrayBackend(cupy, name="cupy")
331
+
332
+
333
+ def real_dtype(name: str) -> np.dtype[Any]:
334
+ """Return the configured real dtype."""
335
+ if name == "float32":
336
+ return np.dtype(np.float32)
337
+ if name == "float64":
338
+ return np.dtype(np.float64)
339
+ raise ValueError(f"unsupported dtype {name!r}")
340
+
341
+
342
+ def complex_dtype(name: str) -> np.dtype[Any]:
343
+ """Return the complex dtype paired with a real dtype."""
344
+ if name == "float32":
345
+ return np.dtype(np.complex64)
346
+ if name == "float64":
347
+ return np.dtype(np.complex128)
348
+ raise ValueError(f"unsupported dtype {name!r}")
349
+
350
+
351
+ def centered_fft2(
352
+ array: NDArray[Any], *, workers: int = 1, backend: ArrayBackend | None = None
353
+ ) -> NDArray[Any]:
354
+ """Centered two-dimensional FFT with unitary normalization."""
355
+ return cast(NDArray[Any], (backend or cpu_backend()).centered_fft2(array, workers=workers))
356
+
357
+
358
+ def centered_ifft2(
359
+ array: NDArray[Any], *, workers: int = 1, backend: ArrayBackend | None = None
360
+ ) -> NDArray[Any]:
361
+ """Centered two-dimensional inverse FFT with unitary normalization."""
362
+ return cast(NDArray[Any], (backend or cpu_backend()).centered_ifft2(array, workers=workers))
363
+
364
+
365
+ def centered_fft_intensity(
366
+ array: NDArray[Any],
367
+ *,
368
+ workers: int = 1,
369
+ backend: ArrayBackend | None = None,
370
+ overwrite_input: bool = False,
371
+ ) -> NDArray[Any]:
372
+ """Return centered unitary FFT intensity without shifting the input field."""
373
+ return cast(
374
+ NDArray[Any],
375
+ (backend or cpu_backend()).centered_fft_intensity(
376
+ array, workers=workers, overwrite_input=overwrite_input
377
+ ),
378
+ )
379
+
380
+
381
+ def next_fast_length(value: int) -> int:
382
+ """Return a convenient CPU FFT length."""
383
+ return cpu_backend().next_fast_length(value)
384
+
385
+
386
+ __all__ = [
387
+ "ArrayBackend",
388
+ "centered_fft2",
389
+ "centered_fft_intensity",
390
+ "centered_ifft2",
391
+ "complex_dtype",
392
+ "cpu_backend",
393
+ "cupy_backend",
394
+ "next_fast_length",
395
+ "real_dtype",
396
+ ]