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 +16 -0
- fastsar/_build.py +59 -0
- fastsar/api.py +487 -0
- fastsar/autofocus.py +136 -0
- fastsar/bp.py +356 -0
- fastsar/burst.py +265 -0
- fastsar/cphd.py +212 -0
- fastsar/exact.py +596 -0
- fastsar/ffbp.py +516 -0
- fastsar/ffbp2.py +924 -0
- fastsar/ffbp_cpu.cpp +363 -0
- fastsar/ffbp_cpu.py +291 -0
- fastsar/ffbp_cuda.py +803 -0
- fastsar/io.py +280 -0
- fastsar/lowfp.py +44 -0
- fastsar/memory.py +164 -0
- fastsar/pallas_ffbp.py +648 -0
- fastsar/patches.py +774 -0
- fastsar/pfa2.py +386 -0
- fastsar/products.py +420 -0
- fastsar/quality.py +99 -0
- fastsar/sim.py +344 -0
- fastsar/stripmap.py +417 -0
- fastsar-0.1.0.dist-info/METADATA +121 -0
- fastsar-0.1.0.dist-info/RECORD +28 -0
- fastsar-0.1.0.dist-info/WHEEL +5 -0
- fastsar-0.1.0.dist-info/licenses/LICENSE +21 -0
- fastsar-0.1.0.dist-info/top_level.txt +1 -0
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)
|