fastsar 0.1.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.
fastsar/__init__.py ADDED
@@ -0,0 +1,16 @@
1
+ """Fast spotlight SAR image formation: factorized backprojection with kernels for Cloud TPUs (Pallas), Nvidia GPUs
2
+ (CUDA) and x86 CPUs (C++/OpenMP), and polar format with its geometric resampling. See fastsar.api.form_image."""
3
+ import os as _os
4
+
5
+ # JAX takes 75% of a GPU's memory when it first runs there unless told otherwise; FastSAR's CUDA kernels (CuPy) and
6
+ # its JAX programs share the GPU, so JAX allocates as it goes (a setting the user made before importing JAX stands);
7
+ # JAX keeps what it has allocated, and the CuPy formers shrink their working sets to what is left
8
+ _os.environ.setdefault('XLA_PYTHON_CLIENT_PREALLOCATE', 'false')
9
+ from .api import form_image, available_backends, ImageFormer # noqa: F401
10
+ from . import io, autofocus, stripmap, burst, patches, quality, products # noqa: F401
11
+ from .bp import backproject, plane_points # noqa: F401
12
+ from .exact import ExactFormer # noqa: F401
13
+ from .cphd import form_cphd # noqa: F401
14
+ from .memory import MemoryWarning # noqa: F401
15
+
16
+ __version__ = '0.1.0'
fastsar/_build.py ADDED
@@ -0,0 +1,59 @@
1
+ """Compiling the C++ kernels on first use: one shared object per source, compiler, flags and host, cached in
2
+ ~/.cache/fastsar. The source is compiled with the given flags and linked in a separate step without them, so that
3
+ -ffast-math, when a kernel is compiled with it, does not link crtfastmath.o (which would switch the whole process to
4
+ flush-to-zero on load). Concurrent first builds each write their own temporary files and the last os.replace wins,
5
+ with identical contents."""
6
+ import hashlib
7
+ import os
8
+ import platform
9
+ import subprocess
10
+
11
+
12
+ def _host_tag():
13
+ """The machine and its CPU's instruction set extensions (a -march=native build must not be loaded on another
14
+ CPU sharing the same home directory)."""
15
+ tag = platform.machine()
16
+ try:
17
+ with open('/proc/cpuinfo') as fh:
18
+ for line in fh:
19
+ if line.startswith(('flags', 'Features')):
20
+ return tag + line
21
+ except OSError:
22
+ pass
23
+ return tag + platform.processor()
24
+
25
+
26
+ def shared_object(name, src, flags):
27
+ """Path of the shared object built from the C++ source text src with the compiler flags (a list), building it
28
+ if it is not cached. The compiler is $CXX (default g++); OpenMP and -fPIC are added."""
29
+ cxx = os.environ.get('CXX', 'g++')
30
+ key = '\0'.join([src, cxx, ' '.join(flags), _host_tag()])
31
+ tag = hashlib.sha1(key.encode()).hexdigest()[:12]
32
+ d = os.path.join(os.path.expanduser('~'), '.cache', 'fastsar')
33
+ os.makedirs(d, exist_ok=True)
34
+ so = os.path.join(d, f'lib{name}_{tag}.so')
35
+ if os.path.exists(so):
36
+ return so
37
+ stem = os.path.join(d, f'{name}_{tag}.{os.getpid()}')
38
+ cpp, obj, tmp = stem + '.cpp', stem + '.o', stem + '.so.tmp'
39
+ with open(cpp, 'w') as fh:
40
+ fh.write(src)
41
+ try:
42
+ for cmd in ([cxx] + list(flags) + ['-fopenmp', '-fPIC', '-c', cpp, '-o', obj],
43
+ [cxx, '-shared', '-fopenmp', obj, '-o', tmp]):
44
+ try:
45
+ subprocess.run(cmd, check=True, capture_output=True, text=True)
46
+ except FileNotFoundError:
47
+ raise RuntimeError(f"fastsar's cpu backend compiles its C++ kernels on first use and needs a C++ compiler "
48
+ f'with OpenMP: {cxx!r} was not found. Install g++ (Debian/Ubuntu: apt install g++; '
49
+ 'RHEL/Fedora: dnf install gcc-c++) or point CXX at one') from None
50
+ except subprocess.CalledProcessError as e:
51
+ raise RuntimeError(f'compiling the {name} kernel failed ({" ".join(cmd)}):\n{e.stderr.strip()}') from None
52
+ os.replace(tmp, so)
53
+ finally:
54
+ for f in (cpp, obj, tmp):
55
+ try:
56
+ os.remove(f)
57
+ except OSError:
58
+ pass
59
+ return so
fastsar/api.py ADDED
@@ -0,0 +1,487 @@
1
+ """One call for spotlight image formation on any of the supported devices.
2
+
3
+ import fastsar
4
+ img = fastsar.form_image(S, ant, fmin, df, nx, ny, spx, spy, e1, e2) # factorized backprojection
5
+ img = fastsar.form_image(..., algorithm='pfa') # polar format
6
+ img = fastsar.form_image(..., algorithm='bp') # exact backprojection
7
+
8
+ S is the phase history [pulses, samples] (complex, frequency domain, motion compensated to the scene reference
9
+ point), ant the antenna phase centers [pulses, 3] in a frame whose origin is the scene reference point, fmin and df
10
+ the first frequency and the sample spacing (Hz), and nx, ny, spx, spy, e1, e2 the output grid: pixel counts and
11
+ spacings (m) along the unit vectors e1 (azimuth) and e2 (range) of the image plane. The result is a complex64 image
12
+ [nx, ny] on that grid.
13
+
14
+ Backends for factorized backprojection (the same plan, filters and float64 geometry on every device):
15
+
16
+ 'tpu' Pallas kernels (JAX on a Cloud TPU)
17
+ 'cuda' CUDA kernels through CuPy (Nvidia GPU)
18
+ 'cpu' C++ kernels with OpenMP (x86-64; compiled with g++ on first use)
19
+ 'jax' the plain JAX program, on whatever device JAX has
20
+ 'auto' tpu if JAX sees a TPU, else cuda if CuPy sees a GPU, else cpu
21
+
22
+ ref [P]: the range (one way) each pulse's samples are referenced to, when it is not |ant| (a bistatic collection's
23
+ half path, a vendor's reference point); both algorithms honor it (polar format re-references the samples to |ant|).
24
+
25
+ pfa_guard: margin (m) around the scene that polar format keeps free of wrap-around; 300 m suits orbital scenes of a
26
+ few kilometers and must be smaller for small simulated scenes.
27
+
28
+ Precision: 'float32' (default), 'float16' (CUDA: float16 storage throughout with float32 accumulation; JAX: float16
29
+ throughout), 'single-pass' (TPU: one bfloat16 pass per product, what the matrix unit does by itself), 'three-pass'
30
+ (TPU: float32-class accuracy). On the TPU 'float32' means three-pass.
31
+
32
+ Tile size: T='auto' (default) picks the largest final tile (32 or 16 pixels) whose predicted error against exact
33
+ backprojection meets target_db (-40 dB by default). The error of the final stage's plane-wave model grows as the
34
+ square of the tile size and falls with range, so orbital collections keep T=32 and short-range (airborne) ones
35
+ drop to 16. The prediction is kept as ImageFormer.predicted_error_db.
36
+ """
37
+ import os
38
+ import threading
39
+
40
+ import numpy as np
41
+
42
+ from . import memory as _mem
43
+ from .sim import Collect
44
+
45
+
46
+ def _window(P, K, sll=35.0, nbar=4):
47
+ from scipy.signal.windows import taylor
48
+ return taylor(P, nbar=nbar, sll=sll, norm=False).astype(np.float32), taylor(K, nbar=nbar, sll=sll, norm=False).astype(np.float32)
49
+
50
+
51
+ def available_backends():
52
+ """The backends this machine can run, in the order 'auto' tries them."""
53
+ out = []
54
+ try:
55
+ import jax
56
+ if any(d.platform == 'tpu' for d in jax.devices()):
57
+ out.append('tpu')
58
+ except Exception:
59
+ pass
60
+ try:
61
+ import cupy
62
+ if cupy.cuda.runtime.getDeviceCount() > 0:
63
+ out.append('cuda')
64
+ except Exception:
65
+ pass
66
+ out.append('cpu')
67
+ return out
68
+
69
+
70
+ _JAX_PROGRAMS = {} # compiled JAX/TPU programs by plan signature (ImageFormer)
71
+ _JAX_LOCK = threading.Lock() # mosaic workers build formers concurrently: one build per signature, no torn eviction
72
+
73
+ BACKENDS = ('auto', 'tpu', 'cuda', 'cpu', 'jax')
74
+
75
+
76
+ def _backend(backend):
77
+ """backend checked against BACKENDS and this machine, 'auto' resolved."""
78
+ if backend not in BACKENDS:
79
+ raise ValueError(f'unknown backend {backend!r}: one of {", ".join(map(repr, BACKENDS))}')
80
+ if backend in ('cpu', 'jax'): # no device probe (a mosaic builds a former per patch)
81
+ return backend
82
+ have = available_backends()
83
+ if backend == 'auto':
84
+ return have[0]
85
+ if backend == 'cuda' and 'cuda' not in have:
86
+ raise ValueError("backend 'cuda' needs CuPy and an Nvidia GPU; this machine has " + ', '.join(have + ['jax']))
87
+ if backend == 'tpu' and 'tpu' not in have and not os.environ.get('FFBP_FORCE_TPU_KERNELS'):
88
+ raise ValueError("backend 'tpu' needs JAX on a Cloud TPU; this machine has " + ', '.join(have + ['jax']))
89
+ return backend
90
+
91
+
92
+ def _check_history(S, P=None, name='phase history', finite=True):
93
+ """S [P, K] complex (numpy or cupy) with at least 2 pulses and 2 samples, all finite (finite=False: the scan
94
+ for NaN and inf is left to the caller); -> S unchanged. Real or integer samples are refused (an I/Q pair
95
+ interleaved along the last axis is S[..., 0] + 1j * S[..., 1])."""
96
+ if not hasattr(S, 'ndim'):
97
+ S = np.asarray(S)
98
+ if S.ndim != 2:
99
+ raise ValueError(f'{name} must be 2-D [pulses, samples], got shape {tuple(S.shape)}')
100
+ if S.dtype.kind != 'c':
101
+ raise TypeError(f'{name} must be complex (complex64 or complex128), got {S.dtype}')
102
+ if P is not None and S.shape[0] != P:
103
+ raise ValueError(f'{name} has {S.shape[0]} pulses but there are {P} antenna positions')
104
+ if S.shape[0] < 2 or S.shape[1] < 2:
105
+ raise ValueError(f'{name} needs at least 2 pulses and 2 samples, got shape {tuple(S.shape)}')
106
+ # a sum per block of rows: one pass, no full-size temporary
107
+ for i in range(0, S.shape[0] if finite else 0, 4096):
108
+ if not np.isfinite(complex(S[i:i + 4096].sum())):
109
+ first = i + int(np.argmax((~np.isfinite(S[i:i + 4096])).any(1)))
110
+ raise ValueError(f'{name} has non-finite samples (NaN or inf), the first in pulse {first}; '
111
+ 'zero them (S[~np.isfinite(S)] = 0) or drop the pulses')
112
+ return S
113
+
114
+
115
+ def _check_positions(a, name='antenna positions', P=None):
116
+ """a [P, 3] finite float64, at least 2 rows (P of them when given)."""
117
+ a = np.asarray(a, np.float64)
118
+ if a.ndim != 2 or a.shape[1] != 3:
119
+ raise ValueError(f'{name} must be [pulses, 3], got shape {a.shape}')
120
+ if P is not None and len(a) != P:
121
+ raise ValueError(f'{name}: {len(a)} rows for {P} pulses')
122
+ if len(a) < 2:
123
+ raise ValueError(f'{name}: at least 2 pulses are needed, got {len(a)}')
124
+ if not np.isfinite(a).all():
125
+ raise ValueError(f'{name}: non-finite values (NaN or inf)')
126
+ return a
127
+
128
+
129
+ def _check_grid(nx, ny, spx, spy, e1, e2):
130
+ """Pixel counts (positive integers), spacings (positive) and the axes (orthonormal 3-vectors) -> nx, ny (int),
131
+ e1, e2 (float64)."""
132
+ for n, v in (('nx', nx), ('ny', ny)):
133
+ if not np.isscalar(v) or not np.isfinite(v) or v < 1 or int(v) != v:
134
+ raise ValueError(f'{n} must be a positive integer, got {v!r}')
135
+ for n, v in (('spx', spx), ('spy', spy)):
136
+ if not np.isfinite(v) or v <= 0:
137
+ raise ValueError(f'{n} must be a positive spacing in metres, got {v!r}')
138
+ e1, e2 = np.asarray(e1, np.float64), np.asarray(e2, np.float64)
139
+ if e1.shape != (3,) or e2.shape != (3,):
140
+ raise ValueError(f'e1 and e2 must be 3-vectors, got shapes {e1.shape} and {e2.shape}')
141
+ if abs(e1 @ e1 - 1) > 1e-6 or abs(e2 @ e2 - 1) > 1e-6 or abs(e1 @ e2) > 1e-6:
142
+ raise ValueError(f'e1 and e2 must be orthonormal (unit length, perpendicular), got {e1.tolist()} and {e2.tolist()}')
143
+ return int(nx), int(ny), e1, e2
144
+
145
+
146
+ def _jax_group(plan, K):
147
+ """First-level children per group on a TPU: memory.TPU_GROUP (groups of 4 are as fast as 8; docs/performance.md),
148
+ halved until memory.full_speed's model (the padded history planes and XLA's temporaries per child) fits in 95% of
149
+ the device's memory limit. The limit, not the memory free at this moment, so that patches of a mosaic with equal
150
+ plans get equal groups (and share a program) however many histories the workers have staged. FASTSAR_TPU_GROUP
151
+ overrides it; ImageFormer halves it again if the device still runs out of memory."""
152
+ forced = _mem._env_number('FASTSAR_TPU_GROUP', 1)
153
+ if forced:
154
+ return forced
155
+ hbm = _mem.tpu_capacity()
156
+ hist, per = _mem.tpu_history_bytes(plan, K) + _mem.TPU_FIXED, _mem.child_bytes(plan, 'tpu')
157
+ ng = _mem.TPU_GROUP
158
+ while ng > 1 and hist + ng * per > 0.95 * hbm:
159
+ ng //= 2
160
+ return ng
161
+
162
+
163
+ def final_weights(plan, weight, P, grad=False, points=None):
164
+ """Weight of each final subaperture at each final tile [ntiles, Pf] (float32), for ImageFormer's aperture_weight:
165
+ weight(points [m, 3], pulses [n] int) -> [n, m], the per-pulse weights at the tile centers, averaged over the
166
+ pulses that each final subaperture spans (its center on the input pulse axis composed through the levels'
167
+ decimations, width the product of the decimation factors) by 3-point Gauss-Legendre quadrature. grad=True
168
+ also returns the gradient along e1 and e2 (per metre) at the tile centers: [3, ntiles, Pf]."""
169
+ cen = np.asarray(plan['final']['cen'], np.float64)
170
+ if grad:
171
+ h1, h2 = 0.5 * plan['T'] * plan['spx'], 0.5 * plan['T'] * plan['spy']
172
+ e1, e2 = np.asarray(plan['e1'], np.float64), np.asarray(plan['e2'], np.float64)
173
+ w = [final_weights(plan, weight, P, points=q) for q in (cen, cen + h1 * e1, cen - h1 * e1, cen + h2 * e2, cen - h2 * e2)]
174
+ return np.stack([w[0], (w[1] - w[2]) / (2 * h1), (w[3] - w[4]) / (2 * h2)]).astype(np.float32)
175
+ q = cen if points is None else points
176
+ idx, D = np.arange(plan['final']['P'], dtype=np.float64), 1.0
177
+ for lv in reversed(plan['levels']):
178
+ idx = lv['Dp'] * idx + float(lv['pidx'][0])
179
+ D *= lv['Dp']
180
+ # pulse p covers [p - 1/2, p + 1/2); a subaperture spans [idx - D/2, idx + D/2] within the collection
181
+ a = np.clip(idx - D / 2, -0.5, P - 0.5); b = np.clip(idx + D / 2, -0.5, P - 0.5)
182
+ t, gw = np.array([-np.sqrt(0.6), 0.0, np.sqrt(0.6)]), np.array([5.0, 8.0, 5.0]) / 18.0
183
+ nodes = np.clip(np.rint(0.5 * (a + b)[:, None] + 0.5 * (b - a)[:, None] * t[None, :]), 0, P - 1).astype(np.int64)
184
+ U, inv = np.unique(nodes.ravel(), return_inverse=True)
185
+ W = np.asarray(weight(q, U), np.float64)[inv].reshape(len(idx), 3, -1) # [Pf, 3, ntiles]
186
+ # a subaperture centred beyond the collection holds the filter tails of the edge pulses (the decimators run m
187
+ # outputs past each end): it takes the weight of the nearest pulse (its nodes clip to it), not zero
188
+ m = np.tensordot(gw, W, axes=(0, 1)) # [Pf, ntiles]
189
+ return np.ascontiguousarray(m.T, np.float32)
190
+
191
+
192
+ def _jax_runtime_error():
193
+ """The exception type JAX raises for a failed computation (JaxRuntimeError; XlaRuntimeError in older releases)."""
194
+ import jax
195
+ err = getattr(jax.errors, 'JaxRuntimeError', None)
196
+ if err is None:
197
+ from jaxlib.xla_extension import XlaRuntimeError as err
198
+ return err
199
+
200
+
201
+ class _Staged:
202
+ """A phase history staged for a JAX or TPU former (ImageFormer.stage): the host array, its device planes and
203
+ scale, and the program they were padded for."""
204
+ __slots__ = ('S', 'hre', 'him', 'scale', 'fn')
205
+
206
+ def __init__(self, S, hre, him, scale, fn):
207
+ self.S, self.hre, self.him, self.scale, self.fn = S, hre, him, scale, fn
208
+
209
+
210
+ class ImageFormer:
211
+ """Factorized backprojection set up once for a collection geometry and output grid, then called on phase
212
+ histories: former = ImageFormer(ant, fmin, df, K, nx, ny, spx, spy, e1, e2); img = former(S).
213
+
214
+ Building plans the tiles and filters, computes the float64 geometry and compiles the kernels; each call then
215
+ pays only the image formation. Reuse one former for repeated images of the same geometry (or for timing);
216
+ a different antenna path needs a new former. Arguments as for form_image (those after e2 by keyword only), and
217
+ aperture_weight(points [m, 3], pulses [n]) -> W [n, m]: a per-pixel weight of pulses (int indices) (a stripmap aperture window), applied
218
+ in the final stage as the mean weight of each final subaperture's pulses at each final tile's center, with its
219
+ first-order variation across the tile (final_weights), and
220
+ ref [P]: the range (one way) each pulse is referenced to, when it is not |ant| (the distance to the origin):
221
+ a bistatic collection's half path |tx| / 2 + |rcv| / 2, or a reference point other than the origin."""
222
+
223
+ def __init__(self, ant, fmin, df, K, nx, ny, spx, spy, e1=(1.0, 0.0, 0.0), e2=(0.0, 1.0, 0.0), *, backend='auto',
224
+ precision='float32', window=True, T='auto', levels=3, pmax=0.4, target_db=-40.0, aperture_weight=None,
225
+ ref=None):
226
+ from . import ffbp2
227
+ import warnings
228
+ self.ant = _check_positions(ant)
229
+ self.P, self.K = self.ant.shape[0], int(K)
230
+ if self.K < 2:
231
+ raise ValueError(f'K must be at least 2 frequency samples, got {K}')
232
+ if not (np.isfinite(fmin) and np.isfinite(df) and fmin > 0 and df > 0):
233
+ raise ValueError(f'fmin and df must be positive frequencies in Hz, got {fmin!r} and {df!r}')
234
+ self.window = window
235
+ col = Collect(fmin=float(fmin), df=float(df), K=self.K, ant=self.ant, res=0.5)
236
+ nx, ny, e1, e2 = _check_grid(nx, ny, spx, spy, e1, e2)
237
+ self.backend = _backend(backend)
238
+ self.precision = precision
239
+ self.predicted_error_db = None
240
+ if T == 'auto':
241
+ fmax = float(fmin) + self.K * float(df)
242
+ if self.backend == 'cuda' and precision == 'float16':
243
+ T, err = 32, ffbp2.final_phase_error(self.ant, fmax, nx, ny, spx, spy, e1, e2, 32)
244
+ if err > target_db:
245
+ warnings.warn(f'cuda float16 needs T=32, predicted error {err:.1f} dB misses target {target_db:.1f} dB; '
246
+ "use precision='float32' for this geometry")
247
+ else:
248
+ T, err = ffbp2.choose_T(self.ant, fmax, nx, ny, spx, spy, e1, e2, target_db)
249
+ if err > target_db:
250
+ warnings.warn(f'predicted error {err:.1f} dB at T={T} misses target {target_db:.1f} dB '
251
+ '(short range for this pixel size)')
252
+ self.predicted_error_db = float(err)
253
+ self.T = T
254
+ plan = ffbp2.make_plan(col, nx, ny, spx, spy, T=T, nlev=levels, pmax=pmax, e1=e1, e2=e2)
255
+ self._plan = plan
256
+ if ref is not None and np.shape(ref) != (self.P,):
257
+ raise ValueError(f'ref must hold one range per pulse ({self.P}), got shape {np.shape(ref)}')
258
+ coll = ffbp2.collection_arrays(plan, self.ant, ref)
259
+ wf = None if aperture_weight is None else final_weights(plan, aperture_weight, self.P, grad=os.environ.get('FASTSAR_WEIGHT_GRAD', '1') == '1')
260
+ if window:
261
+ self.wp, self.wk = _window(self.P, self.K)
262
+ if self.backend == 'cuda':
263
+ from . import ffbp_cuda
264
+ if precision not in ('float32', 'float16'):
265
+ raise ValueError("cuda precision: 'float32' or 'float16'")
266
+ if precision == 'float16' and T != 32:
267
+ raise ValueError('cuda float16 uses the tensor-core final stage, which is built for T=32')
268
+ self._form = ffbp_cuda.make_ffbp_cuda(plan, coll, final_mode='f16tc' if precision == 'float16' else 'fp32',
269
+ store='f16' if precision == 'float16' else 'fp32', wf=wf)
270
+ elif self.backend == 'cpu':
271
+ from . import ffbp_cpu
272
+ if T not in (16, 32):
273
+ raise ValueError(f'the cpu backend supports T=16 or T=32, not {T}')
274
+ if precision != 'float32':
275
+ raise ValueError("cpu precision: 'float32'")
276
+ self._form = ffbp_cpu.make_ffbp_cpu(plan, coll, wf=wf)
277
+ elif self.backend in ('tpu', 'jax'):
278
+ pol = {'float32': 'fp32_high' if self.backend == 'tpu' else 'fp32', 'three-pass': 'fp32_high',
279
+ 'single-pass': 'fp32_fast', 'float16': 'f16'}.get(precision)
280
+ if pol is None:
281
+ raise ValueError(f'unknown precision {precision!r}')
282
+ filt = 'pallas2' if self.backend == 'tpu' else 'conv'
283
+ self._pol = pol
284
+ # one compiled program per plan signature: patches of a mosaic with equal shapes and filters share it
285
+ # (with the device arrays that depend on the plan alone: filters, tile geometry)
286
+ self._filt = filt
287
+ ng = _jax_group(plan, self.K) if filt == 'pallas2' else 1
288
+ if filt == 'pallas2' and ng < min(_mem.TPU_GROUP, plan['levels'][0]['C']) and not os.environ.get('FASTSAR_TPU_GROUP'):
289
+ _mem._warn(f'tpu: first-level groups of {_mem._nchildren(ng)} instead of {_mem.TPU_GROUP} for lack of device '
290
+ f'memory, which is slower; full speed needs about {_mem._gb(_mem.full_speed(plan, self.K, "tpu")[0])} '
291
+ f'of TPU memory, the device has {_mem._gb(_mem.tpu_capacity())}')
292
+ self._fn, static = self._program(ng)
293
+ self._arrs = ffbp2.device_arrays(pol, plan, coll, static)
294
+ if wf is not None:
295
+ import jax.numpy as jnp
296
+ self._arrs['final']['w'] = jnp.asarray(wf if wf.ndim == 3 else np.stack([wf, 0 * wf, 0 * wf]))
297
+ else:
298
+ raise ValueError(f'unknown backend {backend!r}')
299
+
300
+ def stage(self, S):
301
+ """The work of a call that precedes formation, done ahead of it: on the JAX and TPU backends the checks of S,
302
+ its scaling to float32 planes and their upload to the device; former(former.stage(S)) equals former(S).
303
+ A mosaic stages the next patches on its worker threads while one forms. Other backends return S."""
304
+ if self.backend not in ('jax', 'tpu'):
305
+ return S
306
+ from . import ffbp2
307
+ S = self._host(S)
308
+ hre, him, scale = ffbp2.prepare(self._pol, S)
309
+ hre, him = self._fn.pad(hre, him)
310
+ return _Staged(S, hre, him, scale, self._fn)
311
+
312
+ def _host(self, S):
313
+ if not (self.backend == 'cuda' and type(S).__module__.startswith('cupy')):
314
+ S = np.asarray(S)
315
+ if S.shape != (self.P, self.K):
316
+ raise ValueError(f'phase history must be {(self.P, self.K)}, got {S.shape}')
317
+ _check_history(S)
318
+ if self.window:
319
+ wp, wk = self.wp, self.wk
320
+ if not isinstance(S, np.ndarray):
321
+ import cupy as cp
322
+ wp, wk = cp.asarray(wp), cp.asarray(wk)
323
+ S = (S * wp[:, None] * wk[None, :]).astype(np.complex64)
324
+ else:
325
+ S = S.astype(np.complex64, copy=False)
326
+ if not S.flags.c_contiguous: # a transposed or strided view (the kernels read rows in place)
327
+ S = S.copy()
328
+ return S
329
+
330
+ def __call__(self, S):
331
+ """S [P, K] complex phase history (numpy, or cupy on the cuda backend), or former.stage(S) -> complex64
332
+ image [nx, ny]."""
333
+ staged = S if isinstance(S, _Staged) else None
334
+ if self.backend == 'cuda' and staged is None and not type(S).__module__.startswith('cupy'):
335
+ # a host history is windowed and checked on the device as its row blocks are uploaded (two passes over
336
+ # 1.75 GB on the host cost a 4-vCPU instance more than the whole formation); one that streams is prepared
337
+ # on the host
338
+ S = np.asarray(S)
339
+ if S.shape != (self.P, self.K):
340
+ raise ValueError(f'phase history must be {(self.P, self.K)}, got {S.shape}')
341
+ _check_history(S, finite=False)
342
+ import cupy as cp
343
+ return cp.asnumpy(self._form(S, window=(self.wp, self.wk) if self.window else None, check=True)).astype(np.complex64, copy=False)
344
+ S = staged.S if staged is not None else self._host(S)
345
+ if self.backend == 'cuda':
346
+ import cupy as cp
347
+ return cp.asnumpy(self._form(S)).astype(np.complex64, copy=False) # a host S may stream (ffbp_cuda)
348
+ if self.backend == 'cpu':
349
+ return self._form(S).astype(np.complex64, copy=False)
350
+ import jax
351
+ from . import ffbp2
352
+ runtime_error = _jax_runtime_error()
353
+ while True:
354
+ if staged is not None and staged.fn is self._fn:
355
+ hre, him, scale = staged.hre, staged.him, staged.scale
356
+ staged = None
357
+ else:
358
+ staged = None
359
+ hre, him, scale = ffbp2.prepare(self._pol, S)
360
+ hre, him = self._fn.pad(hre, him)
361
+ try:
362
+ # dispatch is asynchronous: a device that runs out of memory may say so only when the result is read
363
+ re, im = jax.block_until_ready(self._fn(hre, him, self._arrs))
364
+ break
365
+ except runtime_error as e: # device memory: fewer first-level children per group, down to one
366
+ # a kernel's on-chip scratch (VMEM) is sized at compile time and does not depend on the group size
367
+ if 'RESOURCE_EXHAUSTED' not in str(e) or 'vmem' in str(e).lower():
368
+ raise
369
+ ng = self._fn.stages['ng'] # the group the program runs (a divisor of the children)
370
+ if ng <= 1:
371
+ need = _mem.tpu_history_bytes(self._plan, self.K) + _mem.child_bytes(self._plan, 'tpu')
372
+ raise MemoryError(f'{self.backend}: the phase history ({_mem._gb(8.0 * self.P * self.K)} as float32 '
373
+ f'planes) and one first-level child do not fit in device memory; this needs about '
374
+ f'{_mem._gb(need)}. Use the cuda or cpu backend (both stream or read the history in '
375
+ 'place), a device with more memory, or fewer pulses.') from e
376
+ del hre, him
377
+ self._fn = self._program(ng // 2)[0]
378
+ _mem._warn(f'{self.backend}: out of device memory; retrying with first-level groups of '
379
+ f'{_mem._nchildren(self._fn.stages["ng"])} (slower); full speed needs about '
380
+ f'{_mem._gb(_mem.full_speed(self._plan, self.K, "tpu")[0])}')
381
+ return ((np.asarray(re) + 1j * np.asarray(im)) * scale).astype(np.complex64)
382
+
383
+ def memory(self):
384
+ """What full speed needs on this former's device and what is free: dict(backend, needed, available,
385
+ full_speed, parts) in bytes (fastsar.memory)."""
386
+ isz = 4 if self.backend == 'cuda' and self.precision == 'float16' else 8
387
+ need, parts = _mem.full_speed(self._plan, self.K, self.backend, isz)
388
+ avail = _mem.device_available(self.backend)
389
+ return dict(backend=self.backend, needed=need, available=avail, full_speed=need <= avail, parts=parts)
390
+
391
+ def _program(self, ng):
392
+ """The compiled JAX/TPU program for ng first-level children per group, with the plan's device arrays: one per
393
+ plan signature, so that patches of a mosaic with equal shapes and filters share it."""
394
+ from . import ffbp2
395
+ key = (self._pol, self._filt, ng, ffbp2.plan_signature(self._plan))
396
+ with _JAX_LOCK:
397
+ hit = _JAX_PROGRAMS.get(key)
398
+ if hit is None:
399
+ hit = (ffbp2.make_ffbp(self._pol, self._plan, self._filt, 1 << 26, 'direct', pallas_pb=256, pallas_nc=8,
400
+ pallas_ng=ng, pallas_final=2, pallas_gen=3),
401
+ ffbp2.static_arrays(self._pol, self._plan, self._filt))
402
+ if len(_JAX_PROGRAMS) >= 32:
403
+ _JAX_PROGRAMS.pop(next(iter(_JAX_PROGRAMS)))
404
+ _JAX_PROGRAMS[key] = hit
405
+ return hit
406
+
407
+
408
+ def form_image(S, ant, fmin, df, nx, ny, spx, spy, e1=(1.0, 0.0, 0.0), e2=(0.0, 1.0, 0.0), algorithm='ffbp',
409
+ backend='auto', precision='float32', window=True, T='auto', levels=3, pmax=0.4, pfa_guard=300.0, target_db=-40.0,
410
+ ref=None, interp=None, upsample=None, center=None):
411
+ """Form the complex image [nx, ny] (complex64). See the module docstring for the arguments. precision, T, levels,
412
+ pmax and target_db apply to factorized backprojection, pfa_guard to polar format, and interp ('cubic' or
413
+ 'linear'), upsample and center to exact backprojection (algorithm='bp'; see ExactFormer). For more than one
414
+ image of the same geometry, build an ImageFormer (or ExactFormer) once and call it; this function sets one up
415
+ on every call."""
416
+ if algorithm not in ('ffbp', 'pfa', 'bp'):
417
+ raise ValueError("algorithm must be 'ffbp', 'pfa' or 'bp'")
418
+ if algorithm != 'bp':
419
+ for n, v in (('interp', interp), ('upsample', upsample), ('center', center)):
420
+ if v is not None:
421
+ raise ValueError(f"{n} applies to exact backprojection (algorithm='bp'), not to algorithm={algorithm!r}")
422
+ # the formers scan S for NaN and inf themselves; polar format is checked here
423
+ S = _check_history(np.asarray(S), len(_check_positions(ant)), finite=algorithm == 'pfa')
424
+ if algorithm == 'bp':
425
+ from .exact import ExactFormer
426
+ return ExactFormer(ant, fmin, df, S.shape[1], nx, ny, spx, spy, e1, e2, backend=backend, window=window,
427
+ interp='cubic' if interp is None else interp, upsample=upsample, ref=ref, center=center)(S)
428
+ if algorithm == 'pfa':
429
+ nx, ny, e1, e2 = _check_grid(nx, ny, spx, spy, e1, e2)
430
+ P, K = S.shape
431
+ col = Collect(fmin=float(fmin), df=float(df), K=K, ant=np.asarray(ant, np.float64), res=0.5)
432
+ if ref is not None: # polar format needs the samples referenced to |ant|: move them there
433
+ from .io import rereference
434
+ if np.shape(ref) != (P,):
435
+ raise ValueError(f'ref must hold one range per pulse ({P}), got shape {np.shape(ref)}')
436
+ S = rereference(S, float(fmin), float(df), np.asarray(ref, np.float64) - np.linalg.norm(col.ant, axis=1))
437
+ return _pfa(S, col, nx, ny, spx, spy, np.asarray(e1, np.float64), np.asarray(e2, np.float64), pfa_guard, window)
438
+ return ImageFormer(ant, fmin, df, S.shape[1], nx, ny, spx, spy, e1, e2, backend=backend, precision=precision,
439
+ window=window, T=T, levels=levels, pmax=pmax, target_db=target_db, ref=ref)(S)
440
+
441
+
442
+ _PFA_PROGRAMS = {}
443
+ _PFA_LOCK = threading.Lock()
444
+
445
+
446
+ def _pfa(S, col, nx, ny, spx, spy, e1, e2, guard=300.0, window=True):
447
+ """Polar format with its final resampling (removes the planar-wavefront displacement), pulse resampling in
448
+ gather form. The frequency samples are weighted by f_c / f_k so that polar format applies the same spectral
449
+ weighting as backprojection (the polar Jacobian). The compiled program and its geometry arrays are kept per
450
+ collection geometry (the last four), so repeated images of one geometry compile once; the window, the weighting
451
+ and the scaling run on the device."""
452
+ import hashlib, sys
453
+ from . import pfa2
454
+ import jax
455
+ import jax.numpy as jnp
456
+ P, K = S.shape
457
+ key = (P, K, float(col.fmin), float(col.df), nx, ny, float(spx), float(spy), tuple(e1), tuple(e2), float(guard), bool(window),
458
+ hashlib.sha1(np.ascontiguousarray(col.ant).tobytes()).hexdigest())
459
+ with _PFA_LOCK:
460
+ prog = _PFA_PROGRAMS.get(key)
461
+ if prog is None:
462
+ geo = pfa2.geometry(col, nx, ny, spx, spy, e1=e1, e2=e2, guard=guard)
463
+ dist = pfa2.distortion(col, nx, ny, spx, spy, e1, e2)
464
+ fn = pfa2.make_pfa(geo, nx, ny, spx, spy, 'taps', None, jax.lax.Precision.HIGHEST, dist=dist)
465
+ arrs = pfa2.arrays(geo, P, 'taps')
466
+ f_k = col.fmin + np.arange(K) * col.df
467
+ wp, wk = _window(P, K) if window else (np.ones(P), np.ones(K))
468
+ wkf = jnp.asarray((wk * (col.fmin + (K // 2) * col.df) / f_k).astype(np.float32))
469
+ wpd = jnp.asarray(np.asarray(wp, np.float32))
470
+
471
+ @jax.jit
472
+ def prep(Sd):
473
+ x = Sd * wpd[:, None] * wkf[None, :]
474
+ scale = jnp.max(jnp.abs(x))
475
+ scale = jnp.where(scale > 0, scale, 1.0) # an all-zero history gives a zero image, not 0/0
476
+ return jnp.real(x) / scale, jnp.imag(x) / scale, scale
477
+ prog = (fn, arrs, prep)
478
+ with _PFA_LOCK:
479
+ if len(_PFA_PROGRAMS) >= 4:
480
+ _PFA_PROGRAMS.pop(next(iter(_PFA_PROGRAMS)))
481
+ _PFA_PROGRAMS[key] = prog
482
+ fn, (W, alpha, eps_r, shift, eps_a), prep = prog
483
+ if 'cupy' in sys.modules and jax.default_backend() == 'gpu':
484
+ sys.modules['cupy'].get_default_memory_pool().free_all_blocks() # blocks CuPy cached are not JAX's to use
485
+ re, im, scale = prep(jax.device_put(np.asarray(S, np.complex64)))
486
+ out = fn(re, im, W, alpha, eps_r, shift, eps_a)
487
+ return (np.asarray(out) * float(scale)).astype(np.complex64)