best-linear-approximation 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.
- best_linear_approximation/__init__.py +43 -0
- best_linear_approximation/_argument_preparation.py +285 -0
- best_linear_approximation/_array_shapes.py +98 -0
- best_linear_approximation/_bla.py +241 -0
- best_linear_approximation/_config.py +5 -0
- best_linear_approximation/_covariance.py +44 -0
- best_linear_approximation/_dataloader.py +359 -0
- best_linear_approximation/_exceptions.py +42 -0
- best_linear_approximation/_frequency_response.py +23 -0
- best_linear_approximation/_linear_algebra.py +110 -0
- best_linear_approximation/_misc.py +37 -0
- best_linear_approximation/_signal_validation.py +111 -0
- best_linear_approximation/_spectra.py +340 -0
- best_linear_approximation/_spectral_validation.py +243 -0
- best_linear_approximation/_typing.py +50 -0
- best_linear_approximation/_uncertainty.py +140 -0
- best_linear_approximation/py.typed +0 -0
- best_linear_approximation/robust/__init__.py +8 -0
- best_linear_approximation/robust/_direct_methods.py +486 -0
- best_linear_approximation/robust/_indirect_methods.py +434 -0
- best_linear_approximation-0.1.0.dist-info/METADATA +178 -0
- best_linear_approximation-0.1.0.dist-info/RECORD +24 -0
- best_linear_approximation-0.1.0.dist-info/WHEEL +4 -0
- best_linear_approximation-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
"""Tools for best linear approximation."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from best_linear_approximation import _dataloader as dataloader
|
|
6
|
+
from best_linear_approximation import robust
|
|
7
|
+
from best_linear_approximation._bla import (
|
|
8
|
+
EstimationMethod,
|
|
9
|
+
ExperimentInfo,
|
|
10
|
+
FrequencyInfo,
|
|
11
|
+
FrequencyResponse,
|
|
12
|
+
NonparametricBLA,
|
|
13
|
+
)
|
|
14
|
+
from best_linear_approximation._dataloader import (
|
|
15
|
+
load_f16,
|
|
16
|
+
load_fine_steering_mirror,
|
|
17
|
+
load_parallel_wiener_hammerstein,
|
|
18
|
+
load_silverbox,
|
|
19
|
+
)
|
|
20
|
+
from best_linear_approximation._spectra import (
|
|
21
|
+
InputSpectrum,
|
|
22
|
+
OutputSpectrum,
|
|
23
|
+
Spectra,
|
|
24
|
+
)
|
|
25
|
+
from best_linear_approximation._uncertainty import Uncertainty
|
|
26
|
+
|
|
27
|
+
__all__ = [
|
|
28
|
+
"EstimationMethod",
|
|
29
|
+
"ExperimentInfo",
|
|
30
|
+
"FrequencyInfo",
|
|
31
|
+
"FrequencyResponse",
|
|
32
|
+
"InputSpectrum",
|
|
33
|
+
"NonparametricBLA",
|
|
34
|
+
"OutputSpectrum",
|
|
35
|
+
"Spectra",
|
|
36
|
+
"Uncertainty",
|
|
37
|
+
"dataloader",
|
|
38
|
+
"load_f16",
|
|
39
|
+
"load_fine_steering_mirror",
|
|
40
|
+
"load_parallel_wiener_hammerstein",
|
|
41
|
+
"load_silverbox",
|
|
42
|
+
"robust",
|
|
43
|
+
]
|
|
@@ -0,0 +1,285 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
from typing import TYPE_CHECKING, cast, overload
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from best_linear_approximation._array_shapes import add_period_axis, to_experiment_layout
|
|
9
|
+
from best_linear_approximation._exceptions import (
|
|
10
|
+
NoiseCovarianceUnavailableWarning,
|
|
11
|
+
PossibleExcitationAmplitudeMismatchWarning,
|
|
12
|
+
PossiblePeriodMismatchWarning,
|
|
13
|
+
TotalCovarianceUnavailableWarning,
|
|
14
|
+
)
|
|
15
|
+
from best_linear_approximation._misc import rms, standardize_channels
|
|
16
|
+
from best_linear_approximation._signal_validation import (
|
|
17
|
+
ContractType,
|
|
18
|
+
SignalContract,
|
|
19
|
+
validate_signal_contract,
|
|
20
|
+
)
|
|
21
|
+
from best_linear_approximation._spectral_validation import (
|
|
22
|
+
resolve_excited_bins,
|
|
23
|
+
validate_sampling_frequency,
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
if TYPE_CHECKING:
|
|
27
|
+
from collections.abc import Mapping
|
|
28
|
+
|
|
29
|
+
from numpy.typing import NDArray
|
|
30
|
+
|
|
31
|
+
from best_linear_approximation._typing import (
|
|
32
|
+
ExcitedBins,
|
|
33
|
+
RealArray,
|
|
34
|
+
SamplingFrequencyHz,
|
|
35
|
+
TimeDomainSignal,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
# Tuned to pass all tests in "tests/test_argument_preparation.py"
|
|
39
|
+
MINIMUM_RELATIVE_PERIOD_MISMATCH = 0.025 # keep in sync with docstring
|
|
40
|
+
MINIMUM_RELATIVE_REALIZATION_MISMATCH = 0.025 # keep in sync with docstring
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
@overload
|
|
44
|
+
def prepare_arguments(
|
|
45
|
+
r: None,
|
|
46
|
+
u: RealArray,
|
|
47
|
+
y: RealArray,
|
|
48
|
+
fs: float,
|
|
49
|
+
excited_bins: NDArray[np.int_] | float,
|
|
50
|
+
contracts: Mapping[ContractType, SignalContract],
|
|
51
|
+
) -> tuple[
|
|
52
|
+
None,
|
|
53
|
+
TimeDomainSignal,
|
|
54
|
+
TimeDomainSignal,
|
|
55
|
+
SamplingFrequencyHz,
|
|
56
|
+
ExcitedBins,
|
|
57
|
+
]: ...
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
@overload
|
|
61
|
+
def prepare_arguments(
|
|
62
|
+
r: RealArray,
|
|
63
|
+
u: RealArray,
|
|
64
|
+
y: RealArray,
|
|
65
|
+
fs: float,
|
|
66
|
+
excited_bins: NDArray[np.int_] | float,
|
|
67
|
+
contracts: Mapping[ContractType, SignalContract],
|
|
68
|
+
) -> tuple[
|
|
69
|
+
TimeDomainSignal,
|
|
70
|
+
TimeDomainSignal,
|
|
71
|
+
TimeDomainSignal,
|
|
72
|
+
SamplingFrequencyHz,
|
|
73
|
+
ExcitedBins,
|
|
74
|
+
]: ...
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def prepare_arguments( # noqa: PLR0913, PLR0917
|
|
78
|
+
r: RealArray | None,
|
|
79
|
+
u: RealArray,
|
|
80
|
+
y: RealArray,
|
|
81
|
+
fs: float,
|
|
82
|
+
excited_bins: NDArray[np.int_] | float,
|
|
83
|
+
contracts: Mapping[ContractType, SignalContract],
|
|
84
|
+
) -> tuple[
|
|
85
|
+
TimeDomainSignal | None,
|
|
86
|
+
TimeDomainSignal,
|
|
87
|
+
TimeDomainSignal,
|
|
88
|
+
SamplingFrequencyHz,
|
|
89
|
+
ExcitedBins,
|
|
90
|
+
]:
|
|
91
|
+
"""Validate and resolve all arguments.
|
|
92
|
+
|
|
93
|
+
Ensures the signals conform to a supported contract and transforms them into the
|
|
94
|
+
canonical five-dimensional experiment layout, verifies the sampling frequency,
|
|
95
|
+
and resolves the excited bins. Warns when the data is insufficient to estimate the
|
|
96
|
+
noise or total covariance, and when adjacent output periods differ substantially
|
|
97
|
+
or excitation amplitudes change between realizations.
|
|
98
|
+
|
|
99
|
+
If ``excited_bins`` is an array, it is validated directly. If it is a float, it is
|
|
100
|
+
interpreted as a threshold for detecting the excited bins, using ``r`` if available
|
|
101
|
+
and ``u`` otherwise.
|
|
102
|
+
"""
|
|
103
|
+
contract_type = validate_signal_contract(r, u, y, contracts)
|
|
104
|
+
|
|
105
|
+
nu = u.shape[1]
|
|
106
|
+
if contract_type is ContractType.REALIZATION:
|
|
107
|
+
r = to_experiment_layout(r, nu) if r is not None else None
|
|
108
|
+
u = to_experiment_layout(u, nu)
|
|
109
|
+
y = to_experiment_layout(y, nu)
|
|
110
|
+
else:
|
|
111
|
+
r = add_period_axis(r) if r is not None else None
|
|
112
|
+
u = add_period_axis(u)
|
|
113
|
+
|
|
114
|
+
r = cast("TimeDomainSignal", r) if r is not None else None
|
|
115
|
+
u = cast("TimeDomainSignal", u)
|
|
116
|
+
y = cast("TimeDomainSignal", y)
|
|
117
|
+
|
|
118
|
+
fs = validate_sampling_frequency(fs)
|
|
119
|
+
excited_bins = resolve_excited_bins(excited_bins, r if r is not None else u, fs)
|
|
120
|
+
|
|
121
|
+
n_experiments, n_periods = y.shape[-2:]
|
|
122
|
+
|
|
123
|
+
if n_experiments == 1:
|
|
124
|
+
_warn_total_covariance_unavailable()
|
|
125
|
+
if n_periods == 1:
|
|
126
|
+
_warn_noise_covariance_unavailable()
|
|
127
|
+
|
|
128
|
+
max_bin = excited_bins[-1]
|
|
129
|
+
if n_experiments > 1:
|
|
130
|
+
_warn_if_excitation_amplitude_mismatch(r if r is not None else u, max_bin)
|
|
131
|
+
if n_periods > 1:
|
|
132
|
+
_warn_if_output_spectra_mismatch(y, max_bin)
|
|
133
|
+
|
|
134
|
+
return r, u, y, fs, excited_bins
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def _warn_if_output_spectra_mismatch(y: RealArray, max_bin: int) -> None:
|
|
138
|
+
"""Warn when the aggregate spectral mismatch between adjacent periods exceeds 2.5%.
|
|
139
|
+
|
|
140
|
+
Requires ``n_periods > 1``.
|
|
141
|
+
|
|
142
|
+
Only ``rfft`` bins from DC through ``max_bin`` are used. An appropriate choice of
|
|
143
|
+
``max_bin`` reduces the influence of higher-frequency measurement noise and
|
|
144
|
+
emphasizes mismatches arising from drift and initial-condition effects.
|
|
145
|
+
"""
|
|
146
|
+
n_periods = y.shape[-1]
|
|
147
|
+
if n_periods <= 1:
|
|
148
|
+
msg = f"y must have more than one period, got n_periods={n_periods}."
|
|
149
|
+
raise ValueError(msg)
|
|
150
|
+
|
|
151
|
+
y = standardize_channels(y)
|
|
152
|
+
|
|
153
|
+
stop_bin = max_bin + 1
|
|
154
|
+
spectrum = np.fft.rfft(y, axis=0)[:stop_bin]
|
|
155
|
+
later_spectrum = spectrum[..., 1:]
|
|
156
|
+
spectral_difference = later_spectrum - spectrum[..., :-1]
|
|
157
|
+
|
|
158
|
+
reduction_axis = tuple(range(y.ndim - 1)) # all axes except the last (period) axis
|
|
159
|
+
difference_rms = rms(spectral_difference, axis=reduction_axis)
|
|
160
|
+
reference_rms = rms(later_spectrum, axis=reduction_axis)
|
|
161
|
+
|
|
162
|
+
relative_difference = np.divide(
|
|
163
|
+
difference_rms,
|
|
164
|
+
reference_rms,
|
|
165
|
+
out=np.where(difference_rms == 0, 0.0, np.inf),
|
|
166
|
+
where=reference_rms != 0,
|
|
167
|
+
)
|
|
168
|
+
|
|
169
|
+
exceeds_threshold = relative_difference > MINIMUM_RELATIVE_PERIOD_MISMATCH
|
|
170
|
+
if np.any(exceeds_threshold):
|
|
171
|
+
values = ", ".join(f"{value:.2%}" for value in relative_difference)
|
|
172
|
+
affected_pairs = np.flatnonzero(exceeds_threshold)
|
|
173
|
+
n_period_pairs = n_periods - 1
|
|
174
|
+
n_affected_pairs = affected_pairs.size
|
|
175
|
+
if n_period_pairs == 1:
|
|
176
|
+
scope = "the"
|
|
177
|
+
period_pair_label = "period pair"
|
|
178
|
+
difference_label = "difference"
|
|
179
|
+
else:
|
|
180
|
+
scope = (
|
|
181
|
+
f"all {n_period_pairs}"
|
|
182
|
+
if n_affected_pairs == n_period_pairs
|
|
183
|
+
else f"{n_affected_pairs} of {n_period_pairs}"
|
|
184
|
+
)
|
|
185
|
+
period_pair_label = "period pairs"
|
|
186
|
+
difference_label = "differences"
|
|
187
|
+
msg = (
|
|
188
|
+
f"The relative spectral mismatch between {scope} adjacent {period_pair_label} "
|
|
189
|
+
f"exceeds the threshold of {MINIMUM_RELATIVE_PERIOD_MISMATCH:.2%}. This may "
|
|
190
|
+
f"indicate transients, drift, changing excitation, or period-alignment "
|
|
191
|
+
f"issues. Aggregated relative {difference_label}: "
|
|
192
|
+
f"[{values}]."
|
|
193
|
+
)
|
|
194
|
+
warnings.warn(msg, PossiblePeriodMismatchWarning, stacklevel=2)
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def _warn_if_excitation_amplitude_mismatch(
|
|
198
|
+
signal: RealArray,
|
|
199
|
+
max_bin: int,
|
|
200
|
+
) -> None:
|
|
201
|
+
"""Warn when excitation magnitude changes between adjacent realizations.
|
|
202
|
+
|
|
203
|
+
Requires ``n_experiments > 1``.
|
|
204
|
+
|
|
205
|
+
Only ``rfft`` bins from DC through ``max_bin`` are used. The input and experiment
|
|
206
|
+
axes are merged into realization layout using column-major ordering. Adjacent
|
|
207
|
+
realizations are compared separately for every physical input channel. Each
|
|
208
|
+
channel is standardized before comparison, so fixed amplitude differences between
|
|
209
|
+
channels do not trigger a warning.
|
|
210
|
+
"""
|
|
211
|
+
n_experiments = signal.shape[-2]
|
|
212
|
+
if n_experiments <= 1:
|
|
213
|
+
msg = (
|
|
214
|
+
"signal must contain more than one experiment, "
|
|
215
|
+
f"got n_experiments={n_experiments}."
|
|
216
|
+
)
|
|
217
|
+
raise ValueError(msg)
|
|
218
|
+
|
|
219
|
+
signal = standardize_channels(signal)
|
|
220
|
+
|
|
221
|
+
stop_bin = max_bin + 1
|
|
222
|
+
spectrum = np.fft.rfft(signal, axis=0)[:stop_bin]
|
|
223
|
+
n_bins, n_channels, nu, _, n_periods = spectrum.shape
|
|
224
|
+
n_realizations = nu * n_experiments
|
|
225
|
+
realization_spectrum = spectrum.reshape(
|
|
226
|
+
n_bins,
|
|
227
|
+
n_channels,
|
|
228
|
+
n_realizations,
|
|
229
|
+
n_periods,
|
|
230
|
+
order="F",
|
|
231
|
+
)
|
|
232
|
+
realization_magnitude_spectrum = np.abs(realization_spectrum)
|
|
233
|
+
|
|
234
|
+
later_magnitude_spectrum = realization_magnitude_spectrum[..., 1:, :]
|
|
235
|
+
spectral_difference = later_magnitude_spectrum - realization_magnitude_spectrum[..., :-1, :]
|
|
236
|
+
|
|
237
|
+
reduction_axes = (0, 3)
|
|
238
|
+
difference_rms = rms(spectral_difference, axis=reduction_axes)
|
|
239
|
+
reference_rms = rms(later_magnitude_spectrum, axis=reduction_axes)
|
|
240
|
+
relative_difference = np.divide(
|
|
241
|
+
difference_rms,
|
|
242
|
+
reference_rms,
|
|
243
|
+
out=np.where(difference_rms == 0, 0.0, np.inf),
|
|
244
|
+
where=reference_rms != 0,
|
|
245
|
+
)
|
|
246
|
+
|
|
247
|
+
exceeds_threshold = relative_difference > MINIMUM_RELATIVE_REALIZATION_MISMATCH
|
|
248
|
+
if np.any(exceeds_threshold):
|
|
249
|
+
values = ", ".join(f"{value:.2%}" for value in relative_difference.ravel())
|
|
250
|
+
affected_pairs = np.flatnonzero(exceeds_threshold)
|
|
251
|
+
n_realization_pairs = n_realizations - 1
|
|
252
|
+
n_affected_pairs = affected_pairs.size
|
|
253
|
+
n_channel_realization_pairs = n_channels * n_realization_pairs
|
|
254
|
+
if n_channel_realization_pairs == 1:
|
|
255
|
+
scope = "the"
|
|
256
|
+
channel_realization_pair_label = "channel-realization pair"
|
|
257
|
+
difference_label = "difference"
|
|
258
|
+
else:
|
|
259
|
+
scope = (
|
|
260
|
+
f"all {n_channel_realization_pairs}"
|
|
261
|
+
if n_affected_pairs == n_channel_realization_pairs
|
|
262
|
+
else f"{n_affected_pairs} of {n_channel_realization_pairs}"
|
|
263
|
+
)
|
|
264
|
+
channel_realization_pair_label = "channel-realization pairs"
|
|
265
|
+
difference_label = "differences"
|
|
266
|
+
msg = (
|
|
267
|
+
f"The relative mismatch between {scope} adjacent {channel_realization_pair_label} "
|
|
268
|
+
f"exceeds the threshold of {MINIMUM_RELATIVE_REALIZATION_MISMATCH:.2%}. This "
|
|
269
|
+
"may indicate an excitation-amplitude change between realizations. Aggregated relative "
|
|
270
|
+
f"{difference_label}: [{values}]."
|
|
271
|
+
)
|
|
272
|
+
warnings.warn(msg, PossibleExcitationAmplitudeMismatchWarning, stacklevel=2)
|
|
273
|
+
|
|
274
|
+
|
|
275
|
+
def _warn_noise_covariance_unavailable() -> None:
|
|
276
|
+
msg = "Only a single period is provided, so the noise covariance cannot be estimated."
|
|
277
|
+
warnings.warn(msg, NoiseCovarianceUnavailableWarning, stacklevel=2)
|
|
278
|
+
|
|
279
|
+
|
|
280
|
+
def _warn_total_covariance_unavailable() -> None:
|
|
281
|
+
msg = (
|
|
282
|
+
"Only a single experiment is provided, so the total covariance "
|
|
283
|
+
"(noise plus nonlinear distortions) cannot be estimated."
|
|
284
|
+
)
|
|
285
|
+
warnings.warn(msg, TotalCovarianceUnavailableWarning, stacklevel=2)
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
from typing import TYPE_CHECKING, cast, overload
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from best_linear_approximation._exceptions import (
|
|
9
|
+
InsufficientExperimentsError,
|
|
10
|
+
RealizationsTruncatedWarning,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
if TYPE_CHECKING:
|
|
14
|
+
from best_linear_approximation._typing import (
|
|
15
|
+
FrequencyDomainSignal,
|
|
16
|
+
RealArray,
|
|
17
|
+
TimeDomainSignal,
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
CANONICAL_SIGNAL_NDIM = 5
|
|
21
|
+
MINIMUM_SIGNAL_NDIM = 3
|
|
22
|
+
EXPERIMENT_LAYOUT_WITHOUT_PERIOD_NDIM = 4
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def to_experiment_layout(
|
|
26
|
+
signal: RealArray,
|
|
27
|
+
nu: int,
|
|
28
|
+
) -> TimeDomainSignal:
|
|
29
|
+
"""Convert a signal from realization layout to five-dimensional experiment layout.
|
|
30
|
+
|
|
31
|
+
Splits the realization axis of a signal with shape
|
|
32
|
+
``(n_samples, n_channels, n_realizations[, n_periods])`` into input and experiment
|
|
33
|
+
axes using column-major ordering. Adds a singleton period axis if none is present.
|
|
34
|
+
The resulting shape is ``(n_samples, n_channels, nu, n_experiments, n_periods)``.
|
|
35
|
+
|
|
36
|
+
Raises an ``InsufficientExperimentsError`` if ``n_realizations < nu``. Warns if
|
|
37
|
+
``n_realizations`` is not divisible by ``nu`` and discards the remaining realizations.
|
|
38
|
+
"""
|
|
39
|
+
n_realizations = signal.shape[2]
|
|
40
|
+
n_experiments = n_realizations // nu
|
|
41
|
+
if n_experiments == 0:
|
|
42
|
+
msg = (
|
|
43
|
+
f"The number of realizations ({n_realizations}) is less than the number of "
|
|
44
|
+
f"input channels ({nu}). No frequency response can be estimated."
|
|
45
|
+
)
|
|
46
|
+
raise InsufficientExperimentsError(msg)
|
|
47
|
+
|
|
48
|
+
n_effective_realizations = n_experiments * nu
|
|
49
|
+
if n_effective_realizations < n_realizations:
|
|
50
|
+
msg = (
|
|
51
|
+
f"The number of realizations ({n_realizations}) is not a multiple of "
|
|
52
|
+
f"the number of input channels ({nu}). Only the first "
|
|
53
|
+
f"{n_effective_realizations} realizations will be used for estimation."
|
|
54
|
+
)
|
|
55
|
+
warnings.warn(msg, RealizationsTruncatedWarning, stacklevel=2)
|
|
56
|
+
|
|
57
|
+
signal = signal[:, :, :n_effective_realizations, ...]
|
|
58
|
+
signal = signal.reshape(*signal.shape[:2], nu, n_experiments, *signal.shape[3:], order="F")
|
|
59
|
+
|
|
60
|
+
return cast(
|
|
61
|
+
"TimeDomainSignal",
|
|
62
|
+
signal
|
|
63
|
+
if signal.ndim == CANONICAL_SIGNAL_NDIM
|
|
64
|
+
else add_period_axis(signal),
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def add_period_axis(signal: RealArray) -> RealArray:
|
|
69
|
+
"""Add a singleton period axis to a four-dimensional experiment signal.
|
|
70
|
+
|
|
71
|
+
Five-dimensional signals are returned unchanged.
|
|
72
|
+
"""
|
|
73
|
+
return signal[..., None] if signal.ndim == EXPERIMENT_LAYOUT_WITHOUT_PERIOD_NDIM else signal
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
@overload
|
|
77
|
+
def as_batched_matrices(signal: TimeDomainSignal) -> TimeDomainSignal: ...
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
@overload
|
|
81
|
+
def as_batched_matrices(signal: FrequencyDomainSignal) -> FrequencyDomainSignal: ...
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def as_batched_matrices(
|
|
85
|
+
signal: TimeDomainSignal | FrequencyDomainSignal,
|
|
86
|
+
) -> TimeDomainSignal | FrequencyDomainSignal:
|
|
87
|
+
"""Arrange a signal as matrices for batched linear algebra.
|
|
88
|
+
|
|
89
|
+
Transforms an array with shape ``(n_leading, n_rows, n_cols, ...)`` into
|
|
90
|
+
one with shape ``(n_leading, ..., n_rows, n_cols)``.
|
|
91
|
+
"""
|
|
92
|
+
if signal.ndim < MINIMUM_SIGNAL_NDIM:
|
|
93
|
+
msg = (
|
|
94
|
+
f"Expected a signal with at least {MINIMUM_SIGNAL_NDIM} dimensions, got {signal.ndim}."
|
|
95
|
+
)
|
|
96
|
+
raise ValueError(msg)
|
|
97
|
+
|
|
98
|
+
return np.moveaxis(signal, (1, 2), (-2, -1))
|
|
@@ -0,0 +1,241 @@
|
|
|
1
|
+
"""Result types for best linear approximation estimates."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
from enum import Enum
|
|
7
|
+
from typing import TYPE_CHECKING, Self
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
from best_linear_approximation._covariance import project_onto_positive_semidefinite
|
|
12
|
+
from best_linear_approximation._uncertainty import Uncertainty
|
|
13
|
+
|
|
14
|
+
if TYPE_CHECKING:
|
|
15
|
+
from numpy.typing import NDArray
|
|
16
|
+
|
|
17
|
+
from best_linear_approximation._spectra import Spectra
|
|
18
|
+
from best_linear_approximation._typing import (
|
|
19
|
+
ComplexArray,
|
|
20
|
+
ExcitedBins,
|
|
21
|
+
RealArray,
|
|
22
|
+
SamplingFrequencyHz,
|
|
23
|
+
TimeDomainSignal,
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True)
|
|
28
|
+
class FrequencyResponse:
|
|
29
|
+
"""Frequency response estimate and its uncertainty estimates."""
|
|
30
|
+
|
|
31
|
+
value: ComplexArray
|
|
32
|
+
noise: Uncertainty
|
|
33
|
+
nonlinear: Uncertainty
|
|
34
|
+
total: Uncertainty
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class EstimationMethod(Enum):
|
|
38
|
+
"""Method used to estimate a best linear approximation."""
|
|
39
|
+
|
|
40
|
+
ROBUST_DIRECT_KNOWN_INPUT = "robust_direct_known_input"
|
|
41
|
+
ROBUST_DIRECT_NOISY_INPUT = "robust_direct_noisy_input"
|
|
42
|
+
ROBUST_INDIRECT = "robust_indirect"
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
@dataclass(frozen=True)
|
|
46
|
+
class ExperimentInfo:
|
|
47
|
+
"""Metadata describing the recordings used to estimate a BLA.
|
|
48
|
+
|
|
49
|
+
Attributes
|
|
50
|
+
----------
|
|
51
|
+
estimation_method : EstimationMethod
|
|
52
|
+
Method used to estimate the BLA.
|
|
53
|
+
n_samples : int
|
|
54
|
+
Number of time-domain samples in each period.
|
|
55
|
+
nu : int
|
|
56
|
+
Number of input channels.
|
|
57
|
+
ny : int
|
|
58
|
+
Number of output channels.
|
|
59
|
+
n_experiments : int
|
|
60
|
+
Number of experiments.
|
|
61
|
+
n_periods : int
|
|
62
|
+
Number of output periods per experiment.
|
|
63
|
+
u_shape : tuple of int
|
|
64
|
+
Canonical time-domain shape of the input signal.
|
|
65
|
+
y_shape : tuple of int
|
|
66
|
+
Canonical time-domain shape of the output signal.
|
|
67
|
+
r_shape : tuple of int or None
|
|
68
|
+
Canonical time-domain shape of the reference signal, when available.
|
|
69
|
+
|
|
70
|
+
"""
|
|
71
|
+
|
|
72
|
+
estimation_method: EstimationMethod
|
|
73
|
+
n_samples: int
|
|
74
|
+
nu: int
|
|
75
|
+
ny: int
|
|
76
|
+
n_experiments: int
|
|
77
|
+
n_periods: int
|
|
78
|
+
u_shape: tuple[int, ...]
|
|
79
|
+
y_shape: tuple[int, ...]
|
|
80
|
+
r_shape: tuple[int, ...] | None = None
|
|
81
|
+
|
|
82
|
+
@property
|
|
83
|
+
def n_realizations(self) -> int:
|
|
84
|
+
"""Number of effective realizations after arranging data by input direction."""
|
|
85
|
+
return self.nu * self.n_experiments
|
|
86
|
+
|
|
87
|
+
@classmethod
|
|
88
|
+
def from_signals(
|
|
89
|
+
cls,
|
|
90
|
+
estimation_method: EstimationMethod,
|
|
91
|
+
u: TimeDomainSignal,
|
|
92
|
+
y: TimeDomainSignal,
|
|
93
|
+
r: TimeDomainSignal | None = None,
|
|
94
|
+
) -> Self:
|
|
95
|
+
"""Create recording metadata from canonical time-domain signals."""
|
|
96
|
+
n_samples, ny, nu, n_experiments, n_periods = y.shape
|
|
97
|
+
return cls(
|
|
98
|
+
estimation_method=estimation_method,
|
|
99
|
+
n_samples=n_samples,
|
|
100
|
+
nu=nu,
|
|
101
|
+
ny=ny,
|
|
102
|
+
n_experiments=n_experiments,
|
|
103
|
+
n_periods=n_periods,
|
|
104
|
+
u_shape=u.shape,
|
|
105
|
+
y_shape=y.shape,
|
|
106
|
+
r_shape=r.shape if r is not None else None,
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
@dataclass(frozen=True)
|
|
111
|
+
class FrequencyInfo:
|
|
112
|
+
"""Metadata for the frequency content of a best linear approximation.
|
|
113
|
+
|
|
114
|
+
Attributes
|
|
115
|
+
----------
|
|
116
|
+
fs : float
|
|
117
|
+
Sampling frequency in Hz.
|
|
118
|
+
f_res : float
|
|
119
|
+
Frequency resolution in Hz.
|
|
120
|
+
f_min : float
|
|
121
|
+
Lowest excited frequency in Hz.
|
|
122
|
+
f_max : float
|
|
123
|
+
Highest excited frequency in Hz.
|
|
124
|
+
freqs : RealArray
|
|
125
|
+
Non-negative DFT frequencies in Hz.
|
|
126
|
+
excited_bins : NDArray[np.int_]
|
|
127
|
+
Indices of the excited DFT frequencies.
|
|
128
|
+
non_excited_bins : NDArray[np.int_]
|
|
129
|
+
Indices of the non-excited DFT frequencies.
|
|
130
|
+
|
|
131
|
+
"""
|
|
132
|
+
|
|
133
|
+
fs: float
|
|
134
|
+
f_res: float
|
|
135
|
+
f_min: float
|
|
136
|
+
f_max: float
|
|
137
|
+
freqs: RealArray
|
|
138
|
+
excited_bins: NDArray[np.int_]
|
|
139
|
+
non_excited_bins: NDArray[np.int_]
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
@dataclass(frozen=True)
|
|
143
|
+
class NonparametricBLA:
|
|
144
|
+
"""Nonparametric best linear approximation estimate.
|
|
145
|
+
|
|
146
|
+
Attributes
|
|
147
|
+
----------
|
|
148
|
+
G : FrequencyResponse
|
|
149
|
+
Frequency response estimate and its uncertainty estimates.
|
|
150
|
+
spectra : Spectra
|
|
151
|
+
Input, output, and optional reference spectra with their uncertainty estimates.
|
|
152
|
+
freq : FrequencyInfo
|
|
153
|
+
Frequency metadata for the estimate.
|
|
154
|
+
experiment : ExperimentInfo
|
|
155
|
+
Recording metadata and estimation method.
|
|
156
|
+
|
|
157
|
+
"""
|
|
158
|
+
|
|
159
|
+
G: FrequencyResponse
|
|
160
|
+
spectra: Spectra
|
|
161
|
+
freq: FrequencyInfo
|
|
162
|
+
experiment: ExperimentInfo
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def create_frequency_info(
|
|
166
|
+
n_samples: int,
|
|
167
|
+
fs: SamplingFrequencyHz,
|
|
168
|
+
excited_bins: ExcitedBins,
|
|
169
|
+
) -> FrequencyInfo:
|
|
170
|
+
"""Create frequency metadata from validated sampling and excitation data."""
|
|
171
|
+
n_bins = n_samples // 2 + 1
|
|
172
|
+
f_res = fs / n_samples
|
|
173
|
+
freqs = np.arange(n_bins) * f_res
|
|
174
|
+
non_excited_bins = np.setdiff1d(np.arange(n_bins), excited_bins)
|
|
175
|
+
f_min = float(freqs[excited_bins[0]])
|
|
176
|
+
f_max = float(freqs[excited_bins[-1]])
|
|
177
|
+
return FrequencyInfo(
|
|
178
|
+
fs=fs,
|
|
179
|
+
f_res=f_res,
|
|
180
|
+
f_min=f_min,
|
|
181
|
+
f_max=f_max,
|
|
182
|
+
freqs=freqs,
|
|
183
|
+
excited_bins=excited_bins,
|
|
184
|
+
non_excited_bins=non_excited_bins,
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def create_bla_frequency_response(
|
|
189
|
+
G_bla: ComplexArray,
|
|
190
|
+
G_total_cov: ComplexArray | None,
|
|
191
|
+
G_noise_cov: ComplexArray | None,
|
|
192
|
+
n_samples: int,
|
|
193
|
+
excited_bins: ExcitedBins,
|
|
194
|
+
) -> FrequencyResponse:
|
|
195
|
+
"""Create a frequency response estimate from its value and covariances."""
|
|
196
|
+
nonlinear_cov = None
|
|
197
|
+
if G_total_cov is not None and G_noise_cov is not None:
|
|
198
|
+
nonlinear_cov = project_onto_positive_semidefinite(G_total_cov - G_noise_cov)
|
|
199
|
+
|
|
200
|
+
marginal_var_shape = G_bla.shape[1:]
|
|
201
|
+
noise = _create_frequency_domain_uncertainty(
|
|
202
|
+
G_noise_cov,
|
|
203
|
+
marginal_var_shape,
|
|
204
|
+
G_bla,
|
|
205
|
+
n_samples,
|
|
206
|
+
excited_bins,
|
|
207
|
+
)
|
|
208
|
+
nonlinear = _create_frequency_domain_uncertainty(
|
|
209
|
+
nonlinear_cov,
|
|
210
|
+
marginal_var_shape,
|
|
211
|
+
G_bla,
|
|
212
|
+
n_samples,
|
|
213
|
+
excited_bins,
|
|
214
|
+
)
|
|
215
|
+
total = _create_frequency_domain_uncertainty(
|
|
216
|
+
G_total_cov,
|
|
217
|
+
marginal_var_shape,
|
|
218
|
+
G_bla,
|
|
219
|
+
n_samples,
|
|
220
|
+
excited_bins,
|
|
221
|
+
)
|
|
222
|
+
return FrequencyResponse(G_bla, noise, nonlinear, total)
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def _create_frequency_domain_uncertainty(
|
|
226
|
+
covariance: ComplexArray | None,
|
|
227
|
+
marginal_var_shape: tuple[int, ...],
|
|
228
|
+
spectrum: ComplexArray,
|
|
229
|
+
n_samples: int,
|
|
230
|
+
frequency_bins: ExcitedBins,
|
|
231
|
+
) -> Uncertainty:
|
|
232
|
+
if covariance is None:
|
|
233
|
+
return Uncertainty.unavailable()
|
|
234
|
+
|
|
235
|
+
return Uncertainty.from_cov(
|
|
236
|
+
covariance,
|
|
237
|
+
marginal_var_shape,
|
|
238
|
+
spectrum=spectrum,
|
|
239
|
+
n_samples=n_samples,
|
|
240
|
+
frequency_bins=frequency_bins,
|
|
241
|
+
)
|